Skip to content

Commit 11641f8

Browse files
Add RSA crypto.Signer support for KMS/HSM
Check public key type instead of private key type to support crypto.Signer implementations (e.g. GCP KMS, AWS KMS, HSM) that aren't concrete *rsa.PrivateKey types. Only RSA keys are supported for crypto.Signer since major SAML IdPs (Azure AD, Auth0, Okta) use RSA signing. Non-RSA crypto.Signer keys return a clear error.
1 parent e839d2c commit 11641f8

5 files changed

Lines changed: 306 additions & 15 deletions

File tree

samlsp/middleware_test.go

Lines changed: 213 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,11 @@ package samlsp
22

33
import (
44
"bytes"
5+
"crypto"
6+
"crypto/ecdsa"
7+
"crypto/ed25519"
8+
"crypto/elliptic"
9+
"crypto/rand"
510
"crypto/rsa"
611
"crypto/sha256"
712
"crypto/x509"
@@ -17,6 +22,7 @@ import (
1722
"testing"
1823
"time"
1924

25+
"github.com/golang-jwt/jwt/v5"
2026
dsig "github.com/russellhaering/goxmldsig"
2127
"gotest.tools/assert"
2228
is "gotest.tools/assert/cmp"
@@ -520,3 +526,210 @@ func TestMiddlewareHandlesInvalidResponse(t *testing.T) {
520526
assert.Check(t, is.Equal("", resp.Header().Get("Location")))
521527
assert.Check(t, is.Equal("", resp.Header().Get("Set-Cookie")))
522528
}
529+
530+
type mockSigner struct {
531+
signer crypto.Signer
532+
}
533+
534+
func (m *mockSigner) Public() crypto.PublicKey {
535+
return m.signer.Public()
536+
}
537+
538+
func (m *mockSigner) Sign(rand io.Reader, digest []byte, opts crypto.SignerOpts) ([]byte, error) {
539+
return m.signer.Sign(rand, digest, opts)
540+
}
541+
542+
func newMockRSASigner(t *testing.T) crypto.Signer {
543+
key := mustParsePrivateKey(golden.Get(t, "key.pem"))
544+
return &mockSigner{signer: key.(crypto.Signer)}
545+
}
546+
547+
func TestMiddleware_WithCryptoSignerE2E(t *testing.T) {
548+
saml.TimeNow = func() time.Time {
549+
rv, _ := time.Parse("Mon Jan 2 15:04:05.999999999 MST 2006", "Mon Dec 1 01:57:09.123456789 UTC 2015")
550+
return rv
551+
}
552+
saml.Clock = dsig.NewFakeClockAt(saml.TimeNow())
553+
saml.RandReader = &testRandomReader{}
554+
555+
cert := mustParseCertificate(golden.Get(t, "cert.pem"))
556+
idpMetadata := golden.Get(t, "idp_metadata.xml")
557+
558+
var metadata saml.EntityDescriptor
559+
if err := xml.Unmarshal(idpMetadata, &metadata); err != nil {
560+
panic(err)
561+
}
562+
563+
mockSigner := newMockRSASigner(t)
564+
565+
opts := Options{
566+
URL: mustParseURL("https://15661444.ngrok.io/"),
567+
Key: mockSigner,
568+
Certificate: cert,
569+
IDPMetadata: &metadata,
570+
}
571+
572+
middleware, err := New(opts)
573+
assert.Check(t, err)
574+
575+
sessionProvider := DefaultSessionProvider(opts)
576+
sessionProvider.Name = "ttt"
577+
sessionProvider.MaxAge = 7200 * time.Second
578+
579+
sessionCodec := sessionProvider.Codec.(JWTSessionCodec)
580+
sessionCodec.MaxAge = 7200 * time.Second
581+
sessionProvider.Codec = sessionCodec
582+
583+
middleware.Session = sessionProvider
584+
middleware.ServiceProvider.MetadataURL.Path = "/saml2/metadata"
585+
middleware.ServiceProvider.AcsURL.Path = "/saml2/acs"
586+
middleware.ServiceProvider.SloURL.Path = "/saml2/slo"
587+
588+
t.Run("SessionEncodeDecode", func(t *testing.T) {
589+
var tc JWTSessionClaims
590+
if err := json.Unmarshal(golden.Get(t, "token.json"), &tc); err != nil {
591+
t.Fatal(err)
592+
}
593+
594+
encoded, err := sessionProvider.Codec.Encode(tc)
595+
assert.Check(t, err)
596+
assert.Assert(t, encoded != "")
597+
598+
decoded, err := sessionProvider.Codec.Decode(encoded)
599+
assert.Check(t, err)
600+
decodedClaims := decoded.(JWTSessionClaims)
601+
assert.Equal(t, tc.Subject, decodedClaims.Subject)
602+
})
603+
604+
t.Run("TrackedRequestEncodeDecode", func(t *testing.T) {
605+
codec := middleware.RequestTracker.(CookieRequestTracker).Codec
606+
trackedReq := TrackedRequest{
607+
Index: "test-index",
608+
SAMLRequestID: "test-request-id",
609+
URI: "/test-uri",
610+
}
611+
612+
encoded, err := codec.Encode(trackedReq)
613+
assert.Check(t, err)
614+
assert.Assert(t, encoded != "")
615+
616+
decoded, err := codec.Decode(encoded)
617+
assert.Check(t, err)
618+
assert.Equal(t, trackedReq.Index, decoded.Index)
619+
assert.Equal(t, trackedReq.SAMLRequestID, decoded.SAMLRequestID)
620+
})
621+
622+
t.Run("RequireAccountFlow", func(t *testing.T) {
623+
handler := middleware.RequireAccount(
624+
http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {
625+
panic("not reached")
626+
}))
627+
628+
req, _ := http.NewRequest("GET", "/protected", nil)
629+
resp := httptest.NewRecorder()
630+
handler.ServeHTTP(resp, req)
631+
632+
assert.Check(t, is.Equal(http.StatusFound, resp.Code))
633+
assert.Assert(t, resp.Header().Get("Location") != "")
634+
assert.Assert(t, resp.Header().Get("Set-Cookie") != "")
635+
})
636+
637+
t.Run("Metadata", func(t *testing.T) {
638+
req, _ := http.NewRequest("GET", "/saml2/metadata", nil)
639+
resp := httptest.NewRecorder()
640+
middleware.ServeHTTP(resp, req)
641+
642+
assert.Check(t, is.Equal(http.StatusOK, resp.Code))
643+
assert.Check(t, is.Equal("application/samlmetadata+xml",
644+
resp.Header().Get("Content-type")))
645+
golden.Assert(t, resp.Body.String(), "expected_middleware_metadata.xml")
646+
})
647+
}
648+
649+
func TestJWTSessionCodec_Ed25519(t *testing.T) {
650+
now := time.Now()
651+
saml.TimeNow = func() time.Time {
652+
return now
653+
}
654+
655+
// Generate Ed25519 key pair
656+
_, privateKey, err := ed25519.GenerateKey(rand.Reader)
657+
assert.Check(t, err)
658+
659+
audience := "https://example.com/"
660+
codec := JWTSessionCodec{
661+
SigningMethod: jwt.SigningMethodEdDSA,
662+
Audience: audience,
663+
Issuer: audience,
664+
MaxAge: time.Hour,
665+
Key: privateKey,
666+
}
667+
668+
// Create test claims directly
669+
tc := JWTSessionClaims{
670+
RegisteredClaims: jwt.RegisteredClaims{
671+
Audience: jwt.ClaimStrings{audience},
672+
Issuer: audience,
673+
Subject: "test-subject-123",
674+
IssuedAt: jwt.NewNumericDate(now),
675+
ExpiresAt: jwt.NewNumericDate(now.Add(time.Hour)),
676+
NotBefore: jwt.NewNumericDate(now),
677+
},
678+
Attributes: Attributes{
679+
"uid": []string{"testuser"},
680+
"givenName": []string{"Test User"},
681+
},
682+
SAMLSession: true,
683+
}
684+
685+
// Test encode
686+
encoded, err := codec.Encode(tc)
687+
assert.Check(t, err)
688+
assert.Assert(t, encoded != "", "encoded token should not be empty")
689+
690+
// Test decode
691+
decoded, err := codec.Decode(encoded)
692+
assert.Check(t, err)
693+
decodedClaims := decoded.(JWTSessionClaims)
694+
695+
// Verify claims match
696+
assert.Equal(t, tc.Subject, decodedClaims.Subject)
697+
assert.Check(t, decodedClaims.SAMLSession, "SAMLSession should be true")
698+
assert.Equal(t, tc.Attributes.Get("uid"), decodedClaims.Attributes.Get("uid"))
699+
}
700+
701+
func TestJWTSessionCodec_NonRSACryptoSignerReturnsError(t *testing.T) {
702+
now := time.Now()
703+
saml.TimeNow = func() time.Time {
704+
return now
705+
}
706+
707+
ecKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
708+
assert.Check(t, err)
709+
710+
signer := &mockSigner{signer: ecKey}
711+
712+
audience := "https://example.com/"
713+
codec := JWTSessionCodec{
714+
SigningMethod: jwt.SigningMethodES256,
715+
Audience: audience,
716+
Issuer: audience,
717+
MaxAge: time.Hour,
718+
Key: signer,
719+
}
720+
721+
tc := JWTSessionClaims{
722+
RegisteredClaims: jwt.RegisteredClaims{
723+
Audience: jwt.ClaimStrings{audience},
724+
Issuer: audience,
725+
Subject: "test",
726+
IssuedAt: jwt.NewNumericDate(now),
727+
ExpiresAt: jwt.NewNumericDate(now.Add(time.Hour)),
728+
NotBefore: jwt.NewNumericDate(now),
729+
},
730+
SAMLSession: true,
731+
}
732+
733+
_, err = codec.Encode(tc)
734+
assert.Check(t, is.ErrorContains(err, "crypto.Signer must hold an RSA key"))
735+
}

samlsp/new.go

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -149,15 +149,18 @@ func DefaultServiceProvider(opts Options) saml.ServiceProvider {
149149
}
150150

151151
func defaultSigningMethodForKey(key crypto.Signer) string {
152-
switch key.(type) {
153-
case *rsa.PrivateKey:
152+
if key == nil {
153+
return ""
154+
}
155+
// Check public key type to support crypto.Signer implementations (KMS/HSM)
156+
// that aren't concrete *rsa.PrivateKey or *ecdsa.PrivateKey types
157+
switch key.Public().(type) {
158+
case *rsa.PublicKey:
154159
return dsig.RSASHA1SignatureMethod
155-
case *ecdsa.PrivateKey:
160+
case *ecdsa.PublicKey:
156161
return dsig.ECDSASHA256SignatureMethod
157-
case nil:
158-
return ""
159162
default:
160-
panic(fmt.Sprintf("programming error: unsupported key type %T", key))
163+
panic(fmt.Sprintf("programming error: unsupported public key type %T", key.Public()))
161164
}
162165
}
163166

samlsp/request_tracker_jwt.go

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,9 @@ package samlsp
22

33
import (
44
"crypto"
5+
"crypto/ecdsa"
6+
"crypto/ed25519"
7+
"crypto/rsa"
58
"fmt"
69
"time"
710

@@ -44,7 +47,19 @@ func (s JWTTrackedRequestCodec) Encode(value TrackedRequest) (string, error) {
4447
SAMLAuthnRequest: true,
4548
}
4649
token := jwt.NewWithClaims(s.SigningMethod, claims)
47-
return token.SignedString(s.Key)
50+
51+
if s.Key == nil {
52+
return "", fmt.Errorf("signing key is nil")
53+
}
54+
55+
// Check if key is a concrete private key type that jwt library can handle directly.
56+
// For crypto.Signer implementations (KMS/HSM), use custom signing.
57+
switch s.Key.(type) {
58+
case *rsa.PrivateKey, *ecdsa.PrivateKey, ed25519.PrivateKey:
59+
return token.SignedString(s.Key)
60+
default:
61+
return signJWTWithCryptoSigner(token, s.Key, s.SigningMethod)
62+
}
4863
}
4964

5065
// Decode returns a Tracked request from an encoded string.

samlsp/session_jwt.go

Lines changed: 60 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,14 @@ package samlsp
22

33
import (
44
"crypto"
5+
"crypto/ecdsa"
6+
"crypto/ed25519"
7+
"crypto/rand"
8+
"crypto/rsa"
9+
"encoding/base64"
510
"errors"
11+
"fmt"
12+
"strings"
613
"time"
714

815
"github.com/golang-jwt/jwt/v5"
@@ -77,12 +84,19 @@ func (c JWTSessionCodec) Encode(s Session) (string, error) {
7784
claims := s.(JWTSessionClaims) // this will panic if you pass the wrong kind of session
7885

7986
token := jwt.NewWithClaims(c.SigningMethod, claims)
80-
signedString, err := token.SignedString(c.Key)
81-
if err != nil {
82-
return "", err
87+
88+
if c.Key == nil {
89+
return "", fmt.Errorf("signing key is nil")
8390
}
8491

85-
return signedString, nil
92+
// Check if key is a concrete private key type that jwt library can handle directly.
93+
// For crypto.Signer implementations (KMS/HSM), use custom signing.
94+
switch c.Key.(type) {
95+
case *rsa.PrivateKey, *ecdsa.PrivateKey, ed25519.PrivateKey:
96+
return token.SignedString(c.Key)
97+
default:
98+
return signJWTWithCryptoSigner(token, c.Key, c.SigningMethod)
99+
}
86100
}
87101

88102
// Decode parses the serialized session that may have been returned by Encode
@@ -137,3 +151,45 @@ func (a Attributes) Get(key string) string {
137151
}
138152
return v[0]
139153
}
154+
155+
// signJWTWithCryptoSigner signs a JWT token using the crypto.Signer interface.
156+
// Only RSA signing methods are supported since major SAML IdPs (Azure AD, Auth0,
157+
// Okta) use RSA. This allows KMS/HSM keys that implement crypto.Signer to sign JWTs.
158+
func signJWTWithCryptoSigner(token *jwt.Token, signer crypto.Signer, method jwt.SigningMethod) (string, error) {
159+
if _, ok := signer.Public().(*rsa.PublicKey); !ok {
160+
return "", fmt.Errorf("crypto.Signer must hold an RSA key, got %T", signer.Public())
161+
}
162+
163+
// Get the signing string (header.payload)
164+
signingString, err := token.SigningString()
165+
if err != nil {
166+
return "", err
167+
}
168+
169+
// Determine hash algorithm based on signing method
170+
var hashFunc crypto.Hash
171+
switch method.Alg() {
172+
case "RS256":
173+
hashFunc = crypto.SHA256
174+
case "RS384":
175+
hashFunc = crypto.SHA384
176+
case "RS512":
177+
hashFunc = crypto.SHA512
178+
default:
179+
return "", fmt.Errorf("unsupported signing algorithm for crypto.Signer: %s", method.Alg())
180+
}
181+
182+
// Hash the signing string
183+
hasher := hashFunc.New()
184+
hasher.Write([]byte(signingString))
185+
digest := hasher.Sum(nil)
186+
187+
// Sign using crypto.Signer
188+
sig, err := signer.Sign(rand.Reader, digest, hashFunc)
189+
if err != nil {
190+
return "", fmt.Errorf("signing with crypto.Signer: %w", err)
191+
}
192+
193+
// Encode signature and return complete JWT
194+
return strings.Join([]string{signingString, base64.RawURLEncoding.EncodeToString(sig)}, "."), nil
195+
}

0 commit comments

Comments
 (0)