diff --git a/src/core/task/Task.ts b/src/core/task/Task.ts index 84f6c2c964..8f1e523784 100644 --- a/src/core/task/Task.ts +++ b/src/core/task/Task.ts @@ -3715,8 +3715,8 @@ export class Task extends EventEmitter implements TaskLike { await abortStream(cancelReason, streamingFailedMessage) if (this.abort) { - // User cancelled - abort the entire task - this.abortReason = cancelReason + // ??= keeps the first reason; a cancel can land during abortStream after cancelReason was already computed. + this.abortReason ??= "user_cancelled" await this.abortTask() } else if (error instanceof OutputTokenLimitError) { // Truncation repeats on an identical request, so never auto-retry it @@ -3751,8 +3751,7 @@ export class Task extends EventEmitter implements TaskLike { console.log( `[Task#${this.taskId}.${this.instanceId}] Task aborted during mid-stream retry backoff`, ) - // Abort the entire task - this.abortReason = "user_cancelled" + this.abortReason ??= "user_cancelled" await this.abortTask() break } diff --git a/src/core/task/__tests__/Task.abort-reason-race.spec.ts b/src/core/task/__tests__/Task.abort-reason-race.spec.ts new file mode 100644 index 0000000000..7b71ee1dd6 --- /dev/null +++ b/src/core/task/__tests__/Task.abort-reason-race.spec.ts @@ -0,0 +1,216 @@ +// npx vitest run core/task/__tests__/Task.abort-reason-race.spec.ts + +import * as os from "os" +import * as path from "path" + +import type { GlobalState, ProviderSettings } from "@roo-code/types" +import { providerIdentifiers } from "@roo-code/types/provider-identifiers" +import { findLast } from "../../../shared/array" +import { TelemetryService } from "@roo-code/telemetry" + +import { Task } from "../Task" +import { ClineProvider } from "../../webview/ClineProvider" + +const { mockSaveTaskMessages, mockSaveApiMessages } = vi.hoisted(() => ({ + mockSaveTaskMessages: vi.fn().mockResolvedValue(undefined), + mockSaveApiMessages: vi.fn().mockResolvedValue(undefined), +})) + +// vscode is globally aliased to __mocks__/vscode.js by vitest.config.ts. +vi.mock("delay", () => ({ __esModule: true, default: vi.fn().mockResolvedValue(undefined) })) +vi.mock("execa", () => ({ execa: vi.fn() })) +vi.mock("p-wait-for", () => ({ default: vi.fn().mockResolvedValue(undefined) })) + +vi.mock("../../task-persistence", async (importOriginal) => { + const mod = await importOriginal() + return { + ...mod, + saveApiMessages: mockSaveApiMessages, + saveTaskMessages: mockSaveTaskMessages, + TaskHistoryStore: vi.fn().mockImplementation(function () { + return { + initialize: vi.fn().mockResolvedValue(undefined), + dispose: vi.fn(), + get: vi.fn(), + getAll: vi.fn().mockReturnValue([]), + upsert: vi.fn().mockResolvedValue([]), + delete: vi.fn().mockResolvedValue(undefined), + deleteMany: vi.fn().mockResolvedValue(undefined), + reconcile: vi.fn().mockResolvedValue(undefined), + initialized: Promise.resolve(), + } + }), + } +}) + +vi.mock("../../mentions", () => ({ + parseMentions: vi + .fn() + .mockImplementation((text) => + Promise.resolve({ text: `processed: ${text}`, mode: undefined, contentBlocks: [] }), + ), + openMention: vi.fn(), + getLatestTerminalOutput: vi.fn(), +})) + +vi.mock("../../mentions/processUserContentMentions", () => ({ + processUserContentMentions: vi.fn().mockImplementation(async ({ userContent }: { userContent: unknown[] }) => ({ + content: userContent, + mode: undefined, + })), +})) + +vi.mock("../../../integrations/misc/extract-text", () => ({ + extractTextFromFile: vi.fn().mockResolvedValue(""), +})) + +vi.mock("../../environment/getEnvironmentDetails", () => ({ + getEnvironmentDetails: vi.fn().mockResolvedValue(""), +})) + +vi.mock("../../ignore/RooIgnoreController") + +vi.mock("../../../utils/storage", () => ({ + getTaskDirectoryPath: vi + .fn() + .mockImplementation((base: string, id: string) => Promise.resolve(`${base}/tasks/${id}`)), + getSettingsDirectoryPath: vi.fn().mockImplementation((base: string) => Promise.resolve(`${base}/settings`)), +})) + +vi.mock("../../../utils/fs", () => ({ + fileExistsAtPath: vi.fn().mockReturnValue(false), +})) + +vi.mock("../../../i18n", () => ({ + t: (key: string) => key, +})) + +function makeMockProvider() { + return { + log: vi.fn(), + taskHistoryStore: { get: () => undefined }, + updateTaskHistory: vi.fn().mockResolvedValue([]), + getState: vi.fn().mockResolvedValue({}), + getSkillsManager: vi.fn().mockReturnValue(undefined), + postStateToWebviewWithoutTaskHistory: vi.fn().mockResolvedValue(undefined), + flushPostStateToWebviewThrottled: vi.fn().mockResolvedValue(undefined), + postMessageToWebview: vi.fn().mockResolvedValue(undefined), + postStateToWebview: vi.fn().mockResolvedValue(undefined), + postStateToWebviewThrottled: vi.fn().mockResolvedValue(undefined), + context: { + globalStorageUri: { fsPath: path.join(os.tmpdir(), "test-storage-abort-race") }, + globalState: { + get: vi.fn().mockImplementation((_key: keyof GlobalState) => undefined), + update: vi.fn().mockResolvedValue(undefined), + keys: vi.fn().mockReturnValue([]), + }, + workspaceState: { + get: vi.fn().mockImplementation(() => undefined), + update: vi.fn().mockResolvedValue(undefined), + keys: vi.fn().mockReturnValue([]), + }, + secrets: { + get: vi.fn().mockResolvedValue(undefined), + store: vi.fn().mockResolvedValue(undefined), + delete: vi.fn().mockResolvedValue(undefined), + }, + extensionUri: { fsPath: "/mock/extension" }, + extension: { packageJSON: { version: "1.0.0" } }, + }, + } +} + +describe("Task abort-reason race", () => { + let mockApiConfig: ProviderSettings + + beforeEach(() => { + vi.clearAllMocks() + mockSaveTaskMessages.mockResolvedValue(undefined) + + if (!TelemetryService.hasInstance()) { + TelemetryService.createInstance([]) + } + + mockApiConfig = { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-3-5-sonnet-20241022", + apiKey: "test-api-key", + } + }) + + it("keeps user_cancelled when a cancel lands during abortStream cleanup", async () => { + const mockProvider = makeMockProvider() + + // Double cast: plain object satisfies only the methods called by this code path. + const task = new Task({ + provider: mockProvider as unknown as ClineProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined) + vi.spyOn(task, "dispose").mockResolvedValue(undefined) + const abortTaskSpy = vi.spyOn(task, "abortTask") + + // eslint-disable-next-line require-yield -- intentional error-only async source + task["attemptApiRequest"] = async function* () { + throw new Error("mid-stream API failure") + } + + // abortStream calls updateApiReqMsg("streaming_failed") synchronously before saving, + // so the api_req_started row has cancelReason="streaming_failed" when we are inside it. + mockSaveTaskMessages.mockImplementation(async () => { + const apiReqMsg = findLast(task.clineMessages, (m) => m.say === "api_req_started") + const parsed = apiReqMsg?.text ? (JSON.parse(apiReqMsg.text) as { cancelReason?: string }) : {} + if (parsed.cancelReason === "streaming_failed" && !task.didFinishAbortingStream) { + // cancelTask sets abortReason before abort=true. + task["abortReason"] = "user_cancelled" + task["abort"] = true + } + }) + + // The outer catch in recursivelyMakeClineRequests swallows the abort throw and returns true. + const result = await task.recursivelyMakeClineRequests([{ type: "text", text: "help me" }]) + expect(result).toBe(true) + expect(task.didFinishAbortingStream).toBe(true) + expect(abortTaskSpy).toHaveBeenCalledOnce() + expect(task.abortReason).toBe("user_cancelled") + }) + + it("keeps pre-existing abortReason when cancel lands during retry backoff", async () => { + const mockProvider = makeMockProvider() + // autoApprovalEnabled causes the else branch to call backoffAndAnnounce. + mockProvider.getState = vi.fn().mockResolvedValue({ autoApprovalEnabled: true }) + + // Double cast: plain object satisfies only the methods called by this code path. + const task = new Task({ + provider: mockProvider as unknown as ClineProvider, + apiConfiguration: mockApiConfig, + task: "test task", + startTask: false, + }) + + vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined) + vi.spyOn(task, "dispose").mockResolvedValue(undefined) + + // eslint-disable-next-line require-yield -- intentional error-only async source + task["attemptApiRequest"] = async function* () { + throw new Error("mid-stream failure") + } + + // Simulate cancel landing during backoff: cancelTask sets abortReason before abort=true. + // Use a distinct seed value so ??= (keeps it) is distinguishable from = (overwrites). + task["backoffAndAnnounce"] = async () => { + task["abortReason"] = "streaming_failed" + task["abort"] = true + } + + const abortTaskSpy = vi.spyOn(task, "abortTask") + // break after backoff exits the while loop; the outer try returns false. + const result = await task.recursivelyMakeClineRequests([{ type: "text", text: "help me" }]) + expect(result).toBe(false) + expect(abortTaskSpy).toHaveBeenCalledOnce() + expect(task.abortReason).toBe("streaming_failed") + }) +}) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 9333329a38..92496e0db2 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -406,35 +406,7 @@ export class ClineProvider } this.emit(RooCodeEventName.TaskCompleted, taskId, tokenUsage, toolUsage) } - const onTaskAborted = async () => { - this.emit(RooCodeEventName.TaskAborted, instance.taskId) - - try { - // Only rehydrate on genuine streaming failures. - // User-initiated cancels are handled by cancelTask(). - if (instance.abortReason === "streaming_failed") { - // Defensive safeguard: if another path already replaced this instance, skip - const current = this.getCurrentTask() - if (current && current.instanceId !== instance.instanceId) { - this.log( - `[onTaskAborted] Skipping rehydrate: current instance ${current.instanceId} != aborted ${instance.instanceId}`, - ) - return - } - - const { historyItem } = await this.getTaskWithId(instance.taskId) - const rootTask = instance.rootTask - const parentTask = instance.parentTask - await this.createTaskWithHistoryItem({ ...historyItem, rootTask, parentTask }) - } - } catch (error) { - this.log( - `[onTaskAborted] Failed to rehydrate after streaming failure: ${ - error instanceof Error ? error.message : String(error) - }`, - ) - } - } + const onTaskAborted = () => this.emit(RooCodeEventName.TaskAborted, instance.taskId) const onTaskFocused = () => this.emit(RooCodeEventName.TaskFocused, instance.taskId) const onTaskUnfocused = () => this.emit(RooCodeEventName.TaskUnfocused, instance.taskId) const onTaskActive = (taskId: string) => this.emit(RooCodeEventName.TaskActive, taskId) diff --git a/src/core/webview/__tests__/ClineProvider.taskHistory.spec.ts b/src/core/webview/__tests__/ClineProvider.taskHistory.spec.ts index 61254a67d6..753332dd5b 100644 --- a/src/core/webview/__tests__/ClineProvider.taskHistory.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.taskHistory.spec.ts @@ -961,6 +961,25 @@ describe("ClineProvider Task History Synchronization", () => { expect(logSpy).toHaveBeenCalledWith(expect.stringContaining("[onTaskCompleted] Failed to write")) }) + it("onTaskAborted does not call createTaskWithHistoryItem", async () => { + // Store the item so getTaskWithId succeeds: without it, the old branch's catch + // swallows the error and the test passes even if the branch comes back. + const existing = createHistoryItem({ id: "task-abort-1", task: "T" }) + await provider.updateTaskHistory(existing, { broadcast: false }) + + const createSpy = vi.spyOn(provider, "createTaskWithHistoryItem") + const abortedListener = vi.fn() + provider.on(RooCodeEventName.TaskAborted, abortedListener) + + const fakeTask = { ...makeFakeTask("task-abort-1"), abortReason: "streaming_failed" } + // Double cast: taskCreationCallback reads only on/taskId/abortReason from the fake. + provider["taskCreationCallback"](fakeTask as unknown as Task) + await fakeTask.emit(RooCodeEventName.TaskAborted) + + expect(abortedListener).toHaveBeenCalledExactlyOnceWith("task-abort-1") + expect(createSpy).not.toHaveBeenCalled() + }) + it("emits delegated completion through the provider after the child is disposed", () => { const listener = vi.fn() provider.on(RooCodeEventName.TaskCompleted, listener)