diff --git a/pkg/crypto/jose/jose.go b/pkg/crypto/jose/jose.go index a9e1012e3..9b3e15c2b 100644 --- a/pkg/crypto/jose/jose.go +++ b/pkg/crypto/jose/jose.go @@ -29,6 +29,7 @@ import ( "encoding/json" "fmt" "math/big" + "strings" ) type ( @@ -103,3 +104,106 @@ func SignJWT(privateKey *rsa.PrivateKey, kid string, claims any) (string, error) return signingInput + "." + signatureB64, nil } + +// RSAPublicKeyFromJWK reconstructs an RSA public key from a JWK. +func RSAPublicKeyFromJWK(jwk JWK) (*rsa.PublicKey, error) { + if jwk.KeyType != "RSA" { + return nil, fmt.Errorf("cannot convert jwk to rsa public key: unsupported key type %q", jwk.KeyType) + } + + nBytes, err := base64.RawURLEncoding.DecodeString(jwk.N) + if err != nil { + return nil, fmt.Errorf("cannot decode rsa modulus: %w", err) + } + + eBytes, err := base64.RawURLEncoding.DecodeString(jwk.E) + if err != nil { + return nil, fmt.Errorf("cannot decode rsa exponent: %w", err) + } + + return &rsa.PublicKey{ + N: new(big.Int).SetBytes(nBytes), + E: int(new(big.Int).SetBytes(eBytes).Int64()), + }, nil +} + +// PublicKeyFromJWKS returns the RSA public key matching the given key ID. +func PublicKeyFromJWKS(jwks *JWKS, kid string) (*rsa.PublicKey, error) { + for _, key := range jwks.Keys { + if key.KeyID == kid { + return RSAPublicKeyFromJWK(key) + } + } + + return nil, fmt.Errorf("cannot find signing key %q in jwks", kid) +} + +// VerifyJWT verifies an RS256 JWT signature and returns the decoded payload. +func VerifyJWT(raw string, pubKey *rsa.PublicKey) ([]byte, error) { + parts := strings.Split(raw, ".") + if len(parts) != 3 { + return nil, fmt.Errorf("cannot verify jwt: invalid format") + } + + headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + return nil, fmt.Errorf("cannot decode jwt header: %w", err) + } + + var header JWTHeader + if err := json.Unmarshal(headerJSON, &header); err != nil { + return nil, fmt.Errorf("cannot parse jwt header: %w", err) + } + + if header.Algorithm != "RS256" { + return nil, fmt.Errorf("cannot verify jwt: unsupported algorithm %q", header.Algorithm) + } + + payload, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return nil, fmt.Errorf("cannot decode jwt payload: %w", err) + } + + signature, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil { + return nil, fmt.Errorf("cannot decode jwt signature: %w", err) + } + + signedContent := parts[0] + "." + parts[1] + hash := sha256.Sum256([]byte(signedContent)) + + if err := rsa.VerifyPKCS1v15(pubKey, crypto.SHA256, hash[:], signature); err != nil { + return nil, fmt.Errorf("cannot verify jwt signature: %w", err) + } + + return payload, nil +} + +// VerifyJWTWithJWKS verifies an RS256 JWT using the matching key from a JWKS. +func VerifyJWTWithJWKS(raw string, jwks *JWKS) ([]byte, error) { + parts := strings.Split(raw, ".") + if len(parts) != 3 { + return nil, fmt.Errorf("cannot verify jwt: invalid format") + } + + headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + return nil, fmt.Errorf("cannot decode jwt header: %w", err) + } + + var header JWTHeader + if err := json.Unmarshal(headerJSON, &header); err != nil { + return nil, fmt.Errorf("cannot parse jwt header: %w", err) + } + + if header.KeyID == "" { + return nil, fmt.Errorf("cannot verify jwt: missing key id") + } + + pubKey, err := PublicKeyFromJWKS(jwks, header.KeyID) + if err != nil { + return nil, err + } + + return VerifyJWT(raw, pubKey) +} diff --git a/pkg/crypto/jose/jose_test.go b/pkg/crypto/jose/jose_test.go index 0379b2e9c..60f097612 100644 --- a/pkg/crypto/jose/jose_test.go +++ b/pkg/crypto/jose/jose_test.go @@ -297,3 +297,266 @@ func TestJWTHeader_JSON(t *testing.T) { }, ) } + +func TestRSAPublicKeyFromJWK(t *testing.T) { + t.Parallel() + + key := testRSAKey(t) + jwk := jose.RSAPublicKeyToJWK(&key.PublicKey, "kid-1") + + t.Run( + "round trips rsa public key", + func(t *testing.T) { + t.Parallel() + + pubKey, err := jose.RSAPublicKeyFromJWK(jwk) + require.NoError(t, err) + assert.Equal(t, key.N, pubKey.N) + assert.Equal(t, key.E, pubKey.E) + }, + ) + + t.Run( + "rejects unsupported key type", + func(t *testing.T) { + t.Parallel() + + _, err := jose.RSAPublicKeyFromJWK(jose.JWK{KeyType: "EC"}) + require.Error(t, err) + }, + ) +} + +func TestPublicKeyFromJWKS(t *testing.T) { + t.Parallel() + + key := testRSAKey(t) + jwks := &jose.JWKS{ + Keys: []jose.JWK{ + jose.RSAPublicKeyToJWK(&key.PublicKey, "kid-1"), + jose.RSAPublicKeyToJWK(&key.PublicKey, "kid-2"), + }, + } + + t.Run( + "finds matching key", + func(t *testing.T) { + t.Parallel() + + pubKey, err := jose.PublicKeyFromJWKS(jwks, "kid-2") + require.NoError(t, err) + assert.Equal(t, key.E, pubKey.E) + }, + ) + + t.Run( + "errors when kid is missing", + func(t *testing.T) { + t.Parallel() + + _, err := jose.PublicKeyFromJWKS(jwks, "missing") + require.Error(t, err) + }, + ) +} + +func TestVerifyJWT(t *testing.T) { + t.Parallel() + + key := testRSAKey(t) + + t.Run( + "verifies signed jwt", + func(t *testing.T) { + t.Parallel() + + claims := map[string]string{"sub": "test"} + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + payload, err := jose.VerifyJWT(token, &key.PublicKey) + require.NoError(t, err) + + var decoded map[string]string + + err = json.Unmarshal(payload, &decoded) + require.NoError(t, err) + assert.Equal(t, "test", decoded["sub"]) + }, + ) + + t.Run( + "rejects malformed jwt", + func(t *testing.T) { + t.Parallel() + + _, err := jose.VerifyJWT("", &key.PublicKey) + require.Error(t, err) + + _, err = jose.VerifyJWT("only-one-part", &key.PublicKey) + require.Error(t, err) + + _, err = jose.VerifyJWT("two.parts", &key.PublicKey) + require.Error(t, err) + + _, err = jose.VerifyJWT("too.many.parts.here", &key.PublicKey) + require.Error(t, err) + }, + ) + + t.Run( + "rejects tampered payload", + func(t *testing.T) { + t.Parallel() + + token, err := jose.SignJWT(key, "kid-1", map[string]string{"sub": "test"}) + require.NoError(t, err) + + parts := strings.Split(token, ".") + parts[1] = base64.RawURLEncoding.EncodeToString([]byte(`{"sub":"evil"}`)) + tampered := strings.Join(parts, ".") + + _, err = jose.VerifyJWT(tampered, &key.PublicKey) + require.Error(t, err) + }, + ) + + t.Run( + "rejects tampered header", + func(t *testing.T) { + t.Parallel() + + token, err := jose.SignJWT(key, "kid-1", map[string]string{"sub": "test"}) + require.NoError(t, err) + + parts := strings.Split(token, ".") + parts[0] = base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"HS256","typ":"JWT","kid":"kid-1"}`)) + tampered := strings.Join(parts, ".") + + _, err = jose.VerifyJWT(tampered, &key.PublicKey) + require.Error(t, err) + }, + ) + + t.Run( + "rejects unsupported algorithm", + func(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + alg string + }{ + {name: "hs256", alg: "HS256"}, + {name: "none", alg: "none"}, + {name: "lowercase rs256", alg: "rs256"}, + } + + for _, tc := range cases { + t.Run( + tc.name, + func(t *testing.T) { + t.Parallel() + + header := base64.RawURLEncoding.EncodeToString( + []byte(`{"alg":"` + tc.alg + `","typ":"JWT","kid":"kid-1"}`), + ) + payload := base64.RawURLEncoding.EncodeToString([]byte(`{"sub":"test"}`)) + token := header + "." + payload + ".c2ln" + + _, err := jose.VerifyJWT(token, &key.PublicKey) + require.Error(t, err) + }, + ) + } + }, + ) + + t.Run( + "rejects signature from different key", + func(t *testing.T) { + t.Parallel() + + otherKey := testRSAKey(t) + + token, err := jose.SignJWT(otherKey, "kid-1", map[string]string{"sub": "test"}) + require.NoError(t, err) + + _, err = jose.VerifyJWT(token, &key.PublicKey) + require.Error(t, err) + }, + ) +} + +func TestVerifyJWTWithJWKS(t *testing.T) { + t.Parallel() + + key := testRSAKey(t) + jwks := &jose.JWKS{ + Keys: []jose.JWK{ + jose.RSAPublicKeyToJWK(&key.PublicKey, "kid-1"), + }, + } + + t.Run( + "verifies signed jwt", + func(t *testing.T) { + t.Parallel() + + token, err := jose.SignJWT(key, "kid-1", map[string]string{"sub": "test"}) + require.NoError(t, err) + + payload, err := jose.VerifyJWTWithJWKS(token, jwks) + require.NoError(t, err) + + var decoded map[string]string + + err = json.Unmarshal(payload, &decoded) + require.NoError(t, err) + assert.Equal(t, "test", decoded["sub"]) + }, + ) + + t.Run( + "rejects missing key id", + func(t *testing.T) { + t.Parallel() + + header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","typ":"JWT"}`)) + payload := base64.RawURLEncoding.EncodeToString([]byte(`{"sub":"test"}`)) + token := header + "." + payload + ".c2ln" + + _, err := jose.VerifyJWTWithJWKS(token, jwks) + require.Error(t, err) + }, + ) + + t.Run( + "rejects unknown key id", + func(t *testing.T) { + t.Parallel() + + token, err := jose.SignJWT(key, "unknown-kid", map[string]string{"sub": "test"}) + require.NoError(t, err) + + _, err = jose.VerifyJWTWithJWKS(token, jwks) + require.Error(t, err) + }, + ) + + t.Run( + "rejects token signed with key outside jwks", + func(t *testing.T) { + t.Parallel() + + otherKey := testRSAKey(t) + + token, err := jose.SignJWT(otherKey, "kid-1", map[string]string{"sub": "test"}) + require.NoError(t, err) + + _, err = jose.VerifyJWTWithJWKS(token, jwks) + require.Error(t, err) + }, + ) +} diff --git a/pkg/iam/oauth2/id_token.go b/pkg/iam/oauth2/id_token.go index 66cf4342a..8ad58db5d 100644 --- a/pkg/iam/oauth2/id_token.go +++ b/pkg/iam/oauth2/id_token.go @@ -30,6 +30,7 @@ import ( "time" "go.probo.inc/probo/pkg/coredata" + "go.probo.inc/probo/pkg/crypto/jose" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/uri" ) @@ -135,20 +136,43 @@ func ParseIDTokenClaims(raw string) (*IDTokenClaims, error) { return &claims, nil } -func ParseIDTokenIdentity(raw string, expectedNonce string) (gid.GID, error) { +func ParseIDTokenIdentity( + raw string, + jwks *jose.JWKS, + expectedNonce string, + expectedIssuer uri.URI, + expectedAudience string, +) (gid.GID, error) { if raw == "" { return gid.GID{}, fmt.Errorf("cannot parse id token: missing token") } - claims, err := ParseIDTokenClaims(raw) + payload, err := jose.VerifyJWTWithJWKS(raw, jwks) if err != nil { - return gid.GID{}, err + return gid.GID{}, fmt.Errorf("cannot verify id token: %w", err) + } + + var claims IDTokenClaims + if err := json.Unmarshal(payload, &claims); err != nil { + return gid.GID{}, fmt.Errorf("cannot decode id token claims: %w", err) + } + + if claims.Issuer != expectedIssuer { + return gid.GID{}, fmt.Errorf("cannot validate issuer: unexpected issuer %q", claims.Issuer) + } + + if claims.Audience != expectedAudience { + return gid.GID{}, fmt.Errorf("cannot validate audience: mismatch") } if claims.Nonce != expectedNonce { return gid.GID{}, fmt.Errorf("cannot validate nonce: mismatch") } + if time.Now().After(time.Unix(claims.ExpiresAt, 0)) { + return gid.GID{}, fmt.Errorf("cannot validate id token: token has expired") + } + identityID, err := gid.ParseGID(claims.Subject) if err != nil { return gid.GID{}, fmt.Errorf("cannot parse identity from id token: %w", err) diff --git a/pkg/iam/oauth2/id_token_test.go b/pkg/iam/oauth2/id_token_test.go index 608026255..060532d79 100644 --- a/pkg/iam/oauth2/id_token_test.go +++ b/pkg/iam/oauth2/id_token_test.go @@ -21,14 +21,17 @@ package oauth2_test import ( + "crypto/rsa" "crypto/sha256" "encoding/base64" + "strings" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.probo.inc/probo/pkg/coredata" + "go.probo.inc/probo/pkg/crypto/jose" "go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/iam/oauth2" "go.probo.inc/probo/pkg/uri" @@ -304,3 +307,301 @@ func TestNewIDTokenClaims(t *testing.T) { }, ) } + +func testSigningKey(t *testing.T) (*rsa.PrivateKey, *jose.JWKS) { + t.Helper() + + key, err := rsa.GenerateKey( + strings.NewReader(strings.Repeat("deterministic-seed-for-test!!!!!", 100)), + 2048, + ) + require.NoError(t, err) + + jwks := &jose.JWKS{ + Keys: []jose.JWK{ + jose.RSAPublicKeyToJWK(&key.PublicKey, "kid-1"), + }, + } + + return key, jwks +} + +func TestParseIDTokenIdentity(t *testing.T) { + t.Parallel() + + identityID := gid.MustParseGID("AAAAAAAAAAAASwAAAAAAAAAAcHJiY2xp") + clientID := gid.MustParseGID("AAAAAAAAAAAASwAAAAAAAAAAcHJiY2xp") + authTime := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + key, jwks := testSigningKey(t) + + parseIDToken := func( + t *testing.T, + token string, + nonce string, + ) (gid.GID, error) { + t.Helper() + + return oauth2.ParseIDTokenIdentity( + token, + jwks, + nonce, + testIssuer, + clientID.String(), + ) + } + + t.Run( + "returns identity when signature nonce and expiry are valid", + func(t *testing.T) { + t.Parallel() + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + got, err := parseIDToken(t, token, "test-nonce") + require.NoError(t, err) + assert.Equal(t, identityID, got) + }, + ) + + t.Run( + "rejects nonce mismatch", + func(t *testing.T) { + t.Parallel() + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "expected-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + _, err = parseIDToken(t, token, "other-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot validate nonce") + }, + ) + + t.Run( + "rejects issuer mismatch", + func(t *testing.T) { + t.Parallel() + + claims := oauth2.NewIDTokenClaims( + uri.URI("https://evil.example.com"), + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + _, err = parseIDToken(t, token, "test-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot validate issuer") + }, + ) + + t.Run( + "rejects audience mismatch", + func(t *testing.T) { + t.Parallel() + + otherClientID := gid.New(clientID.TenantID(), coredata.OAuth2ClientEntityType) + require.NotEqual(t, clientID, otherClientID) + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + otherClientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + _, err = parseIDToken(t, token, "test-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot validate audience") + }, + ) + + t.Run( + "rejects expired token", + func(t *testing.T) { + t.Parallel() + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + -1*time.Hour, + ) + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + _, err = parseIDToken(t, token, "test-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "token has expired") + }, + ) + + t.Run( + "rejects invalid signature", + func(t *testing.T) { + t.Parallel() + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + parts := strings.Split(token, ".") + parts[2] = base64.RawURLEncoding.EncodeToString([]byte("bad-signature")) + tampered := strings.Join(parts, ".") + + _, err = parseIDToken(t, tampered, "test-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot verify id token") + }, + ) + + t.Run( + "rejects empty token", + func(t *testing.T) { + t.Parallel() + + _, err := oauth2.ParseIDTokenIdentity( + "", + jwks, + "test-nonce", + testIssuer, + clientID.String(), + ) + require.Error(t, err) + assert.Contains(t, err.Error(), "missing token") + }, + ) + + t.Run( + "rejects invalid subject", + func(t *testing.T) { + t.Parallel() + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + claims.Subject = "not-a-gid" + + token, err := jose.SignJWT(key, "kid-1", claims) + require.NoError(t, err) + + _, err = parseIDToken(t, token, "test-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot parse identity from id token") + }, + ) + + t.Run( + "rejects token signed with key outside jwks", + func(t *testing.T) { + t.Parallel() + + otherKey, err := rsa.GenerateKey( + strings.NewReader(strings.Repeat("other-deterministic-seed!!!!!!!!", 100)), + 2048, + ) + require.NoError(t, err) + + claims := oauth2.NewIDTokenClaims( + testIssuer, + identityID, + clientID, + authTime, + coredata.OAuth2Scopes{oauth2.ScopeOpenID}, + "test-nonce", + "", + "", + false, + "", + 1*time.Hour, + ) + + token, err := jose.SignJWT(otherKey, "kid-1", claims) + require.NoError(t, err) + + _, err = parseIDToken(t, token, "test-nonce") + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot verify id token") + }, + ) +} diff --git a/pkg/iam/oauth2/service.go b/pkg/iam/oauth2/service.go index d8ff129d5..13f242708 100644 --- a/pkg/iam/oauth2/service.go +++ b/pkg/iam/oauth2/service.go @@ -122,6 +122,7 @@ type ( RefreshToken string IDToken string Scope string + ClientID gid.GID } IntrospectResult struct { @@ -246,6 +247,11 @@ func (s *Service) JWKS() *jose.JWKS { return jwks } +// Issuer returns the OAuth2 issuer URI embedded in ID tokens. +func (s *Service) Issuer() uri.URI { + return s.baseURL +} + func (s *Service) CreateAccessToken( ctx context.Context, clientID gid.GID, @@ -496,6 +502,7 @@ func (s *Service) ExchangeAuthorizationCode( RefreshToken: refreshTokenValue, Scope: code.Scopes.String(), IDToken: idToken, + ClientID: client.ID, }, nil } @@ -684,6 +691,7 @@ func (s *Service) RefreshToken( RefreshToken: refreshTokenValueNew, Scope: previousRefreshToken.Scopes.String(), IDToken: idToken, + ClientID: client.ID, }, nil } @@ -959,6 +967,7 @@ func (s *Service) PollDeviceCode( RefreshToken: refreshTokenValue, Scope: deviceCode.Scopes.String(), IDToken: idToken, + ClientID: client.ID, }, nil } diff --git a/pkg/server/api/complianceportal/v1/oauth_callback_handler.go b/pkg/server/api/complianceportal/v1/oauth_callback_handler.go index 994eecf59..63439edc9 100644 --- a/pkg/server/api/complianceportal/v1/oauth_callback_handler.go +++ b/pkg/server/api/complianceportal/v1/oauth_callback_handler.go @@ -128,7 +128,13 @@ func (h *OAuthCallbackHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) return } - identityID, err := oauth2.ParseIDTokenIdentity(tokenResult.IDToken, state.Nonce) + identityID, err := oauth2.ParseIDTokenIdentity( + tokenResult.IDToken, + h.iam.OAuth2ServerService.JWKS(), + state.Nonce, + h.iam.OAuth2ServerService.Issuer(), + tokenResult.ClientID.String(), + ) if err != nil { h.logger.WarnCtx(ctx, "cannot validate id token", log.Error(err)) httpserver.RenderError(w, http.StatusBadRequest, errInvalidOAuthRequest)