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 的 jti(RegisteredClaims.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 来自 DeviceInfo,reason 正确)。 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) }