@@ -2,6 +2,11 @@ package samlsp
22
33import (
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+ }
0 commit comments