fix(api): 禁用用户与吊销会话合为一个事务,passkey 签名计数只增不减,引导检查遇到数据库故障返回 503
This commit is contained in:
5 files changed
+164
-19
No files matched your search
@@ -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"
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user