From 93a7dd056746975d964ef09c5767464bd68fd80a Mon Sep 17 00:00:00 2001 From: Kyle Mendell Date: Mon, 31 Aug 2026 14:29:56 -0500 Subject: [PATCH] refactor: migrate JWT and JWKS handling to jwx v4 --- backend/go.mod | 13 +- backend/go.sum | 18 + backend/internal/auth/huma_middleware_test.go | 19 +- backend/internal/auth/service.go | 216 +++--- backend/internal/auth/service_test.go | 152 +++-- backend/internal/auth/token.go | 103 +++ backend/internal/di/di.go | 1 + backend/internal/di/di_test.go | 2 + backend/internal/di/providers.go | 11 +- backend/internal/federated/credential.go | 191 ++++++ backend/internal/federated/exchange.go | 365 ++++++++++ backend/internal/federated/service.go | 624 +----------------- backend/internal/federated/service_test.go | 57 +- backend/internal/oidc/service.go | 41 +- backend/pkg/utils/jwtclaims/jwt.go | 37 +- backend/pkg/utils/mldsajose/keyset.go | 105 --- backend/pkg/utils/mldsajose/mldsajose.go | 168 ----- backend/pkg/utils/oidcjwk/algorithms.go | 23 + backend/pkg/utils/oidcjwk/keyset.go | 110 +++ backend/pkg/utils/oidcjwk/manager.go | 192 ++++++ .../oidcjwk_test.go} | 34 +- go.work.sum | 9 - 22 files changed, 1271 insertions(+), 1220 deletions(-) create mode 100644 backend/internal/auth/token.go create mode 100644 backend/internal/federated/credential.go create mode 100644 backend/internal/federated/exchange.go delete mode 100644 backend/pkg/utils/mldsajose/keyset.go delete mode 100644 backend/pkg/utils/mldsajose/mldsajose.go create mode 100644 backend/pkg/utils/oidcjwk/algorithms.go create mode 100644 backend/pkg/utils/oidcjwk/keyset.go create mode 100644 backend/pkg/utils/oidcjwk/manager.go rename backend/pkg/utils/{mldsajose/mldsajose_test.go => oidcjwk/oidcjwk_test.go} (71%) diff --git a/backend/go.mod b/backend/go.mod index e7bff53d65..6e121516b6 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -34,15 +34,16 @@ require ( github.com/getarcaneapp/arcane/cli/v2 v2.9.0 github.com/getarcaneapp/arcane/types/v2 v2.9.0 github.com/go-git/go-git/v5 v5.19.2 - github.com/go-jose/go-jose/v4 v4.1.4 github.com/go-webauthn/webauthn v0.18.0 github.com/gofrs/flock v0.13.1 - github.com/golang-jwt/jwt/v5 v5.3.1 github.com/google/go-containerregistry v0.22.0 github.com/google/uuid v1.6.0 github.com/jinzhu/copier v0.4.0 + github.com/jwx-go/jwkfetch/v4 v4.0.4 github.com/klauspost/compress v1.19.2 github.com/labstack/echo/v5 v5.3.1 + github.com/lestrrat-go/httprc/v3 v3.0.6 + github.com/lestrrat-go/jwx/v4 v4.4.0 github.com/libtnb/sqlite v1.2.2 github.com/lmittmann/tint v1.2.0 github.com/moby/buildkit v0.32.2 @@ -160,11 +161,13 @@ require ( github.com/fxamacker/cbor/v2 v2.9.3 // indirect github.com/go-git/gcfg v1.5.1-0.20230307220236-3a3c6141e376 // indirect github.com/go-git/go-billy/v5 v5.9.0 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-logr/logr v1.4.4 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-ole/go-ole v1.3.0 // indirect github.com/go-viper/mapstructure/v2 v2.5.0 // indirect github.com/go-webauthn/x v0.3.0 // indirect + github.com/golang-jwt/jwt/v5 v5.3.1 // indirect github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 // indirect github.com/google/go-cmp v0.7.0 // indirect github.com/google/go-tpm v0.9.8 // indirect @@ -195,6 +198,11 @@ require ( github.com/knqyf263/go-rpm-version v0.0.0-20220614171824-631e686d1075 // indirect github.com/labstack/echo/v4 v4.15.4 // indirect github.com/labstack/gommon v0.5.0 // indirect + github.com/lestrrat-go/blackmagic v1.0.4 // indirect + github.com/lestrrat-go/dsig v1.4.0 // indirect + github.com/lestrrat-go/httpcc v1.0.1 // indirect + github.com/lestrrat-go/option/v2 v2.0.0 // indirect + github.com/lestrrat-go/option/v3 v3.0.0-alpha1 // indirect github.com/lucasb-eyer/go-colorful v1.4.1 // indirect github.com/lufia/plan9stats v0.0.0-20260330125221-c963978e514e // indirect github.com/mattn/go-colorable v0.1.15 // indirect @@ -266,6 +274,7 @@ require ( github.com/tonistiigi/units v0.0.0-20180711220420-6950e57a87ea // indirect github.com/tonistiigi/vt100 v0.0.0-20240514184818-90bafcd6abab // indirect github.com/valyala/bytebufferpool v1.0.0 // indirect + github.com/valyala/fastjson v1.6.10 // indirect github.com/valyala/fasttemplate v1.2.2 // indirect github.com/vito/midterm v0.1.4 // indirect github.com/vito/progrock v0.10.1 // indirect diff --git a/backend/go.sum b/backend/go.sum index fc9b7cef58..2d4d33b23e 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -485,6 +485,8 @@ github.com/jonboulle/clockwork v0.5.0 h1:Hyh9A8u51kptdkR+cqRpT1EebBwTn1oK9YfGYbd github.com/jonboulle/clockwork v0.5.0/go.mod h1:3mZlmanh0g2NDKO5TWZVJAfofYk64M7XN3SzBPjZF60= github.com/jstemmer/go-junit-report v0.0.0-20190106144839-af01ea7f8024/go.mod h1:6v2b51hI/fHJwM22ozAgKL4VKDeJcHhJFhtBdhmNjmU= github.com/jstemmer/go-junit-report v0.9.1/go.mod h1:Brl9GWCQeLvo8nXZwPNNblvFj/XSXhF0NWZEnDohbsk= +github.com/jwx-go/jwkfetch/v4 v4.0.4 h1:fKgCdegz9WTnBAemDCj1QzPZYIuobg9k2qo+HZkXgmA= +github.com/jwx-go/jwkfetch/v4 v4.0.4/go.mod h1:dTGEkaGuxg9vcCR1F2Dk+TnwmXcZpeUGahNWKtbcKu8= github.com/kevinburke/ssh_config v1.6.0 h1:J1FBfmuVosPHf5GRdltRLhPJtJpTlMdKTBjRgTaQBFY= github.com/kevinburke/ssh_config v1.6.0/go.mod h1:q2RIzfka+BXARoNexmF9gkxEX7DmvbW9P4hIVx2Kg4M= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= @@ -511,6 +513,20 @@ github.com/labstack/echo/v5 v5.3.1 h1:75maCxkQVGualckLc/5s/ihgpH1a1Dc6AuGWNVNs6b github.com/labstack/echo/v5 v5.3.1/go.mod h1:4iEGNQiPPZnkfYpNR/L6fINd3NLiGWUD5+eBotFALas= github.com/labstack/gommon v0.5.0 h1:6VSQ2NOzsnEJ5W6+84E0RbcaDDmgB6NIAzWCczTEe6c= github.com/labstack/gommon v0.5.0/go.mod h1:Rzlg7HHy1maLfzBYGg9NZcVuz1sA68HHhLjhcEllYE0= +github.com/lestrrat-go/blackmagic v1.0.4 h1:IwQibdnf8l2KoO+qC3uT4OaTWsW7tuRQXy9TRN9QanA= +github.com/lestrrat-go/blackmagic v1.0.4/go.mod h1:6AWFyKNNj0zEXQYfTMPfZrAXUWUfTIZ5ECEUEJaijtw= +github.com/lestrrat-go/dsig v1.4.0 h1:g7LUjK8cT74A5DzBXJI5HzsJuLhoYN0Wzj4nuOMIrH8= +github.com/lestrrat-go/dsig v1.4.0/go.mod h1:I8Nddg/vN2cUl/h8N7SRRApLnNNeyZPIqLYpvpOtGGo= +github.com/lestrrat-go/httpcc v1.0.1 h1:ydWCStUeJLkpYyjLDHihupbn2tYmZ7m22BGkcvZZrIE= +github.com/lestrrat-go/httpcc v1.0.1/go.mod h1:qiltp3Mt56+55GPVCbTdM9MlqhvzyuL6W/NMDA8vA5E= +github.com/lestrrat-go/httprc/v3 v3.0.6 h1:4FpLQ18KK/ypPbVU3NLWJNRvH3kcYiqKqWfKGqNWxxI= +github.com/lestrrat-go/httprc/v3 v3.0.6/go.mod h1:mSMtkZW92Z98M5YoNNztbRGxbXHql7tSitCvaxvo9l0= +github.com/lestrrat-go/jwx/v4 v4.4.0 h1:CzoK8+u++WF7vVEmxx9fB8VaheeXWZ698F6HZbrl6SI= +github.com/lestrrat-go/jwx/v4 v4.4.0/go.mod h1:65utsGK/iSrjgGfu6iqj/TAvSfia6SSXkRpjHcKcTyg= +github.com/lestrrat-go/option/v2 v2.0.0 h1:XxrcaJESE1fokHy3FpaQ/cXW8ZsIdWcdFzzLOcID3Ss= +github.com/lestrrat-go/option/v2 v2.0.0/go.mod h1:oSySsmzMoR0iRzCDCaUfsCzxQHUEuhOViQObyy7S6Vg= +github.com/lestrrat-go/option/v3 v3.0.0-alpha1 h1:dvdzLwm/Ba5CJUF3jQP7w/iNYSLfy7yyh9XXNa1WjxI= +github.com/lestrrat-go/option/v3 v3.0.0-alpha1/go.mod h1:5KSg20dfsKkNJtjDmaQRLZVXuUrzuCCcz/gbDK0pfKk= github.com/libtnb/sqlite v1.2.2 h1:Ku5hAPP5B3A4kQcDn4Z3qyNMafKzuAbnFZX4oD9wR4I= github.com/libtnb/sqlite v1.2.2/go.mod h1:JkAuxM7HHo0tc7dQENjIGpp/yIfrLX5nKr5ME9wuvmI= github.com/lmittmann/tint v1.2.0 h1:AogHRHy8HUJUnNJBHJlYa+fR4YY8mko2cnCp67xn9JY= @@ -763,6 +779,8 @@ github.com/ulikunitz/xz v0.5.15 h1:9DNdB5s+SgV3bQ2ApL10xRc35ck0DuIX/isZvIk+ubY= github.com/ulikunitz/xz v0.5.15/go.mod h1:nbz6k7qbPmH4IRqmfOplQw/tblSgqTqBwxkY0oWt/14= github.com/valyala/bytebufferpool v1.0.0 h1:GqA5TC/0021Y/b9FG4Oi9Mr3q7XYx6KllzawFIhcdPw= github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc= +github.com/valyala/fastjson v1.6.10 h1:/yjJg8jaVQdYR3arGxPE2X5z89xrlhS0eGXdv+ADTh4= +github.com/valyala/fastjson v1.6.10/go.mod h1:e6FubmQouUNP73jtMLmcbxS6ydWIpOfhz34TSfO3JaE= github.com/valyala/fasttemplate v1.2.2 h1:lxLXG0uE3Qnshl9QyaK6XJxMXlQZELvChBOCmQD0Loo= github.com/valyala/fasttemplate v1.2.2/go.mod h1:KHLXt3tVN2HBp8eijSv/kGJopbvo7S+qRAEEKiv+SiQ= github.com/vbatts/tar-split v0.12.3 h1:Cd46rkGXI3Td4yrVNwU8ripbxFaQbmesqhjBUUYAJSw= diff --git a/backend/internal/auth/huma_middleware_test.go b/backend/internal/auth/huma_middleware_test.go index 5c9987be9a..644e0bf619 100644 --- a/backend/internal/auth/huma_middleware_test.go +++ b/backend/internal/auth/huma_middleware_test.go @@ -23,9 +23,7 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/user" "github.com/getarcaneapp/arcane/backend/v2/pkg/authz" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/mldsajose" authtypes "github.com/getarcaneapp/arcane/types/v2/auth" - "github.com/golang-jwt/jwt/v5" "github.com/labstack/echo/v5" "github.com/libtnb/sqlite" "github.com/stretchr/testify/require" @@ -281,7 +279,7 @@ func TestNewHumaMiddleware_OpportunisticAuthOnPublicRoute(t *testing.T) { session, _, err := sessionSvc.CreateSession(context.Background(), "u-logout", exp, authtypes.SessionMeta{}) require.NoError(t, err) - claims := jwt.MapClaims{ + claims := map[string]any{ "jti": "u-logout", "sub": "access", "iat": time.Now().Unix(), @@ -291,8 +289,7 @@ func TestNewHumaMiddleware_OpportunisticAuthOnPublicRoute(t *testing.T) { "username": "logouttest", "roles": []string{"user"}, } - token, err := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, claims).SignedString(signingKey) - require.NoError(t, err) + token := signJWXTokenInternal(t, signingKey, claims) router := echo.New() apiGroup := router.Group("/api") @@ -371,7 +368,7 @@ func TestNewHumaMiddleware_VersionMismatchIsRecoverable(t *testing.T) { // An empty appVersion omits the claim, which passes the version check (no pin). mintToken := func(appVersion string) string { - claims := jwt.MapClaims{ + claims := map[string]any{ "jti": "u-ver", "sub": "access", "iat": time.Now().Unix(), @@ -383,9 +380,7 @@ func TestNewHumaMiddleware_VersionMismatchIsRecoverable(t *testing.T) { if appVersion != "" { claims["app_version"] = appVersion } - token, signErr := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, claims).SignedString(signingKey) - require.NoError(t, signErr) - return token + return signJWXTokenInternal(t, signingKey, claims) } router := echo.New() @@ -498,7 +493,7 @@ func mintHumaMiddlewareTestTokenInternal(t *testing.T, userSvc *user.UserService session, _, err := sessionSvc.CreateSession(context.Background(), userID, exp, authtypes.SessionMeta{}) require.NoError(t, err) - claims := jwt.MapClaims{ + claims := map[string]any{ "jti": userID, "sub": "access", "iat": time.Now().Unix(), @@ -507,9 +502,7 @@ func mintHumaMiddlewareTestTokenInternal(t *testing.T, userSvc *user.UserService "user_id": userID, "username": userID, } - token, err := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, claims).SignedString(signingKey) - require.NoError(t, err) - return token + return signJWXTokenInternal(t, signingKey, claims) } func (r staticPermissionResolverInternal) ResolvePermissions(_ context.Context, _ *common.User) (*authz.PermissionSet, error) { diff --git a/backend/internal/auth/service.go b/backend/internal/auth/service.go index 9f588fb616..cd8a666bd2 100644 --- a/backend/internal/auth/service.go +++ b/backend/internal/auth/service.go @@ -27,10 +27,10 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/settings" "github.com/getarcaneapp/arcane/backend/v2/internal/user" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/jwtclaims" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/mldsajose" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/validation" "github.com/getarcaneapp/arcane/types/v2/auth" - "github.com/golang-jwt/jwt/v5" + "github.com/lestrrat-go/jwx/v4/jwa" + "github.com/lestrrat-go/jwx/v4/jwt" "github.com/samber/hot" ) @@ -54,27 +54,6 @@ type AuthSettings struct { Oidc *settings.OidcConfig `json:"oidc,omitempty"` } -type userClaims struct { - jwt.RegisteredClaims - - SessionID string `json:"sid,omitempty"` - UserID string `json:"user_id"` - Username string `json:"username"` - Email string `json:"email,omitempty"` - DisplayName string `json:"display_name,omitempty"` - AppVersion string `json:"app_version,omitempty"` - TokenType string `json:"token_type,omitempty"` - FederatedCredentialID string `json:"federated_credential_id,omitempty"` -} - -type refreshClaims struct { - jwt.RegisteredClaims - - UserID string `json:"user_id"` - SessionID string `json:"sid,omitempty"` - AppVersion string `json:"app_version,omitempty"` -} - type verifiedTokenEntry struct { User common.User SessionID string @@ -729,43 +708,15 @@ func (s *AuthService) RefreshToken(ctx context.Context, refreshToken string, met if err != nil { return nil, err } - token, err := jwt.ParseWithClaims(refreshToken, &refreshClaims{}, - func(t *jwt.Token) (any, error) { - if _, ok := t.Method.(*mldsajose.SigningMethodMLDSA); !ok { - return nil, errors.Errorf("unexpected signing method: %v", t.Header["alg"]) - } - return signingKey.PublicKey(), nil - }, jwt.WithValidMethods([]string{mldsajose.AlgMLDSA87})) + claims, err := parseRefreshTokenInternal(ctx, refreshToken, signingKey.PublicKey()) if err != nil { - return nil, common.ErrInvalidToken - } - - if !token.Valid { - return nil, common.ErrInvalidToken - } - - claims, ok := token.Claims.(*refreshClaims) - if !ok { - return nil, common.Classify(common.ErrTokenValidation, errors.New("Invalid token claims")) - } - - if claims.Subject != "refresh" { - return nil, common.Classify(common.ErrTokenValidation, errors.New("Not a refresh token")) + return nil, err } if claims.AppVersion != "" && claims.AppVersion != config.Version { slog.InfoContext(ctx, "Refresh token version mismatch — rotating to current version", "tokenVersion", claims.AppVersion, "currentVersion", config.Version) } - if claims.UserID == "" { - return nil, common.Classify(common.ErrTokenValidation, errors.New("Missing user ID in token")) - } - if claims.ID == "" { - return nil, common.Classify(common.ErrTokenValidation, errors.New("Missing refresh token ID")) - } - if claims.SessionID == "" { - return nil, common.Classify(common.ErrTokenValidation, errors.New("Missing session ID in token")) - } if s.sessionService == nil { return nil, common.Classify(common.ErrUnavailable, errors.New("Session service is not configured")) } @@ -799,51 +750,22 @@ func (s *AuthService) VerifyToken(ctx context.Context, accessToken string) (*com if err != nil { return nil, "", err } - token, err := jwt.ParseWithClaims(accessToken, &userClaims{}, - func(t *jwt.Token) (any, error) { - if _, ok := t.Method.(*mldsajose.SigningMethodMLDSA); !ok { - return nil, errors.Errorf("unexpected signing method: %v", t.Header["alg"]) - } - return signingKey.PublicKey(), nil - }, jwt.WithValidMethods([]string{mldsajose.AlgMLDSA87})) + claims, err := parseAccessTokenInternal(ctx, accessToken, signingKey.PublicKey()) if err != nil { - if strings.Contains(err.Error(), "token is expired") { - return nil, "", common.ErrExpiredToken - } - return nil, "", common.ErrInvalidToken - } - - if !token.Valid { - return nil, "", common.ErrInvalidToken - } - - claims, ok := token.Claims.(*userClaims) - if !ok { - return nil, "", common.Classify(common.ErrTokenValidation, errors.New("Invalid token claims")) - } - - if claims.Subject != "access" { - return nil, "", common.Classify(common.ErrTokenValidation, errors.New("Not an access token")) - } - - if claims.ID == "" { - return nil, "", common.Classify(common.ErrTokenValidation, errors.New("Missing user ID in token")) + return nil, "", err } if claims.AppVersion != "" && claims.AppVersion != config.Version { slog.InfoContext(ctx, "Token version mismatch detected", "tokenVersion", claims.AppVersion, "currentVersion", config.Version, "user", claims.Username) return nil, "", common.ErrTokenVersionMismatch } - if claims.SessionID == "" { - return nil, "", common.Classify(common.ErrTokenValidation, errors.New("Missing session ID in token")) - } if s.sessionService == nil { return nil, "", common.Classify(common.ErrUnavailable, errors.New("Session service is not configured")) } tokenHash := hashTokenInternal(accessToken) if cached, ok, _ := s.tokenCache.Get(tokenHash); ok { - if cached.User.ID != claims.ID || cached.SessionID != claims.SessionID { + if cached.User.ID != claims.UserID || cached.SessionID != claims.SessionID { s.tokenCache.Delete(tokenHash) return nil, "", common.ErrInvalidToken } @@ -868,7 +790,7 @@ func (s *AuthService) VerifyToken(ctx context.Context, accessToken string) (*com // Verify user exists in DB // This ensures that if the database is wiped or user is deleted, the token becomes invalid // even if the JWT signature is still valid (e.g. same signing key). - dbUser, err := s.userService.GetUserByID(ctx, claims.ID) + dbUser, err := s.userService.GetUserByID(ctx, claims.UserID) if err != nil { if errors.Is(err, common.ErrUserNotFound) { return nil, "", common.ErrInvalidToken @@ -998,58 +920,64 @@ func (s *AuthService) createSessionAndTokensInternal(ctx context.Context, user * func (s *AuthService) buildTokenPairInternal(ctx context.Context, user *common.User, session *session.UserSession, refreshJTI string) (*TokenPair, error) { sessionTimeout, _ := s.GetSessionTimeout(ctx) + now := time.Now() + accessTokenExpiry := now.Add(time.Duration(sessionTimeout) * time.Minute) - accessTokenExpiry := time.Now().Add(time.Duration(sessionTimeout) * time.Minute) - - userClaims := userClaims{ - ID: user.ID, - Subject: "access", - IssuedAt: jwt.NewNumericDate(time.Now()), - ExpiresAt: jwt.NewNumericDate(accessTokenExpiry), - SessionID: session.ID, - UserID: user.ID, - Username: user.Username, - AppVersion: config.Version, + signingKey, err := s.signingKeyInternal(ctx) + if err != nil { + return nil, err } - if user.Email != nil { - userClaims.Email = *user.Email + accessBuilder := jwt.NewBuilder(). + JwtID(user.ID). + Subject(accessTokenSubject). + IssuedAt(now). + Expiration(accessTokenExpiry). + Claim(claimSessionID, session.ID). + Claim(claimUserID, user.ID) + if user.Username != "" { + accessBuilder.Claim(claimUsername, user.Username) } - - if user.DisplayName != nil { - userClaims.DisplayName = *user.DisplayName + if user.Email != nil && *user.Email != "" { + accessBuilder.Claim(claimEmail, *user.Email) } - - signingKey, err := s.signingKeyInternal(ctx) + if user.DisplayName != nil && *user.DisplayName != "" { + accessBuilder.Claim(claimDisplayName, *user.DisplayName) + } + if config.Version != "" { + accessBuilder.Claim(claimAppVersion, config.Version) + } + accessToken, err := accessBuilder.Build() if err != nil { return nil, err } - - accessToken := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, userClaims) - - accessTokenString, err := accessToken.SignedString(signingKey) + accessTokenBytes, err := jwt.Sign(accessToken, jwt.WithKey(jwa.MLDSA87(), signingKey)) if err != nil { return nil, err } - refreshToken := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, refreshClaims{ - ID: refreshJTI, - Subject: "refresh", - IssuedAt: jwt.NewNumericDate(time.Now()), - ExpiresAt: jwt.NewNumericDate(session.ExpiresAt), - UserID: user.ID, - SessionID: session.ID, - AppVersion: config.Version, - }) - - refreshTokenString, err := refreshToken.SignedString(signingKey) + refreshBuilder := jwt.NewBuilder(). + JwtID(refreshJTI). + Subject(refreshTokenSubject). + IssuedAt(now). + Expiration(session.ExpiresAt). + Claim(claimUserID, user.ID). + Claim(claimSessionID, session.ID) + if config.Version != "" { + refreshBuilder.Claim(claimAppVersion, config.Version) + } + refreshToken, err := refreshBuilder.Build() + if err != nil { + return nil, err + } + refreshTokenBytes, err := jwt.Sign(refreshToken, jwt.WithKey(jwa.MLDSA87(), signingKey)) if err != nil { return nil, err } return &TokenPair{ - AccessToken: accessTokenString, - RefreshToken: refreshTokenString, + AccessToken: string(accessTokenBytes), + RefreshToken: string(refreshTokenBytes), ExpiresAt: accessTokenExpiry, }, nil } @@ -1071,38 +999,42 @@ func (s *AuthService) IssueFederatedToken(ctx context.Context, user *common.User return nil, err } - claims := userClaims{ - ID: user.ID, - Subject: "access", - IssuedAt: jwt.NewNumericDate(now), - ExpiresAt: jwt.NewNumericDate(accessTokenExpiry), - SessionID: federatedSession.ID, - UserID: user.ID, - Username: user.Username, - AppVersion: config.Version, - TokenType: session.UserSessionSourceFederated, - FederatedCredentialID: credentialID, + signingKey, err := s.signingKeyInternal(ctx) + if err != nil { + return nil, err } - - if user.Email != nil { - claims.Email = *user.Email + builder := jwt.NewBuilder(). + JwtID(user.ID). + Subject(accessTokenSubject). + IssuedAt(now). + Expiration(accessTokenExpiry). + Claim(claimSessionID, federatedSession.ID). + Claim(claimUserID, user.ID). + Claim(claimTokenType, session.UserSessionSourceFederated). + Claim(claimFederatedCredentialID, credentialID) + if user.Username != "" { + builder.Claim(claimUsername, user.Username) } - if user.DisplayName != nil { - claims.DisplayName = *user.DisplayName + if user.Email != nil && *user.Email != "" { + builder.Claim(claimEmail, *user.Email) } - - signingKey, err := s.signingKeyInternal(ctx) + if user.DisplayName != nil && *user.DisplayName != "" { + builder.Claim(claimDisplayName, *user.DisplayName) + } + if config.Version != "" { + builder.Claim(claimAppVersion, config.Version) + } + accessToken, err := builder.Build() if err != nil { return nil, err } - accessToken := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, claims) - accessTokenString, err := accessToken.SignedString(signingKey) + accessTokenBytes, err := jwt.Sign(accessToken, jwt.WithKey(jwa.MLDSA87(), signingKey)) if err != nil { return nil, err } return &TokenPair{ - AccessToken: accessTokenString, + AccessToken: string(accessTokenBytes), ExpiresAt: accessTokenExpiry, }, nil } diff --git a/backend/internal/auth/service_test.go b/backend/internal/auth/service_test.go index c8e0084aea..dbe4ab7493 100644 --- a/backend/internal/auth/service_test.go +++ b/backend/internal/auth/service_test.go @@ -9,6 +9,8 @@ import ( "context" "crypto/mldsa" + "encoding/base64" + "encoding/json/v2" "testing" "time" @@ -24,9 +26,9 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/session" "github.com/getarcaneapp/arcane/backend/v2/internal/settings" "github.com/getarcaneapp/arcane/backend/v2/internal/user" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/mldsajose" "github.com/getarcaneapp/arcane/types/v2/auth" - "github.com/golang-jwt/jwt/v5" + "github.com/lestrrat-go/jwx/v4/jwa" + "github.com/lestrrat-go/jwx/v4/jwt" "github.com/samber/hot" "github.com/stretchr/testify/assert" ) @@ -93,31 +95,59 @@ func newTestAuthService() *AuthService { } } +func signLegacyTokenInternal(t testing.TB, key *mldsa.PrivateKey, claims map[string]any) string { + t.Helper() + header, err := json.Marshal(map[string]string{"alg": jwa.MLDSA87().String(), "typ": "JWT"}) + require.NoError(t, err) + payload, err := json.Marshal(claims) + require.NoError(t, err) + signingInput := base64.RawURLEncoding.EncodeToString(header) + "." + base64.RawURLEncoding.EncodeToString(payload) + signature, err := key.Sign(nil, []byte(signingInput), nil) + require.NoError(t, err) + return signingInput + "." + base64.RawURLEncoding.EncodeToString(signature) +} + +func signJWXTokenInternal(t testing.TB, key *mldsa.PrivateKey, claims map[string]any) string { + t.Helper() + payload, err := json.Marshal(claims) + require.NoError(t, err) + token, err := jwt.ParseInsecure(payload) + require.NoError(t, err) + signed, err := jwt.Sign(token, jwt.WithKey(jwa.MLDSA87(), key)) + require.NoError(t, err) + return string(signed) +} + +func signUnsignedTokenInternal(t testing.TB, claims map[string]any) string { + t.Helper() + payload, err := json.Marshal(claims) + require.NoError(t, err) + token, err := jwt.ParseInsecure(payload) + require.NoError(t, err) + signed, err := jwt.Sign(token, jwt.WithInsecureNoSignature()) + require.NoError(t, err) + return string(signed) +} + func makeAccessToken(t *testing.T, key *mldsa.PrivateKey, subject string, id string, username string, _ []string, email, displayName string, exp time.Time, sessionIDs ...string) string { t.Helper() sessionID := "" if len(sessionIDs) > 0 { sessionID = sessionIDs[0] } - claims := userClaims{ - ID: id, - Subject: subject, - IssuedAt: jwt.NewNumericDate(time.Now()), - ExpiresAt: jwt.NewNumericDate(exp), - SessionID: sessionID, - UserID: id, - Username: username, - Email: email, - DisplayName: displayName, - AppVersion: config.Version, + claims := map[string]any{ + "jti": id, + "sub": subject, + "iat": time.Now().Unix(), + "exp": exp.Unix(), + "sid": sessionID, + "user_id": id, + "username": username, + "email": email, + "display_name": displayName, + "app_version": config.Version, } - tok := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, claims) - signed, err := tok.SignedString(key) - - require.NoError(t, err, - "sign: %v", err) - - return signed + return signLegacyTokenInternal(t, key, claims) } func makeRefreshToken(t *testing.T, key *mldsa.PrivateKey, subject string, id string, exp time.Time, userIDAndSessionID ...string) string { @@ -130,22 +160,16 @@ func makeRefreshToken(t *testing.T, key *mldsa.PrivateKey, subject string, id st if len(userIDAndSessionID) > 1 { sessionID = userIDAndSessionID[1] } - claims := refreshClaims{ - ID: id, - Subject: subject, - IssuedAt: jwt.NewNumericDate(time.Now()), - ExpiresAt: jwt.NewNumericDate(exp), - UserID: userID, - SessionID: sessionID, - AppVersion: config.Version, + claims := map[string]any{ + "jti": id, + "sub": subject, + "iat": time.Now().Unix(), + "exp": exp.Unix(), + "user_id": userID, + "sid": sessionID, + "app_version": config.Version, } - tok := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, claims) - signed, err := tok.SignedString(key) - - require.NoError(t, err, - "sign: %v", err) - - return signed + return signLegacyTokenInternal(t, key, claims) } func createTestSession(t *testing.T, db *database.DB, userID string, expiresAt time.Time) (*session.UserSession, string) { @@ -159,15 +183,9 @@ func createTestSession(t *testing.T, db *database.DB, userID string, expiresAt t return session, refreshJTI } -func makeUnsignedToken(t *testing.T, claims jwt.Claims) string { +func makeUnsignedToken(t *testing.T, claims map[string]any) string { t.Helper() - tok := jwt.NewWithClaims(jwt.SigningMethodNone, claims) - signed, err := tok.SignedString(jwt.UnsafeAllowNoneSignatureType) - - require.NoError(t, err, - "sign none: %v", err) - - return signed + return signUnsignedTokenInternal(t, claims) } func TestVerifyToken_ValidClaims(t *testing.T) { @@ -213,14 +231,14 @@ func TestVerifyToken_ValidClaims(t *testing.T) { func TestVerifyToken_RejectsNonMLDSAAlg(t *testing.T) { s := newTestAuthService() exp := time.Now().Add(5 * time.Minute) - token := makeUnsignedToken(t, userClaims{ - ID: "u1", - Subject: "access", - IssuedAt: jwt.NewNumericDate(time.Now()), - ExpiresAt: jwt.NewNumericDate(exp), - UserID: "u1", - Username: "bob", - AppVersion: config.Version, + token := makeUnsignedToken(t, map[string]any{ + "jti": "u1", + "sub": "access", + "iat": time.Now().Unix(), + "exp": exp.Unix(), + "user_id": "u1", + "username": "bob", + "app_version": config.Version, }) _, _, err := s.VerifyToken(context.Background(), token) @@ -353,7 +371,7 @@ func TestVerifyToken_VersionMismatch(t *testing.T) { oldVersion := config.Version config.Version = "1.0.0" - token := makeAccessToken(t, s.signingKey, "access", "u1", "bob", []string{"user"}, "", "", exp) + token := makeAccessToken(t, s.signingKey, "access", "u1", "bob", []string{"user"}, "", "", exp, "session-version-mismatch") config.Version = "2.0.0" _, _, err := s.VerifyToken(context.Background(), token) @@ -427,21 +445,17 @@ func TestRefreshToken_VersionMismatchRotates(t *testing.T) { require.NotEmpty(t, tokenPair.AccessToken) require.NotEmpty(t, tokenPair.RefreshToken) - parsedAccess, err := jwt.ParseWithClaims(tokenPair.AccessToken, &userClaims{}, func(*jwt.Token) (any, error) { - return s.signingKey.PublicKey(), nil - }) + parsedAccess, err := jwt.ParseString(tokenPair.AccessToken, jwt.WithKey(jwa.MLDSA87(), s.signingKey.PublicKey())) + require.NoError(t, err) + accessVersion, err := jwt.Get[string](parsedAccess, claimAppVersion) require.NoError(t, err) - accessClaims, ok := parsedAccess.Claims.(*userClaims) - require.True(t, ok) - require.Equal(t, "2.0.0", accessClaims.AppVersion) + require.Equal(t, "2.0.0", accessVersion) - parsedRefresh, err := jwt.ParseWithClaims(tokenPair.RefreshToken, &refreshClaims{}, func(*jwt.Token) (any, error) { - return s.signingKey.PublicKey(), nil - }) + parsedRefresh, err := jwt.ParseString(tokenPair.RefreshToken, jwt.WithKey(jwa.MLDSA87(), s.signingKey.PublicKey())) + require.NoError(t, err) + refreshVersion, err := jwt.Get[string](parsedRefresh, claimAppVersion) require.NoError(t, err) - rClaims, ok := parsedRefresh.Claims.(*refreshClaims) - require.True(t, ok) - require.Equal(t, "2.0.0", rClaims.AppVersion) + require.Equal(t, "2.0.0", refreshVersion) } func TestVerifyToken_RejectsRevokedSession(t *testing.T) { @@ -654,11 +668,11 @@ func TestChangePassword_KeepsCurrentSessionAlive(t *testing.T) { func TestRefreshToken_RejectsNonHMACAlg(t *testing.T) { s := newTestAuthService() exp := time.Now().Add(5 * time.Minute) - token := makeUnsignedToken(t, jwt.RegisteredClaims{ - ID: "u1", - Subject: "refresh", - IssuedAt: jwt.NewNumericDate(time.Now()), - ExpiresAt: jwt.NewNumericDate(exp), + token := makeUnsignedToken(t, map[string]any{ + "jti": "u1", + "sub": "refresh", + "iat": time.Now().Unix(), + "exp": exp.Unix(), }) _, err := s.RefreshToken(context.Background(), token, auth.SessionMeta{}) diff --git a/backend/internal/auth/token.go b/backend/internal/auth/token.go new file mode 100644 index 0000000000..00409c8463 --- /dev/null +++ b/backend/internal/auth/token.go @@ -0,0 +1,103 @@ +package auth + +import ( + "context" + "crypto/mldsa" + + "emperror.dev/errors" + "github.com/lestrrat-go/jwx/v4/jwa" + "github.com/lestrrat-go/jwx/v4/jwt" + + "github.com/getarcaneapp/arcane/backend/v2/internal/common" +) + +const ( + accessTokenSubject = "access" + refreshTokenSubject = "refresh" + + claimSessionID = "sid" + claimUserID = "user_id" + claimUsername = "username" + claimEmail = "email" + claimDisplayName = "display_name" + claimAppVersion = "app_version" + claimTokenType = "token_type" + claimFederatedCredentialID = "federated_credential_id" +) + +type accessTokenClaims struct { + UserID string + SessionID string + Username string + AppVersion string +} + +type refreshTokenClaims struct { + ID string + UserID string + SessionID string + AppVersion string +} + +var tokenParseOptions = []jwt.ParseOption{ + jwt.WithStrictStringClaims(true), + jwt.WithRequiredClaim(jwt.SubjectKey), + jwt.WithRequiredClaim(jwt.JwtIDKey), + jwt.WithRequiredClaim(jwt.IssuedAtKey), + jwt.WithRequiredClaim(jwt.ExpirationKey), + jwt.WithRequiredClaim(claimSessionID), + jwt.WithRequiredClaim(claimUserID), +} + +func parseAccessTokenInternal(ctx context.Context, rawToken string, key *mldsa.PublicKey) (*accessTokenClaims, error) { + token, err := jwt.ParseString(rawToken, append(tokenParseOptions, jwt.WithKey(jwa.MLDSA87(), key), jwt.WithContext(ctx))...) + if err != nil { + switch { + case errors.Is(err, jwt.TokenExpiredError{}): + return nil, common.ErrExpiredToken + case errors.Is(err, jwt.MissingRequiredClaimError{}), errors.Is(err, jwt.ClaimValidationError{}): + return nil, common.ErrTokenValidation + default: + return nil, common.ErrInvalidToken + } + } + + subject, _ := token.Subject() + id, _ := token.JwtID() + userID, _ := jwt.Get[string](token, claimUserID) + sessionID, _ := jwt.Get[string](token, claimSessionID) + if subject != accessTokenSubject || userID == "" || sessionID == "" || id != userID { + return nil, common.ErrTokenValidation + } + + username, usernameErr := jwt.Get[string](token, claimUsername) + appVersion, appVersionErr := jwt.Get[string](token, claimAppVersion) + if errors.Is(usernameErr, jwt.ClaimTypeMismatchError{}) || errors.Is(appVersionErr, jwt.ClaimTypeMismatchError{}) { + return nil, common.ErrTokenValidation + } + return new(accessTokenClaims{UserID: userID, SessionID: sessionID, Username: username, AppVersion: appVersion}), nil +} + +func parseRefreshTokenInternal(ctx context.Context, rawToken string, key *mldsa.PublicKey) (*refreshTokenClaims, error) { + token, err := jwt.ParseString(rawToken, append(tokenParseOptions, jwt.WithKey(jwa.MLDSA87(), key), jwt.WithContext(ctx))...) + if err != nil { + if errors.Is(err, jwt.MissingRequiredClaimError{}) || errors.Is(err, jwt.ClaimValidationError{}) { + return nil, common.ErrTokenValidation + } + return nil, common.ErrInvalidToken + } + + subject, _ := token.Subject() + id, _ := token.JwtID() + userID, _ := jwt.Get[string](token, claimUserID) + sessionID, _ := jwt.Get[string](token, claimSessionID) + if subject != refreshTokenSubject || id == "" || userID == "" || sessionID == "" { + return nil, common.ErrTokenValidation + } + + appVersion, appVersionErr := jwt.Get[string](token, claimAppVersion) + if errors.Is(appVersionErr, jwt.ClaimTypeMismatchError{}) { + return nil, common.ErrTokenValidation + } + return new(refreshTokenClaims{ID: id, UserID: userID, SessionID: sessionID, AppVersion: appVersion}), nil +} diff --git a/backend/internal/di/di.go b/backend/internal/di/di.go index 3287b0d106..3e7f14265a 100644 --- a/backend/internal/di/di.go +++ b/backend/internal/di/di.go @@ -48,6 +48,7 @@ var ServiceOptions = fx.Options( fx.Provide( // Infrastructure values consumed by services. provideResourcesFSInternal, + provideJWKSetManagerInternal, // Services constructed directly through their public constructors. provideEventModuleInternal, diff --git a/backend/internal/di/di_test.go b/backend/internal/di/di_test.go index 00e583336b..d0d3d4644b 100644 --- a/backend/internal/di/di_test.go +++ b/backend/internal/di/di_test.go @@ -50,6 +50,7 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/volume" "github.com/getarcaneapp/arcane/backend/v2/internal/webhook" "github.com/getarcaneapp/arcane/backend/v2/pkg/scheduler" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" "github.com/stretchr/testify/require" "go.uber.org/fx" ) @@ -102,6 +103,7 @@ type graphParams struct { Role *role.RoleService Variable *variable.VariableService AuthMiddleware *auth.AuthMiddleware + JWKSetManager *oidcjwk.KeySetManager AutoUpdate *scheduler.AutoUpdateJob ImageUpdateWatcher *scheduler.ImageUpdateWatcher diff --git a/backend/internal/di/providers.go b/backend/internal/di/providers.go index 9cd144e1bf..429aed4581 100644 --- a/backend/internal/di/providers.go +++ b/backend/internal/di/providers.go @@ -49,6 +49,7 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/webhook" "github.com/getarcaneapp/arcane/backend/v2/pkg/libarcane/edge" "github.com/getarcaneapp/arcane/backend/v2/pkg/scheduler" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" "github.com/getarcaneapp/arcane/backend/v2/resources" "go.uber.org/fx" ) @@ -363,8 +364,14 @@ func provideApiKeyServiceInternal(module *apikey.Module) *apikey.ApiKeyService { return module.Service() } -func provideFederatedCredentialServiceInternal(db *database.DB, auth *auth.AuthService, user *user.UserService, settings *settings.SettingsService, event *event.EventService, httpClient *http.Client, role *role.RoleService) *federated.FederatedCredentialService { - return federated.NewFederatedCredentialService(db, auth, user, settings, event, httpClient).WithRoleService(role) +func provideJWKSetManagerInternal(ctx context.Context, lc fx.Lifecycle) *oidcjwk.KeySetManager { + manager := oidcjwk.NewKeySetManager(ctx) + lc.Append(fx.Hook{OnStop: manager.Shutdown}) + return manager +} + +func provideFederatedCredentialServiceInternal(db *database.DB, auth *auth.AuthService, user *user.UserService, settings *settings.SettingsService, event *event.EventService, httpClient *http.Client, role *role.RoleService, keySetManager *oidcjwk.KeySetManager) *federated.FederatedCredentialService { + return federated.NewFederatedCredentialService(db, auth, user, settings, event, httpClient, keySetManager).WithRoleService(role) } func provideAuthMiddlewareInternal(authService *auth.AuthService, apiKey *apikey.ApiKeyService, env *environment.EnvironmentService, role *role.RoleService, cfg *config.Config) *auth.AuthMiddleware { diff --git a/backend/internal/federated/credential.go b/backend/internal/federated/credential.go new file mode 100644 index 0000000000..ec63de4cc8 --- /dev/null +++ b/backend/internal/federated/credential.go @@ -0,0 +1,191 @@ +package federated + +import ( + "context" + "net/url" + "strings" + + "emperror.dev/errors" + "github.com/samber/mo" + + "github.com/getarcaneapp/arcane/backend/v2/internal/auth" + "github.com/getarcaneapp/arcane/backend/v2/internal/common" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils" + federatedtypes "github.com/getarcaneapp/arcane/types/v2/federated" +) + +func normalizeCreateFederatedCredentialInternal(req federatedtypes.CreateFederatedCredential) (federatedtypes.CreateFederatedCredential, error) { + req.Name = strings.TrimSpace(req.Name) + req.IssuerURL = strings.TrimRight(strings.TrimSpace(req.IssuerURL), "/") + req.SubjectClaim = strings.TrimSpace(req.SubjectClaim) + req.SubjectMatch = strings.TrimSpace(req.SubjectMatch) + req.MatchType = normalizeMatchTypeInternal(req.MatchType) + req.Audiences = utils.UniqueNonEmptyStrings(req.Audiences) + req.EnvironmentID = mo.EmptyableToOption(strings.TrimSpace(mo.PointerToOption(req.EnvironmentID).OrEmpty())).ToPointer() + req.TokenTTLSeconds = auth.ClampFederatedTokenTTLSeconds(req.TokenTTLSeconds) + + if req.SubjectClaim == "" { + req.SubjectClaim = defaultFederatedSubjectClaim + } + if req.Name == "" || req.SubjectMatch == "" || req.RoleID == "" || len(req.Audiences) == 0 { + return req, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) + } + if err := validateIssuerURLInternal(req.IssuerURL); err != nil { + return req, err + } + if err := validateSubjectMatchInternal(req.MatchType, req.SubjectMatch); err != nil { + return req, err + } + return req, nil +} + +func applyFederatedCredentialUpdateInternal(existing FederatedCredential, req federatedtypes.UpdateFederatedCredential) (FederatedCredential, bool, error) { + if req.Name != nil { + name := strings.TrimSpace(*req.Name) + if name == "" { + return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) + } + existing.Name = name + } + if req.Description != nil { + existing.Description = req.Description + } + if req.Enabled != nil { + existing.Enabled = *req.Enabled + } + if req.IssuerURL != nil { + issuerURL := strings.TrimRight(strings.TrimSpace(*req.IssuerURL), "/") + if err := validateIssuerURLInternal(issuerURL); err != nil { + return existing, false, err + } + existing.IssuerURL = issuerURL + } + if req.Audiences != nil { + audiences := utils.UniqueNonEmptyStrings(req.Audiences) + if len(audiences) == 0 { + return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) + } + existing.Audiences = audiences + } + if req.SubjectClaim != nil { + subjectClaim := strings.TrimSpace(*req.SubjectClaim) + if subjectClaim == "" { + subjectClaim = defaultFederatedSubjectClaim + } + existing.SubjectClaim = subjectClaim + } + if req.SubjectMatch != nil { + subjectMatch := strings.TrimSpace(*req.SubjectMatch) + if subjectMatch == "" { + return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) + } + existing.SubjectMatch = subjectMatch + } + if req.MatchType != nil { + existing.MatchType = normalizeMatchTypeInternal(*req.MatchType) + } + if err := validateSubjectMatchInternal(existing.MatchType, existing.SubjectMatch); err != nil { + return existing, false, err + } + + roleChanged := false + if req.RoleID != nil { + roleID := strings.TrimSpace(*req.RoleID) + if roleID == "" { + return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) + } + roleChanged = roleID != existing.RoleID + existing.RoleID = roleID + } + if req.EnvironmentID != nil { + environmentID := mo.EmptyableToOption(strings.TrimSpace(*req.EnvironmentID)).ToPointer() + roleChanged = roleChanged || mo.PointerToOption(existing.EnvironmentID).OrEmpty() != mo.PointerToOption(environmentID).OrEmpty() + existing.EnvironmentID = environmentID + } + if req.TokenTTLSeconds != nil { + existing.TokenTTLSeconds = auth.ClampFederatedTokenTTLSeconds(*req.TokenTTLSeconds) + } + if req.ExpiresAt != nil { + existing.ExpiresAt = req.ExpiresAt + } + return existing, roleChanged, nil +} + +func normalizeMatchTypeInternal(matchType string) string { + if strings.EqualFold(strings.TrimSpace(matchType), federatedtypes.MatchTypeGlob) { + return federatedtypes.MatchTypeGlob + } + return federatedtypes.MatchTypeExact +} + +func validateIssuerURLInternal(rawURL string) error { + parsed, err := url.Parse(rawURL) + if err != nil || parsed == nil || parsed.Host == "" || parsed.Scheme != "https" { + return common.Classify(common.ErrFederatedCredentialInvalid, errors.WithStackIf(errors.New("invalid federated credential: issuerUrl must be an HTTPS URL"))) + } + return nil +} + +func validateSubjectMatchInternal(matchType, subjectMatch string) error { + if strings.TrimSpace(subjectMatch) == "" || normalizeMatchTypeInternal(matchType) == federatedtypes.MatchTypeGlob && strings.TrimSpace(subjectMatch) == "*" { + return common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) + } + return nil +} + +func (s *FederatedCredentialService) validateRoleGrantAgainstUserInternal(ctx context.Context, userID, roleID string, environmentID *string) error { + if s.roleService == nil || strings.TrimSpace(userID) == "" { + return nil + } + + user, err := s.userService.GetUserByID(ctx, userID) + if err != nil { + return errors.WrapIf(err, "load user for federated role validation") + } + permissions, err := s.roleService.ResolvePermissions(ctx, user) + if err != nil { + return errors.WrapIf(err, "resolve user permissions") + } + if err := s.roleService.ValidateRoleAssignmentAgainstCaller(ctx, permissions, roleID, environmentID); err != nil { + if errors.Is(err, common.ErrRolePermissionEscalation) { + return common.Classify(common.ErrFederatedCredentialPermissionEscalation, errors.WrapIf(err, "cannot map a federated credential to a role you do not hold")) + } + return common.Classify(common.ErrFederatedCredentialInvalid, errors.WrapIf(err, "invalid federated credential")) + } + return nil +} + +func toFederatedCredentialDTOInternal(credential *FederatedCredential) federatedtypes.FederatedCredential { + if credential == nil { + return federatedtypes.FederatedCredential{} + } + dto := federatedtypes.FederatedCredential{ + ID: credential.ID, + Name: credential.Name, + Description: credential.Description, + Enabled: credential.Enabled, + IssuerURL: credential.IssuerURL, + Audiences: []string(credential.Audiences), + SubjectClaim: credential.SubjectClaim, + SubjectMatch: credential.SubjectMatch, + MatchType: credential.MatchType, + RoleID: credential.RoleID, + EnvironmentID: credential.EnvironmentID, + IdentityUserID: credential.IdentityUserID, + TokenTTLSeconds: credential.TokenTTLSeconds, + LastUsedAt: credential.LastUsedAt, + ExpiresAt: credential.ExpiresAt, + CreatedAt: credential.CreatedAt, + UpdatedAt: credential.UpdatedAt, + } + if credential.IdentityUser != nil { + dto.ServiceUsername = credential.IdentityUser.Username + } + if credential.Role != nil { + dto.RoleName = credential.Role.Name + } + if credential.Environment != nil { + dto.EnvironmentName = credential.Environment.Name + } + return dto +} diff --git a/backend/internal/federated/exchange.go b/backend/internal/federated/exchange.go new file mode 100644 index 0000000000..6e3b99ec50 --- /dev/null +++ b/backend/internal/federated/exchange.go @@ -0,0 +1,365 @@ +package federated + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "log/slog" + "regexp" + "strings" + "time" + + "emperror.dev/errors" + "github.com/coreos/go-oidc/v3/oidc" + "github.com/samber/mo" + + "github.com/getarcaneapp/arcane/backend/v2/internal/common" + "github.com/getarcaneapp/arcane/backend/v2/internal/database" + "github.com/getarcaneapp/arcane/backend/v2/internal/event" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/jwtclaims" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" + federatedtypes "github.com/getarcaneapp/arcane/types/v2/federated" +) + +func (s *FederatedCredentialService) ExchangeToken(ctx context.Context, req federatedtypes.TokenExchangeRequest) (*federatedtypes.FederatedTokenResponse, error) { + claims := jwtclaims.ParseJWTClaims(req.SubjectToken) + issuer := "" + subject := "" + var audiences []string + if claims != nil { + issuer = utils.ToString(jwtclaims.GetByPath(claims, "iss").OrEmpty()) + subject = utils.ToString(jwtclaims.GetByPath(claims, "sub").OrEmpty()) + audiences = utils.UniqueNonEmptyStrings(jwtclaims.StringSliceFromValue(jwtclaims.GetByPath(claims, "aud").OrEmpty())) + } + + logResult := "failure" + logReason := "" + var matchedCredential *FederatedCredential + var matchedUser *common.User + defer func() { + s.logExchangeInternal(ctx, logResult, logReason, issuer, subject, audiences, matchedCredential, matchedUser) + }() + + if req.GrantType != federatedtypes.TokenExchangeGrantType || strings.TrimSpace(req.SubjectToken) == "" { + logReason = "invalid_request" + return nil, common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) + } + switch req.SubjectTokenType { + case federatedtypes.SubjectTokenTypeJWT, federatedtypes.SubjectTokenTypeIDToken: + default: + logReason = "invalid_request" + return nil, common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) + } + if req.RequestedTokenType != "" && req.RequestedTokenType != federatedtypes.RequestedTokenTypeAccessJWT { + logReason = "invalid_request" + return nil, common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) + } + if issuer == "" { + logReason = "missing_issuer" + return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) + } + + var credentials []FederatedCredential + if err := s.db.WithContext(ctx). + Where("issuer_url = ? AND enabled = ?", issuer, true). + Order("created_at ASC"). + Order("id ASC"). + Find(&credentials).Error; err != nil { + logReason = "credential_lookup_failed" + return nil, errors.WrapIf(err, "failed to list federated credentials for issuer") + } + now := time.Now() + active := credentials[:0] + for _, credential := range credentials { + if credential.ExpiresAt == nil || !now.After(*credential.ExpiresAt) { + active = append(active, credential) + } + } + credentials = active + if len(credentials) == 0 { + logReason = "issuer_not_allowed" + return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) + } + + verifiedToken, verifiedClaims, err := s.verifySubjectTokenInternal(ctx, issuer, req.SubjectToken) + if err != nil { + logReason = "token_verification_failed" + return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.WrapIf(err, "invalid federated token grant")) + } + if subject == "" { + subject = utils.ToString(jwtclaims.GetByPath(verifiedClaims, defaultFederatedSubjectClaim).OrEmpty()) + } + if len(audiences) == 0 { + audiences = append([]string{}, verifiedToken.Audience...) + } + + credential := selectMatchingCredentialInternal(credentials, verifiedToken.Audience, verifiedClaims) + if credential == nil { + logReason = "no_matching_credential" + return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) + } + matchedCredential = credential + if err := s.recordTokenReplayGuardInternal(ctx, issuer, req.SubjectToken, verifiedClaims, verifiedToken.Expiry); err != nil { + logReason = "token_replay_rejected" + return nil, err + } + + user, err := s.userService.GetUserByID(ctx, credential.IdentityUserID) + if err != nil { + logReason = "identity_user_missing" + return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.WrapIf(err, "invalid federated token grant")) + } + matchedUser = user + + tokenPair, err := s.authService.IssueFederatedToken(ctx, user, credential.ID, credential.TokenTTLSeconds) + if err != nil { + logReason = "token_issue_failed" + return nil, err + } + + go func() { + bgCtx := context.WithoutCancel(ctx) + now := time.Now() + cutoff := now.Add(-federatedCredentialLastUsedWriteWindow) + if err := s.db.WithContext(bgCtx). + Model(&FederatedCredential{}). + Where("id = ? AND (last_used_at IS NULL OR last_used_at < ?)", credential.ID, cutoff). + Update("last_used_at", now).Error; err != nil { + slog.WarnContext(bgCtx, "failed to update federated credential last_used_at", "credential_id", credential.ID, "error", err) + } + }() + + logResult = "success" + logReason = "matched" + return &federatedtypes.FederatedTokenResponse{ + AccessToken: tokenPair.AccessToken, + TokenType: "Bearer", + ExpiresIn: max(int(time.Until(tokenPair.ExpiresAt).Seconds()), 0), + IssuedTokenType: federatedtypes.IssuedTokenTypeAccessToken, + }, nil +} + +func (s *FederatedCredentialService) verifySubjectTokenInternal(ctx context.Context, issuer, rawToken string) (*oidc.IDToken, map[string]any, error) { + keySet, err := s.keySetForIssuerInternal(ctx, issuer) + if err != nil { + return nil, nil, err + } + + providerCtx := oidc.ClientContext(ctx, s.httpClient) + verifier := oidc.NewVerifier(issuer, keySet, &oidc.Config{ + SkipClientIDCheck: true, + SupportedSigningAlgs: oidcjwk.SupportedSigningAlgs(), + }) + idToken, err := verifier.Verify(providerCtx, rawToken) + if err != nil { + return nil, nil, err + } + + claims := map[string]any{} + if err := idToken.Claims(&claims); err != nil { + return nil, nil, err + } + return idToken, claims, nil +} + +func (s *FederatedCredentialService) recordTokenReplayGuardInternal(ctx context.Context, issuer, rawToken string, claims map[string]any, expiresAt time.Time) error { + if expiresAt.IsZero() || time.Now().After(expiresAt) { + return common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) + } + + now := time.Now() + if err := s.db.WithContext(ctx). + Where("expires_at < ?", now). + Delete(&FederatedTokenReplay{}).Error; err != nil { + return errors.WrapIf(err, "failed to prune federated token replay records") + } + + tokenID := strings.TrimSpace(utils.ToString(jwtclaims.GetByPath(claims, "jti").OrEmpty())) + tokenKind := "jti" + if tokenID == "" { + tokenID = rawToken + tokenKind = "token" + } + sum := sha256.Sum256([]byte(issuer + "\x00" + tokenKind + "\x00" + tokenID)) + replay := FederatedTokenReplay{ + TokenHash: hex.EncodeToString(sum[:]), + IssuerURL: issuer, + ExpiresAt: expiresAt, + } + if err := s.db.WithContext(ctx).Create(&replay).Error; err != nil { + message := strings.ToLower(err.Error()) + if strings.Contains(message, "unique") || strings.Contains(message, "duplicate key") { + return common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) + } + return errors.WrapIf(err, "failed to record federated token replay guard") + } + return nil +} + +func (s *FederatedCredentialService) keySetForIssuerInternal(ctx context.Context, issuer string) (oidc.KeySet, error) { + s.providerMu.RLock() + if keySet := s.keySets[issuer]; keySet != nil { + s.providerMu.RUnlock() + return keySet, nil + } + s.providerMu.RUnlock() + + value, err, _ := s.providerGroup.Do(issuer, func() (any, error) { + providerCtx := oidc.ClientContext(context.WithoutCancel(ctx), s.httpClient) + provider, err := oidc.NewProvider(providerCtx, issuer) + if err != nil { + return nil, errors.WrapIf(err, "failed to discover federated issuer") + } + + var metadata struct { + JWKSURL string `json:"jwks_uri"` + } + if err := provider.Claims(&metadata); err != nil { + return nil, errors.WrapIf(err, "failed to read federated issuer metadata") + } + if metadata.JWKSURL == "" { + return nil, errors.New("federated issuer metadata is missing jwks_uri") + } + if s.keySetManager == nil { + return nil, errors.New("JWK set manager is not configured") + } + + keySet, err := s.keySetManager.KeySet(context.WithoutCancel(ctx), s.httpClient, metadata.JWKSURL) + if err != nil { + return nil, errors.WrapIf(err, "failed to configure federated issuer JWK set") + } + s.providerMu.Lock() + s.keySets[issuer] = keySet + s.providerMu.Unlock() + return keySet, nil + }) + if err != nil { + return nil, err + } + + keySet, ok := value.(oidc.KeySet) + if !ok || keySet == nil { + return nil, errors.New("federated issuer discovery returned invalid key set") + } + return keySet, nil +} + +func selectMatchingCredentialInternal(credentials []FederatedCredential, tokenAudiences []string, claims map[string]any) *FederatedCredential { + for i := range credentials { + credential := &credentials[i] + if credentialMatchesTokenInternal(credential, tokenAudiences, claims) { + return credential + } + } + return nil +} + +func credentialMatchesTokenInternal(credential *FederatedCredential, tokenAudiences []string, claims map[string]any) bool { + audiences := make(map[string]struct{}, len(credential.Audiences)) + for _, audience := range credential.Audiences { + if audience = strings.TrimSpace(audience); audience != "" { + audiences[audience] = struct{}{} + } + } + audienceMatched := false + for _, audience := range tokenAudiences { + if _, audienceMatched = audiences[audience]; audienceMatched { + break + } + } + if !audienceMatched { + return false + } + + subjectClaim := strings.TrimSpace(credential.SubjectClaim) + if subjectClaim == "" { + subjectClaim = defaultFederatedSubjectClaim + } + subject := utils.ToString(jwtclaims.GetByPath(claims, subjectClaim).OrEmpty()) + if subject == "" { + return false + } + if normalizeMatchTypeInternal(credential.MatchType) != federatedtypes.MatchTypeGlob { + return subject == credential.SubjectMatch + } + + var expression strings.Builder + expression.WriteString("^") + for _, character := range credential.SubjectMatch { + switch character { + case '*': + expression.WriteString(".*") + case '?': + expression.WriteByte('.') + default: + expression.WriteString(regexp.QuoteMeta(string(character))) + } + } + expression.WriteString("$") + matched, err := regexp.MatchString(expression.String(), subject) + return err == nil && matched +} + +func (s *FederatedCredentialService) logExchangeInternal(ctx context.Context, result, reason, issuer, subject string, audiences []string, credential *FederatedCredential, user *common.User) { + credentialID := "" + credentialName := "" + if credential != nil { + credentialID = credential.ID + credentialName = credential.Name + } + slog.InfoContext(ctx, "Federated credential token exchange", + "result", result, + "reason", reason, + "issuer", issuer, + "subject", subject, + "audiences", audiences, + "credential_id", credentialID, + ) + + if s.eventService == nil { + return + } + + metadata := database.JSON{ + "action": "federated_token_exchange", + "result": result, + "reason": reason, + "issuer": issuer, + "subject": subject, + "audiences": audiences, + "credentialId": credentialID, + } + + userID := "" + username := "" + if user != nil { + userID = user.ID + username = user.Username + } + severity := event.EventSeverityInfo + title := "Federated credential token exchange" + if result != "success" { + severity = event.EventSeverityWarning + title = "Federated credential token exchange rejected" + } + + go func() { + bgCtx := context.WithoutCancel(ctx) + _, err := s.eventService.CreateEvent(bgCtx, event.CreateEventRequest{ + Type: event.EventTypeFederatedExchange, + Severity: severity, + Title: title, + Description: "Workload identity federation token exchange", + ResourceType: mo.EmptyableToOption(strings.TrimSpace("federated_credential")).ToPointer(), + ResourceID: mo.EmptyableToOption(strings.TrimSpace(credentialID)).ToPointer(), + ResourceName: mo.EmptyableToOption(strings.TrimSpace(credentialName)).ToPointer(), + UserID: mo.EmptyableToOption(strings.TrimSpace(userID)).ToPointer(), + Username: mo.EmptyableToOption(strings.TrimSpace(username)).ToPointer(), + Metadata: metadata, + }) + if err != nil { + slog.WarnContext(bgCtx, "failed to audit federated credential token exchange", "error", err) + } + }() +} diff --git a/backend/internal/federated/service.go b/backend/internal/federated/service.go index 680002baf9..26564e2df6 100644 --- a/backend/internal/federated/service.go +++ b/backend/internal/federated/service.go @@ -1,22 +1,14 @@ package federated import ( - "github.com/getarcaneapp/arcane/backend/v2/internal/session" - "context" - "crypto/sha256" - "encoding/hex" - "log/slog" "net/http" - "net/url" - "regexp" "strings" "sync" "time" "uuid" "emperror.dev/errors" - "github.com/coreos/go-oidc/v3/oidc" "github.com/samber/mo" "golang.org/x/sync/singleflight" @@ -27,14 +19,13 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/database" "github.com/getarcaneapp/arcane/backend/v2/internal/event" "github.com/getarcaneapp/arcane/backend/v2/internal/role" + "github.com/getarcaneapp/arcane/backend/v2/internal/session" "github.com/getarcaneapp/arcane/backend/v2/internal/settings" "github.com/getarcaneapp/arcane/backend/v2/internal/user" "github.com/getarcaneapp/arcane/backend/v2/pkg/pagination" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/dbutil" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/httpx" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/jwtclaims" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/mldsajose" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" federatedtypes "github.com/getarcaneapp/arcane/types/v2/federated" ) @@ -51,8 +42,9 @@ type FederatedCredentialService struct { eventService *event.EventService roleService *role.RoleService httpClient *http.Client + keySetManager *oidcjwk.KeySetManager providerMu sync.RWMutex - keySets map[string]*mldsajose.KeySet + keySets map[string]oidc.KeySet providerGroup singleflight.Group } @@ -63,6 +55,7 @@ func NewFederatedCredentialService( settingsService *settings.SettingsService, eventService *event.EventService, httpClient *http.Client, + keySetManager *oidcjwk.KeySetManager, ) *FederatedCredentialService { if httpClient == nil { httpClient = httpx.NewHTTPClientWithTimeout(15 * time.Second) @@ -75,7 +68,8 @@ func NewFederatedCredentialService( settingsService: settingsService, eventService: eventService, httpClient: httpClient, - keySets: make(map[string]*mldsajose.KeySet), + keySetManager: keySetManager, + keySets: make(map[string]oidc.KeySet), } } @@ -84,83 +78,6 @@ func (s *FederatedCredentialService) WithRoleService(roleService *role.RoleServi return s } -func (s *FederatedCredentialService) ExchangeToken(ctx context.Context, req federatedtypes.TokenExchangeRequest) (*federatedtypes.FederatedTokenResponse, error) { - issuer, subject, audiences := unverifiedTokenExchangeMetadataInternal(req.SubjectToken) - logResult := "failure" - logReason := "" - var matchedCredential *FederatedCredential - var matchedUser *common.User - defer func() { - s.logExchangeInternal(ctx, logResult, logReason, issuer, subject, audiences, matchedCredential, matchedUser) - }() - - if err := validateTokenExchangeRequestInternal(req); err != nil { - logReason = "invalid_request" - return nil, err - } - if issuer == "" { - logReason = "missing_issuer" - return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) - } - - credentials, err := s.listEnabledCredentialsForIssuerInternal(ctx, issuer) - if err != nil { - logReason = "credential_lookup_failed" - return nil, err - } - if len(credentials) == 0 { - logReason = "issuer_not_allowed" - return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) - } - - verifiedToken, verifiedClaims, err := s.verifySubjectTokenInternal(ctx, issuer, req.SubjectToken) - if err != nil { - logReason = "token_verification_failed" - return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.WrapIf(err, "invalid federated token grant")) - } - if subject == "" { - subject = stringClaimByPathInternal(verifiedClaims, defaultFederatedSubjectClaim) - } - if len(audiences) == 0 { - audiences = append([]string{}, verifiedToken.Audience...) - } - - credential := selectMatchingCredentialInternal(credentials, verifiedToken.Audience, verifiedClaims) - if credential == nil { - logReason = "no_matching_credential" - return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) - } - matchedCredential = credential - if err := s.recordTokenReplayGuardInternal(ctx, issuer, req.SubjectToken, verifiedClaims, verifiedToken.Expiry); err != nil { - logReason = "token_replay_rejected" - return nil, err - } - - user, err := s.userService.GetUserByID(ctx, credential.IdentityUserID) - if err != nil { - logReason = "identity_user_missing" - return nil, common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.WrapIf(err, "invalid federated token grant")) - } - matchedUser = user - - tokenPair, err := s.authService.IssueFederatedToken(ctx, user, credential.ID, credential.TokenTTLSeconds) - if err != nil { - logReason = "token_issue_failed" - return nil, err - } - - s.markCredentialUsedAsyncInternal(ctx, credential.ID) - logResult = "success" - logReason = "matched" - - return &federatedtypes.FederatedTokenResponse{ - AccessToken: tokenPair.AccessToken, - TokenType: "Bearer", - ExpiresIn: max(int(time.Until(tokenPair.ExpiresAt).Seconds()), 0), - IssuedTokenType: federatedtypes.IssuedTokenTypeAccessToken, - }, nil -} - func (s *FederatedCredentialService) Create(ctx context.Context, callerUserID string, req federatedtypes.CreateFederatedCredential) (*federatedtypes.FederatedCredential, error) { normalized, err := normalizeCreateFederatedCredentialInternal(req) if err != nil { @@ -209,7 +126,6 @@ func (s *FederatedCredentialService) Create(ctx context.Context, callerUserID st if err := tx.Create(&assignment).Error; err != nil { return errors.WrapIf(err, "failed to create federated role assignment") } - return nil }) if err != nil { @@ -267,7 +183,7 @@ func (s *FederatedCredentialService) Get(ctx context.Context, id string) (*feder return new(toFederatedCredentialDTOInternal(&credential)), nil } -func (s *FederatedCredentialService) Update(ctx context.Context, callerUserID string, id string, req federatedtypes.UpdateFederatedCredential) (*federatedtypes.FederatedCredential, error) { +func (s *FederatedCredentialService) Update(ctx context.Context, callerUserID, id string, req federatedtypes.UpdateFederatedCredential) (*federatedtypes.FederatedCredential, error) { var credential FederatedCredential if err := s.db.WithContext(ctx).Where("id = ?", id).First(&credential).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { @@ -292,8 +208,11 @@ func (s *FederatedCredentialService) Update(ctx context.Context, callerUserID st return errors.WrapIf(err, "failed to update federated credential") } if revokeActiveSessions { - if err := revokeFederatedCredentialSessionsInternal(tx, updated.ID); err != nil { - return err + now := time.Now() + if err := tx.Model(&session.UserSession{}). + Where("federated_credential_id = ? AND revoked_at IS NULL", updated.ID). + Updates(map[string]any{"revoked_at": now, "updated_at": now}).Error; err != nil { + return errors.WrapIf(err, "failed to revoke federated credential sessions") } } if roleChanged { @@ -319,7 +238,6 @@ func (s *FederatedCredentialService) Update(ctx context.Context, callerUserID st if roleChanged && s.roleService != nil { s.roleService.InvalidateUser(updated.IdentityUserID) } - return s.Get(ctx, id) } @@ -349,519 +267,3 @@ func (s *FederatedCredentialService) Delete(ctx context.Context, id string) erro } return nil } - -func validateTokenExchangeRequestInternal(req federatedtypes.TokenExchangeRequest) error { - if req.GrantType != federatedtypes.TokenExchangeGrantType { - return common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) - } - if strings.TrimSpace(req.SubjectToken) == "" { - return common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) - } - switch req.SubjectTokenType { - case federatedtypes.SubjectTokenTypeJWT, federatedtypes.SubjectTokenTypeIDToken: - default: - return common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) - } - if req.RequestedTokenType != "" && req.RequestedTokenType != federatedtypes.RequestedTokenTypeAccessJWT { - return common.Classify(common.ErrFederatedCredentialInvalidRequest, errors.New("invalid federated token exchange request")) - } - return nil -} - -func (s *FederatedCredentialService) listEnabledCredentialsForIssuerInternal(ctx context.Context, issuer string) ([]FederatedCredential, error) { - var credentials []FederatedCredential - if err := s.db.WithContext(ctx). - Where("issuer_url = ? AND enabled = ?", issuer, true). - Order("created_at ASC"). - Order("id ASC"). - Find(&credentials).Error; err != nil { - return nil, errors.WrapIf(err, "failed to list federated credentials for issuer") - } - - now := time.Now() - active := credentials[:0] - for _, credential := range credentials { - if credential.ExpiresAt != nil && now.After(*credential.ExpiresAt) { - continue - } - active = append(active, credential) - } - return active, nil -} - -func (s *FederatedCredentialService) verifySubjectTokenInternal(ctx context.Context, issuer string, rawToken string) (*oidc.IDToken, map[string]any, error) { - keySet, err := s.keySetForIssuerInternal(ctx, issuer) - if err != nil { - return nil, nil, err - } - - providerCtx := oidc.ClientContext(ctx, s.httpClient) - verifier := oidc.NewVerifier(issuer, keySet, &oidc.Config{ - SkipClientIDCheck: true, - SupportedSigningAlgs: mldsajose.SupportedSigningAlgs(), - }) - idToken, err := verifier.Verify(providerCtx, rawToken) - if err != nil { - return nil, nil, err - } - - claims := map[string]any{} - if err := idToken.Claims(&claims); err != nil { - return nil, nil, err - } - return idToken, claims, nil -} - -func (s *FederatedCredentialService) recordTokenReplayGuardInternal(ctx context.Context, issuer string, rawToken string, claims map[string]any, expiresAt time.Time) error { - if expiresAt.IsZero() || time.Now().After(expiresAt) { - return common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) - } - - now := time.Now() - if err := s.db.WithContext(ctx). - Where("expires_at < ?", now). - Delete(&FederatedTokenReplay{}).Error; err != nil { - return errors.WrapIf(err, "failed to prune federated token replay records") - } - - replay := FederatedTokenReplay{ - TokenHash: federatedTokenReplayHashInternal(issuer, rawToken, claims), - IssuerURL: issuer, - ExpiresAt: expiresAt, - } - if err := s.db.WithContext(ctx).Create(&replay).Error; err != nil { - if isUniqueConstraintErrorInternal(err) { - return common.Classify(common.ErrFederatedCredentialInvalidGrant, errors.New("invalid federated token grant")) - } - return errors.WrapIf(err, "failed to record federated token replay guard") - } - return nil -} - -func federatedTokenReplayHashInternal(issuer string, rawToken string, claims map[string]any) string { - tokenID := strings.TrimSpace(stringClaimByPathInternal(claims, "jti")) - tokenKind := "jti" - if tokenID == "" { - tokenID = rawToken - tokenKind = "token" - } - - sum := sha256.Sum256([]byte(issuer + "\x00" + tokenKind + "\x00" + tokenID)) - return hex.EncodeToString(sum[:]) -} - -func isUniqueConstraintErrorInternal(err error) bool { - if err == nil { - return false - } - msg := strings.ToLower(err.Error()) - return strings.Contains(msg, "unique") || strings.Contains(msg, "duplicate key") -} - -func (s *FederatedCredentialService) keySetForIssuerInternal(ctx context.Context, issuer string) (*mldsajose.KeySet, error) { - s.providerMu.RLock() - if keySet := s.keySets[issuer]; keySet != nil { - s.providerMu.RUnlock() - return keySet, nil - } - s.providerMu.RUnlock() - - v, err, _ := s.providerGroup.Do(issuer, func() (any, error) { - providerCtx := oidc.ClientContext(context.WithoutCancel(ctx), s.httpClient) - provider, err := oidc.NewProvider(providerCtx, issuer) - if err != nil { - return nil, errors.WrapIf(err, "failed to discover federated issuer") - } - - var meta struct { - JWKSURL string `json:"jwks_uri"` - } - if err := provider.Claims(&meta); err != nil { - return nil, errors.WrapIf(err, "failed to read federated issuer metadata") - } - if meta.JWKSURL == "" { - return nil, errors.New("federated issuer metadata is missing jwks_uri") - } - - keySet := mldsajose.NewKeySet(providerCtx, meta.JWKSURL) - s.providerMu.Lock() - s.keySets[issuer] = keySet - s.providerMu.Unlock() - return keySet, nil - }) - if err != nil { - return nil, err - } - - keySet, ok := v.(*mldsajose.KeySet) - if !ok || keySet == nil { - return nil, errors.New("federated issuer discovery returned invalid key set") - } - return keySet, nil -} - -func selectMatchingCredentialInternal(credentials []FederatedCredential, tokenAudiences []string, claims map[string]any) *FederatedCredential { - for i := range credentials { - credential := &credentials[i] - if !audienceMatchesInternal(tokenAudiences, credential.Audiences) { - continue - } - subjectClaim := strings.TrimSpace(credential.SubjectClaim) - if subjectClaim == "" { - subjectClaim = defaultFederatedSubjectClaim - } - subject := stringClaimByPathInternal(claims, subjectClaim) - if !subjectMatchesInternal(credential.MatchType, credential.SubjectMatch, subject) { - continue - } - return credential - } - return nil -} - -func audienceMatchesInternal(tokenAudiences, credentialAudiences []string) bool { - allowed := make(map[string]struct{}, len(credentialAudiences)) - for _, audience := range credentialAudiences { - audience = strings.TrimSpace(audience) - if audience != "" { - allowed[audience] = struct{}{} - } - } - for _, audience := range tokenAudiences { - if _, ok := allowed[audience]; ok { - return true - } - } - return false -} - -func subjectMatchesInternal(matchType, pattern, subject string) bool { - if subject == "" { - return false - } - switch normalizeMatchTypeInternal(matchType) { - case federatedtypes.MatchTypeGlob: - return anchoredGlobMatchesInternal(pattern, subject) - default: - return subject == pattern - } -} - -func anchoredGlobMatchesInternal(pattern, value string) bool { - var b strings.Builder - b.WriteString("^") - for _, r := range pattern { - switch r { - case '*': - b.WriteString(".*") - case '?': - b.WriteByte('.') - default: - b.WriteString(regexp.QuoteMeta(string(r))) - } - } - b.WriteString("$") - matched, err := regexp.MatchString(b.String(), value) - return err == nil && matched -} - -func unverifiedTokenExchangeMetadataInternal(rawToken string) (string, string, []string) { - claims := jwtclaims.ParseJWTClaims(rawToken) - if claims == nil { - return "", "", nil - } - return stringClaimByPathInternal(claims, "iss"), stringClaimByPathInternal(claims, "sub"), utils.UniqueNonEmptyStrings(jwtclaims.StringSliceFromValue(jwtclaims.GetByPath(claims, "aud").OrEmpty())) -} - -func stringClaimByPathInternal(claims map[string]any, path string) string { - value, ok := jwtclaims.GetByPath(claims, path).Get() - if !ok { - return "" - } - return utils.ToString(value) -} - -func normalizeCreateFederatedCredentialInternal(req federatedtypes.CreateFederatedCredential) (federatedtypes.CreateFederatedCredential, error) { - req.Name = strings.TrimSpace(req.Name) - req.IssuerURL = strings.TrimRight(strings.TrimSpace(req.IssuerURL), "/") - req.SubjectClaim = strings.TrimSpace(req.SubjectClaim) - req.SubjectMatch = strings.TrimSpace(req.SubjectMatch) - req.MatchType = normalizeMatchTypeInternal(req.MatchType) - req.Audiences = utils.UniqueNonEmptyStrings(req.Audiences) - req.EnvironmentID = mo.EmptyableToOption(strings.TrimSpace(mo.PointerToOption(req.EnvironmentID).OrEmpty())).ToPointer() - req.TokenTTLSeconds = auth.ClampFederatedTokenTTLSeconds(req.TokenTTLSeconds) - - if req.SubjectClaim == "" { - req.SubjectClaim = defaultFederatedSubjectClaim - } - if req.Name == "" || req.SubjectMatch == "" || req.RoleID == "" || len(req.Audiences) == 0 { - return req, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - if err := validateIssuerURLInternal(req.IssuerURL); err != nil { - return req, err - } - if err := validateSubjectMatchInternal(req.MatchType, req.SubjectMatch); err != nil { - return req, err - } - return req, nil -} - -func applyFederatedCredentialUpdateInternal(existing FederatedCredential, req federatedtypes.UpdateFederatedCredential) (FederatedCredential, bool, error) { - if req.Name != nil { - name := strings.TrimSpace(*req.Name) - if name == "" { - return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - existing.Name = name - } - if req.Description != nil { - existing.Description = req.Description - } - if req.Enabled != nil { - existing.Enabled = *req.Enabled - } - if req.IssuerURL != nil { - issuerURL := strings.TrimRight(strings.TrimSpace(*req.IssuerURL), "/") - if err := validateIssuerURLInternal(issuerURL); err != nil { - return existing, false, err - } - existing.IssuerURL = issuerURL - } - if req.Audiences != nil { - audiences := utils.UniqueNonEmptyStrings(req.Audiences) - if len(audiences) == 0 { - return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - existing.Audiences = audiences - } - if req.SubjectClaim != nil { - subjectClaim := strings.TrimSpace(*req.SubjectClaim) - if subjectClaim == "" { - subjectClaim = defaultFederatedSubjectClaim - } - existing.SubjectClaim = subjectClaim - } - if req.SubjectMatch != nil { - subjectMatch := strings.TrimSpace(*req.SubjectMatch) - if subjectMatch == "" { - return existing, false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - existing.SubjectMatch = subjectMatch - } - if req.MatchType != nil { - existing.MatchType = normalizeMatchTypeInternal(*req.MatchType) - } - if err := validateSubjectMatchInternal(existing.MatchType, existing.SubjectMatch); err != nil { - return existing, false, err - } - roleChanged, err := applyFederatedRoleScopeUpdateInternal(&existing, req.RoleID, req.EnvironmentID) - if err != nil { - return existing, false, err - } - if req.TokenTTLSeconds != nil { - existing.TokenTTLSeconds = auth.ClampFederatedTokenTTLSeconds(*req.TokenTTLSeconds) - } - if req.ExpiresAt != nil { - existing.ExpiresAt = req.ExpiresAt - } - return existing, roleChanged, nil -} - -func applyFederatedRoleScopeUpdateInternal(existing *FederatedCredential, roleID *string, environmentID *string) (bool, error) { - if existing == nil { - return false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - - roleChanged := false - if roleID != nil { - normalizedRoleID := strings.TrimSpace(*roleID) - if normalizedRoleID == "" { - return false, common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - roleChanged = normalizedRoleID != existing.RoleID - existing.RoleID = normalizedRoleID - } - if environmentID != nil { - normalized := mo.EmptyableToOption(strings.TrimSpace(*environmentID)).ToPointer() - roleChanged = roleChanged || mo.PointerToOption(existing.EnvironmentID).OrEmpty() != mo.PointerToOption(normalized).OrEmpty() - existing.EnvironmentID = normalized - } - return roleChanged, nil -} - -func normalizeMatchTypeInternal(matchType string) string { - switch strings.ToLower(strings.TrimSpace(matchType)) { - case federatedtypes.MatchTypeGlob: - return federatedtypes.MatchTypeGlob - default: - return federatedtypes.MatchTypeExact - } -} - -func validateIssuerURLInternal(rawURL string) error { - parsed, err := url.Parse(rawURL) - if err != nil || parsed == nil || parsed.Host == "" || parsed.Scheme != "https" { - return common.Classify(common.ErrFederatedCredentialInvalid, errors.WithStackIf(errors.New("invalid federated credential: issuerUrl must be an HTTPS URL"))) - } - return nil -} - -func validateSubjectMatchInternal(matchType, subjectMatch string) error { - if strings.TrimSpace(subjectMatch) == "" { - return common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - if normalizeMatchTypeInternal(matchType) == federatedtypes.MatchTypeGlob && strings.TrimSpace(subjectMatch) == "*" { - return common.Classify(common.ErrFederatedCredentialInvalid, errors.New("invalid federated credential")) - } - return nil -} - -func (s *FederatedCredentialService) validateRoleGrantAgainstUserInternal(ctx context.Context, userID, roleID string, environmentID *string) error { - if s.roleService == nil || strings.TrimSpace(userID) == "" { - return nil - } - - user, err := s.userService.GetUserByID(ctx, userID) - if err != nil { - return errors.WrapIf(err, "load user for federated role validation") - } - - ps, err := s.roleService.ResolvePermissions(ctx, user) - if err != nil { - return errors.WrapIf(err, "resolve user permissions") - } - - if err := s.roleService.ValidateRoleAssignmentAgainstCaller(ctx, ps, roleID, environmentID); err != nil { - if errors.Is(err, common.ErrRolePermissionEscalation) { - return common.Classify(common.ErrFederatedCredentialPermissionEscalation, errors.WrapIf(err, "cannot map a federated credential to a role you do not hold")) - } - return common.Classify(common.ErrFederatedCredentialInvalid, errors.WrapIf(err, "invalid federated credential")) - } - return nil -} - -func revokeFederatedCredentialSessionsInternal(tx *gorm.DB, credentialID string) error { - if tx == nil || strings.TrimSpace(credentialID) == "" { - return nil - } - - now := time.Now() - if err := tx.Model(&session.UserSession{}). - Where("federated_credential_id = ? AND revoked_at IS NULL", credentialID). - Updates(map[string]any{"revoked_at": now, "updated_at": now}).Error; err != nil { - return errors.WrapIf(err, "failed to revoke federated credential sessions") - } - return nil -} - -func toFederatedCredentialDTOInternal(credential *FederatedCredential) federatedtypes.FederatedCredential { - if credential == nil { - return federatedtypes.FederatedCredential{} - } - dto := federatedtypes.FederatedCredential{ - ID: credential.ID, - Name: credential.Name, - Description: credential.Description, - Enabled: credential.Enabled, - IssuerURL: credential.IssuerURL, - Audiences: []string(credential.Audiences), - SubjectClaim: credential.SubjectClaim, - SubjectMatch: credential.SubjectMatch, - MatchType: credential.MatchType, - RoleID: credential.RoleID, - EnvironmentID: credential.EnvironmentID, - IdentityUserID: credential.IdentityUserID, - TokenTTLSeconds: credential.TokenTTLSeconds, - LastUsedAt: credential.LastUsedAt, - ExpiresAt: credential.ExpiresAt, - CreatedAt: credential.CreatedAt, - UpdatedAt: credential.UpdatedAt, - } - if credential.IdentityUser != nil { - dto.ServiceUsername = credential.IdentityUser.Username - } - if credential.Role != nil { - dto.RoleName = credential.Role.Name - } - if credential.Environment != nil { - dto.EnvironmentName = credential.Environment.Name - } - return dto -} - -func (s *FederatedCredentialService) markCredentialUsedAsyncInternal(ctx context.Context, credentialID string) { - go func() { - bgCtx := context.WithoutCancel(ctx) - now := time.Now() - cutoff := now.Add(-federatedCredentialLastUsedWriteWindow) - if err := s.db.WithContext(bgCtx). - Model(&FederatedCredential{}). - Where("id = ? AND (last_used_at IS NULL OR last_used_at < ?)", credentialID, cutoff). - Update("last_used_at", now).Error; err != nil { - slog.WarnContext(bgCtx, "failed to update federated credential last_used_at", "credential_id", credentialID, "error", err) - } - }() -} - -func (s *FederatedCredentialService) logExchangeInternal(ctx context.Context, result, reason, issuer, subject string, audiences []string, credential *FederatedCredential, user *common.User) { - credentialID := "" - credentialName := "" - if credential != nil { - credentialID = credential.ID - credentialName = credential.Name - } - slog.InfoContext(ctx, "Federated credential token exchange", - "result", result, - "reason", reason, - "issuer", issuer, - "subject", subject, - "audiences", audiences, - "credential_id", credentialID, - ) - - if s.eventService == nil { - return - } - - metadata := database.JSON{ - "action": "federated_token_exchange", - "result": result, - "reason": reason, - "issuer": issuer, - "subject": subject, - "audiences": audiences, - "credentialId": credentialID, - } - - userID := "" - username := "" - if user != nil { - userID = user.ID - username = user.Username - } - severity := event.EventSeverityInfo - title := "Federated credential token exchange" - if result != "success" { - severity = event.EventSeverityWarning - title = "Federated credential token exchange rejected" - } - - go func() { - bgCtx := context.WithoutCancel(ctx) - _, err := s.eventService.CreateEvent(bgCtx, event.CreateEventRequest{ - Type: event.EventTypeFederatedExchange, - Severity: severity, - Title: title, - Description: "Workload identity federation token exchange", - ResourceType: mo.EmptyableToOption(strings.TrimSpace("federated_credential")).ToPointer(), - ResourceID: mo.EmptyableToOption(strings.TrimSpace(credentialID)).ToPointer(), - ResourceName: mo.EmptyableToOption(strings.TrimSpace(credentialName)).ToPointer(), - UserID: mo.EmptyableToOption(strings.TrimSpace(userID)).ToPointer(), - Username: mo.EmptyableToOption(strings.TrimSpace(username)).ToPointer(), - Metadata: metadata, - }) - if err != nil { - slog.WarnContext(bgCtx, "failed to audit federated credential token exchange", "error", err) - } - }() -} diff --git a/backend/internal/federated/service_test.go b/backend/internal/federated/service_test.go index 9f09eaa4c9..a606170088 100644 --- a/backend/internal/federated/service_test.go +++ b/backend/internal/federated/service_test.go @@ -3,12 +3,9 @@ package federated import ( "context" "crypto/mldsa" - "crypto/rand" - "crypto/rsa" "encoding/base64" "encoding/json" "fmt" - "math/big" "net/http" "net/http/httptest" "strings" @@ -16,7 +13,9 @@ import ( "time" "emperror.dev/errors" - "github.com/golang-jwt/jwt/v5" + "github.com/lestrrat-go/jwx/v4/jwa" + "github.com/lestrrat-go/jwx/v4/jws" + "github.com/lestrrat-go/jwx/v4/jwt" "github.com/libtnb/sqlite" "github.com/stretchr/testify/require" "go.uber.org/fx/fxtest" @@ -33,13 +32,14 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/settings" "github.com/getarcaneapp/arcane/backend/v2/internal/user" "github.com/getarcaneapp/arcane/backend/v2/pkg/authz" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" federatedtypes "github.com/getarcaneapp/arcane/types/v2/federated" "github.com/stretchr/testify/assert" ) type federatedTestIssuerInternal struct { IssuerURL string - private *rsa.PrivateKey + private *mldsa.PrivateKey keyID string server *httptest.Server } @@ -47,7 +47,7 @@ type federatedTestIssuerInternal struct { func newFederatedTestIssuerInternal(t *testing.T) *federatedTestIssuerInternal { t.Helper() - privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + privateKey, err := mldsa.GenerateKey(mldsa.MLDSA87()) require.NoError(t, err) issuer := &federatedTestIssuerInternal{ @@ -63,22 +63,21 @@ func newFederatedTestIssuerInternal(t *testing.T) *federatedTestIssuerInternal { "authorization_endpoint": issuer.IssuerURL + "/authorize", "token_endpoint": issuer.IssuerURL + "/token", "subject_types_supported": []string{"public"}, - "id_token_signing_alg_values_supported": []string{"RS256"}, + "id_token_signing_alg_values_supported": []string{jwa.MLDSA87().String()}, })) { return } }) mux.HandleFunc("/jwks", func(w http.ResponseWriter, _ *http.Request) { - pub := privateKey.PublicKey + pub := privateKey.PublicKey() if !assert.NoError(t, json.NewEncoder(w).Encode(map[string]any{ "keys": []map[string]any{ { - "kty": "RSA", + "kty": "AKP", "use": "sig", "kid": issuer.keyID, - "alg": "RS256", - "n": base64.RawURLEncoding.EncodeToString(pub.N.Bytes()), - "e": base64.RawURLEncoding.EncodeToString(big.NewInt(int64(pub.E)).Bytes()), + "alg": jwa.MLDSA87().String(), + "pub": base64.RawURLEncoding.EncodeToString(pub.Bytes()), }, }, })) { @@ -97,20 +96,20 @@ func (i *federatedTestIssuerInternal) tokenInternal(t *testing.T, subject string t.Helper() now := time.Now() - claims := jwt.MapClaims{ - "iss": i.IssuerURL, - "sub": subject, - "aud": audience, - "iat": now.Unix(), - "nbf": now.Add(-time.Minute).Unix(), - "exp": now.Add(5 * time.Minute).Unix(), - } - - token := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) - token.Header["kid"] = i.keyID - signed, err := token.SignedString(i.private) + token, err := jwt.NewBuilder(). + Issuer(i.IssuerURL). + Subject(subject). + Audience(audience). + IssuedAt(now). + NotBefore(now.Add(-time.Minute)). + Expiration(now.Add(5 * time.Minute)). + Build() + require.NoError(t, err) + headers := jws.NewHeaders() + require.NoError(t, headers.Set(jws.KeyIDKey, i.keyID)) + signed, err := jwt.Sign(token, jwt.WithKey(jwa.MLDSA87(), i.private, jws.WithProtectedHeaders(headers))) require.NoError(t, err) - return signed + return string(signed) } func setupFederatedCredentialServiceTestDBInternal(t *testing.T) *database.DB { @@ -155,7 +154,13 @@ func setupFederatedCredentialServiceInternal(t *testing.T, issuer *federatedTest JWTRefreshExpiry: 24 * time.Hour, }, nil).WithSigningKey(signingKey) - service := NewFederatedCredentialService(db, authSvc, userSvc, settingsSvc, eventSvc, issuer.server.Client()).WithRoleService(roleSvc) + keySetManager := oidcjwk.NewKeySetManager(t.Context()) + t.Cleanup(func() { + shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + require.NoError(t, keySetManager.Shutdown(shutdownCtx)) + }) + service := NewFederatedCredentialService(db, authSvc, userSvc, settingsSvc, eventSvc, issuer.server.Client(), keySetManager).WithRoleService(roleSvc) viewerRole := role.Role{ ID: "role-federated-viewer", diff --git a/backend/internal/oidc/service.go b/backend/internal/oidc/service.go index c5fc5aa71c..4426a751a1 100644 --- a/backend/internal/oidc/service.go +++ b/backend/internal/oidc/service.go @@ -27,7 +27,7 @@ import ( "github.com/getarcaneapp/arcane/backend/v2/internal/settings" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils" "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/jwtclaims" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/mldsajose" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" authtypes "github.com/getarcaneapp/arcane/types/v2/auth" ) @@ -39,7 +39,7 @@ type OidcService struct { insecureHttpClient *http.Client providerMutex sync.RWMutex providerCache *hot.HotCache[oidcProviderKey, *oidc.Provider] - keySets map[oidcKeySetKey]*mldsajose.KeySet + keySetManager *oidcjwk.KeySetManager } type oidcProviderKey struct { @@ -47,11 +47,6 @@ type oidcProviderKey struct { skipTLS bool } -type oidcKeySetKey struct { - jwksURL string - skipTLS bool -} - type OidcState struct { State string `json:"state"` Nonce string `json:"nonce"` @@ -60,7 +55,7 @@ type OidcState struct { CreatedAt time.Time `json:"created_at"` } -func NewOidcService(authService *auth.AuthService, settingsService *settings.SettingsService, cfg *config.Config, httpClient *http.Client) *OidcService { +func NewOidcService(authService *auth.AuthService, settingsService *settings.SettingsService, cfg *config.Config, httpClient *http.Client, keySetManager *oidcjwk.KeySetManager) *OidcService { if httpClient == nil { httpClient = http.DefaultClient } @@ -75,29 +70,17 @@ func NewOidcService(authService *auth.AuthService, settingsService *settings.Set settingsService: settingsService, config: cfg, httpClient: &oidcClient, - keySets: map[oidcKeySetKey]*mldsajose.KeySet{}, + keySetManager: keySetManager, } service.providerCache = hot.NewHotCache[oidcProviderKey, *oidc.Provider](hot.LRU, 4).Build() return service } -func (s *OidcService) keySetInternal(ctx context.Context, jwksURL string, skipTLS bool) *mldsajose.KeySet { - key := oidcKeySetKey{jwksURL: jwksURL, skipTLS: skipTLS} - s.providerMutex.RLock() - keySet := s.keySets[key] - s.providerMutex.RUnlock() - if keySet != nil { - return keySet +func (s *OidcService) keySetInternal(ctx context.Context, jwksURL string, skipTLS bool) (oidc.KeySet, error) { + if s.keySetManager == nil { + return nil, errors.New("JWK set manager is not configured") } - - clientCtx := oidc.ClientContext(context.WithoutCancel(ctx), s.getHttpClientInternal(skipTLS)) - s.providerMutex.Lock() - defer s.providerMutex.Unlock() - if keySet = s.keySets[key]; keySet == nil { - keySet = mldsajose.NewKeySet(clientCtx, jwksURL) - s.keySets[key] = keySet - } - return keySet + return s.keySetManager.KeySet(context.WithoutCancel(ctx), s.getHttpClientInternal(skipTLS), jwksURL) } func (s *OidcService) getEffectiveConfigInternal(ctx context.Context) (*settings.OidcConfig, error) { @@ -541,7 +524,7 @@ func (s *OidcService) verifyIDTokenInternal(ctx context.Context, provider *oidc. verifierConfig := &oidc.Config{ ClientID: cfg.ClientID, - SupportedSigningAlgs: mldsajose.SupportedSigningAlgs(), + SupportedSigningAlgs: oidcjwk.SupportedSigningAlgs(), } var issuer, jwksURL string @@ -561,8 +544,12 @@ func (s *OidcService) verifyIDTokenInternal(ctx context.Context, provider *oidc. return nil, "", errors.New("jwks URI must be configured when using manual OIDC endpoints") } + keySet, err := s.keySetInternal(ctx, jwksURL, cfg.SkipTlsVerify) + if err != nil { + return nil, "", errors.WrapIf(err, "failed to configure provider JWK set") + } providerCtx := oidc.ClientContext(ctx, s.getHttpClientInternal(cfg.SkipTlsVerify)) - verifier := oidc.NewVerifier(issuer, s.keySetInternal(ctx, jwksURL, cfg.SkipTlsVerify), verifierConfig) + verifier := oidc.NewVerifier(issuer, keySet, verifierConfig) idToken, err := verifier.Verify(providerCtx, rawIDToken) if err != nil { diff --git a/backend/pkg/utils/jwtclaims/jwt.go b/backend/pkg/utils/jwtclaims/jwt.go index d1ae7f5dd9..23c072552b 100644 --- a/backend/pkg/utils/jwtclaims/jwt.go +++ b/backend/pkg/utils/jwtclaims/jwt.go @@ -1,16 +1,13 @@ package jwtclaims import ( - "encoding/base64" - "encoding/json/v2" "fmt" + "maps" "strings" - "emperror.dev/errors" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils" + "github.com/lestrrat-go/jwx/v4/jwt" "github.com/samber/mo" - "github.com/samber/mo/result" ) // GetStringClaim extracts a string claim from a map @@ -116,26 +113,18 @@ func stringSliceFromInterfacesInternal[T any](items []T) []string { return utils.UniqueNonEmptyStrings(out) } -// ParseJWTClaims decodes and unmarshals the payload part of a JWT +// ParseJWTClaims parses unverified JWT metadata for pre-verification routing. func ParseJWTClaims(idToken string) map[string]any { - return result.Pipe3( - mo.Ok(idToken), - result.FlatMap(func(token string) mo.Result[[]string] { - parts := strings.Split(token, ".") - if len(parts) < 2 { - return mo.Err[[]string](errors.New("JWT has no payload")) - } - return mo.Ok(parts) - }), - result.FlatMap(func(parts []string) mo.Result[[]byte] { - return mo.TupleToResult(base64.RawURLEncoding.DecodeString(parts[1])) - }), - result.FlatMap(func(payload []byte) mo.Result[map[string]any] { - var claims map[string]any - err := json.Unmarshal(payload, &claims) - return mo.TupleToResult(claims, err) - }), - ).OrElse(nil) + token, err := jwt.ParseInsecure([]byte(idToken), jwt.WithStrictStringClaims(true)) + if err != nil { + return nil + } + + claims := maps.Collect(token.Claims()) + if audience, ok := token.Audience(); ok { + claims[jwt.AudienceKey] = audience + } + return claims } // GetByPath extracts a value from a nested map using a dot-separated path diff --git a/backend/pkg/utils/mldsajose/keyset.go b/backend/pkg/utils/mldsajose/keyset.go deleted file mode 100644 index a30bc33bba..0000000000 --- a/backend/pkg/utils/mldsajose/keyset.go +++ /dev/null @@ -1,105 +0,0 @@ -package mldsajose - -import ( - "context" - "io" - "net/http" - "sync" - - "emperror.dev/errors" - "github.com/coreos/go-oidc/v3/oidc" - "github.com/go-jose/go-jose/v4" - "golang.org/x/oauth2" - "golang.org/x/sync/singleflight" -) - -type KeySet struct { - client *http.Client - jwksURL string - remote *oidc.RemoteKeySet - mu sync.RWMutex - keys []Key - group singleflight.Group -} - -func NewKeySet(ctx context.Context, jwksURL string) *KeySet { - client := http.DefaultClient - if c, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && c != nil { - client = c - } - return &KeySet{ - client: client, - jwksURL: jwksURL, - remote: oidc.NewRemoteKeySet(ctx, jwksURL), - } -} - -func (k *KeySet) VerifySignature(ctx context.Context, jwt string) ([]byte, error) { - jws, err := jose.ParseSigned(jwt, joseAlgorithmsInternal(SupportedSigningAlgs())) - if err != nil { - return nil, errors.WrapIf(err, "malformed jwt") - } - if len(jws.Signatures) != 1 { - return nil, errors.New("expected exactly one signature") - } - if !IsAlg(jws.Signatures[0].Header.Algorithm) { - return k.remote.VerifySignature(ctx, jwt) - } - - k.mu.RLock() - keys := k.keys - k.mu.RUnlock() - if len(keys) > 0 { - if payload, err := VerifyCompact(jwt, keys); err == nil { - return payload, nil - } - } - - keys, err = k.refreshInternal(ctx) - if err != nil { - return nil, err - } - return VerifyCompact(jwt, keys) -} - -func (k *KeySet) refreshInternal(ctx context.Context) ([]Key, error) { - v, err, _ := k.group.Do("jwks", func() (any, error) { - req, err := http.NewRequestWithContext(ctx, http.MethodGet, k.jwksURL, nil) - if err != nil { - return nil, errors.WrapIf(err, "failed to build jwks request") - } - req.Header.Set("Cache-Control", "no-cache") - - resp, err := k.client.Do(req) - if err != nil { - return nil, errors.WrapIf(err, "failed to fetch jwks") - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return nil, errors.WrapIf(err, "failed to read jwks response") - } - if resp.StatusCode != http.StatusOK { - return nil, errors.Errorf("failed to fetch jwks: %s", resp.Status) - } - - keys, err := ParseJWKS(body) - if err != nil { - return nil, err - } - - k.mu.Lock() - k.keys = keys - k.mu.Unlock() - return keys, nil - }) - if err != nil { - return nil, err - } - keys, ok := v.([]Key) - if !ok { - return nil, errors.New("jwks refresh returned invalid keys") - } - return keys, nil -} diff --git a/backend/pkg/utils/mldsajose/mldsajose.go b/backend/pkg/utils/mldsajose/mldsajose.go deleted file mode 100644 index e3adfb81d4..0000000000 --- a/backend/pkg/utils/mldsajose/mldsajose.go +++ /dev/null @@ -1,168 +0,0 @@ -package mldsajose - -import ( - "crypto/mldsa" - "encoding/base64" - "encoding/json/v2" - "strings" - - "emperror.dev/errors" - "github.com/coreos/go-oidc/v3/oidc" - "github.com/go-jose/go-jose/v4" - "github.com/golang-jwt/jwt/v5" -) - -const ( - AlgMLDSA44 = "ML-DSA-44" - AlgMLDSA65 = "ML-DSA-65" - AlgMLDSA87 = "ML-DSA-87" - KeyTypeAKP = "AKP" - - ErrInvalidKeyType = errors.Sentinel("mldsajose: invalid key type") - ErrVerification = errors.Sentinel("mldsajose: verification error") -) - -var Parameters = mldsa.MLDSA87() - -func Algorithms() []string { - return []string{AlgMLDSA44, AlgMLDSA65, AlgMLDSA87} -} - -func SupportedSigningAlgs() []string { - return append([]string{ - oidc.RS256, oidc.RS384, oidc.RS512, - oidc.ES256, oidc.ES384, oidc.ES512, - oidc.PS256, oidc.PS384, oidc.PS512, - oidc.EdDSA, - }, Algorithms()...) -} - -func IsAlg(alg string) bool { - _, ok := ParametersForAlg(alg) - return ok -} - -func ParametersForAlg(alg string) (mldsa.Parameters, bool) { - switch alg { - case AlgMLDSA44: - return mldsa.MLDSA44(), true - case AlgMLDSA65: - return mldsa.MLDSA65(), true - case AlgMLDSA87: - return mldsa.MLDSA87(), true - } - return mldsa.Parameters{}, false -} - -type Key struct { - KeyID string - Alg string - Public *mldsa.PublicKey -} - -func ParseJWKS(body []byte) ([]Key, error) { - var set struct { - Keys []struct { - Kty string `json:"kty"` - Kid string `json:"kid"` - Alg string `json:"alg"` - Pub string `json:"pub"` - } `json:"keys"` - } - if err := json.Unmarshal(body, &set); err != nil { - return nil, errors.WrapIf(err, "failed to parse jwks") - } - - var keys []Key - for _, entry := range set.Keys { - if entry.Kty != KeyTypeAKP { - continue - } - params, ok := ParametersForAlg(entry.Alg) - if !ok { - continue - } - raw, err := base64.RawURLEncoding.DecodeString(entry.Pub) - if err != nil { - return nil, errors.WrapIff(err, "invalid AKP key %q", entry.Kid) - } - public, err := mldsa.NewPublicKey(params, raw) - if err != nil { - return nil, errors.WrapIff(err, "invalid AKP key %q", entry.Kid) - } - keys = append(keys, Key{KeyID: entry.Kid, Alg: entry.Alg, Public: public}) - } - return keys, nil -} - -func VerifyCompact(token string, keys []Key) ([]byte, error) { - jws, err := jose.ParseSignedCompact(token, joseAlgorithmsInternal(Algorithms())) - if err != nil { - return nil, errors.WrapIf(err, "malformed jws") - } - if len(jws.Signatures) != 1 { - return nil, errors.New("expected exactly one signature") - } - sig := jws.Signatures[0] - - dot := strings.LastIndexByte(token, '.') - if dot < 0 { - return nil, errors.New("malformed compact jws") - } - signingInput := []byte(token[:dot]) - - for _, key := range keys { - if key.Alg != sig.Header.Algorithm || (sig.Header.KeyID != "" && key.KeyID != sig.Header.KeyID) { - continue - } - if mldsa.Verify(key.Public, signingInput, sig.Signature, nil) == nil { - return jws.UnsafePayloadWithoutVerification(), nil - } - } - return nil, ErrVerification -} - -func joseAlgorithmsInternal(algs []string) []jose.SignatureAlgorithm { - out := make([]jose.SignatureAlgorithm, 0, len(algs)) - for _, alg := range algs { - out = append(out, jose.SignatureAlgorithm(alg)) - } - return out -} - -type SigningMethodMLDSA struct { - params mldsa.Parameters - alg string -} - -var SigningMethodMLDSA87 = &SigningMethodMLDSA{params: mldsa.MLDSA87(), alg: AlgMLDSA87} - -func init() { - jwt.RegisterSigningMethod(AlgMLDSA87, func() jwt.SigningMethod { return SigningMethodMLDSA87 }) -} - -func (m *SigningMethodMLDSA) Alg() string { - return m.alg -} - -func (m *SigningMethodMLDSA) Sign(signingString string, key any) ([]byte, error) { - sk, ok := key.(*mldsa.PrivateKey) - if !ok || sk.PublicKey().Parameters() != m.params { - return nil, ErrInvalidKeyType - } - return sk.Sign(nil, []byte(signingString), nil) -} - -func (m *SigningMethodMLDSA) Verify(signingString string, sig []byte, key any) error { - pk, ok := key.(*mldsa.PublicKey) - if !ok || pk.Parameters() != m.params { - return ErrInvalidKeyType - } - if len(sig) != m.params.SignatureSize() { - return ErrVerification - } - if err := mldsa.Verify(pk, []byte(signingString), sig, nil); err != nil { - return ErrVerification - } - return nil -} diff --git a/backend/pkg/utils/oidcjwk/algorithms.go b/backend/pkg/utils/oidcjwk/algorithms.go new file mode 100644 index 0000000000..2d18ce17b9 --- /dev/null +++ b/backend/pkg/utils/oidcjwk/algorithms.go @@ -0,0 +1,23 @@ +package oidcjwk + +import ( + "github.com/coreos/go-oidc/v3/oidc" + "github.com/lestrrat-go/jwx/v4/jwa" +) + +func Algorithms() []string { + return []string{ + jwa.MLDSA44().String(), + jwa.MLDSA65().String(), + jwa.MLDSA87().String(), + } +} + +func SupportedSigningAlgs() []string { + return append([]string{ + oidc.RS256, oidc.RS384, oidc.RS512, + oidc.ES256, oidc.ES384, oidc.ES512, + oidc.PS256, oidc.PS384, oidc.PS512, + oidc.EdDSA, + }, Algorithms()...) +} diff --git a/backend/pkg/utils/oidcjwk/keyset.go b/backend/pkg/utils/oidcjwk/keyset.go new file mode 100644 index 0000000000..3eb6f599be --- /dev/null +++ b/backend/pkg/utils/oidcjwk/keyset.go @@ -0,0 +1,110 @@ +package oidcjwk + +import ( + "context" + "slices" + "sync/atomic" + "time" + + "emperror.dev/errors" + "github.com/jwx-go/jwkfetch/v4" + "github.com/lestrrat-go/jwx/v4/jwk" + "github.com/lestrrat-go/jwx/v4/jws" + "golang.org/x/sync/singleflight" +) + +const forcedRefreshInterval = 30 * time.Second + +var errForcedRefreshThrottled = errors.Sentinel("oidcjwk: forced refresh throttled") + +type keySet struct { + cache *jwkfetch.Cache + jwksURL string + set jwk.Set + nextRefresh atomic.Int64 + group singleflight.Group +} + +func (k *keySet) VerifySignature(ctx context.Context, rawToken string) ([]byte, error) { + if err := validateTokenAlgorithmInternal(rawToken); err != nil { + return nil, err + } + + payload, verifyErr := k.verifyInternal(ctx, rawToken) + if verifyErr == nil { + return payload, nil + } + + // Providers may rotate key material without changing kid, so every failure earns + // one throttled refresh before the token is rejected. + refreshResult := k.group.DoChan("refresh", func() (any, error) { + if !k.reserveForcedRefreshInternal(time.Now()) { + return nil, errForcedRefreshThrottled + } + refreshCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), initialReadyTimeout) + defer cancel() + _, refreshErr := k.cache.Refresh(refreshCtx, k.jwksURL) + if refreshErr != nil { + return nil, errors.WrapIf(refreshErr, "failed to refresh JWKS") + } + return nil, nil + }) + var refreshErr error + select { + case result := <-refreshResult: + refreshErr = result.Err + case <-ctx.Done(): + return nil, ctx.Err() + } + if refreshErr != nil { + if errors.Is(refreshErr, errForcedRefreshThrottled) { + return k.verifyInternal(ctx, rawToken) + } + return nil, refreshErr + } + return k.verifyInternal(ctx, rawToken) +} + +func (k *keySet) verifyInternal(ctx context.Context, rawToken string) ([]byte, error) { + return jws.Verify([]byte(rawToken), + jws.WithCompact(), + jws.WithContext(ctx), + jws.WithKeySet(k.set, + jws.WithRequireKid(false), + jws.WithInferAlgorithmFromKey(true), + ), + ) +} + +func (k *keySet) reserveForcedRefreshInternal(now time.Time) bool { + nowUnixNano := now.UnixNano() + for { + nextRefresh := k.nextRefresh.Load() + if nowUnixNano < nextRefresh { + return false + } + if k.nextRefresh.CompareAndSwap(nextRefresh, now.Add(forcedRefreshInterval).UnixNano()) { + return true + } + } +} + +func validateTokenAlgorithmInternal(rawToken string) error { + message, err := jws.Parse([]byte(rawToken), jws.WithCompact()) + if err != nil { + return errors.WrapIf(err, "malformed JWT") + } + signatures := message.Signatures() + if len(signatures) != 1 { + return errors.New("expected exactly one signature") + } + headers := signatures[0].ProtectedHeaders() + if headers == nil { + return errors.New("JWT is missing protected headers") + } + algorithm, ok := headers.Algorithm() + if !ok || !slices.Contains(SupportedSigningAlgs(), algorithm.String()) { + return errors.New("JWT uses an unsupported signing algorithm") + } + return nil +} diff --git a/backend/pkg/utils/oidcjwk/manager.go b/backend/pkg/utils/oidcjwk/manager.go new file mode 100644 index 0000000000..9ebb6a11b9 --- /dev/null +++ b/backend/pkg/utils/oidcjwk/manager.go @@ -0,0 +1,192 @@ +package oidcjwk + +import ( + "context" + "fmt" + "net/http" + "sync" + "sync/atomic" + "time" + + "emperror.dev/errors" + "github.com/coreos/go-oidc/v3/oidc" + "github.com/jwx-go/jwkfetch/v4" + "github.com/lestrrat-go/httprc/v3" + "github.com/lestrrat-go/jwx/v4/jwk" + "github.com/samber/hot" + "golang.org/x/sync/singleflight" +) + +const ( + maxJWKSBodySize = 1 << 20 + maxJWKSKeys = 100 + maxHTTPClientPolicies = 16 + initialReadyTimeout = 30 * time.Second + minimumRefreshInterval = 5 * time.Minute + maximumRefreshInterval = time.Hour +) + +var ErrManagerShutdown = errors.Sentinel("oidcjwk: key set manager is shut down") + +type managedCache struct { + cache *jwkfetch.Cache + keySets map[string]*keySet +} + +type KeySetManager struct { + ctx context.Context + mu sync.Mutex + caches *hot.HotCache[*http.Client, *managedCache] + group singleflight.Group + active sync.WaitGroup + shutdown atomic.Bool +} + +func NewKeySetManager(ctx context.Context) *KeySetManager { //nolint:contextcheck // cache workers inherit the application lifecycle context, not request contexts. + if ctx == nil { + ctx = context.Background() + } + return &KeySetManager{ + ctx: ctx, + caches: hot.NewHotCache[*http.Client, *managedCache](hot.LRU, maxHTTPClientPolicies).Build(), + } +} + +func (m *KeySetManager) KeySet(ctx context.Context, client *http.Client, jwksURL string) (oidc.KeySet, error) { + if client == nil { + client = http.DefaultClient + } + if jwksURL == "" { + return nil, errors.New("oidcjwk: JWKS URL is empty") + } + if m.shutdown.Load() { + return nil, ErrManagerShutdown + } + + m.mu.Lock() + if m.shutdown.Load() { + m.mu.Unlock() + return nil, ErrManagerShutdown + } + m.active.Add(1) + if managed, found, _ := m.caches.Get(client); found { + if keySet := managed.keySets[jwksURL]; keySet != nil { + m.active.Done() + m.mu.Unlock() + return keySet, nil + } + } + m.mu.Unlock() + defer m.active.Done() + + groupKey := fmt.Sprintf("%p\x00%s", client, jwksURL) + value, err, _ := m.group.Do(groupKey, func() (any, error) { + m.mu.Lock() + if m.shutdown.Load() { + m.mu.Unlock() + return nil, ErrManagerShutdown + } + managed, found, _ := m.caches.Get(client) + if found { + if keySet := managed.keySets[jwksURL]; keySet != nil { + m.mu.Unlock() + return keySet, nil + } + } else { + var cacheErr error + managed, cacheErr = m.createManagedCacheLockedInternal(client) + if cacheErr != nil { + m.mu.Unlock() + return nil, cacheErr + } + } + m.mu.Unlock() + + registerCtx, cancel := context.WithTimeout(ctx, initialReadyTimeout) + defer cancel() + if registerErr := managed.cache.Register(registerCtx, jwksURL, + jwkfetch.WithWaitReady(true), + jwkfetch.WithMinInterval(minimumRefreshInterval), + jwkfetch.WithMaxInterval(maximumRefreshInterval), + ); registerErr != nil { + cleanupCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + if managed.cache.IsRegistered(cleanupCtx, jwksURL) { + _ = managed.cache.Unregister(cleanupCtx, jwksURL) + } + return nil, errors.WrapIf(registerErr, "failed to register JWKS URL") + } + + set, setErr := managed.cache.CachedSet(jwksURL) + if setErr != nil { + return nil, errors.WrapIf(setErr, "failed to create cached JWK set") + } + keySet := &keySet{ + cache: managed.cache, + jwksURL: jwksURL, + set: set, + } + + m.mu.Lock() + defer m.mu.Unlock() + if m.shutdown.Load() { + return nil, ErrManagerShutdown + } + managed.keySets[jwksURL] = keySet + return keySet, nil + }) + if err != nil { + return nil, err + } + keySet, ok := value.(*keySet) + if !ok || keySet == nil { + return nil, errors.New("oidcjwk: invalid key set") + } + return keySet, nil +} + +func (m *KeySetManager) createManagedCacheLockedInternal(client *http.Client) (*managedCache, error) { + if m.caches.Len() >= maxHTTPClientPolicies { + return nil, errors.New("oidcjwk: too many HTTP client policies") + } + cache, err := jwkfetch.NewCache(m.ctx, httprc.NewClient(), + jwkfetch.WithHTTPClient(jwkfetch.WrapHTTPClientDefaults(client)), + jwkfetch.WithMaxBodySize(maxJWKSBodySize), + jwkfetch.WithParseOptions( + jwk.WithMaxKeys(maxJWKSKeys), + jwk.WithRejectDuplicateKID(true), + ), + ) + if err != nil { + return nil, errors.WrapIf(err, "failed to create JWKS cache") + } + managed := &managedCache{cache: cache, keySets: make(map[string]*keySet)} + m.caches.Set(client, managed) + return managed, nil +} + +func (m *KeySetManager) Shutdown(ctx context.Context) error { + m.mu.Lock() + if m.shutdown.Swap(true) { + m.mu.Unlock() + return nil + } + m.mu.Unlock() + + m.active.Wait() + + m.mu.Lock() + managedCaches := m.caches.Values() + caches := make([]*jwkfetch.Cache, 0, len(managedCaches)) + for _, managed := range managedCaches { + caches = append(caches, managed.cache) + } + m.caches.Purge() + m.mu.Unlock() + + shutdownErrors := make([]error, 0, len(caches)) + for _, cache := range caches { + shutdownErrors = append(shutdownErrors, cache.Shutdown(ctx)) + } + return errors.Combine(shutdownErrors...) +} diff --git a/backend/pkg/utils/mldsajose/mldsajose_test.go b/backend/pkg/utils/oidcjwk/oidcjwk_test.go similarity index 71% rename from backend/pkg/utils/mldsajose/mldsajose_test.go rename to backend/pkg/utils/oidcjwk/oidcjwk_test.go index 72ee43dba1..d46920df33 100644 --- a/backend/pkg/utils/mldsajose/mldsajose_test.go +++ b/backend/pkg/utils/oidcjwk/oidcjwk_test.go @@ -1,6 +1,7 @@ -package mldsajose_test +package oidcjwk_test import ( + "context" "crypto/mldsa" "encoding/base64" "encoding/json/v2" @@ -11,8 +12,7 @@ import ( "time" "github.com/coreos/go-oidc/v3/oidc" - "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/mldsajose" - "github.com/golang-jwt/jwt/v5" + "github.com/getarcaneapp/arcane/backend/v2/pkg/utils/oidcjwk" "github.com/stretchr/testify/require" ) @@ -67,10 +67,17 @@ func TestKeySetVerifySignature(t *testing.T) { } ctx := t.Context() - keySet := mldsajose.NewKeySet(oidc.ClientContext(ctx, srv.Client()), srv.URL) + manager := oidcjwk.NewKeySetManager(ctx) + t.Cleanup(func() { + shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + require.NoError(t, manager.Shutdown(shutdownCtx)) + }) + keySet, err := manager.KeySet(ctx, srv.Client(), srv.URL) + require.NoError(t, err) verifier := oidc.NewVerifier(srv.URL, keySet, &oidc.Config{ ClientID: "arcane", - SupportedSigningAlgs: mldsajose.SupportedSigningAlgs(), + SupportedSigningAlgs: oidcjwk.SupportedSigningAlgs(), }) idToken, err := verifier.Verify(ctx, token) @@ -83,20 +90,3 @@ func TestKeySetVerifySignature(t *testing.T) { }) } } - -func TestSigningMethodMLDSA87(t *testing.T) { - sk, err := mldsa.GenerateKey(mldsajose.Parameters) - require.NoError(t, err) - - signed, err := jwt.NewWithClaims(mldsajose.SigningMethodMLDSA87, jwt.MapClaims{"sub": "x"}).SignedString(sk) - require.NoError(t, err) - - parsed, err := jwt.Parse(signed, func(t *jwt.Token) (any, error) { - return sk.PublicKey(), nil - }, jwt.WithValidMethods([]string{mldsajose.AlgMLDSA87})) - require.NoError(t, err) - require.True(t, parsed.Valid) - - _, err = mldsajose.SigningMethodMLDSA87.Sign("payload", []byte("not-a-key")) - require.ErrorIs(t, err, mldsajose.ErrInvalidKeyType) -} diff --git a/go.work.sum b/go.work.sum index cc02aac4e5..098ec7bd45 100644 --- a/go.work.sum +++ b/go.work.sum @@ -1469,7 +1469,6 @@ github.com/google/pprof v0.0.0-20210407192527-94a9f03dee38/go.mod h1:kpwsk12EmLe github.com/google/pprof v0.0.0-20211214055906-6f57359322fd/go.mod h1:KgnwoLYCZ8IQu3XUZ8Nc/bM9CCZFOyjUNOSygVozoDg= github.com/google/pprof v0.0.0-20250820193118-f64d9cf942d6/go.mod h1:I6V7YzU0XDpsHqbsyrghnFZLO1gwK6NPTNvmetQIk9U= github.com/google/pprof v0.0.0-20260402051712-545e8a4df936/go.mod h1:MxpfABSjhmINe3F1It9d+8exIHFvUqtLIRCdOGNXqiI= -github.com/google/pprof v0.0.0-20260825171938-4d453200e7d9/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk= github.com/google/renameio v0.1.0 h1:GOZbcHa3HfsPKPlmyPyN2KEohoMXOhdMbHrvbpl2QaA= github.com/google/rpmpack v0.7.1 h1:YdWh1IpzOjBz60Wvdw0TU0A5NWP+JTVHA5poDqwMO2o= github.com/google/rpmpack v0.7.1/go.mod h1:h1JL16sUTWCLI/c39ox1rDaTBo3BXUQGjczVJyK4toU= @@ -1729,16 +1728,12 @@ github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjS github.com/lestrrat-go/backoff/v2 v2.0.8 h1:oNb5E5isby2kiro9AgdHLv5N5tint1AnDVVf2E2un5A= github.com/lestrrat-go/backoff/v2 v2.0.8/go.mod h1:rHP/q/r9aT27n24JQLa7JhSQZCKBBOiM/uP402WwN8Y= github.com/lestrrat-go/blackmagic v1.0.2/go.mod h1:UrEqBzIR2U6CnzVyUtfM6oZNMt/7O7Vohk2J0OGSAtU= -github.com/lestrrat-go/blackmagic v1.0.4 h1:IwQibdnf8l2KoO+qC3uT4OaTWsW7tuRQXy9TRN9QanA= -github.com/lestrrat-go/blackmagic v1.0.4/go.mod h1:6AWFyKNNj0zEXQYfTMPfZrAXUWUfTIZ5ECEUEJaijtw= github.com/lestrrat-go/dsig v1.0.0 h1:OE09s2r9Z81kxzJYRn07TFM9XA4akrUdoMwr0L8xj38= github.com/lestrrat-go/dsig v1.0.0/go.mod h1:dEgoOYYEJvW6XGbLasr8TFcAxoWrKlbQvmJgCR0qkDo= github.com/lestrrat-go/dsig v1.2.1 h1:MwxzZhE4+4fguHi+uDALKVlC3Cn+O1QU1Q/F8D7hVIc= github.com/lestrrat-go/dsig v1.2.1/go.mod h1:RD2eOaidyPvpc7IJQoO3Qq52RWdy8ZcJs8lrOnoa1Kc= github.com/lestrrat-go/dsig-secp256k1 v1.0.0 h1:JpDe4Aybfl0soBvoVwjqDbp+9S1Y2OM7gcrVVMFPOzY= github.com/lestrrat-go/dsig-secp256k1 v1.0.0/go.mod h1:CxUgAhssb8FToqbL8NjSPoGQlnO4w3LG1P0qPWQm/NU= -github.com/lestrrat-go/httpcc v1.0.1 h1:ydWCStUeJLkpYyjLDHihupbn2tYmZ7m22BGkcvZZrIE= -github.com/lestrrat-go/httpcc v1.0.1/go.mod h1:qiltp3Mt56+55GPVCbTdM9MlqhvzyuL6W/NMDA8vA5E= github.com/lestrrat-go/httprc/v3 v3.0.1 h1:3n7Es68YYGZb2Jf+k//llA4FTZMl3yCwIjFIk4ubevI= github.com/lestrrat-go/httprc/v3 v3.0.1/go.mod h1:2uAvmbXE4Xq8kAUjVrZOq1tZVYYYs5iP62Cmtru00xk= github.com/lestrrat-go/httprc/v3 v3.0.2 h1:7u4HUaD0NQbf2/n5+fyp+T10hNCsAnwKfqn4A4Baif0= @@ -1758,8 +1753,6 @@ github.com/lestrrat-go/jwx/v3 v3.1.1/go.mod h1:uw/MN2M/Xiu4FhwcIwH11Zsh9JWx9SWzg github.com/lestrrat-go/option v1.0.0/go.mod h1:5ZHFbivi4xwXxhxY9XHDe2FHo6/Z7WWmtT7T5nBBp3I= github.com/lestrrat-go/option v1.0.1 h1:oAzP2fvZGQKWkvHa1/SAcFolBEca1oN+mQ7eooNBEYU= github.com/lestrrat-go/option v1.0.1/go.mod h1:5ZHFbivi4xwXxhxY9XHDe2FHo6/Z7WWmtT7T5nBBp3I= -github.com/lestrrat-go/option/v2 v2.0.0 h1:XxrcaJESE1fokHy3FpaQ/cXW8ZsIdWcdFzzLOcID3Ss= -github.com/lestrrat-go/option/v2 v2.0.0/go.mod h1:oSySsmzMoR0iRzCDCaUfsCzxQHUEuhOViQObyy7S6Vg= github.com/letsencrypt/borp v0.0.0-20251118150929-89c6927051ae h1:yFuF5yRIwaandcuNMi1A4he4FMWJsGRv38rsizIaxJA= github.com/letsencrypt/borp v0.0.0-20251118150929-89c6927051ae/go.mod h1:gMSMCNKhxox/ccR923EJsIvHeVVYfCABGbirqa0EwuM= github.com/letsencrypt/boulder v0.0.0-20240620165639-de9c06129bec h1:2tTW6cDth2TSgRbAhD7yjZzTQmcN25sDRPEeinR51yQ= @@ -2367,8 +2360,6 @@ github.com/valyala/fastjson v1.6.4 h1:uAUNq9Z6ymTgGhcm0UynUAB6tlbakBrz6CQFax3BXV github.com/valyala/fastjson v1.6.4/go.mod h1:CLCAqky6SMuOcxStkYQvblddUtoRxhYMGLrsQns1aXY= github.com/valyala/fastjson v1.6.7 h1:ZE4tRy0CIkh+qDc5McjatheGX2czdn8slQjomexVpBM= github.com/valyala/fastjson v1.6.7/go.mod h1:CLCAqky6SMuOcxStkYQvblddUtoRxhYMGLrsQns1aXY= -github.com/valyala/fastjson v1.6.10 h1:/yjJg8jaVQdYR3arGxPE2X5z89xrlhS0eGXdv+ADTh4= -github.com/valyala/fastjson v1.6.10/go.mod h1:e6FubmQouUNP73jtMLmcbxS6ydWIpOfhz34TSfO3JaE= github.com/vbatts/tar-split v0.11.5/go.mod h1:yZbwRsSeGjusneWgA781EKej9HF8vme8okylkAeNKLk= github.com/vbauerster/mpb/v8 v8.10.2 h1:2uBykSHAYHekE11YvJhKxYmLATKHAGorZwFlyNw4hHM= github.com/vbauerster/mpb/v8 v8.10.2/go.mod h1:+Ja4P92E3/CorSZgfDtK46D7AVbDqmBQRTmyTqPElo0=