From 75eb8b831cbb8068272d7cd413c69575a77a2d61 Mon Sep 17 00:00:00 2001 From: Mathias Date: Mon, 6 Jul 2026 23:14:56 +0200 Subject: [PATCH] =?UTF-8?q?feat(auth):=20multi-issuer=20JWT=20validation?= =?UTF-8?q?=20=E2=80=94=20accept=20k8s=20SA=20tokens=20(infra=20ADR-0011)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit JWTValidator trusted a single issuer (iss matched exactly against DEX_ISSUER_URL), which is why every in-cluster caller fell back to a static bearer. Add support for a LIST of trusted issuers so one MCP server can accept Authentik JWTs (interactive/web) AND k3s cluster ServiceAccount tokens (in-cluster, audience- bound, kubelet-rotated) at once — the enabler for ADR-0011. - Add IssuerConfig{IssuerURL, Audience} + NewMultiJWTValidator([]IssuerConfig). - Refactor JWTValidator to hold []issuerEntry + a shared jwk.Cache; Validate tries each trusted issuer, returns the subject on first success. All-issuers- JWKS-unreachable -> ErrUnavailable (503); any definitive reject -> 401. - NewJWTValidator (single issuer) and BearerMiddleware signatures UNCHANGED — existing consumers (gitea-mcp, ingestion) compile and behave identically. - Add auth/jwt_test.go: an in-process OIDC issuer harness (discovery + JWKS + token minting) — the chassis previously had NO happy-path JWT test. Covers accept-either-trusted-issuer, reject-untrusted (401 not 503), per-issuer audience enforcement, single-issuer backward compat. Proven viable by the ADR-0011 live-cluster spike (k3s OIDC/JWKS validates a projected SA token). Follow-up: wire gitea-mcp/brain-mcp/ingestion to also trust the k3s issuer + mount an audience-scoped projected SA token on one pod. Refs infra ADR-0011; supersedes the sidecar in infra#183. (Tracked as hyperguild#77, which was misfiled — the chassis is this repo, not hyperguild.) Co-Authored-By: Claude Opus 4.8 (1M context) --- auth/bearer_test.go | 9 +-- auth/jwt.go | 172 +++++++++++++++++++++++++++++--------------- auth/jwt_test.go | 143 ++++++++++++++++++++++++++++++++++++ 3 files changed, 261 insertions(+), 63 deletions(-) create mode 100644 auth/jwt_test.go diff --git a/auth/bearer_test.go b/auth/bearer_test.go index d13a0dc..cd1dca2 100644 --- a/auth/bearer_test.go +++ b/auth/bearer_test.go @@ -18,10 +18,11 @@ import ( func unavailableValidator(t *testing.T) *JWTValidator { t.Helper() return &JWTValidator{ - issuer: "https://dex.example", - jwksURI: "http://127.0.0.1:0/jwks", // unregistered in cache → Get fails - cache: jwk.NewCache(context.Background()), - audience: "", + entries: []issuerEntry{{ + issuer: "https://dex.example", + jwksURI: "http://127.0.0.1:0/jwks", // unregistered in cache → Get fails + }}, + cache: jwk.NewCache(context.Background()), } } diff --git a/auth/jwt.go b/auth/jwt.go index c3af027..f7a9724 100644 --- a/auth/jwt.go +++ b/auth/jwt.go @@ -1,12 +1,16 @@ -// Package auth provides the Dex-JWT + static-Bearer authentication primitives -// shared by every Mathias-owned MCP server (gitea-mcp, brain-mcp / ingestion, -// future MCPs spawned from template-go-agent). +// Package auth provides the Dex/Authentik-JWT + static-Bearer authentication +// primitives shared by every Mathias-owned MCP server (gitea-mcp, brain-mcp / +// ingestion, future MCPs spawned from template-go-agent). // // Replaces ~80 LOC of near-identical jwt.go in each consumer, ~50 LOC of // Bearer middleware, and ~25 LOC of RFC 9728 protected-resource metadata // handler. See `gitea.d-ma.be/mathias/infra` docs/superpowers/handoffs/ // 2026-05-22-mcp-chassis-spike.md for the design rationale and the // abort-criterion check. +// +// A validator trusts a LIST of OIDC issuers (infra ADR-0011): one MCP server can +// accept Authentik JWTs (interactive/web) AND k3s cluster ServiceAccount tokens +// (in-cluster, audience-bound, kubelet-rotated) at once. package auth import ( @@ -22,105 +26,155 @@ import ( ) // ErrUnavailable indicates JWT validation could not be COMPLETED because the -// JWKS / Dex endpoint was unreachable — a transient condition — as opposed to +// JWKS / issuer endpoint was unreachable — a transient condition — as opposed to // the token being present and invalid. Callers (e.g. BearerMiddleware) map it -// to HTTP 503 temporarily_unavailable rather than a generic 401, so a Dex +// to HTTP 503 temporarily_unavailable rather than a generic 401, so an issuer // outage is distinguishable from a bad token (gitea-mcp#6). var ErrUnavailable = errors.New("jwt validation temporarily unavailable") -// JWTValidator validates Bearer JWTs issued by a Dex (OIDC) authorization server. -// Audience is optional; leave empty to skip audience validation. -// -// A nil *JWTValidator behaves as "JWT auth disabled" — Validate returns an -// error without panicking. Callers can construct one validator at startup -// keyed on whether DEX_ISSUER_URL is set, and pass nil through the rest of -// the codebase without further branching. -type JWTValidator struct { +// IssuerConfig names one trusted OIDC issuer and its (optional) required +// audience. Trusting a list of these is what lets a single MCP server accept +// tokens from Authentik (D3/D4) and the k3s cluster OIDC issuer (D1 SA tokens) +// simultaneously — infra ADR-0011 Decision 5. +type IssuerConfig struct { + IssuerURL string + Audience string // "" = skip audience validation for this issuer +} + +// issuerEntry is a resolved IssuerConfig: OIDC discovery has run and jwks_uri is +// registered in the shared cache. +type issuerEntry struct { issuer string audience string jwksURI string - cache *jwk.Cache } -// NewJWTValidator fetches the OIDC discovery document from issuerURL, -// extracts jwks_uri, warms the JWKS cache, and returns a ready validator. -// Empty issuerURL returns (nil, nil) so callers can use a single -// constructor regardless of whether Dex is configured. +// JWTValidator validates Bearer JWTs against one or more trusted OIDC issuers. +// +// A nil *JWTValidator behaves as "JWT auth disabled" — Validate returns an +// error without panicking. Callers can construct one validator at startup keyed +// on whether any issuer is configured, and pass nil through the rest of the +// codebase without further branching. +type JWTValidator struct { + entries []issuerEntry + cache *jwk.Cache +} + +// NewJWTValidator builds a single-issuer validator — the common case, and +// backward compatible with every existing caller. Empty issuerURL returns +// (nil, nil) so callers can use one constructor regardless of whether an issuer +// is configured. func NewJWTValidator(ctx context.Context, issuerURL, audience string) (*JWTValidator, error) { if issuerURL == "" { return nil, nil } + return NewMultiJWTValidator(ctx, []IssuerConfig{{IssuerURL: issuerURL, Audience: audience}}) +} +// NewMultiJWTValidator builds a validator that trusts every issuer in the list. +// For each, it fetches the OIDC discovery document, registers jwks_uri in a +// shared cache, and warms it. Entries with an empty IssuerURL are skipped; an +// empty/nil resulting set returns (nil, nil) — "JWT auth disabled". +func NewMultiJWTValidator(ctx context.Context, issuers []IssuerConfig) (*JWTValidator, error) { + cache := jwk.NewCache(ctx) + var entries []issuerEntry + for _, ic := range issuers { + if ic.IssuerURL == "" { + continue + } + jwksURI, err := discoverJWKSURI(ctx, ic.IssuerURL) + if err != nil { + return nil, fmt.Errorf("issuer %s: %w", ic.IssuerURL, err) + } + if err := cache.Register(jwksURI, jwk.WithMinRefreshInterval(time.Hour)); err != nil { + return nil, fmt.Errorf("register jwks cache (%s): %w", ic.IssuerURL, err) + } + if _, err := cache.Refresh(ctx, jwksURI); err != nil { + return nil, fmt.Errorf("initial jwks fetch (%s): %w", ic.IssuerURL, err) + } + entries = append(entries, issuerEntry{issuer: ic.IssuerURL, audience: ic.Audience, jwksURI: jwksURI}) + } + if len(entries) == 0 { + return nil, nil + } + return &JWTValidator{entries: entries, cache: cache}, nil +} + +// discoverJWKSURI fetches the OIDC discovery document from issuerURL and returns +// its jwks_uri. +func discoverJWKSURI(ctx context.Context, issuerURL string) (string, error) { req, err := http.NewRequestWithContext(ctx, http.MethodGet, issuerURL+"/.well-known/openid-configuration", nil) if err != nil { - return nil, fmt.Errorf("build oidc discovery request: %w", err) + return "", fmt.Errorf("build oidc discovery request: %w", err) } resp, err := http.DefaultClient.Do(req) if err != nil { - return nil, fmt.Errorf("fetch oidc discovery: %w", err) + return "", fmt.Errorf("fetch oidc discovery: %w", err) } defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("oidc discovery: status %d", resp.StatusCode) + return "", fmt.Errorf("oidc discovery: status %d", resp.StatusCode) } var doc struct { JWKSURI string `json:"jwks_uri"` } if err := json.NewDecoder(resp.Body).Decode(&doc); err != nil { - return nil, fmt.Errorf("decode oidc discovery: %w", err) + return "", fmt.Errorf("decode oidc discovery: %w", err) } if doc.JWKSURI == "" { - return nil, fmt.Errorf("oidc discovery: empty jwks_uri") + return "", fmt.Errorf("oidc discovery: empty jwks_uri") } - - cache := jwk.NewCache(ctx) - if err := cache.Register(doc.JWKSURI, jwk.WithMinRefreshInterval(time.Hour)); err != nil { - return nil, fmt.Errorf("register jwks cache: %w", err) - } - if _, err := cache.Refresh(ctx, doc.JWKSURI); err != nil { - return nil, fmt.Errorf("initial jwks fetch: %w", err) - } - - return &JWTValidator{ - issuer: issuerURL, - audience: audience, - jwksURI: doc.JWKSURI, - cache: cache, - }, nil + return doc.JWKSURI, nil } -// Validate parses and validates rawToken against the OIDC issuer. Returns -// the subject claim on success. A nil receiver returns -// (`""`, errDisabled) so callers can dispatch on `err != nil` without -// nil-checks at every call site. +// Validate parses rawToken against each trusted issuer and returns the subject +// claim on the first success. If no issuer accepts it: when EVERY issuer's JWKS +// was unreachable the error wraps ErrUnavailable (caller answers 503); otherwise +// at least one issuer reached a signature/claims verdict and rejected it, so the +// error is a definitive rejection (401). A nil receiver returns errDisabled. func (v *JWTValidator) Validate(ctx context.Context, rawToken string) (string, error) { if v == nil { return "", errDisabled } - keySet, err := v.cache.Get(ctx, v.jwksURI) - if err != nil { - // The JWKS could not be fetched — Dex/JWKS is unreachable, not a bad - // token. Tag it ErrUnavailable so the caller can answer 503, not 401. - return "", fmt.Errorf("%w: get jwks: %v", ErrUnavailable, err) + var lastErr error + sawDecision := false // at least one issuer reached a real verify verdict + for _, e := range v.entries { + keySet, err := v.cache.Get(ctx, e.jwksURI) + if err != nil { + // This issuer's JWKS is unreachable — transient. Try the next; only + // if ALL issuers are unreachable do we surface ErrUnavailable. + lastErr = fmt.Errorf("%w: get jwks (%s): %v", ErrUnavailable, e.issuer, err) + continue + } + + opts := []jwt.ParseOption{ + jwt.WithKeySet(keySet), + jwt.WithValidate(true), + jwt.WithIssuer(e.issuer), + } + if e.audience != "" { + opts = append(opts, jwt.WithAudience(e.audience)) + } + + tok, err := jwt.ParseString(rawToken, opts...) + if err == nil { + return tok.Subject(), nil + } + sawDecision = true + lastErr = fmt.Errorf("validate jwt: %w", err) } - opts := []jwt.ParseOption{ - jwt.WithKeySet(keySet), - jwt.WithValidate(true), - jwt.WithIssuer(v.issuer), + if lastErr == nil { + lastErr = errDisabled } - if v.audience != "" { - opts = append(opts, jwt.WithAudience(v.audience)) + if !sawDecision { + // Every issuer's JWKS was unreachable — lastErr already wraps ErrUnavailable. + return "", lastErr } - - tok, err := jwt.ParseString(rawToken, opts...) - if err != nil { - return "", fmt.Errorf("validate jwt: %w", err) - } - return tok.Subject(), nil + return "", lastErr } // errDisabled is the sentinel returned by Validate on a nil receiver. diff --git a/auth/jwt_test.go b/auth/jwt_test.go new file mode 100644 index 0000000..7bb1256 --- /dev/null +++ b/auth/jwt_test.go @@ -0,0 +1,143 @@ +package auth + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/lestrrat-go/jwx/v2/jwa" + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jwt" + "github.com/stretchr/testify/require" +) + +// testIssuer is an in-process OIDC issuer: it serves an openid-configuration +// discovery doc + a JWKS, and mints signed JWTs — enough to exercise the real +// signature/issuer/audience validation path (which the chassis previously had +// no happy-path test for). Used to prove multi-issuer validation (ADR-0011). +type testIssuer struct { + url string + priv jwk.Key +} + +func newTestIssuer(t *testing.T) *testIssuer { + t.Helper() + raw, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + priv, err := jwk.FromRaw(raw) + require.NoError(t, err) + require.NoError(t, priv.Set(jwk.KeyIDKey, "test-kid")) + require.NoError(t, priv.Set(jwk.AlgorithmKey, jwa.RS256)) + pub, err := priv.PublicKey() + require.NoError(t, err) + set := jwk.NewSet() + require.NoError(t, set.AddKey(pub)) + + ti := &testIssuer{priv: priv} + mux := http.NewServeMux() + mux.HandleFunc("/jwks", func(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(set) + }) + mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, _ *http.Request) { + _ = json.NewEncoder(w).Encode(map[string]string{ + "issuer": ti.url, + "jwks_uri": ti.url + "/jwks", + }) + }) + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + ti.url = srv.URL + return ti +} + +func (ti *testIssuer) mint(t *testing.T, aud, sub string) string { + t.Helper() + tok, err := jwt.NewBuilder(). + Issuer(ti.url). + Subject(sub). + Audience([]string{aud}). + IssuedAt(time.Now()). + Expiration(time.Now().Add(time.Hour)). + Build() + require.NoError(t, err) + signed, err := jwt.Sign(tok, jwt.WithKey(jwa.RS256, ti.priv)) + require.NoError(t, err) + return string(signed) +} + +// The core ADR-0011 claim: one validator, several trusted issuers (e.g. +// Authentik + the k3s cluster OIDC), a token from any trusted issuer validates. +func TestNewMultiJWTValidator_AcceptsEitherTrustedIssuer(t *testing.T) { + ctx := context.Background() + a := newTestIssuer(t) + b := newTestIssuer(t) + + v, err := NewMultiJWTValidator(ctx, []IssuerConfig{ + {IssuerURL: a.url, Audience: "brain-mcp"}, + {IssuerURL: b.url, Audience: "brain-mcp"}, + }) + require.NoError(t, err) + require.NotNil(t, v) + + subA, err := v.Validate(ctx, a.mint(t, "brain-mcp", "alice")) + require.NoError(t, err) + require.Equal(t, "alice", subA) + + subB, err := v.Validate(ctx, b.mint(t, "brain-mcp", "bob")) + require.NoError(t, err) + require.Equal(t, "bob", subB) +} + +// A token from an issuer NOT in the trusted list is rejected — definitively +// (401), not as a transient JWKS outage (503). +func TestNewMultiJWTValidator_RejectsUntrustedIssuer(t *testing.T) { + ctx := context.Background() + trusted := newTestIssuer(t) + untrusted := newTestIssuer(t) + + v, err := NewMultiJWTValidator(ctx, []IssuerConfig{{IssuerURL: trusted.url, Audience: "brain-mcp"}}) + require.NoError(t, err) + + _, err = v.Validate(ctx, untrusted.mint(t, "brain-mcp", "eve")) + require.Error(t, err) + require.NotErrorIs(t, err, ErrUnavailable) +} + +// Audience is enforced per issuer — the k8s-SA-token replay guard from the ADR +// spike. A token minted for a different audience is rejected. +func TestNewMultiJWTValidator_EnforcesPerIssuerAudience(t *testing.T) { + ctx := context.Background() + a := newTestIssuer(t) + + v, err := NewMultiJWTValidator(ctx, []IssuerConfig{{IssuerURL: a.url, Audience: "brain-mcp"}}) + require.NoError(t, err) + + _, err = v.Validate(ctx, a.mint(t, "some-other-service", "alice")) + require.Error(t, err) +} + +// Backward compatibility: the existing single-issuer constructor still works and +// validates a real signed token end-to-end. +func TestNewJWTValidator_SingleIssuer_BackwardCompatible(t *testing.T) { + ctx := context.Background() + a := newTestIssuer(t) + + v, err := NewJWTValidator(ctx, a.url, "brain-mcp") + require.NoError(t, err) + require.NotNil(t, v) + + sub, err := v.Validate(ctx, a.mint(t, "brain-mcp", "carol")) + require.NoError(t, err) + require.Equal(t, "carol", sub) +} + +func TestNewMultiJWTValidator_EmptyList_ReturnsNilNil(t *testing.T) { + v, err := NewMultiJWTValidator(context.Background(), nil) + require.NoError(t, err) + require.Nil(t, v) +}