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

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)
}