Files
Felis/internal/pgint/challenges_test.go
T

513 lines
20 KiB
Go

//go:build pgint
package pgint
import (
"context"
"errors"
"sync"
"testing"
"time"
"felis.lolicon.best/internal/api"
)
// ---- login codes and ceremonies that coexist (migration 0029) --------------------
func liveLoginCodes(t *testing.T, userID, purpose string, now time.Time) int {
t.Helper()
var n int
if err := db.QueryRow(`SELECT count(*) FROM email_otps
WHERE user_id = $1 AND purpose = $2 AND consumed_at IS NULL AND expires_at > $3`,
userID, purpose, now).Scan(&n); err != nil {
t.Fatalf("count live codes: %v", err)
}
return n
}
// A login start keeps the three newest live codes: a fourth drops the oldest, and
// redeeming any live one spends its siblings too.
func TestAddLoginEmailOTPKeepsNewestThree(t *testing.T) {
ctx := context.Background()
u := newUser(t, "user", "otp-keep")
purpose := "login_email"
addr := "keep-" + suffix(t) + "@example.net"
t0 := mustNow().Truncate(time.Second)
add := func(hash string, at time.Time) {
t.Helper()
if err := repo.AddLoginEmailOTP(ctx, "keep-"+suffix(t), u.ID, addr, hash, purpose, at, at.Add(10*time.Minute)); err != nil {
t.Fatalf("AddLoginEmailOTP(%s): %v", hash, err)
}
}
add("h1", t0)
add("h2", t0.Add(time.Minute))
add("h3", t0.Add(2*time.Minute))
if n := liveLoginCodes(t, u.ID, purpose, t0.Add(2*time.Minute)); n != 3 {
t.Fatalf("live codes after three starts = %d, want 3", n)
}
add("h4", t0.Add(3*time.Minute))
at := t0.Add(3 * time.Minute)
if n := liveLoginCodes(t, u.ID, purpose, at); n != 3 {
t.Fatalf("live codes after four starts = %d, want 3", n)
}
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h1", at); !errors.Is(err, api.ErrOTPInvalid) {
t.Fatalf("oldest code after a fourth start = %v, want ErrOTPInvalid", err)
}
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h2", at); err != nil {
t.Fatalf("second code = %v, want nil", err)
}
for _, h := range []string{"h3", "h4"} {
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, h, at); !errors.Is(err, api.ErrOTPInvalid) {
t.Fatalf("sibling %s after a sign-in = %v, want ErrOTPInvalid", h, err)
}
}
// Expired codes leave the allowance: a start after they lapse is the only live one.
add("h5", t0.Add(4*time.Minute))
add("h6", t0.Add(5*time.Minute))
later := t0.Add(20 * time.Minute)
add("h7", later)
if n := liveLoginCodes(t, u.ID, purpose, later); n != 1 {
t.Fatalf("live codes after the others expired = %d, want 1", n)
}
var unconsumed int
if err := db.QueryRow(`SELECT count(*) FROM email_otps WHERE user_id = $1 AND purpose = $2 AND consumed_at IS NULL`,
u.ID, purpose).Scan(&unconsumed); err != nil {
t.Fatalf("count unconsumed: %v", err)
}
if unconsumed != 1 {
t.Fatalf("unconsumed rows = %d, want 1 (expired codes are deleted, not kept)", unconsumed)
}
}
// A wrong guess with several live codes costs each open code one attempt and the
// account one failure; when every live code is spent the door answers ErrOTPLocked,
// and a newer code still signs in.
func TestConsumeLoginEmailOTPAcrossLiveCodes(t *testing.T) {
ctx := context.Background()
u := newUser(t, "user", "otp-multi")
purpose := "op_login"
addr := "multi-" + suffix(t) + "@example.net"
t0 := mustNow().Truncate(time.Second)
add := func(hash string, at time.Time) {
t.Helper()
if err := repo.AddLoginEmailOTP(ctx, "multi-"+suffix(t), u.ID, addr, hash, purpose, at, at.Add(10*time.Minute)); err != nil {
t.Fatalf("AddLoginEmailOTP(%s): %v", hash, err)
}
}
add("h-a", t0)
add("h-b", t0.Add(time.Minute))
at := t0.Add(time.Minute)
attempts := func() map[string]int {
t.Helper()
rows, err := db.Query(`SELECT code_hash, attempts FROM email_otps WHERE user_id = $1 AND purpose = $2`, u.ID, purpose)
if err != nil {
t.Fatalf("read attempts: %v", err)
}
defer rows.Close()
got := map[string]int{}
for rows.Next() {
var h string
var n int
if err := rows.Scan(&h, &n); err != nil {
t.Fatalf("scan: %v", err)
}
got[h] = n
}
return got
}
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h-wrong", at); !errors.Is(err, api.ErrOTPInvalid) {
t.Fatalf("wrong guess = %v, want ErrOTPInvalid", err)
}
if got := attempts(); got["h-a"] != 1 || got["h-b"] != 1 {
t.Fatalf("attempts after one wrong guess = %v, want h-a:1 h-b:1", got)
}
for i := 2; i <= 5; i++ {
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h-wrong", at); !errors.Is(err, api.ErrOTPInvalid) {
t.Fatalf("wrong guess %d = %v, want ErrOTPInvalid", i, err)
}
}
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h-a", at); !errors.Is(err, api.ErrOTPLocked) {
t.Fatalf("right code once every live code is spent = %v, want ErrOTPLocked", err)
}
var failures int
if err := db.QueryRow(`SELECT failures FROM otp_failure_windows WHERE user_id = $1 AND purpose = $2`,
u.ID, purpose).Scan(&failures); err != nil {
t.Fatalf("read budget row: %v", err)
}
if failures != 5 {
t.Fatalf("account failures = %d, want 5 (one per wrong guess, not one per code)", failures)
}
add("h-c", t0.Add(2*time.Minute))
at = t0.Add(2 * time.Minute)
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h-b", at); !errors.Is(err, api.ErrOTPInvalid) {
t.Fatalf("spent code beside a fresh one = %v, want ErrOTPInvalid", err)
}
if got := attempts(); got["h-c"] != 1 || got["h-a"] != 5 || got["h-b"] != 5 {
t.Fatalf("attempts after a guess of a spent code = %v, want h-c:1 and the spent codes left at 5", got)
}
if err := repo.ConsumeLoginEmailOTP(ctx, u.ID, purpose, "h-c", at); err != nil {
t.Fatalf("fresh code = %v, want nil", err)
}
if n := liveLoginCodes(t, u.ID, purpose, at); n != 0 {
t.Fatalf("live codes after a sign-in = %d, want 0", n)
}
}
// Email-first passkey ceremonies of one account coexist and are redeemed by the
// challenge the browser signed; one network holds at most 32 live ones.
func TestPasskeyLoginChallengesCoexist(t *testing.T) {
ctx := context.Background()
u := newUser(t, "user", "pk-login")
purpose := "passkey_login"
source := "203.0.113." + suffix(t)
now := mustNow()
add := func(id, source, challenge, session string, expiresAt time.Time) error {
return repo.AddPasskeyLoginChallenge(ctx, id+"-"+suffix(t), u.ID, purpose, source, challenge, []byte(session), now, expiresAt)
}
if err := add("c1", source, "Y2hhbGxlbmdlMQ", "s1", now.Add(5*time.Minute)); err != nil {
t.Fatalf("add c1: %v", err)
}
if err := add("c2", source, "Y2hhbGxlbmdlMg", "s2", now.Add(5*time.Minute)); err != nil {
t.Fatalf("add c2: %v", err)
}
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, "bm9ib2R5", now); !errors.Is(err, api.ErrPasskeyChallengeInvalid) {
t.Fatalf("unissued challenge = %v, want ErrPasskeyChallengeInvalid", err)
}
if sd, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, "Y2hhbGxlbmdlMQ", now); err != nil || string(sd) != "s1" {
t.Fatalf("earlier ceremony = %q, %v; want s1 (a later begin must not cancel it)", sd, err)
}
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, "Y2hhbGxlbmdlMQ", now); !errors.Is(err, api.ErrPasskeyChallengeInvalid) {
t.Fatalf("replay = %v, want ErrPasskeyChallengeInvalid", err)
}
if sd, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, "Y2hhbGxlbmdlMg", now); err != nil || string(sd) != "s2" {
t.Fatalf("later ceremony = %q, %v; want s2", sd, err)
}
// The same challenge under another purpose is a different door.
if err := add("c3", source, "Y2hhbGxlbmdlMw", "s3", now.Add(5*time.Minute)); err != nil {
t.Fatalf("add c3: %v", err)
}
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, "passkey_register", "Y2hhbGxlbmdlMw", now); !errors.Is(err, api.ErrPasskeyChallengeInvalid) {
t.Fatalf("other purpose = %v, want ErrPasskeyChallengeInvalid", err)
}
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, "Y2hhbGxlbmdlMw", now.Add(6*time.Minute)); !errors.Is(err, api.ErrPasskeyChallengeInvalid) {
t.Fatalf("expired ceremony = %v, want ErrPasskeyChallengeInvalid", err)
}
// Other accounts' spent rows from the same source hold no slot: one expired, one
// redeemed. Two accounts, since an account's own begin reaps its spent rows.
other := newUser(t, "user", "pk-login-other")
if err := repo.AddPasskeyLoginChallenge(ctx, "ox-"+suffix(t), other.ID, purpose, source, "ZXhwaXJlZA", []byte("s"), now, now.Add(-time.Second)); err != nil {
t.Fatalf("other account's expired begin: %v", err)
}
third := newUser(t, "user", "pk-login-third")
if err := repo.AddPasskeyLoginChallenge(ctx, "oc-"+suffix(t), third.ID, purpose, source, "cmVkZWVtZWQ", []byte("s"), now, now.Add(5*time.Minute)); err != nil {
t.Fatalf("third account's begin: %v", err)
}
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, third.ID, purpose, "cmVkZWVtZWQ", now); err != nil {
t.Fatalf("third account's finish: %v", err)
}
// Per-source bound: c3 is still live at now, so 31 more fill the allowance.
for i := 0; i < 31; i++ {
if err := add("cap", source, "Y2Fw", "s", now.Add(5*time.Minute)); err != nil {
t.Fatalf("add %d from the source: %v", i+2, err)
}
}
if err := add("over", source, "b3Zlcg", "s", now.Add(5*time.Minute)); !errors.Is(err, api.ErrTooManyPasskeyChallenges) {
t.Fatalf("33rd live challenge from one source = %v, want ErrTooManyPasskeyChallenges", err)
}
if err := repo.AddPasskeyLoginChallenge(ctx, "x-"+suffix(t), other.ID, purpose, source, "eA", []byte("s"), now, now.Add(5*time.Minute)); !errors.Is(err, api.ErrTooManyPasskeyChallenges) {
t.Fatalf("another account's begin from the full source = %v, want ErrTooManyPasskeyChallenges", err)
}
if err := add("elsewhere", "198.51.100."+suffix(t), "ZWxzZXdoZXJl", "s", now.Add(5*time.Minute)); err != nil {
t.Fatalf("begin from another source: %v", err)
}
// A redeemed ceremony frees its slot.
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, "b3Zlcg", now); !errors.Is(err, api.ErrPasskeyChallengeInvalid) {
t.Fatalf("the refused begin left a row: %v", err)
}
var capID string
if err := db.QueryRow(`SELECT challenge FROM webauthn_challenges WHERE user_id = $1 AND source = $2 AND consumed_at IS NULL LIMIT 1`,
u.ID, source).Scan(&capID); err != nil {
t.Fatalf("pick a live challenge: %v", err)
}
if _, err := repo.ConsumePasskeyLoginChallenge(ctx, u.ID, purpose, capID, now); err != nil {
t.Fatalf("redeem one from the full source: %v", err)
}
if err := add("after", source, "YWZ0ZXI", "s", now.Add(5*time.Minute)); err != nil {
t.Fatalf("begin after one was redeemed = %v, want nil", err)
}
}
// The usernameless store holds each source to 32 live challenges and the whole
// table to 16384.
func TestDiscoverableChallengeBounds(t *testing.T) {
ctx := context.Background()
now := mustNow()
source := "2001:db8:" + suffix(t)[:4] + "::/48"
t.Cleanup(func() {
db.Exec(`DELETE FROM webauthn_discoverable_challenges WHERE source LIKE 'pgint-bulk-%' OR source = $1`, source) //nolint:errcheck
})
var ids []string
for i := 0; i < 32; i++ {
id := "disc-" + suffix(t)
if err := repo.CreateDiscoverableChallenge(ctx, id, source, []byte("s"), now, now.Add(5*time.Minute)); err != nil {
t.Fatalf("create %d: %v", i+1, err)
}
ids = append(ids, id)
}
if err := repo.CreateDiscoverableChallenge(ctx, "disc-over-"+suffix(t), source, []byte("s"), now, now.Add(5*time.Minute)); !errors.Is(err, api.ErrTooManyPasskeyChallenges) {
t.Fatalf("33rd from one source = %v, want ErrTooManyPasskeyChallenges", err)
}
if err := repo.CreateDiscoverableChallenge(ctx, "disc-else-"+suffix(t), "pgint-bulk-else", []byte("s"), now, now.Add(5*time.Minute)); err != nil {
t.Fatalf("another source: %v", err)
}
if _, err := repo.ConsumeDiscoverableChallenge(ctx, ids[0], now); err != nil {
t.Fatalf("consume: %v", err)
}
if err := repo.CreateDiscoverableChallenge(ctx, "disc-after-"+suffix(t), source, []byte("s"), now, now.Add(5*time.Minute)); err != nil {
t.Fatalf("after a redeem freed a slot = %v, want nil", err)
}
// Fill the table to its store-wide bound from many sources, none of them full.
var live int
if err := db.QueryRow(`SELECT count(*) FROM webauthn_discoverable_challenges WHERE consumed_at IS NULL AND expires_at > $1`, now).Scan(&live); err != nil {
t.Fatalf("count: %v", err)
}
if _, err := db.Exec(`INSERT INTO webauthn_discoverable_challenges (id, session_data, expires_at, source)
SELECT 'bulk-' || $1 || '-' || i, '\x00'::bytea, $2, 'pgint-bulk-' || (i % 1000)
FROM generate_series(1, $3) AS i`, suffix(t), now.Add(5*time.Minute), 16384-live); err != nil {
t.Fatalf("bulk insert: %v", err)
}
if err := repo.CreateDiscoverableChallenge(ctx, "disc-full-"+suffix(t), "pgint-bulk-fresh", []byte("s"), now, now.Add(5*time.Minute)); !errors.Is(err, api.ErrTooManyPasskeyChallenges) {
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)
}
})
}