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