From 1cf405f032ab2f4c7e908f23e950a9d343d6a00d Mon Sep 17 00:00:00 2001 From: Peter Steinberger Date: Tue, 25 Aug 2026 02:16:04 -0700 Subject: [PATCH] refactor(agents): consolidate extension runner callbacks (#129086) --- config/assertion-safety-baseline.txt | 2 +- src/agents/sessions/extensions/runner.test.ts | 185 +++++++- src/agents/sessions/extensions/runner.ts | 410 ++++++------------ 3 files changed, 293 insertions(+), 304 deletions(-) diff --git a/config/assertion-safety-baseline.txt b/config/assertion-safety-baseline.txt index 341f8cfa1e54..8fd04350dae6 100644 --- a/config/assertion-safety-baseline.txt +++ b/config/assertion-safety-baseline.txt @@ -2023,7 +2023,7 @@ src/agents/sessions/agent-session-tree.ts 2 src/agents/sessions/agent-session-utils.ts 1 src/agents/sessions/auth-storage.ts 4 src/agents/sessions/extensions/loader.ts 7 -src/agents/sessions/extensions/runner.ts 19 +src/agents/sessions/extensions/runner.ts 17 src/agents/sessions/extensions/types.ts 1 src/agents/sessions/keybindings.ts 2 src/agents/sessions/model-registry.ts 13 diff --git a/src/agents/sessions/extensions/runner.test.ts b/src/agents/sessions/extensions/runner.test.ts index 2bc291762709..ebe187250e2b 100644 --- a/src/agents/sessions/extensions/runner.test.ts +++ b/src/agents/sessions/extensions/runner.test.ts @@ -1,7 +1,4 @@ -// Focused tests for emitContext clone gating: the per-turn deep clone of the -// session history must be skipped when no extension registered a "context" -// handler, while handler runs keep receiving an isolated clone. -import { describe, expect, it } from "vitest"; +import { describe, expect, it, vi } from "vitest"; import type { AgentMessage } from "../../runtime/index.js"; import type { ModelRegistry } from "../model-registry.js"; import type { SessionManager } from "../session-manager.js"; @@ -9,17 +6,17 @@ import { ExtensionRunner } from "./runner.js"; import type { Extension, ExtensionRuntime } from "./types.js"; type TestHandler = (...args: unknown[]) => Promise; +type TestHandlers = Record; -function buildExtension(handlers?: Record): Extension { +async function reject(error: Error): Promise { + throw error; +} + +function buildExtension(handlers?: TestHandlers, path = "/tmp/test-extension.ts"): Extension { return { - path: "/tmp/test-extension.ts", - resolvedPath: "/tmp/test-extension.ts", - sourceInfo: { - path: "/tmp/test-extension.ts", - source: "test", - scope: "temporary", - origin: "top-level", - }, + path, + resolvedPath: path, + sourceInfo: { path, source: "test", scope: "temporary", origin: "top-level" }, handlers: new Map(Object.entries(handlers ?? {})), tools: new Map(), messageRenderers: new Map(), @@ -39,11 +36,11 @@ function buildRunner(extensions: Extension[]): ExtensionRunner { ); } -function buildMessages(): AgentMessage[] { +function buildMessages(): [AgentMessage, AgentMessage] { return [ { role: "user", content: [{ type: "text", text: "hello" }] }, { role: "assistant", content: [{ type: "text", text: "hi" }] }, - ] as AgentMessage[]; + ] as [AgentMessage, AgentMessage]; } describe("ExtensionRunner.emitContext", () => { @@ -76,13 +73,157 @@ describe("ExtensionRunner.emitContext", () => { expect(messages).toHaveLength(2); }); - it("applies replacement messages returned by a context handler", async () => { - const replacement = [ - { role: "user", content: [{ type: "text", text: "replaced" }] }, - ] as AgentMessage[]; - const handler: TestHandler = async () => ({ messages: replacement }); - const runner = buildRunner([buildExtension({ context: [handler] })]); + it("chains replacement messages through later context handlers", async () => { + const [user, assistant] = buildMessages(); + const replacement = [user]; + const final = [assistant]; + const runner = buildRunner([ + buildExtension({ + context: [ + async () => ({ messages: replacement }), + async (event) => { + expect((event as { messages: AgentMessage[] }).messages).toBe(replacement); + return { messages: final }; + }, + ], + }), + ]); - expect(await runner.emitContext(buildMessages())).toBe(replacement); + expect(await runner.emitContext(buildMessages())).toBe(final); + }); +}); + +const catchAndContinueCases: Array< + [event: string, invoke: (runner: ExtensionRunner) => Promise] +> = [ + [ + "session_before_switch", + (runner) => runner.emit({ type: "session_before_switch", reason: "new" }), + ], + [ + "message_end", + (runner) => runner.emitMessageEnd({ type: "message_end", message: buildMessages()[0] }), + ], + [ + "tool_result", + (runner) => + runner.emitToolResult({ + type: "tool_result", + toolName: "custom", + toolCallId: "call-1", + input: {}, + content: [], + details: undefined, + isError: false, + }), + ], + [ + "user_bash", + (runner) => + runner.emitUserBash({ + type: "user_bash", + command: "pwd", + cwd: "/tmp", + excludeFromContext: false, + }), + ], + ["context", (runner) => runner.emitContext(buildMessages())], + ["before_provider_request", (runner) => runner.emitBeforeProviderRequest({})], + [ + "before_agent_start", + (runner) => runner.emitBeforeAgentStart("hello", undefined, "system", { cwd: "/tmp" }), + ], + ["resources_discover", (runner) => runner.emitResourcesDiscover("/tmp", "startup")], + ["input", (runner) => runner.emitInput("hello", undefined, "interactive")], +]; + +describe("ExtensionRunner handler dispatch", () => { + it.each(catchAndContinueCases)( + "isolates %s handler failures and reports their extension", + async (event, invoke) => { + const failure = new Error(`${event} failed`); + const continued: string[] = []; + const continueHandler = async () => { + continued.push(event); + if (event === "session_before_switch") { + return { cancel: true }; + } + if (event === "input") { + return { action: "transform", text: "transformed" }; + } + if (event === "message_end") { + return { message: buildMessages()[1] }; + } + return undefined; + }; + const nextHandlers: TestHandler[] = [continueHandler]; + if (event === "input") { + nextHandlers.push( + async (input) => { + continued.push((input as { text: string }).text); + return { action: "handled" }; + }, + async () => void continued.push("unreachable"), + ); + } else if (event === "message_end") { + nextHandlers.push( + async (input) => void continued.push((input as { message: AgentMessage }).message.role), + ); + } else if (event === "session_before_switch") { + nextHandlers.push(async () => void continued.push("unreachable")); + } + const runner = buildRunner([ + buildExtension({ [event]: [() => reject(failure)] }, "/tmp/failing.ts"), + buildExtension({ [event]: nextHandlers }, "/tmp/next.ts"), + ]); + const errors: unknown[] = []; + runner.onError((error) => errors.push(error)); + + await invoke(runner); + + expect(continued).toEqual( + event === "input" + ? [event, "transformed"] + : event === "message_end" + ? [event, "user"] + : [event], + ); + const expectedErrors: unknown[] = [ + { + extensionPath: "/tmp/failing.ts", + event, + error: failure.message, + stack: failure.stack, + }, + ]; + if (event === "message_end") { + expectedErrors.push({ + extensionPath: "/tmp/next.ts", + event, + error: "message_end handlers must return a message with the same role", + }); + } + expect(errors).toEqual(expectedErrors); + }, + ); + + it("lets tool_call handler failures escape and blocks later handlers", async () => { + const failure = new Error("block tool execution"); + const laterHandler = vi.fn(async () => undefined); + const runner = buildRunner([ + buildExtension({ + tool_call: [() => reject(failure), laterHandler], + }), + ]); + + await expect( + runner.emitToolCall({ + type: "tool_call", + toolName: "custom", + toolCallId: "call-1", + input: {}, + }), + ).rejects.toBe(failure); + expect(laterHandler).not.toHaveBeenCalled(); }); }); diff --git a/src/agents/sessions/extensions/runner.ts b/src/agents/sessions/extensions/runner.ts index 28514ee0932c..445966ae9511 100644 --- a/src/agents/sessions/extensions/runner.ts +++ b/src/agents/sessions/extensions/runner.ts @@ -714,134 +714,95 @@ export class ExtensionRunner { ); } - async emit(event: TEvent): Promise> { - const ctx = this.createContext(); - let result: SessionBeforeEventResult | undefined; - + private async dispatchHandlers( + eventType: Exclude, + invoke: ( + handler: NonNullable>[number], + ctx: ExtensionContext, + extensionPath: string, + ) => Promise, + ctx = this.createContext(), + ): Promise { for (const ext of this.extensions) { - const handlers = ext.handlers.get(event.type); - if (!handlers || handlers.length === 0) { - continue; - } - - for (const handler of handlers) { + for (const handler of ext.handlers.get(eventType) ?? []) { try { - const handlerResult = await handler(event, ctx); - - if (this.isSessionBeforeEvent(event) && handlerResult) { - result = handlerResult as SessionBeforeEventResult; - if (result.cancel) { - return result as RunnerEmitResult; - } + const result = await invoke(handler, ctx, ext.path); + if (result !== undefined) { + return result; } } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; this.emitError({ extensionPath: ext.path, - event: event.type, - error: message, - stack, + event: eventType, + error: coerceErrorMessage(err), + stack: err instanceof Error ? err.stack : undefined, }); } } } + return undefined; + } - return result as RunnerEmitResult; + async emit(event: TEvent): Promise> { + let result: SessionBeforeEventResult | undefined; + + const cancelled = await this.dispatchHandlers(event.type, async (handler, ctx) => { + const handlerResult = await handler(event, ctx); + if (this.isSessionBeforeEvent(event) && handlerResult) { + result = handlerResult as SessionBeforeEventResult; + if (result.cancel) { + return result; + } + } + return undefined; + }); + + return (cancelled ?? result) as RunnerEmitResult; } async emitMessageEnd(event: MessageEndEvent): Promise { - const ctx = this.createContext(); let currentMessage = event.message; let modified = false; - for (const ext of this.extensions) { - const handlers = ext.handlers.get("message_end"); - if (!handlers || handlers.length === 0) { - continue; - } - - for (const handler of handlers) { - try { - const currentEvent: MessageEndEvent = { ...event, message: currentMessage }; - const handlerResult = (await handler(currentEvent, ctx)) as - | MessageEndEventResult - | undefined; - if (!handlerResult?.message) { - continue; - } - - if (handlerResult.message.role !== currentMessage.role) { - this.emitError({ - extensionPath: ext.path, - event: "message_end", - error: "message_end handlers must return a message with the same role", - }); - continue; - } - + await this.dispatchHandlers("message_end", async (handler, ctx, extensionPath) => { + const currentEvent: MessageEndEvent = { ...event, message: currentMessage }; + const handlerResult = (await handler(currentEvent, ctx)) as MessageEndEventResult | undefined; + if (handlerResult?.message) { + if (handlerResult.message.role !== currentMessage.role) { + this.emitError({ + extensionPath, + event: "message_end", + error: "message_end handlers must return a message with the same role", + }); + } else { currentMessage = handlerResult.message; modified = true; - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "message_end", - error: message, - stack, - }); } } - } + }); return modified ? currentMessage : undefined; } async emitToolResult(event: ToolResultEvent): Promise { - const ctx = this.createContext(); const currentEvent: ToolResultEvent = { ...event }; let modified = false; - for (const ext of this.extensions) { - const handlers = ext.handlers.get("tool_result"); - if (!handlers || handlers.length === 0) { - continue; + await this.dispatchHandlers("tool_result", async (handler, ctx) => { + const handlerResult = (await handler(currentEvent, ctx)) as ToolResultEventResult | undefined; + if (handlerResult?.content !== undefined) { + currentEvent.content = handlerResult.content; + modified = true; } - - for (const handler of handlers) { - try { - const handlerResult = (await handler(currentEvent, ctx)) as - | ToolResultEventResult - | undefined; - if (!handlerResult) { - continue; - } - - if (handlerResult.content !== undefined) { - currentEvent.content = handlerResult.content; - modified = true; - } - if (handlerResult.details !== undefined) { - currentEvent.details = handlerResult.details; - modified = true; - } - if (handlerResult.isError !== undefined) { - currentEvent.isError = handlerResult.isError; - modified = true; - } - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "tool_result", - error: message, - stack, - }); - } + if (handlerResult?.details !== undefined) { + currentEvent.details = handlerResult.details; + modified = true; } - } + if (handlerResult?.isError !== undefined) { + currentEvent.isError = handlerResult.isError; + modified = true; + } + }); if (!modified) { return undefined; @@ -880,34 +841,10 @@ export class ExtensionRunner { } async emitUserBash(event: UserBashEvent): Promise { - const ctx = this.createContext(); - - for (const ext of this.extensions) { - const handlers = ext.handlers.get("user_bash"); - if (!handlers || handlers.length === 0) { - continue; - } - - for (const handler of handlers) { - try { - const handlerResult = await handler(event, ctx); - if (handlerResult) { - return handlerResult as UserBashEventResult; - } - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "user_bash", - error: message, - stack, - }); - } - } - } - - return undefined; + return await this.dispatchHandlers("user_bash", async (handler, ctx) => { + const handlerResult = await handler(event, ctx); + return handlerResult ? (handlerResult as UserBashEventResult) : undefined; + }); } async emitContext(messages: AgentMessage[]): Promise { @@ -917,71 +854,32 @@ export class ExtensionRunner { if (!this.hasHandlers("context")) { return messages; } - const ctx = this.createContext(); let currentMessages = structuredClone(messages); - for (const ext of this.extensions) { - const handlers = ext.handlers.get("context"); - if (!handlers || handlers.length === 0) { - continue; + await this.dispatchHandlers("context", async (handler, ctx) => { + const event: ContextEvent = { type: "context", messages: currentMessages }; + const handlerResult = (await handler(event, ctx)) as ContextEventResult | undefined; + if (handlerResult?.messages) { + currentMessages = handlerResult.messages; } - - for (const handler of handlers) { - try { - const event: ContextEvent = { type: "context", messages: currentMessages }; - const handlerResult = await handler(event, ctx); - - if (handlerResult && (handlerResult as ContextEventResult).messages) { - currentMessages = (handlerResult as ContextEventResult).messages!; - } - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "context", - error: message, - stack, - }); - } - } - } + }); return currentMessages; } async emitBeforeProviderRequest(payload: unknown): Promise { - const ctx = this.createContext(); let currentPayload = payload; - for (const ext of this.extensions) { - const handlers = ext.handlers.get("before_provider_request"); - if (!handlers || handlers.length === 0) { - continue; + await this.dispatchHandlers("before_provider_request", async (handler, ctx) => { + const event: BeforeProviderRequestEvent = { + type: "before_provider_request", + payload: currentPayload, + }; + const handlerResult = await handler(event, ctx); + if (handlerResult !== undefined) { + currentPayload = handlerResult; } - - for (const handler of handlers) { - try { - const event: BeforeProviderRequestEvent = { - type: "before_provider_request", - payload: currentPayload, - }; - const handlerResult = await handler(event, ctx); - if (handlerResult !== undefined) { - currentPayload = handlerResult; - } - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "before_provider_request", - error: message, - stack, - }); - } - } - } + }); return currentPayload; } @@ -1004,45 +902,29 @@ export class ExtensionRunner { const messages: NonNullable[] = []; let systemPromptModified = false; - for (const ext of this.extensions) { - const handlers = ext.handlers.get("before_agent_start"); - if (!handlers || handlers.length === 0) { - continue; - } - - for (const handler of handlers) { - try { - const event: BeforeAgentStartEvent = { - type: "before_agent_start", - prompt, - images, - systemPrompt: currentSystemPrompt, - systemPromptOptions, - }; - const handlerResult = await handler(event, ctx); - - if (handlerResult) { - const result = handlerResult as BeforeAgentStartEventResult; - if (result.message) { - messages.push(result.message); - } - if (result.systemPrompt !== undefined) { - currentSystemPrompt = result.systemPrompt; - systemPromptModified = true; - } - } - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "before_agent_start", - error: message, - stack, - }); + await this.dispatchHandlers( + "before_agent_start", + async (handler, handlerCtx) => { + const event: BeforeAgentStartEvent = { + type: "before_agent_start", + prompt, + images, + systemPrompt: currentSystemPrompt, + systemPromptOptions, + }; + const handlerResult = (await handler(event, handlerCtx)) as + | BeforeAgentStartEventResult + | undefined; + if (handlerResult?.message) { + messages.push(handlerResult.message); } - } - } + if (handlerResult?.systemPrompt !== undefined) { + currentSystemPrompt = handlerResult.systemPrompt; + systemPromptModified = true; + } + }, + ctx, + ); if (messages.length > 0 || systemPromptModified) { return { @@ -1062,50 +944,24 @@ export class ExtensionRunner { promptPaths: Array<{ path: string; extensionPath: string }>; themePaths: Array<{ path: string; extensionPath: string }>; }> { - const ctx = this.createContext(); const skillPaths: Array<{ path: string; extensionPath: string }> = []; const promptPaths: Array<{ path: string; extensionPath: string }> = []; const themePaths: Array<{ path: string; extensionPath: string }> = []; - for (const ext of this.extensions) { - const handlers = ext.handlers.get("resources_discover"); - if (!handlers || handlers.length === 0) { - continue; - } + await this.dispatchHandlers("resources_discover", async (handler, ctx, extensionPath) => { + const event: ResourcesDiscoverEvent = { type: "resources_discover", cwd, reason }; + const result = (await handler(event, ctx)) as ResourcesDiscoverResult | undefined; - for (const handler of handlers) { - try { - const event: ResourcesDiscoverEvent = { type: "resources_discover", cwd, reason }; - const handlerResult = await handler(event, ctx); - const result = handlerResult as ResourcesDiscoverResult | undefined; - - if (result?.skillPaths?.length) { - skillPaths.push( - ...result.skillPaths.map((path) => ({ path, extensionPath: ext.path })), - ); - } - if (result?.promptPaths?.length) { - promptPaths.push( - ...result.promptPaths.map((path) => ({ path, extensionPath: ext.path })), - ); - } - if (result?.themePaths?.length) { - themePaths.push( - ...result.themePaths.map((path) => ({ path, extensionPath: ext.path })), - ); - } - } catch (err) { - const message = coerceErrorMessage(err); - const stack = err instanceof Error ? err.stack : undefined; - this.emitError({ - extensionPath: ext.path, - event: "resources_discover", - error: message, - stack, - }); - } + if (result?.skillPaths?.length) { + skillPaths.push(...result.skillPaths.map((path) => ({ path, extensionPath }))); } - } + if (result?.promptPaths?.length) { + promptPaths.push(...result.promptPaths.map((path) => ({ path, extensionPath }))); + } + if (result?.themePaths?.length) { + themePaths.push(...result.themePaths.map((path) => ({ path, extensionPath }))); + } + }); return { skillPaths, promptPaths, themePaths }; } @@ -1116,36 +972,28 @@ export class ExtensionRunner { images: ImageContent[] | undefined, source: InputSource, ): Promise { - const ctx = this.createContext(); let currentText = text; let currentImages = images; - for (const ext of this.extensions) { - for (const handler of ext.handlers.get("input") ?? []) { - try { - const event: InputEvent = { - type: "input", - text: currentText, - images: currentImages, - source, - }; - const result = (await handler(event, ctx)) as InputEventResult | undefined; - if (result?.action === "handled") { - return result; - } - if (result?.action === "transform") { - currentText = result.text; - currentImages = result.images ?? currentImages; - } - } catch (err) { - this.emitError({ - extensionPath: ext.path, - event: "input", - error: coerceErrorMessage(err), - stack: err instanceof Error ? err.stack : undefined, - }); - } + const handled = await this.dispatchHandlers("input", async (handler, ctx) => { + const event: InputEvent = { + type: "input", + text: currentText, + images: currentImages, + source, + }; + const result = (await handler(event, ctx)) as InputEventResult | undefined; + if (result?.action === "handled") { + return result; } + if (result?.action === "transform") { + currentText = result.text; + currentImages = result.images ?? currentImages; + } + return undefined; + }); + if (handled) { + return handled; } return currentText !== text || currentImages !== images ? { action: "transform", text: currentText, images: currentImages }