From c91423a19bd41540187f72e9c22b7186134edcd7 Mon Sep 17 00:00:00 2001 From: yml2213 Date: Tue, 18 Aug 2026 20:02:21 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E5=86=BB=E7=BB=93=E7=94=A8?= =?UTF-8?q?=E6=88=B7=E5=BC=BA=E5=88=B6=E4=B8=8B=E7=BA=BF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/middleware/auth.go | 15 +++- backend/internal/middleware/auth_test.go | 36 ++++++++++ backend/internal/model/user.go | 1 + .../internal/modules/adminuser/repository.go | 33 ++++++--- .../modules/adminuser/repository_test.go | 32 +++++++++ backend/internal/modules/adminuser/service.go | 7 ++ backend/internal/modules/auth/handler.go | 6 +- backend/internal/modules/auth/repository.go | 20 ++++++ backend/internal/modules/auth/service.go | 25 +++++-- .../modules/auth/token_version_test.go | 71 +++++++++++++++++++ backend/internal/modules/chathub/hub.go | 16 +++++ backend/internal/modules/chathub/hub_test.go | 25 +++++++ backend/internal/router/router.go | 10 ++- .../migrations/000052_user_token_version.sql | 9 +++ .../features/chats/composables/useChatSSE.ts | 8 ++- frontend/src/shared/api/client.ts | 2 +- 16 files changed, 297 insertions(+), 19 deletions(-) create mode 100644 backend/internal/middleware/auth_test.go create mode 100644 backend/internal/modules/auth/token_version_test.go create mode 100644 backend/internal/modules/chathub/hub_test.go create mode 100644 backend/migrations/000052_user_token_version.sql diff --git a/backend/internal/middleware/auth.go b/backend/internal/middleware/auth.go index 1ffa72f..9549ed8 100644 --- a/backend/internal/middleware/auth.go +++ b/backend/internal/middleware/auth.go @@ -34,6 +34,8 @@ type AdminTokenContext struct { 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 { header := c.GetHeader("Authorization") tokenText := strings.TrimSpace(strings.TrimPrefix(header, "Bearer ")) @@ -50,7 +52,7 @@ func extractToken(c *gin.Context) string { 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) { tokenText := extractToken(c) if tokenText == "" { @@ -65,6 +67,17 @@ func Auth(jwtManager *auth.JWTManager) gin.HandlerFunc { c.Abort() 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(ContextPhone, claims.Phone) diff --git a/backend/internal/middleware/auth_test.go b/backend/internal/middleware/auth_test.go new file mode 100644 index 0000000..e236f06 --- /dev/null +++ b/backend/internal/middleware/auth_test.go @@ -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) + } +} diff --git a/backend/internal/model/user.go b/backend/internal/model/user.go index 96d9407..6610291 100644 --- a/backend/internal/model/user.go +++ b/backend/internal/model/user.go @@ -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"` 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"` + TokenVersion int64 `gorm:"not null;default:0" json:"-"` LastLoginAt *time.Time `json:"last_login_at"` CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` diff --git a/backend/internal/modules/adminuser/repository.go b/backend/internal/modules/adminuser/repository.go index 5ced310..51970bd 100644 --- a/backend/internal/modules/adminuser/repository.go +++ b/backend/internal/modules/adminuser/repository.go @@ -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) } diff --git a/backend/internal/modules/adminuser/repository_test.go b/backend/internal/modules/adminuser/repository_test.go index 5359fcf..0d4cadf 100644 --- a/backend/internal/modules/adminuser/repository_test.go +++ b/backend/internal/modules/adminuser/repository_test.go @@ -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{}) diff --git a/backend/internal/modules/adminuser/service.go b/backend/internal/modules/adminuser/service.go index e7a4ed4..79718b8 100644 --- a/backend/internal/modules/adminuser/service.go +++ b/backend/internal/modules/adminuser/service.go @@ -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 diff --git a/backend/internal/modules/auth/handler.go b/backend/internal/modules/auth/handler.go index 2c2b908..bdc7ed5 100644 --- a/backend/internal/modules/auth/handler.go +++ b/backend/internal/modules/auth/handler.go @@ -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 } diff --git a/backend/internal/modules/auth/repository.go b/backend/internal/modules/auth/repository.go index f0d8721..2027f89 100644 --- a/backend/internal/modules/auth/repository.go +++ b/backend/internal/modules/auth/repository.go @@ -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, diff --git a/backend/internal/modules/auth/service.go b/backend/internal/modules/auth/service.go index 0d362c0..6a37b4c 100644 --- a/backend/internal/modules/auth/service.go +++ b/backend/internal/modules/auth/service.go @@ -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 } diff --git a/backend/internal/modules/auth/token_version_test.go b/backend/internal/modules/auth/token_version_test.go new file mode 100644 index 0000000..0dbfdcf --- /dev/null +++ b/backend/internal/modules/auth/token_version_test.go @@ -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) + } +} diff --git a/backend/internal/modules/chathub/hub.go b/backend/internal/modules/chathub/hub.go index 22f1d11..a99e53b 100644 --- a/backend/internal/modules/chathub/hub.go +++ b/backend/internal/modules/chathub/hub.go @@ -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 { diff --git a/backend/internal/modules/chathub/hub_test.go b/backend/internal/modules/chathub/hub_test.go new file mode 100644 index 0000000..9d701e8 --- /dev/null +++ b/backend/internal/modules/chathub/hub_test.go @@ -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) + } +} diff --git a/backend/internal/router/router.go b/backend/internal/router/router.go index 99c74a7..a6914a5 100644 --- a/backend/internal/router/router.go +++ b/backend/internal/router/router.go @@ -173,6 +173,7 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine { var chatHub *chathub.Hub if deps.DB != nil { chatHub = chathub.NewHub(deps.DB) + adminUserService.SetSessionRevoker(chatHub.DisconnectUser) } var chatRepo *chat.Repository if deps.DB != nil { @@ -369,7 +370,14 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine { } announcementService := announcement.NewService(announcementRepo) 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 if adminAuthRepo != nil { validateAdminToken = func(ctx context.Context, adminID uint64, tokenVersion int64) (middleware.AdminTokenContext, error) { diff --git a/backend/migrations/000052_user_token_version.sql b/backend/migrations/000052_user_token_version.sql new file mode 100644 index 0000000..3ce5e9e --- /dev/null +++ b/backend/migrations/000052_user_token_version.sql @@ -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; diff --git a/frontend/src/features/chats/composables/useChatSSE.ts b/frontend/src/features/chats/composables/useChatSSE.ts index c8922b0..6c800ac 100644 --- a/frontend/src/features/chats/composables/useChatSSE.ts +++ b/frontend/src/features/chats/composables/useChatSSE.ts @@ -1,5 +1,6 @@ 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' export interface SSEMessage { @@ -93,8 +94,11 @@ export function useChatSSE(scope: AuthScope, endpoint: string) { if (!stopped) { reconnectTimer = setTimeout(connect, reconnectDelay) } - } catch { + } catch (error) { closeSource() + if (axios.isAxiosError(error) && error.response?.status === 401) { + redirectToLogin(scope) + } } finally { refreshing = false } diff --git a/frontend/src/shared/api/client.ts b/frontend/src/shared/api/client.ts index 1002ce2..a7703d9 100644 --- a/frontend/src/shared/api/client.ts +++ b/frontend/src/shared/api/client.ts @@ -174,7 +174,7 @@ function getRequestScope(url = ''): AuthScope { return url.startsWith('/admin') ? 'admin' : 'user' } -function redirectToLogin(scope: AuthScope) { +export function redirectToLogin(scope: AuthScope) { clearAuthStorage(scope) const currentPath = window.location.pathname + window.location.search const loginPath = getLoginPath(scope, currentPath)