From 278d0746a28ab0d0c22d70cc63e30073df113276 Mon Sep 17 00:00:00 2001 From: Macson Obeko <278668028+macsonfleek@users.noreply.github.com> Date: Mon, 31 Aug 2026 15:57:39 +0000 Subject: [PATCH] fix: eliminate TOCTOU race in select-winners by moving eligibility check inside transaction MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Move the post status re-read and winner eligibility computation inside the $transaction with Serializable isolation level. This prevents two concurrent selection requests from both passing the "already completed" guard and over-selecting winners beyond maxWinners. Key changes: - Re-read post status + existing winners inside the Serializable transaction - Compute remainingSlots (maxWinners - existingWinners) inside the transaction - Cap all selection methods at remainingSlots to prevent over-assignment - Surface domain errors as 400 responses instead of generic 500s - Add concurrency test asserting exactly one request succeeds under contention Closes #426 🤖 Generated with Codebuff Co-Authored-By: Codebuff --- .../api/posts/[id]/select-winners/route.ts | 298 +++++++++++------- .../api/select-winners-concurrency.test.ts | 206 ++++++++++++ 2 files changed, 383 insertions(+), 121 deletions(-) create mode 100644 app/tests/api/select-winners-concurrency.test.ts diff --git a/app/app/api/posts/[id]/select-winners/route.ts b/app/app/api/posts/[id]/select-winners/route.ts index 68f39cc..af53a04 100644 --- a/app/app/api/posts/[id]/select-winners/route.ts +++ b/app/app/api/posts/[id]/select-winners/route.ts @@ -68,162 +68,191 @@ export const POST = async ( } const body = parsed.data; - const post = await prisma.post.findUnique({ + // ── Pre-transaction ownership & moderation checks (read-only) ─────── + const preCheck = await prisma.post.findUnique({ where: { id }, - include: { - entries: { - include: { burns: true }, - orderBy: { createdAt: "asc" }, - }, - winners: true, + select: { + userId: true, + moderationStatus: true, }, }); - if (!post) return apiError("Post not found", 404); - if (post.userId !== user.id) return apiError("Forbidden", 403); - if (["suspended", "banned"].includes(post.moderationStatus)) { + if (!preCheck) return apiError("Post not found", 404); + if (preCheck.userId !== user.id) return apiError("Forbidden", 403); + if (["suspended", "banned"].includes(preCheck.moderationStatus)) { return apiError("Cannot select winners for moderated content", 403); } - if (post.status === "completed") { - return apiError("Winners already selected for this post", 400); - } - if (!["open", "active", "in_progress"].includes(post.status)) { - return apiError( - `Cannot select winners for a post with status "${post.status}"`, - 400, - ); - } - if (post.entries.length === 0) { - return apiError("No entries to select from", 400); - } + // ── Atomic transaction: re-read status, select winners, persist ───── + // Serializable isolation prevents two concurrent requests from both + // passing the "already completed" guard and over-selecting winners. + const result = await prisma.$transaction( + async (tx) => { + // Re-read the post with fresh status + winner count inside the transaction + const post = await tx.post.findUnique({ + where: { id }, + include: { + entries: { + include: { burns: true }, + orderBy: { createdAt: "asc" }, + }, + winners: true, + }, + }); - // Exclude users who are already winners (prevents duplicates across calls) - const existingWinnerUserIds = new Set(post.winners.map((w) => w.userId)); - const eligibleEntries = post.entries.filter( - (e) => !existingWinnerUserIds.has(e.userId), - ); + if (!post) throw new Error("Post not found"); - const maxWinners = post.maxWinners ?? 1; - let selectedEntries: typeof eligibleEntries = []; + if (post.status === "completed") { + throw new Error("ALREADY_COMPLETED"); + } + if (!["open", "active", "in_progress"].includes(post.status)) { + throw new Error( + `Cannot select winners for a post with status "${post.status}"`, + ); + } + if (post.entries.length === 0) { + throw new Error("No entries to select from"); + } - switch (body.method) { - case "random": { - const count = Math.min( - body.count ?? maxWinners, - eligibleEntries.length, + // Compute eligibility and remaining slots inside the transaction + const existingWinnerUserIds = new Set(post.winners.map((w) => w.userId)); + const eligibleEntries = post.entries.filter( + (e) => !existingWinnerUserIds.has(e.userId), ); - selectedEntries = shuffle(eligibleEntries).slice(0, count); - break; - } - case "manual": { - const { entryIds } = body; + const maxWinners = post.maxWinners ?? 1; + const remainingSlots = Math.max(0, maxWinners - post.winners.length); - // All supplied IDs must belong to this post - const validEntryIds = new Set(post.entries.map((e) => e.id)); - const invalidIds = entryIds.filter((eid) => !validEntryIds.has(eid)); - if (invalidIds.length > 0) { - return apiError( - `Entry IDs not found on this post: ${invalidIds.join(", ")}`, - 400, - ); + if (remainingSlots === 0) { + throw new Error("ALREADY_COMPLETED"); } - // Deduplicate supplied IDs and cap at maxWinners - const uniqueIds = [...new Set(entryIds)].slice(0, maxWinners); - selectedEntries = eligibleEntries.filter((e) => - uniqueIds.includes(e.id), - ); + let selectedEntries: typeof eligibleEntries = []; + + switch (body.method) { + case "random": { + const count = Math.min( + body.count ?? remainingSlots, + remainingSlots, + eligibleEntries.length, + ); + selectedEntries = shuffle(eligibleEntries).slice(0, count); + break; + } + + case "manual": { + const { entryIds } = body; + + // All supplied IDs must belong to this post + const validEntryIds = new Set(post.entries.map((e) => e.id)); + const invalidIds = entryIds.filter((eid) => !validEntryIds.has(eid)); + if (invalidIds.length > 0) { + throw new Error( + `Entry IDs not found on this post: ${invalidIds.join(", ")}`, + ); + } + + // Deduplicate supplied IDs and cap at remainingSlots + const uniqueIds = [...new Set(entryIds)].slice(0, remainingSlots); + selectedEntries = eligibleEntries.filter((e) => + uniqueIds.includes(e.id), + ); + + if (selectedEntries.length === 0) { + throw new Error( + "None of the provided entry IDs belong to eligible entries", + ); + } + break; + } + + case "merit_based": { + // Rank by burn count (descending), then entry age (ascending) as tiebreaker + const count = Math.min( + body.count ?? remainingSlots, + remainingSlots, + eligibleEntries.length, + ); + selectedEntries = [...eligibleEntries] + .sort((a, b) => { + const burnDiff = b.burns.length - a.burns.length; + if (burnDiff !== 0) return burnDiff; + return a.createdAt.getTime() - b.createdAt.getTime(); + }) + .slice(0, count); + break; + } + + case "firstcome": { + // Entries are already ordered by createdAt asc + const count = Math.min( + body.count ?? remainingSlots, + remainingSlots, + eligibleEntries.length, + ); + selectedEntries = eligibleEntries.slice(0, count); + break; + } + } if (selectedEntries.length === 0) { - return apiError( - "None of the provided entry IDs belong to eligible entries", - 400, + throw new Error( + "No eligible entries found for the requested selection", ); } - break; - } - case "merit_based": { - // Rank by burn count (descending), then entry age (ascending) as tiebreaker - const count = Math.min( - body.count ?? maxWinners, - eligibleEntries.length, - ); - selectedEntries = [...eligibleEntries] - .sort((a, b) => { - const burnDiff = b.burns.length - a.burns.length; - if (burnDiff !== 0) return burnDiff; - return a.createdAt.getTime() - b.createdAt.getTime(); - }) - .slice(0, count); - break; - } + const entryIds = selectedEntries.map((e) => e.id); - case "firstcome": { - // Entries are already ordered by createdAt asc - const count = Math.min( - body.count ?? maxWinners, - eligibleEntries.length, - ); - selectedEntries = eligibleEntries.slice(0, count); - break; - } - } + await tx.entry.updateMany({ + where: { id: { in: entryIds } }, + data: { isWinner: true }, + }); - if (selectedEntries.length === 0) { - return apiError( - "No eligible entries found for the requested selection", - 400, - ); - } + await tx.postWinner.createMany({ + data: selectedEntries.map((e) => ({ + postId: post.id, + userId: e.userId, + assignedBy: user.id, + })), + skipDuplicates: true, + }); - // ── Persist in a single transaction ────────────────────────────────── - await prisma.$transaction(async (tx) => { - const entryIds = selectedEntries.map((e) => e.id); + await tx.post.update({ + where: { id: post.id }, + data: { status: "completed" }, + }); - await tx.entry.updateMany({ - where: { id: { in: entryIds } }, - data: { isWinner: true }, - }); + // Notify each winner using delivery layer + const fanOut = fanOutNotificationsInTransaction(tx); + await fanOut({ + userIds: selectedEntries.map((e) => e.userId), + type: "giveaway_win", + message: `Congratulations! You won the giveaway "${post.title}".`, + link: `/posts/${post.id}`, + }); - await tx.postWinner.createMany({ - data: selectedEntries.map((e) => ({ + return { postId: post.id, - userId: e.userId, - assignedBy: user.id, - })), - skipDuplicates: true, - }); - - await tx.post.update({ - where: { id: post.id }, - data: { status: "completed" }, - }); - - // Notify each winner using delivery layer - const fanOut = fanOutNotificationsInTransaction(tx); - await fanOut({ - userIds: selectedEntries.map((e) => e.userId), - type: "giveaway_win", - message: `Congratulations! You won the giveaway "${post.title}".`, - link: `/posts/${post.id}`, - }); - }); + selectedEntries, + }; + }, + { + isolationLevel: "Serializable", + }, + ); // Award badges to winners async (best-effort) - for (const entry of selectedEntries) { + for (const entry of result.selectedEntries) { checkAndAwardBadges(entry.userId).catch(console.error); } return apiSuccess( { method: body.method, - postId: post.id, + postId: result.postId, postStatus: "completed", - totalSelected: selectedEntries.length, - winners: selectedEntries.map((e) => ({ + totalSelected: result.selectedEntries.length, + winners: result.selectedEntries.map((e) => ({ entryId: e.id, userId: e.userId, })), @@ -231,6 +260,33 @@ export const POST = async ( "Winners selected successfully", ); } catch (error) { + // Surface domain errors as 400 responses instead of 500 + if (error instanceof Error) { + switch (error.message) { + case "Post not found": + return apiError("Post not found", 404); + case "ALREADY_COMPLETED": + return apiError("Winners already selected for this post", 400); + case "No entries to select from": + return apiError("No entries to select from", 400); + case "No eligible entries found for the requested selection": + return apiError( + "No eligible entries found for the requested selection", + 400, + ); + default: + if (error.message.startsWith("Cannot select winners")) { + return apiError(error.message, 400); + } + if (error.message.startsWith("Entry IDs not found")) { + return apiError(error.message, 400); + } + if (error.message.startsWith("None of the provided")) { + return apiError(error.message, 400); + } + break; + } + } console.error("Select winners error:", error); return apiError("Failed to select winners", 500); } diff --git a/app/tests/api/select-winners-concurrency.test.ts b/app/tests/api/select-winners-concurrency.test.ts new file mode 100644 index 0000000..0c073b3 --- /dev/null +++ b/app/tests/api/select-winners-concurrency.test.ts @@ -0,0 +1,206 @@ +import { describe, it, expect, beforeEach, vi } from "vitest"; +import { POST } from "@/app/api/posts/[id]/select-winners/route"; +import { prisma } from "@/lib/prisma"; +import { getCurrentUser } from "@/lib/auth"; + +// ── Helpers ────────────────────────────────────────────────────────────────── + +function makeRequest(body: Record) { + return new Request("http://localhost/api/posts/test-post-id/select-winners", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }); +} + +// ── Mocks ──────────────────────────────────────────────────────────────────── + +vi.mock("@/lib/auth", () => ({ + getCurrentUser: vi.fn(), +})); + +vi.mock("@/lib/notifications", () => ({ + fanOutNotificationsInTransaction: vi.fn(() => vi.fn()), +})); + +vi.mock("@/lib/badges", () => ({ + checkAndAwardBadges: vi.fn(), +})); + +// ── Tests ──────────────────────────────────────────────────────────────────── + +describe("POST /api/posts/[id]/select-winners – concurrency protection", () => { + const userId = "owner-user-id"; + const postId = "test-post-id"; + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(getCurrentUser).mockResolvedValue({ id: userId } as any); + }); + + it("rejects the second concurrent request when the first completes first", async () => { + // Arrange: a post with maxWinners=1 and 2 eligible entries + const entry1 = { + id: "entry-1", + userId: "user-a", + burns: [], + createdAt: new Date("2025-01-01"), + }; + const entry2 = { + id: "entry-2", + userId: "user-b", + burns: [], + createdAt: new Date("2025-01-02"), + }; + + const openPost = { + id: postId, + status: "open", + userId, + moderationStatus: "approved", + title: "Test Giveaway", + maxWinners: 1, + entries: [entry1, entry2], + winners: [], + }; + + const completedPost = { + ...openPost, + status: "completed", + winners: [{ userId: "user-a", postId, assignedBy: userId }], + }; + + // The prisma mock tracks calls to simulate Serializable isolation: + // first tx $transaction callback reads "open", second reads "completed" + const txCallbacks: Array<() => Promise> = []; + const originalTransaction = (prisma as any).$transaction; + + let txCallCount = 0; + (prisma as any).$transaction = vi.fn( + async (fn: (tx: any) => Promise, _opts?: any) => { + txCallCount++; + const isFirstTx = txCallCount <= 1; + + // Build a mock `tx` that re-reads the post + const txMock = { + post: { + findUnique: vi.fn().mockResolvedValue( + isFirstTx ? { ...openPost } : { ...completedPost }, + ), + update: vi.fn().mockResolvedValue({}), + }, + entry: { + updateMany: vi.fn().mockResolvedValue({}), + }, + postWinner: { + createMany: vi.fn().mockResolvedValue({}), + }, + }; + + const result = await fn(txMock); + return result; + }, + ); + + // Act: fire two concurrent "random" selections + const [res1, res2] = await Promise.allSettled([ + POST(makeRequest({ method: "random" }), { + params: Promise.resolve({ id: postId }), + }), + POST(makeRequest({ method: "random" }), { + params: Promise.resolve({ id: postId }), + }), + ]); + + // Restore original + (prisma as any).$transaction = originalTransaction; + + // Assert: one should succeed, the other should fail with 400 + const statuses = [ + res1.status === "fulfilled" ? res1.value.status : null, + res2.status === "fulfilled" ? res2.value.status : null, + ]; + + expect(statuses).toContain(200); + expect(statuses).toContain(400); + }, 10_000); + + it("never exceeds maxWinners even with concurrent requests", async () => { + // Arrange: post with maxWinners=2 and 4 entries + const entries = Array.from({ length: 4 }, (_, i) => ({ + id: `entry-${i + 1}`, + userId: `user-${i + 1}`, + burns: [], + createdAt: new Date(`2025-01-0${i + 1}`), + })); + + const openPost = { + id: postId, + status: "open", + userId, + moderationStatus: "approved", + title: "Test Giveaway", + maxWinners: 2, + entries, + winners: [], + }; + + const postAfterFirstTx = { + ...openPost, + status: "completed", + winners: [ + { userId: "user-1", postId, assignedBy: userId }, + { userId: "user-2", postId, assignedBy: userId }, + ], + }; + + let txCallCount = 0; + const originalTransaction = (prisma as any).$transaction; + + (prisma as any).$transaction = vi.fn( + async (fn: (tx: any) => Promise, _opts?: any) => { + txCallCount++; + const isFirstTx = txCallCount <= 1; + + const txMock = { + post: { + findUnique: vi.fn().mockResolvedValue( + isFirstTx ? { ...openPost } : { ...postAfterFirstTx }, + ), + update: vi.fn().mockResolvedValue({}), + }, + entry: { + updateMany: vi.fn().mockResolvedValue({}), + }, + postWinner: { + createMany: vi.fn().mockResolvedValue({}), + }, + }; + + return await fn(txMock); + }, + ); + + // Act: fire two concurrent random selections + const [res1, res2] = await Promise.allSettled([ + POST(makeRequest({ method: "random" }), { + params: Promise.resolve({ id: postId }), + }), + POST(makeRequest({ method: "random" }), { + params: Promise.resolve({ id: postId }), + }), + ]); + + (prisma as any).$transaction = originalTransaction; + + // Count how many winners were attempted via createMany + const results = [res1, res2].filter( + (r): r is PromiseFulfilledResult => + r.status === "fulfilled", + ); + + // At most one should have succeeded + const successCount = results.filter((r) => r.value.status === 200).length; + expect(successCount).toBeLessThanOrEqual(1); + }, 10_000); +});