package service import ( "context" "crypto/hmac" "crypto/sha1" "encoding/base32" "encoding/binary" "fmt" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" "github.com/gochat/gochat/internal/auth" "github.com/gochat/gochat/internal/model" "github.com/gochat/gochat/pkg/crypto" ) // ========== Helper: create confirmed user with account ========== // createConfirmedUserWithAccount 创建已确认邮箱的用户,并为其分配一个账户。 // 返回 user、account 和 accountUser 关联记录。 func createConfirmedUserWithAccount(t *testing.T, db *gorm.DB, email, password, name string) (*model.User, *model.Account, *model.AccountUser) { t.Helper() account := &model.Account{Name: name + "'s Account", Status: "active"} require.NoError(t, db.Create(account).Error, "无法创建账户") digest, err := crypto.HashPassword(password) require.NoError(t, err, "无法 hash 密码") now := time.Now() user := &model.User{ AccountID: account.ID, Name: name, Email: email, PasswordDigest: digest, Password: digest, Provider: "email", Active: true, ConfirmedAt: &now, } require.NoError(t, db.Create(user).Error, "无法创建用户") accountUser := &model.AccountUser{ UserID: user.ID, AccountID: account.ID, Role: "administrator", } require.NoError(t, db.Create(accountUser).Error, "无法创建 AccountUser") return user, account, accountUser } // ========== TOTP code generation helper ========== // computeTOTPCode 根据 secret 和时间步计算 6 位 TOTP code (HMAC-SHA1)。 func computeTOTPCode(secret string, timeStep uint64, digits int) string { // Decode base32 secret (no padding) secretBytes, err := decodeBase32NoPadding(secret) if err != nil { // fallback: 如果 secret 无法 decode,返回 000000 return "000000" } // 时间步编码为 8 字节大端 stepBytes := make([]byte, 8) binary.BigEndian.PutUint64(stepBytes, timeStep) // HMAC-SHA1 hmacResult := hmac.New(sha1.New, secretBytes) hmacResult.Write(stepBytes) hash := hmacResult.Sum(nil) // Dynamic truncation per RFC 4226 offset := hash[len(hash)-1] & 0x0f code := binary.BigEndian.Uint32(hash[offset:offset+4]) & 0x7fffffff // 取 digits 位数字 codeMod := code % uint32(pow10(digits)) return fmt.Sprintf("%0*d", digits, codeMod) } // decodeBase32NoPadding 解码 base32 字符串(无 padding)。 func decodeBase32NoPadding(s string) ([]byte, error) { return base32.StdEncoding.WithPadding(base32.NoPadding).DecodeString(s) } // pow10 计算 10^n (简单整数幂)。 func pow10(n int) uint32 { result := uint32(1) for i := 0; i < n; i++ { result *= 10 } return result } // generateTOTPCode 根据 secret 和当前时间生成 TOTP code。 func generateTOTPCode(secret string) string { cfg := auth.DefaultTOTPConfig() timeStep := uint64(time.Now().Unix()) / cfg.Period return computeTOTPCode(secret, timeStep, cfg.Digits) } // ========== Login Tests ========== // TestAuthService_Login_Success 验证正常登录流程返回 token pair func TestAuthService_Login_Success(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, svc.db, "login@example.com", "password123", "LoginUser") output, err := svc.Login(ctx, &LoginInput{Email: "login@example.com", Password: "password123"}) require.NoError(t, err, "Login 不应返回错误") assert.NotNil(t, output.User) assert.NotNil(t, output.TokenPair) assert.NotEmpty(t, output.TokenPair.AccessToken) assert.NotEmpty(t, output.TokenPair.RefreshToken) assert.Equal(t, user.ID, output.User.ID) assert.True(t, output.User.SignInCount >= 1) } // TestAuthService_Login_WrongPassword 验证密码错误时返回错误 func TestAuthService_Login_WrongPassword(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() _, _, _ = createConfirmedUserWithAccount(t, svc.db, "wrongpw@example.com", "password123", "WrongPWUser") output, err := svc.Login(ctx, &LoginInput{Email: "wrongpw@example.com", Password: "wrong-password"}) assert.Nil(t, output) assert.EqualError(t, err, "invalid email or password") } // TestAuthService_Login_UserNotFound 验证用户不存在时返回错误 func TestAuthService_Login_UserNotFound(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() output, err := svc.Login(ctx, &LoginInput{Email: "nonexistent@example.com", Password: "any-password"}) assert.Nil(t, output) assert.EqualError(t, err, "invalid email or password") } // TestAuthService_Login_InactiveUser 验证 inactive 用户无法登录 func TestAuthService_Login_InactiveUser(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() now := time.Now() account := &model.Account{Name: "InactiveAccount", Status: "active"} require.NoError(t, svc.db.Create(account).Error) digest, err := crypto.HashPassword("password123") require.NoError(t, err) user := &model.User{ AccountID: account.ID, Name: "InactiveUser", Email: "inactive@example.com", PasswordDigest: digest, Password: digest, Provider: "email", ConfirmedAt: &now, } // GORM零值bool问题:先Create Active=true,再Update Active=false require.NoError(t, svc.db.Create(user).Error) require.NoError(t, svc.db.Model(user).Update("Active", false).Error) au := &model.AccountUser{UserID: user.ID, AccountID: account.ID, Role: "agent"} require.NoError(t, svc.db.Create(au).Error) output, err := svc.Login(ctx, &LoginInput{Email: "inactive@example.com", Password: "password123"}) assert.Nil(t, output) assert.EqualError(t, err, "user account is inactive") } // TestAuthService_Login_UnconfirmedEmail 验证未确认邮箱的用户无法登录 func TestAuthService_Login_UnconfirmedEmail(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() account := &model.Account{Name: "UnconfirmedAccount", Status: "active"} require.NoError(t, svc.db.Create(account).Error) digest, err := crypto.HashPassword("password123") require.NoError(t, err) user := &model.User{ AccountID: account.ID, Name: "UnconfirmedUser", Email: "unconfirmed@example.com", PasswordDigest: digest, Password: digest, Provider: "email", Active: true, // ConfirmedAt intentionally nil — not confirmed yet } require.NoError(t, svc.db.Create(user).Error) output, err := svc.Login(ctx, &LoginInput{Email: "unconfirmed@example.com", Password: "password123"}) assert.Nil(t, output) assert.EqualError(t, err, "email not confirmed, please verify your email first") } // TestAuthService_Login_OAuthOnlyUser 验证 OAuth-only 用户无法用密码登录 func TestAuthService_Login_OAuthOnlyUser(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() now := time.Now() account := &model.Account{Name: "OAuthAccount", Status: "active"} require.NoError(t, svc.db.Create(account).Error) user := &model.User{ AccountID: account.ID, Name: "OAuthUser", Email: "oauth@example.com", Provider: "google", UID: "google-12345", Active: true, ConfirmedAt: &now, } require.NoError(t, svc.db.Create(user).Error) au := &model.AccountUser{UserID: user.ID, AccountID: account.ID, Role: "agent"} require.NoError(t, svc.db.Create(au).Error) output, err := svc.Login(ctx, &LoginInput{Email: "oauth@example.com", Password: "any-password"}) assert.Nil(t, output) assert.Contains(t, err.Error(), "this account uses google authentication") } // TestAuthService_Login_MFARequired 验证启用 MFA 的用户登录返回 MFARequired=true func TestAuthService_Login_MFARequired(t *testing.T) { _, _, _, _, mfaSvc, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, svc.db, "mfa@example.com", "password123", "MFAUser") secret, _, err := mfaSvc.GenerateTOTPSecret(user.ID) require.NoError(t, err) require.NoError(t, mfaSvc.EnableTOTP(user.ID, secret)) output, err := svc.Login(ctx, &LoginInput{Email: "mfa@example.com", Password: "password123"}) require.NoError(t, err) assert.True(t, output.MFARequired) assert.Nil(t, output.TokenPair) } // TestAuthService_Login_NoAccount 验证没有 AccountUser 关联的用户登录失败 func TestAuthService_Login_NoAccount(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() now := time.Now() digest, err := crypto.HashPassword("password123") require.NoError(t, err) user := &model.User{ Name: "NoAccountUser", Email: "noaccount@example.com", PasswordDigest: digest, Password: digest, Provider: "email", Active: true, ConfirmedAt: &now, } require.NoError(t, svc.db.Create(user).Error) output, err := svc.Login(ctx, &LoginInput{Email: "noaccount@example.com", Password: "password123"}) assert.Nil(t, output) assert.Contains(t, err.Error(), "failed to get user account") } // ========== Register Tests ========== // TestAuthService_Register_Success 验证正常注册流程 func TestAuthService_Register_Success(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() user, err := svc.Register(ctx, &RegisterInput{ Name: "NewUser", Email: "register@example.com", Password: "password123", }) require.NoError(t, err) assert.Equal(t, "register@example.com", user.Email) assert.Equal(t, "NewUser", user.Name) assert.Equal(t, "email", user.Provider) assert.NotEmpty(t, user.PasswordDigest) assert.NotNil(t, user.ConfirmedAt) // auto-confirmed in dev mode } // TestAuthService_Register_DuplicateEmail 验证重复邮箱注册失败 func TestAuthService_Register_DuplicateEmail(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() _, _, _ = createConfirmedUserWithAccount(t, svc.db, "duplicate@example.com", "password123", "ExistingUser") user, err := svc.Register(ctx, &RegisterInput{ Name: "AnotherUser", Email: "duplicate@example.com", Password: "password456", }) assert.Nil(t, user) assert.EqualError(t, err, "email already registered") } // ========== LoginWithMFA Tests ========== // TestAuthService_LoginWithMFA_Success 验证 MFA 登录验证成功返回 token pair func TestAuthService_LoginWithMFA_Success(t *testing.T) { _, _, _, _, mfaSvc, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, svc.db, "mfalogin@example.com", "password123", "MFALoginUser") secret, _, err := mfaSvc.GenerateTOTPSecret(user.ID) require.NoError(t, err) require.NoError(t, mfaSvc.EnableTOTP(user.ID, secret)) totpCode := generateTOTPCode(secret) output, err := svc.LoginWithMFA(ctx, user.ID, totpCode) require.NoError(t, err) assert.NotNil(t, output.TokenPair) assert.NotNil(t, output.User) assert.False(t, output.MFARequired) } // TestAuthService_LoginWithMFA_InvalidTOTPCode 验证无效 TOTP code 导致 MFA 登录失败 func TestAuthService_LoginWithMFA_InvalidTOTPCode(t *testing.T) { _, _, _, _, mfaSvc, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, svc.db, "mfalogin2@example.com", "password123", "MFALogin2User") secret, _, err := mfaSvc.GenerateTOTPSecret(user.ID) require.NoError(t, err) require.NoError(t, mfaSvc.EnableTOTP(user.ID, secret)) output, err := svc.LoginWithMFA(ctx, user.ID, "000000") assert.Nil(t, output) assert.Contains(t, err.Error(), "invalid totp code") } // TestAuthService_LoginWithMFA_UserNotFound 验证用户不存在时 MFA 登录失败 func TestAuthService_LoginWithMFA_UserNotFound(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() output, err := svc.LoginWithMFA(ctx, 9999, "123456") assert.Nil(t, output) assert.Contains(t, err.Error(), "user not found") } // ========== RefreshToken Tests ========== // TestAuthService_Refresh_Success 验证正常 refresh token 可以刷新 func TestAuthService_Refresh_Success(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() _, _, _ = createConfirmedUserWithAccount(t, svc.db, "refresh@example.com", "password123", "RefreshUser") loginOut, err := svc.Login(ctx, &LoginInput{Email: "refresh@example.com", Password: "password123"}) require.NoError(t, err) // Refresh token 已在 Login 时存储到 Redis,可直接刷新 output, err := svc.Refresh(ctx, &RefreshInput{RefreshToken: loginOut.TokenPair.RefreshToken}) require.NoError(t, err) assert.NotNil(t, output.TokenPair) assert.NotEmpty(t, output.TokenPair.AccessToken) assert.NotEmpty(t, output.TokenPair.RefreshToken) } // TestAuthService_Refresh_InvalidToken 验证无效 refresh token 无法刷新 func TestAuthService_Refresh_InvalidToken(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() output, err := svc.Refresh(ctx, &RefreshInput{RefreshToken: "invalid-refresh-token"}) assert.Nil(t, output) assert.Contains(t, err.Error(), "invalid refresh token") } // TestAuthService_Refresh_RevokedToken 验证已撤销的 refresh token 无法刷新 func TestAuthService_Refresh_RevokedToken(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() _, _, _ = createConfirmedUserWithAccount(t, svc.db, "revoked@example.com", "password123", "RevokedUser") loginOut, err := svc.Login(ctx, &LoginInput{Email: "revoked@example.com", Password: "password123"}) require.NoError(t, err) // Logout 撤销 refresh token require.NoError(t, svc.Logout(ctx, loginOut.User.ID)) output, err := svc.Refresh(ctx, &RefreshInput{RefreshToken: loginOut.TokenPair.RefreshToken}) assert.Nil(t, output) assert.Contains(t, err.Error(), "refresh token expired or revoked") } // ========== Logout Tests ========== // TestAuthService_Logout_Success 验证 logout 撤销 refresh token func TestAuthService_Logout_Success(t *testing.T) { db, _, refreshStore, _, _, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, db, "logout@example.com", "password123", "LogoutUser") loginOut, err := svc.Login(ctx, &LoginInput{Email: "logout@example.com", Password: "password123"}) require.NoError(t, err) // 确保 refresh token 存在于 Redis require.NoError(t, refreshStore.Store(ctx, user.ID, loginOut.TokenPair.RefreshToken)) err = svc.Logout(ctx, user.ID) require.NoError(t, err) valid, err := refreshStore.Validate(ctx, user.ID, loginOut.TokenPair.RefreshToken) require.NoError(t, err) assert.False(t, valid) } // ========== SwitchAccount Tests ========== // TestAuthService_SwitchAccount_Success 验证切换账户返回新 token pair func TestAuthService_SwitchAccount_Success(t *testing.T) { db, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, db, "switch@example.com", "password123", "SwitchUser") // 创建第二个账户并关联 account2 := &model.Account{Name: "Second Account", Status: "active"} require.NoError(t, db.Create(account2).Error) au2 := &model.AccountUser{UserID: user.ID, AccountID: account2.ID, Role: "agent"} require.NoError(t, db.Create(au2).Error) output, err := svc.SwitchAccount(ctx, &SwitchAccountInput{UserID: user.ID, AccountID: account2.ID}) require.NoError(t, err) assert.NotNil(t, output.TokenPair) assert.Equal(t, account2.ID, output.AccountID) assert.Equal(t, "agent", output.Role) } // TestAuthService_SwitchAccount_NotMember 验证不属于目标账户时切换失败 func TestAuthService_SwitchAccount_NotMember(t *testing.T) { db, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() user, _, _ := createConfirmedUserWithAccount(t, db, "switch2@example.com", "password123", "Switch2User") otherAccount := &model.Account{Name: "Other Account", Status: "active"} require.NoError(t, db.Create(otherAccount).Error) output, err := svc.SwitchAccount(ctx, &SwitchAccountInput{UserID: user.ID, AccountID: otherAccount.ID}) assert.Nil(t, output) assert.Contains(t, err.Error(), fmt.Sprintf("user does not belong to account %d", otherAccount.ID)) } // ========== ResetPassword Tests ========== // TestAuthService_ResetPassword_Success 验证密码重置请求不报错 func TestAuthService_ResetPassword_Success(t *testing.T) { db, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() _, _, _ = createConfirmedUserWithAccount(t, db, "reset@example.com", "password123", "ResetUser") err := svc.ResetPassword(ctx, &ResetPasswordInput{Email: "reset@example.com"}) assert.NoError(t, err) } // TestAuthService_ResetPassword_EmailNotFound 验证不存在邮箱也不报错(安全最佳实践) func TestAuthService_ResetPassword_EmailNotFound(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() err := svc.ResetPassword(ctx, &ResetPasswordInput{Email: "nonexistent@example.com"}) assert.NoError(t, err) } // ========== ConfirmResetPassword Tests ========== // TestAuthService_ConfirmResetPassword_NotImplemented 验证当前返回未实现错误 func TestAuthService_ConfirmResetPassword_NotImplemented(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() err := svc.ConfirmResetPassword(ctx, &ConfirmResetPasswordInput{Token: "some-token", Password: "new-password"}) assert.Error(t, err) assert.Contains(t, err.Error(), "not yet implemented") } // ========== ConfirmEmail Tests ========== // TestAuthService_ConfirmEmail_NotImplemented 验证当前返回未实现错误 func TestAuthService_ConfirmEmail_NotImplemented(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() user, err := svc.ConfirmEmail(ctx, &ConfirmEmailInput{Token: "some-token"}) assert.Nil(t, user) assert.Error(t, err) assert.Contains(t, err.Error(), "not yet implemented") } // ========== OAuthLogin Tests ========== // TestAuthService_OAuthLogin_ExchangeCodeFailed 验证 OAuth code exchange 失败 func TestAuthService_OAuthLogin_ExchangeCodeFailed(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() output, err := svc.OAuthLogin(ctx, &OAuthLoginInput{ Provider: auth.OAuthProviderGoogle, Code: "invalid-oauth-code", State: "test-state", }) assert.Nil(t, output) assert.Contains(t, err.Error(), "oauth exchange failed") } // TestAuthService_OAuthLogin_UnconfiguredProvider 验证未配置的 OAuth provider func TestAuthService_OAuthLogin_UnconfiguredProvider(t *testing.T) { _, _, _, _, _, svc := setupAuthService(t) ctx := context.Background() output, err := svc.OAuthLogin(ctx, &OAuthLoginInput{ Provider: auth.OAuthProviderCustom, Code: "some-code", State: "test-state", }) assert.Nil(t, output) assert.Contains(t, err.Error(), "oauth") } // ========== ValidateToken Tests ========== // TestAuthService_ValidateAccessToken_Success 验证生成的 access token 可被正确解析 func TestAuthService_ValidateAccessToken_Success(t *testing.T) { db, jwtSvc, _, _, _, svc := setupAuthService(t) ctx := context.Background() user, account, _ := createConfirmedUserWithAccount(t, db, "validate@example.com", "password123", "ValidateUser") loginOut, err := svc.Login(ctx, &LoginInput{Email: "validate@example.com", Password: "password123"}) require.NoError(t, err) claims, err := jwtSvc.ValidateAccessToken(loginOut.TokenPair.AccessToken) require.NoError(t, err) assert.Equal(t, user.ID, claims.UserID) assert.Equal(t, account.ID, claims.AccountID) } // TestAuthService_ValidateAccessToken_InvalidToken 验证无效 access token 解析失败 func TestAuthService_ValidateAccessToken_InvalidToken(t *testing.T) { _, jwtSvc, _, _, _, _ := setupAuthService(t) claims, err := jwtSvc.ValidateAccessToken("invalid-access-token") assert.Nil(t, claims) assert.Error(t, err) } // ========== createDefaultAccount Tests ========== // TestAuthService_CreateDefaultAccount_Success 验证为新用户自动创建默认账户 func TestAuthService_CreateDefaultAccount_Success(t *testing.T) { db, _, _, _, _, svc := setupAuthService(t) now := time.Now() digest, err := crypto.HashPassword("password123") require.NoError(t, err) user := &model.User{ Name: "NewOAuthUser", Email: "newoauth@example.com", PasswordDigest: digest, Password: digest, Provider: "google", UID: "google-999", Active: true, ConfirmedAt: &now, } require.NoError(t, db.Create(user).Error) accountID, role, err := svc.createDefaultAccount(user) require.NoError(t, err) assert.NotZero(t, accountID) assert.Equal(t, "administrator", role) var au AccountUser require.NoError(t, db.Where("user_id = ? AND account_id = ?", user.ID, accountID).First(&au).Error) assert.Equal(t, "administrator", au.Role) } // ========== getUserDefaultAccount Tests ========== // TestAuthService_GetUserDefaultAccount_Success 验证获取用户默认账户 func TestAuthService_GetUserDefaultAccount_Success(t *testing.T) { db, _, _, _, _, svc := setupAuthService(t) user, account, au := createConfirmedUserWithAccount(t, db, "default@example.com", "password123", "DefaultUser") accountID, role, err := svc.getUserDefaultAccount(user) require.NoError(t, err) assert.Equal(t, account.ID, accountID) assert.Equal(t, au.Role, role) } // TestAuthService_GetUserDefaultAccount_NoAccount 验证没有账户关联时返回错误 func TestAuthService_GetUserDefaultAccount_NoAccount(t *testing.T) { db, _, _, _, _, svc := setupAuthService(t) now := time.Now() digest, err := crypto.HashPassword("password123") require.NoError(t, err) user := &model.User{ Name: "NoAccountDefault", Email: "noaccountdefault@example.com", PasswordDigest: digest, Password: digest, Provider: "email", Active: true, ConfirmedAt: &now, } require.NoError(t, db.Create(user).Error) accountID, role, err := svc.getUserDefaultAccount(user) assert.Zero(t, accountID) assert.Empty(t, role) assert.EqualError(t, err, "user has no account") }