diff --git a/src/core/assistant-message/__tests__/presentAssistantMessage-custom-tool.spec.ts b/src/core/assistant-message/__tests__/presentAssistantMessage-custom-tool.spec.ts index 1ef25e852b..3c253459fd 100644 --- a/src/core/assistant-message/__tests__/presentAssistantMessage-custom-tool.spec.ts +++ b/src/core/assistant-message/__tests__/presentAssistantMessage-custom-tool.spec.ts @@ -58,6 +58,7 @@ describe("presentAssistantMessage - Custom Tool Recording", () => { didAlreadyUseTool: false, consecutiveMistakeCount: 0, clineMessages: [], + getTaskMode: vi.fn().mockResolvedValue("code"), api: { getModel: () => ({ id: "test-model", info: {} }), }, @@ -120,6 +121,37 @@ describe("presentAssistantMessage - Custom Tool Recording", () => { // Should record as "custom_tool", not "my_custom_tool" expect(mockTask.recordToolUsage).toHaveBeenCalledWith("custom_tool") }) + + it("passes the task-local mode to custom tool execution", async () => { + mockTask.getTaskMode.mockResolvedValue("code") + mockTask.providerRef.deref = () => ({ + getState: vi.fn().mockResolvedValue({ + mode: "orchestrator", + customModes: [], + experiments: { customTools: true }, + }), + }) + mockTask.assistantMessageContent = [ + { + type: "tool_use", + id: "tool_call_task_mode", + name: "my_custom_tool", + params: {}, + partial: false, + }, + ] + const execute = vi.fn().mockResolvedValue("Custom tool result") + vi.mocked(customToolRegistry.has).mockReturnValue(true) + vi.mocked(customToolRegistry.get).mockReturnValue({ + name: "my_custom_tool", + description: "A custom tool", + execute, + }) + + await presentAssistantMessage(mockTask) + + expect(execute).toHaveBeenCalledWith(undefined, { mode: "code", task: mockTask }) + }) }) describe("Custom tool error recording", () => { diff --git a/src/core/assistant-message/__tests__/presentAssistantMessage-images.spec.ts b/src/core/assistant-message/__tests__/presentAssistantMessage-images.spec.ts index fcf778b8f8..ed90759cdc 100644 --- a/src/core/assistant-message/__tests__/presentAssistantMessage-images.spec.ts +++ b/src/core/assistant-message/__tests__/presentAssistantMessage-images.spec.ts @@ -42,6 +42,7 @@ describe("presentAssistantMessage - Image Handling in Native Tool Calling", () = didRejectTool: false, didAlreadyUseTool: false, consecutiveMistakeCount: 0, + getTaskMode: vi.fn().mockResolvedValue("code"), api: { getModel: () => ({ id: "test-model", info: {} }), }, diff --git a/src/core/assistant-message/__tests__/presentAssistantMessage-tool-usage-attribution.spec.ts b/src/core/assistant-message/__tests__/presentAssistantMessage-tool-usage-attribution.spec.ts index c75eb6ee18..9cb802a621 100644 --- a/src/core/assistant-message/__tests__/presentAssistantMessage-tool-usage-attribution.spec.ts +++ b/src/core/assistant-message/__tests__/presentAssistantMessage-tool-usage-attribution.spec.ts @@ -1,7 +1,7 @@ // npx vitest src/core/assistant-message/__tests__/presentAssistantMessage-tool-usage-attribution.spec.ts import type { Anthropic } from "@anthropic-ai/sdk" -import { describe, it, expect, beforeEach, vi } from "vitest" +import { describe, it, expect, beforeEach, vi, type Mock } from "vitest" import { presentAssistantMessage } from "../presentAssistantMessage" import { validateToolUse } from "../../tools/validateToolUse" import { getModeBySlug } from "../../../shared/modes" @@ -60,6 +60,7 @@ interface MockTask { didAlreadyUseTool: boolean consecutiveMistakeCount: number clineMessages: unknown[] + getTaskMode: Mock<() => Promise> api: { getModel: () => { id: string; info: Record } } recordToolUsage: ReturnType recordToolError: ReturnType @@ -96,6 +97,7 @@ describe("presentAssistantMessage - tool usage attribution", () => { didAlreadyUseTool: false, consecutiveMistakeCount: 0, clineMessages: [], + getTaskMode: vi.fn().mockResolvedValue("code"), api: { getModel: () => ({ id: "test-model", info: {} }), }, @@ -187,6 +189,35 @@ describe("presentAssistantMessage - tool usage attribution", () => { expect(mockTask.recordToolUsage).not.toHaveBeenCalledWith("mcp_") }) + it("validates tools against the task-local mode when provider state differs", async () => { + mockTask.getTaskMode.mockResolvedValue("code") + mockTask.providerRef.deref = () => ({ + getState: vi.fn().mockResolvedValue({ mode: "orchestrator", customModes: [] }), + }) + mockTask.assistantMessageContent = [ + { + type: "tool_use", + id: "call_task_mode", + name: "read_file", + params: { path: "test.txt" }, + nativeArgs: { path: "test.txt" }, + partial: false, + }, + ] + + await presentAssistantMessage(mockTask as unknown as Task) + + expect(validateToolUse).toHaveBeenCalledWith( + "read_file", + "code", + [], + {}, + { path: "test.txt" }, + undefined, + undefined, + ) + }) + it("records a safe failure key without leaking the raw tool name when validation fails", async () => { vi.mocked(validateToolUse).mockImplementation(() => { throw new Error('Tool "read_file" is not allowed in this mode.') diff --git a/src/core/assistant-message/__tests__/presentAssistantMessage-unknown-tool.spec.ts b/src/core/assistant-message/__tests__/presentAssistantMessage-unknown-tool.spec.ts index 78a4a19e91..8cfaa10972 100644 --- a/src/core/assistant-message/__tests__/presentAssistantMessage-unknown-tool.spec.ts +++ b/src/core/assistant-message/__tests__/presentAssistantMessage-unknown-tool.spec.ts @@ -44,6 +44,7 @@ describe("presentAssistantMessage - Unknown Tool Handling", () => { didAlreadyUseTool: false, consecutiveMistakeCount: 0, clineMessages: [], + getTaskMode: vi.fn().mockResolvedValue("code"), api: { getModel: () => ({ id: "test-model", info: {} }), }, diff --git a/src/core/assistant-message/presentAssistantMessage.ts b/src/core/assistant-message/presentAssistantMessage.ts index 7b25db4e66..b7aeec9657 100644 --- a/src/core/assistant-message/presentAssistantMessage.ts +++ b/src/core/assistant-message/presentAssistantMessage.ts @@ -342,9 +342,10 @@ export async function presentAssistantMessage(cline: Task) { break } - // Fetch state early so it's available for toolDescription and validation + // Shared provider state supplies global settings; mode is owned by the task. const state = await cline.providerRef.deref()?.getState() - const { mode, customModes, experiments: stateExperiments, disabledTools } = state ?? {} + const { customModes, experiments: stateExperiments, disabledTools } = state ?? {} + const mode = await cline.getTaskMode() const toolDescription = (): string => { switch (block.name) { @@ -617,7 +618,7 @@ export async function presentAssistantMessage(cline: Task) { validateToolUse( block.name as ToolName, - mode ?? defaultModeSlug, + mode, customModes ?? [], toolRequirements, block.params, @@ -924,7 +925,7 @@ export async function presentAssistantMessage(cline: Task) { } const result = await customTool.execute(customToolArgs, { - mode: mode ?? defaultModeSlug, + mode, task: cline, }) diff --git a/src/core/environment/__tests__/getEnvironmentDetails.spec.ts b/src/core/environment/__tests__/getEnvironmentDetails.spec.ts index df47e83c21..f894f2a7ee 100644 --- a/src/core/environment/__tests__/getEnvironmentDetails.spec.ts +++ b/src/core/environment/__tests__/getEnvironmentDetails.spec.ts @@ -75,7 +75,7 @@ describe("getEnvironmentDetails", () => { terminalOutputLineLimit: 100, maxWorkspaceFiles: 50, maxOpenTabsContext: 10, - mode: "code", + mode: "orchestrator", customModes: [], experiments: {}, customInstructions: "test instructions", @@ -91,6 +91,7 @@ describe("getEnvironmentDetails", () => { cwd: mockCwd, taskId: mockTaskId, didEditFile: false, + getTaskMode: vi.fn().mockResolvedValue("code"), fileContextTracker: { getAndClearRecentlyModifiedFiles: vi.fn().mockReturnValue([]), } as unknown as FileContextTracker, @@ -156,6 +157,8 @@ describe("getEnvironmentDetails", () => { expect(mockProvider.getState).toHaveBeenCalled() + expect(mockCline.getTaskMode).toHaveBeenCalled() + expect(result).toContain("code") expect(getFullModeDetails).toHaveBeenCalledWith("code", [], undefined, { cwd: mockCwd, globalCustomInstructions: "test instructions", diff --git a/src/core/environment/getEnvironmentDetails.ts b/src/core/environment/getEnvironmentDetails.ts index 0e7d18a57a..c2f1ae7f49 100644 --- a/src/core/environment/getEnvironmentDetails.ts +++ b/src/core/environment/getEnvironmentDetails.ts @@ -8,7 +8,7 @@ import delay from "delay" import type { ExperimentId } from "@roo-code/types" import { formatLanguage } from "../../shared/language" -import { defaultModeSlug, getFullModeDetails } from "../../shared/modes" +import { getFullModeDetails } from "../../shared/modes" import { getApiMetrics } from "../../shared/getApiMetrics" import { listFiles } from "../../services/glob/list-files" import { TerminalRegistry } from "../../integrations/terminal/TerminalRegistry" @@ -205,7 +205,6 @@ export async function getEnvironmentDetails(cline: Task, includeFileDetails: boo // Add current mode and any mode-specific warnings. const { - mode, customModes, customModePrompts, experiments = {} as Record, @@ -213,7 +212,7 @@ export async function getEnvironmentDetails(cline: Task, includeFileDetails: boo language, } = state ?? {} - const currentMode = mode ?? defaultModeSlug + const currentMode = await cline.getTaskMode() const modeDetails = await getFullModeDetails(currentMode, customModes, customModePrompts, { cwd: cline.cwd, diff --git a/src/core/tools/RunSlashCommandTool.ts b/src/core/tools/RunSlashCommandTool.ts index c6fd48665d..44663f0c75 100644 --- a/src/core/tools/RunSlashCommandTool.ts +++ b/src/core/tools/RunSlashCommandTool.ts @@ -55,7 +55,7 @@ export class RunSlashCommandTool extends BaseTool<"run_slash_command"> { const command = await getCommand(task.cwd, commandName) if (!command) { - const currentMode = state?.mode ?? "code" + const currentMode = await task.getTaskMode() const skillsManager = provider?.getSkillsManager() const skillContent = await resolveSkillContentForMode(skillsManager, commandName, currentMode) diff --git a/src/core/tools/SkillTool.ts b/src/core/tools/SkillTool.ts index fa696e5b73..5a4dfdfd8e 100644 --- a/src/core/tools/SkillTool.ts +++ b/src/core/tools/SkillTool.ts @@ -43,9 +43,8 @@ export class SkillTool extends BaseTool<"skill"> { return } - // Get current mode for skill resolution - const state = await provider?.getState() - const currentMode = state?.mode ?? "code" + // Resolve skills against the task's mode, not shared provider state. + const currentMode = await task.getTaskMode() // Fetch skill content const skillContent = await resolveSkillContentForMode(skillsManager, skillName, currentMode) diff --git a/src/core/tools/SwitchModeTool.ts b/src/core/tools/SwitchModeTool.ts index a60ce63bde..e1216fcaad 100644 --- a/src/core/tools/SwitchModeTool.ts +++ b/src/core/tools/SwitchModeTool.ts @@ -2,7 +2,7 @@ import delay from "delay" import { Task } from "../task/Task" import { formatResponse } from "../prompts/responses" -import { defaultModeSlug, getModeBySlug } from "../../shared/modes" +import { getModeBySlug } from "../../shared/modes" import { BaseTool, ToolCallbacks } from "./BaseTool" import type { ToolUse } from "../../shared/tools" @@ -38,8 +38,8 @@ export class SwitchModeTool extends BaseTool<"switch_mode"> { return } - // Check if already in requested mode - const currentMode = (await task.providerRef.deref()?.getState())?.mode ?? defaultModeSlug + // Mode belongs to the task; provider state may still reflect its parent. + const currentMode = await task.getTaskMode() if (currentMode === mode_slug) { task.recordToolError("switch_mode") diff --git a/src/core/tools/__tests__/mcpServerRestriction.spec.ts b/src/core/tools/__tests__/mcpServerRestriction.spec.ts index 6455e1e4b3..669f167846 100644 --- a/src/core/tools/__tests__/mcpServerRestriction.spec.ts +++ b/src/core/tools/__tests__/mcpServerRestriction.spec.ts @@ -7,7 +7,6 @@ vi.mock("../../../shared/modes", async (importOriginal) => { const actual = await importOriginal() return { ...actual, - defaultModeSlug: "code", getModeBySlug: vi.fn(), } }) @@ -16,8 +15,14 @@ import { getModeBySlug } from "../../../shared/modes" const toolError = (error: string) => `ERR:${error}` -function makeTask(state: any): Task { +type ProviderModeState = { + mode?: string + customModes?: [] +} + +function makeTask(state: ProviderModeState, taskMode = "code"): Task { return { + getTaskMode: vi.fn().mockResolvedValue(taskMode), providerRef: { deref: () => ({ getState: vi.fn().mockResolvedValue(state), @@ -60,8 +65,9 @@ describe("getAllowedMcpServersForTask", () => { groups: ["mcp"], allowedMcpServers: ["srv-a"], } as any) - const task = makeTask({ mode: "code", customModes: [] }) + const task = makeTask({ mode: "orchestrator", customModes: [] }, "code") await expect(getAllowedMcpServersForTask(task)).resolves.toEqual(["srv-a"]) + expect(getModeBySlug).toHaveBeenCalledWith("code", []) }) it("returns undefined when the mode does not restrict servers", async () => { @@ -77,7 +83,7 @@ describe("getAllowedMcpServersForTask", () => { it("returns undefined when the mode cannot be resolved", async () => { vi.mocked(getModeBySlug).mockReturnValue(undefined as any) - const task = makeTask({ mode: "missing", customModes: [] }) + const task = makeTask({ mode: "code", customModes: [] }, "missing") await expect(getAllowedMcpServersForTask(task)).resolves.toBeUndefined() }) }) diff --git a/src/core/tools/__tests__/runSlashCommandTool.spec.ts b/src/core/tools/__tests__/runSlashCommandTool.spec.ts index e3d135b45f..9a140931eb 100644 --- a/src/core/tools/__tests__/runSlashCommandTool.spec.ts +++ b/src/core/tools/__tests__/runSlashCommandTool.spec.ts @@ -19,6 +19,7 @@ describe("runSlashCommandTool", () => { vi.clearAllMocks() mockTask = { + getTaskMode: vi.fn().mockResolvedValue("code"), consecutiveMistakeCount: 0, recordToolError: vi.fn(), sayAndCreateMissingParamError: vi.fn().mockResolvedValue("Missing parameter error"), @@ -96,6 +97,7 @@ describe("runSlashCommandTool", () => { }, } + mockTask.getTaskMode.mockResolvedValue("code") const getSkillContent = vi.fn().mockResolvedValue({ name: "skill-only", description: "Skill-generated command", @@ -109,7 +111,7 @@ describe("runSlashCommandTool", () => { experiments: { runSlashCommand: true, }, - mode: "code", + mode: "orchestrator", }), getSkillsManager: vi.fn().mockReturnValue({ getSkillContent, diff --git a/src/core/tools/__tests__/skillTool.spec.ts b/src/core/tools/__tests__/skillTool.spec.ts index 037507c6a5..395dd6f157 100644 --- a/src/core/tools/__tests__/skillTool.spec.ts +++ b/src/core/tools/__tests__/skillTool.spec.ts @@ -18,6 +18,7 @@ describe("skillTool", () => { } mockTask = { + getTaskMode: vi.fn().mockResolvedValue("code"), consecutiveMistakeCount: 0, recordToolError: vi.fn(), didToolFailInCurrentTurn: false, @@ -67,12 +68,20 @@ describe("skillTool", () => { skill: "non-existent", }, } + mockTask.getTaskMode.mockResolvedValue("code") + mockTask.providerRef.deref = vi.fn().mockReturnValue({ + getState: vi.fn().mockResolvedValue({ mode: "orchestrator" }), + getSkillsManager: vi.fn().mockReturnValue(mockSkillsManager), + }) mockSkillsManager.getSkillContent.mockResolvedValue(null) mockSkillsManager.getSkillsForMode.mockReturnValue([{ name: "create-mcp-server" }]) await skillTool.handle(mockTask as Task, block, mockCallbacks) + expect(mockSkillsManager.getSkillContent).toHaveBeenCalledWith("non-existent", "code") + expect(mockSkillsManager.getSkillsForMode).toHaveBeenCalledWith("code") + expect(mockCallbacks.pushToolResult).toHaveBeenCalledWith( formatResponse.toolError("Skill 'non-existent' not found. Available skills: create-mcp-server"), ) @@ -109,6 +118,11 @@ describe("skillTool", () => { skill: "create-mcp-server", }, } + mockTask.getTaskMode.mockResolvedValue("code") + mockTask.providerRef.deref = vi.fn().mockReturnValue({ + getState: vi.fn().mockResolvedValue({ mode: "orchestrator" }), + getSkillsManager: vi.fn().mockReturnValue(mockSkillsManager), + }) const mockSkillContent = { name: "create-mcp-server", @@ -121,6 +135,8 @@ describe("skillTool", () => { await skillTool.handle(mockTask as Task, block, mockCallbacks) + expect(mockSkillsManager.getSkillContent).toHaveBeenCalledWith("create-mcp-server", "code") + expect(mockCallbacks.askApproval).toHaveBeenCalledWith( "tool", JSON.stringify({ diff --git a/src/core/tools/__tests__/switchModeTool.spec.ts b/src/core/tools/__tests__/switchModeTool.spec.ts index a82429ac7c..182fcfa2d1 100644 --- a/src/core/tools/__tests__/switchModeTool.spec.ts +++ b/src/core/tools/__tests__/switchModeTool.spec.ts @@ -41,6 +41,7 @@ describe("SwitchModeTool", () => { mockGetState = vi.fn().mockResolvedValue({ mode: "code", customModes: [] }) mockTask = { + getTaskMode: vi.fn().mockResolvedValue("code"), consecutiveMistakeCount: 0, recordToolError: vi.fn(), didToolFailInCurrentTurn: false, @@ -325,11 +326,9 @@ describe("SwitchModeTool", () => { expect(mockCallbacks.askApproval).toHaveBeenCalledWith("tool", expectedMessage) }) - // ===== getState with custom modes ===== - - it("should read current mode from providerRef state", async () => { - // Set current mode to "architect" - mockGetState.mockResolvedValue({ mode: "architect", customModes: [] }) + it("reads the current mode from the task when provider state differs", async () => { + vi.mocked(mockTask.getTaskMode).mockResolvedValue("architect") + mockGetState.mockResolvedValue({ mode: "orchestrator", customModes: [] }) const block = createBlock({ mode_slug: "code", reason: "switching back" }) @@ -340,18 +339,4 @@ describe("SwitchModeTool", () => { "Successfully switched from Architect mode to Code mode because: switching back.", ) }) - - it("should use defaultModeSlug when getState returns no mode", async () => { - mockGetState.mockResolvedValue({}) - - const block = createBlock({ mode_slug: "ask", reason: "test" }) - - await switchModeTool.handle(mockTask, block, mockCallbacks) - - // defaultModeSlug is "code" (from mock) - // Should report switching from Code mode - expect(mockCallbacks.pushToolResult).toHaveBeenCalledWith( - "Successfully switched from Code mode to Ask mode because: test.", - ) - }) }) diff --git a/src/core/tools/mcpServerRestriction.ts b/src/core/tools/mcpServerRestriction.ts index aa88125733..3bba963fb2 100644 --- a/src/core/tools/mcpServerRestriction.ts +++ b/src/core/tools/mcpServerRestriction.ts @@ -1,4 +1,4 @@ -import { getModeBySlug, defaultModeSlug } from "../../shared/modes" +import { getModeBySlug } from "../../shared/modes" import { Task } from "../task/Task" /** @@ -23,12 +23,12 @@ export function isMcpServerAllowed(serverName: string, allowedMcpServers?: strin } /** - * Resolves the current mode's MCP server allowlist from provider state. + * Resolves the task-local mode's MCP server allowlist. * * Returns `undefined` when the mode does not restrict MCP servers (or when the mode/state * cannot be resolved), which the predicate treats as "unrestricted". * - * @param task The current task, used to reach provider state. + * @param task The task whose execution mode controls the allowlist. * @returns The mode's `allowedMcpServers` allowlist, or `undefined` when unrestricted. */ export async function getAllowedMcpServersForTask(task: Task): Promise { @@ -43,7 +43,7 @@ export async function getAllowedMcpServersForTask(task: Task): Promise