Loading internal/api/api.go +10 −1 Changes for internal/api/api.go: 10 added lines, 1 removed line. Original line number Diff line number Diff line Loading @@ -14,6 +14,7 @@ package api import ( "context" "log" "log/slog" "net/http" "strings" Loading Loading @@ -776,7 +777,15 @@ func (a *API) requireOnboarded(h http.HandlerFunc) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { p := principalFromContext(r.Context()) if p != nil && p.ViaSession && !p.EmailVerified { creds, _ := a.Repo.PasskeyCredentialsForUser(r.Context(), p.UserID) creds, err := a.Repo.PasskeyCredentialsForUser(r.Context(), p.UserID) if err != nil { // A store outage reads as retry-later; setup_required would send // the caller off to enroll a passkey they may already have. log.Printf("api: %s %s: onboarding check (request_id=%s): %v", r.Method, r.URL.Path, requestIDFromContext(r.Context()), err) writeError(w, r, errAuthUnavailable) return } if len(creds) == 0 { writeError(w, r, newError(http.StatusForbidden, "setup_required", "passkey enrollment is required before this action is available")) Loading internal/api/api_test.go +1 −1 Changes for internal/api/api_test.go: 1 added line, 1 removed line. Original line number Diff line number Diff line Loading @@ -583,7 +583,7 @@ func (f *fakeRepo) DeleteAllPasskeyCredentialsForUser(_ context.Context, userID func (f *fakeRepo) AdvanceCredentialSignCount(_ context.Context, credentialID string, newSignCount uint32, usedAt time.Time) error { for id, c := range f.passkeyCreds { if c.CredentialID == credentialID { c.SignCount = newSignCount c.SignCount = max(c.SignCount, newSignCount) t := usedAt c.LastUsedAt = &t f.passkeyCreds[id] = c Loading internal/api/handlers_setup_lockdown_test.go +31 −0 Changes for internal/api/handlers_setup_lockdown_test.go: 31 added lines, 0 removed lines. Original line number Diff line number Diff line Loading @@ -2,6 +2,7 @@ package api import ( "context" "errors" "net/http" "net/http/httptest" "testing" Loading Loading @@ -138,3 +139,33 @@ func TestAdvRequireOnboardedPredicate(t *testing.T) { }) } } // credsDown is a store whose passkey lookup fails, as during a Postgres outage. type credsDown struct{ *fakeRepo } func (credsDown) PasskeyCredentialsForUser(context.Context, string) ([]PasskeyCredential, error) { return nil, errors.New("connection refused") } // A failed passkey lookup is an outage: the caller gets 503 auth_unavailable to // retry, never setup_required, which would send a player with a passkey off to // enroll another. func TestRequireOnboardedReportsAStoreOutage(t *testing.T) { repo := newFakeRepo() api := newTestAPI(repo, newFakeCluster()) api.Repo = credsDown{repo} ran := false h := api.requireOnboarded(func(http.ResponseWriter, *http.Request) { ran = true }) r := httptest.NewRequest("POST", "/api/v1/servers/demo2/wake", nil) r = r.WithContext(context.WithValue(r.Context(), ctxKeyPrincipal, &Principal{UserID: "has-pk", ViaSession: true, EmailVerified: false})) w := httptest.NewRecorder() h(w, r) if ran { t.Fatal("the handler ran although the onboarding check could not be made") } if w.Code != http.StatusServiceUnavailable || errCode(w.Body.Bytes()) != "auth_unavailable" { t.Fatalf("got %d %s, want 503 auth_unavailable", w.Code, w.Body.String()) } } internal/api/pgrepo.go +25 −16 Changes for internal/api/pgrepo.go: 25 added lines, 16 removed lines. Original line number Diff line number Diff line Loading @@ -1401,9 +1401,14 @@ func (p *PGRepo) PasskeyCredentialsForUser(ctx context.Context, userID string) ( // last_used_at. credential_id is UNIQUE so exactly one row is touched; a missing row (the // credential was unbound mid-ceremony) affects zero rows and is a successful no-op, never an // error — the assertion is already cryptographically complete by the time this runs. // // The counter only moves forward: two assertions verified at once against the same stored // value land in either order, and the lower one must not overwrite the higher, or a clone // replaying the count in between would pass the next check. func (p *PGRepo) AdvanceCredentialSignCount(ctx context.Context, credentialID string, newSignCount uint32, usedAt time.Time) error { _, err := p.db.ExecContext(ctx, `UPDATE webauthn_credentials SET sign_count = $2, last_used_at = $3 WHERE credential_id = $1`, `UPDATE webauthn_credentials SET sign_count = GREATEST(sign_count, $2), last_used_at = $3 WHERE credential_id = $1`, credentialID, int64(newSignCount), usedAt) return err } Loading Loading @@ -1738,31 +1743,35 @@ func (p *PGRepo) DeleteUser(ctx context.Context, userID, _ string) error { } // SetUserDisabled flips the disabled flag. Setting disabled→true additionally // revokes every live session so the account is immediately locked out. // revokes every live session so the account is immediately locked out. Both land // in one transaction: a disable whose revocation failed is rolled back and // reported, so the admin never reads "disabled" while a session still stands. func (p *PGRepo) SetUserDisabled(ctx context.Context, userID string, disabled bool) error { // Guard: the user must exist and not be deleted. var ok bool if err := p.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM users WHERE id = $1 AND deleted_at IS NULL)`, userID).Scan(&ok); err != nil { tx, err := p.db.BeginTx(ctx, nil) if err != nil { return err } if !ok { return ErrNotFound } defer func() { _ = tx.Rollback() }() if _, err := p.db.ExecContext(ctx, `UPDATE users SET disabled = $2 WHERE id = $1`, userID, disabled); err != nil { res, err := tx.ExecContext(ctx, `UPDATE users SET disabled = $2 WHERE id = $1 AND deleted_at IS NULL`, userID, disabled) if err != nil { return err } if n, err := res.RowsAffected(); err != nil { return err } else if n == 0 { return ErrNotFound } if disabled { // Revoke every live session so the lockout is immediate. _, _ = p.db.ExecContext(ctx, if _, err := tx.ExecContext(ctx, `UPDATE sessions SET revoked_at = now() WHERE user_id = $1 AND revoked_at IS NULL`, userID) userID); err != nil { return err } return nil } return tx.Commit() } // ---- quota admin ---- Loading internal/pgint/store_test.go +96 −0 Changes for internal/pgint/store_test.go: 96 added lines, 0 removed lines. Original line number Diff line number Diff line Loading @@ -4,11 +4,14 @@ package pgint import ( "context" "database/sql" "errors" "os" "strings" "testing" "time" "felis.lolicon.best/internal/api" "felis.lolicon.best/internal/store" "github.com/jackc/pgx/v5/pgconn" Loading Loading @@ -133,3 +136,96 @@ func TestMigrationBodyHasNoStatementLimit(t *testing.T) { t.Fatalf("statement_timeout inside a migration = %q, want 0", v) } } // Two assertions verified against the same stored counter can land in either // order; the lower one never pulls the counter back, and both stamp the use. func TestSignCountNeverMovesBack(t *testing.T) { ctx := context.Background() u := newUser(t, "user", "signcount") cid := "cid-" + suffix(t) if _, err := db.ExecContext(ctx, `INSERT INTO webauthn_credentials (id, user_id, credential_id, public_key, sign_count) VALUES ($1,$2,$3,'pk',4)`, "cred-"+suffix(t), u.ID, cid); err != nil { t.Fatal(err) } later := mustNow().Add(time.Minute).Truncate(time.Microsecond) if err := repo.AdvanceCredentialSignCount(ctx, cid, 6, mustNow()); err != nil { t.Fatal(err) } if err := repo.AdvanceCredentialSignCount(ctx, cid, 5, later); err != nil { t.Fatal(err) } var n int64 var used time.Time if err := db.QueryRowContext(ctx, `SELECT sign_count, last_used_at FROM webauthn_credentials WHERE credential_id = $1`, cid).Scan(&n, &used); err != nil { t.Fatal(err) } if n != 6 { t.Errorf("sign_count = %d after 6 then 5, want 6", n) } if !used.Equal(later) { t.Errorf("last_used_at = %v, want the later use %v", used, later) } } // Disabling revokes the sessions in the same transaction: when the revocation // fails, the flag is not left set, and the caller hears about it. func TestSetUserDisabledIsAtomic(t *testing.T) { ctx := context.Background() u := newUser(t, "user", "disable") hash := "h-" + suffix(t) if err := repo.CreateSession(ctx, hash, u.ID, mustNow().Add(time.Hour)); err != nil { t.Fatal(err) } // Make every session update fail, as a lost connection or a lock timeout would. fn := "pgint_refuse_" + suffix(t) for _, stmt := range []string{ `CREATE FUNCTION ` + fn + `() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN RAISE EXCEPTION 'refused'; END $$`, `CREATE TRIGGER ` + fn + ` BEFORE UPDATE ON sessions FOR EACH ROW EXECUTE FUNCTION ` + fn + `()`, } { if _, err := db.ExecContext(ctx, stmt); err != nil { t.Fatal(err) } } dropped := false drop := func() { if dropped { return } dropped = true _, _ = db.ExecContext(ctx, `DROP TRIGGER `+fn+` ON sessions`) _, _ = db.ExecContext(ctx, `DROP FUNCTION `+fn+`()`) } t.Cleanup(drop) if err := repo.SetUserDisabled(ctx, u.ID, true); err == nil { t.Fatal("SetUserDisabled succeeded although its sessions could not be revoked") } var disabled bool if err := db.QueryRowContext(ctx, `SELECT disabled FROM users WHERE id = $1`, u.ID).Scan(&disabled); err != nil { t.Fatal(err) } if disabled { t.Fatal("the account was left disabled with its sessions still live") } drop() if err := repo.SetUserDisabled(ctx, u.ID, true); err != nil { t.Fatal(err) } if _, err := repo.SessionUser(ctx, hash, mustNow()); !errors.Is(err, api.ErrNotFound) { t.Fatalf("session after disable = %v, want ErrNotFound", err) } var revoked sql.NullTime if err := db.QueryRowContext(ctx, `SELECT revoked_at FROM sessions WHERE token_hash = $1`, hash).Scan(&revoked); err != nil { t.Fatal(err) } if !revoked.Valid { t.Fatal("disable left the session unrevoked") } if err := repo.SetUserDisabled(ctx, "no-such-"+suffix(t), true); !errors.Is(err, api.ErrNotFound) { t.Fatalf("SetUserDisabled(missing) = %v, want ErrNotFound", err) } } Loading
internal/api/api.go +10 −1 Changes for internal/api/api.go: 10 added lines, 1 removed line. Original line number Diff line number Diff line Loading @@ -14,6 +14,7 @@ package api import ( "context" "log" "log/slog" "net/http" "strings" Loading Loading @@ -776,7 +777,15 @@ func (a *API) requireOnboarded(h http.HandlerFunc) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { p := principalFromContext(r.Context()) if p != nil && p.ViaSession && !p.EmailVerified { creds, _ := a.Repo.PasskeyCredentialsForUser(r.Context(), p.UserID) creds, err := a.Repo.PasskeyCredentialsForUser(r.Context(), p.UserID) if err != nil { // A store outage reads as retry-later; setup_required would send // the caller off to enroll a passkey they may already have. log.Printf("api: %s %s: onboarding check (request_id=%s): %v", r.Method, r.URL.Path, requestIDFromContext(r.Context()), err) writeError(w, r, errAuthUnavailable) return } if len(creds) == 0 { writeError(w, r, newError(http.StatusForbidden, "setup_required", "passkey enrollment is required before this action is available")) Loading
internal/api/api_test.go +1 −1 Changes for internal/api/api_test.go: 1 added line, 1 removed line. Original line number Diff line number Diff line Loading @@ -583,7 +583,7 @@ func (f *fakeRepo) DeleteAllPasskeyCredentialsForUser(_ context.Context, userID func (f *fakeRepo) AdvanceCredentialSignCount(_ context.Context, credentialID string, newSignCount uint32, usedAt time.Time) error { for id, c := range f.passkeyCreds { if c.CredentialID == credentialID { c.SignCount = newSignCount c.SignCount = max(c.SignCount, newSignCount) t := usedAt c.LastUsedAt = &t f.passkeyCreds[id] = c Loading
internal/api/handlers_setup_lockdown_test.go +31 −0 Changes for internal/api/handlers_setup_lockdown_test.go: 31 added lines, 0 removed lines. Original line number Diff line number Diff line Loading @@ -2,6 +2,7 @@ package api import ( "context" "errors" "net/http" "net/http/httptest" "testing" Loading Loading @@ -138,3 +139,33 @@ func TestAdvRequireOnboardedPredicate(t *testing.T) { }) } } // credsDown is a store whose passkey lookup fails, as during a Postgres outage. type credsDown struct{ *fakeRepo } func (credsDown) PasskeyCredentialsForUser(context.Context, string) ([]PasskeyCredential, error) { return nil, errors.New("connection refused") } // A failed passkey lookup is an outage: the caller gets 503 auth_unavailable to // retry, never setup_required, which would send a player with a passkey off to // enroll another. func TestRequireOnboardedReportsAStoreOutage(t *testing.T) { repo := newFakeRepo() api := newTestAPI(repo, newFakeCluster()) api.Repo = credsDown{repo} ran := false h := api.requireOnboarded(func(http.ResponseWriter, *http.Request) { ran = true }) r := httptest.NewRequest("POST", "/api/v1/servers/demo2/wake", nil) r = r.WithContext(context.WithValue(r.Context(), ctxKeyPrincipal, &Principal{UserID: "has-pk", ViaSession: true, EmailVerified: false})) w := httptest.NewRecorder() h(w, r) if ran { t.Fatal("the handler ran although the onboarding check could not be made") } if w.Code != http.StatusServiceUnavailable || errCode(w.Body.Bytes()) != "auth_unavailable" { t.Fatalf("got %d %s, want 503 auth_unavailable", w.Code, w.Body.String()) } }
internal/api/pgrepo.go +25 −16 Changes for internal/api/pgrepo.go: 25 added lines, 16 removed lines. Original line number Diff line number Diff line Loading @@ -1401,9 +1401,14 @@ func (p *PGRepo) PasskeyCredentialsForUser(ctx context.Context, userID string) ( // last_used_at. credential_id is UNIQUE so exactly one row is touched; a missing row (the // credential was unbound mid-ceremony) affects zero rows and is a successful no-op, never an // error — the assertion is already cryptographically complete by the time this runs. // // The counter only moves forward: two assertions verified at once against the same stored // value land in either order, and the lower one must not overwrite the higher, or a clone // replaying the count in between would pass the next check. func (p *PGRepo) AdvanceCredentialSignCount(ctx context.Context, credentialID string, newSignCount uint32, usedAt time.Time) error { _, err := p.db.ExecContext(ctx, `UPDATE webauthn_credentials SET sign_count = $2, last_used_at = $3 WHERE credential_id = $1`, `UPDATE webauthn_credentials SET sign_count = GREATEST(sign_count, $2), last_used_at = $3 WHERE credential_id = $1`, credentialID, int64(newSignCount), usedAt) return err } Loading Loading @@ -1738,31 +1743,35 @@ func (p *PGRepo) DeleteUser(ctx context.Context, userID, _ string) error { } // SetUserDisabled flips the disabled flag. Setting disabled→true additionally // revokes every live session so the account is immediately locked out. // revokes every live session so the account is immediately locked out. Both land // in one transaction: a disable whose revocation failed is rolled back and // reported, so the admin never reads "disabled" while a session still stands. func (p *PGRepo) SetUserDisabled(ctx context.Context, userID string, disabled bool) error { // Guard: the user must exist and not be deleted. var ok bool if err := p.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM users WHERE id = $1 AND deleted_at IS NULL)`, userID).Scan(&ok); err != nil { tx, err := p.db.BeginTx(ctx, nil) if err != nil { return err } if !ok { return ErrNotFound } defer func() { _ = tx.Rollback() }() if _, err := p.db.ExecContext(ctx, `UPDATE users SET disabled = $2 WHERE id = $1`, userID, disabled); err != nil { res, err := tx.ExecContext(ctx, `UPDATE users SET disabled = $2 WHERE id = $1 AND deleted_at IS NULL`, userID, disabled) if err != nil { return err } if n, err := res.RowsAffected(); err != nil { return err } else if n == 0 { return ErrNotFound } if disabled { // Revoke every live session so the lockout is immediate. _, _ = p.db.ExecContext(ctx, if _, err := tx.ExecContext(ctx, `UPDATE sessions SET revoked_at = now() WHERE user_id = $1 AND revoked_at IS NULL`, userID) userID); err != nil { return err } return nil } return tx.Commit() } // ---- quota admin ---- Loading
internal/pgint/store_test.go +96 −0 Changes for internal/pgint/store_test.go: 96 added lines, 0 removed lines. Original line number Diff line number Diff line Loading @@ -4,11 +4,14 @@ package pgint import ( "context" "database/sql" "errors" "os" "strings" "testing" "time" "felis.lolicon.best/internal/api" "felis.lolicon.best/internal/store" "github.com/jackc/pgx/v5/pgconn" Loading Loading @@ -133,3 +136,96 @@ func TestMigrationBodyHasNoStatementLimit(t *testing.T) { t.Fatalf("statement_timeout inside a migration = %q, want 0", v) } } // Two assertions verified against the same stored counter can land in either // order; the lower one never pulls the counter back, and both stamp the use. func TestSignCountNeverMovesBack(t *testing.T) { ctx := context.Background() u := newUser(t, "user", "signcount") cid := "cid-" + suffix(t) if _, err := db.ExecContext(ctx, `INSERT INTO webauthn_credentials (id, user_id, credential_id, public_key, sign_count) VALUES ($1,$2,$3,'pk',4)`, "cred-"+suffix(t), u.ID, cid); err != nil { t.Fatal(err) } later := mustNow().Add(time.Minute).Truncate(time.Microsecond) if err := repo.AdvanceCredentialSignCount(ctx, cid, 6, mustNow()); err != nil { t.Fatal(err) } if err := repo.AdvanceCredentialSignCount(ctx, cid, 5, later); err != nil { t.Fatal(err) } var n int64 var used time.Time if err := db.QueryRowContext(ctx, `SELECT sign_count, last_used_at FROM webauthn_credentials WHERE credential_id = $1`, cid).Scan(&n, &used); err != nil { t.Fatal(err) } if n != 6 { t.Errorf("sign_count = %d after 6 then 5, want 6", n) } if !used.Equal(later) { t.Errorf("last_used_at = %v, want the later use %v", used, later) } } // Disabling revokes the sessions in the same transaction: when the revocation // fails, the flag is not left set, and the caller hears about it. func TestSetUserDisabledIsAtomic(t *testing.T) { ctx := context.Background() u := newUser(t, "user", "disable") hash := "h-" + suffix(t) if err := repo.CreateSession(ctx, hash, u.ID, mustNow().Add(time.Hour)); err != nil { t.Fatal(err) } // Make every session update fail, as a lost connection or a lock timeout would. fn := "pgint_refuse_" + suffix(t) for _, stmt := range []string{ `CREATE FUNCTION ` + fn + `() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN RAISE EXCEPTION 'refused'; END $$`, `CREATE TRIGGER ` + fn + ` BEFORE UPDATE ON sessions FOR EACH ROW EXECUTE FUNCTION ` + fn + `()`, } { if _, err := db.ExecContext(ctx, stmt); err != nil { t.Fatal(err) } } dropped := false drop := func() { if dropped { return } dropped = true _, _ = db.ExecContext(ctx, `DROP TRIGGER `+fn+` ON sessions`) _, _ = db.ExecContext(ctx, `DROP FUNCTION `+fn+`()`) } t.Cleanup(drop) if err := repo.SetUserDisabled(ctx, u.ID, true); err == nil { t.Fatal("SetUserDisabled succeeded although its sessions could not be revoked") } var disabled bool if err := db.QueryRowContext(ctx, `SELECT disabled FROM users WHERE id = $1`, u.ID).Scan(&disabled); err != nil { t.Fatal(err) } if disabled { t.Fatal("the account was left disabled with its sessions still live") } drop() if err := repo.SetUserDisabled(ctx, u.ID, true); err != nil { t.Fatal(err) } if _, err := repo.SessionUser(ctx, hash, mustNow()); !errors.Is(err, api.ErrNotFound) { t.Fatalf("session after disable = %v, want ErrNotFound", err) } var revoked sql.NullTime if err := db.QueryRowContext(ctx, `SELECT revoked_at FROM sessions WHERE token_hash = $1`, hash).Scan(&revoked); err != nil { t.Fatal(err) } if !revoked.Valid { t.Fatal("disable left the session unrevoked") } if err := repo.SetUserDisabled(ctx, "no-such-"+suffix(t), true); !errors.Is(err, api.ErrNotFound) { t.Fatalf("SetUserDisabled(missing) = %v, want ErrNotFound", err) } }