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 {