fix(polls): canonicalize transaction lock order

This commit is contained in:
Simo committed 2026-09-02 19:09:50 +02:00
1 parent e9680775f9
commit de7cc42ddf
4 files changed
+275 -38

No files matched your search

+200 -26
View File
@@ -1,22 +1,127 @@
// @ts-nocheck
import { beforeEach, describe, expect, it, vi } from "vitest";
const { selectLimit, updateSet, updateWhere } = vi.hoisted(() => ({
selectLimit: vi.fn(),
updateSet: vi.fn(() => ({ where: updateWhere })),
updateWhere: vi.fn(),
const database = vi.hoisted(() => {
let selectedQueue: unknown[][] = [];
let mutationSteps: string[] = [];
let validationVoteQuestionIds: number[] = [];
let insertedVoteQuestionIds: number[] = [];
const nextSelected = () => selectedQueue.shift() ?? [];
const conditionPart = (condition: unknown, key: "column" | "value") => {
if (!condition || typeof condition !== "object") return "none";
return String((condition as Record<string, unknown>)[key] ?? "none");
};
const conditionOperator = (condition: unknown) => {
if (!condition || typeof condition !== "object") return "none";
return String((condition as Record<string, unknown>).operator ?? "none");
};
const select = vi.fn(() => ({
from: vi.fn((table: { tableName?: string }) => ({
where: vi.fn((condition: unknown) => {
const tableName = table.tableName ?? "unknown";
const run = async () => nextSelected();
return {
limit: vi.fn(async () => {
if (tableName === "WebsitePollVote") {
const clauses =
(condition as { clauses?: Array<Record<string, unknown>> })
?.clauses ?? [];
const question = clauses.find(
(clause) => clause.column === "questionId",
);
if (question) {
validationVoteQuestionIds.push(Number(question.value));
}
}
return run();
}),
for: vi.fn(async (strength: string) => {
mutationSteps.push(
`lock:${tableName}:${conditionOperator(condition)}:${conditionPart(condition, "column")}:${conditionPart(condition, "value")}:${strength}`,
);
return run();
}),
orderBy: vi.fn(async (column: unknown) => {
mutationSteps.push(
`read:${tableName}:${conditionOperator(condition)}:${conditionPart(condition, "column")}:orderBy:${String(column)}`,
);
return run();
}),
// biome-ignore lint/suspicious/noThenProperty: This query-builder mock must be awaitable like Drizzle.
then(
resolve: (value: unknown[]) => unknown,
reject: (error: unknown) => unknown,
) {
return run().then(resolve, reject);
},
};
}),
})),
}));
const updateWhere = vi.fn();
const updateSet = vi.fn(() => ({ where: updateWhere }));
const update = vi.fn(() => ({ set: updateSet }));
const insert = vi.fn((table: { tableName?: string }) => ({
values: vi.fn(async (values: Record<string, unknown>) => {
mutationSteps.push(`insert:${table.tableName ?? "unknown"}`);
if (table.tableName === "WebsitePollVote") {
insertedVoteQuestionIds.push(Number(values.questionId));
}
return [{ insertId: 1 }];
}),
}));
const deleteFrom = vi.fn((table: { tableName?: string }) => ({
where: vi.fn(async () => {
mutationSteps.push(`delete:${table.tableName ?? "unknown"}`);
}),
}));
const tx = { select, insert, delete: deleteFrom };
const transaction = vi.fn(
async (
run: (transaction: typeof tx) => Promise<unknown>,
_config?: unknown,
) => run(tx),
);
return {
db: { select, update, insert, delete: deleteFrom, transaction },
updateSet,
updateWhere,
transaction,
queueSelected(...values: unknown[][]) {
selectedQueue = [...values];
},
mutationSteps: () => [...mutationSteps],
validationVoteQuestionIds: () => [...validationVoteQuestionIds],
insertedVoteQuestionIds: () => [...insertedVoteQuestionIds],
reset() {
selectedQueue = [];
mutationSteps = [];
validationVoteQuestionIds = [];
insertedVoteQuestionIds = [];
},
};
});
vi.mock("drizzle-orm", () => ({
and: vi.fn((...clauses: unknown[]) => ({ operator: "and", clauses })),
eq: vi.fn((column: unknown, value: unknown) => ({
operator: "eq",
column,
value,
})),
inArray: vi.fn((column: unknown, value: unknown) => ({
operator: "inArray",
column,
value,
})),
}));
vi.mock("@/lib/db", () => ({
db: {
select: vi.fn(() => ({
from: vi.fn(() => ({
where: vi.fn(() => ({ limit: selectLimit })),
})),
})),
update: vi.fn(() => ({ set: updateSet })),
},
db: database.db,
WebsitePoll: {
tableName: "WebsitePoll",
id: "id",
title: "title",
description: "description",
@@ -27,6 +132,7 @@ vi.mock("@/lib/db", () => ({
endsAt: "endsAt",
},
WebsitePollQuestion: {
tableName: "WebsitePollQuestion",
id: "id",
pollId: "pollId",
question: "question",
@@ -34,14 +140,19 @@ vi.mock("@/lib/db", () => ({
sortOrder: "sortOrder",
options: "options",
},
WebsitePollVote: { id: "id", questionId: "questionId", userId: "userId" },
WebsitePollVote: {
tableName: "WebsitePollVote",
id: "id",
questionId: "questionId",
userId: "userId",
},
}));
vi.mock("@/lib/permissions", () => ({ PERMS: { POLLS_EDIT: "polls.edit" } }));
vi.mock("@/lib/services/audit", () => ({ logAudit: vi.fn() }));
vi.mock("next/cache", () => ({ revalidatePath: vi.fn() }));
vi.mock("@/lib/safe-action", () => ({
adminAction: (options, handler) => async (input) => {
vi.mock("@/lib/safe-action", () => {
const wrap = (options, handler) => async (input) => {
const parsed = options.schema.safeParse(input);
if (!parsed.success) return { ok: false, error: "Validation failed" };
try {
@@ -55,9 +166,9 @@ vi.mock("@/lib/safe-action", () => ({
error: error instanceof Error ? error.message : "Internal server error",
};
}
},
authAction: vi.fn(),
}));
};
return { adminAction: wrap, authAction: wrap };
});
vi.mock("@/lib/safe-action-shared", () => ({
actionOk: (data) => ({ ok: true, data: data ?? {} }),
@@ -70,18 +181,21 @@ vi.mock("@/lib/safe-action-shared", () => ({
},
}));
import { updatePoll, updatePollQuestion } from "./polls";
import {
deletePoll,
updatePoll,
updatePollQuestion,
voteOnPoll,
} from "./polls";
beforeEach(() => {
vi.clearAllMocks();
selectLimit.mockReset();
updateSet.mockClear();
updateWhere.mockClear();
database.reset();
});
describe("poll update merge validation", () => {
it("rejects a partial end time before the persisted start without updating", async () => {
selectLimit.mockResolvedValueOnce([
database.queueSelected([
{
id: 1,
title: "Schedule",
@@ -100,11 +214,11 @@ describe("poll update merge validation", () => {
});
expect(result).toEqual({ ok: false, error: "Invalid poll schedule" });
expect(updateSet).not.toHaveBeenCalled();
expect(database.updateSet).not.toHaveBeenCalled();
});
it("rejects a partial options patch for a persisted text question without updating", async () => {
selectLimit.mockResolvedValueOnce([
database.queueSelected([
{
pollId: 1,
question: "Why?",
@@ -117,6 +231,66 @@ describe("poll update merge validation", () => {
const result = await updatePollQuestion({ id: 1, options: "Not allowed" });
expect(result).toEqual({ ok: false, error: "Invalid poll question" });
expect(updateSet).not.toHaveBeenCalled();
expect(database.updateSet).not.toHaveBeenCalled();
});
});
describe("poll transaction lock order", () => {
it("deletes a legacy poll through ordered primary-key locks and child-first writes", async () => {
database.queueSelected(
[{ id: 7, title: "Legacy poll" }],
[{ id: 4 }, { id: 9 }],
[{ id: 4 }],
[{ id: 9 }],
);
const result = await deletePoll({ id: 7 });
expect(result).toEqual({ ok: true, data: {} });
expect(database.transaction).toHaveBeenCalledWith(expect.any(Function), {
isolationLevel: "read committed",
});
expect(database.mutationSteps()).toEqual([
"lock:WebsitePoll:eq:id:7:update",
"read:WebsitePollQuestion:eq:pollId:orderBy:id",
"lock:WebsitePollQuestion:eq:id:4:update",
"lock:WebsitePollQuestion:eq:id:9:update",
"delete:WebsitePollVote",
"delete:WebsitePollQuestion",
"delete:WebsitePoll",
]);
});
it("validates votes in request order but inserts a sorted copy", async () => {
database.queueSelected(
[
{
id: 7,
status: "active",
startsAt: null,
endsAt: null,
},
],
[
{ id: 9, pollId: 7, type: "single", options: "A\nB" },
{ id: 4, pollId: 7, type: "single", options: "A\nB" },
],
[],
[],
);
const input = {
pollId: 7,
votes: [
{ questionId: 9, answer: " A " },
{ questionId: 4, answer: "B" },
],
};
const result = await voteOnPoll(input);
expect(result).toEqual({ ok: true, data: { pollId: 7 } });
expect(database.validationVoteQuestionIds()).toEqual([9, 4]);
expect(database.insertedVoteQuestionIds()).toEqual([4, 9]);
expect(input.votes.map(({ questionId }) => questionId)).toEqual([9, 4]);
});
});
+41 -9
View File
@@ -1,6 +1,6 @@
"use server";
import { and, eq } from "drizzle-orm";
import { and, eq, inArray } from "drizzle-orm";
import { revalidatePath } from "next/cache";
import { z } from "zod";
import {
@@ -97,14 +97,43 @@ const deletePollInput = z.object({
export const deletePoll = adminAction(
{ permission: PERMS.POLLS_EDIT, schema: deletePollInput },
async (ctx) => {
const [existing] = await db
.select({ id: WebsitePoll.id, title: WebsitePoll.title })
.from(WebsitePoll)
.where(eq(WebsitePoll.id, ctx.data.id))
.limit(1);
if (!existing) throw new ActionError("Poll not found");
const existing = await db.transaction(
async (tx) => {
const [poll] = await tx
.select({ id: WebsitePoll.id, title: WebsitePoll.title })
.from(WebsitePoll)
.where(eq(WebsitePoll.id, ctx.data.id))
.for("update");
if (!poll) throw new ActionError("Poll not found");
const questionRows = await tx
.select({ id: WebsitePollQuestion.id })
.from(WebsitePollQuestion)
.where(eq(WebsitePollQuestion.pollId, poll.id))
.orderBy(WebsitePollQuestion.id);
for (const question of questionRows) {
await tx
.select({ id: WebsitePollQuestion.id })
.from(WebsitePollQuestion)
.where(eq(WebsitePollQuestion.id, question.id))
.for("update");
}
const questionIds = questionRows.map(({ id }) => id);
if (questionIds.length > 0) {
await tx
.delete(WebsitePollVote)
.where(inArray(WebsitePollVote.questionId, questionIds));
}
await tx
.delete(WebsitePollQuestion)
.where(eq(WebsitePollQuestion.pollId, poll.id));
await tx.delete(WebsitePoll).where(eq(WebsitePoll.id, poll.id));
return poll;
},
{ isolationLevel: "read committed" },
);
await db.delete(WebsitePoll).where(eq(WebsitePoll.id, ctx.data.id));
logAudit({
userId: ctx.session.user.id,
action: "poll_delete",
@@ -272,8 +301,11 @@ export const voteOnPoll = authAction(
}
}
const votesToInsert = [...ctx.data.votes].sort(
(left, right) => left.questionId - right.questionId,
);
await db.transaction(async (tx) => {
for (const vote of ctx.data.votes) {
for (const vote of votesToInsert) {
await tx.insert(WebsitePollVote).values({
questionId: vote.questionId,
userId,
@@ -101,6 +101,28 @@ describe("Content production mutation adapter", () => {
);
});
it("uses read committed isolation only for poll mutations", async () => {
const pollDeps = dependencies();
await pollDeps.adapter.execute(
"poll.change",
{ action: "delete", id: 7 },
mutationContext,
);
expect(pollDeps.transaction).toHaveBeenCalledWith(expect.any(Function), {
isolationLevel: "read committed",
});
const articleDeps = dependencies();
await articleDeps.adapter.execute(
"article.change",
{ action: "update", id: 7 },
mutationContext,
);
expect(articleDeps.transaction).toHaveBeenCalledWith(
expect.any(Function),
undefined,
);
});
it("persists correlated intent and truthful outcome around external storage writes", async () => {
const deps = dependencies();
await deps.adapter.execute(
@@ -46,9 +46,14 @@ export const CONTENT_MIXED_OPERATIONS = [
type TransactionToken = unknown;
interface ContentTransactionOptions {
isolationLevel: "read committed";
}
export interface ContentProductionMutationDependencies {
transaction<T>(
run: (transaction: TransactionToken) => Promise<T>,
options?: ContentTransactionOptions,
): Promise<T>;
writeAudit(entry: AuditEntry, transaction?: TransactionToken): Promise<void>;
executeOperation(
@@ -126,6 +131,10 @@ export function createContentProductionMutationAdapter(
return {
async execute(operation, input, context) {
if (includesOperation(CONTENT_DATABASE_OPERATIONS, operation)) {
const transactionOptions =
operation === "poll.change"
? ({ isolationLevel: "read committed" } as const)
: undefined;
return dependencies.transaction(async (transaction) => {
const snapshot = await dependencies.executeOperation(
operation,
@@ -138,7 +147,7 @@ export function createContentProductionMutationAdapter(
transaction,
);
return snapshot;
});
}, transactionOptions);
}
await dependencies.writeAudit(auditEntry(operation, context, "intent"));
@@ -192,9 +201,9 @@ export function createContentProductionMutationAdapter(
export const contentProductionMutationAdapter =
createContentProductionMutationAdapter({
async transaction(run) {
async transaction(run, options) {
const { db } = await import("@/lib/db");
return db.transaction((transaction) => run(transaction));
return db.transaction((transaction) => run(transaction), options);
},
async writeAudit(entry, transaction) {
const { logAudit } = await import("@/lib/services/audit");