Skip to content
Draft
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
102 changes: 88 additions & 14 deletions pkg/authserver/upstream/oidc.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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.
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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")
Expand All @@ -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
}
}
}

Expand Down
149 changes: 147 additions & 2 deletions pkg/authserver/upstream/oidc_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
"crypto/rsa"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"math/big"
"net"
"net/http"
Expand Down Expand Up @@ -44,6 +46,7 @@
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 {
Expand Down Expand Up @@ -131,8 +134,15 @@
}
}

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{
{
Expand Down Expand Up @@ -1527,6 +1537,141 @@
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 {

Check failure on line 1635 in pkg/authserver/upstream/oidc_test.go

View workflow job for this annotation

GitHub Actions / Linting / Lint Go Code

context-as-argument: context.Context should be the first parameter of a function (revive)
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
Expand Down
Loading