diff --git a/apps/vscode-e2e/src/fixtures/queued-api-input.ts b/apps/vscode-e2e/src/fixtures/queued-api-input.ts new file mode 100644 index 0000000000..752dabf901 --- /dev/null +++ b/apps/vscode-e2e/src/fixtures/queued-api-input.ts @@ -0,0 +1,49 @@ +import { LLMock } from "@copilotkit/aimock" +import type { ChatCompletionRequest } from "@copilotkit/aimock" + +export const QUEUED_API_INPUT_PROMPT = "QUEUED_API_INPUT_APPROVAL: Run the marker command." +export const QUEUED_API_INPUT_MARKER_FILE = "queued-api-input-marker.txt" +export const QUEUED_API_INPUT_MESSAGE = "Steering note from the API." +export const QUEUED_API_INPUT_RESPONSE_LATENCY_MS = 2_000 + +const COMMAND_CALL_ID = "call_queued_api_input_command_001" + +export function addQueuedApiInputFixtures(mock: InstanceType) { + mock.addFixture({ + match: { + predicate: (req: ChatCompletionRequest) => { + const messages = Array.isArray(req?.messages) ? req.messages : [] + const lastUser = messages.filter((m) => m?.role === "user").at(-1) + return JSON.stringify(lastUser?.content ?? "").includes(QUEUED_API_INPUT_PROMPT) + }, + }, + streamingProfile: { ttft: QUEUED_API_INPUT_RESPONSE_LATENCY_MS }, + response: { + toolCalls: [ + { + name: "execute_command", + arguments: JSON.stringify({ command: `touch ${QUEUED_API_INPUT_MARKER_FILE}` }), + id: COMMAND_CALL_ID, + }, + ], + }, + }) + + mock.addFixture({ + match: { + predicate: (req: ChatCompletionRequest) => { + const messages = Array.isArray(req?.messages) ? req.messages : [] + return messages.filter((m) => m?.role === "tool").at(-1)?.tool_call_id === COMMAND_CALL_ID + }, + }, + response: { + toolCalls: [ + { + name: "attempt_completion", + arguments: JSON.stringify({ result: "Ran the marker command." }), + id: "call_queued_api_input_completion_002", + }, + ], + }, + }) +} diff --git a/apps/vscode-e2e/src/fixtures/subtasks.ts b/apps/vscode-e2e/src/fixtures/subtasks.ts index ebfd94324e..7ce8122185 100644 --- a/apps/vscode-e2e/src/fixtures/subtasks.ts +++ b/apps/vscode-e2e/src/fixtures/subtasks.ts @@ -16,6 +16,8 @@ const SUBTASK_APPROVAL_RESTORE_CHILD_MARKER = "SUBTASK_CHILD_APPROVAL_RESTORE" const SUBTASK_XPROFILE_PARENT_MARKER = "SUBTASK_PARENT_CROSS_PROFILE" const SUBTASK_XPROFILE_SAME_CHILD_MARKER = "SUBTASK_CHILD_SAME_PROFILE" const SUBTASK_XPROFILE_DIFFERENT_CHILD_MARKER = "SUBTASK_CHILD_DIFFERENT_PROFILE" +export const SUBTASK_QUEUED_INPUT_PARENT_MARKER = "SUBTASK_PARENT_QUEUED_INPUT" +export const SUBTASK_QUEUED_INPUT_CHILD_MARKER = "SUBTASK_CHILD_QUEUED_INPUT" const SUBTASK_CHILD_PROMPT = `${SUBTASK_CHILD_MARKER}: Ask the user exactly this follow-up question: What is the square root of 81? After the user answers, complete with only the answer.` export const SUBTASK_PARENT_PROMPT = `${SUBTASK_PARENT_MARKER}: Use the new_task tool exactly once. Create an ask-mode subtask with this exact message: "${SUBTASK_CHILD_PROMPT}" Do not answer directly.` @@ -59,6 +61,14 @@ export const SUBTASK_XPROFILE_SAME_CHILD_RESULT = "Same-profile child completed" export const SUBTASK_XPROFILE_DIFFERENT_CHILD_RESULT = "Different-profile child completed" export const SUBTASK_XPROFILE_PARENT_RESULT = "Sequential cross-profile parent resumed" +const SUBTASK_QUEUED_INPUT_INITIAL_RESULT = "Child completed before queued input" +export const SUBTASK_QUEUED_INPUT_MESSAGE = "Use the queued instruction before completing." +export const SUBTASK_QUEUED_INPUT_CHILD_RESULT = "Child processed queued input" +export const SUBTASK_QUEUED_INPUT_PARENT_RESULT = "Parent resumed after queued input" +const SUBTASK_QUEUED_INPUT_CHILD_PROMPT = `${SUBTASK_QUEUED_INPUT_CHILD_MARKER}: Complete immediately with the exact result "${SUBTASK_QUEUED_INPUT_INITIAL_RESULT}".` +export const SUBTASK_QUEUED_INPUT_PARENT_PROMPT = `${SUBTASK_QUEUED_INPUT_PARENT_MARKER}: Use the new_task tool exactly once. Create an ask-mode subtask with this exact message: "${SUBTASK_QUEUED_INPUT_CHILD_PROMPT}" Do not answer directly. When the subtask returns, complete with the exact result "${SUBTASK_QUEUED_INPUT_PARENT_RESULT}".` +export const SUBTASK_QUEUED_INPUT_RESPONSE_LATENCY_MS = 2_000 + // Scheduler regression tests — exercises TaskScheduler + run() dispatch post-CodeRabbit fix. // Separate markers to avoid collisions with the other subtask fixtures. const SCHED_STANDALONE_MARKER = "SCHED_STANDALONE_INTERRUPT_RESUME" @@ -179,6 +189,81 @@ export function addSubtaskFixtures(mock: InstanceType) { }, }) + mock.addFixture({ + match: { + userMessage: new RegExp(SUBTASK_QUEUED_INPUT_PARENT_MARKER), + sequenceIndex: 0, + }, + response: { + toolCalls: [ + { + name: "new_task", + arguments: JSON.stringify({ + mode: "ask", + message: SUBTASK_QUEUED_INPUT_CHILD_PROMPT, + }), + id: "call_queued_input_parent_new_task_001", + }, + ], + }, + }) + + mock.addFixture({ + match: { + predicate: (req: ChatCompletionRequest) => + lastUserMessageContains(req, SUBTASK_QUEUED_INPUT_CHILD_MARKER) && + !requestContains(req, [SUBTASK_QUEUED_INPUT_PARENT_MARKER]) && + !requestContains(req, [SUBTASK_QUEUED_INPUT_MESSAGE]), + }, + streamingProfile: { ttft: SUBTASK_QUEUED_INPUT_RESPONSE_LATENCY_MS }, + response: { + toolCalls: [ + { + name: "attempt_completion", + arguments: JSON.stringify({ result: SUBTASK_QUEUED_INPUT_INITIAL_RESULT }), + id: "call_queued_input_child_initial_completion_002", + }, + ], + }, + }) + + mock.addFixture({ + match: { + predicate: (req: ChatCompletionRequest) => + requestContains(req, [SUBTASK_QUEUED_INPUT_CHILD_MARKER, SUBTASK_QUEUED_INPUT_MESSAGE]) && + !requestContains(req, [SUBTASK_QUEUED_INPUT_PARENT_MARKER]), + }, + response: { + toolCalls: [ + { + name: "attempt_completion", + arguments: JSON.stringify({ result: SUBTASK_QUEUED_INPUT_CHILD_RESULT }), + id: "call_queued_input_child_revised_completion_003", + }, + ], + }, + }) + + mock.addFixture({ + match: { + predicate: (req: ChatCompletionRequest) => + requestContains(req, [ + SUBTASK_QUEUED_INPUT_PARENT_MARKER, + SUBTASK_RESULT_INJECTION, + SUBTASK_QUEUED_INPUT_CHILD_RESULT, + ]), + }, + response: { + toolCalls: [ + { + name: "attempt_completion", + arguments: JSON.stringify({ result: SUBTASK_QUEUED_INPUT_PARENT_RESULT }), + id: "call_queued_input_parent_completion_004", + }, + ], + }, + }) + mock.addFixture({ match: { userMessage: new RegExp(SUBTASK_FAST_PARENT_MARKER), diff --git a/apps/vscode-e2e/src/runTest.ts b/apps/vscode-e2e/src/runTest.ts index 88c687bc76..e942649894 100644 --- a/apps/vscode-e2e/src/runTest.ts +++ b/apps/vscode-e2e/src/runTest.ts @@ -18,6 +18,7 @@ import { addTerminalProfileResultFixtures } from "./fixtures/terminal-profile" import { addListFilesResultFixtures } from "./fixtures/list-files" import { addReadFileResultFixtures } from "./fixtures/read-file" import { addSearchFilesResultFixtures } from "./fixtures/search-files" +import { addQueuedApiInputFixtures } from "./fixtures/queued-api-input" import { addSubtaskFixtures } from "./fixtures/subtasks" import { addUseMcpToolResultFixtures } from "./fixtures/use-mcp-tool" import { addWriteToFileResultFixtures } from "./fixtures/write-to-file" @@ -140,6 +141,7 @@ async function main() { addReadFileResultFixtures(mock) addSearchFilesResultFixtures(mock) addSubtaskFixtures(mock) + addQueuedApiInputFixtures(mock) addUseMcpToolResultFixtures(mock) addWriteToFileResultFixtures(mock) addDeepSeekV4Fixtures(mock) diff --git a/apps/vscode-e2e/src/suite/queued-api-input.test.ts b/apps/vscode-e2e/src/suite/queued-api-input.test.ts new file mode 100644 index 0000000000..cc1654505b --- /dev/null +++ b/apps/vscode-e2e/src/suite/queued-api-input.test.ts @@ -0,0 +1,99 @@ +import { providerIdentifiers } from "@roo-code/types" +import * as assert from "assert" +import * as fs from "fs/promises" +import * as path from "path" +import * as vscode from "vscode" + +import { RooCodeEventName, type ClineMessage, type QueuedMessage } from "@roo-code/types" + +import { setDefaultSuiteTimeout } from "./test-utils" +import { sleep, waitFor } from "./utils" +import { + QUEUED_API_INPUT_MARKER_FILE, + QUEUED_API_INPUT_MESSAGE, + QUEUED_API_INPUT_PROMPT, +} from "../fixtures/queued-api-input" + +suite("Roo Code queued API input", function () { + setDefaultSuiteTimeout(this) + + let markerPath: string + + suiteSetup(async () => { + const aimockUrl = process.env.AIMOCK_URL + await globalThis.api.setConfiguration({ + apiProvider: providerIdentifiers.openrouter, + openRouterApiKey: aimockUrl ? "mock-key" : process.env.OPENROUTER_API_KEY!, + openRouterModelId: "anthropic/claude-sonnet-4.5", + ...(aimockUrl && { openRouterBaseUrl: `${aimockUrl}/v1` }), + }) + markerPath = path.join(vscode.workspace.workspaceFolders![0]!.uri.fsPath, QUEUED_API_INPUT_MARKER_FILE) + }) + + teardown(async () => { + await fs.rm(markerPath, { force: true }) + while (globalThis.api.getCurrentTaskStack().length > 0) { + await globalThis.api.clearCurrentTask() + } + }) + + test("queued API input does not approve a protected command ask", async () => { + const api = globalThis.api + await fs.rm(markerPath, { force: true }) + + let queue: QueuedMessage[] = [] + let commandAsk: ClineMessage | undefined + let interactive = false + + const onQueue = (_taskId: string, messages: QueuedMessage[]) => { + if (messages.length > 0) queue = messages + } + const onMessage = ({ message }: { message: ClineMessage }) => { + if (message.type === "ask" && message.ask === "command" && message.partial !== true) { + commandAsk = message + } + } + const onInteractive = () => { + interactive = true + } + + api.on(RooCodeEventName.QueuedMessagesUpdated, onQueue) + api.on(RooCodeEventName.Message, onMessage) + api.on(RooCodeEventName.TaskInteractive, onInteractive) + + try { + const taskId = await api.startNewTask({ + configuration: { mode: "code", autoApprovalEnabled: false, enableCheckpoints: false }, + text: QUEUED_API_INPUT_PROMPT, + }) + + // Wait until the first request reaches the mock. The mock then delays its + // response, so the task is still streaming when the API input arrives. + await waitFor(async () => { + const response = await fetch(`${process.env.AIMOCK_URL}/__aimock/journal`) + return JSON.stringify(await response.json()).includes(QUEUED_API_INPUT_PROMPT) + }) + await api.sendMessage(QUEUED_API_INPUT_MESSAGE) + + await waitFor(() => commandAsk !== undefined) + await waitFor(() => interactive) + await sleep(1_500) + + assert.ok(queue.some((m) => m.text === QUEUED_API_INPUT_MESSAGE && m.origin === "api")) + await assert.rejects(fs.access(markerPath), "Queued API input must not run the command") + assert.strictEqual(api.getCurrentTaskStack().at(-1), taskId) + + await api.approveCurrentAsk() + await waitFor(() => + fs.access(markerPath).then( + () => true, + () => false, + ), + ) + } finally { + api.off(RooCodeEventName.QueuedMessagesUpdated, onQueue) + api.off(RooCodeEventName.Message, onMessage) + api.off(RooCodeEventName.TaskInteractive, onInteractive) + } + }) +}) diff --git a/apps/vscode-e2e/src/suite/subtasks.test.ts b/apps/vscode-e2e/src/suite/subtasks.test.ts index 857c8accc5..0da28c1aca 100644 --- a/apps/vscode-e2e/src/suite/subtasks.test.ts +++ b/apps/vscode-e2e/src/suite/subtasks.test.ts @@ -27,6 +27,12 @@ import { SUBTASK_INTERRUPT_PARENT_PROMPT, SUBTASK_INTERRUPT_PARENT_RESULT, SUBTASK_PARENT_PROMPT, + SUBTASK_QUEUED_INPUT_CHILD_MARKER, + SUBTASK_QUEUED_INPUT_CHILD_RESULT, + SUBTASK_QUEUED_INPUT_MESSAGE, + SUBTASK_QUEUED_INPUT_PARENT_MARKER, + SUBTASK_QUEUED_INPUT_PARENT_PROMPT, + SUBTASK_QUEUED_INPUT_PARENT_RESULT, SUBTASK_XPROFILE_DIFFERENT_CHILD_RESULT, SUBTASK_XPROFILE_PARENT_PROMPT, SUBTASK_XPROFILE_PARENT_RESULT, @@ -260,6 +266,73 @@ suite("Roo Code Subtasks", function () { } }) + test("queued input interrupts child completion before the parent resumes", async () => { + const api = globalThis.api + const says: Record = {} + + const messageHandler = ({ taskId, message }: { taskId: string; message: ClineMessage }) => { + if (message.type === "say" && message.partial !== true) { + says[taskId] = says[taskId] || [] + says[taskId].push(message) + } + } + + api.on(RooCodeEventName.Message, messageHandler) + + try { + const parentTaskId = await api.startNewTask({ + configuration: { + mode: "ask", + alwaysAllowModeSwitch: true, + alwaysAllowSubtasks: true, + autoApprovalEnabled: true, + enableCheckpoints: false, + }, + text: SUBTASK_QUEUED_INPUT_PARENT_PROMPT, + }) + + let childTaskId: string | undefined + await waitFor(() => { + const current = api.getCurrentTaskStack().at(-1) + if (current && current !== parentTaskId) { + childTaskId = current + return true + } + return false + }) + + await waitForAimockRequestContaining(SUBTASK_QUEUED_INPUT_CHILD_MARKER, SUBTASK_QUEUED_INPUT_PARENT_MARKER) + + const completedParentTaskId = await waitUntilCompleted({ + api, + start: async () => { + await api.sendMessage(SUBTASK_QUEUED_INPUT_MESSAGE) + return parentTaskId + }, + }) + + assert.strictEqual(completedParentTaskId, parentTaskId) + assert.ok( + says[childTaskId!]?.some( + ({ say, text }) => + say === "completion_result" && text?.trim() === SUBTASK_QUEUED_INPUT_CHILD_RESULT, + ), + "Child should process the queued instruction before returning to its parent", + ) + assert.strictEqual( + says[parentTaskId]?.find(({ say }) => say === "completion_result")?.text?.trim(), + SUBTASK_QUEUED_INPUT_PARENT_RESULT, + "Parent should resume only after the child processes the queued instruction", + ) + } finally { + api.off(RooCodeEventName.Message, messageHandler) + while (api.getCurrentTaskStack().length > 0) { + await api.clearCurrentTask() + } + await waitFor(() => api.getCurrentTaskStack().length === 0).catch(() => {}) + } + }) + // Smoke: child completing normally must resume the parent task. test("child task returns to parent after normal completion", async () => { const api = globalThis.api diff --git a/packages/types/src/api.ts b/packages/types/src/api.ts index de23f67491..85c53d57c5 100644 --- a/packages/types/src/api.ts +++ b/packages/types/src/api.ts @@ -90,7 +90,10 @@ export interface RooCodeAPI extends EventEmitter { */ abandonSubtask(childTaskId: string): Promise /** - * Sends a message to the current task. + * Sends a message to the current task as conversational input. + * If the task is busy the message is queued and becomes the next user + * turn. Queued input never approves a pending or later tool, command, or + * MCP ask; use approveCurrentAsk() for explicit approval. * @param message Optional message to send. * @param images Optional array of image data URIs (e.g., "data:image/webp;base64,..."). */ diff --git a/packages/types/src/message.ts b/packages/types/src/message.ts index 01d7962266..f8c3c1bcd3 100644 --- a/packages/types/src/message.ts +++ b/packages/types/src/message.ts @@ -372,6 +372,12 @@ export const queuedMessageSchema = z.object({ id: z.string(), text: z.string(), images: z.array(z.string()).optional(), + /** + * Where the message was queued from. Absent means the interactive + * webview, whose queued input may answer approval asks. "api" input is + * conversational steering and never answers an approval ask. + */ + origin: z.enum(["webview", "api"]).optional(), }) export type QueuedMessage = z.infer diff --git a/packages/types/src/vscode-extension-host.ts b/packages/types/src/vscode-extension-host.ts index c0e8509105..ba6e2f8a4a 100644 --- a/packages/types/src/vscode-extension-host.ts +++ b/packages/types/src/vscode-extension-host.ts @@ -129,6 +129,8 @@ export interface ExtensionMessage { | "switchTab" | "toggleAutoApprove" invoke?: "newChat" | "sendMessage" | "primaryButtonClick" | "secondaryButtonClick" | "setChatBoxMessage" + /** Origin of an `invoke: "sendMessage"` request. Absent means the interactive webview. */ + origin?: QueuedMessage["origin"] /** * Partial state updates are allowed to reduce message size (e.g. omit large fields like taskHistory). * The webview is responsible for merging. @@ -662,6 +664,8 @@ export interface WebviewMessage { askResponse?: ClineAskResponse apiConfiguration?: ProviderSettings images?: string[] + /** Origin of a `queueMessage` request. Absent means the interactive webview. */ + origin?: QueuedMessage["origin"] bool?: boolean value?: number stepIndex?: number diff --git a/src/core/message-queue/MessageQueueService.ts b/src/core/message-queue/MessageQueueService.ts index 85a2192217..e0162b8338 100644 --- a/src/core/message-queue/MessageQueueService.ts +++ b/src/core/message-queue/MessageQueueService.ts @@ -4,6 +4,11 @@ import { v4 as uuidv4 } from "uuid" import { QueuedMessage } from "@roo-code/types" +export interface AddMessageOptions { + /** Origin of the input. See {@link QueuedMessage.origin}. */ + origin?: QueuedMessage["origin"] +} + export interface MessageQueueState { messages: QueuedMessage[] isProcessing: boolean @@ -34,7 +39,7 @@ export class MessageQueueService extends EventEmitter { return { index, message: this._messages[index] } } - public addMessage(text: string, images?: string[]): QueuedMessage | undefined { + public addMessage(text: string, images?: string[], options?: AddMessageOptions): QueuedMessage | undefined { if (!text && !images?.length) { return undefined } @@ -44,6 +49,7 @@ export class MessageQueueService extends EventEmitter { id: uuidv4(), text, images, + origin: options?.origin, } this._messages.push(message) diff --git a/src/core/task/Task.ts b/src/core/task/Task.ts index 5ef3e54b56..84f077271b 100644 --- a/src/core/task/Task.ts +++ b/src/core/task/Task.ts @@ -210,7 +210,16 @@ export function isBlanketDenyEngaged( ) } -function queuedResponseForAsk(type: ClineAsk, text?: string): QueuedAskResolution | undefined { +function queuedResponseForAsk( + type: ClineAsk, + text: string | undefined, + origin: QueuedMessage["origin"], +): QueuedAskResolution | undefined { + // API-origin queued input is conversational steering, not an approval. + // Returning undefined keeps the message queued for the next turn and + // leaves the ask to auto-approval settings or an explicit decision. + const canApprove = origin !== "api" + if (type === "command_output") { return undefined } @@ -223,12 +232,13 @@ function queuedResponseForAsk(type: ClineAsk, text?: string): QueuedAskResolutio } } catch { // Malformed tool asks retain the existing approve-with-feedback behavior. + return canApprove ? { response: "yesButtonClicked", requiresDurableAck: false } : undefined } - return { response: "yesButtonClicked", requiresDurableAck: false } + return canApprove ? { response: "yesButtonClicked", requiresDurableAck: false } : undefined } if (type === "command" || type === "use_mcp_server") { - return { response: "yesButtonClicked", requiresDurableAck: false } + return canApprove ? { response: "yesButtonClicked", requiresDurableAck: false } : undefined } return { response: "messageResponse", requiresDurableAck: type === "completion_result" } @@ -1319,9 +1329,9 @@ export class Task extends EventEmitter implements TaskLike { * blanket-denied (it answers that denial, not the current ask), and * `hasUnclaimed()` replaces the length-only `isEmpty()`, which reports a * queue containing nothing but claims as available for a new consumer. - * `isMessageQueued`/`isStatusMutable` keep `isEmpty()` semantics on purpose: - * flipping those would re-enable interactive prompt timers whenever a claim - * is outstanding. + * `isStatusMutable` does not read queue length: it keys on whether a claimed + * message will answer the ask, so a claim that answers suppresses the + * interactive prompt timers and a released claim leaves them armed. */ private mayDrainQueuedMessageForAsk(): boolean { return !this.blanketDeniedCommandThisTurn && this.messageQueueService.hasUnclaimed() @@ -1859,7 +1869,14 @@ export class Task extends EventEmitter implements TaskLike { !this.mayDrainQueuedMessageForAsk() ? undefined : this.messageQueueService.claimNextMessage() - const queuedAskResolution = queuedMessage ? queuedResponseForAsk(type, text) : undefined + const queuedAskResolution = queuedMessage ? queuedResponseForAsk(type, text, queuedMessage.origin) : undefined + + if (queuedMessage && !queuedAskResolution) { + // API-origin input cannot answer this ask. Release the claim so the + // message stays queued as the next conversational turn. + this.messageQueueService.releaseMessage(queuedMessage.id) + } + // `this.cwd`, not `provider.cwd`: // The path inside `text` was made relative to this task's workspace, // which for a resumed or child task need not be the one the provider @@ -2032,20 +2049,23 @@ export class Task extends EventEmitter implements TaskLike { // The state is mutable if the message is complete and the task will // block (via the `pWaitFor`). const isBlocking = !(this.askResponse !== undefined || this.lastMessageTs !== askTs) - const isMessageQueued = !this.messageQueueService.isEmpty() // Keep queued user messages intact during command_output asks. Those asks // are terminal flow-control, not conversational turns. const shouldDrainQueuedMessageForAsk = type !== "command_output" - const isStatusMutable = !partial && isBlocking && !isMessageQueued && approval.decision === "ask" + // The FIFO drain answers this ask only from the claimed head message. + // API-origin steering released above stays queued for the next turn, so + // queue emptiness cannot decide whether the ask waits for a response. + const isStatusMutable = + !partial && isBlocking && !(queuedMessage && queuedAskResolution) && approval.decision === "ask" let queuedMessageId: string | undefined // Arm the interactive/resumable/idle status timers for this ask: the // single source of that arm, shared between the queue-free case and the // claim-gated and queued-release paths below. A gated or released claim - // keeps the message in the queue, so `isMessageQueued` stays true and - // `isStatusMutable` — which requires an empty queue — stays false while - // the ask waits for the user; arming only from `isStatusMutable` would - // leave hands-free/API consumers seeing `Running` with no + // keeps the message in the queue, but `isStatusMutable` no longer reads + // queue length: with no claimed message answering the ask it stays true + // while the ask waits for the user, and arming only from the drain sites + // would leave hands-free/API consumers seeing `Running` with no // `TaskInteractive`/`interactionRequired` for a prompt that is in fact // pending. Idempotent: several arm sites can fire for one ask (e.g. the // queue-free arm, then a drain-site release), and a second arm would @@ -2183,13 +2203,6 @@ export class Task extends EventEmitter implements TaskLike { } else { queuedMessageId = this.handleQueuedAskResponse(queuedMessage, queuedAskResolution) } - } else if (shouldDrainQueuedMessageForAsk && isMessageQueued) { - // The claim gate (per-turn latch, or blanket deny engaged for a command - // ask) left the queued message untouched. If the policy still leaves the - // prompt pending, the non-empty queue keeps `isStatusMutable` false, so - // the interactive arm must run from here — the same reason a release - // re-arms. For an auto-answered ask the arm's pending check declines. - armAskStatusTimers() } // At most one drain-site policy re-check is in flight per ask; the @@ -2225,9 +2238,9 @@ export class Task extends EventEmitter implements TaskLike { ) if (!this.abort && this.askResponse === undefined && this.lastMessageTs === askTs) { queuedMessageId = this.applyQueuedCommandPolicyAction(action, message, resolution) - // A "release" outcome leaves the ask pending while the queue stays - // non-empty, so the arm that `isStatusMutable` gates — computed - // once, before the claim — must run here. + // A "release" outcome leaves the ask pending, and `isStatusMutable` + // — computed once, before the claim — still counts the message as + // answering the ask, so the arm must run here. armAskStatusTimers() } } finally { @@ -2263,7 +2276,7 @@ export class Task extends EventEmitter implements TaskLike { this.mayDrainQueuedMessageForAsk() ) { const message = this.messageQueueService.claimNextMessage() - const resolution = message ? queuedResponseForAsk(type, text) : undefined + const resolution = message ? queuedResponseForAsk(type, text, message.origin) : undefined if (message && resolution) { if (type === "command") { // Claim first, then verify the policy off-predicate: a @@ -2283,6 +2296,10 @@ export class Task extends EventEmitter implements TaskLike { } else { queuedMessageId = this.handleQueuedAskResponse(message, resolution) } + } else if (message) { + // API-origin input cannot answer this ask. Release the + // claim so the message stays queued for the next turn. + this.messageQueueService.releaseMessage(message.id) } } diff --git a/src/core/task/__tests__/ask-queued-message-drain.spec.ts b/src/core/task/__tests__/ask-queued-message-drain.spec.ts index b137130174..d83cf4ab59 100644 --- a/src/core/task/__tests__/ask-queued-message-drain.spec.ts +++ b/src/core/task/__tests__/ask-queued-message-drain.spec.ts @@ -1,9 +1,14 @@ +import { RooCodeEventName } from "@roo-code/types" + import { Task } from "../Task" type QueueTaskTestAccess = { say: Task["say"] saveClineMessages: () => Promise - addToClineMessages: () => Promise + addToClineMessages: (message?: unknown) => Promise + clineMessages: unknown[] + taskId: string + emit: (event: string, ...args: unknown[]) => boolean lastMessageTs?: number abort: boolean } @@ -266,6 +271,162 @@ describe("Task.ask queued message drain", () => { } }) + describe("API-origin queued input", () => { + it.each([ + ["command", "npm publish"], + ["use_mcp_server", '{"server_name":"fs","tool_name":"write_file"}'], + ["tool", JSON.stringify({ tool: "readFile", path: "src/a.ts" })], + ["tool", "not-json"], + ] as const)( + "does not approve a protected %s ask from API-origin input queued before the ask", + async (type, text) => { + const task = await createTask() + task.messageQueueService.addMessage("Steer the next turn", undefined, { origin: "api" }) + + const askPromise = task.ask(type, text, false) + setTimeout(() => task.denyAsk(), 0) + const result = await askPromise + + expect(result.response).not.toBe("yesButtonClicked") + expect(result.response).toBe("noButtonClicked") + expect(task.messageQueueService.messages).toHaveLength(1) + expect(task.messageQueueService.claimNextMessage()?.text).toBe("Steer the next turn") + }, + ) + + it("does not approve a protected ask from API-origin input arriving while the ask waits", async () => { + const task = await createTask() + const askPromise = task.ask("command", "npm publish", false) + await vi.waitFor(() => expect(getQueueTaskTestAccess(task).lastMessageTs).toBeDefined()) + + task.messageQueueService.addMessage("Late steering", undefined, { origin: "api" }) + setTimeout(() => task.denyAsk(), 150) + const result = await askPromise + + expect(result.response).toBe("noButtonClicked") + expect(task.messageQueueService.messages).toHaveLength(1) + expect(task.messageQueueService.claimNextMessage()?.text).toBe("Late steering") + }) + + it("still answers conversational asks from API-origin input", async () => { + const task = await createTask() + task.messageQueueService.addMessage("Use the queue module", undefined, { origin: "api" }) + + const result = await task.ask("followup", "Where should this go?", false) + + expect(result).toMatchObject({ response: "messageResponse", text: "Use the queue module" }) + expect(task.messageQueueService.isEmpty()).toBe(true) + }) + + it("emits TaskInteractive after the status delay when API-origin input cannot answer a protected ask", async () => { + vi.useFakeTimers() + try { + const task = await createTask() + const access = getQueueTaskTestAccess(task) + const taskId = "status-regression" + access.taskId = taskId + access.addToClineMessages = vi.fn(async (message: unknown) => { + access.clineMessages.push(message) + }) + const emit = access.emit + task.messageQueueService.addMessage("Steer the next turn", undefined, { origin: "api" }) + + const askPromise = task.ask("command", "npm publish", false) + await vi.advanceTimersByTimeAsync(2_000) + + expect(emit).toHaveBeenCalledTimes(1) + expect(emit).toHaveBeenCalledWith(RooCodeEventName.TaskInteractive, taskId) + expect(task.messageQueueService.messages).toMatchObject([ + { text: "Steer the next turn", origin: "api" }, + ]) + + task.denyAsk() + await vi.advanceTimersByTimeAsync(1_000) + const result = await askPromise + + expect(result.response).toBe("noButtonClicked") + expect(result.text).toBeUndefined() + expect(task.messageQueueService.messages).toMatchObject([ + { text: "Steer the next turn", origin: "api" }, + ]) + } finally { + vi.useRealTimers() + } + }) + + it("emits TaskInteractive when API-origin input precedes a webview message that could answer", async () => { + vi.useFakeTimers() + try { + const task = await createTask() + const access = getQueueTaskTestAccess(task) + const taskId = "status-regression-fifo" + access.taskId = taskId + access.addToClineMessages = vi.fn(async (message: unknown) => { + access.clineMessages.push(message) + }) + const emit = access.emit + task.messageQueueService.addMessage("Steer the next turn", undefined, { origin: "api" }) + task.messageQueueService.addMessage("Approve this one") + + const askPromise = task.ask("command", "npm publish", false) + await vi.advanceTimersByTimeAsync(2_000) + + expect(emit).toHaveBeenCalledWith(RooCodeEventName.TaskInteractive, taskId) + + task.denyAsk() + await vi.advanceTimersByTimeAsync(1_000) + const result = await askPromise + + expect(result.response).toBe("noButtonClicked") + expect(task.messageQueueService.messages.map((message) => message.text)).toEqual([ + "Steer the next turn", + "Approve this one", + ]) + } finally { + vi.useRealTimers() + } + }) + + it("does not emit TaskInteractive when webview-origin input answers the protected ask", async () => { + const task = await createTask() + const access = getQueueTaskTestAccess(task) + const taskId = "status-drained" + access.taskId = taskId + const emit = access.emit + task.messageQueueService.addMessage("Approval context") + + const result = await task.ask("command", "npm publish", false) + + expect(result).toMatchObject({ response: "yesButtonClicked", text: "Approval context" }) + expect(task.messageQueueService.isEmpty()).toBe(true) + expect(emit).not.toHaveBeenCalledWith(RooCodeEventName.TaskInteractive, taskId) + }) + + it.each([ + ["resume_task", RooCodeEventName.TaskResumable], + ["completion_result", RooCodeEventName.TaskIdle], + ] as const)( + "still answers %s from API-origin input without emitting its status event", + async (type, statusEvent) => { + const task = await createTask() + const access = getQueueTaskTestAccess(task) + const taskId = "status-mapping" + access.taskId = taskId + const emit = access.emit + task.messageQueueService.addMessage("Continue with this", undefined, { origin: "api" }) + + const result = await task.ask(type, "Done", false) + + expect(result).toMatchObject({ response: "messageResponse", text: "Continue with this" }) + expect(emit).not.toHaveBeenCalledWith(statusEvent, taskId) + if (result.queuedMessageId) { + task.messageQueueService.removeMessage(result.queuedMessageId) + } + expect(task.messageQueueService.isEmpty()).toBe(true) + }, + ) + }) + it("releases durable queued feedback when the task aborts during retry backoff", async () => { vi.useFakeTimers() try { diff --git a/src/core/webview/__tests__/webviewMessageHandler.queueMessage.spec.ts b/src/core/webview/__tests__/webviewMessageHandler.queueMessage.spec.ts new file mode 100644 index 0000000000..b3858a727d --- /dev/null +++ b/src/core/webview/__tests__/webviewMessageHandler.queueMessage.spec.ts @@ -0,0 +1,79 @@ +import type { Mock } from "vitest" +import { describe, it, expect, vi, beforeEach } from "vitest" + +vi.mock("vscode", () => ({ + window: { + showWarningMessage: vi.fn(), + showErrorMessage: vi.fn(), + }, + workspace: { + workspaceFolders: [{ uri: { fsPath: "/mock/workspace" } }], + getConfiguration: vi.fn().mockReturnValue({ + get: vi.fn(), + update: vi.fn(), + }), + }, + Uri: { + file: vi.fn((path) => ({ fsPath: path })), + }, + env: { + uriScheme: "vscode", + }, +})) + +vi.mock("../../../mentions/resolveImageMentions", () => ({ + resolveImageMentions: vi.fn(async (payload: { text: string; images?: string[] }) => ({ + text: payload.text, + images: payload.images, + })), +})) + +import { webviewMessageHandler } from "../webviewMessageHandler" +import type { ClineProvider } from "../ClineProvider" +import type { WebviewMessage } from "@roo-code/types" +import { MessageQueueService } from "../../message-queue/MessageQueueService" + +describe("webviewMessageHandler - queueMessage origin", () => { + let mockClineProvider: ClineProvider + let messageQueueService: MessageQueueService + + beforeEach(() => { + vi.clearAllMocks() + + messageQueueService = new MessageQueueService() + + mockClineProvider = { + getCurrentTask: vi.fn().mockReturnValue({ messageQueueService }), + postMessageToWebview: vi.fn(), + contextProxy: { + getValue: vi.fn(), + setValue: vi.fn(), + globalStorageUri: { fsPath: "/mock/storage" }, + }, + getState: vi.fn().mockResolvedValue({ + maxImageFileSize: 5, + maxTotalImageSize: 20, + }), + log: vi.fn(), + } as unknown as ClineProvider + }) + + it("stores api origin on the queued message", async () => { + const message: WebviewMessage = { type: "queueMessage", text: "steer the task", images: [], origin: "api" } + + await webviewMessageHandler(mockClineProvider, message) + + expect((mockClineProvider.getCurrentTask as Mock)().messageQueueService).toBe(messageQueueService) + expect(messageQueueService.messages).toMatchObject([{ text: "steer the task", origin: "api" }]) + }) + + it("keeps webview input unattributed so it can still answer approval asks", async () => { + const message: WebviewMessage = { type: "queueMessage", text: "human note", images: [] } + + await webviewMessageHandler(mockClineProvider, message) + + expect(messageQueueService.messages).toHaveLength(1) + expect(messageQueueService.messages[0].text).toBe("human note") + expect(messageQueueService.messages[0].origin).toBeUndefined() + }) +}) diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 193540455b..7c19ab0e75 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -3826,7 +3826,9 @@ export const webviewMessageHandler = async ( case "queueMessage": { const resolved = await resolveIncomingImages({ text: message.text, images: message.images }) - provider.getCurrentTask()?.messageQueueService.addMessage(resolved.text, resolved.images) + provider.getCurrentTask()?.messageQueueService.addMessage(resolved.text, resolved.images, { + origin: message.origin, + }) break } case "removeQueuedMessage": { diff --git a/src/extension/__tests__/api-send-message.spec.ts b/src/extension/__tests__/api-send-message.spec.ts index 5158a218e7..77b0b2a0e1 100644 --- a/src/extension/__tests__/api-send-message.spec.ts +++ b/src/extension/__tests__/api-send-message.spec.ts @@ -52,6 +52,7 @@ describe("API - SendMessage Command", () => { expect(mockPostMessageToWebview).toHaveBeenCalledWith({ type: "invoke", invoke: "sendMessage", + origin: "api", text: messageText, images: undefined, }) @@ -71,6 +72,7 @@ describe("API - SendMessage Command", () => { expect(mockPostMessageToWebview).toHaveBeenCalledWith({ type: "invoke", invoke: "sendMessage", + origin: "api", text: messageText, images, }) @@ -89,6 +91,7 @@ describe("API - SendMessage Command", () => { expect(mockPostMessageToWebview).toHaveBeenCalledWith({ type: "invoke", invoke: "sendMessage", + origin: "api", text: undefined, images, }) @@ -102,6 +105,7 @@ describe("API - SendMessage Command", () => { expect(mockPostMessageToWebview).toHaveBeenCalledWith({ type: "invoke", invoke: "sendMessage", + origin: "api", text: undefined, images: undefined, }) @@ -128,6 +132,7 @@ describe("API - SendMessage Command", () => { expect(mockPostMessageToWebview).toHaveBeenCalledWith({ type: "invoke", invoke: "sendMessage", + origin: "api", text: messageText, images: undefined, }) @@ -149,6 +154,7 @@ describe("API - SendMessage Command", () => { expect(mockPostMessageToWebview).toHaveBeenCalledWith({ type: "invoke", invoke: "sendMessage", + origin: "api", text: messageText, images, }) diff --git a/src/extension/__tests__/api.spec.ts b/src/extension/__tests__/api.spec.ts new file mode 100644 index 0000000000..689e08227d --- /dev/null +++ b/src/extension/__tests__/api.spec.ts @@ -0,0 +1,234 @@ +import type * as vscode from "vscode" +import { IpcMessageType, TaskCommandName, type ClineMessage, type IpcMessage } from "@roo-code/types" + +import { API } from "../api" +import type { ClineProvider } from "../../core/webview/ClineProvider" +import { MessageQueueService } from "../../core/message-queue/MessageQueueService" +import { Task } from "../../core/task/Task" +import { makeClineProviderFactory } from "../../test-utils/provider" + +vi.mock("vscode") +vi.mock("../../core/webview/ClineProvider") + +type TaskCommandHandler = ( + clientId: string, + command: Extract["data"], +) => Promise + +let taskCommandHandler: TaskCommandHandler | undefined + +type TaskTestAccess = { + addToClineMessages: (message: ClineMessage) => Promise +} + +const createStreamingTask = (provider: object) => { + const task = Object.create(Task.prototype) as Task + Object.assign(task, { + abort: false, + clineMessages: [], + taskId: "task-1", + instanceId: "instance-1", + isStreaming: true, + messageQueueService: new MessageQueueService(), + providerRef: { deref: () => provider }, + addToClineMessages: vi.fn(async () => {}), + saveClineMessages: vi.fn(async () => true), + updateClineMessage: vi.fn(async () => {}), + cancelAutoApprovalTimeout: vi.fn(), + checkpointSave: vi.fn(async () => {}), + emit: vi.fn(), + }) + vi.spyOn(task as unknown as TaskTestAccess, "addToClineMessages").mockImplementation(async (message) => { + task.clineMessages.push(message) + }) + return task +} + +vi.mock("@roo-code/ipc", () => ({ + IpcServer: class { + listen() {} + on(messageType: IpcMessageType, handler: TaskCommandHandler) { + if (messageType === IpcMessageType.TaskCommand) { + taskCommandHandler = handler + } + } + }, +})) + +describe("API.sendMessage", () => { + it("enqueues directly with api origin when the current webview task is streaming", async () => { + const addMessage = vi.fn() + const postMessageToWebview = vi.fn() + const provider = { + viewLaunched: true, + getCurrentTask: vi.fn().mockReturnValue({ + isStreaming: true, + messageQueueService: { addMessage }, + }), + getCurrentTaskStack: vi.fn().mockReturnValue([]), + postMessageToWebview, + on: vi.fn(), + } as unknown as ClineProvider + const api = new API({} as vscode.OutputChannel, provider, makeClineProviderFactory()) + const images = ["data:image/png;base64,image1data"] + + await api.sendMessage("Use this before completing", images) + + expect(addMessage).toHaveBeenCalledWith("Use this before completing", images, { origin: "api" }) + expect(postMessageToWebview).not.toHaveBeenCalled() + + addMessage.mockClear() + await api.sendMessage(undefined, images) + expect(addMessage).toHaveBeenCalledWith("", images, { origin: "api" }) + }) + + it("falls through to the webview invoke when the current task is not streaming", async () => { + const addMessage = vi.fn() + const postMessageToWebview = vi.fn() + const provider = { + viewLaunched: true, + getCurrentTask: vi.fn().mockReturnValue({ + isStreaming: false, + messageQueueService: { addMessage }, + }), + getCurrentTaskStack: vi.fn().mockReturnValue([]), + postMessageToWebview, + on: vi.fn(), + } as unknown as ClineProvider + const api = new API({} as vscode.OutputChannel, provider, makeClineProviderFactory()) + + await api.sendMessage("Done with the follow-up") + + expect(addMessage).not.toHaveBeenCalled() + expect(postMessageToWebview).toHaveBeenCalledWith({ + type: "invoke", + invoke: "sendMessage", + origin: "api", + text: "Done with the follow-up", + images: undefined, + }) + }) + + it("falls through to the webview invoke when there is no current task", async () => { + const postMessageToWebview = vi.fn() + const provider = { + viewLaunched: true, + getCurrentTask: vi.fn().mockReturnValue(undefined), + getCurrentTaskStack: vi.fn().mockReturnValue([]), + postMessageToWebview, + on: vi.fn(), + } as unknown as ClineProvider + const api = new API({} as vscode.OutputChannel, provider, makeClineProviderFactory()) + + await api.sendMessage("Start over") + + expect(postMessageToWebview).toHaveBeenCalledWith({ + type: "invoke", + invoke: "sendMessage", + origin: "api", + text: "Start over", + images: undefined, + }) + }) + + it("does not approve a protected ask from streaming IPC input queued before the ask", async () => { + const provider = { + context: {}, + cwd: "/test/cwd", + viewLaunched: true, + getState: vi.fn().mockResolvedValue({ autoApprovalEnabled: false }), + getCurrentTask: vi.fn(), + getCurrentTaskStack: vi.fn().mockReturnValue([]), + postMessageToWebview: vi.fn(), + on: vi.fn(), + } as unknown as ClineProvider + const task = createStreamingTask(provider) + vi.mocked(provider.getCurrentTask).mockReturnValue(task) + const api = new API({} as vscode.OutputChannel, provider, makeClineProviderFactory()) + const images = ["data:image/png;base64,image1data"] + + await api.sendMessage("Steer the next turn", images) + expect(task.messageQueueService.messages[0]?.origin).toBe("api") + + task.isStreaming = false + const ask = task.ask("command", "npm publish", false) + await vi.waitFor(() => expect(task.clineMessages).toHaveLength(1)) + setTimeout(() => task.denyAsk(), 0) + const result = await ask + + expect(result.response).not.toBe("yesButtonClicked") + expect(result.response).toBe("noButtonClicked") + expect(task.messageQueueService.messages).toHaveLength(1) + expect(task.messageQueueService.claimNextMessage()?.text).toBe("Steer the next turn") + }) + + it.each([ + ["command", "npm publish"], + ["use_mcp_server", '{"server_name":"filesystem","tool_name":"write_file"}'], + ] as const)("does not approve a protected headless %s ask from queued IPC input", async (askType, askText) => { + const appendLine = vi.fn() + const provider = { + context: {}, + cwd: "/test/cwd", + viewLaunched: false, + getState: vi.fn().mockResolvedValue({ autoApprovalEnabled: false }), + getCurrentTask: vi.fn(), + getCurrentTaskStack: vi.fn().mockReturnValue([]), + on: vi.fn(), + } as unknown as ClineProvider + const task = createStreamingTask(provider) + vi.mocked(provider.getCurrentTask).mockReturnValue(task) + new API( + { appendLine } as unknown as vscode.OutputChannel, + provider, + makeClineProviderFactory(), + "/tmp/roo-test.sock", + true, + ) + const images = ["data:image/png;base64,image1data"] + const executeProtectedTool = vi.fn() + const ask = task.ask(askType, askText, false) + await vi.waitFor(() => expect(task.clineMessages).toHaveLength(1)) + + await taskCommandHandler?.("client-1", { + commandName: TaskCommandName.SendMessage, + data: { text: "Use this before completing", images }, + }) + const result = await ask + if (result.response === "yesButtonClicked") { + executeProtectedTool() + } + + expect(appendLine).toHaveBeenCalledWith("[API] SendMessage -> Use this before completing") + expect(result).toMatchObject({ response: "messageResponse", text: "Use this before completing", images }) + expect(task.messageQueueService.isEmpty()).toBe(true) + expect(executeProtectedTool).not.toHaveBeenCalled() + }) + + it("logs rejected SendMessage commands without rejecting the IPC handler", async () => { + const appendLine = vi.fn() + const provider = { + context: {}, + cwd: "/test/cwd", + getCurrentTask: vi.fn(), + getCurrentTaskStack: vi.fn().mockReturnValue([]), + on: vi.fn(), + } as unknown as ClineProvider + const api = new API( + { appendLine } as unknown as vscode.OutputChannel, + provider, + makeClineProviderFactory(), + "/tmp/roo-test.sock", + true, + ) + vi.spyOn(api, "sendMessage").mockRejectedValue(new Error("invalid input")) + + await expect( + taskCommandHandler?.("client-1", { + commandName: TaskCommandName.SendMessage, + data: { text: "" }, + }), + ).resolves.toBeUndefined() + expect(appendLine).toHaveBeenCalledWith("[API] SendMessage failed: invalid input") + }) +}) diff --git a/src/extension/api.ts b/src/extension/api.ts index 9ae696adb1..642ebb9ccb 100644 --- a/src/extension/api.ts +++ b/src/extension/api.ts @@ -109,7 +109,12 @@ export class API extends EventEmitter implements RooCodeAPI { break case TaskCommandName.SendMessage: this.log(`[API] SendMessage -> ${command.data.text}`) - await this.sendMessage(command.data.text, command.data.images) + try { + await this.sendMessage(command.data.text, command.data.images) + } catch (error) { + const errorMessage = error instanceof Error ? error.message : String(error) + this.log(`[API] SendMessage failed: ${errorMessage}`) + } break case TaskCommandName.GetCommands: try { @@ -319,12 +324,16 @@ export class API extends EventEmitter implements RooCodeAPI { return this.sidebarProvider.abandonSubtask(childTaskId) } + /** + * Sends conversational input to the current task. Queued input never + * approves a protected ask; use approveCurrentAsk() for explicit approval. + */ public async sendMessage(text?: string, images?: string[]) { const currentTask = this.sidebarProvider.getCurrentTask() // In headless/sandbox flows the webview may not be launched, so routing - // through invoke=sendMessage drops the message. Deliver directly to the - // task ask-response channel instead. + // through invoke=sendMessage drops the message. Keep this path on the task + // ask-response channel; it resolves asks as feedback, never as approval. if (!this.sidebarProvider.viewLaunched) { if (!currentTask) { this.log("[API#sendMessage] no current task in headless mode; message dropped") @@ -335,7 +344,20 @@ export class API extends EventEmitter implements RooCodeAPI { return } - await this.sidebarProvider.postMessageToWebview({ type: "invoke", invoke: "sendMessage", text, images }) + // Ensure steering input reaches the active task before it can finish. + // Origin "api" keeps the queued message from answering protected asks. + if (currentTask?.isStreaming) { + currentTask.messageQueueService.addMessage(text ?? "", images, { origin: "api" }) + return + } + + await this.sidebarProvider.postMessageToWebview({ + type: "invoke", + invoke: "sendMessage", + text, + images, + origin: "api", + }) } public deleteQueuedMessage(messageId: string) { diff --git a/webview-ui/src/components/chat/ChatView.tsx b/webview-ui/src/components/chat/ChatView.tsx index bb690c1f03..bbea6e9a66 100644 --- a/webview-ui/src/components/chat/ChatView.tsx +++ b/webview-ui/src/components/chat/ChatView.tsx @@ -20,7 +20,15 @@ import { getCostBreakdownIfNeeded } from "@src/utils/costFormatting" import { batchNearby } from "@src/utils/batchNearby" import { isBoundary, isIgnorableBetweenTargets } from "@src/utils/chatBatchingPredicates" -import type { ClineAsk, ClineSayTool, ClineMessage, ExtensionMessage, AudioType, SuggestionItem } from "@roo-code/types" +import type { + ClineAsk, + ClineSayTool, + ClineMessage, + ExtensionMessage, + AudioType, + SuggestionItem, + QueuedMessage, +} from "@roo-code/types" import { getCompletionCheckpoint, getSuggestionMode, hasUsableAnswer, isRetiredProvider } from "@roo-code/types" import { findLast } from "@roo/array" @@ -640,7 +648,7 @@ const ChatViewComponent: React.ForwardRefRenderFunction { + (text: string, images: string[], origin?: QueuedMessage["origin"]) => { text = text.trim() if (text || images.length > 0) { @@ -664,7 +672,7 @@ const ChatViewComponent: React.ForwardRefRenderFunction { ) }) + it("keeps api origin when an API sendMessage request is queued", async () => { + const { getByTestId } = renderChatView() + + mockPostMessage({ + clineMessages: [ + { type: "say", say: "task", ts: Date.now() - 2000, text: "Initial task" }, + { + type: "say", + say: "api_req_started", + ts: Date.now(), + text: JSON.stringify({ apiProtocol: "anthropic" }), // No cost = still streaming + }, + ], + }) + + await waitFor(() => { + expect(getByTestId("chat-textarea")).toBeInTheDocument() + }) + vscodePostMessageMock.cleanup() + + await dispatchExtensionMessage({ type: "invoke", invoke: "sendMessage", text: "steer", origin: "api" }) + + await waitFor(() => { + expect(vscode.postMessage).toHaveBeenCalledWith({ + type: "queueMessage", + text: "steer", + images: [], + origin: "api", + }) + }) + }) + + it("queues an image-only API sendMessage request with empty text", async () => { + const { getByTestId } = renderChatView() + + mockPostMessage({ + clineMessages: [ + { type: "say", say: "task", ts: Date.now() - 2000, text: "Initial task" }, + { + type: "say", + say: "api_req_started", + ts: Date.now(), + text: JSON.stringify({ apiProtocol: "anthropic" }), // No cost = still streaming + }, + ], + }) + + await waitFor(() => { + expect(getByTestId("chat-textarea")).toBeInTheDocument() + }) + vscodePostMessageMock.cleanup() + + const images = ["data:image/png;base64,abc"] + await dispatchExtensionMessage({ type: "invoke", invoke: "sendMessage", images, origin: "api" }) + + await waitFor(() => { + expect(vscode.postMessage).toHaveBeenCalledWith({ type: "queueMessage", text: "", images, origin: "api" }) + }) + }) + it("sends messages normally when API request is complete (cost present)", async () => { const { getByTestId } = renderChatView()