fix(api): 禁用用户与吊销会话合为一个事务,passkey 签名计数只增不减,引导检查遇到数据库故障返回 503
This commit is contained in:
5 files changed
+164
-19
No files matched your search
+10
-1
@@ -14,6 +14,7 @@ package api
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"log"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -776,7 +777,15 @@ func (a *API) requireOnboarded(h http.HandlerFunc) http.HandlerFunc {
|
|||||||
return func(w http.ResponseWriter, r *http.Request) {
|
return func(w http.ResponseWriter, r *http.Request) {
|
||||||
p := principalFromContext(r.Context())
|
p := principalFromContext(r.Context())
|
||||||
if p != nil && p.ViaSession && !p.EmailVerified {
|
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 {
|
if len(creds) == 0 {
|
||||||
writeError(w, r, newError(http.StatusForbidden, "setup_required",
|
writeError(w, r, newError(http.StatusForbidden, "setup_required",
|
||||||
"passkey enrollment is required before this action is available"))
|
"passkey enrollment is required before this action is available"))
|
||||||
|
|||||||
@@ -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 {
|
func (f *fakeRepo) AdvanceCredentialSignCount(_ context.Context, credentialID string, newSignCount uint32, usedAt time.Time) error {
|
||||||
for id, c := range f.passkeyCreds {
|
for id, c := range f.passkeyCreds {
|
||||||
if c.CredentialID == credentialID {
|
if c.CredentialID == credentialID {
|
||||||
c.SignCount = newSignCount
|
c.SignCount = max(c.SignCount, newSignCount)
|
||||||
t := usedAt
|
t := usedAt
|
||||||
c.LastUsedAt = &t
|
c.LastUsedAt = &t
|
||||||
f.passkeyCreds[id] = c
|
f.passkeyCreds[id] = c
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package api
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
|
}
|
||||||
+26
-17
@@ -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
|
// 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
|
// 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.
|
// 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 {
|
func (p *PGRepo) AdvanceCredentialSignCount(ctx context.Context, credentialID string, newSignCount uint32, usedAt time.Time) error {
|
||||||
_, err := p.db.ExecContext(ctx,
|
_, 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)
|
credentialID, int64(newSignCount), usedAt)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -1738,31 +1743,35 @@ func (p *PGRepo) DeleteUser(ctx context.Context, userID, _ string) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// SetUserDisabled flips the disabled flag. Setting disabled→true additionally
|
// 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 {
|
func (p *PGRepo) SetUserDisabled(ctx context.Context, userID string, disabled bool) error {
|
||||||
// Guard: the user must exist and not be deleted.
|
tx, err := p.db.BeginTx(ctx, nil)
|
||||||
var ok bool
|
if err != nil {
|
||||||
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 {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if !ok {
|
defer func() { _ = tx.Rollback() }()
|
||||||
|
|
||||||
|
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
|
return ErrNotFound
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := p.db.ExecContext(ctx,
|
|
||||||
`UPDATE users SET disabled = $2 WHERE id = $1`, userID, disabled); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if disabled {
|
if disabled {
|
||||||
// Revoke every live session so the lockout is immediate.
|
if _, err := tx.ExecContext(ctx,
|
||||||
_, _ = p.db.ExecContext(ctx,
|
|
||||||
`UPDATE sessions SET revoked_at = now() WHERE user_id = $1 AND revoked_at IS NULL`,
|
`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 ----
|
// ---- quota admin ----
|
||||||
|
|||||||
@@ -4,11 +4,14 @@ package pgint
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"database/sql"
|
||||||
"errors"
|
"errors"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"felis.lolicon.best/internal/api"
|
||||||
"felis.lolicon.best/internal/store"
|
"felis.lolicon.best/internal/store"
|
||||||
|
|
||||||
"github.com/jackc/pgx/v5/pgconn"
|
"github.com/jackc/pgx/v5/pgconn"
|
||||||
@@ -133,3 +136,96 @@ func TestMigrationBodyHasNoStatementLimit(t *testing.T) {
|
|||||||
t.Fatalf("statement_timeout inside a migration = %q, want 0", v)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in new issue
Block a user