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) } }