diff --git a/internal/api/pgrepo.go b/internal/api/pgrepo.go index 4854c99..e240c3d 100644 --- a/internal/api/pgrepo.go +++ b/internal/api/pgrepo.go @@ -1247,6 +1247,9 @@ func (p *PGRepo) SetSetting(ctx context.Context, key string, value []byte) error // behind. That bounds the table at one row per (user, purpose): the begin→finish loop // nets zero growth, since each begin sweeps the consumed row the previous finish stamped. // (Deleting a consumed row is safe: it has already been redeemed and nothing reads it.) +// Begins for one (user, purpose) take a transaction-scoped advisory lock first: without +// it two racing begins each delete what the other has not committed yet and both insert, +// leaving extra rows that only the next begin sweeps. // The opaque SessionData is held server-side so the client cannot forge the challenge // it must answer at finish. func (p *PGRepo) CreatePasskeyChallenge(ctx context.Context, id, userID, purpose string, sessionData []byte, expiresAt time.Time) error { @@ -1256,6 +1259,11 @@ func (p *PGRepo) CreatePasskeyChallenge(ctx context.Context, id, userID, purpose } defer tx.Rollback() //nolint:errcheck // no-op after commit + if _, err := tx.ExecContext(ctx, + `SELECT pg_advisory_xact_lock(hashtext('passkey-challenge:' || $1 || ':' || $2))`, + userID, purpose); err != nil { + return fmt.Errorf("lock passkey challenges: %w", err) + } if _, err := tx.ExecContext(ctx, `DELETE FROM webauthn_challenges WHERE user_id = $1 AND purpose = $2`, userID, purpose); err != nil { @@ -1316,16 +1324,27 @@ func (p *PGRepo) ConsumePasskeyChallengeByUser(ctx context.Context, userID, purp // address or IPv6 /48, see challengeSource) holds in each login store. A ceremony // takes seconds and a challenge lives passkeyChallengeTTL, so a network with a few // dozen people signing in at once stays well inside it, while a flood of begins -// fills its own allowance and leaves the other networks their sign-ins. The count -// and the insert are not serialised, so a burst of truly concurrent begins can pass -// it together; the per-source token bucket in front of the door caps that burst. +// fills its own allowance and leaves the other networks their sign-ins. Each store's +// begins from one source serialise on a transaction-scoped advisory lock +// (lockChallengeSource), so racing begins cannot all pass one count. const maxLiveChallengesPerSource = 32 +// lockChallengeSource takes the advisory lock that makes a login store's per-source +// count and insert one decision. store names the table, so the two stores' allowances +// for one source never wait on each other. The lock is released at commit or rollback. +func lockChallengeSource(ctx context.Context, tx *sql.Tx, store, source string) error { + if _, err := tx.ExecContext(ctx, + `SELECT pg_advisory_xact_lock(hashtext($1 || ':' || $2))`, store, source); err != nil { + return fmt.Errorf("lock %s source: %w", store, err) + } + return nil +} + // AddPasskeyLoginChallenge stores an email-first login ceremony beside the ones // already live for (user, purpose); finish finds it by the challenge the browser // signed (ConsumePasskeyLoginChallenge), so a begin by anyone who knows the address -// never cancels its owner's ceremony. It reaps the account's spent login rows, -// refuses with ErrTooManyPasskeyChallenges once source holds +// never cancels its owner's ceremony. Under the source's lock it reaps the account's +// spent login rows, refuses with ErrTooManyPasskeyChallenges once source holds // maxLiveChallengesPerSource live login challenges, then inserts. func (p *PGRepo) AddPasskeyLoginChallenge(ctx context.Context, id, userID, purpose, source, challenge string, sessionData []byte, now, expiresAt time.Time) error { tx, err := p.db.BeginTx(ctx, nil) @@ -1334,6 +1353,9 @@ func (p *PGRepo) AddPasskeyLoginChallenge(ctx context.Context, id, userID, purpo } defer tx.Rollback() //nolint:errcheck // no-op after commit + if err := lockChallengeSource(ctx, tx, "webauthn_challenges", source); err != nil { + return err + } if _, err := tx.ExecContext(ctx, `DELETE FROM webauthn_challenges WHERE user_id = $1 AND purpose = $2 AND (expires_at <= $3 OR consumed_at IS NOT NULL)`, @@ -1408,12 +1430,14 @@ func (p *PGRepo) ConsumePasskeyLoginChallenge(ctx context.Context, userID, purpo const maxLiveDiscoverableChallenges = 16384 // CreateDiscoverableChallenge stashes a discoverable-login ceremony under an opaque handle, -// bounding the table in one transaction (see the Repo interface for the full contract). It -// reaps expired/consumed rows first, then refuses when source already holds -// maxLiveChallengesPerSource live rows or the table holds maxLiveDiscoverableChallenges. -// Because the reap ran first, the counts are exactly the live rows, so the bounds hold under -// an adversarial begin-flood (which a reap alone cannot: a burst inside the TTL leaves every -// fresh row live). +// bounding the table in one transaction (see the Repo interface for the full contract). +// Under the source's lock it reaps expired/consumed rows, then refuses when source already +// holds maxLiveChallengesPerSource live rows or the table holds +// maxLiveDiscoverableChallenges. Because the reap ran first, the counts are exactly the live +// rows, so the bounds hold under an adversarial begin-flood (which a reap alone cannot: a +// burst inside the TTL leaves every fresh row live). The per-source bound is exact; the +// table bound can be passed by begins from different sources racing the same count, at +// most one row per open connection, which leaves the few-MB ceiling where it was. func (p *PGRepo) CreateDiscoverableChallenge(ctx context.Context, id, source string, sessionData []byte, now, expiresAt time.Time) error { tx, err := p.db.BeginTx(ctx, nil) if err != nil { @@ -1421,6 +1445,9 @@ func (p *PGRepo) CreateDiscoverableChallenge(ctx context.Context, id, source str } defer tx.Rollback() //nolint:errcheck // no-op after commit + if err := lockChallengeSource(ctx, tx, "webauthn_discoverable_challenges", source); err != nil { + return err + } if _, err := tx.ExecContext(ctx, `DELETE FROM webauthn_discoverable_challenges WHERE expires_at <= $1 OR consumed_at IS NOT NULL`, now); err != nil { @@ -2146,6 +2173,12 @@ func (p *PGRepo) LinkAccount(ctx context.Context, userID, mcUUID, authSource str // earlier unfinished migration for the source (so re-running /felis migrate restarts // cleanly, invalidating a prior outstanding code) and inserts a fresh 'initiated' row, // both under one transaction so the partial unique index never sees two live rows. +// +// Starts for one source serialise on a transaction-scoped advisory lock, so racing +// starts each supersede the one before instead of tripping the unique index. The +// liveness check rides on the INSERT, after the supersede: a redeem of this source holds +// the code_issued row until it commits, the supersede waits on that row, and the INSERT's +// fresh snapshot then sees the source retired (a check before the wait would not). func (p *PGRepo) StartMigration(ctx context.Context, id, sourceUserID string, now time.Time) error { tx, err := p.db.BeginTx(ctx, nil) if err != nil { @@ -2153,29 +2186,32 @@ func (p *PGRepo) StartMigration(ctx context.Context, id, sourceUserID string, no } defer tx.Rollback() //nolint:errcheck // no-op after commit - // The source must be a live (non-deleted) account; a retired one can never - // re-initiate a migration. - var live bool - if err := tx.QueryRowContext(ctx, - `SELECT EXISTS(SELECT 1 FROM users WHERE id = $1 AND deleted_at IS NULL)`, - sourceUserID).Scan(&live); err != nil { - return err + if _, err := tx.ExecContext(ctx, + `SELECT pg_advisory_xact_lock(hashtext('account-migration:' || $1))`, sourceUserID); err != nil { + return fmt.Errorf("lock migration source: %w", err) } - if !live { - return ErrNotFound - } - if _, err := tx.ExecContext(ctx, `DELETE FROM account_migrations WHERE source_user_id = $1 AND state <> 'redeemed'`, sourceUserID); err != nil { return err } - if _, err := tx.ExecContext(ctx, + // The source must be a live (non-deleted) account; a retired one can never + // re-initiate a migration. + res, err := tx.ExecContext(ctx, `INSERT INTO account_migrations (id, source_user_id, state, created_at, updated_at) - VALUES ($1, $2, 'initiated', $3, $3)`, - id, sourceUserID, now); err != nil { + SELECT $1, $2, 'initiated', $3, $3 + WHERE EXISTS (SELECT 1 FROM users WHERE id = $2 AND deleted_at IS NULL)`, + id, sourceUserID, now) + if err != nil { return err } + n, err := res.RowsAffected() + if err != nil { + return err + } + if n == 0 { + return ErrNotFound + } return tx.Commit() } diff --git a/internal/pgint/accounts_test.go b/internal/pgint/accounts_test.go new file mode 100644 index 0000000..cec0a63 --- /dev/null +++ b/internal/pgint/accounts_test.go @@ -0,0 +1,518 @@ +//go:build pgint + +package pgint + +import ( + "context" + "database/sql" + "errors" + "sort" + "strings" + "sync" + "testing" + "time" + + "felis.lolicon.best/internal/api" +) + +// ---- account migration (migration 0015) ------------------------------------------- + +func seedOwnedServer(t *testing.T, name, ownerID string, deleted bool) { + t.Helper() + var deletedAt any + if deleted { + deletedAt = mustNow() + } + mustExec(t, `INSERT INTO servers (name, cached_cpu_milli, cached_memory_mb, cached_storage_mb, owner_id, deleted_at) + VALUES ($1, 100, 128, 1, $2, $3)`, name, ownerID, deletedAt) +} + +func serverOwner(t *testing.T, name string) string { + t.Helper() + var owner string + if err := db.QueryRow(`SELECT COALESCE(owner_id, '') FROM servers WHERE name = $1`, name).Scan(&owner); err != nil { + t.Fatalf("owner of %s: %v", name, err) + } + return owner +} + +func liveMigrations(t *testing.T, sourceID string) int { + t.Helper() + var n int + if err := db.QueryRow(`SELECT count(*) FROM account_migrations WHERE source_user_id = $1 AND state <> 'redeemed'`, + sourceID).Scan(&n); err != nil { + t.Fatalf("count live migrations: %v", err) + } + return n +} + +// startToCode drives a source from nothing to an issued code for target. +func startToCode(t *testing.T, sourceID, targetID, codeHash string, now, expiresAt time.Time) { + t.Helper() + ctx := context.Background() + if err := repo.StartMigration(ctx, "mig-"+suffix(t), sourceID, now); err != nil { + t.Fatalf("StartMigration: %v", err) + } + if err := repo.ConfirmMigration(ctx, sourceID, "passkey", now); err != nil { + t.Fatalf("ConfirmMigration: %v", err) + } + if err := repo.IssueMigrationCode(ctx, sourceID, targetID, codeHash, expiresAt); err != nil { + t.Fatalf("IssueMigrationCode: %v", err) + } +} + +// Each step advances from exactly one state, a restart kills the outstanding code, +// and the redeem moves the source's live servers to the named target, retires the +// source and ends its sessions, in one go and once. +func TestMigrationStateMachine(t *testing.T) { + ctx := context.Background() + src := newUser(t, "user", "mig-src") + dst := newUser(t, "user", "mig-dst") + bystander := newUser(t, "user", "mig-by") + sfx := suffix(t) + liveA, liveB, gone, theirs := "ma-"+sfx, "mb-"+sfx, "mgone-"+sfx, "mtheirs-"+sfx + seedOwnedServer(t, liveA, src.ID, false) + seedOwnedServer(t, liveB, src.ID, false) + seedOwnedServer(t, gone, src.ID, true) + seedOwnedServer(t, theirs, bystander.ID, false) + t0 := mustNow().Truncate(time.Second) + srcSession := newSession(t, src.ID, "mig-src", t0.Add(time.Hour)) + dstSession := newSession(t, dst.ID, "mig-dst", t0.Add(time.Hour)) + + if _, err := repo.MigrationForSource(ctx, src.ID); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("status before start = %v, want ErrNotFound", err) + } + if err := repo.ConfirmMigration(ctx, src.ID, "passkey", t0); !errors.Is(err, api.ErrConflict) { + t.Fatalf("confirm before start = %v, want ErrConflict", err) + } + if err := repo.StartMigration(ctx, "mig1-"+sfx, src.ID, t0); err != nil { + t.Fatalf("StartMigration: %v", err) + } + m, err := repo.MigrationForSource(ctx, src.ID) + if err != nil || m.ID != "mig1-"+sfx || m.State != "initiated" || m.TargetUserID != "" || m.ConfirmedAt != nil { + t.Fatalf("after start = %+v, %v; want mig1 initiated, no target, unconfirmed", m, err) + } + if err := repo.IssueMigrationCode(ctx, src.ID, dst.ID, "h-skip", t0.Add(10*time.Minute)); !errors.Is(err, api.ErrConflict) { + t.Fatalf("issue before confirm = %v, want ErrConflict", err) + } + if err := repo.ConfirmMigration(ctx, src.ID, "email_otp", t0); err != nil { + t.Fatalf("ConfirmMigration: %v", err) + } + if err := repo.ConfirmMigration(ctx, src.ID, "passkey", t0.Add(time.Minute)); !errors.Is(err, api.ErrConflict) { + t.Fatalf("second confirm = %v, want ErrConflict", err) + } + m, err = repo.MigrationForSource(ctx, src.ID) + if err != nil || m.State != "confirmed" || m.ConfirmFactor != "email_otp" || m.ConfirmedAt == nil || !m.ConfirmedAt.Equal(t0) { + t.Fatalf("after confirm = %+v, %v; want confirmed by email_otp at %v", m, err, t0) + } + if err := repo.IssueMigrationCode(ctx, src.ID, dst.ID, "h-first", t0.Add(10*time.Minute)); err != nil { + t.Fatalf("IssueMigrationCode: %v", err) + } + if err := repo.IssueMigrationCode(ctx, src.ID, bystander.ID, "h-again", t0.Add(10*time.Minute)); !errors.Is(err, api.ErrConflict) { + t.Fatalf("second issue = %v, want ErrConflict", err) + } + m, err = repo.MigrationForSource(ctx, src.ID) + if err != nil || m.State != "code_issued" || m.TargetUserID != dst.ID || m.CodeExpiresAt == nil || !m.CodeExpiresAt.Equal(t0.Add(10*time.Minute)) { + t.Fatalf("after issue = %+v, %v; want code_issued for the target, expiring at +10m", m, err) + } + + // Restarting supersedes the whole attempt: the code it issued is dead. + if err := repo.StartMigration(ctx, "mig2-"+sfx, src.ID, t0); err != nil { + t.Fatalf("restart: %v", err) + } + if m, err := repo.MigrationForSource(ctx, src.ID); err != nil || m.ID != "mig2-"+sfx || m.State != "initiated" { + t.Fatalf("after restart = %+v, %v; want mig2 initiated", m, err) + } + if n := liveMigrations(t, src.ID); n != 1 { + t.Fatalf("live migrations after restart = %d, want 1", n) + } + if _, _, err := repo.RedeemMigration(ctx, dst.ID, "h-first", t0); !errors.Is(err, api.ErrLinkCodeInvalid) { + t.Fatalf("code from the superseded attempt = %v, want ErrLinkCodeInvalid", err) + } + + if err := repo.ConfirmMigration(ctx, src.ID, "passkey", t0); err != nil { + t.Fatalf("confirm the restart: %v", err) + } + if err := repo.IssueMigrationCode(ctx, src.ID, dst.ID, "h-second", t0.Add(10*time.Minute)); err != nil { + t.Fatalf("issue the restart: %v", err) + } + for _, bad := range []struct { + what, target, hash string + at time.Time + }{ + {"another user with the code", bystander.ID, "h-second", t0}, + {"the target with a wrong code", dst.ID, "h-wrong", t0}, + {"the target at code expiry", dst.ID, "h-second", t0.Add(10 * time.Minute)}, + {"the source itself", src.ID, "h-second", t0}, + } { + if _, _, err := repo.RedeemMigration(ctx, bad.target, bad.hash, bad.at); !errors.Is(err, api.ErrLinkCodeInvalid) { + t.Fatalf("redeem by %s = %v, want ErrLinkCodeInvalid", bad.what, err) + } + } + if owner := serverOwner(t, liveA); owner != src.ID { + t.Fatalf("a refused redeem moved %s to %s", liveA, owner) + } + + source, moved, err := repo.RedeemMigration(ctx, dst.ID, "h-second", t0.Add(10*time.Minute-time.Second)) + if err != nil { + t.Fatalf("RedeemMigration: %v", err) + } + sort.Strings(moved) + if source != src.ID || strings.Join(moved, ",") != liveA+","+liveB { + t.Fatalf("redeem = %s, %v; want %s, [%s %s]", source, moved, src.ID, liveA, liveB) + } + for name, want := range map[string]string{liveA: dst.ID, liveB: dst.ID, gone: src.ID, theirs: bystander.ID} { + if got := serverOwner(t, name); got != want { + t.Fatalf("owner of %s = %s, want %s", name, got, want) + } + } + var disabled, deleted bool + if err := db.QueryRow(`SELECT disabled, deleted_at IS NOT NULL FROM users WHERE id = $1`, src.ID).Scan(&disabled, &deleted); err != nil { + t.Fatalf("read source: %v", err) + } + if !disabled || !deleted { + t.Fatalf("source after redeem: disabled=%v deleted=%v, want both", disabled, deleted) + } + // The session is revoked outright, beyond failing through the dead account. + var srcRevoked bool + if err := db.QueryRow(`SELECT revoked_at IS NOT NULL FROM sessions WHERE token_hash = $1`, srcSession).Scan(&srcRevoked); err != nil { + t.Fatalf("read source session: %v", err) + } + if !srcRevoked { + t.Fatalf("source session after redeem is not revoked") + } + if _, err := repo.SessionUser(ctx, dstSession, t0); err != nil { + t.Fatalf("target session after redeem: %v", err) + } + var state string + var redeemedAt sql.NullTime + if err := db.QueryRow(`SELECT state, redeemed_at FROM account_migrations WHERE id = $1`, "mig2-"+sfx).Scan(&state, &redeemedAt); err != nil { + t.Fatalf("read migration: %v", err) + } + if state != "redeemed" || !redeemedAt.Valid || !redeemedAt.Time.Equal(t0.Add(10*time.Minute-time.Second)) { + t.Fatalf("migration after redeem: %s at %v, want redeemed at %v", state, redeemedAt, t0.Add(10*time.Minute-time.Second)) + } + + if _, _, err := repo.RedeemMigration(ctx, dst.ID, "h-second", t0); !errors.Is(err, api.ErrLinkCodeInvalid) { + t.Fatalf("replayed redeem = %v, want ErrLinkCodeInvalid", err) + } + if _, err := repo.MigrationForSource(ctx, src.ID); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("status after redeem = %v, want ErrNotFound", err) + } + if err := repo.StartMigration(ctx, "mig3-"+sfx, src.ID, t0); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("retired source starting again = %v, want ErrNotFound", err) + } +} + +// Racing redeems of one code move the servers once; racing starts for one source +// all succeed and leave one live attempt. +func TestMigrationUnderConcurrency(t *testing.T) { + ctx := context.Background() + now := mustNow() + for round := 0; round < 5; round++ { + src := newUser(t, "user", "migr-src") + dst := newUser(t, "user", "migr-dst") + srv := "mr-" + suffix(t) + seedOwnedServer(t, srv, src.ID, false) + startToCode(t, src.ID, dst.ID, "h-race", now, now.Add(10*time.Minute)) + + var wg sync.WaitGroup + errs := make([]error, 4) + moved := make([][]string, 4) + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + _, moved[i], errs[i] = repo.RedeemMigration(ctx, dst.ID, "h-race", now) + }(i) + } + wg.Wait() + won, lost := 0, 0 + for i, err := range errs { + switch { + case err == nil && len(moved[i]) == 1 && moved[i][0] == srv: + won++ + case errors.Is(err, api.ErrLinkCodeInvalid): + lost++ + default: + t.Fatalf("round %d: redeem %d = %v, %v", round, i, moved[i], err) + } + } + if won != 1 || lost != 3 { + t.Fatalf("round %d: racing redeems: %d won, %d refused; want 1, 3", round, won, lost) + } + if owner := serverOwner(t, srv); owner != dst.ID { + t.Fatalf("round %d: server owner = %s, want the target", round, owner) + } + } + + src := newUser(t, "user", "migs-src") + var wg sync.WaitGroup + errs := make([]error, 8) + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + errs[i] = repo.StartMigration(ctx, "migs-"+suffix(t), src.ID, now) + }(i) + } + wg.Wait() + for i, err := range errs { + if err != nil { + t.Fatalf("racing start %d: %v", i, err) + } + } + if n := liveMigrations(t, src.ID); n != 1 { + t.Fatalf("live migrations after 8 racing starts = %d, want 1", n) + } +} + +// waitForLockWait polls until a statement starting with prefix is waiting on a lock. +func waitForLockWait(t *testing.T, prefix string) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + var n int + if err := db.QueryRow(`SELECT count(*) FROM pg_stat_activity + WHERE datname = current_database() AND wait_event_type = 'Lock' AND ltrim(query) LIKE $1 || '%'`, + prefix).Scan(&n); err != nil { + t.Fatalf("read pg_stat_activity: %v", err) + } + if n > 0 { + return + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("no statement %q waited on a lock within 5s", prefix) +} + +// A restart that reaches the source while its redeem is in flight must not leave the +// retired account a live attempt: the redeem holds the migration row, the restart's +// supersede waits on it, and by the time the restart writes, the source is gone. +func TestStartMigrationBehindARedeem(t *testing.T) { + ctx := context.Background() + now := mustNow() + src := newUser(t, "user", "migb-src") + dst := newUser(t, "user", "migb-dst") + srv := "mbr-" + suffix(t) + seedOwnedServer(t, srv, src.ID, false) + startToCode(t, src.ID, dst.ID, "h-behind", now, now.Add(10*time.Minute)) + + // Hold the source's server so the redeem stops after locking the migration row. + hold, err := db.BeginTx(ctx, nil) + if err != nil { + t.Fatalf("begin: %v", err) + } + defer hold.Rollback() //nolint:errcheck // no-op after commit + if _, err := hold.Exec(`SELECT 1 FROM servers WHERE name = $1 FOR UPDATE`, srv); err != nil { + t.Fatalf("hold server: %v", err) + } + redeemed := make(chan error, 1) + go func() { + _, _, err := repo.RedeemMigration(ctx, dst.ID, "h-behind", now) + redeemed <- err + }() + waitForLockWait(t, "UPDATE servers SET owner_id") + started := make(chan error, 1) + go func() { started <- repo.StartMigration(ctx, "migb-"+suffix(t), src.ID, now) }() + waitForLockWait(t, "DELETE FROM account_migrations") + if err := hold.Commit(); err != nil { + t.Fatalf("release server: %v", err) + } + if err := <-redeemed; err != nil { + t.Fatalf("redeem: %v", err) + } + if err := <-started; !errors.Is(err, api.ErrNotFound) { + t.Fatalf("restart behind the redeem = %v, want ErrNotFound", err) + } + if n := liveMigrations(t, src.ID); n != 0 { + t.Fatalf("live migrations for the retired source = %d, want 0", n) + } +} + +// ---- setup tokens (migration 0012) -------------------------------------------------- + +// A setup token redeems once, never at or after expiry, and racing redeems of one +// token yield one session. +func TestConsumeSetupTokenContract(t *testing.T) { + ctx := context.Background() + u := newUser(t, "user", "setup") + t0 := mustNow().Truncate(time.Second) + create := func(hash string) { + t.Helper() + if err := repo.CreateSetupToken(ctx, hash, u.ID, t0.Add(10*time.Minute)); err != nil { + t.Fatalf("CreateSetupToken: %v", err) + } + } + once := "st-once-" + suffix(t) + create(once) + if got, err := repo.ConsumeSetupToken(ctx, once, t0); err != nil || got != u.ID { + t.Fatalf("redeem = %q, %v; want %s", got, err, u.ID) + } + if _, err := repo.ConsumeSetupToken(ctx, once, t0); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("replay = %v, want ErrNotFound", err) + } + if _, err := repo.ConsumeSetupToken(ctx, "st-never-"+suffix(t), t0); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("unknown token = %v, want ErrNotFound", err) + } + + late := "st-late-" + suffix(t) + create(late) + if _, err := repo.ConsumeSetupToken(ctx, late, t0.Add(10*time.Minute)); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("redeem at expiry = %v, want ErrNotFound", err) + } + if got, err := repo.ConsumeSetupToken(ctx, late, t0.Add(10*time.Minute-time.Second)); err != nil || got != u.ID { + t.Fatalf("redeem a second before expiry = %q, %v; want %s (the refused try must not spend it)", got, err, u.ID) + } + + raced := "st-race-" + suffix(t) + create(raced) + var wg sync.WaitGroup + errs := make([]error, 8) + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + _, errs[i] = repo.ConsumeSetupToken(ctx, raced, t0) + }(i) + } + wg.Wait() + won, lost := 0, 0 + for i, err := range errs { + switch { + case err == nil: + won++ + case errors.Is(err, api.ErrNotFound): + lost++ + default: + t.Fatalf("racing redeem %d: %v", i, err) + } + } + if won != 1 || lost != 7 { + t.Fatalf("racing redeems: %d won, %d refused; want 1, 7", won, lost) + } +} + +// ---- remediation: end every session, drop every passkey ----------------------------- + +func TestRevokeAllUserSessionsIsScoped(t *testing.T) { + ctx := context.Background() + now := mustNow() + u := newUser(t, "user", "rall") + other := newUser(t, "user", "rall-other") + a := newSession(t, u.ID, "rall-a", now.Add(time.Hour)) + b := newSession(t, u.ID, "rall-b", now.Add(time.Hour)) + old := newSession(t, u.ID, "rall-old", now.Add(time.Hour)) + theirs := newSession(t, other.ID, "rall-theirs", now.Add(time.Hour)) + earlier := now.Add(-time.Hour).Truncate(time.Microsecond) + mustExec(t, `UPDATE sessions SET revoked_at = $2 WHERE token_hash = $1`, old, earlier) + + if err := repo.RevokeAllUserSessions(ctx, u.ID); err != nil { + t.Fatalf("RevokeAllUserSessions: %v", err) + } + for _, h := range []string{a, b} { + if _, err := repo.SessionUser(ctx, h, now); !errors.Is(err, api.ErrNotFound) { + t.Fatalf("session %s after revoke-all = %v, want ErrNotFound", h, err) + } + } + if _, err := repo.SessionUser(ctx, theirs, now); err != nil { + t.Fatalf("another account's session after revoke-all: %v", err) + } + var revokedAt time.Time + if err := db.QueryRow(`SELECT revoked_at FROM sessions WHERE token_hash = $1`, old).Scan(&revokedAt); err != nil { + t.Fatalf("read old session: %v", err) + } + if !sameMicro(revokedAt, earlier) { + t.Fatalf("an already-revoked session's revoked_at moved from %v to %v", earlier, revokedAt) + } + if err := repo.RevokeAllUserSessions(ctx, newUser(t, "user", "rall-none").ID); err != nil { + t.Fatalf("revoke-all with no sessions = %v, want nil", err) + } +} + +func TestDeleteAllPasskeyCredentialsIsScoped(t *testing.T) { + ctx := context.Background() + seed := func(userID string) { + t.Helper() + mustExec(t, `INSERT INTO webauthn_credentials (id, user_id, credential_id, public_key) VALUES ($1, $2, $3, 'pk')`, + "cred-"+suffix(t), userID, "cid-"+suffix(t)) + } + count := func(userID string) int { + t.Helper() + var n int + if err := db.QueryRow(`SELECT count(*) FROM webauthn_credentials WHERE user_id = $1`, userID).Scan(&n); err != nil { + t.Fatalf("count: %v", err) + } + return n + } + u := newUser(t, "user", "pkall") + other := newUser(t, "user", "pkall-other") + seed(u.ID) + seed(u.ID) + seed(other.ID) + if err := repo.DeleteAllPasskeyCredentialsForUser(ctx, u.ID); err != nil { + t.Fatalf("DeleteAllPasskeyCredentialsForUser: %v", err) + } + if got, theirs := count(u.ID), count(other.ID); got != 0 || theirs != 1 { + t.Fatalf("after delete-all: %d left for the account, %d for the other; want 0, 1", got, theirs) + } + if err := repo.DeleteAllPasskeyCredentialsForUser(ctx, u.ID); err != nil { + t.Fatalf("delete-all with none left = %v, want nil", err) + } +} + +// ---- username reclaim (migration 0006) ---------------------------------------------- + +// A reclaim bars the squatter's UUID (never the contested name) and stashes its data +// once: a retry keeps the first window and handle, and a failed stash leaves no bar. +func TestReclaimUsernameContract(t *testing.T) { + ctx := context.Background() + t0 := mustNow().Truncate(time.Second) + squatter, genuine := testUUID(t), testUUID(t) + name := "Notch" + suffix(t)[:6] + holdID := "hold-" + suffix(t) + + got, err := repo.ReclaimUsername(ctx, holdID, squatter, name, "ref-1", t0.Add(30*24*time.Hour)) + if err != nil || !got.Equal(t0.Add(30*24*time.Hour)) { + t.Fatalf("reclaim = %v, %v; want %v", got, err, t0.Add(30*24*time.Hour)) + } + for uuid, want := range map[string]bool{squatter: true, genuine: false} { + if barred, err := repo.IsUsernameBlacklisted(ctx, uuid); err != nil || barred != want { + t.Fatalf("IsUsernameBlacklisted(%s) = %v, %v; want %v", uuid, barred, err, want) + } + } + + got, err = repo.ReclaimUsername(ctx, "hold-retry-"+suffix(t), squatter, name, "ref-2", t0.Add(60*24*time.Hour)) + if err != nil || !got.Equal(t0.Add(30*24*time.Hour)) { + t.Fatalf("retried reclaim = %v, %v; want the first window %v", got, err, t0.Add(30*24*time.Hour)) + } + var holds int + var ref string + if err := db.QueryRow(`SELECT count(*), max(data_ref) FROM player_data_holds WHERE mc_uuid = $1`, squatter).Scan(&holds, &ref); err != nil { + t.Fatalf("read holds: %v", err) + } + if holds != 1 || ref != "ref-1" { + t.Fatalf("holds after a retry: %d with ref %q, want 1 with ref-1", holds, ref) + } + + unarchived := testUUID(t) + if _, err := repo.ReclaimUsername(ctx, "hold-bare-"+suffix(t), unarchived, name, "", t0.Add(30*24*time.Hour)); err != nil { + t.Fatalf("reclaim without a data ref: %v", err) + } + var bare sql.NullString + if err := db.QueryRow(`SELECT data_ref FROM player_data_holds WHERE mc_uuid = $1`, unarchived).Scan(&bare); err != nil { + t.Fatalf("read bare hold: %v", err) + } + if bare.Valid { + t.Fatalf("empty data ref stored as %q, want NULL", bare.String) + } + + // The stash fails (its id is taken), so the bar written before it rolls back. + orphan := testUUID(t) + if _, err := repo.ReclaimUsername(ctx, holdID, orphan, name, "ref-3", t0.Add(30*24*time.Hour)); err == nil { + t.Fatalf("reclaim reusing a hold id = nil, want an error") + } + if barred, err := repo.IsUsernameBlacklisted(ctx, orphan); err != nil || barred { + t.Fatalf("UUID barred by a reclaim whose stash failed: %v, %v; want false", barred, err) + } +} diff --git a/internal/pgint/challenges_test.go b/internal/pgint/challenges_test.go index 4650359..660d1d2 100644 --- a/internal/pgint/challenges_test.go +++ b/internal/pgint/challenges_test.go @@ -5,6 +5,7 @@ package pgint import ( "context" "errors" + "sync" "testing" "time" @@ -285,3 +286,227 @@ func TestDiscoverableChallengeBounds(t *testing.T) { t.Fatalf("begin with the table full = %v, want ErrTooManyPasskeyChallenges", err) } } + +// ---- enrollment and step-up ceremonies (migration 0007) --------------------------- + +func challengeRows(t *testing.T, userID, purpose string) int { + t.Helper() + var n int + if err := db.QueryRow(`SELECT count(*) FROM webauthn_challenges WHERE user_id = $1 AND purpose = $2`, + userID, purpose).Scan(&n); err != nil { + t.Fatalf("count challenges: %v", err) + } + return n +} + +// An enrollment begin supersedes the account's earlier one for the same purpose and +// sweeps the row its last finish spent, so (user, purpose) holds one row; finish is +// single-use, dies at expires_at, and never crosses purposes or accounts. +func TestPasskeyChallengeByUserContract(t *testing.T) { + ctx := context.Background() + u := newUser(t, "user", "pk-reg") + other := newUser(t, "user", "pk-reg-other") + const reg = "passkey_register" + t0 := mustNow().Truncate(time.Second) + begin := func(userID, purpose, session string, expiresAt time.Time) { + t.Helper() + if err := repo.CreatePasskeyChallenge(ctx, "reg-"+suffix(t), userID, purpose, []byte(session), expiresAt); err != nil { + t.Fatalf("CreatePasskeyChallenge(%s): %v", session, err) + } + } + finish := func(userID, purpose string, at time.Time) (string, error) { + sd, err := repo.ConsumePasskeyChallengeByUser(ctx, userID, purpose, at) + return string(sd), err + } + + if _, err := finish(u.ID, reg, t0); !errors.Is(err, api.ErrPasskeyChallengeInvalid) { + t.Fatalf("finish with no begin = %v, want ErrPasskeyChallengeInvalid", err) + } + begin(u.ID, reg, "s1", t0.Add(5*time.Minute)) + begin(other.ID, reg, "o1", t0.Add(5*time.Minute)) + begin(u.ID, reg, "s2", t0.Add(5*time.Minute)) + if n := challengeRows(t, u.ID, reg); n != 1 { + t.Fatalf("rows after two begins = %d, want 1", n) + } + if sd, err := finish(u.ID, reg, t0); err != nil || sd != "s2" { + t.Fatalf("finish = %q, %v; want s2", sd, err) + } + if _, err := finish(u.ID, reg, t0); !errors.Is(err, api.ErrPasskeyChallengeInvalid) { + t.Fatalf("replayed finish = %v, want ErrPasskeyChallengeInvalid", err) + } + if sd, err := finish(other.ID, reg, t0); err != nil || sd != "o1" { + t.Fatalf("other account's finish = %q, %v; want o1 (a begin never touches another account)", sd, err) + } + + // The spent row goes with the next begin. + begin(u.ID, reg, "s3", t0.Add(5*time.Minute)) + if n := challengeRows(t, u.ID, reg); n != 1 { + t.Fatalf("rows after finish then begin = %d, want 1", n) + } + + // A step-up ceremony is its own door. + begin(u.ID, "reauth", "r1", t0.Add(5*time.Minute)) + if n := challengeRows(t, u.ID, reg); n != 1 { + t.Fatalf("a reauth begin touched the enrollment row: %d rows, want 1", n) + } + if sd, err := finish(u.ID, "reauth", t0); err != nil || sd != "r1" { + t.Fatalf("reauth finish = %q, %v; want r1", sd, err) + } + if _, err := finish(u.ID, "reauth", t0); !errors.Is(err, api.ErrPasskeyChallengeInvalid) { + t.Fatalf("reauth replay = %v, want ErrPasskeyChallengeInvalid", err) + } + + // Expiry: at expires_at the ceremony is dead; a second earlier it still redeems. + if _, err := finish(u.ID, reg, t0.Add(5*time.Minute)); !errors.Is(err, api.ErrPasskeyChallengeInvalid) { + t.Fatalf("finish at expires_at = %v, want ErrPasskeyChallengeInvalid", err) + } + if sd, err := finish(u.ID, reg, t0.Add(5*time.Minute-time.Second)); err != nil || sd != "s3" { + t.Fatalf("finish a second before expiry = %q, %v; want s3", sd, err) + } +} + +// Begins that race for one (user, purpose) leave exactly one row, and finishes that +// race for one ceremony redeem it once. +func TestPasskeyChallengeByUserUnderConcurrency(t *testing.T) { + ctx := context.Background() + const reg = "passkey_register" + now := mustNow() + for round := 0; round < 5; round++ { + u := newUser(t, "user", "pk-race") + var wg sync.WaitGroup + errs := make([]error, 8) + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + errs[i] = repo.CreatePasskeyChallenge(ctx, "race-"+suffix(t), u.ID, reg, []byte("s"), now.Add(5*time.Minute)) + }(i) + } + wg.Wait() + for i, err := range errs { + if err != nil { + t.Fatalf("round %d: begin %d: %v", round, i, err) + } + } + if n := challengeRows(t, u.ID, reg); n != 1 { + t.Fatalf("round %d: rows after 8 racing begins = %d, want 1", round, n) + } + + won, lost := 0, 0 + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + _, errs[i] = repo.ConsumePasskeyChallengeByUser(ctx, u.ID, reg, now) + }(i) + } + wg.Wait() + for i, err := range errs { + switch { + case err == nil: + won++ + case errors.Is(err, api.ErrPasskeyChallengeInvalid): + lost++ + default: + t.Fatalf("round %d: finish %d: %v", round, i, err) + } + } + if won != 1 || lost != 7 { + t.Fatalf("round %d: racing finishes: %d redeemed, %d refused; want 1, 7", round, won, lost) + } + } +} + +// Racing begins from one source still stop at 32 live challenges, in both login +// stores: the count and the insert are one decision. +func TestLoginChallengeSourceBoundUnderConcurrency(t *testing.T) { + ctx := context.Background() + now := mustNow() + race := func(t *testing.T, begin func(i int) error) (ok, refused int) { + t.Helper() + var wg sync.WaitGroup + errs := make([]error, 64) + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + errs[i] = begin(i) + }(i) + } + wg.Wait() + for i, err := range errs { + switch { + case err == nil: + ok++ + case errors.Is(err, api.ErrTooManyPasskeyChallenges): + refused++ + default: + t.Fatalf("begin %d: %v", i, err) + } + } + return ok, refused + } + + t.Run("discoverable", func(t *testing.T) { + source := "2001:db8:" + suffix(t)[:4] + "::/48" + t.Cleanup(func() { + db.Exec(`DELETE FROM webauthn_discoverable_challenges WHERE source = $1`, source) //nolint:errcheck + }) + ok, refused := race(t, func(int) error { + return repo.CreateDiscoverableChallenge(ctx, "drace-"+suffix(t), source, []byte("s"), now, now.Add(5*time.Minute)) + }) + var stored int + if err := db.QueryRow(`SELECT count(*) FROM webauthn_discoverable_challenges WHERE source = $1`, source).Scan(&stored); err != nil { + t.Fatalf("count: %v", err) + } + if ok != 32 || refused != 32 || stored != 32 { + t.Fatalf("64 racing begins: %d stored, %d refused, %d rows; want 32, 32, 32", ok, refused, stored) + } + }) + + t.Run("email-first", func(t *testing.T) { + u := newUser(t, "user", "pk-login-race") + source := "203.0.113." + suffix(t) + ok, refused := race(t, func(int) error { + return repo.AddPasskeyLoginChallenge(ctx, "lrace-"+suffix(t), u.ID, "passkey_login", source, "c-"+suffix(t), []byte("s"), now, now.Add(5*time.Minute)) + }) + var stored int + if err := db.QueryRow(`SELECT count(*) FROM webauthn_challenges WHERE source = $1`, source).Scan(&stored); err != nil { + t.Fatalf("count: %v", err) + } + if ok != 32 || refused != 32 || stored != 32 { + t.Fatalf("64 racing begins: %d stored, %d refused, %d rows; want 32, 32, 32", ok, refused, stored) + } + }) + + t.Run("discoverable finish", func(t *testing.T) { + id := "dfin-" + suffix(t) + if err := repo.CreateDiscoverableChallenge(ctx, id, "pgint-fin-"+suffix(t), []byte("s"), now, now.Add(5*time.Minute)); err != nil { + t.Fatalf("begin: %v", err) + } + var wg sync.WaitGroup + errs := make([]error, 8) + for i := range errs { + wg.Add(1) + go func(i int) { + defer wg.Done() + _, errs[i] = repo.ConsumeDiscoverableChallenge(ctx, id, now) + }(i) + } + wg.Wait() + won, lost := 0, 0 + for i, err := range errs { + switch { + case err == nil: + won++ + case errors.Is(err, api.ErrPasskeyChallengeInvalid): + lost++ + default: + t.Fatalf("finish %d: %v", i, err) + } + } + if won != 1 || lost != 7 { + t.Fatalf("racing finishes: %d redeemed, %d refused; want 1, 7", won, lost) + } + }) +}