package oidc import ( "context" "crypto/tls" "errors" "fmt" "log/slog" "net/http" "net/url" "strings" "time" "github.com/jwx-go/jwkfetch/v4" "github.com/lestrrat-go/httprc/v3" "github.com/lestrrat-go/httprc/v3/errsink" "github.com/lestrrat-go/jwx/v4/jwk" "github.com/lestrrat-go/jwx/v4/jws" "github.com/lestrrat-go/jwx/v4/jwt" "github.com/ory/fosite" "github.com/pocket-id/pocket-id/backend/internal/model" jwkutils "github.com/pocket-id/pocket-id/backend/internal/utils/jwk" ) const clientAssertionTypeJWTBearer = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" // #nosec G101 -- OAuth assertion type identifier, not a credential var errNoFederatedClientAssertion = errors.New("no federated client assertion") // federatedClientStore is the subset of the store the federated authenticator needs. type federatedClientStore interface { GetClient(ctx context.Context, id string) (fosite.Client, error) ClientAssertionJWTValid(ctx context.Context, jti string) error SetClientAssertionJWT(ctx context.Context, jti string, exp time.Time) error } // federatedClientAuthenticator authenticates clients via JWT bearer assertions issued // by a federated identity provider configured per client. type federatedClientAuthenticator struct { clients federatedClientStore httpClient *http.Client jwksCache *jwkfetch.Cache defaultAudience string } func newFederatedClientAuthenticator(ctx context.Context, clients federatedClientStore, httpClient *http.Client, defaultAudience string) (*federatedClientAuthenticator, error) { authenticator := &federatedClientAuthenticator{ clients: clients, httpClient: httpClient, defaultAudience: defaultAudience, } jwksCache, err := authenticator.getJWKCache(ctx) if err != nil { return nil, err } authenticator.jwksCache = jwksCache return authenticator, nil } func (a *federatedClientAuthenticator) getJWKCache(ctx context.Context) (*jwkfetch.Cache, error) { // We need to create a custom HTTP client to set a timeout. client := a.httpClient if client == nil { client = &http.Client{ Timeout: 10 * time.Second, } defaultTransport, ok := http.DefaultTransport.(*http.Transport) if !ok { // Indicates a development-time error panic("Default transport is not of type *http.Transport") } transport := defaultTransport.Clone() transport.TLSClientConfig.MinVersion = tls.VersionTLS12 client.Transport = transport } return jwkfetch.NewCache(ctx, httprc.NewClient( httprc.WithErrorSink(errsink.NewSlog(slog.Default())), httprc.WithHTTPClient(client), ), jwkfetch.WithHTTPClient(client), ) } // newClientAuthenticationStrategy accepts federated client assertions before falling // back to fosite's default client authentication. func newClientAuthenticationStrategy(authenticator *federatedClientAuthenticator, provider *fosite.Fosite) fosite.ClientAuthenticationStrategy { return func(ctx context.Context, r *http.Request, form url.Values) (fosite.Client, error) { client, err := authenticator.authenticateForm(ctx, form) if err == nil { return client, nil } if !errors.Is(err, errNoFederatedClientAssertion) { return nil, err } return provider.DefaultClientAuthenticationStrategy(ctx, r, form) } } // authenticateForm returns errNoFederatedClientAssertion when the form carries no // federated assertion, so the caller can fall back to other authentication methods. func (a *federatedClientAuthenticator) authenticateForm(ctx context.Context, form url.Values) (fosite.Client, error) { if form.Get("client_assertion_type") != clientAssertionTypeJWTBearer || form.Get("client_assertion") == "" { return nil, errNoFederatedClientAssertion } return a.authenticateAssertion(ctx, form.Get("client_assertion"), form.Get("client_id")) } // authenticateAssertion validates the assertion JWT against the client's configured // federated identity. An empty clientID falls back to the assertion's subject. func (a *federatedClientAuthenticator) authenticateAssertion(ctx context.Context, assertion string, clientID string) (fosite.Client, error) { rawAssertion := []byte(assertion) insecureToken, err := jwt.ParseInsecure(rawAssertion) if err != nil { return nil, fosite.ErrInvalidClient.WithHint("Invalid client assertion.").WithWrap(err) } issuer, _ := insecureToken.Issuer() if issuer == "" { return nil, fosite.ErrInvalidClient.WithHint("Client assertion is missing issuer.") } if clientID == "" { clientID, _ = insecureToken.Subject() } if clientID == "" { return nil, fosite.ErrInvalidClient.WithHint("Client assertion is missing subject.") } client, err := a.clients.GetClient(ctx, clientID) if err != nil { return nil, fosite.ErrInvalidClient.WithWrap(err) } oidcClient, ok := client.(Client) if !ok { return nil, errNoFederatedClientAssertion } federatedIdentity, ok := oidcClient.Credentials.FederatedIdentityForIssuer(issuer) if !ok { return nil, errNoFederatedClientAssertion } jwks, err := a.keySetForIdentity(ctx, federatedIdentity) if err != nil { return nil, err } audience := federatedIdentity.Audience if audience == "" { audience = a.defaultAudience } subject := federatedIdentity.Subject if subject == "" { subject = client.GetID() } parsed, err := jwt.Parse(rawAssertion, jwt.WithValidate(true), jwt.WithAcceptableSkew(30*time.Second), jwt.WithRequiredClaim(jwt.ExpirationKey), jwt.WithIssuer(issuer), jwt.WithSubject(subject), jwt.WithAudience(audience), jwt.WithKeySet(jwks, jws.WithInferAlgorithmFromKey(true), jws.WithUseDefault(true)), ) if err != nil { return nil, fosite.ErrInvalidClient.WithHint("Invalid client assertion.").WithWrap(err) } if federatedIdentity.ReplayProtection { jti, ok := parsed.JwtID() if !ok || jti == "" { return nil, fosite.ErrInvalidClient.WithHint("Client assertion is missing jti claim, which is required for replay protection.") } // Check if the jti has been used before if err := a.clients.ClientAssertionJWTValid(ctx, jti); err != nil { return nil, fosite.ErrInvalidClient.WithHint("Client assertion has already been used.").WithWrap(err) } // Store the jti to prevent future reuse exp, _ := parsed.Expiration() if err := a.clients.SetClientAssertionJWT(ctx, jti, exp); err != nil { return nil, fosite.ErrInvalidClient.WithWrap(err) } } return client, nil } // keySetForIdentity returns the keys that may have signed an assertion for the given identity. // Identities with public keys configured are verified against those alone, so no JWKS is fetched over the network. func (a *federatedClientAuthenticator) keySetForIdentity(ctx context.Context, federatedIdentity model.OidcClientFederatedIdentity) (jwk.Set, error) { if len(federatedIdentity.PublicKeys) > 0 { jwks, err := jwkutils.ParsePublicKeySet(federatedIdentity.PublicKeys) if err != nil { return nil, fosite.ErrInvalidClient.WithHint("Unable to load the public keys configured for the client assertion.").WithWrap(err) } return jwks, nil } jwksURL := federatedIdentity.JWKS if jwksURL == "" { jwksURL = strings.TrimRight(federatedIdentity.Issuer, "/") + "/.well-known/jwks.json" } jwks, err := a.fetchJWKSet(ctx, jwksURL) if err != nil { return nil, fosite.ErrInvalidClient.WithHint("Unable to fetch client assertion JWKS.").WithWrap(err) } return jwks, nil } func (a *federatedClientAuthenticator) fetchJWKSet(ctx context.Context, jwksURL string) (jwk.Set, error) { if !a.jwksCache.IsRegistered(ctx, jwksURL) { // We set a timeout because otherwise Register will keep trying in case of errors registerCtx, registerCancel := context.WithTimeout(ctx, 15*time.Second) defer registerCancel() // We need to register the URL err := a.jwksCache.Register(registerCtx, jwksURL, jwkfetch.WithMaxInterval(24*time.Hour), jwkfetch.WithMinInterval(15*time.Minute), jwkfetch.WithWaitReady(true), ) // In case of race conditions (two goroutines calling jwkCache.Register at the same time), it's possible we can get a conflict anyways, so we ignore that error if err != nil && !errors.Is(err, httprc.ErrResourceAlreadyExists()) { return nil, fmt.Errorf("failed to register JWK set: %w", err) } } jwks, err := a.jwksCache.CachedSet(jwksURL) if err != nil { return nil, fmt.Errorf("failed to get cached JWK set: %w", err) } return jwks, nil }