Skip to content
Open
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
29 changes: 29 additions & 0 deletions src/router/RouterRuntime.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import {
DEFAULT_SUBAGENT_POLICY,
type RouterConfig,
type RouterModelRef,
type RouterTokenSaverConfig,
} from "./config/schema.js";
import type {
PilotDeckCustomRouter,
Expand Down Expand Up @@ -99,6 +100,22 @@ export type RouterRuntime = {
shutdown(): Promise<void>;
};

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,
Expand Down Expand Up @@ -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"] } {
Expand All @@ -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;
Expand Down Expand Up @@ -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 }
Expand Down Expand Up @@ -444,6 +471,8 @@ export function createRouterRuntime(
const cacheAware = maybePreserveStickyForCache(
previousStickySelection,
selection,
previousTokenSaverTier,
tokenSaver.tier,
input.request.messages,
baseUsage,
);
Expand Down
181 changes: 181 additions & 0 deletions tests/router/cache-aware-switching.spec.ts
Original file line number Diff line number Diff line change
@@ -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<RouterConfig["stats"]>["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>${tier}</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<RouterConfig["stats"]>["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 },
};
}