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/docs b/docs index 480f237..2bf6ba4 160000 --- a/docs +++ b/docs @@ -1 +1 @@ -Subproject commit 480f237ef416a2a3cbdac92868a8c6a4687c51f2 +Subproject commit 2bf6ba451ede5b6852d301ec3e6f8c92d938b229 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..63d701b 100644 --- a/internal/services/auth_service.go +++ b/internal/services/auth_service.go @@ -81,18 +81,24 @@ func (s *authService) RequirePermission(ctx iris.Context, permission constants.P func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, clientIP, userAgent string) (*response.LoginResponse, error) { username := strings.TrimSpace(req.Username) + principal := normalizeLoginPrincipal(username) password := req.Password if username == "" || strings.TrimSpace(password) == "" { return nil, errorsx.InvalidParam("用户名和密码不能为空") } + if s.isCredentialLocked(principal, authCfg) { + _ = s.createLoginCredentialLog(principal, 0, false, clientIP, userAgent, "credential locked") + return nil, errorsx.CredentialLocked("登录失败次数过多,请稍后再试") + } + user := UserService.GetByUsername(username) if user == nil || user.Status != enums.StatusOk { - _ = s.createLoginCredentialLog(username, 0, false, clientIP, userAgent, "user not found") + _ = s.createLoginCredentialLog(principal, 0, false, clientIP, userAgent, "user not found") return nil, errorsx.InvalidAccount("用户名或密码错误") } if strs.IsBlank(user.Password) || bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password)) != nil { - _ = s.createLoginCredentialLog(username, user.ID, false, clientIP, userAgent, "password mismatch") + _ = s.createLoginCredentialLog(principal, user.ID, false, clientIP, userAgent, "password mismatch") return nil, errorsx.InvalidAccount("用户名或密码错误") } @@ -117,60 +123,15 @@ func (s *authService) Login(req request.LoginRequest, authCfg config.AuthConfig, return nil, err } - _ = s.createLoginCredentialLog(username, user.ID, true, clientIP, userAgent, "") + _ = s.createLoginCredentialLog(principal, user.ID, true, clientIP, userAgent, "") 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 +156,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 +223,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 +251,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 +266,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 不能为空") + return nil, errorsx.Unauthorized("未登录或登录已过期") } - 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 +406,27 @@ func (s *authService) createLoginCredentialLog(principal string, userID int64, s }) } +func (s *authService) isCredentialLocked(principal string, authCfg config.AuthConfig) bool { + maxFailedAttempts := authCfg.MaxFailedAttempts + if maxFailedAttempts <= 0 { + return false + } + lockMinute := authCfg.CredentialLockMinute + if lockMinute <= 0 { + lockMinute = 15 + } + since := time.Now().Add(-time.Duration(lockMinute) * time.Minute) + return LoginCredentialLogService.Count(sqls.NewCnd(). + Eq("principal", normalizeLoginPrincipal(principal)). + Eq("success", false). + NotEq("reason", "credential locked"). + Where("created_at >= ?", since)) >= int64(maxFailedAttempts) +} + +func normalizeLoginPrincipal(principal string) string { + return strings.ToLower(strings.TrimSpace(principal)) +} + 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..f30a9b0 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,409 @@ 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 TestAuthServiceCredentialLockoutDoesNotExtendWhileLocked(t *testing.T) { + db := setupAuthServiceTestDB(t) + createAuthTestUser(t, db, "admin", "secret") + now := time.Now() + entries := []models.LoginCredentialLog{ + { + Principal: "admin", + UserID: 1, + Success: false, + Reason: "password mismatch", + CreatedAt: now.Add(-2 * time.Minute), + }, + { + Principal: "admin", + UserID: 0, + Success: false, + Reason: "credential locked", + CreatedAt: now.Add(-1 * time.Minute), + }, + } + if err := db.Create(&entries).Error; err != nil { + t.Fatalf("seed credential logs: %v", err) + } + + ret, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, config.AuthConfig{ + TokenTTLHours: 2, + MaxFailedAttempts: 2, + CredentialLockMinute: 15, + }, "127.0.0.1", "go-test") + if err != nil { + t.Fatalf("expected locked attempt logs not to extend lockout, got %v", err) + } + if ret == nil || !strings.HasPrefix(ret.AccessToken, "ak_") { + t.Fatalf("expected login response with ak_ token, got %+v", ret) + } +} + +func TestAuthServiceCredentialLockoutNormalizesPrincipalCase(t *testing.T) { + db := setupAuthServiceTestDB(t) + createAuthTestUser(t, db, "admin", "secret") + if err := db.Create(&models.LoginCredentialLog{ + Principal: "admin", + UserID: 1, + Success: false, + Reason: "password mismatch", + CreatedAt: time.Now().Add(-time.Minute), + }).Error; err != nil { + t.Fatalf("seed credential log: %v", err) + } + + _, err := newAuthService().Login(request.LoginRequest{Username: "ADMIN", Password: "secret"}, config.AuthConfig{ + TokenTTLHours: 2, + MaxFailedAttempts: 1, + CredentialLockMinute: 15, + }, "127.0.0.1", "go-test") + if !hasCode(err, errorsx.CodeAuthCredentialLocked) { + t.Fatalf("expected normalized principal to be locked, 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.Principal != "admin" || lockedLog.Reason != "credential locked" { + t.Fatalf("unexpected locked log: %+v", lockedLog) + } +} + +func TestAuthServiceCredentialLockoutDisabledWhenMaxAttemptsNonPositive(t *testing.T) { + db := setupAuthServiceTestDB(t) + user := createAuthTestUser(t, db, "admin", "secret") + now := time.Now() + for i := 0; i < 3; i++ { + if err := db.Create(&models.LoginCredentialLog{ + Principal: "admin", + UserID: user.ID, + Success: false, + Reason: "credential locked", + CreatedAt: now.Add(-time.Duration(i+1) * time.Minute), + }).Error; err != nil { + t.Fatalf("seed credential log: %v", err) + } + } + + ret, err := newAuthService().Login(request.LoginRequest{Username: "admin", Password: "secret"}, config.AuthConfig{ + TokenTTLHours: 2, + MaxFailedAttempts: 0, + CredentialLockMinute: 15, + }, "127.0.0.1", "go-test") + if err != nil { + t.Fatalf("expected lockout to be disabled, got %v", err) + } + if ret == nil || ret.AccessToken == "" { + t.Fatalf("expected login response with access token, got %+v", ret) + } +} + +func TestValidateSessionTokenStates(t *testing.T) { + db := setupAuthServiceTestDB(t) + svc := newAuthService() + now := time.Now() + + if _, err := svc.validateSessionToken(" "); !hasCode(err, errorsx.CodeAuthUnauthorized) { + t.Fatalf("expected unauthorized for empty token, got %v", err) + } + 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 +} diff --git a/web/components/auth-provider.tsx b/web/components/auth-provider.tsx index aaa0f43..f72fbc4 100644 --- a/web/components/auth-provider.tsx +++ b/web/components/auth-provider.tsx @@ -13,6 +13,7 @@ import { usePathname, useRouter } from "next/navigation" import { fetchProfile, logout } from "@/lib/api/auth" import { + AUTH_SESSION_EXPIRED_EVENT, clearSession, readSession, writeSession, @@ -68,17 +69,35 @@ export function AuthProvider({ children }: { children: ReactNode }) { }, [requiresAuth, router]) async function signOut() { - const current = readSession() - await logout(current?.refreshToken) - setSession(null) - startTransition(() => { - router.replace("/dashboard/login") - }) + try { + await logout() + } finally { + setSession(null) + startTransition(() => { + router.replace("/dashboard/login") + }) + } } + useEffect(() => { + function handleAuthExpired() { + setSession(null) + if (requiresAuth) { + startTransition(() => { + router.replace("/dashboard/login") + }) + } + } + + window.addEventListener(AUTH_SESSION_EXPIRED_EVENT, handleAuthExpired) + return () => { + window.removeEventListener(AUTH_SESSION_EXPIRED_EVENT, handleAuthExpired) + } + }, [requiresAuth, router]) + useEffect(() => { const stored = readSession() - setSession(stored) + setSession(stored) if (stored) { void refreshProfile() return diff --git a/web/lib/api/auth.ts b/web/lib/api/auth.ts index bbd5205..f696f23 100644 --- a/web/lib/api/auth.ts +++ b/web/lib/api/auth.ts @@ -30,13 +30,10 @@ export async function fetchProfile() { return request("/api/auth/profile") } -export async function logout(refreshToken?: string) { +export async function logout() { try { await request("/api/auth/logout", { method: "POST", - body: JSON.stringify({ - refreshToken, - }), }) } finally { clearSession() diff --git a/web/lib/api/client.ts b/web/lib/api/client.ts index 3f98a2e..56a4daf 100644 --- a/web/lib/api/client.ts +++ b/web/lib/api/client.ts @@ -1,4 +1,4 @@ -import { clearSession, readSession, writeSession, type AuthSession } from "@/lib/auth" +import { expireSession, readSession } from "@/lib/auth" const API_BASE_URL = process.env.NEXT_PUBLIC_API_BASE_URL?.trim() || "" @@ -12,7 +12,6 @@ type JsonResult = { type RequestOptions = RequestInit & { skipAuth?: boolean - retryOnAuthError?: boolean baseUrl?: string onResponse?: (response: Response) => void } @@ -20,6 +19,9 @@ type RequestOptions = RequestInit & { async function parseResult(response: Response) { const payload = (await response.json()) as JsonResult if (!response.ok || !payload.success) { + if (payload.errorCode === 3000 || payload.errorCode === 3002) { + expireSession() + } const error = new Error(payload.message || "请求失败") ;(error as Error & { errorCode?: number }).errorCode = payload.errorCode throw error @@ -27,40 +29,11 @@ async function parseResult(response: Response) { return payload.data } -async function refreshAccessToken() { - const session = readSession() - if (!session?.refreshToken) { - clearSession() - return null - } - - const data = await request( - "/api/auth/refresh_token", - { - method: "POST", - body: JSON.stringify({ refreshToken: session.refreshToken }), - skipAuth: true, - headers: { - "Content-Type": "application/json", - }, - }, - false - ) - const merged = { - ...data, - refreshToken: data.refreshToken || session.refreshToken, - } - writeSession(merged) - return merged -} - export async function request( path: string, - options: RequestOptions = {}, - retryOnAuthError = true + options: RequestOptions = {} ): Promise { const { headers, skipAuth, baseUrl, onResponse, ...rest } = options - delete (rest as RequestOptions).retryOnAuthError delete (rest as RequestOptions).baseUrl delete (rest as RequestOptions).onResponse const session = readSession() @@ -85,26 +58,5 @@ export async function request( }) onResponse?.(response) - try { - return await parseResult(response) - } catch (error) { - const errorCode = (error as Error & { errorCode?: number }).errorCode - if ( - !skipAuth && - retryOnAuthError && - (errorCode === 3000 || errorCode === 3002) && - session?.refreshToken - ) { - const refreshed = await refreshAccessToken() - if (!refreshed) { - throw error - } - return request(path, options, false) - } - - if (errorCode === 3000 || errorCode === 3002) { - clearSession() - } - throw error - } + return parseResult(response) } diff --git a/web/lib/auth.ts b/web/lib/auth.ts index b1a88db..e59d8cf 100644 --- a/web/lib/auth.ts +++ b/web/lib/auth.ts @@ -9,7 +9,6 @@ export type AuthUser = { export type AuthSession = { accessToken: string - refreshToken: string expiresAt?: string user: AuthUser permissions: string[] @@ -17,6 +16,7 @@ export type AuthSession = { } const SESSION_STORAGE_KEY = "cs-ai-agent-session" +export const AUTH_SESSION_EXPIRED_EVENT = "cs-ai-agent-auth-expired" function hasWindow() { return typeof window !== "undefined" @@ -53,3 +53,11 @@ export function clearSession() { } window.localStorage.removeItem(SESSION_STORAGE_KEY) } + +export function expireSession() { + if (!hasWindow()) { + return + } + clearSession() + window.dispatchEvent(new Event(AUTH_SESSION_EXPIRED_EVENT)) +}