diff --git a/src/core/task/Task.ts b/src/core/task/Task.ts index 88bd8141b8..e74d32bda7 100644 --- a/src/core/task/Task.ts +++ b/src/core/task/Task.ts @@ -3236,6 +3236,7 @@ export class Task extends EventEmitter implements TaskLike { this.didRejectTool = false this.didAlreadyUseTool = false this.assistantMessageSavedToHistory = false + this.didFinishAbortingStream = false this.resetAssistantMessagePersistence() // Reset tool failure flag for each new assistant turn - this ensures that tool failures // only prevent attempt_completion within the same assistant message, not across turns diff --git a/src/core/task/__tests__/Task.spec.ts b/src/core/task/__tests__/Task.spec.ts index 72d0641efd..eae85cf9f2 100644 --- a/src/core/task/__tests__/Task.spec.ts +++ b/src/core/task/__tests__/Task.spec.ts @@ -86,6 +86,7 @@ vi.mock("delay", () => ({ })) import delay from "delay" +import pWaitFor from "p-wait-for" vi.mock("uuid", async (importOriginal) => { const actual = await importOriginal() @@ -713,6 +714,98 @@ describe("Cline", () => { ]) }) + it("clears didFinishAbortingStream on retry so cancelTask waits for the new stream (#1801)", async () => { + const task = new Task({ + provider: mockProvider, + apiConfiguration: mockApiConfig, + task: "abort flag reset test", + startTask: false, + }) + + vi.spyOn(task.diffViewProvider, "reset").mockResolvedValue(undefined) + vi.spyOn(getTaskTestAccess(task), "safeEnsureModelFetched").mockResolvedValue(stubModelInfo) + vi.spyOn(getTaskTestAccess(task), "presentAssistantMessageSafe").mockImplementation(() => {}) + + const failingStream = async function* (): AsyncGenerator { + yield { type: "text", text: "partial" } + throw new Error("simulated mid-stream failure") + } + + let releaseRetryStream: (() => void) | undefined + let markRetryStreamPaused: (() => void) | undefined + const retryStreamRelease = new Promise((resolve) => { + releaseRetryStream = resolve + }) + const retryStreamPaused = new Promise((resolve) => { + markRetryStreamPaused = resolve + }) + const pausedRetryStream = async function* (): AsyncGenerator { + yield { type: "text", text: "retry" } + markRetryStreamPaused?.() + await retryStreamRelease + } + + const attemptApiRequestSpy = vi + .spyOn(task, "attemptApiRequest") + .mockImplementationOnce(() => failingStream()) + .mockImplementationOnce(() => pausedRetryStream()) + .mockImplementation(() => { + throw new Error("stop after retry response") + }) + + const { default: realPWaitFor } = await vi.importActual("p-wait-for") + vi.mocked(pWaitFor).mockImplementation(realPWaitFor) + let request: Promise | undefined + let cancel: Promise | undefined + try { + await mockProvider.addClineToStack(task) + request = task.recursivelyMakeClineRequests([{ type: "text", text: "abort flag reset test" }]) + await Promise.race([ + retryStreamPaused, + request.then(() => { + throw new Error("request settled before the retry stream paused") + }), + ]) + + const firstApiReq = task.clineMessages.find((message) => message.say === "api_req_started") + expect(JSON.parse(firstApiReq?.text ?? "{}")).toMatchObject({ cancelReason: "streaming_failed" }) + expect(attemptApiRequestSpy).toHaveBeenCalledTimes(2) + + let markCancelWaiting: (() => void) | undefined + const cancelWaiting = new Promise((resolve) => { + markCancelWaiting = resolve + }) + vi.mocked(pWaitFor).mockImplementation((condition, options) => + realPWaitFor(async () => { + const done = await condition() + if (!done) markCancelWaiting?.() + return done + }, options), + ) + + vi.spyOn(mockProvider, "createTaskWithHistoryItem").mockImplementation((historyItem) => + ClineProvider.prototype.createTaskWithHistoryItem.call(mockProvider, historyItem, { + startTask: false, + }), + ) + + let cancelSettled = false + cancel = mockProvider.cancelTask().then(() => { + cancelSettled = true + }) + await Promise.race([cancelWaiting, cancel]) + expect(cancelSettled, "cancelTask must wait while the retry stream is live").toBe(false) + + releaseRetryStream?.() + await cancel + await expect(request).resolves.toBe(true) + } finally { + releaseRetryStream?.() + await Promise.allSettled([request, cancel]) + vi.mocked(pWaitFor).mockImplementation(async () => {}) + } + }) + it("finalizes MCP tool call using the request-scoped parser state", async () => { const task = new Task({ provider: mockProvider,