762 lines
25 KiB
Go
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"])
|
|
} |