fix(coding-agent): gate rpc extension turn tracking on send success
This commit is contained in:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user