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:
can1357
2026-02-04 08:26:17 +01:00
parent c95f502464
commit 7d2c8efc5a
7 changed files with 224 additions and 88 deletions
+8
View File
@@ -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
+1 -1
View File
@@ -47,6 +47,6 @@
},
"devDependencies": {
"@sinclair/typebox": "^0.34.48",
"@types/node": "^25.0.10"
"@types/node": "^25.2.0"
}
}
+92 -53
View File
@@ -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;
}
/**
+6
View File
@@ -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>,
+117 -3
View File
@@ -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";
-9
View File
@@ -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)