feat(backend): 会话安全加固 + 授权实时 phase + 首次使用自动试用

会话安全(jti 轮换 / 重用检测 / 改密吊销 / 禁用即时下线 / 清理 / 失败登录落库):
- refresh token 轮换 jti + token-family 重用检测,旧 token 重放即吊销整条会话
- 改密码、停用用户即时吊销其全部活跃会话(revoked_by 审计)
- 中间件 session JOIN user 校验,禁用/删除用户带 token 请求返回 401 USER_DISABLED
- 新增 login_attempts 失败登录落库 + 会话保留期清理 goroutine

授权实时 phase + 心跳回带:
- LicenseGuard 改为按当前 DB 实时计算 phase(30s 每店缓存),续费/过期/被改 ~30s 内对写操作生效,无需重登
- /auth/ping 回带授权概况(ShopInfoView,与 /license/info 同构),客户端一次心跳即刷新横幅/门禁

首次使用自动试用 + code-review 修复:
- 门店首次登录/续期无有效授权时自动签发 30 天 trial(快路径无锁 Count,仅首用走 FOR UPDATE 事务)
- ShopInfo 区分「确无授权」与瞬时 DB 错误,避免误降级
- trial 签发后改为在事务提交后再失效 phase 缓存(修复早于提交的竞态)
- 存量无 sid token 续期纳入显式上限,legacy 会话不再游离于并发配额之外

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
wangjia
2026-06-19 07:34:04 +08:00
parent 2d84bda99a
commit e41085a878
23 changed files with 1248 additions and 74 deletions
+185 -20
View File
@@ -24,6 +24,10 @@ var (
ErrPlatformNotAllowed = errors.New("该平台不允许登录")
ErrTooManyAttempts = errors.New("登录失败次数过多,账号已临时锁定,请稍后再试")
ErrSessionRevoked = errors.New("session revoked")
// errRefreshReuse 内部哨兵:在续期事务内检测到 refresh token 重用,
// 用于让调用方在事务回滚后于事务外提交「吊销整条会话」。
errRefreshReuse = errors.New("refresh token reuse detected")
)
// DeviceInfo 登录请求携带的设备信息,用于会话记录与按平台限并发。
@@ -99,12 +103,14 @@ type TokenPair struct {
func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo) (*TokenPair, *model.User, error) {
limiterKey := shopCode + "|" + username
if loginLim.locked(limiterKey) {
s.recordLoginAttempt(shopCode, username, dev, false, "locked")
return nil, nil, ErrTooManyAttempts
}
var shop model.Shop
if err := s.db.Where("code = ?", shopCode).First(&shop).Error; err != nil {
loginLim.recordFailure(limiterKey)
s.recordLoginAttempt(shopCode, username, dev, false, "invalid_shop")
return nil, nil, ErrInvalidCredentials
}
@@ -112,15 +118,18 @@ func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo)
if err := s.db.Where("shop_id = ? AND username = ? AND deleted_at IS NULL", shop.ID, username).
First(&user).Error; err != nil {
loginLim.recordFailure(limiterKey)
s.recordLoginAttempt(shopCode, username, dev, false, "invalid_user")
return nil, nil, ErrInvalidCredentials
}
if !user.IsActive {
s.recordLoginAttempt(shopCode, username, dev, false, "inactive")
return nil, nil, ErrUserInactive
}
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil {
loginLim.recordFailure(limiterKey)
s.recordLoginAttempt(shopCode, username, dev, false, "bad_password")
return nil, nil, ErrInvalidCredentials
}
@@ -137,10 +146,12 @@ func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo)
pclass := model.PlatformClass(dev.Platform)
quota := s.effectiveQuota(&shop, pclass)
if quota <= 0 {
s.recordLoginAttempt(shopCode, username, dev, false, "platform_not_allowed")
return nil, nil, ErrPlatformNotAllowed
}
sid := uuid.New().String()
jti := uuid.New().String() // 初始 refresh token jti,后续每次续期轮换
now := time.Now()
if err := s.db.Transaction(func(tx *gorm.DB) error {
// 统计该 user 在该 class 的活跃会话;超额踢最旧
@@ -169,6 +180,7 @@ func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo)
PlatformClass: pclass,
IP: dev.IP,
UserAgent: dev.UserAgent,
RefreshJTI: jti,
LastSeenAt: now,
RefreshExpAt: now.Add(time.Duration(config.C.JWT.RefreshExpireH) * time.Hour),
}
@@ -183,13 +195,28 @@ func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo)
loginLim.reset(limiterKey)
user.LastLoginAt = &now
pair, err := s.issueTokens(user.ID, shop.ID, user.Role, sid)
pair, err := s.issueTokens(user.ID, shop.ID, user.Role, sid, jti)
if err != nil {
return nil, nil, err
}
return pair, &user, nil
}
// recordLoginAttempt 记录一次登录尝试(当前只用于失败审计),写库失败仅记日志不阻断登录。
func (s *AuthService) recordLoginAttempt(shopCode, username string, dev DeviceInfo, success bool, reason string) {
att := model.LoginAttempt{
ShopCode: shopCode,
Username: username,
IP: dev.IP,
UserAgent: dev.UserAgent,
Success: success,
Reason: reason,
}
if err := s.db.Create(&att).Error; err != nil {
log.Printf("[auth] record login attempt failed: %v", err)
}
}
// effectiveQuota 返回某店某平台类的有效并发配额:优先每店 session_policy 覆盖,否则全局默认。
func (s *AuthService) effectiveQuota(shop *model.Shop, pclass string) int {
def := map[string]int{
@@ -289,11 +316,11 @@ func (s *AuthService) ListSessions(shopID uint64, currentSID string) ([]SessionV
return views, nil
}
// ForceLogout 管理员强制下线本店某会话(按 id + shop_id 隔离)。
func (s *AuthService) ForceLogout(shopID, sessionID uint64) error {
// ForceLogout 管理员强制下线本店某会话(按 id + shop_id 隔离)。byUserID 记入 revoked_by 供审计。
func (s *AuthService) ForceLogout(shopID, sessionID, byUserID uint64) error {
res := s.db.Model(&model.UserSession{}).
Where("id = ? AND shop_id = ? AND revoked_at IS NULL", sessionID, shopID).
Updates(map[string]interface{}{"revoked_at": time.Now(), "revoked_reason": "admin"})
Updates(map[string]interface{}{"revoked_at": time.Now(), "revoked_reason": "admin", "revoked_by": byUserID})
if res.Error != nil {
return res.Error
}
@@ -303,6 +330,18 @@ func (s *AuthService) ForceLogout(shopID, sessionID uint64) error {
return nil
}
// RevokeUserSessions 吊销某用户在本店的全部活跃会话(改密/禁用时调用)。byUserID 为操作人(记入 revoked_by)。
// db 入参以便 handler 在不持有 AuthService 时也能复用同一语义。
func RevokeUserSessions(db *gorm.DB, shopID, userID, byUserID uint64, reason string) error {
updates := map[string]interface{}{"revoked_at": time.Now(), "revoked_reason": reason}
if byUserID != 0 {
updates["revoked_by"] = byUserID
}
return db.Model(&model.UserSession{}).
Where("shop_id = ? AND user_id = ? AND revoked_at IS NULL", shopID, userID).
Updates(updates).Error
}
// RegisterInput 注册新门店所需参数
type RegisterInput struct {
ShopName string `json:"shop_name" binding:"required"`
@@ -392,14 +431,6 @@ func (s *AuthService) RefreshTokens(refreshToken string) (*TokenPair, error) {
return nil, errors.New("invalid refresh token")
}
// 带 sid 的 token 必须对应未撤销会话(被踢/登出后无法续期)。
if claims.SID != "" {
var sess model.UserSession
if err := s.db.Where("sid = ?", claims.SID).First(&sess).Error; err != nil || sess.RevokedAt != nil {
return nil, ErrSessionRevoked
}
}
// 重新加载用户:不存在或已禁用 → 拒绝续期(修复漏洞)。
var user model.User
if err := s.db.Where("id = ? AND deleted_at IS NULL", claims.UserID).First(&user).Error; err != nil {
@@ -413,10 +444,109 @@ func (s *AuthService) RefreshTokens(refreshToken string) (*TokenPair, error) {
return nil, err
}
// 刷新会话存活时间(同 sid 续期)。
if claims.SID != "" {
s.db.Model(&model.UserSession{}).Where("sid = ?", claims.SID).
Update("last_seen_at", time.Now())
// 带 sid 的会话:在事务内 FOR UPDATE 锁行做「校验 → 重用检测 → 轮换」(check-then-act)。
// 无 sid 的存量 token:首刷即建立可吊销会话并签发带 sid 的新 token,使其转入受控状态
//(自此可被强制下线/踢人/改密吊销),不再永久游离于会话治理之外。
newJTI := uuid.New().String()
now := time.Now()
sid := claims.SID
if sid == "" {
// 存量无 sid token:复用「该用户唯一的 legacy 会话」,而非每次续期都新建一条。
// 事务 + FOR UPDATE 锁住既有 legacy 行、串行化并发的存量续期;找不到才创建。
// 这样重复/并发呈递同一存量 token 始终只对应一条可吊销会话(有界、纳入会话治理、
// 可被强制下线/改密吊销),不再每刷一条无限膨胀。正常客户端首刷后即拿到带 sid 的新
// token,自然转入受控分支,此 legacy 行随其过期由清理任务回收。
exp := now.Add(time.Duration(config.C.JWT.RefreshExpireH) * time.Hour)
if err := s.db.Transaction(func(tx *gorm.DB) error {
// 锁住该用户全部活跃 legacy 会话。legacy class 不在 effectiveQuota 的
// desktop/mobile/web 映射内(会返回 0),若不在此显式约束,legacy 会话便游离于
// 并发配额之外。这里把「每用户至多一条 legacy 会话」从 find-or-create 的隐式产物
// 提升为显式上限:复用最早的一条(轮换 jti),其余多余 legacy 会话一并吊销。
var sessions []model.UserSession
if err := tx.Set("gorm:query_option", "FOR UPDATE").
Where("shop_id = ? AND user_id = ? AND platform_class = ? AND revoked_at IS NULL",
user.ShopID, user.ID, "legacy").
Order("id ASC").Find(&sessions).Error; err != nil {
return err
}
if len(sessions) > 0 {
keep := sessions[0]
sid = keep.SID
if err := tx.Model(&model.UserSession{}).Where("id = ?", keep.ID).
Updates(map[string]interface{}{
"refresh_jti": newJTI,
"last_seen_at": now,
"refresh_exp_at": exp,
}).Error; err != nil {
return err
}
if len(sessions) > 1 {
extraIDs := make([]uint64, 0, len(sessions)-1)
for _, ex := range sessions[1:] {
extraIDs = append(extraIDs, ex.ID)
}
if err := tx.Model(&model.UserSession{}).Where("id IN ?", extraIDs).
Updates(map[string]interface{}{"revoked_at": now, "revoked_reason": "kicked"}).Error; err != nil {
return err
}
}
return nil
}
newSID := uuid.New().String()
sess := model.UserSession{
ShopID: user.ShopID,
UserID: user.ID,
SID: newSID,
Platform: "legacy",
PlatformClass: "legacy",
RefreshJTI: newJTI,
LastSeenAt: now,
RefreshExpAt: exp,
}
if err := tx.Create(&sess).Error; err != nil {
return err
}
sid = newSID
return nil
}); err != nil {
return nil, err
}
} else {
var reuseSessID uint64
txErr := s.db.Transaction(func(tx *gorm.DB) error {
var sess model.UserSession
if err := tx.Set("gorm:query_option", "FOR UPDATE").
Where("sid = ?", claims.SID).First(&sess).Error; err != nil {
return ErrSessionRevoked
}
if sess.RevokedAt != nil {
return ErrSessionRevoked
}
// 重用检测:会话已确立 jti,但呈递的 refresh token 携带的 jti 与当前值不符 →
// 说明这是已被轮换取代的旧 token 重放(盗用信号),吊销整条会话。
// 向后兼容:sess.RefreshJTI=="" 为存量会话,首刷直接采纳新 jti,不报重用。
// 注意:吊销动作放到事务外执行——此处一旦 return 非 nil 错误,事务会回滚,
// 在事务内写吊销会被一并回滚掉。
if sess.RefreshJTI != "" && claims.ID != sess.RefreshJTI {
reuseSessID = sess.ID
return errRefreshReuse
}
// 轮换 jti + 刷新存活时间。
return tx.Model(&model.UserSession{}).Where("id = ?", sess.ID).
Updates(map[string]interface{}{"refresh_jti": newJTI, "last_seen_at": now}).Error
})
if errors.Is(txErr, errRefreshReuse) {
// 事务外提交吊销(不随回滚丢失),吊销整条会话。写库失败必须记日志——
// 否则盗用信号被静默吞掉,被取代的旧 token 仍可继续续期,功能形同虚设。
if err := s.db.Model(&model.UserSession{}).Where("id = ?", reuseSessID).
Updates(map[string]interface{}{"revoked_at": now, "revoked_reason": "reuse"}).Error; err != nil {
log.Printf("[auth] revoke reused session %d failed: %v", reuseSessID, err)
}
return nil, ErrSessionRevoked
}
if txErr != nil {
return nil, txErr
}
}
// 自动续登(refresh)路径同样补发首次试用:老用户用本地 refresh token 自动登录、
@@ -424,7 +554,7 @@ func (s *AuthService) RefreshTokens(refreshToken string) (*TokenPair, error) {
// lic_exp 带上试用到期日。
s.ensureTrialOnFirstUse(user.ShopID)
return s.issueTokens(user.ID, user.ShopID, user.Role, claims.SID)
return s.issueTokens(user.ID, user.ShopID, user.Role, sid, newJTI)
}
// ensureTrialOnFirstUse 门店首次使用(尚无任何 is_active 授权)时自动签发 30 天 trial,
@@ -432,17 +562,51 @@ func (s *AuthService) RefreshTokens(refreshToken string) (*TokenPair, error) {
// 已有有效授权(含已过期但未锁定的 trial/付费)则跳过,不重复发放。
// 签发失败仅记日志、不阻断登录(保持可用,门店维持未激活)。
func (s *AuthService) ensureTrialOnFirstUse(shopID uint64) {
// 快路径:绝大多数登录/续期门店已有有效授权。先做一次无锁 Count 直接返回,
// 避免每次都开事务锁 shop 行——该函数在每次 Login/RefreshTokens 都被调用,
// 高频路径上的行锁会无谓串行化同店并发请求。
var count int64
if err := s.db.Model(&model.License{}).
Where("shop_id = ? AND is_active = 1", shopID).Count(&count).Error; err != nil {
log.Printf("[license] count licenses for shop %d failed: %v", shopID, err)
log.Printf("[license] auto-trial precheck for shop %d failed: %v", shopID, err)
return
}
if count > 0 {
return
}
if err := issueTrialLicense(s.db, shopID); err != nil {
// 慢路径(首次使用):事务内锁住门店行后**重新** Count,串行化同店并发的「首次试用」
// 判断:无既有 license 行可锁,故锁父级 shop 行,避免两个并发请求都读到 count==0
// 各发一条 trial(重复授权)。
issued := false
err := s.db.Transaction(func(tx *gorm.DB) error {
var shop model.Shop
if err := tx.Set("gorm:query_option", "FOR UPDATE").
Where("id = ?", shopID).First(&shop).Error; err != nil {
return err
}
var c int64
if err := tx.Model(&model.License{}).
Where("shop_id = ? AND is_active = 1", shopID).Count(&c).Error; err != nil {
return err
}
if c > 0 {
return nil
}
if err := issueTrialLicense(tx, shopID); err != nil {
return err
}
issued = true
return nil
})
if err != nil {
log.Printf("[license] auto-trial for shop %d skipped: %v", shopID, err)
return
}
// 事务提交后再失效 phase 缓存:此刻新签发的 trial 行对其它连接已可见,
// 不会在提交前的窗口里被旧 phase 重新填充(修复缓存失效早于提交的竞态)。
if issued {
middleware.InvalidateLicensePhase(shopID)
}
}
@@ -458,7 +622,7 @@ func (s *AuthService) checkLicenseNotLocked(shopID uint64) error {
return nil
}
func (s *AuthService) issueTokens(userID, shopID uint64, role, sid string) (*TokenPair, error) {
func (s *AuthService) issueTokens(userID, shopID uint64, role, sid, refreshJTI string) (*TokenPair, error) {
cfg := config.C.JWT
now := time.Now()
@@ -496,6 +660,7 @@ func (s *AuthService) issueTokens(userID, shopID uint64, role, sid string) (*Tok
SID: sid,
LicenseExpiresAt: licExpAt,
RegisteredClaims: jwt.RegisteredClaims{
ID: refreshJTI, // jtirefresh token 轮换与重用检测的依据
ExpiresAt: jwt.NewNumericDate(refreshExp),
IssuedAt: jwt.NewNumericDate(now),
},
+53 -2
View File
@@ -13,6 +13,7 @@ import (
"gorm.io/gorm"
"github.com/wangjia/jiu/backend/config"
"github.com/wangjia/jiu/backend/internal/middleware"
"github.com/wangjia/jiu/backend/internal/model"
"github.com/wangjia/jiu/backend/internal/util"
)
@@ -62,6 +63,9 @@ func (s *LicenseService) Activate(shopID uint64, licenseKey, deviceID, deviceNam
return nil, ErrLicenseExpired
}
// 激活成功即清除该店 phase 缓存:续费/换新授权码后写权限即时恢复,不必等 30s TTL。
defer middleware.InvalidateLicensePhase(shopID)
var existing model.LicenseDevice
err := s.db.Where("license_id = ? AND device_id = ?", lic.ID, deviceID).First(&existing).Error
if err == nil {
@@ -116,11 +120,52 @@ func (s *LicenseService) ShopInfo(shopID uint64) (*model.License, error) {
var lic model.License
if err := s.db.Where("shop_id = ? AND is_active = 1", shopID).
Order("id DESC").First(&lic).Error; err != nil {
return nil, ErrLicenseNotFound
// 仅「确无记录」才算无授权;DB 不可达等瞬时错误必须上抛,
// 否则会被误判为「门店无授权」,把客户端横幅/门禁错误降级。
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, ErrLicenseNotFound
}
return nil, err
}
return &lic, nil
}
// LicenseInfoView 门店授权概况(含设备数与实时 phase),供 /license/info 与心跳 /auth/ping 复用,
// 保证两条路径返回结构一致。
type LicenseInfoView struct {
ID uint64 `json:"id"`
Type string `json:"type"`
IsActive bool `json:"is_active"`
MaxDevices int `json:"max_devices"`
DeviceCount int64 `json:"device_count"`
ExpiresAt *time.Time `json:"expires_at"`
Phase string `json:"phase"`
}
// ShopInfoView 返回门店授权概况;无有效授权时返回 (nil, nil),仅在统计设备数等查询出错时返回 error。
func (s *LicenseService) ShopInfoView(shopID uint64) (*LicenseInfoView, error) {
lic, err := s.ShopInfo(shopID)
if err != nil {
if errors.Is(err, ErrLicenseNotFound) {
return nil, nil // 确无有效授权
}
return nil, err // 瞬时错误上抛:Ping 据此省略 license 字段,客户端保留上次状态
}
count, err := s.CountDevices(lic.ID)
if err != nil {
return nil, err
}
return &LicenseInfoView{
ID: lic.ID,
Type: lic.Type,
IsActive: lic.IsActive,
MaxDevices: lic.MaxDevices,
DeviceCount: count,
ExpiresAt: lic.ExpiresAt,
Phase: middleware.CalcLicensePhase(lic.ExpiresAt),
}, nil
}
// CountDevices 返回指定 license 下已绑定设备数。
func (s *LicenseService) CountDevices(licenseID uint64) (int64, error) {
var count int64
@@ -174,7 +219,13 @@ func issueTrialLicense(db *gorm.DB, shopID uint64) error {
IsActive: true,
MaxDevices: 1,
}
return db.Create(&lic).Error
if err := db.Create(&lic).Error; err != nil {
return err
}
// 注意:phase 缓存失效不在此处做——本函数运行在调用方事务内,提交前失效会留下
// 30s 窗口:并发请求可能在新 license 行可见前用旧 phase 重新填充缓存。
// 失效改由调用方在事务提交后执行(见 ensureTrialOnFirstUse)。
return nil
}
// createTrialLicense 在注册事务中为新门店签发 30 天 trial license。
@@ -0,0 +1,55 @@
package service
import (
"log"
"time"
"gorm.io/gorm"
"github.com/wangjia/jiu/backend/internal/model"
)
// sessionCleanupInterval 清理任务执行周期。
const sessionCleanupInterval = 24 * time.Hour
// StartSessionCleanup 启动后台清理 goroutine:启动即跑一次,之后每 24h 跑一次。
// 删除已撤销/已过期的会话行与过旧的失败登录记录,防止表无限膨胀、IP/UA 长期滞留。
// retentionDays<=0 时视为关闭清理(直接返回,不启动 goroutine)。
func StartSessionCleanup(db *gorm.DB, retentionDays int) {
if retentionDays <= 0 {
log.Printf("[cleanup] session cleanup disabled (retention_days=%d)", retentionDays)
return
}
go func() {
cleanupOnce(db, retentionDays)
ticker := time.NewTicker(sessionCleanupInterval)
defer ticker.Stop()
for range ticker.C {
cleanupOnce(db, retentionDays)
}
}()
}
// cleanupOnce 执行一轮清理,返回各表删除行数(供测试断言)。
func cleanupOnce(db *gorm.DB, retentionDays int) (sessions, attempts int64) {
cutoff := time.Now().AddDate(0, 0, -retentionDays)
r1 := db.Where(
"(revoked_at IS NOT NULL AND revoked_at < ?) OR (refresh_exp_at IS NOT NULL AND refresh_exp_at < ?)",
cutoff, cutoff,
).Delete(&model.UserSession{})
if r1.Error != nil {
log.Printf("[cleanup] purge user_sessions failed: %v", r1.Error)
}
r2 := db.Where("created_at < ?", cutoff).Delete(&model.LoginAttempt{})
if r2.Error != nil {
log.Printf("[cleanup] purge login_attempts failed: %v", r2.Error)
}
if r1.RowsAffected > 0 || r2.RowsAffected > 0 {
log.Printf("[cleanup] purged %d sessions, %d login_attempts (older than %dd)",
r1.RowsAffected, r2.RowsAffected, retentionDays)
}
return r1.RowsAffected, r2.RowsAffected
}
@@ -0,0 +1,292 @@
package service
import (
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/google/uuid"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/wangjia/jiu/backend/config"
"github.com/wangjia/jiu/backend/internal/middleware"
"github.com/wangjia/jiu/backend/internal/model"
"github.com/wangjia/jiu/backend/testutil"
)
// parseRefreshClaims 解析 refresh token 的 claims。
func parseRefreshClaims(t *testing.T, token string) *middleware.Claims {
t.Helper()
claims := &middleware.Claims{}
_, err := jwt.ParseWithClaims(token, claims, func(*jwt.Token) (interface{}, error) {
return []byte(config.C.JWT.Secret), nil
})
require.NoError(t, err)
return claims
}
// parseRefreshJTI 取出 refresh token 的 jtiRegisteredClaims.ID)。
func parseRefreshJTI(t *testing.T, token string) string {
return parseRefreshClaims(t, token).ID
}
// signLegacyRefresh 签一个不带 jti 的 refresh token(模拟发版前的存量 token)。
func signLegacyRefresh(t *testing.T, userID, shopID uint64, role, sid string) string {
t.Helper()
now := time.Now()
claims := middleware.Claims{
UserID: userID, ShopID: shopID, Role: role, SID: sid,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(now.Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(now),
},
}
s, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(config.C.JWT.Secret))
require.NoError(t, err)
return s
}
// #1 续期轮换 jti;旧 refresh token 重放 → 判定盗用 → 吊销整条会话。
func TestRefreshTokens_RotationAndReuseDetection(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD01")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
pair1, _, err := svc.Login("HARD01", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
jti0 := parseRefreshJTI(t, pair1.RefreshToken)
require.NotEmpty(t, jti0)
// 首次续期成功,jti 轮换。
pair2, err := svc.RefreshTokens(pair1.RefreshToken)
require.NoError(t, err)
jti1 := parseRefreshJTI(t, pair2.RefreshToken)
assert.NotEqual(t, jti0, jti1, "续期应轮换 jti")
// 重放已被取代的旧 refresh token → 盗用信号 → ErrSessionRevoked。
_, err = svc.RefreshTokens(pair1.RefreshToken)
assert.ErrorIs(t, err, ErrSessionRevoked)
// 整条会话被吊销,reason=reuse。
var sess model.UserSession
require.NoError(t, db.Where("user_id = ?", user.ID).First(&sess).Error)
assert.NotNil(t, sess.RevokedAt)
assert.Equal(t, "reuse", sess.RevokedReason)
// 即便是「最新」的 refresh token,此后也无法再续期(family 已撤销)。
_, err = svc.RefreshTokens(pair2.RefreshToken)
assert.ErrorIs(t, err, ErrSessionRevoked)
}
// #1 向后兼容:存量会话(refresh_jti 为空)+ 不带 jti 的旧 refresh token,首刷应放行并采纳新 jti。
func TestRefreshTokens_LegacyTokenBackwardCompat(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD02")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
sid := uuid.New().String()
now := time.Now()
require.NoError(t, db.Create(&model.UserSession{
ShopID: shop.ID, UserID: user.ID, SID: sid,
Platform: "windows", PlatformClass: "desktop",
RefreshJTI: "", // 存量会话无 jti
LastSeenAt: now,
RefreshExpAt: now.Add(time.Hour),
}).Error)
legacy := signLegacyRefresh(t, user.ID, shop.ID, "admin", sid)
pair, err := svc.RefreshTokens(legacy)
require.NoError(t, err, "存量 token 首刷应放行")
// 采纳新 jti 写回会话。
newJTI := parseRefreshJTI(t, pair.RefreshToken)
assert.NotEmpty(t, newJTI)
var sess model.UserSession
require.NoError(t, db.Where("sid = ?", sid).First(&sess).Error)
assert.Equal(t, newJTI, sess.RefreshJTI)
assert.Nil(t, sess.RevokedAt)
}
// #4 无 sid 的存量 token 首刷应自建可吊销会话,从此纳入会话治理(可被强制下线)。
func TestRefreshTokens_LegacyNoSidAdoptsSession(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD08")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
// SID 为空的存量 refresh token(发版前签发,从无会话行)。
legacy := signLegacyRefresh(t, user.ID, shop.ID, "admin", "")
pair, err := svc.RefreshTokens(legacy)
require.NoError(t, err, "无 sid 存量 token 首刷应放行")
// 新 token 带上自建的 sid,且已落库一条会话。
newSID := parseRefreshClaims(t, pair.RefreshToken).SID
require.NotEmpty(t, newSID, "首刷应签发带 sid 的新 token")
var sess model.UserSession
require.NoError(t, db.Where("user_id = ? AND sid = ?", user.ID, newSID).First(&sess).Error)
assert.Equal(t, "legacy", sess.PlatformClass)
assert.Nil(t, sess.RevokedAt)
// 自此可被治理:管理员强制下线后,新 token 无法再续期。
views, _ := svc.ListSessions(shop.ID, "")
require.Len(t, views, 1)
require.NoError(t, svc.ForceLogout(shop.ID, views[0].ID, 0))
_, err = svc.RefreshTokens(pair.RefreshToken)
assert.ErrorIs(t, err, ErrSessionRevoked)
}
// #1 同一无 sid 存量 token 重复续期,应复用唯一 legacy 会话而非每次新建(防无界膨胀 + 配额规避)。
func TestRefreshTokens_LegacyNoSidReusesSingleSession(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD09")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
legacy := signLegacyRefresh(t, user.ID, shop.ID, "admin", "")
// 反复呈递「同一」无 sid 存量 token(模拟未采纳新 token 的客户端/重放)。
var firstSID string
for i := 0; i < 5; i++ {
pair, err := svc.RefreshTokens(legacy)
require.NoError(t, err, "存量 token 续期应放行")
sid := parseRefreshClaims(t, pair.RefreshToken).SID
require.NotEmpty(t, sid)
if i == 0 {
firstSID = sid
} else {
assert.Equal(t, firstSID, sid, "重复续期应复用同一 legacy 会话的 sid")
}
}
// 始终只有一条 legacy 会话,而非 5 条。
var count int64
require.NoError(t, db.Model(&model.UserSession{}).
Where("shop_id = ? AND user_id = ? AND platform_class = ?", shop.ID, user.ID, "legacy").
Count(&count).Error)
assert.EqualValues(t, 1, count, "重复存量续期不应无界新建会话")
}
// #4 预存多条活跃 legacy 会话(legacy class 不在并发配额内)时,一次无 sid 续期应把它们
// 收敛为一条:复用最早的一条、吊销其余,使「每用户至多一条 legacy 会话」成为显式强制的上限。
func TestRefreshTokens_LegacyNoSidCollapsesExtraSessions(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD10")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
now := time.Now()
sids := []string{uuid.New().String(), uuid.New().String(), uuid.New().String()}
for _, sid := range sids {
require.NoError(t, db.Create(&model.UserSession{
ShopID: shop.ID, UserID: user.ID, SID: sid,
Platform: "legacy", PlatformClass: "legacy",
RefreshJTI: "", LastSeenAt: now, RefreshExpAt: now.Add(time.Hour),
}).Error)
}
legacy := signLegacyRefresh(t, user.ID, shop.ID, "admin", "")
pair, err := svc.RefreshTokens(legacy)
require.NoError(t, err)
// 复用最早创建(id 最小)的那条会话。
keptSID := parseRefreshClaims(t, pair.RefreshToken).SID
assert.Equal(t, sids[0], keptSID, "应复用最早的一条 legacy 会话")
// 仅剩一条活跃 legacy 会话,其余被吊销(reason=kicked)。
var active int64
require.NoError(t, db.Model(&model.UserSession{}).
Where("shop_id = ? AND user_id = ? AND platform_class = ? AND revoked_at IS NULL",
shop.ID, user.ID, "legacy").Count(&active).Error)
assert.EqualValues(t, 1, active, "多余 legacy 会话应被收敛为一条")
var revoked int64
require.NoError(t, db.Model(&model.UserSession{}).
Where("shop_id = ? AND user_id = ? AND revoked_reason = ?",
shop.ID, user.ID, "kicked").Count(&revoked).Error)
assert.EqualValues(t, 2, revoked, "其余两条应以 kicked 吊销")
}
// #4 ForceLogout 写入 revoked_by。
func TestForceLogout_RecordsRevokedBy(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD04")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
_, _, err := svc.Login("HARD04", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
views, _ := svc.ListSessions(shop.ID, "")
require.Len(t, views, 1)
const adminID = uint64(42)
require.NoError(t, svc.ForceLogout(shop.ID, views[0].ID, adminID))
var sess model.UserSession
require.NoError(t, db.Where("id = ?", views[0].ID).First(&sess).Error)
require.NotNil(t, sess.RevokedBy)
assert.Equal(t, adminID, *sess.RevokedBy)
assert.Equal(t, "admin", sess.RevokedReason)
}
// #5 清理:删除已撤销/过期会话与过旧失败登录,保留新鲜行。
func TestCleanupOnce_PurgesStaleRows(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD05")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
old := time.Now().AddDate(0, 0, -100) // 早于 90 天保留期
fresh := time.Now()
revokedOld := old
// 1) 久前撤销的会话 → 删
require.NoError(t, db.Create(&model.UserSession{
ShopID: shop.ID, UserID: user.ID, SID: "s-revoked-old",
LastSeenAt: old, RefreshExpAt: fresh.Add(time.Hour), RevokedAt: &revokedOld,
}).Error)
// 2) refresh 久前过期的会话 → 删
require.NoError(t, db.Create(&model.UserSession{
ShopID: shop.ID, UserID: user.ID, SID: "s-expired-old",
LastSeenAt: old, RefreshExpAt: old,
}).Error)
// 3) 新鲜活跃会话 → 保留
require.NoError(t, db.Create(&model.UserSession{
ShopID: shop.ID, UserID: user.ID, SID: "s-fresh",
LastSeenAt: fresh, RefreshExpAt: fresh.Add(time.Hour),
}).Error)
// 失败登录:旧 → 删;新 → 留
require.NoError(t, db.Create(&model.LoginAttempt{Username: "x", Reason: "bad_password", CreatedAt: old}).Error)
require.NoError(t, db.Create(&model.LoginAttempt{Username: "y", Reason: "bad_password", CreatedAt: fresh}).Error)
sessions, attempts := cleanupOnce(db, 90)
assert.Equal(t, int64(2), sessions)
assert.Equal(t, int64(1), attempts)
var sessLeft, attLeft int64
db.Model(&model.UserSession{}).Count(&sessLeft)
db.Model(&model.LoginAttempt{}).Count(&attLeft)
assert.Equal(t, int64(1), sessLeft)
assert.Equal(t, int64(1), attLeft)
}
// #7 失败登录落库(ip/ua 来自 DeviceInforeason 正确)。
func TestLogin_RecordsFailedAttempt(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "HARD07")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
dev := DeviceInfo{Platform: "windows", IP: "1.2.3.4", UserAgent: "curl/8.1"}
_, _, err := svc.Login("HARD07", "admin", "wrong-password", dev)
require.Error(t, err)
var att model.LoginAttempt
require.NoError(t, db.Where("username = ?", "admin").First(&att).Error)
assert.False(t, att.Success)
assert.Equal(t, "bad_password", att.Reason)
assert.Equal(t, "1.2.3.4", att.IP)
assert.Equal(t, "curl/8.1", att.UserAgent)
assert.Equal(t, "HARD07", att.ShopCode)
}
+2 -2
View File
@@ -123,7 +123,7 @@ func TestForceLogout_RevokesSession(t *testing.T) {
require.Len(t, views, 1)
assert.True(t, views[0].Online)
require.NoError(t, svc.ForceLogout(shop.ID, views[0].ID))
require.NoError(t, svc.ForceLogout(shop.ID, views[0].ID, 0))
_, err = svc.RefreshTokens(pair.RefreshToken)
assert.ErrorIs(t, err, ErrSessionRevoked)
@@ -148,7 +148,7 @@ func TestForceLogout_TenantIsolation(t *testing.T) {
require.Len(t, views, 1)
// 用 shopB 的 shopID 尝试下线 shopA 的会话 → 找不到
err = svc.ForceLogout(shopB.ID, views[0].ID)
err = svc.ForceLogout(shopB.ID, views[0].ID, 0)
assert.Error(t, err)
}