This is an automated email from the ASF dual-hosted git repository. lahirujayathilake pushed a commit to branch auth-endpoints in repository https://gitbox.apache.org/repos/asf/airavata-custos.git
commit 6b8e7d66b5259fe8050a2207b01c2c49f6b67b56 Author: lahiruj <[email protected]> AuthorDate: Wed Jun 17 15:55:52 2026 -0400 Refresh JWKS on unknown signing key and tighten URL lock invariant --- internal/server/middleware/auth.go | 87 +++++++++++++++++++++++--- internal/server/middleware/auth_test.go | 107 ++++++++++++++++++++++++++++++++ 2 files changed, 186 insertions(+), 8 deletions(-) diff --git a/internal/server/middleware/auth.go b/internal/server/middleware/auth.go index 67165e653..065ec921d 100644 --- a/internal/server/middleware/auth.go +++ b/internal/server/middleware/auth.go @@ -33,6 +33,7 @@ import ( "time" "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jws" "github.com/lestrrat-go/jwx/v2/jwt" "golang.org/x/sync/singleflight" @@ -59,12 +60,18 @@ type Auth struct { httpClient *http.Client jwksCacheTTL time.Duration - mu sync.RWMutex - jwksURL string - cache *jwksEntry - sfGroup singleflight.Group + mu sync.RWMutex + jwksURL string + cache *jwksEntry + sfGroup singleflight.Group + lastForceRefresh time.Time } +// forceRefreshCooldown caps how often a kid-miss can trigger an unscheduled +// JWKS refresh. Bounds the work a malicious client can induce by sending +// tokens with random kid headers. +const forceRefreshCooldown = time.Minute + type jwksEntry struct { set jwk.Set fetchedAt time.Time @@ -169,6 +176,17 @@ func (a *Auth) verify(ctx context.Context, tokenString string) (*identity.Caller jwt.WithIssuer(a.issuer), jwt.WithAudience(a.audience), ) + if err != nil && a.shouldRefreshForKidMiss(tokenString, keys) { + fresh, refreshErr := a.refreshJWKS(ctx) + if refreshErr == nil { + tok, err = jwt.Parse([]byte(tokenString), + jwt.WithKeySet(fresh), + jwt.WithValidate(true), + jwt.WithIssuer(a.issuer), + jwt.WithAudience(a.audience), + ) + } + } if err != nil { return nil, errors.New(reject) } @@ -197,6 +215,54 @@ func (a *Auth) verify(ctx context.Context, tokenString string) (*identity.Caller }, nil } +// shouldRefreshForKidMiss reports whether a parse failure was caused by an +// unknown signing key, which happens after the IdP rotates and the cached set +// hasn't expired yet. Gated by a cooldown so a malicious client can't drive +// JWKS fetches by sending tokens with random kids. +func (a *Auth) shouldRefreshForKidMiss(tokenString string, current jwk.Set) bool { + parsed, err := jws.Parse([]byte(tokenString)) + if err != nil { + return false + } + sigs := parsed.Signatures() + if len(sigs) == 0 { + return false + } + kid := sigs[0].ProtectedHeaders().KeyID() + if kid == "" { + return false + } + if _, ok := current.LookupKeyID(kid); ok { + return false + } + a.mu.Lock() + defer a.mu.Unlock() + if time.Since(a.lastForceRefresh) < forceRefreshCooldown { + return false + } + a.lastForceRefresh = time.Now() + return true +} + +// refreshJWKS bypasses the TTL and fetches a fresh keyset. Single-flighted +// so concurrent kid-miss requests collapse to one fetch. +func (a *Auth) refreshJWKS(ctx context.Context) (jwk.Set, error) { + result, err, _ := a.sfGroup.Do("jwks-refresh", func() (interface{}, error) { + set, err := a.fetchJWKS(ctx) + if err != nil { + return nil, err + } + a.mu.Lock() + a.cache = &jwksEntry{set: set, fetchedAt: time.Now()} + a.mu.Unlock() + return set, nil + }) + if err != nil { + return nil, err + } + return result.(jwk.Set), nil +} + func (a *Auth) getJWKS(ctx context.Context) (jwk.Set, error) { a.mu.RLock() entry := a.cache @@ -234,16 +300,21 @@ func (a *Auth) fetchJWKS(ctx context.Context) (jwk.Set, error) { ctx, cancel := context.WithTimeout(ctx, a.httpClient.Timeout) defer cancel() - if a.jwksURL == "" { - url, err := a.discoverJWKSURL(ctx) + a.mu.RLock() + url := a.jwksURL + a.mu.RUnlock() + + if url == "" { + discovered, err := a.discoverJWKSURL(ctx) if err != nil { return nil, err } a.mu.Lock() - a.jwksURL = url + a.jwksURL = discovered + url = discovered a.mu.Unlock() } - return jwk.Fetch(ctx, a.jwksURL, jwk.WithHTTPClient(a.httpClient)) + return jwk.Fetch(ctx, url, jwk.WithHTTPClient(a.httpClient)) } func (a *Auth) discoverJWKSURL(ctx context.Context) (string, error) { diff --git a/internal/server/middleware/auth_test.go b/internal/server/middleware/auth_test.go index 993f7db5b..2341d156c 100644 --- a/internal/server/middleware/auth_test.go +++ b/internal/server/middleware/auth_test.go @@ -25,6 +25,7 @@ import ( "errors" "net/http" "net/http/httptest" + "sync" "testing" "time" @@ -361,3 +362,109 @@ func TestAuth_DiscoversJWKSURLFromIssuer(t *testing.T) { t.Fatalf("discovery path failed: code=%d called=%v", rec.Code, called) } } + +func newRSAKey(t *testing.T, kid string) (priv jwk.Key, pub jwk.Key) { + t.Helper() + raw, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("generate rsa: %v", err) + } + priv, err = jwk.FromRaw(raw) + if err != nil { + t.Fatalf("priv from raw: %v", err) + } + if err := priv.Set(jwk.KeyIDKey, kid); err != nil { + t.Fatalf("set kid: %v", err) + } + if err := priv.Set(jwk.AlgorithmKey, jwa.RS256); err != nil { + t.Fatalf("set alg: %v", err) + } + pub, err = priv.PublicKey() + if err != nil { + t.Fatalf("public from priv: %v", err) + } + return priv, pub +} + +func mintWith(t *testing.T, key jwk.Key, o tokenOpts) string { + t.Helper() + builder := jwt.NewBuilder(). + Issuer(o.issuer). + Audience([]string{o.audience}). + Subject(o.subject). + Expiration(o.expiresAt). + IssuedAt(time.Now()) + tok, err := builder.Build() + if err != nil { + t.Fatalf("build token: %v", err) + } + signed, err := jwt.Sign(tok, jwt.WithKey(jwa.RS256, key)) + if err != nil { + t.Fatalf("sign token: %v", err) + } + return string(signed) +} + +// TestAuth_ForceRefreshesOnKidMiss simulates IdP key rotation: the cache +// holds keyA but the next token is signed by keyB. The middleware must +// detect the kid miss, refresh the JWKS, and accept the token without +// waiting for the cache TTL. +func TestAuth_ForceRefreshesOnKidMiss(t *testing.T) { + privA, pubA := newRSAKey(t, "key-a") + privB, pubB := newRSAKey(t, "key-b") + + setA := jwk.NewSet() + if err := setA.AddKey(pubA); err != nil { + t.Fatalf("add A: %v", err) + } + setB := jwk.NewSet() + if err := setB.AddKey(pubB); err != nil { + t.Fatalf("add B: %v", err) + } + + var mu sync.Mutex + served := setA + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + mu.Lock() + current := served + mu.Unlock() + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(current) + })) + t.Cleanup(srv.Close) + + a, err := middleware.NewAuth( + config.AuthConfig{Issuer: testIssuer, Audience: testAudience}, + resolverAlways("user-1"), + middleware.WithJWKSURL(srv.URL), + ) + if err != nil { + t.Fatalf("new auth: %v", err) + } + + good := tokenOpts{issuer: testIssuer, audience: testAudience, subject: "sub-1", expiresAt: time.Now().Add(time.Minute)} + + // Warm the cache with a keyA-signed token; auth should accept it. + rec := httptest.NewRecorder() + a.Wrap(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})).ServeHTTP(rec, newRequest(mintWith(t, privA, good))) + if rec.Code != http.StatusOK { + t.Fatalf("warmup with key-a should succeed: code=%d body=%q", rec.Code, rec.Body.String()) + } + + // Rotate: server now serves only keyB. Cached set still has keyA. + mu.Lock() + served = setB + mu.Unlock() + + // keyB-signed token: middleware should kid-miss, force-refresh, retry. + rec = httptest.NewRecorder() + called := false + a.Wrap(http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { + called = true + })).ServeHTTP(rec, newRequest(mintWith(t, privB, good))) + + if rec.Code != http.StatusOK || !called { + t.Fatalf("kid-miss force-refresh did not recover: code=%d called=%v body=%q", + rec.Code, called, rec.Body.String()) + } +}
