diff --git a/packages/api/src/features/agent/service.test.ts b/packages/api/src/features/agent/service.test.ts index a3e30a2aa..6177bc579 100644 --- a/packages/api/src/features/agent/service.test.ts +++ b/packages/api/src/features/agent/service.test.ts @@ -1227,7 +1227,7 @@ describe("agentService.attachments.create", () => { select: vi .fn() .mockImplementationOnce(() => ({ from: () => ({ where: () => lockQuery }) })) - .mockImplementationOnce(() => { + .mockImplementation(() => { quotaRead(); return selectWhereResult([{ total, totalBytes: String(totalBytes) }]); }), @@ -1256,7 +1256,8 @@ describe("agentService.attachments.create", () => { agentService.attachments.create(upload), ]); await vi.waitFor(() => expect(storageServiceMock.write).toHaveBeenCalledTimes(1)); - expect(quotaRead).toHaveBeenCalledTimes(1); + // The first upload's two quota reads (unsent count, thread bytes); the second waits on the lock. + expect(quotaRead).toHaveBeenCalledTimes(2); storageWrite.resolve(); const settled = await results; expect(settled[0]?.status).toBe("fulfilled"); @@ -1266,6 +1267,41 @@ describe("agentService.attachments.create", () => { expect(storageServiceMock.delete).not.toHaveBeenCalled(); }); + it("counts only unsent attachments toward the per-message limit", async () => { + // Ten files already went out with earlier messages; none is waiting to be sent. + const rows = Array.from({ length: 10 }, (_, index) => ({ + "agent_attachments.thread_id": input.threadId, + "agent_attachments.user_id": input.userId, + "agent_attachments.message_id": `message-${index}`, + })); + type Condition = { type: string; conditions?: Condition[]; left?: string; right?: unknown; value?: string }; + // Evaluates the mocked drizzle conditions, so each count follows the query's own filter. + const matches = (row: Record, condition: Condition): boolean => + condition.type === "and" + ? (condition.conditions ?? []).every((part) => matches(row, part)) + : condition.type === "isNull" + ? row[condition.value ?? ""] == null + : row[condition.left ?? ""] === condition.right; + dbMock.select + .mockReturnValueOnce({ from: () => ({ where: () => ({ for: async () => [{ id: input.threadId }] }) }) }) + .mockImplementation((columns: Record) => ({ + from: () => ({ + where: (condition: Condition) => { + const count = rows.filter((row) => matches(row, condition)).length; + return Promise.resolve([ + Object.fromEntries(Object.entries(columns).map(([key, { type }]) => [key, type === "count" ? count : 0])), + ]); + }, + }), + })); + dbMock.insert.mockReturnValue({ + values: (value: object) => ({ returning: async () => [{ ...value, createdAt: new Date() }] }), + }); + const { agentService } = await import("./service"); + + await expect(agentService.attachments.create(input)).resolves.toMatchObject({ filename: input.filename }); + }); + it("requires an owned, active, undeleted thread under the lock", async () => { const where = vi.fn(() => ({ for: vi.fn(async () => []) })); dbMock.select.mockReturnValue({ from: vi.fn(() => ({ where })) }); @@ -1296,7 +1332,7 @@ describe("agentService.attachments.create", () => { it("removes uploaded bytes when metadata insertion fails", async () => { dbMock.select .mockReturnValueOnce({ from: () => ({ where: () => ({ for: async () => [{ id: input.threadId }] }) }) }) - .mockImplementationOnce(() => selectWhereResult([{ total: 0, totalBytes: 0 }])); + .mockImplementation(() => selectWhereResult([{ total: 0, totalBytes: 0 }])); storageServiceMock.write.mockResolvedValue(undefined); storageServiceMock.delete.mockResolvedValue(true); dbMock.insert.mockReturnValue({ values: () => ({ returning: () => Promise.reject(new Error("Insert failed")) }) }); diff --git a/packages/api/src/features/agent/service.ts b/packages/api/src/features/agent/service.ts index 9a78e4a8e..74f8c4e16 100644 --- a/packages/api/src/features/agent/service.ts +++ b/packages/api/src/features/agent/service.ts @@ -1310,16 +1310,25 @@ export const agentService = { .for("update"); if (!thread) throw new ORPCError("NOT_FOUND"); + // Files already sent belong to their messages: only unsent ones count toward the next message. + const [unsent] = await tx + .select({ total: count() }) + .from(schema.agentAttachment) + .where( + and( + eq(schema.agentAttachment.threadId, input.threadId), + eq(schema.agentAttachment.userId, input.userId), + isNull(schema.agentAttachment.messageId), + ), + ); + if ((unsent?.total ?? 0) >= MAX_ATTACHMENTS_PER_MESSAGE) throw new ORPCError("BAD_REQUEST"); + const [stats] = await tx - .select({ - totalBytes: sql`coalesce(sum(${schema.agentAttachment.size}), 0)`, - total: count(), - }) + .select({ totalBytes: sql`coalesce(sum(${schema.agentAttachment.size}), 0)` }) .from(schema.agentAttachment) .where( and(eq(schema.agentAttachment.threadId, input.threadId), eq(schema.agentAttachment.userId, input.userId)), ); - if ((stats?.total ?? 0) >= MAX_ATTACHMENTS_PER_MESSAGE) throw new ORPCError("BAD_REQUEST"); if (Number(stats?.totalBytes ?? 0) + input.data.byteLength > MAX_THREAD_ATTACHMENT_BYTES) { throw new ORPCError("BAD_REQUEST"); }