From c805514c758f797689254a4ebd28d0f7bc828d37 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Korm=C3=A1kur?= Date: Sat, 13 Jun 2026 01:26:59 +0000 Subject: [PATCH] fix(coding-agent): gate rpc extension turn tracking on send success --- .../coding-agent/src/modes/rpc/rpc-mode.ts | 45 ++++++- .../coding-agent/src/modes/runtime-init.ts | 29 ++++- .../test/rpc-prompt-result.test.ts | 113 ++++++++++++++++++ 3 files changed, 177 insertions(+), 10 deletions(-) diff --git a/packages/coding-agent/src/modes/rpc/rpc-mode.ts b/packages/coding-agent/src/modes/rpc/rpc-mode.ts index 237edfa63..1c16665f7 100644 --- a/packages/coding-agent/src/modes/rpc/rpc-mode.ts +++ b/packages/coding-agent/src/modes/rpc/rpc-mode.ts @@ -111,10 +111,13 @@ export function reportLocalOnlyPromptResult(input: { output: (obj: object) => void; onError: (error: Error) => void; hasExtensionAgentMessageTask?: () => boolean; + waitForExtensionAgentMessageTasks?: () => Promise; }): 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>; }; /** @@ -142,11 +146,40 @@ export class RpcExtensionUserMessageTracker { } } + trackAgentMessageTask(task: Promise): void { + for (const scope of this.#activePromptScopes) { + this.#trackAgentMessageTaskForScope(scope, task); + } + } + + #trackAgentMessageTaskForScope(scope: RpcExtensionUserMessageScope, task: Promise): void { + const scopedTask = task.then( + () => { + scope.hasAgentMessageTask = true; + }, + () => {}, + ); + scope.pendingAgentMessageTasks.add(scopedTask); + void scopedTask.finally(() => { + scope.pendingAgentMessageTasks.delete(scopedTask); + }); + } + + async #waitForAgentMessageTasks(scope: RpcExtensionUserMessageScope): Promise { + while (scope.pendingAgentMessageTasks.size > 0) { + await Promise.allSettled(Array.from(scope.pendingAgentMessageTasks)); + } + } + watchPrompt(startPrompt: () => Promise): { prompt: Promise; hasAgentMessageTask: () => boolean; + waitForAgentMessageTasks: () => Promise; } { - const scope: RpcExtensionUserMessageScope = { hasAgentMessageTask: false }; + const scope: RpcExtensionUserMessageScope = { + hasAgentMessageTask: false, + pendingAgentMessageTasks: new Set(), + }; this.#activePromptScopes.add(scope); let prompt: Promise; 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, }); diff --git a/packages/coding-agent/src/modes/runtime-init.ts b/packages/coding-agent/src/modes/runtime-init.ts index d78196b29..aee58426e 100644 --- a/packages/coding-agent/src/modes/runtime-init.ts +++ b/packages/coding-agent/src/modes/runtime-init.ts @@ -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; } /** @@ -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))); }); }, diff --git a/packages/coding-agent/test/rpc-prompt-result.test.ts b/packages/coding-agent/test/rpc-prompt-result.test.ts index a0a9b61a4..ca4643cff 100644 --- a/packages/coding-agent/test/rpc-prompt-result.test.ts +++ b/packages/coding-agent/test/rpc-prompt-result.test.ts @@ -13,6 +13,16 @@ async function waitForPromptHandlers(prompt: Promise): Promise { await Promise.resolve(); } +async function waitForTrackedPromptHandlers(trackedPrompt: { + prompt: Promise; + waitForAgentMessageTasks: () => Promise; +}): Promise { + 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);