fix(api): count only unsent attachments toward the per-message limit

attachments.create counted every file in the thread against MAX_ATTACHMENTS_PER_MESSAGE, so a conversation could never hold more than 10 files. Files already sent with a message still count toward the thread's byte quota.
This commit is contained in:
Amruth Pillai
2026-09-29 22:10:14 +02:00
parent 03c3a88841
commit 2622eeff12
2 changed files with 53 additions and 8 deletions
@@ -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<string, unknown>, 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<string, { type: string }>) => ({
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")) }) });
+14 -5
View File
@@ -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<number>`coalesce(sum(${schema.agentAttachment.size}), 0)`,
total: count(),
})
.select({ totalBytes: sql<number>`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");
}