1379 lines
43 KiB
Go
1379 lines
43 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/alicebob/miniredis/v2"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/golang-jwt/jwt/v5"
|
|
"github.com/redis/go-redis/v9"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/driver/sqlite"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/logger"
|
|
|
|
"github.com/gochat/gochat/internal/config"
|
|
"github.com/gochat/gochat/internal/model"
|
|
)
|
|
|
|
// ========== Helpers ==========
|
|
|
|
func newSSOTestDB(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.AccountLDAPSettings{},
|
|
&model.AccountOIDCSettings{},
|
|
&model.User{},
|
|
&model.Account{},
|
|
&model.AccountUser{},
|
|
)
|
|
require.NoError(t, err)
|
|
return db
|
|
}
|
|
|
|
func newSSOTestRedis(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 newSSOTestConfig(samlEnabled, ldapEnabled, oidcEnabled bool) *config.Config {
|
|
return &config.Config{
|
|
JWT: config.JWTConfig{
|
|
Secret: "test-sso-jwt-secret",
|
|
ExpiryHours: 24,
|
|
Issuer: "gochat-test",
|
|
},
|
|
Session: config.SessionConfig{
|
|
ExpirySeconds: 86400,
|
|
},
|
|
SAML: config.SAMLConfig{
|
|
Enabled: samlEnabled,
|
|
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"},
|
|
},
|
|
LDAP: config.LDAPConfig{
|
|
Enabled: ldapEnabled,
|
|
DefaultHost: "localhost",
|
|
DefaultPort: 389,
|
|
DefaultUseTLS: false,
|
|
DefaultBaseDN: "dc=example,dc=com",
|
|
DefaultBindDN: "cn=admin,dc=example,dc=com",
|
|
DefaultBindPassword: "adminpassword",
|
|
DefaultUserFilter: "(uid=%s)",
|
|
DefaultEmailAttribute: "mail",
|
|
DefaultNameAttribute: "cn",
|
|
DefaultGroupAttribute: "memberOf",
|
|
SyncInterval: 300,
|
|
},
|
|
OIDC: config.OIDCConfig{
|
|
Enabled: oidcEnabled,
|
|
DefaultClientID: "test-client-id",
|
|
DefaultClientSecret: "test-client-secret",
|
|
DefaultRedirectURL: "https://app.example.com/callback",
|
|
DefaultIssuerURL: "https://issuer.example.com",
|
|
DefaultScopes: []string{"openid", "profile", "email"},
|
|
},
|
|
}
|
|
}
|
|
|
|
func newSSOMiddlewareForTest(t *testing.T, db *gorm.DB, rdb redis.Cmdable, cfg *config.Config) *SSOMiddleware {
|
|
t.Helper()
|
|
return NewSSOMiddleware(db, rdb, cfg, nil, nil, nil)
|
|
}
|
|
|
|
func seedAccountSAMLSettings(t *testing.T, db *gorm.DB, accountID uint, active bool) {
|
|
t.Helper()
|
|
settings := model.AccountSamlSettings{
|
|
AccountID: accountID,
|
|
IdpEntityID: "https://idp.example.com/saml",
|
|
IdpSsoTargetURL: "https://idp.example.com/saml/sso",
|
|
IdpCertificate: "MIIC...",
|
|
SpEntityID: "https://sp.example.com/saml",
|
|
// Use true on Create so GORM includes it (zero-value bools are skipped).
|
|
Active: true,
|
|
}
|
|
require.NoError(t, db.Create(&settings).Error)
|
|
// Now update to the desired value using map (bypasses GORM zero-value skipping).
|
|
if !active {
|
|
require.NoError(t, db.Model(&settings).Update("active", false).Error)
|
|
}
|
|
}
|
|
|
|
func seedAccountLDAPSettings(t *testing.T, db *gorm.DB, accountID uint, active bool, autoProvision bool) {
|
|
t.Helper()
|
|
settings := model.AccountLDAPSettings{
|
|
AccountID: accountID,
|
|
Host: "ldap.example.com",
|
|
Port: 389,
|
|
BaseDN: "dc=example,dc=com",
|
|
UserFilter: "(uid=%s)",
|
|
EmailAttribute: "mail",
|
|
NameAttribute: "cn",
|
|
GroupAttribute: "memberOf",
|
|
// Use true for both bools on Create so GORM includes them (zero-value bools are skipped).
|
|
AutoProvision: true,
|
|
Active: true,
|
|
}
|
|
require.NoError(t, db.Create(&settings).Error)
|
|
// Now update to the desired values using map (bypasses GORM zero-value skipping).
|
|
updates := map[string]interface{}{}
|
|
if !autoProvision {
|
|
updates["auto_provision"] = false
|
|
}
|
|
if !active {
|
|
updates["active"] = false
|
|
}
|
|
if len(updates) > 0 {
|
|
require.NoError(t, db.Model(&settings).Updates(updates).Error)
|
|
}
|
|
}
|
|
|
|
func seedAccountOIDCSettings(t *testing.T, db *gorm.DB, accountID uint, active bool, autoProvision bool) {
|
|
t.Helper()
|
|
settings := model.AccountOIDCSettings{
|
|
AccountID: accountID,
|
|
ClientID: "oidc-client-id",
|
|
ClientSecret: "oidc-client-secret",
|
|
RedirectURL: "https://app.example.com/oidc/callback",
|
|
IssuerURL: "https://oidc.example.com",
|
|
// Use true for both bools on Create so GORM includes them (zero-value bools are skipped).
|
|
AutoProvision: true,
|
|
Active: true,
|
|
}
|
|
require.NoError(t, db.Create(&settings).Error)
|
|
// Now update to the desired values using map (bypasses GORM zero-value skipping).
|
|
updates := map[string]interface{}{}
|
|
if !autoProvision {
|
|
updates["auto_provision"] = false
|
|
}
|
|
if !active {
|
|
updates["active"] = false
|
|
}
|
|
if len(updates) > 0 {
|
|
require.NoError(t, db.Model(&settings).Updates(updates).Error)
|
|
}
|
|
}
|
|
|
|
func seedAccountAndUser(t *testing.T, db *gorm.DB, accountName, userName, userEmail, provider string) (*model.Account, *model.User) {
|
|
t.Helper()
|
|
account := &model.Account{Name: accountName}
|
|
require.NoError(t, db.Create(account).Error)
|
|
|
|
user := &model.User{
|
|
AccountID: account.ID,
|
|
Name: userName,
|
|
Email: userEmail,
|
|
Provider: provider,
|
|
UID: userEmail,
|
|
}
|
|
require.NoError(t, db.Create(user).Error)
|
|
|
|
return account, user
|
|
}
|
|
|
|
// ========== Test 1: NewSSOMiddleware constructor ==========
|
|
|
|
func TestNewSSOMiddleware_Constructor(t *testing.T) {
|
|
mr, rdb := newSSOTestRedis(t)
|
|
db := newSSOTestDB(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
|
|
mw := NewSSOMiddleware(db, rdb, cfg, nil, nil, nil)
|
|
|
|
require.NotNil(t, mw)
|
|
assert.Equal(t, db, mw.db)
|
|
assert.Equal(t, rdb, mw.rdb)
|
|
assert.Equal(t, cfg, mw.cfg)
|
|
assert.Equal(t, "test-sso-jwt-secret", mw.jwtSecret)
|
|
// Session.ExpirySeconds=86400 => jwtExpiry = 86400s = 24h
|
|
assert.Equal(t, 24*time.Hour, mw.jwtExpiry)
|
|
// Services are nil since we passed nil
|
|
assert.Nil(t, mw.samlService)
|
|
assert.Nil(t, mw.ldapService)
|
|
assert.Nil(t, mw.oidcService)
|
|
|
|
// Verify miniredis is accessible
|
|
assert.Equal(t, mr.Addr(), rdb.(*redis.Client).Options().Addr)
|
|
}
|
|
|
|
func TestNewSSOMiddleware_DefaultJWTExpiry(t *testing.T) {
|
|
_, rdb := newSSOTestRedis(t)
|
|
db := newSSOTestDB(t)
|
|
// No Session.ExpirySeconds => defaults to 24h
|
|
cfg := &config.Config{
|
|
JWT: config.JWTConfig{Secret: "secret"},
|
|
Session: config.SessionConfig{ExpirySeconds: 0}, // zero means use default
|
|
}
|
|
|
|
mw := NewSSOMiddleware(db, rdb, cfg, nil, nil, nil)
|
|
assert.Equal(t, 24*time.Hour, mw.jwtExpiry)
|
|
}
|
|
|
|
func TestNewSSOMiddleware_CustomJWTExpiry(t *testing.T) {
|
|
_, rdb := newSSOTestRedis(t)
|
|
db := newSSOTestDB(t)
|
|
cfg := &config.Config{
|
|
JWT: config.JWTConfig{Secret: "secret"},
|
|
Session: config.SessionConfig{ExpirySeconds: 7200}, // 2 hours
|
|
}
|
|
|
|
mw := NewSSOMiddleware(db, rdb, cfg, nil, nil, nil)
|
|
assert.Equal(t, 7200*time.Second, mw.jwtExpiry)
|
|
}
|
|
|
|
// ========== Test 2: ResolveProvider with explicit hint ==========
|
|
|
|
func TestResolveProvider_ExplicitHintSAML(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Seed SAML settings for account 1
|
|
seedAccountSAMLSettings(t, db, 1, true)
|
|
|
|
provider, err := mw.ResolveProvider(context.Background(), 1, "saml")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderSAML, provider)
|
|
}
|
|
|
|
func TestResolveProvider_ExplicitHintLDAP(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Seed LDAP settings for account 1
|
|
seedAccountLDAPSettings(t, db, 1, true, true)
|
|
|
|
provider, err := mw.ResolveProvider(context.Background(), 1, "ldap")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderLDAP, provider)
|
|
}
|
|
|
|
func TestResolveProvider_ExplicitHintOIDC(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Seed OIDC settings for account 1
|
|
seedAccountOIDCSettings(t, db, 1, true, true)
|
|
|
|
provider, err := mw.ResolveProvider(context.Background(), 1, "oidc")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderOIDC, provider)
|
|
}
|
|
|
|
func TestResolveProvider_ExplicitHintCaseInsensitive(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
seedAccountLDAPSettings(t, db, 1, true, true)
|
|
|
|
provider, err := mw.ResolveProvider(context.Background(), 1, "LDAP")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderLDAP, provider)
|
|
}
|
|
|
|
func TestResolveProvider_ExplicitHintNotActive(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
// LDAP is globally enabled, but no per-account LDAP settings
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// No per-account LDAP settings seeded for account 99
|
|
// isProviderActive for LDAP: global enabled + no per-account => falls back to global enabled = true
|
|
provider, err := mw.ResolveProvider(context.Background(), 99, "ldap")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderLDAP, provider)
|
|
}
|
|
|
|
func TestResolveProvider_ExplicitHintDisabledGlobally(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
// All providers globally disabled
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
_, err := mw.ResolveProvider(context.Background(), 1, "saml")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "not active")
|
|
}
|
|
|
|
func TestResolveProvider_ExplicitHintUnknownProvider(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
_, err := mw.ResolveProvider(context.Background(), 1, "unknown_provider")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "unknown SSO provider")
|
|
}
|
|
|
|
// ========== Test 3: ResolveProvider with auto-detection ==========
|
|
|
|
func TestResolveProvider_AutoDetection_SAMLPriority(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Seed all three provider settings for account 1
|
|
seedAccountSAMLSettings(t, db, 1, true)
|
|
seedAccountOIDCSettings(t, db, 1, true, true)
|
|
seedAccountLDAPSettings(t, db, 1, true, true)
|
|
|
|
// SAML should win (priority: SAML > OIDC > LDAP)
|
|
provider, err := mw.ResolveProvider(context.Background(), 1, "")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderSAML, provider)
|
|
}
|
|
|
|
func TestResolveProvider_AutoDetection_OIDCWhenNoSAML(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
_ = newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// SAML globally enabled but no per-account SAML => fallback to global
|
|
// Since SAML is globally enabled but no per-account record, it falls back to cfg.SAML.Enabled=true
|
|
// Actually, isProviderActive checks: if cfg.SAML.Enabled, then look for per-account settings
|
|
// If no per-account record (ErrRecordNotFound), falls back to cfg.SAML.Enabled = true
|
|
// So SAML would still be detected as active via fallback
|
|
// Let's make SAML globally disabled to test OIDC priority
|
|
|
|
cfg2 := newSSOTestConfig(false, true, true)
|
|
mw2 := newSSOMiddlewareForTest(t, db, rdb, cfg2)
|
|
|
|
seedAccountOIDCSettings(t, db, 2, true, true)
|
|
seedAccountLDAPSettings(t, db, 2, true, true)
|
|
|
|
provider, err := mw2.ResolveProvider(context.Background(), 2, "")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderOIDC, provider)
|
|
}
|
|
|
|
func TestResolveProvider_AutoDetection_LDAPWhenNoOthers(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
seedAccountLDAPSettings(t, db, 3, true, true)
|
|
|
|
provider, err := mw.ResolveProvider(context.Background(), 3, "")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderLDAP, provider)
|
|
}
|
|
|
|
func TestResolveProvider_AutoDetection_GlobalFallback(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
// SAML globally enabled, no per-account settings => falls back to global config
|
|
cfg := newSSOTestConfig(true, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// No per-account settings seeded — should fall back to global SAML.Enabled=true
|
|
provider, err := mw.ResolveProvider(context.Background(), 99, "")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, SSOProviderSAML, provider)
|
|
}
|
|
|
|
// ========== Test 4: ResolveProvider when no provider is configured ==========
|
|
|
|
func TestResolveProvider_NoProviderConfigured(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
_, err := mw.ResolveProvider(context.Background(), 1, "")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "no SSO provider configured")
|
|
}
|
|
|
|
func TestResolveProvider_PerAccountSAMLInactive(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Seed inactive SAML settings — active=false
|
|
seedAccountSAMLSettings(t, db, 1, false)
|
|
|
|
// Per-account settings exist but are inactive, so isProviderActive won't find active=true record
|
|
// ErrRecordNotFound for "active=true" query, falls back to global SAML.Enabled=true
|
|
provider, err := mw.ResolveProvider(context.Background(), 1, "")
|
|
require.NoError(t, err)
|
|
// Falls back to global enabled
|
|
assert.Equal(t, SSOProviderSAML, provider)
|
|
}
|
|
|
|
// ========== Test 5: isProviderActive ==========
|
|
|
|
func TestIsProviderActive_SAML_EnabledWithPerAccount(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
seedAccountSAMLSettings(t, db, 1, true)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderSAML)
|
|
require.NoError(t, err)
|
|
assert.True(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_SAML_EnabledNoPerAccount(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// No per-account settings — falls back to global Enabled=true
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderSAML)
|
|
require.NoError(t, err)
|
|
assert.True(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_SAML_DisabledGlobally(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderSAML)
|
|
require.NoError(t, err)
|
|
assert.False(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_LDAP_EnabledWithPerAccount(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
seedAccountLDAPSettings(t, db, 1, true, true)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderLDAP)
|
|
require.NoError(t, err)
|
|
assert.True(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_LDAP_EnabledNoPerAccount(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Falls back to global Enabled=true
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderLDAP)
|
|
require.NoError(t, err)
|
|
assert.True(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_LDAP_DisabledGlobally(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderLDAP)
|
|
require.NoError(t, err)
|
|
assert.False(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_OIDC_EnabledWithPerAccount(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
seedAccountOIDCSettings(t, db, 1, true, true)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderOIDC)
|
|
require.NoError(t, err)
|
|
assert.True(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_OIDC_EnabledNoPerAccount(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Falls back to global Enabled=true
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderOIDC)
|
|
require.NoError(t, err)
|
|
assert.True(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_OIDC_DisabledGlobally(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderOIDC)
|
|
require.NoError(t, err)
|
|
assert.False(t, active)
|
|
}
|
|
|
|
func TestIsProviderActive_UnknownProvider(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
active, err := mw.isProviderActive(context.Background(), 1, SSOProviderType("unknown"))
|
|
require.Error(t, err)
|
|
assert.False(t, active)
|
|
assert.Contains(t, err.Error(), "unknown SSO provider")
|
|
}
|
|
|
|
// ========== Test 6: SSOProviderType constants ==========
|
|
|
|
func TestSSOProviderType_Constants(t *testing.T) {
|
|
assert.Equal(t, SSOProviderType("saml"), SSOProviderSAML)
|
|
assert.Equal(t, SSOProviderType("ldap"), SSOProviderLDAP)
|
|
assert.Equal(t, SSOProviderType("oidc"), SSOProviderOIDC)
|
|
|
|
// Verify string representation
|
|
assert.Equal(t, "saml", string(SSOProviderSAML))
|
|
assert.Equal(t, "ldap", string(SSOProviderLDAP))
|
|
assert.Equal(t, "oidc", string(SSOProviderOIDC))
|
|
}
|
|
|
|
// ========== Test 7: SSOAuthResult struct field validation ==========
|
|
|
|
func TestSSOAuthResult_Fields(t *testing.T) {
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderSAML,
|
|
UserID: 42,
|
|
AccountID: 1,
|
|
Email: "user@example.com",
|
|
Name: "Test User",
|
|
FirstName: "Test",
|
|
LastName: "User",
|
|
Subject: "nameid-123",
|
|
Role: "administrator",
|
|
Groups: []string{"admins", "developers"},
|
|
AutoProvision: false,
|
|
}
|
|
|
|
assert.Equal(t, SSOProviderSAML, result.Provider)
|
|
assert.Equal(t, uint(42), result.UserID)
|
|
assert.Equal(t, uint(1), result.AccountID)
|
|
assert.Equal(t, "user@example.com", result.Email)
|
|
assert.Equal(t, "Test User", result.Name)
|
|
assert.Equal(t, "Test", result.FirstName)
|
|
assert.Equal(t, "User", result.LastName)
|
|
assert.Equal(t, "nameid-123", result.Subject)
|
|
assert.Equal(t, "administrator", result.Role)
|
|
assert.Equal(t, []string{"admins", "developers"}, result.Groups)
|
|
assert.False(t, result.AutoProvision)
|
|
}
|
|
|
|
func TestSSOAuthResult_DefaultValues(t *testing.T) {
|
|
result := &SSOAuthResult{}
|
|
|
|
assert.Equal(t, SSOProviderType(""), result.Provider)
|
|
assert.Equal(t, uint(0), result.UserID)
|
|
assert.Equal(t, uint(0), result.AccountID)
|
|
assert.Equal(t, "", result.Email)
|
|
assert.Equal(t, "", result.Name)
|
|
assert.Equal(t, "", result.FirstName)
|
|
assert.Equal(t, "", result.LastName)
|
|
assert.Equal(t, "", result.Subject)
|
|
assert.Equal(t, "", result.Role)
|
|
assert.Nil(t, result.Groups)
|
|
assert.False(t, result.AutoProvision)
|
|
}
|
|
|
|
func TestSSOAuthResult_ProviderTypeVariants(t *testing.T) {
|
|
samlResult := &SSOAuthResult{Provider: SSOProviderSAML}
|
|
assert.Equal(t, "saml", string(samlResult.Provider))
|
|
|
|
ldapResult := &SSOAuthResult{Provider: SSOProviderLDAP}
|
|
assert.Equal(t, "ldap", string(ldapResult.Provider))
|
|
|
|
oidcResult := &SSOAuthResult{Provider: SSOProviderOIDC}
|
|
assert.Equal(t, "oidc", string(oidcResult.Provider))
|
|
}
|
|
|
|
// ========== Test 8: IssueJWT with valid SSOAuthResult ==========
|
|
|
|
func TestIssueJWT_ValidResult(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
mr, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderSAML,
|
|
UserID: 42,
|
|
AccountID: 1,
|
|
Email: "user@example.com",
|
|
Name: "Test User",
|
|
Subject: "nameid-123",
|
|
Role: "administrator",
|
|
}
|
|
|
|
tokenStr, err := mw.IssueJWT(result)
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, tokenStr)
|
|
|
|
// Parse and verify the JWT
|
|
token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (interface{}, error) {
|
|
// Verify signing method
|
|
assert.Equal(t, jwt.SigningMethodHS256, token.Method)
|
|
return []byte("test-sso-jwt-secret"), nil
|
|
})
|
|
require.NoError(t, err)
|
|
require.True(t, token.Valid)
|
|
|
|
claims, ok := token.Claims.(jwt.MapClaims)
|
|
require.True(t, ok)
|
|
|
|
// Verify claim values
|
|
assert.Equal(t, float64(42), claims["user_id"])
|
|
assert.Equal(t, float64(1), claims["account_id"])
|
|
assert.Equal(t, "administrator", claims["role"])
|
|
assert.Equal(t, "saml", claims["provider"])
|
|
assert.Equal(t, "nameid-123", claims["subject"])
|
|
assert.Equal(t, "user@example.com", claims["email"])
|
|
|
|
// Verify expiry claim exists
|
|
assert.NotNil(t, claims["exp"])
|
|
assert.NotNil(t, claims["iat"])
|
|
|
|
// Verify miniredis didn't interfere
|
|
assert.NotEmpty(t, mr.Addr())
|
|
}
|
|
|
|
func TestIssueJWT_OIDCResult(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderOIDC,
|
|
UserID: 100,
|
|
AccountID: 5,
|
|
Email: "oidcuser@example.com",
|
|
Name: "OIDC User",
|
|
Subject: "sub-oidc-456",
|
|
Role: "agent",
|
|
}
|
|
|
|
tokenStr, err := mw.IssueJWT(result)
|
|
require.NoError(t, err)
|
|
|
|
token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (interface{}, error) {
|
|
return []byte("test-sso-jwt-secret"), nil
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
claims := token.Claims.(jwt.MapClaims)
|
|
assert.Equal(t, "oidc", claims["provider"])
|
|
assert.Equal(t, float64(100), claims["user_id"])
|
|
assert.Equal(t, float64(5), claims["account_id"])
|
|
}
|
|
|
|
// ========== Test 9: IssueJWT with missing fields ==========
|
|
|
|
func TestIssueJWT_ZeroUserID(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderSAML,
|
|
UserID: 0, // zero — unprovisioned
|
|
AccountID: 1,
|
|
Email: "newuser@example.com",
|
|
}
|
|
|
|
tokenStr, err := mw.IssueJWT(result)
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "unprovisioned user")
|
|
assert.Empty(t, tokenStr)
|
|
}
|
|
|
|
func TestIssueJWT_EmptyProvider(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Even with empty provider, IssueJWT should succeed if UserID > 0
|
|
result := &SSOAuthResult{
|
|
Provider: "",
|
|
UserID: 10,
|
|
AccountID: 1,
|
|
Email: "user@example.com",
|
|
Subject: "sub-1",
|
|
Role: "agent",
|
|
}
|
|
|
|
tokenStr, err := mw.IssueJWT(result)
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, tokenStr)
|
|
|
|
token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (interface{}, error) {
|
|
return []byte("test-sso-jwt-secret"), nil
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
claims := token.Claims.(jwt.MapClaims)
|
|
assert.Equal(t, "", claims["provider"])
|
|
}
|
|
|
|
func TestIssueJWT_EmptyEmail(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// IssueJWT doesn't validate email; it just puts whatever fields are there in claims
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderLDAP,
|
|
UserID: 10,
|
|
AccountID: 1,
|
|
Email: "",
|
|
Subject: "cn=user,dc=example",
|
|
Role: "agent",
|
|
}
|
|
|
|
tokenStr, err := mw.IssueJWT(result)
|
|
require.NoError(t, err)
|
|
|
|
token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (interface{}, error) {
|
|
return []byte("test-sso-jwt-secret"), nil
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
claims := token.Claims.(jwt.MapClaims)
|
|
assert.Equal(t, "", claims["email"])
|
|
}
|
|
|
|
// ========== Test 10: SSOSessionValidator middleware ==========
|
|
|
|
func TestSSOSessionValidator_ValidSession(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Create an SSO session in Redis manually
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderSAML,
|
|
UserID: 42,
|
|
AccountID: 1,
|
|
Email: "user@example.com",
|
|
Subject: "nameid-123",
|
|
Role: "administrator",
|
|
}
|
|
sessionID, err := mw.CreateSSOSession(context.Background(), result)
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, sessionID)
|
|
|
|
// Set up gin context with the session
|
|
gin.SetMode(gin.TestMode)
|
|
w := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(w)
|
|
c.Request = httptest.NewRequest(http.MethodGet, "/", nil)
|
|
c.Request.Header.Set("X-SSO-Session", sessionID)
|
|
|
|
// Execute the middleware
|
|
handler := mw.SSOSessionValidator()
|
|
handler(c)
|
|
|
|
// Should NOT abort — session is valid
|
|
assert.False(t, c.IsAborted())
|
|
assert.Equal(t, http.StatusOK, w.Code) // default status, not set by middleware
|
|
|
|
// Verify context values were set
|
|
ssoSession, exists := c.Get("sso_session")
|
|
assert.True(t, exists)
|
|
sessionData := ssoSession.(SSOSessionData)
|
|
assert.Equal(t, sessionID, sessionData.SessionID)
|
|
assert.Equal(t, uint(42), sessionData.UserID)
|
|
assert.Equal(t, uint(1), sessionData.AccountID)
|
|
assert.Equal(t, "saml", sessionData.Provider)
|
|
assert.Equal(t, "administrator", sessionData.Role)
|
|
|
|
ssoProvider, exists := c.Get("sso_provider")
|
|
assert.True(t, exists)
|
|
assert.Equal(t, "saml", ssoProvider)
|
|
|
|
userID, exists := c.Get("user_id")
|
|
assert.True(t, exists)
|
|
assert.Equal(t, uint(42), userID)
|
|
|
|
accountID, exists := c.Get("account_id")
|
|
assert.True(t, exists)
|
|
assert.Equal(t, uint(1), accountID)
|
|
|
|
role, exists := c.Get("role")
|
|
assert.True(t, exists)
|
|
assert.Equal(t, "administrator", role)
|
|
}
|
|
|
|
func TestSSOSessionValidator_QueryParam(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderOIDC,
|
|
UserID: 10,
|
|
AccountID: 2,
|
|
Email: "oidc@example.com",
|
|
Subject: "sub-oidc",
|
|
Role: "agent",
|
|
}
|
|
sessionID, err := mw.CreateSSOSession(context.Background(), result)
|
|
require.NoError(t, err)
|
|
|
|
gin.SetMode(gin.TestMode)
|
|
w := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(w)
|
|
c.Request = httptest.NewRequest(http.MethodGet, "/?sso_session="+sessionID, nil)
|
|
// No X-SSO-Session header — session from query param
|
|
|
|
handler := mw.SSOSessionValidator()
|
|
handler(c)
|
|
|
|
assert.False(t, c.IsAborted())
|
|
|
|
ssoProvider, _ := c.Get("sso_provider")
|
|
assert.Equal(t, "oidc", ssoProvider)
|
|
}
|
|
|
|
func TestSSOSessionValidator_NoSessionProvided(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
gin.SetMode(gin.TestMode)
|
|
w := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(w)
|
|
c.Request = httptest.NewRequest(http.MethodGet, "/", nil)
|
|
// No session header or query param
|
|
|
|
handler := mw.SSOSessionValidator()
|
|
handler(c)
|
|
|
|
// Should NOT abort — just skips SSO validation (c.Next())
|
|
assert.False(t, c.IsAborted())
|
|
}
|
|
|
|
func TestSSOSessionValidator_ExpiredSession(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
mr, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderSAML,
|
|
UserID: 42,
|
|
AccountID: 1,
|
|
Email: "user@example.com",
|
|
Subject: "nameid-123",
|
|
Role: "administrator",
|
|
}
|
|
sessionID, err := mw.CreateSSOSession(context.Background(), result)
|
|
require.NoError(t, err)
|
|
|
|
// Expire the session in miniredis
|
|
mr.Del("sso:session:" + sessionID)
|
|
|
|
gin.SetMode(gin.TestMode)
|
|
w := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(w)
|
|
c.Request = httptest.NewRequest(http.MethodGet, "/", nil)
|
|
c.Request.Header.Set("X-SSO-Session", sessionID)
|
|
|
|
handler := mw.SSOSessionValidator()
|
|
handler(c)
|
|
|
|
// Should abort with 401
|
|
assert.True(t, c.IsAborted())
|
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
|
}
|
|
|
|
func TestSSOSessionValidator_CorruptSessionData(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
mr, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Put corrupt JSON in Redis
|
|
corruptSessionID := "corrupt-session-123"
|
|
mr.Set("sso:session:" + corruptSessionID, "not-valid-json{broken")
|
|
|
|
gin.SetMode(gin.TestMode)
|
|
w := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(w)
|
|
c.Request = httptest.NewRequest(http.MethodGet, "/", nil)
|
|
c.Request.Header.Set("X-SSO-Session", corruptSessionID)
|
|
|
|
handler := mw.SSOSessionValidator()
|
|
handler(c)
|
|
|
|
// Should abort with 401 — corrupt data
|
|
assert.True(t, c.IsAborted())
|
|
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
|
}
|
|
|
|
// ========== Test 11: findOrCreateUser logic ==========
|
|
|
|
func TestFindOrCreateUser_FindExisting(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Pre-create a user
|
|
account, user := seedAccountAndUser(t, db, "Test Account", "Existing User", "existing@example.com", "saml")
|
|
|
|
// findOrCreateUser should find the existing user
|
|
foundUser, err := mw.findOrCreateUser(context.Background(), account.ID, "existing@example.com", "Existing User", "nameid-1", "saml", "agent")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, user.ID, foundUser.ID)
|
|
assert.Equal(t, "existing@example.com", foundUser.Email)
|
|
assert.Equal(t, "Existing User", foundUser.Name)
|
|
}
|
|
|
|
func TestFindOrCreateUser_CreateNewUser(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Create an account
|
|
account := &model.Account{Name: "New Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
|
|
// findOrCreateUser should create a new user
|
|
newUser, err := mw.findOrCreateUser(context.Background(), account.ID, "newuser@example.com", "New User", "uid-new", "ldap", "agent")
|
|
require.NoError(t, err)
|
|
assert.NotEqual(t, uint(0), newUser.ID)
|
|
assert.Equal(t, "newuser@example.com", newUser.Email)
|
|
assert.Equal(t, "New User", newUser.Name)
|
|
assert.Equal(t, "ldap", newUser.Provider)
|
|
assert.Equal(t, "uid-new", newUser.UID)
|
|
|
|
// Verify account membership was created
|
|
var au model.AccountUser
|
|
err = db.Where("account_id = ? AND user_id = ?", account.ID, newUser.ID).First(&au).Error
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "agent", au.Role)
|
|
}
|
|
|
|
func TestFindOrCreateUser_LDAPAutoProvisionDisabled(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Create account with LDAP settings where AutoProvision=false
|
|
account := &model.Account{Name: "LDAP No Provision"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
seedAccountLDAPSettings(t, db, account.ID, true, false) // autoProvision=false
|
|
|
|
_, err := mw.findOrCreateUser(context.Background(), account.ID, "newldap@example.com", "LDAP User", "cn=user", "ldap", "agent")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "auto-provision is disabled")
|
|
}
|
|
|
|
func TestFindOrCreateUser_OIDCAutoProvisionDisabled(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
// OIDC service is nil — getAccountSettings will be called on nil oidcService
|
|
// findOrCreateUser for oidc: oidcService.getAccountSettings(accountID) — oidcService is nil
|
|
// This would panic. So for OIDC auto-provision tests, we can only test the fallback (nil settings).
|
|
// When oidcSettings is nil, autoProvision defaults to true.
|
|
// We'll test that default behavior here.
|
|
|
|
cfg := newSSOTestConfig(false, false, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg) // oidcService is nil
|
|
|
|
account := &model.Account{Name: "OIDC Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
|
|
// With nil OIDCService, oidcSettings will be nil, autoProvision defaults to true
|
|
newUser, err := mw.findOrCreateUser(context.Background(), account.ID, "oidcuser@example.com", "OIDC User", "sub-oidc", "oidc", "agent")
|
|
require.NoError(t, err)
|
|
assert.NotEqual(t, uint(0), newUser.ID)
|
|
assert.Equal(t, "oidcuser@example.com", newUser.Email)
|
|
}
|
|
|
|
func TestFindOrCreateUser_DBNil(t *testing.T) {
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
// SSOMiddleware with nil DB
|
|
mw := NewSSOMiddleware(nil, rdb, cfg, nil, nil, nil)
|
|
|
|
_, err := mw.findOrCreateUser(context.Background(), 1, "test@example.com", "Test", "uid-1", "saml", "agent")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "database not available")
|
|
}
|
|
|
|
func TestFindOrCreateUser_SAMLProvider(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, false, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
account := &model.Account{Name: "SAML Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
|
|
// SAML always auto-provisions (no AutoProvision field on AccountSamlSettings)
|
|
newUser, err := mw.findOrCreateUser(context.Background(), account.ID, "samluser@example.com", "SAML User", "nameid-saml", "saml", "administrator")
|
|
require.NoError(t, err)
|
|
assert.NotEqual(t, uint(0), newUser.ID)
|
|
assert.Equal(t, "samluser@example.com", newUser.Email)
|
|
assert.Equal(t, "saml", newUser.Provider)
|
|
|
|
// Verify account membership with administrator role
|
|
var au model.AccountUser
|
|
err = db.Where("account_id = ? AND user_id = ?", account.ID, newUser.ID).First(&au).Error
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "administrator", au.Role)
|
|
}
|
|
|
|
func TestFindOrCreateUser_LDAPMapsProviderInfo(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
account := &model.Account{Name: "LDAP Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
seedAccountLDAPSettings(t, db, account.ID, true, true)
|
|
|
|
newUser, err := mw.findOrCreateUser(context.Background(), account.ID, "ldapuser@example.com", "LDAP User", "cn=ldapuser,dc=example,dc=com", "ldap", "agent")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "ldapuser@example.com", newUser.Email)
|
|
assert.Equal(t, "LDAP User", newUser.Name)
|
|
assert.Equal(t, "ldap", newUser.Provider)
|
|
assert.Equal(t, "cn=ldapuser,dc=example,dc=com", newUser.UID)
|
|
}
|
|
|
|
// ========== Additional tests: CreateSSOSession ==========
|
|
|
|
func TestCreateSSOSession_StoresInRedis(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
mr, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderSAML,
|
|
UserID: 42,
|
|
AccountID: 1,
|
|
Email: "user@example.com",
|
|
Subject: "nameid-123",
|
|
Role: "administrator",
|
|
}
|
|
|
|
sessionID, err := mw.CreateSSOSession(context.Background(), result)
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, sessionID)
|
|
|
|
// Verify session data in Redis
|
|
key := fmt.Sprintf("sso:session:%s", sessionID)
|
|
assert.True(t, mr.Exists(key))
|
|
|
|
// Verify user sessions set
|
|
userKey := fmt.Sprintf("sso:user_sessions:%d", result.UserID)
|
|
assert.True(t, mr.Exists(userKey))
|
|
}
|
|
|
|
func TestCreateSSOSession_SessionDataContent(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderLDAP,
|
|
UserID: 99,
|
|
AccountID: 5,
|
|
Email: "ldap@example.com",
|
|
Subject: "cn=user,dc=example",
|
|
Role: "agent",
|
|
}
|
|
|
|
sessionID, err := mw.CreateSSOSession(context.Background(), result)
|
|
require.NoError(t, err)
|
|
|
|
// Retrieve and parse the session data from Redis
|
|
key := fmt.Sprintf("sso:session:%s", sessionID)
|
|
val, err := rdb.Get(context.Background(), key).Bytes()
|
|
require.NoError(t, err)
|
|
|
|
var sessionData SSOSessionData
|
|
require.NoError(t, json.Unmarshal(val, &sessionData))
|
|
|
|
assert.Equal(t, sessionID, sessionData.SessionID)
|
|
assert.Equal(t, uint(99), sessionData.UserID)
|
|
assert.Equal(t, uint(5), sessionData.AccountID)
|
|
assert.Equal(t, "ldap", sessionData.Provider)
|
|
assert.Equal(t, "cn=user,dc=example", sessionData.NameID)
|
|
assert.Equal(t, "agent", sessionData.Role)
|
|
assert.NotZero(t, sessionData.CreatedAt)
|
|
assert.NotZero(t, sessionData.ExpiresAt)
|
|
}
|
|
|
|
// ========== Additional tests: AuthenticateLDAP with nil service ==========
|
|
|
|
func TestAuthenticateLDAP_NilService(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg) // ldapService is nil
|
|
|
|
_, err := mw.AuthenticateLDAP(context.Background(), 1, "user", "password")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "LDAP service not initialized")
|
|
}
|
|
|
|
// ========== Additional tests: AuthenticateOIDC with nil service ==========
|
|
|
|
func TestAuthenticateOIDC_NilService(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, false, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg) // oidcService is nil
|
|
|
|
_, err := mw.AuthenticateOIDC(context.Background(), "state", "code")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "OIDC service not initialized")
|
|
}
|
|
|
|
// ========== Additional tests: ensureAccountMembership ==========
|
|
|
|
func TestEnsureAccountMembership_NewMembership(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
account := &model.Account{Name: "Membership Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
|
|
user := &model.User{
|
|
AccountID: account.ID,
|
|
Name: "Member User",
|
|
Email: "member@example.com",
|
|
Provider: "email",
|
|
}
|
|
require.NoError(t, db.Create(user).Error)
|
|
|
|
mw.ensureAccountMembership(context.Background(), user.ID, account.ID, "agent")
|
|
|
|
var au model.AccountUser
|
|
err := db.Where("account_id = ? AND user_id = ?", account.ID, user.ID).First(&au).Error
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "agent", au.Role)
|
|
}
|
|
|
|
func TestEnsureAccountMembership_UpdateRole(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(true, true, true)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
account := &model.Account{Name: "Role Update Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
|
|
user := &model.User{
|
|
AccountID: account.ID,
|
|
Name: "Role User",
|
|
Email: "role@example.com",
|
|
Provider: "email",
|
|
}
|
|
require.NoError(t, db.Create(user).Error)
|
|
|
|
// First add with agent role
|
|
mw.ensureAccountMembership(context.Background(), user.ID, account.ID, "agent")
|
|
|
|
// Update to administrator role
|
|
mw.ensureAccountMembership(context.Background(), user.ID, account.ID, "administrator")
|
|
|
|
var au model.AccountUser
|
|
err := db.Where("account_id = ? AND user_id = ?", account.ID, user.ID).First(&au).Error
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "administrator", au.Role)
|
|
}
|
|
|
|
// ========== Additional: SSOSessionData struct ==========
|
|
|
|
func TestSSOSessionData_JSONRoundTrip(t *testing.T) {
|
|
session := SSOSessionData{
|
|
SessionID: "session-abc-123",
|
|
UserID: 42,
|
|
Provider: "saml",
|
|
IdPEntityID: "https://idp.example.com",
|
|
NameID: "nameid-456",
|
|
AccountID: 1,
|
|
Role: "administrator",
|
|
CreatedAt: time.Now().Unix(),
|
|
ExpiresAt: time.Now().Add(24 * time.Hour).Unix(),
|
|
}
|
|
|
|
data, err := json.Marshal(session)
|
|
require.NoError(t, err)
|
|
|
|
var decoded SSOSessionData
|
|
require.NoError(t, json.Unmarshal(data, &decoded))
|
|
|
|
assert.Equal(t, session.SessionID, decoded.SessionID)
|
|
assert.Equal(t, session.UserID, decoded.UserID)
|
|
assert.Equal(t, session.Provider, decoded.Provider)
|
|
assert.Equal(t, session.IdPEntityID, decoded.IdPEntityID)
|
|
assert.Equal(t, session.NameID, decoded.NameID)
|
|
assert.Equal(t, session.AccountID, decoded.AccountID)
|
|
assert.Equal(t, session.Role, decoded.Role)
|
|
assert.Equal(t, session.CreatedAt, decoded.CreatedAt)
|
|
assert.Equal(t, session.ExpiresAt, decoded.ExpiresAt)
|
|
}
|
|
|
|
// ========== Additional: getLDAPSettings ==========
|
|
|
|
func TestGetLDAPSettings_Active(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
seedAccountLDAPSettings(t, db, 1, true, true)
|
|
|
|
settings := mw.getLDAPSettings(1)
|
|
require.NotNil(t, settings)
|
|
assert.Equal(t, "ldap.example.com", settings.Host)
|
|
assert.Equal(t, uint(1), settings.AccountID)
|
|
assert.True(t, settings.Active)
|
|
}
|
|
|
|
func TestGetLDAPSettings_Inactive(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
// Seed inactive settings — getLDAPSettings only looks for active=true
|
|
seedAccountLDAPSettings(t, db, 1, false, true)
|
|
|
|
settings := mw.getLDAPSettings(1)
|
|
assert.Nil(t, settings) // inactive settings not found
|
|
}
|
|
|
|
func TestGetLDAPSettings_NoSettings(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
settings := mw.getLDAPSettings(99)
|
|
assert.Nil(t, settings) // no settings for account 99
|
|
}
|
|
|
|
func TestGetLDAPSettings_NilDB(t *testing.T) {
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := NewSSOMiddleware(nil, rdb, cfg, nil, nil, nil)
|
|
|
|
settings := mw.getLDAPSettings(1)
|
|
assert.Nil(t, settings)
|
|
}
|
|
|
|
// ========== Comprehensive integration: full SSO flow with LDAP ==========
|
|
|
|
func TestSSOFlow_LDAPFindOrCreateThenSession(t *testing.T) {
|
|
db := newSSOTestDB(t)
|
|
_, rdb := newSSOTestRedis(t)
|
|
cfg := newSSOTestConfig(false, true, false)
|
|
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
|
|
|
|
account := &model.Account{Name: "Flow Account"}
|
|
require.NoError(t, db.Create(account).Error)
|
|
seedAccountLDAPSettings(t, db, account.ID, true, true)
|
|
|
|
// Step 1: Find/create user
|
|
user, err := mw.findOrCreateUser(context.Background(), account.ID, "flowuser@example.com", "Flow User", "cn=flowuser", "ldap", "agent")
|
|
require.NoError(t, err)
|
|
assert.NotEqual(t, uint(0), user.ID)
|
|
|
|
// Step 2: Issue JWT
|
|
result := &SSOAuthResult{
|
|
Provider: SSOProviderLDAP,
|
|
UserID: user.ID,
|
|
AccountID: account.ID,
|
|
Email: "flowuser@example.com",
|
|
Name: "Flow User",
|
|
Subject: "cn=flowuser",
|
|
Role: "agent",
|
|
}
|
|
|
|
tokenStr, err := mw.IssueJWT(result)
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, tokenStr)
|
|
|
|
// Step 3: Parse and verify JWT
|
|
token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (interface{}, error) {
|
|
return []byte("test-sso-jwt-secret"), nil
|
|
})
|
|
require.NoError(t, err)
|
|
claims := token.Claims.(jwt.MapClaims)
|
|
assert.Equal(t, float64(user.ID), claims["user_id"])
|
|
assert.Equal(t, "ldap", claims["provider"])
|
|
|
|
// Step 4: Create SSO session
|
|
sessionID, err := mw.CreateSSOSession(context.Background(), result)
|
|
require.NoError(t, err)
|
|
require.NotEmpty(t, sessionID)
|
|
|
|
// Step 5: Validate session via middleware
|
|
gin.SetMode(gin.TestMode)
|
|
w := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(w)
|
|
c.Request = httptest.NewRequest(http.MethodGet, "/", nil)
|
|
c.Request.Header.Set("X-SSO-Session", sessionID)
|
|
|
|
handler := mw.SSOSessionValidator()
|
|
handler(c)
|
|
assert.False(t, c.IsAborted())
|
|
} |