refactor(agents): consolidate extension runner callbacks (#129086)

This commit is contained in:
Peter Steinberger
2026-08-25 02:16:04 -07:00
committed by GitHub
parent c5d1cb38e2
commit 1cf405f032
3 changed files with 293 additions and 304 deletions
+1 -1
View File
@@ -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
+163 -22
View File
@@ -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<unknown>;
type TestHandlers = Record<string, TestHandler[]>;
function buildExtension(handlers?: Record<string, TestHandler[]>): Extension {
async function reject(error: Error): Promise<never> {
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<unknown>]
> = [
[
"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();
});
});
+129 -281
View File
@@ -714,134 +714,95 @@ export class ExtensionRunner {
);
}
async emit<TEvent extends RunnerEmitEvent>(event: TEvent): Promise<RunnerEmitResult<TEvent>> {
const ctx = this.createContext();
let result: SessionBeforeEventResult | undefined;
private async dispatchHandlers<TResult>(
eventType: Exclude<ExtensionEvent["type"], "tool_call">,
invoke: (
handler: NonNullable<ReturnType<Extension["handlers"]["get"]>>[number],
ctx: ExtensionContext,
extensionPath: string,
) => Promise<TResult | undefined>,
ctx = this.createContext(),
): Promise<TResult | undefined> {
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<TEvent>;
}
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<TEvent>;
async emit<TEvent extends RunnerEmitEvent>(event: TEvent): Promise<RunnerEmitResult<TEvent>> {
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<TEvent>;
}
async emitMessageEnd(event: MessageEndEvent): Promise<AgentMessage | undefined> {
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<ToolResultEventResult | undefined> {
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<UserBashEventResult | undefined> {
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<AgentMessage[]> {
@@ -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<unknown> {
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<BeforeAgentStartEventResult["message"]>[] = [];
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<InputEventResult> {
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 }