Files
krow_backend/go-api/internal/oauth/cleanup_test.go
Aravind f2aa3b3ad8
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled
mcp connection
2026-09-22 10:58:02 +05:30

241 lines
7.5 KiB
Go

package oauth
import (
"context"
"sync"
"testing"
"time"
)
/* ── What cleanup removes ───────────────────────────────────────────────── */
func TestCleanupRemovesOnlyDeadRows(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
// A live pair, and a spent code.
code := h.authorizeOK(clientID, verifier43)
live := decodeTokens(t, h.exchange(clientID, code, verifier43))
// A second, which we let expire.
oldCode := h.authorizeOK(clientID, verifier43)
old := decodeTokens(t, h.exchange(clientID, oldCode, verifier43))
before := countRows(t, h)
// Past the access token TTL and the grace, but well inside the refresh
// token's thirty days.
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
result, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("Cleanup: %v", err)
}
// Both access tokens and both codes are dead; both refresh tokens are not.
if result.AccessTokens != 2 {
t.Errorf("removed %d access tokens, want 2", result.AccessTokens)
}
if result.Grants != 2 {
t.Errorf("removed %d grants, want 2", result.Grants)
}
if result.RefreshTokens != 0 {
t.Errorf("removed %d refresh tokens, want 0 — they live thirty days", result.RefreshTokens)
}
if !result.CompletedInOne {
t.Error("a small cleanup reported that more remained")
}
after := countRows(t, h)
if after.tokens >= before.tokens {
t.Error("cleanup removed nothing")
}
// The live refresh tokens must still work. This is the property that
// matters: cleanup must not disconnect anybody.
for name, token := range map[string]string{"live": live.RefreshToken, "old": old.RefreshToken} {
if _, err := h.store.RedeemRefreshToken(ctx, token); err != nil {
t.Errorf("the %s refresh token stopped working after cleanup: %v", name, err)
}
}
}
// The subtle one: a CONSUMED refresh token must survive until its expiry,
// because it is what reuse detection matches against. Delete it early and a
// replayed stolen token looks unknown rather than stolen, and the family is
// never revoked.
func TestCleanupKeepsConsumedRefreshTokensForReuseDetection(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
// Rotate: `first` is now consumed.
if _, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken); err != nil {
t.Fatalf("rotate: %v", err)
}
// Cleanup well past the access token TTL, but inside the refresh TTL.
h.advance(AccessTokenTTL + CleanupGrace + time.Hour)
if _, err := h.store.Cleanup(ctx); err != nil {
t.Fatalf("Cleanup: %v", err)
}
// Replaying the consumed token must STILL be detected as reuse.
_, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken)
if err != ErrRefreshReuse {
t.Errorf("err = %v, want ErrRefreshReuse — cleanup destroyed the evidence "+
"that makes theft detectable", err)
}
}
// Revoked rows are kept until expiry too: "revoked at 14:02 for
// refresh_token_reuse" is an answer somebody will eventually need.
func TestCleanupKeepsRevokedRowsUntilExpiry(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
if err := h.store.RevokeToken(ctx, tokens.AccessToken, "test"); err != nil {
t.Fatalf("revoke: %v", err)
}
// Just past the access TTL: the access row goes, the refresh row stays.
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
if _, err := h.store.Cleanup(ctx); err != nil {
t.Fatalf("Cleanup: %v", err)
}
var revokedRefresh int
if err := h.h.Pool.QueryRow(ctx,
`SELECT count(*) FROM oauth_tokens WHERE token_type='refresh' AND revoked_at IS NOT NULL`).
Scan(&revokedRefresh); err != nil {
t.Fatalf("count: %v", err)
}
if revokedRefresh != 1 {
t.Errorf("%d revoked refresh rows kept, want 1 — the audit trail was swept", revokedRefresh)
}
}
// Nothing is deleted before the grace period, so clock skew between instances
// cannot destroy a row another instance still considers live.
func TestCleanupHonoursTheGracePeriod(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
decodeTokens(t, h.exchange(clientID, code, verifier43))
// Expired, but inside the grace.
h.advance(AccessTokenTTL + time.Minute)
result, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("Cleanup: %v", err)
}
if result.AccessTokens != 0 {
t.Errorf("removed %d access tokens inside the grace period, want 0", result.AccessTokens)
}
}
/* ── Safety ─────────────────────────────────────────────────────────────── */
func TestCleanupIsIdempotent(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
decodeTokens(t, h.exchange(clientID, code, verifier43))
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
first, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("first: %v", err)
}
second, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("second: %v", err)
}
if second.AccessTokens != 0 || second.Grants != 0 || second.RefreshTokens != 0 {
t.Errorf("a second cleanup removed more rows: %+v (first was %+v)", second, first)
}
}
// Two workers running cleanup at once must not error and must not
// double-count. Run with -race.
func TestConcurrentCleanupIsSafe(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
for i := 0; i < 6; i++ {
code := h.authorizeOK(clientID, verifier43)
decodeTokens(t, h.exchange(clientID, code, verifier43))
}
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
const workers = 4
var wg sync.WaitGroup
var mu sync.Mutex
var total int64
errs := make([]error, 0, workers)
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
r, err := h.store.Cleanup(ctx)
mu.Lock()
defer mu.Unlock()
if err != nil {
errs = append(errs, err)
return
}
total += r.AccessTokens
}()
}
wg.Wait()
if len(errs) > 0 {
t.Fatalf("concurrent cleanup errored: %v", errs)
}
// Six access tokens existed; between them the workers removed exactly six.
// More would mean a row was counted twice.
if total != 6 {
t.Errorf("workers removed %d access tokens between them, want 6", total)
}
}
func TestCleanupOnAnEmptyDatabaseIsHarmless(t *testing.T) {
h := newHarness(t)
result, err := h.store.Cleanup(context.Background())
if err != nil {
t.Fatalf("Cleanup on empty: %v", err)
}
if result.Grants != 0 || result.AccessTokens != 0 || result.RefreshTokens != 0 {
t.Errorf("cleanup on an empty database removed %+v", result)
}
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
type rowCounts struct{ grants, tokens int }
func countRows(t *testing.T, h *harness) rowCounts {
t.Helper()
var c rowCounts
ctx := context.Background()
if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_grants`).Scan(&c.grants); err != nil {
t.Fatalf("count grants: %v", err)
}
if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_tokens`).Scan(&c.tokens); err != nil {
t.Fatalf("count tokens: %v", err)
}
return c
}