|
| 1 | +// npx vitest run core/task/__tests__/Task.abort-reason-race.spec.ts |
| 2 | + |
| 3 | +import * as os from "os" |
| 4 | +import * as path from "path" |
| 5 | + |
| 6 | +import type { GlobalState, ProviderSettings } from "@roo-code/types" |
| 7 | +import { providerIdentifiers } from "@roo-code/types/provider-identifiers" |
| 8 | +import { findLast } from "../../../shared/array" |
| 9 | +import { TelemetryService } from "@roo-code/telemetry" |
| 10 | + |
| 11 | +import { Task } from "../Task" |
| 12 | +import { ClineProvider } from "../../webview/ClineProvider" |
| 13 | + |
| 14 | +const { mockSaveTaskMessages, mockSaveApiMessages } = vi.hoisted(() => ({ |
| 15 | + mockSaveTaskMessages: vi.fn().mockResolvedValue(undefined), |
| 16 | + mockSaveApiMessages: vi.fn().mockResolvedValue(undefined), |
| 17 | +})) |
| 18 | + |
| 19 | +// vscode is globally aliased to __mocks__/vscode.js by vitest.config.ts. |
| 20 | +vi.mock("delay", () => ({ __esModule: true, default: vi.fn().mockResolvedValue(undefined) })) |
| 21 | +vi.mock("execa", () => ({ execa: vi.fn() })) |
| 22 | +vi.mock("p-wait-for", () => ({ default: vi.fn().mockResolvedValue(undefined) })) |
| 23 | + |
| 24 | +vi.mock("../../task-persistence", async (importOriginal) => { |
| 25 | + const mod = await importOriginal<typeof import("../../task-persistence")>() |
| 26 | + return { |
| 27 | + ...mod, |
| 28 | + saveApiMessages: mockSaveApiMessages, |
| 29 | + saveTaskMessages: mockSaveTaskMessages, |
| 30 | + TaskHistoryStore: vi.fn().mockImplementation(function () { |
| 31 | + return { |
| 32 | + initialize: vi.fn().mockResolvedValue(undefined), |
| 33 | + dispose: vi.fn(), |
| 34 | + get: vi.fn(), |
| 35 | + getAll: vi.fn().mockReturnValue([]), |
| 36 | + upsert: vi.fn().mockResolvedValue([]), |
| 37 | + delete: vi.fn().mockResolvedValue(undefined), |
| 38 | + deleteMany: vi.fn().mockResolvedValue(undefined), |
| 39 | + reconcile: vi.fn().mockResolvedValue(undefined), |
| 40 | + initialized: Promise.resolve(), |
| 41 | + } |
| 42 | + }), |
| 43 | + } |
| 44 | +}) |
| 45 | + |
| 46 | +vi.mock("../../mentions", () => ({ |
| 47 | + parseMentions: vi |
| 48 | + .fn() |
| 49 | + .mockImplementation((text) => |
| 50 | + Promise.resolve({ text: `processed: ${text}`, mode: undefined, contentBlocks: [] }), |
| 51 | + ), |
| 52 | + openMention: vi.fn(), |
| 53 | + getLatestTerminalOutput: vi.fn(), |
| 54 | +})) |
| 55 | + |
| 56 | +vi.mock("../../mentions/processUserContentMentions", () => ({ |
| 57 | + processUserContentMentions: vi.fn().mockImplementation(async ({ userContent }: { userContent: unknown[] }) => ({ |
| 58 | + content: userContent, |
| 59 | + mode: undefined, |
| 60 | + })), |
| 61 | +})) |
| 62 | + |
| 63 | +vi.mock("../../../integrations/misc/extract-text", () => ({ |
| 64 | + extractTextFromFile: vi.fn().mockResolvedValue(""), |
| 65 | +})) |
| 66 | + |
| 67 | +vi.mock("../../environment/getEnvironmentDetails", () => ({ |
| 68 | + getEnvironmentDetails: vi.fn().mockResolvedValue(""), |
| 69 | +})) |
| 70 | + |
| 71 | +vi.mock("../../ignore/RooIgnoreController") |
| 72 | + |
| 73 | +vi.mock("../../../utils/storage", () => ({ |
| 74 | + getTaskDirectoryPath: vi |
| 75 | + .fn() |
| 76 | + .mockImplementation((base: string, id: string) => Promise.resolve(`${base}/tasks/${id}`)), |
| 77 | + getSettingsDirectoryPath: vi.fn().mockImplementation((base: string) => Promise.resolve(`${base}/settings`)), |
| 78 | +})) |
| 79 | + |
| 80 | +vi.mock("../../../utils/fs", () => ({ |
| 81 | + fileExistsAtPath: vi.fn().mockReturnValue(false), |
| 82 | +})) |
| 83 | + |
| 84 | +vi.mock("../../../i18n", () => ({ |
| 85 | + t: (key: string) => key, |
| 86 | +})) |
| 87 | + |
| 88 | +function makeMockProvider() { |
| 89 | + return { |
| 90 | + log: vi.fn(), |
| 91 | + taskHistoryStore: { get: () => undefined }, |
| 92 | + updateTaskHistory: vi.fn().mockResolvedValue([]), |
| 93 | + getState: vi.fn().mockResolvedValue({}), |
| 94 | + getSkillsManager: vi.fn().mockReturnValue(undefined), |
| 95 | + postStateToWebviewWithoutTaskHistory: vi.fn().mockResolvedValue(undefined), |
| 96 | + flushPostStateToWebviewThrottled: vi.fn().mockResolvedValue(undefined), |
| 97 | + postMessageToWebview: vi.fn().mockResolvedValue(undefined), |
| 98 | + postStateToWebview: vi.fn().mockResolvedValue(undefined), |
| 99 | + postStateToWebviewThrottled: vi.fn().mockResolvedValue(undefined), |
| 100 | + context: { |
| 101 | + globalStorageUri: { fsPath: path.join(os.tmpdir(), "test-storage-abort-race") }, |
| 102 | + globalState: { |
| 103 | + get: vi.fn().mockImplementation((_key: keyof GlobalState) => undefined), |
| 104 | + update: vi.fn().mockResolvedValue(undefined), |
| 105 | + keys: vi.fn().mockReturnValue([]), |
| 106 | + }, |
| 107 | + workspaceState: { |
| 108 | + get: vi.fn().mockImplementation(() => undefined), |
| 109 | + update: vi.fn().mockResolvedValue(undefined), |
| 110 | + keys: vi.fn().mockReturnValue([]), |
| 111 | + }, |
| 112 | + secrets: { |
| 113 | + get: vi.fn().mockResolvedValue(undefined), |
| 114 | + store: vi.fn().mockResolvedValue(undefined), |
| 115 | + delete: vi.fn().mockResolvedValue(undefined), |
| 116 | + }, |
| 117 | + extensionUri: { fsPath: "/mock/extension" }, |
| 118 | + extension: { packageJSON: { version: "1.0.0" } }, |
| 119 | + }, |
| 120 | + } |
| 121 | +} |
| 122 | + |
| 123 | +describe("Task abort-reason race", () => { |
| 124 | + let mockApiConfig: ProviderSettings |
| 125 | + |
| 126 | + beforeEach(() => { |
| 127 | + vi.clearAllMocks() |
| 128 | + mockSaveTaskMessages.mockResolvedValue(undefined) |
| 129 | + |
| 130 | + if (!TelemetryService.hasInstance()) { |
| 131 | + TelemetryService.createInstance([]) |
| 132 | + } |
| 133 | + |
| 134 | + mockApiConfig = { |
| 135 | + apiProvider: providerIdentifiers.anthropic, |
| 136 | + apiModelId: "claude-3-5-sonnet-20241022", |
| 137 | + apiKey: "test-api-key", |
| 138 | + } |
| 139 | + }) |
| 140 | + |
| 141 | + it("keeps user_cancelled when a cancel lands during abortStream cleanup", async () => { |
| 142 | + const mockProvider = makeMockProvider() |
| 143 | + |
| 144 | + // Double cast: plain object satisfies only the methods called by this code path. |
| 145 | + const task = new Task({ |
| 146 | + provider: mockProvider as unknown as ClineProvider, |
| 147 | + apiConfiguration: mockApiConfig, |
| 148 | + task: "test task", |
| 149 | + startTask: false, |
| 150 | + }) |
| 151 | + |
| 152 | + vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined) |
| 153 | + vi.spyOn(task, "dispose").mockResolvedValue(undefined) |
| 154 | + const abortTaskSpy = vi.spyOn(task, "abortTask") |
| 155 | + |
| 156 | + // eslint-disable-next-line require-yield -- intentional error-only async source |
| 157 | + task["attemptApiRequest"] = async function* () { |
| 158 | + throw new Error("mid-stream API failure") |
| 159 | + } |
| 160 | + |
| 161 | + // abortStream calls updateApiReqMsg("streaming_failed") synchronously before saving, |
| 162 | + // so the api_req_started row has cancelReason="streaming_failed" when we are inside it. |
| 163 | + mockSaveTaskMessages.mockImplementation(async () => { |
| 164 | + const apiReqMsg = findLast(task.clineMessages, (m) => m.say === "api_req_started") |
| 165 | + const parsed = apiReqMsg?.text ? (JSON.parse(apiReqMsg.text) as { cancelReason?: string }) : {} |
| 166 | + if (parsed.cancelReason === "streaming_failed" && !task.didFinishAbortingStream) { |
| 167 | + // cancelTask sets abortReason before abort=true. |
| 168 | + task["abortReason"] = "user_cancelled" |
| 169 | + task["abort"] = true |
| 170 | + } |
| 171 | + }) |
| 172 | + |
| 173 | + // The outer catch in recursivelyMakeClineRequests swallows the abort throw and returns true. |
| 174 | + const result = await task.recursivelyMakeClineRequests([{ type: "text", text: "help me" }]) |
| 175 | + expect(result).toBe(true) |
| 176 | + expect(task.didFinishAbortingStream).toBe(true) |
| 177 | + expect(abortTaskSpy).toHaveBeenCalledOnce() |
| 178 | + expect(task.abortReason).toBe("user_cancelled") |
| 179 | + }) |
| 180 | + |
| 181 | + it("keeps pre-existing abortReason when cancel lands during retry backoff", async () => { |
| 182 | + const mockProvider = makeMockProvider() |
| 183 | + // autoApprovalEnabled causes the else branch to call backoffAndAnnounce. |
| 184 | + mockProvider.getState = vi.fn().mockResolvedValue({ autoApprovalEnabled: true }) |
| 185 | + |
| 186 | + // Double cast: plain object satisfies only the methods called by this code path. |
| 187 | + const task = new Task({ |
| 188 | + provider: mockProvider as unknown as ClineProvider, |
| 189 | + apiConfiguration: mockApiConfig, |
| 190 | + task: "test task", |
| 191 | + startTask: false, |
| 192 | + }) |
| 193 | + |
| 194 | + vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined) |
| 195 | + vi.spyOn(task, "dispose").mockResolvedValue(undefined) |
| 196 | + |
| 197 | + // eslint-disable-next-line require-yield -- intentional error-only async source |
| 198 | + task["attemptApiRequest"] = async function* () { |
| 199 | + throw new Error("mid-stream failure") |
| 200 | + } |
| 201 | + |
| 202 | + // Simulate cancel landing during backoff: cancelTask sets abortReason before abort=true. |
| 203 | + // Use a distinct seed value so ??= (keeps it) is distinguishable from = (overwrites). |
| 204 | + task["backoffAndAnnounce"] = async () => { |
| 205 | + task["abortReason"] = "streaming_failed" |
| 206 | + task["abort"] = true |
| 207 | + } |
| 208 | + |
| 209 | + const abortTaskSpy = vi.spyOn(task, "abortTask") |
| 210 | + // break after backoff exits the while loop; the outer try returns false. |
| 211 | + const result = await task.recursivelyMakeClineRequests([{ type: "text", text: "help me" }]) |
| 212 | + expect(result).toBe(false) |
| 213 | + expect(abortTaskSpy).toHaveBeenCalledOnce() |
| 214 | + expect(task.abortReason).toBe("streaming_failed") |
| 215 | + }) |
| 216 | +}) |
0 commit comments