refactor(auth): simplify admin session token flow

This commit is contained in:
mlogclub
2026-04-30 17:42:24 +08:00
parent e763a75f78
commit fbafb13a8a
12 changed files with 411 additions and 179 deletions
+44 -99
View File
@@ -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 {
+325 -1
View File
@@ -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(&current, "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
}