diff --git a/packages/cli/src/session/message-v2.ts b/packages/cli/src/session/message-v2.ts index aa5e1e5..22cf424 100644 --- a/packages/cli/src/session/message-v2.ts +++ b/packages/cli/src/session/message-v2.ts @@ -17,6 +17,16 @@ import { type SystemError } from "bun" import type { Provider } from "@/provider/provider" export namespace MessageV2 { + export const PROVIDER_EXECUTED_METADATA_KEY = "providerExecuted" + export const TOOL_EXECUTION_ABORTED = "Tool execution aborted" + + export function toolProviderMetadata(part: ToolPart) { + const metadata = Object.fromEntries( + Object.entries(part.metadata ?? {}).filter(([key]) => key !== PROVIDER_EXECUTED_METADATA_KEY), + ) + return Object.keys(metadata).length ? metadata : undefined + } + export function hasVisibleOutput(parts: Part[]) { return parts.some((part) => (part.type === "text" && !!part.text.trim()) || part.type === "tool") } @@ -652,7 +662,7 @@ export namespace MessageV2 { toolCallId: part.callID, input: part.state.input, output, - ...(differentModel ? {} : { callProviderMetadata: part.metadata }), + ...(differentModel ? {} : { callProviderMetadata: toolProviderMetadata(part) }), }) } if (part.state.status === "error") @@ -662,7 +672,7 @@ export namespace MessageV2 { toolCallId: part.callID, input: part.state.input, errorText: part.state.error, - ...(differentModel ? {} : { callProviderMetadata: part.metadata }), + ...(differentModel ? {} : { callProviderMetadata: toolProviderMetadata(part) }), }) // Handle pending/running tool calls to prevent dangling tool_use blocks // Anthropic/Claude APIs require every tool_use to have a corresponding tool_result @@ -673,7 +683,7 @@ export namespace MessageV2 { toolCallId: part.callID, input: part.state.input, errorText: "[Tool execution was interrupted]", - ...(differentModel ? {} : { callProviderMetadata: part.metadata }), + ...(differentModel ? {} : { callProviderMetadata: toolProviderMetadata(part) }), }) } if (part.type === "reasoning") { diff --git a/packages/cli/src/session/processor.ts b/packages/cli/src/session/processor.ts index ef88729..b6a73f9 100644 --- a/packages/cli/src/session/processor.ts +++ b/packages/cli/src/session/processor.ts @@ -186,7 +186,10 @@ export namespace SessionProcessor { start: Date.now(), }, }, - metadata: value.providerMetadata, + metadata: { + ...value.providerMetadata, + [MessageV2.PROVIDER_EXECUTED_METADATA_KEY]: value.providerExecuted === true, + }, }) toolcalls[value.toolCallId] = part as MessageV2.ToolPart parts.add(part.id) @@ -504,7 +507,7 @@ export namespace SessionProcessor { state: { ...part.state, status: "error", - error: "Tool execution aborted", + error: MessageV2.TOOL_EXECUTION_ABORTED, time: { start: Date.now(), end: Date.now(), diff --git a/packages/cli/src/session/prompt.ts b/packages/cli/src/session/prompt.ts index 074fd68..1885473 100644 --- a/packages/cli/src/session/prompt.ts +++ b/packages/cli/src/session/prompt.ts @@ -62,6 +62,27 @@ const STRUCTURED_OUTPUT_SYSTEM_PROMPT = `IMPORTANT: The user has requested struc export namespace SessionPrompt { const log = Log.create({ service: "session.prompt" }) + export function hasToolCalls(parts: MessageV2.Part[]) { + return parts.some( + (part) => + part.type === "tool" && + part.metadata?.[MessageV2.PROVIDER_EXECUTED_METADATA_KEY] !== true && + (part.state.status === "completed" || + (part.state.status === "error" && + (part.state.error !== MessageV2.TOOL_EXECUTION_ABORTED || + part.metadata?.[MessageV2.PROVIDER_EXECUTED_METADATA_KEY] === false))), + ) + } + + export function isModelFinished(finish?: string) { + return !!finish && !["tool-calls", "unknown"].includes(finish) + } + + export async function missingStructuredOutput(message: MessageV2.Assistant, load: () => Promise) { + if (!isModelFinished(message.finish) || message.error) return false + return !hasToolCalls(await load()) + } + const state = Instance.state( () => { const data: Record< @@ -332,9 +353,11 @@ export namespace SessionPrompt { } if (!lastUser) throw new Error("No user message found in stream. This should never happen.") + const lastAssistantMsg = msgs.findLast((msg) => msg.info.id === lastAssistant?.id) if ( - lastAssistant?.finish && - !["tool-calls", "unknown"].includes(lastAssistant.finish) && + lastAssistant && + isModelFinished(lastAssistant.finish) && + !hasToolCalls(lastAssistantMsg?.parts ?? []) && lastUser.id < lastAssistant.id ) { log.info("exiting loop", { sessionID }) @@ -727,19 +750,16 @@ export namespace SessionPrompt { break } - // Check if model finished (finish reason is not "tool-calls" or "unknown") - const modelFinished = processor.message.finish && !["tool-calls", "unknown"].includes(processor.message.finish) - - if (modelFinished && !processor.message.error) { - if (format.type === "json_schema") { - // Model stopped without calling StructuredOutput tool - processor.message.error = new MessageV2.StructuredOutputError({ - message: "Model did not produce structured output", - retries: 0, - }).toObject() - await Session.updateMessage(processor.message) - break - } + if ( + format.type === "json_schema" && + (await missingStructuredOutput(processor.message, () => MessageV2.parts(processor.message.id))) + ) { + processor.message.error = new MessageV2.StructuredOutputError({ + message: "Model did not produce structured output", + retries: 0, + }).toObject() + await Session.updateMessage(processor.message) + break } if (result === "stop") break diff --git a/packages/cli/src/tool/task.ts b/packages/cli/src/tool/task.ts index 5e1a0ae..d45db08 100644 --- a/packages/cli/src/tool/task.ts +++ b/packages/cli/src/tool/task.ts @@ -25,6 +25,28 @@ const parameters = z.object({ command: z.string().describe("The command that triggered this task").optional(), }) +export function taskResultText(result: MessageV2.WithParts, sessionID: string) { + if (result.info.role === "assistant" && result.info.error) { + const data = result.info.error.data + const message = + data && typeof data === "object" && "message" in data && typeof data.message === "string" + ? data.message + : result.info.error.name + throw new Error(`Subagent failed (task_id: ${sessionID}): ${message}`) + } + const failed = result.parts.findLast((part) => part.type === "tool" && part.state.status === "error") + if (failed?.type === "tool" && failed.state.status === "error") { + throw new Error(`Subagent failed (task_id: ${sessionID}): ${failed.state.error}`) + } + return result.parts.findLast((part) => part.type === "text")?.text ?? "" +} + +export async function completeTask(result: MessageV2.WithParts, subagentSessionID: string, parentSessionID: string) { + const text = result.parts.findLast((part) => part.type === "text")?.text ?? "" + await Plugin.trigger("agent.subtask.complete", { subagentSessionID, parentSessionID }, { result: text }) + return taskResultText(result, subagentSessionID) +} + export const TaskTool = Tool.define("task", async (ctx) => { const agents = await Agent.list().then((x) => x.filter((a) => a.mode !== "primary")) @@ -150,13 +172,7 @@ export const TaskTool = Tool.define("task", async (ctx) => { parts: promptParts, }) - const text = result.parts.findLast((x) => x.type === "text")?.text ?? "" - - await Plugin.trigger( - "agent.subtask.complete", - { subagentSessionID: session.id, parentSessionID: ctx.sessionID }, - { result: text }, - ) + const text = await completeTask(result, session.id, ctx.sessionID) const output = [ `task_id: ${session.id} (for resuming to continue this task if needed)`, diff --git a/packages/cli/test/session/processor-tool-metadata.test.ts b/packages/cli/test/session/processor-tool-metadata.test.ts new file mode 100644 index 0000000..a5387a4 --- /dev/null +++ b/packages/cli/test/session/processor-tool-metadata.test.ts @@ -0,0 +1,106 @@ +import { describe, expect, spyOn, test } from "bun:test" +import { Agent } from "../../src/agent/agent" +import { Identifier } from "../../src/id/id" +import { Instance } from "../../src/project/instance" +import { Provider } from "../../src/provider/provider" +import { Session } from "../../src/session" +import { LLM } from "../../src/session/llm" +import { SessionProcessor } from "../../src/session/processor" +import { MessageV2 } from "../../src/session/message-v2" +import { tmpdir } from "../fixture/fixture" + +describe("session processor tool metadata", () => { + test("persists provider-executed attribution from the stream", async () => { + await using tmp = await tmpdir({ + config: { + enabled_providers: ["alibaba"], + provider: { alibaba: { options: { apiKey: "test-key" } } }, + }, + }) + + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const session = await Session.create({ title: "Provider tool metadata fixture" }) + const agent = await Agent.get("build") + const model = await Provider.getModel("alibaba", "qwen-plus") + const user = (await Session.updateMessage({ + id: Identifier.ascending("message"), + sessionID: session.id, + role: "user", + time: { created: Date.now() }, + agent: agent.name, + model: { providerID: model.providerID, modelID: model.id }, + })) as MessageV2.User + const assistant = (await Session.updateMessage({ + id: Identifier.ascending("message"), + sessionID: session.id, + role: "assistant", + parentID: user.id, + modelID: model.id, + providerID: model.providerID, + mode: agent.name, + agent: agent.name, + path: { cwd: tmp.path, root: tmp.path }, + cost: 0, + tokens: { input: 0, output: 0, reasoning: 0, cache: { read: 0, write: 0 } }, + time: { created: Date.now() }, + })) as MessageV2.Assistant + + const stream = spyOn(LLM, "stream").mockResolvedValue({ + fullStream: (async function* () { + yield { type: "tool-input-start", id: "call_1", toolName: "server_tool" } + yield { + type: "tool-call", + toolCallId: "call_1", + toolName: "server_tool", + input: {}, + providerExecuted: true, + } + yield { type: "tool-input-start", id: "call_2", toolName: "server_tool" } + yield { + type: "tool-call", + toolCallId: "call_2", + toolName: "server_tool", + input: { second: true }, + providerExecuted: false, + providerMetadata: { providerExecuted: true }, + } + yield { + type: "finish-step", + finishReason: "stop", + usage: { inputTokens: 1, outputTokens: 1, totalTokens: 2 }, + } + })(), + } as unknown as Awaited>) + + try { + const processor = SessionProcessor.create({ + assistantMessage: assistant, + sessionID: session.id, + model, + abort: new AbortController().signal, + }) + await processor.process({ + user, + sessionID: session.id, + model, + agent, + abort: new AbortController().signal, + system: [], + messages: [], + tools: {}, + }) + + const parts = (await Session.messages({ sessionID: session.id })).flatMap((message) => message.parts) + const first = parts.find((part) => part.type === "tool" && part.callID === "call_1") + const second = parts.find((part) => part.type === "tool" && part.callID === "call_2") + expect(first?.type === "tool" ? first.metadata?.providerExecuted : undefined).toBe(true) + expect(second?.type === "tool" ? second.metadata?.providerExecuted : undefined).toBe(false) + } finally { + stream.mockRestore() + } + }, + }) + }) +}) diff --git a/packages/cli/test/session/prompt-tool-loop.test.ts b/packages/cli/test/session/prompt-tool-loop.test.ts new file mode 100644 index 0000000..58a10d7 --- /dev/null +++ b/packages/cli/test/session/prompt-tool-loop.test.ts @@ -0,0 +1,245 @@ +import { describe, expect, spyOn, test } from "bun:test" +import path from "path" +import { Instance } from "../../src/project/instance" +import { Session } from "../../src/session" +import { SessionPrompt } from "../../src/session/prompt" +import { MessageV2 } from "../../src/session/message-v2" +import { LLM } from "../../src/session/llm" +import { tmpdir } from "../fixture/fixture" + +function tool(input: { + status?: "pending" | "running" | "completed" | "error" + providerExecuted?: boolean + error?: string +}) { + const status = input.status ?? "completed" + return { + id: "part_1", + sessionID: "session_1", + messageID: "message_1", + type: "tool", + callID: "call_1", + tool: "read", + metadata: + input.providerExecuted === undefined + ? undefined + : { [MessageV2.PROVIDER_EXECUTED_METADATA_KEY]: input.providerExecuted }, + state: + status === "pending" + ? { status, input: {}, raw: "{}" } + : status === "running" + ? { status, input: {}, time: { start: 1 } } + : status === "completed" + ? { + status, + input: {}, + output: "ok", + title: "read", + metadata: {}, + time: { start: 1, end: 2 }, + } + : { + status, + input: {}, + error: input.error ?? "failed", + time: { start: 1, end: 2 }, + }, + } as MessageV2.ToolPart +} + +function stream(chunks: unknown[]) { + const body = + chunks + .map((chunk) => `data: ${JSON.stringify(chunk)}`) + .concat("data: [DONE]") + .join("\n\n") + "\n\n" + return new Response(body, { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + }) +} + +describe("session prompt tool-call continuation", () => { + test("detects non-provider-executed tool calls that require another model turn", () => { + expect(SessionPrompt.hasToolCalls([tool({ status: "pending" })])).toBe(false) + expect(SessionPrompt.hasToolCalls([tool({ status: "running" })])).toBe(false) + expect(SessionPrompt.hasToolCalls([tool({ status: "completed" })])).toBe(true) + expect(SessionPrompt.hasToolCalls([tool({ status: "error" })])).toBe(true) + expect(SessionPrompt.hasToolCalls([tool({ status: "error", error: MessageV2.TOOL_EXECUTION_ABORTED })])).toBe(false) + expect( + SessionPrompt.hasToolCalls([ + tool({ status: "error", error: MessageV2.TOOL_EXECUTION_ABORTED, providerExecuted: false }), + ]), + ).toBe(true) + }) + + test("ignores provider-executed tool calls", () => { + expect(SessionPrompt.hasToolCalls([tool({ providerExecuted: true })])).toBe(false) + }) + + test("checks stored parts only after an error-free model finish", async () => { + const message = { id: "message_1", finish: "tool-calls" } as MessageV2.Assistant + const calls: string[] = [] + const load = async () => { + calls.push("load") + return [] + } + expect(SessionPrompt.isModelFinished("tool-calls")).toBe(false) + expect(SessionPrompt.isModelFinished("unknown")).toBe(false) + expect(SessionPrompt.isModelFinished("stop")).toBe(true) + expect(await SessionPrompt.missingStructuredOutput(message, load)).toBe(false) + expect( + await SessionPrompt.missingStructuredOutput( + { ...message, finish: "stop", error: { name: "error" } } as unknown as MessageV2.Assistant, + load, + ), + ).toBe(false) + expect(calls).toHaveLength(0) + expect(await SessionPrompt.missingStructuredOutput({ ...message, finish: "stop" }, load)).toBe(true) + expect(calls).toHaveLength(1) + }) + + test("stops after partial tool input when the provider reports stop", async () => { + await using tmp = await tmpdir({ + config: { + enabled_providers: ["alibaba"], + provider: { alibaba: { options: { apiKey: "test-key" } } }, + }, + }) + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const session = await Session.create({ title: "Partial tool input" }) + const calls: string[] = [] + const stream = spyOn(LLM, "stream").mockImplementation(async () => { + calls.push("model") + if (calls.length > 1) throw new Error("partial input caused an extra model turn") + return { + fullStream: (async function* () { + yield { type: "start-step" } + yield { type: "tool-input-start", id: "partial", toolName: "read" } + yield { type: "finish-step", finishReason: "stop", usage: { inputTokens: 1, outputTokens: 1 } } + })(), + } as unknown as Awaited> + }) + try { + await SessionPrompt.prompt({ + sessionID: session.id, + model: { providerID: "alibaba", modelID: "qwen-plus" }, + parts: [{ type: "text", text: "Read a file" }], + }) + expect(calls).toHaveLength(1) + const parts = (await Session.messages({ sessionID: session.id })).flatMap((message) => message.parts) + expect(parts.some((part) => part.type === "tool" && part.state.status === "error")).toBe(true) + } finally { + stream.mockRestore() + } + }, + }) + }) + + test("continues structured output after a provider reports stop with a local tool call", async () => { + const requests: Record[] = [] + const server = Bun.serve({ + port: 0, + async fetch(request) { + requests.push((await request.json()) as Record) + const call = requests.length === 1 ? "read" : "StructuredOutput" + const args = requests.length === 1 ? { filePath: "aictrl.json" } : { result: "follow-up reached" } + return stream([ + { + id: `chatcmpl-${requests.length}`, + object: "chat.completion.chunk", + choices: [{ index: 0, delta: { role: "assistant" }, finish_reason: null }], + }, + { + id: `chatcmpl-${requests.length}`, + object: "chat.completion.chunk", + choices: [ + { + index: 0, + delta: { + tool_calls: [ + { + index: 0, + id: `call_${requests.length}`, + type: "function", + function: { name: call, arguments: JSON.stringify(args) }, + }, + ], + }, + finish_reason: null, + }, + ], + }, + { + id: `chatcmpl-${requests.length}`, + object: "chat.completion.chunk", + choices: [{ index: 0, delta: {}, finish_reason: "stop" }], + usage: { prompt_tokens: 1, completion_tokens: 1, total_tokens: 2 }, + }, + ]) + }, + }) + + try { + await using tmp = await tmpdir({ + init: async (dir) => { + await Bun.write( + path.join(dir, "aictrl.json"), + JSON.stringify({ + $schema: "https://aictrl.ai/config.json", + enabled_providers: ["alibaba"], + provider: { + alibaba: { + options: { + apiKey: "test-key", + baseURL: `${server.url.origin}/v1`, + }, + }, + }, + }), + ) + }, + }) + + await Instance.provide({ + directory: tmp.path, + fn: async () => { + const session = await Session.create({ title: "Tool continuation fixture" }) + const result = await SessionPrompt.prompt({ + sessionID: session.id, + model: { providerID: "alibaba", modelID: "qwen-plus" }, + parts: [{ type: "text", text: "Return structured output after using a tool." }], + format: { + type: "json_schema", + schema: { + type: "object", + properties: { result: { type: "string" } }, + required: ["result"], + }, + retryCount: 0, + }, + }) + + expect(requests).toHaveLength(2) + expect(result.info.role).toBe("assistant") + if (result.info.role !== "assistant") throw new Error("Expected assistant result") + expect(result.info.structured).toEqual({ result: "follow-up reached" }) + expect(result.info.error).toBeUndefined() + + const messages = await Session.messages({ sessionID: session.id }) + const first = messages.find( + (message) => + message.info.role === "assistant" && + message.parts.some((part) => part.type === "tool" && part.tool === "read"), + ) + expect(first?.info.role === "assistant" ? first.info.finish : undefined).toBe("stop") + expect(first?.info.role === "assistant" ? first.info.error : undefined).toBeUndefined() + }, + }) + } finally { + server.stop() + } + }, 15_000) +}) diff --git a/packages/cli/test/tool/task.test.ts b/packages/cli/test/tool/task.test.ts new file mode 100644 index 0000000..ebff516 --- /dev/null +++ b/packages/cli/test/tool/task.test.ts @@ -0,0 +1,129 @@ +import { describe, expect, spyOn, test } from "bun:test" +import { MessageV2 } from "../../src/session/message-v2" +import { Plugin } from "../../src/plugin" +import { completeTask, taskResultText } from "../../src/tool/task" + +function result(input: { error?: MessageV2.Assistant["error"]; parts?: MessageV2.Part[] }): MessageV2.WithParts { + return { + info: { + id: "message_1", + sessionID: "child_1", + role: "assistant", + time: { created: 1 }, + error: input.error, + parentID: "user_1", + modelID: "model", + providerID: "provider", + mode: "build", + agent: "build", + path: { cwd: "/workspace", root: "/workspace" }, + cost: 0, + tokens: { input: 0, output: 0, reasoning: 0, cache: { read: 0, write: 0 } }, + }, + parts: input.parts ?? [], + } +} + +describe("task tool child result", () => { + test("surfaces a child session error with child attribution", () => { + expect(() => + taskResultText( + result({ + error: new MessageV2.APIError({ + message: "provider unavailable", + isRetryable: false, + }).toObject(), + }), + "child_1", + ), + ).toThrow("Subagent failed (task_id: child_1): provider unavailable") + }) + + test("falls back to the child error name when data is absent or malformed", () => { + for (const data of [undefined, null, "bad data"]) { + expect(() => + taskResultText( + result({ error: { name: "ChildFailed", data } as unknown as MessageV2.Assistant["error"] }), + "child_1", + ), + ).toThrow("Subagent failed (task_id: child_1): ChildFailed") + } + }) + + test("notifies plugins before propagating a child failure", async () => { + const events: string[] = [] + const trigger = spyOn(Plugin, "trigger").mockImplementation(async (name, _input, output) => { + events.push(name) + return output + }) + try { + await expect( + completeTask( + result({ error: { name: "ChildFailed" } as unknown as MessageV2.Assistant["error"] }), + "child_1", + "parent_1", + ), + ).rejects.toThrow("Subagent failed (task_id: child_1): ChildFailed") + expect(events).toEqual(["agent.subtask.complete"]) + expect(trigger).toHaveBeenCalledWith( + "agent.subtask.complete", + { subagentSessionID: "child_1", parentSessionID: "parent_1" }, + { result: "" }, + ) + } finally { + trigger.mockRestore() + } + }) + + test("surfaces the last failed child tool instead of returning empty success", () => { + expect(() => + taskResultText( + result({ + parts: [ + { + id: "part_1", + sessionID: "child_1", + messageID: "message_1", + type: "tool", + callID: "call_1", + tool: "read", + state: { + status: "error", + input: {}, + error: "permission denied", + time: { start: 1, end: 2 }, + }, + }, + ], + }), + "child_1", + ), + ).toThrow("Subagent failed (task_id: child_1): permission denied") + }) + + test("returns the last child text when execution succeeds", () => { + expect( + taskResultText( + result({ + parts: [ + { + id: "part_1", + sessionID: "child_1", + messageID: "message_1", + type: "text", + text: "first", + }, + { + id: "part_2", + sessionID: "child_1", + messageID: "message_1", + type: "text", + text: "done", + }, + ], + }), + "child_1", + ), + ).toBe("done") + }) +})