修复冻结用户强制下线
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user