Skip to content
Open
Show file tree
Hide file tree
Changes from 3 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
84 changes: 19 additions & 65 deletions src/core/task/Task.ts
Original file line number Diff line number Diff line change
Expand Up @@ -410,7 +410,6 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
private readonly _isHistoryTask: boolean
// No streaming parser is required.
assistantMessageParser?: undefined
private providerProfileChangeListener?: (config: { name: string; provider?: string }) => void

// Native tool call streaming state (track which index each tool is at)
private streamingToolCallIndices: Map<string, number> = new Map()
Expand Down Expand Up @@ -562,9 +561,6 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {

this.messageQueueService.on("stateChanged", this.messageQueueStateChangedHandler)

// Listen for provider profile changes to update parser state
this.setupProviderProfileChangeListener(provider)

// Set up diff strategy
this.diffStrategy = new MultiSearchReplaceDiffStrategy(diffFuzzyThreshold)

Expand Down Expand Up @@ -690,35 +686,6 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
}
}

/**
* Sets up a listener for provider profile changes.
*
* @private
* @param provider - The ClineProvider instance to listen to
*/
private setupProviderProfileChangeListener(provider: ClineProvider): void {
// Only set up listener if provider has the on method (may not exist in test mocks)
if (typeof provider.on !== "function") {
return
}

this.providerProfileChangeListener = async () => {
try {
const newState = await provider.getState()
if (newState?.apiConfiguration) {
this.updateApiConfiguration(newState.apiConfiguration)
}
} catch (error) {
console.error(
`[Task#${this.taskId}.${this.instanceId}] Failed to update API configuration on profile change:`,
error,
)
}
}

provider.on(RooCodeEventName.ProviderProfileChanged, this.providerProfileChangeListener)
}

