Files
2026-07-30 12:23:05 +08:00

526 lines
16 KiB
Go

package service
import (
"context"
"errors"
"fmt"
"testing"
"time"
"gorm.io/gorm"
"oci-portal/internal/model"
)
// loginSession 用密码登录建一条会话,返回令牌。
func loginSession(t *testing.T, auth *AuthService, ip, ua string) string {
t.Helper()
token, _, err := auth.Login(context.Background(), "admin", "pass123", "",
SessionMeta{ClientIP: ip, UserAgent: ua})
if err != nil {
t.Fatalf("login: %v", err)
}
return token
}
// mustProof 从有效令牌解析版本/jti 快照。
func mustProof(t *testing.T, auth *AuthService, token string) TokenProof {
t.Helper()
_, proof, err := auth.ParseTokenProof(context.Background(), token)
if err != nil {
t.Fatalf("parse token proof: %v", err)
}
return proof
}
func sessionRows(t *testing.T, auth *AuthService) []model.UserSession {
t.Helper()
rows := []model.UserSession{}
if err := auth.db.Order("id").Find(&rows).Error; err != nil {
t.Fatalf("load sessions: %v", err)
}
return rows
}
func TestSessionRecordedOnLogin(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
loginSession(t, auth, "10.9.0.1", "TestAgent/1.0")
rows := sessionRows(t, auth)
if len(rows) != 1 {
t.Fatalf("sessions = %d, want 1", len(rows))
}
r := rows[0]
if r.Method != "password" || r.ClientIP != "10.9.0.1" || r.UserAgent != "TestAgent/1.0" {
t.Errorf("row = %+v, want password/10.9.0.1/TestAgent", r)
}
if r.TokenID == "" || r.ExpiresAt.Before(time.Now()) {
t.Errorf("row token/expiry invalid: %+v", r)
}
}
func TestSessionRenewContinuity(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
token := loginSession(t, auth, "10.9.0.2", "UA")
before := sessionRows(t, auth)[0]
// 敏感变更路径:版本递增 + 换发接续(RevokeSessions 即 bump+renew)
fresh, _, err := auth.RevokeSessions(ctx, "admin", token, SessionMeta{ClientIP: "10.9.0.2", UserAgent: "UA"}, mustProof(t, auth, token))
if err != nil {
t.Fatalf("RevokeSessions: %v", err)
}
rows := sessionRows(t, auth)
if len(rows) != 1 {
t.Fatalf("sessions after renew = %d, want 1 (continuity, no new row)", len(rows))
}
after := rows[0]
if after.ID != before.ID || after.TokenID != before.TokenID {
t.Errorf("renew should keep row id and jti: before %+v after %+v", before, after)
}
if after.Method != "password" {
t.Errorf("method = %q, want inherited password", after.Method)
}
if _, err := auth.ParseToken(ctx, token); err == nil {
t.Error("old token still valid after version bump")
}
if _, err := auth.ParseToken(ctx, fresh); err != nil {
t.Errorf("fresh token invalid: %v", err)
}
// 列表只剩接续会话且标记 current
list, err := auth.ListSessions(ctx, "admin", fresh)
if err != nil || len(list) != 1 || !list[0].Current {
t.Errorf("ListSessions = (%+v, %v), want single current session", list, err)
}
}
func TestSessionRevokeSingle(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
token1 := loginSession(t, auth, "10.9.0.3", "Laptop")
token2 := loginSession(t, auth, "10.9.0.4", "Phone")
// token2 已被 ParseToken 校验过(节流缓存生效)后再撤销,验证缓存被清、即时失效
if _, err := auth.ParseToken(ctx, token2); err != nil {
t.Fatalf("token2 parse: %v", err)
}
list, err := auth.ListSessions(ctx, "admin", token1)
if err != nil || len(list) != 2 {
t.Fatalf("ListSessions = (%d, %v), want 2", len(list), err)
}
var otherID uint
for _, it := range list {
if it.Current {
continue
}
otherID = it.ID
}
tests := []struct {
name string
id uint
wantErr error
}{
{name: "撤销当前会话被拒", id: currentSessionID(t, list), wantErr: ErrSessionCurrent},
{name: "撤销其他会话成功", id: otherID},
{name: "重复撤销已不可见", id: otherID, wantErr: gorm.ErrRecordNotFound},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := auth.RevokeSession(ctx, "admin", token1, tt.id)
if tt.wantErr != nil {
if !errors.Is(err, tt.wantErr) && !(tt.wantErr == gorm.ErrRecordNotFound && err != nil) {
t.Fatalf("RevokeSession err = %v, want %v", err, tt.wantErr)
}
return
}
if err != nil {
t.Fatalf("RevokeSession: %v", err)
}
})
}
if _, err := auth.ParseToken(ctx, token2); err == nil {
t.Error("revoked session token still valid")
}
if _, err := auth.ParseToken(ctx, token1); err != nil {
t.Errorf("current token broken by revoking another session: %v", err)
}
}
func TestSessionRevokeTombstone(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
token1 := loginSession(t, auth, "10.9.0.7", "Laptop")
token2 := loginSession(t, auth, "10.9.0.8", "Phone")
list, err := auth.ListSessions(ctx, "admin", token1)
if err != nil || len(list) != 2 {
t.Fatalf("ListSessions = (%d, %v), want 2", len(list), err)
}
var otherID uint
for _, it := range list {
if !it.Current {
otherID = it.ID
}
}
if err := auth.RevokeSession(ctx, "admin", token1, otherID); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
// 模拟「校验读库通过→撤销→校验回写 seen」的竞态:手工回写节流缓存,
// 撤销负缓存必须仍然盖过它,令牌不得复活
jti := auth.signedJti(token2)
auth.seen.Set("seen|"+jti, struct{}{}, sessionSeenTTL)
if _, err := auth.ParseToken(ctx, token2); err == nil {
t.Error("revoked token revived by racing seen-cache write")
}
}
// TestSensitiveOpStaleProof 验证在途绕过防线:敏感请求鉴权后挂起,期间发生
// 「撤销全部」(版本递增),恢复后的敏感事务复核快照失败,拒绝执行、不发新令牌。
// 同一防线也使并发敏感操作串行化(后到者复核失败)。
func TestSensitiveOpStaleProof(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
tokenA := loginSession(t, auth, "10.9.1.1", "A")
proofA := mustProof(t, auth, tokenA)
// 另一设备撤销全部:版本递增,tokenA 的快照随之过期
tokenB := loginSession(t, auth, "10.9.1.2", "B")
if _, _, err := auth.RevokeSessions(ctx, "admin", tokenB, SessionMeta{ClientIP: "10.9.1.2"}, mustProof(t, auth, tokenB)); err != nil {
t.Fatalf("RevokeSessions: %v", err)
}
// 挂起的旧请求恢复:携带过期快照的敏感操作必须被拒
if _, _, err := auth.RevokeSessions(ctx, "admin", tokenA, SessionMeta{ClientIP: "10.9.1.1"}, proofA); !errors.Is(err, ErrTokenStale) {
t.Fatalf("stale-proof revoke err = %v, want ErrTokenStale", err)
}
}
// TestSensitiveOpRevokedJtiProof 验证快照的 jti 维度:定点撤销(不递增版本)
// 同样令该令牌的在途敏感请求失效。
func TestSensitiveOpRevokedJtiProof(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
token1 := loginSession(t, auth, "10.9.2.1", "Laptop")
token2 := loginSession(t, auth, "10.9.2.2", "Phone")
proof2 := mustProof(t, auth, token2)
list, err := auth.ListSessions(ctx, "admin", token1)
if err != nil || len(list) != 2 {
t.Fatalf("ListSessions = (%d, %v), want 2", len(list), err)
}
for _, it := range list {
if !it.Current {
if err := auth.RevokeSession(ctx, "admin", token1, it.ID); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
}
}
// token2 已被定点撤销(版本未变):其在途敏感请求恢复后必须被拒
if _, _, err := auth.RevokeSessions(ctx, "admin", token2, SessionMeta{ClientIP: "10.9.2.2"}, proof2); !errors.Is(err, ErrTokenStale) {
t.Fatalf("revoked-jti proof err = %v, want ErrTokenStale", err)
}
}
func TestCheckedSessionCannotRenewAfterTargetedRevoke(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
stale := loginSession(t, auth, "10.9.3.1", "stale")
current := loginSession(t, auth, "10.9.3.2", "current")
proof := mustProof(t, auth, stale)
assertProofCurrentTx(t, auth, proof)
staleID := otherSessionID(t, auth, current)
if err := auth.bumpTokenVersion(ctx, "admin"); err != nil {
t.Fatalf("simulate sensitive mutation: %v", err)
}
if err := auth.RevokeSession(ctx, "admin", current, staleID); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
if _, _, err := auth.RenewToken(
ctx, "admin", stale, SessionMeta{ClientIP: "10.9.3.1"},
); !errors.Is(err, ErrTokenStale) {
t.Fatalf("RenewToken err = %v, want ErrTokenStale", err)
}
if got := len(sessionRows(t, auth)); got != 2 {
t.Fatalf("session rows = %d, want 2 (不得 fallback CREATE)", got)
}
}
func otherSessionID(t *testing.T, auth *AuthService, current string) uint {
t.Helper()
for _, item := range mustSessions(t, auth, current) {
if !item.Current {
return item.ID
}
}
t.Fatal("no other session")
return 0
}
func TestLegacyLogoutBlocksInflightProofAndRenew(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
token, _, _, err := auth.signToken("admin", 0)
if err != nil {
t.Fatalf("signToken: %v", err)
}
proof := mustProof(t, auth, token)
assertProofCurrentTx(t, auth, proof)
auth.Logout(ctx, token)
assertProofStaleTx(t, auth, proof)
if err := auth.bumpTokenVersion(ctx, "admin"); err != nil {
t.Fatalf("simulate already committed mutation: %v", err)
}
if _, _, err := auth.RenewToken(
ctx, "admin", token, SessionMeta{ClientIP: "10.9.4.1"},
); !errors.Is(err, ErrTokenStale) {
t.Fatalf("RenewToken err = %v, want ErrTokenStale", err)
}
if got := len(sessionRows(t, auth)); got != 0 {
t.Fatalf("session rows = %d, want 0", got)
}
}
type logoutAfterRenameCase struct {
name string
legacy bool
meta SessionMeta
}
func TestLogoutOldTokenAfterCredentialRename(t *testing.T) {
tests := []logoutAfterRenameCase{
{name: "recorded session", meta: SessionMeta{ClientIP: "10.9.5.1", UserAgent: "recorded"}},
{name: "legacy session without row", legacy: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
testLogoutOldTokenAfterCredentialRename(t, tt)
})
}
}
func testLogoutOldTokenAfterCredentialRename(t *testing.T, tt logoutAfterRenameCase) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
old, oldExpires := tokenForRenewalTest(t, auth, tt.legacy)
fresh, expires := renewAfterCredentialRename(t, auth, old, tt.meta)
if auth.signedJti(old) != auth.signedJti(fresh) {
t.Fatal("renewal changed jti; logout lineage would be lost")
}
if tt.legacy && len(sessionRows(t, auth)) != 0 {
t.Fatal("legacy renewal unexpectedly created a session row")
}
auth.Logout(context.Background(), old)
if _, err := auth.ParseToken(context.Background(), fresh); err == nil {
t.Fatal("fresh token survived logout of its pre-renewal token")
}
if tt.legacy {
assertTombstoneCovers(t, auth, old, oldExpires, expires)
return
}
assertSessionJTIRevoked(t, auth, old)
}
func renewAfterCredentialRename(
t *testing.T, auth *AuthService, old string, meta SessionMeta,
) (string, time.Time) {
t.Helper()
ctx := context.Background()
finalName, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
NewUsername: "root", CurrentPassword: "pass123",
}, mustProof(t, auth, old))
if err != nil {
t.Fatalf("UpdateCredentials: %v", err)
}
fresh, expires, err := auth.RenewToken(ctx, finalName, old, meta)
if err != nil {
t.Fatalf("RenewToken: %v", err)
}
return fresh, expires
}
func assertSessionJTIRevoked(t *testing.T, auth *AuthService, token string) {
t.Helper()
var row model.UserSession
err := auth.db.Where("token_id = ?", auth.signedJti(token)).First(&row).Error
if err != nil {
t.Fatalf("find session: %v", err)
}
if row.RevokedAt == nil {
t.Fatal("renamed user's session row was not revoked")
}
}
func assertTombstoneCovers(
t *testing.T, auth *AuthService, token string, oldExpires, freshExpires time.Time,
) {
t.Helper()
jti := auth.signedJti(token)
auth.revokedJti.mu.RLock()
tombstoneExpires, ok := auth.revokedJti.m[jti]
auth.revokedJti.mu.RUnlock()
if !ok {
t.Fatal("logout jti tombstone missing")
}
requiredUntil := oldExpires.Truncate(time.Second).Add(tokenTTL)
if tombstoneExpires.Before(requiredUntil) || tombstoneExpires.Before(freshExpires) {
t.Fatalf("tombstone expires %v before required horizon %v", tombstoneExpires, requiredUntil)
}
}
func tokenForRenewalTest(t *testing.T, auth *AuthService, legacy bool) (string, time.Time) {
t.Helper()
if !legacy {
token, expires, err := auth.Login(context.Background(), "admin", "pass123", "",
SessionMeta{ClientIP: "10.9.5.1", UserAgent: "recorded"})
if err != nil {
t.Fatalf("Login: %v", err)
}
return token, expires
}
token, expires, _, err := auth.signToken("admin", 0)
if err != nil {
t.Fatalf("signToken: %v", err)
}
return token, expires
}
func assertProofCurrentTx(t *testing.T, auth *AuthService, proof TokenProof) {
t.Helper()
err := auth.db.Transaction(func(tx *gorm.DB) error {
user, err := lockUserForAuthChange(tx, "admin")
if err != nil {
return err
}
return auth.ensureTokenCurrentTx(tx, user, proof)
})
if err != nil {
t.Fatalf("proof should be current: %v", err)
}
}
func assertProofStaleTx(t *testing.T, auth *AuthService, proof TokenProof) {
t.Helper()
err := auth.db.Transaction(func(tx *gorm.DB) error {
user, lockErr := lockUserForAuthChange(tx, "admin")
if lockErr != nil {
return lockErr
}
return auth.ensureTokenCurrentTx(tx, user, proof)
})
if !errors.Is(err, ErrTokenStale) {
t.Fatalf("proof err = %v, want ErrTokenStale", err)
}
}
func mustSessions(t *testing.T, auth *AuthService, current string) []SessionInfo {
t.Helper()
items, err := auth.ListSessions(context.Background(), "admin", current)
if err != nil {
t.Fatalf("ListSessions: %v", err)
}
return items
}
// TestTombstoneNoEviction 验证撤销负缓存无容量上限:大量撤销标记全部存活,
// 不存在「满载淘汰导致已撤销令牌复活」。
func TestTombstoneNoEviction(t *testing.T) {
ts := newJtiTombstones()
for i := 0; i < 600; i++ {
ts.put(fmt.Sprintf("jti-%d", i), time.Minute)
}
for i := 0; i < 600; i++ {
if !ts.has(fmt.Sprintf("jti-%d", i)) {
t.Fatalf("tombstone jti-%d evicted", i)
}
}
if ts.has("jti-none") {
t.Error("unknown jti reported revoked")
}
}
// currentSessionID 取列表中 current 条目的 ID。
func currentSessionID(t *testing.T, list []SessionInfo) uint {
t.Helper()
for _, it := range list {
if it.Current {
return it.ID
}
}
t.Fatal("no current session in list")
return 0
}
func TestSessionLegacyTokenAllowed(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
// 直签令牌(无会话行)模拟升级前存量令牌:仍可通过校验
token, _, _, err := auth.signToken("admin", 0)
if err != nil {
t.Fatalf("signToken: %v", err)
}
if _, err := auth.ParseToken(context.Background(), token); err != nil {
t.Errorf("legacy token rejected: %v", err)
}
}
func TestSessionLogoutMarksRow(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
token := loginSession(t, auth, "10.9.0.5", "UA")
auth.Logout(ctx, token)
rows := sessionRows(t, auth)
if len(rows) != 1 || rows[0].RevokedAt == nil {
t.Errorf("logout should mark session revoked: %+v", rows)
}
}
func TestSessionCleanup(t *testing.T) {
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
now := time.Now()
old := now.Add(-8 * 24 * time.Hour)
seed := []model.UserSession{
{UserID: 1, TokenID: "expired", ExpiresAt: now.Add(-25 * time.Hour), LastSeenAt: old},
{UserID: 1, TokenID: "revoked-old", ExpiresAt: now.Add(time.Hour), RevokedAt: &old, LastSeenAt: old},
{UserID: 1, TokenID: "alive", ExpiresAt: now.Add(time.Hour), LastSeenAt: now},
}
for i := range seed {
if err := auth.db.Create(&seed[i]).Error; err != nil {
t.Fatalf("seed: %v", err)
}
}
auth.cleanupSessionsOnce(ctx)
rows := sessionRows(t, auth)
if len(rows) != 1 || rows[0].TokenID != "alive" {
t.Errorf("after cleanup rows = %+v, want only alive", rows)
}
}