diff --git a/packages/agent/CHANGELOG.md b/packages/agent/CHANGELOG.md index 671206816..b4b54a593 100644 --- a/packages/agent/CHANGELOG.md +++ b/packages/agent/CHANGELOG.md @@ -1,6 +1,14 @@ # Changelog ## [Unreleased] +### Added + +- Added `concurrency` option to `AgentTool` to control tool scheduling: "shared" (default, runs in parallel) or "exclusive" (runs alone) +- Implemented parallel execution of shared tools within a single agent turn for improved performance + +### Changed + +- Refactored tool execution to support concurrent scheduling with proper interrupt handling and steering message checks ## [9.2.2] - 2026-01-31 diff --git a/packages/agent/package.json b/packages/agent/package.json index 176140df8..94efe0b50 100644 --- a/packages/agent/package.json +++ b/packages/agent/package.json @@ -47,6 +47,6 @@ }, "devDependencies": { "@sinclair/typebox": "^0.34.48", - "@types/node": "^25.0.10" + "@types/node": "^25.2.0" } } diff --git a/packages/agent/src/agent-loop.ts b/packages/agent/src/agent-loop.ts index 8ad3f37cd..8948a4043 100644 --- a/packages/agent/src/agent-loop.ts +++ b/packages/agent/src/agent-loop.ts @@ -372,17 +372,58 @@ async function executeToolCalls( getToolContext?: AgentLoopConfig["getToolContext"], interruptMode: AgentLoopConfig["interruptMode"] = "immediate", ): Promise<{ toolResults: ToolResultMessage[]; steeringMessages?: AgentMessage[] }> { - const toolCalls = assistantMessage.content.filter(c => c.type === "toolCall"); + type ToolCallContent = Extract; + const toolCalls = assistantMessage.content.filter((c): c is ToolCallContent => c.type === "toolCall"); const results: ToolResultMessage[] = []; let steeringMessages: AgentMessage[] | undefined; const shouldInterruptImmediately = interruptMode !== "wait"; const toolCallInfos = toolCalls.map(call => ({ id: call.id, name: call.name })); const batchId = `${assistantMessage.timestamp ?? Date.now()}_${toolCalls[0]?.id ?? "batch"}`; + const steeringAbortController = new AbortController(); + const toolSignal = signal + ? AbortSignal.any([signal, steeringAbortController.signal]) + : steeringAbortController.signal; + const interruptState = { triggered: false }; + let steeringCheck: Promise | null = null; - for (let index = 0; index < toolCalls.length; index++) { - const toolCall = toolCalls[index]; - const tool = tools?.find(t => t.name === toolCall.name); + const checkSteering = async (): Promise => { + if (!shouldInterruptImmediately || !getSteeringMessages || interruptState.triggered) { + return; + } + if (steeringCheck) { + await steeringCheck; + return; + } + steeringCheck = (async () => { + const steering = await getSteeringMessages(); + if (steering.length > 0) { + steeringMessages = steering; + interruptState.triggered = true; + steeringAbortController.abort(); + } + })().finally(() => { + steeringCheck = null; + }); + await steeringCheck; + }; + const records = toolCalls.map(toolCall => ({ + toolCall, + tool: tools?.find(t => t.name === toolCall.name), + started: false, + result: undefined as AgentToolResult | undefined, + isError: false, + skipped: false, + })); + + const runTool = async (record: (typeof records)[number], index: number): Promise => { + if (interruptState.triggered) { + record.skipped = true; + return; + } + + const { toolCall, tool } = record; + record.started = true; stream.push({ type: "tool_execution_start", toolCallId: toolCall.id, @@ -397,7 +438,6 @@ async function executeToolCalls( if (!tool) throw new Error(`Tool ${toolCall.name} not found`); const validatedArgs = validateToolArguments(tool, toolCall); - const toolContext = getToolContext ? getToolContext({ batchId, @@ -409,8 +449,9 @@ async function executeToolCalls( result = await tool.execute( toolCall.id, validatedArgs, - tool.nonAbortable ? undefined : signal, + tool.nonAbortable ? undefined : toolSignal, partialResult => { + if (interruptState.triggered) return; stream.push({ type: "tool_execution_update", toolCallId: toolCall.id, @@ -429,6 +470,49 @@ async function executeToolCalls( isError = true; } + if (!interruptState.triggered) { + record.result = result; + record.isError = isError; + } else { + record.skipped = true; + } + + await checkSteering(); + }; + + let lastExclusive: Promise = Promise.resolve(); + let sharedTasks: Promise[] = []; + const tasks: Promise[] = []; + + for (let index = 0; index < records.length; index++) { + const record = records[index]; + const concurrency = record.tool?.concurrency ?? "shared"; + const start = concurrency === "exclusive" ? Promise.all([lastExclusive, ...sharedTasks]) : lastExclusive; + const task = start.then(() => runTool(record, index)); + tasks.push(task); + if (concurrency === "exclusive") { + lastExclusive = task; + sharedTasks = []; + } else { + sharedTasks.push(task); + } + } + + await Promise.allSettled(tasks); + + for (const record of records) { + const toolCall = record.toolCall; + const shouldSkip = record.skipped || record.result === undefined; + const result = shouldSkip ? createSkippedToolResult() : record.result; + const isError = shouldSkip ? true : record.isError; + if (!record.started) { + stream.push({ + type: "tool_execution_start", + toolCallId: toolCall.id, + toolName: toolCall.name, + args: toolCall.arguments, + }); + } stream.push({ type: "tool_execution_end", toolCallId: toolCall.id, @@ -450,61 +534,16 @@ async function executeToolCalls( results.push(toolResultMessage); stream.push({ type: "message_start", message: toolResultMessage }); stream.push({ type: "message_end", message: toolResultMessage }); - - // Check for steering messages - skip remaining tools if user interrupted - if (shouldInterruptImmediately && getSteeringMessages) { - const steering = await getSteeringMessages(); - if (steering.length > 0) { - steeringMessages = steering; - const remainingCalls = toolCalls.slice(index + 1); - for (const skipped of remainingCalls) { - results.push(skipToolCall(skipped, stream)); - } - break; - } - } } return { toolResults: results, steeringMessages }; } -function skipToolCall( - toolCall: Extract, - stream: EventStream, -): ToolResultMessage { - const result: AgentToolResult = { +function createSkippedToolResult(): AgentToolResult { + return { content: [{ type: "text", text: "Skipped due to queued user message." }], details: {}, }; - - stream.push({ - type: "tool_execution_start", - toolCallId: toolCall.id, - toolName: toolCall.name, - args: toolCall.arguments, - }); - stream.push({ - type: "tool_execution_end", - toolCallId: toolCall.id, - toolName: toolCall.name, - result, - isError: true, - }); - - const toolResultMessage: ToolResultMessage = { - role: "toolResult", - toolCallId: toolCall.id, - toolName: toolCall.name, - content: result.content, - details: {}, - isError: true, - timestamp: Date.now(), - }; - - stream.push({ type: "message_start", message: toolResultMessage }); - stream.push({ type: "message_end", message: toolResultMessage }); - - return toolResultMessage; } /** diff --git a/packages/agent/src/types.ts b/packages/agent/src/types.ts index 083115840..e24cce836 100644 --- a/packages/agent/src/types.ts +++ b/packages/agent/src/types.ts @@ -206,6 +206,12 @@ export interface AgentTool, diff --git a/packages/agent/test/agent-loop.test.ts b/packages/agent/test/agent-loop.test.ts index 29a028064..96ee46eec 100644 --- a/packages/agent/test/agent-loop.test.ts +++ b/packages/agent/test/agent-loop.test.ts @@ -9,7 +9,7 @@ import type { AgentToolContext, ToolCallContext, } from "@oh-my-pi/pi-agent-core/types"; -import type { AssistantMessage, Message, Model, UserMessage } from "@oh-my-pi/pi-ai"; +import type { AssistantMessage, Message, Model, ToolResultMessage, UserMessage } from "@oh-my-pi/pi-ai"; import { AssistantMessageEventStream } from "@oh-my-pi/pi-ai/utils/event-stream"; import { Type } from "@sinclair/typebox"; @@ -371,15 +371,128 @@ describe("agentLoop with AgentMessage", () => { } }); - it("should inject queued messages and skip remaining tool calls", async () => { + it("runs shared tools in parallel but emits ordered results", async () => { const toolSchema = Type.Object({ value: Type.String() }); - const executed: string[] = []; + const startTimes: Record = {}; + const finishTimes: Record = {}; + const { promise: slowContinue, resolve: slowResolve } = Promise.withResolvers(); + const { promise: slowStarted, resolve: slowStartedResolve } = Promise.withResolvers(); + const { promise: fastFinished, resolve: fastFinishedResolve } = Promise.withResolvers(); + const tool: AgentTool = { name: "echo", label: "Echo", description: "Echo tool", parameters: toolSchema, async execute(_toolCallId, params) { + if (params.value === "slow") { + startTimes.slow = performance.now(); + slowStartedResolve(); + await slowContinue; + finishTimes.slow = performance.now(); + } else { + await slowStarted; + startTimes.fast = performance.now(); + finishTimes.fast = performance.now(); + fastFinishedResolve(); + } + return { + content: [{ type: "text", text: `echoed: ${params.value}` }], + details: { value: params.value }, + }; + }, + }; + + const context: AgentContext = { + systemPrompt: "", + messages: [], + tools: [tool], + }; + + const userPrompt: AgentMessage = createUserMessage("start"); + + const config: AgentLoopConfig = { + model: createModel(), + convertToLlm: identityConverter, + }; + + let callIndex = 0; + const streamFn = () => { + const stream = new MockAssistantStream(); + queueMicrotask(() => { + if (callIndex === 0) { + const message = createAssistantMessage( + [ + { type: "toolCall", id: "tool-1", name: "echo", arguments: { value: "slow" } }, + { type: "toolCall", id: "tool-2", name: "echo", arguments: { value: "fast" } }, + ], + "toolUse", + ); + stream.push({ type: "done", reason: "toolUse", message }); + } else { + const message = createAssistantMessage([{ type: "text", text: "done" }]); + stream.push({ type: "done", reason: "stop", message }); + } + callIndex++; + }); + return stream; + }; + + const events: AgentEvent[] = []; + const stream = agentLoop([userPrompt], context, config, undefined, streamFn); + const streamTask = (async () => { + for await (const event of stream) { + events.push(event); + } + })(); + + await fastFinished; + slowResolve(); + await streamTask; + + expect(startTimes.fast).toBeDefined(); + expect(startTimes.slow).toBeDefined(); + expect(finishTimes.fast).toBeDefined(); + expect(finishTimes.slow).toBeDefined(); + expect(startTimes.fast).toBeLessThan(finishTimes.slow); + expect(finishTimes.fast).toBeLessThan(finishTimes.slow); + + const toolResultStarts = events.filter( + (e): e is Extract => + e.type === "message_start" && e.message.role === "toolResult", + ); + expect(toolResultStarts).toHaveLength(2); + expect((toolResultStarts[0].message as ToolResultMessage).toolCallId).toBe("tool-1"); + expect((toolResultStarts[1].message as ToolResultMessage).toolCallId).toBe("tool-2"); + }); + + it("should inject queued messages and skip remaining tool calls", async () => { + const toolSchema = Type.Object({ value: Type.String() }); + const executed: string[] = []; + const { promise: allowSecond, resolve: allowSecondResolve } = Promise.withResolvers(); + const tool: AgentTool = { + name: "echo", + label: "Echo", + description: "Echo tool", + parameters: toolSchema, + async execute(_toolCallId, params, signal) { + if (params.value === "second") { + await new Promise((resolve, reject) => { + if (signal?.aborted) { + reject(new Error("Tool aborted")); + return; + } + const onAbort = () => reject(new Error("Tool aborted")); + signal?.addEventListener("abort", onAbort, { once: true }); + allowSecond.then(() => { + signal?.removeEventListener("abort", onAbort); + resolve(); + }); + }); + if (signal?.aborted) { + throw new Error("Tool aborted"); + } + } executed.push(params.value); return { content: [{ type: "text", text: `ok:${params.value}` }], @@ -408,6 +521,7 @@ describe("agentLoop with AgentMessage", () => { // Return queued message after first tool executes if (executed.length === 1 && !queuedDelivered) { queuedDelivered = true; + allowSecondResolve(); return [queuedUserMessage]; } return []; diff --git a/packages/ai/src/providers/amazon-bedrock.ts b/packages/ai/src/providers/amazon-bedrock.ts index 3f0256cd8..146a4ce99 100644 --- a/packages/ai/src/providers/amazon-bedrock.ts +++ b/packages/ai/src/providers/amazon-bedrock.ts @@ -98,28 +98,6 @@ export const streamBedrock: StreamFunction<"bedrock-converse-stream"> = ( // in Node.js/Bun environment only if (typeof process !== "undefined" && (process.versions?.node || process.versions?.bun)) { config.region = config.region || process.env.AWS_REGION || process.env.AWS_DEFAULT_REGION; - - if ( - process.env.HTTP_PROXY || - process.env.HTTPS_PROXY || - process.env.NO_PROXY || - process.env.http_proxy || - process.env.https_proxy || - process.env.no_proxy - ) { - const nodeHttpHandler = await import("@smithy/node-http-handler"); - const proxyAgent = await import("proxy-agent"); - - const agent = new proxyAgent.ProxyAgent(); - - // Bedrock runtime uses NodeHttp2Handler by default since v3.798.0, which is based - // on `http2` module and has no support for http agent. - // Use NodeHttpHandler to support http agent. - config.requestHandler = new nodeHttpHandler.NodeHttpHandler({ - httpAgent: agent, - httpsAgent: agent, - }); - } } config.region = config.region || "us-east-1"; diff --git a/packages/ai/src/stream.ts b/packages/ai/src/stream.ts index f6d36c67f..dfab63ae6 100644 --- a/packages/ai/src/stream.ts +++ b/packages/ai/src/stream.ts @@ -30,15 +30,6 @@ import type { ToolChoice, } from "./types"; -// Set up http proxy according to env variables for `fetch` based SDKs in Node.js. -// Bun has builtin support for this. -if (typeof process !== "undefined" && process.versions?.node) { - import("undici").then(m => { - const { EnvHttpProxyAgent, setGlobalDispatcher } = m; - setGlobalDispatcher(new EnvHttpProxyAgent()); - }); -} - let cachedVertexAdcCredentialsExists: boolean | null = null; // Cached .env file contents (parsed once per process)