diff --git a/packages/ai/src/providers/cursor.ts b/packages/ai/src/providers/cursor.ts index 47d4b3c8f..0bb0e1273 100644 --- a/packages/ai/src/providers/cursor.ts +++ b/packages/ai/src/providers/cursor.ts @@ -12,6 +12,7 @@ import type { CursorExecHandlerResult, CursorExecHandlers, CursorMcpCall, + CursorShellStreamCallbacks, CursorToolResultHandler, ImageContent, Message, @@ -51,8 +52,10 @@ import { DiagnosticsRejectedSchema, DiagnosticsResultSchema, DiagnosticsSuccessSchema, + ExecClientControlMessageSchema, type ExecClientMessage, ExecClientMessageSchema, + ExecClientStreamCloseSchema, type ExecServerMessage, FetchErrorSchema, FetchResultSchema, @@ -397,6 +400,8 @@ export const streamCursor: StreamFunction<"cursor-agent"> = ( conversationStateCache.set(conversationId, checkpoint); }; + let resolveH2: (() => void) | undefined; + h2Request.on("data", (chunk: Buffer) => { pendingBuffer = Buffer.concat([pendingBuffer, chunk]); @@ -419,6 +424,9 @@ export const streamCursor: StreamFunction<"cursor-agent"> = ( try { const serverMessage = fromBinary(AgentServerMessageSchema, messageBytes); + const isTurnEnded = + serverMessage.message.case === "interactionUpdate" && + serverMessage.message.value.message?.case === "turnEnded"; void handleServerMessage( serverMessage, output, @@ -434,6 +442,14 @@ export const streamCursor: StreamFunction<"cursor-agent"> = ( ).catch(error => { log("error", "handleServerMessage", { error: String(error) }); }); + + // Resolve only on explicit turnEnded. stopReason defaults to "stop" + // and is not a reliable signal for stream completion. + if (isTurnEnded && resolveH2) { + const r = resolveH2; + resolveH2 = undefined; + r(); + } } catch (e) { log("error", "parseServerMessage", { error: String(e) }); } @@ -456,6 +472,8 @@ export const streamCursor: StreamFunction<"cursor-agent"> = ( heartbeatTimer = setInterval(sendHeartbeat, 5000); await new Promise((resolve, reject) => { + resolveH2 = resolve; + h2Request!.on("trailers", trailers => { const status = trailers["grpc-status"]; const msg = trailers["grpc-message"]; @@ -465,6 +483,7 @@ export const streamCursor: StreamFunction<"cursor-agent"> = ( }); h2Request!.on("end", () => { + resolveH2 = undefined; if (endStreamError) { reject(endStreamError); return; @@ -662,32 +681,87 @@ async function handleShellStreamArgs( execHandlers: CursorExecHandlers | undefined, onToolResult: CursorToolResultHandler | undefined, ): Promise { - const { execResult } = await resolveExecHandler( - args as any, - execHandlers?.shell?.bind(execHandlers), - onToolResult, - toolResult => buildShellResultFromToolResult(args as any, toolResult), - reason => buildShellRejectedResult((args as any).command, (args as any).workingDirectory, reason), - error => buildShellFailureResult((args as any).command, (args as any).workingDirectory, error), - ); + const normalizedWorkingDirectory = args.workingDirectory || process.cwd(); + const normalizedArgs: ShellArgs = { ...args, workingDirectory: normalizedWorkingDirectory }; + const startTs = Date.now(); + log("shellStream", "start", { + command: (args as any).command, + workingDirectory: normalizedWorkingDirectory, + execId: execMsg.execId, + hasExecHandlers: !!execHandlers, + hasShell: !!execHandlers?.shell, + hasShellStream: !!execHandlers?.shellStream, + }); sendShellStreamEvent(h2Request, execMsg, { case: "start", value: create(ShellStreamStartSchema, {}) }); + const streamCallbacks: CursorShellStreamCallbacks = { + onStdout(data: string) { + sendShellStreamEvent(h2Request, execMsg, { + case: "stdout", + value: create(ShellStreamStdoutSchema, { data }), + }); + }, + onStderr(data: string) { + sendShellStreamEvent(h2Request, execMsg, { + case: "stderr", + value: create(ShellStreamStderrSchema, { data }), + }); + }, + }; + + // Prefer the streaming handler — it forwards output chunks in real time. + // Falls back to the batch shell handler otherwise. + const streamHandler = execHandlers?.shellStream?.bind(execHandlers); + const batchHandler = execHandlers?.shell?.bind(execHandlers); + const handler = streamHandler ? (shellArgs: ShellArgs) => streamHandler(shellArgs, streamCallbacks) : batchHandler; + + const { execResult } = await resolveExecHandler( + args as any, + handler as typeof batchHandler, + onToolResult, + toolResult => buildShellResultFromToolResult(normalizedArgs as any, toolResult), + reason => + buildShellRejectedResult((normalizedArgs as any).command, (normalizedArgs as any).workingDirectory, reason), + error => + buildShellFailureResult((normalizedArgs as any).command, (normalizedArgs as any).workingDirectory, error), + ); + + // When using the batch handler (no shellStream), send buffered stdout/stderr + // after execution completes. With shellStream these were already sent in real time. + const sendBufferedOutput = !streamHandler; + sendShellStreamExitFromResult(h2Request, execMsg, execResult, sendBufferedOutput); + // Cursor can keep the turn pending when it receives only stream deltas. + // Send the final structured shellResult as completion acknowledgement. + sendExecClientMessage(h2Request, execMsg, "shellResult", execResult); + sendExecClientStreamClose(h2Request, execMsg); + + log("shellStream", "done", { elapsed: Date.now() - startTs }); +} + +function sendShellStreamExitFromResult( + h2Request: http2.ClientHttp2Stream, + execMsg: ExecServerMessage, + execResult: { result: { case?: string; value?: any } }, + sendBufferedOutput: boolean, +): void { const result = execResult.result; switch (result.case) { case "success": { const value = result.value; - if (value.stdout) { - sendShellStreamEvent(h2Request, execMsg, { - case: "stdout", - value: create(ShellStreamStdoutSchema, { data: value.stdout }), - }); - } - if (value.stderr) { - sendShellStreamEvent(h2Request, execMsg, { - case: "stderr", - value: create(ShellStreamStderrSchema, { data: value.stderr }), - }); + if (sendBufferedOutput) { + if (value.stdout) { + sendShellStreamEvent(h2Request, execMsg, { + case: "stdout", + value: create(ShellStreamStdoutSchema, { data: value.stdout }), + }); + } + if (value.stderr) { + sendShellStreamEvent(h2Request, execMsg, { + case: "stderr", + value: create(ShellStreamStderrSchema, { data: value.stderr }), + }); + } } sendShellStreamEvent(h2Request, execMsg, { case: "exit", @@ -701,17 +775,19 @@ async function handleShellStreamArgs( } case "failure": { const value = result.value; - if (value.stdout) { - sendShellStreamEvent(h2Request, execMsg, { - case: "stdout", - value: create(ShellStreamStdoutSchema, { data: value.stdout }), - }); - } - if (value.stderr) { - sendShellStreamEvent(h2Request, execMsg, { - case: "stderr", - value: create(ShellStreamStderrSchema, { data: value.stderr }), - }); + if (sendBufferedOutput) { + if (value.stdout) { + sendShellStreamEvent(h2Request, execMsg, { + case: "stdout", + value: create(ShellStreamStdoutSchema, { data: value.stdout }), + }); + } + if (value.stderr) { + sendShellStreamEvent(h2Request, execMsg, { + case: "stderr", + value: create(ShellStreamStderrSchema, { data: value.stderr }), + }); + } } sendShellStreamEvent(h2Request, execMsg, { case: "exit", @@ -779,6 +855,7 @@ async function handleExecServerMessage( requestContextTools: McpToolDefinition[], ): Promise { const execCase = execMsg.message.case; + log("exec", "dispatch", { execCase, execId: execMsg.execId, hasHandlers: !!execHandlers }); if (execCase === "requestContextArgs") { const requestContext = create(RequestContextSchema, { rules: [], @@ -884,13 +961,14 @@ async function handleExecServerMessage( } case "shellArgs": { const args = execMsg.message.value; + const normalizedArgs: ShellArgs = { ...args, workingDirectory: args.workingDirectory || process.cwd() }; const { execResult } = await resolveExecHandler( args, execHandlers?.shell?.bind(execHandlers), onToolResult, - toolResult => buildShellResultFromToolResult(args, toolResult), - reason => buildShellRejectedResult(args.command, args.workingDirectory, reason), - error => buildShellFailureResult(args.command, args.workingDirectory, error), + toolResult => buildShellResultFromToolResult(normalizedArgs, toolResult), + reason => buildShellRejectedResult(normalizedArgs.command, normalizedArgs.workingDirectory, reason), + error => buildShellFailureResult(normalizedArgs.command, normalizedArgs.workingDirectory, error), ); sendExecClientMessage(h2Request, execMsg, "shellResult", execResult); return; @@ -989,8 +1067,19 @@ async function handleExecServerMessage( sendExecClientMessage(h2Request, execMsg, "computerUseResult", execResult); return; } - default: + default: { log("warn", "unhandledExecMessage", { execCase }); + // Send a bare ExecClientMessage (id + execId only, no typed result) so the + // server gets an acknowledgement and doesn't hang waiting forever. + const ack = create(ExecClientMessageSchema, { + id: execMsg.id, + execId: execMsg.execId, + }); + const clientMessage = create(AgentClientMessageSchema, { + message: { case: "execClientMessage", value: ack }, + }); + h2Request.write(frameConnectMessage(toBinary(AgentClientMessageSchema, clientMessage))); + } } } @@ -1019,6 +1108,23 @@ function sendExecClientMessage( log("execClientMessage", messageCase, value); } +function sendExecClientStreamClose(h2Request: http2.ClientHttp2Stream, execMsg: ExecServerMessage): void { + const closeMessage = create(ExecClientControlMessageSchema, { + message: { + case: "streamClose", + value: create(ExecClientStreamCloseSchema, { + id: execMsg.id, + }), + }, + }); + const clientMessage = create(AgentClientMessageSchema, { + message: { case: "execClientControlMessage", value: closeMessage }, + }); + const responseBytes = toBinary(AgentClientMessageSchema, clientMessage); + h2Request.write(frameConnectMessage(responseBytes)); + log("execClientControl", "streamClose", { id: execMsg.id, execId: execMsg.execId }); +} + /** Exported for tests: verifies handler is invoked with correct `this` when passed as bound. */ export async function resolveExecHandler( args: TArgs, diff --git a/packages/ai/src/types.ts b/packages/ai/src/types.ts index e38224e60..bcae7f4eb 100644 --- a/packages/ai/src/types.ts +++ b/packages/ai/src/types.ts @@ -380,6 +380,11 @@ export interface CursorMcpCall { rawArgs: Record; } +export interface CursorShellStreamCallbacks { + onStdout(data: string): void; + onStderr(data: string): void; +} + export interface CursorExecHandlers { read?: (args: ReadArgs) => Promise>; ls?: (args: LsArgs) => Promise>; @@ -387,6 +392,10 @@ export interface CursorExecHandlers { write?: (args: WriteArgs) => Promise>; delete?: (args: DeleteArgs) => Promise>; shell?: (args: ShellArgs) => Promise>; + shellStream?: ( + args: ShellArgs, + callbacks: CursorShellStreamCallbacks, + ) => Promise>; diagnostics?: (args: DiagnosticsArgs) => Promise>; mcp?: (call: CursorMcpCall) => Promise>; onToolResult?: CursorToolResultHandler; diff --git a/packages/coding-agent/src/cursor.ts b/packages/coding-agent/src/cursor.ts index 3c47f235e..7548da262 100644 --- a/packages/coding-agent/src/cursor.ts +++ b/packages/coding-agent/src/cursor.ts @@ -7,7 +7,12 @@ import type { AgentToolResult, AgentToolUpdateCallback, } from "@oh-my-pi/pi-agent-core"; -import type { CursorMcpCall, CursorExecHandlers as ICursorExecHandlers, ToolResultMessage } from "@oh-my-pi/pi-ai"; +import type { + CursorMcpCall, + CursorShellStreamCallbacks, + CursorExecHandlers as ICursorExecHandlers, + ToolResultMessage, +} from "@oh-my-pi/pi-ai"; import { resolveToCwd } from "./tools/path-utils"; interface CursorExecBridgeOptions { @@ -204,6 +209,66 @@ export class CursorExecHandlers implements ICursorExecHandlers { return toolResultMessage; } + async shellStream( + args: Parameters>[0], + callbacks: CursorShellStreamCallbacks, + ) { + const toolCallId = decodeToolCallId(args.toolCallId); + const toolName = "bash"; + const tool = this.options.tools.get(toolName); + if (!tool) { + const result = buildToolErrorResult(`Tool "${toolName}" not available`); + return createToolResultMessage(toolCallId, toolName, result, true); + } + + const timeoutSeconds = args.timeout && args.timeout > 0 ? args.timeout : undefined; + const toolArgs: Record = { + command: args.command, + cwd: args.workingDirectory || undefined, + timeout: timeoutSeconds, + }; + + this.options.emitEvent?.({ type: "tool_execution_start", toolCallId, toolName, args: toolArgs }); + + let result: AgentToolResult; + let isError = false; + + // Track previously streamed text so we only forward deltas. + let streamedLen = 0; + const onUpdate: AgentToolUpdateCallback = partialResult => { + this.options.emitEvent?.({ + type: "tool_execution_update", + toolCallId, + toolName, + args: toolArgs, + partialResult, + }); + const text = partialResult.content.map(c => (c.type === "text" ? c.text : "")).join(""); + if (text.length > streamedLen) { + callbacks.onStdout(text.slice(streamedLen)); + streamedLen = text.length; + } + }; + + try { + result = await tool.execute(toolCallId, toolArgs, undefined, onUpdate, this.options.getToolContext?.()); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + result = buildToolErrorResult(message); + isError = true; + } + + // onUpdate may not fire for every chunk — flush any remaining output + // from the final result that wasn't already streamed. + const finalText = result.content.map(c => (c.type === "text" ? c.text : "")).join(""); + if (finalText.length > streamedLen) { + callbacks.onStdout(finalText.slice(streamedLen)); + } + + this.options.emitEvent?.({ type: "tool_execution_end", toolCallId, toolName, result, isError }); + return createToolResultMessage(toolCallId, toolName, result, isError); + } + async diagnostics(args: Parameters>[0]) { const toolCallId = decodeToolCallId(args.toolCallId); const toolResultMessage = await executeTool(this.options, "lsp", toolCallId, {