Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 3 additions & 4 deletions src/core/task/Task.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3715,8 +3715,8 @@
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"

Check warning on line 3719 in src/core/task/Task.ts

View workflow job for this annotation

GitHub Actions / mutation-diff

Mutation test advisory

src/core/task/Task.ts:3719: 2 mutation test gaps; example: NoCoverage StringLiteral mutant (replacement: ""). See the job summary for the complete list and resolution guidance.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we apply the same ??= in the sibling abort path below (the mid-stream retry backoff branch, this.abortReason = "user_cancelled" at line 3755)? It's a harmless no-op overwrite today, but using ??= in both places keeps the "first abort reason wins" invariant uniform across this whole catch block and avoids future readers wondering whether the two paths intentionally differ.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done. Also removed the now-redundant comment on that line since the code is self-explanatory.

await this.abortTask()
} else if (error instanceof OutputTokenLimitError) {
// Truncation repeats on an identical request, so never auto-retry it
Expand Down Expand Up @@ -3751,8 +3751,7 @@
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"

Check warning on line 3754 in src/core/task/Task.ts

View workflow job for this annotation

GitHub Actions / mutation-diff

Mutation test advisory

src/core/task/Task.ts:3754: NoCoverage StringLiteral mutant (replacement: ""). See the job summary for the complete list and resolution guidance.
await this.abortTask()
break
}
Expand Down
216 changes: 216 additions & 0 deletions src/core/task/__tests__/Task.abort-reason-race.spec.ts
Original file line number Diff line number Diff line change
@@ -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<typeof import("../../task-persistence")>()
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)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
expect(abortTaskSpy).toHaveBeenCalledOnce()
expect(task.abortReason).toBe("user_cancelled")
Comment thread
coderabbitai[bot] marked this conversation as resolved.
})

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
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

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")
})
})
30 changes: 1 addition & 29 deletions src/core/webview/ClineProvider.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
19 changes: 19 additions & 0 deletions src/core/webview/__tests__/ClineProvider.taskHistory.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Comment thread
coderabbitai[bot] marked this conversation as resolved.
})

it("emits delegated completion through the provider after the child is disposed", () => {
const listener = vi.fn()
provider.on(RooCodeEventName.TaskCompleted, listener)
Expand Down
Loading