492 lines
14 KiB
Plaintext
492 lines
14 KiB
Plaintext
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 `<?xml version="1.0"?>
|
|
<EntityDescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" entityID="https://idp.test/saml">
|
|
<IDPSSODescriptor WantAuthnRequestsSigned="false" protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
|
|
<NameIDFormat>urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress</NameIDFormat>
|
|
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" Location="https://idp.test/sso"/>
|
|
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="https://idp.test/sso"/>
|
|
</IDPSSODescriptor>
|
|
</EntityDescriptor>`
|
|
}
|
|
|
|
// --- 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("<samlp:AuthnRequest/>")
|
|
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)
|
|
} |