feat(agent): added concurrency modes and parallel tool execution with interrupt handling
- Added `concurrency` option to `AgentTool` interface supporting 'shared' (default, parallel execution) and 'exclusive' (solo execution) modes. - Implemented parallel execution of shared tools within a single agent turn with ordered result emission. - Refactored tool execution to support concurrent scheduling with interrupt handling via steering messages and AbortSignal. - Added steering message checking mechanism that interrupts tool execution when user messages arrive during tool processing. - Removed HTTP proxy setup code from stream.ts that was conditionally setting up undici global dispatcher.
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -47,6 +47,6 @@
|
||||
},
|
||||
"devDependencies": {
|
||||
"@sinclair/typebox": "^0.34.48",
|
||||
"@types/node": "^25.0.10"
|
||||
"@types/node": "^25.2.0"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<AssistantMessage["content"][number], { type: "toolCall" }>;
|
||||
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<void> | 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<void> => {
|
||||
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<any> | undefined,
|
||||
isError: false,
|
||||
skipped: false,
|
||||
}));
|
||||
|
||||
const runTool = async (record: (typeof records)[number], index: number): Promise<void> => {
|
||||
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<void> = Promise.resolve();
|
||||
let sharedTasks: Promise<void>[] = [];
|
||||
const tasks: Promise<void>[] = [];
|
||||
|
||||
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<AssistantMessage["content"][number], { type: "toolCall" }>,
|
||||
stream: EventStream<AgentEvent, AgentMessage[]>,
|
||||
): ToolResultMessage {
|
||||
const result: AgentToolResult<any> = {
|
||||
function createSkippedToolResult(): AgentToolResult<any> {
|
||||
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;
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -206,6 +206,12 @@ export interface AgentTool<TParameters extends TSchema = TSchema, TDetails = any
|
||||
hidden?: boolean;
|
||||
/** If true, tool execution ignores abort signals (runs to completion) */
|
||||
nonAbortable?: boolean;
|
||||
/**
|
||||
* Concurrency mode for tool scheduling when multiple calls are in one turn.
|
||||
* - "shared": can run alongside other shared tools (default)
|
||||
* - "exclusive": runs alone; other tools wait until it finishes
|
||||
*/
|
||||
concurrency?: "shared" | "exclusive";
|
||||
execute: (
|
||||
toolCallId: string,
|
||||
params: Static<TParameters>,
|
||||
|
||||
@@ -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<string, number> = {};
|
||||
const finishTimes: Record<string, number> = {};
|
||||
const { promise: slowContinue, resolve: slowResolve } = Promise.withResolvers<void>();
|
||||
const { promise: slowStarted, resolve: slowStartedResolve } = Promise.withResolvers<void>();
|
||||
const { promise: fastFinished, resolve: fastFinishedResolve } = Promise.withResolvers<void>();
|
||||
|
||||
const tool: AgentTool<typeof toolSchema, { value: string }> = {
|
||||
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<AgentEvent, { type: "message_start" }> =>
|
||||
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<void>();
|
||||
const tool: AgentTool<typeof toolSchema, { value: string }> = {
|
||||
name: "echo",
|
||||
label: "Echo",
|
||||
description: "Echo tool",
|
||||
parameters: toolSchema,
|
||||
async execute(_toolCallId, params, signal) {
|
||||
if (params.value === "second") {
|
||||
await new Promise<void>((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 [];
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user