diff --git a/src/router/RouterRuntime.ts b/src/router/RouterRuntime.ts index 6965b3796..10255e268 100644 --- a/src/router/RouterRuntime.ts +++ b/src/router/RouterRuntime.ts @@ -16,6 +16,7 @@ import { DEFAULT_SUBAGENT_POLICY, type RouterConfig, type RouterModelRef, + type RouterTokenSaverConfig, } from "./config/schema.js"; import type { PilotDeckCustomRouter, @@ -99,6 +100,22 @@ export type RouterRuntime = { shutdown(): Promise; }; +function isTierDowngrade( + tiers: RouterTokenSaverConfig["tiers"] | undefined, + previousTier: string | undefined, + nextTier: string, +): boolean { + if (!tiers || !previousTier) { + return false; + } + + // Tier declaration order is the router's low-to-high capability order. + const tierOrder = Object.keys(tiers); + const previousIndex = tierOrder.indexOf(previousTier); + const nextIndex = tierOrder.indexOf(nextTier); + return previousIndex >= 0 && nextIndex >= 0 && nextIndex < previousIndex; +} + export function createRouterRuntime( config: RouterConfig, deps: RouterRuntimeDeps, @@ -218,6 +235,8 @@ export function createRouterRuntime( function maybePreserveStickyForCache( current: RouterModelRef | undefined, next: RouterModelRef, + previousTier: string | undefined, + nextTier: string, messages: CanonicalModelRequest["messages"], lastUsage: import("../model/index.js").CanonicalUsage | undefined, ): { selection: RouterModelRef; mutation?: RouterMutationsLog["cacheAwareSwitch"] } { @@ -229,6 +248,13 @@ export function createRouterRuntime( return { selection: next }; } + // Cache savings may justify retaining a stronger model on a downgrade, + // but must never veto a judge-requested capability upgrade. If direction + // cannot be established, prefer the fresh judge decision over stale state. + if (!isTierDowngrade(config.tokenSaver?.tiers, previousTier, nextTier)) { + return { selection: next }; + } + const estimatedInputTokens = countMessagesTokens(messages); const observedInputTokens = lastUsage?.inputTokens ?? 0; const observedCacheReadTokens = lastUsage?.cacheReadTokens ?? 0; @@ -364,6 +390,7 @@ export function createRouterRuntime( : sticky?.stickyProvider && sticky.stickyModel ? { id: `${sticky.stickyProvider}/${sticky.stickyModel}`, provider: sticky.stickyProvider, model: sticky.stickyModel } : undefined; + const previousTokenSaverTier = sticky?.tokenSaverTier ?? input.metadata?.previousTier; let selection: RouterModelRef | undefined = custom?.provider && custom.model ? { id: `${custom.provider}/${custom.model}`, provider: custom.provider, model: custom.model } @@ -444,6 +471,8 @@ export function createRouterRuntime( const cacheAware = maybePreserveStickyForCache( previousStickySelection, selection, + previousTokenSaverTier, + tokenSaver.tier, input.request.messages, baseUsage, ); diff --git a/tests/router/cache-aware-switching.spec.ts b/tests/router/cache-aware-switching.spec.ts new file mode 100644 index 000000000..c8d033323 --- /dev/null +++ b/tests/router/cache-aware-switching.spec.ts @@ -0,0 +1,181 @@ +import assert from "node:assert/strict"; +import test from "node:test"; + +import type { CanonicalModelRequest, ModelRuntime, ModelRuntimeOptions } from "../../src/model/index.js"; +import { createRouterRuntime } from "../../src/router/RouterRuntime.js"; +import type { RouterConfig } from "../../src/router/config/schema.js"; + +const capabilities = { + supportsToolUse: true, + supportsStreaming: true, + supportsParallelToolCalls: false, + supportsThinking: false, + supportsJsonSchema: false, + supportsSystemPrompt: true, + supportsPromptCache: true, + maxContextTokens: 8192, + maxOutputTokens: 1024, +}; + +const modelRuntime: ModelRuntime = { + async *stream(_request: CanonicalModelRequest, _options?: ModelRuntimeOptions) {}, + async complete() { + throw new Error("not used"); + }, + getCapabilities() { + return capabilities; + }, + getMultimodal() { + return { input: ["text"] }; + }, + getProviderProtocol() { + return "openai"; + }, + getProviderBaseUrl(provider: string) { + return `https://${provider}.invalid`; + }, +}; + +test("cache-aware switching never blocks a judge-requested tier upgrade", async () => { + const decision = await decideAcrossTurns({ + previousTier: "simple", + nextTier: "reasoning", + modelPricing: { + "test/simple": { input: 1, cacheRead: 0.1 }, + "test/reasoning": { input: 10, cacheRead: 1 }, + }, + }); + + assert.equal(decision.tokenSaverTier, "reasoning"); + assert.equal(decision.model, "reasoning"); + assert.equal(decision.mutations.cacheAwareSwitch, undefined); +}); + +test("cache-aware switching can retain the current model on a tier downgrade", async () => { + const decision = await decideAcrossTurns({ + previousTier: "reasoning", + nextTier: "simple", + modelPricing: { + "test/reasoning": { input: 10, cacheRead: 0.1 }, + "test/simple": { input: 1, cacheRead: 0.5 }, + }, + }); + + assert.equal(decision.tokenSaverTier, "reasoning"); + assert.equal(decision.model, "reasoning"); + assert.equal(decision.mutations.cacheAwareSwitch?.action, "kept_sticky"); +}); + +test("cache-aware switching still allows a cost-effective tier downgrade", async () => { + const decision = await decideAcrossTurns({ + previousTier: "reasoning", + nextTier: "simple", + modelPricing: { + "test/reasoning": { input: 10, cacheRead: 2 }, + "test/simple": { input: 1, cacheRead: 0.5 }, + }, + }); + + assert.equal(decision.tokenSaverTier, "simple"); + assert.equal(decision.model, "simple"); + assert.equal(decision.mutations.cacheAwareSwitch?.action, "switched"); +}); + +async function decideAcrossTurns(input: { + previousTier: "simple" | "medium" | "complex" | "reasoning"; + nextTier: "simple" | "medium" | "complex" | "reasoning"; + modelPricing: NonNullable["modelPricing"]; +}) { + const judgeTiers = [input.previousTier, input.nextTier]; + let judgeCall = 0; + const judgeRuntime = { + async complete() { + const tier = judgeTiers[judgeCall++]; + return { + role: "assistant" as const, + content: [{ type: "text" as const, text: `${tier}` }], + finishReason: "stop" as const, + }; + }, + } as unknown as ModelRuntime; + + const router = createRouterRuntime(createConfig(input.modelPricing), { + modelRuntime, + judgeRuntime, + }); + + try { + const firstDecision = await router.decide({ + sessionId: "cache-aware-tier-direction", + isMainAgent: true, + request: createRequest([ + { role: "user", content: [{ type: "text", text: "first task" }] }, + ]), + }); + assert.equal(firstDecision.tokenSaverTier, input.previousTier); + + router.observeUsage("cache-aware-tier-direction", { + inputTokens: 1_000, + outputTokens: 10, + cacheReadTokens: 1_000, + totalTokens: 2_010, + }); + const previous = router.invalidateSticky("cache-aware-tier-direction"); + + const decision = await router.decide({ + sessionId: "cache-aware-tier-direction", + isMainAgent: true, + metadata: { + previousTier: previous.previousTier, + previousProvider: previous.previousProvider, + previousModel: previous.previousModel, + }, + request: createRequest([ + { role: "user", content: [{ type: "text", text: "first task" }] }, + { role: "assistant", content: [{ type: "text", text: "first response" }] }, + { role: "user", content: [{ type: "text", text: "second task" }] }, + ]), + }); + assert.equal(judgeCall, 2); + return decision; + } finally { + await router.shutdown(); + } +} + +function createRequest(messages: CanonicalModelRequest["messages"]): CanonicalModelRequest { + return { + provider: "test", + model: "default", + messages, + }; +} + +function createConfig( + modelPricing: NonNullable["modelPricing"], +): RouterConfig { + const tierModel = (tier: string) => ({ + model: { id: `test/${tier}`, provider: "test", model: tier }, + }); + + return { + enabled: true, + scenarios: { + default: { id: "test/default", provider: "test", model: "default" }, + }, + tokenSaver: { + enabled: true, + judge: { id: "test/judge", provider: "test", model: "judge" }, + defaultTier: "medium", + judgeTimeoutMs: 5_000, + tiers: { + simple: tierModel("simple"), + medium: tierModel("medium"), + complex: tierModel("complex"), + reasoning: tierModel("reasoning"), + }, + cacheAwareSwitching: { enabled: true, minSavingsRatio: 0 }, + }, + stats: { enabled: false, modelPricing }, + }; +}