Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
---
repos:
- repo: https://github.com/pre-commit/pre-commit-hooks
rev: v5.0.0
rev: v6.0.0
hooks:
- id: check-ast
- id: check-json
Expand All @@ -14,6 +14,7 @@ repos:
# locally; the hook still scans staged files for leaked AWS keys.
args: [--allow-missing-credentials]
- id: detect-private-key
exclude: "pkg/auth/saml/provider.go"
- id: check-yaml
- id: end-of-file-fixer
- id: trailing-whitespace
Expand All @@ -23,14 +24,14 @@ repos:
- id: requirements-txt-fixer

- repo: https://github.com/Bahjat/pre-commit-golang
rev: v1.0.5
rev: v1.0.6
hooks:
- id: go-fmt-import
- id: go-static-check # install https://staticcheck.io/docs/
- id: go-unit-tests

- repo: https://github.com/golangci/golangci-lint
rev: v2.11.4
rev: v2.13.2
hooks:
- id: golangci-lint
args: [--config=.golangci.yml]
8 changes: 4 additions & 4 deletions cmd/api/handlers/auth_oidc_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -29,10 +29,10 @@ import (
// happy-path verification is already covered by the pkg-level tests,
// so we don't need a working /token endpoint here.
type fakeIdP struct {
srv *httptest.Server
key *rsa.PrivateKey
keyID string
issuer string // overridden after srv.URL is known
srv *httptest.Server
key *rsa.PrivateKey
keyID string
issuer string // overridden after srv.URL is known
}

func newFakeIdP(t *testing.T) *fakeIdP {
Expand Down
6 changes: 3 additions & 3 deletions cmd/api/handlers/auth_resolve.go
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ func (h *HandlersApi) resolveFederatedUser(identity auth.ResolvedIdentity, polic
if existing.AuthSource != policy.authSource {
linkedLocal := existing.AuthSource == ""
if err := h.Users.ChangeAuthSource(existing.Username, policy.authSource); err != nil {
return users.AdminUser{}, fmt.Errorf("%w: updating auth source: %v", ErrAuthUserRejected, err)
return users.AdminUser{}, fmt.Errorf("%w: updating auth source: %w", ErrAuthUserRejected, err)
}
existing.AuthSource = policy.authSource
if linkedLocal {
Expand Down Expand Up @@ -135,14 +135,14 @@ func (h *HandlersApi) resolveFederatedUser(identity auth.ResolvedIdentity, polic
false, // service = false
)
if err != nil {
return users.AdminUser{}, fmt.Errorf("%w: new user: %v", ErrAuthUserRejected, err)
return users.AdminUser{}, fmt.Errorf("%w: new user: %w", ErrAuthUserRejected, err)
}
// Tag the row with the provider type (oidc / saml) so the Users
// page can display the right badge. Purely informational; the auth
// flow itself doesn't gate on this field.
u.AuthSource = policy.authSource
if err := h.Users.Create(u); err != nil {
return users.AdminUser{}, fmt.Errorf("%w: create user: %v", ErrAuthUserRejected, err)
return users.AdminUser{}, fmt.Errorf("%w: create user: %w", ErrAuthUserRejected, err)
}
return u, nil
}
2 changes: 1 addition & 1 deletion cmd/api/handlers/events.go
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,7 @@ func (h *HandlersApi) EventsHandler(w http.ResponseWriter, r *http.Request) {
}
environmentID := uint(0)
environmentUUID := "all"
if !(globalOnly && envSelector == "all") {
if !globalOnly || envSelector != "all" {
env, err := h.Envs.Get(envSelector)
if err != nil {
apiErrorResponse(w, "environment not found", http.StatusNotFound, nil)
Expand Down
2 changes: 1 addition & 1 deletion cmd/api/handlers/handlers.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package handlers

import (
"github.com/jmpsec/osctrl/pkg/events"
"net/http"
"time"

Expand All @@ -14,6 +13,7 @@ import (
"github.com/jmpsec/osctrl/pkg/config"
"github.com/jmpsec/osctrl/pkg/console"
"github.com/jmpsec/osctrl/pkg/environments"
"github.com/jmpsec/osctrl/pkg/events"
"github.com/jmpsec/osctrl/pkg/fileexplorer"
"github.com/jmpsec/osctrl/pkg/geoip"
"github.com/jmpsec/osctrl/pkg/health"
Expand Down
7 changes: 4 additions & 3 deletions cmd/api/handlers/nodes_inactivity_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -80,14 +80,15 @@ func TestNodeConsumersEnvironmentThresholds(t *testing.T) {
if tc.active == (i == 1) {
want = 1
}
if tc.paged {
switch {
case tc.paged:
require.Equal(t, http.StatusOK, rr.Code, rr.Body.String())
var page types.NodesPagedResponse
require.NoError(t, json.Unmarshal(rr.Body.Bytes(), &page))
require.Equal(t, want, page.TotalItems)
} else if want == 0 {
case want == 0:
require.Equal(t, http.StatusNotFound, rr.Code, rr.Body.String())
} else {
default:
require.Equal(t, http.StatusOK, rr.Code, rr.Body.String())
var got []nodes.OsqueryNode
require.NoError(t, json.Unmarshal(rr.Body.Bytes(), &got))
Expand Down
2 changes: 1 addition & 1 deletion cmd/api/handlers/query_dispatch_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@ func TestCreateQueryInvalidatesDispatchAfterCommit(t *testing.T) {
require.NoError(t, err)
require.True(t, cached, "do not invalidate while query creation is uncommitted")
if rollback {
tx.AddError(errors.New("target write failed"))
_ = tx.AddError(errors.New("target write failed"))
}
if name == "canceled-request" {
cancel()
Expand Down
12 changes: 6 additions & 6 deletions cmd/api/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -611,7 +611,6 @@ func osctrlAPIService() {
if err != nil {
log.Fatal().Err(err).Msg("invalid events configuration")
}
defer eventBus.Close()
queriesmgr.Events = eventBus
filecarves.Events = eventBus
if alertsMgr != nil {
Expand Down Expand Up @@ -1352,9 +1351,8 @@ func osctrlAPIService() {
}
if tlsTermination {
srv.TLSConfig = &tls.Config{
MinVersion: tls.VersionTLS12,
CurvePreferences: []tls.CurveID{tls.CurveP521, tls.CurveP384, tls.CurveP256},
PreferServerCipherSuites: true,
MinVersion: tls.VersionTLS12,
CurvePreferences: []tls.CurveID{tls.CurveP521, tls.CurveP384, tls.CurveP256},
CipherSuites: []uint16{
tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
tls.TLS_ECDHE_RSA_WITH_AES_256_CBC_SHA,
Expand All @@ -1378,22 +1376,24 @@ func osctrlAPIService() {
// Stop live streams before draining on process shutdown or restart.
shutdownSignals := make(chan os.Signal, 1)
signal.Notify(shutdownSignals, syscall.SIGINT, syscall.SIGTERM)
defer signal.Stop(shutdownSignals)
// Wait for either a server error or a restart signal.
select {
case err := <-serverErr:
if err != nil && !errors.Is(err, http.ErrServerClosed) {
if eventBus != nil {
eventBus.Close()
}
log.Fatal().Msgf("ListenAndServe: %v", err)
}
case <-shutdownSignals:
if eventBus != nil {
eventBus.Close()
}
drainCtx, cancelDrain := context.WithTimeout(context.Background(), restartDrainTimeout)
defer cancelDrain()
if err := srv.Shutdown(drainCtx); err != nil {
log.Err(err).Msg("error draining HTTP server")
}
cancelDrain()
case <-restartCh:
if eventBus != nil {
eventBus.Close()
Expand Down
2 changes: 1 addition & 1 deletion cmd/cli/alert.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,12 @@ import (
"encoding/csv"
"encoding/json"
"fmt"
"github.com/jmpsec/osctrl/pkg/apiclient"
"os"
"strconv"
"strings"

"github.com/jmpsec/osctrl/pkg/alerts"
"github.com/jmpsec/osctrl/pkg/apiclient"
"github.com/olekukonko/tablewriter"
"github.com/urfave/cli/v3"
)
Expand Down
2 changes: 1 addition & 1 deletion cmd/cli/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,11 @@ package main
import (
"context"
"fmt"
"github.com/jmpsec/osctrl/pkg/apiclient"
"os"
"path/filepath"
"strconv"

"github.com/jmpsec/osctrl/pkg/apiclient"
"github.com/jmpsec/osctrl/pkg/auditlog"
"github.com/jmpsec/osctrl/pkg/backend"
"github.com/jmpsec/osctrl/pkg/carves"
Expand Down
2 changes: 1 addition & 1 deletion cmd/cli/shell_module_extras.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,12 @@ package main
import (
"encoding/json"
"fmt"
"github.com/jmpsec/osctrl/pkg/apiclient"
"sort"
"strconv"
"strings"
"time"

"github.com/jmpsec/osctrl/pkg/apiclient"
"github.com/jmpsec/osctrl/pkg/console"
"github.com/jmpsec/osctrl/pkg/fileexplorer"
"github.com/jmpsec/osctrl/pkg/posture"
Expand Down
2 changes: 1 addition & 1 deletion cmd/cli/shell_store.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,12 @@ import (
"bytes"
"encoding/json"
"fmt"
"github.com/jmpsec/osctrl/pkg/apiclient"
"path"
"strconv"
"strings"
"time"

"github.com/jmpsec/osctrl/pkg/apiclient"
"github.com/jmpsec/osctrl/pkg/auditlog"
"github.com/jmpsec/osctrl/pkg/carves"
"github.com/jmpsec/osctrl/pkg/config"
Expand Down
2 changes: 1 addition & 1 deletion cmd/cli/shell_store_extras.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,8 @@ package main

import (
"fmt"
"github.com/jmpsec/osctrl/pkg/apiclient"

"github.com/jmpsec/osctrl/pkg/apiclient"
"github.com/jmpsec/osctrl/pkg/console"
"github.com/jmpsec/osctrl/pkg/fileexplorer"
"github.com/jmpsec/osctrl/pkg/posture"
Expand Down
12 changes: 8 additions & 4 deletions cmd/tls/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -469,7 +469,6 @@ func osctrlService() {
if err != nil {
log.Fatal().Err(err).Msg("invalid events configuration")
}
defer eventBus.Close()
queriesmgr.Events = eventBus
filecarves.Events = eventBus
if alertsMgr != nil {
Expand Down Expand Up @@ -586,9 +585,8 @@ func osctrlService() {
if flagParams.TLS.Termination {
log.Info().Msg("TLS Termination is enabled")
srv.TLSConfig = &tls.Config{
MinVersion: tls.VersionTLS12,
CurvePreferences: []tls.CurveID{tls.CurveP521, tls.CurveP384, tls.CurveP256},
PreferServerCipherSuites: true,
MinVersion: tls.VersionTLS12,
CurvePreferences: []tls.CurveID{tls.CurveP521, tls.CurveP384, tls.CurveP256},
CipherSuites: []uint16{
tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
tls.TLS_ECDHE_RSA_WITH_AES_256_CBC_SHA,
Expand Down Expand Up @@ -616,12 +614,18 @@ func osctrlService() {
sinkStatsWriter.Stop()
stopAlerts(alertsRefreshStop, alertsInactiveStop, alertsRetentionStop, alertsWorker)
if err != nil && !errors.Is(err, http.ErrServerClosed) {
if eventBus != nil {
eventBus.Close()
}
log.Fatal().Msgf("ListenAndServe: %v", err)
}
case <-restartCh:
stopCommandWatcher()
sinkStatsWriter.Stop()
stopAlerts(alertsRefreshStop, alertsInactiveStop, alertsRetentionStop, alertsWorker)
if eventBus != nil {
eventBus.Close()
}
log.Info().Msg("TLS service command consumed — exiting for restart")
os.Exit(1)
}
Expand Down
6 changes: 3 additions & 3 deletions pkg/auth/oidc/provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -253,7 +253,7 @@ func (p *Provider) HandleCallback(parentCtx context.Context, r *http.Request, st
if err != nil {
// Wrap, don't merge — callers may want to log err
// server-side without exposing it to clients.
return auth.ResolvedIdentity{}, fmt.Errorf("%w: %v", ErrTokenExchange, err)
return auth.ResolvedIdentity{}, fmt.Errorf("%w: %w", ErrTokenExchange, err)
}

// (6) id_token present.
Expand All @@ -265,7 +265,7 @@ func (p *Provider) HandleCallback(parentCtx context.Context, r *http.Request, st
// (7) Verify signature + iss + aud + exp + nbf.
idToken, err := p.verifier.Verify(ctx, rawIDToken)
if err != nil {
return auth.ResolvedIdentity{}, fmt.Errorf("%w: %v", ErrIDTokenVerify, err)
return auth.ResolvedIdentity{}, fmt.Errorf("%w: %w", ErrIDTokenVerify, err)
}

// (8) Nonce match.
Expand All @@ -276,7 +276,7 @@ func (p *Provider) HandleCallback(parentCtx context.Context, r *http.Request, st
// Decode the claims we care about.
var claims idTokenClaims
if err := idToken.Claims(&claims); err != nil {
return auth.ResolvedIdentity{}, fmt.Errorf("%w: claims decode: %v", ErrIDTokenVerify, err)
return auth.ResolvedIdentity{}, fmt.Errorf("%w: claims decode: %w", ErrIDTokenVerify, err)
}
// Also decode into a generic map so pickUsername can look up
// custom claim names like "nickname" (Auth0) or "upn" (Entra)
Expand Down
4 changes: 2 additions & 2 deletions pkg/auth/saml/config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,12 +58,12 @@ func TestConfigValidate_Failures(t *testing.T) {
},
{
name: "no entity id",
mut: func(c *Config) { c.EntityID = "" },
mut: func(c *Config) { c.EntityID = "" },
want: "EntityID is required",
},
{
name: "no acs url",
mut: func(c *Config) { c.ACSURL = "" },
mut: func(c *Config) { c.ACSURL = "" },
want: "ACSURL is required",
},
{
Expand Down
7 changes: 4 additions & 3 deletions pkg/auth/saml/provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -288,7 +288,7 @@ func (p *Provider) loginURLAndRequestID(state auth.State) (string, string, error
// above. This is the security perimeter.
func (p *Provider) HandleCallback(_ context.Context, r *http.Request, state auth.State) (auth.ResolvedIdentity, error) {
if err := r.ParseForm(); err != nil {
return auth.ResolvedIdentity{}, fmt.Errorf("%w: parse form: %v", ErrParseResponse, err)
return auth.ResolvedIdentity{}, fmt.Errorf("%w: parse form: %w", ErrParseResponse, err)
}

// (1) SAMLResponse field must be present.
Expand Down Expand Up @@ -326,11 +326,12 @@ func (p *Provider) HandleCallback(_ context.Context, r *http.Request, state auth
// the PrivateErr at WARN so operators can diagnose IdP / cert /
// audience mismatches; the client still sees only the sentinel.
logEvent := log.Warn().Err(err)
if ire, ok := err.(*crewjam.InvalidResponseError); ok && ire != nil && ire.PrivateErr != nil {
var ire *crewjam.InvalidResponseError
if errors.As(err, &ire) && ire != nil && ire.PrivateErr != nil {
logEvent = logEvent.AnErr("private", ire.PrivateErr)
}
logEvent.Msg("saml: ParseResponse failed")
return auth.ResolvedIdentity{}, fmt.Errorf("%w: %v", ErrParseResponse, err)
return auth.ResolvedIdentity{}, fmt.Errorf("%w: %w", ErrParseResponse, err)
}

// (4) Replay defense — the assertion just passed signature +
Expand Down
7 changes: 4 additions & 3 deletions pkg/auth/saml/provider_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package saml

import (
"context"
"errors"
"net/http"
"net/http/httptest"
"net/url"
Expand Down Expand Up @@ -128,7 +129,7 @@ func TestHandleCallback_StateMismatch(t *testing.T) {

_, err := p.HandleCallback(context.Background(), r,
auth.State{EnvUUID: "global", Nonce: "n-the-real-nonce", OAuthState: "the-real-nonce"})
if err != ErrStateMismatch {
if !errors.Is(err, ErrStateMismatch) {
t.Errorf("expected ErrStateMismatch, got %v", err)
}
}
Expand All @@ -147,7 +148,7 @@ func TestHandleCallback_MissingSAMLResponse(t *testing.T) {

_, err := p.HandleCallback(context.Background(), r,
auth.State{EnvUUID: "global", Nonce: "n-the-real-nonce", OAuthState: "the-real-nonce"})
if err != ErrMissingSAMLResponse {
if !errors.Is(err, ErrMissingSAMLResponse) {
t.Errorf("expected ErrMissingSAMLResponse, got %v", err)
}
}
Expand Down Expand Up @@ -239,7 +240,7 @@ func errIs(got, target error) bool {
if got == nil {
return false
}
if got == target {
if errors.Is(got, target) {
return true
}
// fmt.Errorf("%w: ...", target) wraps the sentinel.
Expand Down
2 changes: 1 addition & 1 deletion pkg/cache/cache_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ func TestRedisConnectionErrorLeavesOtherErrorsAlone(t *testing.T) {
baseErr := errors.New("dial tcp 127.0.0.1:6379: connect: connection refused")

err := redisConnectionError(config.YAMLConfigurationRedis{}, baseErr)
if err != baseErr {
if !errors.Is(err, baseErr) {
t.Fatalf("error = %v, want original error", err)
}
}
Loading
Loading