From fbafb13a8aa6eea67bd61552404b812f41fff420 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Thu, 30 Apr 2026 17:42:24 +0800 Subject: [PATCH] refactor(auth): simplify admin session token flow --- config/config.example.yaml | 3 +- internal/controllers/api/auth_controller.go | 23 +- .../dashboard/session_controller.go | 2 - internal/models/models.go | 39 +-- internal/pkg/config/config.go | 3 +- internal/pkg/constants/auth.go | 8 +- internal/pkg/dto/request/auth_request.go | 8 - internal/pkg/dto/response/admin_response.go | 7 +- internal/pkg/dto/response/auth_response.go | 11 +- internal/pkg/errorsx/errors.go | 17 +- internal/services/auth_service.go | 143 +++----- internal/services/auth_service_test.go | 326 +++++++++++++++++- 12 files changed, 411 insertions(+), 179 deletions(-) diff --git a/config/config.example.yaml b/config/config.example.yaml index be88c2f..fb044f6 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -15,8 +15,7 @@ logger: addSource: false auth: - accessTokenTTLHours: 12 - refreshTokenTTLDays: 7 + tokenTTLHours: 12 maxFailedAttempts: 5 credentialLockMinute: 15 diff --git a/internal/controllers/api/auth_controller.go b/internal/controllers/api/auth_controller.go index 5166aac..95091de 100644 --- a/internal/controllers/api/auth_controller.go +++ b/internal/controllers/api/auth_controller.go @@ -30,20 +30,6 @@ func (c *AuthController) PostLogin() *web.JsonResult { return web.JsonData(ret) } -func (c *AuthController) PostRefresh_token() *web.JsonResult { - cfg := config.Current() - req := request.RefreshTokenRequest{} - if err := params.ReadJSON(c.Ctx, &req); err != nil { - return web.JsonError(err) - } - - ret, err := services.AuthService.RefreshToken(req.RefreshToken, cfg.Auth, c.Ctx.RemoteAddr(), c.Ctx.GetHeader("User-Agent")) - if err != nil { - return web.JsonError(err) - } - return web.JsonData(ret) -} - func (c *AuthController) GetWxwork_login() { loginURL, err := services.WxWorkLoginService.BuildWxWorkLoginURL(c.Ctx.URLParam("next")) if err != nil { @@ -91,14 +77,7 @@ func (c *AuthController) PostWxwork_exchange() *web.JsonResult { } func (c *AuthController) PostLogout() *web.JsonResult { - req := request.LogoutRequest{} - if c.Ctx.GetContentLength() > 0 { - if err := params.ReadJSON(c.Ctx, &req); err != nil { - return web.JsonError(err) - } - } - - if err := services.AuthService.Logout(c.Ctx.GetHeader("Authorization"), req.RefreshToken); err != nil { + if err := services.AuthService.Logout(c.Ctx.GetHeader("Authorization")); err != nil { return web.JsonError(err) } return web.JsonSuccess() diff --git a/internal/controllers/dashboard/session_controller.go b/internal/controllers/dashboard/session_controller.go index b8362f1..7417b46 100644 --- a/internal/controllers/dashboard/session_controller.go +++ b/internal/controllers/dashboard/session_controller.go @@ -24,7 +24,6 @@ func (c *SessionController) AnyList() *web.JsonResult { cnd := params.NewPagedSqlCnd(c.Ctx, params.QueryFilter{ParamName: "userId"}, - params.QueryFilter{ParamName: "tokenType"}, params.QueryFilter{ParamName: "clientType"}, ).Desc("id") list, paging := services.LoginSessionService.FindPageByCnd(cnd) @@ -38,7 +37,6 @@ func (c *SessionController) AnyList() *web.JsonResult { ID: item.ID, UserID: item.UserID, Username: username, - TokenType: item.TokenType, ClientType: item.ClientType, ClientIP: item.ClientIP, UserAgent: item.UserAgent, diff --git a/internal/models/models.go b/internal/models/models.go index fe95898..7fff0b8 100644 --- a/internal/models/models.go +++ b/internal/models/models.go @@ -287,31 +287,30 @@ type UserPermission struct { AuditFields } -// LoginSession 登录会话或刷新令牌记录。 +// LoginSession 表示一次后台登录会话。 type LoginSession struct { - ID int64 `gorm:"primaryKey;autoIncrement"` - UserID int64 `gorm:"type:bigint;not null;index"` - TokenID string `gorm:"type:varchar(128);not null;uniqueIndex"` - TokenType string `gorm:"type:varchar(20);not null;default:'';index"` - ClientType string `gorm:"type:varchar(50);not null;default:'';index"` - ClientIP string `gorm:"type:varchar(64);not null;default:''"` - UserAgent string `gorm:"type:varchar(255);not null;default:''"` - ExpiredAt time.Time `gorm:"type:datetime;not null;index"` - RevokedAt *time.Time `gorm:"type:datetime;index"` - LastSeenAt *time.Time `gorm:"type:datetime"` + ID int64 `gorm:"primaryKey;autoIncrement"` // ID 为登录会话主键。 + UserID int64 `gorm:"type:bigint;not null;index"` // UserID 为登录用户 ID。 + Token string `gorm:"type:varchar(128);not null;uniqueIndex"` // Token 为随机不透明登录凭证,使用 ak_ 前缀。 + ClientType string `gorm:"type:varchar(50);not null;default:'';index"` // ClientType 为客户端类型,后台 Web 端固定为 admin_web。 + ClientIP string `gorm:"type:varchar(64);not null;default:''"` // ClientIP 为登录请求来源 IP。 + UserAgent string `gorm:"type:varchar(255);not null;default:''"` // UserAgent 为登录请求浏览器或客户端 UA。 + ExpiredAt time.Time `gorm:"type:datetime;not null;index"` // ExpiredAt 为 token 过期时间。 + RevokedAt *time.Time `gorm:"type:datetime;index"` // RevokedAt 为主动注销或踢下线时间,非空表示已失效。 + LastSeenAt *time.Time `gorm:"type:datetime"` // LastSeenAt 为最近一次成功鉴权时间。 AuditFields } -// LoginCredentialLog 登录凭证校验日志。 +// LoginCredentialLog 记录一次后台登录凭证校验结果。 type LoginCredentialLog struct { - ID int64 `gorm:"primaryKey;autoIncrement"` - Principal string `gorm:"type:varchar(100);not null;default:'';index"` - UserID int64 `gorm:"type:bigint;not null;default:0;index"` - Success bool `gorm:"not null;default:false;index"` - ClientIP string `gorm:"type:varchar(64);not null;default:''"` - UserAgent string `gorm:"type:varchar(255);not null;default:''"` - Reason string `gorm:"type:varchar(255);not null;default:''"` - CreatedAt time.Time `gorm:"type:datetime;not null;index"` + ID int64 `gorm:"primaryKey;autoIncrement"` // ID 为登录凭证日志主键。 + Principal string `gorm:"type:varchar(100);not null;default:'';index"` // Principal 为用户输入的登录名。 + UserID int64 `gorm:"type:bigint;not null;default:0;index"` // UserID 为匹配到的用户 ID,未匹配时为 0。 + Success bool `gorm:"not null;default:false;index"` // Success 表示本次凭证校验是否成功。 + ClientIP string `gorm:"type:varchar(64);not null;default:''"` // ClientIP 为登录请求来源 IP。 + UserAgent string `gorm:"type:varchar(255);not null;default:''"` // UserAgent 为登录请求浏览器或客户端 UA。 + Reason string `gorm:"type:varchar(255);not null;default:''"` // Reason 为校验结果原因。 + CreatedAt time.Time `gorm:"type:datetime;not null;index"` // CreatedAt 为日志创建时间。 } // Asset 存储的文件资源,如上传的附件等。 diff --git a/internal/pkg/config/config.go b/internal/pkg/config/config.go index 7e555fd..1558f34 100644 --- a/internal/pkg/config/config.go +++ b/internal/pkg/config/config.go @@ -55,8 +55,7 @@ type LoggerConfig struct { } type AuthConfig struct { - AccessTokenTTLHours int `yaml:"accessTokenTTLHours"` - RefreshTokenTTLDays int `yaml:"refreshTokenTTLDays"` + TokenTTLHours int `yaml:"tokenTTLHours"` MaxFailedAttempts int `yaml:"maxFailedAttempts"` CredentialLockMinute int `yaml:"credentialLockMinute"` } diff --git a/internal/pkg/constants/auth.go b/internal/pkg/constants/auth.go index 0e53636..ec5c8fd 100644 --- a/internal/pkg/constants/auth.go +++ b/internal/pkg/constants/auth.go @@ -8,13 +8,7 @@ const ( ) const ( - AccessTokenPrefix = "atk_" - RefreshTokenPrefix = "rtk_" -) - -const ( - TokenTypeAccess = "access" - TokenTypeRefresh = "refresh" + AuthTokenPrefix = "ak_" ) const ( diff --git a/internal/pkg/dto/request/auth_request.go b/internal/pkg/dto/request/auth_request.go index 3f38c3b..112bd87 100644 --- a/internal/pkg/dto/request/auth_request.go +++ b/internal/pkg/dto/request/auth_request.go @@ -5,14 +5,6 @@ type LoginRequest struct { Password string `json:"password"` } -type RefreshTokenRequest struct { - RefreshToken string `json:"refreshToken"` -} - -type LogoutRequest struct { - RefreshToken string `json:"refreshToken"` -} - type WxWorkExchangeRequest struct { Ticket string `json:"ticket"` } diff --git a/internal/pkg/dto/response/admin_response.go b/internal/pkg/dto/response/admin_response.go index 8d6592d..ae585d6 100644 --- a/internal/pkg/dto/response/admin_response.go +++ b/internal/pkg/dto/response/admin_response.go @@ -47,12 +47,11 @@ type CreateUserResultResponse struct { type SessionResponse struct { ID int64 `json:"id"` UserID int64 `json:"userId"` - Username string `json:"username,omitempty"` - TokenType string `json:"tokenType"` + Username string `json:"username"` ClientType string `json:"clientType"` ClientIP string `json:"clientIp"` UserAgent string `json:"userAgent"` ExpiredAt string `json:"expiredAt"` - RevokedAt string `json:"revokedAt,omitempty"` - LastSeenAt string `json:"lastSeenAt,omitempty"` + RevokedAt string `json:"revokedAt"` + LastSeenAt string `json:"lastSeenAt"` } diff --git a/internal/pkg/dto/response/auth_response.go b/internal/pkg/dto/response/auth_response.go index 96fd2aa..e0c2a1d 100644 --- a/internal/pkg/dto/response/auth_response.go +++ b/internal/pkg/dto/response/auth_response.go @@ -12,10 +12,9 @@ type AuthUserResponse struct { } type LoginResponse struct { - AccessToken string `json:"accessToken"` - RefreshToken string `json:"refreshToken"` - ExpiresAt string `json:"expiresAt"` - User *AuthUserResponse `json:"user"` - Permissions []string `json:"permissions"` - Roles []string `json:"roles"` + AccessToken string `json:"accessToken"` + ExpiresAt string `json:"expiresAt"` + User *AuthUserResponse `json:"user"` + Permissions []string `json:"permissions"` + Roles []string `json:"roles"` } diff --git a/internal/pkg/errorsx/errors.go b/internal/pkg/errorsx/errors.go index 938cead..12bb2de 100644 --- a/internal/pkg/errorsx/errors.go +++ b/internal/pkg/errorsx/errors.go @@ -3,12 +3,13 @@ package errorsx import "github.com/mlogclub/simple/web" const ( - CodeInvalidParam = 1000 - CodeBusinessError = 2000 - CodeAuthUnauthorized = 3000 - CodeAuthForbidden = 3001 - CodeAuthInvalidToken = 3002 - CodeAuthInvalidAccount = 3003 + CodeInvalidParam = 1000 + CodeBusinessError = 2000 + CodeAuthUnauthorized = 3000 + CodeAuthForbidden = 3001 + CodeAuthInvalidToken = 3002 + CodeAuthInvalidAccount = 3003 + CodeAuthCredentialLocked = 3004 ) func InvalidParam(message string) error { @@ -34,3 +35,7 @@ func InvalidToken(message string) error { func InvalidAccount(message string) error { return web.NewError(CodeAuthInvalidAccount, message) } + +func CredentialLocked(message string) error { + return web.NewError(CodeAuthCredentialLocked, message) +} diff --git a/internal/services/auth_service.go b/internal/services/auth_service.go index 17a3deb..64d2e23 100644 --- a/internal/services/auth_service.go +++ b/internal/services/auth_service.go @@ -87,6 +87,14 @@ func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, } user := UserService.GetByUsername(username) + if s.isCredentialLocked(username, authCfg) { + userID := int64(0) + if user != nil { + userID = user.ID + } + _ = s.createLoginCredentialLog(username, userID, false, clientIP, userAgent, "credential locked") + return nil, errorsx.CredentialLocked("登录失败次数过多,请稍后再试") + } if user == nil || user.Status != enums.StatusOk { _ = s.createLoginCredentialLog(username, 0, false, clientIP, userAgent, "user not found") return nil, errorsx.InvalidAccount("用户名或密码错误") @@ -121,56 +129,11 @@ func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, return ret, nil } -func (s *authService) RefreshToken(refreshToken string, authCfg config.AuthConfig, clientIP, userAgent string) (*response.LoginResponse, error) { - session, err := s.validateSessionToken(refreshToken, constants.TokenTypeRefresh) - if err != nil { - return nil, err - } - if session.RevokedAt != nil { - return nil, errorsx.InvalidToken("refresh token 已失效") - } - - user := UserService.Get(session.UserID) - if user == nil || user.Status != enums.StatusOk { - return nil, errorsx.Unauthorized("用户不存在或已被禁用") - } - - var ret *response.LoginResponse - if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error { - var dbErr error - if dbErr = repositories.LoginSessionRepository.Updates(ctx.Tx, session.ID, map[string]any{ - "revoked_at": time.Now(), - "update_user_id": user.ID, - "update_user_name": user.Username, - "updated_at": time.Now(), - }); dbErr != nil { - return dbErr - } - if ret, dbErr = s.issueTokens(ctx, user, clientIP, userAgent, authCfg); dbErr != nil { - return dbErr - } - return nil - }); err != nil { - return nil, err - } - return ret, nil -} - -func (s *authService) Logout(accessToken, refreshToken string) error { +func (s *authService) Logout(accessToken string) error { accessToken = s.extractBearerToken(accessToken) now := time.Now() if accessToken != "" { - if session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token_id", accessToken).Eq("token_type", constants.TokenTypeAccess)); session != nil && session.RevokedAt == nil { - if err := LoginSessionService.Updates(session.ID, map[string]any{ - "revoked_at": now, - "updated_at": now, - }); err != nil { - return err - } - } - } - if refreshToken != "" { - if session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token_id", refreshToken).Eq("token_type", constants.TokenTypeRefresh)); session != nil && session.RevokedAt == nil { + if session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token", accessToken)); session != nil && session.RevokedAt == nil { if err := LoginSessionService.Updates(session.ID, map[string]any{ "revoked_at": now, "updated_at": now, @@ -195,7 +158,7 @@ func (s *authService) Authenticate(ctx iris.Context) (*dto.AuthPrincipal, error) return nil, errorsx.Unauthorized("未登录或登录已过期") } - session, err := s.validateSessionToken(token, constants.TokenTypeAccess) + session, err := s.validateSessionToken(token) if err != nil { return nil, err } @@ -262,48 +225,20 @@ func (s *authService) issueTokens(ctx *sqls.TxContext, user *models.User, client return nil, err } - accessTTL, refreshTTL := s.resolveTokenTTL(authCfg) - accessToken, err := randomToken(constants.AccessTokenPrefix) - if err != nil { - return nil, err - } - refreshToken, err := randomToken(constants.RefreshTokenPrefix) + tokenTTL := s.resolveTokenTTL(authCfg) + accessToken, err := randomToken(constants.AuthTokenPrefix) if err != nil { return nil, err } now := time.Now() - // accessSession if err := repositories.LoginSessionRepository.Create(ctx.Tx, &models.LoginSession{ UserID: user.ID, - TokenID: accessToken, - TokenType: constants.TokenTypeAccess, + Token: accessToken, ClientType: constants.ClientTypeAdminWeb, ClientIP: clientIP, UserAgent: userAgent, - ExpiredAt: now.Add(accessTTL), - LastSeenAt: &now, - AuditFields: models.AuditFields{ - CreatedAt: now, - CreateUserID: user.ID, - CreateUserName: user.Username, - UpdatedAt: now, - UpdateUserID: user.ID, - UpdateUserName: user.Username, - }, - }); err != nil { - return nil, err - } - - // refreshSession - if err := repositories.LoginSessionRepository.Create(ctx.Tx, &models.LoginSession{ - UserID: user.ID, - TokenID: refreshToken, - TokenType: constants.TokenTypeRefresh, - ClientType: constants.ClientTypeAdminWeb, - ClientIP: clientIP, - UserAgent: userAgent, - ExpiredAt: now.Add(refreshTTL), + ExpiredAt: now.Add(tokenTTL), LastSeenAt: &now, AuditFields: models.AuditFields{ CreatedAt: now, @@ -318,9 +253,8 @@ func (s *authService) issueTokens(ctx *sqls.TxContext, user *models.User, client } return &response.LoginResponse{ - AccessToken: accessToken, - RefreshToken: refreshToken, - ExpiresAt: now.Add(accessTTL).Format(time.DateTime), + AccessToken: accessToken, + ExpiresAt: now.Add(tokenTTL).Format(time.DateTime), User: &response.AuthUserResponse{ ID: user.ID, Username: user.Username, @@ -334,33 +268,27 @@ func (s *authService) issueTokens(ctx *sqls.TxContext, user *models.User, client }, nil } -func (s *authService) resolveTokenTTL(authCfg config.AuthConfig) (time.Duration, time.Duration) { - accessTTL := 12 * time.Hour - refreshTTL := 7 * 24 * time.Hour - if authCfg.AccessTokenTTLHours > 0 { - accessTTL = time.Duration(authCfg.AccessTokenTTLHours) * time.Hour +func (s *authService) resolveTokenTTL(authCfg config.AuthConfig) time.Duration { + tokenTTL := 12 * time.Hour + if authCfg.TokenTTLHours > 0 { + tokenTTL = time.Duration(authCfg.TokenTTLHours) * time.Hour } - if authCfg.RefreshTokenTTLDays > 0 { - refreshTTL = time.Duration(authCfg.RefreshTokenTTLDays) * 24 * time.Hour - } - return accessTTL, refreshTTL + return tokenTTL } -func (s *authService) validateSessionToken(token, tokenType string) (*models.LoginSession, error) { +func (s *authService) validateSessionToken(token string) (*models.LoginSession, error) { if strings.TrimSpace(token) == "" { return nil, errorsx.InvalidToken("token 不能为空") } - session := LoginSessionService.FindOne(sqls.NewCnd(). - Eq("token_id", token). - Eq("token_type", tokenType)) + session := LoginSessionService.FindOne(sqls.NewCnd().Eq("token", token)) if session == nil { - return nil, errorsx.InvalidToken("token 无效") + return nil, errorsx.InvalidToken("登录凭证无效") } if session.RevokedAt != nil { - return nil, errorsx.InvalidToken("token 已失效") + return nil, errorsx.InvalidToken("登录凭证已失效") } if time.Now().After(session.ExpiredAt) { - return nil, errorsx.InvalidToken("token 已过期") + return nil, errorsx.InvalidToken("登录凭证已过期") } return session, nil } @@ -480,6 +408,23 @@ func (s *authService) createLoginCredentialLog(principal string, userID int64, s }) } +func (s *authService) isCredentialLocked(username string, authCfg config.AuthConfig) bool { + maxFailedAttempts := authCfg.MaxFailedAttempts + if maxFailedAttempts <= 0 { + maxFailedAttempts = 5 + } + lockMinute := authCfg.CredentialLockMinute + if lockMinute <= 0 { + lockMinute = 15 + } + since := time.Now().Add(-time.Duration(lockMinute) * time.Minute) + return LoginCredentialLogService.Count(sqls.NewCnd(). + Eq("principal", username). + Eq("success", false). + NotEq("reason", "credential locked"). + Where("created_at >= ?", since)) >= int64(maxFailedAttempts) +} + func randomToken(prefix string) (string, error) { buf := make([]byte, 24) if _, err := rand.Read(buf); err != nil { diff --git a/internal/services/auth_service_test.go b/internal/services/auth_service_test.go index ca67347..3c250c5 100644 --- a/internal/services/auth_service_test.go +++ b/internal/services/auth_service_test.go @@ -1,6 +1,24 @@ package services -import "testing" +import ( + "errors" + "strings" + "testing" + "time" + + "cs-agent/internal/models" + "cs-agent/internal/pkg/config" + "cs-agent/internal/pkg/dto/request" + "cs-agent/internal/pkg/enums" + "cs-agent/internal/pkg/errorsx" + + "github.com/glebarez/sqlite" + "github.com/mlogclub/simple/sqls" + "github.com/mlogclub/simple/web" + "golang.org/x/crypto/bcrypt" + "gorm.io/gorm" + "gorm.io/gorm/schema" +) func TestExtractBearerToken(t *testing.T) { svc := newAuthService() @@ -13,3 +31,309 @@ func TestExtractBearerToken(t *testing.T) { t.Fatalf("expected raw token to be rejected by bearer extractor, got %q", got) } } + +func TestAuthServiceLoginCreatesSingleAccessSession(t *testing.T) { + db := setupAuthServiceTestDB(t) + user := createAuthTestUser(t, db, "admin", "secret") + svc := newAuthService() + + ret, err := svc.Login(request.LoginRequest{ + Username: " admin ", + Password: "secret", + }, config.AuthConfig{TokenTTLHours: 2, MaxFailedAttempts: 5, CredentialLockMinute: 15}, "127.0.0.1", "go-test") + if err != nil { + t.Fatalf("login failed: %v", err) + } + + if ret.AccessToken == "" || !strings.HasPrefix(ret.AccessToken, "ak_") { + t.Fatalf("expected ak_ access token, got %q", ret.AccessToken) + } + if ret.ExpiresAt == "" { + t.Fatal("expected expiresAt to be returned") + } + + var sessions []models.LoginSession + if err := db.Find(&sessions).Error; err != nil { + t.Fatalf("query login sessions: %v", err) + } + if len(sessions) != 1 { + t.Fatalf("expected exactly one session, got %d", len(sessions)) + } + if sessions[0].Token != ret.AccessToken { + t.Fatalf("expected session token %q, got %q", ret.AccessToken, sessions[0].Token) + } + if sessions[0].UserID != user.ID { + t.Fatalf("expected session user %d, got %d", user.ID, sessions[0].UserID) + } + if sessions[0].ClientType != "admin_web" { + t.Fatalf("expected admin_web client type, got %q", sessions[0].ClientType) + } + + logs := findCredentialLogs(t, db) + if len(logs) != 1 { + t.Fatalf("expected one credential log, got %d", len(logs)) + } + if !logs[0].Success || logs[0].Principal != "admin" || logs[0].UserID != user.ID { + t.Fatalf("unexpected success credential log: %+v", logs[0]) + } +} + +func TestAuthServiceLoginFailureWritesCredentialLogs(t *testing.T) { + db := setupAuthServiceTestDB(t) + createAuthTestUser(t, db, "admin", "secret") + svc := newAuthService() + authCfg := config.AuthConfig{TokenTTLHours: 2, MaxFailedAttempts: 5, CredentialLockMinute: 15} + + if _, err := svc.Login(request.LoginRequest{Username: "missing", Password: "secret"}, authCfg, "127.0.0.1", "go-test"); !hasCode(err, errorsx.CodeAuthInvalidAccount) { + t.Fatalf("expected invalid account for missing user, got %v", err) + } + if _, err := svc.Login(request.LoginRequest{Username: "admin", Password: "wrong"}, authCfg, "127.0.0.1", "go-test"); !hasCode(err, errorsx.CodeAuthInvalidAccount) { + t.Fatalf("expected invalid account for password mismatch, got %v", err) + } + + logs := findCredentialLogs(t, db) + if len(logs) != 2 { + t.Fatalf("expected two credential logs, got %d", len(logs)) + } + if logs[0].Reason != "user not found" || logs[0].Success { + t.Fatalf("unexpected missing-user log: %+v", logs[0]) + } + if logs[1].Reason != "password mismatch" || logs[1].Success { + t.Fatalf("unexpected password-mismatch log: %+v", logs[1]) + } +} + +func TestAuthServiceLoginCredentialLockout(t *testing.T) { + db := setupAuthServiceTestDB(t) + user := createAuthTestUser(t, db, "admin", "secret") + now := time.Now() + for i := 0; i < 2; i++ { + if err := db.Create(&models.LoginCredentialLog{ + Principal: "admin", + UserID: user.ID, + Success: false, + Reason: "password mismatch", + CreatedAt: now.Add(-time.Duration(i+1) * time.Minute), + }).Error; err != nil { + t.Fatalf("seed credential log: %v", err) + } + } + if err := db.Create(&models.LoginCredentialLog{ + Principal: "admin", + UserID: user.ID, + Success: false, + Reason: "password mismatch", + CreatedAt: now.Add(-30 * time.Minute), + }).Error; err != nil { + t.Fatalf("seed old credential log: %v", err) + } + + svc := newAuthService() + _, err := svc.Login(request.LoginRequest{Username: "admin", Password: "secret"}, config.AuthConfig{ + TokenTTLHours: 2, + MaxFailedAttempts: 2, + CredentialLockMinute: 15, + }, "127.0.0.1", "go-test") + if !hasCode(err, errorsx.CodeAuthCredentialLocked) { + t.Fatalf("expected credential locked error, got %v", err) + } + + var lockedLog models.LoginCredentialLog + if err := db.Order("id DESC").Take(&lockedLog).Error; err != nil { + t.Fatalf("query latest credential log: %v", err) + } + if lockedLog.Reason != "credential locked" || lockedLog.Success { + t.Fatalf("unexpected locked credential log: %+v", lockedLog) + } + + var sessionCount int64 + if err := db.Model(&models.LoginSession{}).Count(&sessionCount).Error; err != nil { + t.Fatalf("count sessions: %v", err) + } + if sessionCount != 0 { + t.Fatalf("expected no session while credential locked, got %d", sessionCount) + } +} + +func TestValidateSessionTokenStates(t *testing.T) { + db := setupAuthServiceTestDB(t) + svc := newAuthService() + now := time.Now() + + if _, err := svc.validateSessionToken("missing"); !hasCode(err, errorsx.CodeAuthInvalidToken) { + t.Fatalf("expected invalid token for missing session, got %v", err) + } + + revokedAt := now + if err := db.Create(&models.LoginSession{ + UserID: 1, + Token: "ak_revoked", + ClientType: "admin_web", + ExpiredAt: now.Add(time.Hour), + RevokedAt: &revokedAt, + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + }).Error; err != nil { + t.Fatalf("seed revoked session: %v", err) + } + if _, err := svc.validateSessionToken("ak_revoked"); !hasCode(err, errorsx.CodeAuthInvalidToken) { + t.Fatalf("expected invalid token for revoked session, got %v", err) + } + + if err := db.Create(&models.LoginSession{ + UserID: 1, + Token: "ak_expired", + ClientType: "admin_web", + ExpiredAt: now.Add(-time.Hour), + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + }).Error; err != nil { + t.Fatalf("seed expired session: %v", err) + } + if _, err := svc.validateSessionToken("ak_expired"); !hasCode(err, errorsx.CodeAuthInvalidToken) { + t.Fatalf("expected invalid token for expired session, got %v", err) + } + + if err := db.Create(&models.LoginSession{ + UserID: 1, + Token: "ak_valid", + ClientType: "admin_web", + ExpiredAt: now.Add(time.Hour), + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + }).Error; err != nil { + t.Fatalf("seed valid session: %v", err) + } + session, err := svc.validateSessionToken("ak_valid") + if err != nil { + t.Fatalf("expected valid session token, got %v", err) + } + if session.Token != "ak_valid" { + t.Fatalf("expected valid session token ak_valid, got %q", session.Token) + } +} + +func TestAuthServiceLogoutRevokesCurrentTokenOnly(t *testing.T) { + db := setupAuthServiceTestDB(t) + now := time.Now() + sessions := []models.LoginSession{ + { + UserID: 1, + Token: "ak_current", + ClientType: "admin_web", + ExpiredAt: now.Add(time.Hour), + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + }, + { + UserID: 1, + Token: "ak_other", + ClientType: "admin_web", + ExpiredAt: now.Add(time.Hour), + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + }, + } + if err := db.Create(&sessions).Error; err != nil { + t.Fatalf("seed sessions: %v", err) + } + + if err := newAuthService().Logout("Bearer ak_current"); err != nil { + t.Fatalf("logout failed: %v", err) + } + + var current models.LoginSession + if err := db.Take(¤t, "token = ?", "ak_current").Error; err != nil { + t.Fatalf("query current session: %v", err) + } + if current.RevokedAt == nil { + t.Fatal("expected current session to be revoked") + } + var other models.LoginSession + if err := db.Take(&other, "token = ?", "ak_other").Error; err != nil { + t.Fatalf("query other session: %v", err) + } + if other.RevokedAt != nil { + t.Fatal("expected other session to remain active") + } +} + +func setupAuthServiceTestDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + TablePrefix: "t_", + SingularTable: true, + }, + }) + if err != nil { + t.Fatalf("open sqlite db: %v", err) + } + if err := db.AutoMigrate( + &models.User{}, + &models.Role{}, + &models.Permission{}, + &models.UserRole{}, + &models.RolePermission{}, + &models.UserPermission{}, + &models.LoginSession{}, + &models.LoginCredentialLog{}, + ); err != nil { + t.Fatalf("migrate auth tables: %v", err) + } + sqls.SetDB(db) + return db +} + +func createAuthTestUser(t *testing.T, db *gorm.DB, username, password string) *models.User { + t.Helper() + passwordHash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) + if err != nil { + t.Fatalf("hash password: %v", err) + } + now := time.Now() + user := &models.User{ + Username: username, + Nickname: username, + Password: string(passwordHash), + Status: enums.StatusOk, + AuditFields: models.AuditFields{ + CreatedAt: now, + UpdatedAt: now, + }, + } + if err := db.Create(user).Error; err != nil { + t.Fatalf("create auth test user: %v", err) + } + return user +} + +func findCredentialLogs(t *testing.T, db *gorm.DB) []models.LoginCredentialLog { + t.Helper() + var logs []models.LoginCredentialLog + if err := db.Order("id ASC").Find(&logs).Error; err != nil { + t.Fatalf("query credential logs: %v", err) + } + return logs +} + +func hasCode(err error, code int) bool { + if err == nil { + return false + } + var codeErr *web.CodeError + if errors.As(err, &codeErr) { + return codeErr.Code == code + } + return false +}