From b7f3a96e8ec8caea2639936428a247a873234d48 Mon Sep 17 00:00:00 2001 From: siddiqui irshad Date: Wed, 9 Sep 2026 12:16:04 +0530 Subject: [PATCH] Retry ID-token validation after OIDC refresh A successful refresh exchange that then fails ID-token verification currently discards the new tokens. Callers retry the exchange with the now-consumed rotating refresh token and the session dies with invalid_grant. Retry verification in place on transient JWKS or network failures. If keys still cannot be fetched, drop the unvalidated ID token so the new access and refresh tokens can be kept. Fixes #6194 Signed-off-by: siddiqui irshad Co-authored-by: Cursor --- pkg/authserver/upstream/oidc.go | 102 +++++++++++++++--- pkg/authserver/upstream/oidc_test.go | 149 ++++++++++++++++++++++++++- 2 files changed, 235 insertions(+), 16 deletions(-) diff --git a/pkg/authserver/upstream/oidc.go b/pkg/authserver/upstream/oidc.go index 38ab8157f0..ebcfed1189 100644 --- a/pkg/authserver/upstream/oidc.go +++ b/pkg/authserver/upstream/oidc.go @@ -8,10 +8,12 @@ import ( "errors" "fmt" "log/slog" + "net" "net/http" "net/url" "regexp" "slices" + "strings" "github.com/coreos/go-oidc/v3/oidc" "golang.org/x/oauth2" @@ -23,6 +25,12 @@ import ( const ( // ProviderTypeOIDC is for OIDC providers that support discovery. ProviderTypeOIDC ProviderType = "oidc" + + // idTokenValidationAttempts is how many times RefreshTokens will try to + // verify a newly issued ID token after a successful token-endpoint + // exchange. The exchange itself is not retried: rotating refresh tokens + // are single-use, so a second token request would send a consumed grant. + idTokenValidationAttempts = 3 ) // OIDCConfig contains configuration for OIDC providers that support discovery. @@ -422,6 +430,65 @@ func (p *OIDCProviderImpl) resolveSubject(token *oidc.IDToken) (string, error) { return value, nil } +// validateRefreshedIDToken verifies the ID token returned by a successful +// refresh exchange. Transient JWKS/network failures are retried in place so +// the already-consumed refresh token is not replayed. If every attempt is +// still a transient fetch failure, the unvalidated ID token is dropped and +// (nil, nil) is returned so the caller can keep the new access/refresh tokens +// (the storage layer already carries forward the previous ID token when the +// new one is empty). Permanent verification failures still fail closed. +func (p *OIDCProviderImpl) validateRefreshedIDToken(ctx context.Context, tokens *Tokens) (*oidc.IDToken, error) { + var lastErr error + for attempt := 1; attempt <= idTokenValidationAttempts; attempt++ { + token, err := p.validateIDToken(ctx, tokens.IDToken, "") + if err == nil { + return token, nil + } + lastErr = err + if !isTransientIDTokenValidationError(err) { + return nil, fmt.Errorf("ID token validation failed: %w", err) + } + slog.Warn("transient ID token validation failure after token refresh; retrying without re-exchanging", + "attempt", attempt, + "error", err, + ) + } + + // The token endpoint already succeeded and, for rotating refresh tokens, + // consumed the grant. Returning an error here would make the caller retry + // the exchange with the old token and get invalid_grant (#6194). Drop the + // unvalidated ID token rather than using it; access and refresh tokens + // still came from the trusted token endpoint. + slog.Warn("ID token validation failed after token refresh; dropping unvalidated ID token to preserve rotated refresh token", + "error", lastErr, + ) + tokens.IDToken = "" + return nil, nil +} + +// isTransientIDTokenValidationError reports whether ID-token verification +// failed because keys could not be fetched, not because the token itself was +// rejected. Permanent failures (bad signature, missing/mismatched nonce) must +// not be retried or bypassed. +func isTransientIDTokenValidationError(err error) bool { + if err == nil { + return false + } + if errors.Is(err, ErrNonceMismatch) || errors.Is(err, ErrNonceMissing) || errors.Is(err, ErrSubjectMismatch) { + return false + } + if errors.Is(err, context.DeadlineExceeded) { + return true + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + // go-oidc wraps a JWKS HTTP failure as "fetching keys oidc: get keys failed". + msg := err.Error() + return strings.Contains(msg, "fetching keys") || strings.Contains(msg, "get keys failed") +} + // validateIDToken validates an ID token and returns the parsed token. func (p *OIDCProviderImpl) validateIDToken(ctx context.Context, idToken, nonce string) (*oidc.IDToken, error) { if p.verifier == nil { @@ -510,6 +577,11 @@ func (p *OIDCProviderImpl) buildOIDCParams() map[string]string { // RefreshTokens refreshes the upstream IDP tokens. // This overrides the base implementation to add OIDC-specific ID token validation. +// After a successful token-endpoint exchange, ID-token verification is retried +// on transient JWKS/network failures so a rotating refresh token is not +// replayed. If verification still cannot fetch keys, the unvalidated ID token +// is dropped and the new access/refresh tokens are returned; the storage layer +// carries forward the previous ID token when the new one is empty. func (p *OIDCProviderImpl) RefreshTokens(ctx context.Context, refreshToken, expectedSubject string) (*Tokens, error) { if p.endpoints == nil { return nil, errors.New("OIDC endpoints not discovered") @@ -533,22 +605,24 @@ func (p *OIDCProviderImpl) RefreshTokens(ctx context.Context, refreshToken, expe // authorization request exists to provide an expected nonce value. // Full nonce validation occurs in ExchangeCodeForIdentity during the initial auth flow. if tokens.IDToken != "" && p.verifier != nil { - token, err := p.validateIDToken(ctx, tokens.IDToken, "") - if err != nil { - return nil, fmt.Errorf("ID token validation failed: %w", err) - } - // The stored expectedSubject is the resolved subject (SubjectClaim, or - // "sub" by default). Resolve the refreshed token through the same path - // before comparing — comparing the raw "sub" would wrongly reject a - // refresh whenever a non-"sub" SubjectClaim is configured. OIDC Core - // Section 12.2 still holds: for the default "sub" this is identical to - // the original, and a custom claim must likewise be identical. - refreshedSubject, err := p.resolveSubject(token) + token, err := p.validateRefreshedIDToken(ctx, tokens) if err != nil { - return nil, fmt.Errorf("%w: %w", ErrIdentityResolutionFailed, err) + return nil, err } - if expectedSubject != "" && refreshedSubject != expectedSubject { - return nil, ErrSubjectMismatch + if token != nil { + // The stored expectedSubject is the resolved subject (SubjectClaim, or + // "sub" by default). Resolve the refreshed token through the same path + // before comparing — comparing the raw "sub" would wrongly reject a + // refresh whenever a non-"sub" SubjectClaim is configured. OIDC Core + // Section 12.2 still holds: for the default "sub" this is identical to + // the original, and a custom claim must likewise be identical. + refreshedSubject, err := p.resolveSubject(token) + if err != nil { + return nil, fmt.Errorf("%w: %w", ErrIdentityResolutionFailed, err) + } + if expectedSubject != "" && refreshedSubject != expectedSubject { + return nil, ErrSubjectMismatch + } } } diff --git a/pkg/authserver/upstream/oidc_test.go b/pkg/authserver/upstream/oidc_test.go index 2960665221..71cec45d48 100644 --- a/pkg/authserver/upstream/oidc_test.go +++ b/pkg/authserver/upstream/oidc_test.go @@ -9,6 +9,8 @@ import ( "crypto/rsa" "encoding/base64" "encoding/json" + "errors" + "fmt" "math/big" "net" "net/http" @@ -44,6 +46,7 @@ type mockOIDCServer struct { privateKey *rsa.PrivateKey keyID string tokenHandler func(w http.ResponseWriter, r *http.Request) + jwksHandler func(w http.ResponseWriter, r *http.Request) } func newMockOIDCServer(t *testing.T) *mockOIDCServer { @@ -131,8 +134,15 @@ func (*mockOIDCServer) handleUserInfo(w http.ResponseWriter, r *http.Request) { } } -func (m *mockOIDCServer) handleJWKS(w http.ResponseWriter, _ *http.Request) { - // Return JWKS with public key +func (m *mockOIDCServer) handleJWKS(w http.ResponseWriter, r *http.Request) { + if m.jwksHandler != nil { + m.jwksHandler(w, r) + return + } + m.writeJWKS(w) +} + +func (m *mockOIDCServer) writeJWKS(w http.ResponseWriter) { jwks := map[string]any{ "keys": []map[string]any{ { @@ -1527,6 +1537,141 @@ func TestOIDCProvider_RefreshTokens(t *testing.T) { require.NoError(t, err) assert.Equal(t, "refreshed-access-token", tokens.AccessToken) }) + + t.Run("transient JWKS failure is retried without a second token exchange", func(t *testing.T) { + t.Parallel() + + mock := newMockOIDCServer(t) + t.Cleanup(mock.Close) + + var tokenCalls, jwksCalls atomic.Int32 + mock.tokenHandler = func(w http.ResponseWriter, _ *http.Request) { + tokenCalls.Add(1) + writeRefreshIDTokenResponse(w, mock.signIDToken(testClientID, "user-123", "", time.Now().Add(time.Hour))) + } + mock.jwksHandler = func(w http.ResponseWriter, _ *http.Request) { + if jwksCalls.Add(1) == 1 { + http.Error(w, "temporarily unavailable", http.StatusInternalServerError) + return + } + mock.writeJWKS(w) + } + + provider := mustNewOIDCProvider(t, ctx, mock.issuer) + tokens, err := provider.RefreshTokens(ctx, "old-refresh-token", "user-123") + require.NoError(t, err) + assert.Equal(t, "refreshed-access-token", tokens.AccessToken) + assert.Equal(t, "new-refresh-token", tokens.RefreshToken) + assert.NotEmpty(t, tokens.IDToken) + assert.Equal(t, int32(1), tokenCalls.Load(), "token endpoint must not be retried after a successful exchange") + assert.GreaterOrEqual(t, jwksCalls.Load(), int32(2), "JWKS fetch should be retried after the first failure") + }) + + t.Run("persistent JWKS failure drops ID token and keeps rotated refresh token", func(t *testing.T) { + t.Parallel() + + mock := newMockOIDCServer(t) + t.Cleanup(mock.Close) + + var tokenCalls, jwksCalls atomic.Int32 + mock.tokenHandler = func(w http.ResponseWriter, _ *http.Request) { + tokenCalls.Add(1) + writeRefreshIDTokenResponse(w, mock.signIDToken(testClientID, "user-123", "", time.Now().Add(time.Hour))) + } + mock.jwksHandler = func(w http.ResponseWriter, _ *http.Request) { + jwksCalls.Add(1) + http.Error(w, "temporarily unavailable", http.StatusInternalServerError) + } + + provider := mustNewOIDCProvider(t, ctx, mock.issuer) + tokens, err := provider.RefreshTokens(ctx, "old-refresh-token", "user-123") + require.NoError(t, err) + assert.Equal(t, "refreshed-access-token", tokens.AccessToken) + assert.Equal(t, "new-refresh-token", tokens.RefreshToken) + assert.Empty(t, tokens.IDToken, "unvalidated ID token must not be returned") + assert.Equal(t, int32(1), tokenCalls.Load(), "token endpoint must not be retried") + assert.Equal(t, int32(idTokenValidationAttempts), jwksCalls.Load()) + }) + + t.Run("signature failure after successful JWKS fetch is not bypassed", func(t *testing.T) { + t.Parallel() + + mock := newMockOIDCServer(t) + t.Cleanup(mock.Close) + + otherKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + var tokenCalls atomic.Int32 + mock.tokenHandler = func(w http.ResponseWriter, _ *http.Request) { + tokenCalls.Add(1) + orig := mock.privateKey + mock.privateKey = otherKey + idToken := mock.signIDToken(testClientID, "user-123", "", time.Now().Add(time.Hour)) + mock.privateKey = orig + writeRefreshIDTokenResponse(w, idToken) + } + + provider := mustNewOIDCProvider(t, ctx, mock.issuer) + _, err = provider.RefreshTokens(ctx, "old-refresh-token", "user-123") + require.Error(t, err) + assert.Contains(t, err.Error(), "ID token validation failed") + assert.NotErrorIs(t, err, ErrSubjectMismatch) + assert.Equal(t, int32(1), tokenCalls.Load(), "token endpoint must not be retried") + }) +} + +func writeRefreshIDTokenResponse(w http.ResponseWriter, idToken string) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(testTokenResponse{ + AccessToken: "refreshed-access-token", + TokenType: "Bearer", + RefreshToken: "new-refresh-token", + ExpiresIn: 3600, + IDToken: idToken, + }) +} + +func mustNewOIDCProvider(t *testing.T, ctx context.Context, issuer string) *OIDCProviderImpl { + t.Helper() + provider, err := NewOIDCProvider(ctx, &OIDCConfig{ + CommonOAuthConfig: CommonOAuthConfig{ + ClientID: testClientID, + ClientSecret: testClientSecret, + RedirectURI: testRedirectURI, + }, + Issuer: issuer, + }) + require.NoError(t, err) + return provider +} + +func TestIsTransientIDTokenValidationError(t *testing.T) { + t.Parallel() + + timeoutErr := &net.DNSError{Err: "i/o timeout", Name: "idp.example", IsTimeout: true} + + tests := []struct { + name string + err error + want bool + }{ + {name: "nil", err: nil, want: false}, + {name: "nonce mismatch", err: ErrNonceMismatch, want: false}, + {name: "nonce missing", err: ErrNonceMissing, want: false}, + {name: "subject mismatch", err: ErrSubjectMismatch, want: false}, + {name: "deadline exceeded", err: context.DeadlineExceeded, want: true}, + {name: "net timeout", err: timeoutErr, want: true}, + {name: "wrapped go-oidc JWKS fetch", err: fmt.Errorf("failed to verify ID token: %w", errors.New("fetching keys oidc: get keys failed")), want: true}, + {name: "bad signature", err: errors.New("failed to verify signature"), want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, tt.want, isTransientIDTokenValidationError(tt.err)) + }) + } } // TestNewOIDCProvider_DrainsOwnClientOnFailure pins that a construction that