修复冻结用户强制下线

This commit is contained in:
yml2213
2026-08-18 20:02:21 +08:00
parent 0bc6cf9f46
commit c91423a19b
16 changed files with 297 additions and 19 deletions
+14 -1
View File
@@ -34,6 +34,8 @@ type AdminTokenContext struct {
type AdminTokenValidatorFunc func(ctx context.Context, adminID uint64, tokenVersion int64) (AdminTokenContext, error) type AdminTokenValidatorFunc func(ctx context.Context, adminID uint64, tokenVersion int64) (AdminTokenContext, error)
type UserTokenValidatorFunc func(ctx context.Context, userID uint64, tokenVersion int64) error
func extractBearerToken(c *gin.Context) string { func extractBearerToken(c *gin.Context) string {
header := c.GetHeader("Authorization") header := c.GetHeader("Authorization")
tokenText := strings.TrimSpace(strings.TrimPrefix(header, "Bearer ")) tokenText := strings.TrimSpace(strings.TrimPrefix(header, "Bearer "))
@@ -50,7 +52,7 @@ func extractToken(c *gin.Context) string {
return c.Query("token") return c.Query("token")
} }
func Auth(jwtManager *auth.JWTManager) gin.HandlerFunc { func Auth(jwtManager *auth.JWTManager, validate UserTokenValidatorFunc) gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
tokenText := extractToken(c) tokenText := extractToken(c)
if tokenText == "" { if tokenText == "" {
@@ -65,6 +67,17 @@ func Auth(jwtManager *auth.JWTManager) gin.HandlerFunc {
c.Abort() c.Abort()
return return
} }
if validate != nil {
if err := validate(c.Request.Context(), claims.UserID, claims.TokenVersion); err != nil {
if errors.Is(err, auth.ErrDependencyUnavailable) {
response.ServiceUnavailable(c, "用户认证服务暂时不可用")
} else {
response.Unauthorized(c, "登录状态已失效,请重新登录")
}
c.Abort()
return
}
}
c.Set(ContextUserID, claims.UserID) c.Set(ContextUserID, claims.UserID)
c.Set(ContextPhone, claims.Phone) c.Set(ContextPhone, claims.Phone)
+36
View File
@@ -0,0 +1,36 @@
package middleware
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"hfb_sys/backend/internal/modules/auth"
"github.com/gin-gonic/gin"
)
func TestAuthRejectsRevokedUserToken(t *testing.T) {
gin.SetMode(gin.TestMode)
manager := auth.NewJWTManager("test-secret")
pair, err := manager.GenerateSubjectPairWithVersion(7, "13900000007", "user", 0)
if err != nil {
t.Fatalf("生成令牌失败:%v", err)
}
engine := gin.New()
engine.GET("/protected", Auth(manager, func(_ context.Context, _ uint64, _ int64) error {
return auth.ErrTokenVersionMismatch
}), func(c *gin.Context) {
c.Status(http.StatusNoContent)
})
request := httptest.NewRequest(http.MethodGet, "/protected", nil)
request.Header.Set("Authorization", "Bearer "+pair.AccessToken)
response := httptest.NewRecorder()
engine.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized {
t.Fatalf("响应状态 = %d, want %d", response.Code, http.StatusUnauthorized)
}
}
+1
View File
@@ -15,6 +15,7 @@ type User struct {
RenterGrowthPoints int64 `gorm:"not null;default:0;index:idx_users_renter_growth_level,priority:2" json:"renter_growth_points"` RenterGrowthPoints int64 `gorm:"not null;default:0;index:idx_users_renter_growth_level,priority:2" json:"renter_growth_points"`
RenterGrowthLevel string `gorm:"size:32;not null;default:'normal';index:idx_users_renter_growth_level,priority:1" json:"renter_growth_level"` RenterGrowthLevel string `gorm:"size:32;not null;default:'normal';index:idx_users_renter_growth_level,priority:1" json:"renter_growth_level"`
Status string `gorm:"size:32;not null;default:'active'" json:"status"` Status string `gorm:"size:32;not null;default:'active'" json:"status"`
TokenVersion int64 `gorm:"not null;default:0" json:"-"`
LastLoginAt *time.Time `json:"last_login_at"` LastLoginAt *time.Time `json:"last_login_at"`
CreatedAt time.Time `json:"created_at"` CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"` UpdatedAt time.Time `json:"updated_at"`
@@ -19,8 +19,9 @@ import (
) )
type Repository struct { type Repository struct {
db *gorm.DB db *gorm.DB
encryptor crypto.Encryptor encryptor crypto.Encryptor
statusChangeNotifier func(userID uint64)
} }
type AuditMeta = auditlog.Meta type AuditMeta = auditlog.Meta
@@ -33,6 +34,11 @@ func NewRepository(db *gorm.DB, encryptors ...crypto.Encryptor) *Repository {
return &Repository{db: db, encryptor: encryptor} return &Repository{db: db, encryptor: encryptor}
} }
// SetStatusChangeNotifier 设置状态变更后的会话撤销通知。
func (r *Repository) SetStatusChangeNotifier(notifier func(userID uint64)) {
r.statusChangeNotifier = notifier
}
func (r *Repository) List(ctx context.Context, page, pageSize int, query ListQuery) (*PaginatedResult, error) { func (r *Repository) List(ctx context.Context, page, pageSize int, query ListQuery) (*PaginatedResult, error) {
growthConfig, err := rentergrowth.ConfigForTx(r.db.WithContext(ctx)) growthConfig, err := rentergrowth.ConfigForTx(r.db.WithContext(ctx))
if err != nil { if err != nil {
@@ -207,6 +213,7 @@ func (r *Repository) AdjustGrowthPoints(ctx context.Context, adminID uint64, use
} }
func (r *Repository) updateStatus(ctx context.Context, adminID uint64, userID uint64, status string, riskStatus string, action string, reason string, meta AuditMeta) (*UserDTO, error) { func (r *Repository) updateStatus(ctx context.Context, adminID uint64, userID uint64, status string, riskStatus string, action string, reason string, meta AuditMeta) (*UserDTO, error) {
statusChanged := false
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var user model.User var user model.User
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&user, userID).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&user, userID).Error; err != nil {
@@ -214,23 +221,33 @@ func (r *Repository) updateStatus(ctx context.Context, adminID uint64, userID ui
} }
beforeStatus := user.Status beforeStatus := user.Status
beforeRisk := user.RiskStatus beforeRisk := user.RiskStatus
beforeTokenVersion := user.TokenVersion
user.Status = status user.Status = status
user.RiskStatus = riskStatus user.RiskStatus = riskStatus
if beforeStatus != status {
user.TokenVersion++
statusChanged = true
}
if err := tx.Save(&user).Error; err != nil { if err := tx.Save(&user).Error; err != nil {
return err return err
} }
return appendAuditLog(tx, adminID, action, user.ID, meta, map[string]any{ return appendAuditLog(tx, adminID, action, user.ID, meta, map[string]any{
"user_id": user.ID, "user_id": user.ID,
"reason": reason, "reason": reason,
"before_status": beforeStatus, "before_status": beforeStatus,
"after_status": status, "after_status": status,
"before_risk_status": beforeRisk, "before_risk_status": beforeRisk,
"after_risk_status": riskStatus, "after_risk_status": riskStatus,
"before_token_version": beforeTokenVersion,
"after_token_version": user.TokenVersion,
}) })
}) })
if err != nil { if err != nil {
return nil, err return nil, err
} }
if statusChanged && r.statusChangeNotifier != nil {
r.statusChangeNotifier(userID)
}
return r.Find(ctx, userID) return r.Find(ctx, userID)
} }
@@ -293,6 +293,38 @@ func TestAdjustWalletInsufficientBalance(t *testing.T) {
} }
} }
func TestFreezeAndUnfreezeBumpUserTokenVersion(t *testing.T) {
db := setupAdminUserTestDB(t)
user := model.User{Phone: "13900000006", Status: "active", TokenVersion: 0}
if err := db.Create(&user).Error; err != nil {
t.Fatalf("创建用户失败:%v", err)
}
repo := NewRepository(db)
got, err := repo.Freeze(t.Context(), 66, user.ID, FreezeRequest{Reason: "风险处置"}, AuditMeta{})
if err != nil {
t.Fatalf("冻结用户失败:%v", err)
}
var saved model.User
if err := db.First(&saved, user.ID).Error; err != nil {
t.Fatalf("查询冻结用户失败:%v", err)
}
if got.Status != "frozen" || saved.TokenVersion != 1 {
t.Fatalf("冻结后状态/版本 = %s/%d, want frozen/1", got.Status, saved.TokenVersion)
}
got, err = repo.Unfreeze(t.Context(), 66, user.ID, AuditMeta{})
if err != nil {
t.Fatalf("解冻用户失败:%v", err)
}
if err := db.First(&saved, user.ID).Error; err != nil {
t.Fatalf("查询解冻用户失败:%v", err)
}
if got.Status != "active" || saved.TokenVersion != 2 {
t.Fatalf("解冻后状态/版本 = %s/%d, want active/2", got.Status, saved.TokenVersion)
}
}
func setupAdminUserTestDB(t *testing.T) *gorm.DB { func setupAdminUserTestDB(t *testing.T) *gorm.DB {
t.Helper() t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
@@ -37,6 +37,13 @@ func NewService(repo *Repository) *Service {
return &Service{repo: repo} return &Service{repo: repo}
} }
// SetSessionRevoker 设置用户会话撤销回调,例如关闭用户的实时连接。
func (s *Service) SetSessionRevoker(revoker func(userID uint64)) {
if s.repo != nil {
s.repo.SetStatusChangeNotifier(revoker)
}
}
func (s *Service) List(ctx context.Context, page, pageSize int, query ListQuery) (*PaginatedResult, error) { func (s *Service) List(ctx context.Context, page, pageSize int, query ListQuery) (*PaginatedResult, error) {
if s.repo == nil { if s.repo == nil {
return nil, ErrDependencyUnavailable return nil, ErrDependencyUnavailable
+5 -1
View File
@@ -131,8 +131,12 @@ func (h *Handler) Refresh(c *gin.Context) {
response.BadRequest(c, "refresh_token 不能为空") response.BadRequest(c, "refresh_token 不能为空")
return return
} }
tokens, err := h.service.RefreshToken(req.RefreshToken) tokens, err := h.service.RefreshToken(c.Request.Context(), req.RefreshToken)
if err != nil { if err != nil {
if errors.Is(err, ErrDependencyUnavailable) {
response.ServiceUnavailable(c, "用户认证服务暂时不可用")
return
}
response.Unauthorized(c, "刷新令牌无效或已过期") response.Unauthorized(c, "刷新令牌无效或已过期")
return return
} }
@@ -27,6 +27,26 @@ func (r *UserRepository) FindByID(ctx context.Context, id uint64) (*model.User,
return &user, nil return &user, nil
} }
// FindActiveForToken 校验用户仍可用且令牌版本未被撤销。
func (r *UserRepository) FindActiveForToken(ctx context.Context, id uint64, tokenVersion int64) (*model.User, error) {
if r == nil || r.db == nil {
return nil, ErrDependencyUnavailable
}
var user model.User
if err := r.db.WithContext(ctx).
Select("id, phone, status, token_version").
First(&user, id).Error; err != nil {
return nil, err
}
if user.Status != "active" {
return nil, ErrUserDisabled
}
if user.TokenVersion != tokenVersion {
return nil, ErrTokenVersionMismatch
}
return &user, nil
}
func (r *UserRepository) UpdateProfile(ctx context.Context, id uint64, nickname string, avatarURL string) (*model.User, error) { func (r *UserRepository) UpdateProfile(ctx context.Context, id uint64, nickname string, avatarURL string) (*model.User, error) {
if err := r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(map[string]any{ if err := r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(map[string]any{
"nickname": nickname, "nickname": nickname,
+20 -5
View File
@@ -33,6 +33,7 @@ var (
ErrPasswordTooWeak = errors.New("password too weak") ErrPasswordTooWeak = errors.New("password too weak")
ErrLoginLocked = errors.New("login locked") ErrLoginLocked = errors.New("login locked")
ErrUserAlreadyExists = errors.New("user already exists") ErrUserAlreadyExists = errors.New("user already exists")
ErrTokenVersionMismatch = errors.New("user token version mismatch")
) )
const ( const (
@@ -180,7 +181,7 @@ func (s *Service) LoginWithSMS(ctx context.Context, phone string, code string) (
return LoginResult{}, ErrUserDisabled return LoginResult{}, ErrUserDisabled
} }
tokens, err := s.jwt.GeneratePair(user.ID, user.Phone) tokens, err := s.generateUserTokenPair(user)
if err != nil { if err != nil {
return LoginResult{}, err return LoginResult{}, err
} }
@@ -189,12 +190,26 @@ func (s *Service) LoginWithSMS(ctx context.Context, phone string, code string) (
return LoginResult{User: user, Tokens: tokens}, nil return LoginResult{User: user, Tokens: tokens}, nil
} }
func (s *Service) RefreshToken(refreshToken string) (TokenPair, error) { func (s *Service) RefreshToken(ctx context.Context, refreshToken string) (TokenPair, error) {
if s.jwt == nil || s.users == nil {
return TokenPair{}, ErrDependencyUnavailable
}
claims, err := s.jwt.ParseSubject(refreshToken, tokenTypeRefresh, "user") claims, err := s.jwt.ParseSubject(refreshToken, tokenTypeRefresh, "user")
if err != nil { if err != nil {
return TokenPair{}, err return TokenPair{}, err
} }
return s.jwt.GeneratePair(claims.UserID, claims.Phone) user, err := s.users.FindActiveForToken(ctx, claims.UserID, claims.TokenVersion)
if err != nil {
return TokenPair{}, err
}
return s.generateUserTokenPair(user)
}
func (s *Service) generateUserTokenPair(user *model.User) (TokenPair, error) {
if s.jwt == nil {
return TokenPair{}, ErrDependencyUnavailable
}
return s.jwt.GenerateSubjectPairWithVersion(user.ID, user.Phone, "user", user.TokenVersion)
} }
func codeKey(phone string) string { func codeKey(phone string) string {
@@ -264,7 +279,7 @@ func (s *Service) LoginWithPassword(ctx context.Context, phone, password, client
_ = clearLoginFailure(ctx, s.redis, phone, clientIP) _ = clearLoginFailure(ctx, s.redis, phone, clientIP)
tokens, err := s.jwt.GeneratePair(user.ID, user.Phone) tokens, err := s.generateUserTokenPair(user)
if err != nil { if err != nil {
return LoginResult{}, err return LoginResult{}, err
} }
@@ -369,7 +384,7 @@ func (s *Service) RegisterWithPassword(ctx context.Context, phone, code, passwor
return LoginResult{}, ErrUserDisabled return LoginResult{}, ErrUserDisabled
} }
tokens, err := s.jwt.GeneratePair(user.ID, user.Phone) tokens, err := s.generateUserTokenPair(user)
if err != nil { if err != nil {
return LoginResult{}, err return LoginResult{}, err
} }
@@ -0,0 +1,71 @@
package auth
import (
"errors"
"testing"
"hfb_sys/backend/internal/model"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
func TestUserTokenVersionRevokesAccessAndRefreshTokens(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil {
t.Fatalf("打开测试数据库失败:%v", err)
}
if err := db.AutoMigrate(&model.User{}); err != nil {
t.Fatalf("数据库迁移失败:%v", err)
}
user := model.User{Phone: "13900000001", Status: "active", TokenVersion: 0}
if err := db.Create(&user).Error; err != nil {
t.Fatalf("创建用户失败:%v", err)
}
repo := NewUserRepository(db)
manager := NewJWTManager("test-secret")
pair, err := manager.GenerateSubjectPairWithVersion(user.ID, user.Phone, "user", user.TokenVersion)
if err != nil {
t.Fatalf("生成令牌失败:%v", err)
}
_, err = repo.FindActiveForToken(t.Context(), user.ID, 0)
if err != nil {
t.Fatalf("有效令牌校验失败:%v", err)
}
if err := db.Model(&model.User{}).Where("id = ?", user.ID).Updates(map[string]any{
"status": "frozen",
"token_version": 1,
}).Error; err != nil {
t.Fatalf("冻结用户失败:%v", err)
}
_, err = repo.FindActiveForToken(t.Context(), user.ID, 0)
if !errors.Is(err, ErrUserDisabled) {
t.Fatalf("冻结用户校验错误 = %v, want %v", err, ErrUserDisabled)
}
service := NewService(repo, nil, manager, nil, nil)
_, err = service.RefreshToken(t.Context(), pair.RefreshToken)
if !errors.Is(err, ErrUserDisabled) {
t.Fatalf("冻结用户刷新错误 = %v, want %v", err, ErrUserDisabled)
}
// 解冻后版本仍然不同,冻结前签发的 token 不能恢复使用。
if err := db.Model(&model.User{}).Where("id = ?", user.ID).Updates(map[string]any{
"status": "active",
"token_version": 2,
}).Error; err != nil {
t.Fatalf("解冻用户失败:%v", err)
}
_, err = repo.FindActiveForToken(t.Context(), user.ID, 0)
if !errors.Is(err, ErrTokenVersionMismatch) {
t.Fatalf("旧令牌版本校验错误 = %v, want %v", err, ErrTokenVersionMismatch)
}
_, err = repo.FindActiveForToken(t.Context(), user.ID, 2)
if err != nil {
t.Fatalf("新令牌版本校验失败:%v", err)
}
}
+16
View File
@@ -86,6 +86,22 @@ func (h *Hub) Unsubscribe(pType string, pID uint64, ch <-chan *ChatEvent) {
h.mu.Unlock() h.mu.Unlock()
} }
// DisconnectUser 关闭指定用户当前进程内的全部实时连接。
func (h *Hub) DisconnectUser(userID uint64) {
h.disconnect("user", userID)
}
func (h *Hub) disconnect(pType string, pID uint64) {
key := principal{Type: pType, ID: pID}
h.mu.Lock()
clients := h.clients[key]
delete(h.clients, key)
for ch := range clients {
close(ch)
}
h.mu.Unlock()
}
// NotifyConversation 查询会话参与者,向所有在线参与者推送事件。 // NotifyConversation 查询会话参与者,向所有在线参与者推送事件。
func (h *Hub) NotifyConversation(conversationID uint64, event *ChatEvent) { func (h *Hub) NotifyConversation(conversationID uint64, event *ChatEvent) {
if h.db == nil { if h.db == nil {
@@ -0,0 +1,25 @@
package chathub
import "testing"
func TestDisconnectUserClosesAllConnections(t *testing.T) {
hub := NewHub(nil)
first := hub.Subscribe("user", 7)
second := hub.Subscribe("user", 7)
hub.Subscribe("user", 8)
hub.DisconnectUser(7)
for name, ch := range map[string]<-chan *ChatEvent{"first": first, "second": second} {
select {
case _, ok := <-ch:
if ok {
t.Fatalf("%s connection was not closed", name)
}
default:
t.Fatalf("%s connection did not close immediately", name)
}
}
if got := hub.OnlineCount(); got != 1 {
t.Fatalf("online connection count = %d, want 1", got)
}
}
+9 -1
View File
@@ -173,6 +173,7 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine {
var chatHub *chathub.Hub var chatHub *chathub.Hub
if deps.DB != nil { if deps.DB != nil {
chatHub = chathub.NewHub(deps.DB) chatHub = chathub.NewHub(deps.DB)
adminUserService.SetSessionRevoker(chatHub.DisconnectUser)
} }
var chatRepo *chat.Repository var chatRepo *chat.Repository
if deps.DB != nil { if deps.DB != nil {
@@ -369,7 +370,14 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine {
} }
announcementService := announcement.NewService(announcementRepo) announcementService := announcement.NewService(announcementRepo)
announcementHandler := announcement.NewHandler(announcementService) announcementHandler := announcement.NewHandler(announcementService)
requireAuth := middleware.Auth(jwtManager) validateUserToken := func(ctx context.Context, userID uint64, tokenVersion int64) error {
if userRepo == nil {
return auth.ErrDependencyUnavailable
}
_, err := userRepo.FindActiveForToken(ctx, userID, tokenVersion)
return err
}
requireAuth := middleware.Auth(jwtManager, validateUserToken)
var validateAdminToken middleware.AdminTokenValidatorFunc var validateAdminToken middleware.AdminTokenValidatorFunc
if adminAuthRepo != nil { if adminAuthRepo != nil {
validateAdminToken = func(ctx context.Context, adminID uint64, tokenVersion int64) (middleware.AdminTokenContext, error) { validateAdminToken = func(ctx context.Context, adminID uint64, tokenVersion int64) (middleware.AdminTokenContext, error) {
@@ -0,0 +1,9 @@
-- +goose Up
ALTER TABLE users
ADD COLUMN token_version BIGINT NOT NULL DEFAULT 0 COMMENT '用户会话版本,递增后旧令牌失效';
-- +goose Down
ALTER TABLE users
DROP COLUMN token_version;
@@ -1,5 +1,6 @@
import { onBeforeUnmount, ref, type Ref } from 'vue' import { onBeforeUnmount, ref, type Ref } from 'vue'
import { refreshAccessToken } from '@/shared/api/client' import axios from 'axios'
import { redirectToLogin, refreshAccessToken } from '@/shared/api/client'
import { getAccessToken, type AuthScope } from '@/shared/utils/authStorage' import { getAccessToken, type AuthScope } from '@/shared/utils/authStorage'
export interface SSEMessage { export interface SSEMessage {
@@ -93,8 +94,11 @@ export function useChatSSE(scope: AuthScope, endpoint: string) {
if (!stopped) { if (!stopped) {
reconnectTimer = setTimeout(connect, reconnectDelay) reconnectTimer = setTimeout(connect, reconnectDelay)
} }
} catch { } catch (error) {
closeSource() closeSource()
if (axios.isAxiosError(error) && error.response?.status === 401) {
redirectToLogin(scope)
}
} finally { } finally {
refreshing = false refreshing = false
} }
+1 -1
View File
@@ -174,7 +174,7 @@ function getRequestScope(url = ''): AuthScope {
return url.startsWith('/admin') ? 'admin' : 'user' return url.startsWith('/admin') ? 'admin' : 'user'
} }
function redirectToLogin(scope: AuthScope) { export function redirectToLogin(scope: AuthScope) {
clearAuthStorage(scope) clearAuthStorage(scope)
const currentPath = window.location.pathname + window.location.search const currentPath = window.location.pathname + window.location.search
const loginPath = getLoginPath(scope, currentPath) const loginPath = getLoginPath(scope, currentPath)