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 }