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())
+       }
+}

Reply via email to