fix(coding-agent): gate rpc extension turn tracking on send success

This commit is contained in:
Kormákur
2026-06-13 01:26:59 +00:00
committed by can1357
parent 62f9cc0ddc
commit c805514c75
3 changed files with 177 additions and 10 deletions
@@ -111,10 +111,13 @@ export function reportLocalOnlyPromptResult(input: {
output: (obj: object) => void;
onError: (error: Error) => void;
hasExtensionAgentMessageTask?: () => boolean;
waitForExtensionAgentMessageTasks?: () => Promise<void>;
}): void {
void input.prompt
.then(agentInvoked => {
if (!agentInvoked && !input.hasExtensionAgentMessageTask?.()) {
.then(async agentInvoked => {
if (agentInvoked) return;
await input.waitForExtensionAgentMessageTasks?.();
if (!input.hasExtensionAgentMessageTask?.()) {
input.output({ type: "prompt_result", id: input.id, agentInvoked: false });
}
})
@@ -125,6 +128,7 @@ export function reportLocalOnlyPromptResult(input: {
type RpcExtensionUserMessageScope = {
hasAgentMessageTask: boolean;
pendingAgentMessageTasks: Set<Promise<void>>;
};
/**
@@ -142,11 +146,40 @@ export class RpcExtensionUserMessageTracker {
}
}
trackAgentMessageTask(task: Promise<void>): void {
for (const scope of this.#activePromptScopes) {
this.#trackAgentMessageTaskForScope(scope, task);
}
}
#trackAgentMessageTaskForScope(scope: RpcExtensionUserMessageScope, task: Promise<void>): void {
const scopedTask = task.then(
() => {
scope.hasAgentMessageTask = true;
},
() => {},
);
scope.pendingAgentMessageTasks.add(scopedTask);
void scopedTask.finally(() => {
scope.pendingAgentMessageTasks.delete(scopedTask);
});
}
async #waitForAgentMessageTasks(scope: RpcExtensionUserMessageScope): Promise<void> {
while (scope.pendingAgentMessageTasks.size > 0) {
await Promise.allSettled(Array.from(scope.pendingAgentMessageTasks));
}
}
watchPrompt<T>(startPrompt: () => Promise<T>): {
prompt: Promise<T>;
hasAgentMessageTask: () => boolean;
waitForAgentMessageTasks: () => Promise<void>;
} {
const scope: RpcExtensionUserMessageScope = { hasAgentMessageTask: false };
const scope: RpcExtensionUserMessageScope = {
hasAgentMessageTask: false,
pendingAgentMessageTasks: new Set(),
};
this.#activePromptScopes.add(scope);
let prompt: Promise<T>;
try {
@@ -160,6 +193,7 @@ export class RpcExtensionUserMessageTracker {
this.#activePromptScopes.delete(scope);
}),
hasAgentMessageTask: () => scope.hasAgentMessageTask,
waitForAgentMessageTasks: () => this.#waitForAgentMessageTasks(scope),
};
}
}
@@ -178,6 +212,7 @@ export function watchAndReportLocalOnlyPromptResult(input: {
output: input.output,
onError: input.onError,
hasExtensionAgentMessageTask: trackedPrompt.hasAgentMessageTask,
waitForExtensionAgentMessageTasks: trackedPrompt.waitForAgentMessageTasks,
});
}
@@ -616,8 +651,8 @@ export async function runRpcMode(
onShutdown: () => {
shutdownState.requested = true;
},
markAgentInvokingMessage: () => {
extensionUserMessageTracker.markAgentMessageTask();
trackAgentInvokingMessage: task => {
extensionUserMessageTracker.trackAgentMessageTask(task);
},
uiContext: rpcUiContext,
});
@@ -25,6 +25,8 @@ export interface InitializeExtensionsOptions {
uiContext?: ExtensionUIContext;
/** Optional lifecycle hook for extension-originated messages that can start an agent turn. */
markAgentInvokingMessage?: () => void;
/** Optional lifecycle hook for extension-originated sends whose success/failure determines turn ownership. */
trackAgentInvokingMessage?: (task: Promise<void>) => void;
}
/**
@@ -37,23 +39,40 @@ export async function initializeExtensions(session: AgentSession, options: Initi
const runner = session.extensionRunner;
if (!runner) return;
const { reportSendError, reportRuntimeError, onShutdown, uiContext, markAgentInvokingMessage } = options;
const {
reportSendError,
reportRuntimeError,
onShutdown,
uiContext,
markAgentInvokingMessage,
trackAgentInvokingMessage,
} = options;
const shutdown = onShutdown ?? (() => {});
runner.initialize(
// ExtensionActions
{
sendMessage: (message, sendOptions) => {
const sendTask = session.sendCustomMessage(message, sendOptions);
if (sendOptions?.triggerTurn) {
markAgentInvokingMessage?.();
if (trackAgentInvokingMessage) {
trackAgentInvokingMessage(sendTask);
} else {
markAgentInvokingMessage?.();
}
}
session.sendCustomMessage(message, sendOptions).catch(e => {
sendTask.catch(e => {
reportSendError("extension_send", e instanceof Error ? e : new Error(String(e)));
});
},
sendUserMessage: (content, sendOptions) => {
markAgentInvokingMessage?.();
session.sendUserMessage(content, sendOptions).catch(e => {
const sendTask = session.sendUserMessage(content, sendOptions);
if (trackAgentInvokingMessage) {
trackAgentInvokingMessage(sendTask);
} else {
markAgentInvokingMessage?.();
}
sendTask.catch(e => {
reportSendError("extension_send_user", e instanceof Error ? e : new Error(String(e)));
});
},
@@ -13,6 +13,16 @@ async function waitForPromptHandlers(prompt: Promise<unknown>): Promise<void> {
await Promise.resolve();
}
async function waitForTrackedPromptHandlers(trackedPrompt: {
prompt: Promise<unknown>;
waitForAgentMessageTasks: () => Promise<void>;
}): Promise<void> {
await trackedPrompt.prompt.catch(() => undefined);
await trackedPrompt.waitForAgentMessageTasks();
await Promise.resolve();
await Promise.resolve();
}
describe("reportLocalOnlyPromptResult", () => {
test("emits prompt_result when prompt resolves without invoking the agent or extension user message", async () => {
const output: object[] = [];
@@ -140,6 +150,109 @@ describe("reportLocalOnlyPromptResult", () => {
expect(sentOptions).toEqual({ triggerTurn: true });
});
test("suppresses prompt_result when extension sendUserMessage succeeds", async () => {
let extensionActions: ExtensionActions | undefined;
let sentContent: unknown;
const output: object[] = [];
const extensionUserMessages = new RpcExtensionUserMessageTracker();
const session = {
extensionRunner: {
initialize: (actions: ExtensionActions) => {
extensionActions = actions;
},
onError: () => {},
emit: async () => {},
},
sendUserMessage: async (content: unknown) => {
sentContent = content;
},
} as unknown as AgentSession;
await initializeExtensions(session, {
reportSendError: (_action, error) => {
throw error;
},
reportRuntimeError: error => {
throw error.error;
},
trackAgentInvokingMessage: task => {
extensionUserMessages.trackAgentMessageTask(task);
},
});
const trackedPrompt = extensionUserMessages.watchPrompt(() => {
if (!extensionActions) throw new Error("extensions not initialized");
extensionActions.sendUserMessage("start work");
return Promise.resolve(false);
});
reportLocalOnlyPromptResult({
id: "req_success",
prompt: trackedPrompt.prompt,
output: frame => output.push(frame),
onError: error => {
throw error;
},
hasExtensionAgentMessageTask: trackedPrompt.hasAgentMessageTask,
waitForExtensionAgentMessageTasks: trackedPrompt.waitForAgentMessageTasks,
});
await waitForTrackedPromptHandlers(trackedPrompt);
expect(sentContent).toBe("start work");
expect(output).toEqual([]);
});
test("emits prompt_result when extension sendUserMessage rejects", async () => {
let extensionActions: ExtensionActions | undefined;
const output: object[] = [];
const reportedErrors: Error[] = [];
const thrown = new Error("missing model");
const extensionUserMessages = new RpcExtensionUserMessageTracker();
const session = {
extensionRunner: {
initialize: (actions: ExtensionActions) => {
extensionActions = actions;
},
onError: () => {},
emit: async () => {},
},
sendUserMessage: async () => {
throw thrown;
},
} as unknown as AgentSession;
await initializeExtensions(session, {
reportSendError: (_action, error) => {
reportedErrors.push(error);
},
reportRuntimeError: error => {
throw error.error;
},
trackAgentInvokingMessage: task => {
extensionUserMessages.trackAgentMessageTask(task);
},
});
const trackedPrompt = extensionUserMessages.watchPrompt(() => {
if (!extensionActions) throw new Error("extensions not initialized");
extensionActions.sendUserMessage("start work");
return Promise.resolve(false);
});
reportLocalOnlyPromptResult({
id: "req_rejected",
prompt: trackedPrompt.prompt,
output: frame => output.push(frame),
onError: error => {
throw error;
},
hasExtensionAgentMessageTask: trackedPrompt.hasAgentMessageTask,
waitForExtensionAgentMessageTasks: trackedPrompt.waitForAgentMessageTasks,
});
await waitForTrackedPromptHandlers(trackedPrompt);
expect(reportedErrors).toEqual([thrown]);
expect(output).toEqual([{ type: "prompt_result", id: "req_rejected", agentInvoked: false }]);
});
test("does not emit when prompt invokes the agent", async () => {
const output: object[] = [];
const prompt = Promise.resolve(true);