241 lines
7.5 KiB
Go
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
|
|
}
|