package auth // Reference: P2E ยง1.6 โ€” SAML 2.0 Service Provider integration tests // Tests cover: config validation, service initialization (enabled/disabled), // attribute mapping, user find/create, replay prevention helpers. import ( "crypto/rand" "crypto/rsa" "crypto/x509" "encoding/pem" "fmt" "math/big" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/gochat/gochat/internal/config" "github.com/gochat/gochat/internal/model" "gorm.io/driver/sqlite" "gorm.io/gorm" ) // --- Test helpers --- func newTestSAMLConfig() *config.SAMLConfig { return &config.SAMLConfig{ Enabled: false, // disabled by default for most tests IdPMetadataURL: "", IdPMetadataXML: testIdPMetadataXML(), SPEntityID: "https://gochat.test/saml", ACSURL: "https://gochat.test/api/v1/saml/acs", SPPrivateKey: "", SPCertificate: "", ClockDriftTolerance: 180, AttributeMap: config.SAMLAttributeMap{ Email: "email", DisplayName: "displayName", FirstName: "firstName", LastName: "lastName", }, } } func newTestSAMLConfigEnabled() *config.SAMLConfig { cfg := newTestSAMLConfig() cfg.Enabled = true cfg.SPPrivateKey = generateTestPrivateKeyPEM() cfg.SPCertificate = generateTestCertificatePEM() return cfg } func newTestDB(t *testing.T) *gorm.DB { db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) require.NoError(t, err) err = db.AutoMigrate(&model.User{}) require.NoError(t, err) return db } // Generate a self-signed RSA key + cert for testing func generateTestRSAPair() (*rsa.PrivateKey, *x509.Certificate) { key, _ := rsa.GenerateKey(rand.Reader, 2048) cert := &x509.Certificate{ SerialNumber: big.NewInt(1), NotBefore: time.Now(), NotAfter: time.Now().Add(365 * 24 * time.Hour), IsCA: true, BasicConstraintsValid: true, } return key, cert } func generateTestPrivateKeyPEM() string { key, _ := rsa.GenerateKey(rand.Reader, 2048) keyBytes := x509.MarshalPKCS1PrivateKey(key) return string(pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: keyBytes})) } func generateTestCertificatePEM() string { key, cert := generateTestRSAPair() certBytes, _ := x509.CreateCertificate(rand.Reader, cert, cert, &key.PublicKey, key) return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certBytes})) } // Minimal IdP metadata XML for testing func testIdPMetadataXML() string { return ` urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress ` } // --- Config Tests --- func TestSAMLConfig_ClockDriftDuration(t *testing.T) { cfg := newTestSAMLConfig() // Default (180 seconds) assert.Equal(t, 180*time.Second, cfg.ClockDriftDuration()) // Explicit value cfg.ClockDriftTolerance = 300 assert.Equal(t, 300*time.Second, cfg.ClockDriftDuration()) // Zero value โ†’ default 180s cfg.ClockDriftTolerance = 0 assert.Equal(t, 180*time.Second, cfg.ClockDriftDuration()) } // --- Service Initialization Tests --- func TestNewSAMLService_Disabled(t *testing.T) { cfg := newTestSAMLConfig() svc, err := NewSAMLService(cfg, nil, nil) require.NoError(t, err) assert.NotNil(t, svc) assert.False(t, svc.cfg.Enabled) assert.Nil(t, svc.idpMetadata) // no IdP metadata when disabled } func TestNewSAMLService_Enabled_InvalidConfig(t *testing.T) { tests := []struct { name string modify func(cfg *config.SAMLConfig) wantErr error }{ { name: "missing sp_entity_id", modify: func(cfg *config.SAMLConfig) { cfg.SPEntityID = "" }, wantErr: ErrSAMLInvalidConfig, }, { name: "missing acs_url", modify: func(cfg *config.SAMLConfig) { cfg.ACSURL = "" }, wantErr: ErrSAMLInvalidConfig, }, { name: "missing both idp metadata sources", modify: func(cfg *config.SAMLConfig) { cfg.IdPMetadataURL = ""; cfg.IdPMetadataXML = "" }, wantErr: ErrSAMLIdPMetadata, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { cfg := newTestSAMLConfigEnabled() tt.modify(cfg) svc, err := NewSAMLService(cfg, nil, nil) assert.Nil(t, svc) assert.ErrorIs(t, err, tt.wantErr) }) } } func TestNewSAMLService_Enabled_ValidConfig(t *testing.T) { cfg := newTestSAMLConfigEnabled() svc, err := NewSAMLService(cfg, nil, nil) require.NoError(t, err) assert.NotNil(t, svc) assert.True(t, svc.cfg.Enabled) assert.NotNil(t, svc.idpMetadata) } // --- Operation Tests (disabled service should reject) --- func TestSAMLService_Disabled_InitiateLogin(t *testing.T) { cfg := newTestSAMLConfig() svc, _ := NewSAMLService(cfg, nil, nil) _, err := svc.InitiateLogin("test-state") assert.ErrorIs(t, err, ErrSAMLEnabled) } func TestSAMLService_Disabled_ProcessResponse(t *testing.T) { cfg := newTestSAMLConfig() svc, _ := NewSAMLService(cfg, nil, nil) _, err := svc.ProcessResponse("fake-response") assert.ErrorIs(t, err, ErrSAMLEnabled) } func TestSAMLService_Disabled_GetSPMetadata(t *testing.T) { cfg := newTestSAMLConfig() svc, _ := NewSAMLService(cfg, nil, nil) _, err := svc.GetSPMetadata() assert.ErrorIs(t, err, ErrSAMLEnabled) } // --- Attribute Extraction Tests --- func TestExtractAttributes(t *testing.T) { statements := []SAMLAttributeStatementXML{ { Attributes: []SAMLAttributeXML{ { Name: "email", FriendlyName: "Email Address", Values: []SAMLAttributeValueXML{{Value: "user@test.com"}}, }, { Name: "displayName", Values: []SAMLAttributeValueXML{{Value: "Test User"}}, }, { Name: "firstName", Values: []SAMLAttributeValueXML{{Value: "Test"}}, }, { Name: "lastName", Values: []SAMLAttributeValueXML{{Value: "User"}}, }, { Name: "orgRole", FriendlyName: "Organization Role", Values: []SAMLAttributeValueXML{{Value: "admin"}}, }, }, }, } attrs := extractAttributesFromXML(statements) assert.Equal(t, "user@test.com", attrs["email"]) assert.Equal(t, "user@test.com", attrs["Email Address"]) // FriendlyName too assert.Equal(t, "Test User", attrs["displayName"]) assert.Equal(t, "Test", attrs["firstName"]) assert.Equal(t, "User", attrs["lastName"]) assert.Equal(t, "admin", attrs["orgRole"]) assert.Equal(t, "admin", attrs["Organization Role"]) } func TestGetAttribute(t *testing.T) { statements := []SAMLAttributeStatementXML{ { Attributes: []SAMLAttributeXML{ { Name: "email", FriendlyName: "mail", Values: []SAMLAttributeValueXML{{Value: "admin@corp.com"}}, }, }, }, } attrs := extractAttributesFromXML(statements) // Lookup by Name assert.Equal(t, "admin@corp.com", getAttributeFromXML(attrs, "email")) // Lookup by FriendlyName assert.Equal(t, "admin@corp.com", getAttributeFromXML(attrs, "mail")) // Missing attribute assert.Equal(t, "", getAttributeFromXML(attrs, "phone")) } func TestGetAttribute_MultipleStatements(t *testing.T) { statements := []SAMLAttributeStatementXML{ { Attributes: []SAMLAttributeXML{ { Name: "firstName", Values: []SAMLAttributeValueXML{{Value: "Alice"}}, }, }, }, { Attributes: []SAMLAttributeXML{ { Name: "lastName", Values: []SAMLAttributeValueXML{{Value: "Smith"}}, }, }, }, } attrs := extractAttributesFromXML(statements) assert.Equal(t, "Alice", getAttributeFromXML(attrs, "firstName")) assert.Equal(t, "Smith", getAttributeFromXML(attrs, "lastName")) } // --- PEM Parsing Tests --- func TestParseRSAPrivateKey_PKCS1(t *testing.T) { pemData := generateTestPrivateKeyPEM() key, err := parseRSAPrivateKey(pemData) require.NoError(t, err) assert.NotNil(t, key) assert.Equal(t, 2048, key.N.BitLen()) } func TestParseRSAPrivateKey_PKCS8(t *testing.T) { key, _ := rsa.GenerateKey(rand.Reader, 2048) keyBytes, _ := x509.MarshalPKCS8PrivateKey(key) pemData := string(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyBytes})) parsedKey, err := parseRSAPrivateKey(pemData) require.NoError(t, err) assert.NotNil(t, parsedKey) } func TestParseRSAPrivateKey_Invalid(t *testing.T) { _, err := parseRSAPrivateKey("not-a-pem") assert.Error(t, err) } func TestParseX509Certificate_Valid(t *testing.T) { pemData := generateTestCertificatePEM() cert, err := parseX509Certificate(pemData) require.NoError(t, err) assert.NotNil(t, cert) } func TestParseX509Certificate_Invalid(t *testing.T) { _, err := parseX509Certificate("not-a-pem") assert.Error(t, err) } // --- IdP Metadata Loading Tests --- func TestLoadIdPMetadata_InlineXML(t *testing.T) { cfg := &config.SAMLConfig{ IdPMetadataXML: testIdPMetadataXML(), } metadata, err := loadIdPMetadata(cfg) require.NoError(t, err) assert.NotNil(t, metadata) assert.Equal(t, "https://idp.test/saml", metadata.EntityID) } func TestLoadIdPMetadata_NoSource(t *testing.T) { cfg := &config.SAMLConfig{ IdPMetadataURL: "", IdPMetadataXML: "", } _, err := loadIdPMetadata(cfg) assert.Error(t, err) } func TestLoadIdPMetadata_InvalidXML(t *testing.T) { cfg := &config.SAMLConfig{ IdPMetadataXML: "not valid xml at all", } _, err := loadIdPMetadata(cfg) assert.Error(t, err) } // --- User FindOrCreate Tests --- func TestSAMLService_FindOrCreateUser_NewUser(t *testing.T) { db := newTestDB(t) cfg := newTestSAMLConfig() // disabled, but FindOrCreateUser only needs db svc := &SAMLService{cfg: cfg, db: db} userInfo := &SAMLUserInfo{ NameID: "alice@saml.test", Email: "alice@example.com", DisplayName: "Alice Smith", FirstName: "Alice", LastName: "Smith", Attributes: map[string]string{"email": "alice@example.com"}, } user, err := svc.FindOrCreateUser(userInfo) require.NoError(t, err) assert.NotZero(t, user.ID) assert.Equal(t, "alice@saml.test", user.UID) assert.Equal(t, "alice@example.com", user.Email) assert.Equal(t, "Alice Smith", user.Name) assert.Equal(t, "saml", user.Provider) assert.Equal(t, "agent", user.Role) assert.True(t, user.Active) } func TestSAMLService_FindOrCreateUser_ExistingByUID(t *testing.T) { db := newTestDB(t) // Create an existing SAML user existing := &model.User{ Name: "Old Name", Email: "bob@example.com", Provider: "saml", UID: "bob@saml.test", Role: "agent", Active: true, } require.NoError(t, db.Create(existing).Error) cfg := newTestSAMLConfig() svc := &SAMLService{cfg: cfg, db: db} userInfo := &SAMLUserInfo{ NameID: "bob@saml.test", Email: "bob@example.com", DisplayName: "Bob Updated", Attributes: map[string]string{}, } user, err := svc.FindOrCreateUser(userInfo) require.NoError(t, err) assert.Equal(t, existing.ID, user.ID) // same user assert.Equal(t, "Bob Updated", user.Name) // name updated } func TestSAMLService_FindOrCreateUser_LinkExistingEmail(t *testing.T) { db := newTestDB(t) // Create an email-authenticated user (no SAML yet) existing := &model.User{ Name: "Charlie Email", Email: "charlie@example.com", Provider: "email", Role: "agent", Active: true, } require.NoError(t, db.Create(existing).Error) cfg := newTestSAMLConfig() svc := &SAMLService{cfg: cfg, db: db} userInfo := &SAMLUserInfo{ NameID: "charlie@saml.test", Email: "charlie@example.com", DisplayName: "Charlie SAML", Attributes: map[string]string{}, } user, err := svc.FindOrCreateUser(userInfo) require.NoError(t, err) assert.Equal(t, existing.ID, user.ID) // linked same user assert.Equal(t, "saml", user.Provider) // provider updated to saml assert.Equal(t, "charlie@saml.test", user.UID) // UID set } func TestSAMLService_FindOrCreateUser_NoDB(t *testing.T) { cfg := newTestSAMLConfig() svc := &SAMLService{cfg: cfg, db: nil} userInfo := &SAMLUserInfo{ NameID: "nodb@test.com", Email: "nodb@test.com", } _, err := svc.FindOrCreateUser(userInfo) assert.Error(t, err) assert.Contains(t, err.Error(), "database not available") } // --- SAMLUserInfo Tests --- func TestSAMLUserInfo_DisplayNameFallback(t *testing.T) { // When displayName attribute is missing, compose from firstName + lastName userInfo := &SAMLUserInfo{ NameID: "fallback@test.com", FirstName: "John", LastName: "Doe", DisplayName: "", // empty โ€” should be composed } // DisplayName composition happens in ProcessResponse, not in the struct itself // But let's verify the logic separately name := "" if userInfo.DisplayName == "" && (userInfo.FirstName != "" || userInfo.LastName != "") { name = fmt.Sprintf("%s %s", userInfo.FirstName, userInfo.LastName) } assert.Equal(t, "John Doe", name) } // --- EncodeSAMLRequest helper test --- func TestEncodeSAMLRequest(t *testing.T) { encoded, err := EncodeSAMLRequest("") require.NoError(t, err) assert.NotEmpty(t, encoded) } // --- SP Metadata Generation Test (with enabled service) --- func TestSAMLService_Enabled_GetSPMetadata(t *testing.T) { cfg := newTestSAMLConfigEnabled() svc, err := NewSAMLService(cfg, nil, nil) require.NoError(t, err) xml, err := svc.GetSPMetadata() require.NoError(t, err) assert.NotEmpty(t, xml) assert.Contains(t, string(xml), "EntityDescriptor") assert.Contains(t, string(xml), cfg.SPEntityID) }