Files
gochat/internal/auth/saml_test.go
T
2026-06-04 15:44:48 +08:00

762 lines
25 KiB
Go

package auth
import (
"encoding/base64"
"encoding/json"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/gochat/gochat/internal/config"
"github.com/gochat/gochat/internal/model"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// --- Helpers for SAML tests ---
func newTestSAMLDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
require.NoError(t, err)
err = db.AutoMigrate(&model.AccountSamlSettings{}, &model.User{}, &model.Account{}, &model.AccountUser{})
require.NoError(t, err)
return db
}
func setupTestRedis(t *testing.T) (*miniredis.Miniredis, redis.Cmdable) {
t.Helper()
mr := miniredis.RunT(t)
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
t.Cleanup(func() { rdb.Close() })
return mr, rdb
}
func newTestSAMLConfig(enabled bool) *config.SAMLConfig {
return &config.SAMLConfig{
Enabled: enabled,
SPEntityID: "https://sp.example.com/saml",
ACSURL: "https://sp.example.com/saml/acs",
IdPMetadataURL: "",
IdPMetadataXML: "",
SPPrivateKey: "",
SPCertificate: "",
AttributeMap: config.SAMLAttributeMap{Email: "email", DisplayName: "displayName", FirstName: "firstName", LastName: "lastName"},
ClockDriftTolerance: 180,
}
}
// Sample IdP metadata XML for testing
const testIdPMetadataXML = `<EntityDescriptor entityID="https://idp.example.com/saml" xmlns="urn:oasis:names:tc:SAML:2.0:metadata">
<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" Location="https://idp.example.com/saml/sso"/>
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="https://idp.example.com/saml/sso/post"/>
<SingleLogoutService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" Location="https://idp.example.com/saml/slo"/>
<SingleLogoutService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="https://idp.example.com/saml/slo/post"/>
<KeyDescriptor use="signing">
<KeyInfo>
<X509Data>
<X509Certificate>MIIDXTCCAkWgAwIBAgIJAJC1HiIAZAiIMA0GCSqGSIb3DQEBCwUA</X509Certificate>
</X509Data>
</KeyInfo>
</KeyDescriptor>
</IDPSSODescriptor>
</EntityDescriptor>`
// --- NewSAMLService tests ---
func TestNewSAMLService_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
assert.NoError(t, err, "disabled SAML should not error")
assert.NotNil(t, svc)
}
func TestNewSAMLService_EnabledMissingSPEntityID(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.SPEntityID = ""
_, err := NewSAMLService(cfg, rdb, db)
assert.ErrorIs(t, err, ErrSAMLInvalidConfig, "missing SPEntityID should return ErrSAMLInvalidConfig")
}
func TestNewSAMLService_EnabledMissingACSURL(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.ACSURL = ""
_, err := NewSAMLService(cfg, rdb, db)
assert.ErrorIs(t, err, ErrSAMLInvalidConfig, "missing ACSURL should return ErrSAMLInvalidConfig")
}
func TestNewSAMLService_EnabledNoMetadataSource(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataURL = ""
cfg.IdPMetadataXML = ""
_, err := NewSAMLService(cfg, rdb, db)
assert.ErrorIs(t, err, ErrSAMLIdPMetadata, "no IdP metadata source should return ErrSAMLIdPMetadata")
}
func TestNewSAMLService_EnabledWithInlineMetadata(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
assert.NoError(t, err, "valid inline metadata should not error")
assert.NotNil(t, svc)
assert.Equal(t, "https://idp.example.com/saml", svc.GetIdPEntityID())
}
// --- GetIdPEntityID tests ---
func TestGetIdPEntityID_WithMetadata(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
entityID := svc.GetIdPEntityID()
assert.Equal(t, "https://idp.example.com/saml", entityID)
}
func TestGetIdPEntityID_NoMetadata(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
entityID := svc.GetIdPEntityID()
assert.Equal(t, "", entityID, "no metadata should return empty entityID")
}
// --- parseIdPMetadataXML tests ---
func TestParseIdPMetadataXML_Valid(t *testing.T) {
metadata, err := parseIdPMetadataXML([]byte(testIdPMetadataXML))
require.NoError(t, err)
assert.Equal(t, "https://idp.example.com/saml", metadata.EntityID)
assert.Equal(t, "https://idp.example.com/saml/sso", metadata.SSOURL, "should prefer HTTP-Redirect binding")
assert.Equal(t, "https://idp.example.com/saml/slo", metadata.SLORedirectURL)
assert.Equal(t, "https://idp.example.com/saml/slo/post", metadata.SLOPostURL)
assert.NotEmpty(t, metadata.Certificates, "should extract signing certificates")
}
func TestParseIdPMetadataXML_MissingEntityID(t *testing.T) {
xml := `<EntityDescriptor entityID="" xmlns="urn:oasis:names:tc:SAML:2.0:metadata">
<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" Location="https://idp.example.com/saml/sso"/>
</IDPSSODescriptor>
</EntityDescriptor>`
_, err := parseIdPMetadataXML([]byte(xml))
assert.Error(t, err, "missing entityID should error")
assert.Contains(t, err.Error(), "entityID")
}
func TestParseIdPMetadataXML_InvalidXML(t *testing.T) {
_, err := parseIdPMetadataXML([]byte("not valid xml"))
assert.Error(t, err, "invalid XML should error")
}
func TestParseIdPMetadataXML_NoIDPSSODescriptor(t *testing.T) {
xml := `<EntityDescriptor entityID="https://idp.example.com/saml" xmlns="urn:oasis:names:tc:SAML:2.0:metadata">
</EntityDescriptor>`
metadata, err := parseIdPMetadataXML([]byte(xml))
require.NoError(t, err)
assert.Equal(t, "https://idp.example.com/saml", metadata.EntityID)
assert.Empty(t, metadata.SSOURL, "no IDPSSODescriptor means no SSO URL")
assert.Empty(t, metadata.Certificates)
}
func TestParseIdPMetadataXML_FallbackSSOBinding(t *testing.T) {
// Only POST binding, no Redirect binding — should fallback to first SSO URL
xml := `<EntityDescriptor entityID="https://idp.example.com/saml" xmlns="urn:oasis:names:tc:SAML:2.0:metadata">
<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="https://idp.example.com/saml/sso/post"/>
</IDPSSODescriptor>
</EntityDescriptor>`
metadata, err := parseIdPMetadataXML([]byte(xml))
require.NoError(t, err)
assert.Equal(t, "https://idp.example.com/saml/sso/post", metadata.SSOURL, "should fallback to POST binding when no Redirect")
}
func TestParseIdPMetadataXML_CertificatePEMFormatting(t *testing.T) {
xml := `<EntityDescriptor entityID="https://idp.example.com/saml" xmlns="urn:oasis:names:tc:SAML:2.0:metadata">
<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" Location="https://idp.example.com/saml/sso"/>
<KeyDescriptor use="signing">
<KeyInfo>
<X509Data>
<X509Certificate>MIIDXTCCAkWgAwIBAgIJAJC1HiIAZAiIMA0GCSqGSIb3DQEBCwUA</X509Certificate>
</X509Data>
</KeyInfo>
</KeyDescriptor>
</IDPSSODescriptor>
</EntityDescriptor>`
metadata, err := parseIdPMetadataXML([]byte(xml))
require.NoError(t, err)
require.NotEmpty(t, metadata.Certificates)
// Certificate should be PEM-formatted
assert.Contains(t, metadata.Certificates[0], "-----BEGIN CERTIFICATE-----")
}
// --- extractAttributesFromXML tests ---
func TestExtractAttributesFromXML(t *testing.T) {
statements := []SAMLAttributeStatementXML{
{
Attributes: []SAMLAttributeXML{
{
Name: "email",
FriendlyName: "EmailAddress",
Values: []SAMLAttributeValueXML{{Value: "user@example.com"}},
},
{
Name: "displayName",
Values: []SAMLAttributeValueXML{{Value: "Test User"}},
},
},
},
{
Attributes: []SAMLAttributeXML{
{
Name: "firstName",
Values: []SAMLAttributeValueXML{{Value: "Test"}, {Value: "Extra"}},
},
},
},
}
attrs := extractAttributesFromXML(statements)
assert.Equal(t, "user@example.com", attrs["email"])
assert.Equal(t, "user@example.com", attrs["EmailAddress"], "should also store by FriendlyName")
assert.Equal(t, "Test User", attrs["displayName"])
assert.Equal(t, "Test", attrs["firstName"], "should use first value for multi-valued attributes")
}
func TestExtractAttributesFromXML_Empty(t *testing.T) {
attrs := extractAttributesFromXML(nil)
assert.Empty(t, attrs)
attrs = extractAttributesFromXML([]SAMLAttributeStatementXML{})
assert.Empty(t, attrs)
}
func TestExtractAttributesFromXML_NoValues(t *testing.T) {
statements := []SAMLAttributeStatementXML{
{
Attributes: []SAMLAttributeXML{
{Name: "email", Values: []SAMLAttributeValueXML{}},
},
},
}
attrs := extractAttributesFromXML(statements)
assert.NotContains(t, attrs, "email", "attribute with no values should not appear in map")
}
// --- getAttributeFromXML tests ---
func TestGetAttributeFromXML(t *testing.T) {
attrs := map[string]string{"email": "user@example.com", "name": "Test"}
assert.Equal(t, "user@example.com", getAttributeFromXML(attrs, "email"))
assert.Equal(t, "Test", getAttributeFromXML(attrs, "name"))
assert.Equal(t, "", getAttributeFromXML(attrs, "missing"), "missing attribute should return empty string")
}
// --- parseSAMLTime tests ---
func TestParseSAMLTime(t *testing.T) {
// RFC3339
ts, err := parseSAMLTime("2024-01-15T10:30:00Z")
require.NoError(t, err)
assert.Equal(t, time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC), ts)
// RFC3339Nano
ts, err = parseSAMLTime("2024-01-15T10:30:00.123456789Z")
require.NoError(t, err)
assert.True(t, ts.Year() == 2024)
// Simple UTC format
ts, err = parseSAMLTime("2024-01-15T10:30:00Z")
require.NoError(t, err)
assert.Equal(t, time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC), ts)
// Invalid format
_, err = parseSAMLTime("not-a-time")
assert.Error(t, err, "invalid time format should error")
}
func TestParseSAMLTime_WithOffset(t *testing.T) {
ts, err := parseSAMLTime("2024-01-15T10:30:00+05:00")
require.NoError(t, err)
// Should normalize to UTC
assert.Equal(t, time.Date(2024, 1, 15, 5, 30, 0, 0, time.UTC), ts.UTC())
}
// --- deflateAndBase64Encode tests ---
func TestDeflateAndBase64Encode(t *testing.T) {
input := "<samlp:AuthnRequest ID=\"test123\" IssueInstant=\"2024-01-15T10:30:00Z\"/>"
encoded, err := deflateAndBase64Encode(input)
require.NoError(t, err)
assert.NotEmpty(t, encoded, "encoded result should not be empty")
}
func TestDeflateAndBase64Encode_Empty(t *testing.T) {
encoded, err := deflateAndBase64Encode("")
require.NoError(t, err)
assert.NotEmpty(t, encoded, "even empty input produces encoded output")
}
// --- EncodeSAMLRequest tests ---
func TestEncodeSAMLRequest(t *testing.T) {
xml := "<samlp:AuthnRequest xmlns:samlp=\"urn:oasis:names:tc:SAML:2.0:protocol\" ID=\"test123\"/>"
encoded, err := EncodeSAMLRequest(xml)
require.NoError(t, err)
assert.NotEmpty(t, encoded)
// Simple base64 — verify it can be decoded
decoded, err := base64.StdEncoding.DecodeString(encoded)
require.NoError(t, err)
assert.Equal(t, xml, string(decoded))
}
// --- generateRandomID tests ---
func TestGenerateRandomID(t *testing.T) {
id1 := generateRandomID()
id2 := generateRandomID()
assert.NotEmpty(t, id1, "random ID should not be empty")
assert.NotEmpty(t, id2, "random ID should not be empty")
assert.NotEqual(t, id1, id2, "two random IDs should differ (extremely unlikely collision)")
}
// --- extractCertBase64FromPEM tests ---
func TestExtractCertBase64FromPEM(t *testing.T) {
pemData := "-----BEGIN CERTIFICATE-----\nMIIDXTCCAkWgAwIBAgIJAJC1HiIAZAiIMA0GCSqGSIb3DQEBCwUA\n-----END CERTIFICATE-----"
result := extractCertBase64FromPEM(pemData)
assert.NotEmpty(t, result, "should extract base64 data from PEM")
}
func TestExtractCertBase64FromPEM_NoPEMBlock(t *testing.T) {
rawData := "not-a-pem-block"
result := extractCertBase64FromPEM(rawData)
assert.Equal(t, rawData, result, "fallback should return raw data when no PEM block found")
}
// --- SAMLConfig ClockDriftDuration tests (in config package) ---
func TestSAMLConfigClockDriftDuration(t *testing.T) {
cfg := config.SAMLConfig{ClockDriftTolerance: 300}
assert.Equal(t, 300*time.Second, cfg.ClockDriftDuration())
cfg = config.SAMLConfig{ClockDriftTolerance: 0}
assert.Equal(t, 180*time.Second, cfg.ClockDriftDuration(), "default should be 180s")
cfg = config.SAMLConfig{ClockDriftTolerance: -1}
assert.Equal(t, 180*time.Second, cfg.ClockDriftDuration(), "negative should use default 180s")
}
// --- SAML InitiateLogin tests (disabled service) ---
func TestSAMLService_InitiateLogin_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
_, err = svc.InitiateLogin("test-state")
assert.ErrorIs(t, err, ErrSAMLEnabled, "disabled service should reject login")
}
// --- SAML ProcessResponse tests (disabled service) ---
func TestSAMLService_ProcessResponse_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
_, err = svc.ProcessResponse("some-response")
assert.ErrorIs(t, err, ErrSAMLEnabled, "disabled service should reject processing")
}
// --- SAML FindOrCreateUser tests ---
func TestSAMLService_FindOrCreateUser_NewUser(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
userInfo := &SAMLUserInfo{
NameID: "user123@idp.example.com",
Email: "samluser@example.com",
DisplayName: "SAML User",
FirstName: "SAML",
LastName: "User",
Attributes: map[string]string{"email": "samluser@example.com"},
}
user, err := svc.FindOrCreateUser(userInfo)
require.NoError(t, err)
assert.NotNil(t, user)
assert.Equal(t, "samluser@example.com", user.Email)
assert.Equal(t, "SAML User", user.Name)
assert.Equal(t, "saml", user.Provider)
assert.Equal(t, "user123@idp.example.com", user.UID)
}
func TestSAMLService_FindOrCreateUser_ExistingUserByEmail(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
// Create a user first
existingUser := &model.User{
Name: "Existing",
Email: "existing@example.com",
Password: "hashed",
Provider: "email",
Active: true,
}
require.NoError(t, db.Create(existingUser).Error)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
userInfo := &SAMLUserInfo{
NameID: "existing@idp.example.com",
Email: "existing@example.com",
DisplayName: "Existing Updated",
Attributes: map[string]string{"email": "existing@example.com"},
}
user, err := svc.FindOrCreateUser(userInfo)
require.NoError(t, err)
assert.NotNil(t, user)
assert.Equal(t, existingUser.ID, user.ID, "should find existing user by email and link to SAML")
assert.Equal(t, "saml", user.Provider, "should update provider to saml")
assert.Equal(t, "existing@idp.example.com", user.UID, "should set UID to SAML NameID")
}
func TestSAMLService_FindOrCreateUser_ExistingSAMLUser(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
// Create a SAML user first
existingUser := &model.User{
Name: "SAML Existing",
Email: "saml.existing@example.com",
Password: "hashed",
Provider: "saml",
UID: "saml123@idp.example.com",
Active: true,
}
require.NoError(t, db.Create(existingUser).Error)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
userInfo := &SAMLUserInfo{
NameID: "saml123@idp.example.com",
Email: "saml.existing@example.com",
DisplayName: "SAML Existing Updated",
Attributes: map[string]string{"email": "saml.existing@example.com"},
}
user, err := svc.FindOrCreateUser(userInfo)
require.NoError(t, err)
assert.NotNil(t, user)
assert.Equal(t, existingUser.ID, user.ID, "should find existing SAML user by provider+UID")
assert.Equal(t, "SAML Existing Updated", user.Name, "should update display name")
}
func TestSAMLService_FindOrCreateUser_NoDB(t *testing.T) {
_, rdb := setupTestRedis(t)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, nil)
require.NoError(t, err)
userInfo := &SAMLUserInfo{
NameID: "nodb@idp.example.com",
Email: "nodb@example.com",
DisplayName: "No DB User",
}
_, err = svc.FindOrCreateUser(userInfo)
assert.Error(t, err, "nil DB should error")
assert.Contains(t, err.Error(), "database not available")
}
// --- SAML GetSPMetadata tests ---
func TestSAMLService_GetSPMetadata(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
metadata, err := svc.GetSPMetadata()
require.NoError(t, err)
assert.NotEmpty(t, metadata, "SP metadata should not be empty")
}
func TestSAMLService_GetSPMetadata_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
_, err = svc.GetSPMetadata()
assert.ErrorIs(t, err, ErrSAMLEnabled, "disabled service should reject SP metadata request")
}
// --- SAML Logout tests (disabled) ---
func TestSAMLService_InitiateLogout_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
_, err = svc.InitiateLogout("session1", "nameid1", "index1", "state1")
assert.ErrorIs(t, err, ErrSAMLEnabled, "disabled service should reject logout initiation")
}
func TestSAMLService_ProcessLogoutResponse_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
err = svc.ProcessLogoutResponse("response", "relayState")
assert.ErrorIs(t, err, ErrSAMLEnabled, "disabled service should reject logout response")
}
func TestSAMLService_ProcessLogoutRequest_Disabled(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
cfg := newTestSAMLConfig(false)
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
_, err = svc.ProcessLogoutRequest("request")
assert.ErrorIs(t, err, ErrSAMLEnabled, "disabled service should reject logout request")
}
// --- SAML error sentinel tests ---
func TestSAMLErrorSentinels(t *testing.T) {
assert.Equal(t, "saml authentication is not enabled", ErrSAMLEnabled.Error())
assert.Equal(t, "saml configuration is invalid", ErrSAMLInvalidConfig.Error())
assert.Equal(t, "saml response validation failed", ErrSAMLInvalidResponse.Error())
assert.Equal(t, "saml response replay detected", ErrSAMLReplay.Error())
assert.Equal(t, "saml response missing NameID", ErrSAMLMissingNameID.Error())
assert.Equal(t, "saml response missing email attribute", ErrSAMLMissingEmail.Error())
assert.Equal(t, "failed to load IdP metadata", ErrSAMLIdPMetadata.Error())
}
// --- SAML attribute mapping tests ---
func TestSAMLAttributeMapping_DefaultConfig(t *testing.T) {
cfg := newTestSAMLConfig(true)
assert.Equal(t, "email", cfg.AttributeMap.Email)
assert.Equal(t, "displayName", cfg.AttributeMap.DisplayName)
assert.Equal(t, "firstName", cfg.AttributeMap.FirstName)
assert.Equal(t, "lastName", cfg.AttributeMap.LastName)
}
// --- validateConditions tests ---
func TestValidateConditions_ValidConditions(t *testing.T) {
now := time.Now().UTC()
drift := 180 * time.Second
conditions := &SAMLConditionsXML{
NotBefore: now.Add(-60 * time.Second).Format(time.RFC3339),
NotOnOrAfter: now.Add(5 * time.Minute).Format(time.RFC3339),
AudienceRestrictions: []SAMLAudienceRestrictionXML{
{Audiences: []SAMLAudienceXML{{Value: "https://sp.example.com/saml"}}},
},
}
err := validateConditions(conditions, "https://sp.example.com/saml", now, drift)
assert.NoError(t, err, "valid conditions should pass")
}
func TestValidateConditions_ExpiredConditions(t *testing.T) {
now := time.Now().UTC()
drift := 180 * time.Second
conditions := &SAMLConditionsXML{
NotBefore: now.Add(-10 * time.Minute).Format(time.RFC3339),
NotOnOrAfter: now.Add(-1 * time.Minute).Format(time.RFC3339), // already expired
AudienceRestrictions: []SAMLAudienceRestrictionXML{
{Audiences: []SAMLAudienceXML{{Value: "https://sp.example.com/saml"}}},
},
}
err := validateConditions(conditions, "https://sp.example.com/saml", now, drift)
assert.Error(t, err, "expired conditions should fail")
}
func TestValidateConditions_WrongAudience(t *testing.T) {
now := time.Now().UTC()
drift := 180 * time.Second
conditions := &SAMLConditionsXML{
NotBefore: now.Add(-60 * time.Second).Format(time.RFC3339),
NotOnOrAfter: now.Add(5 * time.Minute).Format(time.RFC3339),
AudienceRestrictions: []SAMLAudienceRestrictionXML{
{Audiences: []SAMLAudienceXML{{Value: "https://wrong-sp.example.com/saml"}}},
},
}
err := validateConditions(conditions, "https://sp.example.com/saml", now, drift)
assert.Error(t, err, "wrong audience should fail")
}
func TestValidateConditions_ClockDrift(t *testing.T) {
now := time.Now().UTC()
drift := 180 * time.Second
// NotBefore is slightly in the future (within drift tolerance)
conditions := &SAMLConditionsXML{
NotBefore: now.Add(60 * time.Second).Format(time.RFC3339), // 60s in future
NotOnOrAfter: now.Add(5 * time.Minute).Format(time.RFC3339),
AudienceRestrictions: []SAMLAudienceRestrictionXML{
{Audiences: []SAMLAudienceXML{{Value: "https://sp.example.com/saml"}}},
},
}
err := validateConditions(conditions, "https://sp.example.com/saml", now, drift)
assert.NoError(t, err, "conditions within drift tolerance should pass")
}
func TestValidateConditions_ExcessiveClockDrift(t *testing.T) {
now := time.Now().UTC()
drift := 30 * time.Second
// NotBefore is 60s in future, drift tolerance only 30s
conditions := &SAMLConditionsXML{
NotBefore: now.Add(60 * time.Second).Format(time.RFC3339),
NotOnOrAfter: now.Add(5 * time.Minute).Format(time.RFC3339),
AudienceRestrictions: []SAMLAudienceRestrictionXML{
{Audiences: []SAMLAudienceXML{{Value: "https://sp.example.com/saml"}}},
},
}
err := validateConditions(conditions, "https://sp.example.com/saml", now, drift)
assert.Error(t, err, "conditions exceeding drift tolerance should fail")
}
func TestValidateConditions_MultipleAudiences(t *testing.T) {
now := time.Now().UTC()
drift := 180 * time.Second
conditions := &SAMLConditionsXML{
NotBefore: now.Add(-60 * time.Second).Format(time.RFC3339),
NotOnOrAfter: now.Add(5 * time.Minute).Format(time.RFC3339),
AudienceRestrictions: []SAMLAudienceRestrictionXML{
{Audiences: []SAMLAudienceXML{
{Value: "https://other-sp.example.com/saml"},
{Value: "https://sp.example.com/saml"},
}},
},
}
err := validateConditions(conditions, "https://sp.example.com/saml", now, drift)
assert.NoError(t, err, "should pass when SP is among multiple audiences")
}
// --- SAML per-account settings tests ---
func TestSAMLService_WithPerAccountSettings(t *testing.T) {
_, rdb := setupTestRedis(t)
db := newTestSAMLDB(t)
// Create an account and SAML settings for it
account := &model.Account{Name: "Test Corp", Active: true}
require.NoError(t, db.Create(account).Error)
settings := &model.AccountSamlSettings{
AccountID: account.ID,
IdpEntityID: "https://corp-idp.example.com/saml",
IdpSsoTargetURL: "https://corp-idp.example.com/saml/sso",
IdpCertificate: "MIIDXTCCAkWgAwIBAgIJAJC1HiIAZAiIMA0GCSqGSIb3DQEBCwUA",
SpEntityID: "https://sp.example.com/saml",
Active: true,
}
require.NoError(t, db.Create(settings).Error)
cfg := newTestSAMLConfig(true)
cfg.IdPMetadataXML = testIdPMetadataXML
svc, err := NewSAMLService(cfg, rdb, db)
require.NoError(t, err)
assert.NotNil(t, svc)
}
// --- Role mapping via SAML settings (JSON RoleMappings) ---
func TestSAMLRoleMappings_JSONParsing(t *testing.T) {
mappings := json.RawMessage(`{"admins": "administrator", "developers": "agent", "managers": "supervisor"}`)
settings := &model.AccountSamlSettings{
RoleMappings: mappings,
}
assert.NotNil(t, settings.RoleMappings)
var parsed map[string]string
err := json.Unmarshal(settings.RoleMappings, &parsed)
assert.NoError(t, err)
assert.Equal(t, "administrator", parsed["admins"])
assert.Equal(t, "agent", parsed["developers"])
}