修复冻结用户强制下线

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
@@ -19,8 +19,9 @@ import (
)
type Repository struct {
db *gorm.DB
encryptor crypto.Encryptor
db *gorm.DB
encryptor crypto.Encryptor
statusChangeNotifier func(userID uint64)
}
type AuditMeta = auditlog.Meta
@@ -33,6 +34,11 @@ func NewRepository(db *gorm.DB, encryptors ...crypto.Encryptor) *Repository {
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) {
growthConfig, err := rentergrowth.ConfigForTx(r.db.WithContext(ctx))
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) {
statusChanged := false
err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
var user model.User
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
beforeRisk := user.RiskStatus
beforeTokenVersion := user.TokenVersion
user.Status = status
user.RiskStatus = riskStatus
if beforeStatus != status {
user.TokenVersion++
statusChanged = true
}
if err := tx.Save(&user).Error; err != nil {
return err
}
return appendAuditLog(tx, adminID, action, user.ID, meta, map[string]any{
"user_id": user.ID,
"reason": reason,
"before_status": beforeStatus,
"after_status": status,
"before_risk_status": beforeRisk,
"after_risk_status": riskStatus,
"user_id": user.ID,
"reason": reason,
"before_status": beforeStatus,
"after_status": status,
"before_risk_status": beforeRisk,
"after_risk_status": riskStatus,
"before_token_version": beforeTokenVersion,
"after_token_version": user.TokenVersion,
})
})
if err != nil {
return nil, err
}
if statusChanged && r.statusChangeNotifier != nil {
r.statusChangeNotifier(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 {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
@@ -37,6 +37,13 @@ func NewService(repo *Repository) *Service {
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) {
if s.repo == nil {
return nil, ErrDependencyUnavailable
+5 -1
View File
@@ -131,8 +131,12 @@ func (h *Handler) Refresh(c *gin.Context) {
response.BadRequest(c, "refresh_token 不能为空")
return
}
tokens, err := h.service.RefreshToken(req.RefreshToken)
tokens, err := h.service.RefreshToken(c.Request.Context(), req.RefreshToken)
if err != nil {
if errors.Is(err, ErrDependencyUnavailable) {
response.ServiceUnavailable(c, "用户认证服务暂时不可用")
return
}
response.Unauthorized(c, "刷新令牌无效或已过期")
return
}
@@ -27,6 +27,26 @@ func (r *UserRepository) FindByID(ctx context.Context, id uint64) (*model.User,
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) {
if err := r.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", id).Updates(map[string]any{
"nickname": nickname,
+20 -5
View File
@@ -33,6 +33,7 @@ var (
ErrPasswordTooWeak = errors.New("password too weak")
ErrLoginLocked = errors.New("login locked")
ErrUserAlreadyExists = errors.New("user already exists")
ErrTokenVersionMismatch = errors.New("user token version mismatch")
)
const (
@@ -180,7 +181,7 @@ func (s *Service) LoginWithSMS(ctx context.Context, phone string, code string) (
return LoginResult{}, ErrUserDisabled
}
tokens, err := s.jwt.GeneratePair(user.ID, user.Phone)
tokens, err := s.generateUserTokenPair(user)
if err != nil {
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
}
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")
if err != nil {
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 {
@@ -264,7 +279,7 @@ func (s *Service) LoginWithPassword(ctx context.Context, phone, password, client
_ = clearLoginFailure(ctx, s.redis, phone, clientIP)
tokens, err := s.jwt.GeneratePair(user.ID, user.Phone)
tokens, err := s.generateUserTokenPair(user)
if err != nil {
return LoginResult{}, err
}
@@ -369,7 +384,7 @@ func (s *Service) RegisterWithPassword(ctx context.Context, phone, code, passwor
return LoginResult{}, ErrUserDisabled
}
tokens, err := s.jwt.GeneratePair(user.ID, user.Phone)
tokens, err := s.generateUserTokenPair(user)
if err != nil {
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()
}
// 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 查询会话参与者,向所有在线参与者推送事件。
func (h *Hub) NotifyConversation(conversationID uint64, event *ChatEvent) {
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)
}
}