修复冻结用户强制下线
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"`
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ 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,8 +221,13 @@ 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
|
||||||
}
|
}
|
||||||
@@ -226,11 +238,16 @@ func (r *Repository) updateStatus(ctx context.Context, adminID uint64, userID ui
|
|||||||
"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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user