Files
jiu/backend/internal/service/session_hardening_test.go
wangjia e41085a878 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>
2026-06-19 07:34:04 +08:00

293 lines
11 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
}