Files
gochat/backend/internal/auth/sso_middleware_test.go
T
Rogee 851ca7e372 refactor: 移除 SAML/LDAP/MFA 登录方式,仅保留本地账号密码和 OIDC
后端移除:
- SAML: auth/saml.go, handler/saml_handler.go, account_saml_settings_handler.go,
  model/account_saml_settings.go, model/saml_idp_config.go, repo/*.go
- LDAP: auth/ldap.go, handler/ldap_handler.go, model/account_ldap_settings.go,
  repo/account_ldap_settings_repo.go
- MFA: auth/mfa.go, handler/mfa_handler.go
- auth_service: 移除 mfaService 依赖、MFARequired 字段、LoginWithMFA 方法
- auth_handler: 移除 LoginMFA handler、MFA 分支逻辑
- bootstrap: 移除 SAML/LDAP/MFA service 初始化和 handler 注册
- sso_middleware: 精简为仅支持 OIDC provider
- router: 移除 SAML/LDAP/MFA 路由注册
- config: 移除 SAMLConfig/LDAPConfig struct 和 defaults

前端移除:
- v3/login: 移除 MFA 验证流程和 SAML 登录入口
- v3/api/auth: 移除 MFA 响应处理
- v3/routes: 移除 SSO login 路由
- dashboard: 移除 MFA 设置页面、SAML 安全设置页面
- i18n: 移除 mfa.json
- featureFlags: 移除 SAML feature flag

.env.example / .env: 移除 SAML/LDAP 配置段
2026-07-29 19:03:04 +08:00

1033 lines
31 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.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(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,
},
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)
}
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)
mw := NewSSOMiddleware(db, rdb, cfg, 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.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)
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)
assert.Equal(t, 7200*time.Second, mw.jwtExpiry)
}
// ========== Test 2: ResolveProvider with explicit hint ==========
func TestResolveProvider_ExplicitHintOIDC(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(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(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
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_ExplicitHintNotActive(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
// OIDC is globally enabled, but no per-account OIDC settings
cfg := newSSOTestConfig(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
// No per-account OIDC settings seeded for account 99
// isProviderActive for OIDC: global enabled + no per-account => falls back to global enabled = true
provider, err := mw.ResolveProvider(context.Background(), 99, "oidc")
require.NoError(t, err)
assert.Equal(t, SSOProviderOIDC, provider)
}
func TestResolveProvider_ExplicitHintDisabledGlobally(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
// OIDC globally disabled
cfg := newSSOTestConfig(false)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
_, err := mw.ResolveProvider(context.Background(), 1, "oidc")
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)
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_OIDC(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
seedAccountOIDCSettings(t, db, 2, true, true)
provider, err := mw.ResolveProvider(context.Background(), 2, "")
require.NoError(t, err)
assert.Equal(t, SSOProviderOIDC, provider)
}
func TestResolveProvider_AutoDetection_GlobalFallback(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
// OIDC globally enabled, no per-account settings => falls back to global config
cfg := newSSOTestConfig(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
// No per-account settings seeded — should fall back to global OIDC.Enabled=true
provider, err := mw.ResolveProvider(context.Background(), 99, "")
require.NoError(t, err)
assert.Equal(t, SSOProviderOIDC, provider)
}
// ========== Test 4: ResolveProvider when no provider is configured ==========
func TestResolveProvider_NoProviderConfigured(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(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_PerAccountOIDCInactive(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
// Seed inactive OIDC settings — active=false
seedAccountOIDCSettings(t, db, 1, false, true)
// Per-account settings exist but are inactive, so isProviderActive won't find active=true record
// ErrRecordNotFound for "active=true" query, falls back to global OIDC.Enabled=true
provider, err := mw.ResolveProvider(context.Background(), 1, "")
require.NoError(t, err)
// Falls back to global enabled
assert.Equal(t, SSOProviderOIDC, provider)
}
// ========== Test 5: isProviderActive ==========
func TestIsProviderActive_OIDC_EnabledWithPerAccount(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(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(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)
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)
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("oidc"), SSOProviderOIDC)
// Verify string representation
assert.Equal(t, "oidc", string(SSOProviderOIDC))
}
// ========== Test 7: SSOAuthResult struct field validation ==========
func TestSSOAuthResult_Fields(t *testing.T) {
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
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, SSOProviderOIDC, 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) {
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
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, "oidc", 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)
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
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)
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
// IssueJWT doesn't validate email; it just puts whatever fields are there in claims
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
UserID: 10,
AccountID: 1,
Email: "",
Subject: "sub-oidc",
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
// Create an SSO session in Redis manually
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
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, "oidc", sessionData.Provider)
assert.Equal(t, "administrator", sessionData.Role)
ssoProvider, exists := c.Get("sso_provider")
assert.True(t, exists)
assert.Equal(t, "oidc", 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)
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)
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
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)
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
// Pre-create a user
account, user := seedAccountAndUser(t, db, "Test Account", "Existing User", "existing@example.com", "oidc")
// findOrCreateUser should find the existing user
foundUser, err := mw.findOrCreateUser(context.Background(), account.ID, "existing@example.com", "Existing User", "nameid-1", "oidc", "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)
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", "oidc", "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, "oidc", 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_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(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)
// SSOMiddleware with nil DB
mw := NewSSOMiddleware(nil, rdb, cfg, nil)
_, err := mw.findOrCreateUser(context.Background(), 1, "test@example.com", "Test", "uid-1", "oidc", "agent")
require.Error(t, err)
assert.Contains(t, err.Error(), "database not available")
}
// ========== Additional tests: CreateSSOSession ==========
func TestCreateSSOSession_StoresInRedis(t *testing.T) {
db := newSSOTestDB(t)
mr, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
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)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
UserID: 99,
AccountID: 5,
Email: "oidc@example.com",
Subject: "sub-oidc",
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, "oidc", sessionData.Provider)
assert.Equal(t, "sub-oidc", sessionData.NameID)
assert.Equal(t, "agent", sessionData.Role)
assert.NotZero(t, sessionData.CreatedAt)
assert.NotZero(t, sessionData.ExpiresAt)
}
// ========== Additional tests: AuthenticateOIDC with nil service ==========
func TestAuthenticateOIDC_NilService(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(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)
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)
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: "oidc",
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)
}
// ========== Comprehensive integration: full SSO flow with OIDC ==========
func TestSSOFlow_OIDCFindOrCreateThenSession(t *testing.T) {
db := newSSOTestDB(t)
_, rdb := newSSOTestRedis(t)
cfg := newSSOTestConfig(true)
mw := newSSOMiddlewareForTest(t, db, rdb, cfg)
account := &model.Account{Name: "Flow Account"}
require.NoError(t, db.Create(account).Error)
// Step 1: Find/create user
user, err := mw.findOrCreateUser(context.Background(), account.ID, "flowuser@example.com", "Flow User", "sub-flowuser", "oidc", "agent")
require.NoError(t, err)
assert.NotEqual(t, uint(0), user.ID)
// Step 2: Issue JWT
result := &SSOAuthResult{
Provider: SSOProviderOIDC,
UserID: user.ID,
AccountID: account.ID,
Email: "flowuser@example.com",
Name: "Flow User",
Subject: "sub-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, "oidc", 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())
}