Verify OAuth2 ID tokens before trusting claims

The compliance portal OAuth callback accepted ID tokens after only
parsing claims, without checking the signature, issuer, audience, or
expiry. Add RS256 verification helpers to the JOSE package, enforce
those checks in ParseIDTokenIdentity, and thread JWKS, issuer, and
client ID through the token response and callback handler.

Signed-off-by: Bryan Frimin <bryan@probo.com>
This commit is contained in:
Bryan Frimin
2026-07-15 14:38:37 +02:00
parent 4e73bd6a97
commit 48dba254ca
6 changed files with 711 additions and 4 deletions

View File

@@ -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)

View File

@@ -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")
},
)
}

View File

@@ -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
}