/**
* Wait for the task mode to be initialized before proceeding.
* This method ensures that any operations depending on the task mode
Expand Down Expand Up @@ -1537,6 +1504,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
if (provider) {
if (mode) {
await provider.setMode(mode)
this._taskMode = mode
}

if (providerProfile) {
Expand Down Expand Up @@ -1591,7 +1559,9 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
// Get condensing configuration
const state = await this.providerRef.deref()?.getState()
const customCondensingPrompt = state?.customSupportPrompts?.CONDENSE
const { mode, apiConfiguration } = state ?? {}
// Use task-local values, not provider state, to prevent cross-task configuration leaks.
const mode = await this.getTaskMode()
const apiConfiguration = this.apiConfiguration

const { contextTokens: prevContextTokens } = this.getTokenUsage()

Expand Down Expand Up @@ -2277,19 +2247,6 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
console.error("Error cancelling current request:", error)
}

// Remove provider profile change listener
try {
if (this.providerProfileChangeListener) {
const provider = this.providerRef.deref()
if (provider) {
provider.off(RooCodeEventName.ProviderProfileChanged, this.providerProfileChangeListener)
}
this.providerProfileChangeListener = undefined
}
} catch (error) {
console.error("Error removing provider profile change listener:", error)
}

// Dispose message queue and remove event listeners.
try {
if (this.messageQueueStateChangedHandler) {
Expand Down Expand Up @@ -2581,7 +2538,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
const showRooIgnoredFiles = state?.showRooIgnoredFiles ?? false
const includeDiagnosticMessages = state?.includeDiagnosticMessages ?? true
const maxDiagnosticMessages = state?.maxDiagnosticMessages ?? 50
const currentMode = state?.mode ?? defaultModeSlug
const currentMode = await this.getTaskMode()

const { content: parsedUserContent, mode: slashCommandMode } = await processUserContentMentions({
userContent: currentUserContent,
Expand Down Expand Up @@ -3782,16 +3739,11 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {

const state = await this.providerRef.deref()?.getState()

const {
mode,
customModes,
customModePrompts,
customInstructions,
experiments,
language,
apiConfiguration,
enableSubfolderRules,
} = state ?? {}
const { customModes, customModePrompts, customInstructions, experiments, language, enableSubfolderRules } =
state ?? {}
// Use task-local values, not provider state, to prevent cross-task configuration leaks.
const mode = await this.getTaskMode()
const apiConfiguration = this.apiConfiguration

return await (async () => {
const provider = this.providerRef.deref()
Expand Down Expand Up @@ -3857,7 +3809,10 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {

private async handleContextWindowExceededError(): Promise<void> {
const state = await this.providerRef.deref()?.getState()
const { profileThresholds = {}, mode, apiConfiguration } = state ?? {}
const { profileThresholds = {} } = state ?? {}
// Use task-local values, not provider state, to prevent cross-task configuration leaks.
const mode = await this.getTaskMode()
const apiConfiguration = this.apiConfiguration

const { contextTokens } = this.getTokenUsage()
await this.safeEnsureModelFetched()
Expand Down Expand Up @@ -3997,9 +3952,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
* the `api_req_rate_limit_wait` say type (not an error).
*/
private async maybeWaitForProviderRateLimit(retryAttempt: number): Promise<void> {
const state = await this.providerRef.deref()?.getState()
const rateLimitSeconds =
state?.apiConfiguration?.rateLimitSeconds ?? this.apiConfiguration?.rateLimitSeconds ?? 0
const rateLimitSeconds = this.apiConfiguration?.rateLimitSeconds ?? 0
Comment thread
coderabbitai[bot] marked this conversation as resolved.

const lastRequestTime = this.rateLimitClock.getLastRequestTime()
if (rateLimitSeconds <= 0 || !lastRequestTime) {
Expand Down Expand Up @@ -4032,14 +3985,15 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
const state = await this.providerRef.deref()?.getState()

const {
apiConfiguration,
autoApprovalEnabled,
requestDelaySeconds,
mode,
autoCondenseContext = true,
autoCondenseContextPercent = 100,
profileThresholds = {},
} = state ?? {}
// Use task-local values, not provider state, to prevent cross-task configuration leaks.
const mode = await this.getTaskMode()
const apiConfiguration = this.apiConfiguration
Comment thread
coderabbitai[bot] marked this conversation as resolved.

// Get condensing configuration for automatic triggers.
const customCondensingPrompt = state?.customSupportPrompts?.CONDENSE
Expand Down Expand Up @@ -4452,7 +4406,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {

// Respect provider rate limit window
let rateLimitDelay = 0
const rateLimit = (state?.apiConfiguration ?? this.apiConfiguration)?.rateLimitSeconds || 0
const rateLimit = this.apiConfiguration?.rateLimitSeconds ?? 0
const lastRequestTime = this.rateLimitClock.getLastRequestTime()
if (lastRequestTime && rateLimit > 0) {
const elapsed = performance.now() - lastRequestTime
Expand Down
143 changes: 140 additions & 3 deletions src/core/task/__tests__/Task.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import {
import { TelemetryService } from "@roo-code/telemetry"

import { Task } from "../Task"
import { SYSTEM_PROMPT } from "../../prompts/system"
import { createRateLimitClock } from "../RateLimitClock"
import { summarizeConversation } from "../../condense"
import { ClineProvider } from "../../webview/ClineProvider"
Expand Down Expand Up @@ -223,6 +224,15 @@ vi.mock("../../condense", async (importOriginal) => {
}),
}
})

vi.mock("../../prompts/system", async (importOriginal) => {
const actual = await importOriginal<typeof import("../../prompts/system")>()
return {
...actual,
SYSTEM_PROMPT: vi.fn(actual.SYSTEM_PROMPT),
}
})

// Mock storagePathManager to prevent dynamic import issues.
vi.mock("../../../utils/storage", () => ({
getTaskDirectoryPath: vi
Expand Down Expand Up @@ -451,6 +461,92 @@ describe("Cline", () => {
})
})

describe("task-local configuration isolation", () => {
it("uses the task mode and API configuration when focused provider state differs", async () => {
const taskApiConfiguration: ProviderSettings = {
...mockApiConfig,
todoListEnabled: true,
}
vi.spyOn(mockProvider, "getState").mockResolvedValue({ mode: "architect", mcpEnabled: false })

const task = new Task({
provider: mockProvider,
apiConfiguration: taskApiConfiguration,
task: "test task",
startTask: false,
})
await task.getTaskMode()

vi.spyOn(mockProvider, "getState").mockResolvedValue({
mode: "code",
mcpEnabled: false,
apiConfiguration: { ...mockApiConfig, todoListEnabled: false },
})
vi.mocked(SYSTEM_PROMPT).mockResolvedValueOnce("mock system prompt")

await getTaskTestAccess(task).getSystemPrompt()

const systemPromptCall = requireDefined(vi.mocked(SYSTEM_PROMPT).mock.calls.at(-1))
const [, , , , , mode, , , , , , , settings] = systemPromptCall
expect(mode).toBe("architect")
expect(settings).toMatchObject({ todoListEnabled: true })
})

it("uses the task mode when manually condensing after focused state changes", async () => {
vi.spyOn(mockProvider, "getState").mockResolvedValue({ mode: "architect", mcpEnabled: false })
const task = new Task({
provider: mockProvider,
apiConfiguration: mockApiConfig,
task: "test task",
startTask: false,
})
await task.getTaskMode()
vi.spyOn(mockProvider, "getState").mockResolvedValue({ mode: "code", mcpEnabled: false })
vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt")

await task.condenseContext()

const [options] = requireDefined(vi.mocked(summarizeConversation).mock.calls.at(-1))
expect(options.metadata?.mode).toBe("architect")
})

it("uses the task mode in request metadata when focused provider state differs", async () => {
vi.spyOn(mockProvider, "getState").mockResolvedValue({
mode: "ask",
mcpEnabled: false,
autoApprovalEnabled: true,
requestDelaySeconds: 0,
})
const task = new Task({
provider: mockProvider,
apiConfiguration: mockApiConfig,
task: "test task",
startTask: false,
})
await task.getTaskMode()
vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt")

vi.spyOn(mockProvider, "getState").mockResolvedValue({
mode: "code",
mcpEnabled: false,
autoApprovalEnabled: true,
requestDelaySeconds: 0,
})
const stream = (async function* () {
yield { type: "text", text: "response" } as ApiStreamChunk
})()
const createMessage = vi.spyOn(task.api, "createMessage").mockReturnValue(stream)
task.apiConversationHistory = [
{ role: "user", content: [{ type: "text", text: "test message" }], ts: Date.now() },
]

await task.attemptApiRequest().next()

const metadata = requireDefined(createMessage.mock.calls[0])[2]
expect(metadata?.mode).toBe("ask")
})
})

describe("sayAndCreateMissingParamError", () => {
it("surfaces a localized error notice and returns the missing-parameter tool error for both relPath branches", async () => {
const cline = new Task({
Expand Down Expand Up @@ -750,7 +846,7 @@ describe("Cline", () => {
expect(mockDelay).toHaveBeenCalledWith(1000)
})

it("should respect rate limit window in retry backoff", async () => {
it("uses the task rate limit in retry backoff when focused provider state differs", async () => {
const clock = createRateLimitClock()
const rateLimitConfig = {
...mockApiConfig,
Expand Down Expand Up @@ -815,18 +911,29 @@ describe("Cline", () => {
const providerState = await mockProvider.getState()
vi.spyOn(mockProvider, "getState").mockResolvedValue({
...providerState,
apiConfiguration: rateLimitConfig,
apiConfiguration: {
...mockApiConfig,
rateLimitSeconds: 1,
},
autoApprovalEnabled: true,
requestDelaySeconds: 3,
})

const iterator = cline.attemptApiRequest(0)
await iterator.next()

// rateLimitSeconds=10 > exponentialDelay=ceil(3*2^0)=3, so
// The task rateLimitSeconds=10 (rather than the focused provider's 1)
// exceeds exponentialDelay=ceil(3*2^0)=3, so
// finalDelay=10 and the countdown loop fires delay(1000) ten times.
expect(mockDelay).toHaveBeenCalledWith(1000)
expect(mockDelay).toHaveBeenCalledTimes(10)
const countdownMessages = saySpy.mock.calls.filter(
([type, text, , partial]) =>
type === "api_req_retry_delayed" && partial && typeof text === "string",
)
expect(countdownMessages.map(([, text]) => text)).toEqual(
Array.from({ length: 10 }, (_, index) => `API Error\n<retry_timer>${10 - index}</retry_timer>`),
)
expect(clock.getLastRequestTime()).toBeDefined()
})

Expand Down Expand Up @@ -1574,6 +1681,33 @@ describe("Cline", () => {
expect(mockProvider.postMessageToWebview).not.toHaveBeenCalled()
})

it("uses a mode selected through submitUserMessage in the next API request", async () => {
vi.spyOn(mockProvider, "getState").mockResolvedValue({ mode: "ask", mcpEnabled: false })
vi.spyOn(mockProvider, "setMode").mockResolvedValue(undefined)
const task = new Task({
provider: mockProvider,
apiConfiguration: mockApiConfig,
task: "initial task",
startTask: false,
})
vi.spyOn(task, "handleWebviewAskResponse").mockImplementation(() => {})

await task.submitUserMessage("switch modes", undefined, "code")
vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt")
const stream = (async function* () {
yield { type: "text", text: "response" } as ApiStreamChunk
})()
const createMessage = vi.spyOn(task.api, "createMessage").mockReturnValue(stream)
task.apiConversationHistory = [
{ role: "user", content: [{ type: "text", text: "test message" }], ts: Date.now() },
]

await task.attemptApiRequest().next()

expect(mockProvider.setMode).toHaveBeenCalledWith("code")
expect(requireDefined(createMessage.mock.calls[0])[2]?.mode).toBe("code")
})

it("should handle empty messages gracefully", async () => {
const task = new Task({
provider: mockProvider,
Expand Down Expand Up @@ -2411,12 +2545,14 @@ describe("Cline", () => {
})

it("should propagate AbortController signal through attemptApiRequest context-window retry path", async () => {
vi.spyOn(mockProvider, "getState").mockResolvedValue({ mode: "architect", mcpEnabled: false })
const task = new Task({
provider: mockProvider,
apiConfiguration: mockApiConfig,
task: "test task",
startTask: false,
})
await task.getTaskMode()

vi.spyOn(getTaskTestAccess(task), "getSystemPrompt").mockResolvedValue("mock system prompt")
vi.spyOn(task, "getTokenUsage").mockReturnValue({
Expand Down Expand Up @@ -2511,6 +2647,7 @@ describe("Cline", () => {
expect(summarizeConversation).toHaveBeenCalled()
const [options] = vi.mocked(summarizeConversation).mock.calls.at(-1)!
expect(options.metadata?.taskId).toBe(task.taskId)
expect(options.metadata?.mode).toBe("architect")
expect(options.metadata?.abortSignal).toBeInstanceOf(AbortSignal)
expect(options.metadata?.abortSignal?.aborted).toBe(false)
})
Expand Down
Loading