Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,7 @@ require (
github.com/felixge/httpsnoop v1.0.4 // indirect
github.com/go-logr/logr v1.4.3 // indirect
github.com/go-logr/stdr v1.2.2 // indirect
github.com/golang-jwt/jwt/v5 v5.3.1 // indirect
github.com/golang-jwt/jwt/v5 v5.3.1
github.com/google/s2a-go v0.1.9 // indirect
github.com/googleapis/enterprise-certificate-proxy v0.3.14 // indirect
github.com/googleapis/gax-go/v2 v2.21.0 // indirect
Expand Down
6 changes: 4 additions & 2 deletions internal/api/handler_oidc_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -88,8 +88,10 @@ func TestHandleOIDCJWKS(t *testing.T) {
t.Fatalf("keys=%d want 1", len(jwks.Keys))
}
k := jwks.Keys[0]
if k.Kty != "RSA" || k.Alg != "RS256" || k.Kid == "" || k.N == "" {
t.Errorf("jwk malformed: %+v", k)
// ES256: EC key with P-256 curve. Positively assert the algorithm
// change -- regression to RSA/RS256 must not pass.
if k.Kty != "EC" || k.Alg != "ES256" || k.Crv != "P-256" || k.Kid == "" || k.X == "" || k.Y == "" {
t.Errorf("jwk malformed (want EC/ES256/P-256): %+v", k)
}
}

Expand Down
3 changes: 2 additions & 1 deletion internal/credentials/resolver_extra_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package credentials

import (
"context"
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
Expand All @@ -24,7 +25,7 @@ type stubOIDCSigner struct{}
func (stubOIDCSigner) Sign(context.Context, []byte) ([]byte, error) {
return nil, assert.AnError
}
func (stubOIDCSigner) PublicKey(context.Context) (*rsa.PublicKey, error) {
func (stubOIDCSigner) PublicKey(context.Context) (crypto.PublicKey, error) {
return nil, assert.AnError
}
func (stubOIDCSigner) KeyID(context.Context) (string, error) {
Expand Down
29 changes: 16 additions & 13 deletions internal/oidc/aws_signer.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,8 @@ package oidc

import (
"context"
"crypto/rsa"
"crypto"
"crypto/ecdsa"
"crypto/x509"
"fmt"
"sync"
Expand All @@ -24,15 +25,16 @@ type AWSKMSClient interface {
type AWSKMSSigner struct {
client AWSKMSClient
err error
pubKey *rsa.PublicKey
pubKey *ecdsa.PublicKey
keyID string
kid string
once sync.Once
}

// NewAWSKMSSigner constructs a signer bound to the given KMS key. The
// keyID may be the key ARN, alias ARN, or alias name — anything
// accepted by kms:Sign and kms:GetPublicKey.
// keyID may be the key ARN, alias ARN, or alias name -- anything
// accepted by kms:Sign and kms:GetPublicKey. The KMS key must be an
// ECC_NIST_P256 key configured for signing (ES256).
func NewAWSKMSSigner(ctx context.Context, keyID string) (*AWSKMSSigner, error) {
if keyID == "" {
return nil, fmt.Errorf("oidc: empty AWS KMS keyID")
Expand All @@ -51,24 +53,25 @@ func NewAWSKMSSignerFromClient(client AWSKMSClient, keyID string) *AWSKMSSigner
}

// Sign calls kms:Sign with the raw SHA-256 digest and the signing
// algorithm RSASSA_PKCS1_V1_5_SHA_256, which matches what RS256 JWS
// signatures expect.
// algorithm ECDSA_SHA_256. AWS KMS returns a DER/ASN.1-encoded ECDSA
// signature, which is converted to the RFC 7518 section 3.4 raw R || S
// form required by ES256 JWS signatures (the Signer.Sign contract).
func (s *AWSKMSSigner) Sign(ctx context.Context, digest []byte) ([]byte, error) {
out, err := s.client.Sign(ctx, &kms.SignInput{
KeyId: &s.keyID,
Message: digest,
MessageType: types.MessageTypeDigest,
SigningAlgorithm: types.SigningAlgorithmSpecRsassaPkcs1V15Sha256,
SigningAlgorithm: types.SigningAlgorithmSpecEcdsaSha256,
})
if err != nil {
return nil, fmt.Errorf("oidc: kms:Sign: %w", err)
}
return out.Signature, nil
return derToRawECDSASignature(out.Signature)
}

// PublicKey fetches the public half of the KMS key once and caches it.
// Subsequent calls return the cached value.
func (s *AWSKMSSigner) PublicKey(ctx context.Context) (*rsa.PublicKey, error) {
func (s *AWSKMSSigner) PublicKey(ctx context.Context) (crypto.PublicKey, error) {
s.resolveOnce(ctx)
return s.pubKey, s.err
}
Expand All @@ -91,17 +94,17 @@ func (s *AWSKMSSigner) resolveOnce(ctx context.Context) {
s.err = fmt.Errorf("oidc: parse kms public key: %w", err)
return
}
rsaPub, ok := pub.(*rsa.PublicKey)
ecPub, ok := pub.(*ecdsa.PublicKey)
if !ok {
s.err = fmt.Errorf("oidc: kms key is not RSA (got %T)", pub)
s.err = fmt.Errorf("oidc: kms key is not ECDSA (got %T); key must be ECC_NIST_P256", pub)
return
}
kid, err := ComputeKeyID(rsaPub)
kid, err := ComputeKeyID(ecPub)
if err != nil {
s.err = err
return
}
s.pubKey = rsaPub
s.pubKey = ecPub
s.kid = kid
})
}
76 changes: 20 additions & 56 deletions internal/oidc/aws_signer_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,28 +2,24 @@ package oidc

import (
"context"
"crypto"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"testing"

"github.com/aws/aws-sdk-go-v2/service/kms"
)

var base64RawURL = base64.RawURLEncoding

// fakeKMSClient is a minimal AWSKMSClient backed by an in-process RSA
// key. It lets TestAWSKMSSigner exercise the Signer contract without
// touching real AWS.
// fakeKMSClient is a minimal AWSKMSClient backed by an in-process P-256
// ECDSA key. It lets TestAWSKMSSigner exercise the Signer contract without
// touching real AWS. Like real AWS KMS, it returns a DER/ASN.1 signature.
type fakeKMSClient struct {
key *rsa.PrivateKey
key *ecdsa.PrivateKey
}

func (f *fakeKMSClient) Sign(_ context.Context, in *kms.SignInput, _ ...func(*kms.Options)) (*kms.SignOutput, error) {
sig, err := rsa.SignPKCS1v15(rand.Reader, f.key, crypto.SHA256, in.Message)
sig, err := ecdsa.SignASN1(rand.Reader, f.key, in.Message)
if err != nil {
return nil, err
}
Expand All @@ -40,19 +36,23 @@ func (f *fakeKMSClient) GetPublicKey(_ context.Context, _ *kms.GetPublicKeyInput

func TestAWSKMSSignerRoundTrip(t *testing.T) {
ctx := context.Background()
key, err := rsa.GenerateKey(rand.Reader, 2048)
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("gen key: %v", err)
}
signer := NewAWSKMSSignerFromClient(&fakeKMSClient{key: key}, "alias/test-key")

// Signer contract: PublicKey returns the RSA pub half.
pub, err := signer.PublicKey(ctx)
// Signer contract: PublicKey returns the ECDSA pub half.
rawPub, err := signer.PublicKey(ctx)
if err != nil {
t.Fatalf("public key: %v", err)
}
if pub.N.Cmp(key.N) != 0 {
t.Fatal("public key modulus mismatch")
ecPub, ok := rawPub.(*ecdsa.PublicKey)
if !ok {
t.Fatalf("public key is not *ecdsa.PublicKey, got %T", rawPub)
}
if !ecPub.Equal(&key.PublicKey) {
t.Fatal("public key point mismatch")
}

// KeyID stable across calls.
Expand All @@ -62,7 +62,7 @@ func TestAWSKMSSignerRoundTrip(t *testing.T) {
t.Errorf("kid unstable or empty: %s vs %s", k1, k2)
}

// Mint a JWT and verify the signature end-to-end.
// Mint a JWT and verify the ECDSA signature end-to-end.
jws, err := Mint(ctx, signer, map[string]any{
"iss": "https://cudly.example.com",
"sub": "cudly-controller",
Expand All @@ -71,43 +71,7 @@ func TestAWSKMSSignerRoundTrip(t *testing.T) {
if err != nil {
t.Fatalf("mint: %v", err)
}
// Verify using the underlying pub half.
parts := splitJWS(t, jws)
signingInput := parts[0] + "." + parts[1]
digest := sha256.Sum256([]byte(signingInput))
if err := rsa.VerifyPKCS1v15(pub, crypto.SHA256, digest[:], decodeB64(t, parts[2])); err != nil {
t.Errorf("signature verify: %v", err)
}
}

// helpers — unit tests only, kept private.
func splitJWS(t *testing.T, jws string) [3]string {
t.Helper()
var out [3]string
last := 0
idx := 0
for i := 0; i < len(jws); i++ {
if jws[i] == '.' {
if idx >= 3 {
t.Fatalf("too many dots in JWS: %q", jws)
}
out[idx] = jws[last:i]
idx++
last = i + 1
}
}
if idx != 2 {
t.Fatalf("expected 2 dots in JWS, got %d: %q", idx, jws)
}
out[2] = jws[last:]
return out
}

func decodeB64(t *testing.T, s string) []byte {
t.Helper()
b, err := base64RawURL.DecodeString(s)
if err != nil {
t.Fatalf("b64 decode: %v", err)
}
return b
// The AWS KMS fake returns DER (ecdsa.SignASN1); the signer must
// convert it to the RFC 7518 raw R||S form before Mint encodes it.
assertRawES256JWS(t, jws, ecPub)
}
95 changes: 63 additions & 32 deletions internal/oidc/azure_factory_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,9 @@ package oidc

import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/rsa"
"math/big"
"testing"

"github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azkeys"
Expand All @@ -13,73 +13,104 @@ import (
)

// fakeAzureKeyVaultClient is a minimal AzureKeyVaultClient backed by an
// in-process RSA key. Used to exercise resolveOnce without a real Key Vault.
// in-process EC key. Used to exercise resolveOnce without a real Key Vault.
type fakeAzureKeyVaultClient struct {
signErr error
key *rsa.PublicKey
eBytes []byte
key *ecdsa.PublicKey
// xBytes and yBytes allow callers to inject nil to simulate incomplete
// responses from the Key Vault API.
xBytes []byte
yBytes []byte
nilKey bool // if true, return a KeyBundle with Key==nil
}

func (f *fakeAzureKeyVaultClient) Sign(_ context.Context, _, _ string, _ azkeys.SignParameters, _ *azkeys.SignOptions) (azkeys.SignResponse, error) {
return azkeys.SignResponse{}, f.signErr
}

func (f *fakeAzureKeyVaultClient) GetKey(_ context.Context, _, _ string, _ *azkeys.GetKeyOptions) (azkeys.GetKeyResponse, error) {
eBytes := f.eBytes
if eBytes == nil {
// Normal path: real exponent from the RSA key.
e := big.NewInt(int64(f.key.E))
eBytes = e.Bytes()
if f.nilKey {
return azkeys.GetKeyResponse{
KeyBundle: azkeys.KeyBundle{Key: nil},
}, nil
}
xBytes := f.xBytes
yBytes := f.yBytes
if xBytes == nil || yBytes == nil {
// Derive the fixed-width coordinates from the uncompressed SEC 1
// point (0x04 || X || Y) via crypto/ecdh, matching ComputeKeyID
// and avoiding the deprecated ecdsa.PublicKey.X/Y fields.
ecdhKey, err := f.key.ECDH()
if err != nil {
return azkeys.GetKeyResponse{}, err
}
uncompressed := ecdhKey.Bytes()
byteLen := (f.key.Curve.Params().BitSize + 7) / 8
if xBytes == nil {
xBytes = uncompressed[1 : 1+byteLen]
}
if yBytes == nil {
yBytes = uncompressed[1+byteLen:]
}
}
// Use sentinel value to represent "caller explicitly passed nil" vs
// "caller didn't override" -- an empty slice signals nil field.
var jwkX, jwkY []byte
if f.xBytes != nil || xBytes != nil {
jwkX = xBytes
}
if f.yBytes != nil || yBytes != nil {
jwkY = yBytes
}
keyBundle := azkeys.JSONWebKey{
N: f.key.N.Bytes(),
E: eBytes,
X: jwkX,
Y: jwkY,
}
return azkeys.GetKeyResponse{
KeyBundle: azkeys.KeyBundle{Key: &keyBundle},
}, nil
}

// ---------------------------------------------------------------------------
// M6 — Azure public-exponent overflow guard
// M6 -- Azure EC key completeness guard (replaces RSA exponent-range check)
// ---------------------------------------------------------------------------

// TestAzureSigner_ExponentRange verifies that resolveOnce rejects oversized or
// non-positive exponents before constructing the rsa.PublicKey (03-M6).
func TestAzureSigner_ExponentRange(t *testing.T) {
key, err := rsa.GenerateKey(rand.Reader, 2048)
// TestAzureSigner_ECKeyCompleteness verifies that resolveOnce rejects Key
// Vault responses that are missing the EC public-key coordinates, which
// would otherwise produce a zero-valued *ecdsa.PublicKey and sign invalid
// tokens. A valid response with both X and Y must be accepted.
func TestAzureSigner_ECKeyCompleteness(t *testing.T) {
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
require.NoError(t, err)

cases := []struct {
client *fakeAzureKeyVaultClient
name string
errSubstr string
eBytes []byte
wantErr bool
}{
{
name: "normal exponent 65537 accepted",
eBytes: big.NewInt(65537).Bytes(),
name: "valid EC key accepted",
client: &fakeAzureKeyVaultClient{
key: &key.PublicKey,
},
wantErr: false,
},
{
name: "exponent 0 rejected",
eBytes: big.NewInt(0).Bytes(),
wantErr: true,
errSubstr: "exponent",
},
{
name: "exponent too large (> MaxInt32) rejected",
eBytes: new(big.Int).Add(big.NewInt(0x7fffffff), big.NewInt(1)).Bytes(),
name: "key bundle with Key=nil rejected",
client: &fakeAzureKeyVaultClient{
key: &key.PublicKey,
nilKey: true,
},
wantErr: true,
errSubstr: "exponent",
errSubstr: "missing X or Y",
},
}

for _, tc := range cases {
tc := tc
t.Run(tc.name, func(t *testing.T) {
client := &fakeAzureKeyVaultClient{key: &key.PublicKey, eBytes: tc.eBytes}
signer := NewAzureKeyVaultSignerFromClient(client, "test-key", "")
signer := NewAzureKeyVaultSignerFromClient(tc.client, "test-key", "")
ctx := context.Background()

_, err := signer.PublicKey(ctx)
Expand All @@ -94,7 +125,7 @@ func TestAzureSigner_ExponentRange(t *testing.T) {
}

// ---------------------------------------------------------------------------
// M7 — factory precise error for half-configured Azure
// M7 -- factory precise error for half-configured Azure
// ---------------------------------------------------------------------------

// TestNewSignerFromEnv_AzureHalfConfigured verifies that exactly one of the
Expand Down
Loading
Loading