mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-07 08:26:51 +00:00
feat: add heuristic-based model routing for cost optimization
Adds an experimental feature that dynamically routes API calls to a lighter/cheaper model when the task is in an information-gathering phase. Addresses Issue #11269 - Choose Model Dynamically Based on Request. ## How it works - New experiment flag: modelRouting (disabled by default) - New setting: modelRoutingLightModelId - the cheaper model ID to use - ModelRouter tracks tool usage per API turn - If previous turn only used "read" group tools (read_file, list_files, search_files, codebase_search), the next API call uses the light model - Edit, command, browser, and MCP tools always use the primary model - First turn always uses the primary model ## Changes - packages/types/src/experiment.ts: Add modelRouting experiment ID - packages/types/src/global-settings.ts: Add modelRoutingLightModelId - src/shared/experiments.ts: Add MODEL_ROUTING config - src/core/task/ModelRouter.ts: New heuristic-based model router - src/core/task/Task.ts: Integrate ModelRouter into task lifecycle - src/core/task/__tests__/ModelRouter.spec.ts: 35 tests (all passing)
This commit is contained in:
parent
8ef61bd32b
commit
6e090e7112
9 changed files with 554 additions and 1 deletions
|
|
@ -6,7 +6,13 @@ import type { Keys, Equals, AssertEqual } from "./type-fu.js"
|
|||
* ExperimentId
|
||||
*/
|
||||
|
||||
export const experimentIds = ["preventFocusDisruption", "imageGeneration", "runSlashCommand", "customTools"] as const
|
||||
export const experimentIds = [
|
||||
"preventFocusDisruption",
|
||||
"imageGeneration",
|
||||
"runSlashCommand",
|
||||
"customTools",
|
||||
"modelRouting",
|
||||
] as const
|
||||
|
||||
export const experimentIdsSchema = z.enum(experimentIds)
|
||||
|
||||
|
|
@ -21,6 +27,7 @@ export const experimentsSchema = z.object({
|
|||
imageGeneration: z.boolean().optional(),
|
||||
runSlashCommand: z.boolean().optional(),
|
||||
customTools: z.boolean().optional(),
|
||||
modelRouting: z.boolean().optional(),
|
||||
})
|
||||
|
||||
export type Experiments = z.infer<typeof experimentsSchema>
|
||||
|
|
|
|||
|
|
@ -232,6 +232,13 @@ export const globalSettingsSchema = z.object({
|
|||
* @default true
|
||||
*/
|
||||
showWorktreesInHomeScreen: z.boolean().optional(),
|
||||
|
||||
/**
|
||||
* The model ID to use for "light" tasks when model routing is enabled.
|
||||
* Must be a model available from the same provider as the primary model.
|
||||
* Requires the "modelRouting" experiment to be enabled.
|
||||
*/
|
||||
modelRoutingLightModelId: z.string().optional(),
|
||||
})
|
||||
|
||||
export type GlobalSettings = z.infer<typeof globalSettingsSchema>
|
||||
|
|
|
|||
|
|
@ -334,6 +334,7 @@ export type ExtensionState = Pick<
|
|||
| "maxGitStatusFiles"
|
||||
| "requestDelaySeconds"
|
||||
| "showWorktreesInHomeScreen"
|
||||
| "modelRoutingLightModelId"
|
||||
> & {
|
||||
version: string
|
||||
clineMessages: ClineMessage[]
|
||||
|
|
|
|||
233
src/core/task/ModelRouter.ts
Normal file
233
src/core/task/ModelRouter.ts
Normal file
|
|
@ -0,0 +1,233 @@
|
|||
import type { ToolName, ProviderSettings, Experiments, ToolGroup } from "@roo-code/types"
|
||||
|
||||
import { modelIdKeysByProvider, isTypicalProvider } from "@roo-code/types"
|
||||
|
||||
import { TOOL_GROUPS, ALWAYS_AVAILABLE_TOOLS } from "../../shared/tools"
|
||||
import { EXPERIMENT_IDS, experiments } from "../../shared/experiments"
|
||||
import { buildApiHandler, type ApiHandler } from "../../api"
|
||||
|
||||
/**
|
||||
* Tool complexity tier used for model routing decisions.
|
||||
*
|
||||
* - "light": Information-gathering tools (read_file, list_files, search_files, codebase_search)
|
||||
* - "standard": All other tools (edit, command, browser, mcp, etc.)
|
||||
*/
|
||||
export type ModelTier = "light" | "standard"
|
||||
|
||||
/**
|
||||
* Set of tool groups considered "light" for routing purposes.
|
||||
* Turns that only use tools from these groups (or always-available tools)
|
||||
* are eligible for routing to a cheaper model.
|
||||
*/
|
||||
const LIGHT_TOOL_GROUPS: ReadonlySet<ToolGroup> = new Set<ToolGroup>(["read"])
|
||||
|
||||
/**
|
||||
* Build a reverse map from tool name to tool group.
|
||||
*/
|
||||
function buildToolToGroupMap(): Map<string, ToolGroup> {
|
||||
const map = new Map<string, ToolGroup>()
|
||||
for (const [groupName, groupConfig] of Object.entries(TOOL_GROUPS)) {
|
||||
for (const tool of groupConfig.tools) {
|
||||
map.set(tool, groupName as ToolGroup)
|
||||
}
|
||||
if (groupConfig.customTools) {
|
||||
for (const tool of groupConfig.customTools) {
|
||||
map.set(tool, groupName as ToolGroup)
|
||||
}
|
||||
}
|
||||
}
|
||||
return map
|
||||
}
|
||||
|
||||
const TOOL_TO_GROUP = buildToolToGroupMap()
|
||||
|
||||
/**
|
||||
* Always-available tools as a Set for fast lookup.
|
||||
* These tools (ask_followup_question, attempt_completion, update_todo_list, etc.)
|
||||
* do not affect model tier classification.
|
||||
*/
|
||||
const ALWAYS_AVAILABLE_SET = new Set<string>(ALWAYS_AVAILABLE_TOOLS)
|
||||
|
||||
/**
|
||||
* Determines the tool group for a given tool name.
|
||||
* Returns undefined for always-available tools (which are tier-neutral).
|
||||
*/
|
||||
export function getToolGroup(toolName: string): ToolGroup | undefined {
|
||||
if (ALWAYS_AVAILABLE_SET.has(toolName)) {
|
||||
return undefined // Tier-neutral
|
||||
}
|
||||
return TOOL_TO_GROUP.get(toolName)
|
||||
}
|
||||
|
||||
/**
|
||||
* Classifies a set of tool names into a model tier.
|
||||
*
|
||||
* - If no tools were used → "standard" (pure reasoning, needs full model)
|
||||
* - If ALL tools are in light groups or always-available → "light"
|
||||
* - If ANY tool is in a non-light group → "standard"
|
||||
*/
|
||||
export function classifyToolUsage(toolNames: ReadonlySet<string>): ModelTier {
|
||||
if (toolNames.size === 0) {
|
||||
return "standard"
|
||||
}
|
||||
|
||||
let hasLightTool = false
|
||||
|
||||
for (const toolName of toolNames) {
|
||||
const group = getToolGroup(toolName)
|
||||
|
||||
if (group === undefined) {
|
||||
// Always-available tool, does not affect classification
|
||||
continue
|
||||
}
|
||||
|
||||
if (LIGHT_TOOL_GROUPS.has(group)) {
|
||||
hasLightTool = true
|
||||
} else {
|
||||
// Found a non-light tool, immediately classify as standard
|
||||
return "standard"
|
||||
}
|
||||
}
|
||||
|
||||
// If we only found light tools (and possibly always-available ones), classify as light
|
||||
return hasLightTool ? "light" : "standard"
|
||||
}
|
||||
|
||||
/**
|
||||
* ModelRouter provides heuristic-based model routing for cost optimization.
|
||||
*
|
||||
* It tracks which tools were used in each API turn and uses that information
|
||||
* to decide whether the next API call should use a lighter (cheaper) model
|
||||
* or the primary (more capable) model.
|
||||
*
|
||||
* ## Heuristic (v1)
|
||||
* - First turn: always use primary model
|
||||
* - If previous turn only used "read" group tools: use light model
|
||||
* - If previous turn used edit/command/browser/mcp tools: use primary model
|
||||
* - If previous turn had no tool calls (pure reasoning): use primary model
|
||||
*
|
||||
* ## Usage
|
||||
* ```typescript
|
||||
* const router = new ModelRouter()
|
||||
*
|
||||
* // Before each API call
|
||||
* if (router.shouldUseLightModel()) {
|
||||
* // Use light model
|
||||
* }
|
||||
*
|
||||
* // During tool execution
|
||||
* router.recordToolUse("read_file")
|
||||
*
|
||||
* // After API turn completes
|
||||
* router.endTurn()
|
||||
* ```
|
||||
*/
|
||||
export class ModelRouter {
|
||||
/** Tools used in the current (ongoing) turn */
|
||||
private currentTurnTools: Set<string> = new Set()
|
||||
|
||||
/** Classification of the previous (completed) turn */
|
||||
private previousTurnTier: ModelTier = "standard"
|
||||
|
||||
/** Whether at least one turn has completed */
|
||||
private hasPreviousTurn = false
|
||||
|
||||
/**
|
||||
* Record that a tool was used in the current turn.
|
||||
*/
|
||||
recordToolUse(toolName: ToolName): void {
|
||||
this.currentTurnTools.add(toolName)
|
||||
}
|
||||
|
||||
/**
|
||||
* Signal that the current turn has completed.
|
||||
* Moves current turn's tool usage to the "previous turn" classification.
|
||||
*/
|
||||
endTurn(): void {
|
||||
this.previousTurnTier = classifyToolUsage(this.currentTurnTools)
|
||||
this.currentTurnTools = new Set()
|
||||
this.hasPreviousTurn = true
|
||||
}
|
||||
|
||||
/**
|
||||
* Check whether the next API call should use the light model.
|
||||
*
|
||||
* Returns true only if:
|
||||
* - At least one turn has completed (never on first turn)
|
||||
* - The previous turn was classified as "light"
|
||||
*/
|
||||
shouldUseLightModel(): boolean {
|
||||
return this.hasPreviousTurn && this.previousTurnTier === "light"
|
||||
}
|
||||
|
||||
/**
|
||||
* Get the current tier classification (for debugging/logging).
|
||||
*/
|
||||
getCurrentTier(): ModelTier {
|
||||
return this.hasPreviousTurn ? this.previousTurnTier : "standard"
|
||||
}
|
||||
|
||||
/**
|
||||
* Reset the router state (e.g., when task is restarted).
|
||||
*/
|
||||
reset(): void {
|
||||
this.currentTurnTools = new Set()
|
||||
this.previousTurnTier = "standard"
|
||||
this.hasPreviousTurn = false
|
||||
}
|
||||
|
||||
/**
|
||||
* Check if model routing is enabled based on experiment settings and configuration.
|
||||
*
|
||||
* @param experimentsConfig - The experiments configuration
|
||||
* @param lightModelId - The light model ID from settings
|
||||
* @returns true if model routing is fully configured and enabled
|
||||
*/
|
||||
static isEnabled(experimentsConfig: Experiments | undefined, lightModelId: string | undefined): boolean {
|
||||
if (!experimentsConfig || !lightModelId || lightModelId.trim() === "") {
|
||||
return false
|
||||
}
|
||||
return experiments.isEnabled(experimentsConfig, EXPERIMENT_IDS.MODEL_ROUTING as any)
|
||||
}
|
||||
|
||||
/**
|
||||
* Build a ProviderSettings with the light model ID substituted in place of
|
||||
* the primary model ID. The provider and all other settings remain the same.
|
||||
*
|
||||
* @param baseConfig - The primary provider settings
|
||||
* @param lightModelId - The model ID to use for light tasks
|
||||
* @returns A new ProviderSettings with the light model, or null if the provider
|
||||
* is not supported for model routing
|
||||
*/
|
||||
static buildLightModelConfig(baseConfig: ProviderSettings, lightModelId: string): ProviderSettings | null {
|
||||
const provider = baseConfig.apiProvider
|
||||
if (!provider || !isTypicalProvider(provider)) {
|
||||
return null
|
||||
}
|
||||
|
||||
const modelIdKey = modelIdKeysByProvider[provider]
|
||||
if (!modelIdKey) {
|
||||
return null
|
||||
}
|
||||
|
||||
return {
|
||||
...baseConfig,
|
||||
[modelIdKey]: lightModelId,
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Build an ApiHandler configured for the light model.
|
||||
*
|
||||
* @param baseConfig - The primary provider settings
|
||||
* @param lightModelId - The model ID to use for light tasks
|
||||
* @returns An ApiHandler for the light model, or null if routing is not possible
|
||||
*/
|
||||
static buildLightModelHandler(baseConfig: ProviderSettings, lightModelId: string): ApiHandler | null {
|
||||
const lightConfig = ModelRouter.buildLightModelConfig(baseConfig, lightModelId)
|
||||
if (!lightConfig) {
|
||||
return null
|
||||
}
|
||||
return buildApiHandler(lightConfig)
|
||||
}
|
||||
}
|
||||
|
|
@ -96,6 +96,7 @@ import { getTaskDirectoryPath } from "../../utils/storage"
|
|||
import { formatResponse } from "../prompts/responses"
|
||||
import { SYSTEM_PROMPT } from "../prompts/system"
|
||||
import { buildNativeToolsArrayWithRestrictions } from "./build-tools"
|
||||
import { ModelRouter } from "./ModelRouter"
|
||||
|
||||
// core modules
|
||||
import { ToolRepetitionDetector } from "../tools/ToolRepetitionDetector"
|
||||
|
|
@ -297,6 +298,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
|
|||
}
|
||||
|
||||
toolRepetitionDetector: ToolRepetitionDetector
|
||||
modelRouter: ModelRouter
|
||||
rooIgnoreController?: RooIgnoreController
|
||||
rooProtectedController?: RooProtectedController
|
||||
fileContextTracker: FileContextTracker
|
||||
|
|
@ -614,6 +616,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
|
|||
|
||||
this.apiConfiguration = apiConfiguration
|
||||
this.api = buildApiHandler(this.apiConfiguration)
|
||||
this.modelRouter = new ModelRouter()
|
||||
this.autoApprovalHandler = new AutoApprovalHandler()
|
||||
|
||||
this.urlContentFetcher = new UrlContentFetcher(provider.context)
|
||||
|
|
@ -2894,6 +2897,28 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
|
|||
|
||||
await this.diffViewProvider.reset()
|
||||
|
||||
// Model routing: temporarily swap to light model if heuristics say so
|
||||
let primaryApiHandler: typeof this.api | undefined
|
||||
{
|
||||
const routingState = await this.providerRef.deref()?.getState()
|
||||
if (
|
||||
ModelRouter.isEnabled(routingState?.experiments, routingState?.modelRoutingLightModelId) &&
|
||||
this.modelRouter.shouldUseLightModel()
|
||||
) {
|
||||
const lightHandler = ModelRouter.buildLightModelHandler(
|
||||
this.apiConfiguration,
|
||||
routingState!.modelRoutingLightModelId!,
|
||||
)
|
||||
if (lightHandler) {
|
||||
primaryApiHandler = this.api
|
||||
this.api = lightHandler
|
||||
console.log(
|
||||
`[Task#${this.taskId}] Model routing: using light model "${routingState!.modelRoutingLightModelId}" for this turn`,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Cache model info once per API request to avoid repeated calls during streaming
|
||||
// This is especially important for tools and background usage collection
|
||||
this.cachedStreamingModel = this.api.getModel()
|
||||
|
|
@ -3598,6 +3623,13 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
|
|||
|
||||
await pWaitFor(() => this.userMessageContentReady)
|
||||
|
||||
// Model routing: end current turn and restore primary handler
|
||||
this.modelRouter.endTurn()
|
||||
if (primaryApiHandler) {
|
||||
this.api = primaryApiHandler
|
||||
primaryApiHandler = undefined
|
||||
}
|
||||
|
||||
// If the model did not tool use, then we need to tell it to
|
||||
// either use a tool or attempt_completion.
|
||||
const didToolUse = this.assistantMessageContent.some(
|
||||
|
|
@ -4639,6 +4671,7 @@ export class Task extends EventEmitter<TaskEvents> implements TaskLike {
|
|||
}
|
||||
|
||||
this.toolUsage[toolName].attempts++
|
||||
this.modelRouter.recordToolUse(toolName)
|
||||
}
|
||||
|
||||
public recordToolError(toolName: ToolName, error?: string) {
|
||||
|
|
|
|||
266
src/core/task/__tests__/ModelRouter.spec.ts
Normal file
266
src/core/task/__tests__/ModelRouter.spec.ts
Normal file
|
|
@ -0,0 +1,266 @@
|
|||
import type { ToolName, ProviderSettings, Experiments } from "@roo-code/types"
|
||||
|
||||
import { ModelRouter, classifyToolUsage, getToolGroup, type ModelTier } from "../ModelRouter"
|
||||
|
||||
describe("getToolGroup", () => {
|
||||
it("returns 'read' for read group tools", () => {
|
||||
expect(getToolGroup("read_file")).toBe("read")
|
||||
expect(getToolGroup("search_files")).toBe("read")
|
||||
expect(getToolGroup("list_files")).toBe("read")
|
||||
expect(getToolGroup("codebase_search")).toBe("read")
|
||||
})
|
||||
|
||||
it("returns 'edit' for edit group tools", () => {
|
||||
expect(getToolGroup("apply_diff")).toBe("edit")
|
||||
expect(getToolGroup("write_to_file")).toBe("edit")
|
||||
})
|
||||
|
||||
it("returns 'command' for command group tools", () => {
|
||||
expect(getToolGroup("execute_command")).toBe("command")
|
||||
expect(getToolGroup("read_command_output")).toBe("command")
|
||||
})
|
||||
|
||||
it("returns 'browser' for browser group tools", () => {
|
||||
expect(getToolGroup("browser_action")).toBe("browser")
|
||||
})
|
||||
|
||||
it("returns 'mcp' for mcp group tools", () => {
|
||||
expect(getToolGroup("use_mcp_tool")).toBe("mcp")
|
||||
expect(getToolGroup("access_mcp_resource")).toBe("mcp")
|
||||
})
|
||||
|
||||
it("returns undefined for always-available tools including mode tools (tier-neutral)", () => {
|
||||
// switch_mode and new_task are in ALWAYS_AVAILABLE_TOOLS, so they are tier-neutral
|
||||
expect(getToolGroup("switch_mode")).toBeUndefined()
|
||||
expect(getToolGroup("new_task")).toBeUndefined()
|
||||
expect(getToolGroup("ask_followup_question")).toBeUndefined()
|
||||
expect(getToolGroup("attempt_completion")).toBeUndefined()
|
||||
expect(getToolGroup("update_todo_list")).toBeUndefined()
|
||||
expect(getToolGroup("run_slash_command")).toBeUndefined()
|
||||
expect(getToolGroup("skill")).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe("classifyToolUsage", () => {
|
||||
it('returns "standard" when no tools were used', () => {
|
||||
expect(classifyToolUsage(new Set())).toBe("standard")
|
||||
})
|
||||
|
||||
it('returns "light" when only read tools were used', () => {
|
||||
expect(classifyToolUsage(new Set(["read_file"]))).toBe("light")
|
||||
expect(classifyToolUsage(new Set(["read_file", "search_files"]))).toBe("light")
|
||||
expect(classifyToolUsage(new Set(["read_file", "list_files", "codebase_search"]))).toBe("light")
|
||||
})
|
||||
|
||||
it('returns "light" when read tools and always-available tools are used', () => {
|
||||
expect(classifyToolUsage(new Set(["read_file", "ask_followup_question"]))).toBe("light")
|
||||
expect(classifyToolUsage(new Set(["search_files", "update_todo_list"]))).toBe("light")
|
||||
})
|
||||
|
||||
it('returns "standard" when only always-available tools are used (no read tools)', () => {
|
||||
// Only always-available tools - no light tools present, so "standard"
|
||||
expect(classifyToolUsage(new Set(["ask_followup_question"]))).toBe("standard")
|
||||
expect(classifyToolUsage(new Set(["update_todo_list"]))).toBe("standard")
|
||||
expect(classifyToolUsage(new Set(["attempt_completion"]))).toBe("standard")
|
||||
})
|
||||
|
||||
it('returns "standard" when any edit tool is used', () => {
|
||||
expect(classifyToolUsage(new Set(["read_file", "apply_diff"]))).toBe("standard")
|
||||
expect(classifyToolUsage(new Set(["write_to_file"]))).toBe("standard")
|
||||
})
|
||||
|
||||
it('returns "standard" when any command tool is used', () => {
|
||||
expect(classifyToolUsage(new Set(["read_file", "execute_command"]))).toBe("standard")
|
||||
})
|
||||
|
||||
it('returns "standard" when any browser tool is used', () => {
|
||||
expect(classifyToolUsage(new Set(["read_file", "browser_action"]))).toBe("standard")
|
||||
})
|
||||
|
||||
it('returns "standard" when any mcp tool is used', () => {
|
||||
expect(classifyToolUsage(new Set(["use_mcp_tool"]))).toBe("standard")
|
||||
})
|
||||
})
|
||||
|
||||
describe("ModelRouter", () => {
|
||||
let router: ModelRouter
|
||||
|
||||
beforeEach(() => {
|
||||
router = new ModelRouter()
|
||||
})
|
||||
|
||||
describe("shouldUseLightModel", () => {
|
||||
it("returns false on first turn (no previous turn)", () => {
|
||||
expect(router.shouldUseLightModel()).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false after first turn with no tools", () => {
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(false)
|
||||
})
|
||||
|
||||
it("returns true after a turn with only read tools", () => {
|
||||
router.recordToolUse("read_file" as ToolName)
|
||||
router.recordToolUse("search_files" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(true)
|
||||
})
|
||||
|
||||
it("returns false after a turn with edit tools", () => {
|
||||
router.recordToolUse("read_file" as ToolName)
|
||||
router.recordToolUse("apply_diff" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false after a turn with command tools", () => {
|
||||
router.recordToolUse("execute_command" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(false)
|
||||
})
|
||||
|
||||
it("returns true when read tools + always-available tools used", () => {
|
||||
router.recordToolUse("read_file" as ToolName)
|
||||
router.recordToolUse("update_todo_list" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(true)
|
||||
})
|
||||
|
||||
it("tracks multiple turns correctly", () => {
|
||||
// Turn 1: read only -> next should use light
|
||||
router.recordToolUse("read_file" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(true)
|
||||
|
||||
// Turn 2: edit -> next should use standard
|
||||
router.recordToolUse("write_to_file" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(false)
|
||||
|
||||
// Turn 3: read only again -> next should use light
|
||||
router.recordToolUse("list_files" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
describe("getCurrentTier", () => {
|
||||
it('returns "standard" before any turn completes', () => {
|
||||
expect(router.getCurrentTier()).toBe("standard")
|
||||
})
|
||||
|
||||
it('returns "light" after a read-only turn', () => {
|
||||
router.recordToolUse("read_file" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.getCurrentTier()).toBe("light")
|
||||
})
|
||||
})
|
||||
|
||||
describe("reset", () => {
|
||||
it("resets router state to initial", () => {
|
||||
router.recordToolUse("read_file" as ToolName)
|
||||
router.endTurn()
|
||||
expect(router.shouldUseLightModel()).toBe(true)
|
||||
|
||||
router.reset()
|
||||
expect(router.shouldUseLightModel()).toBe(false)
|
||||
expect(router.getCurrentTier()).toBe("standard")
|
||||
})
|
||||
})
|
||||
|
||||
describe("isEnabled", () => {
|
||||
it("returns false when experiments is undefined", () => {
|
||||
expect(ModelRouter.isEnabled(undefined, "some-model")).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false when lightModelId is undefined", () => {
|
||||
const experiments: Experiments = { modelRouting: true }
|
||||
expect(ModelRouter.isEnabled(experiments, undefined)).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false when lightModelId is empty string", () => {
|
||||
const experiments: Experiments = { modelRouting: true }
|
||||
expect(ModelRouter.isEnabled(experiments, "")).toBe(false)
|
||||
expect(ModelRouter.isEnabled(experiments, " ")).toBe(false)
|
||||
})
|
||||
|
||||
it("returns false when experiment is not enabled", () => {
|
||||
const experiments: Experiments = { modelRouting: false }
|
||||
expect(ModelRouter.isEnabled(experiments, "some-model")).toBe(false)
|
||||
})
|
||||
|
||||
it("returns true when experiment is enabled and lightModelId is set", () => {
|
||||
const experiments: Experiments = { modelRouting: true }
|
||||
expect(ModelRouter.isEnabled(experiments, "claude-3-haiku-20241022")).toBe(true)
|
||||
})
|
||||
})
|
||||
|
||||
describe("buildLightModelConfig", () => {
|
||||
it("returns null for unsupported provider types", () => {
|
||||
const config: ProviderSettings = {
|
||||
apiProvider: "fake-ai" as any,
|
||||
}
|
||||
expect(ModelRouter.buildLightModelConfig(config, "some-model")).toBeNull()
|
||||
})
|
||||
|
||||
it("returns null when apiProvider is not set", () => {
|
||||
const config: ProviderSettings = {}
|
||||
expect(ModelRouter.buildLightModelConfig(config, "some-model")).toBeNull()
|
||||
})
|
||||
|
||||
it("creates config with light model ID for anthropic provider", () => {
|
||||
const config: ProviderSettings = {
|
||||
apiProvider: "anthropic",
|
||||
apiModelId: "claude-sonnet-4-20250514",
|
||||
apiKey: "test-key",
|
||||
}
|
||||
const result = ModelRouter.buildLightModelConfig(config, "claude-3-haiku-20241022")
|
||||
expect(result).not.toBeNull()
|
||||
expect(result!.apiProvider).toBe("anthropic")
|
||||
expect(result!.apiModelId).toBe("claude-3-haiku-20241022")
|
||||
expect(result!.apiKey).toBe("test-key")
|
||||
})
|
||||
|
||||
it("creates config with light model ID for openrouter provider", () => {
|
||||
const config: ProviderSettings = {
|
||||
apiProvider: "openrouter",
|
||||
openRouterModelId: "anthropic/claude-sonnet-4-20250514",
|
||||
openRouterApiKey: "test-key",
|
||||
}
|
||||
const result = ModelRouter.buildLightModelConfig(config, "anthropic/claude-3-haiku-20241022")
|
||||
expect(result).not.toBeNull()
|
||||
expect(result!.apiProvider).toBe("openrouter")
|
||||
expect(result!.openRouterModelId).toBe("anthropic/claude-3-haiku-20241022")
|
||||
expect(result!.openRouterApiKey).toBe("test-key")
|
||||
})
|
||||
|
||||
it("creates config with light model ID for gemini provider", () => {
|
||||
const config: ProviderSettings = {
|
||||
apiProvider: "gemini",
|
||||
apiModelId: "gemini-2.5-pro",
|
||||
geminiApiKey: "test-key",
|
||||
}
|
||||
const result = ModelRouter.buildLightModelConfig(config, "gemini-2.0-flash")
|
||||
expect(result).not.toBeNull()
|
||||
expect(result!.apiProvider).toBe("gemini")
|
||||
expect(result!.apiModelId).toBe("gemini-2.0-flash")
|
||||
expect(result!.geminiApiKey).toBe("test-key")
|
||||
})
|
||||
|
||||
it("preserves all other settings from base config", () => {
|
||||
const config: ProviderSettings = {
|
||||
apiProvider: "anthropic",
|
||||
apiModelId: "claude-sonnet-4-20250514",
|
||||
apiKey: "test-key",
|
||||
modelTemperature: 0.5,
|
||||
enableReasoningEffort: true,
|
||||
reasoningEffort: "medium",
|
||||
}
|
||||
const result = ModelRouter.buildLightModelConfig(config, "claude-3-haiku-20241022")
|
||||
expect(result).not.toBeNull()
|
||||
expect(result!.modelTemperature).toBe(0.5)
|
||||
expect(result!.enableReasoningEffort).toBe(true)
|
||||
expect(result!.reasoningEffort).toBe("medium")
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
@ -2411,6 +2411,7 @@ export class ClineProvider
|
|||
customSupportPrompts: stateValues.customSupportPrompts ?? {},
|
||||
enhancementApiConfigId: stateValues.enhancementApiConfigId,
|
||||
experiments: stateValues.experiments ?? experimentDefault,
|
||||
modelRoutingLightModelId: stateValues.modelRoutingLightModelId,
|
||||
autoApprovalEnabled: stateValues.autoApprovalEnabled ?? false,
|
||||
customModes,
|
||||
maxOpenTabsContext: stateValues.maxOpenTabsContext ?? 20,
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ describe("experiments", () => {
|
|||
imageGeneration: false,
|
||||
runSlashCommand: false,
|
||||
customTools: false,
|
||||
modelRouting: false,
|
||||
}
|
||||
expect(Experiments.isEnabled(experiments, EXPERIMENT_IDS.PREVENT_FOCUS_DISRUPTION)).toBe(false)
|
||||
})
|
||||
|
|
@ -31,6 +32,7 @@ describe("experiments", () => {
|
|||
imageGeneration: false,
|
||||
runSlashCommand: false,
|
||||
customTools: false,
|
||||
modelRouting: false,
|
||||
}
|
||||
expect(Experiments.isEnabled(experiments, EXPERIMENT_IDS.PREVENT_FOCUS_DISRUPTION)).toBe(true)
|
||||
})
|
||||
|
|
@ -41,6 +43,7 @@ describe("experiments", () => {
|
|||
imageGeneration: false,
|
||||
runSlashCommand: false,
|
||||
customTools: false,
|
||||
modelRouting: false,
|
||||
}
|
||||
expect(Experiments.isEnabled(experiments, EXPERIMENT_IDS.PREVENT_FOCUS_DISRUPTION)).toBe(false)
|
||||
})
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ export const EXPERIMENT_IDS = {
|
|||
IMAGE_GENERATION: "imageGeneration",
|
||||
RUN_SLASH_COMMAND: "runSlashCommand",
|
||||
CUSTOM_TOOLS: "customTools",
|
||||
MODEL_ROUTING: "modelRouting",
|
||||
} as const satisfies Record<string, ExperimentId>
|
||||
|
||||
type _AssertExperimentIds = AssertEqual<Equals<ExperimentId, Values<typeof EXPERIMENT_IDS>>>
|
||||
|
|
@ -20,6 +21,7 @@ export const experimentConfigsMap: Record<ExperimentKey, ExperimentConfig> = {
|
|||
IMAGE_GENERATION: { enabled: false },
|
||||
RUN_SLASH_COMMAND: { enabled: false },
|
||||
CUSTOM_TOOLS: { enabled: false },
|
||||
MODEL_ROUTING: { enabled: false },
|
||||
}
|
||||
|
||||
export const experimentDefault = Object.fromEntries(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue