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:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user