From fcf5dd539d65a96171fa7b24eb34c58ce2d8ec0b Mon Sep 17 00:00:00 2001 From: Rudy Celekli Date: Wed, 7 Oct 2026 03:00:02 -0400 Subject: [PATCH] fix: index concurrent first chat saves atomically Signed-off-by: Rudy Celekli --- db/services/chats.ts | 46 ++++++----------- web/trpc/chat-concurrency.test.ts | 82 +++++++++++++++++++++++++++++++ 2 files changed, 98 insertions(+), 30 deletions(-) create mode 100644 web/trpc/chat-concurrency.test.ts diff --git a/db/services/chats.ts b/db/services/chats.ts index ebac3bdb..b2b52560 100644 --- a/db/services/chats.ts +++ b/db/services/chats.ts @@ -68,17 +68,17 @@ export async function saveChat( if (!(await waitForSessionOwnership(scope, chat.sessionId))) return; await ensureScope(scope); const now = new Date(); - const existing = await db - .select({ sessionId: chats.sessionId }) - .from(chats) - .where( - and( - eq(chats.workspaceId, scope.workspaceId), - eq(chats.sessionId, chat.sessionId) - ) - ); - if (existing.length === 0) { - await db.insert(chats).values({ + const updates: Partial = { updatedAt: now }; + if (chat.channel !== undefined) updates.channel = chat.channel; + if (chat.title !== undefined) updates.title = chat.title; + if (chat.usage !== undefined) { + updates.costUsd = chat.usage.costUsd; + updates.inputTokens = chat.usage.inputTokens; + updates.outputTokens = chat.usage.outputTokens; + } + await db + .insert(chats) + .values({ channel: chat.channel ?? null, costUsd: chat.usage?.costUsd ?? null, createdAt: now, @@ -88,24 +88,10 @@ export async function saveChat( title: chat.title ?? "New chat", updatedAt: now, workspaceId: scope.workspaceId, + }) + .onConflictDoUpdate({ + target: chats.sessionId, + set: updates, + setWhere: eq(chats.workspaceId, scope.workspaceId), }); - return; - } - const updates: Partial = { updatedAt: now }; - if (chat.channel !== undefined) updates.channel = chat.channel; - if (chat.title !== undefined) updates.title = chat.title; - if (chat.usage !== undefined) { - updates.costUsd = chat.usage.costUsd; - updates.inputTokens = chat.usage.inputTokens; - updates.outputTokens = chat.usage.outputTokens; - } - await db - .update(chats) - .set(updates) - .where( - and( - eq(chats.workspaceId, scope.workspaceId), - eq(chats.sessionId, chat.sessionId) - ) - ); } diff --git a/web/trpc/chat-concurrency.test.ts b/web/trpc/chat-concurrency.test.ts new file mode 100644 index 00000000..66852fb4 --- /dev/null +++ b/web/trpc/chat-concurrency.test.ts @@ -0,0 +1,82 @@ +import { PGlite } from "@electric-sql/pglite"; +import { drizzle } from "drizzle-orm/pglite"; +import { migrate } from "drizzle-orm/pglite/migrator"; +import { + afterAll, + beforeAll, + beforeEach, + describe, + expect, + it, + vi, +} from "vitest"; +import * as Database from "@db"; +import * as schema from "@db/schema"; +import { readChat, saveChat } from "@db/services/chats"; +import { ensureScope } from "@db/services/scope"; +import { claimSession } from "@db/services/sessions"; +import { appRouter } from "@web/trpc/router"; + +const client = new PGlite(); +const database = drizzle(client, { schema }); +const scope = { userId: "chat-fixture", workspaceId: "workspace:chat-fixture" }; +const caller = appRouter.createCaller({ origin: "https://example.com", scope }); + +beforeAll(async () => { + await migrate(database, { migrationsFolder: "db/migrations" }); + // SAFETY: PGlite supplies the real Drizzle query-builder contract used by the service; only the database driver changes. + // oxlint-disable-next-line typescript/no-unsafe-type-assertion -- Exercise the committed migrations and real services against an isolated PostgreSQL-compatible database. + vi.spyOn(Database, "db", "get").mockReturnValue(database as never); +}, 20_000); + +beforeEach(async () => { + await database.delete(schema.workspaces); + await ensureScope(scope); + await claimSession(scope, "chat-session"); +}); + +afterAll(async () => { + vi.restoreAllMocks(); + await client.close(); +}); + +describe("chat indexing", () => { + it("preserves title and usage when the first UI saves overlap", async () => { + await Promise.all([ + caller.chats.save({ sessionId: "chat-session", title: "Travel plans" }), + caller.chats.save({ + sessionId: "chat-session", + usage: { costUsd: 0.25, inputTokens: 10, outputTokens: 4 }, + }), + ]); + expect(await readChat(scope, "chat-session")).toMatchObject({ + title: "Travel plans", + usage: { costUsd: 0.25, inputTokens: 10, outputTokens: 4 }, + }); + }); + + it("preserves sparse fields during sequential saves and denies another workspace", async () => { + await caller.chats.save({ + sessionId: "chat-session", + title: "Travel plans", + }); + await saveChat(scope, { channel: "http", sessionId: "chat-session" }); + await caller.chats.save({ + sessionId: "chat-session", + usage: { costUsd: 0.25, inputTokens: 10, outputTokens: 4 }, + }); + const before = await readChat(scope, "chat-session"); + expect(before).toMatchObject({ + channel: "http", + title: "Travel plans", + usage: { costUsd: 0.25, inputTokens: 10, outputTokens: 4 }, + }); + await appRouter + .createCaller({ + origin: "https://example.com", + scope: { userId: "other-user", workspaceId: "workspace:other-user" }, + }) + .chats.save({ sessionId: "chat-session", title: "Other title" }); + expect(await readChat(scope, "chat-session")).toEqual(before); + }); +});