diff --git a/docs/references/architecture-agent.md b/docs/references/architecture-agent.md index bf9164240..8defc00b4 100644 --- a/docs/references/architecture-agent.md +++ b/docs/references/architecture-agent.md @@ -61,7 +61,7 @@ session ends — no explicit `unregister` loop is required. [`ToolLoopOrchestrator`](../../src/app/service/agent/service_worker/tool_loop_orchestrator.ts) drives one conversation turn: call the model, execute any tool calls the model requested, feed results back, and repeat -until the model stops calling tools or `maxIterations` is hit. It depends on injected `callLLM` and +until the model stops calling tools (or the user stops it via Loop Guard / cancellation). It depends on injected `callLLM` and `autoCompact` functions (rather than importing a concrete client) so tests can substitute spies. [`retry_utils.ts`](../../src/app/service/agent/service_worker/retry_utils.ts)'s `isRetryableError` matches an error message containing `429`, a `5xx` code, or a network-ish signal (`network`/`fetch`/`ECONNRESET`), then diff --git a/packages/message/types.ts b/packages/message/types.ts index fea7719e6..1b323b8de 100644 --- a/packages/message/types.ts +++ b/packages/message/types.ts @@ -11,6 +11,8 @@ export type TMessagQueueUnit = { export type TMessageCommAction = { action: string; data?: NonNullable; + /** 可选请求关联 ID,用于连接上的请求/响应批次匹配。 */ + requestId?: string; msgQueue?: never; code?: never; }; diff --git a/src/app/repo/agent_chat.test.ts b/src/app/repo/agent_chat.test.ts index f3b1e2c6e..95b0dc1ce 100644 --- a/src/app/repo/agent_chat.test.ts +++ b/src/app/repo/agent_chat.test.ts @@ -10,6 +10,8 @@ function createMockOPFS() { data = content; }), close: vi.fn(async () => {}), + // 真实 OPFS 的 abort() 放弃这次写入的临时副本,不影响 dir 里已提交的旧内容 + abort: vi.fn(async () => {}), getData: () => data, }; } @@ -32,6 +34,11 @@ function createMockOPFS() { const written = writable.getData(); dir.set(name, written); await origClose(); + const forcedError = dir.get("__close_error__"); + if (forcedError) { + dir.delete("__close_error__"); + throw forcedError; + } }); return writable; }), @@ -46,14 +53,16 @@ function createMockOPFS() { if (opts?.create) { store.set("__dir__" + name, new Map()); } else { - throw new Error("Not found"); + throw new DOMException("A requested file or directory could not be found.", "NotFoundError"); } } return createMockDirHandle(store.get("__dir__" + name)); }), getFileHandle: vi.fn(async (name: string, opts?: { create?: boolean }) => { + const forcedError = store.get("__get_error__"); + if (forcedError) throw forcedError; if (!store.has(name) && !opts?.create) { - throw new Error("Not found"); + throw new DOMException("A requested file or directory could not be found.", "NotFoundError"); } if (!store.has(name)) { store.set(name, ""); @@ -61,6 +70,8 @@ function createMockOPFS() { return createMockFileHandle(name, store); }), removeEntry: vi.fn(async (name: string) => { + const forcedError = store.get("__remove_error__"); + if (forcedError) throw forcedError; store.delete(name); store.delete("__dir__" + name); }), @@ -202,51 +213,958 @@ describe("AgentChatRepo 附件存储", () => { expect(await repo.getAttachment("att-c")).toBeNull(); }); + it("附件读取与删除不得把权限错误伪装成不存在", async () => { + const uploadsDir = navigateDir(rootStore, "agents", "workspace", "uploads"); + uploadsDir.set("__get_error__", new DOMException("read denied", "NotAllowedError")); + await expect(repo.getAttachment("private.png")).rejects.toThrow("read denied"); + + uploadsDir.delete("__get_error__"); + uploadsDir.set("__remove_error__", new DOMException("delete denied", "NotAllowedError")); + await expect(repo.deleteAttachment("private.png")).rejects.toThrow("delete denied"); + }); + it("saveAttachment 纯文本(非 data URL)应作为 octet-stream 存储", async () => { const size = await repo.saveAttachment("att-5", "plain text content"); expect(size).toBeGreaterThan(0); }); + it("附件 close 已提交后报错时应通过大小读回确认成功", async () => { + const uploadsDir = navigateDir(rootStore, "agents", "workspace", "uploads"); + uploadsDir.set("__close_error__", new Error("ambiguous attachment close")); + + await expect(repo.saveAttachment("ambiguous.bin", new Blob(["durable"]))).resolves.toBe(7); + expect(await (await repo.getAttachment("ambiguous.bin"))!.text()).toBe("durable"); + }); + it("deleteConversation 应清理关联的附件", async () => { // 先保存会话和消息(含附件) const convId = "conv-1"; - await repo.saveConversation({ + const conversation = await repo.createConversation({ id: convId, title: "Test", modelId: "m1", createtime: Date.now(), updatetime: Date.now(), }); - await repo.saveMessages(convId, [ - { - id: "msg-1", - conversationId: convId, - role: "assistant", - content: "", - toolCalls: [ - { - id: "tc-1", - name: "screenshot", - arguments: "{}", - attachments: [ - { id: "att-del-1", type: "image", name: "img.jpg", mimeType: "image/jpeg" }, - { id: "att-del-2", type: "file", name: "file.zip", mimeType: "application/zip" }, - ], - }, - ], - createtime: Date.now(), - }, - ]); + await repo.saveMessages( + convId, + [ + { + id: "msg-1", + conversationId: convId, + role: "assistant", + content: "", + toolCalls: [ + { + id: "tc-1", + name: "screenshot", + arguments: "{}", + attachments: [ + { id: "att-del-1", type: "image", name: "img.jpg", mimeType: "image/jpeg" }, + { id: "att-del-2", type: "file", name: "file.zip", mimeType: "application/zip" }, + { id: "att-borrowed", type: "image", name: "borrowed.jpg", mimeType: "image/jpeg" }, + ], + ownedAttachmentIds: ["att-del-1", "att-del-2"], + }, + ], + createtime: Date.now(), + }, + { + id: "msg-borrowed", + conversationId: convId, + role: "user", + content: [{ type: "image", attachmentId: "att-borrowed", mimeType: "image/jpeg" }], + createtime: Date.now(), + }, + ], + undefined, + { generation: conversation.generation! } + ); // 保存附件数据 await repo.saveAttachment("att-del-1", new Blob(["img"])); await repo.saveAttachment("att-del-2", new Blob(["zip"])); + await repo.saveAttachment("att-borrowed", new Blob(["borrowed"])); // 删除会话 - await repo.deleteConversation(convId); + await repo.deleteConversation(convId, { + generation: conversation.generation!, + expectedRevision: conversation.revision, + }); // 附件应被清理 expect(await repo.getAttachment("att-del-1")).toBeNull(); expect(await repo.getAttachment("att-del-2")).toBeNull(); + expect(await repo.getAttachment("att-borrowed")).not.toBeNull(); + }); + + it("删除一个会话时不应删除其它会话借用的附件", async () => { + const owner = await repo.createConversation({ + id: "conv-attachment-owner", + title: "Owner", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const borrower = await repo.createConversation({ + id: "conv-attachment-borrower", + title: "Borrower", + modelId: "m1", + createtime: 2, + updatetime: 2, + }); + await repo.saveMessages( + owner.id, + [ + { + id: "owner-message", + conversationId: owner.id, + role: "user", + content: [{ type: "image", attachmentId: "shared.png", mimeType: "image/png" }], + ownedAttachmentIds: ["shared.png"], + createtime: 1, + }, + ], + undefined, + { generation: owner.generation! } + ); + await repo.saveMessages( + borrower.id, + [ + { + id: "borrower-message", + conversationId: borrower.id, + role: "user", + content: [{ type: "image", attachmentId: "shared.png", mimeType: "image/png" }], + createtime: 2, + }, + ], + undefined, + { generation: borrower.generation! } + ); + await repo.saveAttachment("shared.png", new Blob(["shared"])); + + await repo.deleteConversation(owner.id, { + generation: owner.generation!, + expectedRevision: owner.revision, + }); + + expect(await repo.getAttachment("shared.png")).not.toBeNull(); + }); + + it("同一会话的新快照借用旧附件时不应在所有权转移期间误删", async () => { + const conversation = await repo.createConversation({ + id: "conv-same-conversation-borrow", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveMessages( + conversation.id, + [ + { + id: "owner-message", + conversationId: conversation.id, + role: "user", + content: [{ type: "image", attachmentId: "same.png", mimeType: "image/png" }], + ownedAttachmentIds: ["same.png"], + createtime: 1, + }, + ], + undefined, + { generation: conversation.generation! } + ); + await repo.saveAttachment("same.png", new Blob(["image"])); + const previous = await repo.getMessageSnapshot(conversation.id, conversation.generation); + + await repo.saveMessages( + conversation.id, + [ + { + id: "borrower-message", + conversationId: conversation.id, + role: "user", + content: [{ type: "image", attachmentId: "same.png", mimeType: "image/png" }], + createtime: 2, + }, + ], + undefined, + { generation: conversation.generation!, expectedRevision: previous.revision } + ); + + expect(await repo.getAttachment("same.png")).not.toBeNull(); + }); + + it("替换历史时应递归清理被移除的子代理附件并保留仍被引用的附件", async () => { + const conversation = await repo.createConversation({ + id: "conv-nested-attachments", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const retained = { + id: "keep", + conversationId: conversation.id, + role: "user" as const, + content: [{ type: "image" as const, attachmentId: "att-keep", mimeType: "image/png" }], + ownedAttachmentIds: ["att-keep"], + createtime: 1, + }; + const nested = { + id: "nested", + conversationId: conversation.id, + role: "assistant" as const, + content: "", + toolCalls: [ + { + id: "tool-1", + name: "agent", + arguments: "{}", + ownedAttachmentIds: ["att-child", "att-child-tool"], + subAgentDetails: { + agentId: "child", + description: "child", + messages: [ + { + content: [{ type: "image", attachmentId: "att-child", mimeType: "image/png" }], + toolCalls: [ + { + id: "child-tool", + name: "image_generation", + arguments: "{}", + attachments: [{ id: "att-child-tool", type: "image", name: "child.png", mimeType: "image/png" }], + ownedAttachmentIds: ["att-child-tool"], + }, + ], + }, + ], + }, + }, + ], + createtime: 2, + } as any; + await repo.saveMessages(conversation.id, [retained, nested], undefined, { + generation: conversation.generation!, + }); + await Promise.all([ + repo.saveAttachment("att-keep", new Blob(["keep"])), + repo.saveAttachment("att-child", new Blob(["child"])), + repo.saveAttachment("att-child-tool", new Blob(["tool"])), + ]); + const snapshot = await repo.getMessageSnapshot(conversation.id, conversation.generation); + + await repo.saveMessages(conversation.id, [retained], undefined, { + generation: conversation.generation!, + expectedRevision: snapshot.revision, + }); + + expect(await repo.getAttachment("att-keep")).not.toBeNull(); + expect(await repo.getAttachment("att-child")).toBeNull(); + expect(await repo.getAttachment("att-child-tool")).toBeNull(); + }); + + it("重新生成转移消息所有权时应暂时保留指定附件", async () => { + const conversation = await repo.createConversation({ + id: "conv-transfer-attachment", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveMessages( + conversation.id, + [ + { + id: "owned-user-message", + conversationId: conversation.id, + role: "user", + content: [{ type: "image", attachmentId: "transfer.png", mimeType: "image/png" }], + ownedAttachmentIds: ["transfer.png"], + createtime: 1, + }, + ], + undefined, + { generation: conversation.generation! } + ); + await repo.saveAttachment("transfer.png", new Blob(["image"])); + const snapshot = await repo.getMessageSnapshot(conversation.id, conversation.generation); + + await repo.saveMessages(conversation.id, [], undefined, { + generation: conversation.generation!, + expectedRevision: snapshot.revision, + preserveAttachmentIds: ["transfer.png"], + }); + + expect(await repo.getAttachment("transfer.png")).not.toBeNull(); + }); +}); + +describe("AgentChatRepo.saveMessages 取消安全", () => { + let repo: AgentChatRepo; + + beforeEach(() => { + createMockOPFS(); + repo = new AgentChatRepo(); + }); + + it("signal 已 abort 时 saveMessages 应 reject 且不覆盖已持久化的旧消息", async () => { + const convId = "conv-cancel"; + const conversation = await repo.createConversation({ + id: convId, + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + const oldMessage = { + id: "msg-old", + conversationId: convId, + role: "user" as const, + content: "原始历史", + createtime: Date.now(), + }; + await repo.saveMessages(convId, [oldMessage], undefined, { generation: conversation.generation! }); + + const controller = new AbortController(); + controller.abort(); + await expect( + repo.saveMessages(convId, [{ ...oldMessage, id: "msg-new", content: "摘要覆盖" }], controller.signal) + ).rejects.toThrow("Aborted"); + + // 旧内容必须完整保留,没有被这次放弃的写入部分覆盖或破坏 + const stored = await repo.getMessages(convId); + expect(stored).toEqual([oldMessage]); + }); +}); + +describe("AgentChatRepo 跨上下文读-改-写安全", () => { + let repo: AgentChatRepo; + let rootStore: Map; + + beforeEach(() => { + ({ rootStore } = createMockOPFS()); + repo = new AgentChatRepo(); + }); + + it("并发 appendMessage 不应互相覆盖丢消息", async () => { + const convId = "conv-race"; + const conversation = await repo.createConversation({ + id: convId, + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + const makeMessage = (id: string) => ({ + id, + conversationId: convId, + role: "user" as const, + content: id, + createtime: Date.now(), + }); + + // 两个并发的读-改-写:无锁时双方都会读到空快照,后写者覆盖先写者 + await Promise.all([ + repo.appendMessage(makeMessage("m1"), conversation.generation), + repo.appendMessage(makeMessage("m2"), conversation.generation), + ]); + + const stored = await repo.getMessages(convId); + expect(stored.map((m) => m.id).sort()).toEqual(["m1", "m2"]); + }); + + it("相同消息 ID 的持久化重试不应生成重复记录", async () => { + const conversation = await repo.createConversation({ + id: "conv-idempotent", + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + const message = { + id: "stable-message", + conversationId: conversation.id, + role: "assistant" as const, + content: "done", + createtime: Date.now(), + }; + + await repo.appendMessage(message, conversation.generation); + await repo.appendMessage(message, conversation.generation); + + expect((await repo.getMessages(conversation.id)).filter((item) => item.id === message.id)).toHaveLength(1); + }); + + it("appendMessage close 已提交后报错时应读回确认消息及附件所有权", async () => { + const conversation = await repo.createConversation({ + id: "conv-ambiguous-append", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const message = { + id: "owned-message", + conversationId: conversation.id, + role: "user" as const, + content: [{ type: "image" as const, attachmentId: "owned.png", mimeType: "image/png" }], + ownedAttachmentIds: ["owned.png"], + createtime: 2, + }; + const originalWrite = (repo as any).writeJsonFile.bind(repo); + vi.spyOn(repo as any, "writeJsonFile").mockImplementationOnce(async (...args: unknown[]) => { + await originalWrite(...args); + throw new Error("ambiguous append close"); + }); + + await expect(repo.appendMessage(message, conversation.generation)).resolves.toMatchObject({ + messages: [message], + }); + }); + + it("saveTasks close 已提交后报错时应读回确认候选任务列表", async () => { + const conversation = await repo.createConversation({ + id: "conv-ambiguous-tasks", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const tasks = [{ id: "1", subject: "persist", status: "pending" as const }]; + const originalWrite = (repo as any).writeJsonFile.bind(repo); + vi.spyOn(repo as any, "writeJsonFile").mockImplementationOnce(async (...args: unknown[]) => { + await originalWrite(...args); + throw new Error("ambiguous task close"); + }); + + await expect(repo.saveTasks(conversation.id, tasks, undefined, conversation.generation)).resolves.toBeUndefined(); + await expect(repo.getTasks(conversation.id, conversation.generation)).resolves.toEqual(tasks); + }); + + it("saveTasks 应拒绝覆盖读取快照后已经变更的任务", async () => { + const conversation = await repo.createConversation({ + id: "conv-task-cas", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const currentTasks = [{ id: "1", subject: "current", status: "pending" as const }]; + await repo.saveTasks(conversation.id, currentTasks, undefined, conversation.generation); + + await expect(repo.saveTasks(conversation.id, [], undefined, conversation.generation, 0)).rejects.toThrow( + 'Tasks for conversation "conv-task-cas" changed' + ); + await expect(repo.getTasks(conversation.id, conversation.generation)).resolves.toEqual(currentTasks); + }); + + it("saveMessages close 已提交后报错时仍应继续清理被移除附件", async () => { + const conversation = await repo.createConversation({ + id: "conv-ambiguous-replace", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveMessages( + conversation.id, + [ + { + id: "owned-old", + conversationId: conversation.id, + role: "user", + content: [{ type: "image", attachmentId: "old.png", mimeType: "image/png" }], + ownedAttachmentIds: ["old.png"], + createtime: 1, + }, + ], + undefined, + { generation: conversation.generation! } + ); + await repo.saveAttachment("old.png", new Blob(["old"])); + const snapshot = await repo.getMessageSnapshot(conversation.id, conversation.generation); + const originalWrite = (repo as any).writeJsonFile.bind(repo); + vi.spyOn(repo as any, "writeJsonFile").mockImplementationOnce(async (...args: unknown[]) => { + await originalWrite(...args); + throw new Error("ambiguous replacement close"); + }); + + await expect( + repo.saveMessages(conversation.id, [], undefined, { + generation: conversation.generation!, + expectedRevision: snapshot.revision, + }) + ).resolves.toMatchObject({ messages: [] }); + expect(await repo.getAttachment("old.png")).toBeNull(); + }); + + it("并发 saveConversation 不应互相覆盖丢会话", async () => { + const makeConv = (id: string) => ({ + id, + title: id, + modelId: "m", + createtime: Date.now(), + updatetime: Date.now(), + }); + + await Promise.all([repo.createConversation(makeConv("c1")), repo.createConversation(makeConv("c2"))]); + + const stored = await repo.listConversations(); + expect(stored.map((c) => c.id).sort()).toEqual(["c1", "c2"]); + }); + + it("消息文件损坏时读取应抛错,而不是让后续写入基于空快照覆盖旧数据", async () => { + const conversation = await repo.createConversation({ + id: "conv-corrupt", + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + const messagesDir = navigateDir(rootStore, "agents", "conversations", "data"); + messagesDir.set("conv-corrupt.json", "{ 损坏的 JSON"); + + await expect(repo.getMessages("conv-corrupt")).rejects.toThrow(); + // appendMessage 的读阶段同样必须失败,绝不能把损坏文件当作空历史整份覆写 + await expect( + repo.appendMessage( + { + id: "m1", + conversationId: "conv-corrupt", + role: "user", + content: "hi", + createtime: Date.now(), + }, + conversation.generation + ) + ).rejects.toThrow(); + expect(messagesDir.get("conv-corrupt.json")).toBe("{ 损坏的 JSON"); + }); + + it("文件尚未创建(NotFoundError)时读取返回默认值", async () => { + await expect(repo.getMessages("conv-none")).resolves.toEqual([]); + }); + + it("支持 Web Locks 的环境下写操作应在 navigator.locks 排它锁内执行", async () => { + const request = vi.fn(async (_name: string, _opts: unknown, fn: () => Promise) => fn()); + Object.defineProperty(navigator, "locks", { + value: { request }, + configurable: true, + writable: true, + }); + try { + const conversation = await repo.createConversation({ + id: "conv-lock", + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + await repo.appendMessage( + { + id: "m1", + conversationId: "conv-lock", + role: "user", + content: "hi", + createtime: Date.now(), + }, + conversation.generation + ); + expect(request).toHaveBeenCalledWith( + expect.stringContaining("conv-lock"), + expect.objectContaining({ mode: "exclusive" }), + expect.any(Function) + ); + } finally { + // @ts-expect-error 清理测试注入的 locks + delete navigator.locks; + } + }); + + it("删除后延迟的元数据和消息写入都不应复活旧 generation", async () => { + const conversation = await repo.createConversation({ + id: "conv-deleted", + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + await repo.deleteConversation(conversation.id, { + generation: conversation.generation!, + expectedRevision: conversation.revision, + }); + + conversation.title = "stale rename"; + await expect(repo.saveConversation(conversation)).rejects.toThrow("deleted"); + await expect( + repo.appendMessage( + { + id: "late", + conversationId: conversation.id, + role: "assistant", + content: "late", + createtime: Date.now(), + }, + conversation.generation + ) + ).rejects.toThrow("deleted"); + expect(await repo.listConversations()).toEqual([]); + expect(await repo.getMessages(conversation.id)).toEqual([]); + }); + + it("会话元数据删除已提交后 close 报错时仍应完成子数据清理", async () => { + const conversation = await repo.createConversation({ + id: "conv-ambiguous-delete", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveMessages( + conversation.id, + [ + { + id: "owned-before-delete", + conversationId: conversation.id, + role: "user", + content: [{ type: "image", attachmentId: "delete-me.png", mimeType: "image/png" }], + ownedAttachmentIds: ["delete-me.png"], + createtime: 1, + }, + ], + undefined, + { generation: conversation.generation! } + ); + await repo.saveAttachment("delete-me.png", new Blob(["delete"])); + const originalWrite = (repo as any).writeJsonFile.bind(repo); + vi.spyOn(repo as any, "writeJsonFile").mockImplementationOnce(async (...args: unknown[]) => { + await originalWrite(...args); + throw new Error("ambiguous delete close"); + }); + + await expect( + repo.deleteConversation(conversation.id, { + generation: conversation.generation!, + expectedRevision: conversation.revision, + }) + ).resolves.toBeUndefined(); + expect(await repo.listConversations()).toEqual([]); + expect(await repo.getAttachment("delete-me.png")).toBeNull(); + }); + + it("deleteConversation 的附件/消息 GC 真正失败时仍应报告删除成功", async () => { + const conversation = await repo.createConversation({ + id: "conv-gc-fail-delete", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + vi.spyOn(repo, "deleteAttachments").mockRejectedValueOnce(new Error("disk error")); + + // 元数据删除本身(主提交)已经完成,附件/消息清理失败只是 GC 债务,不应报告整个删除失败 + await expect( + repo.deleteConversation(conversation.id, { + generation: conversation.generation!, + expectedRevision: conversation.revision, + }) + ).resolves.toBeUndefined(); + expect(await repo.listConversations()).toEqual([]); + }); + + it("createConversation 复用 ID 时旧数据清理失败不应报告创建失败", async () => { + const original = await repo.createConversation({ + id: "conv-reused", + title: "Old", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + // 先正常删除,留下同 id 可被复用;被删除的旧一代消息/任务文件在真实场景里可能残留 + await repo.deleteConversation("conv-reused", { + generation: original.generation!, + expectedRevision: original.revision, + }); + vi.spyOn(repo as any, "deleteFile").mockRejectedValue(new Error("disk error")); + + // 复用 id 时旧一代的消息/任务文件清理失败仅是 GC 债务;conversations.json 的插入才是主提交 + const recreated = await repo.createConversation({ + id: "conv-reused", + title: "New", + modelId: "m1", + createtime: 2, + updatetime: 2, + }); + expect(recreated.id).toBe("conv-reused"); + const list = await repo.listConversations(); + expect(list).toHaveLength(1); + expect(list[0].title).toBe("New"); + }); + + it("会话 ID 复用且旧任务文件清理失败时,getTasks 不应读到上一代的任务", async () => { + const original = await repo.createConversation({ + id: "conv-task-reused", + title: "Old", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveTasks( + original.id, + [{ id: "task-1", title: "旧一代任务", status: "pending" } as any], + undefined, + original.generation + ); + // 删除会话本身成功,但旧任务文件的清理失败(GC 债务),文件残留在磁盘上 + vi.spyOn(repo as any, "deleteFile").mockRejectedValue(new Error("disk error")); + await repo.deleteConversation("conv-task-reused", { + generation: original.generation!, + expectedRevision: original.revision, + }); + + // 复用同一个 ID 创建新一代会话 + const recreated = await repo.createConversation({ + id: "conv-task-reused", + title: "New", + modelId: "m1", + createtime: 2, + updatetime: 2, + }); + expect(recreated.generation).not.toBe(original.generation); + + // 新一代会话读取任务时不应看到残留的旧一代任务,而应得到空列表 + const tasks = await repo.getTasks(recreated.id, recreated.generation); + expect(tasks).toEqual([]); + }); + + it("会话 ID 复用且残留原始 Task 数组时,getTasks 不应复活上一代任务", async () => { + const original = await repo.createConversation({ + id: "conv-raw-task-reused", + title: "Old", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const tasksDir = navigateDir(rootStore, "agents", "conversations", "tasks"); + tasksDir.set(`${original.id}.json`, JSON.stringify([{ id: "task-raw", title: "旧原始任务", status: "pending" }])); + vi.spyOn(repo as any, "deleteFile").mockRejectedValue(new Error("disk error")); + await repo.deleteConversation(original.id, { + generation: original.generation!, + expectedRevision: original.revision, + }); + + const recreated = await repo.createConversation({ + id: original.id, + title: "New", + modelId: "m1", + createtime: 2, + updatetime: 2, + }); + + await expect(repo.getTasks(recreated.id, recreated.generation)).resolves.toEqual([]); + }); + + it("saveMessages 提交新快照后附件 GC 失败不应报告 clear/compact 失败", async () => { + const conversation = await repo.createConversation({ + id: "conv-gc-fail-save", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveMessages( + conversation.id, + [ + { + id: "msg-with-attachment", + conversationId: conversation.id, + role: "user", + content: [{ type: "image", attachmentId: "removed.png", mimeType: "image/png" }], + ownedAttachmentIds: ["removed.png"], + createtime: 1, + }, + ], + undefined, + { generation: conversation.generation! } + ); + vi.spyOn(repo, "deleteAttachments").mockRejectedValueOnce(new Error("disk error")); + + // 新的(空)快照已经提交;被移除消息引用的附件清理失败不应让 clear 报告失败 + const saved = await repo.saveMessages(conversation.id, [], undefined, { generation: conversation.generation! }); + expect(saved.messages).toEqual([]); + expect(await repo.getMessages(conversation.id)).toEqual([]); + }); + + it("升级前的历史会话(legacy generation)删除时应按 content block 推断清理旧附件", async () => { + // 直接写入没有 generation/revision 字段的会话记录,模拟所有权模型引入之前创建的历史数据 + await (repo as any).writeJsonFile("conversations.json", [ + { id: "conv-legacy-del", title: "Test", modelId: "m1", createtime: 1, updatetime: 1 }, + ]); + const [conv] = await repo.listConversations(); + expect(conv.generation).toBe("legacy:conv-legacy-del"); + + // 历史消息只有 content block 引用附件,从未写入过 ownedAttachmentIds(该字段是本次新增的) + await repo.saveMessages( + conv.id, + [ + { + id: "legacy-msg", + conversationId: conv.id, + role: "user", + content: [{ type: "image", attachmentId: "legacy-owned.png", mimeType: "image/png" }], + createtime: 1, + }, + ], + undefined, + { generation: conv.generation! } + ); + await repo.saveAttachment("legacy-owned.png", new Blob(["legacy"])); + + await repo.deleteConversation(conv.id, { generation: conv.generation!, expectedRevision: conv.revision }); + + // 升级前的历史必须能按 content block 推断出所有权,否则这类附件永远不会被清理 + expect(await repo.getAttachment("legacy-owned.png")).toBeNull(); + }); + + it("非 legacy 会话中未声明所有权的 content block 引用应保持借用语义,不因清理被误删", async () => { + const conversation = await repo.createConversation({ + id: "conv-current-borrow", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + await repo.saveMessages( + conversation.id, + [ + { + id: "msg-borrowed", + conversationId: conversation.id, + role: "user", + // 当前模型下 undefined ownedAttachmentIds 合法地表示"借用",不代表遗留数据 + content: [{ type: "image", attachmentId: "shared.png", mimeType: "image/png" }], + createtime: 1, + }, + ], + undefined, + { generation: conversation.generation! } + ); + await repo.saveAttachment("shared.png", new Blob(["shared"])); + + // 用一次空快照替换(模拟 clear):借用引用不应被当作该会话的"已拥有附件"而删除 + await repo.saveMessages(conversation.id, [], undefined, { generation: conversation.generation! }); + + expect(await repo.getAttachment("shared.png")).not.toBeNull(); + }); + + it("历史替换应以 revision 做 CAS,不能覆盖并发追加", async () => { + const conversation = await repo.createConversation({ + id: "conv-cas", + title: "Test", + modelId: "m1", + createtime: Date.now(), + updatetime: Date.now(), + }); + await repo.appendMessage( + { id: "m1", conversationId: conversation.id, role: "user", content: "old", createtime: 1 }, + conversation.generation + ); + const stale = await repo.getMessageSnapshot(conversation.id, conversation.generation); + await repo.appendMessage( + { id: "m2", conversationId: conversation.id, role: "assistant", content: "fresh", createtime: 2 }, + conversation.generation + ); + + await expect( + repo.saveMessages(conversation.id, [], undefined, { + generation: conversation.generation!, + expectedRevision: stale.revision, + }) + ).rejects.toThrow("changed"); + expect((await repo.getMessages(conversation.id)).map((message) => message.id)).toEqual(["m1", "m2"]); + }); + + it("工具调用 assistant 与全部 tool 结果应在一次历史 revision 中提交", async () => { + const conversation = await repo.createConversation({ + id: "conv-tool-round", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const before = await repo.getMessageSnapshot(conversation.id, conversation.generation); + await repo.commitToolRound( + { + id: "assistant", + conversationId: conversation.id, + role: "assistant", + content: "", + toolCalls: [ + { id: "call-1", name: "one", arguments: "{}", status: "completed" }, + { id: "call-2", name: "two", arguments: "{}", status: "error" }, + ], + createtime: 1, + }, + [ + { + id: "tool-1", + conversationId: conversation.id, + role: "tool", + content: "ok", + toolCallId: "call-1", + createtime: 2, + }, + { + id: "tool-2", + conversationId: conversation.id, + role: "tool", + content: "failed", + toolCallId: "call-2", + createtime: 3, + }, + ], + conversation.generation + ); + + const after = await repo.getMessageSnapshot(conversation.id, conversation.generation); + expect(after.revision).toBe(before.revision + 1); + expect(after.messages.map((message) => message.id)).toEqual(["assistant", "tool-1", "tool-2"]); + }); + + it("工具轮次 close 报错但读回确认完整时应按已提交成功处理", async () => { + const conversation = await repo.createConversation({ + id: "conv-ambiguous-tool-round", + title: "Test", + modelId: "m1", + createtime: 1, + updatetime: 1, + }); + const assistant = { + id: "assistant-ambiguous", + conversationId: conversation.id, + role: "assistant" as const, + content: "", + toolCalls: [{ id: "call-1", name: "tool", arguments: "{}" }], + createtime: 2, + }; + const toolMessage = { + id: "tool-ambiguous", + conversationId: conversation.id, + role: "tool" as const, + content: "done", + toolCallId: "call-1", + createtime: 3, + }; + const originalWrite = (repo as any).writeJsonFile.bind(repo); + vi.spyOn(repo as any, "writeJsonFile").mockImplementationOnce(async (...args: unknown[]) => { + await originalWrite(...args); + throw new Error("ambiguous close failure"); + }); + + await expect(repo.commitToolRound(assistant, [toolMessage], conversation.generation)).resolves.toMatchObject({ + messages: [assistant, toolMessage], + }); }); }); diff --git a/src/app/repo/agent_chat.ts b/src/app/repo/agent_chat.ts index b91d1da8d..88fbaf1ff 100644 --- a/src/app/repo/agent_chat.ts +++ b/src/app/repo/agent_chat.ts @@ -1,13 +1,144 @@ -import type { Conversation, ChatMessage } from "@App/app/service/agent/core/types"; +import type { Conversation, ChatMessage, MessageContent } from "@App/app/service/agent/core/types"; import type { Task } from "@App/app/service/agent/core/tools/task_tools"; +import { isLegacyGeneration } from "@App/app/service/agent/core/persisted_messages"; import { OPFSRepo } from "./opfs_repo"; -import { writeWorkspaceFile, getWorkspaceRoot, getDirectory } from "@App/app/service/agent/core/opfs_helpers"; +import { + decodeDataUrl, + getDirectory, + getWorkspaceRoot, + isDataUrl, + writeWorkspaceFile, +} from "@App/app/service/agent/core/opfs_helpers"; +import { uuidv4 } from "@App/pkg/utils/uuid"; +import { RevisionConflictError } from "./revision"; const CONVERSATIONS_FILE = "conversations.json"; const MESSAGES_DIR = "data"; const ATTACHMENTS_DIR = "attachments"; const TASKS_DIR = "tasks"; +function isNotFoundError(error: unknown): boolean { + return (error as { name?: string })?.name === "NotFoundError"; +} + +export type MessageSnapshot = { + generation: string; + revision: number; + messages: ChatMessage[]; +}; + +export type TaskSnapshot = { + generation: string; + revision: number; + tasks: Task[]; +}; + +export type ConversationMutationGuard = { + generation: string; + expectedRevision?: number; + preserveAttachmentIds?: string[]; +}; + +function normalizeConversation(conversation: Conversation): Conversation { + return { + ...conversation, + generation: conversation.generation || `legacy:${conversation.id}`, + revision: conversation.revision ?? 0, + }; +} + +function isMessageSnapshot(value: ChatMessage[] | MessageSnapshot): value is MessageSnapshot { + return !Array.isArray(value); +} + +function isTaskSnapshot(value: Task[] | TaskSnapshot): value is TaskSnapshot { + return !Array.isArray(value); +} + +// content block 里携带的附件 id(image/file/audio block 均以 attachmentId 引用) +function contentBlockAttachmentIds(content: MessageContent): string[] { + if (typeof content === "string") return []; + return content + .filter((block) => block.type !== "text") + .map((block) => (block as { attachmentId: string }).attachmentId); +} + +// ownedAttachmentIds 是本次所有权模型新增的可选字段。当前模型里 undefined 与显式空数组 +// 都合法地表示"这条消息/工具调用不拥有任何附件、content block 引用的都是借用"(比如引用 +// 另一条消息已拥有的图片),因此不能仅凭字段缺失就判断为"升级前的历史数据"从而反推所有权, +// 否则会把当前模型里正常的借用误判为拥有,导致这类附件被其它消息删除时连带误删(回归)。 +// 只有当调用方能证明这条历史确实创建于所有权模型引入之前(例如 generation 是本仓库为 +// 无 generation 的旧数据统一回填的 "legacy:" 前缀)时,才允许把 ownedAttachmentIds 缺失 +// 解读为"从未写入过",退化为按 content block / 工具附件元数据推断所有权。 +export function collectMessageAttachmentIds(messages: ChatMessage[], legacy = false): Set { + const result = new Set(); + const collectToolCalls = (toolCalls: NonNullable) => { + for (const toolCall of toolCalls) { + if (toolCall.ownedAttachmentIds !== undefined) { + for (const attachmentId of toolCall.ownedAttachmentIds) result.add(attachmentId); + } else if (legacy) { + for (const attachment of toolCall.attachments || []) result.add(attachment.id); + } + for (const subMessage of toolCall.subAgentDetails?.messages || []) { + // SubAgentMessage 从未有过 ownedAttachmentIds 字段,只有 legacy 历史才按引用推断, + // 避免把当前模型里子代理消息中正常借用的附件误判为拥有 + if (legacy) { + for (const attachmentId of contentBlockAttachmentIds(subMessage.content)) result.add(attachmentId); + } + collectToolCalls(subMessage.toolCalls); + } + } + }; + for (const message of messages) { + if (message.ownedAttachmentIds !== undefined) { + for (const attachmentId of message.ownedAttachmentIds) result.add(attachmentId); + } else if (legacy) { + for (const attachmentId of contentBlockAttachmentIds(message.content)) result.add(attachmentId); + } + collectToolCalls(message.toolCalls || []); + } + return result; +} + +// GC 引用扫描需要包含借用附件;与所有权收集不同,这里记录消息内容和工具结果中的全部引用。 +function collectMessageAttachmentReferences(messages: ChatMessage[], legacy = false): Set { + const result = collectMessageAttachmentIds(messages, legacy); + const collectToolCalls = (toolCalls: NonNullable) => { + for (const toolCall of toolCalls) { + for (const attachment of toolCall.attachments || []) result.add(attachment.id); + for (const subMessage of toolCall.subAgentDetails?.messages || []) { + for (const attachmentId of contentBlockAttachmentIds(subMessage.content)) result.add(attachmentId); + collectToolCalls(subMessage.toolCalls); + } + } + }; + for (const message of messages) { + for (const attachmentId of contentBlockAttachmentIds(message.content)) result.add(attachmentId); + collectToolCalls(message.toolCalls || []); + } + return result; +} + +function containsCompleteToolRound( + messages: ChatMessage[], + assistantMessage: ChatMessage, + toolMessages: ChatMessage[] +): boolean { + const assistant = messages.find((message) => message.id === assistantMessage.id && message.role === "assistant"); + if (!assistant) return false; + const expectedToolCallIds = new Set((assistantMessage.toolCalls || []).map((toolCall) => toolCall.id)); + if (toolMessages.length !== expectedToolCallIds.size) return false; + return toolMessages.every( + (expected) => + expected.role === "tool" && + expected.toolCallId !== undefined && + expectedToolCallIds.has(expected.toolCallId) && + messages.some( + (message) => message.id === expected.id && message.role === "tool" && message.toolCallId === expected.toolCallId + ) + ); +} + // 目录结构:agents/conversations/ // agents/conversations/conversations.json - 会话列表 // agents/conversations/data/{id}.json - 每个会话的消息 @@ -18,91 +149,314 @@ export class AgentChatRepo extends OPFSRepo { super("conversations"); } + private async writeJsonFileConfirmed( + filename: string, + data: T, + defaultValue: T, + dir?: FileSystemDirectoryHandle, + signal?: AbortSignal + ): Promise { + try { + await this.writeJsonFile(filename, data, dir, signal); + } catch (error) { + try { + const durable = await this.readJsonFile(filename, defaultValue, dir); + if (JSON.stringify(durable) === JSON.stringify(data)) return; + } catch { + // Preserve the original write error when read-back confirmation itself fails. + } + throw error; + } + } + // 获取所有会话列表 async listConversations(): Promise { - return this.readJsonFile(CONVERSATIONS_FILE, []); + return (await this.readJsonFile(CONVERSATIONS_FILE, [])).map(normalizeConversation); } - // 保存/更新会话 - async saveConversation(conversation: Conversation): Promise { - const conversations = await this.readJsonFile(CONVERSATIONS_FILE, []); - const index = conversations.findIndex((c) => c.id === conversation.id); - if (index >= 0) { - conversations[index] = conversation; - } else { - conversations.unshift(conversation); - } - await this.writeJsonFile(CONVERSATIONS_FILE, conversations); + async createConversation(conversation: Conversation): Promise { + return this.withFileLock(`lifecycle:${conversation.id}`, async () => { + const created = normalizeConversation({ ...conversation, generation: uuidv4(), revision: 1 }); + await this.withFileLock(CONVERSATIONS_FILE, async () => { + const conversations = (await this.readJsonFile(CONVERSATIONS_FILE, [])).map( + normalizeConversation + ); + if (conversations.some((item) => item.id === created.id)) { + throw new RevisionConflictError(`Conversation "${created.id}" already exists`); + } + conversations.unshift(created); + await this.writeJsonFileConfirmed(CONVERSATIONS_FILE, conversations, []); + }); + + // Reusing an explicitly supplied ID starts a fresh generation with no legacy child state. + // The conversation record is already durably committed above; a failure here is garbage-collection + // debt, not a creation failure — swallow it so the caller (and any retry) sees the conversation as + // created instead of conflicting with the record that already exists. The stale files + // stay inert: readMessageSnapshot/getTasks compare against the new generation and ignore them. + try { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + const tasksDir = await this.getChildDir(TASKS_DIR); + await this.deleteFile(`${created.id}.json`, messagesDir); + await this.deleteFile(`${created.id}.json`, tasksDir); + } catch { + // best-effort cleanup of the previous generation's leftover files + } + return created; + }); + } + + // 更新现有会话。generation/revision 都必须匹配,绝不以 upsert 语义复活已删除记录。 + // conversations.json 被 Options 页与 Service Worker 两个上下文共享,所有读-改-写 + // 都必须在同一把跨上下文排它锁内执行,否则双方会基于同一旧快照互相覆盖 + async saveConversation(conversation: Conversation): Promise { + return this.withFileLock(`lifecycle:${conversation.id}`, async () => { + return this.withFileLock(CONVERSATIONS_FILE, async () => { + const conversations = (await this.readJsonFile(CONVERSATIONS_FILE, [])).map( + normalizeConversation + ); + const index = conversations.findIndex((item) => item.id === conversation.id); + const current = index >= 0 ? conversations[index] : undefined; + if ( + !current || + !conversation.generation || + conversation.generation !== current.generation || + conversation.revision !== current.revision + ) { + throw new RevisionConflictError(`Conversation "${conversation.id}" changed or was deleted`); + } + const saved = normalizeConversation({ ...conversation, revision: current.revision! + 1 }); + conversations[index] = saved; + await this.writeJsonFileConfirmed(CONVERSATIONS_FILE, conversations, []); + Object.assign(conversation, saved); + return saved; + }); + }); } // 删除会话及其消息和附件 - async deleteConversation(id: string): Promise { - // 清理会话关联的附件 - const messages = await this.getMessages(id); - const attachmentIds: string[] = []; - for (const msg of messages) { - // 扫描 toolCalls 中的附件 - if (msg.toolCalls) { - for (const tc of msg.toolCalls) { - if (tc.attachments) { - for (const att of tc.attachments) { - attachmentIds.push(att.id); - } - } + async deleteConversation(id: string, guard?: ConversationMutationGuard): Promise { + await this.withFileLock(`lifecycle:${id}`, async () => { + let deleted: Conversation | undefined; + await this.withFileLock(CONVERSATIONS_FILE, async () => { + const conversations = (await this.readJsonFile(CONVERSATIONS_FILE, [])).map( + normalizeConversation + ); + const index = conversations.findIndex((item) => item.id === id); + const current = index >= 0 ? conversations[index] : undefined; + if (!current) return; + if ( + guard && + (current.generation !== guard.generation || + (guard.expectedRevision !== undefined && current.revision !== guard.expectedRevision)) + ) { + throw new RevisionConflictError(`Conversation "${id}" changed before deletion`); } + deleted = current; + conversations.splice(index, 1); + await this.writeJsonFileConfirmed(CONVERSATIONS_FILE, conversations, []); + }); + if (!deleted) return; + + // 会话记录已经在上面提交删除,这里是善后 GC:失败不能让 deleteConversation 报告失败—— + // 那样调用方重试时 conversations.json 里已经找不到该 id(见上面 `if (!current) return`), + // 会静默当作删除成功返回,剩余的消息/任务/附件却永远不会被清理。 + try { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + const stored = await this.readMessageSnapshot(id, deleted.generation!, messagesDir); + await this.deleteUnreferencedAttachments( + [...collectMessageAttachmentIds(stored.messages, isLegacyGeneration(deleted.generation))], + id + ); + await this.deleteFile(`${id}.json`, messagesDir); + const tasksDir = await this.getChildDir(TASKS_DIR); + await this.deleteFile(`${id}.json`, tasksDir); + } catch { + // best-effort cleanup of messages/tasks/attachments for the now-deleted conversation } - // 扫描 ContentBlock[] 中的附件 - if (Array.isArray(msg.content)) { - for (const block of msg.content) { - if (block.type !== "text" && "attachmentId" in block) { - attachmentIds.push(block.attachmentId); + }); + } + + // 获取指定会话的所有消息 + async getMessages(conversationId: string): Promise { + try { + return (await this.getMessageSnapshot(conversationId)).messages; + } catch (error) { + if (error instanceof RevisionConflictError) return []; + throw error; + } + } + + async getMessageSnapshot(conversationId: string, generation?: string): Promise { + return this.withFileLock(`lifecycle:${conversationId}`, async () => { + const current = await this.requireConversation(conversationId, generation); + return this.withFileLock(`messages:${conversationId}`, async () => { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + return this.readMessageSnapshot(conversationId, current.generation!, messagesDir); + }); + }); + } + + // 追加消息(读-改-写,须持有该会话消息文件的跨上下文排它锁) + async appendMessage(message: ChatMessage, generation?: string): Promise { + return this.withFileLock(`lifecycle:${message.conversationId}`, async () => { + const current = await this.requireConversation(message.conversationId, generation); + return this.withFileLock(`messages:${message.conversationId}`, async () => { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + const snapshot = await this.readMessageSnapshot(message.conversationId, current.generation!, messagesDir); + // Callers retry final-message persistence with the same stable ID. If close() committed but surfaced an + // ambiguous error, the retry must observe the committed message instead of appending a duplicate. + if (snapshot.messages.some((item) => item.id === message.id)) return snapshot; + const saved = { ...snapshot, revision: snapshot.revision + 1, messages: [...snapshot.messages, message] }; + await this.writeJsonFileConfirmed( + `${message.conversationId}.json`, + saved, + { generation: current.generation!, revision: 0, messages: [] }, + messagesDir + ); + return saved; + }); + }); + } + + // 更新消息(按 id 匹配;读-改-写,同上须持锁) + async updateMessage(message: ChatMessage, generation?: string): Promise { + return this.withFileLock(`lifecycle:${message.conversationId}`, async () => { + const current = await this.requireConversation(message.conversationId, generation); + return this.withFileLock(`messages:${message.conversationId}`, async () => { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + const snapshot = await this.readMessageSnapshot(message.conversationId, current.generation!, messagesDir); + const messages = [...snapshot.messages]; + const index = messages.findIndex((item) => item.id === message.id); + if (index < 0) return snapshot; + messages[index] = message; + const saved = { ...snapshot, revision: snapshot.revision + 1, messages }; + await this.writeJsonFileConfirmed( + `${message.conversationId}.json`, + saved, + { generation: current.generation!, revision: 0, messages: [] }, + messagesDir + ); + return saved; + }); + }); + } + + /** Persist one assistant tool-call message and its complete tool-result group in one file commit. */ + async commitToolRound( + assistantMessage: ChatMessage, + toolMessages: ChatMessage[], + generation?: string + ): Promise { + const conversationId = assistantMessage.conversationId; + return this.withFileLock(`lifecycle:${conversationId}`, async () => { + const current = await this.requireConversation(conversationId, generation); + return this.withFileLock(`messages:${conversationId}`, async () => { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + const snapshot = await this.readMessageSnapshot(conversationId, current.generation!, messagesDir); + const groupIds = new Set([assistantMessage.id, ...toolMessages.map((message) => message.id)]); + const messages = snapshot.messages.filter((message) => !groupIds.has(message.id)); + messages.push(assistantMessage, ...toolMessages); + const saved = { ...snapshot, revision: snapshot.revision + 1, messages }; + try { + await this.writeJsonFile(`${conversationId}.json`, saved, messagesDir); + return saved; + } catch (error) { + // OPFS close() may atomically commit and still surface an error. Read back before reporting failure so + // callers never delete attachments that the durable round already references. + try { + const committed = await this.readMessageSnapshot(conversationId, current.generation!, messagesDir); + if (containsCompleteToolRound(committed.messages, assistantMessage, toolMessages)) return committed; + } catch { + // Preserve the original commit error when confirmation itself fails. } + throw error; } - } - } - if (attachmentIds.length > 0) { - await this.deleteAttachments(attachmentIds); - } + }); + }); + } - const conversations = await this.readJsonFile(CONVERSATIONS_FILE, []); - const filtered = conversations.filter((c) => c.id !== id); - await this.writeJsonFile(CONVERSATIONS_FILE, filtered); - // 删除对应消息文件 - const messagesDir = await this.getChildDir(MESSAGES_DIR); - await this.deleteFile(`${id}.json`, messagesDir); - // 删除关联的任务数据 - await this.deleteTasks(id).catch(() => {}); + // 保存整个消息列表(用于批量更新)。整份覆写虽无读阶段,但仍须与其它读-改-写同锁排队, + // 否则可能穿插进别人临界区的读与写之间。 + // signal 可选:传入时若在写入落定前已 abort,则放弃这次整份覆写而不是让它继续提交 + // (OPFS createWritable() 本身是事务性的,写入的是临时副本,abort 不会影响已持久化的旧内容) + async saveMessages( + conversationId: string, + messages: ChatMessage[], + signal?: AbortSignal, + guard?: ConversationMutationGuard + ): Promise { + return this.withFileLock(`lifecycle:${conversationId}`, async () => { + const current = await this.requireConversation(conversationId, guard?.generation); + return this.withFileLock(`messages:${conversationId}`, async () => { + const messagesDir = await this.getChildDir(MESSAGES_DIR); + const snapshot = await this.readMessageSnapshot(conversationId, current.generation!, messagesDir); + if (guard?.expectedRevision !== undefined && snapshot.revision !== guard.expectedRevision) { + throw new RevisionConflictError(`Messages for conversation "${conversationId}" changed`); + } + const saved = { generation: current.generation!, revision: snapshot.revision + 1, messages }; + await this.writeJsonFileConfirmed( + `${conversationId}.json`, + saved, + { generation: current.generation!, revision: 0, messages: [] }, + messagesDir, + signal + ); + const legacy = isLegacyGeneration(current.generation); + const retainedAttachments = collectMessageAttachmentIds(messages, legacy); + for (const attachmentId of guard?.preserveAttachmentIds || []) retainedAttachments.add(attachmentId); + const removedAttachments = [...collectMessageAttachmentIds(snapshot.messages, legacy)].filter( + (id) => !retainedAttachments.has(id) + ); + // 新快照已经提交在上面;删除不再被引用的旧附件属于 GC,失败不能让 clear/compact/deleteMessages + // 报告失败——历史其实已经替换成功了。 + try { + await this.deleteUnreferencedAttachments( + removedAttachments, + conversationId, + collectMessageAttachmentReferences(messages, legacy) + ); + } catch { + // best-effort cleanup of attachments no longer referenced by the saved history + } + return saved; + }); + }); } - // 获取指定会话的所有消息 - async getMessages(conversationId: string): Promise { - const messagesDir = await this.getChildDir(MESSAGES_DIR); - return this.readJsonFile(`${conversationId}.json`, [], messagesDir); - } - - // 追加消息 - async appendMessage(message: ChatMessage): Promise { - const messagesDir = await this.getChildDir(MESSAGES_DIR); - const messages = await this.readJsonFile(`${message.conversationId}.json`, [], messagesDir); - messages.push(message); - await this.writeJsonFile(`${message.conversationId}.json`, messages, messagesDir); - } - - // 更新消息(按 id 匹配) - async updateMessage(message: ChatMessage): Promise { - const messagesDir = await this.getChildDir(MESSAGES_DIR); - const messages = await this.readJsonFile(`${message.conversationId}.json`, [], messagesDir); - const index = messages.findIndex((m) => m.id === message.id); - if (index >= 0) { - messages[index] = message; - await this.writeJsonFile(`${message.conversationId}.json`, messages, messagesDir); + private async requireConversation(id: string, generation?: string): Promise { + const conversations = (await this.readJsonFile(CONVERSATIONS_FILE, [])).map(normalizeConversation); + const current = conversations.find((item) => item.id === id); + if (!current || (generation !== undefined && current.generation !== generation)) { + throw new RevisionConflictError(`Conversation "${id}" changed or was deleted`); } + return current; + } + + private async readMessageSnapshot( + conversationId: string, + generation: string, + messagesDir: FileSystemDirectoryHandle + ): Promise { + const stored = await this.readJsonFile(`${conversationId}.json`, [], messagesDir); + if (!isMessageSnapshot(stored)) return { generation, revision: 0, messages: stored }; + if (stored.generation !== generation) return { generation, revision: 0, messages: [] }; + return stored; } - // 保存整个消息列表(用于批量更新) - async saveMessages(conversationId: string, messages: ChatMessage[]): Promise { - const messagesDir = await this.getChildDir(MESSAGES_DIR); - await this.writeJsonFile(`${conversationId}.json`, messages, messagesDir); + private async readTaskSnapshot( + conversationId: string, + generation: string, + tasksDir: FileSystemDirectoryHandle + ): Promise { + const stored = await this.readJsonFile(`${conversationId}.json`, [], tasksDir); + // 旧格式只有在 legacy 会话中才属于当前会话;新 generation 不得复活同 ID 的旧文件。 + if (!isTaskSnapshot(stored)) { + return isLegacyGeneration(generation) + ? { generation, revision: 0, tasks: stored } + : { generation, revision: 0, tasks: [] }; + } + if (stored.generation !== generation) return { generation, revision: 0, tasks: [] }; + return stored; } // ---- 附件存储 ---- @@ -111,8 +465,19 @@ export class AgentChatRepo extends OPFSRepo { // 保存附件数据到 workspace/uploads(支持 base64/data URL 字符串或 Blob) async saveAttachment(id: string, data: string | Blob): Promise { - const result = await writeWorkspaceFile(`uploads/${id}`, data); - return result.size; + const expectedSize = + data instanceof Blob ? data.size : isDataUrl(data) ? decodeDataUrl(data).data.byteLength : new Blob([data]).size; + try { + const result = await writeWorkspaceFile(`uploads/${id}`, data); + return result.size; + } catch (error) { + try { + if ((await this.getAttachment(id))?.size === expectedSize) return expectedSize; + } catch { + // Preserve the original write error when read-back confirmation itself fails. + } + throw error; + } } // 读取附件数据为 Blob(先查 workspace 新路径,fallback 旧路径) @@ -122,15 +487,16 @@ export class AgentChatRepo extends OPFSRepo { const workspace = await getWorkspaceRoot(); const dir = await getDirectory(workspace, "uploads"); return await (await dir.getFileHandle(id)).getFile(); - } catch { - // 新路径不存在,尝试旧路径 + } catch (error) { + if (!isNotFoundError(error)) throw error; } // 旧路径回退: agents/conversations/attachments/{id} try { const dir = await this.getChildDir(ATTACHMENTS_DIR); return await (await dir.getFileHandle(id)).getFile(); - } catch { - return null; + } catch (error) { + if (isNotFoundError(error)) return null; + throw error; } } @@ -141,15 +507,15 @@ export class AgentChatRepo extends OPFSRepo { const workspace = await getWorkspaceRoot(); const dir = await getDirectory(workspace, "uploads"); await dir.removeEntry(id); - } catch { - // 新路径不存在则忽略 + } catch (error) { + if (!isNotFoundError(error)) throw error; } // 旧路径: agents/conversations/attachments/{id} try { const dir = await this.getChildDir(ATTACHMENTS_DIR); await dir.removeEntry(id); - } catch { - // 旧路径不存在则忽略 + } catch (error) { + if (!isNotFoundError(error)) throw error; } } @@ -160,24 +526,88 @@ export class AgentChatRepo extends OPFSRepo { } } + // 附件存储在全局 workspace 中;删除前必须确认其它会话没有借用这些附件。 + private async deleteUnreferencedAttachments( + ids: string[], + excludedConversationId: string, + protectedReferences: Set = new Set() + ): Promise { + if (ids.length === 0) return; + const candidates = new Set(ids); + const referenced = new Set([...protectedReferences].filter((id) => candidates.has(id))); + try { + for (const conversation of await this.listConversations()) { + if (conversation.id === excludedConversationId) continue; + const snapshot = await this.getMessageSnapshot(conversation.id, conversation.generation); + for (const id of collectMessageAttachmentReferences( + snapshot.messages, + isLegacyGeneration(conversation.generation) + )) { + if (candidates.has(id)) referenced.add(id); + } + if (referenced.size === candidates.size) return; + } + } catch { + // 无法完整确认引用关系时安全地保留附件,避免误删其它会话仍在使用的数据。 + return; + } + await this.deleteAttachments(ids.filter((id) => !referenced.has(id))); + } + // ---- 任务 (task_tools) 存储 ---- // 获取会话关联的任务列表 - async getTasks(conversationId: string): Promise { - const tasksDir = await this.getChildDir(TASKS_DIR); - return this.readJsonFile(`${conversationId}.json`, [], tasksDir); + async getTasks(conversationId: string, generation?: string): Promise { + return (await this.getTaskSnapshot(conversationId, generation)).tasks; + } + + // 获取会话关联的任务快照(含 revision,供跨上下文写入做 CAS) + async getTaskSnapshot(conversationId: string, generation?: string): Promise { + return this.withFileLock(`lifecycle:${conversationId}`, async () => { + const current = await this.requireConversation(conversationId, generation); + const tasksDir = await this.getChildDir(TASKS_DIR); + // generation 不匹配(例如会话 ID 被复用、旧任务文件清理失败残留)时返回空快照, + // 而不是把上一代会话的任务列表当作当前会话的任务,与 readMessageSnapshot 的语义一致。 + const snapshot = await this.readTaskSnapshot(conversationId, current.generation!, tasksDir); + return snapshot; + }); } // 保存会话关联的任务列表 - async saveTasks(conversationId: string, tasks: Task[]): Promise { - const tasksDir = await this.getChildDir(TASKS_DIR); - await this.writeJsonFile(`${conversationId}.json`, tasks, tasksDir); + async saveTasks( + conversationId: string, + tasks: Task[], + signal?: AbortSignal, + generation?: string, + expectedRevision?: number + ): Promise { + await this.withFileLock(`lifecycle:${conversationId}`, async () => { + const current = await this.requireConversation(conversationId, generation); + await this.withFileLock(`tasks:${conversationId}`, async () => { + const tasksDir = await this.getChildDir(TASKS_DIR); + const previous = await this.readTaskSnapshot(conversationId, current.generation!, tasksDir); + if (expectedRevision !== undefined && previous.revision !== expectedRevision) { + throw new RevisionConflictError(`Tasks for conversation "${conversationId}" changed`); + } + const saved: TaskSnapshot = { generation: current.generation!, revision: previous.revision + 1, tasks }; + await this.writeJsonFileConfirmed( + `${conversationId}.json`, + saved, + { generation: current.generation!, revision: 0, tasks: [] }, + tasksDir, + signal + ); + }); + }); } // 删除会话关联的任务 async deleteTasks(conversationId: string): Promise { - const tasksDir = await this.getChildDir(TASKS_DIR); - await this.deleteFile(`${conversationId}.json`, tasksDir); + await this.withFileLock(`lifecycle:${conversationId}`, async () => { + await this.requireConversation(conversationId); + const tasksDir = await this.getChildDir(TASKS_DIR); + await this.deleteFile(`${conversationId}.json`, tasksDir); + }); } } diff --git a/src/app/repo/agent_model.ts b/src/app/repo/agent_model.ts index aa6525842..336d336e8 100644 --- a/src/app/repo/agent_model.ts +++ b/src/app/repo/agent_model.ts @@ -1,4 +1,5 @@ import type { AgentModelConfig } from "@App/app/service/agent/core/types"; +import { normalizeModelLimits } from "@App/app/service/agent/core/model_context"; import { Repo, loadCache } from "./repo"; const DEFAULT_MODEL_KEY = "agent_model:__default__"; @@ -21,9 +22,9 @@ export class AgentModelRepo extends Repo { return this.get(id); } - // 保存模型 + // 保存模型:存储边界统一归一化 maxTokens/contextWindow(见 model_context.ts) async saveModel(model: AgentModelConfig): Promise { - await this._save(model.id, model); + await this._save(model.id, normalizeModelLimits(model)); } // 删除模型 diff --git a/src/app/repo/agent_task.test.ts b/src/app/repo/agent_task.test.ts index 068cac8ae..f8b4e100a 100644 --- a/src/app/repo/agent_task.test.ts +++ b/src/app/repo/agent_task.test.ts @@ -1,4 +1,4 @@ -import { describe, expect, it, beforeEach } from "vitest"; +import { describe, expect, it, beforeEach, vi } from "vitest"; import { AgentTaskRepo, AgentTaskRunRepo } from "./agent_task"; import type { AgentTask, AgentTaskRun } from "@App/app/service/agent/core/types"; import { createMockOPFS } from "./test-helpers"; @@ -74,6 +74,91 @@ describe("AgentTaskRepo", () => { const runs = await runRepo.listRuns(taskId); expect(runs).toHaveLength(0); }); + + it("任务已删除但 runs 清理失败时重试仍应完成孤儿历史清理", async () => { + const taskId = "t-clean-retry"; + await repo.saveTask(makeTask({ id: taskId })); + const runRepo = new AgentTaskRunRepo(); + await runRepo.appendRun(makeRun({ id: "retry-run", taskId })); + const originalClearRuns = AgentTaskRunRepo.prototype.clearRuns; + const clearRuns = vi + .spyOn(AgentTaskRunRepo.prototype, "clearRuns") + .mockRejectedValueOnce(new Error("temporary OPFS failure")) + .mockImplementation(function (this: AgentTaskRunRepo, id: string) { + return originalClearRuns.call(this, id); + }); + + await expect(repo.removeTask(taskId)).rejects.toThrow("temporary OPFS failure"); + expect(await repo.getTask(taskId)).toBeUndefined(); + await repo.removeTask(taskId); + + expect(await runRepo.listRuns(taskId)).toHaveLength(0); + expect(clearRuns).toHaveBeenCalledTimes(2); + }); + + it("删除后旧 generation 的完成写入不应复活任务", async () => { + const task = await repo.saveTask(makeTask({ id: "stale-task" })); + await repo.removeTask(task.id, task.generation, task.revision); + + task.lastRunStatus = "success"; + await expect(repo.saveTask(task)).rejects.toThrow("deleted"); + await expect( + repo.updateRunState(task.id, task.generation!, { + lastruntime: Date.now(), + lastRunStatus: "success", + lastRunError: undefined, + }) + ).rejects.toThrow("deleted"); + expect(await repo.getTask(task.id)).toBeUndefined(); + }); + + it("旧 revision 不应覆盖用户刚保存的新配置", async () => { + const task = await repo.saveTask(makeTask({ id: "cas-task" })); + const stale = { ...task } as AgentTask; + task.name = "新配置"; + await repo.saveTask(task); + + stale.name = "旧配置"; + await expect(repo.saveTask(stale)).rejects.toThrow("changed"); + expect((await repo.getTask(task.id))?.name).toBe("新配置"); + }); + + it("支持 Web Locks 时任务生命周期写操作应使用跨上下文排它锁", async () => { + const request = vi.fn(async (_name: string, _opts: unknown, fn: () => Promise) => fn()); + Object.defineProperty(navigator, "locks", { + value: { request }, + configurable: true, + writable: true, + }); + try { + const task = await repo.createTask(makeTask({ id: "cross-context-lock" })); + task.name = "已更新"; + await repo.saveTask(task); + + expect(request).toHaveBeenCalledWith( + expect.stringContaining("cross-context-lock"), + expect.objectContaining({ mode: "exclusive" }), + expect.any(Function) + ); + } finally { + // @ts-expect-error 清理测试注入的 locks + delete navigator.locks; + } + }); + + it("从备份导入时应忽略外部 generation 并以本机版本覆盖同 ID 任务", async () => { + const local = await repo.saveTask(makeTask({ id: "import-task", name: "本机" })); + + const imported = await repo.importTask( + makeTask({ id: "import-task", name: "备份", generation: "foreign-generation", revision: 99 }) + ); + + expect(imported).toMatchObject({ + name: "备份", + generation: local.generation, + revision: local.revision! + 1, + }); + }); }); describe("AgentTaskRunRepo", () => { @@ -139,6 +224,27 @@ describe("AgentTaskRunRepo", () => { expect(runs[0].status).toBe("running"); }); + it("写入已提交但 close 报错时应以 durable read-back 确认 append/update/remove", async () => { + const taskId = "ambiguous-run-write"; + const originalWrite = (repo as any).writeJsonFile.bind(repo); + const failAfterCommit = () => + vi.spyOn(repo as any, "writeJsonFile").mockImplementationOnce(async (...args: unknown[]) => { + await originalWrite(...args); + throw new Error("close failed after commit"); + }); + + failAfterCommit(); + await expect(repo.appendRun(makeRun({ id: "ambiguous-run", taskId }))).resolves.toBe(true); + + failAfterCommit(); + await expect(repo.updateRun(taskId, "ambiguous-run", { status: "success" })).resolves.toBeUndefined(); + expect((await repo.listRuns(taskId))[0].status).toBe("success"); + + failAfterCommit(); + await expect(repo.removeRun(taskId, "ambiguous-run")).resolves.toBeUndefined(); + expect(await repo.listRuns(taskId)).toHaveLength(0); + }); + it("appendRun 超过 MAX_RUNS_PER_TASK 时裁剪最老记录", async () => { const taskId = "task-ring"; // 预填 500 条数据(最新在前),避免逐条 append 超时 diff --git a/src/app/repo/agent_task.ts b/src/app/repo/agent_task.ts index 6720b7cb6..882ba6e0c 100644 --- a/src/app/repo/agent_task.ts +++ b/src/app/repo/agent_task.ts @@ -1,30 +1,152 @@ import type { AgentTask, AgentTaskRun } from "@App/app/service/agent/core/types"; import { Repo } from "./repo"; import { OPFSRepo } from "./opfs_repo"; +import { stackAsyncTask } from "@App/pkg/utils/async_queue"; +import { uuidv4 } from "@App/pkg/utils/uuid"; +import { RevisionConflictError } from "./revision"; +import { nextTimeInfo } from "@App/pkg/utils/cron"; + +function normalizeTask(task: AgentTask): AgentTask { + return { + ...task, + generation: task.generation || `legacy:${task.id}`, + revision: task.revision ?? 0, + }; +} + +function withTaskLock(id: string, fn: () => Promise): Promise { + const key = `agent-task:${id}`; + const locks = (globalThis as { navigator?: { locks?: LockManager } }).navigator?.locks; + if (locks?.request) return locks.request(key, { mode: "exclusive" }, fn) as Promise; + return stackAsyncTask(key, fn); +} export class AgentTaskRepo extends Repo { constructor() { super("agent_task:"); - this.enableCache(); } async listTasks(): Promise { - return this.find(); + return (await this.find()).map(normalizeTask); } async getTask(id: string): Promise { - return this.get(id); + const task = await this.get(id); + return task ? normalizeTask(task) : undefined; } - async saveTask(task: AgentTask): Promise { - await this._save(task.id, task); + async createTask(task: AgentTask): Promise { + return withTaskLock(task.id, async () => { + if (await this.getTask(task.id)) throw new RevisionConflictError(`Task "${task.id}" already exists`); + const created = normalizeTask({ ...task, generation: uuidv4(), revision: 1 }); + await this._save(created.id, created); + Object.assign(task, created); + return created; + }); } - async removeTask(id: string): Promise { - await this.delete(id); - // 同时清理关联的 runs - const runRepo = new AgentTaskRunRepo(); - await runRepo.clearRuns(id); + async saveTask(task: AgentTask): Promise { + return withTaskLock(task.id, async () => { + const current = await this.getTask(task.id); + if (!current) { + // Backward-compatible import path for legacy unversioned records. Versioned stale writes never create. + if (task.generation !== undefined || task.revision !== undefined) { + throw new RevisionConflictError(`Task "${task.id}" was deleted`); + } + const created = normalizeTask({ ...task, generation: uuidv4(), revision: 1 }); + await this._save(created.id, created); + Object.assign(task, created); + return created; + } + if (task.generation !== current.generation || task.revision !== current.revision) { + throw new RevisionConflictError(`Task "${task.id}" changed or was deleted`); + } + const saved = normalizeTask({ ...task, revision: current.revision! + 1 }); + await this._save(saved.id, saved); + Object.assign(task, saved); + return saved; + }); + } + + /** Restore a task from backup while keeping generations local to this installation. */ + async importTask(task: AgentTask): Promise { + return withTaskLock(task.id, async () => { + const current = await this.getTask(task.id); + const imported = normalizeTask({ + ...task, + generation: current?.generation || uuidv4(), + revision: current ? current.revision! + 1 : 1, + }); + await this._save(imported.id, imported); + return imported; + }); + } + + async updateRunState( + id: string, + generation: string, + state: Pick, + advanceSchedule = false + ): Promise { + return withTaskLock(id, async () => { + const current = await this.getTask(id); + if (!current || current.generation !== generation) { + throw new RevisionConflictError(`Task "${id}" changed or was deleted`); + } + const saved = normalizeTask({ + ...current, + ...state, + nextruntime: advanceSchedule ? nextTimeInfo(current.crontab).next.toMillis() : current.nextruntime, + revision: current.revision! + 1, + updatetime: Date.now(), + }); + await this._save(id, saved); + return saved; + }); + } + + async claimDueTask(id: string, generation: string, now: number): Promise { + return withTaskLock(id, async () => { + const current = await this.getTask(id); + if ( + !current || + current.generation !== generation || + !current.enabled || + !current.nextruntime || + current.nextruntime > now + ) { + return null; + } + const saved = normalizeTask({ + ...current, + nextruntime: nextTimeInfo(current.crontab).next.toMillis(), + revision: current.revision! + 1, + updatetime: Date.now(), + }); + await this._save(id, saved); + return saved; + }); + } + + async removeTask(id: string, generation?: string, expectedRevision?: number): Promise { + await withTaskLock(id, async () => { + const current = await this.getTask(id); + const runRepo = new AgentTaskRunRepo(); + if (!current) { + // A previous attempt may have removed chrome.storage state before OPFS run cleanup failed. + await runRepo.clearRuns(id); + return; + } + if ( + (generation !== undefined && current.generation !== generation) || + (expectedRevision !== undefined && current.revision !== expectedRevision) + ) { + throw new RevisionConflictError(`Task "${id}" changed before deletion`); + } + await this.delete(id); + // 同时清理关联的 runs + await runRepo.clearRuns(id); + }); } } @@ -39,22 +161,50 @@ export class AgentTaskRunRepo extends OPFSRepo { return `${taskId}.json`; } - async appendRun(run: AgentTaskRun): Promise { - const runs = await this.readJsonFile(this.filename(run.taskId), []); - runs.unshift(run); - // 环形缓冲:超过上限时裁剪最老的记录 - if (runs.length > MAX_RUNS_PER_TASK) { - runs.length = MAX_RUNS_PER_TASK; + private async writeRunsConfirmed(taskId: string, runs: AgentTaskRun[]): Promise { + const filename = this.filename(taskId); + try { + await this.writeJsonFile(filename, runs); + } catch (error) { + try { + const durable = await this.readJsonFile(filename, []); + if (JSON.stringify(durable) === JSON.stringify(runs)) return; + } catch { + // Preserve the original write error when read-back confirmation itself fails. + } + throw error; } - await this.writeJsonFile(this.filename(run.taskId), runs); + } + + async appendRun(run: AgentTaskRun): Promise { + return this.withFileLock(`runs:${run.taskId}`, async () => { + const runs = await this.readJsonFile(this.filename(run.taskId), []); + if (runs.some((existing) => existing.id === run.id)) return false; + runs.unshift(run); + // 环形缓冲:超过上限时裁剪最老的记录 + if (runs.length > MAX_RUNS_PER_TASK) runs.length = MAX_RUNS_PER_TASK; + await this.writeRunsConfirmed(run.taskId, runs); + return true; + }); } async updateRun(taskId: string, id: string, data: Partial): Promise { - const runs = await this.readJsonFile(this.filename(taskId), []); - const idx = runs.findIndex((r) => r.id === id); - if (idx < 0) return; - Object.assign(runs[idx], data); - await this.writeJsonFile(this.filename(taskId), runs); + await this.withFileLock(`runs:${taskId}`, async () => { + const runs = await this.readJsonFile(this.filename(taskId), []); + const idx = runs.findIndex((run) => run.id === id); + if (idx < 0) return; + Object.assign(runs[idx], data); + await this.writeRunsConfirmed(taskId, runs); + }); + } + + async removeRun(taskId: string, id: string): Promise { + await this.withFileLock(`runs:${taskId}`, async () => { + const runs = await this.readJsonFile(this.filename(taskId), []); + const retained = runs.filter((run) => run.id !== id); + if (retained.length === runs.length) return; + await this.writeRunsConfirmed(taskId, retained); + }); } async listRuns(taskId: string, limit = 50): Promise { @@ -63,6 +213,6 @@ export class AgentTaskRunRepo extends OPFSRepo { } async clearRuns(taskId: string): Promise { - await this.deleteFile(this.filename(taskId)); + await this.withFileLock(`runs:${taskId}`, () => this.deleteFile(this.filename(taskId))); } } diff --git a/src/app/repo/opfs_repo.ts b/src/app/repo/opfs_repo.ts index 626059777..b3192bfa3 100644 --- a/src/app/repo/opfs_repo.ts +++ b/src/app/repo/opfs_repo.ts @@ -1,8 +1,26 @@ // OPFS(Origin Private File System)通用 Repo 基类 // 所有 Agent 相关的持久化数据统一存储在 agents/ 目录下 +import { stackAsyncTask } from "@App/pkg/utils/async_queue"; + const AGENTS_ROOT = "agents"; +function isNotFoundError(error: unknown): boolean { + return (error as { name?: string })?.name === "NotFoundError"; +} + +// 跨上下文互斥:Options 页与 Service Worker 都会直接读写同一份 OPFS JSON 文件, +// 进程内队列(stackAsyncTask)覆盖不了跨上下文的读-改-写竞争。Web Locks 按 origin +// 全局生效(扩展页与 MV3 SW 同源),是两者之间唯一共享的互斥原语;不支持 Web Locks +// 的环境(单元测试 jsdom)退化为进程内按 key 排队。 +function withExclusiveFileLock(key: string, fn: () => Promise): Promise { + const locks = (globalThis as { navigator?: { locks?: LockManager } }).navigator?.locks; + if (locks?.request) { + return locks.request(key, { mode: "exclusive" }, fn) as Promise; + } + return stackAsyncTask(key, fn); +} + // 获取 agents 根目录 async function getAgentsRoot(): Promise { const root = await navigator.storage.getDirectory(); @@ -43,25 +61,57 @@ export class OPFSRepo { return getSubDir(dir, childPath); } - // 读取 JSON 文件,文件不存在时返回默认值 + // 以当前 Repo + 逻辑范围(scope)为粒度的排它锁,供子类把"读-改-写"包成互斥临界区 + protected withFileLock(scope: string, fn: () => Promise): Promise { + return withExclusiveFileLock(`opfs-repo:${this.subPath}:${scope}`, fn); + } + + // 读取 JSON 文件。文件尚未创建(NotFoundError)是预期状态,返回默认值; + // 解析失败、权限或 I/O 错误一律抛出——把这类失败静默转成默认值,会让后续的 + // 读-改-写(appendMessage / saveConversation 等)基于空快照把仍然存在的旧数据 + // 整份覆写掉。 protected async readJsonFile(filename: string, defaultValue: T, dir?: FileSystemDirectoryHandle): Promise { + const targetDir = dir || (await this.getDir()); + let fileHandle: FileSystemFileHandle; try { - const targetDir = dir || (await this.getDir()); - const fileHandle = await targetDir.getFileHandle(filename); - const file = await fileHandle.getFile(); - const text = await file.text(); - return JSON.parse(text) as T; - } catch { - return defaultValue; + fileHandle = await targetDir.getFileHandle(filename); + } catch (error) { + if (isNotFoundError(error)) return defaultValue; + throw error; } + const file = await fileHandle.getFile(); + const text = await file.text(); + // 空文件:createWritable() 事务性写入不会留下半截内容,空文件只可能是 create 后从未写入, + // 与"文件不存在"同义,按默认值处理 + if (!text) return defaultValue; + return JSON.parse(text) as T; } - // 写入 JSON 文件 - protected async writeJsonFile(filename: string, data: unknown, dir?: FileSystemDirectoryHandle): Promise { + // 写入 JSON 文件。 + // OPFS 的 createWritable() 本身是事务性的:write() 写入的是临时副本,只有 close() 成功 + // 才会原子替换原文件;调用方持有的旧内容在此之前始终完整可读。 + // 传入 signal 时,若在"调用 close() 之前"已 abort,则改为 writable.abort() 放弃这次写入。 + // 诚实说明这里的边界:这只保证 close() 调用前的 abort 一定不会提交;一旦 close() 已经 + // 发出,FSA 规范不提供可靠的方式中途取消它,abort 恰好落在 close() 进行期间这个极窄窗口 + // 理论上仍可能提交。调用方(compact_service.ts / chat_service.ts)都会在 + // saveMessages() resolve 之后再次检查 signal,因此即使命中这个窗口,也不会对外报告 + // 虚假的成功事件(compact_done/done)——唯一的残留风险是磁盘内容被替换但会话已判定为 + // 取消,这是一个已知的、极窄的边界情况,未做完整的事务回滚。 + protected async writeJsonFile( + filename: string, + data: unknown, + dir?: FileSystemDirectoryHandle, + signal?: AbortSignal + ): Promise { + if (signal?.aborted) throw new Error("Aborted"); const targetDir = dir || (await this.getDir()); const fileHandle = await targetDir.getFileHandle(filename, { create: true }); const writable = await fileHandle.createWritable(); await writable.write(JSON.stringify(data)); + if (signal?.aborted) { + await writable.abort().catch(() => {}); + throw new Error("Aborted"); + } await writable.close(); } @@ -70,8 +120,8 @@ export class OPFSRepo { try { const targetDir = dir || (await this.getDir()); await targetDir.removeEntry(filename); - } catch { - // 文件不存在则忽略 + } catch (error) { + if (!isNotFoundError(error)) throw error; } } @@ -80,8 +130,8 @@ export class OPFSRepo { try { const targetDir = dir || (await this.getDir()); await targetDir.removeEntry(name, { recursive: true }); - } catch { - // 目录不存在则忽略 + } catch (error) { + if (!isNotFoundError(error)) throw error; } } diff --git a/src/app/repo/repo.test.ts b/src/app/repo/repo.test.ts index 2a6d14cbd..384848c39 100644 --- a/src/app/repo/repo.test.ts +++ b/src/app/repo/repo.test.ts @@ -1,5 +1,6 @@ -import { describe, it, expect, beforeEach } from "vitest"; +import { describe, it, expect, beforeEach, vi } from "vitest"; import { Repo } from "./repo"; +import { OPFSRepo } from "./opfs_repo"; // 定义测试数据类型 interface TestItem { @@ -25,6 +26,37 @@ class TestRepo extends Repo { } } +class TestOPFSRepo extends OPFSRepo { + constructor() { + super("test"); + } + + removeFile(dir: FileSystemDirectoryHandle) { + return this.deleteFile("item.json", dir); + } + + removeDir(dir: FileSystemDirectoryHandle) { + return this.removeDirectory("items", dir); + } +} + +describe("OPFSRepo 删除错误", () => { + it("仅忽略 NotFoundError,权限或 I/O 错误必须向上传播", async () => { + const repo = new TestOPFSRepo(); + const denied = { + removeEntry: vi.fn().mockRejectedValue(new DOMException("denied", "NotAllowedError")), + } as unknown as FileSystemDirectoryHandle; + const missing = { + removeEntry: vi.fn().mockRejectedValue(new DOMException("missing", "NotFoundError")), + } as unknown as FileSystemDirectoryHandle; + + await expect(repo.removeFile(denied)).rejects.toThrow("denied"); + await expect(repo.removeDir(denied)).rejects.toThrow("denied"); + await expect(repo.removeFile(missing)).resolves.toBeUndefined(); + await expect(repo.removeDir(missing)).resolves.toBeUndefined(); + }); +}); + describe("Repo", () => { let repo: TestRepo; @@ -70,6 +102,39 @@ describe("Repo", () => { const result = await repo.get("1"); expect(result).toEqual(item2); }); + + it("chrome.storage 写入失败时应 reject 而不是报告幽灵成功", async () => { + const setSpy = vi.spyOn(chrome.storage.local, "set").mockImplementation((_items, callback) => { + Object.defineProperty(chrome.runtime, "lastError", { + value: { message: "quota exceeded" }, + configurable: true, + }); + callback?.(); + delete (chrome.runtime as { lastError?: chrome.runtime.LastError }).lastError; + }); + + await expect(repo.save("failed", { id: "failed", name: "失败", value: 1 })).rejects.toThrow("quota exceeded"); + setSpy.mockRestore(); + await expect(repo.get("failed")).resolves.toBeUndefined(); + }); + + it("chrome.storage 读取失败时所有查询入口都应 reject", async () => { + const getSpy = vi.spyOn(chrome.storage.local, "get").mockImplementation((_keys?: unknown, callback?: any) => { + const cb = typeof _keys === "function" ? _keys : callback; + Object.defineProperty(chrome.runtime, "lastError", { + value: { message: "storage unavailable" }, + configurable: true, + }); + cb?.({}); + delete (chrome.runtime as { lastError?: chrome.runtime.LastError }).lastError; + }); + + await expect(repo.get("failed")).rejects.toThrow("storage unavailable"); + await expect(repo.gets(["failed"])).rejects.toThrow("storage unavailable"); + await expect(repo.getRecord(["failed"])).rejects.toThrow("storage unavailable"); + await expect(repo.find()).rejects.toThrow("storage unavailable"); + getSpy.mockRestore(); + }); }); describe("gets", () => { diff --git a/src/app/repo/repo.ts b/src/app/repo/repo.ts index ccf32f114..b4fea975c 100644 --- a/src/app/repo/repo.ts +++ b/src/app/repo/repo.ts @@ -9,12 +9,13 @@ export function loadCache(): Promise>> { return Promise.resolve(cache); } if (!loadCachePromise) { - loadCachePromise = new Promise>>((resolve) => { + loadCachePromise = new Promise>>((resolve, reject) => { chrome.storage.local.get((result: Partial> | undefined) => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.get:", lastError); - // 无视storage API错误,继续执行 + loadCachePromise = undefined; + reject(new Error(lastError.message || "chrome.storage.local.get failed")); + return; } cache = result || {}; loadCachePromise = undefined; @@ -29,43 +30,38 @@ function saveCacheAndStorage(key: string, value: T): Promise; function saveCacheAndStorage(items: Record): Promise; function saveCacheAndStorage(keyOrItems: string | Record, value?: T): Promise { if (typeof keyOrItems === "string") { - return Promise.all([ - loadCache().then((cache) => { - cache[keyOrItems] = value; - }), - new Promise((resolve) => { - chrome.storage.local.set({ [keyOrItems]: value }, () => { - const lastError = chrome.runtime.lastError; - if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.set:", lastError); - // 无视storage API错误,继续执行 - } - resolve(); - }); - }), - ]).then(() => value); + return new Promise((resolve, reject) => { + chrome.storage.local.set({ [keyOrItems]: value }, () => { + const lastError = chrome.runtime.lastError; + if (lastError) { + reject(new Error(lastError.message || "chrome.storage.local.set failed")); + return; + } + resolve(); + }); + }).then(async () => { + (await loadCache())[keyOrItems] = value; + return value as T; + }); } else { const items = keyOrItems; - return Promise.all([ - loadCache().then((cache) => { - Object.assign(cache, items); - }), - new Promise((resolve) => { - chrome.storage.local.set(items, () => { - const lastError = chrome.runtime.lastError; - if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.set:", lastError); - // 无视storage API错误,继续执行 - } - resolve(); - }); - }), - ]).then(() => undefined); + return new Promise((resolve, reject) => { + chrome.storage.local.set(items, () => { + const lastError = chrome.runtime.lastError; + if (lastError) { + reject(new Error(lastError.message || "chrome.storage.local.set failed")); + return; + } + resolve(); + }); + }).then(async () => { + Object.assign(await loadCache(), items); + }); } } function saveStorage(key: string, value: T): Promise { - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.set( { [key]: value, @@ -73,8 +69,8 @@ function saveStorage(key: string, value: T): Promise { () => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.set:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.set failed")); + return; } resolve(value); } @@ -83,12 +79,12 @@ function saveStorage(key: string, value: T): Promise { } function saveStorageRecord(record: Partial>): Promise { - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.set(record, () => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.set:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.set failed")); + return; } resolve(); }); @@ -105,12 +101,12 @@ function getCache(key: string): Promise { } function getStorage(key: string): Promise { - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.get(key, (result) => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.get:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.get failed")); + return; } resolve(result[key]); }); @@ -118,12 +114,12 @@ function getStorage(key: string): Promise { } function getStorageRecord(keys: string[]): Promise>> { - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.get(keys, (result) => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.get:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.get failed")); + return; } resolve(result); }); @@ -137,12 +133,12 @@ function deleteCache(key: string) { } function deleteStorage(key: string) { - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.remove(key, () => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.remove:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.remove failed")); + return; } resolve(); }); @@ -150,20 +146,15 @@ function deleteStorage(key: string) { } export function deletesStorage(keys: string[]) { - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.remove(keys, () => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.remove:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.remove failed")); + return; } resolve(); }); - }).catch(async () => { - // fallback - for (const key of keys) { - await deleteStorage(key); - } }); } @@ -213,12 +204,12 @@ export abstract class Repo { }); }); } - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.get(keys, (result) => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.get:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.get failed")); + return; } resolve(keys.map((key) => result[key] as T | undefined)); }); @@ -240,12 +231,12 @@ export abstract class Repo { return record; }); } - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.get(keys, (result) => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.get:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.get failed")); + return; } resolve(result as Partial>); }); @@ -275,12 +266,12 @@ export abstract class Repo { }); }); } - return new Promise((resolve) => { + return new Promise((resolve, reject) => { chrome.storage.local.get((result: { [key: string]: T }) => { const lastError = chrome.runtime.lastError; if (lastError) { - console.error("chrome.runtime.lastError in chrome.storage.local.get:", lastError); - // 无视storage API错误,继续执行 + reject(new Error(lastError.message || "chrome.storage.local.get failed")); + return; } resolve(this.filter(result, filters)); }); @@ -299,7 +290,7 @@ export abstract class Repo { public delete(key: string): Promise { key = this.joinKey(key); if (this.useCache) { - return Promise.all([deleteCache(key), deleteStorage(key)]).then(() => undefined); + return deleteStorage(key).then(() => deleteCache(key)); } return deleteStorage(key); } @@ -307,12 +298,13 @@ export abstract class Repo { public deletes(keys: string[]): Promise { keys = keys.map((key) => this.joinKey(key)); if (this.useCache) { - return loadCache().then((cache) => { - for (const key of keys) { - delete cache[key]; - } - return deletesStorage(keys); - }); + return deletesStorage(keys) + .then(() => loadCache()) + .then((cache) => { + for (const key of keys) { + delete cache[key]; + } + }); } return deletesStorage(keys); } diff --git a/src/app/repo/revision.ts b/src/app/repo/revision.ts new file mode 100644 index 000000000..89b1dc8ea --- /dev/null +++ b/src/app/repo/revision.ts @@ -0,0 +1,10 @@ +export class RevisionConflictError extends Error { + constructor(message: string) { + super(message); + this.name = "RevisionConflictError"; + } +} + +export function isRevisionConflict(error: unknown): error is RevisionConflictError { + return error instanceof RevisionConflictError || (error as { name?: string })?.name === "RevisionConflictError"; +} diff --git a/src/app/service/agent/core/abort_utils.ts b/src/app/service/agent/core/abort_utils.ts new file mode 100644 index 000000000..819b6670e --- /dev/null +++ b/src/app/service/agent/core/abort_utils.ts @@ -0,0 +1,32 @@ +/** Create a standard abort error for tool execution flows. */ +export function createAbortError(message = "Aborted"): Error { + return new Error(message); +} + +/** Throw immediately when the signal is already aborted. */ +export function throwIfAborted(signal?: AbortSignal): void { + if (signal?.aborted) { + throw createAbortError(); + } +} + +/** Reject as soon as the signal aborts, while leaving the original promise untouched. */ +export function raceWithAbort(promise: Promise, signal?: AbortSignal): Promise { + if (!signal) return promise; + if (signal.aborted) return Promise.reject(createAbortError()); + + return new Promise((resolve, reject) => { + const onAbort = () => reject(createAbortError()); + signal.addEventListener("abort", onAbort, { once: true }); + promise.then( + (value) => { + signal.removeEventListener("abort", onAbort); + resolve(value); + }, + (error) => { + signal.removeEventListener("abort", onAbort); + reject(error); + } + ); + }); +} diff --git a/src/app/service/agent/core/agent.test.ts b/src/app/service/agent/core/agent.test.ts index 369d2543f..542c8104f 100644 --- a/src/app/service/agent/core/agent.test.ts +++ b/src/app/service/agent/core/agent.test.ts @@ -268,7 +268,7 @@ describe("OpenAI Provider", () => { it("应正确解析带 usage 的 done 事件", async () => { const events = await collectEvents(parseOpenAIStream, [ - 'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n', + 'data: {"choices":[{"delta":{"content":"hi"},"finish_reason":"stop"}]}\n\n', 'data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n', ]); @@ -328,18 +328,17 @@ describe("OpenAI Provider", () => { } }); - it("signal 中断时不应产生错误事件", async () => { + it("signal 中断时应 reject 而不是静默完成(否则外层 callLLM 永远挂起)", async () => { const abortController = new AbortController(); abortController.abort(); - const events = await collectEvents( - parseOpenAIStream, - ['data: {"choices":[{"delta":{"content":"hello"}}]}\n\n'], - abortController.signal - ); - - // signal 已中断,不应处理任何数据 - expect(events).toHaveLength(0); + await expect( + collectEvents( + parseOpenAIStream, + ['data: {"choices":[{"delta":{"content":"hello"}}]}\n\n'], + abortController.signal + ) + ).rejects.toThrow("Aborted"); }); }); }); @@ -548,7 +547,7 @@ function buildSSEResponse(sseChunks: string[]): Response { // 辅助:构造纯文本 SSE 数据(OpenAI 格式) function makeTextSSE(text: string, usage?: { prompt_tokens: number; completion_tokens: number }): string[] { const chunks: string[] = []; - chunks.push(`data: {"choices":[{"delta":{"content":"${text}"}}]}\n\n`); + chunks.push(`data: {"choices":[{"delta":{"content":"${text}"},"finish_reason":"stop"}]}\n\n`); if (usage) { chunks.push(`data: {"usage":${JSON.stringify(usage)}}\n\n`); } else { @@ -569,7 +568,7 @@ function makeToolCallSSE( `data: {"choices":[{"delta":{"tool_calls":[{"id":"${toolId}","function":{"name":"${toolName}","arguments":""}}]}}]}\n\n` ); chunks.push( - `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":"${args.replace(/"/g, '\\"')}"}}]}}]}\n\n` + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":"${args.replace(/"/g, '\\"')}"}}]}, "finish_reason":"tool_calls"}]}\n\n` ); if (usage) { chunks.push(`data: {"usage":${JSON.stringify(usage)}}\n\n`); @@ -588,6 +587,8 @@ function createTestService() { listConversations: vi.fn().mockResolvedValue([]), saveConversation: vi.fn().mockResolvedValue(undefined), saveMessages: vi.fn().mockResolvedValue(undefined), + updateMessage: vi.fn().mockResolvedValue(undefined), + commitToolRound: vi.fn().mockResolvedValue(undefined), getAttachment: vi.fn().mockResolvedValue(null), saveAttachment: vi.fn().mockResolvedValue(0), }); @@ -638,7 +639,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "你好" }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -689,7 +689,6 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages, tools: [{ name: "get_weather", description: "获取天气", parameters: {} }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -699,7 +698,7 @@ describe("callLLMWithToolLoop", () => { expect(fetchSpy).toHaveBeenCalledTimes(2); // 工具应被执行 - expect(mockExecutor.execute).toHaveBeenCalledWith({ city: "北京" }); + expect(mockExecutor.execute).toHaveBeenCalledWith({ city: "北京" }, expect.any(AbortSignal), "call_1"); // messages 应包含 assistant(tool_call) + tool(result) + user 原始消息 expect(messages.length).toBe(3); // user + assistant(toolCalls) + tool(result) @@ -746,7 +745,6 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages, tools: [{ name: "search", description: "搜索", parameters: {} }], - maxIterations: 10, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -759,41 +757,6 @@ describe("callLLMWithToolLoop", () => { expect(events.find((e) => e.type === "done")).toBeDefined(); }); - it("超过 maxIterations 限制", async () => { - const { service, toolRegistry } = createTestService(); - const events: ChatStreamEvent[] = []; - - toolRegistry.registerBuiltin( - { name: "loop_tool", description: "循环工具", parameters: { type: "object", properties: {} } }, - { execute: vi.fn().mockResolvedValue("ok") } - ); - - // 每次都返回 tool_call - fetchSpy.mockImplementation(() => Promise.resolve(buildSSEResponse(makeToolCallSSE("call_x", "loop_tool", "{}")))); - - await (service as any).callLLMWithToolLoop({ - toolRegistry: (service as any).toolRegistry, - model: openaiConfig, - messages: [{ role: "user", content: "test" }], - tools: [{ name: "loop_tool", description: "循环工具", parameters: {} }], - maxIterations: 3, - sendEvent: (e: ChatStreamEvent) => events.push(e), - signal: new AbortController().signal, - scriptToolCallback: null, - }); - - // fetch 应被调用 3 次(maxIterations) - expect(fetchSpy).toHaveBeenCalledTimes(3); - - // 应收到 error 事件 - const errorEvent = events.find((e) => e.type === "error"); - expect(errorEvent).toBeDefined(); - if (errorEvent?.type === "error") { - expect(errorEvent.message).toContain("maximum iterations"); - expect(errorEvent.message).toContain("3"); - } - }); - it("signal 中止后应提前退出", async () => { const { service } = createTestService(); const events: ChatStreamEvent[] = []; @@ -810,7 +773,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "test" }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: abortController.signal, scriptToolCallback: null, @@ -837,7 +799,6 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages: [{ role: "user", content: "test" }], tools: [{ name: "script_tool", description: "脚本工具", parameters: {} }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: scriptCallback, @@ -845,7 +806,10 @@ describe("callLLMWithToolLoop", () => { // scriptCallback 应被调用 expect(scriptCallback).toHaveBeenCalledTimes(1); - expect(scriptCallback).toHaveBeenCalledWith([expect.objectContaining({ id: "call_1", name: "script_tool" })]); + expect(scriptCallback).toHaveBeenCalledWith( + [expect.objectContaining({ id: "call_1", name: "script_tool" })], + expect.any(AbortSignal) + ); expect(events.find((e) => e.type === "done")).toBeDefined(); }); @@ -860,7 +824,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "问题" }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -874,7 +837,8 @@ describe("callLLMWithToolLoop", () => { conversationId: "conv-123", role: "assistant", content: "回答", - }) + }), + undefined ); }); @@ -888,7 +852,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "问题" }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -917,32 +880,30 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages: [{ role: "user", content: "问题" }], tools: [{ name: "my_tool", description: "工具", parameters: {} }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, conversationId: "conv-456", }); - // 应持久化 3 条消息:assistant(tool_call) + tool(result) + assistant(final) - expect(mockRepo.appendMessage).toHaveBeenCalledTimes(3); + // 工具轮次必须作为一个原子组提交,最终回答再单独追加。 + expect(mockRepo.commitToolRound).toHaveBeenCalledTimes(1); + expect(mockRepo.appendMessage).toHaveBeenCalledTimes(1); - // 第一次:assistant with toolCalls - expect(mockRepo.appendMessage.mock.calls[0][0]).toMatchObject({ + const [toolAssistant, toolMessages] = mockRepo.commitToolRound.mock.calls[0]; + expect(toolAssistant).toMatchObject({ conversationId: "conv-456", role: "assistant", toolCalls: expect.arrayContaining([expect.objectContaining({ id: "call_1", name: "my_tool" })]), }); - - // 第二次:tool result - expect(mockRepo.appendMessage.mock.calls[1][0]).toMatchObject({ + expect(toolMessages).toHaveLength(1); + expect(toolMessages[0]).toMatchObject({ conversationId: "conv-456", role: "tool", toolCallId: "call_1", }); - // 第三次:最终 assistant 回答 - expect(mockRepo.appendMessage.mock.calls[2][0]).toMatchObject({ + expect(mockRepo.appendMessage.mock.calls[0][0]).toMatchObject({ conversationId: "conv-456", role: "assistant", content: "最终回答", @@ -961,7 +922,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "test" }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -983,7 +943,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "test" }], - maxIterations: 5, sendEvent: () => {}, signal: new AbortController().signal, scriptToolCallback: null, @@ -1003,7 +962,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "test" }], - maxIterations: 5, sendEvent: () => {}, signal: new AbortController().signal, scriptToolCallback: null, @@ -1024,7 +982,6 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages: [{ role: "user", content: "test" }], // 不传 tools,allToolDefs 为空 - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -1048,7 +1005,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "test" }], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -1107,7 +1063,6 @@ describe("callLLMWithToolLoop", () => { { name: "tool_a", description: "工具A", parameters: {} }, { name: "tool_b", description: "工具B", parameters: {} }, ], - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -1125,10 +1080,13 @@ describe("callLLMWithToolLoop", () => { expect(messages[2].role).toBe("tool"); expect(messages[3].role).toBe("tool"); - // 持久化:assistant(toolCalls) + 2个tool结果 + assistant(final) = 4 - expect(mockRepo.appendMessage).toHaveBeenCalledTimes(4); - expect(mockRepo.appendMessage.mock.calls[1][0]).toMatchObject({ role: "tool", toolCallId: "call_a" }); - expect(mockRepo.appendMessage.mock.calls[2][0]).toMatchObject({ role: "tool", toolCallId: "call_b" }); + // 工具 assistant 与两个结果原子提交,最终 assistant 单独追加。 + expect(mockRepo.commitToolRound).toHaveBeenCalledTimes(1); + expect(mockRepo.commitToolRound.mock.calls[0][1]).toEqual([ + expect.objectContaining({ role: "tool", toolCallId: "call_a" }), + expect.objectContaining({ role: "tool", toolCallId: "call_b" }), + ]); + expect(mockRepo.appendMessage).toHaveBeenCalledTimes(1); expect(events.find((e) => e.type === "done")).toBeDefined(); callLLMSpy.mockRestore(); @@ -1169,7 +1127,6 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages: [{ role: "user", content: "test" }], tools: scriptTools, - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -1207,7 +1164,6 @@ describe("callLLMWithToolLoop", () => { model: openaiConfig, messages: [{ role: "user", content: "test" }], // 不传 tools - maxIterations: 5, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, @@ -1251,7 +1207,6 @@ describe("callLLMWithToolLoop", () => { toolRegistry: (service as any).toolRegistry, model: openaiConfig, messages: [{ role: "user", content: "test" }], - maxIterations: 10, sendEvent: (e: ChatStreamEvent) => events.push(e), signal: new AbortController().signal, scriptToolCallback: null, diff --git a/src/app/service/agent/core/attachment_resolver.test.ts b/src/app/service/agent/core/attachment_resolver.test.ts new file mode 100644 index 000000000..dc4191c90 --- /dev/null +++ b/src/app/service/agent/core/attachment_resolver.test.ts @@ -0,0 +1,67 @@ +import { describe, expect, it, vi } from "vitest"; +import { prepareAttachmentSnapshot } from "./attachment_resolver"; +import type { AgentModelConfig, ChatRequest } from "./types"; + +const MODEL: AgentModelConfig = { + id: "vision", + name: "Vision", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", +}; + +const MESSAGES: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "image-1", mimeType: "image/png" }] }, +]; + +describe("附件请求快照", () => { + it("预检大小与 provider payload 应复用同一次不可变读取", async () => { + const getAttachment = vi + .fn() + .mockResolvedValueOnce(new Blob([new Uint8Array([1, 2, 3])], { type: "image/png" })) + .mockResolvedValueOnce(new Blob([new Uint8Array(1_000_000)], { type: "image/png" })); + + const snapshot = await prepareAttachmentSnapshot(MESSAGES, MODEL, getAttachment); + + expect(getAttachment).toHaveBeenCalledOnce(); + expect(snapshot.sizes.get("image-1")?.bytes).toBe(3); + expect(snapshot.resolver("image-1")).toContain("AQID"); + }); + + it("Stop 应中断正在进行的附件读取且不继续构建 payload", async () => { + const deferred = new Promise(() => {}); + const blob = { size: 3, type: "image/png", arrayBuffer: () => deferred } as Blob; + const controller = new AbortController(); + const pending = prepareAttachmentSnapshot(MESSAGES, MODEL, vi.fn().mockResolvedValue(blob), controller.signal); + + await Promise.resolve(); + controller.abort(); + + await expect(pending).rejects.toThrow("Aborted"); + }); + + it("Stop 应立即中断尚未完成的 OPFS 附件查询", async () => { + const controller = new AbortController(); + const pending = prepareAttachmentSnapshot( + MESSAGES, + MODEL, + vi.fn().mockReturnValue(new Promise(() => {})), + controller.signal + ); + + controller.abort(); + + await expect(pending).rejects.toThrow("Aborted"); + }); + + it("OPFS 权限或 I/O 错误不得伪装成附件不存在", async () => { + await expect( + prepareAttachmentSnapshot( + MESSAGES, + MODEL, + vi.fn().mockRejectedValue(new DOMException("denied", "NotAllowedError")) + ) + ).rejects.toThrow("denied"); + }); +}); diff --git a/src/app/service/agent/core/attachment_resolver.ts b/src/app/service/agent/core/attachment_resolver.ts index 2808d19aa..1c75177f6 100644 --- a/src/app/service/agent/core/attachment_resolver.ts +++ b/src/app/service/agent/core/attachment_resolver.ts @@ -1,61 +1,86 @@ import type { ChatRequest, AgentModelConfig } from "./types"; import { isContentBlocks } from "./content_utils"; import { supportsVision } from "./model_capabilities"; +import type { AttachmentSizeInfo } from "./context_elision"; +import { raceWithAbort, throwIfAborted } from "./abort_utils"; -/** - * 解析消息中 image+vision 的 attachmentId → base64 data URL - * file/audio/image(无vision) 不加载,provider 使用 OPFS 路径引用 - * @param messages 待解析的消息列表 - * @param model 当前模型配置(用于判断是否支持 vision) - * @param getAttachment 通过 attachmentId 异步获取 Blob 的函数(未找到返回 null/undefined) - * @returns resolver 函数:给定 attachmentId 返回 data URL 或 null - */ -export async function resolveAttachments( +export type AttachmentSnapshot = { + resolver: (id: string) => string | null; + sizes: Map; +}; + +export async function prepareAttachmentSnapshot( messages: ChatRequest["messages"], model: AgentModelConfig, - getAttachment: (id: string) => Promise -): Promise<(id: string) => string | null> { + getAttachment: (id: string) => Promise, + signal?: AbortSignal +): Promise { const resolved = new Map(); + const sizes = new Map(); const mimeTypes = new Map(); const ids = new Set(); - const hasVision = supportsVision(model); - - for (const m of messages) { - if (isContentBlocks(m.content)) { - for (const block of m.content) { - // 只收集 image + vision 的 attachmentId - if (block.type === "image" && hasVision && "attachmentId" in block) { - ids.add(block.attachmentId); - if (block.mimeType) { - mimeTypes.set(block.attachmentId, block.mimeType); - } - } + if (supportsVision(model)) { + for (const message of messages) { + if (!isContentBlocks(message.content)) continue; + for (const block of message.content) { + if (block.type !== "image") continue; + ids.add(block.attachmentId); + if (block.mimeType) mimeTypes.set(block.attachmentId, block.mimeType); } } } - if (ids.size === 0) return () => null; - for (const id of ids) { - try { - const blob = await getAttachment(id); - if (blob) { - // Blob → base64 data URL(分块拼接,避免 O(n²) 字符串拼接) - const buffer = await blob.arrayBuffer(); - const bytes = new Uint8Array(buffer); - const CHUNK_SIZE = 8192; - const chunks: string[] = []; - for (let i = 0; i < bytes.length; i += CHUNK_SIZE) { - chunks.push(String.fromCharCode(...bytes.subarray(i, Math.min(i + CHUNK_SIZE, bytes.length)))); - } - const b64 = btoa(chunks.join("")); - const mime = mimeTypes.get(id) || blob.type || "application/octet-stream"; - resolved.set(id, `data:${mime};base64,${b64}`); + throwIfAborted(signal); + const blob = await raceWithAbort(getAttachment(id), signal); + throwIfAborted(signal); + if (!blob) continue; + const info: AttachmentSizeInfo = { bytes: blob.size }; + if (typeof createImageBitmap === "function") { + try { + const bitmapPromise = createImageBitmap(blob); + // raceWithAbort cannot cancel browser decoding. If it finishes after Stop, close the abandoned bitmap. + void bitmapPromise.then( + (bitmap) => { + if (signal?.aborted) bitmap.close(); + }, + () => {} + ); + const bitmap = await raceWithAbort(bitmapPromise, signal); + info.width = bitmap.width; + info.height = bitmap.height; + bitmap.close(); + } catch (error) { + if (signal?.aborted) throw error; + // Unsupported or damaged image: size remains a conservative byte-based estimate. } - } catch { - // 加载失败,跳过 } + throwIfAborted(signal); + const bytes = new Uint8Array(await raceWithAbort(blob.arrayBuffer(), signal)); + throwIfAborted(signal); + const chunks: string[] = []; + for (let index = 0; index < bytes.length; index += 8192) { + chunks.push(String.fromCharCode(...bytes.subarray(index, Math.min(index + 8192, bytes.length)))); + } + const mime = mimeTypes.get(id) || blob.type || "application/octet-stream"; + resolved.set(id, `data:${mime};base64,${btoa(chunks.join(""))}`); + sizes.set(id, info); } + return { resolver: (id) => resolved.get(id) ?? null, sizes }; +} - return (id: string) => resolved.get(id) ?? null; +/** + * 解析消息中 image+vision 的 attachmentId → base64 data URL + * file/audio/image(无vision) 不加载,provider 使用 OPFS 路径引用 + * @param messages 待解析的消息列表 + * @param model 当前模型配置(用于判断是否支持 vision) + * @param getAttachment 通过 attachmentId 异步获取 Blob 的函数(未找到返回 null/undefined) + * @returns resolver 函数:给定 attachmentId 返回 data URL 或 null + */ +export async function resolveAttachments( + messages: ChatRequest["messages"], + model: AgentModelConfig, + getAttachment: (id: string) => Promise +): Promise<(id: string) => string | null> { + return (await prepareAttachmentSnapshot(messages, model, getAttachment)).resolver; } diff --git a/src/app/service/agent/core/context_elision.test.ts b/src/app/service/agent/core/context_elision.test.ts new file mode 100644 index 000000000..87d7a381b --- /dev/null +++ b/src/app/service/agent/core/context_elision.test.ts @@ -0,0 +1,353 @@ +import { describe, expect, it, vi } from "vitest"; +import { + elideOldToolResults, + elideOldAttachments, + ELIDED_TOOL_RESULT_STUB, + elideUntilWithinBudget, + estimateRequestTokens, + loadAttachmentSizes, +} from "./context_elision"; +import type { AgentModelConfig, ChatRequest, ToolCall } from "./types"; + +const VISION_MODEL: AgentModelConfig = { + id: "m-vision", + name: "Vision", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", +}; + +const NON_VISION_MODEL: AgentModelConfig = { + id: "m-text", + name: "Text", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4", +}; + +// 构造一轮 assistant(带 toolCalls) + N 条 tool 结果 +function round(index: number, toolCount = 1): ChatRequest["messages"] { + const assistantMsg: ChatRequest["messages"][number] = { + role: "assistant", + content: `assistant-${index}`, + toolCalls: Array.from({ length: toolCount }, (_, i) => ({ + id: `t${index}-${i}`, + name: "execute_script", + arguments: "{}", + })), + }; + const toolMsgs: ChatRequest["messages"] = Array.from({ length: toolCount }, (_, i) => ({ + role: "tool" as const, + content: `tool-result-${index}-${i}`, + toolCallId: `t${index}-${i}`, + })); + return [assistantMsg, ...toolMsgs]; +} + +describe("elideOldToolResults", () => { + it("轮次数不超过保留窗口时不应裁剪任何 tool 结果", () => { + const messages: ChatRequest["messages"] = [{ role: "user", content: "开始" }, ...round(1), ...round(2)]; + elideOldToolResults(messages, 5); + expect(messages.every((m) => m.role !== "tool" || m.content !== ELIDED_TOOL_RESULT_STUB)).toBe(true); + }); + + it("超过保留窗口的旧 tool 结果应被替换为占位文本,最近 K 轮保持不变", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "开始" }, + ...round(1), + ...round(2), + ...round(3), + ]; + elideOldToolResults(messages, 2); + + // 第 1 轮(最旧,超出保留窗口)应被裁剪 + const round1Tool = messages.find((m) => m.toolCallId === "t1-0"); + expect(round1Tool?.content).toBe(ELIDED_TOOL_RESULT_STUB); + + // 第 2、3 轮(最近 2 轮)应保持原文 + const round2Tool = messages.find((m) => m.toolCallId === "t2-0"); + const round3Tool = messages.find((m) => m.toolCallId === "t3-0"); + expect(round2Tool?.content).toBe("tool-result-2-0"); + expect(round3Tool?.content).toBe("tool-result-3-0"); + }); + + it("assistant 消息文本与 toolCalls 不应被裁剪,只处理 tool 角色消息", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "开始" }, + ...round(1), + ...round(2), + ...round(3), + ]; + elideOldToolResults(messages, 2); + + const round1Assistant = messages.find((m) => m.role === "assistant" && m.content === "assistant-1"); + expect(round1Assistant).toBeDefined(); + expect(round1Assistant?.toolCalls).toHaveLength(1); + }); + + it("已裁剪过的 tool 结果再次裁剪应保持幂等", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "开始" }, + ...round(1), + ...round(2), + ...round(3), + ]; + elideOldToolResults(messages, 2); + elideOldToolResults(messages, 2); + + const round1Tool = messages.find((m) => m.toolCallId === "t1-0"); + expect(round1Tool?.content).toBe(ELIDED_TOOL_RESULT_STUB); + }); + + it("单轮内多个 tool 结果应一并裁剪", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "开始" }, + ...round(1, 3), + ...round(2), + ...round(3), + ]; + elideOldToolResults(messages, 2); + + expect(messages.find((m) => m.toolCallId === "t1-0")?.content).toBe(ELIDED_TOOL_RESULT_STUB); + expect(messages.find((m) => m.toolCallId === "t1-1")?.content).toBe(ELIDED_TOOL_RESULT_STUB); + expect(messages.find((m) => m.toolCallId === "t1-2")?.content).toBe(ELIDED_TOOL_RESULT_STUB); + }); +}); + +describe("上下文预算估算与裁剪", () => { + it("按 UTF-8 字节和工具定义保守估算请求 token", () => { + const messages: ChatRequest["messages"] = [{ role: "user", content: "你好世界" }]; + const small = estimateRequestTokens(messages, []); + const large = estimateRequestTokens(messages, [{ name: "工具", description: "x".repeat(1000) }]); + + expect(large).toBeGreaterThan(small); + }); + + it("完整历史超过预算时裁剪工具结果,预算内的小历史保持原文", () => { + const smallHistory: ChatRequest["messages"] = [{ role: "user", content: "你好" }, ...round(1), ...round(2)]; + elideUntilWithinBudget(smallHistory, 1000, [], 0.6); + expect(smallHistory.find((message) => message.role === "tool")?.content).toBe("tool-result-1-0"); + + const largeHistory: ChatRequest["messages"] = []; + for (let i = 0; i < 5; i++) { + largeHistory.push(...round(i)); + const tool = largeHistory[largeHistory.length - 1]; + if (tool.role === "tool") tool.content = "结果".repeat(5000); + } + elideUntilWithinBudget(largeHistory, 1000, [], 0.6); + expect( + largeHistory + .filter((message) => message.role === "tool") + .every((message) => message.content === ELIDED_TOOL_RESULT_STUB) + ).toBe(true); + }); + + it("vision 模型下按图片附件实际字节估算,并可只省略较旧的多模态块", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "old", mimeType: "image/png" }] }, + { role: "assistant", content: "已看到图片" }, + { role: "user", content: "继续" }, + ]; + // 用较大的真实照片量级(600KB)而不是 6000 字节:折算后(/40)约 2 万 token, + // 明显超出 1000 的预算才能触发裁剪;6000 字节这种量级折算后不到 200 token,不足以撑爆预算。 + const sizes = new Map([["old", { bytes: 600_000 }]]); + // 图片按 IMAGE_CONSERVATIVE_BYTES_PER_TOKEN(40)折算为 token(不能再把 + // base64 字节数 1:1 当 token 数,否则普通照片会被判定为超出上下文)。 + // 验证 base64 展开确实被计入——若只是把原始字节数朴素除以换算系数(600000/40=15000), + // 结果不会低于它。 + expect(estimateRequestTokens(messages, [], sizes, VISION_MODEL)).toBeGreaterThan(15_000); + expect(elideUntilWithinBudget(messages, 1000, [], 0.9, sizes, VISION_MODEL)).toBe(true); + expect(messages[0].content).toEqual([{ type: "text", text: expect.stringContaining("attachment elided") }]); + expect((messages[0].content as any)[0].text).toContain("uploads/old"); + expect(messages[2].content).toBe("继续"); + }); + + it("已知尺寸的图片按 宽×高/750(长边 1568 缩放)估算,而不是按压缩后的字节数", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "dim", mimeType: "image/png" }] }, + ]; + // 高度可压缩的大图:压缩后只有 10KB,但 1920×1080 的真实视觉计费由尺寸决定 + // (1568×882 / 750 ≈ 1845 token);按字节换算只有 ~340 token,会被严重低估 + const sizes = new Map([["dim", { bytes: 10_000, width: 1920, height: 1080 }]]); + const tokens = estimateRequestTokens(messages, [], sizes, VISION_MODEL); + expect(tokens).toBeGreaterThan(1500); + expect(tokens).toBeLessThan(3000); + }); + + it("解码不出尺寸的图片退回按字节数的保守换算", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "nodim", mimeType: "image/png" }] }, + ]; + const sizes = new Map([["nodim", { bytes: 120_000 }]]); + // base64 展开约 16 万字节,按 40 字节/token 折算约 4000 token + expect(estimateRequestTokens(messages, [], sizes, VISION_MODEL)).toBeGreaterThan(3500); + }); + + it("loadAttachmentSizes 在支持 createImageBitmap 的环境下附带图片尺寸", async () => { + const close = vi.fn(); + (globalThis as any).createImageBitmap = vi.fn(async () => ({ width: 800, height: 600, close })); + try { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "img1", mimeType: "image/png" }] }, + ]; + const sizes = await loadAttachmentSizes(messages, async () => new Blob(["x".repeat(100)])); + expect(sizes.get("img1")).toMatchObject({ bytes: 100, width: 800, height: 600 }); + expect(close).toHaveBeenCalled(); + } finally { + delete (globalThis as any).createImageBitmap; + } + }); + + it("loadAttachmentSizes 在解码失败时仍返回字节数(退回字节换算)", async () => { + (globalThis as any).createImageBitmap = vi.fn(async () => { + throw new Error("decode failed"); + }); + try { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "img2", mimeType: "image/png" }] }, + ]; + const sizes = await loadAttachmentSizes(messages, async () => new Blob(["x".repeat(64)])); + expect(sizes.get("img2")).toEqual({ bytes: 64 }); + } finally { + delete (globalThis as any).createImageBitmap; + } + }); + + it("vision 模型下缺失图片大小时应使用文本降级估算而不是 Infinity", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "missing", mimeType: "image/png" }] }, + ]; + + const estimate = estimateRequestTokens(messages, [], undefined, VISION_MODEL); + + expect(Number.isFinite(estimate)).toBe(true); + expect(estimate).toBeGreaterThan(0); + }); + + it("非 vision 模型下不解析图片,不应因缺失大小把预算估算撑爆", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "old", mimeType: "image/png" }] }, + ]; + // 没有提供 size(模拟无法读取该附件),非 vision 模型下图片从不内联,不应导致 Infinity + expect(estimateRequestTokens(messages, [], undefined, NON_VISION_MODEL)).toBeLessThan(Number.POSITIVE_INFINITY); + }); + + it("file 块从不内联为二进制,缺失大小也不应导致估算为 Infinity", () => { + const messages: ChatRequest["messages"] = [ + { + role: "user", + content: [{ type: "file", attachmentId: "missing-file", mimeType: "application/pdf", name: "a.pdf" }], + }, + ]; + expect(estimateRequestTokens(messages, [], undefined, VISION_MODEL)).toBeLessThan(Number.POSITIVE_INFINITY); + }); + + it("audio 块从不内联为二进制,即使是 vision 模型也不应计入大文件字节", () => { + const bigSize = 50_000_000; + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "audio", attachmentId: "audio1", mimeType: "audio/mpeg" }] }, + ]; + const sizes = new Map([["audio1", { bytes: bigSize }]]); + const estimate = estimateRequestTokens(messages, [], sizes, VISION_MODEL); + expect(estimate).toBeLessThan(bigSize); + }); + + it("普通 100KB 照片不应被判定为超出未配置 contextWindow 模型(128K)的输入预算", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "帮我看看这张截图" }, + { role: "user", content: [{ type: "image", attachmentId: "screenshot", mimeType: "image/png" }] }, + ]; + const sizes = new Map([["screenshot", { bytes: 100_000 }]]); + // VISION_MODEL 未显式配置 contextWindow,按 gpt-4o 前缀推断为 128_000 + const estimate = estimateRequestTokens(messages, [], sizes, VISION_MODEL); + // 128_000 * 0.9 的预检阈值 ≈ 115_200;一张普通照片折算后的 token 数应远低于这个预算, + // 而不是像按 1 字节 = 1 token 估算那样膨胀到十几万 token 直接把预算撑爆 + expect(estimate).toBeLessThan(115_200 * 0.5); + }); +}); + +describe("estimateRequestTokens 的按消息缓存不应产生陈旧结果", () => { + it("elideOldToolResults 原地改写 content 后,同一批 message 对象的估算值应立即反映裁剪后的内容", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "你好" }, + { + role: "assistant", + content: "", + toolCalls: [{ id: "c1", name: "tool", arguments: "{}", status: "completed" }], + }, + { role: "tool", content: "x".repeat(5000), toolCallId: "c1" }, + ]; + + const before = estimateRequestTokens(messages); + // keepLastAssistantTurns=0 会把所有 tool 结果裁剪为占位文本 + elideOldToolResults(messages, 0); + const after = estimateRequestTokens(messages); + + expect(after).toBeLessThan(before); + expect(messages[2].content).toBe(ELIDED_TOOL_RESULT_STUB); + }); + + it("elideOldAttachments 原地改写 content 后,估算值应立即反映占位后的内容", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: [{ type: "image", attachmentId: "img1", mimeType: "image/png" }] }, + { role: "user", content: "占位" }, + { role: "user", content: "占位" }, + ]; + const sizes = new Map([["img1", { bytes: 6_000_000 }]]); + + const before = estimateRequestTokens(messages, [], sizes, VISION_MODEL); + elideOldAttachments(messages, 2); + const after = estimateRequestTokens(messages, [], sizes, VISION_MODEL); + + expect(after).toBeLessThan(before); + expect(messages[0].content).toEqual([{ type: "text", text: expect.stringContaining("attachment elided") }]); + }); + + it("assistant 的 UI-only 工具状态和结果不应进入 provider wire 估算", () => { + const toolCall: ToolCall = { id: "c1", name: "tool", arguments: "{}", status: "running" }; + const messages: ChatRequest["messages"] = [{ role: "assistant", content: "", toolCalls: [toolCall] }]; + + const runningEstimate = estimateRequestTokens(messages); + // 工具执行完成后原地回写 status(tool_loop_orchestrator.ts 的 applyToolUpdates 同款操作) + toolCall.status = "completed"; + toolCall.result = "a longer completed result string that changes the byte size"; + const completedEstimate = estimateRequestTokens(messages); + + expect(completedEstimate).toBe(runningEstimate); + }); + + it("大型 subAgentDetails 不应夸大 provider 实际只发送工具身份字段的请求", () => { + const toolCall: ToolCall = { + id: "c1", + name: "agent", + arguments: "{}", + subAgentDetails: { + agentId: "child", + description: "child", + messages: [{ content: "x".repeat(1_000_000), toolCalls: [] }], + }, + }; + const messages: ChatRequest["messages"] = [{ role: "assistant", content: "", toolCalls: [toolCall] }]; + const withDetails = estimateRequestTokens(messages); + delete toolCall.subAgentDetails; + const withoutDetails = estimateRequestTokens(messages); + + expect(withDetails).toBe(withoutDetails); + }); + + it("反复调用应返回稳定一致的结果(缓存命中不改变估算值)", () => { + const messages: ChatRequest["messages"] = [ + { role: "user", content: "稳定内容".repeat(100) }, + { role: "assistant", content: "回复内容".repeat(100) }, + ]; + + const first = estimateRequestTokens(messages); + const second = estimateRequestTokens(messages); + const third = estimateRequestTokens(messages); + + expect(second).toBe(first); + expect(third).toBe(first); + }); +}); diff --git a/src/app/service/agent/core/context_elision.ts b/src/app/service/agent/core/context_elision.ts new file mode 100644 index 000000000..c5a0192a2 --- /dev/null +++ b/src/app/service/agent/core/context_elision.ts @@ -0,0 +1,248 @@ +// 滑动窗口裁剪:仅裁剪内存中传给 LLM 的 messages(不影响 chatRepo 持久化/UI 历史), +// 用于在触及 autoCompact 的 80% 阈值之前,减少长 tool loop 中旧 tool 结果的重复计费。 +import type { AgentModelConfig, ChatRequest, ContentBlock } from "./types"; +import { supportsVision } from "./model_capabilities"; +import { imageBlockFallbackText } from "./providers/content_utils"; + +export const ELIDED_TOOL_RESULT_STUB = + "[tool result elided to save context — re-run the tool if you need this data again]"; + +/** 附件裁剪占位文本:保留 type/attachmentId(OPFS 路径),供模型按需重新读取,而非丢弃全部标识信息。 */ +function elidedAttachmentStub(block: Exclude): string { + return `[attachment elided to save context, type: ${block.type}, OPFS path: uploads/${block.attachmentId} — re-open the attachment if needed]`; +} + +/** 附件的预算估算信息:字节规模 + (图片能解码时的)像素尺寸。 */ +export type AttachmentSizeInfo = { bytes: number; width?: number; height?: number }; + +/** 读取 provider 会把附件展开成 data URL 后的实际字节规模; + * 图片附件在环境支持时(Service Worker 有 createImageBitmap)额外解码出像素尺寸, + * 供按 provider 尺寸计费规则估算视觉 token。 */ +export async function loadAttachmentSizes( + messages: ChatRequest["messages"], + getAttachment: (id: string) => Promise +): Promise> { + // 同一 attachmentId 可能被多个块引用;只要任一处以 image 块出现就按图片解码尺寸 + const kinds = new Map(); + for (const message of messages) { + if (!Array.isArray(message.content)) continue; + for (const block of message.content) { + if (block.type === "text") continue; + if (block.type === "image" || !kinds.has(block.attachmentId)) { + kinds.set(block.attachmentId, block.type === "image" ? "image" : "other"); + } + } + } + const sizes = new Map(); + await Promise.all( + [...kinds].map(async ([id, kind]) => { + try { + const blob = await getAttachment(id); + if (!blob) return; + const info: AttachmentSizeInfo = { bytes: blob.size }; + if (kind === "image" && typeof createImageBitmap === "function") { + try { + const bitmap = await createImageBitmap(blob); + info.width = bitmap.width; + info.height = bitmap.height; + bitmap.close(); + } catch { + // 解码失败(损坏/不支持的格式):退回按字节数的保守换算 + } + } + sizes.set(id, info); + } catch { + // 无法读取的附件保留为未知大小,由预算检查决定是否省略或拒绝请求。 + } + }) + ); + return sizes; +} + +// 字节→token 的保守换算:常见 tokenizer 在 UTF-8 文本上很少低于 2 字节/token +// (英文/Latin 文本通常 ~4 字节/token,CJK 每字符 3 字节通常对应 1-2 token)。 +// 直接把字节数当 token 数会比真实 token 数偏大数倍,导致远低于模型真实上限就被裁剪/拒绝; +// 除以该常量后仍保持保守(不会低估),同时不再把每个 UTF-8 字节当作独立 token。 +const CONSERVATIVE_BYTES_PER_TOKEN = 2; + +// 每条 message 的 role/content/toolCallId 序列化字节数缓存。 +// tool loop 每轮都要重新估算整段 messages 的预算占用;不缓存的话,每轮都要把(随对话增长的) +// 完整历史重新 JSON.stringify 一遍,R 轮下来是 O(R·N) —— 长对话里最耗时的部分。 +// 而 content/toolCallId 一旦写入某条 message 就几乎不再变化,唯二的例外是 +// elideOldToolResults / elideOldAttachments 原地改写 m.content,这两处都会显式调用 +// invalidateMessageByteCache() 使缓存失效;因此按 message 对象身份缓存是安全的。 +// toolCalls 字段(status/attachments/subAgentDetails 会在工具执行后原地变化)不缓存, +// 每次都单独重算——但它只随"该条消息自身的工具调用数"增长,不随对话长度增长,代价很小。 +const stableContentByteCache = new WeakMap(); + +/** 消息 content 被原地改写后必须调用,否则会读到改写前的缓存字节数。 */ +export function invalidateMessageByteCache(message: object): void { + stableContentByteCache.delete(message); +} + +function getStableContentBytes(message: { role: string; content: unknown; toolCallId?: string }): number { + let bytes = stableContentByteCache.get(message); + if (bytes === undefined) { + const stable: Record = { role: message.role, content: message.content }; + if (message.toolCallId) stable.toolCallId = message.toolCallId; + bytes = new TextEncoder().encode(JSON.stringify(stable)).byteLength; + stableContentByteCache.set(message, bytes); + } + return bytes; +} + +// vision 图片的字节→token 保守换算(仅在解不出像素尺寸时使用)。真实 provider 计费和 +// base64 字节数没有线性对应关系(OpenAI 按分块计费,一张图约 85~1105 token;Anthropic 按 +// 宽×高/750),典型压缩照片的 base64 体积换算下来大约是每 100~300 字节 1 token。 +// 每 40 字节算 1 token 比真实比例保守 2.5~7.5 倍,不会把图片开销算得比实际便宜, +// 也不会把一张普通照片(几十到一百多 KB)放大到几万 token 从而被拒绝在预算之外。 +const IMAGE_CONSERVATIVE_BYTES_PER_TOKEN = 40; + +// 能解出像素尺寸时按 provider 的尺寸计费规则估算:Anthropic 为 宽×高/750(长边超过 1568 +// 先等比缩小);OpenAI 按 512px tile 计费,同尺寸下低于该公式。取更保守的 Anthropic 公式, +// 并以 85(OpenAI 低清底价)为下限。相比压缩字节换算,尺寸公式不受压缩率影响: +// 高度可压缩的大图不会被低估,噪点大的小图不会被高估。 +const IMAGE_MAX_DIMENSION = 1568; +const IMAGE_PIXELS_PER_TOKEN = 750; +const IMAGE_MIN_TOKENS = 85; + +function estimateImageTokensFromDimensions(width: number, height: number): number { + const scale = Math.min(1, IMAGE_MAX_DIMENSION / Math.max(width, height, 1)); + const scaledWidth = Math.max(1, Math.floor(width * scale)); + const scaledHeight = Math.max(1, Math.floor(height * scale)); + return Math.max(IMAGE_MIN_TOKENS, Math.ceil((scaledWidth * scaledHeight) / IMAGE_PIXELS_PER_TOKEN)); +} + +/** + * 请求 token 的启发式估算,不是任何 provider 的精确 tokenizer(本仓库未接入 + * tiktoken/Anthropic 官方计数器,引入新依赖超出当前改动范围)。分两部分独立估算, + * 因为文本/JSON 与 vision 图片的字节→token 比例机制完全不同,不应共用同一个换算系数: + * + * - 文本 + JSON 结构(messages/tools/toolCalls):按 UTF-8 字节数 / CONSERVATIVE_BYTES_PER_TOKEN + * 保守折算。这仍然是启发式而非保证:高熵内容或极端 tokenizer 差异下仍可能被低估, + * 调用方(preflight 预算检查、elideUntilWithinBudget)应把结果当作"大致上界"而非精确值。 + * - vision 图片:按 IMAGE_CONSERVATIVE_BYTES_PER_TOKEN 折算(见上方常量注释),而不是直接把 + * base64 字节数当 token 数——后者会让一张普通照片的估算膨胀到几万 token,超出未配置 + * contextWindow 的模型(如 128K)的输入预算,把正常截图/照片当作"超出上下文"拒绝掉。 + * + * 只对 provider 实际会内联展开为 base64 的块计入二进制字节: + * 当前仅 vision 模型的 image 块会被 resolveAttachments 加载;file/audio 及非 vision 模型的 + * image 均降级为纯文本描述(见 providers/content_utils.ts),其体积已包含在下方的 JSON 基线字节里。 + */ +export function estimateRequestTokens( + messages: ChatRequest["messages"], + tools?: unknown[], + attachmentSizes?: Map, + model?: AgentModelConfig +): number { + const hasVision = model ? supportsVision(model) : false; + const attachmentTokens = messages.reduce((sum, message) => { + if (!Array.isArray(message.content)) return sum; + return ( + sum + + message.content.reduce((blockSum, block) => { + if (block.type === "text") return blockSum; + // file/audio 从不被内联;image 只在 vision 模型上才会被解析为 data URL + if (block.type === "file" || block.type === "audio" || !hasVision) return blockSum; + const info = attachmentSizes?.get(block.attachmentId); + if (info == null) { + // 附件大小未知时降级为纯文本描述——这部分本身就是文本,按文本换算系数折算, + // 不能套用图片换算系数(那是给真实二进制图片数据用的) + return ( + blockSum + + Math.ceil(new TextEncoder().encode(imageBlockFallbackText(block)).byteLength / CONSERVATIVE_BYTES_PER_TOKEN) + ); + } + // 优先按像素尺寸估算(provider 的真实计费维度);解不出尺寸时退回字节换算 + if (info.width != null && info.height != null) { + return blockSum + estimateImageTokensFromDimensions(info.width, info.height); + } + const base64Bytes = Math.ceil(info.bytes / 3) * 4 + 128; + return blockSum + Math.ceil(base64Bytes / IMAGE_CONSERVATIVE_BYTES_PER_TOKEN); + }, 0) + ); + }, 0); + if (!Number.isFinite(attachmentTokens)) return Number.POSITIVE_INFINITY; + + let bytes = 0; + for (const message of messages) { + bytes += getStableContentBytes(message); + if (message.toolCalls && message.toolCalls.length > 0) { + // Providers only send the wire identity/function fields. UI-only result/status/attachments/ + // subAgentDetails must not inflate admission estimates. + const wireToolCalls = message.toolCalls.map((toolCall) => ({ + id: toolCall.id, + name: toolCall.name, + arguments: toolCall.arguments, + })); + bytes += new TextEncoder().encode(JSON.stringify(wireToolCalls)).byteLength; + } + } + if (tools && tools.length > 0) { + bytes += new TextEncoder().encode(JSON.stringify(tools)).byteLength; + } + // 文本/JSON 部分按文本换算系数折算;图片部分已经在上面按图片换算系数折算成 token,两者相加 + return Math.ceil(bytes / CONSERVATIVE_BYTES_PER_TOKEN) + attachmentTokens; +} + +/** + * 保留最近 keepLastAssistantTurns 轮 assistant(带 toolCalls) 及其后的消息原文, + * 将更早的 tool 角色消息内容替换为占位文本。 + */ +export function elideOldToolResults(messages: ChatRequest["messages"], keepLastAssistantTurns: number): void { + let assistantTurnsSeen = 0; + let cutoffIndex = -1; + for (let i = messages.length - 1; i >= 0; i--) { + const m = messages[i]; + if (m.role === "assistant" && m.toolCalls && m.toolCalls.length > 0) { + assistantTurnsSeen++; + if (assistantTurnsSeen === keepLastAssistantTurns) { + cutoffIndex = i; + break; + } + } + } + + const limit = keepLastAssistantTurns === 0 ? messages.length : cutoffIndex; + for (let i = 0; i < limit; i++) { + const m = messages[i]; + if (m.role === "tool" && m.content !== ELIDED_TOOL_RESULT_STUB) { + m.content = ELIDED_TOOL_RESULT_STUB; + invalidateMessageByteCache(m); + } + } +} + +/** 将较旧消息中的多模态块替换为可恢复的文本占位,保留最近消息的附件。 */ +export function elideOldAttachments(messages: ChatRequest["messages"], keepLastMessages = 2): void { + const cutoff = Math.max(0, messages.length - keepLastMessages); + for (let i = 0; i < cutoff; i++) { + const message = messages[i]; + if (!Array.isArray(message.content)) continue; + message.content = message.content.map((block) => + block.type === "text" ? block : { type: "text", text: elidedAttachmentStub(block) } + ); + invalidateMessageByteCache(message); + } +} + +/** 在安全预算内尽量保留最近的 tool 结果;必要时裁剪全部旧结果。 */ +export function elideUntilWithinBudget( + messages: ChatRequest["messages"], + contextWindow: number, + tools?: unknown[], + budgetRatio = 0.9, + attachmentSizes?: Map, + model?: AgentModelConfig +): boolean { + // 原先按 hasToolResults 分两条 if 判断,但两分支条件完全相同(结果都只取决于 estimate()), + // 白白多做一次 O(n) 的 messages.some() 扫描;合并为一次判断,语义不变。 + const estimate = () => estimateRequestTokens(messages, tools, attachmentSizes, model); + if (estimate() / contextWindow < budgetRatio) return true; + for (let keep = 5; keep >= 0; keep--) { + elideOldToolResults(messages, keep); + if (estimate() / contextWindow < budgetRatio) return true; + } + elideOldAttachments(messages); + return estimate() / contextWindow < budgetRatio; +} diff --git a/src/app/service/agent/core/mcp_client.test.ts b/src/app/service/agent/core/mcp_client.test.ts index e6a4e5965..3def824e0 100644 --- a/src/app/service/agent/core/mcp_client.test.ts +++ b/src/app/service/agent/core/mcp_client.test.ts @@ -160,6 +160,52 @@ describe("MCPClient", () => { await expect(client.callTool("search")).rejects.toThrow("Tool failed"); }); + it("工具返回 structuredContent 时不应丢失结构化结果", async () => { + installServer({ + "tools/call": (request) => + jsonResponse(request.id, { + content: [{ type: "text", text: "Readable result" }], + structuredContent: { value: 42 }, + }), + }); + const client = new MCPClient(createConfig()); + await client.initialize(); + + await expect(client.callTool("search")).resolves.toEqual({ + content: [{ type: "text", text: "Readable result" }], + structuredContent: { value: 42 }, + }); + }); + + it("工具返回非文本错误内容时应保留诊断信息", async () => { + installServer({ + "tools/call": (request) => + jsonResponse(request.id, { + content: [{ type: "image", data: "base64", mimeType: "image/png" }], + isError: true, + }), + }); + const client = new MCPClient(createConfig()); + await client.initialize(); + + await expect(client.callTool("search")).rejects.toThrow('"mimeType":"image/png"'); + }); + + it("仅由 structuredContent 携带错误详情时也应保留诊断信息", async () => { + installServer({ + "tools/call": (request) => + jsonResponse(request.id, { + content: [], + structuredContent: { error: "quota exceeded" }, + isError: true, + }), + }); + const client = new MCPClient(createConfig()); + await client.initialize(); + + await expect(client.callTool("search")).rejects.toThrow('"quota exceeded"'); + }); + it("可列出和读取资源", async () => { const client = new MCPClient(createConfig()); await client.initialize(); diff --git a/src/app/service/agent/core/mcp_client.ts b/src/app/service/agent/core/mcp_client.ts index c67e398bc..7392113c7 100644 --- a/src/app/service/agent/core/mcp_client.ts +++ b/src/app/service/agent/core/mcp_client.ts @@ -2,6 +2,13 @@ import { Client } from "@modelcontextprotocol/sdk/client/index.js"; import { StreamableHTTPClientTransport } from "@modelcontextprotocol/sdk/client/streamableHttp.js"; import type { MCPServerConfig, MCPTool, MCPResource, MCPPrompt, MCPPromptMessage } from "./types"; +type MCPContent = { type: string; text?: string; [key: string]: unknown }; + +export type MCPToolCallResult = { + content: MCPContent[]; + structuredContent?: unknown; +}; + export class MCPClient { private readonly client: Client; private readonly transport: StreamableHTTPClientTransport; @@ -35,23 +42,43 @@ export class MCPClient { })); } - async callTool(name: string, args?: Record): Promise { + async callTool(name: string, args?: Record, signal?: AbortSignal): Promise { this.ensureInitialized(); - const result = (await this.client.callTool({ name, arguments: args ?? {} })) as { - content: Array<{ type: string; text?: string; [key: string]: unknown }>; - isError?: boolean; - }; + const result = (await this.client.callTool( + { name, arguments: args ?? {} }, + undefined, + signal ? { signal } : undefined + )) as MCPToolCallResult & { isError?: boolean }; const content = result.content; if (result.isError) { - const errorText = content.map((item) => (item.type === "text" ? item.text : "")).join("\n"); - throw new Error(errorText || "Tool call failed"); + const diagnostics = content + .map((item) => { + if (item.type === "text") return item.text || ""; + try { + return JSON.stringify(item) || item.type; + } catch { + return item.type; + } + }) + .filter(Boolean); + if (result.structuredContent !== undefined) { + try { + diagnostics.push(JSON.stringify(result.structuredContent) || String(result.structuredContent)); + } catch { + diagnostics.push(String(result.structuredContent)); + } + } + throw new Error(diagnostics.join("\n") || "Tool call failed"); } - if (content.length === 1 && content[0].type === "text") { + if (content.length === 1 && content[0].type === "text" && result.structuredContent === undefined) { return content[0].text; } - return content; + return { + content, + ...(result.structuredContent !== undefined ? { structuredContent: result.structuredContent } : {}), + }; } async listResources(): Promise { diff --git a/src/app/service/agent/core/mcp_tool_executor.test.ts b/src/app/service/agent/core/mcp_tool_executor.test.ts index e8e5c1e74..af5ce843f 100644 --- a/src/app/service/agent/core/mcp_tool_executor.test.ts +++ b/src/app/service/agent/core/mcp_tool_executor.test.ts @@ -16,7 +16,17 @@ describe("MCPToolExecutor", () => { const result = await executor.execute({ query: "hello" }); expect(result).toBe("tool result"); - expect(client.callTool).toHaveBeenCalledWith("search", { query: "hello" }); + expect(client.callTool).toHaveBeenCalledWith("search", { query: "hello" }, undefined); + }); + + it("应正确传递工具名", async () => { + const client = createMockClient({ data: [1, 2, 3] }); + const executor = new MCPToolExecutor(client, "fetch_data"); + + const result = await executor.execute({ limit: 10 }); + + expect(result).toEqual({ data: [1, 2, 3] }); + expect(client.callTool).toHaveBeenCalledWith("fetch_data", { limit: 10 }, undefined); }); it("callTool 抛出异常时应向上传播", async () => { @@ -92,4 +102,22 @@ describe("MCPToolExecutor", () => { expect(result.attachments[0].mimeType).toBe("image/png"); expect(result.attachments[0].data).toBe("data:image/png;base64,abc123"); }); + + it("包含 structuredContent 的 image 结果应保留结构化诊断", async () => { + const client = createMockClient({ + content: [{ type: "image", data: "abc123", mimeType: "image/png" }], + structuredContent: { caption: "chart" }, + }); + const executor = new MCPToolExecutor(client, "structured_image"); + + const result = (await executor.execute({})) as { + content: string; + attachments: unknown[]; + structuredContent?: unknown; + }; + + expect(result.content).toBe('{"caption":"chart"}'); + expect(result.attachments).toHaveLength(1); + expect(result.structuredContent).toEqual({ caption: "chart" }); + }); }); diff --git a/src/app/service/agent/core/mcp_tool_executor.ts b/src/app/service/agent/core/mcp_tool_executor.ts index 88dbafc6d..2974a9782 100644 --- a/src/app/service/agent/core/mcp_tool_executor.ts +++ b/src/app/service/agent/core/mcp_tool_executor.ts @@ -1,4 +1,4 @@ -import type { MCPClient } from "./mcp_client"; +import type { MCPClient, MCPToolCallResult } from "./mcp_client"; import type { ToolExecutor } from "./tool_registry"; import type { ToolResultWithAttachments } from "./types"; @@ -9,15 +9,23 @@ export class MCPToolExecutor implements ToolExecutor { private toolName: string ) {} - async execute(args: Record): Promise { - const result = await this.client.callTool(this.toolName, args); + async execute(args: Record, signal?: AbortSignal): Promise { + const result = await this.client.callTool(this.toolName, args, signal); // 检测 MCP 返回的 content 数组是否包含 image 类型 - if (Array.isArray(result)) { + const structuredResult = + !Array.isArray(result) && + typeof result === "object" && + result !== null && + Array.isArray((result as { content?: unknown }).content) + ? (result as MCPToolCallResult) + : undefined; + const content = Array.isArray(result) ? result : structuredResult?.content; + if (content) { const textParts: string[] = []; const attachments: ToolResultWithAttachments["attachments"] = []; - for (const item of result) { + for (const item of content) { if (item.type === "text" && item.text) { textParts.push(item.text); } else if (item.type === "image" && item.data) { @@ -32,8 +40,15 @@ export class MCPToolExecutor implements ToolExecutor { if (attachments.length > 0) { return { - content: textParts.join("\n") || "Tool completed.", + content: + textParts.join("\n") || + (structuredResult?.structuredContent !== undefined + ? JSON.stringify(structuredResult.structuredContent) + : "Tool completed."), attachments, + ...(structuredResult?.structuredContent !== undefined + ? { structuredContent: structuredResult.structuredContent } + : {}), } as ToolResultWithAttachments; } } diff --git a/src/app/service/agent/core/model_context.test.ts b/src/app/service/agent/core/model_context.test.ts index 39d0e8085..a58951949 100644 --- a/src/app/service/agent/core/model_context.test.ts +++ b/src/app/service/agent/core/model_context.test.ts @@ -1,5 +1,13 @@ import { describe, expect, it } from "vitest"; -import { getContextWindow, inferContextWindow, DEFAULT_CONTEXT_WINDOW } from "./model_context"; +import { + getContextWindow, + getInputTokenBudget, + getReservedOutputTokens, + inferContextWindow, + normalizeModelLimits, + DEFAULT_CONTEXT_WINDOW, + DEFAULT_ANTHROPIC_MAX_TOKENS, +} from "./model_context"; describe("getContextWindow", () => { it("returns user-configured contextWindow when provided", () => { @@ -53,6 +61,71 @@ describe("getContextWindow", () => { it("returns default for unknown models", () => { expect(getContextWindow({ model: "my-custom-model" })).toBe(DEFAULT_CONTEXT_WINDOW); }); + + it("负数 contextWindow 是 truthy,不能被原样返回,否则 getInputTokenBudget 会塌缩为 0", () => { + expect(getContextWindow({ model: "gpt-4o", contextWindow: -1 })).toBe(128_000); + expect(getInputTokenBudget({ model: "gpt-4o", contextWindow: -1, provider: "openai" } as any)).toBeGreaterThan(0); + }); + + it("非有限数 contextWindow 应回退到前缀匹配", () => { + expect(getContextWindow({ model: "gpt-4o", contextWindow: Infinity })).toBe(128_000); + expect(getContextWindow({ model: "gpt-4o", contextWindow: NaN })).toBe(128_000); + }); +}); + +describe("getReservedOutputTokens / getInputTokenBudget(小上下文窗口不应把输入预算压成 0)", () => { + it("未配置 maxTokens 时,大窗口模型仍保留默认输出预留", () => { + expect(getReservedOutputTokens({ model: "claude-3-haiku", provider: "anthropic" } as any)).toBe( + DEFAULT_ANTHROPIC_MAX_TOKENS + ); + expect(getReservedOutputTokens({ model: "gpt-4o", provider: "openai" } as any)).toBe(DEFAULT_ANTHROPIC_MAX_TOKENS); + }); + + it("未配置 maxTokens 时,默认输出预留不应吃掉小窗口模型的全部输入预算", () => { + // gpt-4 基础版窗口 8192:预留 16384 会让输入预算塌缩为 0,导致每次对话都直接报 context_too_large + for (const model of ["gpt-4-0613", "gpt-3.5-turbo", "phi-3-mini"]) { + const budget = getInputTokenBudget({ model, provider: "openai" } as any); + expect(budget, `${model} 的输入预算不应为 0`).toBeGreaterThan(0); + } + }); + + it("用户配置的小 contextWindow(本地小模型)同样保留可用的输入预算", () => { + const budget = getInputTokenBudget({ model: "my-local-model", contextWindow: 8192, provider: "openai" } as any); + expect(budget).toBeGreaterThan(0); + }); + + it("显式配置的 maxTokens 原样生效,不被默认值改写", () => { + expect(getReservedOutputTokens({ model: "gpt-4-0613", maxTokens: 4096, provider: "openai" } as any)).toBe(4096); + expect(getInputTokenBudget({ model: "gpt-4-0613", maxTokens: 4096, provider: "openai" } as any)).toBeGreaterThan(0); + }); +}); + +describe("normalizeModelLimits(存储边界统一归一化)", () => { + it("负数 / 非有限数 / 超出合理范围一律归一化为 undefined", () => { + expect(normalizeModelLimits({ maxTokens: -5, contextWindow: -1 })).toEqual({ + maxTokens: undefined, + contextWindow: undefined, + }); + expect(normalizeModelLimits({ maxTokens: Infinity, contextWindow: NaN })).toEqual({ + maxTokens: undefined, + contextWindow: undefined, + }); + expect(normalizeModelLimits({ maxTokens: 50_000_000, contextWindow: 50_000_000 })).toEqual({ + maxTokens: undefined, + contextWindow: undefined, + }); + }); + + it("合法正整数保持不变(向下取整)", () => { + expect(normalizeModelLimits({ maxTokens: 4096.7, contextWindow: 128_000 })).toEqual({ + maxTokens: 4096, + contextWindow: 128_000, + }); + }); + + it("未配置(undefined)保持 undefined,交给下游默认值", () => { + expect(normalizeModelLimits({})).toEqual({ maxTokens: undefined, contextWindow: undefined }); + }); }); describe("inferContextWindow", () => { diff --git a/src/app/service/agent/core/model_context.ts b/src/app/service/agent/core/model_context.ts index 46ea4e748..ee0e204da 100644 --- a/src/app/service/agent/core/model_context.ts +++ b/src/app/service/agent/core/model_context.ts @@ -1,3 +1,5 @@ +import type { AgentModelConfig } from "./types"; + // 模型上下文窗口大小映射表 // [前缀, 上下文窗口大小],按前缀长度降序排列以确保最精确匹配优先 const MODEL_CONTEXT_PREFIXES: Array<[string, number]> = [ @@ -41,10 +43,16 @@ const MODEL_CONTEXT_PREFIXES: Array<[string, number]> = [ ]; export const DEFAULT_CONTEXT_WINDOW = 128_000; +export const DEFAULT_ANTHROPIC_MAX_TOKENS = 16_384; +export const CONTEXT_SAFETY_MARGIN_RATIO = 0.1; /** 获取模型的上下文窗口大小,优先使用用户配置,否则按前缀匹配 */ export function getContextWindow(config: { model: string; contextWindow?: number }): number { - if (config.contextWindow) return config.contextWindow; + // > 0 而非直接 truthy 判断:负数是 truthy,会原样返回并让 getInputTokenBudget() 的预算计算 + // 塌缩为 0 + if (typeof config.contextWindow === "number" && Number.isFinite(config.contextWindow) && config.contextWindow > 0) { + return config.contextWindow; + } const modelLower = config.model.toLowerCase(); for (const [prefix, size] of MODEL_CONTEXT_PREFIXES) { if (modelLower.startsWith(prefix)) return size; @@ -52,6 +60,37 @@ export function getContextWindow(config: { model: string; contextWindow?: number return DEFAULT_CONTEXT_WINDOW; } +/** + * 获取 provider 实际会请求的最大输出 token 数。 + * 未显式配置 maxTokens 时,不能按 0 预留:OpenAI 兼容请求体在这种情况下会直接省略 + * max_tokens 字段(见 providers/openai.ts),provider 侧会套用它自己的默认输出上限 + * (通常有实质数值,不是 0)。这里统一用一个有记录的保守默认值兜底,而不是假装输出不占预算, + * 否则输入预算会把整个 contextWindow 都算给输入,实际请求的 输入+输出 可能超出真实上限。 + */ +export function getReservedOutputTokens(config: AgentModelConfig): number { + const contextWindow = getContextWindow(config); + const configured = + typeof config.maxTokens === "number" && Number.isFinite(config.maxTokens) + ? Math.max(0, Math.floor(config.maxTokens)) + : 0; + // 未显式配置时,默认预留不能超过窗口的 1/4:否则小窗口模型(gpt-4 8K、gpt-3.5 16K、 + // 本地小模型等)会被 16384 的默认预留 + 安全边际吃掉全部输入预算, + // getInputTokenBudget() 塌缩为 0,每次对话都直接报 context_too_large。 + const requested = + configured > 0 ? configured : Math.min(DEFAULT_ANTHROPIC_MAX_TOKENS, Math.max(1, Math.floor(contextWindow / 4))); + return Math.min(contextWindow, requested); +} + +/** + * 发送请求前允许输入占用的最大 token 预算。 + * contextWindow 是输入与输出的总和,因此需同时预留 provider 请求的输出额度和安全边际。 + */ +export function getInputTokenBudget(config: AgentModelConfig): number { + const contextWindow = getContextWindow(config); + const safetyMargin = Math.ceil(contextWindow * CONTEXT_SAFETY_MARGIN_RATIO); + return Math.max(0, contextWindow - getReservedOutputTokens(config) - safetyMargin); +} + /** 根据模型名称推断上下文窗口大小(不考虑用户配置) */ export function inferContextWindow(model: string): number { const modelLower = model.toLowerCase(); @@ -60,3 +99,29 @@ export function inferContextWindow(model: string): number { } return DEFAULT_CONTEXT_WINDOW; } + +// 允许的 token 数上限:远超已知最大模型(Gemini/GPT-4.1 约 1M),为未来更大模型留余量, +// 同时拒绝明显异常的值(Infinity、Number.MAX_SAFE_INTEGER 等) +const MAX_REASONABLE_TOKEN_LIMIT = 10_000_000; + +/** + * 把用户可能填入的 maxTokens/contextWindow 归一化为有限正整数(或 undefined,交给下游默认值)。 + * 非有限数、非正数、超出合理范围一律视为未配置——防止负数因为 JS 的 truthy 判断被当作"已配置" + * 直接发给 provider(如 Anthropic 的 max_tokens = config.maxTokens || 16384 会把负数原样发出), + * 也防止负的 contextWindow 让 getInputTokenBudget() 的预算计算塌缩为 0。 + */ +function normalizeTokenLimit(value: number | undefined): number | undefined { + if (typeof value !== "number" || !Number.isFinite(value)) return undefined; + const normalized = Math.floor(value); + if (normalized <= 0 || normalized > MAX_REASONABLE_TOKEN_LIMIT) return undefined; + return normalized; +} + +/** 在模型配置持久化前统一归一化 maxTokens/contextWindow,作为存储边界的唯一校验点。 */ +export function normalizeModelLimits(config: T): T { + return { + ...config, + maxTokens: normalizeTokenLimit(config.maxTokens), + contextWindow: normalizeTokenLimit(config.contextWindow), + }; +} diff --git a/src/app/service/agent/core/opfs_helpers.ts b/src/app/service/agent/core/opfs_helpers.ts index 452110751..88ac18f77 100644 --- a/src/app/service/agent/core/opfs_helpers.ts +++ b/src/app/service/agent/core/opfs_helpers.ts @@ -1,6 +1,8 @@ // OPFS 工作区公共辅助函数 // 供 opfs_tools、agent_dom 等模块复用 +import { throwIfAborted } from "./abort_utils"; + export const WORKSPACE_ROOT = "agents/workspace"; /** Strip leading `/`, reject `..` segments */ @@ -19,20 +21,24 @@ export function sanitizePath(raw: string): string { export async function getDirectory( root: FileSystemDirectoryHandle, path: string, - create = false + create = false, + signal?: AbortSignal ): Promise { const segments = path.split("/").filter(Boolean); let dir = root; for (const seg of segments) { + throwIfAborted(signal); dir = await dir.getDirectoryHandle(seg, { create }); } return dir; } /** Get the workspace root directory handle */ -export async function getWorkspaceRoot(create = false): Promise { +export async function getWorkspaceRoot(create = false, signal?: AbortSignal): Promise { + throwIfAborted(signal); const opfsRoot = await navigator.storage.getDirectory(); - return getDirectory(opfsRoot, WORKSPACE_ROOT, create); + throwIfAborted(signal); + return getDirectory(opfsRoot, WORKSPACE_ROOT, create, signal); } /** Split a sanitized path into parent directory path and filename */ @@ -71,8 +77,10 @@ export function isDataUrl(str: string): boolean { /** 将二进制数据写入 OPFS workspace 指定路径 */ export async function writeWorkspaceFile( path: string, - data: Uint8Array | Blob | string + data: Uint8Array | Blob | string, + signal?: AbortSignal ): Promise<{ path: string; size: number }> { + throwIfAborted(signal); const safePath = sanitizePath(path); if (!safePath) throw new Error("path is required"); @@ -82,21 +90,33 @@ export async function writeWorkspaceFile( data = decoded.data; } - const workspace = await getWorkspaceRoot(true); + const workspace = await getWorkspaceRoot(true, signal); const { dirPath, fileName } = splitPath(safePath); - const dir = dirPath ? await getDirectory(workspace, dirPath, true) : workspace; + const dir = dirPath ? await getDirectory(workspace, dirPath, true, signal) : workspace; + throwIfAborted(signal); const fileHandle = await dir.getFileHandle(fileName, { create: true }); + throwIfAborted(signal); const writable = await fileHandle.createWritable(); - - if (data instanceof Blob) { - await writable.write(data); - } else if (data instanceof Uint8Array) { - // 精确截取视图对应的字节段,避免切片视图写入整个底层 buffer - await writable.write((data.buffer as ArrayBuffer).slice(data.byteOffset, data.byteOffset + data.byteLength)); - } else { - await writable.write(data); + try { + throwIfAborted(signal); + if (data instanceof Blob) { + await writable.write(data); + } else if (data instanceof Uint8Array) { + // 精确截取视图对应的字节段,避免切片视图写入整个底层 buffer + await writable.write((data.buffer as ArrayBuffer).slice(data.byteOffset, data.byteOffset + data.byteLength)); + } else { + await writable.write(data); + } + throwIfAborted(signal); + await writable.close(); + } catch (error) { + try { + await writable.abort(); + } catch { + // 写入流已经关闭或终止时无需额外处理 + } + throw error; } - await writable.close(); let size: number; if (data instanceof Blob) { diff --git a/src/app/service/agent/core/persisted_messages.test.ts b/src/app/service/agent/core/persisted_messages.test.ts new file mode 100644 index 000000000..d686a8134 --- /dev/null +++ b/src/app/service/agent/core/persisted_messages.test.ts @@ -0,0 +1,44 @@ +import { describe, expect, it } from "vitest"; +import { toLLMMessages } from "./persisted_messages"; +import type { ChatMessage } from "./types"; + +describe("持久化工具协议恢复", () => { + it("应为缺失结果补错并丢弃孤立、重复和未知工具结果", () => { + const messages: ChatMessage[] = [ + { + id: "assistant", + conversationId: "conv", + role: "assistant", + content: "", + toolCalls: [ + { id: "call-1", name: "one", arguments: "{}" }, + { id: "call-2", name: "two", arguments: "{}" }, + ], + createtime: 1, + }, + { id: "result-1", conversationId: "conv", role: "tool", content: "first", toolCallId: "call-1", createtime: 2 }, + { + id: "duplicate", + conversationId: "conv", + role: "tool", + content: "duplicate", + toolCallId: "call-1", + createtime: 3, + }, + { id: "unknown", conversationId: "conv", role: "tool", content: "unknown", toolCallId: "other", createtime: 4 }, + { id: "orphan", conversationId: "conv", role: "tool", content: "orphan", toolCallId: "orphan", createtime: 5 }, + { id: "user", conversationId: "conv", role: "user", content: "next", createtime: 6 }, + ]; + + const normalized = toLLMMessages(messages); + + expect(normalized.map((message) => [message.role, message.toolCallId])).toEqual([ + ["assistant", undefined], + ["tool", "call-1"], + ["tool", "call-2"], + ["user", undefined], + ]); + expect(normalized[1].content).toBe("first"); + expect(normalized[2].content).toContain("recovery"); + }); +}); diff --git a/src/app/service/agent/core/persisted_messages.ts b/src/app/service/agent/core/persisted_messages.ts new file mode 100644 index 000000000..7ce9c5c6f --- /dev/null +++ b/src/app/service/agent/core/persisted_messages.ts @@ -0,0 +1,87 @@ +import type { ChatMessage, ChatRequest, MessageContent } from "./types"; + +// 会话创建于所有权模型引入之前:agent_chat.ts 的 normalizeConversation 为这类记录统一回填 +// "legacy:" 前缀 generation。放在这里(而非 repo 层)是因为大量测试会整体 mock agent_chat 模块, +// 从中导入纯函数会在那些测试里解析为 undefined。 +export function isLegacyGeneration(generation: string | undefined): boolean { + return generation?.startsWith("legacy:") ?? false; +} + +// content block 里携带的附件 id(image/file/audio block 均以 attachmentId 引用) +function contentBlockAttachmentIds(content: MessageContent): string[] { + if (typeof content === "string") return []; + return content + .filter((block) => block.type !== "text") + .map((block) => (block as { attachmentId: string }).attachmentId); +} + +/** Keep ownership only for durable artifacts that a compacted summary explicitly keeps addressable. + * undefined/空的 ownedAttachmentIds 在当前模型里合法地表示"不拥有,content block 引用的都是借用", + * 因此只有 legacy=true(会话创建于所有权模型引入之前,见 agent_chat.ts 的 "legacy:" generation 前缀) + * 时才退化为按 content block / 工具附件元数据推断候选集合;否则升级前对话摘要保留的旧附件 + * 永远无法被正确识别为候选,即使摘要文本里仍然显式引用了它。*/ +export function retainedSummaryAttachmentIds(summary: string, messages: ChatMessage[], legacy = false): string[] { + const owned = new Set(); + const collectToolCalls = (toolCalls: NonNullable) => { + for (const toolCall of toolCalls) { + if (toolCall.ownedAttachmentIds !== undefined) { + for (const id of toolCall.ownedAttachmentIds) owned.add(id); + } else if (legacy) { + for (const attachment of toolCall.attachments || []) owned.add(attachment.id); + } + for (const message of toolCall.subAgentDetails?.messages || []) { + if (legacy) { + for (const id of contentBlockAttachmentIds(message.content)) owned.add(id); + } + collectToolCalls(message.toolCalls); + } + } + }; + for (const message of messages) { + if (message.ownedAttachmentIds !== undefined) { + for (const id of message.ownedAttachmentIds) owned.add(id); + } else if (legacy) { + for (const id of contentBlockAttachmentIds(message.content)) owned.add(id); + } + collectToolCalls(message.toolCalls || []); + } + return [...owned].filter((id) => summary.includes(`uploads/${id}`)); +} + +/** 将持久化消息转换为 LLM 消息,错误占位消息只用于 UI 展示,不应重放。 */ +export function toLLMMessages( + messages: Array> +): ChatRequest["messages"] { + const normalized: ChatRequest["messages"] = []; + for (let index = 0; index < messages.length; index++) { + const message = messages[index]; + if (message.error || message.role === "tool") continue; + normalized.push({ + role: message.role, + content: message.content, + toolCallId: message.toolCallId, + toolCalls: message.toolCalls, + }); + if (message.role !== "assistant" || !message.toolCalls?.length) continue; + + const results = new Map(); + let cursor = index + 1; + while (cursor < messages.length && messages[cursor].role === "tool") { + const toolMessage = messages[cursor]; + if (!toolMessage.error && toolMessage.toolCallId && !results.has(toolMessage.toolCallId)) { + results.set(toolMessage.toolCallId, toolMessage); + } + cursor++; + } + for (const toolCall of message.toolCalls) { + const result = results.get(toolCall.id); + normalized.push({ + role: "tool", + content: result?.content ?? JSON.stringify({ error: "Tool result unavailable after recovery" }), + toolCallId: toolCall.id, + }); + } + index = cursor - 1; + } + return normalized; +} diff --git a/src/app/service/agent/core/providers/anthropic.test.ts b/src/app/service/agent/core/providers/anthropic.test.ts index 36697e7b9..a2c20d039 100644 --- a/src/app/service/agent/core/providers/anthropic.test.ts +++ b/src/app/service/agent/core/providers/anthropic.test.ts @@ -170,6 +170,95 @@ describe("buildAnthropicRequest", () => { expect(body.system[0].cache_control).toBeUndefined(); // tool 不应有 cache_control expect(body.tools[0].cache_control).toBeUndefined(); + // 消息历史也不应有 cache_control + expect(body.messages[0].content).toBe("hi"); + }); + + describe("消息历史的 cache_control 断点(用于长 tool loop 中复用已缓存前缀)", () => { + it("最后一条纯文本消息应转换为带 cache_control 的 text block", () => { + const request: ChatRequest = { + conversationId: "c1", + modelId: "test", + messages: [ + { role: "user", content: "第一句" }, + { role: "assistant", content: "第一句回复" }, + { role: "user", content: "最后一句" }, + ], + }; + + const { init } = buildAnthropicRequest(config, request); + const body = JSON.parse(init.body as string); + + // 非最后一条消息保持原样(字符串),不应被转换 + expect(body.messages[0].content).toBe("第一句"); + expect(body.messages[1].content).toBe("第一句回复"); + // 最后一条转换为带 cache_control 的 text block + expect(body.messages[2].content).toEqual([ + { type: "text", text: "最后一句", cache_control: { type: "ephemeral" } }, + ]); + }); + + it("最后一条消息已是 tool_result content block 时应在最后一个 block 上加 cache_control", () => { + const request: ChatRequest = { + conversationId: "c1", + modelId: "test", + messages: [ + { role: "user", content: "天气" }, + { + role: "assistant", + content: "让我查一下", + toolCalls: [{ id: "toolu_1", name: "get_weather", arguments: "{}" }], + }, + { role: "tool", content: '{"temp":25}', toolCallId: "toolu_1" }, + ], + }; + + const { init } = buildAnthropicRequest(config, request); + const body = JSON.parse(init.body as string); + + const lastMsg = body.messages[body.messages.length - 1]; + expect(lastMsg.role).toBe("user"); + expect(lastMsg.content[0].type).toBe("tool_result"); + expect(lastMsg.content[0].cache_control).toEqual({ type: "ephemeral" }); + }); + + it("cache: false 时最后一条消息不应转换或添加 cache_control", () => { + const request: ChatRequest = { + conversationId: "c1", + modelId: "test", + messages: [{ role: "user", content: "最后一句" }], + cache: false, + }; + + const { init } = buildAnthropicRequest(config, request); + const body = JSON.parse(init.body as string); + + expect(body.messages[0].content).toBe("最后一句"); + }); + + it("最后一条消息内容为空字符串时不应添加 cache_control(无内容块可挂载)", () => { + const request: ChatRequest = { + conversationId: "c1", + modelId: "test", + messages: [ + { role: "user", content: "天气" }, + { + role: "assistant", + content: "", + toolCalls: [{ id: "toolu_1", name: "get_weather", arguments: "{}" }], + }, + ], + }; + + const { init } = buildAnthropicRequest(config, request); + const body = JSON.parse(init.body as string); + + const lastMsg = body.messages[body.messages.length - 1]; + // 仅有 tool_use block,没有额外的空 text block + expect(lastMsg.content).toHaveLength(1); + expect(lastMsg.content[0].type).toBe("tool_use"); + expect(lastMsg.content[0].cache_control).toEqual({ type: "ephemeral" }); + }); }); it("默认 max_tokens 为 16384,应设置 stream", () => { @@ -201,6 +290,7 @@ describe("buildAnthropicRequest", () => { messages: [ { role: "user", content: "hi" }, { role: "tool", content: "result" }, // 无 toolCallId + { role: "user", content: "继续" }, // 占位,避免上一条被当作最后一条消息加上 cache_control ], }; @@ -374,6 +464,33 @@ describe("parseAnthropicStream", () => { } }); + it("message_start 后的 error 终态应保留已知输入与缓存 usage", async () => { + const reader = createMockReader([ + 'event: message_start\ndata: {"message":{"usage":{"input_tokens":20,"cache_read_input_tokens":5}}}\n\n', + 'event: error\ndata: {"error":{"message":"Overloaded"}}\n\n', + ]); + const events: ChatStreamEvent[] = []; + + await parseAnthropicStream(reader, (event) => events.push(event), new AbortController().signal); + + expect(events.at(-1)).toMatchObject({ + type: "error", + usage: { inputTokens: 20, outputTokens: 0, cacheReadInputTokens: 5 }, + }); + }); + + it("message_start 后直接 message_stop 也应报告已知输入 usage", async () => { + const reader = createMockReader([ + 'event: message_start\ndata: {"message":{"usage":{"input_tokens":9}}}\n\n', + "event: message_stop\ndata: {}\n\n", + ]); + const events: ChatStreamEvent[] = []; + + await parseAnthropicStream(reader, (event) => events.push(event), new AbortController().signal); + + expect(events.at(-1)).toMatchObject({ type: "done", usage: { inputTokens: 9, outputTokens: 0 } }); + }); + it("error 事件无 message 时使用默认错误信息", async () => { const reader = createMockReader(['event: error\ndata: {"error":{}}\n\n']); @@ -388,7 +505,7 @@ describe("parseAnthropicStream", () => { } }); - it("signal 已中止时应停止读取", async () => { + it("signal 已中止时应停止读取并 reject,而不是静默 resolve(否则外层 callLLM 永远挂起)", async () => { const controller = new AbortController(); controller.abort(); @@ -397,11 +514,65 @@ describe("parseAnthropicStream", () => { ]); const events: ChatStreamEvent[] = []; - await parseAnthropicStream(reader, (e) => events.push(e), controller.signal); + await expect(parseAnthropicStream(reader, (e) => events.push(e), controller.signal)).rejects.toThrow("Aborted"); expect(events).toHaveLength(0); }); + it("abort 前已经收到过 message_start usage 时,reject 的错误应携带这部分已知 usage", async () => { + const controller = new AbortController(); + const encoder = new TextEncoder(); + let index = 0; + const chunks = [ + 'event: message_start\ndata: {"message":{"usage":{"input_tokens":20,"cache_read_input_tokens":5}}}\n\n', + ]; + const reader = { + read: async () => { + if (index < chunks.length) { + return { done: false, value: encoder.encode(chunks[index++]) }; + } + controller.abort(); + throw new Error("Aborted"); + }, + cancel: async () => {}, + closed: Promise.resolve(undefined), + releaseLock: () => {}, + } as any; + + const events: ChatStreamEvent[] = []; + await expect(parseAnthropicStream(reader, (e) => events.push(e), controller.signal)).rejects.toMatchObject({ + message: "Aborted", + usage: { inputTokens: 20, outputTokens: 0, cacheReadInputTokens: 5 }, + }); + }); + + it("流正常结束但没有 message_stop/message_delta(usage)/error 终态帧时应补发 error,而不是静默完成", async () => { + // 只有 content_block_delta,reader 提前 done,模拟连接在终态帧之前被服务端关闭 + const reader = createMockReader([ + 'event: content_block_delta\ndata: {"delta":{"type":"text_delta","text":"hi"}}\n\n', + ]); + + const events: ChatStreamEvent[] = []; + const controller = new AbortController(); + await parseAnthropicStream(reader, (e) => events.push(e), controller.signal); + + expect(events.some((e) => e.type === "error")).toBe(true); + }); + + it("message_start 后意外 EOF 的 error 应保留已知 usage", async () => { + const reader = createMockReader([ + 'event: message_start\ndata: {"message":{"usage":{"input_tokens":7,"cache_creation_input_tokens":2}}}\n\n', + ]); + const events: ChatStreamEvent[] = []; + + await parseAnthropicStream(reader, (event) => events.push(event), new AbortController().signal); + + expect(events.at(-1)).toMatchObject({ + type: "error", + usage: { inputTokens: 7, outputTokens: 0, cacheCreationInputTokens: 2 }, + }); + }); + it("读取错误时应发送 error 事件", async () => { const reader = { read: async () => { @@ -424,7 +595,7 @@ describe("parseAnthropicStream", () => { } }); - it("读取错误但 signal 已中止时不应发送 error", async () => { + it("读取错误但 signal 已中止时应 reject 而不是发送 error", async () => { const controller = new AbortController(); const reader = { read: async () => { @@ -437,7 +608,7 @@ describe("parseAnthropicStream", () => { } as any; const events: ChatStreamEvent[] = []; - await parseAnthropicStream(reader, (e) => events.push(e), controller.signal); + await expect(parseAnthropicStream(reader, (e) => events.push(e), controller.signal)).rejects.toThrow("Aborted"); expect(events).toHaveLength(0); }); diff --git a/src/app/service/agent/core/providers/anthropic.ts b/src/app/service/agent/core/providers/anthropic.ts index 27486e164..c268c96f2 100644 --- a/src/app/service/agent/core/providers/anthropic.ts +++ b/src/app/service/agent/core/providers/anthropic.ts @@ -1,6 +1,7 @@ import type { ChatStreamEvent, ChatRequest, ContentBlock } from "../types"; import type { AgentModelConfig } from "../types"; import { isContentBlocks } from "../content_utils"; +import { getReservedOutputTokens } from "../model_context"; import { generateAttachmentId, convertTextBlock, @@ -117,13 +118,30 @@ export function buildAnthropicRequest( return { role: m.role, content: m.content }; }); + // 最后一条消息追加 cache_control:使本轮完整历史被缓存。 + // 下一轮请求中该消息不再是最后一条,但 Anthropic 按前缀匹配复用已缓存部分, + // 只需为新增的增量消息计费,从而避免长 tool loop 下每轮都全量计费输入 token。 + if (useCache && messages.length > 0) { + const lastMessage = messages[messages.length - 1]; + const content: unknown = lastMessage.content; + if (Array.isArray(content) && content.length > 0) { + (content[content.length - 1] as Record).cache_control = { type: "ephemeral" }; + } else if (typeof content === "string" && content.length > 0) { + (lastMessage as { content: unknown }).content = [ + { type: "text", text: content, cache_control: { type: "ephemeral" } }, + ]; + } + } + const body: Record = { model: config.model, messages, stream: true, }; - body.max_tokens = config.maxTokens || 16384; + // 用归一化后的 getReservedOutputTokens 而不是 config.maxTokens || 16384: + // 负数在 JS 里是 truthy,会绕过 || 兜底原样发给 provider + body.max_tokens = getReservedOutputTokens(config) || 16384; if (systemMessages.length > 0) { const systemBlocks = systemMessages.map((m) => ({ @@ -187,6 +205,10 @@ export function parseAnthropicStream( // 跟踪图片块的累积 base64 数据 let imageBlockData: { index: number; mediaType: string; base64Chunks: string[] } | null = null; + // 标记是否已发出终态事件(message_stop / message_delta 带 usage / error), + // 避免连接在收到终态帧前就断开时,调用方(LLMClient.callLLM)的外层 Promise 永远不 settle + let doneSent = false; + const toolUseByIndex = new Map(); return readSSEStream( @@ -288,6 +310,7 @@ export function parseAnthropicStream( case "message_delta": { // 消息结束,合并 message_start 的 input usage 和 message_delta 的 output usage if (json.usage) { + doneSent = true; onEvent({ type: "done", usage: { @@ -303,13 +326,19 @@ export function parseAnthropicStream( } case "message_stop": { toolUseByIndex.clear(); - onEvent({ type: "done" }); + doneSent = true; + onEvent({ + type: "done", + usage: cachedUsage ? { ...cachedUsage, outputTokens: 0 } : undefined, + }); return true; } case "error": { + doneSent = true; onEvent({ type: "error", message: json.error?.message || "Anthropic API error", + usage: cachedUsage ? { ...cachedUsage, outputTokens: 0 } : undefined, }); return true; } @@ -319,8 +348,34 @@ export function parseAnthropicStream( } return false; }, - (message) => onEvent({ type: "error", message }) - ); + (message) => { + doneSent = true; + onEvent({ + type: "error", + message, + usage: cachedUsage ? { ...cachedUsage, outputTokens: 0 } : undefined, + }); + } + ) + .then(() => { + // 流正常结束(reader done)但没收到 message_stop / message_delta(usage) / error 终态帧, + // 必须补发终态事件,否则调用方的外层 Promise 会永远挂起 + if (!signal.aborted && !doneSent) { + onEvent({ + type: "error", + message: "Stream ended unexpectedly without a terminal frame", + usage: cachedUsage ? { ...cachedUsage, outputTokens: 0 } : undefined, + }); + } + }) + .catch((error) => { + // readSSEStream 只在 abort 时才 reject(见 content_utils.ts);message_start 可能已经 + // 带回了 input/cache usage,把它带在 abort 错误上,避免取消时把这部分已知花费从终态 + // usage 里丢掉。outputTokens 在 message_delta 之前始终未知,记 0。 + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { + usage: cachedUsage ? { ...cachedUsage, outputTokens: 0 } : undefined, + }); + }); } // ---- LLMProvider 接口适配 ---- diff --git a/src/app/service/agent/core/providers/content_utils.ts b/src/app/service/agent/core/providers/content_utils.ts index 69bacc3f5..b96316f8b 100644 --- a/src/app/service/agent/core/providers/content_utils.ts +++ b/src/app/service/agent/core/providers/content_utils.ts @@ -1,6 +1,7 @@ import type { ContentBlock } from "../types"; import { SSEParser } from "../sse_parser"; import type { SSEEvent } from "../sse_parser"; +import { createAbortError } from "../abort_utils"; /** * 生成图片附件 ID,格式:img_{时间戳}_{随机串}.{扩展名} @@ -26,13 +27,20 @@ export function convertFileBlock(block: Extract) }; } +/** + * image 块无法解析时的文本降级描述 + */ +export function imageBlockFallbackText(block: { name?: string; attachmentId: string }): string { + return `[Image: ${block.name || "image"}, OPFS path: uploads/${block.attachmentId}]`; +} + /** * image 块无法解析时的文本降级描述 */ export function imageBlockFallback(block: Extract): Record { return { type: "text", - text: `[Image: ${block.name || "image"}, OPFS path: uploads/${block.attachmentId}]`, + text: imageBlockFallbackText(block), }; } @@ -58,7 +66,16 @@ export async function readSSEStream( ): Promise { const parser = new SSEParser(); const decoder = new TextDecoder(); + // abort 时主动 cancel reader,唤醒可能卡在 reader.read() 上的等待 + const onAbort = () => { + reader.cancel().catch(() => {}); + }; + signal.addEventListener("abort", onAbort, { once: true }); + // onEvent 提前终止(返回 true)时,body 里可能还有未读完的数据(例如 Anthropic 在 + // message_delta 就终止本地处理,message_stop 及之后的数据从未被读取);finally 里需要 + // 据此决定是否主动 cancel 释放底层连接资源,避免 reader lock 一直占用到 GC + let earlyExit = false; try { while (!signal.aborted) { const { done, value } = await reader.read(); @@ -69,11 +86,26 @@ export async function readSSEStream( for (const sseEvent of events) { // onEvent 返回 true 代表流处理完毕,提前退出 - if (onEvent(sseEvent)) return; + if (onEvent(sseEvent)) { + earlyExit = true; + return; + } } } + // abort 导致循环退出:必须 reject 而不是静默 resolve,否则调用方(LLMClient.callLLM) + // 的外层 Promise 永远不会 settle,取消会一直挂起 + if (signal.aborted) throw createAbortError(); } catch (e: any) { - if (signal.aborted) return; + if (signal.aborted) throw createAbortError(); onError(e.message || "Stream read error"); + } finally { + signal.removeEventListener("abort", onAbort); + if (earlyExit) { + try { + await reader.cancel(); + } catch { + // 已经关闭/abort 时 cancel 可能抛错,忽略 + } + } } } diff --git a/src/app/service/agent/core/providers/openai.test.ts b/src/app/service/agent/core/providers/openai.test.ts index b145895fc..305e89216 100644 --- a/src/app/service/agent/core/providers/openai.test.ts +++ b/src/app/service/agent/core/providers/openai.test.ts @@ -95,6 +95,28 @@ describe("buildOpenAIRequest", () => { expect(body.stream).toBe(true); expect(body.stream_options).toEqual({ include_usage: true }); }); + + it("未建模音频能力时应稳定使用 OPFS 文本引用而不内联二进制", () => { + const { init } = buildOpenAIRequest( + config, + { + conversationId: "c1", + modelId: "test", + messages: [ + { + role: "user", + content: [{ type: "audio", attachmentId: "shared.wav", mimeType: "audio/wav", name: "sample.wav" }], + }, + ], + }, + () => "data:audio/wav;base64,AAAA" + ); + + const body = JSON.parse(init.body as string); + expect(body.messages[0].content).toEqual([ + { type: "text", text: "[Audio: sample.wav, OPFS path: uploads/shared.wav]" }, + ]); + }); }); // 辅助函数:创建 mock ReadableStreamDefaultReader @@ -157,7 +179,7 @@ describe("parseOpenAIStream", () => { it("应正确处理 usage 信息", async () => { const reader = createMockReader([ - 'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n', + 'data: {"choices":[{"delta":{"content":"hi"},"finish_reason":"stop"}]}\n\n', 'data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n', ]); @@ -175,7 +197,7 @@ describe("parseOpenAIStream", () => { it("应正确处理含 cached_tokens 的 usage 信息", async () => { const reader = createMockReader([ - 'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n', + 'data: {"choices":[{"delta":{"content":"hi"},"finish_reason":"stop"}]}\n\n', 'data: {"usage":{"prompt_tokens":100,"completion_tokens":20,"prompt_tokens_details":{"cached_tokens":80}}}\n\n', ]); @@ -206,6 +228,21 @@ describe("parseOpenAIStream", () => { } }); + it("已知 usage 后收到 API 错误帧时应把用量保留在终态事件", async () => { + const reader = createMockReader([ + 'data: {"usage":{"prompt_tokens":12,"completion_tokens":4}}\n\n', + 'data: {"error":{"message":"Rate limit exceeded"}}\n\n', + ]); + const events: ChatStreamEvent[] = []; + + await parseOpenAIStream(reader, (event) => events.push(event), new AbortController().signal); + + expect(events.at(-1)).toMatchObject({ + type: "error", + usage: { inputTokens: 12, outputTokens: 4 }, + }); + }); + it("应忽略无 choices 的事件", async () => { const reader = createMockReader([ 'data: {"id":"chatcmpl-xxx","object":"chat.completion.chunk"}\n\n', @@ -238,18 +275,46 @@ describe("parseOpenAIStream", () => { expect(events[0]).toEqual({ type: "content_delta", delta: "ok" }); }); - it("signal 已中止时应停止读取", async () => { + it("signal 已中止时应停止读取并 reject,而不是静默 resolve(否则外层 callLLM 永远挂起)", async () => { const controller = new AbortController(); controller.abort(); const reader = createMockReader(['data: {"choices":[{"delta":{"content":"hello"}}]}\n\n']); const events: ChatStreamEvent[] = []; - await parseOpenAIStream(reader, (e) => events.push(e), controller.signal); + await expect(parseOpenAIStream(reader, (e) => events.push(e), controller.signal)).rejects.toThrow("Aborted"); expect(events).toHaveLength(0); }); + it("abort 前已经收到过 usage chunk 时,reject 的错误应携带这部分已知 usage", async () => { + const controller = new AbortController(); + const encoder = new TextEncoder(); + let index = 0; + const chunks = [ + 'data: {"choices":[{"delta":{"content":"hi"}}],"usage":{"prompt_tokens":10,"completion_tokens":3}}\n\n', + ]; + const reader = { + read: async () => { + if (index < chunks.length) { + return { done: false, value: encoder.encode(chunks[index++]) }; + } + // 第一个带 usage 的 chunk 读取完毕后才 abort,模拟"取消发生在已经拿到部分 usage 之后" + controller.abort(); + throw new Error("Aborted"); + }, + cancel: async () => {}, + closed: Promise.resolve(undefined), + releaseLock: () => {}, + } as any; + + const events: ChatStreamEvent[] = []; + await expect(parseOpenAIStream(reader, (e) => events.push(e), controller.signal)).rejects.toMatchObject({ + message: "Aborted", + usage: { inputTokens: 10, outputTokens: 3 }, + }); + }); + it("读取错误时应发送 error 事件", async () => { const reader = { read: async () => { @@ -272,7 +337,7 @@ describe("parseOpenAIStream", () => { } }); - it("读取错误但 signal 已中止时不应发送 error 事件", async () => { + it("读取错误但 signal 已中止时应 reject 而不是发送 error 事件", async () => { const controller = new AbortController(); const reader = { read: async () => { @@ -285,7 +350,7 @@ describe("parseOpenAIStream", () => { } as any; const events: ChatStreamEvent[] = []; - await parseOpenAIStream(reader, (e) => events.push(e), controller.signal); + await expect(parseOpenAIStream(reader, (e) => events.push(e), controller.signal)).rejects.toThrow("Aborted"); expect(events).toHaveLength(0); }); @@ -431,6 +496,22 @@ describe("parseOpenAIStream", () => { } }); + it("EOF 前从未见过 finish_reason 时不应当作成功完成——网络中断和正常结束无法区分", async () => { + // 只有内容 delta,没有任何 chunk 带 finish_reason,也没有 [DONE]:模拟网络中断 + const reader = createMockReader([ + 'data: {"choices":[{"delta":{"content":"部分"}}]}\n\n', + 'data: {"choices":[{"delta":{"content":"回答"}}]}\n\n', + ]); + + const events: ChatStreamEvent[] = []; + const controller = new AbortController(); + + await parseOpenAIStream(reader, (e) => events.push(e), controller.signal); + + expect(events).toHaveLength(3); + expect(events[2].type).toBe("error"); + }); + it("应解析单个 chunk 内的 ... 标签", async () => { const reader = createMockReader([ 'data: {"choices":[{"delta":{"content":"beforereasoningafter"}}]}\n\n', diff --git a/src/app/service/agent/core/providers/openai.ts b/src/app/service/agent/core/providers/openai.ts index 619c14a17..799766c8e 100644 --- a/src/app/service/agent/core/providers/openai.ts +++ b/src/app/service/agent/core/providers/openai.ts @@ -35,19 +35,9 @@ function convertContentBlocks( result.push(convertFileBlock(block)); break; case "audio": { - const data = attachmentResolver?.(block.attachmentId); - if (data) { - const match = data.match(/^data:([^;]+);base64,(.+)$/s); - if (match) { - // 从 mimeType 提取格式 (e.g. "audio/wav" → "wav") - const format = block.mimeType.split("/")[1] || "wav"; - result.push({ type: "input_audio", input_audio: { data: match[2], format } }); - } else { - result.push(audioBlockFallback(block)); - } - } else { - result.push(audioBlockFallback(block)); - } + // Audio capability is not modeled or admitted by attachment preflight. Keep the durable OPFS reference + // as text instead of exposing an unreachable binary branch that request accounting cannot budget. + result.push(audioBlockFallback(block)); break; } } @@ -100,8 +90,9 @@ export function buildOpenAIRequest( stream_options: { include_usage: true }, }; - if (config.maxTokens) { - body.max_tokens = config.maxTokens; + // config.maxTokens > 0 而非直接 truthy 判断:负数在 JS 里是 truthy,会绕过校验原样发给 provider + if (typeof config.maxTokens === "number" && Number.isFinite(config.maxTokens) && config.maxTokens > 0) { + body.max_tokens = Math.floor(config.maxTokens); } // 添加工具定义 @@ -149,6 +140,11 @@ export function parseOpenAIStream( | undefined; // 标记是否已通过 [DONE] 信号发出了 done 事件,避免 .then() 再次发出 let doneSent = false; + // 标记是否观察到过 provider 的完成标记(finish_reason 非空)。 + // reader EOF 但没有 [DONE] 时,只有见过 finish_reason 才能确认是"提前不发 [DONE] 的 + // 非标准 provider";否则无法区分"网络中断"与"正常结束",不应把截断的部分内容当作 + // 完整答案持久化。 + let sawFinishReason = false; // 跨 chunk 追踪 ... 块状态(用于把思考混在 content 里的模型) let inThinkBlock = false; @@ -186,12 +182,14 @@ export function parseOpenAIStream( onEvent({ type: "error", message: json.error.message || JSON.stringify(json.error), + usage: lastUsage, }); return true; } const choice = json.choices?.[0]; if (choice) { + if (choice.finish_reason != null) sawFinishReason = true; const delta = choice.delta; if (delta) { // 思考过程增量(reasoning_content 兼容 deepseek / openai o-series) @@ -304,15 +302,28 @@ export function parseOpenAIStream( }, (message) => { doneSent = true; - onEvent({ type: "error", message }); - } - ).then(() => { - // 流正常结束但没收到 [DONE](某些 API 可能如此) - if (!signal.aborted && !doneSent) { - flushThinkCarry(); - onEvent({ type: "done", usage: lastUsage }); + onEvent({ type: "error", message, usage: lastUsage }); } - }); + ) + .then(() => { + // 流正常结束但没收到 [DONE]。只有见过 finish_reason(某些 API 提前不发 [DONE] 但仍会 + // 标出完成原因)才能确认这是真正的完整回答;否则无法与网络中断区分,按未预期断连报错, + // 不能把可能截断的部分内容当作成功答案持久化 + if (!signal.aborted && !doneSent) { + flushThinkCarry(); + if (sawFinishReason) { + onEvent({ type: "done", usage: lastUsage }); + } else { + onEvent({ type: "error", message: "Stream ended unexpectedly without a finish reason", usage: lastUsage }); + } + } + }) + .catch((error) => { + // readSSEStream 只在 abort 时才 reject(见 content_utils.ts);把已知的部分 usage + // (如某些 provider 每个 chunk 都带 usage)带在 abort 错误上,避免取消时把这部分 + // 已经产生的花费从终态 usage 里丢掉 + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { usage: lastUsage }); + }); } // ---- LLMProvider 接口适配 ---- diff --git a/src/app/service/agent/core/session_tool_registry.test.ts b/src/app/service/agent/core/session_tool_registry.test.ts index 97a5704d0..9de97a18a 100644 --- a/src/app/service/agent/core/session_tool_registry.test.ts +++ b/src/app/service/agent/core/session_tool_registry.test.ts @@ -163,6 +163,23 @@ describe("SessionToolRegistry", () => { expect(parentSpy).not.toHaveBeenCalled(); }); + it("应将 signal 传递给 session 工具执行器", async () => { + const parent = new ToolRegistry(); + const session = new SessionToolRegistry(parent); + const executeSpy = vi.fn().mockResolvedValue("session_result"); + session.register("session", taskDef, { execute: executeSpy }); + const controller = new AbortController(); + + await session.execute( + [{ id: "t1", name: "create_task", arguments: "{}" }], + undefined, + undefined, + controller.signal + ); + + expect(executeSpy).toHaveBeenCalledWith({}, controller.signal, "t1"); + }); + it("session 无该工具时回退到 parent 工具", async () => { const parent = new ToolRegistry(); parent.registerBuiltin( diff --git a/src/app/service/agent/core/session_tool_registry.ts b/src/app/service/agent/core/session_tool_registry.ts index 9d0f8186f..7ae060ebf 100644 --- a/src/app/service/agent/core/session_tool_registry.ts +++ b/src/app/service/agent/core/session_tool_registry.ts @@ -85,7 +85,8 @@ export class SessionToolRegistry implements ToolExecutorLike { async execute( toolCalls: ToolCall[], scriptCallback?: ScriptToolCallback | null, - excludeTools?: Set + excludeTools?: Set, + signal?: AbortSignal ): Promise { // 构建合并 Map:parent 在下、session 在上(session 覆盖 parent 同名工具) const merged = new Map(); @@ -95,6 +96,6 @@ export class SessionToolRegistry implements ToolExecutorLike { for (const [name, entry] of this.sessionTools) { merged.set(name, entry); } - return this.parent.executeTools(merged, toolCalls, scriptCallback, excludeTools); + return this.parent.executeTools(merged, toolCalls, scriptCallback, excludeTools, signal); } } diff --git a/src/app/service/agent/core/skill_script_executor.test.ts b/src/app/service/agent/core/skill_script_executor.test.ts index 033bae2f1..524caf4e8 100644 --- a/src/app/service/agent/core/skill_script_executor.test.ts +++ b/src/app/service/agent/core/skill_script_executor.test.ts @@ -500,7 +500,9 @@ describe("SkillScriptExecutor 超时处理", () => { let capturedUuid = ""; const sender = { sendMessage: vi.fn().mockImplementation((msg: any) => { - capturedUuid = msg.data.uuid; + if (msg.action === "offscreen/executeSkillScript") { + capturedUuid = msg.data.uuid; + } return new Promise(() => {}); }), } as any; @@ -559,4 +561,33 @@ describe("SkillScriptExecutor 超时处理", () => { await expect(execPromise).resolves.toBeDefined(); }); + + it("signal 中止时应停止脚本并返回 Aborted", async () => { + let capturedUuid = ""; + const sender = { + sendMessage: vi.fn().mockImplementation((msg: any) => { + if (msg.action === "offscreen/executeSkillScript") { + capturedUuid = msg.data.uuid; + return new Promise(() => {}); + } + if (msg.action === "offscreen/script/stopScript") { + return Promise.resolve({ data: undefined }); + } + return Promise.resolve({ data: "mock_result" }); + }), + } as any; + + const record = createRecord([], { name: "abort_tool" }); + const executor = new SkillScriptExecutor(record, sender); + const controller = new AbortController(); + + const execPromise = executor.execute({}, controller.signal); + controller.abort(); + + await expect(execPromise).rejects.toThrow("Aborted"); + expect(sender.sendMessage.mock.calls.some((call: any[]) => call[0].action === "offscreen/script/stopScript")).toBe( + true + ); + expect(getSkillScriptNameByUuid(capturedUuid)).toBe(""); + }); }); diff --git a/src/app/service/agent/core/skill_script_executor.ts b/src/app/service/agent/core/skill_script_executor.ts index ff8f332de..f9ee7ebb4 100644 --- a/src/app/service/agent/core/skill_script_executor.ts +++ b/src/app/service/agent/core/skill_script_executor.ts @@ -2,9 +2,9 @@ import type { MessageSend } from "@Packages/message/types"; import type { SkillScriptRecord, JsonValue } from "./types"; import type { ToolExecutor } from "./tool_registry"; import { getSkillScriptBody } from "@App/pkg/utils/skill_script"; -import { executeSkillScript } from "@App/app/service/offscreen/client"; +import { executeSkillScript, stopScript } from "@App/app/service/offscreen/client"; import { uuidv4 } from "@App/pkg/utils/uuid"; -import { withTimeout } from "@App/pkg/utils/with_timeout"; +import { createAbortError, throwIfAborted } from "./abort_utils"; // Skill Script UUID 前缀,用于在 GM API 请求中识别 Skill Script export const SKILL_SCRIPT_UUID_PREFIX = "skillscript-"; @@ -40,7 +40,9 @@ export class SkillScriptExecutor implements ToolExecutor { private configValues?: Record ) {} - async execute(args: Record): Promise { + async execute(args: Record, signal?: AbortSignal): Promise { + throwIfAborted(signal); + // 根据 @param 定义做基本的类型转换 const typedArgs: Record = {}; for (const param of this.record.params) { @@ -58,10 +60,6 @@ export class SkillScriptExecutor implements ToolExecutor { } } - // 在 service worker 端生成 UUID 并注册映射 - const uuid = SKILL_SCRIPT_UUID_PREFIX + uuidv4(); - skillScriptUuidMap.set(uuid, { name: this.record.name, grants: this.record.grants }); - // 加载 @require 资源内容 let requires: Array<{ url: string; content: string }> | undefined; if (this.record.requires?.length && this.requireLoader) { @@ -80,22 +78,66 @@ export class SkillScriptExecutor implements ToolExecutor { const code = getSkillScriptBody(this.record.code); const timeoutMs = this.record.timeout ? this.record.timeout * 1000 : SKILL_SCRIPT_DEFAULT_TIMEOUT_MS; const timeoutSec = timeoutMs / 1000; - try { - const execPromise = executeSkillScript(this.sender, { - uuid, - code, - args: typedArgs, - grants: this.record.grants, - name: this.record.name, - requires, - configValues: this.configValues, + throwIfAborted(signal); + + // 在 service worker 端生成 UUID 并注册映射 + const uuid = SKILL_SCRIPT_UUID_PREFIX + uuidv4(); + skillScriptUuidMap.set(uuid, { name: this.record.name, grants: this.record.grants }); + + const execPromise = executeSkillScript(this.sender, { + uuid, + code, + args: typedArgs, + grants: this.record.grants, + name: this.record.name, + requires, + configValues: this.configValues, + }); + + let timeoutId: ReturnType | undefined; + let abortCleanup = () => {}; + let stopped = false; + const stopExecution = async () => { + if (stopped) return; + stopped = true; + try { + await stopScript(this.sender, uuid); + } catch { + // 停止脚本失败不覆盖原始中止/超时错误 + } + }; + const timeoutPromise = new Promise((_, reject) => { + timeoutId = setTimeout(() => { + void stopExecution(); + reject( + Object.assign(new Error(`SkillScript "${this.record.name}" timed out after ${timeoutSec}s`), { + errorCode: "tool_timeout", + }) + ); + }, timeoutMs); + }); + const abortPromise = + signal && + new Promise((_, reject) => { + const onAbort = () => { + void stopExecution(); + reject(createAbortError()); + }; + abortCleanup = () => signal.removeEventListener("abort", onAbort); + signal.addEventListener("abort", onAbort, { once: true }); + if (signal.aborted) onAbort(); }); - return await withTimeout(execPromise, timeoutMs, () => - Object.assign(new Error(`SkillScript "${this.record.name}" timed out after ${timeoutSec}s`), { - errorCode: "tool_timeout", - }) - ); + + try { + const races: Promise[] = [execPromise, timeoutPromise]; + if (abortPromise) { + races.push(abortPromise); + } + const result = await Promise.race(races); + return result as JsonValue; } finally { + if (timeoutId) clearTimeout(timeoutId); + abortCleanup(); // 执行完毕后清理映射 skillScriptUuidMap.delete(uuid); } diff --git a/src/app/service/agent/core/sub_agent_types.test.ts b/src/app/service/agent/core/sub_agent_types.test.ts index 90aa80bc0..5bd1b71fa 100644 --- a/src/app/service/agent/core/sub_agent_types.test.ts +++ b/src/app/service/agent/core/sub_agent_types.test.ts @@ -100,7 +100,7 @@ describe("Sub-Agent 类型系统", () => { }); it.concurrent("allowedTools 和 excludeTools 都未指定时返回空数组", () => { - const config: any = { name: "empty", maxIterations: 10, timeoutMs: 60000, systemPromptAddition: "" }; + const config: any = { name: "empty", timeoutMs: 60000, systemPromptAddition: "" }; const excluded = getExcludeToolsForType(config, allTools); expect(excluded).toEqual([]); }); @@ -110,7 +110,6 @@ describe("Sub-Agent 类型系统", () => { name: "test", allowedTools: ["web_fetch"], excludeTools: ["web_search"], - maxIterations: 10, timeoutMs: 60000, systemPromptAddition: "", }; diff --git a/src/app/service/agent/core/sub_agent_types.ts b/src/app/service/agent/core/sub_agent_types.ts index aca7ff8f0..b597c0e22 100644 --- a/src/app/service/agent/core/sub_agent_types.ts +++ b/src/app/service/agent/core/sub_agent_types.ts @@ -5,7 +5,6 @@ export interface SubAgentTypeConfig { description: string; // 英文,写入 agent tool 描述供 LLM 选择 allowedTools?: string[]; // 白名单模式(优先于 excludeTools) excludeTools?: string[]; // 黑名单模式 - maxIterations: number; timeoutMs: number; systemPromptAddition: string; // 注入 sub-agent system prompt 的角色说明 } @@ -30,7 +29,6 @@ export const SUB_AGENT_TYPES: Record = { "opfs_list", "opfs_delete", ], - maxIterations: 20, timeoutMs: 600_000, systemPromptAddition: `## Role: Researcher @@ -65,7 +63,6 @@ You are a research-focused sub-agent. Your job is to search, fetch, read, and su "opfs_list", "opfs_delete", ], - maxIterations: 30, timeoutMs: 600_000, systemPromptAddition: `## Role: Page Operator @@ -86,7 +83,6 @@ You are a page interaction sub-agent. Your job is to navigate web pages, interac name: "general", description: "All tools, general-purpose", excludeTools: ["ask_user", "agent"], - maxIterations: 30, timeoutMs: 600_000, systemPromptAddition: `## Role: General Sub-Agent diff --git a/src/app/service/agent/core/task_scheduler.test.ts b/src/app/service/agent/core/task_scheduler.test.ts index a67c400ee..1fed88872 100644 --- a/src/app/service/agent/core/task_scheduler.test.ts +++ b/src/app/service/agent/core/task_scheduler.test.ts @@ -1,7 +1,7 @@ import { describe, expect, it, beforeEach, vi } from "vitest"; import { AgentTaskScheduler } from "./task_scheduler"; import { AgentTaskRepo, AgentTaskRunRepo } from "@App/app/repo/agent_task"; -import type { AgentTask } from "@App/app/service/agent/core/types"; +import type { AgentTask, InternalAgentTask } from "@App/app/service/agent/core/types"; // Mock OPFS 文件系统(AgentTaskRunRepo 使用 OPFS 存储) function createMockOPFS() { @@ -40,12 +40,13 @@ function createMockOPFS() { getDirectoryHandle: vi.fn(async (name: string, opts?: { create?: boolean }) => { if (!store.has("__dir__" + name)) { if (opts?.create) store.set("__dir__" + name, new Map()); - else throw new Error("Not found"); + else throw new DOMException("A requested file or directory could not be found.", "NotFoundError"); } return createMockDirHandle(store.get("__dir__" + name)); }), getFileHandle: vi.fn(async (name: string, opts?: { create?: boolean }) => { - if (!store.has(name) && !opts?.create) throw new Error("Not found"); + if (!store.has(name) && !opts?.create) + throw new DOMException("A requested file or directory could not be found.", "NotFoundError"); if (!store.has(name)) store.set(name, ""); return createMockFileHandle(name, store); }), @@ -84,13 +85,17 @@ describe("AgentTaskScheduler", () => { let runRepo: AgentTaskRunRepo; let internalExecutor: ReturnType< typeof vi.fn< - (task: AgentTask) => Promise<{ conversationId: string; usage?: { inputTokens: number; outputTokens: number } }> + ( + task: AgentTask, + signal?: AbortSignal + ) => Promise<{ conversationId: string; usage?: { inputTokens: number; outputTokens: number } }> > >; - let eventEmitter: ReturnType Promise>>; + let eventEmitter: ReturnType Promise>>; let scheduler: AgentTaskScheduler; - beforeEach(() => { + beforeEach(async () => { + await chrome.storage.local.clear(); createMockOPFS(); repo = new AgentTaskRepo(); runRepo = new AgentTaskRunRepo(); @@ -112,6 +117,25 @@ describe("AgentTaskScheduler", () => { expect(updated!.nextruntime).toBeGreaterThan(Date.now() - 1000); }); + it("init 应把上次 Service Worker 中断留下的 running 记录收敛为结果未知错误", async () => { + const task = await repo.saveTask(makeTask({ id: "interrupted-run", nextruntime: Date.now() + 60_000 })); + await runRepo.appendRun({ + id: `scheduled:${task.id}:${task.generation}:${Date.now() - 1_000}`, + taskId: task.id, + starttime: Date.now() - 5_000, + status: "running", + }); + + await scheduler.init(); + + expect((await runRepo.listRuns(task.id))[0]).toMatchObject({ + status: "error", + error: expect.stringContaining("interrupted"), + endtime: expect.any(Number), + }); + expect((await repo.getTask(task.id))!.lastRunStatus).toBe("error"); + }); + it("tick 执行到期任务", async () => { const task = makeTask({ id: "tick-1", nextruntime: Date.now() - 1000 }); await repo.saveTask(task); @@ -122,6 +146,7 @@ describe("AgentTaskScheduler", () => { await vi.waitFor(async () => { expect(internalExecutor).toHaveBeenCalledTimes(1); }); + await vi.waitFor(() => expect(scheduler.isRunning("tick-1")).toBe(false)); }); it("tick 跳过未到期任务", async () => { @@ -169,7 +194,10 @@ describe("AgentTaskScheduler", () => { await scheduler.executeTask(task); expect(internalExecutor).toHaveBeenCalledTimes(1); - expect(internalExecutor).toHaveBeenCalledWith(expect.objectContaining({ id: "internal-1" })); + expect(internalExecutor).toHaveBeenCalledWith( + expect.objectContaining({ id: "internal-1" }), + expect.any(AbortSignal) + ); expect(eventEmitter).not.toHaveBeenCalled(); // 检查 run 记录 @@ -212,6 +240,24 @@ describe("AgentTaskScheduler", () => { expect(updatedTask!.lastRunError).toBe("LLM 调用失败"); }); + it("执行失败时保留已累计的 usage", async () => { + internalExecutor.mockRejectedValue( + Object.assign(new Error("请求失败"), { + usage: { inputTokens: 120, outputTokens: 40 }, + conversationId: "conv-failed", + }) + ); + + const task = makeTask({ id: "error-usage-1", nextruntime: Date.now() - 1000 }); + await repo.saveTask(task); + await scheduler.executeTask(task); + + const runs = await runRepo.listRuns("error-usage-1"); + expect(runs[0].status).toBe("error"); + expect(runs[0].usage).toEqual({ inputTokens: 120, outputTokens: 40 }); + expect(runs[0].conversationId).toBe("conv-failed"); + }); + it("执行完成后更新 nextruntime", async () => { const task = makeTask({ id: "next-1", nextruntime: Date.now() - 1000 }); await repo.saveTask(task); @@ -243,4 +289,90 @@ describe("AgentTaskScheduler", () => { await scheduler.executeTask(task); expect(internalExecutor).toHaveBeenCalled(); }); + + it("调度运行记录无法持久化时不得提前消耗到期槽位", async () => { + const due = Date.now() - 1_000; + const task = await repo.saveTask(makeTask({ id: "append-before-claim", nextruntime: due })); + runRepo.appendRun = vi.fn().mockRejectedValue(new Error("run storage unavailable")); + + await expect(scheduler.executeTask(task, true, Date.now())).rejects.toThrow("run storage unavailable"); + + expect((await repo.getTask(task.id))!.nextruntime).toBe(due); + expect(internalExecutor).not.toHaveBeenCalled(); + }); + + it("取消活动任务应中止传给执行器的 signal 并释放运行占位", async () => { + let observedSignal: AbortSignal | undefined; + internalExecutor.mockImplementation( + async (_task, signal) => + new Promise((_resolve, reject) => { + observedSignal = signal; + signal?.addEventListener("abort", () => reject(new Error("Task cancelled")), { once: true }); + }) + ); + const task = await repo.saveTask(makeTask({ id: "cancel-active" })); + const execution = scheduler.executeTask(task); + await vi.waitFor(() => expect(internalExecutor).toHaveBeenCalledOnce()); + + scheduler.cancelTask(task.id); + await execution; + + expect(observedSignal?.aborted).toBe(true); + expect(scheduler.isRunning(task.id)).toBe(false); + expect((await runRepo.listRuns(task.id))[0]).toMatchObject({ status: "error", error: "Task cancelled" }); + }); + + it("运行期间的用户编辑应保留,完成时只合并运行遥测", async () => { + let finish!: () => void; + internalExecutor.mockReturnValue( + new Promise((resolve) => { + finish = () => resolve({ conversationId: "conv-edited" }); + }) + ); + const task = await repo.saveTask(makeTask({ id: "edited-during-run", name: "旧名称" })); + const execution = scheduler.executeTask(task); + await vi.waitFor(() => expect(internalExecutor).toHaveBeenCalledOnce()); + + const edited = (await repo.getTask(task.id))!; + edited.name = "用户的新名称"; + (edited as InternalAgentTask).prompt = "用户的新提示词"; + await repo.saveTask(edited); + finish(); + await execution; + + const stored = (await repo.getTask(task.id))!; + expect(stored.name).toBe("用户的新名称"); + expect((stored as InternalAgentTask).prompt).toBe("用户的新提示词"); + expect(stored.lastRunStatus).toBe("success"); + }); + + it("运行期间删除任务后完成回写不应复活旧 generation", async () => { + let finish!: () => void; + internalExecutor.mockReturnValue( + new Promise((resolve) => { + finish = () => resolve({ conversationId: "conv-deleted" }); + }) + ); + const task = await repo.saveTask(makeTask({ id: "deleted-during-run" })); + const execution = scheduler.executeTask(task); + await vi.waitFor(() => expect(internalExecutor).toHaveBeenCalledOnce()); + + const current = (await repo.getTask(task.id))!; + await repo.removeTask(current.id, current.generation, current.revision); + finish(); + await execution; + + expect(await repo.getTask(task.id)).toBeUndefined(); + }); + + it("已领取的调度槽即使运行记录更新失败也不应在下一次 tick 重复执行", async () => { + const task = await repo.saveTask(makeTask({ id: "claimed-slot", nextruntime: Date.now() - 1000 })); + runRepo.updateRun = vi.fn().mockRejectedValueOnce(new Error("run storage failed")); + + await expect(scheduler.executeTask(task, true)).rejects.toThrow("run storage failed"); + await scheduler.tick(); + await new Promise((resolve) => setTimeout(resolve, 0)); + + expect(internalExecutor).toHaveBeenCalledTimes(1); + }); }); diff --git a/src/app/service/agent/core/task_scheduler.ts b/src/app/service/agent/core/task_scheduler.ts index 280354ac5..a55f6349b 100644 --- a/src/app/service/agent/core/task_scheduler.ts +++ b/src/app/service/agent/core/task_scheduler.ts @@ -2,16 +2,21 @@ import type { AgentTask, AgentTaskRun, EventAgentTask, InternalAgentTask } from import type { AgentTaskRepo, AgentTaskRunRepo } from "@App/app/repo/agent_task"; import { nextTimeInfo } from "@App/pkg/utils/cron"; import { uuidv4 } from "@App/pkg/utils/uuid"; +import { isRevisionConflict } from "@App/app/repo/revision"; -export type InternalExecutor = (task: InternalAgentTask) => Promise<{ +export type InternalExecutor = ( + task: InternalAgentTask, + signal: AbortSignal +) => Promise<{ conversationId: string; usage?: { inputTokens: number; outputTokens: number }; }>; -export type EventEmitter = (task: EventAgentTask) => Promise; +export type EventEmitter = (task: EventAgentTask, signal: AbortSignal) => Promise; export class AgentTaskScheduler { - private runningTasks = new Set(); + private runningTasks = new Map(); + private initialization?: Promise; constructor( private repo: AgentTaskRepo, @@ -21,23 +26,60 @@ export class AgentTaskScheduler { ) {} async init(): Promise { + if (!this.initialization) { + this.initialization = this.initialize().catch((error) => { + this.initialization = undefined; + throw error; + }); + } + return this.initialization; + } + + private async initialize(): Promise { // 加载所有 enabled 任务,计算 nextruntime const tasks = await this.repo.listTasks(); for (const task of tasks) { + let currentTask = task; if (task.enabled && !task.nextruntime) { + let nextRuntime: number | undefined; try { - const info = nextTimeInfo(task.crontab); - task.nextruntime = info.next.toMillis(); - task.updatetime = Date.now(); - await this.repo.saveTask(task); + nextRuntime = nextTimeInfo(task.crontab).next.toMillis(); } catch { - // cron 表达式无效,跳过 + // cron 表达式无效,保留未调度状态 + } + if (nextRuntime !== undefined) { + task.nextruntime = nextRuntime; + task.updatetime = Date.now(); + currentTask = await this.repo.saveTask(task); } } + + // A Service Worker restart can terminate JavaScript after the durable claim advanced nextruntime but + // before the executor outcome was recorded. Replaying may duplicate irreversible side effects, so use an + // explicit at-most-once recovery policy: close stale running telemetry as outcome-unknown and keep the + // already advanced schedule. + const interruptedRuns = (await this.runRepo.listRuns(currentTask.id, 500)).filter( + (run) => run.status === "running" + ); + if (interruptedRuns.length > 0) { + const endtime = Date.now(); + const error = "Task execution interrupted by service restart; outcome unknown"; + await Promise.all( + interruptedRuns.map((run) => + this.runRepo.updateRun(currentTask.id, run.id, { status: "error", error, endtime }) + ) + ); + await this.repo.updateRunState(currentTask.id, currentTask.generation!, { + lastruntime: Math.max(...interruptedRuns.map((run) => run.starttime)), + lastRunStatus: "error", + lastRunError: error, + }); + } } } async tick(now?: number): Promise { + await this.init(); const currentTime = now ?? Date.now(); const tasks = await this.repo.listTasks(); @@ -47,74 +89,97 @@ export class AgentTaskScheduler { if (!task.nextruntime || task.nextruntime > currentTime) continue; // 不 await,并行执行多个任务 - this.executeTask(task).catch(() => { + this.executeTask(task, true, currentTime).catch(() => { // 错误已在 executeTask 内部处理 }); } } - async executeTask(task: AgentTask): Promise { + async executeTask(task: AgentTask, claimScheduled = false, now = Date.now()): Promise { if (this.runningTasks.has(task.id)) return; - this.runningTasks.add(task.id); + const abortController = new AbortController(); + this.runningTasks.set(task.id, abortController); try { + await this.init(); const run: AgentTaskRun = { - id: uuidv4(), + id: + claimScheduled && task.nextruntime ? `scheduled:${task.id}:${task.generation}:${task.nextruntime}` : uuidv4(), taskId: task.id, starttime: Date.now(), status: "running", }; - await this.runRepo.appendRun(run); + // Record the slot before advancing nextruntime. A deterministic scheduled-run ID makes a restart between + // these two durable writes idempotent instead of either losing the slot or creating duplicate run records. + const createdRun = (await this.runRepo.appendRun(run)) !== false; + let executionTask = task; + if (claimScheduled) { + const claimed = await this.repo.claimDueTask(task.id, task.generation!, now); + if (!claimed) { + if (createdRun) await this.runRepo.removeRun(task.id, run.id); + return; + } + executionTask = claimed; + } try { - if (task.mode === "internal") { - const result = await this.internalExecutor(task); + if (abortController.signal.aborted) throw new Error("Task cancelled"); + if (executionTask.mode === "internal") { + const result = await this.internalExecutor(executionTask, abortController.signal); run.conversationId = result.conversationId; run.usage = result.usage; } else { - await this.eventEmitter(task); + await this.eventEmitter(executionTask, abortController.signal); } run.status = "success"; run.endtime = Date.now(); - - task.lastRunStatus = "success"; - task.lastRunError = undefined; } catch (e: any) { run.status = "error"; run.error = e.message || "Unknown error"; + if (e.usage) run.usage = e.usage; + if (e.conversationId) run.conversationId = e.conversationId; run.endtime = Date.now(); + } - task.lastRunStatus = "error"; - task.lastRunError = run.error; - } finally { - // 更新 run 记录 - await this.runRepo.updateRun(task.id, run.id, { - status: run.status, - endtime: run.endtime, - error: run.error, - conversationId: run.conversationId, - usage: run.usage, - }); + await this.runRepo.updateRun(executionTask.id, run.id, { + status: run.status, + endtime: run.endtime, + error: run.error, + conversationId: run.conversationId, + usage: run.usage, + }); - // 更新 task 状态 - task.lastruntime = run.starttime; - try { - const info = nextTimeInfo(task.crontab); - task.nextruntime = info.next.toMillis(); - } catch { - task.nextruntime = undefined; - } - task.updatetime = Date.now(); - await this.repo.saveTask(task); + // 只合并运行遥测,不保存启动时的旧配置快照;运行期间的编辑/禁用必须保留。 + try { + await this.repo.updateRunState( + executionTask.id, + executionTask.generation!, + { + lastruntime: run.starttime, + lastRunStatus: run.status === "success" ? "success" : "error", + lastRunError: run.error, + }, + !claimScheduled + ); + } catch (error) { + // 删除或重建发生在运行期间时,旧 generation 的完成回写应被丢弃,而不是复活任务。 + if (!isRevisionConflict(error)) throw error; } } finally { // 必须在最外层 finally 确保任何异常都能清理 runningTasks - this.runningTasks.delete(task.id); + if (this.runningTasks.get(task.id) === abortController) this.runningTasks.delete(task.id); } } isRunning(taskId: string): boolean { return this.runningTasks.has(taskId); } + + cancelTask(taskId: string): boolean { + const abortController = this.runningTasks.get(taskId); + if (!abortController) return false; + abortController.abort(); + return true; + } } diff --git a/src/app/service/agent/core/tool_registry.test.ts b/src/app/service/agent/core/tool_registry.test.ts index d9d5992c9..87a148065 100644 --- a/src/app/service/agent/core/tool_registry.test.ts +++ b/src/app/service/agent/core/tool_registry.test.ts @@ -5,7 +5,7 @@ import type { ToolCall, ToolDefinition, ToolResultWithAttachments } from "./type import type { AgentChatRepo } from "@App/app/repo/agent_chat"; // 创建一个简单的 mock executor -function createExecutor(fn: (args: Record) => Promise): ToolExecutor { +function createExecutor(fn: (args: Record, signal?: AbortSignal) => Promise): ToolExecutor { return { execute: fn }; } @@ -105,6 +105,30 @@ describe("ToolRegistry", () => { expect(results[0].result).toBe('{"temp":25,"unit":"C"}'); }); + it("内置工具返回 undefined 或不可序列化值时应返回已定义的字符串", async () => { + const registry = new ToolRegistry(); + const cyclic: { self?: unknown } = {}; + cyclic.self = cyclic; + registry.registerBuiltin( + weatherDef, + createExecutor(async () => undefined) + ); + registry.registerBuiltin( + calcDef, + createExecutor(async () => cyclic) + ); + + const results = await registry.execute([ + { id: "tc_1", name: "get_weather", arguments: "{}" }, + { id: "tc_2", name: "calc", arguments: "{}" }, + ]); + + expect(results).toEqual([ + { id: "tc_1", result: "null" }, + { id: "tc_2", result: "null" }, + ]); + }); + it("空 arguments 时应传入空对象", async () => { const registry = new ToolRegistry(); const executeSpy = vi.fn().mockResolvedValue("ok"); @@ -112,7 +136,17 @@ describe("ToolRegistry", () => { await registry.execute([{ id: "tc_1", name: "get_weather", arguments: "" }]); - expect(executeSpy).toHaveBeenCalledWith({}); + expect(executeSpy).toHaveBeenCalledWith({}, undefined, "tc_1"); + }); + + it("应把 tool call 的 id 作为 toolCallId 传给 executor(供 agent 等需要区分并发调用的工具使用)", async () => { + const registry = new ToolRegistry(); + const executeSpy = vi.fn().mockResolvedValue("ok"); + registry.registerBuiltin(weatherDef, { execute: executeSpy }); + + await registry.execute([{ id: "tc_unique_42", name: "get_weather", arguments: "{}" }]); + + expect(executeSpy).toHaveBeenCalledWith({}, undefined, "tc_unique_42"); }); it("内置工具抛出异常时应返回错误信息", async () => { @@ -162,7 +196,7 @@ describe("ToolRegistry", () => { const results = await registry.execute([{ id: "tc_1", name: "unknown_tool", arguments: "{}" }], scriptCallback); - expect(scriptCallback).toHaveBeenCalledWith([{ id: "tc_1", name: "unknown_tool", arguments: "{}" }]); + expect(scriptCallback).toHaveBeenCalledWith([{ id: "tc_1", name: "unknown_tool", arguments: "{}" }], undefined); expect(results[0].result).toBe("script result"); }); @@ -205,7 +239,7 @@ describe("ToolRegistry", () => { // 脚本工具结果 expect(results.find((r) => r.id === "tc_2")?.result).toBe("script_result"); // scriptCallback 只收到脚本工具 - expect(scriptCallback).toHaveBeenCalledWith([toolCalls[1]]); + expect(scriptCallback).toHaveBeenCalledWith([toolCalls[1]], undefined); }); it("空 toolCalls 数组时应返回空结果", async () => { @@ -214,6 +248,53 @@ describe("ToolRegistry", () => { expect(results).toHaveLength(0); }); + it("应将 signal 传递给内置工具执行器", async () => { + const registry = new ToolRegistry(); + const executeSpy = vi.fn().mockResolvedValue("ok"); + registry.registerBuiltin(weatherDef, { execute: executeSpy }); + const controller = new AbortController(); + + await registry.execute( + [{ id: "tc_1", name: "get_weather", arguments: "{}" }], + undefined, + undefined, + controller.signal + ); + + expect(executeSpy).toHaveBeenCalledWith({}, controller.signal, "tc_1"); + }); + + it("取消后应等待内置工具抵达提交边界并保留其已提交成功结果", async () => { + const registry = new ToolRegistry(); + let finishExecution!: () => void; + const executeSpy = vi.fn().mockReturnValue( + new Promise((resolve) => { + finishExecution = () => resolve("late success"); + }) + ); + registry.registerBuiltin(weatherDef, { execute: executeSpy }); + const controller = new AbortController(); + + const resultPromise = registry.execute( + [{ id: "tc_1", name: "get_weather", arguments: "{}" }], + undefined, + undefined, + controller.signal + ); + controller.abort(); + + let settled = false; + void resultPromise.finally(() => { + settled = true; + }); + await Promise.resolve(); + expect(settled).toBe(false); + + finishExecution(); + await expect(resultPromise).resolves.toEqual([{ id: "tc_1", result: "late success" }]); + expect(executeSpy).toHaveBeenCalledWith({}, controller.signal, "tc_1"); + }); + it("JSON 解析失败时应返回错误", async () => { const registry = new ToolRegistry(); const executor = createExecutor(async () => "ok"); @@ -358,6 +439,63 @@ describe("ToolRegistry", () => { expect(mockRepo.saveAttachment).toHaveBeenCalledWith(expect.any(String), "data:image/jpeg;base64,/9j/abc"); }); + it("MCP 结构化结果经过附件保存后不应丢失", async () => { + const registry = new ToolRegistry(); + const mockRepo = createMockChatRepo(); + registry.setChatRepo(mockRepo); + registry.registerBuiltin( + weatherDef, + createExecutor(async () => ({ + content: "Screenshot captured.", + attachments: [], + structuredContent: { caption: "chart" }, + })) + ); + + const results = await registry.execute([{ id: "tc-structured", name: "get_weather", arguments: "{}" }]); + + expect(results[0].structuredContent).toEqual({ caption: "chart" }); + }); + + it("附件写入期间被取消时,应回收本批已保存的附件并返回错误结果", async () => { + const registry = new ToolRegistry(); + const mockRepo = createMockChatRepo(); + registry.setChatRepo(mockRepo); + const controller = new AbortController(); + + const savedIds: string[] = []; + // 第二个附件写入完成的同时 Stop 到达:写入已提交,但结果不能再按成功上报 + vi.mocked(mockRepo.saveAttachment).mockImplementation(async (id: string) => { + savedIds.push(id); + if (savedIds.length === 2) controller.abort(); + return 1024; + }); + + const structuredResult: ToolResultWithAttachments = { + content: "Files generated.", + attachments: [ + { type: "image", name: "a.png", mimeType: "image/png", data: "data:image/png;base64,a" }, + { type: "image", name: "b.png", mimeType: "image/png", data: "data:image/png;base64,b" }, + ], + }; + const executor = createExecutor(async () => structuredResult); + registry.registerBuiltin(weatherDef, executor); + + const results = await registry.execute( + [{ id: "tc_1", name: "get_weather", arguments: "{}" }], + null, + undefined, + controller.signal + ); + + expect(results[0].error).toBe(true); + expect(results[0].attachments).toBeUndefined(); + // 本批两个已落盘的附件都必须被回收,不能只删最后一个 + for (const id of savedIds) { + expect(mockRepo.deleteAttachment).toHaveBeenCalledWith(id); + } + }); + it("内置工具返回 ToolResultWithAttachments 含多个附件时应全部保存", async () => { const registry = new ToolRegistry(); const mockRepo = createMockChatRepo(); @@ -386,6 +524,118 @@ describe("ToolRegistry", () => { expect(mockRepo.saveAttachment).toHaveBeenCalledTimes(2); }); + it("保存附件期间 signal abort:应停止继续写入并把该 toolCall 标为 error", async () => { + const registry = new ToolRegistry(); + const mockRepo = createMockChatRepo(); + const controller = new AbortController(); + // 第一个附件保存"成功"的同时触发 abort,模拟 Stop 恰好落在多附件保存期间 + (mockRepo.saveAttachment as any).mockImplementationOnce(async () => { + controller.abort(); + return 1024; + }); + registry.setChatRepo(mockRepo); + + const structuredResult: ToolResultWithAttachments = { + content: "Files generated.", + attachments: [ + { type: "image", name: "img1.png", mimeType: "image/png", data: "data:image/png;base64,abc" }, + { type: "file", name: "report.xlsx", mimeType: "application/octet-stream", data: "base64data" }, + ], + }; + const executor = createExecutor(async () => structuredResult); + registry.registerBuiltin(weatherDef, executor); + + const results = await registry.execute( + [{ id: "tc_1", name: "get_weather", arguments: "{}" }], + undefined, + undefined, + controller.signal + ); + + // 只应写入第一个附件,第二个在 abort 后不再保存 + expect(mockRepo.saveAttachment).toHaveBeenCalledTimes(1); + expect(results[0].error).toBe(true); + }); + + it("子代理在取消边界返回时应保留已生成附件的所有权以供提交或回收", async () => { + const registry = new ToolRegistry(); + registry.setChatRepo(createMockChatRepo()); + const controller = new AbortController(); + registry.registerBuiltin(weatherDef, { + execute: async () => { + controller.abort(); + return { + content: "partial", + attachments: [ + { + attachmentId: "sub-image.png", + type: "image" as const, + name: "sub-image.png", + mimeType: "image/png", + }, + ], + ownedAttachmentIds: ["sub-image.png"], + subAgentDetails: { agentId: "child", description: "child", messages: [] }, + usage: { inputTokens: 8, outputTokens: 2 }, + }; + }, + }); + + const [result] = await registry.execute( + [{ id: "tc_1", name: "get_weather", arguments: "{}" }], + undefined, + undefined, + controller.signal + ); + + expect(result).toMatchObject({ + attachments: [{ id: "sub-image.png" }], + ownedAttachmentIds: ["sub-image.png"], + subAgentDetails: { agentId: "child" }, + usage: { inputTokens: 8, outputTokens: 2 }, + }); + expect(result.error).toBeUndefined(); + }); + + it("混合批次在脚本回调期间取消时应保留已完成内置工具结果", async () => { + const registry = new ToolRegistry(); + registry.setChatRepo(createMockChatRepo()); + registry.registerBuiltin(weatherDef, { + execute: async () => ({ + content: "builtin complete", + attachments: [ + { + attachmentId: "builtin-owned.png", + type: "image" as const, + name: "builtin-owned.png", + mimeType: "image/png", + }, + ], + ownedAttachmentIds: ["builtin-owned.png"], + usage: { inputTokens: 3, outputTokens: 1 }, + }), + }); + const controller = new AbortController(); + const scriptCallback = vi.fn(() => new Promise>(() => {})); + + const pending = registry.execute( + [ + { id: "builtin", name: "get_weather", arguments: "{}" }, + { id: "script", name: "script_tool", arguments: "{}" }, + ], + scriptCallback, + undefined, + controller.signal + ); + await vi.waitFor(() => expect(scriptCallback).toHaveBeenCalledOnce()); + controller.abort(); + + await expect(pending).resolves.toEqual([ + expect.objectContaining({ id: "builtin", ownedAttachmentIds: ["builtin-owned.png"] }), + expect.objectContaining({ id: "script", error: true }), + ]); + }); + it("内置工具返回 Blob 附件时应正确保存", async () => { const registry = new ToolRegistry(); const mockRepo = createMockChatRepo(); @@ -475,6 +725,24 @@ describe("ToolRegistry", () => { expect(mockRepo.saveAttachment).toHaveBeenCalledTimes(1); }); + it("脚本工具附件写入报错时,应回收可能已提交但无法确认的附件", async () => { + const registry = new ToolRegistry(); + const mockRepo = createMockChatRepo(); + vi.mocked(mockRepo.saveAttachment).mockRejectedValue(new Error("ambiguous close failure")); + registry.setChatRepo(mockRepo); + const structuredResult: ToolResultWithAttachments = { + content: "File generated.", + attachments: [{ type: "file", name: "output.zip", mimeType: "application/zip", data: "base64zipdata" }], + }; + const scriptCallback = vi.fn().mockResolvedValue([{ id: "tc_1", result: JSON.stringify(structuredResult) }]); + + const results = await registry.execute([{ id: "tc_1", name: "script_tool", arguments: "{}" }], scriptCallback); + + const attemptedId = vi.mocked(mockRepo.saveAttachment).mock.calls[0][0]; + expect(results[0].error).toBe(true); + expect(mockRepo.deleteAttachment).toHaveBeenCalledWith(attemptedId); + }); + it("脚本工具返回普通 JSON 时不应产生附件", async () => { const registry = new ToolRegistry(); const mockRepo = createMockChatRepo(); diff --git a/src/app/service/agent/core/tool_registry.ts b/src/app/service/agent/core/tool_registry.ts index 72ff81b15..7eed84349 100644 --- a/src/app/service/agent/core/tool_registry.ts +++ b/src/app/service/agent/core/tool_registry.ts @@ -1,11 +1,21 @@ -import type { Attachment, SubAgentDetails, ToolCall, ToolDefinition, ToolResultWithAttachments } from "./types"; +import type { + Attachment, + SubAgentDetails, + TokenUsage, + ToolCall, + ToolDefinition, + ToolResultWithAttachments, +} from "./types"; import type { AgentChatRepo } from "@App/app/repo/agent_chat"; import { uuidv4 } from "@App/pkg/utils/uuid"; import { getExtFromMime } from "./content_utils"; +import { raceWithAbort, throwIfAborted } from "./abort_utils"; // 工具执行器接口 export interface ToolExecutor { - execute(args: Record): Promise; + // toolCallId: 本次调用的 tool_call id(如 LLM 未提供则为 undefined)。 + // 供需要区分并发调用的工具(如 agent 子代理)关联自身产生的事件与结果。 + execute(args: Record, signal?: AbortSignal, toolCallId?: string): Promise; } // 工具来源分类 @@ -24,14 +34,22 @@ export interface ToolEntry { } // 脚本工具回调类型:将 tool calls 发送到 Sandbox 执行 -export type ScriptToolCallback = (toolCalls: ToolCall[]) => Promise>; +export type ScriptToolCallback = ( + toolCalls: ToolCall[], + signal?: AbortSignal +) => Promise>; // 工具执行结果(可能含附件和子代理详情) export type ToolExecuteResult = { id: string; result: string; + structuredContent?: unknown; + error?: boolean; attachments?: Attachment[]; subAgentDetails?: SubAgentDetails; + /** Attachment files created by this execution and safe to release if the owning round cannot commit. */ + ownedAttachmentIds?: string[]; + usage?: TokenUsage; }; // 可执行工具的最小接口,供 ToolLoopOrchestrator / SubAgentService 按接口接收 @@ -41,7 +59,8 @@ export interface ToolExecutorLike { execute( toolCalls: ToolCall[], scriptCallback?: ScriptToolCallback | null, - excludeTools?: Set + excludeTools?: Set, + signal?: AbortSignal ): Promise; } @@ -52,6 +71,15 @@ function extractErrorMessage(e: unknown): string { return String(e) || "Tool execution failed"; } +function normalizeToolResult(value: unknown): string { + if (typeof value === "string") return value; + try { + return JSON.stringify(value) ?? "null"; + } catch { + return "null"; + } +} + // 判断返回值是否是带附件的结构化结果 function isToolResultWithAttachments(value: unknown): value is ToolResultWithAttachments { if (typeof value !== "object" || value === null) return false; @@ -59,11 +87,43 @@ function isToolResultWithAttachments(value: unknown): value is ToolResultWithAtt return typeof obj.content === "string" && Array.isArray(obj.attachments); } -// 判断返回值是否包含子代理详情 -function isToolResultWithSubAgent(value: unknown): value is { content: string; subAgentDetails: SubAgentDetails } { +function isStructuredToolResult(value: unknown): value is { + content: string; + attachments?: ToolResultWithAttachments["attachments"]; + structuredContent?: unknown; + subAgentDetails?: SubAgentDetails; + ownedAttachmentIds?: string[]; + usage?: TokenUsage; +} { if (typeof value !== "object" || value === null) return false; const obj = value as Record; - return typeof obj.content === "string" && typeof obj.subAgentDetails === "object" && obj.subAgentDetails !== null; + return ( + typeof obj.content === "string" && + (Array.isArray(obj.attachments) || + obj.structuredContent !== undefined || + (typeof obj.subAgentDetails === "object" && obj.subAgentDetails !== null) || + typeof obj.usage === "object") + ); +} + +function persistedAttachmentReferences( + attachments?: ToolResultWithAttachments["attachments"] +): Attachment[] | undefined { + if (!attachments) return undefined; + const references = attachments.flatMap((attachment) => + attachment.data == null && attachment.attachmentId + ? [ + { + id: attachment.attachmentId, + type: attachment.type, + name: attachment.name, + mimeType: attachment.mimeType, + size: attachment.size, + }, + ] + : [] + ); + return references.length > 0 ? references : undefined; } // 工具注册表,管理内置工具和脚本工具的统一执行 @@ -171,9 +231,10 @@ export class ToolRegistry implements ToolExecutorLike { async execute( toolCalls: ToolCall[], scriptCallback?: ScriptToolCallback | null, - excludeTools?: Set + excludeTools?: Set, + signal?: AbortSignal ): Promise { - return this.executeTools(this.tools, toolCalls, scriptCallback, excludeTools); + return this.executeTools(this.tools, toolCalls, scriptCallback, excludeTools, signal); } // 执行工具调用(接收外部 tools Map),供 SessionToolRegistry 复用附件保存等共享逻辑 @@ -182,7 +243,8 @@ export class ToolRegistry implements ToolExecutorLike { tools: ReadonlyMap, toolCalls: ToolCall[], scriptCallback?: ScriptToolCallback | null, - excludeTools?: Set + excludeTools?: Set, + signal?: AbortSignal ): Promise { const results: ToolExecuteResult[] = []; const builtinCalls: ToolCall[] = []; @@ -194,6 +256,7 @@ export class ToolRegistry implements ToolExecutorLike { results.push({ id: tc.id, result: JSON.stringify({ error: `Tool "${tc.name}" is not available in this context` }), + error: true, }); continue; } @@ -209,46 +272,105 @@ export class ToolRegistry implements ToolExecutorLike { builtinCalls.map(async (tc): Promise => { const tool = tools.get(tc.name)!; try { + throwIfAborted(signal); let args: Record = {}; if (tc.arguments) { args = JSON.parse(tc.arguments); } - const rawResult = await tool.executor.execute(args); - - // 检查是否带附件或子代理详情 - if (isToolResultWithAttachments(rawResult)) { - const attachments = await this.saveAttachments(rawResult.attachments); - return { id: tc.id, result: rawResult.content, attachments }; - } else if (isToolResultWithSubAgent(rawResult)) { - return { id: tc.id, result: rawResult.content, subAgentDetails: rawResult.subAgentDetails }; - } else { - return { id: tc.id, result: typeof rawResult === "string" ? rawResult : JSON.stringify(rawResult) }; + // Registered executors receive the signal and define their own commit boundary. Abandoning the promise + // with raceWithAbort can report failure while a non-cancellable storage close commits in the background. + const rawResult = await tool.executor.execute(args, signal, tc.id); + + // 附件与子代理详情可以同时存在,统一保留两类元数据。 + if (isStructuredToolResult(rawResult)) { + let saved: { attachments: Attachment[]; ownedAttachmentIds: string[] }; + try { + saved = rawResult.attachments + ? await this.saveAttachments(rawResult.attachments, signal) + : { attachments: [], ownedAttachmentIds: [] }; + } catch (error) { + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { + subAgentDetails: rawResult.subAgentDetails, + attachments: persistedAttachmentReferences(rawResult.attachments), + ownedAttachmentIds: rawResult.ownedAttachmentIds, + usage: rawResult.usage, + }); + } + const structured = { + id: tc.id, + result: rawResult.content, + attachments: saved.attachments, + ...(rawResult.structuredContent !== undefined ? { structuredContent: rawResult.structuredContent } : {}), + subAgentDetails: rawResult.subAgentDetails, + ownedAttachmentIds: [...saved.ownedAttachmentIds, ...(rawResult.ownedAttachmentIds || [])], + usage: rawResult.usage, + }; + return structured; } + return { id: tc.id, result: normalizeToolResult(rawResult) }; } catch (e: any) { console.error(`[ToolRegistry] tool "${tc.name}" execution failed:`, e); - return { id: tc.id, result: JSON.stringify({ error: extractErrorMessage(e) }) }; + return { + id: tc.id, + result: JSON.stringify({ error: extractErrorMessage(e) }), + error: true, + subAgentDetails: e.subAgentDetails, + attachments: e.attachments, + ownedAttachmentIds: e.ownedAttachmentIds, + usage: e.usage, + }; } }) ); results.push(...builtinResults); + if (signal?.aborted) { + return results; + } + // 执行脚本工具 if (scriptCalls.length > 0) { if (scriptCallback) { - const scriptResults = await scriptCallback(scriptCalls); + let scriptResults: Array<{ id: string; result: string; error?: boolean }>; + try { + scriptResults = await raceWithAbort(scriptCallback(scriptCalls, signal), signal); + } catch (error) { + if (!signal?.aborted) throw error; + // Builtin calls in the same batch may already own durable attachments. Preserve those completed + // results and synthesize only the interrupted script results so the caller can commit or reclaim all + // ownership metadata deterministically. + scriptResults = scriptCalls.map((toolCall) => ({ + id: toolCall.id, + result: JSON.stringify({ error: "Tool execution cancelled" }), + error: true, + })); + } // 脚本工具也可能返回带附件的结构化结果 for (const sr of scriptResults) { + let parsed: unknown; try { - const parsed = JSON.parse(sr.result); - if (isToolResultWithAttachments(parsed)) { - const attachments = await this.saveAttachments(parsed.attachments); - results.push({ id: sr.id, result: parsed.content, attachments }); - continue; - } + parsed = JSON.parse(sr.result); } catch { - // 不是 JSON 或不是结构化结果,按原始字符串处理 + // 不是 JSON,按原始字符串处理 + results.push({ id: sr.id, result: sr.result, error: sr.error }); + continue; + } + if (isToolResultWithAttachments(parsed)) { + // saveAttachments 失败(含 abort 及 OPFS 写入错误等其他原因)必须落为 error 结果, + // 不能落回“按原始字符串处理”分支——那会让脚本工具声明产出的附件在实际未写入时仍被当作已完成上报 + try { + const saved = await this.saveAttachments(parsed.attachments, signal); + results.push({ id: sr.id, result: parsed.content, ...saved, error: sr.error }); + } catch (e: any) { + results.push({ + id: sr.id, + result: JSON.stringify({ error: extractErrorMessage(e) }), + error: true, + }); + } + continue; } - results.push({ id: sr.id, result: sr.result }); + results.push({ id: sr.id, result: sr.result, error: sr.error }); } } else { // 没有脚本回调,返回错误并列出可用工具名,引导 LLM 自我纠正 @@ -262,6 +384,7 @@ export class ToolRegistry implements ToolExecutorLike { result: JSON.stringify({ error: `Tool "${tc.name}" not found. Available tools: [${availableNames.join(", ")}].${hint}`, }), + error: true, }); } } @@ -270,36 +393,54 @@ export class ToolRegistry implements ToolExecutorLike { return results; } - // 保存附件数据到 OPFS,返回 Attachment 元数据 - private async saveAttachments(attachmentDataList: ToolResultWithAttachments["attachments"]): Promise { - if (!this.chatRepo || attachmentDataList.length === 0) return []; + // 保存附件数据到 OPFS,返回 Attachment 元数据。 + // 传入 signal 时在每个附件写入前检查,abort 时中止剩余写入并抛错——调用方的 catch 块会把这 + // 转成该 toolCall 的 error 结果,避免 Stop 之后仍继续写多个文件。 + private async saveAttachments( + attachmentDataList: ToolResultWithAttachments["attachments"], + signal?: AbortSignal + ): Promise<{ attachments: Attachment[]; ownedAttachmentIds: string[] }> { + if (!this.chatRepo || attachmentDataList.length === 0) return { attachments: [], ownedAttachmentIds: [] }; const attachments: Attachment[] = []; - for (const ad of attachmentDataList) { - if (!ad.data) { - // 无 data 的附件是已保存的引用(如 skill script 返回的 imageBlock),直接透传元数据 - if ("attachmentId" in ad && (ad as any).attachmentId) { - attachments.push({ - id: (ad as any).attachmentId, - type: ad.type, - name: ad.name, - mimeType: ad.mimeType, - size: (ad as any).size, - }); + // 本批尝试写入的附件 id(不含无 data 的已保存引用):写入报错仍可能已经提交,必须整批回收, + // 否则该 toolCall 以 error 结果收场后,这些文件不再被任何消息引用。 + const savedIds: string[] = []; + try { + for (const ad of attachmentDataList) { + if (ad.data == null) { + // 无 data 的附件是已保存的引用(如 skill script 返回的 imageBlock),直接透传元数据 + if (ad.attachmentId) { + attachments.push({ + id: ad.attachmentId, + type: ad.type, + name: ad.name, + mimeType: ad.mimeType, + size: ad.size, + }); + } + continue; } - continue; + throwIfAborted(signal); + const ext = getExtFromMime(ad.mimeType); + const id = `${uuidv4()}.${ext}`; + savedIds.push(id); + const size = await this.chatRepo.saveAttachment(id, ad.data); + // 写入期间可能已被 Stop:不能把这次结果当作成功返回,进入 catch 统一回收 + throwIfAborted(signal); + attachments.push({ + id, + type: ad.type, + name: ad.name, + mimeType: ad.mimeType, + size, + }); } - const ext = getExtFromMime(ad.mimeType); - const id = `${uuidv4()}.${ext}`; - const size = await this.chatRepo.saveAttachment(id, ad.data); - attachments.push({ - id, - type: ad.type, - name: ad.name, - mimeType: ad.mimeType, - size, - }); + } catch (error) { + const repo = this.chatRepo; + await Promise.all(savedIds.map((id) => repo.deleteAttachment(id).catch(() => {}))); + throw error; } - return attachments; + return { attachments, ownedAttachmentIds: savedIds }; } } diff --git a/src/app/service/agent/core/tools/ask_user.test.ts b/src/app/service/agent/core/tools/ask_user.test.ts index bb0c1fd70..62bb733d9 100644 --- a/src/app/service/agent/core/tools/ask_user.test.ts +++ b/src/app/service/agent/core/tools/ask_user.test.ts @@ -27,6 +27,7 @@ describe("ask_user", () => { const result = await resultPromise; expect(JSON.parse(result as string)).toEqual({ answer: "Blue" }); expect(resolvers.size).toBe(0); + expect(events.filter((event) => event.type === "ask_user_resolved")).toHaveLength(1); }); it("should throw if question is missing", async () => { @@ -51,10 +52,25 @@ describe("ask_user", () => { const result = JSON.parse((await resultPromise) as string); expect(result).toEqual({ answer: null, reason: "timeout" }); expect(resolvers.size).toBe(0); + expect(sendEvent).toHaveBeenLastCalledWith(expect.objectContaining({ type: "ask_user_expired" })); vi.useRealTimers(); }); + it("should settle and emit expiration when aborted", async () => { + const controller = new AbortController(); + const sendEvent = vi.fn(); + const resolvers = new Map void>(); + const { executor } = createAskUserTool(sendEvent, resolvers, controller.signal); + + const resultPromise = executor.execute({ question: "Waiting..." }); + controller.abort(); + + expect(JSON.parse((await resultPromise) as string)).toEqual({ answer: null, reason: "aborted" }); + expect(resolvers.size).toBe(0); + expect(sendEvent).toHaveBeenLastCalledWith(expect.objectContaining({ type: "ask_user_expired" })); + }); + it("should generate unique ask IDs", async () => { const events: ChatStreamEvent[] = []; const sendEvent = (event: ChatStreamEvent) => events.push(event); diff --git a/src/app/service/agent/core/tools/ask_user.ts b/src/app/service/agent/core/tools/ask_user.ts index a473dd48b..8f5000d57 100644 --- a/src/app/service/agent/core/tools/ask_user.ts +++ b/src/app/service/agent/core/tools/ask_user.ts @@ -31,7 +31,8 @@ const ASK_USER_TIMEOUT_MS = 5 * 60 * 1000; export function createAskUserTool( sendEvent: (event: ChatStreamEvent) => void, - resolvers: Map void> + resolvers: Map void>, + signal: AbortSignal = new AbortController().signal ): { definition: ToolDefinition; executor: ToolExecutor } { let askCounter = 0; @@ -44,21 +45,34 @@ export function createAskUserTool( const askId = `ask_${Date.now()}_${++askCounter}`; + if (signal.aborted) return JSON.stringify({ answer: null, reason: "aborted" }); + // 通知 UI 显示提问 sendEvent({ type: "ask_user", id: askId, question, options, multiple }); // 等待用户回复 return new Promise((resolve) => { - const timer = setTimeout(() => { + let settled = false; + const finish = (result: string, event: ChatStreamEvent) => { + if (settled) return; + settled = true; + clearTimeout(timer); + signal.removeEventListener("abort", onAbort); resolvers.delete(askId); - resolve(JSON.stringify({ answer: null, reason: "timeout" })); + sendEvent(event); + resolve(result); + }; + const onAbort = () => { + finish(JSON.stringify({ answer: null, reason: "aborted" }), { type: "ask_user_expired", id: askId }); + }; + const timer = setTimeout(() => { + finish(JSON.stringify({ answer: null, reason: "timeout" }), { type: "ask_user_expired", id: askId }); }, ASK_USER_TIMEOUT_MS); - + signal.addEventListener("abort", onAbort, { once: true }); resolvers.set(askId, (answer: string) => { - clearTimeout(timer); - resolvers.delete(askId); - resolve(JSON.stringify({ answer })); + finish(JSON.stringify({ answer }), { type: "ask_user_resolved", id: askId }); }); + if (signal.aborted) onAbort(); }); }, }; diff --git a/src/app/service/agent/core/tools/execute_script.test.ts b/src/app/service/agent/core/tools/execute_script.test.ts index 49a4fcd75..fed061882 100644 --- a/src/app/service/agent/core/tools/execute_script.test.ts +++ b/src/app/service/agent/core/tools/execute_script.test.ts @@ -80,7 +80,7 @@ describe("execute_script 工具", () => { expect(parsed).toEqual({ result: { sum: 42 }, target: "sandbox" }); expect(parsed).not.toHaveProperty("tab_id"); - expect(mockExecuteInSandbox).toHaveBeenCalledWith("return 1+2"); + expect(mockExecuteInSandbox).toHaveBeenCalledWith("return 1+2", expect.any(AbortSignal)); }); it.concurrent("返回值为 undefined 时应转为 null", async () => { @@ -102,18 +102,73 @@ describe("execute_script 工具", () => { const { executor } = createExecuteScriptTool(deps); await expect(executor.execute({ code: "while(true){}", target: "page" })).rejects.toThrow( - "execute_script timed out after 0.05s" + "execute_script (target=page) timed out after 0.05s" ); }); it.concurrent("sandbox 模式超时应报错", async () => { - const mockExecuteInSandbox = vi.fn().mockReturnValue(new Promise(() => {})); + const onAbort = vi.fn(); + const mockExecuteInSandbox = vi.fn().mockImplementation((_code: string, signal?: AbortSignal) => { + signal?.addEventListener("abort", onAbort, { once: true }); + return new Promise(() => {}); + }); const deps = makeDeps({ executeInSandbox: mockExecuteInSandbox, timeoutMs: 50 }); const { executor } = createExecuteScriptTool(deps); await expect(executor.execute({ code: "while(true){}", target: "sandbox" })).rejects.toThrow( "execute_script timed out after 0.05s" ); + expect(onAbort).toHaveBeenCalledOnce(); + }); + + it.concurrent("signal 已中止时应立即中断并且不执行脚本", async () => { + const mockExecuteInSandbox = vi.fn().mockReturnValue(new Promise(() => {})); + const deps = makeDeps({ executeInSandbox: mockExecuteInSandbox }); + const { executor } = createExecuteScriptTool(deps); + const controller = new AbortController(); + controller.abort(); + + await expect(executor.execute({ code: "return 1", target: "sandbox" }, controller.signal)).rejects.toThrow( + "Aborted" + ); + expect(mockExecuteInSandbox).not.toHaveBeenCalled(); + }); + }); + + describe("返回值序列化", () => { + it.concurrent("page 模式应完整保留大型返回值", async () => { + const bigString = "x".repeat(50_000); + const mockExecuteInPage = vi.fn().mockResolvedValue({ result: bigString, tabId: 1 }); + const deps = makeDeps({ executeInPage: mockExecuteInPage }); + const { executor } = createExecuteScriptTool(deps); + + const result = await executor.execute({ code: "return bigString", target: "page" }); + const parsed = JSON.parse(result as string); + + expect(parsed).toEqual({ result: bigString, target: "page", tab_id: 1 }); + }); + + it.concurrent("sandbox 模式应完整保留大型结构化返回值", async () => { + const bigArray = Array.from({ length: 10_000 }, (_, i) => i); + const mockExecuteInSandbox = vi.fn().mockResolvedValue(bigArray); + const deps = makeDeps({ executeInSandbox: mockExecuteInSandbox }); + const { executor } = createExecuteScriptTool(deps); + + const result = await executor.execute({ code: "return bigArray", target: "sandbox" }); + const parsed = JSON.parse(result as string); + + expect(parsed).toEqual({ result: bigArray, target: "sandbox" }); + }); + + it.concurrent("正常大小的返回值保持原有 envelope", async () => { + const mockExecuteInPage = vi.fn().mockResolvedValue({ result: { count: 5 }, tabId: 1 }); + const deps = makeDeps({ executeInPage: mockExecuteInPage }); + const { executor } = createExecuteScriptTool(deps); + + const result = await executor.execute({ code: "return {count:5}", target: "page" }); + const parsed = JSON.parse(result as string); + + expect(parsed).toEqual({ result: { count: 5 }, target: "page", tab_id: 1 }); }); }); diff --git a/src/app/service/agent/core/tools/execute_script.ts b/src/app/service/agent/core/tools/execute_script.ts index 8091ef9f0..231654433 100644 --- a/src/app/service/agent/core/tools/execute_script.ts +++ b/src/app/service/agent/core/tools/execute_script.ts @@ -1,6 +1,7 @@ import type { ToolDefinition } from "@App/app/service/agent/core/types"; import type { ToolExecutor } from "@App/app/service/agent/core/tool_registry"; import { withTimeout } from "@App/pkg/utils/with_timeout"; +import { createAbortError, throwIfAborted } from "../abort_utils"; import { requireString } from "./param_utils"; export const EXECUTE_SCRIPT_DEFINITION: ToolDefinition = { @@ -8,7 +9,10 @@ export const EXECUTE_SCRIPT_DEFINITION: ToolDefinition = { description: "Execute JavaScript code. " + "target='page': run in a browser tab (MAIN world) with full DOM access, shares page's window/globals — can access page JS variables and call page functions. Cannot access extension blob URLs. " + - "target='sandbox': isolated computation environment, no DOM. " + + "chrome.scripting.executeScript has no cancellation API: on timeout/stop this tool stops WAITING and returns an error, " + + "but the injected page code keeps running to completion in the tab (it is not actually terminated). " + + "Avoid long-running or blocking code with target='page'. " + + "target='sandbox': isolated computation environment, no DOM, and IS genuinely cancelled on timeout/stop. " + "Use `return` to return a value. Timeout: 30 seconds.", parameters: { type: "object", @@ -30,9 +34,60 @@ export const EXECUTE_SCRIPT_DEFINITION: ToolDefinition = { const EXECUTE_SCRIPT_TIMEOUT_MS = 30_000; +function executeSandboxWithTimeout( + execute: (signal: AbortSignal) => Promise, + parentSignal: AbortSignal | undefined, + timeoutMs: number +): Promise { + const controller = new AbortController(); + const onParentAbort = () => controller.abort(); + parentSignal?.addEventListener("abort", onParentAbort, { once: true }); + if (parentSignal?.aborted) controller.abort(); + + return new Promise((resolve, reject) => { + let settled = false; + let timedOut = false; + const cleanup = () => { + clearTimeout(timer); + controller.signal.removeEventListener("abort", onAbort); + parentSignal?.removeEventListener("abort", onParentAbort); + }; + const finish = (callback: () => void) => { + if (settled) return; + settled = true; + cleanup(); + callback(); + }; + const timeoutError = () => new Error(`execute_script timed out after ${timeoutMs / 1000}s`); + const onAbort = () => finish(() => reject(timedOut ? timeoutError() : createAbortError())); + + controller.signal.addEventListener("abort", onAbort, { once: true }); + const timer = setTimeout(() => { + timedOut = true; + controller.abort(); + }, timeoutMs); + if (controller.signal.aborted) { + onAbort(); + return; + } + + let execution: Promise; + try { + execution = execute(controller.signal); + } catch (error) { + finish(() => reject(error)); + return; + } + execution.then( + (result) => finish(() => resolve(result)), + (error) => finish(() => reject(error)) + ); + }); +} + export type ExecuteScriptDeps = { executeInPage: (code: string, options?: { tabId?: number }) => Promise<{ result: unknown; tabId: number }>; - executeInSandbox: (code: string) => Promise; + executeInSandbox: (code: string, signal?: AbortSignal) => Promise; timeoutMs?: number; // 可选超时(ms),默认 30s,测试用 }; @@ -43,7 +98,8 @@ export function createExecuteScriptTool(deps: ExecuteScriptDeps): { const timeoutMs = deps.timeoutMs ?? EXECUTE_SCRIPT_TIMEOUT_MS; const executor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const code = requireString(args, "code"); const target = requireString(args, "target"); @@ -52,20 +108,28 @@ export function createExecuteScriptTool(deps: ExecuteScriptDeps): { } if (target === "page") { + // chrome.scripting.executeScript 无法被真正中止:withTimeout 只能让调用方停止等待, + // 注入到页面的代码仍会在 tab 内继续跑到自然结束。错误信息必须说明"调用方停止等待" + // 与"页面脚本已停止"是两回事,避免误导上层以为页面副作用已经终止。 const tabId = args.tab_id as number | undefined; const { result, tabId: actualTabId } = await withTimeout( deps.executeInPage(code, { tabId }), timeoutMs, - () => new Error(`execute_script timed out after ${timeoutMs / 1000}s`) + () => + new Error( + `execute_script (target=page) timed out after ${timeoutMs / 1000}s waiting for a response. ` + + `The page code cannot be forcibly terminated and may still be running in the tab.` + ), + signal ); return JSON.stringify({ result: result ?? null, target: "page", tab_id: actualTabId }); } // sandbox - const result = await withTimeout( - deps.executeInSandbox(code), - timeoutMs, - () => new Error(`execute_script timed out after ${timeoutMs / 1000}s`) + const result = await executeSandboxWithTimeout( + (executionSignal) => deps.executeInSandbox(code, executionSignal), + signal, + timeoutMs ); return JSON.stringify({ result: result ?? null, target: "sandbox" }); }, diff --git a/src/app/service/agent/core/tools/opfs_tools.test.ts b/src/app/service/agent/core/tools/opfs_tools.test.ts index ce12bd9bb..77cab1d9a 100644 --- a/src/app/service/agent/core/tools/opfs_tools.test.ts +++ b/src/app/service/agent/core/tools/opfs_tools.test.ts @@ -270,6 +270,18 @@ describe("opfs_tools", () => { '".." is not allowed' ); }); + + it("中止时不应写入文件", async () => { + const write = getTool("opfs_write"); + const read = getTool("opfs_read"); + const controller = new AbortController(); + controller.abort(); + + await expect( + write.executor.execute({ path: "cancelled.txt", content: "bad" }, controller.signal) + ).rejects.toThrow("Aborted"); + await expect(read.executor.execute({ path: "cancelled.txt" })).rejects.toThrow(); + }); }); describe("opfs_read 文本读取", () => { diff --git a/src/app/service/agent/core/tools/opfs_tools.ts b/src/app/service/agent/core/tools/opfs_tools.ts index f2810e8ba..e4cb3643c 100644 --- a/src/app/service/agent/core/tools/opfs_tools.ts +++ b/src/app/service/agent/core/tools/opfs_tools.ts @@ -8,6 +8,7 @@ import { writeWorkspaceFile, } from "@App/app/service/agent/core/opfs_helpers"; import { isText } from "@App/pkg/utils/istextorbinary"; +import { throwIfAborted } from "../abort_utils"; import { requireString } from "./param_utils"; // re-export sanitizePath 供外部使用 @@ -134,25 +135,30 @@ export function createOPFSTools(): { tools: Array<{ definition: ToolDefinition; executor: ToolExecutor }>; } { const writeExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const path = requireString(args, "path"); - const result = await writeWorkspaceFile(path, args.content as string | Blob); + const result = await writeWorkspaceFile(path, args.content as string | Blob, signal); + throwIfAborted(signal); return JSON.stringify(result); }, }; const readExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const safePath = sanitizePath(requireString(args, "path")); if (!safePath) throw new Error("path is required"); - const workspace = await getWorkspaceRoot(); + const workspace = await getWorkspaceRoot(false, signal); const { dirPath, fileName } = splitPath(safePath); - const dir = dirPath ? await getDirectory(workspace, dirPath) : workspace; + const dir = dirPath ? await getDirectory(workspace, dirPath, false, signal) : workspace; + throwIfAborted(signal); const fileHandle = await dir.getFileHandle(fileName); const file = await fileHandle.getFile(); const mimeType = guessMimeType(safePath); const arrayBuffer = await file.arrayBuffer(); + throwIfAborted(signal); // 确定返回模式:auto 通过内容字节检测文本/二进制 const mode = (args.mode as string) || "auto"; @@ -163,7 +169,9 @@ export function createOPFSTools(): { if (!createBlobUrlFn) { throw new Error("Blob URL creation not available (Offscreen not initialized)"); } + throwIfAborted(signal); const blobUrl = await createBlobUrlFn(arrayBuffer, mimeType); + throwIfAborted(signal); return JSON.stringify({ path: safePath, blobUrl, size: file.size, mimeType, type: "binary" }); } @@ -201,17 +209,20 @@ export function createOPFSTools(): { }; const listExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const rawPath = (args.path as string) || ""; const safePath = sanitizePath(rawPath); - const workspace = await getWorkspaceRoot(true); - const dir = safePath ? await getDirectory(workspace, safePath) : workspace; + const workspace = await getWorkspaceRoot(true, signal); + const dir = safePath ? await getDirectory(workspace, safePath, false, signal) : workspace; const entries: Array<{ name: string; type: "file" | "directory"; size?: number }> = []; for await (const [name, handle] of dir as unknown as AsyncIterable<[string, FileSystemHandle]>) { + throwIfAborted(signal); if (handle.kind === "file") { const f = await (handle as FileSystemFileHandle).getFile(); + throwIfAborted(signal); entries.push({ name, type: "file", size: f.size }); } else { entries.push({ name, type: "directory" }); @@ -223,13 +234,15 @@ export function createOPFSTools(): { }; const deleteExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const safePath = sanitizePath(requireString(args, "path")); if (!safePath) throw new Error("path is required"); - const workspace = await getWorkspaceRoot(); + const workspace = await getWorkspaceRoot(false, signal); const { dirPath, fileName } = splitPath(safePath); - const dir = dirPath ? await getDirectory(workspace, dirPath) : workspace; + const dir = dirPath ? await getDirectory(workspace, dirPath, false, signal) : workspace; + throwIfAborted(signal); await dir.removeEntry(fileName, { recursive: true }); return JSON.stringify({ success: true }); diff --git a/src/app/service/agent/core/tools/sub_agent.test.ts b/src/app/service/agent/core/tools/sub_agent.test.ts index 81edc42cb..8c4fca00e 100644 --- a/src/app/service/agent/core/tools/sub_agent.test.ts +++ b/src/app/service/agent/core/tools/sub_agent.test.ts @@ -44,6 +44,19 @@ describe("sub_agent", () => { }); }); + it("should forward the invoking tool call's id as toolCallId", async () => { + const mockRunSubAgent = vi.fn().mockResolvedValue({ agentId: "id4", result: "ok" }); + const { executor } = createSubAgentTool({ runSubAgent: mockRunSubAgent }); + + await executor.execute({ prompt: "Research X" }, undefined, "tool-call-42"); + expect(mockRunSubAgent).toHaveBeenCalledWith({ + prompt: "Research X", + description: "Sub-agent task", + type: undefined, + toolCallId: "tool-call-42", + }); + }); + it("should throw if prompt is missing", async () => { const mockRunSubAgent = vi.fn(); const { executor } = createSubAgentTool({ runSubAgent: mockRunSubAgent }); @@ -58,4 +71,24 @@ describe("sub_agent", () => { await expect(executor.execute({ prompt: "fail" })).rejects.toThrow("Agent failed"); }); + + it("应把子代理 usage 暴露给父工具循环", async () => { + const mockRunSubAgent = vi.fn().mockResolvedValue({ + agentId: "usage-child", + result: "done", + usage: { inputTokens: 100, outputTokens: 20 }, + details: { + agentId: "usage-child", + description: "usage", + messages: [], + usage: { inputTokens: 100, outputTokens: 20 }, + }, + }); + const { executor } = createSubAgentTool({ runSubAgent: mockRunSubAgent }); + + await expect(executor.execute({ prompt: "count usage" })).resolves.toMatchObject({ + usage: { inputTokens: 100, outputTokens: 20 }, + subAgentDetails: { usage: { inputTokens: 100, outputTokens: 20 } }, + }); + }); }); diff --git a/src/app/service/agent/core/tools/sub_agent.ts b/src/app/service/agent/core/tools/sub_agent.ts index 8a93d8c98..739034c2f 100644 --- a/src/app/service/agent/core/tools/sub_agent.ts +++ b/src/app/service/agent/core/tools/sub_agent.ts @@ -1,4 +1,4 @@ -import type { SubAgentDetails, ToolDefinition } from "@App/app/service/agent/core/types"; +import type { Attachment, SubAgentDetails, TokenUsage, ToolDefinition } from "@App/app/service/agent/core/types"; import type { ToolExecutor } from "@App/app/service/agent/core/tool_registry"; import { SUB_AGENT_TYPES } from "@App/app/service/agent/core/sub_agent_types"; import { requireString } from "./param_utils"; @@ -9,6 +9,9 @@ export type SubAgentRunOptions = { description: string; type?: string; tabId?: number; // 父代理传递的标签页上下文 + // 发起本次子代理的 agent 工具调用 ID。并行 agent 调用各自独立, + // 供上层把子代理事件与自己的 toolCallId 关联,避免并发时 UI 误匹配到另一个子代理。 + toolCallId?: string; }; // 子代理运行结果 @@ -16,6 +19,9 @@ export type SubAgentRunResult = { agentId: string; result: string; details?: SubAgentDetails; // 执行详情(用于持久化) + usage?: TokenUsage; + attachments?: Attachment[]; + ownedAttachmentIds?: string[]; }; // 在模块加载时固化一次可用 type 列表,供 provider 做 JSON Schema 强校验 @@ -60,18 +66,30 @@ export function createSubAgentTool(params: { executor: ToolExecutor; } { const executor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, _signal?: AbortSignal, toolCallId?: string) => { const prompt = requireString(args, "prompt"); const description = (args.description as string) || "Sub-agent task"; const type = args.type as string | undefined; const tabId = args.tab_id as number | undefined; - const result = await params.runSubAgent({ prompt, description, type, tabId }); + const result = await params.runSubAgent({ prompt, description, type, tabId, toolCallId }); // 返回结构化结果,附带子代理执行详情用于持久化 const content = `[agentId: ${result.agentId}]\n\n${result.result}`; if (result.details) { - return { content, subAgentDetails: result.details }; + return { + content, + subAgentDetails: result.details, + usage: result.usage, + attachments: result.attachments?.map((attachment) => ({ + attachmentId: attachment.id, + type: attachment.type, + name: attachment.name, + mimeType: attachment.mimeType, + size: attachment.size, + })), + ownedAttachmentIds: result.ownedAttachmentIds, + }; } return content; }, diff --git a/src/app/service/agent/core/tools/tab_tools.test.ts b/src/app/service/agent/core/tools/tab_tools.test.ts index 52b309cb1..83d2ba00e 100644 --- a/src/app/service/agent/core/tools/tab_tools.test.ts +++ b/src/app/service/agent/core/tools/tab_tools.test.ts @@ -150,7 +150,7 @@ describe("get_tab_content", () => { const raw = (await executor.execute({ tab_id: 1, prompt: "What is the price?" })) as string; const result = JSON.parse(raw); - expect(mockSummarize).toHaveBeenCalledWith(expect.any(String), "What is the price?"); + expect(mockSummarize).toHaveBeenCalledWith(expect.any(String), "What is the price?", undefined); expect(result.content).toBe("Summarized content"); expect(result.truncated).toBe(false); }); diff --git a/src/app/service/agent/core/tools/tab_tools.ts b/src/app/service/agent/core/tools/tab_tools.ts index 0e82ecb9f..a3accc1ab 100644 --- a/src/app/service/agent/core/tools/tab_tools.ts +++ b/src/app/service/agent/core/tools/tab_tools.ts @@ -5,6 +5,10 @@ import { extractHtmlWithSelectors } from "@App/app/service/offscreen/client"; import { assertDomUrlAllowed } from "@App/app/service/agent/service_worker/dom_policy"; import { stripHtmlTags } from "./web_fetch"; import { requireNumber, requireString, optionalString, optionalNumber, optionalBoolean } from "./param_utils"; +import { raceWithAbort, throwIfAborted } from "../abort_utils"; +import type { TokenUsage } from "../types"; + +type SummaryResult = string | { content: string; usage?: TokenUsage }; // ---- Tool Definitions ---- @@ -98,12 +102,13 @@ const ACTIVATE_TAB_DEFINITION: ToolDefinition = { export function createTabTools(deps: { sender: MessageSend; - summarize: (content: string, prompt: string) => Promise; + summarize: (content: string, prompt: string, signal?: AbortSignal) => Promise; }): { tools: Array<{ definition: ToolDefinition; executor: ToolExecutor }> } { const { sender, summarize } = deps; const getTabContentExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const tabId = requireNumber(args, "tab_id"); const prompt = args.prompt as string | undefined; const selector = optionalString(args, "selector"); @@ -156,11 +161,15 @@ export function createTabTools(deps: { } // 通过 Offscreen 提取 markdown(带 selector 标注) + throwIfAborted(signal); let content: string; try { - const extracted = await extractHtmlWithSelectors(sender, pageData.html); + const extracted = await raceWithAbort(extractHtmlWithSelectors(sender, pageData.html), signal); content = extracted && extracted.length > 20 ? extracted : pageData.html; - } catch { + } catch (error) { + if (error instanceof Error && error.message === "Aborted") { + throw error; + } // 降级:简单去标签 content = stripHtmlTags(pageData.html); } @@ -173,12 +182,15 @@ export function createTabTools(deps: { } // LLM 摘要 + let summaryUsage: TokenUsage | undefined; if (prompt) { - content = await summarize(content, prompt); + const summary = await summarize(content, prompt, signal); + content = typeof summary === "string" ? summary : summary.content; + summaryUsage = typeof summary === "string" ? undefined : summary.usage; truncated = false; // 摘要后不再截断 } - return JSON.stringify({ + const serialized = JSON.stringify({ tab_id: tabId, url: pageData.url, title: pageData.title, @@ -186,11 +198,13 @@ export function createTabTools(deps: { truncated, used_selector: selector || null, }); + return summaryUsage ? { content: serialized, usage: summaryUsage } : serialized; }, }; const listTabsExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const urlPattern = optionalString(args, "url_pattern"); const titlePattern = optionalString(args, "title_pattern"); const active = optionalBoolean(args, "active"); @@ -242,7 +256,8 @@ export function createTabTools(deps: { }; const openTabExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const url = requireString(args, "url"); const tabId = optionalNumber(args, "tab_id"); @@ -253,19 +268,27 @@ export function createTabTools(deps: { await chrome.tabs.update(tabId, { url }); if (waitUntilLoaded) { - await new Promise((resolve) => { + await new Promise((resolve, reject) => { const timerId = setTimeout(() => { chrome.tabs.onUpdated.removeListener(listener); resolve(); }, 30_000); + const onAbort = () => { + clearTimeout(timerId); + chrome.tabs.onUpdated.removeListener(listener); + reject(new Error("Aborted")); + }; const listener = (updatedTabId: number, changeInfo: { status?: string }) => { if (updatedTabId === tabId && changeInfo.status === "complete") { clearTimeout(timerId); + signal?.removeEventListener("abort", onAbort); chrome.tabs.onUpdated.removeListener(listener); resolve(); } }; chrome.tabs.onUpdated.addListener(listener); + signal?.addEventListener("abort", onAbort, { once: true }); + if (signal?.aborted) onAbort(); }); } @@ -297,7 +320,8 @@ export function createTabTools(deps: { }; const closeTabExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const tabId = requireNumber(args, "tab_id"); await chrome.tabs.remove(tabId); return JSON.stringify({ success: true, tab_id: tabId }); @@ -305,7 +329,8 @@ export function createTabTools(deps: { }; const activateTabExecutor: ToolExecutor = { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const tabId = requireNumber(args, "tab_id"); const tab = await chrome.tabs.update(tabId, { active: true }); if (!tab) throw new Error(`Tab ${tabId} not found`); diff --git a/src/app/service/agent/core/tools/task_tools.test.ts b/src/app/service/agent/core/tools/task_tools.test.ts index ac18fa40b..3b20d0d08 100644 --- a/src/app/service/agent/core/tools/task_tools.test.ts +++ b/src/app/service/agent/core/tools/task_tools.test.ts @@ -22,6 +22,56 @@ describe("task_tools", () => { expect(result2).toEqual({ id: "2", subject: "Task 2", description: "Details", status: "pending" }); }); + it("中止时不应创建任务、持久化或发送更新", async () => { + const onSave = vi.fn().mockResolvedValue(undefined); + const sendEvent = vi.fn(); + const { tools, tasks } = createTaskTools({ onSave, sendEvent }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + const controller = new AbortController(); + controller.abort(); + + await expect(create.executor.execute({ subject: "不应创建" }, controller.signal)).rejects.toThrow("Aborted"); + expect(tasks.size).toBe(0); + expect(onSave).not.toHaveBeenCalled(); + expect(sendEvent).not.toHaveBeenCalled(); + }); + + it("onSave 应收到透传的 AbortSignal,Stop 后底层写入可拒绝提交", async () => { + const onSave = vi.fn().mockResolvedValue(undefined); + const { tools } = createTaskTools({ onSave }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + const controller = new AbortController(); + + await create.executor.execute({ subject: "任务" }, controller.signal); + + expect(onSave).toHaveBeenCalledWith(expect.any(Array), controller.signal); + }); + + it("带初始 revision 时每次保存都应透传 CAS 版本", async () => { + const onSave = vi.fn().mockResolvedValue(undefined); + const { tools } = createTaskTools({ onSave, initialRevision: 7 }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + const update = tools.find((tool) => tool.definition.name === "update_task")!; + + await create.executor.execute({ subject: "任务" }); + await update.executor.execute({ task_id: "1", status: "completed" }); + + expect(onSave).toHaveBeenNthCalledWith(1, [{ id: "1", subject: "任务", status: "pending" }], undefined, 7); + expect(onSave).toHaveBeenNthCalledWith(2, [{ id: "1", subject: "任务", status: "completed" }], undefined, 8); + }); + + it("create_task 持久化失败时不应把未提交任务留在内存或消耗 ID", async () => { + const onSave = vi.fn().mockRejectedValueOnce(new Error("disk full")).mockResolvedValue(undefined); + const { tools, tasks } = createTaskTools({ onSave }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + + await expect(create.executor.execute({ subject: "失败任务" })).rejects.toThrow("disk full"); + expect(tasks.size).toBe(0); + + const committed = JSON.parse((await create.executor.execute({ subject: "成功任务" })) as string); + expect(committed.id).toBe("1"); + }); + it("update_task 应更新任务字段", async () => { const { tools } = createTaskTools(); const create = tools.find((t) => t.definition.name === "create_task")!; @@ -36,6 +86,69 @@ describe("task_tools", () => { expect(result.subject).toBe("Updated"); }); + it("update_task 持久化失败时不应污染内存中的已提交任务", async () => { + const onSave = vi.fn().mockResolvedValueOnce(undefined).mockRejectedValueOnce(new Error("disk full")); + const { tools, tasks } = createTaskTools({ onSave }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + const update = tools.find((tool) => tool.definition.name === "update_task")!; + + await create.executor.execute({ subject: "Original" }); + await expect( + update.executor.execute({ task_id: "1", status: "completed", subject: "Uncommitted" }) + ).rejects.toThrow("disk full"); + + expect(tasks.get("1")).toEqual({ id: "1", subject: "Original", status: "pending" }); + }); + + it("持久化落定时同时发生中止也应保持磁盘与内存状态一致", async () => { + const controller = new AbortController(); + let persisted: Task[] = []; + const onSave = vi.fn(async (candidate: Task[]) => { + persisted = candidate.map((task) => ({ ...task })); + controller.abort(); + }); + const { tools, tasks } = createTaskTools({ onSave }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + + await create.executor.execute({ subject: "已提交任务" }, controller.signal); + + expect(persisted).toEqual([{ id: "1", subject: "已提交任务", status: "pending" }]); + expect(Array.from(tasks.values())).toEqual(persisted); + }); + + it("通知失败不应回滚已提交任务或复用已消耗的 ID", async () => { + const onSave = vi.fn().mockResolvedValue(undefined); + const sendEvent = vi.fn().mockImplementationOnce(() => { + throw new Error("port closed"); + }); + const { tools, tasks } = createTaskTools({ onSave, sendEvent }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + + await create.executor.execute({ subject: "任务一" }); + const second = JSON.parse((await create.executor.execute({ subject: "任务二" })) as string); + + expect(second.id).toBe("2"); + expect(Array.from(tasks.keys())).toEqual(["1", "2"]); + }); + + it("同一批并发创建应串行分配 ID 并持久化包含全部任务的快照", async () => { + const snapshots: Task[][] = []; + const onSave = vi.fn(async (candidate: Task[]) => { + snapshots.push(candidate.map((task) => ({ ...task }))); + }); + const { tools, tasks } = createTaskTools({ onSave }); + const create = tools.find((tool) => tool.definition.name === "create_task")!; + + const [first, second] = await Promise.all([ + create.executor.execute({ subject: "任务一" }), + create.executor.execute({ subject: "任务二" }), + ]); + + expect([JSON.parse(first as string).id, JSON.parse(second as string).id]).toEqual(["1", "2"]); + expect(snapshots.at(-1)?.map((task) => task.id)).toEqual(["1", "2"]); + expect(Array.from(tasks.keys())).toEqual(["1", "2"]); + }); + it("update_task 应对不存在的任务抛错", async () => { const { tools } = createTaskTools(); const update = tools.find((t) => t.definition.name === "update_task")!; @@ -89,7 +202,8 @@ describe("task_tools", () => { await create.executor.execute({ subject: "Test" }); expect(onSave).toHaveBeenCalledOnce(); - expect(onSave).toHaveBeenCalledWith([{ id: "1", subject: "Test", status: "pending" }]); + // 第二个参数是透传给底层写入的 AbortSignal,未传 signal 时为 undefined + expect(onSave).toHaveBeenCalledWith([{ id: "1", subject: "Test", status: "pending" }], undefined); expect(sendEvent).toHaveBeenCalledOnce(); expect(sendEvent).toHaveBeenCalledWith({ diff --git a/src/app/service/agent/core/tools/task_tools.ts b/src/app/service/agent/core/tools/task_tools.ts index 5dd4d8690..3c128b04b 100644 --- a/src/app/service/agent/core/tools/task_tools.ts +++ b/src/app/service/agent/core/tools/task_tools.ts @@ -1,5 +1,6 @@ import type { ToolDefinition, ChatStreamEvent } from "@App/app/service/agent/core/types"; import type { ToolExecutor } from "@App/app/service/agent/core/tool_registry"; +import { throwIfAborted } from "../abort_utils"; import { requireString, optionalString } from "./param_utils"; export type Task = { @@ -59,8 +60,10 @@ const LIST_TASKS_DEFINITION: ToolDefinition = { export type TaskToolsOptions = { // 初始任务列表(从持久化加载) initialTasks?: Task[]; - // 任务变更时的持久化回调 - onSave?: (tasks: Task[]) => Promise; + // 初始任务快照 revision;提供时 onSave 会携带 expectedRevision 做 CAS + initialRevision?: number; + // 任务变更时的持久化回调;signal 透传到底层写入,Stop 后不再提交任务快照 + onSave?: (tasks: Task[], signal?: AbortSignal, expectedRevision?: number) => Promise; // 任务变更时的事件推送回调(推送到 UI) sendEvent?: (event: ChatStreamEvent) => void; }; @@ -71,11 +74,13 @@ export function createTaskTools(options?: TaskToolsOptions): { } { const tasks = new Map(); let nextId = 1; + let revision = options?.initialRevision; + let mutationQueue = Promise.resolve(); // 从持久化数据恢复 if (options?.initialTasks) { for (const task of options.initialTasks) { - tasks.set(task.id, task); + tasks.set(task.id, { ...task }); const numId = parseInt(task.id, 10); if (!isNaN(numId) && numId >= nextId) { nextId = numId + 1; @@ -83,56 +88,93 @@ export function createTaskTools(options?: TaskToolsOptions): { } } - // 持久化并推送事件 - const emitUpdate = async () => { - const taskList = Array.from(tasks.values()); - if (options?.onSave) { - await options.onSave(taskList); - } - if (options?.sendEvent) { - options.sendEvent({ + const runMutation = (mutation: () => Promise): Promise => { + const result = mutationQueue.then(mutation, mutation); + mutationQueue = result.then( + () => undefined, + () => undefined + ); + return result; + }; + + // 先持久化候选快照,成功后才替换内存状态并推送事件。 + const emitTaskUpdate = (taskList: Task[]) => { + try { + options?.sendEvent?.({ type: "task_update", - tasks: taskList.map((t) => ({ - id: t.id, - subject: t.subject, - status: t.status, - description: t.description, + tasks: taskList.map((task) => ({ + id: task.id, + subject: task.subject, + status: task.status, + description: task.description, })), }); + } catch { + // 持久化已经落定;连接关闭等通知失败不能把已提交的工具操作伪装成失败。 + } + }; + + const commitUpdate = async (candidate: Map, signal?: AbortSignal): Promise => { + const taskList = Array.from(candidate.values(), (task) => ({ ...task })); + throwIfAborted(signal); + if (options?.onSave) { + if (revision === undefined) { + await options.onSave(taskList, signal); + } else { + await options.onSave(taskList, signal, revision); + revision++; + } } + // onSave resolve 表示候选快照已经提交。即使 signal 恰好在 close() 落定后 abort, + // 内存也必须接受同一份快照,不能制造“磁盘已提交、内存仍回滚”的分叉状态。 + tasks.clear(); + for (const task of taskList) tasks.set(task.id, task); + return taskList; }; const createExecutor: ToolExecutor = { - execute: async (args: Record) => { - const task: Task = { - id: String(nextId++), - subject: requireString(args, "subject"), - description: optionalString(args, "description"), - status: "pending", - }; - tasks.set(task.id, task); - await emitUpdate(); - return JSON.stringify(task); - }, + execute: (args: Record, signal?: AbortSignal) => + runMutation(async () => { + throwIfAborted(signal); + const task: Task = { + id: String(nextId), + subject: requireString(args, "subject"), + description: optionalString(args, "description"), + status: "pending", + }; + const candidate = new Map(tasks); + candidate.set(task.id, task); + const committed = await commitUpdate(candidate, signal); + nextId++; + emitTaskUpdate(committed); + return JSON.stringify(task); + }), }; const updateExecutor: ToolExecutor = { - execute: async (args: Record) => { - const taskId = requireString(args, "task_id"); - const task = tasks.get(taskId); - if (!task) { - throw new Error(`Task "${taskId}" not found`); - } - if (args.status) task.status = args.status as Task["status"]; - if (args.subject) task.subject = args.subject as string; - if (args.description !== undefined) task.description = args.description as string; - await emitUpdate(); - return JSON.stringify(task); - }, + execute: (args: Record, signal?: AbortSignal) => + runMutation(async () => { + throwIfAborted(signal); + const taskId = requireString(args, "task_id"); + const existing = tasks.get(taskId); + if (!existing) { + throw new Error(`Task "${taskId}" not found`); + } + const task = { ...existing }; + if (args.status) task.status = args.status as Task["status"]; + if (args.subject) task.subject = args.subject as string; + if (args.description !== undefined) task.description = args.description as string; + const candidate = new Map(tasks); + candidate.set(taskId, task); + const committed = await commitUpdate(candidate, signal); + emitTaskUpdate(committed); + return JSON.stringify(task); + }), }; const listExecutor: ToolExecutor = { - execute: async () => { + execute: async (_args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); return JSON.stringify(Array.from(tasks.values())); }, }; diff --git a/src/app/service/agent/core/tools/web_fetch.ts b/src/app/service/agent/core/tools/web_fetch.ts index 2d18efd52..b668ba2fc 100644 --- a/src/app/service/agent/core/tools/web_fetch.ts +++ b/src/app/service/agent/core/tools/web_fetch.ts @@ -3,6 +3,10 @@ import type { ToolExecutor } from "@App/app/service/agent/core/tool_registry"; import type { MessageSend } from "@Packages/message/types"; import { extractHtmlContent } from "@App/app/service/offscreen/client"; import { requireString, optionalNumber } from "./param_utils"; +import { raceWithAbort, throwIfAborted } from "../abort_utils"; +import type { TokenUsage } from "../types"; + +type SummaryResult = string | { content: string; usage?: TokenUsage }; // Agent User-Agent 字符串 const AGENT_USER_AGENT = "Mozilla/5.0 (compatible; ScriptCat Agent)"; @@ -38,16 +42,16 @@ export function stripHtmlTags(html: string): string { } export class WebFetchExecutor implements ToolExecutor { - private summarize?: (content: string, prompt: string) => Promise; + private summarize?: (content: string, prompt: string, signal?: AbortSignal) => Promise; constructor( private sender: MessageSend, - deps?: { summarize?: (content: string, prompt: string) => Promise } + deps?: { summarize?: (content: string, prompt: string, signal?: AbortSignal) => Promise } ) { this.summarize = deps?.summarize; } - async execute(args: Record): Promise { + async execute(args: Record, signal?: AbortSignal): Promise { const url = requireString(args, "url"); const prompt = requireString(args, "prompt"); const maxLength = optionalNumber(args, "max_length"); @@ -63,9 +67,11 @@ export class WebFetchExecutor implements ToolExecutor { throw new Error("Only http/https URLs are supported"); } + throwIfAborted(signal); + const response = await fetch(url, { headers: { "User-Agent": AGENT_USER_AGENT }, - signal: AbortSignal.timeout(30_000), + signal: signal ? AbortSignal.any([signal, AbortSignal.timeout(30_000)]) : AbortSignal.timeout(30_000), }); if (!response.ok) { throw new Error(`HTTP ${response.status}: ${response.statusText}`); @@ -73,9 +79,11 @@ export class WebFetchExecutor implements ToolExecutor { // 检测重定向:最终 URL 与请求 URL 不同 const finalUrl = response.url && response.url !== url ? response.url : undefined; + throwIfAborted(signal); const contentType = response.headers.get("content-type") || ""; const text = await response.text(); + throwIfAborted(signal); let content: string; let detectedType: string; @@ -94,7 +102,7 @@ export class WebFetchExecutor implements ToolExecutor { // b) Content-Type 含 html 或未知 → 送 Offscreen extractHtmlContent else if (contentType.includes("html") || !contentType) { try { - const extracted = await extractHtmlContent(this.sender, text); + const extracted = await raceWithAbort(extractHtmlContent(this.sender, text), signal); if (extracted && extracted.length > 50) { content = extracted; detectedType = "html"; @@ -103,7 +111,10 @@ export class WebFetchExecutor implements ToolExecutor { content = stripHtmlTags(text); detectedType = "text"; } - } catch { + } catch (error) { + if (error instanceof Error && error.message === "Aborted") { + throw error; + } // Offscreen 提取失败,降级 content = stripHtmlTags(text); detectedType = "text"; @@ -122,8 +133,11 @@ export class WebFetchExecutor implements ToolExecutor { } // LLM 摘要 + let summaryUsage: TokenUsage | undefined; if (this.summarize) { - content = await this.summarize(content, prompt); + const summary = await this.summarize(content, prompt, signal); + content = typeof summary === "string" ? summary : summary.content; + summaryUsage = typeof summary === "string" ? undefined : summary.usage; truncated = false; } @@ -136,6 +150,7 @@ export class WebFetchExecutor implements ToolExecutor { if (finalUrl) { result.final_url = finalUrl; } - return JSON.stringify(result); + const serialized = JSON.stringify(result); + return summaryUsage ? { content: serialized, usage: summaryUsage } : serialized; } } diff --git a/src/app/service/agent/core/tools/web_search.ts b/src/app/service/agent/core/tools/web_search.ts index ee52ddbb0..fe1637b6e 100644 --- a/src/app/service/agent/core/tools/web_search.ts +++ b/src/app/service/agent/core/tools/web_search.ts @@ -5,6 +5,7 @@ import type { SearchConfigRepo } from "./search_config"; import { extractSearchResults, extractBingResults, extractBaiduResults } from "@App/app/service/offscreen/client"; import { withTimeout } from "@App/pkg/utils/with_timeout"; import { requireString, optionalNumber } from "./param_utils"; +import { throwIfAborted } from "../abort_utils"; // Agent User-Agent 字符串 const AGENT_USER_AGENT = "Mozilla/5.0 (compatible; ScriptCat Agent)"; @@ -49,7 +50,7 @@ export class WebSearchExecutor implements ToolExecutor { private configRepo: SearchConfigRepo ) {} - async execute(args: Record): Promise { + async execute(args: Record, signal?: AbortSignal): Promise { const query = requireString(args, "query"); const maxResults = Math.min(optionalNumber(args, "max_results") ?? 5, 10); @@ -57,14 +58,14 @@ export class WebSearchExecutor implements ToolExecutor { switch (config.engine) { case "google_custom": - return this.searchGoogle(query, maxResults, config.googleApiKey || "", config.googleCseId || ""); + return this.searchGoogle(query, maxResults, config.googleApiKey || "", config.googleCseId || "", signal); case "duckduckgo": - return this.searchDuckDuckGo(query, maxResults); + return this.searchDuckDuckGo(query, maxResults, signal); case "baidu": - return this.searchBaidu(query, maxResults); + return this.searchBaidu(query, maxResults, signal); case "bing": default: - return this.searchBing(query, maxResults); + return this.searchBing(query, maxResults, signal); } } @@ -79,49 +80,74 @@ export class WebSearchExecutor implements ToolExecutor { url: string, extractFn: (html: string) => Promise, engineName: string, - maxResults: number + maxResults: number, + signal?: AbortSignal ): Promise { + throwIfAborted(signal); + const response = await fetch(url, { headers: { "User-Agent": AGENT_USER_AGENT }, - signal: AbortSignal.timeout(SEARCH_TIMEOUT_MS), + signal: signal + ? AbortSignal.any([signal, AbortSignal.timeout(SEARCH_TIMEOUT_MS)]) + : AbortSignal.timeout(SEARCH_TIMEOUT_MS), }); if (!response.ok) { throw new Error(`${engineName} search failed: HTTP ${response.status}`); } const html = await response.text(); + throwIfAborted(signal); let results: SearchResult[] = []; let extractionFailed = false; try { // 提取函数走 Offscreen 通道,加 10s 超时防卡死 - results = await withTimeout(extractFn(html), 10_000, () => new Error("extract timeout")); - } catch { + results = await withTimeout(extractFn(html), 10_000, () => new Error("extract timeout"), signal); + } catch (error) { + if (error instanceof Error && error.message === "Aborted") { + throw error; + } extractionFailed = true; } return formatSearchResults(results.slice(0, maxResults), extractionFailed, engineName); } - private async searchDuckDuckGo(query: string, maxResults: number): Promise { + private async searchDuckDuckGo(query: string, maxResults: number, signal?: AbortSignal): Promise { const url = `https://html.duckduckgo.com/html/?q=${encodeURIComponent(query)}`; - return this.fetchAndExtract(url, (html) => extractSearchResults(this.sender, html), "DuckDuckGo", maxResults); + return this.fetchAndExtract( + url, + (html) => extractSearchResults(this.sender, html), + "DuckDuckGo", + maxResults, + signal + ); } - private async searchBing(query: string, maxResults: number): Promise { + private async searchBing(query: string, maxResults: number, signal?: AbortSignal): Promise { const url = `https://www.bing.com/search?q=${encodeURIComponent(query)}`; - return this.fetchAndExtract(url, (html) => extractBingResults(this.sender, html), "Bing", maxResults); + return this.fetchAndExtract(url, (html) => extractBingResults(this.sender, html), "Bing", maxResults, signal); } - private async searchBaidu(query: string, maxResults: number): Promise { + private async searchBaidu(query: string, maxResults: number, signal?: AbortSignal): Promise { const url = `https://www.baidu.com/s?wd=${encodeURIComponent(query)}&rn=${maxResults}`; - return this.fetchAndExtract(url, (html) => extractBaiduResults(this.sender, html), "Baidu", maxResults); + return this.fetchAndExtract(url, (html) => extractBaiduResults(this.sender, html), "Baidu", maxResults, signal); } - private async searchGoogle(query: string, maxResults: number, apiKey: string, cseId: string): Promise { + private async searchGoogle( + query: string, + maxResults: number, + apiKey: string, + cseId: string, + signal?: AbortSignal + ): Promise { if (!apiKey || !cseId) { throw new Error("Google Custom Search requires API Key and CSE ID. Configure them in Agent Tool Settings."); } const url = `https://www.googleapis.com/customsearch/v1?key=${encodeURIComponent(apiKey)}&cx=${encodeURIComponent(cseId)}&q=${encodeURIComponent(query)}&num=${maxResults}`; - const response = await fetch(url, { signal: AbortSignal.timeout(SEARCH_TIMEOUT_MS) }); + const response = await fetch(url, { + signal: signal + ? AbortSignal.any([signal, AbortSignal.timeout(SEARCH_TIMEOUT_MS)]) + : AbortSignal.timeout(SEARCH_TIMEOUT_MS), + }); if (!response.ok) { const text = await response.text().catch(() => ""); diff --git a/src/app/service/agent/core/types.ts b/src/app/service/agent/core/types.ts index 54574d295..c2f197d5e 100644 --- a/src/app/service/agent/core/types.ts +++ b/src/app/service/agent/core/types.ts @@ -22,6 +22,10 @@ export type MessageContent = string | ContentBlock[]; export type Conversation = { id: string; + /** Immutable identity for this incarnation of an ID. Filled when legacy records are loaded. */ + generation?: string; + /** Optimistic-concurrency version. Filled when legacy records are loaded. */ + revision?: number; title: string; modelId: string; system?: string; @@ -46,19 +50,24 @@ export type AttachmentData = { type: "image" | "file" | "audio"; name: string; mimeType: string; - data: string | Blob; // base64/data URL 或 Blob + data?: string | Blob; // base64/data URL 或 Blob;省略时表示已持久化引用 + attachmentId?: string; + size?: number; }; export type ToolResultWithAttachments = { content: string; // 文本结果(发给 LLM) attachments: AttachmentData[]; // 附件数据(仅存储+展示) + structuredContent?: unknown; // MCP 等工具提供的结构化结果(保留给后续处理) }; // 子代理单轮消息(持久化用) export type SubAgentMessage = { - content: string; + content: MessageContent; thinking?: string; toolCalls: ToolCall[]; + /** 生成数据丢失等非致命警告(如图片保存失败),需与最终回复分支同样持久化,避免刷新后丢失提示 */ + warning?: string; }; // 子代理执行详情(持久化到 ToolCall) @@ -76,6 +85,8 @@ export type ToolCall = { arguments: string; result?: string; attachments?: Attachment[]; + /** Attachment files created by this tool call and owned by the persisted conversation. */ + ownedAttachmentIds?: string[]; subAgentDetails?: SubAgentDetails; status?: "pending" | "running" | "completed" | "error"; }; @@ -89,11 +100,15 @@ export type ChatMessage = { conversationId: string; role: MessageRole; content: MessageContent; + /** Attachment files owned by this persisted message. ContentBlock references are borrowed unless listed here. */ + ownedAttachmentIds?: string[]; thinking?: ThinkingBlock; toolCalls?: ToolCall[]; // tool 角色的消息需要关联到对应的 tool_call toolCallId?: string; error?: string; + // 错误分类码,用于 UI 判断错误类型并展示针对性操作 + errorCode?: string; warning?: string; modelId?: string; usage?: TokenUsage; @@ -110,6 +125,9 @@ export type SubAgentEventInfo = { agentId: string; description: string; subAgentType?: string; + // 发起该子代理的父级 agent 工具调用 ID。并行 agent 调用各自持有独立的 toolCallId, + // UI 据此做显式 toolCallId -> agentId 匹配,避免并发时错误关联到另一个子代理。 + toolCallId?: string; }; // LLM 流式输出事件 @@ -118,7 +136,15 @@ export type LLMStreamEvent = | { type: "thinking_delta"; delta: string } | { type: "tool_call_start"; toolCall: Omit } | { type: "tool_call_delta"; id: string; delta: string; index?: number } - | { type: "tool_call_complete"; id: string; result: string; attachments?: Attachment[] } + | { + type: "tool_call_complete"; + id: string; + result: string; + status?: "completed" | "error"; + attachments?: Attachment[]; + /** Internal ownership handoff used to clean nested tool attachments if the parent round cannot commit. */ + ownedAttachmentIds?: string[]; + } | { type: "content_block_start"; block: Omit } | { type: "content_block_complete"; block: ImageBlock | FileBlock | AudioBlock; data?: string }; @@ -136,15 +162,31 @@ export type ForwardableEvent = }; durationMs?: number; } - | { type: "error"; message: string; errorCode?: string } - | { type: "retry"; attempt: number; maxRetries: number; error: string; delayMs: number }; + | { + type: "error"; + message: string; + errorCode?: string; + usage?: TokenUsage; + durationMs?: number; + } + | { type: "retry"; attempt: number; maxRetries: number; error: string; delayMs: number } + | { type: "system_warning"; message: string }; // Service Worker -> UI/Sandbox 的流式事件(通过 MessageConnect 的 sendMessage 传输) // ForwardableEvent 携带可选 subAgent 标识(扁平化子代理事件,消除递归包装) export type ChatStreamEvent = | (ForwardableEvent & { subAgent?: SubAgentEventInfo }) - | { type: "ask_user"; id: string; question: string; options?: string[]; multiple?: boolean } - | { type: "system_warning"; message: string } + | { + type: "ask_user"; + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + } + | { type: "ask_user_expired"; id: string } + | { type: "ask_user_resolved"; id: string } | { type: "task_update"; tasks: Array<{ @@ -158,7 +200,14 @@ export type ChatStreamEvent = | { type: "sync"; streamingMessage?: { content: string; thinking?: string; toolCalls: ToolCall[] }; - pendingAskUser?: { id: string; question: string; options?: string[]; multiple?: boolean }; + pendingAskUser?: { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + }; tasks: Array<{ id: string; subject: string; @@ -220,9 +269,8 @@ export type ConversationCreateOptions = { id?: string; system?: string; model?: string; // modelId,不传则使用默认模型 - maxIterations?: number; // tool calling 最大循环次数,默认 20 skills?: "auto" | string[]; // 加载的 Skill,"auto" 加载全部,数组指定名称 - tools?: Array) => Promise }>; + tools?: Array, signal: AbortSignal) => Promise }>; commands?: Record; // 自定义命令处理器,以 / 开头 ephemeral?: boolean; // 临时会话:不持久化、不加载内置资源、工具由脚本提供 cache?: boolean; // 是否启用 prompt caching,默认 true @@ -231,7 +279,7 @@ export type ConversationCreateOptions = { // conv.chat() 的参数 export type ChatOptions = { - tools?: Array) => Promise }>; + tools?: Array, signal: AbortSignal) => Promise }>; }; // conv.chat() 的返回值 @@ -245,12 +293,24 @@ export type ChatReply = { cacheCreationInputTokens?: number; cacheReadInputTokens?: number; }; + durationMs?: number; command?: boolean; // 标识该回复来自命令处理 + /** 生成数据丢失等非致命警告(如图片保存失败),与 UI 侧的消息 warning 字段同一语义 */ + warning?: string; }; // conv.chatStream() 的流式 chunk export type StreamChunk = { - type: "content_delta" | "thinking_delta" | "tool_call" | "content_block" | "done" | "error"; + type: + | "content_delta" + | "thinking_delta" + | "tool_call" + | "tool_call_complete" + | "content_block" + | "new_message" + | "system_warning" + | "done" + | "error"; content?: string; block?: ContentBlock; toolCall?: ToolCall; @@ -260,10 +320,13 @@ export type StreamChunk = { cacheCreationInputTokens?: number; cacheReadInputTokens?: number; }; + durationMs?: number; error?: string; - /** 错误分类码:"rate_limit" | "auth" | "tool_timeout" | "max_iterations" | "api_error" */ + /** 错误分类码:"rate_limit" | "auth" | "tool_timeout" | "context_too_large" | "api_error" */ errorCode?: string; command?: boolean; // 标识该 chunk 来自命令处理 + /** type 为 "system_warning" 时携带的警告文本;type 为 "done" 时携带本轮累计警告 */ + warning?: string; }; // ---- Skill 类型 ---- @@ -577,6 +640,10 @@ export type MCPApiRequest = /** 定时任务基础字段(两种模式共用) */ type AgentTaskBase = { id: string; + /** Immutable identity for this incarnation of the task ID. */ + generation?: string; + /** Optimistic-concurrency version. */ + revision?: number; name: string; crontab: string; // cron 表达式(复用 cron.ts 格式) enabled: boolean; @@ -596,8 +663,9 @@ export type InternalAgentTask = AgentTaskBase & { prompt: string; // 每次触发发送的消息 modelId?: string; // 使用的模型 ID conversationId?: string; // 可选:续接已有对话 + /** conversationId 指向对话被绑定时的 generation;执行时若当前 generation 不一致(会话已被删除重建)则拒绝续接 */ + conversationGeneration?: string; skills?: "auto" | string[]; - maxIterations?: number; // 工具循环上限,默认 10 }; /** 事件模式:通知用户脚本处理 */ @@ -630,9 +698,9 @@ export type AgentTaskApiRequest = | { action: "list" } | { action: "get"; id: string } | { action: "create"; task: Omit } - | { action: "update"; id: string; task: Partial } - | { action: "delete"; id: string } - | { action: "enable"; id: string; enabled: boolean } + | { action: "update"; id: string; generation: string; revision: number; task: Partial } + | { action: "delete"; id: string; generation: string; revision: number } + | { action: "enable"; id: string; generation: string; revision: number; enabled: boolean } | { action: "runNow"; id: string } | { action: "listRuns"; taskId: string; limit?: number } | { action: "clearRuns"; taskId: string }; @@ -644,6 +712,9 @@ export type ConversationApiRequest = | { action: "chat"; conversationId: string; + // 调用方持有的会话 generation;若与当前存储的 generation 不一致(会话已被删除重建), + // 服务端拒绝该次操作而不是静默作用于新的一代会话 + generation?: string; message: MessageContent; tools?: ToolDefinition[]; scriptUuid: string; @@ -653,6 +724,14 @@ export type ConversationApiRequest = system?: string; modelId?: string; } - | { action: "getMessages"; conversationId: string; scriptUuid: string } - | { action: "save"; conversationId: string; scriptUuid: string } - | { action: "clearMessages"; conversationId: string; scriptUuid: string }; + | { action: "getMessages"; conversationId: string; generation?: string; scriptUuid: string } + | { action: "save"; conversationId: string; generation?: string; scriptUuid: string } + | { action: "clearMessages"; conversationId: string; generation?: string; scriptUuid?: string } + | { + action: "deleteMessages"; + conversationId: string; + generation?: string; + messageIds: string[]; + preserveAttachmentIds?: string[]; + } + | { action: "delete"; conversationId: string; generation: string; revision?: number }; diff --git a/src/app/service/agent/service_worker/agent.ts b/src/app/service/agent/service_worker/agent.ts index 37661d45d..4ef1420f7 100644 --- a/src/app/service/agent/service_worker/agent.ts +++ b/src/app/service/agent/service_worker/agent.ts @@ -32,6 +32,7 @@ import { SkillService } from "./skill_service"; import { AgentTaskService } from "./task_service"; import { AgentModelService } from "./model_service"; import { AgentTaskRepo, AgentTaskRunRepo } from "@App/app/repo/agent_task"; +import { ScriptDAO } from "@App/app/repo/scripts"; import { AgentTaskScheduler } from "@App/app/service/agent/core/task_scheduler"; import { WEB_FETCH_DEFINITION, WebFetchExecutor } from "@App/app/service/agent/core/tools/web_fetch"; import { WEB_SEARCH_DEFINITION, WebSearchExecutor } from "@App/app/service/agent/core/tools/web_search"; @@ -39,11 +40,12 @@ import { SearchConfigRepo, type SearchEngineConfig } from "@App/app/service/agen import { SubAgentService } from "./sub_agent_service"; import { BackgroundSessionManager } from "./background_session_manager"; import { createOPFSTools, setCreateBlobUrlFn } from "@App/app/service/agent/core/tools/opfs_tools"; -import { createObjectURL } from "@App/app/service/offscreen/client"; import { AgentOPFSService } from "./opfs_service"; -import { executeSkillScript } from "@App/app/service/offscreen/client"; +import { createObjectURL, executeSkillScript, stopScript } from "@App/app/service/offscreen/client"; import { createTabTools } from "@App/app/service/agent/core/tools/tab_tools"; +import type { AttachmentSnapshot } from "@App/app/service/agent/core/attachment_resolver"; import { ChatService } from "./chat_service"; +import { createAbortError, throwIfAborted } from "@App/app/service/agent/core/abort_utils"; // 保留对外 API(测试文件直接从 "./agent" import 这三个函数) export { isRetryableError, withRetry, classifyErrorCode } from "./retry_utils"; @@ -62,6 +64,7 @@ export class AgentService { private opfsService!: AgentOPFSService; private taskRepo = new AgentTaskRepo(); private taskRunRepo = new AgentTaskRunRepo(); + private scriptDAO = new ScriptDAO(); private taskScheduler!: AgentTaskScheduler; // 定时任务逻辑委托给 AgentTaskService private agentTaskService!: AgentTaskService; @@ -101,8 +104,8 @@ export class AgentService { { // callLLM 通过 lambda 注入,确保测试 spy 可以拦截 service.callLLM callLLM: (model, params, sendEvent, signal) => this.callLLM(model, params, sendEvent, signal), - autoCompact: (convId, model, msgs, sendEvent, signal) => - this.compactService.autoCompact(convId, model, msgs, sendEvent, signal), + autoCompact: (convId, generation, model, msgs, sendEvent, signal) => + this.compactService.autoCompact(convId, generation, model, msgs, sendEvent, signal), }, agentChatRepo ); @@ -118,15 +121,32 @@ export class AgentService { this.subAgentService, { executeInPage: (code, options) => this.domService.executeScript(code, options), - executeInSandbox: (code) => { + executeInSandbox: (code: string, signal?: AbortSignal) => { + throwIfAborted(signal); const uuid = SKILL_SCRIPT_UUID_PREFIX + uuidv4(); - return executeSkillScript(this.sender, { + const execPromise = executeSkillScript(this.sender, { uuid, code, args: {}, grants: [], name: "execute_script", }); + if (!signal) return execPromise; + return new Promise((resolve, reject) => { + let aborted = false; + const onAbort = () => { + if (aborted) return; + aborted = true; + void stopScript(this.sender, uuid); + reject(createAbortError()); + }; + signal.addEventListener("abort", onAbort, { once: true }); + if (signal.aborted) { + onAbort(); + return; + } + execPromise.then(resolve, reject).finally(() => signal.removeEventListener("abort", onAbort)); + }); }, }, { @@ -199,16 +219,19 @@ export class AgentService { callLLMWithToolLoop: (params) => this.callLLMWithToolLoop(params), }, this.taskRepo, - this.taskRunRepo + this.taskRunRepo, + this.scriptDAO ); // 初始化定时任务调度器 this.taskScheduler = new AgentTaskScheduler( this.taskRepo, this.taskRunRepo, - (task) => this.agentTaskService.executeInternalTask(task), - (task) => this.agentTaskService.emitTaskEvent(task) + (task, signal) => this.agentTaskService.executeInternalTask(task, signal), + (task, signal) => this.agentTaskService.emitTaskEvent(task, signal) ); - this.taskScheduler.init(); + void this.taskScheduler + .init() + .catch((error) => console.error("[AgentTaskScheduler] initialization failed:", error)); // 注入 scheduler 到 AgentTaskService(解决循环依赖) this.agentTaskService.setScheduler(this.taskScheduler); // 搜索配置 API(供 Options UI 调用) @@ -218,7 +241,7 @@ export class AgentService { this.toolRegistry.registerBuiltin( WEB_FETCH_DEFINITION, new WebFetchExecutor(this.sender, { - summarize: (content, prompt) => this.summarizeContent(content, prompt), + summarize: (content, prompt, signal) => this.summarizeContent(content, prompt, signal), }) ); this.toolRegistry.registerBuiltin(WEB_SEARCH_DEFINITION, new WebSearchExecutor(this.sender, this.searchConfigRepo)); @@ -235,7 +258,7 @@ export class AgentService { // 注册 Tab 操作工具 const tabTools = createTabTools({ sender: this.sender, - summarize: (content, prompt) => this.summarizeContent(content, prompt), + summarize: (content, prompt, signal) => this.summarizeContent(content, prompt, signal), }); for (const t of tabTools.tools) { this.toolRegistry.registerBuiltin(t.definition, t.executor); @@ -344,9 +367,9 @@ export class AgentService { async handleConversationChatFromGmApi( params: { conversationId: string; + generation?: string; message: MessageContent; tools?: ToolDefinition[]; - maxIterations?: number; scriptUuid: string; // ephemeral 会话专用字段 ephemeral?: boolean; @@ -362,7 +385,10 @@ export class AgentService { } // 附加到后台运行会话,供 GMApi 调用 - async handleAttachToConversationFromGmApi(params: { conversationId: string }, sender: IGetSender) { + async handleAttachToConversationFromGmApi( + params: { conversationId: string; generation?: string }, + sender: IGetSender + ) { return this.handleAttachToConversation(params, sender); } @@ -372,7 +398,10 @@ export class AgentService { } // 附加到后台运行中的会话(委托给 BackgroundSessionManager) - private async handleAttachToConversation(params: { conversationId: string }, sender: IGetSender) { + private async handleAttachToConversation( + params: { conversationId: string; generation?: string }, + sender: IGetSender + ) { return this.bgSessionManager.handleAttach(params, sender); } @@ -396,14 +425,19 @@ export class AgentService { // 对内容做摘要/提取(供 tab 工具使用) // 优先使用摘要模型,fallback 到默认模型 - private async summarizeContent(content: string, prompt: string): Promise { - return this.compactService.summarizeContent(content, prompt); + private async summarizeContent(content: string, prompt: string, signal?: AbortSignal) { + return this.compactService.summarizeContent(content, prompt, signal); } // 调用 LLM 并收集完整响应(委托给 LLMClient) private async callLLM( model: AgentModelConfig, - params: { messages: ChatRequest["messages"]; tools?: ToolDefinition[]; cache?: boolean }, + params: { + messages: ChatRequest["messages"]; + tools?: ToolDefinition[]; + cache?: boolean; + attachmentSnapshot?: AttachmentSnapshot; + }, sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal ) { diff --git a/src/app/service/agent/service_worker/autocompact.test.ts b/src/app/service/agent/service_worker/autocompact.test.ts index 9544fac08..70a4c9bb2 100644 --- a/src/app/service/agent/service_worker/autocompact.test.ts +++ b/src/app/service/agent/service_worker/autocompact.test.ts @@ -1,5 +1,5 @@ import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; -import { createTestService, makeSSEResponse } from "./test-helpers"; +import { createTestService, createMockSenderWithCallbacks, makeSSEResponse } from "./test-helpers"; // ---- Compact 功能测试 ---- @@ -16,7 +16,7 @@ describe("Compact 功能", () => { function makeTextResponseWithTokens(text: string, promptTokens = 10): Response { return makeSSEResponse([ - `data: {"choices":[{"delta":{"content":${JSON.stringify(text)}}}]}\n\n`, + `data: {"choices":[{"delta":{"content":${JSON.stringify(text)}},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":${promptTokens},"completion_tokens":5}}\n\n`, ]); } @@ -76,6 +76,35 @@ describe("Compact 功能", () => { expect(events.some((e: any) => e.type === "done")).toBe(true); }); + it("手动 compact:Stop 恰好落在摘要提交之后时不应以旧快照回滚新历史", async () => { + const { service, mockRepo } = createTestService(); + const { sender, sentMessages, simulateMessage } = createMockSenderWithCallbacks(); + + const original = [ + { id: "m1", conversationId: "conv-1", role: "user", content: "Hello", createtime: 1 }, + { id: "m2", conversationId: "conv-1", role: "assistant", content: "Hi there!", createtime: 2 }, + ]; + mockRepo.listConversations.mockResolvedValue([BASE_CONV]); + mockRepo.getMessages.mockResolvedValue(original); + + const saveCalls: any[][] = []; + mockRepo.saveMessages.mockImplementation(async (_id: string, messages: any[]) => { + saveCalls.push(messages); + // 模拟 Stop 恰好落在 close() 提交窗口:写入已生效,signal 事后才被观察到 + if (saveCalls.length === 1) simulateMessage({ action: "stop" }); + }); + + fetchSpy.mockResolvedValueOnce(makeTextResponseWithTokens("迟到的压缩")); + + await (service as any).handleConversationChat({ conversationId: "conv-1", message: "", compact: true }, sender); + + // 摘要已通过 CAS 提交,取消只能影响终态;旧历史不能再覆盖这次已提交写入。 + expect(saveCalls).toHaveLength(1); + expect(saveCalls[0]).toEqual([expect.objectContaining({ content: expect.stringContaining("迟到的压缩") })]); + const events = sentMessages.map((m: any) => m.data); + expect(events.some((e: any) => e.type === "compact_done")).toBe(false); + }); + it("手动 compact:带自定义指令", async () => { const { service, mockRepo } = createTestService(); const { sender } = createMockSender(); @@ -99,6 +128,66 @@ describe("Compact 功能", () => { expect(lastUserMsg.content).toContain("只保留代码"); }); + it("手动 compact:先过滤错误占位消息并裁剪超出上下文阈值的旧工具结果", async () => { + const { service, mockRepo, mockModelRepo } = createTestService(); + const { sender } = createMockSender(); + mockModelRepo.getModel.mockResolvedValue({ + id: "test-openai", + name: "Test", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", + // 40000 窗口 + 未配置 maxTokens:默认输出预留 min(16384, 窗口/4)=10000(见 model_context.ts), + // 输入预算 = 40000 - 10000 - 4000(10% 安全边际) = 26000 token;下面 6 条约 7200 token 的 + // 工具结果远超预算,保证"裁剪旧工具结果"场景稳定触发 + contextWindow: 40000, + }); + mockRepo.listConversations.mockResolvedValue([BASE_CONV]); + const messages: any[] = []; + for (let i = 0; i < 6; i++) { + messages.push({ + id: `a${i}`, + conversationId: "conv-1", + role: "assistant", + content: "", + toolCalls: [{ id: `tc${i}`, name: "tool", arguments: "{}" }], + createtime: i, + }); + messages.push({ + id: `t${i}`, + conversationId: "conv-1", + role: "tool", + content: "工具结果".repeat(1200), + toolCallId: `tc${i}`, + createtime: i, + }); + } + messages.push({ + id: "error", + conversationId: "conv-1", + role: "assistant", + content: "", + error: "Reply was generated but failed to save", + errorCode: "persist_failed", + createtime: 99, + }); + mockRepo.getMessages.mockResolvedValue(messages); + fetchSpy.mockResolvedValueOnce(makeTextResponseWithTokens("压缩完成")); + + await (service as any).handleConversationChat({ conversationId: "conv-1", message: "", compact: true }, sender); + + const body = JSON.parse(fetchSpy.mock.calls[0][1].body as string); + expect( + body.messages.some( + (message: any) => message.role === "assistant" && message.content === "" && !message.tool_calls + ) + ).toBe(false); + expect( + body.messages.some((message: any) => typeof message.content === "string" && message.content.includes("elided")) + ).toBe(true); + }); + it("手动 compact:空消息时返回错误", async () => { const { service, mockRepo } = createTestService(); const { sender, sentMessages } = createMockSender(); @@ -145,7 +234,7 @@ describe("Compact 功能", () => { // 第一次 LLM 调用:返回文本但 inputTokens 超过 80% (110000/128000 ≈ 86%) fetchSpy.mockResolvedValueOnce( makeSSEResponse([ - `data: {"choices":[{"delta":{"content":"Some response"}}]}\n\n`, + `data: {"choices":[{"delta":{"content":"Some response"},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":110000,"completion_tokens":100}}\n\n`, ]) ); @@ -178,7 +267,7 @@ describe("Compact 功能", () => { // inputTokens = 50000, contextWindow(gpt-4o) = 128000, 39% < 80% fetchSpy.mockResolvedValueOnce( makeSSEResponse([ - `data: {"choices":[{"delta":{"content":"OK"}}]}\n\n`, + `data: {"choices":[{"delta":{"content":"OK"},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":50000,"completion_tokens":5}}\n\n`, ]) ); diff --git a/src/app/service/agent/service_worker/background.test.ts b/src/app/service/agent/service_worker/background.test.ts index 82db11843..d0751c252 100644 --- a/src/app/service/agent/service_worker/background.test.ts +++ b/src/app/service/agent/service_worker/background.test.ts @@ -1,5 +1,21 @@ import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; -import { createTestService, makeTextResponse, createRunningConversation } from "./test-helpers"; +import { createTestService, makeTextResponse, makeSSEResponse, createRunningConversation } from "./test-helpers"; + +// 与 llm.test.ts 中的同名辅助函数一致:构造带 tool_call 的 OpenAI SSE 响应 +function makeToolCallResponse(toolCalls: Array<{ id: string; name: string; arguments: string }>): Response { + const chunks: string[] = []; + toolCalls.forEach((tc, i) => { + chunks.push( + `data: {"choices":[{"delta":{"tool_calls":[{"id":"${tc.id}","function":{"name":"${tc.name}","arguments":""}}]}}]}\n\n` + ); + const isLast = i === toolCalls.length - 1; + chunks.push( + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":${JSON.stringify(tc.arguments)}}}]}${isLast ? ', "finish_reason":"tool_calls"' : ""}}]}\n\n` + ); + }); + chunks.push(`data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`); + return makeSSEResponse(chunks); +} // ---- updateStreamingState 快照状态管理 ---- @@ -312,6 +328,55 @@ describe("handleAttachToConversation 重连逻辑", () => { (service as any).bgSessionManager.delete("conv-done"); }); + it("cancelling 阶段 attach:sync 仍报 running 并继续订阅 listener,直到真正的终态事件", async () => { + const { service } = createTestService(); + const rc = createRunningConversation({ status: "cancelling" }); + (service as any).bgSessionManager.set("conv-cancelling", rc); + + const { sender, sentMessages } = createMockSender(); + await (service as any).handleAttachToConversation({ conversationId: "conv-cancelling" }, sender); + + const syncEvent = sentMessages.find((m: any) => m.action === "event" && m.data.type === "sync"); + // 不能报 error:那会让客户端立即断开重载,永远收不到 orchestrator 后续广播的真正终态事件 + expect(syncEvent.data.status).toBe("running"); + expect(rc.listeners.size).toBe(1); + + // 之后 orchestrator 广播真正的终态事件时,这个迟到的 listener 应该能收到 + (service as any).bgSessionManager.broadcastEvent(rc, { + type: "error", + message: "Conversation cancelled", + errorCode: "cancelled", + usage: { inputTokens: 1, outputTokens: 2 }, + }); + const terminalEvent = sentMessages.find((m: any) => m.data?.errorCode === "cancelled"); + expect(terminalEvent).toBeDefined(); + expect(terminalEvent.data.usage).toBeDefined(); + + (service as any).bgSessionManager.delete("conv-cancelling"); + }); + + it("finalizeCancelled 在 status 已被终态事件改写为 error 后仍应调度清理", async () => { + vi.useFakeTimers(); + try { + const { service } = createTestService(); + const rc = createRunningConversation({ status: "cancelling" }); + (service as any).bgSessionManager.set("conv-late-error", rc); + + // 模拟 emitCancelled() 广播的终态事件已经先经过 sendEvent → updateStreamingState + // 把 status 从 cancelling 改写成 error(真实时序:chat_service.ts 的 sendEvent + // 总是先 updateStreamingState 再 finalizeCancelled 才被调用) + (service as any).bgSessionManager.updateStreamingState(rc, { type: "error", message: "x" }); + expect(rc.status).toBe("error"); + + (service as any).bgSessionManager.finalizeCancelled("conv-late-error", rc); + + await vi.advanceTimersByTimeAsync(30_000); + expect((service as any).bgSessionManager.get("conv-late-error")).toBeUndefined(); + } finally { + vi.useRealTimers(); + } + }); + it("运行中的会话添加 listener,断开时移除", async () => { const { service } = createTestService(); const rc = createRunningConversation({ status: "running" }); @@ -329,19 +394,21 @@ describe("handleAttachToConversation 重连逻辑", () => { (service as any).bgSessionManager.delete("conv-run"); }); - it("通过 attach 发送 askUserResponse 能正确 resolve", async () => { + it("通过 attach 发送 askUserResponse 能正确 resolve,且只广播一次 ask_user_resolved", async () => { const { service } = createTestService(); const rc = createRunningConversation({ status: "running" }); - // 注册一个 askResolver + // 真实的 resolver(ask_user.ts / askUserForGuard)在 resolve 时会自行广播终态事件, + // attach 不应该再额外广播一次,否则同一次回答会产生两条 ask_user_resolved let resolvedAnswer: string | undefined; rc.askResolvers.set("ask-1", (answer: string) => { resolvedAnswer = answer; + (service as any).bgSessionManager.broadcastEvent(rc, { type: "ask_user_resolved", id: "ask-1" }); }); rc.pendingAskUser = { id: "ask-1", question: "选择" }; (service as any).bgSessionManager.set("conv-ask", rc); - const { sender, simulateMessage } = createMockSender(); + const { sender, sentMessages, simulateMessage } = createMockSender(); await (service as any).handleAttachToConversation({ conversationId: "conv-ask" }, sender); // 模拟 UI 回复 ask_user @@ -350,11 +417,13 @@ describe("handleAttachToConversation 重连逻辑", () => { expect(resolvedAnswer).toBe("红色"); expect(rc.pendingAskUser).toBeUndefined(); expect(rc.askResolvers.has("ask-1")).toBe(false); + const resolvedEvents = sentMessages.filter((message) => message.data?.type === "ask_user_resolved"); + expect(resolvedEvents).toHaveLength(1); (service as any).bgSessionManager.delete("conv-ask"); }); - it("多个 listener 回复同一 ask_user 只有第一个生效", async () => { + it("多个 listener 回复同一 ask_user 只有第一个生效,且每个 listener 只收到一次 ask_user_resolved", async () => { const { service } = createTestService(); const rc = createRunningConversation({ status: "running" }); @@ -363,12 +432,13 @@ describe("handleAttachToConversation 重连逻辑", () => { rc.askResolvers.set("ask-1", (answer: string) => { resolveCount++; lastAnswer = answer; + (service as any).bgSessionManager.broadcastEvent(rc, { type: "ask_user_resolved", id: "ask-1" }); }); rc.pendingAskUser = { id: "ask-1", question: "选择" }; (service as any).bgSessionManager.set("conv-multi", rc); - const { sender: sender1, simulateMessage: sim1 } = createMockSender(); - const { sender: sender2, simulateMessage: sim2 } = createMockSender(); + const { sender: sender1, sentMessages: sent1, simulateMessage: sim1 } = createMockSender(); + const { sender: sender2, sentMessages: sent2, simulateMessage: sim2 } = createMockSender(); await (service as any).handleAttachToConversation({ conversationId: "conv-multi" }, sender1); await (service as any).handleAttachToConversation({ conversationId: "conv-multi" }, sender2); @@ -378,23 +448,65 @@ describe("handleAttachToConversation 重连逻辑", () => { expect(resolveCount).toBe(1); expect(lastAnswer).toBe("第一个"); + expect(sent1.filter((message) => message.data?.type === "ask_user_resolved")).toHaveLength(1); + expect(sent2.filter((message) => message.data?.type === "ask_user_resolved")).toHaveLength(1); (service as any).bgSessionManager.delete("conv-multi"); }); - it("通过 attach 发送 stop 能中止会话", async () => { - const { service } = createTestService(); - const rc = createRunningConversation({ status: "running" }); - (service as any).bgSessionManager.set("conv-stop", rc); - - const { sender, simulateMessage } = createMockSender(); - await (service as any).handleAttachToConversation({ conversationId: "conv-stop" }, sender); + it("通过 attach 发送 stop 只置为 cancelling 并 abort,不自行广播终态(终态事件唯一来源是执行方 emitCancelled)", async () => { + vi.useFakeTimers(); + try { + const { service } = createTestService(); + const rc = createRunningConversation({ + status: "running", + pendingAskUser: { id: "ask-1", question: "继续吗" }, + }); + rc.askResolvers.set("ask-1", vi.fn()); + (service as any).bgSessionManager.set("conv-stop", rc); + + const { sender, sentMessages, simulateMessage } = createMockSender(); + await (service as any).handleAttachToConversation({ conversationId: "conv-stop" }, sender); + + simulateMessage({ action: "stop" }); + + expect(rc.abortController.signal.aborted).toBe(true); + // stop() 只置为 cancelling:真正的终态由持有该 rc 的执行方 promise 落定后写入, + // 避免同一 conversationId 在旧执行尚未退出时就被新会话顶替 + expect(rc.status).toBe("cancelling"); + expect(rc.pendingAskUser).toBeUndefined(); + expect(rc.askResolvers.size).toBe(0); + expect((service as any).bgSessionManager.has("conv-stop")).toBe(true); + // cancelling 现在也纳入发现列表:UI 侧(Options 页面刷新/重连)需要据此判断是否 + // 应该继续 attach 等待真正的终态事件,而不是误判为"已经不在运行" + expect(service.getRunningConversationIds()).toContain("conv-stop"); + // stop() 本身不再广播任何终态事件:唯一的终态事件来自执行方(orchestrator 的 + // emitCancelled,走正常 sendEvent → updateStreamingState → broadcastEvent 路径), + // 避免"先广播一条不带 usage 的事件,UI 断开后丢失后到的真实终态事件"的竞态 + expect(sentMessages.some((message) => message.data?.type === "error")).toBe(false); + + // 模拟执行方在 abort 落定后调用 finalizeCancelled + (service as any).bgSessionManager.finalizeCancelled("conv-stop", rc); + expect(rc.status).toBe("error"); + expect((service as any).bgSessionManager.has("conv-stop")).toBe(false); + + await vi.advanceTimersByTimeAsync(30_000); + expect((service as any).bgSessionManager.get("conv-stop")).toBeUndefined(); + } finally { + vi.useRealTimers(); + } + }); - simulateMessage({ action: "stop" }); + it("stop() 传入 expectedRc 与当前会话实例不符时应忽略,避免旧连接的延迟 Stop 误伤新会话", async () => { + const { service } = createTestService(); + const staleRc = createRunningConversation({ status: "running" }); + const currentRc = createRunningConversation({ status: "running" }); + (service as any).bgSessionManager.set("conv-race", currentRc); - expect(rc.abortController.signal.aborted).toBe(true); + (service as any).bgSessionManager.stop("conv-race", staleRc); - (service as any).bgSessionManager.delete("conv-stop"); + expect(currentRc.abortController.signal.aborted).toBe(false); + expect(currentRc.status).toBe("running"); }); it("空 streamingState 的 sync 不包含 streamingMessage 字段", async () => { @@ -494,7 +606,7 @@ describe("后台运行会话 集成测试", () => { if (readCalled === 1) { return { done: false, - value: encoder.encode(`data: {"choices":[{"delta":{"content":"hello"}}]}\n\n`), + value: encoder.encode(`data: {"choices":[{"delta":{"content":"hello"},"finish_reason":"stop"}]}\n\n`), }; } if (readCalled === 2) { @@ -595,11 +707,11 @@ describe("后台运行会话 集成测试", () => { expect(events.some((e: any) => e.type === "done")).toBe(true); }); - it("后台模式:stop 指令中止会话后不抛未捕获异常", async () => { + it("后台模式:stop 指令中止会话后会广播终态而且不抛未捕获异常", async () => { const { service, mockRepo } = createTestService(); setupConversation(mockRepo); - const { sender, simulateMessage } = createMockSender(); + const { sender, sentMessages, simulateMessage } = createMockSender(); fetchSpy.mockImplementation(async (_url: any, init: any) => { if (init?.signal?.aborted) { @@ -621,7 +733,59 @@ describe("后台运行会话 集成测试", () => { simulateMessage({ action: "stop" }); await chatPromise; - // 不应抛异常 + const rc = (service as any).bgSessionManager.get("conv-bg"); + expect(rc.status).toBe("error"); + expect( + sentMessages.some((message) => message.data?.type === "error" && message.data.errorCode === "cancelled") + ).toBe(true); + }); + + it("循环检测升级提问超时(5 分钟无人应答)后,应清除 rc.pendingAskUser,避免后续 attach 看到过期提问", async () => { + vi.useFakeTimers(); + try { + const { service, mockRepo } = createTestService(); + setupConversation(mockRepo); + const { sender, sentMessages } = createMockSender(); + + const registry = (service as any).toolRegistry; + registry.registerBuiltin( + { name: "dup", description: "dup", parameters: { type: "object", properties: {} } }, + { execute: async () => "ok" } + ); + + // 连续 4 轮相同参数调用同一工具,触发两次循环检测告警(第 2、4 轮), + // 第二次告警会暂停并调用 askUserForGuard;第 5 轮永不 resolve, + // 用于在会话真正结束(发出 done 事件)之前观测超时后的瞬时状态—— + // 如果只等到 done 事件才检查,done 处理器本身就会清掉 pendingAskUser, + // 从而掩盖“超时未主动清除”这个问题 + fetchSpy + .mockResolvedValueOnce(makeToolCallResponse([{ id: "c1", name: "dup", arguments: "{}" }])) + .mockResolvedValueOnce(makeToolCallResponse([{ id: "c2", name: "dup", arguments: "{}" }])) + .mockResolvedValueOnce(makeToolCallResponse([{ id: "c3", name: "dup", arguments: "{}" }])) + .mockResolvedValueOnce(makeToolCallResponse([{ id: "c4", name: "dup", arguments: "{}" }])) + .mockReturnValueOnce(new Promise(() => {})); + + void (service as any).handleConversationChat( + { conversationId: "conv-bg", message: "test", background: true }, + sender + ); + + // 推进 5 分钟,让 askUserForGuard 的超时触发(默认按 Continue 继续), + // 此时第 5 轮请求已发出但永不 resolve,会话仍处于 running,尚未发出 done 事件 + await vi.advanceTimersByTimeAsync(5 * 60 * 1000); + + const rc = (service as any).bgSessionManager.get("conv-bg"); + expect(rc).toBeDefined(); + expect(rc.status).toBe("running"); + expect(rc.pendingAskUser).toBeUndefined(); + expect(sentMessages.some((message) => message.data?.type === "ask_user_expired")).toBe(true); + // 超时属于“过期”而非“已回答”:不应在发出 ask_user_expired 后又紧接着发出 ask_user_resolved + expect(sentMessages.some((message) => message.data?.type === "ask_user_resolved")).toBe(false); + + registry.unregisterBuiltin("dup"); + } finally { + vi.useRealTimers(); + } }); it("getRunningConversationIds 返回正确的 ID 列表", () => { @@ -633,9 +797,9 @@ describe("后台运行会话 集成测试", () => { (service as any).bgSessionManager.set("conv-2", { status: "done" }); const ids = service.getRunningConversationIds(); - expect(ids).toHaveLength(2); + expect(ids).toHaveLength(1); expect(ids).toContain("conv-1"); - expect(ids).toContain("conv-2"); + expect(ids).not.toContain("conv-2"); (service as any).bgSessionManager.delete("conv-1"); (service as any).bgSessionManager.delete("conv-2"); diff --git a/src/app/service/agent/service_worker/background_session_manager.ts b/src/app/service/agent/service_worker/background_session_manager.ts index de92f558e..1378502fa 100644 --- a/src/app/service/agent/service_worker/background_session_manager.ts +++ b/src/app/service/agent/service_worker/background_session_manager.ts @@ -10,13 +10,25 @@ export type ListenerEntry = { // 后台运行会话状态 export type RunningConversation = { conversationId: string; + // 该次运行绑定的会话 generation;attach() 的调用方必须持有同一 generation 才允许附加, + // 否则会静默观察到删除重建后无关的新一代会话 + generation: string; abortController: AbortController; listeners: Set; streamingState: { content: string; thinking: string; toolCalls: ToolCall[] }; - pendingAskUser?: { id: string; question: string; options?: string[]; multiple?: boolean }; + pendingAskUser?: { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + }; askResolvers: Map void>; tasks: Array<{ id: string; subject: string; status: "pending" | "in_progress" | "completed"; description?: string }>; - status: "running" | "done" | "error"; + // cancelling:stop() 已触发但执行方尚未真正退出(promise 未 settle); + // 在此期间该 conversationId 仍视为"占用中",避免同 ID 的替换会话被过早放行。 + status: "running" | "cancelling" | "done" | "error"; }; // 后台会话注册表:管理流式状态快照、listener 广播、UI 附加逻辑 @@ -24,7 +36,8 @@ export class BackgroundSessionManager { private runningConversations = new Map(); has(conversationId: string): boolean { - return this.runningConversations.has(conversationId); + const status = this.runningConversations.get(conversationId)?.status; + return status === "running" || status === "cancelling"; } get(conversationId: string): RunningConversation | undefined { @@ -39,8 +52,13 @@ export class BackgroundSessionManager { this.runningConversations.delete(conversationId); } + // cancelling 也算在内:handleAttach() 已经支持对 cancelling 会话继续订阅直到真正的终态 + // 事件,但如果这里不把 cancelling 纳入发现列表,UI 侧(Options 页面刷新/重连)永远不会 + // 尝试 attach,也就永远等不到那条真正携带取消原因/usage 的终态事件 listIds(): string[] { - return Array.from(this.runningConversations.keys()); + return Array.from(this.runningConversations.entries()) + .filter(([, conversation]) => conversation.status === "running" || conversation.status === "cancelling") + .map(([conversationId]) => conversationId); } // 更新后台会话的流式状态快照 @@ -89,7 +107,7 @@ export class BackgroundSessionManager { case "tool_call_complete": { const tc = rc.streamingState.toolCalls.find((t) => t.id === event.id); if (tc) { - tc.status = "completed"; + tc.status = event.status || "completed"; tc.result = event.result; tc.attachments = event.attachments; } @@ -104,9 +122,17 @@ export class BackgroundSessionManager { id: event.id, question: event.question, options: event.options, + optionValues: event.optionValues, multiple: event.multiple, + allowCustom: event.allowCustom, }; break; + case "ask_user_expired": + if (rc.pendingAskUser?.id === event.id) rc.pendingAskUser = undefined; + break; + case "ask_user_resolved": + if (rc.pendingAskUser?.id === event.id) rc.pendingAskUser = undefined; + break; case "task_update": rc.tasks = event.tasks; break; @@ -132,8 +158,40 @@ export class BackgroundSessionManager { } } + // 停止后台会话:仅置为 cancelling(占用态,阻止同 ID 会话被过早顶替)并 abort, + // 不在这里广播终态事件。终态事件的唯一来源是执行方(orchestrator 的 emitCancelled, + // 见 tool_loop_orchestrator.ts)在 promise 真正落定后发出的那一条——它携带完整的 + // 累计 usage/耗时。此前 stop() 自己先广播一条不带 usage 的 error:cancelled,会导致 UI + // 连接在收到 orchestrator 发出的、携带真实 usage 的终态事件之前就已断开丢弃它, + // 也会让 finalizeCancelled() 因为 status 已被后到的 error 事件提前改写而永远不触发清理。 + // expectedRc 用于校验调用方持有的会话实例仍是当前会话,防止旧连接的延迟 Stop 误伤已顶替上位的新会话。 + stop(conversationId: string, expectedRc?: RunningConversation): void { + const rc = this.runningConversations.get(conversationId); + if (!rc || (expectedRc && rc !== expectedRc)) return; + if (rc.status !== "running") return; + + rc.status = "cancelling"; + rc.pendingAskUser = undefined; + rc.askResolvers.clear(); + rc.abortController.abort(); + } + + // 执行方在 abort 落定、promise 真正 settle 后调用,把 cancelling 收敛为终态并调度清理。 + // 幂等:emitCancelled() 广播的终态事件会先经过正常的 sendEvent → updateStreamingState + // 把 status 从 cancelling 改写成 error,等这里再执行时 status 已经不是 cancelling 了—— + // 之前的实现只认 cancelling,导致这个真实存在的时序下 cleanupIfDone() 永远不会被调度, + // 记录永久留在 runningConversations 里。这里改为:cancelling 或已经落定的终态都视为 + // 可以安全调度清理(cleanupIfDone 自身按实例比对 + status!=="running" 做二次确认,幂等安全)。 + // 通过实例比对避免误将已被新会话顶替的 map 条目错误终态化。 + finalizeCancelled(conversationId: string, rc: RunningConversation): void { + if (this.runningConversations.get(conversationId) !== rc) return; + if (rc.status === "running") return; + if (rc.status === "cancelling") rc.status = "error"; + this.cleanupIfDone(conversationId); + } + // 附加 UI 连接到后台运行中的会话(同步快照 + listener + askUser resolver + stop) - async handleAttach(params: { conversationId: string }, sender: IGetSender): Promise { + async handleAttach(params: { conversationId: string; generation?: string }, sender: IGetSender): Promise { if (!sender.isType(GetSenderType.CONNECT)) { throw new Error("attachToConversation requires connect mode"); } @@ -151,6 +209,13 @@ export class BackgroundSessionManager { return; } + // 调用方持有的 generation 与实际运行中的会话不一致:会话已被删除重建, + // 不能让旧一代的调用方附加到无关的新一代会话上 + if (params.generation !== undefined && rc.generation !== params.generation) { + sendEvent({ type: "sync", tasks: [], status: "done" }); + return; + } + // 发送 sync 快照 const syncEvent: ChatStreamEvent = { type: "sync", @@ -164,16 +229,21 @@ export class BackgroundSessionManager { : undefined, pendingAskUser: rc.pendingAskUser, tasks: rc.tasks, - status: rc.status, + // cancelling 还没有产生真正携带取消原因/usage/耗时的终态事件(那条事件由 orchestrator 的 + // emitCancelled() 在 promise 落定后才广播,见 tool_loop_orchestrator.ts)。 + // 之前这里直接报 "error" 会让客户端立即断开重载,永远收不到那条真正完整的终态事件, + // 也可能在取消记录落库前就重载;改为对外仍按 "running" 处理并继续订阅 listener, + // 客户端会一直等到下面广播的那条真正的终态事件再断开。 + status: rc.status === "cancelling" ? "running" : rc.status, }; sendEvent(syncEvent); - // 已完成则不需要添加 listener - if (rc.status !== "running") { + // 只有已产生真正终态事件(done/error)才不需要 listener;cancelling 必须继续订阅 + // 直到 orchestrator 广播出那条真正的终态事件(见上面的注释) + if (rc.status !== "running" && rc.status !== "cancelling") { return; } - // 添加 listener const listener: ListenerEntry = { sendEvent }; rc.listeners.add(listener); @@ -184,11 +254,13 @@ export class BackgroundSessionManager { if (resolver) { rc.askResolvers.delete(msg.data.id); rc.pendingAskUser = undefined; + // resolver 自身(ask_user.ts / askUserForGuard)负责广播其终态事件, + // 这里不再重复广播,否则同一次回答会产生两条 ask_user_resolved resolver(msg.data.answer); } } if (msg.action === "stop") { - rc.abortController.abort(); + this.stop(params.conversationId, rc); } }); diff --git a/src/app/service/agent/service_worker/chat.test.ts b/src/app/service/agent/service_worker/chat.test.ts index f37dd5a80..97ff823b6 100644 --- a/src/app/service/agent/service_worker/chat.test.ts +++ b/src/app/service/agent/service_worker/chat.test.ts @@ -24,6 +24,7 @@ describe("handleConversationChat skipSaveUserMessage", () => { const sender = { isType: (type: any) => type === 1, // GetSenderType.CONNECT getConnect: () => mockConn, + getSender: () => ({ url: chrome.runtime.getURL("src/options.html#/agent/chat") }), }; return { sender, sentMessages }; } @@ -64,6 +65,26 @@ describe("handleConversationChat skipSaveUserMessage", () => { expect(userCall![0].content).toBe("你好"); }); + it("仅 UI 明确声明的新上传附件应随用户消息持久化所有权", async () => { + const { service, mockRepo } = createTestService(); + const { sender } = createMockSender(); + mockRepo.listConversations.mockResolvedValue([BASE_CONV]); + mockRepo.getMessages.mockResolvedValue([]); + fetchSpy.mockResolvedValueOnce(makeTextResponse("收到")); + + await (service as any).handleConversationChat( + { + conversationId: "conv-1", + message: [{ type: "image", attachmentId: "upload.png", mimeType: "image/png" }], + ownedAttachmentIds: ["upload.png"], + }, + sender + ); + + const userMessage = mockRepo.appendMessage.mock.calls.find((call: any[]) => call[0].role === "user")?.[0]; + expect(userMessage.ownedAttachmentIds).toEqual(["upload.png"]); + }); + it("【bug 回归】skipSaveUserMessage=true:用户消息不应再次保存到 storage", async () => { const { service, mockRepo } = createTestService(); const { sender } = createMockSender(); @@ -185,6 +206,36 @@ describe("handleConversationChat skipSaveUserMessage", () => { }); }); +describe("userscript 会话工具隔离", () => { + it("携带 scriptUuid 时不注册无法交互的 ask_user 工具", async () => { + const { service } = createTestService(); + const chatService = (service as any).chatService; + const result = await chatService.buildSessionToolRegistry({ + conv: { + id: "conv-script", + title: "Script", + modelId: "test-openai", + createtime: 1, + updatetime: 1, + }, + model: { + id: "test-openai", + name: "Test", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", + }, + params: { conversationId: "conv-script", message: "hi", scriptUuid: "script-1" }, + sendEvent: vi.fn(), + abortController: new AbortController(), + askResolvers: new Map(), + }); + + expect(result.sessionRegistry.getDefinitions().some((tool: any) => tool.name === "ask_user")).toBe(false); + }); +}); + // ---- handleConversationChat 场景补充 ---- describe("handleConversationChat 场景补充", () => { @@ -208,6 +259,7 @@ describe("handleConversationChat 场景补充", () => { const sender = { isType: (type: any) => type === 1, getConnect: () => mockConn, + getSender: () => ({ url: chrome.runtime.getURL("src/options.html#/agent/chat") }), }; return { sender, sentMessages }; } @@ -321,12 +373,189 @@ describe("handleConversationChat 场景补充", () => { mockRepo.listConversations.mockResolvedValue([]); // 空 - await (service as any).handleConversationChat({ conversationId: "not-exist", message: "hi" }, sender); + await (service as any).handleConversationChat( + { conversationId: "not-exist", message: "hi", ownedAttachmentIds: ["provisional.png"] }, + sender + ); const events = sentMessages.map((m) => m.data); const errorEvents = events.filter((e: any) => e.type === "error"); expect(errorEvents).toHaveLength(1); expect(errorEvents[0].message).toContain("Conversation not found"); + expect(mockRepo.deleteAttachment).toHaveBeenCalledWith("provisional.png"); + }); + + it("调用方持有的 generation 与当前存储不一致时应拒绝 chat,而不是作用于新一代会话", async () => { + const { service, mockRepo } = createTestService(); + const { sender, sentMessages } = createMockSender(); + + // conv-1 的 ID 被删除重建,当前存储的 generation 已经变成 "gen-b" + const conv = { + id: "conv-1", + title: "Test", + modelId: "test-openai", + generation: "gen-b", + createtime: Date.now(), + updatetime: Date.now(), + }; + mockRepo.listConversations.mockResolvedValue([conv]); + mockRepo.getMessages.mockResolvedValue([]); + + // 陈旧的 ConversationInstance 仍持有创建时的 generation "gen-a" + await (service as any).handleConversationChat( + { conversationId: "conv-1", generation: "gen-a", message: "hi" }, + sender + ); + + const events = sentMessages.map((m) => m.data); + const errorEvents = events.filter((e: any) => e.type === "error"); + expect(errorEvents).toHaveLength(1); + expect(errorEvents[0].errorCode).toBe("conversation_generation_mismatch"); + // 不应触发任何 LLM 调用或持久化写入 + expect(fetchSpy).not.toHaveBeenCalled(); + expect(mockRepo.appendMessage).not.toHaveBeenCalled(); + }); + + it("compact 模式下调用方持有的 generation 与当前存储不一致时也应拒绝,而不是压缩新一代会话的历史", async () => { + const { service, mockRepo } = createTestService(); + const { sender, sentMessages } = createMockSender(); + + // conv-1 的 ID 被删除重建,当前存储的 generation 已经变成 "gen-b" + const conv = { + id: "conv-1", + title: "Test", + modelId: "test-openai", + generation: "gen-b", + createtime: Date.now(), + updatetime: Date.now(), + }; + mockRepo.listConversations.mockResolvedValue([conv]); + + // 陈旧的 Options 标签页仍持有创建时的 generation "gen-a" + await (service as any).handleConversationChat( + { conversationId: "conv-1", generation: "gen-a", message: "", compact: true }, + sender + ); + + const events = sentMessages.map((m) => m.data); + const errorEvents = events.filter((e: any) => e.type === "error"); + expect(errorEvents).toHaveLength(1); + expect(errorEvents[0].errorCode).toBe("conversation_generation_mismatch"); + // 不应读取消息快照或触发任何 LLM 调用 + expect(mockRepo.getMessageSnapshot).not.toHaveBeenCalled(); + expect(fetchSpy).not.toHaveBeenCalled(); + }); + + it("handleConversation 的 getMessages/clearMessages 在 generation 不一致时应拒绝而非作用于新一代会话", async () => { + const { service, mockRepo } = createTestService(); + + const conv = { + id: "conv-1", + title: "Test", + modelId: "test-openai", + generation: "gen-b", + createtime: Date.now(), + updatetime: Date.now(), + }; + mockRepo.listConversations.mockResolvedValue([conv]); + // 模拟真实 repo 行为:generation 提供且与当前存储不一致时拒绝 + mockRepo.getMessageSnapshot.mockImplementation(async (_conversationId: string, generation?: string) => { + if (generation !== undefined && generation !== "gen-b") { + throw new Error(`Conversation "conv-1" changed or was deleted`); + } + return { generation: "gen-b", revision: 3, messages: [] }; + }); + + await expect( + (service as any).handleConversation({ + action: "getMessages", + conversationId: "conv-1", + generation: "gen-a", + scriptUuid: "script-1", + }) + ).rejects.toThrow(); + + await expect( + (service as any).handleConversation({ + action: "clearMessages", + conversationId: "conv-1", + generation: "gen-a", + scriptUuid: "script-1", + }) + ).rejects.toThrow(); + + // 拒绝路径不应触及持久化写入 + expect(mockRepo.saveMessages).not.toHaveBeenCalled(); + }); + + it("用户消息 append 报错但确认读证实已落盘时,不应删除刚上传的附件", async () => { + const { service, mockRepo } = createTestService(); + const { sender } = createMockSender(); + + const conv = { + id: "conv-1", + title: "Test", + modelId: "test-openai", + generation: "gen-1", + createtime: Date.now(), + updatetime: Date.now(), + }; + mockRepo.listConversations.mockResolvedValue([conv]); + mockRepo.getMessages.mockResolvedValue([]); + fetchSpy.mockResolvedValueOnce(makeTextResponse("收到")); + + let appended: any; + mockRepo.appendMessage.mockImplementationOnce(async (message: any) => { + // 写入其实已经落盘(模拟 OPFS close 报告二义性错误前已经 commit) + appended = message; + throw new Error("ambiguous close failure"); + }); + mockRepo.getMessageSnapshot.mockImplementation(async () => ({ + generation: "gen-1", + revision: 1, + messages: appended ? [appended] : [], + })); + + await (service as any).handleConversationChat( + { + conversationId: "conv-1", + message: [{ type: "image", attachmentId: "upload.png", mimeType: "image/png" }], + ownedAttachmentIds: ["upload.png"], + }, + sender + ); + + // 已确认落盘:不能把二义性错误当作"未持久化"从而删除消息实际引用的附件 + expect(mockRepo.deleteAttachment).not.toHaveBeenCalledWith("upload.png"); + }); + + it("用户消息 append 报错且确认读也失败时,不应删除可能已被消息引用的附件", async () => { + const { service, mockRepo } = createTestService(); + const { sender } = createMockSender(); + mockRepo.listConversations.mockResolvedValue([ + { + id: "conv-1", + title: "Test", + modelId: "test-openai", + generation: "gen-1", + createtime: Date.now(), + updatetime: Date.now(), + }, + ]); + mockRepo.getMessages.mockResolvedValue([]); + mockRepo.appendMessage.mockRejectedValueOnce(new Error("ambiguous close failure")); + mockRepo.getMessageSnapshot.mockRejectedValueOnce(new Error("confirmation read failed")); + + await (service as any).handleConversationChat( + { + conversationId: "conv-1", + message: [{ type: "image", attachmentId: "upload-uncertain.png", mimeType: "image/png" }], + ownedAttachmentIds: ["upload-uncertain.png"], + }, + sender + ); + + expect(mockRepo.deleteAttachment).not.toHaveBeenCalledWith("upload-uncertain.png"); }); it("skill 预加载:历史消息含 load_skill 调用时预执行以标记 skill 已加载", async () => { @@ -508,3 +737,734 @@ describe.concurrent("handleModelApi", () => { ); }); }); + +// ---- 脚本工具回调:abort/disconnect/超时应结束等待,而不是永远挂起 ---- + +describe("scriptToolCallback 的 abort/disconnect/超时处理", () => { + let fetchSpy: ReturnType; + + beforeEach(() => { + fetchSpy = vi.spyOn(globalThis, "fetch"); + }); + + afterEach(() => { + fetchSpy.mockRestore(); + }); + + // 带 message/disconnect 模拟回调的 mock sender(脚本工具不会主动回复 toolResults) + function createMockSender() { + const sentMessages: any[] = []; + let messageHandler: ((msg: any) => void) | null = null; + let disconnectHandler: (() => void) | null = null; + const mockConn = { + sendMessage: (msg: any) => sentMessages.push(msg), + onMessage: vi.fn((handler: any) => { + messageHandler = handler; + }), + onDisconnect: vi.fn((handler: any) => { + disconnectHandler = handler; + }), + }; + const sender = { + isType: (type: any) => type === 1, + getConnect: () => mockConn, + getSender: () => ({ url: chrome.runtime.getURL("src/options.html#/agent/chat") }), + }; + return { + sender, + sentMessages, + simulateMessage: (msg: any) => messageHandler?.(msg), + simulateDisconnect: () => disconnectHandler?.(), + }; + } + + // 构造带脚本自定义工具调用(非内置工具,走 scriptCallback)的 OpenAI SSE 响应 + function makeScriptToolCallResponse(toolId: string, toolName: string, args: string): Response { + const encoder = new TextEncoder(); + const chunks = [ + `data: {"choices":[{"delta":{"tool_calls":[{"id":"${toolId}","function":{"name":"${toolName}","arguments":""}}]}}]}\n\n`, + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":${JSON.stringify(args)}}}]}, "finish_reason":"tool_calls"}]}\n\n`, + `data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`, + ]; + let i = 0; + return { + ok: true, + status: 200, + body: { + getReader() { + return { + async read() { + if (i < chunks.length) return { done: false, value: encoder.encode(chunks[i++]) }; + return { done: true, value: undefined }; + }, + cancel: async () => {}, + }; + }, + }, + } as unknown as Response; + } + + function setupConversation(mockRepo: any) { + mockRepo.listConversations.mockResolvedValue([ + { id: "conv-script", title: "Chat", modelId: "test-openai", createtime: Date.now(), updatetime: Date.now() }, + ]); + mockRepo.getMessages.mockResolvedValue([]); + } + + it("脚本工具返回结果时,应继续工具循环并完成对话", async () => { + const { service, mockRepo } = createTestService(); + setupConversation(mockRepo); + const { sender, sentMessages, simulateMessage } = createMockSender(); + + fetchSpy + .mockResolvedValueOnce(makeScriptToolCallResponse("call-success", "my_tool", "{}")) + .mockResolvedValueOnce(makeTextResponse("最终回复")); + + const chatPromise = (service as any).handleConversationChat( + { + conversationId: "conv-script", + message: "使用工具", + tools: [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} } }], + }, + sender + ); + + await vi.waitFor(() => { + expect(sentMessages.some((message) => message.action === "executeTools")).toBe(true); + }); + const executeMessage = sentMessages.find((message) => message.action === "executeTools"); + simulateMessage({ + action: "toolResults", + requestId: executeMessage.requestId, + data: [{ id: "call-success", result: "ok" }], + }); + + await expect(chatPromise).resolves.toBeUndefined(); + expect(sentMessages.some((message) => message.data?.type === "done")).toBe(true); + }); + + it("非后台会话断开(触发 abort)时,等待中的脚本工具调用应立即结束,而不是永远挂起", async () => { + const { service, mockRepo } = createTestService(); + setupConversation(mockRepo); + const { sender, simulateDisconnect } = createMockSender(); + + fetchSpy.mockResolvedValueOnce(makeScriptToolCallResponse("call-1", "my_tool", "{}")); + + const chatPromise = (service as any).handleConversationChat( + { + conversationId: "conv-script", + message: "使用工具", + tools: [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} } }], + }, + sender + ); + + // 等待 SSE 响应被消费、executeTools 消息发出,进入"等待 toolResults"状态 + await new Promise((r) => setTimeout(r, 20)); + simulateDisconnect(); + + // 若脚本工具回调未结束,chatPromise 永远不会 resolve,下面的 race 会超时 + const TIMEOUT = Symbol("timeout"); + const result = await Promise.race([ + chatPromise.then(() => "done"), + new Promise((r) => setTimeout(() => r(TIMEOUT), 500)), + ]); + expect(result).toBe("done"); + }); + + it("脚本连接长时间无响应(超时)时,应主动结束该轮脚本工具调用而不是无限期等待", async () => { + vi.useFakeTimers(); + try { + const { service, mockRepo } = createTestService(); + setupConversation(mockRepo); + const { sender, sentMessages } = createMockSender(); + + fetchSpy + .mockResolvedValueOnce(makeScriptToolCallResponse("call-2", "my_tool", "{}")) + .mockResolvedValueOnce(makeTextResponse("最终回复")); + + const chatPromise = (service as any).handleConversationChat( + { + conversationId: "conv-script", + message: "使用工具", + tools: [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} } }], + }, + sender + ); + + // 推进到脚本工具调用的超时阈值(5 分钟),期间脚本从未回复 toolResults + await vi.advanceTimersByTimeAsync(5 * 60 * 1000); + await chatPromise; + + const doneEvents = sentMessages.filter((m) => m.data?.type === "done"); + expect(doneEvents).toHaveLength(1); + + // 超时批次必须向客户端发送带 requestId 的作废通知:客户端可能仍在串行执行 + // 该批次剩余 handler,不通知会让其副作用与下一批次交叠 + const executeMessage = sentMessages.find((m) => m.action === "executeTools"); + const cancelMessage = sentMessages.find((m) => m.action === "cancelToolBatch"); + expect(cancelMessage).toBeDefined(); + expect(cancelMessage.requestId).toBe(executeMessage.requestId); + } finally { + vi.useRealTimers(); + } + }); +}); + +describe("同一 conversationId 的 chat/compact/clear 必须串行执行", () => { + function createMockSender() { + const sentMessages: any[] = []; + const mockConn = { + sendMessage: (msg: any) => sentMessages.push(msg), + onMessage: vi.fn(), + onDisconnect: vi.fn(), + }; + const sender = { + isType: (type: any) => type === 1, + getConnect: () => mockConn, + }; + return { sender, sentMessages }; + } + + it("clearMessages 与另一个并发的 clearMessages 请求不应交叉执行(按 conversationId 排队)", async () => { + const { service, mockRepo } = createTestService(); + + const order: string[] = []; + mockRepo.saveMessages.mockImplementation(async (_id: string, _msgs: any[]) => { + order.push("start"); + // 人为延迟第一次调用,暴露"若未排队,第二次调用会在第一次完成前插入"的竞态 + await new Promise((r) => setTimeout(r, 20)); + order.push("end"); + }); + + const call1 = (service as any).handleConversation({ action: "clearMessages", conversationId: "conv-race" }); + const call2 = (service as any).handleConversation({ action: "clearMessages", conversationId: "conv-race" }); + + await Promise.all([call1, call2]); + + // 排队生效:必须是 start,end,start,end,而不是 start,start,end,end(交叉执行) + expect(order).toEqual(["start", "end", "start", "end"]); + }); + + it("不同 conversationId 之间不应互相阻塞排队", async () => { + const { service, mockRepo } = createTestService(); + + const order: string[] = []; + mockRepo.saveMessages.mockImplementation(async (id: string) => { + order.push(`start:${id}`); + await new Promise((r) => setTimeout(r, 20)); + order.push(`end:${id}`); + }); + + const call1 = (service as any).handleConversation({ action: "clearMessages", conversationId: "conv-a" }); + const call2 = (service as any).handleConversation({ action: "clearMessages", conversationId: "conv-b" }); + + await Promise.all([call1, call2]); + + // 不同会话应并发执行,两个 start 都先于任意一个 end 出现 + expect(order.indexOf("start:conv-a")).toBeLessThan(order.indexOf("end:conv-a")); + expect(order.indexOf("start:conv-b")).toBeLessThan(order.indexOf("end:conv-b")); + expect(order.slice(0, 2).sort()).toEqual(["start:conv-a", "start:conv-b"]); + }); + + it("正在进行的 chat 与随后到达的 clearMessages 不应交叉:clear 必须等 chat 落库完成", async () => { + const fetchSpy = vi.spyOn(globalThis, "fetch"); + try { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([ + { id: "conv-lock", title: "Chat", modelId: "test-openai", createtime: Date.now(), updatetime: Date.now() }, + ]); + mockRepo.getMessages.mockResolvedValue([]); + + const order: string[] = []; + mockRepo.appendMessage.mockImplementation(async () => { + order.push("chat-write"); + }); + mockRepo.saveMessages.mockImplementation(async () => { + order.push("clear-write"); + }); + + fetchSpy.mockImplementation(async () => { + // 模拟一次有延迟的 LLM 响应,给 clearMessages 制造"抢在 chat 落库前执行"的窗口 + await new Promise((r) => setTimeout(r, 20)); + return makeTextResponse("ok"); + }); + + const { sender } = createMockSender(); + const chatPromise = (service as any).handleConversationChat( + { conversationId: "conv-lock", message: "hi" }, + sender + ); + const clearPromise = (service as any).handleConversation({ + action: "clearMessages", + conversationId: "conv-lock", + }); + + await Promise.all([chatPromise, clearPromise]); + + // chat 的落库(appendMessage)必须先于随后到达的 clear 的落库(saveMessages)完成, + // 而不是被 clear 抢先覆盖掉正在写入的历史 + expect(order.indexOf("chat-write")).toBeLessThan(order.indexOf("clear-write")); + } finally { + fetchSpy.mockRestore(); + } + }); +}); + +describe("会话队列的连接感知与重入策略", () => { + let fetchSpy: ReturnType; + + beforeEach(() => { + fetchSpy = vi.spyOn(globalThis, "fetch"); + }); + + afterEach(() => { + fetchSpy.mockRestore(); + }); + + function createMockSender() { + const sentMessages: any[] = []; + let messageHandler: ((msg: any) => void) | null = null; + let disconnectHandler: (() => void) | null = null; + const mockConn = { + sendMessage: (msg: any) => sentMessages.push(msg), + onMessage: vi.fn((handler: any) => { + messageHandler = handler; + }), + onDisconnect: vi.fn((handler: any) => { + disconnectHandler = handler; + }), + }; + const sender = { + isType: (type: any) => type === 1, + getConnect: () => mockConn, + getSender: () => ({ url: chrome.runtime.getURL("src/options.html#/agent/chat") }), + }; + return { + sender, + sentMessages, + simulateMessage: (msg: any) => messageHandler?.(msg), + simulateDisconnect: () => disconnectHandler?.(), + }; + } + + function conv(id: string) { + return { id, title: "Chat", modelId: "test-openai", createtime: Date.now(), updatetime: Date.now() }; + } + + it("排队等待期间收到 stop 的请求:入锁后直接以取消收尾,不再调用 LLM", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-q")]); + mockRepo.getMessages.mockResolvedValue([]); + + // chat1 占住会话队列(fetch 挂起直到手动放行) + let release!: (response: Response) => void; + fetchSpy.mockImplementationOnce(() => new Promise((resolve) => (release = resolve))); + + const s1 = createMockSender(); + const p1 = (service as any).handleConversationChat({ conversationId: "conv-q", message: "第一条" }, s1.sender); + await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledTimes(1)); + + // chat2 入队等待;等待期间用户点了 Stop——回调必须在入队前就已注册,否则这次 stop 会被丢掉 + const s2 = createMockSender(); + const p2 = (service as any).handleConversationChat({ conversationId: "conv-q", message: "第二条" }, s2.sender); + s2.simulateMessage({ action: "stop" }); + + release(makeTextResponse("第一条回复")); + await Promise.all([p1, p2]); + + expect(fetchSpy).toHaveBeenCalledTimes(1); + const events2 = s2.sentMessages.filter((m: any) => m.action === "event").map((m: any) => m.data); + expect(events2.find((e: any) => e.type === "error")?.errorCode).toBe("cancelled"); + }); + + it("删除会话应取消已入队但尚未开始的请求", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-delete-queued")]); + mockRepo.getMessages.mockResolvedValue([]); + + let rejectFirst!: (error: Error) => void; + fetchSpy.mockImplementationOnce( + async (_url: RequestInfo | URL, init?: RequestInit) => + new Promise((_resolve, reject) => { + rejectFirst = reject; + (init?.signal as AbortSignal).addEventListener( + "abort", + () => reject(new DOMException("Aborted", "AbortError")), + { once: true } + ); + }) + ); + + const first = (service as any).handleConversationChat( + { conversationId: "conv-delete-queued", message: "第一条" }, + createMockSender().sender + ); + await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledOnce()); + const second = (service as any).handleConversationChat( + { conversationId: "conv-delete-queued", message: "第二条" }, + createMockSender().sender + ); + + const deletion = (service as any).handleConversation({ + action: "delete", + conversationId: "conv-delete-queued", + generation: "legacy:conv-delete-queued", + }); + await Promise.all([first, second, deletion]); + + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(mockRepo.deleteConversation).toHaveBeenCalledWith("conv-delete-queued", { + generation: "legacy:conv-delete-queued", + }); + void rejectFirst; + }); + + it("排队前快速拒绝应清理尚未被消息采用的上传附件", async () => { + const { service, mockRepo } = createTestService(); + (service as any).chatService.conversationsAwaitingScriptTools.add("conv-busy-upload"); + + await (service as any).handleConversationChat( + { + conversationId: "conv-busy-upload", + message: [{ type: "image", attachmentId: "busy-upload.png", mimeType: "image/png" }], + ownedAttachmentIds: ["busy-upload.png"], + }, + createMockSender().sender + ); + + expect(mockRepo.deleteAttachment).toHaveBeenCalledWith("busy-upload.png"); + expect(fetchSpy).not.toHaveBeenCalled(); + }); + + it("非 Options 调用方不得借 ownedAttachmentIds 删除借用附件", async () => { + const { service, mockRepo } = createTestService(); + (service as any).chatService.conversationsAwaitingScriptTools.add("conv-untrusted-upload"); + const connection = createMockSender(); + (connection.sender as any).getSender = () => ({ url: chrome.runtime.getURL("src/content.html") }); + + await (service as any).handleConversationChat( + { + conversationId: "conv-untrusted-upload", + message: [{ type: "image", attachmentId: "victim.png", mimeType: "image/png" }], + ownedAttachmentIds: ["victim.png"], + }, + connection.sender + ); + + expect(mockRepo.deleteAttachment).not.toHaveBeenCalledWith("victim.png"); + }); + + it("排队等待期间客户端已断开的前台请求:不启动、不调用 LLM、不发送事件", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-q2")]); + mockRepo.getMessages.mockResolvedValue([]); + + let release!: (response: Response) => void; + fetchSpy.mockImplementationOnce(() => new Promise((resolve) => (release = resolve))); + + const s1 = createMockSender(); + const p1 = (service as any).handleConversationChat({ conversationId: "conv-q2", message: "第一条" }, s1.sender); + await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledTimes(1)); + + const s2 = createMockSender(); + const p2 = (service as any).handleConversationChat({ conversationId: "conv-q2", message: "第二条" }, s2.sender); + s2.simulateDisconnect(); + + release(makeTextResponse("第一条回复")); + await Promise.all([p1, p2]); + + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(s2.sentMessages.filter((m: any) => m.action === "event")).toHaveLength(0); + }); + + it("端口在入队前已死(注册回调即抛错):请求安全返回,不留下卡死的后台会话记录", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-dead")]); + mockRepo.getMessages.mockResolvedValue([]); + const mockConn = { + sendMessage: vi.fn(), + onMessage: vi.fn(() => { + throw new Error("onMessage Invalid Port"); + }), + onDisconnect: vi.fn(() => { + throw new Error("onDisconnect Invalid Port"); + }), + }; + const sender = { isType: () => true, getConnect: () => mockConn }; + + await expect( + (service as any).handleConversationChat({ conversationId: "conv-dead", message: "hi", background: true }, sender) + ).resolves.toBeUndefined(); + + expect(fetchSpy).not.toHaveBeenCalled(); + // 不能把后台会话记录留在 running/cancelling 占位状态 + expect((service as any).bgSessionManager.has("conv-dead")).toBe(false); + }); + + it("会话等待脚本工具结果期间 clearMessages 应显式拒绝(重入死锁窗口),工具完成后恢复可用", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-reent")]); + mockRepo.getMessages.mockResolvedValue([]); + + const encoder = new TextEncoder(); + const toolCallChunks = [ + `data: {"choices":[{"delta":{"tool_calls":[{"id":"call-1","function":{"name":"my_tool","arguments":""}}]}}]}\n\n`, + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":"{}"}}]}, "finish_reason":"tool_calls"}]}\n\n`, + `data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`, + ]; + let i = 0; + const toolCallResponse = { + ok: true, + status: 200, + body: { + getReader() { + return { + async read() { + if (i < toolCallChunks.length) return { done: false, value: encoder.encode(toolCallChunks[i++]) }; + return { done: true, value: undefined }; + }, + cancel: async () => {}, + }; + }, + }, + } as unknown as Response; + + fetchSpy.mockResolvedValueOnce(toolCallResponse).mockResolvedValueOnce(makeTextResponse("最终回复")); + + const s = createMockSender(); + const chatPromise = (service as any).handleConversationChat( + { + conversationId: "conv-reent", + message: "使用工具", + tools: [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} } }], + }, + s.sender + ); + + await vi.waitFor(() => { + expect(s.sentMessages.some((message: any) => message.action === "executeTools")).toBe(true); + }); + + // 死锁窗口:chat 持有会话队列锁等待 toolResults;此刻的 clear 若排队会形成相互等待 + await expect( + (service as any).handleConversation({ action: "clearMessages", conversationId: "conv-reent" }) + ).rejects.toThrow(); + + const executeMessage = s.sentMessages.find((message: any) => message.action === "executeTools"); + s.simulateMessage({ + action: "toolResults", + requestId: executeMessage.requestId, + data: [{ id: "call-1", result: "ok" }], + }); + await chatPromise; + + // 工具等待结束后,clear 恢复正常排队语义 + await expect( + (service as any).handleConversation({ action: "clearMessages", conversationId: "conv-reent" }) + ).resolves.toBe(true); + }); + + it("会话等待脚本工具结果期间,同会话的重入 chat 应立即拒绝而不是排队死锁", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-reent-chat")]); + mockRepo.getMessages.mockResolvedValue([]); + + const encoder = new TextEncoder(); + let index = 0; + const chunks = [ + `data: {"choices":[{"delta":{"tool_calls":[{"id":"call-1","function":{"name":"my_tool","arguments":""}}]}}]}\n\n`, + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":"{}"}}]},"finish_reason":"tool_calls"}]}\n\n`, + `data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`, + ]; + fetchSpy + .mockResolvedValueOnce({ + ok: true, + status: 200, + body: { + getReader: () => ({ + read: async () => + index < chunks.length + ? { done: false, value: encoder.encode(chunks[index++]) } + : { done: true, value: undefined }, + cancel: async () => {}, + }), + }, + } as unknown as Response) + .mockResolvedValueOnce(makeTextResponse("最终回复")) + .mockResolvedValueOnce(makeTextResponse("不应执行的重入回复")); + + const outer = createMockSender(); + const outerPromise = (service as any).handleConversationChat( + { + conversationId: "conv-reent-chat", + message: "使用工具", + tools: [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} } }], + }, + outer.sender + ); + await vi.waitFor(() => + expect(outer.sentMessages.some((message: any) => message.action === "executeTools")).toBe(true) + ); + + const nested = createMockSender(); + const nestedPromise = (service as any).handleConversationChat( + { conversationId: "conv-reent-chat", message: "nested" }, + nested.sender + ); + const nestedOutcome = await Promise.race([ + nestedPromise.then(() => "returned"), + new Promise((resolve) => setTimeout(() => resolve("blocked"), 20)), + ]); + + const executeMessage = outer.sentMessages.find((message: any) => message.action === "executeTools"); + outer.simulateMessage({ + action: "toolResults", + requestId: executeMessage.requestId, + data: [{ id: "call-1", result: "ok" }], + }); + await Promise.all([outerPromise, nestedPromise]); + + expect(nestedOutcome).toBe("returned"); + expect(nested.sentMessages).toContainEqual({ + action: "event", + data: expect.objectContaining({ type: "error" }), + }); + expect(fetchSpy).toHaveBeenCalledTimes(2); + }); + + it("手动 compact 在 LLM 调用期间 Stop 时应恰好发送一次 cancelled 终态", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-compact-stop")]); + mockRepo.getMessages.mockResolvedValue([ + { id: "m1", conversationId: "conv-compact-stop", role: "user", content: "history", createtime: 1 }, + ]); + fetchSpy.mockImplementation( + async (_url: RequestInfo | URL, init?: RequestInit) => + new Promise((_resolve, reject) => { + (init?.signal as AbortSignal).addEventListener( + "abort", + () => reject(new DOMException("Aborted", "AbortError")), + { once: true } + ); + }) + ); + + const connection = createMockSender(); + const compactPromise = (service as any).handleConversationChat( + { conversationId: "conv-compact-stop", message: "", compact: true }, + connection.sender + ); + await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledOnce()); + connection.simulateMessage({ action: "stop" }); + await compactPromise; + + const terminals = connection.sentMessages + .filter((message: any) => message.action === "event") + .map((message: any) => message.data) + .filter((event: any) => event.type === "done" || event.type === "error"); + expect(terminals).toEqual([expect.objectContaining({ type: "error", errorCode: "cancelled" })]); + }); + + it("手动 compact 忽略模型生成 block 时应清理对应附件", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-compact-block")]); + mockRepo.getMessages.mockResolvedValue([ + { id: "m1", conversationId: "conv-compact-block", role: "user", content: "history", createtime: 1 }, + ]); + const chatService = (service as any).chatService; + chatService.llmDeps.callLLM = vi.fn().mockResolvedValue({ + content: "摘要", + contentBlocks: [{ type: "image", attachmentId: "manual-orphan.png", mimeType: "image/png" }], + }); + + await (service as any).handleConversationChat( + { conversationId: "conv-compact-block", message: "", compact: true }, + createMockSender().sender + ); + + expect(mockRepo.deleteAttachment).toHaveBeenCalledWith("manual-orphan.png"); + }); + + it("手动 compact 成功时应发送耗时", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-compact-duration")]); + mockRepo.getMessages.mockResolvedValue([ + { id: "m1", conversationId: "conv-compact-duration", role: "user", content: "history", createtime: 1 }, + ]); + fetchSpy.mockResolvedValueOnce(makeTextResponse("摘要")); + + const connection = createMockSender(); + await (service as any).handleConversationChat( + { conversationId: "conv-compact-duration", message: "", compact: true }, + connection.sender + ); + + const done = connection.sentMessages + .map((message: any) => message.data) + .find((event: any) => event.type === "done"); + expect(done).toEqual(expect.objectContaining({ type: "done", durationMs: expect.any(Number) })); + }); + + it("子代理终态不应吞掉父对话自己的终态", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-sub-terminal")]); + mockRepo.getMessages.mockResolvedValue([]); + const chatService = (service as any).chatService; + chatService.llmDeps.callLLMWithToolLoop = vi.fn(async ({ sendEvent }: any) => { + sendEvent({ + type: "done", + subAgent: { agentId: "child-1", description: "子任务" }, + }); + sendEvent({ type: "done", usage: { inputTokens: 3, outputTokens: 2 } }); + }); + + const connection = createMockSender(); + await (service as any).handleConversationChat( + { conversationId: "conv-sub-terminal", message: "run" }, + connection.sender + ); + + const terminals = connection.sentMessages + .filter((message: any) => message.action === "event") + .map((message: any) => message.data) + .filter((event: any) => event.type === "done" || event.type === "error"); + expect(terminals).toEqual([ + expect.objectContaining({ type: "done", subAgent: expect.objectContaining({ agentId: "child-1" }) }), + expect.objectContaining({ type: "done", usage: { inputTokens: 3, outputTokens: 2 } }), + ]); + }); + + it("删除活动会话应先取消执行并等待落定,再删除对应 generation", async () => { + const { service, mockRepo } = createTestService(); + mockRepo.listConversations.mockResolvedValue([conv("conv-delete-active")]); + mockRepo.getMessages.mockResolvedValue([]); + fetchSpy.mockImplementation( + async (_url: RequestInfo | URL, init?: RequestInit) => + new Promise((_resolve, reject) => { + (init?.signal as AbortSignal).addEventListener( + "abort", + () => reject(new DOMException("Aborted", "AbortError")), + { once: true } + ); + }) + ); + const connection = createMockSender(); + const chat = (service as any).handleConversationChat( + { conversationId: "conv-delete-active", message: "run" }, + connection.sender + ); + await vi.waitFor(() => expect(fetchSpy).toHaveBeenCalledOnce()); + + const deletion = (service as any).handleConversation({ + action: "delete", + conversationId: "conv-delete-active", + generation: "legacy:conv-delete-active", + }); + await Promise.all([chat, deletion]); + + expect(mockRepo.deleteConversation).toHaveBeenCalledWith("conv-delete-active", { + generation: "legacy:conv-delete-active", + }); + }); +}); diff --git a/src/app/service/agent/service_worker/chat_service.ts b/src/app/service/agent/service_worker/chat_service.ts index 8ce7a833b..89f79bd28 100644 --- a/src/app/service/agent/service_worker/chat_service.ts +++ b/src/app/service/agent/service_worker/chat_service.ts @@ -1,5 +1,6 @@ import type { IGetSender } from "@Packages/message/server"; import { GetSenderType } from "@Packages/message/server"; +import type { MessageConnect } from "@Packages/message/types"; import type { AgentModelConfig, ChatRequest, @@ -7,6 +8,7 @@ import type { Conversation, ConversationApiRequest, MessageContent, + TokenUsage, ToolDefinition, } from "@App/app/service/agent/core/types"; import type { ScriptToolCallback, ToolExecutor } from "@App/app/service/agent/core/tool_registry"; @@ -33,20 +35,37 @@ import { createExecuteScriptTool } from "@App/app/service/agent/core/tools/execu import { resolveSubAgentType } from "@App/app/service/agent/core/sub_agent_types"; import { classifyErrorCode } from "./retry_utils"; import { getTextContent } from "@App/app/service/agent/core/content_utils"; +import { + isLegacyGeneration, + retainedSummaryAttachmentIds, + toLLMMessages, +} from "@App/app/service/agent/core/persisted_messages"; import { uuidv4 } from "@App/pkg/utils/uuid"; +import { stackAsyncTask } from "@App/pkg/utils/async_queue"; +import { t } from "@App/locales/locales"; +import { elideUntilWithinBudget, estimateRequestTokens } from "@App/app/service/agent/core/context_elision"; +import { HEURISTIC_HARD_REJECT_RATIO } from "./tool_loop_orchestrator"; +import { getInputTokenBudget } from "@App/app/service/agent/core/model_context"; import type { LLMCallResult } from "./llm_client"; +import { prepareAttachmentSnapshot, type AttachmentSnapshot } from "@App/app/service/agent/core/attachment_resolver"; +import { RevisionConflictError } from "@App/app/repo/revision"; /** ChatService 需要的 execute_script 工具依赖 */ export interface ChatServiceExecuteScriptDeps { executeInPage: (code: string, options?: { tabId?: number }) => Promise<{ result: unknown; tabId: number }>; - executeInSandbox: (code: string) => Promise; + executeInSandbox: (code: string, signal?: AbortSignal) => Promise; } /** ChatService 需要的 LLM 调用依赖 */ export interface ChatServiceLLMDeps { callLLM: ( model: AgentModelConfig, - params: { messages: ChatRequest["messages"]; tools?: ToolDefinition[]; cache?: boolean }, + params: { + messages: ChatRequest["messages"]; + tools?: ToolDefinition[]; + cache?: boolean; + attachmentSnapshot?: AttachmentSnapshot; + }, sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal ) => Promise; @@ -56,9 +75,11 @@ export interface ChatServiceLLMDeps { /** handleConversationChat 参数类型 */ type ConversationChatParams = { conversationId: string; + // 调用方持有的会话 generation;提供时服务端会校验与当前存储一致 + generation?: string; message: MessageContent; + ownedAttachmentIds?: string[]; tools?: ToolDefinition[]; - maxIterations?: number; scriptUuid?: string; modelId?: string; enableTools?: boolean; // 是否携带 tools,undefined 表示不覆盖 @@ -89,7 +110,129 @@ interface BuildMessagesResult { messages: ChatRequest["messages"]; } +/** chat / compact / clearMessages / 定时任务续接共用的按 conversationId 队列锁 key。 + * 所有对同一会话持久化消息的读改写路径都必须经由这把锁串行,否则 appendMessage 的 + * 读-改-写会互相覆盖丢消息。 */ +export function conversationChatLockKey(conversationId: string): string { + return `agent-chat:${conversationId}`; +} + +function canAssertAttachmentOwnership(sender: IGetSender): boolean { + const senderUrl = sender.getSender?.()?.url; + if (!senderUrl) return false; + try { + return senderUrl.startsWith(chrome.runtime.getURL("src/options.html")); + } catch { + return false; + } +} + +/** 一次 conversation chat 连接的共享状态。 + * 连接回调(stop / askUserResponse / toolResults / 断开)必须在排队【之前】注册: + * 断开的端口上注册回调会直接抛错(见 extension_message.ts),而且排队等待期间到达的 + * 断开与 Stop 必须被如实记录,否则临界区开始后要么在死连接上白跑一整条会话, + * 要么漏掉停止指令。 */ +interface ChatConnectionSession { + msgConn: MessageConnect; + isBackground: boolean; + abortController: AbortController; + askResolvers: Map void>; + scriptToolCallback: ScriptToolCallback; + isDisconnected: () => boolean; + userAttachmentsAdopted: boolean; + releaseProvisionalUserAttachments: () => Promise; + /** 后台模式:临界区内创建 RunningConversation 后回填,供排队期间注册的连接回调路由 stop/askUser */ + rc?: RunningConversation; + bgListener?: ListenerEntry; +} + +/** 为脚本工具连接增加 executeTools 请求批次关联(requestId),并在后台客户端离线后 + * 直接返回结构化错误结果:过期批次的 toolResults 被丢弃,不会被下一个批次误认领; + * 连接已断开时不再把 executeTools 扔进黑洞,而是立刻回填整批 error 结果,让 tool loop + * 能继续走统一的失败路径而不是挂起等待。 */ +function wrapScriptToolConnection(original: MessageConnect): MessageConnect { + let disconnected = false; + let activeRequestId: string | undefined; + const inboundHandlers: Array<(message: any) => void> = []; + + const failBatch = (message: any, reason: string) => { + const toolCalls: ToolCall[] = message.data || []; + queueMicrotask(() => { + const response = { + action: "toolResults", + requestId: message.requestId, + data: toolCalls.map((toolCall) => ({ + id: toolCall.id, + result: JSON.stringify({ error: reason }), + error: true, + })), + }; + for (const handler of inboundHandlers) handler(response); + }); + }; + + return { + onMessage(callback) { + inboundHandlers.push(callback as (message: any) => void); + original.onMessage((message: any) => { + if (message.action === "toolResults") { + if (!message.requestId || message.requestId !== activeRequestId) return; + activeRequestId = undefined; + } + callback(message); + }); + }, + sendMessage(message: any) { + if (message.action === "cancelToolBatch") { + // 由包装层补上当前批次的 requestId:批次超时后客户端可能仍在串行执行剩余 handler, + // 明确通知其作废该批次;没有进行中的批次则无事可做 + if (!activeRequestId || disconnected) return; + try { + original.sendMessage({ ...message, requestId: activeRequestId }); + } catch { + disconnected = true; + } + return; + } + if (message.action !== "executeTools") { + if (!disconnected) original.sendMessage(message); + return; + } + + const correlated = { ...message, requestId: uuidv4() }; + activeRequestId = correlated.requestId; + if (disconnected) { + failBatch(correlated, "Script tool client is unavailable"); + return; + } + try { + original.sendMessage(correlated); + } catch (error) { + disconnected = true; + failBatch( + correlated, + error instanceof Error && error.message ? error.message : "Script tool client is unavailable" + ); + } + }, + disconnect(ignoreAlreadyDisconnected?: boolean) { + original.disconnect(ignoreAlreadyDisconnected); + }, + onDisconnect(callback) { + original.onDisconnect((isSelfDisconnected) => { + disconnected = true; + callback(isSelfDisconnected); + }); + }, + }; +} + export class ChatService { + // 正在等待 Sandbox 回复脚本工具结果的会话:此窗口内 chat 持有会话队列锁等待 toolResults, + // 来自工具 handler 内部的 await conv.clear() 若照常排队会形成"锁等我、我等锁"的死锁 + private conversationsAwaitingScriptTools = new Set(); + private admittedChats = new Map>(); + constructor( private toolRegistry: ToolRegistry, private modelService: AgentModelService, @@ -101,6 +244,25 @@ export class ChatService { private chatRepo: AgentChatRepo ) {} + private admitChat(conversationId: string, abortController: AbortController) { + const controllers = this.admittedChats.get(conversationId) || new Set(); + controllers.add(abortController); + this.admittedChats.set(conversationId, controllers); + } + + private releaseAdmittedChat(conversationId: string, abortController: AbortController) { + const controllers = this.admittedChats.get(conversationId); + if (!controllers) return; + controllers.delete(abortController); + if (controllers.size === 0) this.admittedChats.delete(conversationId); + } + + private abortAdmittedChats(conversationId: string) { + for (const abortController of this.admittedChats.get(conversationId) || []) { + abortController.abort(); + } + } + // 处理 Sandbox conversation API 请求(非流式) async handleConversation(params: ConversationApiRequest): Promise { switch (params.action) { @@ -109,13 +271,78 @@ export class ChatService { case "get": return this.getConversation(params.id); case "getMessages": - return this.chatRepo.getMessages(params.conversationId); - case "save": - // 对话已经在 chat 过程中持久化,这里确保元数据也保存 + // params.generation 提供时,与当前存储不一致(会话已被删除重建)则拒绝而非返回无关一代的消息; + // 未提供 generation 时保留旧行为:会话不存在则返回空数组 + try { + return (await this.chatRepo.getMessageSnapshot(params.conversationId, params.generation)).messages; + } catch (error) { + if (params.generation === undefined && error instanceof RevisionConflictError) return []; + throw error; + } + case "save": { + // 对话已经在 chat 过程中持久化,这里确保元数据也保存;仍需校验调用方持有的 generation + if (params.generation !== undefined) { + const conv = await this.getConversation(params.conversationId); + if (!conv || conv.generation !== params.generation) { + throw new Error("Conversation generation mismatch"); + } + } return true; + } case "clearMessages": - await this.chatRepo.saveMessages(params.conversationId, []); - return true; + // 会话正在等待脚本工具结果时,这个 clear 很可能来自该工具 handler 内部的 + // await conv.clear():chat 持有会话队列锁等待 toolResults,clear 排队等锁, + // 相互等待成死锁。对这个窗口显式拒绝(fail fast);其余时刻仍与 chat/compact + // 共用同一把按 conversationId 的队列锁排队执行,避免互相覆盖写入 + if (this.conversationsAwaitingScriptTools.has(params.conversationId)) { + throw new Error( + "Conversation is waiting for script tool results; clearing messages now would deadlock. Finish or stop the chat first." + ); + } + return stackAsyncTask(conversationChatLockKey(params.conversationId), async () => { + // 与 getMessages 同理,generation 提供时须匹配当前存储 + const snapshot = await this.chatRepo.getMessageSnapshot(params.conversationId, params.generation); + await this.chatRepo.saveMessages(params.conversationId, [], undefined, { + generation: snapshot.generation, + expectedRevision: snapshot.revision, + }); + const taskSnapshot = await this.chatRepo.getTaskSnapshot(params.conversationId, snapshot.generation); + await this.chatRepo.saveTasks( + params.conversationId, + [], + undefined, + snapshot.generation, + taskSnapshot.revision + ); + return true; + }); + case "deleteMessages": + return stackAsyncTask(conversationChatLockKey(params.conversationId), async () => { + const snapshot = await this.chatRepo.getMessageSnapshot(params.conversationId, params.generation); + const ids = new Set(params.messageIds); + await this.chatRepo.saveMessages( + params.conversationId, + snapshot.messages.filter((message) => !ids.has(message.id)), + undefined, + { + generation: snapshot.generation, + expectedRevision: snapshot.revision, + preserveAttachmentIds: params.preserveAttachmentIds, + } + ); + return true; + }); + case "delete": { + this.abortAdmittedChats(params.conversationId); + this.bgSessionManager.stop(params.conversationId); + return stackAsyncTask(conversationChatLockKey(params.conversationId), async () => { + await this.chatRepo.deleteConversation(params.conversationId, { + generation: params.generation, + ...(params.revision === undefined ? {} : { expectedRevision: params.revision }), + }); + return true; + }); + } default: throw new Error(`Unknown conversation action: ${(params as any).action}`); } @@ -132,119 +359,373 @@ export class ChatService { createtime: Date.now(), updatetime: Date.now(), }; - await this.chatRepo.saveConversation(conv); - return conv; + return this.chatRepo.createConversation(conv); } private async getConversation(id: string): Promise { const conversations = await this.chatRepo.listConversations(); - return conversations.find((c) => c.id === id) || null; + const conversation = conversations.find((item) => item.id === id); + if (!conversation) return null; + return { + ...conversation, + generation: conversation.generation || `legacy:${conversation.id}`, + revision: conversation.revision ?? 0, + }; } // 统一的流式 conversation chat(UI 和脚本 API 共用) + // 同一 conversationId 的 chat / compact(compact 复用本方法的 params.compact 分支)都必须与 + // clearMessages 串行执行,避免并发读改写互相覆盖对方的持久化写入。 + // "会话正在运行中" 的快速拒绝在排队之前完成,避免重复的后台请求白白卡在队列里等待。 async handleConversationChat(params: ConversationChatParams, sender: IGetSender) { + if (params.ownedAttachmentIds?.length && !canAssertAttachmentOwnership(sender)) { + params = { ...params, ownedAttachmentIds: undefined }; + } + let userAttachmentsAdopted = false; + let releasePromise: Promise | undefined; + const releaseProvisionalUserAttachments = () => { + if (params.ephemeral || userAttachmentsAdopted || !params.ownedAttachmentIds?.length) { + return Promise.resolve(); + } + releasePromise ||= Promise.all( + [...new Set(params.ownedAttachmentIds)].map((id) => this.chatRepo.deleteAttachment(id).catch(() => {})) + ).then(() => undefined); + return releasePromise; + }; + if (!sender.isType(GetSenderType.CONNECT)) { + await releaseProvisionalUserAttachments(); throw new Error("Conversation chat requires connect mode"); } - const msgConn = sender.getConnect()!; + const msgConn = wrapScriptToolConnection(sender.getConnect()!); // 后台模式:非 ephemeral、非 compact 时可用 const isBackground = params.background === true && !params.ephemeral && !params.compact; - // 检查是否已有后台运行的同一会话 + if (!params.ephemeral && this.conversationsAwaitingScriptTools.has(params.conversationId)) { + try { + msgConn.sendMessage({ + action: "event", + data: { + type: "error", + message: "Conversation is waiting for script tool results; reentrant chat is not allowed", + errorCode: "conversation_busy", + } as ChatStreamEvent, + }); + } catch { + // 端口已断开,无需通知 + } + await releaseProvisionalUserAttachments(); + return; + } + + // 检查是否已有后台运行的同一会话(排队前快速拒绝,入锁后还会复查一次) if (isBackground && this.bgSessionManager.has(params.conversationId)) { - msgConn.sendMessage({ - action: "event", - data: { type: "error", message: "会话正在运行中" } as ChatStreamEvent, - }); + try { + msgConn.sendMessage({ + action: "event", + data: { type: "error", message: "会话正在运行中" } as ChatStreamEvent, + }); + } catch { + // 端口已断开,无需通知 + } + await releaseProvisionalUserAttachments(); return; } const abortController = new AbortController(); + const askResolvers = new Map void>(); let isDisconnected = false; - // 后台模式:创建 RunningConversation + const session: ChatConnectionSession = { + msgConn, + isBackground, + abortController, + askResolvers, + isDisconnected: () => isDisconnected, + // 立即在下方赋值;提前占位以便连接回调闭包引用 session 对象本身 + scriptToolCallback: null as unknown as ScriptToolCallback, + get userAttachmentsAdopted() { + return userAttachmentsAdopted; + }, + set userAttachmentsAdopted(value: boolean) { + userAttachmentsAdopted = value; + }, + releaseProvisionalUserAttachments, + }; + + // 等待中的脚本工具调用:MessageConnect 断开或 abortController 触发时, + // 必须主动结束这个 pending promise —— 只有这个连接对应的 Sandbox 能回复 toolResults, + // 连接一旦断开该调用永远不会有结果,不结束会让整条 tool loop 挂起。 + let pendingScriptCall: { + toolCalls: ToolCall[]; + settle: (results: Array<{ id: string; result: string; error?: boolean }>) => void; + } | null = null; + + const scriptToolAbortedResults = (message: string) => + pendingScriptCall!.toolCalls.map((tc) => ({ + id: tc.id, + result: JSON.stringify({ error: message }), + error: true, + })); + + const settlePendingScriptCall = (message: string) => { + if (!pendingScriptCall) return; + pendingScriptCall.settle(scriptToolAbortedResults(message)); + }; + + // 连接回调必须在排队之前注册:断开的端口上注册会直接抛错(见 extension_message.ts), + // 而且排队等待期间到达的断开/Stop 必须被记录,否则临界区开始后要么在死连接上白跑 + // 一整条会话,要么漏掉停止指令 + try { + msgConn.onDisconnect(() => { + isDisconnected = true; + if (isBackground) { + // 后台模式:只移除 listener,不 abort 整条会话; + // 但这条连接对应的脚本工具调用必须结束,否则永远等不到 toolResults + if (session.rc && session.bgListener) session.rc.listeners.delete(session.bgListener); + settlePendingScriptCall("Script connection disconnected"); + } else { + abortController.abort(); + } + }); + + msgConn.onMessage((msg: any) => { + if (msg.action === "toolResults" && pendingScriptCall) { + const { settle } = pendingScriptCall; + settle(msg.data); + } + if (msg.action === "askUserResponse" && msg.data) { + const resolver = askResolvers.get(msg.data.id); + if (resolver) { + askResolvers.delete(msg.data.id); + if (session.rc) session.rc.pendingAskUser = undefined; + resolver(msg.data.answer); + } + } + if (msg.action === "stop") { + if (session.rc) { + this.bgSessionManager.stop(params.conversationId, session.rc); + } else { + abortController.abort(); + } + } + }); + } catch { + // 连接在注册回调前就已断开:请求方已不存在,直接不入队 + await releaseProvisionalUserAttachments(); + return; + } + + // abort(stop / 非后台断开)时结束等待中的脚本工具调用,避免循环卡死 + abortController.signal.addEventListener("abort", () => { + if (pendingScriptCall && !isDisconnected) { + try { + msgConn.sendMessage({ action: "cancelToolBatch" }); + } catch { + // 端口在 abort 竞态中关闭,pending 仍会在下面本地 settle。 + } + } + settlePendingScriptCall("Tool execution aborted"); + }); + + // 脚本工具单次调用的最长等待时间:Sandbox 长时间无响应(如脚本卡死)时, + // 主动结束这轮 tool call 而不是无限期挂起整条对话 + const SCRIPT_TOOL_TIMEOUT_MS = 5 * 60 * 1000; + + session.scriptToolCallback = (toolCalls: ToolCall[]) => { + return new Promise((resolve) => { + const settle = (results: Array<{ id: string; result: string; error?: boolean }>) => { + if (pendingScriptCall?.settle !== settle) return; + pendingScriptCall = null; + if (!params.ephemeral) this.conversationsAwaitingScriptTools.delete(params.conversationId); + clearTimeout(timer); + resolve(results); + }; + const timer = setTimeout(() => { + // 先通知客户端作废该批次(包装层补 requestId):超时后客户端可能仍在串行执行 + // 剩余 handler,其副作用会与下一批次交叠 + try { + msgConn.sendMessage({ action: "cancelToolBatch" }); + } catch { + // 端口已断开,无需通知 + } + settle( + toolCalls.map((tc) => ({ + id: tc.id, + result: JSON.stringify({ error: "Tool execution timed out" }), + error: true, + })) + ); + }, SCRIPT_TOOL_TIMEOUT_MS); + + pendingScriptCall = { toolCalls, settle }; + if (!params.ephemeral) this.conversationsAwaitingScriptTools.add(params.conversationId); + try { + msgConn.sendMessage({ action: "executeTools", data: toolCalls }); + } catch { + // 包装层已把断开的批次转成 failBatch 错误结果回填,这里仅防御极端时序下的底层抛错 + } + + if (abortController.signal.aborted) { + settle(scriptToolAbortedResults("Tool execution aborted")); + } + }); + }; + + // ephemeral 不读写 chatRepo(消息历史由调用方在内存中维护),没有跨请求的持久化竞争,无需排队 + if (!params.ephemeral) this.admitChat(params.conversationId, abortController); + try { + if (params.ephemeral) return await this.handleConversationChatLocked(params, session); + return await stackAsyncTask(conversationChatLockKey(params.conversationId), () => + this.handleConversationChatLocked(params, session) + ); + } finally { + if (!params.ephemeral) this.releaseAdmittedChat(params.conversationId, abortController); + await releaseProvisionalUserAttachments(); + } + } + + private async handleConversationChatLocked(params: ConversationChatParams, session: ChatConnectionSession) { + const { msgConn, isBackground, abortController, askResolvers, scriptToolCallback } = session; + + const sendEventDirect = (event: ChatStreamEvent) => { + if (session.isDisconnected()) return; + try { + msgConn.sendMessage({ action: "event", data: event }); + } catch { + // 端口在竞态下刚好断开,事件无处可送 + } + }; + + const { releaseProvisionalUserAttachments } = session; + + // 排队等待期间已被 Stop(前台断开也会 abort):不再启动,回发终态取消事件收尾 + if (abortController.signal.aborted) { + await releaseProvisionalUserAttachments(); + sendEventDirect({ type: "error", message: "Conversation cancelled", errorCode: "cancelled" }); + return; + } + + // 入锁后复查后台占用:排队前的快速拒绝与真正入锁之间存在时间窗 + if (isBackground && this.bgSessionManager.has(params.conversationId)) { + await releaseProvisionalUserAttachments(); + sendEventDirect({ type: "error", message: "会话正在运行中" }); + return; + } + + // 后台模式:创建 RunningConversation(askResolvers 与排队前注册的连接回调共享同一个 Map) let rc: RunningConversation | undefined; if (isBackground) { + // 后台会话必须先确认调用方持有的 generation 与当前存储一致,否则一次删除重建后的 + // 陈旧调用会静默附加到无关的新一代会话上 + const conv = await this.getConversation(params.conversationId); + if (!conv) { + await releaseProvisionalUserAttachments(); + sendEventDirect({ type: "error", message: "Conversation not found" }); + return; + } + if (params.generation !== undefined && conv.generation !== params.generation) { + await releaseProvisionalUserAttachments(); + sendEventDirect({ + type: "error", + message: "Conversation generation mismatch", + errorCode: "conversation_generation_mismatch", + }); + return; + } rc = { conversationId: params.conversationId, + generation: conv.generation!, abortController, listeners: new Set(), streamingState: { content: "", thinking: "", toolCalls: [] }, - askResolvers: new Map(), + askResolvers, tasks: [], status: "running", }; this.bgSessionManager.set(params.conversationId, rc); - } + session.rc = rc; - // ask_user resolvers(后台模式挂在 rc 上,普通模式本地) - const askResolvers = rc ? rc.askResolvers : new Map void>(); + // 初始 listener;排队期间客户端已断开的后台会话照常运行,只是不再挂 listener + const listener: ListenerEntry = { sendEvent: sendEventDirect }; + session.bgListener = listener; + if (!session.isDisconnected()) rc.listeners.add(listener); + } + let terminalEventSent = false; const sendEvent = (event: ChatStreamEvent) => { + const isParentTerminal = (event.type === "done" || event.type === "error") && !event.subAgent; + if (isParentTerminal) { + if (terminalEventSent) return; + terminalEventSent = true; + } if (rc) { // 后台模式:先更新快照,再广播到所有 listener this.bgSessionManager.updateStreamingState(rc, event); this.bgSessionManager.broadcastEvent(rc, event); } else { - if (!isDisconnected) { - msgConn.sendMessage({ action: "event", data: event }); - } + sendEventDirect(event); } }; - if (rc) { - // 后台模式:初始 listener - const listener: ListenerEntry = { - sendEvent: (event) => { - if (!isDisconnected) { - msgConn.sendMessage({ action: "event", data: event }); - } - }, - }; - rc.listeners.add(listener); - - msgConn.onDisconnect(() => { - isDisconnected = true; - // 后台模式:只移除 listener,不 abort - rc!.listeners.delete(listener); + const emitCancelledOnce = (error?: { usage?: TokenUsage; durationMs?: number }) => { + sendEvent({ + type: "error", + message: "Conversation cancelled", + errorCode: "cancelled", + usage: error?.usage, + durationMs: error?.durationMs, }); - } else { - msgConn.onDisconnect(() => { - isDisconnected = true; - abortController.abort(); - }); - } - - // 构建脚本工具回调:通过 MessageConnect 让 Sandbox 执行 handler - let toolResultResolve: ((results: Array<{ id: string; result: string }>) => void) | null = null; - - msgConn.onMessage((msg: any) => { - if (msg.action === "toolResults" && toolResultResolve) { - const resolve = toolResultResolve; - toolResultResolve = null; - resolve(msg.data); - } - if (msg.action === "askUserResponse" && msg.data) { - const resolver = askResolvers.get(msg.data.id); - if (resolver) { - askResolvers.delete(msg.data.id); - if (rc) rc.pendingAskUser = undefined; - resolver(msg.data.answer); - } - } - if (msg.action === "stop") { - abortController.abort(); - } - }); + }; - const scriptToolCallback: ScriptToolCallback = (toolCalls: ToolCall[]) => { + // 循环检测(tool_call_guard)连续命中时暂停询问用户是否继续;复用 ask_user 的事件/resolver 机制, + // 5 分钟无人应答时默认"继续",避免无 UI 监听的后台会话被无限期挂起 + const askUserForGuard = (strikeCount: number): Promise => { return new Promise((resolve) => { - toolResultResolve = resolve; - msgConn.sendMessage({ action: "executeTools", data: toolCalls }); + if (abortController.signal.aborted) { + resolve("stop"); + return; + } + const askId = `guard_${uuidv4()}`; + const cleanup = () => abortController.signal.removeEventListener("abort", onAbort); + // settle 只负责清理与 resolve,不发送任何终态事件; + // 终态事件(resolved / expired)由每个触发路径各自发送且只发一次,避免重复广播 + const settle = (answer: string) => { + clearTimeout(timer); + askResolvers.delete(askId); + cleanup(); + resolve(answer); + }; + const onAbort = () => { + sendEvent({ type: "ask_user_expired", id: askId }); + settle("stop"); + }; + const timer = setTimeout( + () => { + sendEvent({ type: "ask_user_expired", id: askId }); + settle("continue"); + }, + 5 * 60 * 1000 + ); + sendEvent({ + type: "ask_user", + id: askId, + question: t("agent:chat_guard_question", { count: strikeCount }), + options: [t("agent:chat_guard_continue"), t("agent:chat_guard_stop")], + optionValues: ["continue", "stop"], + multiple: false, + allowCustom: false, + }); + abortController.signal.addEventListener("abort", onAbort, { once: true }); + askResolvers.set(askId, (answer: string) => { + sendEvent({ type: "ask_user_resolved", id: askId }); + settle(answer); + }); }); }; + let conversationGeneration: string | undefined; try { // ephemeral 模式:无状态处理,不从 repo 加载/持久化 if (params.ephemeral) { @@ -255,6 +736,7 @@ export class ChatService { // compact 模式:压缩对话历史 if (params.compact) { await this.handleCompactChat(params, sendEvent, abortController); + if (abortController.signal.aborted) emitCancelledOnce(); return; } @@ -264,6 +746,16 @@ export class ChatService { sendEvent({ type: "error", message: "Conversation not found" }); return; } + // 调用方持有的 generation 与当前存储不一致:会话已被删除重建,拒绝作用于无关的新一代会话 + if (params.generation !== undefined && conv.generation !== params.generation) { + sendEvent({ + type: "error", + message: "Conversation generation mismatch", + errorCode: "conversation_generation_mismatch", + }); + return; + } + conversationGeneration = conv.generation; // UI 传入 modelId / enableTools 时覆盖 conversation 的配置 let needSave = false; @@ -297,7 +789,7 @@ export class ChatService { // 预加载历史中已使用过的 skill 工具 if (enableTools) { - await this.preloadSkillsFromHistory(existingMessages, metaTools); + await this.preloadSkillsFromHistory(existingMessages, metaTools, abortController.signal); } // 构建消息列表并持久化用户消息 @@ -307,6 +799,9 @@ export class ChatService { existingMessages, enableTools, promptSuffix, + onUserMessagePersisted: () => { + session.userAttachmentsAdopted = true; + }, }); try { @@ -316,44 +811,67 @@ export class ChatService { model, messages, tools: enableTools ? params.tools : undefined, - maxIterations: params.maxIterations || 50, sendEvent, signal: abortController.signal, scriptToolCallback: enableTools && params.tools && params.tools.length > 0 ? scriptToolCallback : null, conversationId: params.conversationId, + conversationGeneration, + rehydratedHistory: true, skipBuiltinTools: !enableTools, + askUserForGuard: params.scriptUuid ? undefined : askUserForGuard, }); - // 后台模式:正常完成后延迟清理 - this.bgSessionManager.cleanupIfDone(params.conversationId); + // callLLMWithToolLoop 在 signal.aborted 时是 return(正常 resolve)而非 throw, + // 因此 abort 落定也会走到这里;必须先收敛 cancelling 为终态,而不是直接当作正常完成清理 + if (rc && abortController.signal.aborted) { + this.bgSessionManager.finalizeCancelled(params.conversationId, rc); + } else { + // 后台模式:正常完成后延迟清理 + this.bgSessionManager.cleanupIfDone(params.conversationId); + } } finally { // sessionRegistry 超出作用域后由 GC 清理,无需手动 unregister // 清理子代理上下文缓存 this.subAgentService.cleanup(params.conversationId); } } catch (e: any) { - // 后台模式:abort 也需要清理注册表 + // 后台模式:abort 后必须等待本次执行 promise 真正落定,才能把 cancelling 收敛为终态, + // 否则 stop() 造成的 cancelling 占位会一直阻塞同 ID 的新会话(见 finalizeCancelled) if (abortController.signal.aborted) { - this.bgSessionManager.cleanupIfDone(params.conversationId); + emitCancelledOnce(e); + if (rc) { + this.bgSessionManager.finalizeCancelled(params.conversationId, rc); + } else { + this.bgSessionManager.cleanupIfDone(params.conversationId); + } return; } const errorMsg = e.message || "Unknown error"; + const errorCode = classifyErrorCode(e); // 持久化错误消息到 OPFS,确保刷新后仍可见 if (params.conversationId && !params.ephemeral) { try { - await this.chatRepo.appendMessage({ - id: uuidv4(), - conversationId: params.conversationId, - role: "assistant", - content: "", - error: errorMsg, - createtime: Date.now(), - }); + await this.chatRepo.appendMessage( + { + id: uuidv4(), + conversationId: params.conversationId, + role: "assistant", + content: "", + error: errorMsg, + errorCode, + usage: e.usage, + durationMs: e.durationMs, + createtime: Date.now(), + }, + conversationGeneration + ); } catch { // 持久化失败不阻塞错误事件发送 } } - sendEvent({ type: "error", message: errorMsg, errorCode: classifyErrorCode(e) }); + sendEvent({ type: "error", message: errorMsg, errorCode, usage: e.usage, durationMs: e.durationMs }); this.bgSessionManager.cleanupIfDone(params.conversationId); + } finally { + await releaseProvisionalUserAttachments(); } } @@ -369,7 +887,6 @@ export class ChatService { ): Promise { const model = await this.modelService.getModel(params.modelId); - // 使用脚本传入的完整消息历史 const messages: ChatRequest["messages"] = []; // 添加 system prompt(内置提示词 + 用户自定义) @@ -379,12 +896,7 @@ export class ChatService { // 添加脚本端维护的消息历史(已含最新 user message) if (params.messages) { for (const msg of params.messages) { - messages.push({ - role: msg.role, - content: msg.content, - toolCallId: msg.toolCallId, - toolCalls: msg.toolCalls, - }); + messages.push(...toLLMMessages([msg])); } } @@ -395,7 +907,6 @@ export class ChatService { model, messages, tools: params.tools, - maxIterations: params.maxIterations || 20, sendEvent, signal: abortController.signal, scriptToolCallback: params.tools && params.tools.length > 0 ? scriptToolCallback : null, @@ -412,43 +923,88 @@ export class ChatService { sendEvent: (event: ChatStreamEvent) => void, abortController: AbortController ): Promise { + const startTime = Date.now(); const conv = await this.getConversation(params.conversationId); if (!conv) { sendEvent({ type: "error", message: "Conversation not found" }); return; } + // 与非 compact 的 chat 分支同样的保护:调用方持有的 generation 与当前存储不一致时 + // (会话已被删除重建),拒绝而不是静默压缩无关的新一代会话历史 + if (params.generation !== undefined && conv.generation !== params.generation) { + sendEvent({ + type: "error", + message: "Conversation generation mismatch", + errorCode: "conversation_generation_mismatch", + }); + return; + } const model = await this.modelService.getModel(params.modelId || conv.modelId); - const existingMessages = await this.chatRepo.getMessages(params.conversationId); + if (!conv.generation) { + sendEvent({ type: "error", message: "Conversation not found" }); + return; + } + const snapshot = await this.chatRepo.getMessageSnapshot(params.conversationId, conv.generation); + const existingMessages = snapshot.messages; + const historyMessages = toLLMMessages(existingMessages).filter((msg) => msg.role !== "system"); - if (existingMessages.filter((m) => m.role !== "system").length === 0) { + if (historyMessages.length === 0) { sendEvent({ type: "error", message: "No messages to compact" }); return; } - // 构建摘要请求 const summaryMessages: ChatRequest["messages"] = []; summaryMessages.push({ role: "system", content: COMPACT_SYSTEM_PROMPT }); - for (const msg of existingMessages) { - if (msg.role === "system") continue; - summaryMessages.push({ - role: msg.role, - content: msg.content, - toolCallId: msg.toolCallId, - toolCalls: msg.toolCalls, - }); - } - + summaryMessages.push(...historyMessages); summaryMessages.push({ role: "user", content: buildCompactUserPrompt(params.compactInstruction) }); + const attachmentSnapshot = await prepareAttachmentSnapshot( + summaryMessages, + model, + (id) => this.chatRepo.getAttachment(id), + abortController.signal + ); + const inputBudget = getInputTokenBudget(model); + const effectiveWindow = Math.max(1, Math.floor(inputBudget / 0.9)); + if (!elideUntilWithinBudget(summaryMessages, effectiveWindow, undefined, 0.9, attachmentSnapshot.sizes, model)) { + // estimateRequestTokens 是启发式估算(固定 2 字节/token),不是真实 tokenizer,可能明显 + // 高估。裁剪到底后仍未达到预算的这个倍数时才在本地硬拒绝,否则仍交给 provider 自行判定—— + // provider 的真实拒绝会经由 classifyErrorCode 识别为 context_too_large + const elidedTokens = estimateRequestTokens(summaryMessages, undefined, attachmentSnapshot.sizes, model); + if (elidedTokens >= inputBudget * HEURISTIC_HARD_REJECT_RATIO) { + sendEvent({ + type: "error", + message: "Conversation history is too large to compact", + errorCode: "context_too_large", + }); + return; + } + } + // 不带 tools 调用 LLM const result = await this.llmDeps.callLLM( model, - { messages: summaryMessages, cache: false }, + { messages: summaryMessages, cache: false, attachmentSnapshot }, sendEvent, abortController.signal ); + await Promise.all( + (result.contentBlocks || []) + .filter((block) => block.type !== "text") + .map((block) => this.chatRepo.deleteAttachment(block.attachmentId).catch(() => {})) + ); + const compactError = (message: string, cause?: unknown) => + Object.assign(new Error(message), { + usage: result.usage, + durationMs: Date.now() - startTime, + cause, + }); + + // LLM 调用期间可能已被 stop:落库/广播终态事件前必须重新检查, + // 否则 cancelled 之后仍可能持久化摘要并发出 compact_done/done + if (abortController.signal.aborted) throw compactError("Aborted"); const summary = extractSummary(result.content); const originalCount = existingMessages.length; @@ -459,12 +1015,27 @@ export class ChatService { conversationId: params.conversationId, role: "user" as const, content: `[Conversation Summary]\n\n${summary}`, + ownedAttachmentIds: retainedSummaryAttachmentIds(summary, existingMessages, isLegacyGeneration(conv.generation)), createtime: Date.now(), }; - await this.chatRepo.saveMessages(params.conversationId, [summaryMessage]); + // 传入 signal:写入落定前若已 abort,则放弃这次整份覆写而不提交 + try { + await this.chatRepo.saveMessages(params.conversationId, [summaryMessage], abortController.signal, { + generation: conv.generation, + expectedRevision: snapshot.revision, + }); + } catch (error) { + throw compactError(error instanceof Error ? error.message : String(error), error); + } + + // 不做旧快照回滚:写入已经通过 revision CAS 线性化;无条件回写会覆盖随后追加的新消息。 + // Stop 只影响终态报告,不得再用进入 compact 时的历史覆盖更新后的状态。 + if (abortController.signal.aborted) { + throw compactError("Aborted"); + } sendEvent({ type: "compact_done", summary, originalCount }); - sendEvent({ type: "done", usage: result.usage }); + sendEvent({ type: "done", usage: result.usage, durationMs: Date.now() - startTime }); } /** @@ -486,9 +1057,6 @@ export class ChatService { // enableTools 默认为 true const enableTools = conv.enableTools !== false; - // 每个 chat 请求一个独立的 SessionToolRegistry(parent = 全局 toolRegistry) - // 会话级 meta-tools(skill / task / ask_user / sub_agent / execute_script)只注册到 session, - // 避免并发会话的闭包互相覆盖。session 超出作用域后由 GC 清理,无需手动 unregister。 const sessionRegistry = new SessionToolRegistry(this.toolRegistry); // 解析 Skills(注入 prompt + 注册 meta-tools),仅在启用 tools 时执行 @@ -499,25 +1067,27 @@ export class ChatService { promptSuffix = resolved.promptSuffix; metaTools = resolved.metaTools; - // 注册 skill meta-tools 到 session for (const mt of metaTools) { sessionRegistry.register("skill", mt.definition, mt.executor); } // Task tools(从持久化加载,变更时保存并推送事件到 UI) - const initialTasks = await this.chatRepo.getTasks(params.conversationId); + const taskSnapshot = await this.chatRepo.getTaskSnapshot(params.conversationId, conv.generation); const { tools: taskToolDefs } = createTaskTools({ - initialTasks, - onSave: (tasks) => this.chatRepo.saveTasks(params.conversationId, tasks), + initialTasks: taskSnapshot.tasks, + initialRevision: taskSnapshot.revision, + onSave: (tasks, signal, expectedRevision) => + this.chatRepo.saveTasks(params.conversationId, tasks, signal, conv.generation, expectedRevision), sendEvent, }); for (const t of taskToolDefs) { sessionRegistry.register("session", t.definition, t.executor); } - // Ask user - const askTool = createAskUserTool(sendEvent, askResolvers); - sessionRegistry.register("session", askTool.definition, askTool.executor); + if (!params.scriptUuid) { + const askTool = createAskUserTool(sendEvent, askResolvers, abortController.signal); + sessionRegistry.register("session", askTool.definition, askTool.executor); + } // Sub-agent const subAgentTool = createSubAgentTool({ @@ -537,6 +1107,7 @@ export class ChatService { agentId, description: options.description || "Sub-agent task", subAgentType: typeConfig.name, + toolCallId: options.toolCallId, }, } as ChatStreamEvent); @@ -546,7 +1117,6 @@ export class ChatService { childRegistry.register("session", t.definition, t.executor); } - // 独立的 execute_script const childExecTool = createExecuteScriptTool(this.executeScriptDeps); childRegistry.register("session", childExecTool.definition, childExecTool.executor); @@ -574,7 +1144,6 @@ export class ChatService { }); sessionRegistry.register("session", subAgentTool.definition, subAgentTool.executor); - // Execute script const executeScriptTool = createExecuteScriptTool(this.executeScriptDeps); sessionRegistry.register("session", executeScriptTool.definition, executeScriptTool.executor); } @@ -588,7 +1157,8 @@ export class ChatService { */ private async preloadSkillsFromHistory( existingMessages: Awaited>, - metaTools: Array<{ definition: ToolDefinition; executor: ToolExecutor }> + metaTools: Array<{ definition: ToolDefinition; executor: ToolExecutor }>, + signal?: AbortSignal ): Promise { const loadSkillMeta = metaTools.find((mt) => mt.definition.name === "load_skill"); if (!loadSkillMeta) return; @@ -611,10 +1181,10 @@ export class ChatService { } } - // 预执行 load_skill 以注册动态工具(结果不需要,只需要副作用) for (const skillName of loadedSkillNames) { + if (signal?.aborted) break; try { - await loadSkillMeta.executor.execute({ skill_name: skillName }); + await loadSkillMeta.executor.execute({ skill_name: skillName }, signal); } catch { // 加载失败,跳过 } @@ -630,10 +1200,10 @@ export class ChatService { existingMessages: Awaited>; enableTools: boolean; promptSuffix: string; + onUserMessagePersisted?: () => void; }): Promise { const { conv, params, existingMessages, enableTools, promptSuffix } = ctx; - // 构建消息列表 const messages: ChatRequest["messages"] = []; // 添加 system 消息(内置提示词 + 用户自定义 + skill prompt) @@ -643,27 +1213,44 @@ export class ChatService { }); messages.push({ role: "system", content: systemContent }); - // 添加历史消息(跳过 system) - for (const msg of existingMessages) { - if (msg.role === "system") continue; - messages.push({ - role: msg.role, - content: msg.content, - toolCallId: msg.toolCallId, - toolCalls: msg.toolCalls, - }); - } + // 添加历史消息(跳过 system;跳过错误占位消息 —— 错误占位仅用于 UI 展示,重放给 LLM + // 会产生空 content 的无意义消息,部分 provider(如 Anthropic)甚至会因空 content 拒绝请求) + messages.push(...toLLMMessages(existingMessages).filter((msg) => msg.role !== "system")); if (!params.skipSaveUserMessage) { // 添加新用户消息到 LLM 上下文并持久化 messages.push({ role: "user", content: params.message }); - await this.chatRepo.appendMessage({ - id: uuidv4(), - conversationId: params.conversationId, - role: "user", - content: params.message, - createtime: Date.now(), - }); + const userMessageId = uuidv4(); + try { + await this.chatRepo.appendMessage( + { + id: userMessageId, + conversationId: params.conversationId, + role: "user", + content: params.message, + ownedAttachmentIds: params.ownedAttachmentIds, + createtime: Date.now(), + }, + conv.generation + ); + } catch (error) { + // OPFS close 报告的错误具有二义性:写入可能已经落盘,只是确认读又恰好失败。 + // 附件所有权只有在这次 append 被判定成功后才会转移给消息(见 onUserMessagePersisted); + // 若在这里把二义性错误当作"未持久化"直接向上抛,外层会删除刚刚可能已经被这条 + // 持久化消息引用的临时附件,导致消息引用悬空文件。因此必须先明确 + // 确认这条消息是否已经真正写入,只有确认"确实未写入"才允许向上传播失败。 + let committed: boolean; + try { + const snapshot = await this.chatRepo.getMessageSnapshot(params.conversationId, conv.generation); + committed = snapshot.messages.some((message) => message.id === userMessageId); + } catch { + // 无法确认是否落盘时,附件可能已经被消息引用;先阻止外层把它当作临时文件删除。 + ctx.onUserMessagePersisted?.(); + throw error; + } + if (!committed) throw error; + } + ctx.onUserMessagePersisted?.(); } // 更新对话标题(如果是第一条消息) diff --git a/src/app/service/agent/service_worker/compact_service.test.ts b/src/app/service/agent/service_worker/compact_service.test.ts new file mode 100644 index 000000000..eaa81a2e8 --- /dev/null +++ b/src/app/service/agent/service_worker/compact_service.test.ts @@ -0,0 +1,224 @@ +import { describe, expect, it, vi } from "vitest"; +import { CompactService } from "./compact_service"; +import type { AgentModelConfig, ChatRequest } from "@App/app/service/agent/core/types"; +import { COMPACT_SYSTEM_PROMPT, buildCompactUserPrompt } from "@App/app/service/agent/core/compact_prompt"; +import { estimateRequestTokens } from "@App/app/service/agent/core/context_elision"; +import { getContextWindow, getInputTokenBudget } from "@App/app/service/agent/core/model_context"; + +const MODEL: AgentModelConfig = { + id: "compact-model", + name: "Compact", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", + contextWindow: 10_000, + maxTokens: 2_000, +}; + +// 二分查找最小长度而非逐字符递增:estimateRequestTokens 引入字节→token 折算后, +// 命中目标区间所需的字符数随之变大,逐字符线性扫描(每次都要 O(len) 的 JSON.stringify) +// 会在字符数翻倍后耗时成倍增长,曾在 CI 上导致该用例超时。 +function makeCurrentMessages(): ChatRequest["messages"] { + const buildMessages = (len: number): ChatRequest["messages"] => [{ role: "user", content: "x".repeat(len) }]; + const estimateFor = (len: number) => { + const summaryMessages: ChatRequest["messages"] = [ + { role: "system", content: COMPACT_SYSTEM_PROMPT }, + ...buildMessages(len), + { role: "user", content: buildCompactUserPrompt() }, + ]; + return estimateRequestTokens(summaryMessages, undefined, undefined, MODEL); + }; + + const budget = getInputTokenBudget(MODEL); + const ceiling = getContextWindow(MODEL) * 0.9; + + let low = 1; + let high = 20_000; + while (low < high) { + const mid = Math.floor((low + high) / 2); + if (estimateFor(mid) > budget) high = mid; + else low = mid + 1; + } + + const estimate = estimateFor(low); + if (estimate > budget && estimate < ceiling) return buildMessages(low); + throw new Error("未能构造出位于输入预算与 90% 预检之间的 compact 测试样例"); +} + +describe("CompactService 自动压缩", () => { + it("自动压缩应返回摘要请求的 token 用量", async () => { + const usage = { inputTokens: 120, outputTokens: 30, cacheCreationInputTokens: 10, cacheReadInputTokens: 5 }; + const modelService = {} as any; + const orchestrator = { + callLLM: vi.fn().mockResolvedValue({ + content: "摘要", + usage, + contentBlocks: [{ type: "image", attachmentId: "auto-orphan.png", mimeType: "image/png" }], + }), + }; + const chatRepo = { + getAttachment: vi.fn().mockResolvedValue(null), + getMessageSnapshot: vi.fn().mockResolvedValue({ generation: "gen-1", revision: 2, messages: [] }), + saveMessages: vi.fn().mockResolvedValue(undefined), + deleteAttachment: vi.fn().mockResolvedValue(undefined), + } as any; + const service = new CompactService(modelService, orchestrator, chatRepo); + + await expect( + service.autoCompact( + "conv-1", + "gen-1", + MODEL, + [{ role: "user", content: "需要摘要的内容" }], + vi.fn(), + new AbortController().signal + ) + ).resolves.toEqual(usage); + expect(chatRepo.deleteAttachment).toHaveBeenCalledWith("auto-orphan.png"); + }); + + it("摘要保留 uploads 路径时应把对应历史附件所有权转移给摘要消息", async () => { + const orchestrator = { + callLLM: vi.fn().mockResolvedValue({ content: "继续使用 uploads/retained.png" }), + }; + const chatRepo = { + getAttachment: vi.fn().mockResolvedValue(new Blob(["image"], { type: "image/png" })), + getMessageSnapshot: vi.fn().mockResolvedValue({ + generation: "gen-1", + revision: 2, + messages: [ + { + id: "m1", + conversationId: "conv-1", + role: "user", + content: [{ type: "image", attachmentId: "retained.png", mimeType: "image/png" }], + ownedAttachmentIds: ["retained.png"], + createtime: 1, + }, + ], + }), + saveMessages: vi.fn().mockResolvedValue(undefined), + deleteAttachment: vi.fn().mockResolvedValue(undefined), + } as any; + const service = new CompactService({} as any, orchestrator, chatRepo); + + await service.autoCompact( + "conv-1", + "gen-1", + MODEL, + [{ role: "user", content: [{ type: "image", attachmentId: "retained.png", mimeType: "image/png" }] }], + vi.fn(), + new AbortController().signal + ); + + expect(chatRepo.saveMessages.mock.calls[0][1][0].ownedAttachmentIds).toEqual(["retained.png"]); + }); + + it("Stop 恰好落在摘要提交之后时不应以旧快照覆盖后续历史", async () => { + const controller = new AbortController(); + const priorMessages = [{ id: "m1", conversationId: "conv-1", role: "user", content: "原始历史", createtime: 1 }]; + const modelService = {} as any; + const orchestrator = { callLLM: vi.fn().mockResolvedValue({ content: "摘要" }) }; + const saveCalls: any[][] = []; + const chatRepo = { + getAttachment: vi.fn().mockResolvedValue(null), + getMessageSnapshot: vi.fn().mockResolvedValue({ generation: "gen-1", revision: 2, messages: priorMessages }), + saveMessages: vi.fn().mockImplementation(async (_id: string, messages: any[]) => { + saveCalls.push(messages); + // 模拟 abort 恰好落在 close() 提交窗口:写入已生效,signal 事后才被观察到 + if (saveCalls.length === 1) controller.abort(); + }), + } as any; + const service = new CompactService(modelService, orchestrator, chatRepo); + const sendEvent = vi.fn(); + + await expect( + service.autoCompact("conv-1", "gen-1", MODEL, [{ role: "user", content: "内容" }], sendEvent, controller.signal) + ).rejects.toThrow("Aborted"); + + // 写入已经通过 revision CAS 线性化;不能再无条件回写旧历史覆盖之后的追加。 + expect(saveCalls).toHaveLength(1); + expect(sendEvent).not.toHaveBeenCalledWith(expect.objectContaining({ type: "compact_done" })); + }); + + it("摘要请求超过输出保留预算时应返回 context_too_large", async () => { + const modelService = {} as any; + const orchestrator = { callLLM: vi.fn() }; + const chatRepo = { + getAttachment: vi.fn().mockResolvedValue(null), + getMessageSnapshot: vi.fn().mockResolvedValue({ generation: "gen-1", revision: 0, messages: [] }), + saveMessages: vi.fn().mockResolvedValue(undefined), + } as any; + const service = new CompactService(modelService, orchestrator, chatRepo); + const sendEvent = vi.fn(); + const signal = new AbortController().signal; + const currentMessages = makeCurrentMessages(); + + await expect( + service.autoCompact("conv-1", "gen-1", MODEL, currentMessages, sendEvent, signal) + ).rejects.toMatchObject({ + errorCode: "context_too_large", + }); + + expect(orchestrator.callLLM).not.toHaveBeenCalled(); + expect(chatRepo.saveMessages).not.toHaveBeenCalled(); + }); + + it("摘要成功后持久化失败时异常应保留摘要调用 usage", async () => { + const usage = { inputTokens: 44, outputTokens: 9 }; + const orchestrator = { callLLM: vi.fn().mockResolvedValue({ content: "摘要", usage }) }; + const chatRepo = { + getAttachment: vi.fn().mockResolvedValue(null), + getMessageSnapshot: vi.fn().mockResolvedValue({ generation: "gen-1", revision: 1, messages: [] }), + saveMessages: vi.fn().mockRejectedValue(new Error("disk full")), + } as any; + const service = new CompactService({} as any, orchestrator, chatRepo); + + await expect( + service.autoCompact( + "conv-1", + "gen-1", + MODEL, + [{ role: "user", content: "content" }], + vi.fn(), + new AbortController().signal + ) + ).rejects.toMatchObject({ message: "disk full", usage }); + }); + + it("摘要内容超过摘要模型预算时应在调用 provider 前返回 context_too_large", async () => { + const modelService = { + getSummaryModel: vi.fn().mockResolvedValue(MODEL), + } as any; + const orchestrator = { callLLM: vi.fn().mockResolvedValue({ content: "ok" }) }; + const chatRepo = { + getAttachment: vi.fn().mockResolvedValue(null), + saveMessages: vi.fn().mockResolvedValue(undefined), + } as any; + const service = new CompactService(modelService, orchestrator, chatRepo); + const hugeContent = "x".repeat(20_000); + + await expect(service.summarizeContent(hugeContent, "extract")).rejects.toMatchObject({ + errorCode: "context_too_large", + estimatedInputTokens: expect.any(Number), + }); + + expect(orchestrator.callLLM).not.toHaveBeenCalled(); + }); + + it("网页摘要忽略模型生成 block 时应清理对应附件", async () => { + const modelService = { getSummaryModel: vi.fn().mockResolvedValue(MODEL) } as any; + const orchestrator = { + callLLM: vi.fn().mockResolvedValue({ + content: "summary", + contentBlocks: [{ type: "image", attachmentId: "summary-orphan.png", mimeType: "image/png" }], + }), + }; + const chatRepo = { deleteAttachment: vi.fn().mockResolvedValue(undefined) } as any; + const service = new CompactService(modelService, orchestrator, chatRepo); + + await expect(service.summarizeContent("content", "extract")).resolves.toMatchObject({ content: "summary" }); + expect(chatRepo.deleteAttachment).toHaveBeenCalledWith("summary-orphan.png"); + }); +}); diff --git a/src/app/service/agent/service_worker/compact_service.ts b/src/app/service/agent/service_worker/compact_service.ts index b09a47362..17ebae48c 100644 --- a/src/app/service/agent/service_worker/compact_service.ts +++ b/src/app/service/agent/service_worker/compact_service.ts @@ -6,6 +6,7 @@ import type { ToolCall, ContentBlock, ToolDefinition, + TokenUsage, } from "@App/app/service/agent/core/types"; import { COMPACT_SYSTEM_PROMPT, @@ -14,6 +15,11 @@ import { } from "@App/app/service/agent/core/compact_prompt"; import { uuidv4 } from "@App/pkg/utils/uuid"; import type { AgentModelService } from "./model_service"; +import { elideUntilWithinBudget, estimateRequestTokens } from "@App/app/service/agent/core/context_elision"; +import { getInputTokenBudget } from "@App/app/service/agent/core/model_context"; +import { throwIfAborted } from "@App/app/service/agent/core/abort_utils"; +import { prepareAttachmentSnapshot, type AttachmentSnapshot } from "@App/app/service/agent/core/attachment_resolver"; +import { isLegacyGeneration, retainedSummaryAttachmentIds } from "@App/app/service/agent/core/persisted_messages"; /** LLM 调用结果(与 AgentService.callLLM 返回值一致) */ interface CompactLLMResult { @@ -33,12 +39,19 @@ interface CompactLLMResult { export interface CompactOrchestrator { callLLM( model: AgentModelConfig, - params: { messages: ChatRequest["messages"]; tools?: ToolDefinition[]; cache?: boolean }, + params: { + messages: ChatRequest["messages"]; + tools?: ToolDefinition[]; + cache?: boolean; + attachmentSnapshot?: AttachmentSnapshot; + }, sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal ): Promise; } +export type SummarizeResult = { content: string; usage?: TokenUsage }; + export class CompactService { constructor( private modelService: AgentModelService, @@ -46,14 +59,26 @@ export class CompactService { private chatRepo: AgentChatRepo ) {} + private async releaseIgnoredContentBlocks(result: CompactLLMResult): Promise { + await Promise.all( + (result.contentBlocks || []) + .filter((block) => block.type !== "text") + .map((block) => this.chatRepo.deleteAttachment(block.attachmentId).catch(() => {})) + ); + } + /** 自动 compact:汇总对话历史为 summary 并替换 currentMessages */ async autoCompact( conversationId: string, + conversationGeneration: string, model: AgentModelConfig, currentMessages: ChatRequest["messages"], sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal - ): Promise { + ): Promise { + throwIfAborted(signal); + const snapshot = await this.chatRepo.getMessageSnapshot(conversationId, conversationGeneration); + // 构建摘要请求(用 currentMessages 而非从 repo 加载,因为可能有未持久化的 tool 消息) const summaryMessages: ChatRequest["messages"] = []; summaryMessages.push({ role: "system", content: COMPACT_SYSTEM_PROMPT }); @@ -64,39 +89,81 @@ export class CompactService { } summaryMessages.push({ role: "user", content: buildCompactUserPrompt() }); + const attachmentSnapshot = await prepareAttachmentSnapshot( + summaryMessages, + model, + (id) => this.chatRepo.getAttachment(id), + signal + ); + const inputBudget = getInputTokenBudget(model); + const effectiveWindow = Math.max(1, Math.floor(inputBudget / 0.9)); + if (!elideUntilWithinBudget(summaryMessages, effectiveWindow, undefined, 0.9, attachmentSnapshot.sizes, model)) { + throw Object.assign(new Error("Conversation history is too large to compact"), { + errorCode: "context_too_large", + }); + } + // 调用 LLM 获取摘要(不带 tools,不发流式事件给 UI) const noopSendEvent = () => {}; const result = await this.orchestrator.callLLM( model, - { messages: summaryMessages, cache: false }, + { messages: summaryMessages, cache: false, attachmentSnapshot }, noopSendEvent, signal ); + await this.releaseIgnoredContentBlocks(result); - const summary = extractSummary(result.content); + const throwWithUsage = (error: unknown): never => { + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { usage: result.usage }); + }; - // 替换 currentMessages(保留 system,替换其余为摘要) - const systemMsg = currentMessages.find((m) => m.role === "system"); - currentMessages.length = 0; - if (systemMsg) currentMessages.push(systemMsg); - currentMessages.push({ role: "user", content: `[Conversation Summary]\n\n${summary}` }); + // LLM 调用期间可能已被 stop:落地前必须重新检查,避免取消之后仍持久化/广播 compact_done + if (signal.aborted) throwWithUsage(new Error("Aborted")); + + const summary = extractSummary(result.content); - // 持久化 + // 持久化:先写盘、成功后才允许覆写内存中的 currentMessages。 + // OPFS 的 createWritable() 是事务性的,signal 在 close() 落定前 abort 时会放弃这次 + // 整份覆写而不提交(见 opfs_repo.ts writeJsonFile);只有 saveMessages 真正成功后, + // 再让内存态 currentMessages 反映同一份摘要,避免"写盘失败但内存已被摘要顶替"的不一致 const summaryMessage = { id: uuidv4(), conversationId, role: "user" as const, content: `[Conversation Summary]\n\n${summary}`, + ownedAttachmentIds: retainedSummaryAttachmentIds( + summary, + snapshot.messages, + isLegacyGeneration(conversationGeneration) + ), createtime: Date.now(), }; - await this.chatRepo.saveMessages(conversationId, [summaryMessage]); + try { + await this.chatRepo.saveMessages(conversationId, [summaryMessage], signal, { + generation: conversationGeneration, + expectedRevision: snapshot.revision, + }); + } catch (error) { + throwWithUsage(error); + } + // 写入由 revision CAS 线性化。提交后 Stop 不得用旧快照回滚,否则会覆盖更新的历史。 + if (signal.aborted) throwWithUsage(new Error("Aborted")); + + // 替换 currentMessages(保留 system,替换其余为摘要)——只有走到这里才说明落盘已提交 + const systemMsg = currentMessages.find((m) => m.role === "system"); + currentMessages.length = 0; + if (systemMsg) currentMessages.push(systemMsg); + currentMessages.push({ role: "user", content: `[Conversation Summary]\n\n${summary}` }); // 通知 UI sendEvent({ type: "compact_done", summary, originalCount: -1 }); + return result.usage; } /** 使用 summary 模型对任意内容做提取/总结(供 tab 工具使用) */ - async summarizeContent(content: string, prompt: string): Promise { + async summarizeContent(content: string, prompt: string, signal?: AbortSignal): Promise { + throwIfAborted(signal); + const model = await this.modelService.getSummaryModel(); const messages: ChatRequest["messages"] = [ @@ -111,18 +178,33 @@ export class CompactService { }, ]; + const inputBudget = getInputTokenBudget(model); + const estimatedInputTokens = estimateRequestTokens(messages, undefined, undefined, model); + if (estimatedInputTokens > inputBudget) { + throw Object.assign(new Error("Summarization content exceeds the summary model context window"), { + errorCode: "context_too_large", + estimatedInputTokens, + }); + } + const noopSendEvent = () => {}; - const controller = new AbortController(); try { const result = await this.orchestrator.callLLM( model, { messages, cache: false }, noopSendEvent, - controller.signal + signal ?? new AbortController().signal ); - return result.content; + await this.releaseIgnoredContentBlocks(result); + return { content: result.content, usage: result.usage }; } catch (e: any) { - throw new Error(`Summarization failed: ${e.message}`); + if (e?.errorCode || e?.message === "Aborted") { + throw e; + } + throw Object.assign(new Error(`Summarization failed: ${e.message}`), { + usage: e?.usage, + cause: e, + }); } } } diff --git a/src/app/service/agent/service_worker/llm.test.ts b/src/app/service/agent/service_worker/llm.test.ts index ff480bf9b..91be84b5f 100644 --- a/src/app/service/agent/service_worker/llm.test.ts +++ b/src/app/service/agent/service_worker/llm.test.ts @@ -1,5 +1,6 @@ import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; import { createTestService, makeSSEResponse, makeTextResponse } from "./test-helpers"; +import { LLMClient } from "./llm_client"; // ---- callLLM 相关测试(通过 callLLMWithToolLoop 间接测试) ---- @@ -69,7 +70,7 @@ describe("callLLM 流式响应解析", () => { fetchSpy.mockResolvedValueOnce( makeSSEResponse([ `data: {"choices":[{"delta":{"content":"你好"}}]}\n\n`, - `data: {"choices":[{"delta":{"content":"世界"}}]}\n\n`, + `data: {"choices":[{"delta":{"content":"世界"},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`, ]) ); @@ -195,7 +196,7 @@ describe("callLLM 流式响应解析", () => { } as unknown as Response); fetchSpy.mockResolvedValueOnce( makeSSEResponse([ - `data: {"choices":[{"delta":{"content":"重试成功"}}]}\n\n`, + `data: {"choices":[{"delta":{"content":"重试成功"},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":5,"completion_tokens":2}}\n\n`, ]) ); @@ -241,7 +242,7 @@ describe("callLLM 流式响应解析", () => { // 第二次成功 fetchSpy.mockResolvedValueOnce( makeSSEResponse([ - `data: {"choices":[{"delta":{"content":"恢复了"}}]}\n\n`, + `data: {"choices":[{"delta":{"content":"恢复了"},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":10,"completion_tokens":3}}\n\n`, ]) ); @@ -333,6 +334,96 @@ describe("callLLM 流式响应解析", () => { const doneEvents = events.filter((e: any) => e.type === "done"); expect(doneEvents).toHaveLength(0); }); + + it("图片落盘期间取消时应保留已知 usage 并清理稍后完成的孤儿附件", async () => { + const controller = new AbortController(); + let finishSave!: () => void; + const savePending = new Promise((resolve) => { + finishSave = resolve; + }); + const repo = { + getAttachment: vi.fn().mockResolvedValue(null), + saveAttachment: vi.fn().mockImplementation(() => savePending), + deleteAttachment: vi.fn().mockResolvedValue(undefined), + } as any; + const client = new LLMClient(repo); + fetchSpy.mockResolvedValueOnce( + makeAnthropicSSEResponse([ + { event: "message_start", data: { message: { usage: { input_tokens: 15 } } } }, + { + event: "content_block_start", + data: { index: 0, content_block: { type: "image", source: { type: "base64", media_type: "image/png" } } }, + }, + { event: "content_block_delta", data: { index: 0, delta: { type: "image_delta", data: "AAAA" } } }, + { event: "content_block_stop", data: { index: 0 } }, + { event: "message_delta", data: { usage: { output_tokens: 4 } } }, + ]) + ); + + const call = client.callLLM( + { + id: "anthropic-image", + name: "Anthropic image", + provider: "anthropic", + apiBaseUrl: "https://api.anthropic.com", + apiKey: "test", + model: "claude-test", + }, + { messages: [{ role: "user", content: "draw" }] }, + vi.fn(), + controller.signal + ); + await vi.waitFor(() => expect(repo.saveAttachment).toHaveBeenCalledOnce()); + const attachmentId = repo.saveAttachment.mock.calls[0][0]; + controller.abort(); + + await expect(call).rejects.toMatchObject({ + message: "Aborted", + usage: { inputTokens: 15, outputTokens: 4 }, + }); + finishSave(); + await vi.waitFor(() => expect(repo.deleteAttachment).toHaveBeenCalledWith(attachmentId)); + }); + + it("生成图片保存失败时应在结果里携带可见 warning,而不是静默丢弃", async () => { + const repo = { + getAttachment: vi.fn().mockResolvedValue(null), + saveAttachment: vi.fn().mockRejectedValue(new Error("disk full")), + deleteAttachment: vi.fn().mockResolvedValue(undefined), + } as any; + const client = new LLMClient(repo); + fetchSpy.mockResolvedValueOnce( + makeAnthropicSSEResponse([ + { event: "message_start", data: { message: { usage: { input_tokens: 15 } } } }, + { + event: "content_block_start", + data: { index: 0, content_block: { type: "image", source: { type: "base64", media_type: "image/png" } } }, + }, + { event: "content_block_delta", data: { index: 0, delta: { type: "image_delta", data: "AAAA" } } }, + { event: "content_block_stop", data: { index: 0 } }, + { event: "message_delta", data: { usage: { output_tokens: 4 } } }, + ]) + ); + + const result = await client.callLLM( + { + id: "anthropic-image", + name: "Anthropic image", + provider: "anthropic", + apiBaseUrl: "https://api.anthropic.com", + apiKey: "test", + model: "claude-test", + }, + { messages: [{ role: "user", content: "draw" }] }, + vi.fn(), + new AbortController().signal + ); + + // usage 不能因为图片保存失败而丢失,结果里必须携带可见 warning + expect(result.usage).toMatchObject({ inputTokens: 15, outputTokens: 4 }); + expect(result.contentBlocks).toBeUndefined(); + expect(result.warning).toMatch(/failed to save/i); + }); }); // ---- callLLMWithToolLoop 场景补充 ---- @@ -350,18 +441,38 @@ describe("callLLMWithToolLoop 工具调用循环", () => { function makeToolCallResponse(toolCalls: Array<{ id: string; name: string; arguments: string }>): Response { const chunks: string[] = []; - for (const tc of toolCalls) { + toolCalls.forEach((tc, i) => { chunks.push( `data: {"choices":[{"delta":{"tool_calls":[{"id":"${tc.id}","function":{"name":"${tc.name}","arguments":""}}]}}]}\n\n` ); + const isLast = i === toolCalls.length - 1; chunks.push( - `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":${JSON.stringify(tc.arguments)}}}]}}]}\n\n` + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":${JSON.stringify(tc.arguments)}}}]}${isLast ? ', "finish_reason":"tool_calls"' : ""}}]}\n\n` ); - } + }); chunks.push(`data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`); return makeSSEResponse(chunks); } + // 与 makeToolCallResponse 相同,但 prompt_tokens 可自定义,用于驱动 usageRatio 跨越裁剪阈值 + function makeToolCallResponseWithUsage( + toolCalls: Array<{ id: string; name: string; arguments: string }>, + promptTokens: number + ): Response { + const chunks: string[] = []; + toolCalls.forEach((tc, i) => { + chunks.push( + `data: {"choices":[{"delta":{"tool_calls":[{"id":"${tc.id}","function":{"name":"${tc.name}","arguments":""}}]}}]}\n\n` + ); + const isLast = i === toolCalls.length - 1; + chunks.push( + `data: {"choices":[{"delta":{"tool_calls":[{"function":{"arguments":${JSON.stringify(tc.arguments)}}}]}${isLast ? ', "finish_reason":"tool_calls"' : ""}}]}\n\n` + ); + }); + chunks.push(`data: {"usage":{"prompt_tokens":${promptTokens},"completion_tokens":5}}\n\n`); + return makeSSEResponse(chunks); + } + function createMockSender() { const sentMessages: any[] = []; const mockConn = { @@ -416,10 +527,10 @@ describe("callLLMWithToolLoop 工具调用循环", () => { expect(events.some((e: any) => e.type === "new_message")).toBe(true); expect(events.some((e: any) => e.type === "done")).toBe(true); - // assistant 消息应持久化(tool_calls 和最终文本各一条) - const appendCalls = mockRepo.appendMessage.mock.calls; - const assistantCalls = appendCalls.filter((c: any) => c[0].role === "assistant"); - expect(assistantCalls).toHaveLength(2); // tool_call + final text + // 工具 assistant 与结果原子提交,最终文本单独追加。 + expect(mockRepo.commitToolRound).toHaveBeenCalledTimes(1); + expect(mockRepo.commitToolRound.mock.calls[0][0]).toMatchObject({ role: "assistant" }); + expect(mockRepo.appendMessage.mock.calls.filter((c: any) => c[0].role === "assistant")).toHaveLength(1); // fetch 应调用 2 次 expect(fetchSpy).toHaveBeenCalledTimes(2); @@ -468,37 +579,111 @@ describe("callLLMWithToolLoop 工具调用循环", () => { registry.unregisterBuiltin("counter"); }); - it("超过 maxIterations:sendEvent 收到 max_iterations 错误", async () => { + it("上下文历史中的错误占位消息不应被重放给 LLM", async () => { const { service, mockRepo } = createTestService(); - const { sender, sentMessages } = createMockSender(); + const { sender } = createMockSender(); + + mockRepo.listConversations.mockResolvedValue([BASE_CONV]); + // 历史中包含一条持久化的错误占位消息(content 为空字符串) + mockRepo.getMessages.mockResolvedValue([ + { id: "u1", conversationId: "conv-1", role: "user", content: "第一条消息", createtime: 1 }, + { + id: "a1", + conversationId: "conv-1", + role: "assistant", + content: "", + error: "Reply was generated but failed to save", + errorCode: "persist_failed", + createtime: 2, + }, + ]); + + fetchSpy.mockResolvedValueOnce(makeTextResponse("好的,继续")); + + await (service as any).handleConversationChat({ conversationId: "conv-1", message: "请继续。" }, sender); + + const reqInit = fetchSpy.mock.calls[0][1] as RequestInit; + const body = JSON.parse(reqInit.body as string); + // 出站请求中不应包含空 content 且无 tool_calls 的 assistant 消息(即错误占位消息) + const emptyAssistantMsgs = body.messages.filter( + (m: any) => m.role === "assistant" && m.content === "" && !m.tool_calls + ); + expect(emptyAssistantMsgs).toHaveLength(0); + + // 正常的历史用户消息与新消息应仍然存在 + const userMsgs = body.messages.filter((m: any) => m.role === "user"); + expect(userMsgs.map((m: any) => m.content)).toEqual(["第一条消息", "请继续。"]); + }); + + it("上下文占用跨过旧裁剪阈值时仍应完整保留全部 tool 结果", async () => { + const { service, mockRepo, mockModelRepo } = createTestService(); + const { sender } = createMockSender(); + + // 最后一轮保持在完整上下文窗口的 80% 以下,避免自动 Compact 干扰工具结果透传断言。 + mockModelRepo.getModel.mockResolvedValue({ + id: "test-openai", + name: "Test", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", + contextWindow: 120000, + }); + + // 使用两个不同名称的工具交替调用,避免触发 tool_call_guard 的重复调用检测 + // (相同工具名连续出现会命中循环检测,暂停询问用户,与本测试无关) const registry = (service as any).toolRegistry; + let callCount = 0; + const execute = async () => { + callCount++; + return `count=${callCount}`; + }; registry.registerBuiltin( - { name: "loop", description: "Loop", parameters: { type: "object", properties: {} } }, - { execute: async () => "ok" } + { name: "counterA", description: "Count A", parameters: { type: "object", properties: {} } }, + { execute } + ); + registry.registerBuiltin( + { name: "counterB", description: "Count B", parameters: { type: "object", properties: {} } }, + { execute } ); mockRepo.listConversations.mockResolvedValue([BASE_CONV]); mockRepo.getMessages.mockResolvedValue([]); - // maxIterations=1 但 LLM 一直返回 tool_call - fetchSpy.mockResolvedValueOnce(makeToolCallResponse([{ id: "c1", name: "loop", arguments: "{}" }])); - - await (service as any).handleConversationChat( - { conversationId: "conv-1", message: "test", maxIterations: 1 }, - sender - ); - - const events = sentMessages.map((m) => m.data); - const errorEvents = events.filter((e: any) => e.type === "error"); - expect(errorEvents).toHaveLength(1); - expect(errorEvents[0].message).toContain("maximum iterations"); - expect(errorEvents[0].errorCode).toBe("max_iterations"); + // 覆盖此前会触发 40%/60% 滑动裁剪的输入占用区间。 + const usages = [10000, 10000, 10000, 10000, 50000, 70000]; + for (let i = 0; i < usages.length; i++) { + const toolName = i % 2 === 0 ? "counterA" : "counterB"; + // 每轮参数不同,避免触发 tool_call_guard 的“相同参数重复调用”检测 + fetchSpy.mockResolvedValueOnce( + makeToolCallResponseWithUsage([{ id: `c${i + 1}`, name: toolName, arguments: `{"round":${i}}` }], usages[i]) + ); + } + // 第 7 轮:最终文本,结束循环 + fetchSpy.mockResolvedValueOnce(makeTextResponse("done")); - // fetch 只调用 1 次(maxIterations=1) - expect(fetchSpy).toHaveBeenCalledTimes(1); + await (service as any).handleConversationChat({ conversationId: "conv-1", message: "test" }, sender); - registry.unregisterBuiltin("loop"); + expect(fetchSpy).toHaveBeenCalledTimes(7); + + // 最后一次请求必须保持 append-only 前缀,所有工具结果均按原值发送。 + const lastCall = fetchSpy.mock.calls[fetchSpy.mock.calls.length - 1]; + const lastBody = JSON.parse((lastCall[1] as RequestInit).body as string); + const toolMessages = lastBody.messages.filter((m: any) => m.role === "tool"); + + expect(toolMessages).toHaveLength(6); + expect(toolMessages.map((message: any) => message.content)).toEqual([ + "count=1", + "count=2", + "count=3", + "count=4", + "count=5", + "count=6", + ]); + + registry.unregisterBuiltin("counterA"); + registry.unregisterBuiltin("counterB"); }); it("工具执行后附件回写:toolCalls 被更新", async () => { @@ -570,6 +755,13 @@ describe("callLLMWithToolLoop 工具调用循环", () => { storedMessages.length = 0; storedMessages.push(...msgs.map((m) => structuredClone(m))); }); + mockRepo.updateMessage.mockImplementation(async (msg: any) => { + const index = storedMessages.findIndex((m) => m.id === msg.id); + if (index >= 0) storedMessages[index] = structuredClone(msg); + }); + mockRepo.commitToolRound.mockImplementation(async (assistant: any, toolMessages: any[]) => { + storedMessages.push(structuredClone(assistant), ...toolMessages.map((message) => structuredClone(message))); + }); fetchSpy.mockResolvedValueOnce(makeToolCallResponse([{ id: "call_1", name: "echo", arguments: '{"msg":"hi"}' }])); fetchSpy.mockResolvedValueOnce(makeTextResponse("done")); @@ -637,19 +829,53 @@ describe("callLLMWithToolLoop 工具调用循环", () => { // 应有两个 tool_call_complete const completeEvents = events.filter((e: any) => e.type === "tool_call_complete"); expect(completeEvents).toHaveLength(2); + expect(completeEvents.every((event: any) => event.status === "completed")).toBe(true); expect(completeEvents.find((e: any) => e.id === "call_a").result).toBe("a: hello"); expect(completeEvents.find((e: any) => e.id === "call_b").result).toBe("b: world"); - // 持久化的 assistant 消息应包含两个 toolCalls - const assistantMsgs = mockRepo.appendMessage.mock.calls - .map((c: any) => c[0]) - .filter((m: any) => m.role === "assistant" && m.toolCalls); - expect(assistantMsgs).toHaveLength(1); - expect(assistantMsgs[0].toolCalls).toHaveLength(2); + // 原子提交的 assistant 消息应包含两个 toolCalls。 + expect(mockRepo.commitToolRound).toHaveBeenCalledTimes(1); + expect(mockRepo.commitToolRound.mock.calls[0][0].toolCalls).toHaveLength(2); + expect(mockRepo.commitToolRound.mock.calls[0][1]).toHaveLength(2); expect(fetchSpy).toHaveBeenCalledTimes(2); registry.unregisterBuiltin("tool_a"); registry.unregisterBuiltin("tool_b"); }); + + it("最终回复持久化多次重试仍失败时应报结构化错误而不是假装 done", async () => { + vi.useFakeTimers(); + try { + const { service, mockRepo } = createTestService(); + const { sender, sentMessages } = createMockSender(); + + mockRepo.listConversations.mockResolvedValue([BASE_CONV]); + mockRepo.getMessages.mockResolvedValue([]); + // 只让最终 assistant 消息持久化失败(模拟 OPFS 写入故障),user 消息正常落库, + // 这样才能验证的是"最终成功回复的落库失败处理"而不是更早的用户消息落库失败 + mockRepo.appendMessage.mockImplementation(async (msg: any) => { + if (msg.role === "assistant") throw new Error("disk write failed"); + }); + + fetchSpy.mockResolvedValueOnce(makeTextResponse("done")); + + const chatPromise = (service as any).handleConversationChat( + { conversationId: "conv-1", message: "test" }, + sender + ); + // 有限重试之间有退避延迟(200ms/400ms),推进假定时器让它们落定 + await vi.advanceTimersByTimeAsync(1000); + await chatPromise; + + const events = sentMessages.map((m) => m.data); + // 不应报告 done:持久化最终失败,不能对外承诺"回复已保存" + expect(events.some((e: any) => e.type === "done")).toBe(false); + const errorEvent = events.find((e: any) => e.type === "error"); + expect(errorEvent).toBeDefined(); + expect(errorEvent.errorCode).toBe("persist_failed"); + } finally { + vi.useRealTimers(); + } + }); }); diff --git a/src/app/service/agent/service_worker/llm_client.ts b/src/app/service/agent/service_worker/llm_client.ts index ec076f115..dbd8481bb 100644 --- a/src/app/service/agent/service_worker/llm_client.ts +++ b/src/app/service/agent/service_worker/llm_client.ts @@ -8,7 +8,7 @@ import type { ToolDefinition, } from "@App/app/service/agent/core/types"; import { providerRegistry } from "@App/app/service/agent/core/providers"; -import { resolveAttachments } from "@App/app/service/agent/core/attachment_resolver"; +import { prepareAttachmentSnapshot, type AttachmentSnapshot } from "@App/app/service/agent/core/attachment_resolver"; import { generateAttachmentId } from "@App/app/service/agent/core/providers/content_utils"; export interface LLMCallResult { @@ -22,6 +22,10 @@ export interface LLMCallResult { cacheReadInputTokens?: number; }; contentBlocks?: ContentBlock[]; + /** Non-fatal issue surfaced alongside an otherwise-successful result, e.g. a generated image that + * failed to persist. The round still resolves so accumulated usage/text aren't discarded, but callers + * should show this to the user and persist it onto the assistant message. */ + warning?: string; } export class LLMClient { @@ -32,7 +36,12 @@ export class LLMClient { */ async callLLM( model: AgentModelConfig, - params: { messages: ChatRequest["messages"]; tools?: ToolDefinition[]; cache?: boolean }, + params: { + messages: ChatRequest["messages"]; + tools?: ToolDefinition[]; + cache?: boolean; + attachmentSnapshot?: AttachmentSnapshot; + }, sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal ): Promise { @@ -45,9 +54,9 @@ export class LLMClient { }; // 预解析消息中 ContentBlock 引用的 attachmentId → base64 - const attachmentResolver = await resolveAttachments(params.messages, model, (id) => - this.chatRepo.getAttachment(id) - ); + const attachmentSnapshot = + params.attachmentSnapshot || + (await prepareAttachmentSnapshot(params.messages, model, (id) => this.chatRepo.getAttachment(id), signal)); const provider = providerRegistry.get(model.provider); if (!provider) { @@ -56,7 +65,7 @@ export class LLMClient { const { url, init } = await provider.buildRequest({ model, request: chatRequest, - resolver: attachmentResolver, + resolver: attachmentSnapshot.resolver, }); // 带重试的 LLM 调用,最多重试 5 次,间隔递增:10s, 10s, 20s, 20s, 30s @@ -135,6 +144,26 @@ export class LLMClient { const pendingImageSaves: Array<{ block: ContentBlock & { type: "image" }; data: string }> = []; return new Promise((resolve, reject) => { + // 最后一道保险:即使 provider parser 出现未预见的静默完成路径,也不能让这个 Promise 永远挂起 + let settled = false; + const settleOnce = (fn: () => void) => { + if (settled) return; + settled = true; + signal.removeEventListener("abort", onAbortSafeguard); + fn(); + }; + // parseStream 自身在 abort 时会 reject,并可能携带这一轮已知的部分 usage(见 openai.ts/ + // anthropic.ts)。这里的 signal 监听只是最后一道保险,不能抢在 parseStream 的 reject 之前 + // 立即 settle——那样会用一个不带 usage 的裸 Error 抢占更有信息量的那个 reject。 + // 延迟到下一个宏任务,给 parseStream 的 reject 一个先落定的机会,本身仍然是安全网, + // 不依赖它必定生效。 + const onAbortSafeguard = () => { + setTimeout(() => settleOnce(() => reject(Object.assign(new Error("Aborted"), { usage }))), 0); + }; + signal.addEventListener("abort", onAbortSafeguard, { once: true }); + const resolveOnce: typeof resolve = (value) => settleOnce(() => resolve(value)); + const rejectOnce: typeof reject = (reason) => settleOnce(() => reject(reason)); + const onEvent = (event: ChatStreamEvent) => { // 只转发流式内容事件,done 和 error 由 callLLMWithToolLoop 统一管理 // 避免在 tool calling 循环中提前发送 done 导致客户端过早 resolve @@ -188,17 +217,33 @@ export class LLMClient { usage = event.usage; } - // 保存模型生成的图片到 OPFS,然后转发事件 + // 保存模型生成的图片到 OPFS,然后转发事件。 + // abort 安全:settled 一旦为 true(外层 Promise 已因 abort 落定),后续保存产生的 + // 附件不会再被任何持久化的 assistant 消息引用,是孤儿文件;已发出的 content_block_complete + // 也会晚于终态事件到达客户端。因此每一步都先检查 settled,中途发现已 settle 就 + // 停止继续保存/发送,并清理这一轮已经落盘但用不上的附件。 const finalize = async () => { const savedBlocks: ContentBlock[] = []; + // 本轮所有成功落盘的附件 id(不止是取消时正在保存的那一个):一旦 settled, + // finalize() 的返回值不会再被使用(resolveOnce/rejectOnce 已是 no-op), + // 这一轮已经保存的所有附件都变成孤儿文件,必须全部清理,不能只删除取消时 + // 正在保存的那一个而漏掉更早已经保存成功的 + const allSavedIds: string[] = []; + // 保存失败的图片没有 markdown 原文可回退:之前静默丢弃,用户会拿到一个成功、 + // 计费的回复但缺图(甚至纯图回复时 content 为空)。记录数量,随结果一起报告 + // 给调用方持久化/展示,而不是假装什么都没发生 + let failedImageSaves = 0; for (const pending of pendingImageSaves) { + if (settled) break; try { await this.chatRepo.saveAttachment(pending.block.attachmentId, pending.data); + allSavedIds.push(pending.block.attachmentId); + if (settled) break; savedBlocks.push(pending.block); // 转发不含 data 的 content_block_complete 事件给 UI sendEvent({ type: "content_block_complete", block: pending.block }); } catch { - // 保存失败忽略 + failedImageSaves++; } } @@ -206,13 +251,15 @@ export class LLMClient { const imgRegex = /!\[([^\]]*)\]\((data:image\/([^;]+);base64,[A-Za-z0-9+/=\s]+)\)/g; let match; let cleanedContent = content; - while ((match = imgRegex.exec(content)) !== null) { + while (!settled && (match = imgRegex.exec(content)) !== null) { const [fullMatch, alt, dataUrl, subtype] = match; const mimeType = `image/${subtype}`; const ext = subtype || "png"; const blockId = generateAttachmentId(ext); try { await this.chatRepo.saveAttachment(blockId, dataUrl); + allSavedIds.push(blockId); + if (settled) break; const block: ContentBlock = { type: "image", attachmentId: blockId, @@ -231,29 +278,47 @@ export class LLMClient { content = cleanedContent.replace(/\n{3,}/g, "\n\n").trim(); } - return savedBlocks.length > 0 ? savedBlocks : undefined; + if (settled && allSavedIds.length > 0) { + await Promise.all(allSavedIds.map((id) => this.chatRepo.deleteAttachment(id).catch(() => {}))); + } + + return { + contentBlocks: savedBlocks.length > 0 ? savedBlocks : undefined, + warning: + failedImageSaves > 0 + ? `${failedImageSaves} generated image(s) failed to save and were lost.` + : undefined, + }; }; finalize() - .then((contentBlocks) => { - resolve({ + .then(({ contentBlocks, warning }) => { + resolveOnce({ content, thinking: thinking || undefined, toolCalls: toolCalls.length > 0 ? toolCalls : undefined, usage, contentBlocks, + warning, }); }) - .catch(reject); + .catch(rejectOnce); break; } case "error": - reject(new Error(event.message)); + // 保留 usage/errorCode/durationMs 等字段,不能转成裸 Error 丢掉这些信息 + rejectOnce( + Object.assign(new Error(event.message), { + errorCode: event.errorCode, + usage: event.usage, + durationMs: event.durationMs, + }) + ); break; } }; - parseStream(reader, onEvent, signal).catch(reject); + parseStream(reader, onEvent, signal).catch(rejectOnce); }); } } diff --git a/src/app/service/agent/service_worker/mcp.test.ts b/src/app/service/agent/service_worker/mcp.test.ts index ade68726d..6d63ee4ed 100644 --- a/src/app/service/agent/service_worker/mcp.test.ts +++ b/src/app/service/agent/service_worker/mcp.test.ts @@ -146,6 +146,153 @@ describe("MCPService", () => { await service.disconnectServer(server.id); expect(toolRegistry.getDefinitions().length).toBe(0); }); + + it("断开服务器应等待客户端关闭完成", async () => { + let releaseClose!: () => void; + const close = vi.fn( + () => + new Promise((resolve) => { + releaseClose = resolve; + }) + ); + const baseFactory = createMockClientFactory(); + const clientFactory: MCPClientFactory = (config) => { + const client = baseFactory(config); + client.close = close; + return client; + }; + const repo = createMockRepo(); + service = new MCPService(toolRegistry, { clientFactory, repo }); + const server = (await service.handleMCPApi({ + action: "addServer", + config: { name: "AsyncClose", url: "https://mcp.test.com", enabled: false }, + scriptUuid: "test", + })) as any; + await service.connectServer(server.id); + + let settled = false; + const disconnect = service.disconnectServer(server.id).then(() => { + settled = true; + }); + await vi.waitFor(() => expect(close).toHaveBeenCalledOnce()); + expect(settled).toBe(false); + + releaseClose(); + await disconnect; + expect(settled).toBe(true); + }); + + it("列出工具失败时应关闭已初始化的客户端", async () => { + const close = vi.fn().mockResolvedValue(undefined); + const baseFactory = createMockClientFactory(); + const clientFactory: MCPClientFactory = (config) => { + const client = baseFactory(config); + client.listTools = vi.fn().mockRejectedValue(new Error("list failed")); + client.close = close; + return client; + }; + const repo = createMockRepo(); + service = new MCPService(toolRegistry, { clientFactory, repo }); + const server = (await service.handleMCPApi({ + action: "addServer", + config: { name: "BrokenTools", url: "https://mcp.test.com", enabled: false }, + scriptUuid: "test", + })) as any; + + await expect(service.connectServer(server.id)).rejects.toThrow("list failed"); + expect(close).toHaveBeenCalledOnce(); + expect(toolRegistry.getDefinitions()).toHaveLength(0); + }); + + it("服务器名称变更后应重连并更新工具注册", async () => { + const close = vi.fn().mockResolvedValue(undefined); + const baseFactory = createMockClientFactory(); + const clientFactory: MCPClientFactory = (config) => { + const client = baseFactory(config); + client.close = close; + return client; + }; + const repo = createMockRepo(); + service = new MCPService(toolRegistry, { clientFactory, repo }); + const server = (await service.handleMCPApi({ + action: "addServer", + config: { name: "Before", url: "https://before.example.com", enabled: true }, + scriptUuid: "test", + })) as any; + expect(toolRegistry.getDefinitions()[0].name).toContain("before"); + + await service.handleMCPApi({ + action: "updateServer", + id: server.id, + config: { name: "After", url: "https://after.example.com" }, + scriptUuid: "test", + }); + + expect(close).toHaveBeenCalledOnce(); + expect(toolRegistry.getDefinitions()).toHaveLength(1); + expect(toolRegistry.getDefinitions()[0].name).toContain("after"); + }); + + it("名称清洗后相同的服务器仍应拥有独立工具名", async () => { + const first = (await service.handleMCPApi({ + action: "addServer", + config: { name: "a-b", url: "https://first.example.com", enabled: false }, + scriptUuid: "test", + })) as any; + const second = (await service.handleMCPApi({ + action: "addServer", + config: { name: "a_b", url: "https://second.example.com", enabled: false }, + scriptUuid: "test", + })) as any; + + await service.connectServer(first.id); + await service.connectServer(second.id); + + const names = toolRegistry.getDefinitions().map((definition) => definition.name); + expect(names).toHaveLength(2); + expect(new Set(names).size).toBe(2); + + await service.disconnectServer(first.id); + expect(toolRegistry.getDefinitions()).toHaveLength(1); + }); + + it("并发连接同一服务器时应复用同一个连接操作", async () => { + const baseFactory = createMockClientFactory(); + const clientFactory = vi.fn((config) => baseFactory(config)); + const repo = createMockRepo(); + service = new MCPService(toolRegistry, { clientFactory, repo }); + const server = (await service.handleMCPApi({ + action: "addServer", + config: { name: "Concurrent", url: "https://mcp.test.com", enabled: false }, + scriptUuid: "test", + })) as any; + + await Promise.all([service.connectServer(server.id), service.connectServer(server.id)]); + + expect(clientFactory).toHaveBeenCalledOnce(); + expect(toolRegistry.getDefinitions()).toHaveLength(1); + }); + + it("原始服务器 ID 不同时即使清洗结果相同也应保持工具名唯一", async () => { + const servers = new Map([ + ["a-b", { id: "a-b", name: "same", url: "https://first.example.com", enabled: false }], + ["a_b", { id: "a_b", name: "same", url: "https://second.example.com", enabled: false }], + ]); + const repo = { + listServers: vi.fn(async () => [...servers.values()]), + getServer: vi.fn(async (id: string) => servers.get(id)), + saveServer: vi.fn(async (config: any) => servers.set(config.id, config)), + removeServer: vi.fn(async (id: string) => servers.delete(id)), + } as unknown as MCPServerRepo; + service = new MCPService(toolRegistry, { clientFactory: createMockClientFactory(), repo }); + + await service.connectServer("a-b"); + await service.connectServer("a_b"); + + const names = toolRegistry.getDefinitions().map((definition) => definition.name); + expect(names).toHaveLength(2); + expect(new Set(names).size).toBe(2); + }); }); describe("handleMCPApi - listTools", () => { diff --git a/src/app/service/agent/service_worker/mcp.ts b/src/app/service/agent/service_worker/mcp.ts index de783228b..be1f4f325 100644 --- a/src/app/service/agent/service_worker/mcp.ts +++ b/src/app/service/agent/service_worker/mcp.ts @@ -5,11 +5,12 @@ import { MCPServerRepo } from "@App/app/repo/mcp_server_repo"; import type { ToolRegistry } from "@App/app/service/agent/core/tool_registry"; import { uuidv4 } from "@App/pkg/utils/uuid"; -// 将服务器名和工具名合成为全局唯一的工具名 -function mcpToolName(serverName: string, toolName: string): string { - // 使用小写字母和下划线,避免特殊字符 +// 将服务器 ID、名称和工具名合成为全局唯一的工具名 +function mcpToolName(serverId: string, serverName: string, toolName: string): string { + // 名称仅用于可读性;服务器 ID 用码点编码,避免 `a-b` 与 `a_b` 之类的碰撞。 const safeName = serverName.replace(/[^a-zA-Z0-9]/g, "_").toLowerCase(); - return `mcp_${safeName}_${toolName}`; + const encodedId = Array.from(serverId, (char) => char.codePointAt(0)!.toString(16)).join("_") || "empty"; + return `mcp_${safeName}_${encodedId}_${toolName}`; } // MCPClient 工厂函数类型 @@ -24,6 +25,7 @@ export class MCPService { private clients = new Map(); // 记录每个服务器注册的工具名,便于注销 private registeredTools = new Map(); + private connecting = new Map>(); private createClient: MCPClientFactory; constructor( @@ -54,6 +56,19 @@ export class MCPService { // 连接服务器:创建 MCPClient,初始化,列出工具,注册到 ToolRegistry async connectServer(id: string): Promise { + const pending = this.connecting.get(id); + if (pending) return pending; + + const connection = this.connectServerInternal(id); + this.connecting.set(id, connection); + try { + return await connection; + } finally { + if (this.connecting.get(id) === connection) this.connecting.delete(id); + } + } + + private async connectServerInternal(id: string): Promise { const config = await this.repo.getServer(id); if (!config) { throw new Error(`MCP server "${id}" not found`); @@ -65,27 +80,32 @@ export class MCPService { } const client = this.createClient(config); - await client.initialize(); - - // 列出工具 - const tools = await client.listTools(); - this.clients.set(id, client); - - // 注册工具到 ToolRegistry - const toolNames: string[] = []; - for (const tool of tools) { - const name = mcpToolName(config.name, tool.name); - const definition: ToolDefinition = { - name, - description: `[MCP: ${config.name}] ${tool.description || tool.name}`, - parameters: tool.inputSchema, - }; - this.toolRegistry.register("mcp", definition, new MCPToolExecutor(client, tool.name)); - toolNames.push(name); - } - this.registeredTools.set(id, toolNames); + try { + await client.initialize(); - return tools; + // 列出工具 + const tools = await client.listTools(); + this.clients.set(id, client); + + // 注册工具到 ToolRegistry + const toolNames: string[] = []; + for (const tool of tools) { + const name = mcpToolName(config.id, config.name, tool.name); + const definition: ToolDefinition = { + name, + description: `[MCP: ${config.name}] ${tool.description || tool.name}`, + parameters: tool.inputSchema, + }; + this.toolRegistry.register("mcp", definition, new MCPToolExecutor(client, tool.name)); + toolNames.push(name); + } + this.registeredTools.set(id, toolNames); + + return tools; + } catch (error) { + await client.close().catch(() => {}); + throw error; + } } // 确保服务器已连接(懒连接) @@ -114,8 +134,11 @@ export class MCPService { const client = this.clients.get(id); if (client) { - client.close(); - this.clients.delete(id); + try { + await client.close(); + } finally { + this.clients.delete(id); + } } } @@ -188,8 +211,14 @@ export class MCPService { }; await this.repo.saveServer(updated); + // 已连接服务器的配置变更必须重连,否则工具仍使用旧的 URL/header/API key。 + const wasConnected = this.clients.has(request.id); + if (wasConnected) { + await this.disconnectServer(request.id); + } + // 处理 enabled 状态变更 - if (updated.enabled && !this.clients.has(request.id)) { + if (updated.enabled) { try { await this.connectServer(request.id); } catch { diff --git a/src/app/service/agent/service_worker/retry.test.ts b/src/app/service/agent/service_worker/retry.test.ts index 7bec6361a..017154a0f 100644 --- a/src/app/service/agent/service_worker/retry.test.ts +++ b/src/app/service/agent/service_worker/retry.test.ts @@ -157,8 +157,38 @@ describe("classifyErrorCode", () => { expect(classifyErrorCode(e)).toBe("tool_timeout"); }); + it("已标注的 context_too_large 应保留结构化错误码", () => { + const e = Object.assign(new Error("Conversation history is too large to compact"), { + errorCode: "context_too_large", + }); + expect(classifyErrorCode(e)).toBe("context_too_large"); + }); + it("其他错误应分类为 api_error", () => { expect(classifyErrorCode(new Error("500 Internal Server Error"))).toBe("api_error"); expect(classifyErrorCode(new Error("Unknown error"))).toBe("api_error"); }); + + // 本地字节数估算可能低估真实 token 数,放行的请求仍可能被 provider 真正拒绝; + // 没有逐 provider 精确计数的前提下,把常见的"上下文超限"错误措辞识别出来分类为 + // context_too_large,是仅有的兜底恢复路径,而不是笼统地归为不透明的 api_error。 + it("OpenAI 风格 context_length_exceeded 应分类为 context_too_large", () => { + expect( + classifyErrorCode( + new Error( + "This model's maximum context length is 128000 tokens. However, your messages resulted in 130000 tokens" + ) + ) + ).toBe("context_too_large"); + }); + + it("Anthropic 风格 prompt is too long 应分类为 context_too_large", () => { + expect(classifyErrorCode(new Error("400 prompt is too long: 250000 tokens > 200000 maximum"))).toBe( + "context_too_large" + ); + }); + + it("含 context window 措辞应分类为 context_too_large", () => { + expect(classifyErrorCode(new Error("Input exceeds the context window of this model"))).toBe("context_too_large"); + }); }); diff --git a/src/app/service/agent/service_worker/retry_utils.ts b/src/app/service/agent/service_worker/retry_utils.ts index c2ec3e87e..2ec0a1d2f 100644 --- a/src/app/service/agent/service_worker/retry_utils.ts +++ b/src/app/service/agent/service_worker/retry_utils.ts @@ -43,11 +43,28 @@ export async function withRetry( throw lastError; } +// provider 侧因上下文过长拒绝请求时的常见措辞(OpenAI/Anthropic 及兼容实现)。 +// 本地的字节数估算是保守启发式,可能低估真实 token 数:估算认为"能放下" +// 而放行的请求仍可能被 provider 真正拒绝。没有逐 provider 精确计数的前提下,把这类错误 +// 识别出来并归到与本地预判一致的 errorCode,是唯一可行的兜底恢复路径——至少能让调用方 +// (UI/自动压缩)用同一套"上下文超限"处理逻辑响应,而不是当成不透明的 api_error。 +const CONTEXT_LENGTH_ERROR_PATTERN = + /context.{0,20}(length|window|too long|exceed)|exceed.{0,20}context|maximum context length|too many tokens|prompt is too long|input is too long/i; + +const STRUCTURED_ERROR_CODES = new Set(["context_too_large", "persist_indeterminate", "tool_timeout"]); + // 将 Error 分类为 errorCode 字符串 export function classifyErrorCode(e: Error): string { + // 受信任的内部错误码比错误文案更精确;未知值不能直接透传,避免把外部错误对象的任意字段 + // 当作对外协议错误码。 + const structuredErrorCode = (e as Error & { errorCode?: unknown }).errorCode; + if (typeof structuredErrorCode === "string" && STRUCTURED_ERROR_CODES.has(structuredErrorCode)) { + return structuredErrorCode; + } const msg = e.message; + if (CONTEXT_LENGTH_ERROR_PATTERN.test(msg)) return "context_too_large"; if (/429/.test(msg)) return "rate_limit"; if (/401|403/.test(msg)) return "auth"; - if (/timed out/.test(msg) || (e as any).errorCode === "tool_timeout") return "tool_timeout"; + if (/timed out/.test(msg)) return "tool_timeout"; return "api_error"; } diff --git a/src/app/service/agent/service_worker/skill_service.ts b/src/app/service/agent/service_worker/skill_service.ts index 404287771..13c2b7357 100644 --- a/src/app/service/agent/service_worker/skill_service.ts +++ b/src/app/service/agent/service_worker/skill_service.ts @@ -18,6 +18,7 @@ import { cacheInstance } from "@App/app/cache"; import type { ToolExecutor } from "@App/app/service/agent/core/tool_registry"; import type { ResourceService } from "@App/app/service/service_worker/resource"; import { versionCompare } from "@App/pkg/utils/semver"; +import { throwIfAborted } from "@App/app/service/agent/core/abort_utils"; // 更新检查结果 export type SkillUpdateInfo = { @@ -459,7 +460,8 @@ export class SkillService { }, }, executor: { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const skillName = args.skill_name as string; const record = this.skillCache.get(skillName); if (!record) { @@ -512,7 +514,8 @@ export class SkillService { }, }, executor: { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const skillName = args.skill as string; const scriptName = args.script as string; const params = (args.params || {}) as Record; @@ -528,7 +531,7 @@ export class SkillService { ? await this.skillRepo.getConfigValues(skillName) : undefined; const executor = new SkillScriptExecutor(scriptRecord, this.sender, this.createRequireLoader(), configValues); - return executor.execute(params); + return executor.execute(params, signal); }, }, }); @@ -550,7 +553,8 @@ export class SkillService { }, }, executor: { - execute: async (args: Record) => { + execute: async (args: Record, signal?: AbortSignal) => { + throwIfAborted(signal); const skillName = args.skill_name as string; const refName = args.reference_name as string; const ref = await this.skillRepo.getReference(skillName, refName); diff --git a/src/app/service/agent/service_worker/sub_agent_service.test.ts b/src/app/service/agent/service_worker/sub_agent_service.test.ts index 04e50b994..73c476186 100644 --- a/src/app/service/agent/service_worker/sub_agent_service.test.ts +++ b/src/app/service/agent/service_worker/sub_agent_service.test.ts @@ -61,6 +61,247 @@ describe("SubAgentService", () => { expect(result.agentId).toBe("test-agent-1"); expect(result.result).toBe("result"); + expect(result.usage).toEqual({ + inputTokens: 0, + outputTokens: 0, + cacheCreationInputTokens: 0, + cacheReadInputTokens: 0, + }); expect(orchestrator.callLLMWithToolLoop).toHaveBeenCalledOnce(); }); + + it("终态错误事件也累计 usage", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ + type: "error", + message: "请求失败", + errorCode: "api_error", + usage: { inputTokens: 12, outputTokens: 4 }, + }); + }); + service = new SubAgentService(subOrchestrator); + + const result = await service.runSubAgent({ + options: { prompt: "做个任务", description: "测试任务" }, + agentId: "test-agent-2", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }); + + expect(result.details?.usage).toEqual(expect.objectContaining({ inputTokens: 12, outputTokens: 4 })); + }); + + it("终态异常保留子代理的部分消息与 usage", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ type: "content_delta", delta: "部分结果" }); + throw Object.assign(new Error("请求失败"), { + errorCode: "api_error", + usage: { inputTokens: 20, outputTokens: 6 }, + }); + }); + service = new SubAgentService(subOrchestrator); + + await expect( + service.runSubAgent({ + options: { prompt: "做个任务", description: "失败任务" }, + agentId: "test-agent-failed", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }) + ).rejects.toMatchObject({ + subAgentDetails: { + messages: [{ content: "部分结果" }], + usage: { inputTokens: 20, outputTokens: 6 }, + }, + }); + }); + + it("编排器抛出的原始异常未经 sendEvent 上报终态时,应补发一次子代理 error 事件,避免 UI 卡在 running", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async () => { + // 模拟 callLLM/autoCompact 的原始失败:orchestrator 直接 throw,从未调用 sendEvent + throw Object.assign(new Error("网络错误"), { usage: { inputTokens: 5, outputTokens: 1 } }); + }); + service = new SubAgentService(subOrchestrator); + + await expect( + service.runSubAgent({ + options: { prompt: "做个任务", description: "失败任务" }, + agentId: "test-agent-raw-fail", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }) + ).rejects.toThrow("网络错误"); + + // 必须补发一次终态事件,否则实时 UI 收不到 done/error,会一直显示 running + const errorEvents = (sendEvent as ReturnType).mock.calls + .map((c) => c[0]) + .filter((e) => e.type === "error"); + expect(errorEvents).toHaveLength(1); + expect(errorEvents[0].message).toBe("网络错误"); + }); + + it("已通过 sendEvent 上报终态错误的异常不应重复补发 error 事件", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ type: "error", message: "请求失败", errorCode: "api_error" }); + throw Object.assign(new Error("请求失败"), { errorCode: "api_error" }); + }); + service = new SubAgentService(subOrchestrator); + + await expect( + service.runSubAgent({ + options: { prompt: "做个任务", description: "失败任务" }, + agentId: "test-agent-reported-fail", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }) + ).rejects.toThrow("请求失败"); + + const errorEvents = (sendEvent as ReturnType).mock.calls + .map((c) => c[0]) + .filter((e) => e.type === "error"); + expect(errorEvents).toHaveLength(1); + }); + + it("父会话取消后子代理应以错误落定并保留部分结果与附件所有权", async () => { + const controller = new AbortController(); + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ type: "content_delta", delta: "部分结果" }); + sendEvent({ + type: "content_block_complete", + block: { type: "image", attachmentId: "partial.png", mimeType: "image/png", name: "partial.png" }, + }); + sendEvent({ + type: "error", + message: "Cancelled", + errorCode: "cancelled", + usage: { inputTokens: 30, outputTokens: 6 }, + }); + controller.abort(); + }); + service = new SubAgentService(subOrchestrator); + + await expect( + service.runSubAgent({ + options: { prompt: "做个任务", description: "取消任务" }, + agentId: "test-agent-cancelled", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal: controller.signal, + }) + ).rejects.toMatchObject({ + message: "Aborted", + usage: { inputTokens: 30, outputTokens: 6 }, + ownedAttachmentIds: ["partial.png"], + subAgentDetails: { + messages: [expect.objectContaining({ content: expect.any(Array) })], + usage: { inputTokens: 30, outputTokens: 6 }, + }, + }); + }); + + it("tool_call_complete 应使用事件自身的 status,失败的嵌套工具重载后仍应显示为 error", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ type: "tool_call_start", toolCall: { id: "t1", name: "web_fetch", arguments: "" } }); + sendEvent({ type: "tool_call_complete", id: "t1", result: "失败原因", status: "error" }); + sendEvent({ type: "done" }); + }); + service = new SubAgentService(subOrchestrator); + + const result = await service.runSubAgent({ + options: { prompt: "做个任务", description: "测试任务" }, + agentId: "test-agent-tool-error", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }); + + const toolCall = result.details?.messages[0]?.toolCalls[0]; + expect(toolCall?.status).toBe("error"); + }); + + it("嵌套工具创建的附件所有权应上交父工具轮次以便提交失败时回收", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ type: "tool_call_start", toolCall: { id: "t1", name: "image_generation", arguments: "" } }); + sendEvent({ + type: "tool_call_complete", + id: "t1", + result: "created", + status: "completed", + attachments: [{ id: "nested.png", type: "image", name: "nested.png", mimeType: "image/png" }], + ownedAttachmentIds: ["nested.png"], + }); + sendEvent({ type: "done" }); + }); + service = new SubAgentService(subOrchestrator); + + const result = await service.runSubAgent({ + options: { prompt: "做个任务", description: "嵌套附件" }, + agentId: "test-agent-owned", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }); + + expect(result.ownedAttachmentIds).toEqual(["nested.png"]); + expect(result.details?.messages[0]?.toolCalls[0].attachments?.[0].id).toBe("nested.png"); + }); + + it("子代理早期轮次生成的图片也应把 OPFS 引用、多模态详情和所有权转交父轮次", async () => { + const subOrchestrator = makeMockOrchestrator(); + (subOrchestrator.callLLMWithToolLoop as ReturnType).mockImplementationOnce(async ({ sendEvent }) => { + sendEvent({ + type: "content_block_complete", + block: { type: "image", attachmentId: "child-image.png", mimeType: "image/png", name: "child.png" }, + }); + sendEvent({ type: "new_message" }); + sendEvent({ type: "content_delta", delta: "最终说明" }); + sendEvent({ type: "done" }); + }); + service = new SubAgentService(subOrchestrator); + + const result = await service.runSubAgent({ + options: { prompt: "画图", description: "生成图片" }, + agentId: "test-agent-image", + model: MODEL, + parentConversationId: "conv-1", + toolRegistry, + sendEvent, + signal, + }); + + expect(result.result).toContain("最终说明"); + expect(result.result).toContain("uploads/child-image.png"); + expect(result.details?.messages[0]?.content).toEqual([ + { type: "image", attachmentId: "child-image.png", mimeType: "image/png", name: "child.png" }, + ]); + expect(result.attachments).toEqual([ + expect.objectContaining({ id: "child-image.png", type: "image", mimeType: "image/png" }), + ]); + expect(result.ownedAttachmentIds).toEqual(["child-image.png"]); + }); }); diff --git a/src/app/service/agent/service_worker/sub_agent_service.ts b/src/app/service/agent/service_worker/sub_agent_service.ts index 0cb7a8194..3ee7b1944 100644 --- a/src/app/service/agent/service_worker/sub_agent_service.ts +++ b/src/app/service/agent/service_worker/sub_agent_service.ts @@ -1,8 +1,11 @@ import type { AgentModelConfig, + Attachment, ChatRequest, ChatStreamEvent, + ContentBlock, SubAgentMessage, + TokenUsage, } from "@App/app/service/agent/core/types"; import type { ToolExecutorLike } from "@App/app/service/agent/core/tool_registry"; import type { SubAgentRunOptions, SubAgentRunResult } from "@App/app/service/agent/core/tools/sub_agent"; @@ -15,12 +18,12 @@ export interface SubAgentOrchestrator { toolRegistry: ToolExecutorLike; model: AgentModelConfig; messages: ChatRequest["messages"]; - maxIterations: number; sendEvent: (event: ChatStreamEvent) => void; signal: AbortSignal; scriptToolCallback: null; excludeTools?: string[]; cache?: boolean; + throwOnTerminalError?: boolean; }): Promise; } @@ -66,19 +69,35 @@ export class SubAgentService { { role: "user", content: userPrompt }, ]; - const { - result, - details, - usage: subUsage, - } = await this.runSubAgentCore({ - toolRegistry, - messages, - model, - excludeTools, - maxIterations: typeConfig.maxIterations, - sendEvent, - signal, - }); + let coreResult: Awaited>; + try { + coreResult = await this.runSubAgentCore({ + toolRegistry, + messages, + model, + excludeTools, + sendEvent, + signal, + }); + } catch (error) { + const terminalError = error instanceof Error ? error : new Error(String(error)); + const partial = terminalError as Error & { + details?: SubAgentMessage[]; + usage?: TokenUsage; + attachments?: Attachment[]; + ownedAttachmentIds?: string[]; + }; + (terminalError as Error & { subAgentDetails?: NonNullable }).subAgentDetails = { + agentId, + description: options.description, + subAgentType: typeConfig.name, + messages: partial.details || [], + usage: partial.usage, + }; + throw terminalError; + } + + const { result, details, usage: subUsage, attachments, ownedAttachmentIds } = coreResult; return { agentId, @@ -90,6 +109,9 @@ export class SubAgentService { messages: details, usage: subUsage, }, + usage: subUsage, + attachments, + ownedAttachmentIds, }; } @@ -99,7 +121,6 @@ export class SubAgentService { messages: ChatRequest["messages"]; model: AgentModelConfig; excludeTools: string[]; - maxIterations: number; sendEvent: (event: ChatStreamEvent) => void; signal: AbortSignal; }): Promise<{ @@ -111,13 +132,33 @@ export class SubAgentService { cacheCreationInputTokens: number; cacheReadInputTokens: number; }; + attachments: Attachment[]; + ownedAttachmentIds: string[]; }> { let resultContent = ""; // 收集子代理执行详情用于持久化 const details: SubAgentMessage[] = []; let currentMsg: SubAgentMessage = { content: "", toolCalls: [] }; + let currentText = ""; + let currentBlocks: ContentBlock[] = []; + const generatedAttachments: Attachment[] = []; + const ownedAttachmentIds = new Set(); + const updateCurrentContent = () => { + currentMsg.content = currentBlocks.length + ? [...(currentText ? [{ type: "text" as const, text: currentText }] : []), ...currentBlocks] + : currentText; + }; + const hasCurrentMessage = () => + currentText.length > 0 || + currentBlocks.length > 0 || + Boolean(currentMsg.thinking) || + Boolean(currentMsg.warning) || + currentMsg.toolCalls.length > 0; // 累计 usage const subUsage = { inputTokens: 0, outputTokens: 0, cacheCreationInputTokens: 0, cacheReadInputTokens: 0 }; + // 是否已通过 sendEvent 转发过终态(done/error)。orchestrator 在 callLLM/autoCompact + // 原生失败时只 throw、不 sendEvent,若不补发,实时 UI 收不到子代理的终态事件,会一直显示 running。 + let terminalEventEmitted = false; const subSendEvent = (event: ChatStreamEvent) => { // 转发事件给父代理 @@ -126,11 +167,32 @@ export class SubAgentService { switch (event.type) { case "content_delta": resultContent += event.delta; - currentMsg.content += event.delta; + currentText += event.delta; + updateCurrentContent(); + break; + case "content_block_complete": { + currentBlocks.push(event.block); + updateCurrentContent(); + generatedAttachments.push({ + id: event.block.attachmentId, + type: event.block.type, + name: event.block.name || event.block.attachmentId, + mimeType: event.block.mimeType, + size: "size" in event.block ? event.block.size : undefined, + }); + ownedAttachmentIds.add(event.block.attachmentId); + const reference = `[Generated ${event.block.type}: uploads/${event.block.attachmentId}]`; + resultContent += resultContent ? `\n\n${reference}` : reference; break; + } case "thinking_delta": currentMsg.thinking = (currentMsg.thinking || "") + event.delta; break; + case "system_warning": + // 生成数据丢失等警告(如图片保存失败)需随当前轮次一起归档,否则子代理详情持久化 + // 后刷新页面就丢失了这条提示——与父级 assistant 消息的 warning 字段同样的语义 + currentMsg.warning = currentMsg.warning ? `${currentMsg.warning}\n${event.message}` : event.message; + break; case "tool_call_start": currentMsg.toolCalls.push({ ...event.toolCall, @@ -156,21 +218,27 @@ export class SubAgentService { case "tool_call_complete": { const tc = currentMsg.toolCalls.find((t) => t.id === event.id); if (tc) { - tc.status = "completed"; + tc.status = event.status ?? "completed"; tc.result = event.result; tc.attachments = event.attachments; + tc.ownedAttachmentIds = event.ownedAttachmentIds; + for (const id of event.ownedAttachmentIds || []) ownedAttachmentIds.add(id); } break; } case "new_message": // 新一轮开始,归档当前消息 resultContent = ""; - if (currentMsg.content || currentMsg.thinking || currentMsg.toolCalls.length > 0) { + if (hasCurrentMessage()) { details.push(currentMsg); } currentMsg = { content: "", toolCalls: [] }; + currentText = ""; + currentBlocks = []; break; case "done": + case "error": + terminalEventEmitted = true; if (event.usage) { subUsage.inputTokens += event.usage.inputTokens; subUsage.outputTokens += event.usage.outputTokens; @@ -181,38 +249,83 @@ export class SubAgentService { } }; - await this.orchestrator.callLLMWithToolLoop({ - toolRegistry: params.toolRegistry, - model: params.model, - messages: params.messages, - maxIterations: params.maxIterations, - sendEvent: subSendEvent, - signal: params.signal, - scriptToolCallback: null, - excludeTools: params.excludeTools, - cache: false, - }); + try { + await this.orchestrator.callLLMWithToolLoop({ + toolRegistry: params.toolRegistry, + model: params.model, + messages: params.messages, + sendEvent: subSendEvent, + signal: params.signal, + scriptToolCallback: null, + excludeTools: params.excludeTools, + cache: false, + throwOnTerminalError: true, + }); + } catch (error) { + const terminalError = error instanceof Error ? error : new Error(String(error)); + if (hasCurrentMessage()) { + details.push(currentMsg); + } + const terminalUsage = (terminalError as Error & { usage?: TokenUsage }).usage; + if (terminalUsage) Object.assign(subUsage, terminalUsage); + // orchestrator 的 callLLM/autoCompact 原生失败只 throw、不 sendEvent 终态; + // 若不在此补发一次,实时 UI(依赖 subAgent 元信息的 done/error 事件)会一直显示 running。 + if (!terminalEventEmitted) { + const errorLike = terminalError as Error & { errorCode?: string; durationMs?: number }; + params.sendEvent({ + type: "error", + message: terminalError.message, + errorCode: errorLike.errorCode, + usage: subUsage, + durationMs: errorLike.durationMs, + }); + } + Object.assign(terminalError, { + details, + usage: subUsage, + attachments: generatedAttachments, + ownedAttachmentIds: [...ownedAttachmentIds], + }); + throw terminalError; + } // 检查是否因超时中止(区分用户主动取消和超时) - if (params.signal.aborted) { - const reason = params.signal.reason; - const isTimeout = reason instanceof DOMException && reason.name === "TimeoutError"; - if (isTimeout) { - resultContent += resultContent - ? "\n\n[Sub-agent timed out. The results above may be incomplete.]" - : "[Sub-agent timed out before producing any output.]"; - } + const abortReason = params.signal.reason; + const isTimeout = abortReason instanceof DOMException && abortReason.name === "TimeoutError"; + if (params.signal.aborted && isTimeout) { + resultContent += resultContent + ? "\n\n[Sub-agent timed out. The results above may be incomplete.]" + : "[Sub-agent timed out before producing any output.]"; } // 归档最后一轮消息 - if (currentMsg.content || currentMsg.thinking || currentMsg.toolCalls.length > 0) { + if (hasCurrentMessage()) { details.push(currentMsg); } + // 父会话 Stop 与超时策略不同:超时允许把已有结果作为部分成功返回;用户取消必须让 + // 父工具调用持久化为 error,不能在实时流已显示 cancelled 后又提交 completed。 + if (params.signal.aborted && !isTimeout) { + throw Object.assign(new Error("Aborted"), { + details, + usage: subUsage, + attachments: generatedAttachments, + ownedAttachmentIds: [...ownedAttachmentIds], + }); + } + + const attachmentReferences = generatedAttachments.map( + (attachment) => `[Generated ${attachment.type}: uploads/${attachment.id}]` + ); + const missingReferences = attachmentReferences.filter((reference) => !resultContent.includes(reference)); + const finalResult = [resultContent, ...missingReferences].filter(Boolean).join("\n\n"); + return { - result: resultContent || "(sub-agent produced no output)", + result: finalResult || "(sub-agent produced no output)", details, usage: subUsage, + attachments: generatedAttachments, + ownedAttachmentIds: [...ownedAttachmentIds], }; } diff --git a/src/app/service/agent/service_worker/task_service.test.ts b/src/app/service/agent/service_worker/task_service.test.ts new file mode 100644 index 000000000..58e62a272 --- /dev/null +++ b/src/app/service/agent/service_worker/task_service.test.ts @@ -0,0 +1,363 @@ +import { describe, expect, it, vi } from "vitest"; +import { AgentTaskService } from "./task_service"; +import { conversationChatLockKey } from "./chat_service"; +import { stackAsyncTask } from "@App/pkg/utils/async_queue"; +import type { EventAgentTask, InternalAgentTask } from "@App/app/service/agent/core/types"; + +function createService(overrides?: { appendMessage?: ReturnType }) { + const appendMessage = overrides?.appendMessage ?? vi.fn().mockResolvedValue(undefined); + const repo = { + appendMessage, + getMessages: vi.fn().mockResolvedValue([]), + listConversations: vi.fn().mockResolvedValue([{ id: "conv-lock", title: "t", modelId: "m1" }]), + createConversation: vi.fn().mockImplementation(async (conversation: any) => ({ + ...conversation, + generation: "gen-created", + revision: 1, + })), + saveConversation: vi.fn().mockResolvedValue(undefined), + } as any; + const orchestrator = { + getModel: vi.fn().mockResolvedValue({ id: "m1", model: "gpt-4o", provider: "openai" }), + callLLMWithToolLoop: vi.fn().mockResolvedValue(undefined), + }; + const skillService = { resolveSkills: vi.fn().mockReturnValue({ promptSuffix: "", metaTools: [] }) } as any; + const scriptDAO = { get: vi.fn().mockResolvedValue(undefined) } as any; + const service = new AgentTaskService( + {} as any, + repo, + {} as any, + skillService, + orchestrator as any, + {} as any, + {} as any, + scriptDAO + ); + return { service, repo, orchestrator, scriptDAO }; +} + +describe("AgentTaskService 定时任务与会话锁", () => { + it("续接已有会话的定时任务必须等待 agent-chat 会话锁释放后才写入消息", async () => { + const { service, repo } = createService(); + + // 模拟一条正在进行的 UI 对话占用同一会话的队列锁 + let releaseLock!: () => void; + const lockHeld = new Promise((resolve) => { + releaseLock = resolve; + }); + const lockTask = stackAsyncTask(conversationChatLockKey("conv-lock"), () => lockHeld); + + const task = { + id: "task-1", + name: "定时任务", + mode: "internal", + prompt: "继续", + conversationId: "conv-lock", + } as unknown as InternalAgentTask; + + const runPromise = service.executeInternalTask(task); + + // 锁被 UI 对话占用期间,定时任务不得写入任何消息(appendMessage 是读改写,会互相覆盖) + await new Promise((resolve) => setTimeout(resolve, 20)); + expect(repo.appendMessage).not.toHaveBeenCalled(); + + releaseLock(); + await lockTask; + const result = await runPromise; + + expect(result.conversationId).toBe("conv-lock"); + expect(repo.appendMessage).toHaveBeenCalledTimes(1); + }); + + it("新建会话的定时任务正常执行并返回新 conversationId", async () => { + const { service, repo, orchestrator } = createService(); + const task = { + id: "task-2", + name: "新任务", + mode: "internal", + prompt: "hello", + } as unknown as InternalAgentTask; + + const result = await service.executeInternalTask(task); + + expect(result.conversationId).toBeTruthy(); + expect(repo.createConversation).toHaveBeenCalled(); + expect(orchestrator.callLLMWithToolLoop).toHaveBeenCalledWith( + expect.objectContaining({ conversationId: result.conversationId, rehydratedHistory: false }) + ); + }); + + it("工具循环取消后正常返回时定时任务仍应报告取消", async () => { + const { service, orchestrator } = createService(); + const controller = new AbortController(); + orchestrator.callLLMWithToolLoop.mockImplementationOnce(async () => { + controller.abort(); + }); + const task = { + id: "task-cancel-after-loop", + name: "可取消任务", + mode: "internal", + prompt: "hello", + } as unknown as InternalAgentTask; + + await expect(service.executeInternalTask(task, controller.signal)).rejects.toThrow("Aborted"); + }); + + it("任务绑定的会话已被删除重建(generation 不一致)时应拒绝续接,而不是静默写入无关会话", async () => { + const { service, repo, orchestrator } = createService(); + // 存储里 conv-lock 当前的 generation 是 "gen-b"(被删除重建过) + repo.listConversations.mockResolvedValue([{ id: "conv-lock", title: "t", modelId: "m1", generation: "gen-b" }]); + + const task = { + id: "task-3", + name: "定时任务", + mode: "internal", + prompt: "继续", + conversationId: "conv-lock", + // 任务创建时绑定的是旧的一代 + conversationGeneration: "gen-a", + } as unknown as InternalAgentTask; + + await expect(service.executeInternalTask(task)).rejects.toThrow(/generation/i); + expect(repo.appendMessage).not.toHaveBeenCalled(); + expect(orchestrator.callLLMWithToolLoop).not.toHaveBeenCalled(); + }); +}); + +describe("AgentTaskService 任务生命周期", () => { + function createMutationService() { + const current = { + id: "task-cas", + generation: "generation-current", + revision: 3, + name: "current", + crontab: "0 9 * * *", + mode: "internal", + prompt: "current", + enabled: true, + notify: false, + nextruntime: Date.now() - 1_000, + createtime: 1, + updatetime: 1, + } as const; + const taskRepo = { + getTask: vi.fn().mockResolvedValue(current), + createTask: vi.fn(async (candidate: any) => candidate), + saveTask: vi.fn(async (candidate: any) => { + if (candidate.generation !== current.generation || candidate.revision !== current.revision) { + throw new Error("revision conflict"); + } + return candidate; + }), + removeTask: vi.fn().mockResolvedValue(undefined), + }; + const scheduler = { + cancelTask: vi.fn(), + executeTask: vi.fn().mockResolvedValue(undefined), + }; + const scriptDAO = { get: vi.fn().mockResolvedValue({ uuid: "installed-script" }) } as any; + const service = new AgentTaskService( + {} as any, + {} as any, + {} as any, + {} as any, + {} as any, + taskRepo as any, + {} as any, + scriptDAO + ); + service.setScheduler(scheduler as any); + return { service, taskRepo, scheduler, current, scriptDAO }; + } + + it("update 与 enable 必须使用客户端看到的 generation/revision 做 CAS", async () => { + const { service, taskRepo } = createMutationService(); + + await expect( + service.handleAgentTask({ + action: "update", + id: "task-cas", + generation: "generation-current", + revision: 2, + task: { name: "stale edit" }, + } as any) + ).rejects.toThrow("revision conflict"); + await expect( + service.handleAgentTask({ + action: "enable", + id: "task-cas", + generation: "generation-current", + revision: 2, + enabled: false, + } as any) + ).rejects.toThrow("revision conflict"); + + expect(taskRepo.saveTask).toHaveBeenCalledWith(expect.objectContaining({ revision: 2 })); + }); + + it("delete 应先取消活动执行并使用客户端版本删除", async () => { + const { service, taskRepo, scheduler } = createMutationService(); + + await service.handleAgentTask({ + action: "delete", + id: "task-cas", + generation: "generation-current", + revision: 3, + } as any); + + expect(scheduler.cancelTask).toHaveBeenCalledWith("task-cas"); + expect(taskRepo.removeTask).toHaveBeenCalledWith("task-cas", "generation-current", 3); + }); + + it("delete 应先 cancelTask 中止执行,再清理元数据/运行记录,即使清理失败", async () => { + const { service, taskRepo, scheduler } = createMutationService(); + const callOrder: string[] = []; + scheduler.cancelTask.mockImplementation(() => { + callOrder.push("cancelTask"); + return true; + }); + taskRepo.removeTask.mockImplementation(async () => { + callOrder.push("removeTask"); + throw new Error("run-history cleanup failed"); + }); + + await expect( + service.handleAgentTask({ + action: "delete", + id: "task-cas", + generation: "generation-current", + revision: 3, + } as any) + ).rejects.toThrow("run-history cleanup failed"); + + // cancelTask 必须先发生:即使 removeTask(含 run-history 清理)失败,正在运行的执行也已经被中止 + expect(callOrder).toEqual(["cancelTask", "removeTask"]); + }); + + it("runNow 遇到已到期任务时应领取当前槽位而不是随后再由 tick 重复执行", async () => { + const { service, scheduler, current } = createMutationService(); + + await service.handleAgentTask({ action: "runNow", id: current.id }); + + expect(scheduler.executeTask).toHaveBeenCalledWith(current, true, expect.any(Number)); + }); + + it("create 拒绝 sourceScriptUuid 为空的事件任务,防止创建无接收脚本的死信任务", async () => { + const { service, taskRepo, scriptDAO } = createMutationService(); + scriptDAO.get.mockResolvedValue(undefined); + + await expect( + service.handleAgentTask({ + action: "create", + task: { + name: "事件任务", + mode: "event", + crontab: "0 9 * * *", + sourceScriptUuid: "", + enabled: true, + notify: false, + }, + } as any) + ).rejects.toThrow(/sourceScriptUuid/); + expect(taskRepo.createTask).not.toHaveBeenCalled(); + }); + + it("create 拒绝 sourceScriptUuid 指向未安装脚本的事件任务", async () => { + const { service, taskRepo, scriptDAO } = createMutationService(); + scriptDAO.get.mockResolvedValue(undefined); + + await expect( + service.handleAgentTask({ + action: "create", + task: { + name: "事件任务", + mode: "event", + crontab: "0 9 * * *", + sourceScriptUuid: "unknown-script", + enabled: true, + notify: false, + }, + } as any) + ).rejects.toThrow(/sourceScriptUuid/); + expect(taskRepo.createTask).not.toHaveBeenCalled(); + }); + + it("create 接受 sourceScriptUuid 指向已安装脚本的事件任务", async () => { + const { service, taskRepo, scriptDAO } = createMutationService(); + scriptDAO.get.mockResolvedValue({ uuid: "installed-script" }); + + await service.handleAgentTask({ + action: "create", + task: { + name: "事件任务", + mode: "event", + crontab: "0 9 * * *", + sourceScriptUuid: "installed-script", + enabled: true, + notify: false, + }, + } as any); + + expect(scriptDAO.get).toHaveBeenCalledWith("installed-script"); + expect(taskRepo.createTask).toHaveBeenCalled(); + }); + + it("update 拒绝把事件任务的 sourceScriptUuid 改成未安装的脚本", async () => { + const { service, taskRepo, scriptDAO } = createMutationService(); + taskRepo.getTask.mockResolvedValue({ + id: "task-cas", + generation: "generation-current", + revision: 3, + name: "事件任务", + crontab: "0 9 * * *", + mode: "event", + sourceScriptUuid: "installed-script", + enabled: true, + notify: false, + createtime: 1, + updatetime: 1, + }); + scriptDAO.get.mockResolvedValue(undefined); + + await expect( + service.handleAgentTask({ + action: "update", + id: "task-cas", + generation: "generation-current", + revision: 3, + task: { sourceScriptUuid: "" }, + } as any) + ).rejects.toThrow(/sourceScriptUuid/); + expect(taskRepo.saveTask).not.toHaveBeenCalled(); + }); + + it("事件派发通道无响应时取消信号应立即终止等待", async () => { + const sender = { sendMessage: vi.fn(() => new Promise(() => {})) } as any; + const service = new AgentTaskService( + sender, + {} as any, + {} as any, + {} as any, + {} as any, + {} as any, + {} as any, + { get: vi.fn().mockResolvedValue({ uuid: "script-1" }) } as any + ); + const task = { + id: "event-cancel", + name: "事件任务", + mode: "event", + crontab: "0 9 * * *", + sourceScriptUuid: "script-1", + enabled: true, + notify: false, + } as EventAgentTask; + const controller = new AbortController(); + + const execution = service.emitTaskEvent(task, controller.signal); + await vi.waitFor(() => expect(sender.sendMessage).toHaveBeenCalledOnce()); + controller.abort(); + + await expect(execution).rejects.toThrow("Aborted"); + }); +}); diff --git a/src/app/service/agent/service_worker/task_service.ts b/src/app/service/agent/service_worker/task_service.ts index f5101b52f..425787c20 100644 --- a/src/app/service/agent/service_worker/task_service.ts +++ b/src/app/service/agent/service_worker/task_service.ts @@ -14,6 +14,7 @@ import type { ScriptToolCallback, ToolExecutorLike, ToolRegistry } from "@App/ap import { SessionToolRegistry } from "@App/app/service/agent/core/session_tool_registry"; import type { SkillService } from "./skill_service"; import type { AgentTaskRepo, AgentTaskRunRepo } from "@App/app/repo/agent_task"; +import type { ScriptDAO } from "@App/app/repo/scripts"; import type { AgentTaskScheduler } from "@App/app/service/agent/core/task_scheduler"; import type { AgentChatRepo } from "@App/app/repo/agent_chat"; import { buildSystemPrompt } from "@App/app/service/agent/core/system_prompt"; @@ -21,6 +22,10 @@ import { uuidv4 } from "@App/pkg/utils/uuid"; import { nextTimeInfo } from "@App/pkg/utils/cron"; import { InfoNotification } from "@App/app/service/service_worker/utils"; import { sendMessage } from "@Packages/message/client"; +import { toLLMMessages } from "@App/app/service/agent/core/persisted_messages"; +import { stackAsyncTask } from "@App/pkg/utils/async_queue"; +import { conversationChatLockKey } from "./chat_service"; +import { raceWithAbort, throwIfAborted } from "@App/app/service/agent/core/abort_utils"; /** 供 TaskService 调用的 orchestrator 能力 */ export interface TaskOrchestrator { @@ -29,11 +34,13 @@ export interface TaskOrchestrator { toolRegistry: ToolExecutorLike; model: AgentModelConfig; messages: ChatRequest["messages"]; - maxIterations: number; sendEvent: (event: ChatStreamEvent) => void; signal: AbortSignal; scriptToolCallback: ScriptToolCallback | null; conversationId: string; + conversationGeneration: string; + rehydratedHistory?: boolean; + throwOnTerminalError?: boolean; }): Promise; } @@ -48,9 +55,22 @@ export class AgentTaskService { private skillService: SkillService, private orchestrator: TaskOrchestrator, private taskRepo: AgentTaskRepo, - private taskRunRepo: AgentTaskRunRepo + private taskRunRepo: AgentTaskRunRepo, + private scriptDAO: ScriptDAO ) {} + // event 模式任务的 sourceScriptUuid 必须指向当前已安装的脚本,否则事件派发时无脚本可接收, + // 任务会显示为"已调度/成功"但实际是死信——校验必须在写入前完成,不能事后依赖执行期报错 + private async assertValidEventSource(sourceScriptUuid: string | undefined): Promise { + if (!sourceScriptUuid) { + throw new Error("Event task requires a non-empty sourceScriptUuid identifying an installed script"); + } + const script = await this.scriptDAO.get(sourceScriptUuid); + if (!script) { + throw new Error(`Event task sourceScriptUuid does not match an installed script: ${sourceScriptUuid}`); + } + } + // 延迟注入 scheduler(避免循环依赖:AgentTaskScheduler ↔ AgentTaskService) setScheduler(scheduler: AgentTaskScheduler) { this.taskScheduler = scheduler; @@ -58,8 +78,10 @@ export class AgentTaskService { // internal 模式定时任务执行:构建对话并调用 LLM async executeInternalTask( - task: InternalAgentTask + task: InternalAgentTask, + signal = new AbortController().signal ): Promise<{ conversationId: string; usage?: { inputTokens: number; outputTokens: number } }> { + throwIfAborted(signal); const model = await this.orchestrator.getModel(task.modelId); // 解析 Skills @@ -71,14 +93,38 @@ export class AgentTaskService { sessionRegistry.register("skill", mt.definition, mt.executor); } + // 与同一会话的 UI 聊天 / compact / clearMessages 共用同一把按 conversationId 的队列锁: + // appendMessage 是读-改-写,续接已有会话的定时任务若不排队,会与进行中的对话互相覆盖丢消息 + const conversationId = task.conversationId || uuidv4(); + return stackAsyncTask(conversationChatLockKey(conversationId), () => + this.executeInternalTaskLocked(task, conversationId, model, promptSuffix, metaTools, sessionRegistry, signal) + ); + } + + private async executeInternalTaskLocked( + task: InternalAgentTask, + conversationId: string, + model: AgentModelConfig, + promptSuffix: string, + metaTools: ReturnType["metaTools"], + sessionRegistry: SessionToolRegistry, + signal: AbortSignal + ): Promise<{ conversationId: string; usage?: { inputTokens: number; outputTokens: number } }> { try { - let conversationId: string; + throwIfAborted(signal); const messages: ChatRequest["messages"] = []; + let conversation: Conversation; if (task.conversationId) { // 续接已有对话 - conversationId = task.conversationId; const conv = await this.getConversation(conversationId); + if (!conv?.generation) throw new Error("Conversation not found"); + // task.conversationGeneration 记录任务绑定该对话时的 generation;若当前 generation + // 不一致,说明该会话已被删除重建为无关的新一代,绝不能静默续接 + if (task.conversationGeneration && conv.generation !== task.conversationGeneration) { + throw new Error("Conversation generation mismatch; the bound conversation was deleted and recreated"); + } + conversation = conv; const systemContent = buildSystemPrompt({ userSystem: conv?.system, @@ -87,7 +133,7 @@ export class AgentTaskService { messages.push({ role: "system", content: systemContent }); // 加载历史消息 - if (conv) { + { const existingMessages = await this.repo.getMessages(conversationId); // 预加载之前已加载的 skill 的工具 @@ -113,19 +159,10 @@ export class AgentTaskService { } } - for (const msg of existingMessages) { - if (msg.role === "system") continue; - messages.push({ - role: msg.role, - content: msg.content, - toolCallId: msg.toolCallId, - toolCalls: msg.toolCalls, - }); - } + messages.push(...toLLMMessages(existingMessages).filter((message) => message.role !== "system")); } } else { // 创建新对话 - conversationId = uuidv4(); const conv: Conversation = { id: conversationId, title: task.name, @@ -134,7 +171,7 @@ export class AgentTaskService { createtime: Date.now(), updatetime: Date.now(), }; - await this.repo.saveConversation(conv); + conversation = await this.repo.createConversation(conv); const systemContent = buildSystemPrompt({ skillSuffix: promptSuffix }); messages.push({ role: "system", content: systemContent }); @@ -143,21 +180,22 @@ export class AgentTaskService { // 添加用户消息(task.prompt) const userContent = task.prompt || task.name; messages.push({ role: "user", content: userContent }); - await this.repo.appendMessage({ - id: uuidv4(), - conversationId, - role: "user", - content: userContent, - createtime: Date.now(), - }); + await this.repo.appendMessage( + { + id: uuidv4(), + conversationId, + role: "user", + content: userContent, + createtime: Date.now(), + }, + conversation.generation + ); // 收集 usage const totalUsage = { inputTokens: 0, outputTokens: 0 }; - const abortController = new AbortController(); - const sendEvent = (event: ChatStreamEvent) => { // 定时任务无 UI 连接,但需要收集 usage - if (event.type === "done" && event.usage) { + if ((event.type === "done" || event.type === "error") && event.usage) { totalUsage.inputTokens += event.usage.inputTokens; totalUsage.outputTokens += event.usage.outputTokens; } @@ -167,12 +205,16 @@ export class AgentTaskService { toolRegistry: sessionRegistry, model, messages, - maxIterations: task.maxIterations || 10, sendEvent, - signal: abortController.signal, + signal, scriptToolCallback: null, conversationId, + conversationGeneration: conversation.generation!, + rehydratedHistory: Boolean(task.conversationId), + throwOnTerminalError: true, }); + // 工具循环在取消时可能正常 return;定时任务必须把这个状态转成失败/取消,避免误记成功。 + throwIfAborted(signal); // 通知 if (task.notify) { @@ -186,7 +228,8 @@ export class AgentTaskService { } // event 模式定时任务:通知脚本 - async emitTaskEvent(task: EventAgentTask): Promise { + async emitTaskEvent(task: EventAgentTask, signal = new AbortController().signal): Promise { + throwIfAborted(signal); const trigger: AgentTaskTrigger = { taskId: task.id, name: task.name, @@ -195,12 +238,16 @@ export class AgentTaskService { }; // 通过 offscreen → sandbox → 脚本 EventEmitter 链路通知脚本 - await sendMessage(this.sender, "offscreen/runtime/emitEvent", { - uuid: task.sourceScriptUuid, - event: "agentTask", - eventId: task.id, - data: trigger, - }); + await raceWithAbort( + sendMessage(this.sender, "offscreen/runtime/emitEvent", { + uuid: task.sourceScriptUuid, + event: "agentTask", + eventId: task.id, + data: trigger, + }), + signal + ); + throwIfAborted(signal); if (task.notify) { InfoNotification(task.name, "定时任务已触发"); @@ -209,7 +256,13 @@ export class AgentTaskService { private async getConversation(id: string): Promise { const conversations = await this.repo.listConversations(); - return conversations.find((c) => c.id === id) || null; + const conversation = conversations.find((item) => item.id === id); + if (!conversation) return null; + return { + ...conversation, + generation: conversation.generation || `legacy:${conversation.id}`, + revision: conversation.revision ?? 0, + }; } // 处理定时任务 CRUD 及 run 操作 @@ -227,6 +280,9 @@ export class AgentTaskService { createtime: now, updatetime: now, } as AgentTask; + if (task.mode === "event") { + await this.assertValidEventSource(task.sourceScriptUuid); + } // 计算 nextruntime if (task.enabled) { try { @@ -236,13 +292,32 @@ export class AgentTaskService { // cron 无效,不设置 nextruntime } } - await this.taskRepo.saveTask(task); - return task; + // 绑定续接对话时记录当时的 generation,执行期据此拒绝已被删除重建的会话 + if (task.mode === "internal" && task.conversationId && !task.conversationGeneration) { + const conv = await this.getConversation(task.conversationId); + if (conv?.generation) task.conversationGeneration = conv.generation; + } + return this.taskRepo.createTask(task); } case "update": { const existing = await this.taskRepo.getTask(params.id); if (!existing) throw new Error("Task not found"); - const updated = { ...existing, ...params.task, updatetime: Date.now() } as AgentTask; + const updated = { + ...existing, + ...params.task, + id: params.id, + generation: params.generation, + revision: params.revision, + updatetime: Date.now(), + } as AgentTask; + // conversationId 变更(或首次绑定)时重新记录 generation,避免沿用旧会话的 generation + if (updated.mode === "internal" && updated.conversationId && "conversationId" in params.task) { + const conv = await this.getConversation(updated.conversationId); + updated.conversationGeneration = conv?.generation; + } + if (updated.mode === "event") { + await this.assertValidEventSource(updated.sourceScriptUuid); + } // 如果 crontab 或 enabled 变化,重新计算 nextruntime if (params.task.crontab !== undefined || params.task.enabled !== undefined) { if (updated.enabled) { @@ -254,33 +329,44 @@ export class AgentTaskService { } } } - await this.taskRepo.saveTask(updated); - return updated; + return this.taskRepo.saveTask(updated); } - case "delete": - await this.taskRepo.removeTask(params.id); + case "delete": { + // 先中止正在运行的执行,再清理元数据/运行记录:cancelTask 是同步的 abort(),必须最先 + // 发生,否则被删除的任务会在 removeTask(含 run-history 清理)完成前继续调用 LLM/工具/ + // 产生外部副作用;若 removeTask 之后才 cancel,一旦 removeTask 因清理失败而抛出, + // cancelTask 根本不会被调用,执行也就永远不会被中止 + this.taskScheduler?.cancelTask(params.id); + await this.taskRepo.removeTask(params.id, params.generation, params.revision); return true; + } case "enable": { const task = await this.taskRepo.getTask(params.id); if (!task) throw new Error("Task not found"); - task.enabled = params.enabled; - task.updatetime = Date.now(); - if (task.enabled) { + const updated = { + ...task, + enabled: params.enabled, + generation: params.generation, + revision: params.revision, + updatetime: Date.now(), + } as AgentTask; + if (updated.enabled) { try { - const info = nextTimeInfo(task.crontab); - task.nextruntime = info.next.toMillis(); + const info = nextTimeInfo(updated.crontab); + updated.nextruntime = info.next.toMillis(); } catch { - task.nextruntime = undefined; + updated.nextruntime = undefined; } } - await this.taskRepo.saveTask(task); - return task; + return this.taskRepo.saveTask(updated); } case "runNow": { const task = await this.taskRepo.getTask(params.id); if (!task) throw new Error("Task not found"); // 不 await,立即返回 - this.taskScheduler?.executeTask(task).catch(() => {}); + const now = Date.now(); + const claimScheduled = Boolean(task.enabled && task.nextruntime && task.nextruntime <= now); + this.taskScheduler?.executeTask(task, claimScheduled, now).catch(() => {}); return true; } case "listRuns": diff --git a/src/app/service/agent/service_worker/test-helpers.ts b/src/app/service/agent/service_worker/test-helpers.ts index 17b86146a..3944c5e4a 100644 --- a/src/app/service/agent/service_worker/test-helpers.ts +++ b/src/app/service/agent/service_worker/test-helpers.ts @@ -39,17 +39,60 @@ export function createTestService() { const mockGroup = { on: vi.fn() } as any; const mockSender = {} as any; + // 消息按 conversationId 存储,供下面的 appendMessage/updateMessage/getMessages mock 共享, + // 让 updateMessage 的按 id 替换语义在测试里也能被后续 getMessages 断言观察到(与真实 repo 行为一致)。 + // 真实 OPFS repo 每次读写都经过 JSON 序列化往返,调用方之间不共享对象引用; + // 这里的 mock 也必须在存取时做深拷贝,否则某次调用记录下的数组引用会被后续调用原地修改 + // (例如 saveMessages 记录的数组之后被 appendMessage push,导致 vi.fn() 的 mock.calls 断言看到"事后被改过"的内容)。 + const messagesByConv = new Map(); + // 重置 agent_chat 单例 mock 方法(保持对象身份不变,只替换 vi.fn) Object.assign(mockChatRepo, { - appendMessage: vi.fn().mockResolvedValue(undefined), - getMessages: vi.fn().mockResolvedValue([]), + appendMessage: vi.fn().mockImplementation(async (message: any) => { + const list = messagesByConv.get(message.conversationId) ?? []; + list.push(structuredClone(message)); + messagesByConv.set(message.conversationId, list); + }), + getMessages: vi + .fn() + .mockImplementation(async (conversationId: string) => + (messagesByConv.get(conversationId) ?? []).map((m) => structuredClone(m)) + ), listConversations: vi.fn().mockResolvedValue([]), + createConversation: vi.fn().mockImplementation(async (conversation: any) => ({ + ...conversation, + generation: conversation.generation || "test-generation", + revision: conversation.revision ?? 1, + })), saveConversation: vi.fn().mockResolvedValue(undefined), - saveMessages: vi.fn().mockResolvedValue(undefined), + deleteConversation: vi.fn().mockResolvedValue(undefined), + getMessageSnapshot: vi.fn().mockImplementation(async (conversationId: string) => ({ + generation: "test-generation", + revision: 0, + messages: await mockChatRepo.getMessages(conversationId), + })), + saveMessages: vi.fn().mockImplementation(async (conversationId: string, messages: any[]) => { + messagesByConv.set( + conversationId, + messages.map((m) => structuredClone(m)) + ); + }), + updateMessage: vi.fn().mockImplementation(async (message: any) => { + const list = messagesByConv.get(message.conversationId) ?? []; + const index = list.findIndex((m) => m.id === message.id); + if (index >= 0) list[index] = structuredClone(message); + }), + commitToolRound: vi.fn().mockImplementation(async (assistant: any, toolMessages: any[]) => { + const list = messagesByConv.get(assistant.conversationId) ?? []; + list.push(structuredClone(assistant), ...toolMessages.map((message) => structuredClone(message))); + messagesByConv.set(assistant.conversationId, list); + }), getTasks: vi.fn().mockResolvedValue([]), + getTaskSnapshot: vi.fn().mockResolvedValue({ generation: "test-generation", revision: 0, tasks: [] }), saveTasks: vi.fn().mockResolvedValue(undefined), getAttachment: vi.fn().mockResolvedValue(null), saveAttachment: vi.fn().mockResolvedValue(0), + deleteAttachment: vi.fn().mockResolvedValue(undefined), }); const service = new AgentService(mockGroup, mockSender); @@ -178,7 +221,7 @@ export function createMockSenderWithCallbacks() { export function makeTextResponse(text: string): Response { const encoder = new TextEncoder(); const chunks = [ - `data: {"choices":[{"delta":{"content":"${text}"}}]}\n\n`, + `data: {"choices":[{"delta":{"content":"${text}"},"finish_reason":"stop"}]}\n\n`, `data: {"usage":{"prompt_tokens":10,"completion_tokens":5}}\n\n`, ]; let i = 0; diff --git a/src/app/service/agent/service_worker/tool_loop_orchestrator.test.ts b/src/app/service/agent/service_worker/tool_loop_orchestrator.test.ts new file mode 100644 index 000000000..782f53c8f --- /dev/null +++ b/src/app/service/agent/service_worker/tool_loop_orchestrator.test.ts @@ -0,0 +1,677 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { ToolLoopOrchestrator, type ToolLoopDeps } from "./tool_loop_orchestrator"; +import type { ToolExecutorLike, ToolExecuteResult } from "@App/app/service/agent/core/tool_registry"; +import type { ToolCall, AgentModelConfig, ChatRequest, ChatStreamEvent } from "@App/app/service/agent/core/types"; +import type { LLMCallResult } from "./llm_client"; + +const MODEL: AgentModelConfig = { + id: "m1", + name: "Test", + provider: "openai", + apiBaseUrl: "", + apiKey: "", + model: "gpt-4o", +}; + +function makeFakeChatRepo() { + return { + appendMessage: vi.fn().mockResolvedValue(undefined), + getMessages: vi.fn().mockResolvedValue([]), + saveMessages: vi.fn().mockResolvedValue(undefined), + updateMessage: vi.fn().mockResolvedValue(undefined), + commitToolRound: vi.fn().mockResolvedValue(undefined), + getMessageSnapshot: vi.fn().mockResolvedValue({ generation: "gen-1", revision: 0, messages: [] }), + getAttachment: vi.fn().mockResolvedValue(null), + deleteAttachment: vi.fn().mockResolvedValue(undefined), + } as any; +} + +function makeFakeToolRegistry(): ToolExecutorLike { + return { + getDefinitions: () => [{ name: "dup", description: "dup", parameters: { type: "object", properties: {} } }], + execute: async (toolCalls: ToolCall[]): Promise => + toolCalls.map((tc) => ({ id: tc.id, result: "ok" })), + }; +} + +// 连续 4 轮调用同一工具 dup 且参数完全相同(触发两次重复调用告警:第 2、4 轮) +function dupToolCallResult(id: string): LLMCallResult { + return { content: "", toolCalls: [{ id, name: "dup", arguments: "{}" }] } as LLMCallResult; +} + +function finalTextResult(text: string): LLMCallResult { + return { content: text } as LLMCallResult; +} + +describe("ToolLoopOrchestrator 循环检测升级(loop-guard escalation)", () => { + let chatRepo: ReturnType; + let toolRegistry: ToolExecutorLike; + let callLLM: ReturnType>; + let autoCompact: ReturnType>; + let deps: ToolLoopDeps; + let orchestrator: ToolLoopOrchestrator; + let sendEvent: ReturnType; + + beforeEach(() => { + chatRepo = makeFakeChatRepo(); + toolRegistry = makeFakeToolRegistry(); + callLLM = vi.fn(); + autoCompact = vi.fn().mockResolvedValue(undefined); + deps = { callLLM, autoCompact }; + orchestrator = new ToolLoopOrchestrator(deps, chatRepo); + sendEvent = vi.fn(); + }); + + function baseParams(overrides: Record = {}) { + return { + toolRegistry, + model: MODEL, + messages: [{ role: "user", content: "开始" }] as ChatRequest["messages"], + sendEvent: sendEvent as (event: ChatStreamEvent) => void, + signal: new AbortController().signal, + scriptToolCallback: null, + conversationId: "conv-1", + conversationGeneration: "gen-1", + rehydratedHistory: true, + ...overrides, + }; + } + + it("未提供 askUserForGuard 时,重复调用告警照常触发但不暂停循环", async () => { + callLLM + .mockResolvedValueOnce(dupToolCallResult("c1")) + .mockResolvedValueOnce(dupToolCallResult("c2")) + .mockResolvedValueOnce(dupToolCallResult("c3")) + .mockResolvedValueOnce(dupToolCallResult("c4")) + .mockResolvedValueOnce(finalTextResult("done")); + + await orchestrator.callLLMWithToolLoop(baseParams()); + + expect(callLLM).toHaveBeenCalledTimes(5); + const warningEvents = sendEvent.mock.calls.filter((c) => c[0].type === "system_warning"); + expect(warningEvents).toHaveLength(2); + const doneEvents = sendEvent.mock.calls.filter((c) => c[0].type === "done"); + expect(doneEvents).toHaveLength(1); + }); + + it("成功工具的内部 usage 应累计到父会话终态", async () => { + toolRegistry = { + getDefinitions: makeFakeToolRegistry().getDefinitions, + execute: async (toolCalls) => + toolCalls.map((toolCall) => ({ + id: toolCall.id, + result: "child done", + usage: { inputTokens: 100, outputTokens: 20 }, + })), + }; + callLLM + .mockResolvedValueOnce({ + content: "", + toolCalls: [{ id: "child-1", name: "dup", arguments: "{}" }], + usage: { inputTokens: 10, outputTokens: 2 }, + }) + .mockResolvedValueOnce(finalTextResult("done")); + + await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })); + + const doneEvent = sendEvent.mock.calls.map((call) => call[0]).find((event) => event.type === "done"); + expect(doneEvent?.usage).toEqual(expect.objectContaining({ inputTokens: 110, outputTokens: 22 })); + }); + + it("auto compact 失败时抛出包含累计 usage 与 conversationId 的结构化错误", async () => { + callLLM.mockResolvedValue({ content: "继续", usage: { inputTokens: 13000, outputTokens: 5 } }); + autoCompact.mockRejectedValue( + Object.assign(new Error("compact failed"), { usage: { inputTokens: 120, outputTokens: 15 } }) + ); + + await expect( + orchestrator.callLLMWithToolLoop( + // 13000/16000 > 80%,按完整上下文窗口触发自动压缩。 + baseParams({ model: { ...MODEL, contextWindow: 16000 }, messages: [{ role: "user", content: "开始" }] }) + ) + ).rejects.toMatchObject({ + message: "compact failed", + conversationId: "conv-1", + usage: { inputTokens: 13120, outputTokens: 20 }, + }); + }); + + it("provider 非取消失败时,错误上挂载的本轮部分 usage 应并入累计 usage 后再抛出", async () => { + // 第 1 轮正常返回 tool call(累计 usage 10/5),第 2 轮 provider 失败, + // 失败错误上带有解析层已知的本轮部分 usage(如逐 chunk 附带 usage 的 OpenAI 兼容 API) + callLLM + .mockResolvedValueOnce({ + content: "", + toolCalls: [{ id: "c1", name: "dup", arguments: "{}" }], + usage: { inputTokens: 10, outputTokens: 5 }, + } as LLMCallResult) + .mockRejectedValueOnce( + Object.assign(new Error("stream truncated"), { usage: { inputTokens: 7, outputTokens: 3 } }) + ); + + await expect(orchestrator.callLLMWithToolLoop(baseParams())).rejects.toMatchObject({ + message: "stream truncated", + conversationId: "conv-1", + // 之前的实现用上一轮累计覆盖错误自带的 usage,本轮已知的 7/3 会丢失 + usage: { inputTokens: 17, outputTokens: 8 }, + }); + }); + + it("LLM 返回后立即取消时,已保存但未被任何消息引用的生成附件应被删除", async () => { + const controller = new AbortController(); + callLLM.mockImplementation(async () => { + // 取消恰好落在 provider 成功返回之后、消息持久化之前 + controller.abort(); + return { + content: "生成了一张图", + contentBlocks: [{ type: "image", attachmentId: "img_orphan.png", mimeType: "image/png" }], + usage: { inputTokens: 10, outputTokens: 5 }, + } as LLMCallResult; + }); + + await orchestrator.callLLMWithToolLoop(baseParams({ signal: controller.signal })); + + // 终态是取消,assistant 消息不会持久化,生成的附件必须回收,否则成为孤儿文件 + expect(chatRepo.deleteAttachment).toHaveBeenCalledWith("img_orphan.png"); + const terminal = sendEvent.mock.calls.map((c) => c[0]).find((e) => e.type === "error"); + expect(terminal?.errorCode).toBe("cancelled"); + }); + + it("模型生成附件应把所有权持久化到 assistant 消息", async () => { + callLLM.mockResolvedValue({ + content: "生成完成", + contentBlocks: [{ type: "image", attachmentId: "generated-owned.png", mimeType: "image/png" }], + }); + + await orchestrator.callLLMWithToolLoop(baseParams()); + + expect(chatRepo.appendMessage).toHaveBeenCalledWith( + expect.objectContaining({ ownedAttachmentIds: ["generated-owned.png"] }), + "gen-1" + ); + }); + + it("LLM 结果携带的图片保存 warning 应持久化到 assistant 消息并广播 system_warning", async () => { + callLLM.mockResolvedValue({ + content: "", + warning: "1 generated image(s) failed to save and were lost.", + }); + + await orchestrator.callLLMWithToolLoop(baseParams()); + + expect(chatRepo.appendMessage).toHaveBeenCalledWith( + expect.objectContaining({ warning: "1 generated image(s) failed to save and were lost." }), + "gen-1" + ); + expect( + sendEvent.mock.calls.some( + (call) => + call[0].type === "system_warning" && call[0].message === "1 generated image(s) failed to save and were lost." + ) + ).toBe(true); + }); + + it("带工具调用的 assistant 消息携带 warning 时也应持久化并广播 system_warning,而不只是最终回复分支", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "image_tool", description: "image", parameters: { type: "object", properties: {} } }, + ], + execute: vi.fn().mockResolvedValue([{ id: "call-1", result: "ok" }]), + }; + callLLM + .mockResolvedValueOnce({ + content: "", + warning: "1 generated image(s) failed to save and were lost.", + toolCalls: [{ id: "call-1", name: "image_tool", arguments: "{}" }], + }) + .mockResolvedValueOnce(finalTextResult("done")); + + await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })); + + expect(chatRepo.commitToolRound).toHaveBeenCalledWith( + expect.objectContaining({ warning: "1 generated image(s) failed to save and were lost." }), + expect.anything(), + "gen-1" + ); + expect( + sendEvent.mock.calls.some( + (call) => + call[0].type === "system_warning" && call[0].message === "1 generated image(s) failed to save and were lost." + ) + ).toBe(true); + }); + + it( + "最终回复持久化失败(persist_failed)时应回收生成附件", + // 持久化重试退避 200ms + 400ms,放宽超时 + { timeout: 3000 }, + async () => { + chatRepo.appendMessage.mockRejectedValue(new Error("disk full")); + callLLM.mockResolvedValue({ + content: "回复", + contentBlocks: [{ type: "image", attachmentId: "img_lost.png", mimeType: "image/png" }], + usage: { inputTokens: 3, outputTokens: 2 }, + } as LLMCallResult); + + await orchestrator.callLLMWithToolLoop(baseParams()); + + const terminal = sendEvent.mock.calls.map((c) => c[0]).find((e) => e.type === "error"); + expect(terminal?.errorCode).toBe("persist_failed"); + expect(chatRepo.deleteAttachment).toHaveBeenCalledWith("img_lost.png"); + } + ); + + it("最终回复持久化报错且确认读也失败时,不应删除可能已被消息引用的生成附件", { timeout: 3000 }, async () => { + chatRepo.appendMessage.mockRejectedValue(new Error("ambiguous close failure")); + chatRepo.getMessageSnapshot.mockRejectedValue(new Error("confirmation read failed")); + callLLM.mockResolvedValue({ + content: "回复", + contentBlocks: [{ type: "image", attachmentId: "img_maybe_committed.png", mimeType: "image/png" }], + usage: { inputTokens: 3, outputTokens: 2 }, + } as LLMCallResult); + + await orchestrator.callLLMWithToolLoop(baseParams()); + + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("img_maybe_committed.png"); + }); + + it("工具结果持久化期间被取消时应立即终态化,而不是带着已取消的信号进入下一轮", async () => { + const controller = new AbortController(); + callLLM.mockResolvedValueOnce(dupToolCallResult("c1")); + // 工具执行正常结束后,tool 结果消息落库完成的同时 Stop 到达(晚于 cancelledDuringTools 采样点) + chatRepo.commitToolRound.mockImplementation(async () => { + controller.abort(); + }); + + await orchestrator.callLLMWithToolLoop(baseParams({ signal: controller.signal })); + + expect(callLLM).toHaveBeenCalledTimes(1); + const events = sendEvent.mock.calls.map((c) => c[0]); + expect(events.find((e) => e.type === "error")?.errorCode).toBe("cancelled"); + expect(events.some((e) => e.type === "new_message")).toBe(false); + }); + + it("脚本遗漏工具结果时应补齐完整结果组,并在原子提交后才发布事件", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "script_tool", description: "script", parameters: { type: "object", properties: {} } }, + ], + execute: vi.fn().mockResolvedValue([{ id: "call-1", result: "ok" }]), + }; + callLLM + .mockResolvedValueOnce({ + content: "", + toolCalls: [ + { id: "call-1", name: "script_tool", arguments: "{}" }, + { id: "call-2", name: "script_tool", arguments: "{}" }, + ], + }) + .mockResolvedValueOnce(finalTextResult("done")); + chatRepo.commitToolRound.mockImplementation(async () => { + expect(sendEvent.mock.calls.some((call) => call[0].type === "tool_call_complete")).toBe(false); + }); + + await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })); + + const [, toolMessages] = chatRepo.commitToolRound.mock.calls[0]; + expect(toolMessages).toHaveLength(2); + expect(toolMessages.map((message: any) => message.toolCallId)).toEqual(["call-1", "call-2"]); + expect(toolMessages[1].content).toContain("missing"); + const secondRequest = callLLM.mock.calls[1][1]; + expect(secondRequest.messages.filter((message: any) => message.role === "tool")).toHaveLength(2); + expect(sendEvent.mock.calls.filter((call) => call[0].type === "tool_call_complete")).toHaveLength(2); + }); + + it("工具结果组提交失败时应回收本轮新建附件且不发布完成事件", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "image_tool", description: "image", parameters: { type: "object", properties: {} } }, + ], + execute: vi.fn().mockResolvedValue([ + { + id: "call-image", + result: "image", + attachments: [{ id: "owned.png", type: "image", name: "owned.png", mimeType: "image/png" }], + ownedAttachmentIds: ["owned.png"], + }, + ]), + }; + callLLM.mockResolvedValueOnce({ + content: "", + contentBlocks: [{ type: "image", attachmentId: "generated.png", mimeType: "image/png" }], + toolCalls: [{ id: "call-image", name: "image_tool", arguments: "{}" }], + usage: { inputTokens: 12, outputTokens: 4 }, + }); + chatRepo.commitToolRound.mockRejectedValueOnce(new Error("disk full")); + + const error = await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })).catch((reason) => reason); + + expect(error).toMatchObject({ + message: "disk full", + usage: expect.objectContaining({ inputTokens: 12, outputTokens: 4 }), + }); + expect(chatRepo.deleteAttachment).toHaveBeenCalledWith("owned.png"); + expect(chatRepo.deleteAttachment).toHaveBeenCalledWith("generated.png"); + expect(sendEvent.mock.calls.some((call) => call[0].type === "tool_call_complete")).toBe(false); + }); + + it("工具结果组提交报错但读回确认已落盘时应保留附件并发布结果", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "image_tool", description: "image", parameters: { type: "object", properties: {} } }, + ], + execute: vi.fn().mockResolvedValue([ + { + id: "call-image", + result: "image", + attachments: [{ id: "owned.png", type: "image", name: "owned.png", mimeType: "image/png" }], + ownedAttachmentIds: ["owned.png"], + }, + ]), + }; + callLLM + .mockResolvedValueOnce({ + content: "", + contentBlocks: [{ type: "image", attachmentId: "generated.png", mimeType: "image/png" }], + toolCalls: [{ id: "call-image", name: "image_tool", arguments: "{}" }], + }) + .mockResolvedValueOnce(finalTextResult("done")); + chatRepo.commitToolRound.mockImplementationOnce(async (assistant: any, toolMessages: any[]) => { + chatRepo.getMessageSnapshot.mockResolvedValueOnce({ + generation: "gen-1", + revision: 1, + messages: [assistant, ...toolMessages], + }); + throw new Error("ambiguous close failure"); + }); + + await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })); + + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("owned.png"); + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("generated.png"); + expect(sendEvent.mock.calls.some((call) => call[0].type === "tool_call_complete")).toBe(true); + }); + + it("提交报错且确认读本身也失败(不确定态)时不应删除附件,而不是当作未落盘", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "image_tool", description: "image", parameters: { type: "object", properties: {} } }, + ], + execute: vi.fn().mockResolvedValue([ + { + id: "call-image", + result: "image", + attachments: [{ id: "owned.png", type: "image", name: "owned.png", mimeType: "image/png" }], + ownedAttachmentIds: ["owned.png"], + }, + ]), + }; + callLLM + .mockResolvedValueOnce({ + content: "", + contentBlocks: [{ type: "image", attachmentId: "generated.png", mimeType: "image/png" }], + toolCalls: [{ id: "call-image", name: "image_tool", arguments: "{}" }], + usage: { inputTokens: 12, outputTokens: 4 }, + }) + .mockResolvedValueOnce(finalTextResult("done")); + chatRepo.commitToolRound.mockRejectedValueOnce(new Error("disk full")); + // 确认读也失败:无法证实写入是否落盘,属于不确定态,不能等同于"确实未落盘" + chatRepo.getMessageSnapshot.mockRejectedValueOnce(new Error("read failed")); + + const error = await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })).catch((reason) => reason); + + expect(error).toBeUndefined(); + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("owned.png"); + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("generated.png"); + }); + + it("提交与幂等重试均报错、确认读本身也持续失败时应以 persist_indeterminate 终止,不发布也不删附件", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "image_tool", description: "image", parameters: { type: "object", properties: {} } }, + ], + execute: vi.fn().mockResolvedValue([ + { + id: "call-image", + result: "image", + attachments: [{ id: "owned.png", type: "image", name: "owned.png", mimeType: "image/png" }], + ownedAttachmentIds: ["owned.png"], + }, + ]), + }; + callLLM.mockResolvedValueOnce({ + content: "", + contentBlocks: [{ type: "image", attachmentId: "generated.png", mimeType: "image/png" }], + toolCalls: [{ id: "call-image", name: "image_tool", arguments: "{}" }], + usage: { inputTokens: 12, outputTokens: 4 }, + }); + // 首次提交失败、确认读失败(不确定态);幂等重试提交依旧失败、重试后的确认读依旧失败—— + // 无法在合理次数内确认落盘状态,必须终止而不是当作成功发布 + chatRepo.commitToolRound + .mockRejectedValueOnce(new Error("disk full")) + .mockRejectedValueOnce(new Error("disk full")); + chatRepo.getMessageSnapshot + .mockRejectedValueOnce(new Error("read failed")) + .mockRejectedValueOnce(new Error("read failed")); + + const error = await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })).catch((reason) => reason); + + expect(error).toMatchObject({ errorCode: "persist_indeterminate" }); + expect(chatRepo.commitToolRound).toHaveBeenCalledTimes(2); + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("owned.png"); + expect(chatRepo.deleteAttachment).not.toHaveBeenCalledWith("generated.png"); + expect(sendEvent.mock.calls.some((call) => call[0].type === "tool_call_complete")).toBe(false); + }); + + it("工具内部摘要 LLM 的 usage 应恰好一次计入父对话终态", async () => { + toolRegistry = { + getDefinitions: () => [ + { name: "web_fetch", description: "fetch", parameters: { type: "object", properties: {} } }, + ], + execute: vi + .fn() + .mockResolvedValue([{ id: "call-summary", result: "summary", usage: { inputTokens: 30, outputTokens: 8 } }]), + }; + callLLM + .mockResolvedValueOnce({ + content: "", + toolCalls: [{ id: "call-summary", name: "web_fetch", arguments: "{}" }], + usage: { inputTokens: 10, outputTokens: 2 }, + }) + .mockResolvedValueOnce({ content: "done", usage: { inputTokens: 5, outputTokens: 1 } }); + + await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })); + + const done = sendEvent.mock.calls.map((call) => call[0]).find((event) => event.type === "done"); + expect(done?.usage).toEqual(expect.objectContaining({ inputTokens: 45, outputTokens: 11 })); + }); + + it("触发 autoCompact 时应保留原始模型,而不是把 contextWindow 预先缩小", async () => { + callLLM.mockResolvedValue({ content: "done", usage: { inputTokens: 8500, outputTokens: 5 } }); + + await orchestrator.callLLMWithToolLoop( + baseParams({ model: { ...MODEL, contextWindow: 10_000, maxTokens: 2_000 } }) + ); + + expect(callLLM).toHaveBeenCalledTimes(1); + expect(callLLM.mock.calls[0][0]).toMatchObject({ contextWindow: 10_000, maxTokens: 2_000 }); + expect(autoCompact).toHaveBeenCalledTimes(1); + expect(autoCompact.mock.calls[0][2]).toMatchObject({ contextWindow: 10_000, maxTokens: 2_000 }); + }); + + it("输入占用未达到完整上下文窗口 80% 时不应提前自动压缩", async () => { + callLLM.mockResolvedValue({ content: "done", usage: { inputTokens: 6000, outputTokens: 5 } }); + + await orchestrator.callLLMWithToolLoop( + baseParams({ model: { ...MODEL, contextWindow: 10_000, maxTokens: 2_000 } }) + ); + + expect(autoCompact).not.toHaveBeenCalled(); + }); + + it("自动压缩的 token 用量应计入最终 done 事件", async () => { + callLLM.mockResolvedValue({ content: "done", usage: { inputTokens: 8500, outputTokens: 5 } }); + autoCompact.mockResolvedValue({ + inputTokens: 120, + outputTokens: 30, + cacheCreationInputTokens: 10, + cacheReadInputTokens: 5, + }); + + await orchestrator.callLLMWithToolLoop( + baseParams({ model: { ...MODEL, contextWindow: 10_000, maxTokens: 2_000 } }) + ); + + const doneEvent = sendEvent.mock.calls.find((call) => call[0].type === "done")?.[0]; + expect(doneEvent).toMatchObject({ + usage: { inputTokens: 8620, outputTokens: 35, cacheCreationInputTokens: 10, cacheReadInputTokens: 5 }, + }); + }); + + it("续接长历史时应完整透传旧工具结果且不修改原始历史", async () => { + const oldToolResult = "完整工具结果".repeat(100); + const messages: ChatRequest["messages"] = []; + for (let i = 0; i < 6; i++) { + messages.push({ + role: "assistant", + content: "", + toolCalls: [{ id: `tc${i}`, name: "dup", arguments: "{}" }], + }); + messages.push({ role: "tool", content: oldToolResult, toolCallId: `tc${i}` }); + } + callLLM.mockImplementation(async (_model, request) => { + expect(request.messages[1].content).toBe(oldToolResult); + expect(messages[1].content).toBe(oldToolResult); + return finalTextResult("done"); + }); + + await orchestrator.callLLMWithToolLoop(baseParams({ messages, model: { ...MODEL, contextWindow: 1000 } })); + }); + + it("第 2 次触发告警时应暂停并询问用户;回答非 Stop 时应继续循环", async () => { + callLLM + .mockResolvedValueOnce(dupToolCallResult("c1")) + .mockResolvedValueOnce(dupToolCallResult("c2")) + .mockResolvedValueOnce(dupToolCallResult("c3")) + .mockResolvedValueOnce(dupToolCallResult("c4")) + .mockResolvedValueOnce(finalTextResult("done")); + + const askUserForGuard = vi.fn().mockResolvedValue("Continue"); + + await orchestrator.callLLMWithToolLoop(baseParams({ askUserForGuard })); + + expect(askUserForGuard).toHaveBeenCalledTimes(1); + expect(askUserForGuard.mock.calls[0][0]).toBe(2); + + expect(callLLM).toHaveBeenCalledTimes(5); + const doneEvents = sendEvent.mock.calls.filter((c) => c[0].type === "done"); + expect(doneEvents).toHaveLength(1); + }); + + it("回答 Stop 时应提前结束循环,以 done(而非 error)收尾", async () => { + callLLM + .mockResolvedValueOnce(dupToolCallResult("c1")) + .mockResolvedValueOnce(dupToolCallResult("c2")) + .mockResolvedValueOnce(dupToolCallResult("c3")) + .mockResolvedValueOnce(dupToolCallResult("c4")) + .mockResolvedValueOnce(finalTextResult("done")); // 若未提前结束,不应被调用到 + + const askUserForGuard = vi.fn().mockResolvedValue("Stop"); + + await orchestrator.callLLMWithToolLoop(baseParams({ askUserForGuard })); + + // 应在第 4 轮后立即停止,不再发起第 5 次 LLM 调用 + expect(callLLM).toHaveBeenCalledTimes(4); + + const errorEvents = sendEvent.mock.calls.filter((c) => c[0].type === "error"); + expect(errorEvents).toHaveLength(0); + const doneEvents = sendEvent.mock.calls.filter((c) => c[0].type === "done"); + expect(doneEvents).toHaveLength(1); + + // 停止信息应被持久化 + expect(chatRepo.appendMessage).toHaveBeenCalledWith( + expect.objectContaining({ role: "assistant", conversationId: "conv-1" }), + "gen-1" + ); + }); + + it("回答 Continue 后应重置命中计数,之后需再次连续命中 2 次才会重新暂停询问", async () => { + // 第 1~4 轮:触发前两次命中(第 2、4 轮),第 4 轮暂停询问,回答 Continue + // 第 5~6 轮:命中第 3 次(重置后的第 1 次),不应再次暂停 + // 第 7~8 轮:命中第 4 次(重置后的第 2 次),应再次暂停询问 + // 第 9 轮:最终文本,结束循环 + for (let i = 1; i <= 8; i++) { + callLLM.mockResolvedValueOnce(dupToolCallResult(`c${i}`)); + } + callLLM.mockResolvedValueOnce(finalTextResult("done")); + + const askUserForGuard = vi.fn().mockResolvedValue("Continue"); + + await orchestrator.callLLMWithToolLoop(baseParams({ askUserForGuard })); + + // 仅在第 4 轮和第 8 轮各暂停一次,中间第 6 轮的命中(重置后第 1 次)不应触发暂停 + expect(askUserForGuard).toHaveBeenCalledTimes(2); + expect(askUserForGuard.mock.calls[0][0]).toBe(2); + expect(askUserForGuard.mock.calls[1][0]).toBe(2); + + expect(callLLM).toHaveBeenCalledTimes(9); + const doneEvents = sendEvent.mock.calls.filter((c) => c[0].type === "done"); + expect(doneEvents).toHaveLength(1); + }); +}); + +describe("ToolLoopOrchestrator 工具结果透传", () => { + let chatRepo: ReturnType; + let toolRegistry: ToolExecutorLike; + let callLLM: ReturnType>; + let autoCompact: ReturnType>; + let deps: ToolLoopDeps; + let orchestrator: ToolLoopOrchestrator; + let sendEvent: ReturnType; + + beforeEach(() => { + chatRepo = makeFakeChatRepo(); + toolRegistry = makeFakeToolRegistry(); + callLLM = vi.fn(); + autoCompact = vi.fn().mockResolvedValue(undefined); + deps = { callLLM, autoCompact }; + orchestrator = new ToolLoopOrchestrator(deps, chatRepo); + sendEvent = vi.fn(); + }); + + function baseParams(overrides: Record = {}) { + return { + toolRegistry, + model: { ...MODEL, contextWindow: 20000 }, + messages: [{ role: "user", content: "开始" }] as ChatRequest["messages"], + sendEvent: sendEvent as (event: ChatStreamEvent) => void, + signal: new AbortController().signal, + scriptToolCallback: null, + conversationId: "conv-1", + ...overrides, + }; + } + + it("大型工具结果应完整传给下一轮模型调用", async () => { + const hugeResult = "巨大的工具结果".repeat(2000); + toolRegistry = { + getDefinitions: () => [{ name: "dup", description: "dup", parameters: { type: "object", properties: {} } }], + execute: async (toolCalls: ToolCall[]): Promise => + toolCalls.map((tc) => ({ id: tc.id, result: hugeResult })), + }; + deps = { callLLM, autoCompact }; + orchestrator = new ToolLoopOrchestrator(deps, chatRepo); + + callLLM + .mockResolvedValueOnce({ content: "", toolCalls: [{ id: "c1", name: "dup", arguments: "{}" }] } as LLMCallResult) + .mockImplementationOnce(async (_model, request) => { + const toolMsg = request.messages.find((m) => m.role === "tool"); + expect(toolMsg?.content).toBe(hugeResult); + return finalTextResult("done"); + }); + + await orchestrator.callLLMWithToolLoop(baseParams({ toolRegistry })); + + expect(callLLM).toHaveBeenCalledTimes(2); + }); +}); diff --git a/src/app/service/agent/service_worker/tool_loop_orchestrator.ts b/src/app/service/agent/service_worker/tool_loop_orchestrator.ts index 09ab85769..30bdf8a5e 100644 --- a/src/app/service/agent/service_worker/tool_loop_orchestrator.ts +++ b/src/app/service/agent/service_worker/tool_loop_orchestrator.ts @@ -4,34 +4,94 @@ import type { AgentModelConfig, ChatRequest, ChatStreamEvent, + ChatMessage, ToolDefinition, ToolCall, Attachment, SubAgentDetails, ContentBlock, MessageContent, + TokenUsage, } from "@App/app/service/agent/core/types"; import { uuidv4 } from "@App/pkg/utils/uuid"; import { getContextWindow } from "@App/app/service/agent/core/model_context"; import { detectToolCallIssues, type ToolCallRecord } from "@App/app/service/agent/core/tool_call_guard"; import type { LLMCallResult } from "./llm_client"; +import { t } from "@App/locales/locales"; +import { prepareAttachmentSnapshot, type AttachmentSnapshot } from "@App/app/service/agent/core/attachment_resolver"; + +// Compact 路径使用启发式估算时的硬拒绝阈值;正常 Tool Loop 不再按估算主动裁剪上下文。 +export const HEURISTIC_HARD_REJECT_RATIO = 2; + +// 循环检测(tool_call_guard)连续命中达到此次数时,暂停并询问用户是否继续(仅当调用方提供 askUserForGuard 时生效) +const GUARD_ESCALATION_STRIKES = 2; +const GUARD_STOP_ANSWER = "stop"; + +type AskUserForGuard = (strikeCount: number) => Promise; + +/** + * Provider 归一化后的实际上下文输入 token。 + * Anthropic 将缓存命中/写入 token 与断点后的 input_tokens 分开返回;OpenAI 的 prompt_tokens 已含缓存部分。 + */ +function getContextInputTokens(model: AgentModelConfig, usage: NonNullable): number { + if (model.provider !== "anthropic") return usage.inputTokens; + return usage.inputTokens + (usage.cacheCreationInputTokens || 0) + (usage.cacheReadInputTokens || 0); +} + +/** 等待 loop-guard 回答;AbortSignal 触发时立即返回 null,不再阻塞会话停止/断开。 */ +function waitForGuardAnswer( + askUserForGuard: AskUserForGuard, + strikeCount: number, + signal: AbortSignal +): Promise { + if (signal.aborted) return Promise.resolve(null); + + return new Promise((resolve, reject) => { + let settled = false; + const cleanup = () => signal.removeEventListener("abort", onAbort); + const resolveOnce = (answer: string | null) => { + if (settled) return; + settled = true; + cleanup(); + resolve(answer); + }; + const rejectOnce = (error: unknown) => { + if (settled) return; + settled = true; + cleanup(); + reject(error); + }; + const onAbort = () => resolveOnce(null); + + signal.addEventListener("abort", onAbort, { once: true }); + Promise.resolve() + .then(() => (settled || signal.aborted ? null : askUserForGuard(strikeCount))) + .then(resolveOnce, rejectOnce); + }); +} /** ToolLoopOrchestrator 所需的外部依赖(由 AgentService 注入) */ export interface ToolLoopDeps { // callLLM 通过 lambda 注入,确保测试 spy 可以拦截 callLLM( model: AgentModelConfig, - params: { messages: ChatRequest["messages"]; tools?: ToolDefinition[]; cache?: boolean }, + params: { + messages: ChatRequest["messages"]; + tools?: ToolDefinition[]; + cache?: boolean; + attachmentSnapshot?: AttachmentSnapshot; + }, sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal ): Promise; autoCompact( conversationId: string, + conversationGeneration: string, model: AgentModelConfig, messages: ChatRequest["messages"], sendEvent: (event: ChatStreamEvent) => void, signal: AbortSignal - ): Promise; + ): Promise; } export class ToolLoopOrchestrator { @@ -43,6 +103,96 @@ export class ToolLoopOrchestrator { private chatRepo: AgentChatRepo ) {} + /** 生成的附件(如模型产出的图片)在被 assistant 消息持久化引用之前只是"本轮租约": + * 任何不会把引用消息落库的退出路径(取消、持久化失败、落库异常)都必须删除这些文件, + * 否则它们会成为无任何消息引用的孤儿附件。 */ + private async releaseGeneratedAttachments(result: LLMCallResult): Promise { + const blocks = result.contentBlocks; + if (!blocks || blocks.length === 0) return; + await Promise.all( + blocks + .filter((block): block is Exclude => block.type !== "text") + .map((block) => this.chatRepo.deleteAttachment(block.attachmentId).catch(() => {})) + ); + } + + /** 工具轮次落盘状态:确认读本身失败时必须与"确实未提交"区分开——前者是不确定态,不能 + * 被当作可以安全删除附件的证据。 */ + private async checkToolRoundDurability( + conversationId: string, + conversationGeneration: string | undefined, + assistantMessage: ChatMessage, + toolMessages: ChatMessage[] + ): Promise<"durable" | "not_durable" | "indeterminate"> { + let snapshot: Awaited>; + try { + snapshot = await this.chatRepo.getMessageSnapshot(conversationId, conversationGeneration); + } catch { + // 确认读失败不代表写入未落盘,只是无法证实——不确定态 + return "indeterminate"; + } + const assistant = snapshot.messages.find( + (message) => message.id === assistantMessage.id && message.role === "assistant" + ); + if (!assistant) return "not_durable"; + const expectedToolCallIds = new Set((assistantMessage.toolCalls || []).map((toolCall) => toolCall.id)); + if (toolMessages.length !== expectedToolCallIds.size) return "not_durable"; + const durable = toolMessages.every( + (expected) => + expected.toolCallId !== undefined && + expectedToolCallIds.has(expected.toolCallId) && + snapshot.messages.some( + (message) => + message.id === expected.id && message.role === "tool" && message.toolCallId === expected.toolCallId + ) + ); + return durable ? "durable" : "not_durable"; + } + + /** 取消(stop)落定时的统一终态化:持久化一条终态记录 + 发送唯一的终态事件,携带累计 usage/耗时。 + * 落库失败不能阻塞事件发送,否则客户端永远收不到终态事件。 */ + private async emitCancelled( + conversationId: string | undefined, + conversationGeneration: string | undefined, + totalUsage: { + inputTokens: number; + outputTokens: number; + cacheCreationInputTokens: number; + cacheReadInputTokens: number; + }, + startTime: number, + sendEvent: (event: ChatStreamEvent) => void + ): Promise { + const durationMs = Date.now() - startTime; + if (conversationId) { + try { + await this.chatRepo.appendMessage( + { + id: uuidv4(), + conversationId, + role: "assistant", + content: "", + error: "Conversation cancelled", + errorCode: "cancelled", + usage: totalUsage, + durationMs, + createtime: Date.now(), + }, + conversationGeneration + ); + } catch { + // 持久化失败不阻塞终态事件发送 + } + } + sendEvent({ + type: "error", + message: "Conversation cancelled", + errorCode: "cancelled", + usage: totalUsage, + durationMs, + }); + } + // 统一的 tool calling 循环,UI 和脚本共用 async callLLMWithToolLoop(params: { // 本次调用使用的工具注册表(SessionToolRegistry 或 ToolRegistry) @@ -50,40 +200,72 @@ export class ToolLoopOrchestrator { model: AgentModelConfig; messages: ChatRequest["messages"]; tools?: ToolDefinition[]; - maxIterations: number; sendEvent: (event: ChatStreamEvent) => void; signal: AbortSignal; // 脚本自定义工具的回调,null 表示只用内置工具 scriptToolCallback: ScriptToolCallback | null; // 对话 ID,用于持久化消息(可选,UI 场景由 hooks 自行持久化) conversationId?: string; + // 持久化会话的不可变 generation;旧执行在删除/重建后不得写入新会话。 + conversationGeneration?: string; // 跳过内置工具,仅使用传入的 tools(ephemeral 模式) skipBuiltinTools?: boolean; // 排除的工具名称列表(子代理不可用 ask_user、agent) excludeTools?: string[]; // 是否启用 prompt caching,默认 true cache?: boolean; + // 消息来自持久化历史时使用独立副本,避免 Tool Loop 修改持久化历史对象。 + rehydratedHistory?: boolean; + // 循环检测连续命中达到 GUARD_ESCALATION_STRIKES 次时调用,暂停循环询问用户是否继续。 + // 仅由 UI 对话(含后台会话)传入;定时任务、子代理不传,保持原有的仅告警不暂停行为。 + askUserForGuard?: AskUserForGuard; + // 定时任务和子代理需要将终态错误(如 context_too_large)作为失败抛给调用方。 + throwOnTerminalError?: boolean; }): Promise { const { toolRegistry, model, - messages, + messages: inputMessages, tools, - maxIterations, sendEvent, signal, scriptToolCallback, conversationId, + conversationGeneration, + rehydratedHistory, } = params; - const startTime = Date.now(); + const totalUsage = { + inputTokens: 0, + outputTokens: 0, + cacheCreationInputTokens: 0, + cacheReadInputTokens: 0, + }; + let attachmentSnapshot = await prepareAttachmentSnapshot( + inputMessages, + model, + (id) => this.chatRepo.getAttachment(id), + signal + ); + // 持久化历史使用独立副本;正常 Tool Loop 保持消息前缀和工具结果完整。 + const messages = rehydratedHistory ? inputMessages.map((message) => ({ ...message })) : inputMessages; + let iterations = 0; - const totalUsage = { inputTokens: 0, outputTokens: 0, cacheCreationInputTokens: 0, cacheReadInputTokens: 0 }; const toolCallHistory: ToolCallRecord[] = []; let guardStartIndex = 0; + // 循环检测命中次数,达到 GUARD_ESCALATION_STRIKES 后暂停询问用户 + let guardStrikeCount = 0; - while (iterations < maxIterations) { + while (true) { iterations++; + if (iterations > 1) { + attachmentSnapshot = await prepareAttachmentSnapshot( + messages, + model, + (id) => this.chatRepo.getAttachment(id), + signal + ); + } // 每轮重新获取工具定义(load_skill 可能动态注册了新工具) let allToolDefs = params.skipBuiltinTools ? tools || [] : toolRegistry.getDefinitions(tools); @@ -92,17 +274,45 @@ export class ToolLoopOrchestrator { allToolDefs = allToolDefs.filter((t) => !excludeSet.has(t.name)); } - // 调用 LLM(重试由 llm_client 内部处理) - const result = await this.deps.callLLM( - model, - { messages, tools: allToolDefs.length > 0 ? allToolDefs : undefined, cache: params.cache }, - sendEvent, - signal - ); - - if (signal.aborted) return; + // 重试由 llm_client 内部处理 + let result: LLMCallResult; + try { + result = await this.deps.callLLM( + model, + { + messages, + tools: allToolDefs.length > 0 ? allToolDefs : undefined, + cache: params.cache, + attachmentSnapshot, + }, + sendEvent, + signal + ); + } catch (error) { + // 无论取消还是真实失败,provider 层挂在错误上的本轮已知部分 usage(如 Anthropic 的 + // message_start、部分 OpenAI 兼容 API 每个 chunk 都带的 usage)都必须并入 totalUsage, + // 否则这部分已经产生的花费会从终态 usage、定时任务与子代理的累计里丢失 + const partialUsage = (error as { usage?: LLMCallResult["usage"] })?.usage; + if (partialUsage) { + totalUsage.inputTokens += partialUsage.inputTokens; + totalUsage.outputTokens += partialUsage.outputTokens; + totalUsage.cacheCreationInputTokens += partialUsage.cacheCreationInputTokens || 0; + totalUsage.cacheReadInputTokens += partialUsage.cacheReadInputTokens || 0; + } + // SSE 解析层现在会在 abort 时 reject(而不是静默挂起,见 content_utils.ts), + // 这类 reject 必须走统一的取消终态化路径,而不是当作真实错误往外抛 + if (signal.aborted) { + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { + usage: totalUsage, + durationMs: Date.now() - startTime, + conversationId, + }); + } - // 累计 usage + // 先累计本轮 usage 再检查 aborted:即使取消发生在这次响应之后,其花费也不应从终态 usage 中丢失 if (result.usage) { totalUsage.inputTokens += result.usage.inputTokens; totalUsage.outputTokens += result.usage.outputTokens; @@ -110,13 +320,60 @@ export class ToolLoopOrchestrator { totalUsage.cacheReadInputTokens += result.usage.cacheReadInputTokens || 0; } - // 自动 compact:当上下文占用超过 80% 时触发 + if (signal.aborted) { + // 本轮生成的附件尚未被任何持久化消息引用,取消退出前必须回收 + await this.releaseGeneratedAttachments(result); + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + + // 自动 compact:按模型完整上下文窗口计算,避免因输出预留和安全边际把“80%”提前触发。 if (result.usage && conversationId) { - const contextWindow = getContextWindow(model); - const usageRatio = result.usage.inputTokens / contextWindow; + const contextUsageRatio = getContextInputTokens(model, result.usage) / getContextWindow(model); - if (usageRatio >= 0.8) { - await this.deps.autoCompact(conversationId, model, messages, sendEvent, signal); + if (contextUsageRatio >= 0.8) { + try { + const compactUsage = await this.deps.autoCompact( + conversationId, + conversationGeneration!, + model, + messages, + sendEvent, + signal + ); + if (compactUsage) { + totalUsage.inputTokens += compactUsage.inputTokens; + totalUsage.outputTokens += compactUsage.outputTokens; + totalUsage.cacheCreationInputTokens += compactUsage.cacheCreationInputTokens || 0; + totalUsage.cacheReadInputTokens += compactUsage.cacheReadInputTokens || 0; + } + } catch (error) { + const compactUsage = (error as { usage?: TokenUsage })?.usage; + if (compactUsage) { + totalUsage.inputTokens += compactUsage.inputTokens; + totalUsage.outputTokens += compactUsage.outputTokens; + totalUsage.cacheCreationInputTokens += compactUsage.cacheCreationInputTokens || 0; + totalUsage.cacheReadInputTokens += compactUsage.cacheReadInputTokens || 0; + } + // 本轮结果的 assistant 消息不会再持久化,生成的附件必须先回收 + await this.releaseGeneratedAttachments(result); + if (signal.aborted) { + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { + usage: totalUsage, + durationMs: Date.now() - startTime, + conversationId, + }); + } + // autoCompact 期间可能已被 stop:继续持久化/发送最终消息前必须重新检查, + // 并统一走取消终态化路径(唯一一条终态事件 + 已回写的累计 usage),而不是静默 return + if (signal.aborted) { + await this.releaseGeneratedAttachments(result); + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } } } @@ -133,93 +390,226 @@ export class ToolLoopOrchestrator { // 如果有 tool calls,需要执行并继续循环 if (result.toolCalls && result.toolCalls.length > 0 && allToolDefs.length > 0) { - // 持久化 assistant 消息(含 tool calls) + // 先构造 assistant 消息,等全部工具结果归一化后与完整 tool 结果组一次性提交。 + let persistedAssistantMessage: ChatMessage | undefined; if (conversationId) { - await this.chatRepo.appendMessage({ + persistedAssistantMessage = { id: uuidv4(), conversationId, role: "assistant", content: buildMessageContent(), + ownedAttachmentIds: result.contentBlocks + ?.filter((block) => block.type !== "text") + .map((block) => block.attachmentId), thinking: result.thinking ? { content: result.thinking } : undefined, toolCalls: result.toolCalls, + warning: result.warning, createtime: Date.now(), - }); + }; } // 将 assistant 消息加入上下文(带 toolCalls,供 provider 构建 tool_calls 字段) - messages.push({ role: "assistant", content: result.content || "", toolCalls: result.toolCalls }); + messages.push({ + role: "assistant", + content: buildMessageContent(), + toolCalls: result.toolCalls, + }); + + // 生成图片保存失败等警告:与最终回复分支同样的即时可见提示,不等到工具轮结束—— + // 消息里的 warning 字段(上面已写入 persistedAssistantMessage)负责刷新后仍可见 + if (result.warning) { + sendEvent({ type: "system_warning", message: result.warning }); + } // 通过 ToolRegistry 执行工具(内置工具直接执行,脚本工具回调 Sandbox) // excludeTools 做后端强校验:被排除的工具名直接返回 error,避免 LLM 盲调绕过白/黑名单 const excludeToolsSet = params.excludeTools && params.excludeTools.length > 0 ? new Set(params.excludeTools) : undefined; - const toolResults = await toolRegistry.execute(result.toolCalls, scriptToolCallback, excludeToolsSet); + // 脚本工具通过 raceWithAbort 包裹:abort 时会直接 reject,而不是返回部分结果, + // 必须捕获后按"取消"处理,复用下面统一的补全逻辑,而不是让异常直接抛出跳过终态化 + let toolResults: Awaited>; + try { + toolResults = await toolRegistry.execute(result.toolCalls, scriptToolCallback, excludeToolsSet, signal); + } catch (error) { + if (!signal.aborted) { + await this.releaseGeneratedAttachments(result); + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { + usage: totalUsage, + durationMs: Date.now() - startTime, + conversationId, + }); + } + toolResults = []; + } + const cancelledDuringTools = signal.aborted; + // 外部脚本可能返回缺失、重复或未知 ID。始终按请求顺序规整为每个 call 恰好一个结果, + // 否则下一轮 provider 请求和持久化历史都会违反 tool-call 协议。 + const returnedById = new Map(); + const requestedIds = new Set(result.toolCalls.map((toolCall) => toolCall.id)); + const discardedOwnedAttachmentIds: string[] = []; + for (const toolResult of toolResults) { + if (requestedIds.has(toolResult.id) && !returnedById.has(toolResult.id)) { + returnedById.set(toolResult.id, toolResult); + } else { + discardedOwnedAttachmentIds.push(...(toolResult.ownedAttachmentIds || [])); + } + } + await Promise.all(discardedOwnedAttachmentIds.map((id) => this.chatRepo.deleteAttachment(id).catch(() => {}))); + toolResults = result.toolCalls.map( + (toolCall) => + returnedById.get(toolCall.id) || { + id: toolCall.id, + result: JSON.stringify({ + error: cancelledDuringTools ? "Tool execution cancelled" : "Tool result missing", + }), + error: true, + } + ); + for (const toolResult of toolResults) { + if (!toolResult.usage) continue; + totalUsage.inputTokens += toolResult.usage.inputTokens; + totalUsage.outputTokens += toolResult.usage.outputTokens; + totalUsage.cacheCreationInputTokens += toolResult.usage.cacheCreationInputTokens || 0; + totalUsage.cacheReadInputTokens += toolResult.usage.cacheReadInputTokens || 0; + } - // 将 tool 结果加入消息,并通知 UI 工具执行完成 // 收集需要回写的 toolCall 元数据(执行状态 / 附件 / 子代理详情) const attachmentUpdates = new Map(); + const ownershipUpdates = new Map(); const subAgentUpdates = new Map(); const completedIds = new Set(); + const failedIds = new Set(); for (const tr of toolResults) { - // LLM 上下文只包含文本结果,不含附件 - messages.push({ role: "tool", content: tr.result, toolCallId: tr.id }); - // 通知 UI 工具执行完成(含附件元数据) - sendEvent({ type: "tool_call_complete", id: tr.id, result: tr.result, attachments: tr.attachments }); - completedIds.add(tr.id); + if (tr.error) failedIds.add(tr.id); if (tr.attachments?.length) { attachmentUpdates.set(tr.id, tr.attachments); } + if (tr.ownedAttachmentIds?.length) { + ownershipUpdates.set(tr.id, tr.ownedAttachmentIds); + } if (tr.subAgentDetails) { subAgentUpdates.set(tr.id, tr.subAgentDetails); } - - // 持久化 tool 结果消息 - if (conversationId) { - await this.chatRepo.appendMessage({ - id: uuidv4(), - conversationId, - role: "tool", - content: tr.result, - toolCallId: tr.id, - createtime: Date.now(), - }); - } } - // 回写工具执行结果到 assistant 消息的 toolCalls(内存 + 持久化)。 - // assistant 消息在执行前已落库,其 toolCalls 的 status 仍是 "running"; - // 必须在此回写为 "completed",否则刷新/重载会从库里读回 running,导致工具图标一直转圈。 + // 工具结果先回写到内存中的 assistant toolCalls,再与全部 tool 消息原子提交; + // 持久化历史因此不会暴露只有 running assistant 或缺少结果的半轮状态。 const applyToolUpdates = (toolCalls: ToolCall[]) => { for (const tc of toolCalls) { - if (completedIds.has(tc.id)) tc.status = "completed"; + if (failedIds.has(tc.id)) tc.status = "error"; + else if (completedIds.has(tc.id)) tc.status = "completed"; const atts = attachmentUpdates.get(tc.id); if (atts) tc.attachments = atts; + const ownedAttachmentIds = ownershipUpdates.get(tc.id); + if (ownedAttachmentIds) tc.ownedAttachmentIds = ownedAttachmentIds; const sad = subAgentUpdates.get(tc.id); if (sad) tc.subAgentDetails = sad; } }; - // 内存上下文中的 assistant 消息 - const assistantMsg = messages.find( - (m) => m.role === "assistant" && m.toolCalls?.some((tc: ToolCall) => completedIds.has(tc.id)) - ); + // 内存上下文中的 assistant 消息:目标消息一定是刚 push 过 toolCalls 的那条, + // 从尾部往回找(本轮只追加了少量消息)比 Array.prototype.find 的正向全量扫描更快。 + // persistedAssistantMessage.toolCalls 与这里的 toolCalls 是同一个数组引用(都来自 + // result.toolCalls),因此这一次 applyToolUpdates 同时完成了内存态与待持久化对象的回写。 + let assistantMsg: (typeof messages)[number] | undefined; + for (let i = messages.length - 1; i >= 0; i--) { + const m = messages[i]; + if (m.role === "assistant" && m.toolCalls?.some((tc: ToolCall) => completedIds.has(tc.id))) { + assistantMsg = m; + break; + } + } if (assistantMsg?.toolCalls) applyToolUpdates(assistantMsg.toolCalls); - // 持久化的 assistant 消息 - if (conversationId) { - const allMessages = await this.chatRepo.getMessages(conversationId); - for (let i = allMessages.length - 1; i >= 0; i--) { - const msg = allMessages[i]; - if (msg.role === "assistant" && msg.toolCalls?.some((tc: ToolCall) => completedIds.has(tc.id))) { - applyToolUpdates(msg.toolCalls!); - await this.chatRepo.saveMessages(conversationId, allMessages); - break; + const persistedToolMessages: ChatMessage[] = conversationId + ? toolResults.map((toolResult) => ({ + id: uuidv4(), + conversationId, + role: "tool" as const, + content: toolResult.result, + toolCallId: toolResult.id, + createtime: Date.now(), + })) + : []; + + if (persistedAssistantMessage && conversationId) { + try { + await this.chatRepo.commitToolRound( + persistedAssistantMessage, + persistedToolMessages, + conversationGeneration + ); + } catch (error) { + let durability = await this.checkToolRoundDurability( + conversationId, + conversationGeneration, + persistedAssistantMessage, + persistedToolMessages + ); + if (durability === "indeterminate") { + // 确认读本身失败,无法判断是否已落盘。commitToolRound 按消息 id 去重(同一 id 的 + // assistant/tool 消息会先被过滤掉再重新写入),用同一批消息重试是幂等的——重试一次 + // 并重新确认,而不是直接把不确定态当成功发布或直接判定失败。 + try { + await this.chatRepo.commitToolRound( + persistedAssistantMessage, + persistedToolMessages, + conversationGeneration + ); + durability = "durable"; + } catch { + durability = await this.checkToolRoundDurability( + conversationId, + conversationGeneration, + persistedAssistantMessage, + persistedToolMessages + ); + } + } + if (durability !== "durable") { + // not_durable:确已证实未落盘,回收本轮租约态附件。 + // indeterminate(重试后仍无法确认):不能证实未落盘,可能仍被引用,不能删除, + // 只能保留租约、以 persist_indeterminate 终止而不是当作成功发布。 + if (durability === "not_durable") { + const ownedAttachmentIds = toolResults.flatMap((toolResult) => toolResult.ownedAttachmentIds || []); + await Promise.all(ownedAttachmentIds.map((id) => this.chatRepo.deleteAttachment(id).catch(() => {}))); + await this.releaseGeneratedAttachments(result); + } + throw Object.assign(error instanceof Error ? error : new Error(String(error)), { + usage: totalUsage, + durationMs: Date.now() - startTime, + conversationId, + errorCode: durability === "indeterminate" ? "persist_indeterminate" : undefined, + }); } + // durable:OPFS close 报告了二义性错误,或重试后确认已落盘,发布结果并保留附件。 } } + // Durability is the publication boundary: only committed results enter the next request or UI stream. + for (const toolResult of toolResults) { + messages.push({ role: "tool", content: toolResult.result, toolCallId: toolResult.id }); + sendEvent({ + type: "tool_call_complete", + id: toolResult.id, + result: toolResult.result, + status: toolResult.error ? "error" : "completed", + attachments: toolResult.attachments, + ownedAttachmentIds: toolResult.ownedAttachmentIds, + }); + } + + // 工具调用状态已全部回写完毕,取消可以安全终态化了:只发一条终态事件,不再进入循环检测/下一轮。 + // cancelledDuringTools 是工具执行结束时的采样;上面的事件发送/tool 消息持久化/状态回写 + // 都是 await,期间到达的 Stop 也必须在这里被观察到,否则会带着已取消的 signal 继续 + // 进入循环检测甚至下一轮 LLM 调用 + if (cancelledDuringTools || signal.aborted) { + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + // 记录工具调用历史用于模式检测 const resultMap = new Map(toolResults.map((r) => [r.id, r])); for (const tc of result.toolCalls) { @@ -239,51 +629,135 @@ export class ToolLoopOrchestrator { guardStartIndex = toolCallHistory.length; messages.push({ role: "user", content: toolCallWarning }); sendEvent({ type: "system_warning", message: toolCallWarning }); + guardStrikeCount++; + + // 连续命中达到阈值时暂停,询问用户是否继续(仅 UI 对话传入 askUserForGuard 时生效) + if (guardStrikeCount >= GUARD_ESCALATION_STRIKES && params.askUserForGuard) { + const answer = await waitForGuardAnswer(params.askUserForGuard, guardStrikeCount, signal); + if (answer === null) { + // waitForGuardAnswer 只在 signal abort 时才会返回 null,必须走取消终态化路径, + // 而不是当作正常完成发 done——否则 Stop 期间恰好卡在等待用户回答会被误报为成功 + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + if (answer.trim().toLowerCase() === GUARD_STOP_ANSWER.toLowerCase()) { + const durationMs = Date.now() - startTime; + if (conversationId) { + try { + await this.chatRepo.appendMessage( + { + id: uuidv4(), + conversationId, + role: "assistant", + content: t("agent:chat_guard_stopped_message"), + usage: totalUsage, + durationMs, + createtime: Date.now(), + }, + conversationGeneration + ); + } catch { + // 持久化失败不阻塞终态事件发送 + } + } + // 持久化期间也可能已被 Stop:取消优先于 done,不能在 Stop 之后仍报告成功 + if (signal.aborted) { + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + sendEvent({ type: "done", usage: totalUsage, durationMs }); + return; + } + // 用户选择继续:重置命中计数,避免此后每一次告警都重新弹出询问 + guardStrikeCount = 0; + } } // 通知 UI 即将开始新一轮 LLM 调用,创建新的 assistant 消息 sendEvent({ type: "new_message" }); - - // 继续循环 continue; } - // 没有 tool calls,对话结束 + // 没有 tool calls,对话结束。与取消/错误终态不同,done 对外承诺"回复已持久化"; + // UI 完成回调会重新从 OPFS 加载消息,静默吞掉写入失败仍报 done 会让回复看起来生成成功、 + // 刷新后又消失。这里有限重试几次,仍失败则改报结构化错误而不是假装成功。 const durationMs = Date.now() - startTime; + let persistFailed = false; + let persistenceIndeterminate = false; if (conversationId) { - await this.chatRepo.appendMessage({ + const assistantMessage = { id: uuidv4(), conversationId, - role: "assistant", + role: "assistant" as const, content: buildMessageContent(), + ownedAttachmentIds: result.contentBlocks + ?.filter((block) => block.type !== "text") + .map((block) => block.attachmentId), thinking: result.thinking ? { content: result.thinking } : undefined, + // 生成图片保存失败等非致命问题:持久化到消息上,刷新后仍然可见,而不是只在这次流式响应里 + // 一闪而过 + warning: result.warning, usage: totalUsage, durationMs, createtime: Date.now(), + }; + const MAX_PERSIST_ATTEMPTS = 3; + for (let attempt = 1; attempt <= MAX_PERSIST_ATTEMPTS; attempt++) { + try { + await this.chatRepo.appendMessage(assistantMessage, conversationGeneration); + persistFailed = false; + break; + } catch { + persistFailed = true; + if (attempt < MAX_PERSIST_ATTEMPTS) { + await new Promise((resolve) => setTimeout(resolve, 200 * attempt)); + } + } + } + // 重试全部失败仍不能断定未落盘:appendMessage 按 id 去重,第一次尝试可能已经真正写入, + // 只是确认读恰好也持续失败。删除生成的图片附件之前必须 positively 证实消息确实不在 + // 存储里,否则会删掉仍被这条消息引用的文件 + if (persistFailed) { + try { + const snapshot = await this.chatRepo.getMessageSnapshot(conversationId, conversationGeneration); + if (snapshot.messages.some((message) => message.id === assistantMessage.id)) persistFailed = false; + } catch { + persistenceIndeterminate = true; + } + } + } + + // 持久化期间也可能已被 stop:内容已经落库(不丢失),但终态事件必须反映取消, + // 不能在 Stop 之后仍对外报告 done(否则后台会话状态会被"成功"覆盖掉 cancelled) + if (signal.aborted) { + // 引用消息未落库(ephemeral 或 persist 失败)时,生成附件没有持久化引用,必须回收 + if (!conversationId || (persistFailed && !persistenceIndeterminate)) { + await this.releaseGeneratedAttachments(result); + } + await this.emitCancelled(conversationId, conversationGeneration, totalUsage, startTime, sendEvent); + return; + } + + if (persistFailed) { + // 只有确认读证实消息未落盘时才能回收;确认读失败属于不确定态,附件可能已被消息引用。 + if (!persistenceIndeterminate) await this.releaseGeneratedAttachments(result); + sendEvent({ + type: "error", + message: "Reply was generated but failed to save. It may be lost after reloading.", + errorCode: "persist_failed", + usage: totalUsage, + durationMs, }); + return; + } + + // 生成图片保存失败:done 之前发一次可见警告,让 UI 立即展示(消息里的 warning 字段负责刷新后仍可见) + if (result.warning) { + sendEvent({ type: "system_warning", message: result.warning }); } - // 发送 done 事件 sendEvent({ type: "done", usage: totalUsage, durationMs }); return; } - - // 超过最大迭代次数 - const maxIterMsg = `Tool calling loop exceeded maximum iterations (${maxIterations})`; - if (conversationId) { - await this.chatRepo.appendMessage({ - id: uuidv4(), - conversationId, - role: "assistant", - content: "", - error: maxIterMsg, - createtime: Date.now(), - }); - } - sendEvent({ - type: "error", - message: maxIterMsg, - errorCode: "max_iterations", - }); } } diff --git a/src/app/service/content/gm_api/cat_agent.test.ts b/src/app/service/content/gm_api/cat_agent.test.ts index 47fce7225..1ea74a921 100644 --- a/src/app/service/content/gm_api/cat_agent.test.ts +++ b/src/app/service/content/gm_api/cat_agent.test.ts @@ -20,7 +20,7 @@ function mockConnect(): MessageConnect { onMessage(cb: (msg: any) => void) { setTimeout(() => { cb({ action: "event", data: { type: "content_delta", delta: "LLM reply" } }); - cb({ action: "event", data: { type: "done", usage: { inputTokens: 10, outputTokens: 5 } } }); + cb({ action: "event", data: { type: "done", usage: { inputTokens: 10, outputTokens: 5 }, durationMs: 123 } }); }, 10); }, onDisconnect() {}, @@ -30,16 +30,33 @@ function mockConnect(): MessageConnect { return conn; } -function createInstance(commands?: Record Promise>) { +function mockConnectWithSequence(events: Array<{ delayMs: number; data: any }>): MessageConnect { + return { + onMessage(cb: (msg: any) => void) { + for (const { delayMs, data } of events) { + setTimeout(() => { + cb({ action: "event", data }); + }, delayMs); + } + }, + onDisconnect() {}, + sendMessage() {}, + disconnect() {}, + }; +} + +function createInstance( + commands?: Record Promise>, + conn: MessageConnect = mockConnect() +) { const gmSendMessage = vi.fn().mockResolvedValue(undefined); - const gmConnect = vi.fn().mockResolvedValue(mockConnect()); + const gmConnect = vi.fn().mockResolvedValue(conn); const instance = new ConversationInstance( mockConversation(), gmSendMessage, gmConnect, "test-script-uuid", - 20, undefined, // initialTools commands ); @@ -83,6 +100,7 @@ describe("ConversationInstance 命令机制", () => { expect(result.command).toBeUndefined(); expect(result.content).toBe("LLM reply"); + expect(result.durationMs).toBe(123); // 应该建立了 LLM 连接 expect(gmConnect).toHaveBeenCalled(); }); @@ -116,6 +134,15 @@ describe("ConversationInstance 命令机制", () => { expect(gmConnect).not.toHaveBeenCalled(); }); + it("chatStream 成功完成时透传 durationMs", async () => { + const { instance } = createInstance(); + const stream = await instance.chatStream("你好"); + const chunks: StreamChunk[] = []; + for await (const chunk of stream) chunks.push(chunk); + + expect(chunks.find((chunk) => chunk.type === "done")?.durationMs).toBe(123); + }); + it("脚本覆盖内置 /new 命令", async () => { const customNewHandler = vi.fn().mockResolvedValue("自定义清空逻辑"); const { instance, gmSendMessage } = createInstance({ @@ -186,7 +213,6 @@ function createEphemeralInstance(options?: { gmSendMessage, gmConnect, "test-script-uuid", - 20, options?.tools, undefined, // commands true, // ephemeral @@ -357,6 +383,324 @@ describe("ConversationInstance ephemeral 模式", () => { }); }); +// ---- tool_call_complete / new_message 事件重建测试 ---- + +// 模拟真实的一轮 tool call 往返:tool_call_start/delta → executeTools(脚本执行)→ toolResults → +// tool_call_complete(带执行结果)→ new_message(下一轮开始)→ 最终文本 → done +function mockConnectWithToolRound(toolId: string, toolName: string): MessageConnect { + let onMsgCb: (msg: any) => void = () => {}; + return { + onMessage(cb: (msg: any) => void) { + onMsgCb = cb; + setTimeout(() => { + cb({ + action: "event", + data: { type: "tool_call_start", toolCall: { id: toolId, name: toolName, arguments: "" } }, + }); + cb({ action: "event", data: { type: "tool_call_delta", id: toolId, delta: '{"a":1}' } }); + cb({ action: "executeTools", data: [{ id: toolId, name: toolName, arguments: '{"a":1}' }] }); + }, 0); + }, + onDisconnect() {}, + sendMessage(msg: any) { + if (msg.action === "toolResults") { + const result = msg.data[0]; + setTimeout(() => { + onMsgCb({ + action: "event", + data: { + type: "tool_call_complete", + id: toolId, + result: result.result, + status: result.error ? "error" : "completed", + }, + }); + onMsgCb({ action: "event", data: { type: "new_message" } }); + onMsgCb({ action: "event", data: { type: "content_delta", delta: "Final answer" } }); + onMsgCb({ + action: "event", + data: { type: "done", usage: { inputTokens: 10, outputTokens: 5 }, durationMs: 50 }, + }); + }, 0); + } + }, + disconnect() {}, + }; +} + +describe("ConversationInstance tool_call_complete / new_message 重建历史", () => { + it("chat():assistant toolCalls 应带执行结果与 completed 状态,最终回复不重复", async () => { + const handler = vi.fn().mockResolvedValue("ok"); + const gmSendMessage = vi.fn().mockResolvedValue(undefined); + const gmConnect = vi.fn().mockResolvedValue(mockConnectWithToolRound("call-1", "my_tool")); + + const instance = new ConversationInstance( + mockConversation({ modelId: "test-model" }), + gmSendMessage, + gmConnect, + "test-script-uuid", + [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} }, handler }], + undefined, + true // ephemeral + ); + + const reply = await instance.chat("使用工具"); + expect(reply.content).toBe("Final answer"); + + const messages = await instance.getMessages(); + const assistantWithTools = messages.find((m) => m.toolCalls && m.toolCalls.length > 0); + expect(assistantWithTools).toBeDefined(); + expect(assistantWithTools!.toolCalls![0]).toMatchObject({ id: "call-1", result: "ok", status: "completed" }); + + // 应有对应的 tool 角色消息 + const toolMsg = messages.find((m) => m.role === "tool" && m.toolCallId === "call-1"); + expect(toolMsg?.content).toBe("ok"); + + // 最终 assistant 内容只应出现一次,不应重复 + const finalAssistantMsgs = messages.filter((m) => m.role === "assistant" && m.content === "Final answer"); + expect(finalAssistantMsgs).toHaveLength(1); + }); + + it("chatStream():ephemeral 历史应包含带结果的 toolCalls 及对应 tool 消息", async () => { + const handler = vi.fn().mockResolvedValue("ok"); + const gmSendMessage = vi.fn().mockResolvedValue(undefined); + const gmConnect = vi.fn().mockResolvedValue(mockConnectWithToolRound("call-2", "my_tool")); + + const instance = new ConversationInstance( + mockConversation({ modelId: "test-model" }), + gmSendMessage, + gmConnect, + "test-script-uuid", + [{ name: "my_tool", description: "d", parameters: { type: "object", properties: {} }, handler }], + undefined, + true // ephemeral + ); + + const stream = await instance.chatStream("使用工具"); + const chunks: any[] = []; + for await (const chunk of stream) { + chunks.push(chunk); + } + expect(chunks.some((c) => c.type === "tool_call_complete")).toBe(true); + + const messages = await instance.getMessages(); + const assistantWithTools = messages.find((m) => m.toolCalls && m.toolCalls.length > 0); + expect(assistantWithTools).toBeDefined(); + expect(assistantWithTools!.toolCalls![0]).toMatchObject({ id: "call-2", result: "ok", status: "completed" }); + + const toolMsg = messages.find((m) => m.role === "tool" && m.toolCallId === "call-2"); + expect(toolMsg?.content).toBe("ok"); + + const finalAssistantMsgs = messages.filter((m) => m.role === "assistant" && m.content === "Final answer"); + expect(finalAssistantMsgs).toHaveLength(1); + }); + + it("chatStream() 提前 break:未完成的 toolCall 不应作为无结果的悬空协议状态记入历史", async () => { + const gmSendMessage = vi.fn().mockResolvedValue(undefined); + // 只发 tool_call_start,永远不发 tool_call_complete,模拟消费方在工具调用完成前就 break + const gmConnect = vi.fn().mockResolvedValue( + mockConnectWithEvents([ + { type: "tool_call_start", toolCall: { id: "call-early", name: "my_tool", arguments: "" } }, + { type: "tool_call_delta", id: "call-early", delta: "{}" }, + ]) + ); + + const instance = new ConversationInstance( + mockConversation({ modelId: "test-model" }), + gmSendMessage, + gmConnect, + "test-script-uuid", + [], + undefined, + true // ephemeral + ); + + const stream = await instance.chatStream("使用工具"); + for await (const chunk of stream) { + if (chunk.type === "tool_call") break; + } + + const messages = await instance.getMessages(); + const assistantWithTools = messages.find((m) => m.toolCalls && m.toolCalls.length > 0); + expect(assistantWithTools).toBeDefined(); + // 没有收到 result 的 toolCall 必须被补成终态(而不是原样带着 status:"running"、result:undefined 记入历史) + expect(assistantWithTools!.toolCalls![0].status).not.toBe("running"); + expect(assistantWithTools!.toolCalls![0].result).toBeDefined(); + + // 必须有配对的 tool 结果消息,否则重放给 provider 时协议状态不完整 + const toolMsg = messages.find((m) => m.role === "tool" && m.toolCallId === "call-early"); + expect(toolMsg).toBeDefined(); + }); +}); + +describe("executeTools:连接 settle 后不应继续执行剩余 handler", () => { + it("连接在第一个 handler 执行期间断开时,第二个 handler 不应被调用", async () => { + let disconnectCb: ((isSelfDisconnected: boolean) => void) | undefined; + + const handlerA = vi.fn().mockImplementation(async () => { + // 模拟 handlerA 执行期间用户点击 Stop / 脚本工具超时:连接断开 + disconnectCb?.(false); + return "result-a"; + }); + const handlerB = vi.fn().mockResolvedValue("result-b"); + + const conn: MessageConnect = { + onMessage(cb: (msg: any) => void) { + // 与文件中其他 mock 一致:用 setTimeout(0) 异步派发,避免手动轮询回调是否已注册 + setTimeout(() => { + cb({ + action: "executeTools", + requestId: "req-1", + data: [ + { id: "call-a", name: "tool_a", arguments: "{}" }, + { id: "call-b", name: "tool_b", arguments: "{}" }, + ], + }); + }, 0); + }, + onDisconnect(cb: (isSelfDisconnected: boolean) => void) { + disconnectCb = cb; + }, + sendMessage() {}, + disconnect() {}, + }; + + const gmSendMessage = vi.fn().mockResolvedValue(undefined); + const gmConnect = vi.fn().mockResolvedValue(conn); + + const instance = new ConversationInstance( + mockConversation({ modelId: "test-model" }), + gmSendMessage, + gmConnect, + "test-script-uuid", + [ + { name: "tool_a", description: "d", parameters: { type: "object", properties: {} }, handler: handlerA }, + { name: "tool_b", description: "d", parameters: { type: "object", properties: {} }, handler: handlerB }, + ], + undefined, + true // ephemeral + ); + + // chat() 因连接断开而 reject,这正是本测试要观察的效果 + await expect(instance.chat("使用工具")).rejects.toThrow(); + + expect(handlerA).toHaveBeenCalledOnce(); + // handlerB 不应被调用:executeTools 在 handlerA 执行期间连接已 settle, + // 后续 toolCall 直接补成取消结果,不再串行往下执行 + expect(handlerB).not.toHaveBeenCalled(); + }); +}); + +describe("executeTools:批次级取消", () => { + it("收到 cancelToolBatch 后,该批次剩余的 handler 不应再执行", async () => { + let messageCb: ((msg: any) => void) | undefined; + + const handlerA = vi.fn().mockImplementation(async () => { + // handlerA 执行期间,SW 端脚本工具批次超时,发来该批次的作废通知 + messageCb?.({ action: "cancelToolBatch", requestId: "req-timeout" }); + return "result-a"; + }); + const handlerB = vi.fn().mockResolvedValue("result-b"); + + const conn: MessageConnect = { + onMessage(cb: (msg: any) => void) { + messageCb = cb; + setTimeout(() => { + cb({ + action: "executeTools", + requestId: "req-timeout", + data: [ + { id: "call-a", name: "tool_a", arguments: "{}" }, + { id: "call-b", name: "tool_b", arguments: "{}" }, + ], + }); + // SW 端已用超时错误结果推进对话,稍后正常完成 + setTimeout(() => { + cb({ action: "event", data: { type: "done", usage: { inputTokens: 1, outputTokens: 1 } } }); + }, 5); + }, 0); + }, + onDisconnect() {}, + sendMessage() {}, + disconnect() {}, + }; + + const gmSendMessage = vi.fn().mockResolvedValue(undefined); + const gmConnect = vi.fn().mockResolvedValue(conn); + + const instance = new ConversationInstance( + mockConversation({ modelId: "test-model" }), + gmSendMessage, + gmConnect, + "test-script-uuid", + [ + { name: "tool_a", description: "d", parameters: { type: "object", properties: {} }, handler: handlerA }, + { name: "tool_b", description: "d", parameters: { type: "object", properties: {} }, handler: handlerB }, + ], + undefined, + true // ephemeral + ); + + await instance.chat("使用工具"); + + expect(handlerA).toHaveBeenCalledOnce(); + // handlerB 不应被调用:批次已被 SW 端超时作废,剩余 handler 的副作用会与下一批次交叠 + expect(handlerB).not.toHaveBeenCalled(); + }); + + it("收到 cancelToolBatch 时应中止当前 handler 的批次级 AbortSignal", async () => { + let messageCb: ((msg: any) => void) | undefined; + let handlerSignal: AbortSignal | undefined; + + const handler = vi.fn().mockImplementation(async (_args: Record, signal?: AbortSignal) => { + handlerSignal = signal; + messageCb?.({ action: "cancelToolBatch", requestId: "req-active" }); + return "late-result"; + }); + + const conn: MessageConnect = { + onMessage(cb: (msg: any) => void) { + messageCb = cb; + setTimeout(() => { + cb({ + action: "executeTools", + requestId: "req-active", + data: [{ id: "call-active", name: "tool_active", arguments: "{}" }], + }); + setTimeout(() => { + cb({ action: "event", data: { type: "done", usage: { inputTokens: 1, outputTokens: 1 } } }); + }, 5); + }, 0); + }, + onDisconnect() {}, + sendMessage() {}, + disconnect() {}, + }; + + const instance = new ConversationInstance( + mockConversation({ modelId: "test-model" }), + vi.fn().mockResolvedValue(undefined), + vi.fn().mockResolvedValue(conn), + "test-script-uuid", + [ + { + name: "tool_active", + description: "d", + parameters: { type: "object", properties: {} }, + handler, + }, + ], + undefined, + true + ); + + await instance.chat("使用工具"); + + expect(handlerSignal).toBeInstanceOf(AbortSignal); + expect(handlerSignal?.aborted).toBe(true); + }); +}); + // ---- errorCode 透传测试 ---- // 创建发送指定事件序列的 mock 连接 @@ -380,21 +724,28 @@ function mockConnectWithEvents(events: any[]): MessageConnect { describe("errorCode 透传:chat()", () => { it("error event 带 errorCode 时,reject 的 Error 应有对应 errorCode", async () => { - const errorEvent = { type: "error", message: "Rate limit exceeded", errorCode: "rate_limit" }; + const errorEvent = { + type: "error", + message: "Rate limit exceeded", + errorCode: "rate_limit", + usage: { inputTokens: 12, outputTokens: 3 }, + durationMs: 321, + }; const gmConnect = vi.fn().mockResolvedValue(mockConnectWithEvents([errorEvent])); const instance = new ConversationInstance( mockConversation(), vi.fn().mockResolvedValue(undefined), gmConnect, - "uuid", - 20 + "uuid" ); const err = await instance.chat("你好").catch((e) => e); expect(err).toBeInstanceOf(Error); expect(err.message).toBe("Rate limit exceeded"); expect((err as any).errorCode).toBe("rate_limit"); + expect((err as any).usage).toEqual({ inputTokens: 12, outputTokens: 3 }); + expect((err as any).durationMs).toBe(321); }); it("error event 无 errorCode 时,errorCode 应为 undefined", async () => { @@ -405,8 +756,7 @@ describe("errorCode 透传:chat()", () => { mockConversation(), vi.fn().mockResolvedValue(undefined), gmConnect, - "uuid", - 20 + "uuid" ); const err = await instance.chat("你好").catch((e) => e); @@ -415,7 +765,7 @@ describe("errorCode 透传:chat()", () => { }); it("各种 errorCode 值均能正确透传", async () => { - const codes = ["rate_limit", "auth", "tool_timeout", "max_iterations", "api_error"]; + const codes = ["rate_limit", "auth", "tool_timeout", "context_too_large", "api_error"]; for (const code of codes) { const gmConnect = vi @@ -426,8 +776,7 @@ describe("errorCode 透传:chat()", () => { mockConversation(), vi.fn().mockResolvedValue(undefined), gmConnect, - "uuid", - 20 + "uuid" ); const err = await instance.chat("test").catch((e) => e); @@ -442,15 +791,20 @@ describe("errorCode 透传:chatStream()", () => { // 因此不会 throw,只需检查 chunk 中的 errorCode 即可。 it("error event 带 errorCode 时,error chunk 应有对应 errorCode", async () => { - const errorEvent = { type: "error", message: "Tool timed out", errorCode: "tool_timeout" }; + const errorEvent = { + type: "error", + message: "Tool timed out", + errorCode: "tool_timeout", + usage: { inputTokens: 12, outputTokens: 3 }, + durationMs: 321, + }; const gmConnect = vi.fn().mockResolvedValue(mockConnectWithEvents([errorEvent])); const instance = new ConversationInstance( mockConversation(), vi.fn().mockResolvedValue(undefined), gmConnect, - "uuid", - 20 + "uuid" ); const stream = await instance.chatStream("你好"); @@ -463,6 +817,8 @@ describe("errorCode 透传:chatStream()", () => { expect(errorChunk).toBeDefined(); expect(errorChunk!.error).toBe("Tool timed out"); expect((errorChunk as any).errorCode).toBe("tool_timeout"); + expect((errorChunk as any).usage).toEqual({ inputTokens: 12, outputTokens: 3 }); + expect((errorChunk as any).durationMs).toBe(321); }); it("error event 无 errorCode 时,chunk.errorCode 应为 undefined", async () => { @@ -473,8 +829,7 @@ describe("errorCode 透传:chatStream()", () => { mockConversation(), vi.fn().mockResolvedValue(undefined), gmConnect, - "uuid", - 20 + "uuid" ); const stream = await instance.chatStream("你好"); @@ -489,7 +844,7 @@ describe("errorCode 透传:chatStream()", () => { }); it("各种 errorCode 均能正确在 chunk 中透传", async () => { - const codes = ["rate_limit", "auth", "tool_timeout", "max_iterations", "api_error"]; + const codes = ["rate_limit", "auth", "tool_timeout", "context_too_large", "api_error"]; for (const code of codes) { const gmConnect = vi @@ -500,8 +855,7 @@ describe("errorCode 透传:chatStream()", () => { mockConversation(), vi.fn().mockResolvedValue(undefined), gmConnect, - "uuid", - 20 + "uuid" ); const stream = await instance.chatStream("test"); @@ -515,3 +869,117 @@ describe("errorCode 透传:chatStream()", () => { } }); }); + +describe("ConversationInstance 子代理事件隔离", () => { + it("chat() 应忽略 subAgent 终态事件并等待父会话结束", async () => { + const subAgent = { agentId: "sa-1", description: "child task", subAgentType: "general" }; + const conn = mockConnectWithSequence([ + { delayMs: 0, data: { type: "content_delta", delta: "child text", subAgent } }, + { + delayMs: 1, + data: { + type: "error", + message: "child failed", + errorCode: "api_error", + subAgent, + }, + }, + { delayMs: 2, data: { type: "content_delta", delta: "parent text" } }, + { + delayMs: 3, + data: { type: "done", usage: { inputTokens: 11, outputTokens: 4 }, durationMs: 91 }, + }, + ]); + const { instance } = createInstance(undefined, conn); + + await expect(instance.chat("hello")).resolves.toMatchObject({ + content: "parent text", + durationMs: 91, + }); + }); + + it("chatStream() 应忽略 subAgent 终态事件并继续输出父会话内容", async () => { + const subAgent = { agentId: "sa-2", description: "child task", subAgentType: "general" }; + const conn = mockConnectWithSequence([ + { delayMs: 0, data: { type: "content_delta", delta: "child text", subAgent } }, + { + delayMs: 1, + data: { + type: "done", + usage: { inputTokens: 7, outputTokens: 2 }, + durationMs: 33, + subAgent, + }, + }, + { delayMs: 2, data: { type: "content_delta", delta: "parent text" } }, + { + delayMs: 3, + data: { type: "done", usage: { inputTokens: 11, outputTokens: 4 }, durationMs: 91 }, + }, + ]); + const { instance } = createInstance(undefined, conn); + + const chunks: StreamChunk[] = []; + const stream = await instance.chatStream("hello"); + for await (const chunk of stream) chunks.push(chunk); + + expect(chunks).toEqual([ + { type: "content_delta", content: "parent text" }, + { type: "done", usage: { inputTokens: 11, outputTokens: 4 }, durationMs: 91 }, + ]); + }); +}); + +describe("ConversationInstance tool_call_complete 结果净化", () => { + it("chat() 返回的 toolCalls 不应携带事件专属字段,且无 status 时默认 completed", async () => { + const conn = mockConnectWithSequence([ + { + delayMs: 0, + data: { type: "tool_call_start", toolCall: { id: "tc-1", name: "my_tool", arguments: "" } }, + }, + { delayMs: 1, data: { type: "tool_call_complete", id: "tc-1", result: "done" } }, + { delayMs: 2, data: { type: "done", usage: { inputTokens: 1, outputTokens: 1 }, durationMs: 12 } }, + ]); + const { instance } = createInstance(undefined, conn); + + const reply = await instance.chat("hello"); + expect(reply.toolCalls).toHaveLength(1); + expect(reply.toolCalls?.[0]).toMatchObject({ + id: "tc-1", + name: "my_tool", + arguments: "", + result: "done", + status: "completed", + }); + expect(reply.toolCalls?.[0]).not.toHaveProperty("type"); + expect(reply.toolCalls?.[0]).not.toHaveProperty("subAgent"); + }); + + it("chatStream() 返回的 tool_call_complete chunk 应只保留工具调用字段", async () => { + const conn = mockConnectWithSequence([ + { + delayMs: 0, + data: { type: "tool_call_start", toolCall: { id: "tc-2", name: "my_tool", arguments: "" } }, + }, + { delayMs: 1, data: { type: "tool_call_complete", id: "tc-2", result: "done" } }, + { delayMs: 2, data: { type: "done", usage: { inputTokens: 1, outputTokens: 1 }, durationMs: 12 } }, + ]); + const { instance } = createInstance(undefined, conn); + + const chunks: StreamChunk[] = []; + const stream = await instance.chatStream("hello"); + for await (const chunk of stream) chunks.push(chunk); + + const completionChunk = chunks.find((chunk) => chunk.type === "tool_call_complete"); + expect(completionChunk).toBeDefined(); + expect(completionChunk?.toolCall).toMatchObject({ + id: "tc-2", + name: "my_tool", + arguments: "", + result: "done", + status: "completed", + }); + expect(completionChunk?.toolCall).not.toHaveProperty("type"); + expect(completionChunk?.toolCall).not.toHaveProperty("subAgent"); + }); +}); diff --git a/src/app/service/content/gm_api/cat_agent.ts b/src/app/service/content/gm_api/cat_agent.ts index e0a484a51..8f63e06cd 100644 --- a/src/app/service/content/gm_api/cat_agent.ts +++ b/src/app/service/content/gm_api/cat_agent.ts @@ -19,16 +19,76 @@ import type { } from "@App/app/service/agent/core/types"; import { getTextContent } from "@App/app/service/agent/core/content_utils"; +export type ConversationStreamChunk = + | StreamChunk + | { + type: "sync"; + streamingMessage?: { + content: string; + thinking?: string; + toolCalls: ToolCall[]; + }; + pendingAskUser?: { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + }; + tasks: Array<{ + id: string; + subject: string; + status: "pending" | "in_progress" | "completed"; + description?: string; + }>; + status: "running" | "done" | "error"; + }; + +type ToolHandler = (args: Record, signal: AbortSignal) => Promise; + +function cloneToolCall(toolCall: ToolCall): ToolCall { + return { + ...toolCall, + attachments: toolCall.attachments ? [...toolCall.attachments] : undefined, + }; +} + +function stringifyToolResult(result: unknown): string { + if (typeof result === "string") return result; + const serialized = JSON.stringify(result); + return serialized === undefined ? "null" : serialized; +} + +function buildContent(text: string, blocks: ContentBlock[]): MessageContent { + if (blocks.length === 0) return text; + return [...(text ? [{ type: "text" as const, text }] : []), ...blocks]; +} + +function resolveToolCall( + ordered: ToolCall[], + byId: Map, + id: string, + index?: number +): ToolCall | undefined { + if (id && byId.has(id)) return byId.get(id); + if (index !== undefined && ordered[index]) return ordered[index]; + for (let i = ordered.length - 1; i >= 0; i--) { + if ((ordered[i].status ?? "running") === "running") return ordered[i]; + } + return undefined; +} + // 对话实例,暴露给用户脚本 // 导出供测试使用 export class ConversationInstance { - private toolHandlers: Map) => Promise> = new Map(); - private toolDefs: ToolDefinition[] = []; + public toolHandlers: Map = new Map(); + public toolDefs: ToolDefinition[] = []; private commandHandlers: Map = new Map(); - private ephemeral: boolean; + public ephemeral: boolean; private cache?: boolean; private systemPrompt?: string; - private messageHistory: Array<{ + public messageHistory: Array<{ role: MessageRole; content: MessageContent; toolCallId?: string; @@ -42,7 +102,6 @@ export class ConversationInstance { private gmSendMessage: (api: string, params: any[]) => Promise, private gmConnect: (api: string, params: any[]) => Promise, private scriptUuid: string, - private maxIterations: number, initialTools?: ConversationCreateOptions["tools"], commands?: Record, ephemeral?: boolean, @@ -104,9 +163,9 @@ export class ConversationInstance { // 通过 GM API connect 建立流式连接 const connectParams: Record = { conversationId: this.conv.id, + generation: this.conv.generation, message: content, tools: toolDefs.length > 0 ? toolDefs : undefined, - maxIterations: this.maxIterations, scriptUuid: this.scriptUuid, }; @@ -127,16 +186,9 @@ export class ConversationInstance { const reply = await this.processChat(conn, handlers); - // ephemeral 模式:收集 assistant 响应到内存历史 + // ephemeral 模式:中间轮次(带 tool calls)已在 processChat 内按 new_message 边界追加到内存历史, + // 这里只需追加不含 tool calls 的最终回复(done 事件保证到达时已无待处理的 tool calls)。 if (this.ephemeral) { - if (reply.toolCalls && reply.toolCalls.length > 0) { - this.messageHistory.push({ role: "assistant", content: reply.content, toolCalls: reply.toolCalls }); - for (const tc of reply.toolCalls) { - if (tc.result !== undefined) { - this.messageHistory.push({ role: "tool", content: tc.result, toolCallId: tc.id }); - } - } - } this.messageHistory.push({ role: "assistant", content: reply.content }); } @@ -177,9 +229,9 @@ export class ConversationInstance { const connectParams: Record = { conversationId: this.conv.id, + generation: this.conv.generation, message: content, tools: toolDefs.length > 0 ? toolDefs : undefined, - maxIterations: this.maxIterations, scriptUuid: this.scriptUuid, }; @@ -198,12 +250,14 @@ export class ConversationInstance { const conn = await this.gmConnect("CAT_agentConversationChat", [connectParams]); + // chat 连接不会收到 sync 事件(sync 快照仅由 attach 的 SW 端发出), + // 公开签名与 scriptcat.d.ts 保持一致:chatStream 只产出 StreamChunk // ephemeral 模式:包装 stream 以收集 assistant 消息到内存历史 if (this.ephemeral) { - return this.processStreamEphemeral(conn, handlers); + return this.processStreamEphemeral(conn, handlers) as AsyncIterable; } - return this.processStream(conn, handlers); + return this.processStream(conn, handlers) as AsyncIterable; } // 解析命令:"/command args" -> { name, args } @@ -228,21 +282,21 @@ export class ConversationInstance { return { content: (result || "") as string, command: true }; } - // 合并实例级别和调用级别的工具定义 - private mergeTools(callTools?: ChatOptions["tools"]) { + // 合并实例级别和调用级别的工具定义(调用级同名工具同时替换 schema 与 handler) + protected mergeTools(callTools?: ChatOptions["tools"]) { const toolDefs: ToolDefinition[] = [...this.toolDefs]; - const handlers = new Map(this.toolHandlers); - - if (callTools) { - for (const tool of callTools) { - // 调用级别的工具覆盖实例级别的同名工具 - if (!handlers.has(tool.name)) { - toolDefs.push({ name: tool.name, description: tool.description, parameters: tool.parameters }); - } - handlers.set(tool.name, tool.handler); - } + const handlers = new Map(this.toolHandlers); + for (const tool of callTools || []) { + const definition = { + name: tool.name, + description: tool.description, + parameters: tool.parameters, + }; + const index = toolDefs.findIndex((item) => item.name === tool.name); + if (index >= 0) toolDefs[index] = definition; + else toolDefs.push(definition); + handlers.set(tool.name, tool.handler); } - return { toolDefs, handlers }; } @@ -264,6 +318,7 @@ export class ConversationInstance { { action: "getMessages", conversationId: this.conv.id, + generation: this.conv.generation, scriptUuid: this.scriptUuid, } as ConversationApiRequest, ]); @@ -280,6 +335,7 @@ export class ConversationInstance { { action: "clearMessages", conversationId: this.conv.id, + generation: this.conv.generation, scriptUuid: this.scriptUuid, } as ConversationApiRequest, ]); @@ -291,44 +347,111 @@ export class ConversationInstance { { action: "save", conversationId: this.conv.id, + generation: this.conv.generation, scriptUuid: this.scriptUuid, } as ConversationApiRequest, ]); } - // 附加到后台运行中的会话,返回流式事件 - async attach(): Promise> { + // 附加到后台运行中的会话,返回流式事件(首个 chunk 为 sync 快照) + async attach(): Promise> { const conn = await this.gmConnect("CAT_agentAttachToConversation", [ - { conversationId: this.conv.id, scriptUuid: this.scriptUuid }, + { conversationId: this.conv.id, generation: this.conv.generation, scriptUuid: this.scriptUuid }, ]); return this.processStream(conn, new Map()); } // 处理非流式 chat 的响应 - private processChat( - conn: MessageConnect, - handlers: Map) => Promise> - ): Promise { + protected processChat(conn: MessageConnect, handlers: Map): Promise { return new Promise((resolve, reject) => { + // 部分实现的 disconnect() 会同步触发 onDisconnect,若在 resolve/reject 之前调用会被 + // "Connection disconnected" 抢先 reject;用 settled 标记确保只有第一次终态判定生效。 + let settled = false; let content = ""; let thinking = ""; - const toolCalls: ToolCall[] = []; - const contentBlocks: ContentBlock[] = []; - let currentToolCall: ToolCall | null = null; - let usage: { inputTokens: number; outputTokens: number } | undefined; - - conn.onMessage(async (msg: any) => { - if (msg.action === "executeTools") { - // Service Worker 请求执行 tools - const requestedToolCalls: ToolCall[] = msg.data; - const results = await this.executeTools(requestedToolCalls, handlers); - conn.sendMessage({ action: "toolResults", data: results }); - return; + let blocks: ContentBlock[] = []; + let ordered: ToolCall[] = []; + let byId = new Map(); + const aggregate: ToolCall[] = []; + let usage: ChatReply["usage"]; + let durationMs: number | undefined; + let warning: string | undefined; + + const finishRound = (record = true): MessageContent => { + const finalContent = buildContent(content, blocks); + const round = ordered.map(cloneToolCall); + aggregate.push(...round); + if (this.ephemeral && record && (content || blocks.length || round.length)) { + this.messageHistory.push({ + role: "assistant", + content: finalContent, + toolCalls: round.length ? round : undefined, + }); + for (const toolCall of round) { + if (toolCall.result !== undefined) { + this.messageHistory.push({ + role: "tool", + content: toolCall.result, + toolCallId: toolCall.id, + }); + } + } } + content = ""; + blocks = []; + ordered = []; + byId = new Map(); + return finalContent; + }; - if (msg.action !== "event") return; - const event: ChatStreamEvent = msg.data; + // SW 端脚本工具批次超时后发来的作废通知(按 requestId 关联): + // 该批次剩余 handler 不再执行,避免其副作用与下一批次交叠 + const cancelledBatches = new Set(); + const batchControllers = new Map(); + const abortBatches = () => { + for (const controller of batchControllers.values()) controller.abort(); + batchControllers.clear(); + }; + conn.onMessage(async (message: any) => { + if (message.action === "cancelToolBatch") { + if (message.requestId) { + cancelledBatches.add(message.requestId); + batchControllers.get(message.requestId)?.abort(); + } + return; + } + if (message.action === "executeTools") { + const batchId: string | undefined = message.requestId; + const controller = new AbortController(); + if (batchId) batchControllers.set(batchId, controller); + const data = await this.executeTools( + message.data, + handlers, + () => { + return settled || (batchId !== undefined && cancelledBatches.has(batchId)); + }, + controller.signal + ); + if (batchId) batchControllers.delete(batchId); + // 工具函数执行期间连接可能已经因 Stop/脚本工具超时而 settle 并断开; + // 断开后的连接 sendMessage 会抛错,且这里是异步回调,事件源不会 await/捕获它, + // 会在用户脚本上下文里变成 unhandled rejection + if (settled || (batchId !== undefined && cancelledBatches.has(batchId))) return; + try { + conn.sendMessage({ + action: "toolResults", + requestId: message.requestId, + data, + }); + } catch { + // 连接已断开,结果无处可送,安全忽略 + } + return; + } + if (message.action !== "event") return; + const event: ChatStreamEvent = message.data; + if ("subAgent" in event && event.subAgent) return; switch (event.type) { case "content_delta": content += event.delta; @@ -337,124 +460,290 @@ export class ConversationInstance { thinking += event.delta; break; case "content_block_complete": - // 收集模型生成的图片/文件/音频 blocks(data 已由 finalize 保存到 attachment 存储) - contentBlocks.push(event.block); + blocks.push(event.block); break; - case "tool_call_start": - if (currentToolCall) toolCalls.push(currentToolCall); - currentToolCall = { ...event.toolCall, arguments: event.toolCall.arguments || "" }; + case "tool_call_start": { + const toolCall: ToolCall = { + ...event.toolCall, + arguments: event.toolCall.arguments || "", + status: "running", + }; + ordered.push(toolCall); + byId.set(toolCall.id, toolCall); break; - case "tool_call_delta": - if (currentToolCall) currentToolCall.arguments += event.delta; + } + case "tool_call_delta": { + const toolCall = resolveToolCall(ordered, byId, event.id, event.index); + if (toolCall) toolCall.arguments += event.delta; break; - case "done": { - if (currentToolCall) { - toolCalls.push(currentToolCall); - currentToolCall = null; - } - if (event.usage) usage = event.usage; - // 合并文本和 content blocks 到 MessageContent - let finalContent: MessageContent = content; - if (contentBlocks.length > 0) { - const blocks: ContentBlock[] = []; - if (content) blocks.push({ type: "text", text: content }); - blocks.push(...contentBlocks); - finalContent = blocks; + } + case "tool_call_complete": { + const toolCall = resolveToolCall(ordered, byId, event.id); + if (toolCall) { + toolCall.result = event.result; + toolCall.status = event.status ?? "completed"; + toolCall.attachments = event.attachments ? [...event.attachments] : undefined; } + break; + } + case "new_message": + finishRound(); + break; + case "system_warning": + warning = warning ? `${warning}\n${event.message}` : event.message; + break; + case "done": + usage = event.usage; + durationMs = event.durationMs; + settled = true; + abortBatches(); resolve({ - content: finalContent, + content: finishRound(false), thinking: thinking || undefined, - toolCalls: toolCalls.length > 0 ? toolCalls : undefined, + toolCalls: aggregate.length ? aggregate : undefined, usage, + durationMs, + warning, }); + conn.disconnect(); break; - } case "error": - reject(Object.assign(new Error(event.message), { errorCode: event.errorCode })); + settled = true; + abortBatches(); + reject(Object.assign(new Error(event.message), event)); + conn.disconnect(); break; } }); - conn.onDisconnect(() => { + if (settled) return; + settled = true; + abortBatches(); reject(new Error("Connection disconnected")); }); }); } // 处理流式 chat 的响应 - private processStream( + protected processStream( conn: MessageConnect, - handlers: Map) => Promise> - ): AsyncIterable { - const chunks: StreamChunk[] = []; - let resolve: (() => void) | null = null; + handlers: Map + ): AsyncIterable { + const chunks: ConversationStreamChunk[] = []; + let wake: (() => void) | undefined; let done = false; - let error: Error | null = null; + let error: Error | undefined; + let surfaced = false; + let ordered: ToolCall[] = []; + let byId = new Map(); + + const reset = (toolCalls: ToolCall[] = []) => { + ordered = toolCalls.map((toolCall) => ({ + ...cloneToolCall(toolCall), + status: toolCall.status ?? "running", + })); + byId = new Map(ordered.map((toolCall) => [toolCall.id, toolCall])); + }; + const push = (chunk: ConversationStreamChunk) => { + chunks.push(chunk); + wake?.(); + }; - conn.onMessage(async (msg: any) => { - if (msg.action === "executeTools") { - const requestedToolCalls: ToolCall[] = msg.data; - const results = await this.executeTools(requestedToolCalls, handlers); - conn.sendMessage({ action: "toolResults", data: results }); + // SW 端脚本工具批次超时后发来的作废通知(按 requestId 关联): + // 该批次剩余 handler 不再执行,避免其副作用与下一批次交叠 + const cancelledBatches = new Set(); + const batchControllers = new Map(); + const abortBatches = () => { + for (const controller of batchControllers.values()) controller.abort(); + batchControllers.clear(); + }; + + conn.onMessage(async (message: any) => { + if (message.action === "cancelToolBatch") { + if (message.requestId) { + cancelledBatches.add(message.requestId); + batchControllers.get(message.requestId)?.abort(); + } return; } - - if (msg.action !== "event") return; - const event: ChatStreamEvent = msg.data; - - let chunk: StreamChunk | null = null; + if (message.action === "executeTools") { + const batchId: string | undefined = message.requestId; + const controller = new AbortController(); + if (batchId) batchControllers.set(batchId, controller); + const data = await this.executeTools( + message.data, + handlers, + () => { + return done || (batchId !== undefined && cancelledBatches.has(batchId)); + }, + controller.signal + ); + if (batchId) batchControllers.delete(batchId); + // 工具函数执行期间连接可能已经因 Stop/脚本工具超时而结束并断开; + // 断开后的连接 sendMessage 会抛错,且这里是异步回调,事件源不会 await/捕获它, + // 会在用户脚本上下文里变成 unhandled rejection + if (done || (batchId !== undefined && cancelledBatches.has(batchId))) return; + try { + conn.sendMessage({ + action: "toolResults", + requestId: message.requestId, + data, + }); + } catch { + // 连接已断开,结果无处可送,安全忽略 + } + return; + } + if (message.action !== "event") return; + const event: ChatStreamEvent = message.data; + if ("subAgent" in event && event.subAgent) return; switch (event.type) { + case "sync": + reset(event.streamingMessage?.toolCalls || []); + push({ + type: "sync", + streamingMessage: event.streamingMessage + ? { + ...event.streamingMessage, + toolCalls: ordered.map(cloneToolCall), + } + : undefined, + pendingAskUser: event.pendingAskUser ? { ...event.pendingAskUser } : undefined, + tasks: event.tasks.map((task) => ({ ...task })), + status: event.status, + }); + done = event.status !== "running"; + // 终态快照:attach 的会话已经结束,SW 侧不会再为这条连接注册 listener, + // 也就不会再有后续事件——必须在这里主动断开,否则 port 会一直挂着 + if (done) conn.disconnect(); + break; case "content_delta": - chunk = { type: "content_delta", content: event.delta }; + push({ type: "content_delta", content: event.delta }); break; case "thinking_delta": - chunk = { type: "thinking_delta", content: event.delta }; + push({ type: "thinking_delta", content: event.delta }); break; case "content_block_complete": - chunk = { type: "content_block", block: event.block }; + push({ type: "content_block", block: event.block }); break; - case "tool_call_start": - chunk = { type: "tool_call", toolCall: { ...event.toolCall, arguments: "" } }; + case "tool_call_start": { + const toolCall: ToolCall = { + ...event.toolCall, + arguments: event.toolCall.arguments || "", + status: "running", + }; + ordered.push(toolCall); + byId.set(toolCall.id, toolCall); + push({ type: "tool_call", toolCall: cloneToolCall(toolCall) }); + break; + } + case "tool_call_delta": { + const toolCall = resolveToolCall(ordered, byId, event.id, event.index); + if (toolCall) { + toolCall.arguments += event.delta; + push({ type: "tool_call", toolCall: cloneToolCall(toolCall) }); + } + break; + } + case "tool_call_complete": { + const toolCall = resolveToolCall(ordered, byId, event.id); + if (toolCall) { + toolCall.result = event.result; + toolCall.status = event.status ?? "completed"; + toolCall.attachments = event.attachments ? [...event.attachments] : undefined; + } + push({ + type: "tool_call_complete", + toolCall: + toolCall || + ({ + id: event.id, + name: "", + arguments: "", + result: event.result, + status: event.status ?? "completed", + attachments: event.attachments ? [...event.attachments] : undefined, + } as ToolCall), + }); + break; + } + case "new_message": + push({ type: "new_message" }); + reset(); + break; + case "system_warning": + push({ type: "system_warning", warning: event.message }); break; case "done": - chunk = { type: "done", usage: event.usage }; + push({ + type: "done", + usage: event.usage, + durationMs: event.durationMs, + }); done = true; + abortBatches(); + conn.disconnect(); break; case "error": - chunk = { type: "error", error: event.message, errorCode: event.errorCode }; - error = Object.assign(new Error(event.message), { errorCode: event.errorCode }); + push({ + type: "error", + error: event.message, + errorCode: event.errorCode, + usage: event.usage, + durationMs: event.durationMs, + }); + error = Object.assign(new Error(event.message), event); done = true; + abortBatches(); + conn.disconnect(); break; } - - if (chunk) { - chunks.push(chunk); - resolve?.(); - } }); - conn.onDisconnect(() => { + if (done) return; done = true; - error = error || new Error("Connection disconnected"); - resolve?.(); + abortBatches(); + error = new Error("Connection disconnected"); + wake?.(); }); + // 提前退出(for await...break、消费方 throw)时必须断开连接并唤醒挂起的 next(), + // 否则 port/listener 会一直挂在 SW 侧,直到脚本上下文销毁 + const closeEarly = () => { + if (done) return; + done = true; + abortBatches(); + try { + conn.disconnect(); + } catch { + // port 可能已断开 + } + wake?.(); + }; + return { [Symbol.asyncIterator]() { return { - async next(): Promise> { - while (chunks.length === 0 && !done) { - await new Promise((r) => { - resolve = r; - }); + async next(): Promise> { + while (!chunks.length && !done) await new Promise((resolve) => (wake = resolve)); + if (chunks.length) { + const chunk = chunks.shift()!; + if (chunk.type === "error") surfaced = true; + return { value: chunk, done: false }; } - - if (chunks.length > 0) { - return { value: chunks.shift()!, done: false }; + if (error && !surfaced) { + surfaced = true; + throw error; } - - if (error && !done) throw error; - return { value: undefined as any, done: true }; + return { value: undefined as never, done: true }; + }, + async return(value?: unknown): Promise> { + closeEarly(); + return { value: value as never, done: true }; + }, + async throw(err?: unknown): Promise> { + closeEarly(); + throw err; }, }; }, @@ -462,83 +751,130 @@ export class ConversationInstance { } // 处理 ephemeral 流式 chat 的响应(收集 assistant 消息到内存历史) - private processStreamEphemeral( + protected processStreamEphemeral( conn: MessageConnect, - handlers: Map) => Promise> - ): AsyncIterable { + handlers: Map + ): AsyncIterable { const inner = this.processStream(conn, handlers); - const messageHistory = this.messageHistory; - let content = ""; - const toolCalls: ToolCall[] = []; - + let text = ""; + let blocks: ContentBlock[] = []; + let toolCalls: ToolCall[] = []; + const finish = () => { + if (text || blocks.length || toolCalls.length) { + // 提前退出(for await...break / 消费方抛错)时,可能有 toolCall 还停在 tool_call_start/ + // delta 阶段就被 return()/throw() 打断,从未收到 tool_call_complete。这类 toolCall 没有 + // result,如果原样把它们的 assistant 消息记入历史重放给 provider,大多数 provider 会 + // 因为"assistant 消息里的 tool_call 缺少对应的 tool 结果消息"而报错。 + // 统一在这里把没有 result 的 toolCall 补成终态 cancelled,并补上配对的 tool 结果消息, + // 保证重放给 provider 的历史里 tool_call/tool_result 协议状态始终完整。 + const finalized = toolCalls.map((toolCall) => { + if (toolCall.result !== undefined) return cloneToolCall(toolCall); + return cloneToolCall({ + ...toolCall, + status: "error", + result: JSON.stringify({ error: "Tool call cancelled: stream ended before it completed" }), + }); + }); + this.messageHistory.push({ + role: "assistant", + content: buildContent(text, blocks), + toolCalls: finalized.length ? finalized : undefined, + }); + for (const toolCall of finalized) { + this.messageHistory.push({ + role: "tool", + content: toolCall.result!, + toolCallId: toolCall.id, + }); + } + } + text = ""; + blocks = []; + toolCalls = []; + }; return { [Symbol.asyncIterator]() { - const iter = inner[Symbol.asyncIterator](); + const iterator = inner[Symbol.asyncIterator](); return { - async next(): Promise> { - const result = await iter.next(); + async next() { + const result = await iterator.next(); if (result.done) { - // 流结束时,追加 assistant 消息到历史 - if (content || toolCalls.length > 0) { - messageHistory.push({ - role: "assistant", - content, - toolCalls: toolCalls.length > 0 ? [...toolCalls] : undefined, - }); - } + finish(); return result; } - const chunk = result.value; - switch (chunk.type) { - case "content_delta": - content += chunk.content || ""; - break; - case "tool_call": - if (chunk.toolCall) { - toolCalls.push(chunk.toolCall); - } - break; - } + if (chunk.type === "content_delta") text += chunk.content || ""; + else if (chunk.type === "content_block" && chunk.block) blocks.push(chunk.block); + else if ((chunk.type === "tool_call" || chunk.type === "tool_call_complete") && chunk.toolCall) { + const index = toolCalls.findIndex((toolCall) => toolCall.id === chunk.toolCall!.id); + if (index >= 0) toolCalls[index] = cloneToolCall(chunk.toolCall); + else toolCalls.push(cloneToolCall(chunk.toolCall)); + } else if (chunk.type === "new_message") finish(); return result; }, + // 转发 return()/throw() 给内层 processStream 的迭代器,否则 for await...break + // 或消费方抛错时内层不会 disconnect,port 会一直挂着。 + // 提前退出时也提交已累积的部分输出到 messageHistory,与正常完成时的行为一致, + // 避免下一轮 chat() 因为丢失这部分历史而导致上下文断裂。 + async return(value?: unknown) { + finish(); + await iterator.return?.(value); + return { value: value as never, done: true }; + }, + async throw(err?: unknown) { + finish(); + await iterator.throw?.(err); + throw err; + }, }; }, }; } - // 执行用户定义的 tool handlers - private async executeTools( + // 执行用户定义的 tool handlers。 + // isSettled:连接/请求批次是否已经 settle(Stop、脚本工具超时、连接断开)。串行执行期间在每个 + // handler 之前检查,而不是只在整批结束后检查一次——否则已经 settle 之后,剩余的 + // handler 仍会继续跑,其副作用可能和后续新批次的 handler 重叠 + protected async executeTools( toolCalls: ToolCall[], - handlers: Map) => Promise> - ): Promise> { - const results: Array<{ id: string; result: string }> = []; - - for (const tc of toolCalls) { - const handler = handlers.get(tc.name); + handlers: Map, + isSettled?: () => boolean, + signal: AbortSignal = new AbortController().signal + ): Promise> { + const results: Array<{ id: string; result: string; error?: boolean }> = []; + for (const toolCall of toolCalls) { + if (signal.aborted || isSettled?.()) { + results.push({ + id: toolCall.id, + result: JSON.stringify({ error: "Tool execution cancelled: connection already settled" }), + error: true, + }); + continue; + } + const handler = handlers.get(toolCall.name); if (!handler) { - results.push({ id: tc.id, result: JSON.stringify({ error: `Tool "${tc.name}" not found` }) }); + results.push({ + id: toolCall.id, + result: JSON.stringify({ error: `Tool "${toolCall.name}" not found` }), + error: true, + }); continue; } - try { - let args: Record = {}; - if (tc.arguments) { - args = JSON.parse(tc.arguments); - } - const result = await handler(args); - results.push({ id: tc.id, result: typeof result === "string" ? result : JSON.stringify(result) }); - } catch (e: any) { - const errorMsg = - e instanceof Error - ? e.message || e.toString() - : typeof e === "string" - ? e - : String(e) || "Tool execution failed"; - results.push({ id: tc.id, result: JSON.stringify({ error: errorMsg }) }); + const args = toolCall.arguments ? JSON.parse(toolCall.arguments) : {}; + results.push({ + id: toolCall.id, + result: stringifyToolResult(await handler(args, signal)), + }); + } catch (error) { + const message = error instanceof Error ? error.message || error.toString() : String(error); + results.push({ + id: toolCall.id, + result: JSON.stringify({ error: message }), + error: true, + }); } } - return results; } } @@ -562,7 +898,6 @@ function buildInstance( ctx.sendMessage.bind(ctx), ctx.connect.bind(ctx), ctx.scriptRes?.uuid || "", - options?.maxIterations || 20, options?.tools, options?.commands, options?.ephemeral, diff --git a/src/app/service/content/gm_api/cat_agent_task.ts b/src/app/service/content/gm_api/cat_agent_task.ts index 94b4df27e..70321ee91 100644 --- a/src/app/service/content/gm_api/cat_agent_task.ts +++ b/src/app/service/content/gm_api/cat_agent_task.ts @@ -63,18 +63,32 @@ export default class CATAgentTaskApi { >; } + // task 必须携带 get()/list() 返回的 generation/revision(乐观并发版本号), + // 否则服务端无法区分"修改的是当前这个任务"还是"ID 被删除重建后的另一个任务" @GMContext.API({ follow: "CAT.agent.task" }) public "CAT.agent.task.update"(id: string, task: Partial): Promise { const ctx = this as unknown as GMBaseContext; + if (task.generation === undefined || task.revision === undefined) { + throw new Error( + "CAT.agent.task.update: task must include the generation/revision returned by CAT.agent.task.get() or list() — spread the fetched task before applying changes." + ); + } return ctx.sendMessage("CAT_agentTask", [ - { action: "update", id, task } as AgentTaskApiRequest, + { action: "update", id, generation: task.generation, revision: task.revision, task } as AgentTaskApiRequest, ]) as Promise; } @GMContext.API({ follow: "CAT.agent.task" }) - public "CAT.agent.task.remove"(id: string): Promise { + public "CAT.agent.task.remove"(id: string, task: Pick): Promise { const ctx = this as unknown as GMBaseContext; - return ctx.sendMessage("CAT_agentTask", [{ action: "delete", id } as AgentTaskApiRequest]) as Promise; + if (task?.generation === undefined || task?.revision === undefined) { + throw new Error( + "CAT.agent.task.remove: task must include the generation/revision returned by CAT.agent.task.get() or list()." + ); + } + return ctx.sendMessage("CAT_agentTask", [ + { action: "delete", id, generation: task.generation, revision: task.revision } as AgentTaskApiRequest, + ]) as Promise; } @GMContext.API({ follow: "CAT.agent.task" }) diff --git a/src/app/service/service_worker/client.ts b/src/app/service/service_worker/client.ts index f770d7e0d..4e33062ce 100644 --- a/src/app/service/service_worker/client.ts +++ b/src/app/service/service_worker/client.ts @@ -13,7 +13,12 @@ import { type ResourceBackup } from "@App/pkg/backup/struct"; import { type ConfigBundle } from "@App/pkg/backup/config_bundle"; import { type VSCodeConnectParam } from "../offscreen/vscode-connect"; import { type ScriptInfo } from "@App/pkg/utils/scriptInstall"; -import type { AgentModelConfig, MCPApiRequest, SkillConfigField } from "@App/app/service/agent/core/types"; +import type { + AgentModelConfig, + AgentTaskApiRequest, + MCPApiRequest, + SkillConfigField, +} from "@App/app/service/agent/core/types"; import type { SearchEngineConfig } from "@App/app/service/agent/core/tools/search_config"; import type { ScriptService, @@ -511,6 +516,10 @@ export class AgentClient extends Client { return this.do("saveSearchConfig", config); } + agentTask(request: AgentTaskApiRequest): Promise { + return this.doThrow("agentTask", request); + } + // MCP API mcpApi(request: MCPApiRequest): Promise { return this.doThrow("mcpApi", request); diff --git a/src/app/service/service_worker/synchronize.ts b/src/app/service/service_worker/synchronize.ts index 4f580a690..a500d50f3 100644 --- a/src/app/service/service_worker/synchronize.ts +++ b/src/app/service/service_worker/synchronize.ts @@ -234,7 +234,7 @@ export class SynchronizeService { ...Object.entries(bundle.systemConfig || {}).map(([k, v]) => sync.set(k, v)), ...(bundle.agent?.models || []).map((m) => modelRepo.saveModel(m)), ...(bundle.agent?.mcp || []).map((m) => mcpRepo.saveServer(m)), - ...(bundle.agent?.tasks || []).map((t) => taskRepo.saveTask(t)), + ...(bundle.agent?.tasks || []).map((t) => taskRepo.importTask(t)), ]); // 仅在备份带出模型选择时覆盖(部分还原未选"AI 模型"时保留本机当前默认/摘要模型) if (bundle.agent?.defaultModelId) await modelRepo.setDefaultModelId(bundle.agent.defaultModelId); diff --git a/src/locales/de-DE/agent.json b/src/locales/de-DE/agent.json index e395c049a..2741e9bc4 100644 --- a/src/locales/de-DE/agent.json +++ b/src/locales/de-DE/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "Kopiert", "chat_regenerate": "Neu generieren", "chat_streaming": "Wird generiert...", + "chat_guard_question": "Die Loop-Guard-Warnung des Agents wurde {{count}} Mal ausgelöst. Fortfahren oder stoppen?", + "chat_guard_continue": "Fortfahren", + "chat_guard_stop": "Stoppen", + "chat_guard_stopped_message": "Auf Wunsch des Benutzers nach wiederholten Loop-Guard-Warnungen angehalten.", "chat_loading_audio": "Audio wird geladen...", "chat_starting": "Wird gestartet...", "chat_retrying": "Wiederholung läuft ({{attempt}}/{{max}})...", @@ -205,7 +209,6 @@ "tasks_run_now": "Jetzt ausführen", "tasks_history": "Verlauf", "tasks_prompt": "Prompt", - "tasks_max_iterations": "Maximale Iterationen", "tasks_notify": "Bei Abschluss benachrichtigen", "tasks_notify_desc": "Bei Abschluss der Aufgabe eine Browser-Benachrichtigung senden", "tasks_no_tasks": "Keine geplanten Aufgaben", diff --git a/src/locales/en-US/agent.json b/src/locales/en-US/agent.json index 7eac89d4f..0b0278633 100644 --- a/src/locales/en-US/agent.json +++ b/src/locales/en-US/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "Copied", "chat_regenerate": "Regenerate", "chat_streaming": "Generating...", + "chat_guard_question": "The Agent has triggered the loop-guard warning {{count}} times. Continue or stop?", + "chat_guard_continue": "Continue", + "chat_guard_stop": "Stop", + "chat_guard_stopped_message": "Stopped at the user's request after repeated loop-guard warnings.", "chat_loading_audio": "Loading audio...", "chat_starting": "Starting...", "chat_retrying": "Retrying ({{attempt}}/{{max}})...", @@ -184,7 +188,6 @@ "tasks_run_now": "Run Now", "tasks_history": "History", "tasks_prompt": "Prompt", - "tasks_max_iterations": "Max Iterations", "tasks_notify": "Notify on Complete", "tasks_notify_desc": "Send a browser notification when the task completes", "tasks_no_tasks": "No scheduled tasks", diff --git a/src/locales/ja-JP/agent.json b/src/locales/ja-JP/agent.json index 29c48cf87..165d6b0c8 100644 --- a/src/locales/ja-JP/agent.json +++ b/src/locales/ja-JP/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "コピーしました", "chat_regenerate": "再生成", "chat_streaming": "生成中...", + "chat_guard_question": "Agent のループガード警告が {{count}} 回発生しました。続行しますか?", + "chat_guard_continue": "続行", + "chat_guard_stop": "停止", + "chat_guard_stopped_message": "ループガード警告が繰り返されたため、ユーザーの要求で停止しました。", "chat_loading_audio": "音声を読み込み中...", "chat_starting": "起動中...", "chat_retrying": "再試行中 ({{attempt}}/{{max}})...", @@ -205,7 +209,6 @@ "tasks_run_now": "今すぐ実行", "tasks_history": "実行履歴", "tasks_prompt": "プロンプト", - "tasks_max_iterations": "最大反復回数", "tasks_notify": "完了通知", "tasks_notify_desc": "タスク完了時にブラウザ通知を送信します", "tasks_no_tasks": "定時タスクがありません", diff --git a/src/locales/ko-KR/agent.json b/src/locales/ko-KR/agent.json index 29ec19f96..de21efed8 100644 --- a/src/locales/ko-KR/agent.json +++ b/src/locales/ko-KR/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "복사되었습니다", "chat_regenerate": "다시 생성", "chat_streaming": "생성 중...", + "chat_guard_question": "AI 에이전트에서 반복 감지 경고가 {{count}}회 발생했습니다. 계속하거나 중지하시겠습니까?", + "chat_guard_continue": "계속", + "chat_guard_stop": "중지", + "chat_guard_stopped_message": "반복 감지 경고가 여러 번 발생한 후 사용자 요청에 따라 중지했습니다.", "chat_loading_audio": "오디오 불러오는 중...", "chat_starting": "시작하는 중...", "chat_retrying": "재시도 중 ({{attempt}}/{{max}})...", @@ -184,7 +188,6 @@ "tasks_run_now": "지금 실행", "tasks_history": "실행 기록", "tasks_prompt": "프롬프트", - "tasks_max_iterations": "최대 반복 횟수", "tasks_notify": "완료 시 알림", "tasks_notify_desc": "작업이 완료되면 브라우저 알림을 보냅니다", "tasks_no_tasks": "예약된 작업이 없습니다", diff --git a/src/locales/pt-BR/agent.json b/src/locales/pt-BR/agent.json index 7b6486de3..4c106138b 100644 --- a/src/locales/pt-BR/agent.json +++ b/src/locales/pt-BR/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "Copiado", "chat_regenerate": "Gerar novamente", "chat_streaming": "Gerando...", + "chat_guard_question": "O Agente de IA acionou o alerta de proteção contra loops {{count}} vezes. Continuar ou parar?", + "chat_guard_continue": "Continuar", + "chat_guard_stop": "Parar", + "chat_guard_stopped_message": "Interrompido a pedido do usuário após alertas repetidos de proteção contra loops.", "chat_loading_audio": "Carregando áudio...", "chat_starting": "Iniciando...", "chat_retrying": "Tentando novamente ({{attempt}}/{{max}})...", @@ -184,7 +188,6 @@ "tasks_run_now": "Executar agora", "tasks_history": "Histórico", "tasks_prompt": "Prompt", - "tasks_max_iterations": "Máximo de iterações", "tasks_notify": "Notificar ao concluir", "tasks_notify_desc": "Enviar uma notificação do navegador quando a tarefa for concluída", "tasks_no_tasks": "Nenhuma tarefa agendada", diff --git a/src/locales/ru-RU/agent.json b/src/locales/ru-RU/agent.json index d43c6e245..1f7d61047 100644 --- a/src/locales/ru-RU/agent.json +++ b/src/locales/ru-RU/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "Скопировано", "chat_regenerate": "Сгенерировать заново", "chat_streaming": "Генерация...", + "chat_guard_question": "Предупреждение защиты от циклов Agent сработало {{count}} раз. Продолжить или остановиться?", + "chat_guard_continue": "Продолжить", + "chat_guard_stop": "Остановить", + "chat_guard_stopped_message": "Остановлено по запросу пользователя после повторных предупреждений о цикле.", "chat_loading_audio": "Загрузка аудио...", "chat_starting": "Запуск...", "chat_retrying": "Повтор ({{attempt}}/{{max}})...", @@ -205,7 +209,6 @@ "tasks_run_now": "Запустить сейчас", "tasks_history": "История", "tasks_prompt": "Промпт", - "tasks_max_iterations": "Максимум итераций", "tasks_notify": "Уведомлять о завершении", "tasks_notify_desc": "Отправлять уведомление браузера после завершения задачи", "tasks_no_tasks": "Нет задач по расписанию", diff --git a/src/locales/tr-TR/agent.json b/src/locales/tr-TR/agent.json index d6798cbe6..bfb18ad36 100644 --- a/src/locales/tr-TR/agent.json +++ b/src/locales/tr-TR/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "Kopyalandı", "chat_regenerate": "Yeniden oluştur", "chat_streaming": "Oluşturuluyor...", + "chat_guard_question": "Agent döngü koruması uyarısı {{count}} kez tetiklendi. Devam edilsin mi?", + "chat_guard_continue": "Devam et", + "chat_guard_stop": "Durdur", + "chat_guard_stopped_message": "Tekrarlanan döngü koruması uyarılarından sonra kullanıcının isteğiyle durduruldu.", "chat_loading_audio": "Ses yükleniyor...", "chat_starting": "Başlatılıyor...", "chat_retrying": "Yeniden deneniyor ({{attempt}}/{{max}})...", @@ -184,7 +188,6 @@ "tasks_run_now": "Şimdi Çalıştır", "tasks_history": "Geçmiş", "tasks_prompt": "İstem", - "tasks_max_iterations": "Maksimum Yineleme", "tasks_notify": "Tamamlandığında Bildir", "tasks_notify_desc": "Görev tamamlandığında bir tarayıcı bildirimi gönder", "tasks_no_tasks": "Zamanlanmış görev yok", diff --git a/src/locales/vi-VN/agent.json b/src/locales/vi-VN/agent.json index 010adfb5e..66f68e1f4 100644 --- a/src/locales/vi-VN/agent.json +++ b/src/locales/vi-VN/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "Đã sao chép", "chat_regenerate": "Tạo lại", "chat_streaming": "Đang tạo...", + "chat_guard_question": "Agent đã kích hoạt cảnh báo vòng lặp {{count}} lần. Tiếp tục hay dừng?", + "chat_guard_continue": "Tiếp tục", + "chat_guard_stop": "Dừng", + "chat_guard_stopped_message": "Đã dừng theo yêu cầu của người dùng sau nhiều cảnh báo vòng lặp.", "chat_loading_audio": "Đang tải âm thanh...", "chat_starting": "Đang khởi động...", "chat_retrying": "Đang thử lại ({{attempt}}/{{max}})...", @@ -205,7 +209,6 @@ "tasks_run_now": "Chạy ngay", "tasks_history": "Lịch sử", "tasks_prompt": "Câu lệnh", - "tasks_max_iterations": "Số lần lặp tối đa", "tasks_notify": "Thông báo khi hoàn tất", "tasks_notify_desc": "Gửi thông báo trình duyệt khi tác vụ hoàn tất", "tasks_no_tasks": "Chưa có tác vụ định kỳ", diff --git a/src/locales/zh-CN/agent.json b/src/locales/zh-CN/agent.json index cb8a6d200..95abba3e0 100644 --- a/src/locales/zh-CN/agent.json +++ b/src/locales/zh-CN/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "已复制", "chat_regenerate": "重新生成", "chat_streaming": "生成中...", + "chat_guard_question": "Agent 已触发循环检测警告 {{count}} 次。是否继续?", + "chat_guard_continue": "继续", + "chat_guard_stop": "停止", + "chat_guard_stopped_message": "用户在多次循环检测警告后请求停止。", "chat_loading_audio": "正在加载音频...", "chat_starting": "正在启动...", "chat_retrying": "正在重试 ({{attempt}}/{{max}})...", @@ -184,7 +188,6 @@ "tasks_run_now": "立即运行", "tasks_history": "运行历史", "tasks_prompt": "提示词", - "tasks_max_iterations": "最大迭代次数", "tasks_notify": "完成通知", "tasks_notify_desc": "任务完成后发送浏览器通知", "tasks_no_tasks": "暂无定时任务", diff --git a/src/locales/zh-TW/agent.json b/src/locales/zh-TW/agent.json index cc83de016..62a562ab4 100644 --- a/src/locales/zh-TW/agent.json +++ b/src/locales/zh-TW/agent.json @@ -91,6 +91,10 @@ "chat_copy_success": "已複製", "chat_regenerate": "重新生成", "chat_streaming": "生成中...", + "chat_guard_question": "Agent 已觸發循環偵測警告 {{count}} 次。是否繼續?", + "chat_guard_continue": "繼續", + "chat_guard_stop": "停止", + "chat_guard_stopped_message": "使用者在多次循環偵測警告後要求停止。", "chat_loading_audio": "正在載入音訊...", "chat_starting": "正在啟動...", "chat_retrying": "正在重試 ({{attempt}}/{{max}})...", @@ -205,7 +209,6 @@ "tasks_run_now": "立即執行", "tasks_history": "執行歷史", "tasks_prompt": "提示詞", - "tasks_max_iterations": "最大迭代次數", "tasks_notify": "完成通知", "tasks_notify_desc": "任務完成後發送瀏覽器通知", "tasks_no_tasks": "尚無定時任務", diff --git a/src/pages/options/routes/Agent/Chat/AskUserBlock.test.tsx b/src/pages/options/routes/Agent/Chat/AskUserBlock.test.tsx index 9af62320e..7d65f509e 100644 --- a/src/pages/options/routes/Agent/Chat/AskUserBlock.test.tsx +++ b/src/pages/options/routes/Agent/Chat/AskUserBlock.test.tsx @@ -38,4 +38,23 @@ describe("用户提问块 AskUserBlock", () => { expect(screen.queryByTestId("ask-input")).toBeNull(); expect(screen.getByText("红")).toBeInTheDocument(); }); + + it("禁用自定义输入时使用稳定选项值", () => { + const onRespond = vi.fn(); + render( + + ); + + expect(screen.queryByTestId("ask-input")).toBeNull(); + fireEvent.click(screen.getByTestId("ask-option-stop")); + expect(onRespond).toHaveBeenCalledWith("guard-1", "stop"); + expect(screen.getByText("停止")).toBeInTheDocument(); + }); }); diff --git a/src/pages/options/routes/Agent/Chat/AskUserBlock.tsx b/src/pages/options/routes/Agent/Chat/AskUserBlock.tsx index a345a5554..fa4268f14 100644 --- a/src/pages/options/routes/Agent/Chat/AskUserBlock.tsx +++ b/src/pages/options/routes/Agent/Chat/AskUserBlock.tsx @@ -7,13 +7,17 @@ export default function AskUserBlock({ id, question, options, + optionValues, multiple, + allowCustom = true, onRespond, }: { id: string; question: string; options?: string[]; + optionValues?: string[]; multiple?: boolean; + allowCustom?: boolean; onRespond: (id: string, answer: string) => void; }) { const { t } = useTranslation(); @@ -58,6 +62,14 @@ export default function AskUserBlock({ const displayAnswer = (() => { if (!answer) return ""; + if (selectedOptions.length > 0) { + return selectedOptions + .map((value) => { + const index = optionValues?.indexOf(value) ?? -1; + return index >= 0 ? options?.[index] || value : value; + }) + .join(", "); + } if (multiple) { try { const arr = JSON.parse(answer); @@ -70,6 +82,7 @@ export default function AskUserBlock({ })(); const hasOptions = options && options.length > 0; + const getOptionValue = (index: number, label: string) => optionValues?.[index] ?? label; // 已提交:紧凑的完成状态 if (submitted) { @@ -104,14 +117,15 @@ export default function AskUserBlock({ {/* 选项 */} {hasOptions && (
- {options.map((opt) => { - const isSelected = selectedOptions.includes(opt); + {options.map((opt, index) => { + const value = getOptionValue(index, opt); + const isSelected = selectedOptions.includes(value); return ( -
+ {allowCustom && ( +
+ setAnswer(e.target.value)} + onKeyDown={(e) => e.key === "Enter" && handleSubmit()} + placeholder={t("agent:chat_input_placeholder")} + className="flex-1 bg-transparent border-none outline-none text-sm text-foreground placeholder:text-muted-foreground min-w-0" + /> + +
+ )} diff --git a/src/pages/options/routes/Agent/Chat/ChatArea.test.tsx b/src/pages/options/routes/Agent/Chat/ChatArea.test.tsx index fced8c37b..fa9137688 100644 --- a/src/pages/options/routes/Agent/Chat/ChatArea.test.tsx +++ b/src/pages/options/routes/Agent/Chat/ChatArea.test.tsx @@ -1,7 +1,7 @@ import { describe, it, expect, vi, beforeAll, beforeEach, afterEach } from "vitest"; -import { render, cleanup, screen, fireEvent } from "@testing-library/react"; +import { render, cleanup, fireEvent, screen, waitFor } from "@testing-library/react"; import { initTestLanguage } from "@Tests/initTestLanguage"; -import type { ChatMessage, AgentModelConfig } from "@App/app/service/agent/core/types"; +import type { ChatMessage, AgentModelConfig, MessageContent } from "@App/app/service/agent/core/types"; const navigate = vi.hoisted(() => vi.fn()); vi.mock("react-router-dom", () => ({ useNavigate: () => navigate })); @@ -14,8 +14,25 @@ const hookState = vi.hoisted(() => ({ askUserPending: null as unknown, })); +const repoMock = vi.hoisted(() => ({ + saveAttachment: vi.fn().mockResolvedValue(1), + deleteAttachment: vi.fn().mockResolvedValue(undefined), + getAttachment: vi.fn().mockResolvedValue(null), + getMessages: vi.fn().mockResolvedValue([]), + getTaskSnapshot: vi.fn().mockResolvedValue({ generation: "test-generation", revision: 0, tasks: [] }), + saveTasks: vi.fn().mockResolvedValue(undefined), +})); + +vi.mock("@App/app/repo/agent_chat", () => ({ agentChatRepo: repoMock })); + vi.mock("./hooks", () => ({ - useMessages: () => ({ messages: hookState.messages, setMessages: vi.fn(), loadMessages: vi.fn() }), + useMessages: () => ({ + messages: hookState.messages, + setMessages: (update: ChatMessage[] | ((messages: ChatMessage[]) => ChatMessage[])) => { + hookState.messages = typeof update === "function" ? update(hookState.messages) : update; + }, + loadMessages: vi.fn(), + }), useStreamingChat: () => ({ isStreaming: hookState.isStreaming, setIsStreaming: vi.fn(), @@ -35,6 +52,35 @@ vi.mock("./hooks", () => ({ clearMessages: vi.fn(() => Promise.resolve()), })); +vi.mock("./ChatInput", () => ({ + default: ({ onSend }: { onSend: (content: MessageContent, files?: Map) => Promise }) => ( + + ), +})); + +vi.mock("./MessageToolbar", () => ({ + default: ({ onDelete }: { onDelete: () => void }) => ( + + ), +})); + import ChatArea from "./ChatArea"; const model: AgentModelConfig = { @@ -57,6 +103,7 @@ const baseProps = { beforeAll(() => initTestLanguage("zh-CN")); beforeEach(() => { + vi.clearAllMocks(); hookState.messages = []; hookState.isStreaming = false; hookState.tasks = []; @@ -115,4 +162,62 @@ describe("聊天主区域 ChatArea", () => { expect(navigate).toHaveBeenCalledWith("/agent/provider"); }); }); + + it("取消排队消息时应删除尚未被持久化消息接管的附件", async () => { + hookState.isStreaming = true; + render(); + + fireEvent.click(screen.getByTestId("queue-file")); + await waitFor(() => expect(repoMock.saveAttachment).toHaveBeenCalledWith("queued.png", expect.any(File))); + fireEvent.click(await screen.findByTestId("cancel-pending-message")); + + await waitFor(() => expect(repoMock.deleteAttachment).toHaveBeenCalledWith("queued.png")); + }); + + it("切换会话时应删除上个会话尚未提交的排队附件", async () => { + hookState.isStreaming = true; + const { rerender } = render(); + fireEvent.click(screen.getByTestId("queue-file")); + await waitFor(() => expect(repoMock.saveAttachment).toHaveBeenCalledWith("queued.png", expect.any(File))); + + rerender(); + + await waitFor(() => expect(repoMock.deleteAttachment).toHaveBeenCalledWith("queued.png")); + }); + + it("排队结束后的历史读取失败时应回收尚未交接的附件", async () => { + hookState.isStreaming = true; + repoMock.getMessages.mockRejectedValueOnce(new Error("read failed")); + const { rerender } = render(); + fireEvent.click(screen.getByTestId("queue-file")); + await waitFor(() => expect(repoMock.saveAttachment).toHaveBeenCalledWith("queued.png", expect.any(File))); + + hookState.isStreaming = false; + rerender(); + + await waitFor(() => expect(repoMock.deleteAttachment).toHaveBeenCalledWith("queued.png")); + }); + + it("删除消息轮次时应同时清空已失去历史依据的会话任务", async () => { + hookState.messages = [msg({ role: "user", content: "问题" }), msg({ role: "assistant", content: "答案" })]; + render(); + + fireEvent.click(screen.getByTestId("toolbar-delete-direct")); + + await waitFor(() => expect(repoMock.saveTasks).toHaveBeenCalledWith("c1", [], undefined, "test-generation", 0)); + }); + + it("新的 ask_user 请求应重置上一个请求的已提交状态", () => { + hookState.askUserPending = { id: "question-1", question: "旧问题", options: ["旧答案"] }; + const { rerender } = render(); + + fireEvent.click(screen.getByTestId("ask-option-旧答案")); + expect(screen.queryByTestId("ask-input")).toBeNull(); + + hookState.askUserPending = { id: "question-2", question: "新问题" }; + rerender(); + + expect(screen.getByText("新问题")).toBeInTheDocument(); + expect(screen.getByTestId("ask-input")).toBeInTheDocument(); + }); }); diff --git a/src/pages/options/routes/Agent/Chat/ChatArea.tsx b/src/pages/options/routes/Agent/Chat/ChatArea.tsx index b9a8e6c78..6a7f51af4 100644 --- a/src/pages/options/routes/Agent/Chat/ChatArea.tsx +++ b/src/pages/options/routes/Agent/Chat/ChatArea.tsx @@ -36,6 +36,15 @@ function genId(): string { return Date.now().toString(36) + Math.random().toString(36).slice(2, 8); } +async function releasePendingAttachments(attachmentIds?: string[]): Promise { + await Promise.all((attachmentIds || []).map((id) => agentChatRepo.deleteAttachment(id).catch(() => {}))); +} + +function buildSubAgentContent(text: string, blocks: ContentBlock[]): MessageContent { + if (blocks.length === 0) return text; + return [...(text ? [{ type: "text" as const, text }] : []), ...blocks]; +} + // 欢迎界面 function WelcomeScreen({ hasConversation }: { hasConversation: boolean }) { const { t } = useTranslation(); @@ -105,6 +114,7 @@ function NoModelBar({ onConfigure }: { onConfigure: () => void }) { export default function ChatArea({ conversationId, + conversationGeneration, models, modelsLoaded, selectedModelId, @@ -120,6 +130,9 @@ export default function ChatArea({ onBackgroundEnabledChange, }: { conversationId: string; + // 当前会话的 generation:随每次 activeConv 变化传入,用于让下面的持久化操作在服务端做 + // 乐观并发校验,避免一个过期的 Options 标签页作用于同 ID 被删除重建后的新会话 + conversationGeneration?: string; models: AgentModelConfig[]; modelsLoaded?: boolean; selectedModelId: string; @@ -146,14 +159,18 @@ export default function ChatArea({ respondToAskUser, attachToConversation, } = useStreamingChat(); - const { tasks, setTasks, handleTaskUpdate, loadTasks } = useConversationTasks(conversationId); + const { tasks, setTasks, handleTaskUpdate, loadTasks } = useConversationTasks(conversationId, conversationGeneration); const messagesEndRef = useRef(null); const streamingMsgRef = useRef(null); const sendStartTimeRef = useRef(0); const firstTokenRecordedRef = useRef(false); const firstTokenMsRef = useRef(undefined); - const pendingMessageRef = useRef<{ content: MessageContent; messageId: string } | null>(null); + const pendingMessageRef = useRef<{ + content: MessageContent; + messageId: string; + ownedAttachmentIds?: string[]; + } | null>(null); const [pendingMessageId, setPendingMessageId] = useState(null); // 切换会话时丢弃上个会话残留的排队消息(渲染期比较上一个会话 id,避免在 effect 中同步 setState) @@ -163,9 +180,13 @@ export default function ChatArea({ setPendingMessageId(null); } - // ref 不能在渲染期写入,故清空排队消息内容放入 effect + // 会话切换或组件卸载时,尚未提交的排队附件仍是临时租约,必须回收。 useEffect(() => { - pendingMessageRef.current = null; + return () => { + const pending = pendingMessageRef.current; + pendingMessageRef.current = null; + if (pending) void releasePendingAttachments(pending.ownedAttachmentIds); + }; }, [conversationId]); const scrollToBottom = useCallback(() => { @@ -216,15 +237,17 @@ export default function ChatArea({ // 子代理事件:扁平化路由 if ("subAgent" in event && event.subAgent) { - const { agentId, description, subAgentType } = event.subAgent; + const { agentId, description, subAgentType, toolCallId } = event.subAgent; let sa = subAgentsRef.current.get(agentId); if (!sa) { sa = { agentId, description, subAgentType, + toolCallId, completedMessages: [], currentContent: "", + currentBlocks: [], currentThinking: "", currentToolCalls: [], isRunning: true, @@ -235,6 +258,10 @@ export default function ChatArea({ case "content_delta": sa.currentContent += event.delta; break; + case "content_block_complete": + sa.currentBlocks ||= []; + sa.currentBlocks.push(event.block); + break; case "thinking_delta": sa.currentThinking += event.delta; break; @@ -259,28 +286,43 @@ export default function ChatArea({ case "tool_call_complete": { const tc = sa.currentToolCalls.find((x) => x.id === event.id); if (tc) { - tc.status = "completed"; + tc.status = event.status || "completed"; tc.result = event.result; tc.attachments = event.attachments; } break; } + case "system_warning": + // 生成数据丢失等警告(如子代理内图片保存失败)需随当前轮次一起归档, + // 否则子代理气泡下刷新页面就丢失了这条提示——与父级消息 warning 字段同样的语义 + sa.currentWarning = sa.currentWarning ? `${sa.currentWarning}\n${event.message}` : event.message; + break; case "new_message": - if (sa.currentContent || sa.currentThinking || sa.currentToolCalls.length > 0) { + if ( + sa.currentContent || + sa.currentBlocks?.length || + sa.currentThinking || + sa.currentWarning || + sa.currentToolCalls.length > 0 + ) { sa.completedMessages.push({ - content: sa.currentContent, + content: buildSubAgentContent(sa.currentContent, sa.currentBlocks || []), thinking: sa.currentThinking || undefined, + warning: sa.currentWarning, toolCalls: [...sa.currentToolCalls], }); } sa.currentContent = ""; + sa.currentBlocks = []; sa.currentThinking = ""; + sa.currentWarning = undefined; sa.currentToolCalls = []; break; case "retry": sa.retryInfo = { attempt: event.attempt, maxRetries: event.maxRetries, error: event.error }; break; case "done": + case "error": if (event.usage) { if (!sa.usage) sa.usage = { inputTokens: 0, outputTokens: 0 }; sa.usage.inputTokens += event.usage.inputTokens; @@ -290,17 +332,24 @@ export default function ChatArea({ sa.usage.cacheReadInputTokens = (sa.usage.cacheReadInputTokens || 0) + (event.usage.cacheReadInputTokens || 0); } - // falls through - case "error": sa.retryInfo = undefined; - if (sa.currentContent || sa.currentThinking || sa.currentToolCalls.length > 0) { + if ( + sa.currentContent || + sa.currentBlocks?.length || + sa.currentThinking || + sa.currentWarning || + sa.currentToolCalls.length > 0 + ) { sa.completedMessages.push({ - content: sa.currentContent, + content: buildSubAgentContent(sa.currentContent, sa.currentBlocks || []), thinking: sa.currentThinking || undefined, + warning: sa.currentWarning, toolCalls: [...sa.currentToolCalls], }); sa.currentContent = ""; + sa.currentBlocks = []; sa.currentThinking = ""; + sa.currentWarning = undefined; sa.currentToolCalls = []; } sa.isRunning = false; @@ -360,7 +409,7 @@ export default function ChatArea({ case "tool_call_complete": { const tc = msg.toolCalls?.find((x) => x.id === event.id); if (tc) { - tc.status = "completed"; + tc.status = event.status || "completed"; tc.result = event.result; tc.attachments = event.attachments; } @@ -433,6 +482,9 @@ export default function ChatArea({ break; case "error": msg.error = event.message; + msg.errorCode = event.errorCode; + if (event.usage) msg.usage = event.usage; + if (event.durationMs != null) msg.durationMs = event.durationMs; break; case "done": if (event.usage) msg.usage = event.usage; @@ -465,8 +517,13 @@ export default function ChatArea({ if (!pending) return; pendingMessageRef.current = null; setPendingMessageId(null); - const freshMsgs = await agentChatRepo.getMessages(conversationId); - startStreamingRef.current(freshMsgs, pending.content); + try { + const freshMsgs = await agentChatRepo.getMessages(conversationId); + startStreamingRef.current(freshMsgs, pending.content, undefined, pending.ownedAttachmentIds); + } catch { + await releasePendingAttachments(pending.ownedAttachmentIds); + setMessages((prev) => prev.filter((message) => message.id !== pending.messageId)); + } }; const createDoneCallback = () => { @@ -486,7 +543,12 @@ export default function ChatArea({ }; }; - const startStreaming = (baseMessages: ChatMessage[], content: MessageContent, skipUserMessage?: boolean) => { + const startStreaming = ( + baseMessages: ChatMessage[], + content: MessageContent, + skipUserMessage?: boolean, + ownedAttachmentIds?: string[] + ) => { sendStartTimeRef.current = Date.now(); setStreamStartTime(sendStartTimeRef.current); firstTokenRecordedRef.current = false; @@ -499,6 +561,7 @@ export default function ChatArea({ conversationId, role: "user", content, + ownedAttachmentIds, createtime: Date.now(), }); } @@ -524,7 +587,7 @@ export default function ChatArea({ selectedModelId, skipUserMessage, enableTools, - { background: backgroundEnabled } + { background: backgroundEnabled, ownedAttachmentIds, generation: conversationGeneration } ); }; @@ -553,7 +616,7 @@ export default function ChatArea({ // assistantMsg 在此处才进入 messages 被渲染,故 streamingMsgId 镜像也在此同步设置 setStreamingMsgId(assistantMsg.id); setMessages((prev) => [...prev, assistantMsg]); - void attachToConversation(conversationId, createStreamCallback(), createDoneCallback()); + void attachToConversation(conversationId, createStreamCallback(), createDoneCallback(), conversationGeneration); }); // eslint-disable-next-line react-hooks/exhaustive-deps }, [conversationId, runningIds]); @@ -564,7 +627,7 @@ export default function ChatArea({ // /new:清空对话上下文及任务 if (typeof content === "string" && content.trim() === "/new") { if (isStreaming) return; - await clearMessages(conversationId); + await clearMessages(conversationId, conversationGeneration); setMessages([]); void loadTasks(); return; @@ -586,31 +649,46 @@ export default function ChatArea({ selectedModelId, undefined, undefined, - { compact: true, compactInstruction: instruction || undefined } + { compact: true, compactInstruction: instruction || undefined, generation: conversationGeneration } ); return; } // 保存附件到 OPFS if (files && files.size > 0) { - for (const [id, file] of files) { - await agentChatRepo.saveAttachment(id, file); + const savedIds: string[] = []; + try { + for (const [id, file] of files) { + await agentChatRepo.saveAttachment(id, file); + savedIds.push(id); + } + } catch (error) { + await releasePendingAttachments(savedIds); + throw error; } } + const ownedAttachmentIds = files && files.size > 0 ? [...files.keys()] : undefined; // LLM 运行中:排队 if (isStreaming) { const msgId = genId(); - pendingMessageRef.current = { content, messageId: msgId }; + pendingMessageRef.current = { content, messageId: msgId, ownedAttachmentIds }; setPendingMessageId(msgId); setMessages((prev) => [ ...prev, - { id: msgId, conversationId, role: "user" as const, content, createtime: Date.now() }, + { + id: msgId, + conversationId, + role: "user" as const, + content, + ownedAttachmentIds, + createtime: Date.now(), + }, ]); return; } - startStreaming(messages, content); + startStreaming(messages, content, undefined, ownedAttachmentIds); }; const handleCopy = useCallback( @@ -627,21 +705,29 @@ export default function ChatArea({ ); const clearTasks = useCallback(async () => { - await agentChatRepo.saveTasks(conversationId, []); + const snapshot = await agentChatRepo.getTaskSnapshot(conversationId, conversationGeneration); + await agentChatRepo.saveTasks(conversationId, [], undefined, snapshot.generation, snapshot.revision); setTasks([]); - }, [conversationId, setTasks]); + }, [conversationGeneration, conversationId, setTasks]); const handleRegenerate = useCallback( async (groups: MessageGroup[], groupIndex: number) => { if (isStreaming) return; const action = computeRegenerateAction(groups, groupIndex, messages); if (!action) return; - await deleteMessages(conversationId, action.idsToDelete); await clearTasks(); - setMessages(action.remainingMessages); - startStreamingRef.current(action.remainingMessages, action.userContent); + let historyDeleted = false; + try { + await deleteMessages(conversationId, action.idsToDelete, action.ownedAttachmentIds, conversationGeneration); + historyDeleted = true; + setMessages(action.remainingMessages); + startStreamingRef.current(action.remainingMessages, action.userContent, undefined, action.ownedAttachmentIds); + } catch (error) { + if (historyDeleted) await releasePendingAttachments(action.ownedAttachmentIds); + throw error; + } }, - [conversationId, isStreaming, messages, setMessages, clearTasks] + [conversationId, conversationGeneration, isStreaming, messages, setMessages, clearTasks] ); const handleRegenerateUserMessage = useCallback( @@ -650,13 +736,13 @@ export default function ChatArea({ const action = computeUserRegenerateAction(messageId, messages); if (!action) return; if (action.idsToDelete.length > 0) { - await deleteMessages(conversationId, action.idsToDelete); + await deleteMessages(conversationId, action.idsToDelete, undefined, conversationGeneration); } await clearTasks(); setMessages(action.remainingMessages); startStreamingRef.current(action.remainingMessages, action.userContent, action.skipUserMessage); }, - [conversationId, isStreaming, messages, setMessages, clearTasks] + [conversationId, conversationGeneration, isStreaming, messages, setMessages, clearTasks] ); const handleDeleteRound = useCallback( @@ -679,10 +765,11 @@ export default function ChatArea({ .map((m) => m.id); idsToDelete.push(...originalToolMsgIds); - await deleteMessages(conversationId, idsToDelete); + await clearTasks(); + await deleteMessages(conversationId, idsToDelete, undefined, conversationGeneration); void loadMessages(); }, - [conversationId, isStreaming, messages, loadMessages] + [conversationId, conversationGeneration, isStreaming, messages, loadMessages, clearTasks] ); const handleEditMessage = useCallback( @@ -690,22 +777,42 @@ export default function ChatArea({ if (isStreaming) return; const action = computeEditAction(messageId, messages); if (!action) return; - if (files && files.size > 0) { - for (const [id, file] of files) { - await agentChatRepo.saveAttachment(id, file); + const referencedAttachmentIds = new Set( + Array.isArray(content) ? content.flatMap((block) => (block.type === "text" ? [] : [block.attachmentId])) : [] + ); + const preservedAttachmentIds = action.ownedAttachmentIds.filter((id) => referencedAttachmentIds.has(id)); + await clearTasks(); + const savedIds: string[] = []; + let historyDeleted = false; + try { + if (files && files.size > 0) { + for (const [id, file] of files) { + await agentChatRepo.saveAttachment(id, file); + savedIds.push(id); + } } + await deleteMessages(conversationId, action.idsToDelete, preservedAttachmentIds, conversationGeneration); + historyDeleted = true; + setMessages(action.remainingMessages); + startStreamingRef.current(action.remainingMessages, content, undefined, [ + ...preservedAttachmentIds, + ...savedIds, + ]); + } catch (error) { + await releasePendingAttachments([...savedIds, ...(historyDeleted ? preservedAttachmentIds : [])]); + throw error; } - await deleteMessages(conversationId, action.idsToDelete); - await clearTasks(); - setMessages(action.remainingMessages); - startStreamingRef.current(action.remainingMessages, content); }, - [conversationId, isStreaming, messages, setMessages, clearTasks] + [conversationId, conversationGeneration, isStreaming, messages, setMessages, clearTasks] ); const handleStop = useCallback(async () => { clearRetryTimer(); stopGeneration(); + // 只做纯 UI 侧的乐观更新(清掉"正在流式"标记、把仍显示 running 的 toolCall 标灰); + // 不在这里处理排队消息或重新加载历史——那必须等真正的终态事件到达、取消记录落库完成后, + // 由 onDone(createDoneCallback,见 stopGeneration 里保留连接直到终态事件到达)统一处理, + // 否则排队消息可能在旧会话仍占用中时被拒绝、且从未持久化就丢失 streamingMsgRef.current = null; setStreamingMsgId(null); setMessages((prev) => { @@ -719,13 +826,7 @@ export default function ChatArea({ }; }); }); - if (pendingMessageRef.current) { - void processPendingMessage(); - } else { - void loadMessages(); - } - // eslint-disable-next-line react-hooks/exhaustive-deps - }, [clearRetryTimer, stopGeneration, setMessages, loadMessages]); + }, [clearRetryTimer, stopGeneration, setMessages]); const handleCancelPending = useCallback(() => { const pending = pendingMessageRef.current; @@ -734,6 +835,7 @@ export default function ChatArea({ pendingMessageRef.current = null; setPendingMessageId(null); setMessages((prev) => prev.filter((m) => m.id !== msgId)); + void releasePendingAttachments(pending.ownedAttachmentIds); }, [setMessages]); // 兜底:连接断开但 done 回调未触发时,处理排队消息 @@ -750,7 +852,7 @@ export default function ChatArea({ const noModel = modelsLoaded === true && models.length === 0; const showWelcome = !conversationId || (messages.length === 0 && !isStreaming); - const mergedMessages = mergeToolResults(messages); + const mergedMessages = mergeToolResults(messages, !isStreaming && !runningIds?.has(conversationId)); const messageGroups = groupMessages(mergedMessages); return ( @@ -797,10 +899,13 @@ export default function ChatArea({ )} {askUserPending && ( )} diff --git a/src/pages/options/routes/Agent/Chat/ChatInput.test.tsx b/src/pages/options/routes/Agent/Chat/ChatInput.test.tsx index 68e1be72f..b9df15963 100644 --- a/src/pages/options/routes/Agent/Chat/ChatInput.test.tsx +++ b/src/pages/options/routes/Agent/Chat/ChatInput.test.tsx @@ -1,10 +1,15 @@ import { describe, it, expect, vi, beforeAll, afterEach } from "vitest"; -import { render, cleanup, screen, fireEvent } from "@testing-library/react"; +import { render, cleanup, screen, fireEvent, waitFor } from "@testing-library/react"; import { t } from "@App/locales/locales"; import { initTestLanguage } from "@Tests/initTestLanguage"; import type { AgentModelConfig, SkillSummary } from "@App/app/service/agent/core/types"; +import { notify } from "@App/pages/components/ui/toast"; import ChatInput from "./ChatInput"; +vi.mock("@App/pages/components/ui/toast", () => ({ + notify: { error: vi.fn(), info: vi.fn() }, +})); + beforeAll(() => initTestLanguage("zh-CN")); afterEach(() => cleanup()); @@ -32,14 +37,14 @@ function setup(over?: Partial>) { } describe("聊天输入框 ChatInput", () => { - it("输入文本点击发送触发 onSend 并清空输入", () => { - const onSend = vi.fn(); + it("输入文本点击发送触发 onSend 并清空输入", async () => { + const onSend = vi.fn().mockResolvedValue(undefined); setup({ onSend }); const ta = screen.getByTestId("chat-textarea") as HTMLTextAreaElement; fireEvent.change(ta, { target: { value: " 你好 " } }); fireEvent.click(screen.getByTestId("chat-send")); expect(onSend).toHaveBeenCalledWith("你好"); - expect(ta.value).toBe(""); + await waitFor(() => expect(ta.value).toBe("")); }); it("空输入时不触发发送", () => { @@ -49,15 +54,15 @@ describe("聊天输入框 ChatInput", () => { expect(onSend).not.toHaveBeenCalled(); }); - it("Enter 发送,Shift+Enter 不发送", () => { - const onSend = vi.fn(); + it("Enter 发送,Shift+Enter 不发送", async () => { + const onSend = vi.fn().mockResolvedValue(undefined); setup({ onSend }); const ta = screen.getByTestId("chat-textarea"); fireEvent.change(ta, { target: { value: "发我" } }); fireEvent.keyDown(ta, { key: "Enter", shiftKey: true }); expect(onSend).not.toHaveBeenCalled(); fireEvent.keyDown(ta, { key: "Enter" }); - expect(onSend).toHaveBeenCalledWith("发我"); + await waitFor(() => expect(onSend).toHaveBeenCalledWith("发我")); }); it("流式中展示停止按钮并触发 onStop", () => { @@ -91,4 +96,49 @@ describe("聊天输入框 ChatInput", () => { fireEvent.mouseDown(screen.getByTestId("slash-item-search")); expect(ta.value).toBe("/search "); }); + + it("发送失败时应提示错误并保留文本与附件草稿", async () => { + const onSend = vi.fn().mockRejectedValue(new Error("storage unavailable")); + setup({ onSend }); + const ta = screen.getByTestId("chat-textarea") as HTMLTextAreaElement; + fireEvent.change(ta, { target: { value: "draft" } }); + const fileInput = document.querySelector('input[type="file"]') as HTMLInputElement; + fireEvent.change(fileInput, { + target: { files: [new File(["file"], "draft.txt", { type: "text/plain" })] }, + }); + + fireEvent.click(screen.getByTestId("chat-send")); + + await waitFor(() => expect(notify.error).toHaveBeenCalledWith(expect.stringContaining("storage unavailable"))); + expect(ta.value).toBe("draft"); + expect(screen.getByTitle("draft.txt")).toBeInTheDocument(); + }); + + it("选中图片附件后卸载组件应 revoke 对应的预览 URL,而不是卸载时的初始空数组", () => { + const createObjectURLSpy = vi.spyOn(URL, "createObjectURL").mockReturnValue("blob:mock-preview"); + const revokeObjectURLSpy = vi.spyOn(URL, "revokeObjectURL").mockImplementation(() => {}); + + const props: React.ComponentProps = { + models: [model("gpt-4o")], + selectedModelId: "gpt-4o", + onModelChange: vi.fn(), + onSend: vi.fn(), + onStop: vi.fn(), + isStreaming: false, + }; + const { unmount } = render(); + + const fileInput = document.querySelector('input[type="file"]') as HTMLInputElement; + fireEvent.change(fileInput, { + target: { files: [new File(["img"], "pic.png", { type: "image/png" })] }, + }); + expect(createObjectURLSpy).toHaveBeenCalled(); + + unmount(); + + expect(revokeObjectURLSpy).toHaveBeenCalledWith("blob:mock-preview"); + + createObjectURLSpy.mockRestore(); + revokeObjectURLSpy.mockRestore(); + }); }); diff --git a/src/pages/options/routes/Agent/Chat/ChatInput.tsx b/src/pages/options/routes/Agent/Chat/ChatInput.tsx index c136faa3e..c1b442bb9 100644 --- a/src/pages/options/routes/Agent/Chat/ChatInput.tsx +++ b/src/pages/options/routes/Agent/Chat/ChatInput.tsx @@ -212,7 +212,7 @@ export default function ChatInput({ models: AgentModelConfig[]; selectedModelId: string; onModelChange: (id: string) => void; - onSend: (content: MessageContent, files?: Map) => void; + onSend: (content: MessageContent, files?: Map) => Promise; onStop: () => void; isStreaming: boolean; disabled?: boolean; @@ -229,9 +229,12 @@ export default function ChatInput({ const [input, setInput] = useState(""); const [attachments, setAttachments] = useState([]); const [isDragging, setIsDragging] = useState(false); + const [isSending, setIsSending] = useState(false); const [slashActiveIndex, setSlashActiveIndex] = useState(0); const textareaRef = useRef(null); const fileInputRef = useRef(null); + // 卸载清理只能在 effect 里读到挂载时那次渲染捕获的 attachments(空数组),必须用 ref 跟踪最新值 + const attachmentsRef = useRef(attachments); // 斜杠命令过滤 const slashQuery = useMemo(() => { @@ -265,12 +268,17 @@ export default function ChatInput({ } }, [input]); - // 卸载时清理 objectURLs + // 保持 ref 跟随最新 attachments,供下面卸载时的 effect 读取 + useEffect(() => { + attachmentsRef.current = attachments; + }, [attachments]); + + // 卸载时清理 objectURLs:读 ref 而非闭包里的 attachments,否则空依赖数组只会捕获挂载时的 + // 初始空数组,之后选中的附件在卸载时永远不会被 revoke useEffect(() => { return () => { - attachments.forEach((a) => a.previewUrl && URL.revokeObjectURL(a.previewUrl)); + attachmentsRef.current.forEach((a) => a.previewUrl && URL.revokeObjectURL(a.previewUrl)); }; - // eslint-disable-next-line react-hooks/exhaustive-deps }, []); const addFiles = useCallback((files: File[]) => { @@ -294,31 +302,47 @@ export default function ChatInput({ }); }, []); - const handleSend = () => { + const handleSend = async () => { const trimmed = input.trim(); - if ((!trimmed && attachments.length === 0) || disabled || hasPendingMessage) return; - - if (attachments.length > 0) { - const blocks: ContentBlock[] = []; - const files = new Map(); - if (trimmed) blocks.push({ type: "text", text: trimmed }); - for (const att of attachments) { - const mime = att.file.type; - if (mime.startsWith("image/")) { - blocks.push({ type: "image", attachmentId: att.id, mimeType: mime, name: att.file.name }); - } else if (mime.startsWith("audio/")) { - blocks.push({ type: "audio", attachmentId: att.id, mimeType: mime, name: att.file.name }); - } else { - blocks.push({ type: "file", attachmentId: att.id, mimeType: mime, name: att.file.name, size: att.file.size }); + if ((!trimmed && attachments.length === 0) || disabled || hasPendingMessage || isSending) return; + + setIsSending(true); + try { + if (attachments.length > 0) { + const blocks: ContentBlock[] = []; + const files = new Map(); + if (trimmed) blocks.push({ type: "text", text: trimmed }); + for (const att of attachments) { + const mime = att.file.type; + if (mime.startsWith("image/")) { + blocks.push({ type: "image", attachmentId: att.id, mimeType: mime, name: att.file.name }); + } else if (mime.startsWith("audio/")) { + blocks.push({ type: "audio", attachmentId: att.id, mimeType: mime, name: att.file.name }); + } else { + blocks.push({ + type: "file", + attachmentId: att.id, + mimeType: mime, + name: att.file.name, + size: att.file.size, + }); + } + files.set(att.id, att.file); + } + await onSend(blocks, files); + for (const attachment of attachments) { + if (attachment.previewUrl) URL.revokeObjectURL(attachment.previewUrl); } - files.set(att.id, att.file); + setAttachments([]); + } else { + await onSend(trimmed); } - onSend(blocks, files); - setAttachments([]); - } else { - onSend(trimmed); + setInput(""); + } catch (error) { + notify.error(`${t("common:error")}: ${error instanceof Error ? error.message : String(error)}`); + } finally { + setIsSending(false); } - setInput(""); }; const handleSlashSelect = useCallback((skill: SkillSummary) => { @@ -354,7 +378,7 @@ export default function ChatInput({ if (e.key === "Enter" && !e.shiftKey) { e.preventDefault(); - handleSend(); + void handleSend(); } }; @@ -385,7 +409,7 @@ export default function ChatInput({ e.target.value = ""; }; - const canSend = !!(input.trim() || attachments.length > 0) && !disabled && !hasPendingMessage; + const canSend = !!(input.trim() || attachments.length > 0) && !disabled && !hasPendingMessage && !isSending; const iconBtn = "size-7 max-md:size-11 rounded flex items-center justify-center bg-transparent border-none cursor-pointer text-muted-foreground hover:text-foreground hover:bg-accent transition-colors"; @@ -543,7 +567,7 @@ export default function ChatInput({ type="button" data-testid="chat-send" aria-label={t("agent:chat_send")} - onClick={handleSend} + onClick={() => void handleSend()} disabled={!canSend} className={cn( "flex size-8 items-center justify-center rounded-full border-none shadow-sm transition-opacity max-md:size-11", diff --git a/src/pages/options/routes/Agent/Chat/MessageItem.test.tsx b/src/pages/options/routes/Agent/Chat/MessageItem.test.tsx index 9b8a131e6..ff0346a84 100644 --- a/src/pages/options/routes/Agent/Chat/MessageItem.test.tsx +++ b/src/pages/options/routes/Agent/Chat/MessageItem.test.tsx @@ -1,5 +1,5 @@ import { describe, it, expect, vi, beforeAll, afterEach } from "vitest"; -import { render, cleanup, screen, fireEvent } from "@testing-library/react"; +import { render, cleanup, screen, fireEvent, waitFor } from "@testing-library/react"; import { initTestLanguage } from "@Tests/initTestLanguage"; import type { ChatMessage } from "@App/app/service/agent/core/types"; import type { SubAgentState } from "./types"; @@ -34,6 +34,51 @@ describe("用户消息 UserMessageItem", () => { fireEvent.click(screen.getByTestId("user-regenerate")); expect(onRegenerate).toHaveBeenCalledOnce(); }); + + it("编辑中添加图片附件并保存成功后应 revoke 预览 URL 并清空待处理附件", async () => { + const createObjectURLSpy = vi.spyOn(URL, "createObjectURL").mockReturnValue("blob:mock-edit-preview"); + const revokeObjectURLSpy = vi.spyOn(URL, "revokeObjectURL").mockImplementation(() => {}); + const onEdit = vi.fn().mockResolvedValue(undefined); + + render(); + fireEvent.click(screen.getByTestId("user-edit")); + + const fileInput = document.querySelector('input[type="file"]') as HTMLInputElement; + fireEvent.change(fileInput, { + target: { files: [new File(["img"], "pic.png", { type: "image/png" })] }, + }); + expect(createObjectURLSpy).toHaveBeenCalled(); + expect(revokeObjectURLSpy).not.toHaveBeenCalled(); + + fireEvent.click(screen.getByTestId("user-edit-save")); + expect(onEdit).toHaveBeenCalled(); + + await waitFor(() => expect(revokeObjectURLSpy).toHaveBeenCalledWith("blob:mock-edit-preview")); + + createObjectURLSpy.mockRestore(); + revokeObjectURLSpy.mockRestore(); + }); + + it("编辑中选中图片附件后卸载组件应 revoke 对应的预览 URL", () => { + const createObjectURLSpy = vi.spyOn(URL, "createObjectURL").mockReturnValue("blob:mock-unmount-preview"); + const revokeObjectURLSpy = vi.spyOn(URL, "revokeObjectURL").mockImplementation(() => {}); + + const { unmount } = render(); + fireEvent.click(screen.getByTestId("user-edit")); + + const fileInput = document.querySelector('input[type="file"]') as HTMLInputElement; + fireEvent.change(fileInput, { + target: { files: [new File(["img"], "pic.png", { type: "image/png" })] }, + }); + expect(createObjectURLSpy).toHaveBeenCalled(); + + unmount(); + + expect(revokeObjectURLSpy).toHaveBeenCalledWith("blob:mock-unmount-preview"); + + createObjectURLSpy.mockRestore(); + revokeObjectURLSpy.mockRestore(); + }); }); describe("助手消息组 AssistantMessageGroup", () => { diff --git a/src/pages/options/routes/Agent/Chat/MessageItem.tsx b/src/pages/options/routes/Agent/Chat/MessageItem.tsx index c0544c9aa..40b6c44dc 100644 --- a/src/pages/options/routes/Agent/Chat/MessageItem.tsx +++ b/src/pages/options/routes/Agent/Chat/MessageItem.tsx @@ -73,7 +73,9 @@ function AssistantMessageContent({ {message.error && (
- {message.error} +
+ {message.error} +
)} @@ -159,7 +161,7 @@ export function UserMessageItem({ onCancel, }: { message: ChatMessage; - onEdit?: (content: MessageContent, files?: Map) => void; + onEdit?: (content: MessageContent, files?: Map) => void | Promise; onRegenerate?: () => void; isStreaming?: boolean; onCancel?: () => void; @@ -171,6 +173,19 @@ export function UserMessageItem({ const [pendingAttachments, setPendingAttachments] = useState([]); const textareaRef = useRef(null); const fileInputRef = useRef(null); + // 卸载清理只能在 effect 里读到挂载时那次渲染捕获的 pendingAttachments(空数组),必须用 ref 跟踪最新值 + const pendingAttachmentsRef = useRef(pendingAttachments); + + useEffect(() => { + pendingAttachmentsRef.current = pendingAttachments; + }, [pendingAttachments]); + + // 卸载时清理未保存的编辑态预览 objectURLs:读 ref 而非闭包里的 pendingAttachments + useEffect(() => { + return () => { + pendingAttachmentsRef.current.forEach((a) => a.previewUrl && URL.revokeObjectURL(a.previewUrl)); + }; + }, []); useEffect(() => { if (editing && textareaRef.current) { @@ -197,14 +212,14 @@ export function UserMessageItem({ setEditBlocks([]); }; - const handleSave = () => { + const handleSave = async () => { const trimmed = editContent.trim(); const hasAttachments = editBlocks.length > 0 || pendingAttachments.length > 0; if (!trimmed && !hasAttachments) return; setEditing(false); if (!hasAttachments) { - onEdit?.(trimmed); + void onEdit?.(trimmed); return; } @@ -225,7 +240,11 @@ export function UserMessageItem({ files.set(att.id, att.file); } - onEdit?.(blocks, files.size > 0 ? files : undefined); + // 附件已随 blocks/files 交给上层保存,交接完成后再 revoke 预览 URL,避免过早释放仍在使用的 Blob + await onEdit?.(blocks, files.size > 0 ? files : undefined); + pendingAttachmentsRef.current.forEach((a) => a.previewUrl && URL.revokeObjectURL(a.previewUrl)); + setPendingAttachments([]); + setEditBlocks([]); }; const addFiles = useCallback((files: File[]) => { @@ -350,7 +369,7 @@ export function UserMessageItem({ if (e.key === "Escape") handleCancel(); if (e.key === "Enter" && !e.shiftKey) { e.preventDefault(); - handleSave(); + void handleSave(); } }} onPaste={handleEditPaste} @@ -404,6 +423,7 @@ export function UserMessageItem({ className={iconBtn} title={t("agent:chat_cancel_message")} aria-label={t("agent:chat_cancel_message")} + data-testid="cancel-pending-message" onClick={onCancel} > diff --git a/src/pages/options/routes/Agent/Chat/SubAgentBlock.tsx b/src/pages/options/routes/Agent/Chat/SubAgentBlock.tsx index d92412f13..0dddb7c65 100644 --- a/src/pages/options/routes/Agent/Chat/SubAgentBlock.tsx +++ b/src/pages/options/routes/Agent/Chat/SubAgentBlock.tsx @@ -1,7 +1,7 @@ import { useState } from "react"; import { useTranslation } from "react-i18next"; import { AlertCircle, Check, ChevronDown, Loader2 } from "lucide-react"; -import type { SubAgentMessage } from "@App/app/service/agent/core/types"; +import type { ContentBlock, SubAgentMessage } from "@App/app/service/agent/core/types"; import { cn } from "@App/pkg/utils/cn"; import type { SubAgentState } from "./types"; import ToolCallBlock from "./ToolCallBlock"; @@ -27,10 +27,23 @@ export default function SubAgentBlock({ state }: { state: SubAgentState }) { // 合并所有消息(已完成 + 当前) const allMessages: SubAgentMessage[] = [...state.completedMessages]; - if (state.currentContent || state.currentThinking || state.currentToolCalls.length > 0) { + if ( + state.currentContent || + state.currentBlocks?.length || + state.currentThinking || + state.currentWarning || + state.currentToolCalls.length > 0 + ) { + const content: SubAgentMessage["content"] = state.currentBlocks?.length + ? [ + ...(state.currentContent ? [{ type: "text" as const, text: state.currentContent }] : []), + ...(state.currentBlocks as ContentBlock[]), + ] + : state.currentContent; allMessages.push({ - content: state.currentContent, + content, thinking: state.currentThinking, + warning: state.currentWarning, toolCalls: state.currentToolCalls, }); } @@ -99,6 +112,12 @@ export default function SubAgentBlock({ state }: { state: SubAgentState }) { {msg.toolCalls.map((tc) => ( ))} + {msg.warning && ( +
+ + {msg.warning} +
+ )} ))} diff --git a/src/pages/options/routes/Agent/Chat/chat_utils.test.ts b/src/pages/options/routes/Agent/Chat/chat_utils.test.ts index 690cb2a62..27b170f3a 100644 --- a/src/pages/options/routes/Agent/Chat/chat_utils.test.ts +++ b/src/pages/options/routes/Agent/Chat/chat_utils.test.ts @@ -87,6 +87,25 @@ describe("mergeToolResults", () => { expect(result[1].toolCalls?.[0].status).toBe("running"); }); + it("会话已确认不活跃时应把缺失结果的历史工具修复为 error", () => { + const messages: ChatMessage[] = [ + makeMsg({ id: "u1", role: "user", content: "hello" }), + makeMsg({ + id: "a1", + role: "assistant", + content: "", + toolCalls: [{ id: "tc1", name: "test", arguments: "{}", status: "running" }], + }), + ]; + + const result = mergeToolResults(messages, true); + + expect(result[1].toolCalls?.[0]).toMatchObject({ + status: "error", + result: expect.stringContaining("unavailable after recovery"), + }); + }); + it("有 tool 结果消息时,error 状态不被覆盖为 completed", () => { const messages: ChatMessage[] = [ makeMsg({ id: "u1", role: "user", content: "hello" }), diff --git a/src/pages/options/routes/Agent/Chat/chat_utils.ts b/src/pages/options/routes/Agent/Chat/chat_utils.ts index 8ad8c7b06..358a87d52 100644 --- a/src/pages/options/routes/Agent/Chat/chat_utils.ts +++ b/src/pages/options/routes/Agent/Chat/chat_utils.ts @@ -5,7 +5,7 @@ import type { SubAgentState } from "./types"; export type MessageGroup = { type: "user"; message: ChatMessage } | { type: "assistant"; messages: ChatMessage[] }; /** 将 tool 角色消息的结果合并进 assistant 的 toolCalls,并过滤掉 tool/system 消息 */ -export function mergeToolResults(messages: ChatMessage[]): ChatMessage[] { +export function mergeToolResults(messages: ChatMessage[], repairInactiveMissingResults = false): ChatMessage[] { const toolResultMap = new Map(); for (const msg of messages) { if (msg.role === "tool" && msg.toolCallId) { @@ -17,7 +17,7 @@ export function mergeToolResults(messages: ChatMessage[]): ChatMessage[] { return messages .filter((msg) => msg.role === "user" || msg.role === "assistant") .map((msg) => { - if (msg.role === "assistant" && msg.toolCalls && toolResultMap.size > 0) { + if (msg.role === "assistant" && msg.toolCalls && (toolResultMap.size > 0 || repairInactiveMissingResults)) { const updatedToolCalls = msg.toolCalls.map((tc) => { const result = toolResultMap.get(tc.id); if (result !== undefined) { @@ -27,6 +27,13 @@ export function mergeToolResults(messages: ChatMessage[]): ChatMessage[] { const status = !tc.status || tc.status === "running" ? "completed" : tc.status; return { ...tc, result, status }; } + if (repairInactiveMissingResults && (!tc.status || tc.status === "pending" || tc.status === "running")) { + return { + ...tc, + result: JSON.stringify({ error: "Tool result unavailable after recovery" }), + status: "error" as const, + }; + } return tc; }); return { ...msg, toolCalls: updatedToolCalls }; @@ -61,7 +68,12 @@ export function computeRegenerateAction( groups: MessageGroup[], assistantGroupIndex: number, allMessages: ChatMessage[] -): { idsToDelete: string[]; remainingMessages: ChatMessage[]; userContent: MessageContent } | null { +): { + idsToDelete: string[]; + remainingMessages: ChatMessage[]; + userContent: MessageContent; + ownedAttachmentIds: string[]; +} | null { const group = groups[assistantGroupIndex]; if (!group || group.type !== "assistant") return null; @@ -92,21 +104,26 @@ export function computeRegenerateAction( const idSet = new Set(idsToDelete); const remainingMessages = allMessages.filter((m) => !idSet.has(m.id)); - return { idsToDelete, remainingMessages, userContent: userMessage.content }; + return { + idsToDelete, + remainingMessages, + userContent: userMessage.content, + ownedAttachmentIds: userMessage.ownedAttachmentIds || [], + }; } /** 计算「编辑用户消息」需要删除的消息(该消息及其后全部)与保留的消息 */ export function computeEditAction( messageId: string, allMessages: ChatMessage[] -): { idsToDelete: string[]; remainingMessages: ChatMessage[] } | null { +): { idsToDelete: string[]; remainingMessages: ChatMessage[]; ownedAttachmentIds: string[] } | null { const idx = allMessages.findIndex((m) => m.id === messageId); if (idx < 0) return null; const idsToDelete = allMessages.slice(idx).map((m) => m.id); const remainingMessages = allMessages.slice(0, idx); - return { idsToDelete, remainingMessages }; + return { idsToDelete, remainingMessages, ownedAttachmentIds: allMessages[idx].ownedAttachmentIds || [] }; } /** 用户消息位于 groups 的 userGroupIndex,返回紧跟其后的 assistant 组索引 */ @@ -142,12 +159,19 @@ export function computeUserRegenerateAction( /** 匹配 agent 工具调用对应的子代理状态(流式 map 优先,回退到持久化 subAgentDetails) */ export function getSubAgentForToolCall( - tc: { name: string; result?: string; arguments?: string; subAgentDetails?: SubAgentDetails }, + tc: { id?: string; name: string; result?: string; arguments?: string; subAgentDetails?: SubAgentDetails }, subAgents?: Map ): SubAgentState | undefined { if (tc.name !== "agent") return undefined; if (subAgents) { + // 0. 显式 toolCallId -> agentId 匹配:并发 agent 调用各自持有独立 toolCallId, + // 一旦子代理事件带上了它就能唯一定位,不再依赖"猜第一个运行中的"这类推断。 + if (tc.id) { + for (const sa of subAgents.values()) { + if (sa.toolCallId === tc.id) return sa; + } + } // 1a. 从已完成结果匹配(格式: "[agentId: xxx]\n\n...") if (tc.result) { const match = tc.result.match(/^\[agentId: ([^\]]+)\]/); @@ -165,15 +189,22 @@ export function getSubAgentForToolCall( // 参数可能仍在流式构建中 } } - // 1c. 无结果时:优先运行中的子代理,回退到已完成的 - // (覆盖 sub-agent done → tool_call_complete 之间的间隙) - if (!tc.result) { + // 1c. 无结果、且这次调用未携带 toolCallId 时(旧数据/尚未收到子代理事件)的兜底: + // 仅当至多一个候选子代理时才可信;一旦并发数 > 1,宁可不匹配也不能张冠李戴。 + if (!tc.result && !tc.id) { let completed: SubAgentState | undefined; + let runningCount = 0; + let running: SubAgentState | undefined; for (const sa of subAgents.values()) { - if (sa.isRunning) return sa; - if (!completed) completed = sa; + if (sa.isRunning) { + runningCount++; + running = sa; + } else if (!completed) { + completed = sa; + } } - if (completed) return completed; + if (runningCount === 1 && running) return running; + if (runningCount === 0 && completed) return completed; } } @@ -186,6 +217,7 @@ export function getSubAgentForToolCall( subAgentType: d.subAgentType, completedMessages: d.messages, currentContent: "", + currentBlocks: [], currentThinking: "", currentToolCalls: [], isRunning: false, diff --git a/src/pages/options/routes/Agent/Chat/export_utils.ts b/src/pages/options/routes/Agent/Chat/export_utils.ts index 9f0f32306..0e9fc40d1 100644 --- a/src/pages/options/routes/Agent/Chat/export_utils.ts +++ b/src/pages/options/routes/Agent/Chat/export_utils.ts @@ -1,7 +1,17 @@ import type { Conversation, ChatMessage, ToolCall, SubAgentDetails } from "@App/app/service/agent/core/types"; -import { getTextContent } from "@App/app/service/agent/core/content_utils"; import { mergeToolResults } from "./chat_utils"; +function renderContent(content: ChatMessage["content"]): string { + if (typeof content === "string") return content; + return content + .map((block) => + block.type === "text" + ? block.text + : `[${block.type}: uploads/${block.attachmentId}${block.name ? ` (${block.name})` : ""}]` + ) + .join("\n"); +} + /** 格式化时间戳 */ function formatDate(ts: number): string { return new Date(ts).toLocaleString(); @@ -53,7 +63,7 @@ function renderSubAgent(details: SubAgentDetails, indent = ""): string { } if (msg.content) { lines.push(""); - lines.push(`${indent}${msg.content}`); + lines.push(`${indent}${renderContent(msg.content)}`); } for (const tc of msg.toolCalls) { lines.push(""); @@ -89,7 +99,7 @@ function renderMessageContent(msg: ChatMessage): string { } } - const text = getTextContent(msg.content); + const text = renderContent(msg.content); if (text) { parts.push(text); } @@ -120,7 +130,7 @@ export function exportToMarkdown(conversation: Conversation, messages: ChatMessa if (msg.role === "system") { lines.push("### 🔧 System"); lines.push(""); - lines.push(getTextContent(msg.content)); + lines.push(renderContent(msg.content)); lines.push(""); lines.push("---"); lines.push(""); diff --git a/src/pages/options/routes/Agent/Chat/hooks.test.ts b/src/pages/options/routes/Agent/Chat/hooks.test.ts index af010b877..1eb89e22d 100644 --- a/src/pages/options/routes/Agent/Chat/hooks.test.ts +++ b/src/pages/options/routes/Agent/Chat/hooks.test.ts @@ -6,12 +6,16 @@ import type { Conversation, ChatMessage } from "@App/app/service/agent/core/type // 会话仓库整体打桩:hooks 只是 repo + 消息总线的薄封装,测试聚焦其状态迁移逻辑。 const repoMock = vi.hoisted(() => ({ listConversations: vi.fn<() => Promise>(() => Promise.resolve([])), + createConversation: vi.fn<(c: Conversation) => Promise>((conversation) => + Promise.resolve({ ...conversation, generation: "gen-new", revision: 1 }) + ), saveConversation: vi.fn<(c: Conversation) => Promise>(() => Promise.resolve()), deleteConversation: vi.fn<(id: string) => Promise>(() => Promise.resolve()), getMessages: vi.fn<(id: string) => Promise>(() => Promise.resolve([])), saveMessages: vi.fn<(id: string, m: ChatMessage[]) => Promise>(() => Promise.resolve()), saveTasks: vi.fn<(id: string, t: unknown[]) => Promise>(() => Promise.resolve()), - getTasks: vi.fn<(id: string) => Promise>(() => Promise.resolve([])), + getTasks: vi.fn<(id: string, generation?: string) => Promise>(() => Promise.resolve([])), + deleteAttachment: vi.fn<(id: string) => Promise>(() => Promise.resolve()), })); vi.mock("@App/app/repo/agent_chat", () => ({ agentChatRepo: repoMock })); @@ -23,12 +27,123 @@ vi.mock("@App/app/repo/skill_repo", () => ({ }, })); vi.mock("@App/pages/store/global", () => ({ message: {} })); -vi.mock("@Packages/message/client", () => ({ connect: vi.fn(), sendMessage: vi.fn(() => Promise.resolve([])) })); -import { useConversations, deleteMessages, clearMessages } from "./hooks"; +// 可控 mock 连接:测试直接驱动 onMessage 回调,模拟"stop 之后终态事件才到达"的真实时序 +function createMockConn() { + let messageHandler: ((msg: any) => void) | null = null; + let disconnectHandler: (() => void) | null = null; + return { + conn: { + sendMessage: vi.fn(), + onMessage: (cb: (msg: any) => void) => { + messageHandler = cb; + }, + onDisconnect: (cb: () => void) => { + disconnectHandler = cb; + }, + disconnect: vi.fn(), + }, + emit: (data: any) => messageHandler?.({ data }), + fireDisconnect: () => disconnectHandler?.(), + }; +} + +const mockConnect = vi.hoisted(() => vi.fn()); +const mockSendMessage = vi.hoisted(() => vi.fn(() => Promise.resolve([]))); +vi.mock("@Packages/message/client", () => ({ + connect: mockConnect, + sendMessage: mockSendMessage, +})); + +import { useConversations, deleteMessages, clearMessages, useConversationTasks, useStreamingChat } from "./hooks"; + +describe("useConversationTasks:会话代际隔离", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("加载任务时应透传 conversation generation", async () => { + const { result } = renderHook(() => useConversationTasks("c1", "gen-c1")); + + await waitFor(() => expect(repoMock.getTasks).toHaveBeenCalledWith("c1", "gen-c1")); + expect(result.current.tasks).toEqual([]); + }); +}); + +describe("useStreamingChat:stop 后仍需放行终态事件", () => { + it("应把本次 UI 新上传附件的所有权发送给 Service Worker", async () => { + const { conn } = createMockConn(); + mockConnect.mockResolvedValue(conn); + const { result } = renderHook(() => useStreamingChat()); + + await act(async () => { + await result.current.sendMessage("conv-1", "hi", vi.fn(), vi.fn(), undefined, undefined, undefined, { + ownedAttachmentIds: ["upload.png"], + }); + }); + + expect(mockConnect).toHaveBeenCalledWith({}, "serviceWorker/agent/conversationChat", { + conversationId: "conv-1", + message: "hi", + modelId: undefined, + skipSaveUserMessage: undefined, + enableTools: undefined, + ownedAttachmentIds: ["upload.png"], + }); + }); + + it("连接建立失败时应回收尚未被 Service Worker 接管的上传附件", async () => { + mockConnect.mockRejectedValueOnce(new Error("connect failed")); + const { result } = renderHook(() => useStreamingChat()); + + await act(async () => { + await result.current.sendMessage("conv-1", "hi", vi.fn(), vi.fn(), undefined, undefined, undefined, { + ownedAttachmentIds: ["upload.png"], + }); + }); + + expect(repoMock.deleteAttachment).toHaveBeenCalledWith("upload.png"); + }); + + it("stopGeneration 之后到达的终态事件仍应触发 onDone 并断开连接,而不是被 abortedRef 吞掉", async () => { + const { conn, emit } = createMockConn(); + mockConnect.mockResolvedValue(conn); + + const { result } = renderHook(() => useStreamingChat()); + const onEvent = vi.fn(); + const onDone = vi.fn(); + + await act(async () => { + await result.current.sendMessage("conv-1", "hi", onEvent, onDone); + }); + + // 用户点击 Stop:只发 stop 消息、置 abortedRef,不立即断开 + act(() => { + result.current.stopGeneration(); + }); + expect(conn.sendMessage).toHaveBeenCalledWith({ action: "stop" }); + expect(conn.disconnect).not.toHaveBeenCalled(); + + // Stop 之后、真正终态事件到达之前的流式增量应被抑制 + act(() => { + emit({ type: "content_delta", delta: "不应该被处理" }); + }); + expect(onEvent).not.toHaveBeenCalledWith({ type: "content_delta", delta: "不应该被处理" }); + + // 真正的终态事件(携带取消原因/usage)到达:必须放行、断开连接、触发 onDone + act(() => { + emit({ type: "error", errorCode: "cancelled", message: "Conversation cancelled", usage: { inputTokens: 1 } }); + }); + expect(onEvent).toHaveBeenCalledWith(expect.objectContaining({ type: "error", errorCode: "cancelled" })); + expect(conn.disconnect).toHaveBeenCalledOnce(); + expect(onDone).toHaveBeenCalledOnce(); + }); +}); const conv = (id: string, title = "c"): Conversation => ({ id, + generation: `gen-${id}`, + revision: 1, title, modelId: "gpt-4o", createtime: 1, @@ -56,7 +171,7 @@ describe("会话管理 Hook useConversations", () => { await act(async () => { created = await result.current.createConversation("gpt-4o"); }); - expect(repoMock.saveConversation).toHaveBeenCalledOnce(); + expect(repoMock.createConversation).toHaveBeenCalledOnce(); expect(created!.conv.modelId).toBe("gpt-4o"); expect(created!.conv.title).toBe("New Chat"); await waitFor(() => expect(result.current.activeId).toBe(created!.conv.id)); @@ -71,7 +186,11 @@ describe("会话管理 Hook useConversations", () => { await act(async () => { await result.current.deleteConversation("a"); }); - expect(repoMock.deleteConversation).toHaveBeenCalledWith("a"); + expect(mockSendMessage).toHaveBeenCalledWith({}, "serviceWorker/agent/conversation", { + action: "delete", + conversationId: "a", + generation: "gen-a", + }); await waitFor(() => expect(result.current.activeId).toBe("b")); }); @@ -89,24 +208,30 @@ describe("会话管理 Hook useConversations", () => { }); describe("消息持久化操作", () => { - it("deleteMessages 过滤掉指定 id 后回写", async () => { - const msgs: ChatMessage[] = [ - { id: "m1", conversationId: "c", role: "user", content: "a", createtime: 1 }, - { id: "m2", conversationId: "c", role: "assistant", content: "b", createtime: 2 }, - { id: "m3", conversationId: "c", role: "user", content: "c", createtime: 3 }, - ]; - repoMock.getMessages.mockResolvedValue(msgs); - + it("deleteMessages 应通过 Service Worker 串行删除指定消息", async () => { await deleteMessages("c", ["m2"]); + expect(mockSendMessage).toHaveBeenCalledWith({}, "serviceWorker/agent/conversation", { + action: "deleteMessages", + conversationId: "c", + messageIds: ["m2"], + }); + }); - const [convId, saved] = repoMock.saveMessages.mock.calls.at(-1)!; - expect(convId).toBe("c"); - expect((saved as ChatMessage[]).map((m) => m.id)).toEqual(["m1", "m3"]); + it("deleteMessages 在重新生成时应保留即将转移所有权的附件", async () => { + await deleteMessages("c", ["m2"], ["keep.png"]); + expect(mockSendMessage).toHaveBeenCalledWith({}, "serviceWorker/agent/conversation", { + action: "deleteMessages", + conversationId: "c", + messageIds: ["m2"], + preserveAttachmentIds: ["keep.png"], + }); }); - it("clearMessages 同时清空消息与任务", async () => { + it("clearMessages 应通过 Service Worker 串行清空消息与任务", async () => { await clearMessages("c"); - expect(repoMock.saveMessages).toHaveBeenCalledWith("c", []); - expect(repoMock.saveTasks).toHaveBeenCalledWith("c", []); + expect(mockSendMessage).toHaveBeenCalledWith({}, "serviceWorker/agent/conversation", { + action: "clearMessages", + conversationId: "c", + }); }); }); diff --git a/src/pages/options/routes/Agent/Chat/hooks.ts b/src/pages/options/routes/Agent/Chat/hooks.ts index dac8f3690..a6e8bbf88 100644 --- a/src/pages/options/routes/Agent/Chat/hooks.ts +++ b/src/pages/options/routes/Agent/Chat/hooks.ts @@ -66,10 +66,10 @@ export function useConversations() { createtime: Date.now(), updatetime: Date.now(), }; - await agentChatRepo.saveConversation(conv); + const created = await agentChatRepo.createConversation(conv); const list = await loadConversations(); - setActiveId(conv.id); - return { conv, list }; + setActiveId(created.id); + return { conv: created, list }; }, [loadConversations, setActiveId] ); @@ -77,13 +77,20 @@ export function useConversations() { // 删除会话 const deleteConversation = useCallback( async (id: string) => { - await agentChatRepo.deleteConversation(id); + const conversation = conversations.find((item) => item.id === id); + if (conversation?.generation) { + await sendMsg(extensionMessage, "serviceWorker/agent/conversation", { + action: "delete", + conversationId: id, + generation: conversation.generation, + }); + } const list = await loadConversations(); if (activeId === id) { setActiveId(list[0]?.id || ""); } }, - [activeId, loadConversations, setActiveId] + [activeId, conversations, loadConversations, setActiveId] ); // 重命名会话 @@ -135,7 +142,14 @@ export function useMessages(conversationId: string) { } // ask_user 待回复状态 -export type AskUserPending = { id: string; question: string; options?: string[]; multiple?: boolean }; +export type AskUserPending = { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; +}; // 流式聊天 hook export function useStreamingChat() { @@ -147,21 +161,18 @@ export function useStreamingChat() { const stopGeneration = useCallback(() => { abortedRef.current = true; const conn = connRef.current; - connRef.current = null; - // 先确保 UI 状态重置,再断开连接(避免 sendMessage/disconnect 抛异常导致状态卡住) - setIsStreaming(false); setAskUserPending(null); + // 不在这里把 isStreaming 置 false、也不立即断开连接——ChatArea 里"连接断开但 done + // 回调未触发时处理排队消息"的兜底逻辑正是监听 isStreaming 由 true 变 false 来触发的; + // 提前置为 false 会让排队消息在旧会话仍处于 cancelling(占用中)时就被处理,进而被拒绝 + // 且从未持久化就丢失。isStreaming 必须留到真正的终态事件到达时 + // (onMessage 的终态分支)或连接意外断开时(onDisconnect)才由那两处统一置 false。 if (conn) { try { conn.sendMessage({ action: "stop" }); } catch { // port 可能已断开 } - try { - conn.disconnect(); - } catch { - // port 可能已断开 - } } }, []); @@ -182,7 +193,15 @@ export function useStreamingChat() { modelId?: string, skipSaveUserMessage?: boolean, enableTools?: boolean, - extra?: { compact?: boolean; compactInstruction?: string; background?: boolean } + extra?: { + compact?: boolean; + compactInstruction?: string; + background?: boolean; + ownedAttachmentIds?: string[]; + // 调用方持有的会话 generation:与当前存储不一致(会话已被删除重建)时, + // 服务端拒绝而不是静默作用于无关的新一代会话(含 compact 分支) + generation?: string; + } ) => { setIsStreaming(true); abortedRef.current = false; @@ -201,22 +220,33 @@ export function useStreamingChat() { connRef.current = conn; conn.onMessage((msg) => { - if (abortedRef.current) return; const event = msg.data as ChatStreamEvent; + const isTerminal = + (event.type === "done" || event.type === "error") && !("subAgent" in event && event.subAgent); + // stop 之后(abortedRef=true)必须继续放行终态事件——那条事件携带真正的取消原因/ + // usage/耗时,且负责断开连接、触发 onDone(进而处理排队消息);只需要抑制中间的 + // 流式增量/ask_user 事件,避免用户点了停止之后 UI 还在继续刷新内容 + if (abortedRef.current && !isTerminal) return; // 处理 ask_user 事件 if (event.type === "ask_user") { setAskUserPending({ id: event.id, question: event.question, options: event.options, + optionValues: event.optionValues, multiple: event.multiple, + allowCustom: event.allowCustom, }); } + if (event.type === "ask_user_expired") setAskUserPending(null); + if (event.type === "ask_user_resolved") setAskUserPending(null); onEvent(event); - if ((event.type === "done" || event.type === "error") && !("subAgent" in event && event.subAgent)) { + if (isTerminal) { setIsStreaming(false); setAskUserPending(null); connRef.current = null; + // 终态事件后必须主动断开,否则 port/listener 会一直挂在 SW 侧,直到用户手动刷新页面 + conn.disconnect(); onDone(); } }); @@ -227,6 +257,9 @@ export function useStreamingChat() { connRef.current = null; }); } catch (e: any) { + await Promise.all( + (extra?.ownedAttachmentIds || []).map((id) => agentChatRepo.deleteAttachment(id).catch(() => {})) + ); setIsStreaming(false); setAskUserPending(null); onEvent({ type: "error", message: e.message || "Connection failed" }); @@ -238,28 +271,42 @@ export function useStreamingChat() { // 附加到后台运行中的会话 const attachToConversation = useCallback( - async (conversationId: string, onEvent: (event: ChatStreamEvent) => void, onDone: () => void) => { + async ( + conversationId: string, + onEvent: (event: ChatStreamEvent) => void, + onDone: () => void, + generation?: string + ) => { abortedRef.current = false; try { const conn = await connect(extensionMessage, "serviceWorker/agent/attachToConversation", { conversationId, + generation, }); connRef.current = conn; conn.onMessage((msg) => { - if (abortedRef.current) return; const event = msg.data as ChatStreamEvent; + const isTerminalEvent = + (event.type === "done" || event.type === "error") && !("subAgent" in event && event.subAgent); + // 与 sendMessage() 同理:stop 之后必须继续放行终态事件(含终态 sync), + // 否则会错过真正携带取消原因/usage 的那条事件,也无法触发 onDone 处理排队消息 + if (abortedRef.current && event.type !== "sync" && !isTerminalEvent) return; if (event.type === "ask_user") { setAskUserPending({ id: event.id, question: event.question, options: event.options, + optionValues: event.optionValues, multiple: event.multiple, + allowCustom: event.allowCustom, }); } + if (event.type === "ask_user_expired") setAskUserPending(null); + if (event.type === "ask_user_resolved") setAskUserPending(null); onEvent(event); @@ -274,15 +321,18 @@ export function useStreamingChat() { // done 或 error,无需保持连接 setIsStreaming(false); connRef.current = null; + // 终态 sync 后必须主动断开,否则 port/listener 会一直挂在 SW 侧 + conn.disconnect(); onDone(); } return; } - if ((event.type === "done" || event.type === "error") && !("subAgent" in event && event.subAgent)) { + if (isTerminalEvent) { setIsStreaming(false); setAskUserPending(null); connRef.current = null; + conn.disconnect(); onDone(); } }); @@ -336,21 +386,32 @@ export function useRunningConversations() { } // 批量删除持久化消息 -export async function deleteMessages(conversationId: string, messageIds: string[]): Promise { - const messages = await agentChatRepo.getMessages(conversationId); - const idSet = new Set(messageIds); - const filtered = messages.filter((m) => !idSet.has(m.id)); - await agentChatRepo.saveMessages(conversationId, filtered); +export async function deleteMessages( + conversationId: string, + messageIds: string[], + preserveAttachmentIds?: string[], + generation?: string +): Promise { + await sendMsg(extensionMessage, "serviceWorker/agent/conversation", { + action: "deleteMessages", + conversationId, + generation, + messageIds, + ...(preserveAttachmentIds?.length ? { preserveAttachmentIds } : {}), + }); } // 清空对话消息及任务 -export async function clearMessages(conversationId: string): Promise { - await agentChatRepo.saveMessages(conversationId, []); - await agentChatRepo.saveTasks(conversationId, []); +export async function clearMessages(conversationId: string, generation?: string): Promise { + await sendMsg(extensionMessage, "serviceWorker/agent/conversation", { + action: "clearMessages", + conversationId, + generation, + }); } // 会话任务列表 hook -export function useConversationTasks(conversationId: string) { +export function useConversationTasks(conversationId: string, generation?: string) { const [tasks, setTasks] = useState([]); const loadTasks = useCallback(async () => { @@ -358,9 +419,9 @@ export function useConversationTasks(conversationId: string) { setTasks([]); return; } - const loaded = await agentChatRepo.getTasks(conversationId); + const loaded = await agentChatRepo.getTasks(conversationId, generation); setTasks(loaded); - }, [conversationId]); + }, [conversationId, generation]); useEffect(() => { void (async () => { diff --git a/src/pages/options/routes/Agent/Chat/index.tsx b/src/pages/options/routes/Agent/Chat/index.tsx index 2c6cccb70..adb80fbec 100644 --- a/src/pages/options/routes/Agent/Chat/index.tsx +++ b/src/pages/options/routes/Agent/Chat/index.tsx @@ -172,6 +172,7 @@ export default function AgentChat() { const chatArea = ( { }); }); + describe("路径 0:显式 toolCallId 匹配(并发 agent 调用)", () => { + it("两个并行 agent 调用各自匹配到自己的子代理,不会都指向同一个", () => { + const saA = makeSA({ agentId: "agent-a", toolCallId: "tc-a", isRunning: true }); + const saB = makeSA({ agentId: "agent-b", toolCallId: "tc-b", isRunning: true }); + const subAgents = new Map([ + ["agent-a", saA], + ["agent-b", saB], + ]); + + // 两个工具调用均无 result(仍在流式中),仅凭各自的 toolCallId 区分 + const tcA = { id: "tc-a", name: "agent" }; + const tcB = { id: "tc-b", name: "agent" }; + + expect(getSubAgentForToolCall(tcA, subAgents)).toBe(saA); + expect(getSubAgentForToolCall(tcB, subAgents)).toBe(saB); + }); + + it("其中一个子代理率先启动运行时,另一个仍未产生事件的调用不会被误配给它", () => { + // 只有 agent-a 已经开始运行(订阅到了流式事件),agent-b 对应的子代理事件尚未到达。 + const saA = makeSA({ agentId: "agent-a", toolCallId: "tc-a", isRunning: true }); + const subAgents = new Map([["agent-a", saA]]); + + const tcB = { id: "tc-b", name: "agent" }; + // 旧实现会把"唯一运行中的子代理"错误地匹配给 tcB;toolCallId 已知时必须拒绝匹配。 + expect(getSubAgentForToolCall(tcB, subAgents)).toBeUndefined(); + }); + + it("toolCallId 已知但尚未匹配到任何子代理时,不回退到 result/arguments 之外的猜测", () => { + const saA = makeSA({ agentId: "agent-a", toolCallId: "tc-a", isRunning: false }); + const subAgents = new Map([["agent-a", saA]]); + + const tcB = { id: "tc-b", name: "agent" }; + expect(getSubAgentForToolCall(tcB, subAgents)).toBeUndefined(); + }); + }); + describe("路径 2:回退到持久化 subAgentDetails", () => { it("无流式 subAgents 时从 subAgentDetails 构建状态", () => { const result = getSubAgentForToolCall({ @@ -165,7 +201,8 @@ describe("getSubAgentForToolCall", () => { let mergedTc = merged.find((m) => m.role === "assistant")!.toolCalls![0]; expect(getSubAgentForToolCall(mergedTc, subAgents)).toBeUndefined(); - const sa = makeSA({ agentId: "sa-1", isRunning: true }); + // 真实流程中,子代理事件从第一次到达起就带着发起它的 toolCallId(见 chat_service.ts 的 subSendEvent) + const sa = makeSA({ agentId: "sa-1", toolCallId: "tc-agent", isRunning: true }); subAgents.set("sa-1", sa); merged = mergeToolResults(streamingMessages); mergedTc = merged.find((m) => m.role === "assistant")!.toolCalls![0]; diff --git a/src/pages/options/routes/Agent/Chat/types.ts b/src/pages/options/routes/Agent/Chat/types.ts index 7d9952a57..bcd9b0fb3 100644 --- a/src/pages/options/routes/Agent/Chat/types.ts +++ b/src/pages/options/routes/Agent/Chat/types.ts @@ -1,4 +1,4 @@ -import type { SubAgentMessage, ToolCall, TokenUsage } from "@App/app/service/agent/core/types"; +import type { ContentBlock, SubAgentMessage, ToolCall, TokenUsage } from "@App/app/service/agent/core/types"; export type { SubAgentMessage }; @@ -7,12 +7,16 @@ export type SubAgentState = { agentId: string; description: string; subAgentType?: string; + /** 发起该子代理的 agent 工具调用 ID,用于并发调用时的显式匹配(而非猜测第一个运行中的子代理) */ + toolCallId?: string; /** 已完成的消息轮次 */ completedMessages: SubAgentMessage[]; /** 当前正在构建的消息内容 */ currentContent: string; + currentBlocks?: ContentBlock[]; currentThinking: string; currentToolCalls: ToolCall[]; + currentWarning?: string; isRunning: boolean; /** 重试信息 */ retryInfo?: { attempt: number; maxRetries: number; error: string }; diff --git a/src/pages/options/routes/Agent/Settings/index.test.tsx b/src/pages/options/routes/Agent/Settings/index.test.tsx index 526d01d8d..256900fcf 100644 --- a/src/pages/options/routes/Agent/Settings/index.test.tsx +++ b/src/pages/options/routes/Agent/Settings/index.test.tsx @@ -16,7 +16,9 @@ vi.mock("@App/pages/options/hooks/useScrollSpy", () => ({ }), })); -const { getSearchConfigMock } = vi.hoisted(() => ({ getSearchConfigMock: vi.fn() })); +const { getSearchConfigMock } = vi.hoisted(() => ({ + getSearchConfigMock: vi.fn(), +})); vi.mock("@App/pages/store/features/script", () => ({ agentClient: { listModels: vi.fn(async () => [ diff --git a/src/pages/options/routes/Agent/Tasks/TaskFormDialog.test.tsx b/src/pages/options/routes/Agent/Tasks/TaskFormDialog.test.tsx index 06afb08be..15e6e4d42 100644 --- a/src/pages/options/routes/Agent/Tasks/TaskFormDialog.test.tsx +++ b/src/pages/options/routes/Agent/Tasks/TaskFormDialog.test.tsx @@ -14,10 +14,36 @@ function setup(props: Record = {}) { } describe("TaskFormDialog 定时任务弹窗", () => { - it("内部模式显示提示词,切换到事件模式后隐藏", () => { + it("新建任务时事件模式不可选,因为 Options 无法像脚本上下文那样自动注入来源脚本 UUID", () => { setup(); expect(screen.getByTestId("task-prompt")).toBeInTheDocument(); + expect(screen.getByTestId("task-mode-event")).toBeDisabled(); fireEvent.click(screen.getByTestId("task-mode-event")); + // 禁用状态下点击不应切换模式,提示词字段应保持可见 + expect(screen.getByTestId("task-prompt")).toBeInTheDocument(); + }); + + it("编辑已有事件任务时事件模式保持可选并隐藏提示词", () => { + const onSubmit = vi.fn(); + const eventTask = { + id: "task-event-1", + name: "事件任务", + mode: "event" as const, + crontab: "0 9 * * *", + enabled: true, + notify: false, + sourceScriptUuid: "script-1", + createtime: 1, + updatetime: 1, + }; + // 组件用「渲染期比较上一次的 open/value」同步外部 prop:初始挂载时 value 与其自身相等不会触发同步, + // 需先以 value=null 挂载,再 rerender 传入编辑值,才能复现真实的「打开编辑」场景 + const { rerender } = render( + {}} onSubmit={onSubmit} /> + ); + rerender( {}} onSubmit={onSubmit} />); + + expect(screen.getByTestId("task-mode-event")).not.toBeDisabled(); expect(screen.queryByTestId("task-prompt")).toBeNull(); }); diff --git a/src/pages/options/routes/Agent/Tasks/TaskFormDialog.tsx b/src/pages/options/routes/Agent/Tasks/TaskFormDialog.tsx index e150b4bd2..f89c0e69f 100644 --- a/src/pages/options/routes/Agent/Tasks/TaskFormDialog.tsx +++ b/src/pages/options/routes/Agent/Tasks/TaskFormDialog.tsx @@ -21,7 +21,10 @@ import { nextRunText } from "./cron"; // 在联合类型每个分支上分别 Omit,保留 internal/event 各自的专有字段 type DistributiveOmit = T extends unknown ? Omit : never; -export type TaskFormValue = DistributiveOmit; +export type TaskFormValue = DistributiveOmit< + AgentTask, + "id" | "generation" | "revision" | "createtime" | "updatetime" | "nextruntime" +>; export function TaskFormDialog({ open, @@ -34,7 +37,7 @@ export function TaskFormDialog({ value: AgentTask | null; models: AgentModelConfig[]; onOpenChange: (v: boolean) => void; - onSubmit: (task: TaskFormValue) => void; + onSubmit: (task: TaskFormValue) => Promise; }) { const { t } = useTranslation(["agent", "common", "script"]); const [name, setName] = useState(""); @@ -44,7 +47,7 @@ export function TaskFormDialog({ const [enabled, setEnabled] = useState(true); const [prompt, setPrompt] = useState(""); const [modelId, setModelId] = useState(""); - const [maxIterations, setMaxIterations] = useState(""); + const [isSubmitting, setIsSubmitting] = useState(false); // 弹窗打开(或打开期间 value 变化)时,依据传入的 value 重置/同步各字段。 // 用「渲染期比较上一次的 open/value 再 setState」模式同步外部 prop,等价于原 useEffect。 @@ -61,11 +64,9 @@ export function TaskFormDialog({ if (value?.mode === "internal") { setPrompt(value.prompt ?? ""); setModelId(value.modelId ?? ""); - setMaxIterations(value.maxIterations != null ? String(value.maxIterations) : ""); } else { setPrompt(""); setModelId(""); - setMaxIterations(""); } } else if (open !== prevOpen || value !== prevValue) { // 弹窗关闭或 value 在关闭状态下变化:仅记录最新值,不触碰表单字段(与原 `if (!open) return` 一致) @@ -77,8 +78,11 @@ export function TaskFormDialog({ const cronInvalid = crontab.trim().length > 0 && !cron.valid; const canSubmit = !!name && cron.valid; const hasModels = models.length > 0; + // Options 无法像脚本上下文那样自动注入创建者的脚本 UUID,事件任务留空 sourceScriptUuid + // 会导致后端拒绝创建;因此事件模式仅对已存在的事件任务(编辑场景)开放,新建任务不可选 + const eventModeSelectable = value?.mode === "event"; - const handleSubmit = () => { + const handleSubmit = async () => { const base = { name, crontab, enabled, notify }; const task: TaskFormValue = mode === "internal" @@ -87,7 +91,6 @@ export function TaskFormDialog({ mode: "internal", prompt, modelId: modelId || undefined, - maxIterations: maxIterations ? Number(maxIterations) : undefined, } : { ...base, @@ -95,7 +98,12 @@ export function TaskFormDialog({ // 事件任务由脚本创建;编辑时保留来源脚本 UUID,新建时留空 sourceScriptUuid: value?.mode === "event" ? value.sourceScriptUuid : "", }; - onSubmit(task); + setIsSubmitting(true); + try { + await onSubmit(task); + } finally { + setIsSubmitting(false); + } }; return ( @@ -120,6 +128,7 @@ export function TaskFormDialog({ value: m, label: m === "internal" ? t("agent:tasks_mode_internal_short") : t("agent:tasks_mode_event_short"), testId: `task-mode-${m}`, + disabled: m === "event" && !eventModeSelectable, }))} /> @@ -173,15 +182,6 @@ export function TaskFormDialog({ - - setMaxIterations(e.target.value)} - /> - )} @@ -197,7 +197,7 @@ export function TaskFormDialog({ - diff --git a/src/pages/options/routes/Agent/Tasks/index.test.tsx b/src/pages/options/routes/Agent/Tasks/index.test.tsx index daeca79fd..aaaa3692d 100644 --- a/src/pages/options/routes/Agent/Tasks/index.test.tsx +++ b/src/pages/options/routes/Agent/Tasks/index.test.tsx @@ -1,23 +1,17 @@ import { describe, it, expect, vi, beforeAll, beforeEach, afterEach } from "vitest"; -import { render, cleanup, screen } from "@testing-library/react"; +import { render, cleanup, screen, fireEvent, waitFor } from "@testing-library/react"; import { t } from "@App/locales/locales"; import { initTestLanguage } from "@Tests/initTestLanguage"; import { useIsMobile } from "@App/pages/components/use-is-mobile"; -const { listTasksMock } = vi.hoisted(() => ({ listTasksMock: vi.fn() })); +const { listTasksMock, agentTaskMock } = vi.hoisted(() => ({ listTasksMock: vi.fn(), agentTaskMock: vi.fn() })); -vi.mock("@App/app/repo/agent_task", () => ({ - AgentTaskRepo: class { - listTasks = listTasksMock; - saveTask = vi.fn(); - removeTask = vi.fn(); - }, - AgentTaskRunRepo: class { - listRuns = vi.fn(async () => []); - clearRuns = vi.fn(); +vi.mock("@App/pages/store/features/script", () => ({ + agentClient: { + listModels: vi.fn(async () => []), + agentTask: agentTaskMock, }, })); -vi.mock("@App/pages/store/features/script", () => ({ agentClient: { listModels: vi.fn(async () => []) } })); // DOM 测试环境默认未实现 matchMedia,useIsMobile 依赖它——默认桌面,移动用例单独覆盖 vi.mock("@App/pages/components/use-is-mobile", () => ({ useIsMobile: vi.fn(() => false) })); @@ -33,13 +27,20 @@ const sampleTask = { prompt: "总结今天", createtime: 0, updatetime: 0, + generation: "generation-1", + revision: 1, }; beforeAll(() => initTestLanguage("zh-CN")); beforeEach(() => { (useIsMobile as unknown as ReturnType).mockReturnValue(false); + listTasksMock.mockReset(); listTasksMock.mockResolvedValue([sampleTask]); + agentTaskMock.mockReset(); + agentTaskMock.mockImplementation(async (request: { action: string }) => + request.action === "list" ? listTasksMock() : [] + ); }); afterEach(() => cleanup()); @@ -78,4 +79,21 @@ describe("AgentTasks 页面", () => { expect(screen.queryByTestId("page-header-docs")).toBeNull(); expect(screen.queryByTestId("tasks-mobile-bar")).toBeNull(); }); + + it("编辑发生 revision 冲突时应关闭绑定旧快照的弹窗并刷新任务", async () => { + agentTaskMock.mockImplementation(async (request: { action: string }) => { + if (request.action === "update") throw new Error("Task changed"); + return request.action === "list" ? listTasksMock() : []; + }); + render(); + await screen.findByText("每日总结"); + + fireEvent.pointerDown(screen.getByTestId("card-menu"), { button: 0 }); + fireEvent.click(screen.getByTestId("card-menu-edit")); + fireEvent.change(await screen.findByTestId("task-name"), { target: { value: "新名称" } }); + fireEvent.click(screen.getByTestId("task-submit")); + + await waitFor(() => expect(screen.queryByTestId("task-submit")).toBeNull()); + expect(listTasksMock).toHaveBeenCalledTimes(2); + }); }); diff --git a/src/pages/options/routes/Agent/Tasks/index.tsx b/src/pages/options/routes/Agent/Tasks/index.tsx index fd69c1f43..889e1e89b 100644 --- a/src/pages/options/routes/Agent/Tasks/index.tsx +++ b/src/pages/options/routes/Agent/Tasks/index.tsx @@ -3,7 +3,6 @@ import { useTranslation } from "react-i18next"; import { Plus, CalendarClock } from "lucide-react"; import { notify } from "@App/pages/components/ui/toast"; import { Button } from "@App/pages/components/ui/button"; -import { AgentTaskRepo, AgentTaskRunRepo } from "@App/app/repo/agent_task"; import { agentClient } from "@App/pages/store/features/script"; import type { AgentTask, AgentModelConfig, AgentTaskRun } from "@App/app/service/agent/core/types"; import { useIsMobile } from "@App/pages/components/use-is-mobile"; @@ -16,9 +15,6 @@ import { TaskFormDialog, type TaskFormValue } from "./TaskFormDialog"; import { TaskHistorySheet } from "./TaskHistorySheet"; import { nextRunText } from "./cron"; -const taskRepo = new AgentTaskRepo(); -const taskRunRepo = new AgentTaskRunRepo(); - export default function AgentTasks() { const { t } = useTranslation(["agent", "common"]); const isMobile = useIsMobile(); @@ -35,7 +31,10 @@ export default function AgentTasks() { const [runs, setRuns] = useState([]); const reload = useCallback(async () => { - const [taskList, modelList] = await Promise.all([taskRepo.listTasks(), agentClient.listModels()]); + const [taskList, modelList] = await Promise.all([ + agentClient.agentTask({ action: "list" }) as Promise, + agentClient.listModels(), + ]); setTasks(taskList); setModels(modelList); setLoading(false); @@ -58,39 +57,67 @@ export default function AgentTasks() { }; const handleSubmit = async (formValue: TaskFormValue) => { - const now = Date.now(); - const task = ( - editing - ? { ...formValue, id: editing.id, createtime: editing.createtime, updatetime: now } - : { ...formValue, id: crypto.randomUUID(), createtime: now, updatetime: now } - ) as AgentTask; - await taskRepo.saveTask(task); - setDialogOpen(false); - notify.success(t("common:save_success")); - await reload(); + try { + if (editing) { + await agentClient.agentTask({ + action: "update", + id: editing.id, + generation: editing.generation!, + revision: editing.revision!, + task: formValue, + }); + } else { + await agentClient.agentTask({ action: "create", task: formValue }); + } + setDialogOpen(false); + notify.success(t("common:save_success")); + await reload(); + } catch (error) { + setDialogOpen(false); + setEditing(null); + notify.error(`${t("common:error")}: ${error instanceof Error ? error.message : String(error)}`); + await reload(); + } }; const handleToggle = useCallback( async (task: AgentTask, enabled: boolean) => { - await taskRepo.saveTask({ ...task, enabled, updatetime: Date.now() }); - await reload(); + try { + await agentClient.agentTask({ + action: "enable", + id: task.id, + generation: task.generation!, + revision: task.revision!, + enabled, + }); + } catch (error) { + notify.error(`${t("common:error")}: ${error instanceof Error ? error.message : String(error)}`); + } finally { + await reload(); + } }, - [reload] + [reload, t] ); const handleDelete = async (task: AgentTask) => { - await taskRepo.removeTask(task.id); - notify.success(t("common:delete_success")); - await reload(); + try { + await agentClient.agentTask({ + action: "delete", + id: task.id, + generation: task.generation!, + revision: task.revision!, + }); + notify.success(t("common:delete_success")); + } catch (error) { + notify.error(`${t("common:error")}: ${error instanceof Error ? error.message : String(error)}`); + } finally { + await reload(); + } }; const handleRunNow = async (task: AgentTask) => { try { - await chrome.runtime.sendMessage({ - channel: "agent", - action: "agentTask", - data: { action: "runNow", id: task.id }, - }); + await agentClient.agentTask({ action: "runNow", id: task.id }); notify.success(t("agent:tasks_run_now")); } catch { // 调度器会在下次 tick 执行 @@ -102,14 +129,14 @@ export default function AgentTasks() { setHistoryTask(task); setHistoryOpen(true); setHistoryLoading(true); - const list = await taskRunRepo.listRuns(task.id); + const list = (await agentClient.agentTask({ action: "listRuns", taskId: task.id })) as AgentTaskRun[]; setRuns(list); setHistoryLoading(false); }; const handleClearRuns = async () => { if (!historyTask) return; - await taskRunRepo.clearRuns(historyTask.id); + await agentClient.agentTask({ action: "clearRuns", taskId: historyTask.id }); setRuns([]); }; diff --git a/src/pkg/utils/with_timeout.ts b/src/pkg/utils/with_timeout.ts index d9904ba2c..d60fe1f56 100644 --- a/src/pkg/utils/with_timeout.ts +++ b/src/pkg/utils/with_timeout.ts @@ -4,12 +4,31 @@ * @param ms 超时毫秒数 * @param onTimeoutError 可选:自定义超时错误构造器;默认抛 Error("operation timed out") */ -export function withTimeout(promise: Promise, ms: number, onTimeoutError?: () => Error): Promise { - let timer: ReturnType; +export function withTimeout( + promise: Promise, + ms: number, + onTimeoutError?: () => Error, + signal?: AbortSignal +): Promise { + let timer: ReturnType | undefined; + let onAbort: (() => void) | undefined; const timeoutPromise = new Promise((_, reject) => { + if (signal?.aborted) { + reject(new Error("Aborted")); + return; + } + if (signal) { + onAbort = () => reject(new Error("Aborted")); + signal.addEventListener("abort", onAbort, { once: true }); + } timer = setTimeout(() => { reject(onTimeoutError ? onTimeoutError() : new Error("operation timed out")); }, ms); }); - return Promise.race([promise, timeoutPromise]).finally(() => clearTimeout(timer)); + return Promise.race([promise, timeoutPromise]).finally(() => { + clearTimeout(timer); + if (signal && onAbort) { + signal.removeEventListener("abort", onAbort); + } + }); } diff --git a/src/types/scriptcat.agent-background.d.ts b/src/types/scriptcat.agent-background.d.ts new file mode 100644 index 000000000..911497416 --- /dev/null +++ b/src/types/scriptcat.agent-background.d.ts @@ -0,0 +1,41 @@ +/** 英文与中文声明文件共用的后台会话扩展。 */ +declare namespace CATAgent { + /** 附加到后台会话时首先返回的状态快照。 */ + interface SyncStreamChunk { + type: "sync"; + /** 附加前已累计的 assistant 输出。 */ + streamingMessage?: { + content: string; + thinking?: string; + toolCalls: ToolCallInfo[]; + }; + /** 会话正在等待输入时的 ask_user 请求。 */ + pendingAskUser?: { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + }; + /** 当前任务快照。 */ + tasks: Array<{ + id: string; + subject: string; + status: "pending" | "in_progress" | "completed"; + description?: string; + }>; + /** 终态快照是 attach() 返回的最后一个数据块。 */ + status: "running" | "done" | "error"; + } + + interface ConversationCreateOptions { + /** 原始页面断开后仍在 Service Worker 中继续运行对话。 */ + background?: boolean; + } + + interface ConversationInstance { + /** 附加到后台会话,先接收必需快照,再接收后续流式数据。 */ + attach(): Promise>; + } +} diff --git a/src/types/scriptcat.d.ts b/src/types/scriptcat.d.ts index 85a45b9d0..f121861a3 100644 --- a/src/types/scriptcat.d.ts +++ b/src/types/scriptcat.d.ts @@ -841,8 +841,8 @@ declare namespace CATAgent { description: string; /** JSON Schema describing the tool parameters. */ parameters: Record; - /** Handler invoked when the LLM calls this tool. */ - handler: (args: Record) => Promise; + /** Handler invoked when the LLM calls this tool. Stop work promptly when the batch signal is aborted. */ + handler: (args: Record, signal: AbortSignal) => Promise; } /** @@ -861,8 +861,6 @@ declare namespace CATAgent { system?: string; /** Model ID; uses the default model if omitted. */ model?: string; - /** Max tool-calling loop iterations (default: 20). */ - maxIterations?: number; /** Skills to load: `"auto"` loads all installed skills, or specify names. */ skills?: "auto" | string[]; /** Tools with inline handlers, available for the lifetime of this conversation. */ @@ -879,6 +877,8 @@ declare namespace CATAgent { ephemeral?: boolean; /** Enable prompt caching. Defaults to true. */ cache?: boolean; + /** Keep the conversation running in the Service Worker after the page disconnects. */ + background?: boolean; } /** Options for a single `chat()` / `chatStream()` call. */ @@ -930,9 +930,18 @@ declare namespace CATAgent { /** Tool calls made during this turn. */ toolCalls?: ToolCallInfo[]; /** Token usage. */ - usage?: { inputTokens: number; outputTokens: number }; + usage?: { + inputTokens: number; + outputTokens: number; + cacheCreationInputTokens?: number; + cacheReadInputTokens?: number; + }; + /** Total response duration in ms. */ + durationMs?: number; /** `true` when the reply was produced by a command handler, not the LLM. */ command?: boolean; + /** Non-fatal warning about generated data loss (e.g. a generated image failed to save). */ + warning?: string; } /** A single chunk emitted during streaming via `chatStream()`. */ @@ -941,26 +950,76 @@ declare namespace CATAgent { * Chunk type: * - `"content_delta"` — incremental text * - `"thinking_delta"` — incremental thinking/reasoning - * - `"tool_call"` — a tool call event + * - `"tool_call"` — a tool call event (start or argument delta) + * - `"tool_call_complete"` — a tool call finished executing, carries result/status/attachments * - `"content_block"` — a complete non-text content block + * - `"new_message"` — the current round ended, the next assistant message is about to start + * - `"system_warning"` — a non-fatal warning about generated data loss * - `"done"` — stream finished * - `"error"` — an error occurred */ - type: "content_delta" | "thinking_delta" | "tool_call" | "content_block" | "done" | "error"; + type: + | "content_delta" + | "thinking_delta" + | "tool_call" + | "tool_call_complete" + | "content_block" + | "new_message" + | "system_warning" + | "done" + | "error"; /** Text delta (for content_delta / thinking_delta). */ content?: string; /** Complete content block (for content_block). */ block?: ContentBlock; - /** Tool call info (for tool_call). */ + /** Tool call info (for tool_call / tool_call_complete). */ toolCall?: ToolCallInfo; /** Token usage (for done). */ - usage?: { inputTokens: number; outputTokens: number }; + usage?: { + inputTokens: number; + outputTokens: number; + cacheCreationInputTokens?: number; + cacheReadInputTokens?: number; + }; + /** Total response duration in ms. */ + durationMs?: number; /** Error message (for error). */ error?: string; - /** Error classification: `"rate_limit"` | `"auth"` | `"tool_timeout"` | `"max_iterations"` | `"api_error"` */ + /** Error classification: `"rate_limit"` | `"auth"` | `"tool_timeout"` | `"context_too_large"` | `"api_error"` */ errorCode?: string; /** `true` when the chunk was produced by a command handler. */ command?: boolean; + /** Warning text (for `"system_warning"`). */ + warning?: string; + } + + /** Initial snapshot returned when attaching to a background conversation. */ + interface SyncStreamChunk { + type: "sync"; + /** Assistant output accumulated before attach(). */ + streamingMessage?: { + content: string; + thinking?: string; + toolCalls: ToolCallInfo[]; + }; + /** Ask-user request when the conversation is waiting for input. */ + pendingAskUser?: { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + }; + /** Current task snapshot. */ + tasks: Array<{ + id: string; + subject: string; + status: "pending" | "in_progress" | "completed"; + description?: string; + }>; + /** Final snapshot status emitted by attach(). */ + status: "running" | "done" | "error"; } // ---- Chat message ---- @@ -983,6 +1042,10 @@ declare namespace CATAgent { toolCallId?: string; /** Error message (if the turn errored). */ error?: string; + /** Error classification code (for example, `context_too_large`). */ + errorCode?: string; + /** Non-fatal warning associated with this message. */ + warning?: string; /** Model ID used for this message. */ modelId?: string; /** Token usage for this message. */ @@ -1024,6 +1087,9 @@ declare namespace CATAgent { /** Send a message and receive a streaming response. */ chatStream(content: MessageContent, options?: ChatOptions): Promise>; + /** Attach to a background conversation and receive the sync snapshot plus stream events. */ + attach(): Promise>; + /** Get all messages in this conversation. */ getMessages(): Promise; @@ -1280,6 +1346,13 @@ declare namespace CATAgentTask { interface AgentTask { /** Task ID. */ id: string; + /** + * Optimistic-concurrency version, returned by `get()`/`list()`. Pass the value you fetched back + * through `update()`/`remove()` so the system can detect the task was changed or recreated in between. + */ + generation?: string; + /** Optimistic-concurrency revision counter, paired with `generation`. */ + revision?: number; /** Task name. */ name: string; /** Cron expression. */ @@ -1302,10 +1375,10 @@ declare namespace CATAgentTask { modelId?: string; /** Existing conversation ID to continue. */ conversationId?: string; + /** Conversation generation captured when this task was bound to a conversation. */ + conversationGeneration?: string; /** Skills to load. */ skills?: "auto" | string[]; - /** Max tool-calling iterations (default: 10). */ - maxIterations?: number; // --- event mode fields --- /** UUID of the script that created this task. */ @@ -1358,11 +1431,18 @@ declare namespace CATAgentTask { /** Get a task by ID. */ get(id: string): Promise; - /** Update a task. */ + /** + * Update a task. `task` must include the `generation`/`revision` returned by `get()`/`list()` — + * spread the fetched task before applying your changes. Throws if the task changed or was + * recreated since you fetched it. + */ update(id: string, task: Partial): Promise; - /** Remove a task by ID. */ - remove(id: string): Promise; + /** + * Remove a task by ID. `task` must include the `generation`/`revision` returned by `get()`/`list()`, + * so a stale reference can't delete a task that was recreated with the same ID in the meantime. + */ + remove(id: string, task: Pick): Promise; /** Immediately trigger a task (regardless of cron schedule). */ runNow(id: string): Promise; @@ -1468,7 +1548,7 @@ declare namespace CATAgentModel { /** User-defined display name (e.g. "GPT-4o", "Claude Sonnet"). */ name: string; /** LLM provider. */ - provider: "openai" | "anthropic"; + provider: "openai" | "anthropic" | "zhipu"; /** API base URL. */ apiBaseUrl: string; /** Model identifier sent to the provider API. */ diff --git a/src/types/scriptcat.zh-CN.d.ts b/src/types/scriptcat.zh-CN.d.ts index 4c99999cb..9529c96de 100644 --- a/src/types/scriptcat.zh-CN.d.ts +++ b/src/types/scriptcat.zh-CN.d.ts @@ -3,7 +3,6 @@ // 此文件为 scriptcat.d.ts 的中文翻译版本,包含所有 GM_*/CAT_*/CAT.agent API。 // 如需接入,请在 tsconfig.json 中替换或追加此文件。 // ============================================================================ - // @copyright https://github.com/silverwzw/Tampermonkey-Typescript-Declaration declare const unsafeWindow: Window; @@ -849,7 +848,8 @@ declare namespace CATAgent { /** 描述工具参数的 JSON Schema。 */ parameters: Record; /** LLM 调用此工具时执行的处理函数。 */ - handler: (args: Record) => Promise; + /** LLM 调用工具时执行;批次 signal 中止后应立即停止副作用。 */ + handler: (args: Record, signal: AbortSignal) => Promise; } /** @@ -868,8 +868,6 @@ declare namespace CATAgent { system?: string; /** 模型 ID,省略则使用默认模型。 */ model?: string; - /** 工具调用循环最大迭代次数(默认 20)。 */ - maxIterations?: number; /** 加载的 Skill:`"auto"` 加载全部已安装 Skill,或指定名称数组。 */ skills?: "auto" | string[]; /** 带内联处理函数的工具,在此对话生命周期内可用。 */ @@ -886,6 +884,8 @@ declare namespace CATAgent { ephemeral?: boolean; /** 是否启用 prompt caching,默认 true。 */ cache?: boolean; + /** 在页面断开后仍让对话继续在 Service Worker 中运行。 */ + background?: boolean; } /** 单次 `chat()` / `chatStream()` 调用的选项。 */ @@ -937,9 +937,18 @@ declare namespace CATAgent { /** 本轮中的工具调用。 */ toolCalls?: ToolCallInfo[]; /** Token 用量。 */ - usage?: { inputTokens: number; outputTokens: number }; + usage?: { + inputTokens: number; + outputTokens: number; + cacheCreationInputTokens?: number; + cacheReadInputTokens?: number; + }; + /** 总响应时长(毫秒)。 */ + durationMs?: number; /** 当回复由命令处理器产生(而非 LLM)时为 `true`。 */ command?: boolean; + /** 生成数据丢失等非致命警告(如生成的图片保存失败)。 */ + warning?: string; } /** 通过 `chatStream()` 流式返回的单个数据块。 */ @@ -948,26 +957,76 @@ declare namespace CATAgent { * 数据块类型: * - `"content_delta"` — 增量文本 * - `"thinking_delta"` — 增量思考/推理 - * - `"tool_call"` — 工具调用事件 + * - `"tool_call"` — 工具调用事件(开始或参数增量) + * - `"tool_call_complete"` — 工具调用执行完成,携带结果/状态/附件 * - `"content_block"` — 完整的非文本内容块 + * - `"new_message"` — 当前轮次结束,下一轮 assistant 消息即将开始 + * - `"system_warning"` — 生成数据丢失等非致命警告 * - `"done"` — 流结束 * - `"error"` — 发生错误 */ - type: "content_delta" | "thinking_delta" | "tool_call" | "content_block" | "done" | "error"; + type: + | "content_delta" + | "thinking_delta" + | "tool_call" + | "tool_call_complete" + | "content_block" + | "new_message" + | "system_warning" + | "done" + | "error"; /** 文本增量(用于 content_delta / thinking_delta)。 */ content?: string; /** 完整内容块(用于 content_block)。 */ block?: ContentBlock; - /** 工具调用信息(用于 tool_call)。 */ + /** 工具调用信息(用于 tool_call / tool_call_complete)。 */ toolCall?: ToolCallInfo; /** Token 用量(用于 done)。 */ - usage?: { inputTokens: number; outputTokens: number }; + usage?: { + inputTokens: number; + outputTokens: number; + cacheCreationInputTokens?: number; + cacheReadInputTokens?: number; + }; + /** 总响应时长(毫秒)。 */ + durationMs?: number; /** 错误信息(用于 error)。 */ error?: string; - /** 错误分类码:`"rate_limit"` | `"auth"` | `"tool_timeout"` | `"max_iterations"` | `"api_error"` */ + /** 错误分类码:`"rate_limit"` | `"auth"` | `"tool_timeout"` | `"context_too_large"` | `"api_error"` */ errorCode?: string; /** 当数据块由命令处理器产生时为 `true`。 */ command?: boolean; + /** 警告文本(用于 `"system_warning"`)。 */ + warning?: string; + } + + /** 附加到后台对话时返回的初始状态快照。 */ + interface SyncStreamChunk { + type: "sync"; + /** 附加前已累计的 assistant 输出。 */ + streamingMessage?: { + content: string; + thinking?: string; + toolCalls: ToolCallInfo[]; + }; + /** 会话正在等待输入时的 ask_user 请求。 */ + pendingAskUser?: { + id: string; + question: string; + options?: string[]; + optionValues?: string[]; + multiple?: boolean; + allowCustom?: boolean; + }; + /** 当前任务快照。 */ + tasks: Array<{ + id: string; + subject: string; + status: "pending" | "in_progress" | "completed"; + description?: string; + }>; + /** attach() 返回的最终快照状态。 */ + status: "running" | "done" | "error"; } // ---- 聊天消息 ---- @@ -990,6 +1049,10 @@ declare namespace CATAgent { toolCallId?: string; /** 错误信息(当轮次出错时)。 */ error?: string; + /** 错误分类码(例如 `context_too_large`)。 */ + errorCode?: string; + /** 与此消息关联的非致命警告。 */ + warning?: string; /** 生成此消息使用的模型 ID。 */ modelId?: string; /** 此消息的 Token 用量。 */ @@ -1031,6 +1094,9 @@ declare namespace CATAgent { /** 发送消息并接收流式响应。 */ chatStream(content: MessageContent, options?: ChatOptions): Promise>; + /** 附加到后台运行中的对话,接收初始快照与后续流式数据。 */ + attach(): Promise>; + /** 获取此对话中的所有消息。 */ getMessages(): Promise; @@ -1287,6 +1353,13 @@ declare namespace CATAgentTask { interface AgentTask { /** 任务 ID。 */ id: string; + /** + * 乐观并发版本号,由 `get()`/`list()` 返回。调用 `update()`/`remove()` 时须传回这个取到的值, + * 系统才能识别出该任务是否已在此期间被修改或重建。 + */ + generation?: string; + /** 乐观并发修订号,与 `generation` 配对使用。 */ + revision?: number; /** 任务名称。 */ name: string; /** Cron 表达式。 */ @@ -1309,10 +1382,10 @@ declare namespace CATAgentTask { modelId?: string; /** 要续接的已有对话 ID。 */ conversationId?: string; + /** 任务绑定对话时记录的会话 generation。 */ + conversationGeneration?: string; /** 加载的 Skill。 */ skills?: "auto" | string[]; - /** 工具调用最大迭代次数(默认 10)。 */ - maxIterations?: number; // --- event 模式字段 --- /** 创建此任务的脚本 UUID。 */ @@ -1365,11 +1438,17 @@ declare namespace CATAgentTask { /** 根据 ID 获取任务。 */ get(id: string): Promise; - /** 更新任务。 */ + /** + * 更新任务。`task` 必须携带 `get()`/`list()` 返回的 `generation`/`revision`——先展开取到的任务对象 + * 再应用改动。若任务在此期间被修改或重建,会抛出错误。 + */ update(id: string, task: Partial): Promise; - /** 根据 ID 删除任务。 */ - remove(id: string): Promise; + /** + * 根据 ID 删除任务。`task` 必须携带 `get()`/`list()` 返回的 `generation`/`revision`, + * 避免持有旧引用的调用方删掉同 ID 被重建后的新任务。 + */ + remove(id: string, task: Pick): Promise; /** 立即触发任务(不受 cron 计划限制)。 */ runNow(id: string): Promise; @@ -1475,7 +1554,7 @@ declare namespace CATAgentModel { /** 用户自定义显示名称(如 "GPT-4o"、"Claude Sonnet")。 */ name: string; /** LLM 提供商。 */ - provider: "openai" | "anthropic"; + provider: "openai" | "anthropic" | "zhipu"; /** API 基础 URL。 */ apiBaseUrl: string; /** 发送给提供商 API 的模型标识符。 */