feat(complexity_router): escalate oversized prompts to a tier that fits before dispatch (#38844)

* feat(complexity_router): escalate oversized prompts to a tier that fits before dispatch

The classifier scores complexity and never prompt size, so a long agentic
session whose newest ask is trivial classifies SIMPLE onto a small-window
tier and the provider rejects it with a context-window 400 that nothing
retries. The gate runs after classification on every decision path
(classify tail and session-affinity pin), estimates prompt tokens
including the out-of-band carriers (top-level system, tools,
instructions), and when the decided tier provably cannot hold the prompt
moves the request to the lowest configured tier with a model whose
declared window fits, restricting the pick to fitting models when the
decided tier can keep it. Models with no resolvable window are never
escalated away from or onto, escalated decisions are never written as
session pins, and the decision records context_escalated plus the
original tier in spend logs.

Resolves LIT-6503

* fix(complexity_router): judge groups by smallest window, bound skips by bytes, filter adaptive picks

Review-round rework, one mechanism per finding. A group is judged by its
smallest resolvable deployment window, since the core router picks within
a group with no fit check. The counting skip is gated on UTF-8 byte
length, which BPE token counts can never exceed, so token-dense scripts
cannot slip past it; only a real tokenizer count ever moves a request and
a failed count leaves the placement alone. The fit facts now filter every
adaptive phase including cold start and the tier fallbacks. Window
questions adopt the declared provider and never resolve authenticating
providers, and a router instance without get_model_list degrades the gate
to a no-op. Tests rebuilt on real Router instances resolving deployment
model_info end to end, plus a full-path test through
async_get_available_deployment
This commit is contained in:
tin-berri 2026-08-31 16:12:59 -07:00 • committed by GitHub
parent 81c8c93bef
commit f93d9b6b67
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 720 additions and 15 deletions

View file

@ -733,11 +733,20 @@ def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bo
of the three: an agent names the conversation on its first turn, so the cheapest tier would be
the pin every session starts with, and the real work that follows would run there for the whole
TTL. It describes what that one call is, never what the session's traffic looks like.
A context-window escalation describes the prompt's size, not the session's complexity, and
size shrinks again the moment the client compacts: pinning the escalated tier would hold the
session on the big-window model long after the oversized context that forced it is gone. The
gate re-fires per request, so leaving these unpinned costs nothing but the classifier call.
"""
return decision is None or decision.get("cause") not in (
"default_model_fallback",
"plan_mode",
"housekeeping",
return decision is None or (
decision.get("cause")
not in (
"default_model_fallback",
"plan_mode",
"housekeeping",
)
and not decision.get("context_escalated")
)
@ -786,6 +795,39 @@ class ClassificationOutcome(NamedTuple):
classifier_cost: float | None = None
def _allowed(models: tuple[str, ...], fit_filter: frozenset[str] | None) -> tuple[str, ...]:
return models if fit_filter is None else tuple(model for model in models if model in fit_filter)
def _apply_context_placement(
tier: ComplexityTier | str, signals: tuple[str, ...], placement: _ContextWindowPlacement | None
) -> tuple[ComplexityTier | str, tuple[str, ...], ComplexityTier | str | None]:
"""(final tier, signals, original tier when the gate escalated, else None)."""
if placement is None:
return tier, signals, None
if _tier_name(placement.tier) == _tier_name(tier):
return placement.tier, signals, None
return placement.tier, (*signals, "context_escalation"), tier
def _window_can_hold(window: int | None, needed: int, buffer: float) -> bool:
return window is None or needed <= int(window * buffer)
def _group_provably_fits(facts: tuple[int | None, bool], needed: int, buffer: float) -> bool:
window, has_unknown = facts
return window is not None and not has_unknown and needed <= int(window * buffer)
class _ContextWindowPlacement(NamedTuple):
"""Where the context-window gate placed the request: the placement tier, the subset of its
pool the pick may use, and every configured group not provably misfit (the adaptive filter)."""
tier: ComplexityTier | str
allowed_models: tuple[str, ...]
holdable_models: frozenset[str]
class _SessionAffinityPin(NamedTuple):
model: str
tier: ComplexityTier | None
@ -1222,6 +1264,7 @@ class ComplexityRouter(CustomLogger):
classifier_cost: float | None = None,
conversation_continuing: bool = True,
tier_litellm_params: Mapping[str, object] | None = None,
context_escalation_original_tier: ComplexityTier | str | None = None,
) -> StandardLoggingRoutingDecision:
"""Assemble the per-request provenance record for this router's decision.
@ -1271,6 +1314,12 @@ class ComplexityRouter(CustomLogger):
decision["classifier_model"] = classifier_model
if classifier_cost is not None:
decision["classifier_cost"] = classifier_cost
if context_escalation_original_tier is not None:
# The pair travels together: the flag says the gate moved the request off its
# decided tier on prompt size, and the original tier names where the decision
# (classifier, keyword rule, or session pin) had placed it before physics did.
decision["context_escalated"] = True
decision["context_escalation_original_tier"] = _tier_name(context_escalation_original_tier)
if tier_litellm_params:
masked_tier_litellm_params: Final = mask_credentials_in_payload(tier_litellm_params)
if isinstance(masked_tier_litellm_params, Mapping):
@ -1671,7 +1720,7 @@ class ComplexityRouter(CustomLogger):
return entry.litellm_params if entry is not None else MappingProxyType({})
@staticmethod
def _pick_from_tier_value(model: str | list[str], tier_key: str) -> str:
def _pick_from_tier_value(model: str | Sequence[str], tier_key: str) -> str:
if isinstance(model, str):
return model
if not model:
@ -1687,15 +1736,21 @@ class ComplexityRouter(CustomLogger):
raw_messages: list[dict[str, Any]] | None,
resolved_messages: list[dict[str, Any]] | None,
request_kwargs: dict,
allowed_models: tuple[str, ...] | None = None,
) -> str:
if not self.config.plugins:
if allowed_models is not None:
return self._pick_from_tier_value(allowed_models, _tier_name(tier))
return self.get_model_for_tier(tier)
from litellm.types.router import RoutingContext
tier_key: Final = _tier_name(tier)
metadata_key: Final = get_metadata_variable_name_from_kwargs(request_kwargs)
pool: Final = tuple(self._tier_pools().get(tier_key, ()))
full_pool: Final = tuple(self._tier_pools().get(tier_key, ()))
pool: Final = (
tuple(model for model in full_pool if model in allowed_models) if allowed_models is not None else full_pool
)
if not pool:
# Nothing for the plugins to filter. Falling through would raise the
# plugin-filtering error below and send the operator hunting for a policy
@ -1789,6 +1844,7 @@ class ComplexityRouter(CustomLogger):
request_kwargs: dict[str, Any] | None = None,
hard_floor: ComplexityTier | str | None = None,
hard_ceiling: ComplexityTier | str | None = None,
fit_filter: frozenset[str] | None = None,
) -> str:
"""hard_floor excludes every candidate whose tiers all sit below it, turning this pick's
soft floors (a distance penalty a high-scoring cheap model can outweigh) into a hard
@ -1801,7 +1857,10 @@ class ComplexityRouter(CustomLogger):
tier because that is all it is worth, so a bandit trading cost for quality has nothing to
win and must not reach above it. Without it the distance penalty is the only thing holding
the tier, and a deployment that lowers tier_distance_penalty silently gets the expensive
model back while the routing decision still reads as the cheapest tier."""
model back while the routing decision still reads as the cheapest tier.
fit_filter excludes candidates the context-window gate proved cannot hold the prompt,
in every phase including cold start and the tier fallbacks."""
from litellm.router_strategy.adaptive_router.bandit import (
normalized_cost,
thompson_sample,
@ -1812,12 +1871,12 @@ class ComplexityRouter(CustomLogger):
if adaptive is None or not isinstance(classified_tier, ComplexityTier):
# Custom tier names have no severity index; adaptive is rejected alongside
# tier_definitions, so this guard is the contract for any future caller.
return self.get_model_for_tier(classified_tier)
return self._fitting_tier_fallback(classified_tier, fit_filter)
request_type: Final = classify_prompt(user_message)
classified_idx: Final = TIER_SEVERITY_ORDER.index(classified_tier)
pools: Final = self._tier_pools()
classified_candidates: Final = tuple(pools.get(_tier_name(classified_tier), ()))
classified_candidates: Final = _allowed(tuple(pools.get(_tier_name(classified_tier), ())), fit_filter)
cold_start_candidates: Final = tuple(
model for model in classified_candidates if adaptive._cells[(request_type, model)].total_samples == 0
)
@ -1847,9 +1906,9 @@ class ComplexityRouter(CustomLogger):
if self.config.adaptive_eligible == "classified_tier":
candidates = list(classified_candidates)
if not candidates:
return self.get_model_for_tier(classified_tier)
return self._fitting_tier_fallback(classified_tier, fit_filter)
else:
candidates = list(adaptive.config.available_models)
candidates = list(_allowed(tuple(adaptive.config.available_models), fit_filter))
all_costs: Final = [adaptive.model_to_cost.get(m, 0.0) for m in candidates]
quality_weight: Final = self.config.adaptive_weights.quality
@ -1896,7 +1955,7 @@ class ComplexityRouter(CustomLogger):
best_score = score
best_model = model
if best_model is None:
return self.get_model_for_tier(classified_tier)
return self._fitting_tier_fallback(classified_tier, fit_filter)
if request_kwargs is not None:
metadata = request_kwargs.setdefault("metadata", {})
if isinstance(metadata, dict):
@ -1913,6 +1972,12 @@ class ComplexityRouter(CustomLogger):
}
return best_model
def _fitting_tier_fallback(self, classified_tier: ComplexityTier | str, fit_filter: frozenset[str] | None) -> str:
fitting: Final = _allowed(tuple(self._tier_pools().get(_tier_name(classified_tier), ())), fit_filter)
if fit_filter is not None and fitting:
return self._pick_from_tier_value(fitting, _tier_name(classified_tier))
return self.get_model_for_tier(classified_tier)
def _resolve_plan_mode_floor(self) -> ComplexityTier | str | None:
"""The configured floor as an active tier: the built-in enum member, or the defined
name itself for a custom tier set; None when the feature is off."""
@ -1983,6 +2048,163 @@ class ComplexityRouter(CustomLogger):
return None
return name if self.config.has_custom_tiers else ComplexityTier(name)
def _deployment_window(self, group: str, deployment: Mapping[str, object]) -> int | None:
from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider
deployment_model_info: Final = deployment.get("model_info")
declared: Final = (
deployment_model_info.get("max_input_tokens") if isinstance(deployment_model_info, Mapping) else None
)
if isinstance(declared, int):
return declared
litellm_params: Final = deployment.get("litellm_params")
params: Final = litellm_params if isinstance(litellm_params, Mapping) else EMPTY_MAPPING
provider_override: Final = params.get("custom_llm_provider")
# get_router_model_info resolves the provider, and get_llm_provider runs the OAuth device
# flow for github_copilot/chatgpt, so a metadata question must never reach it for those.
if declared_authenticating_provider(
str(params.get("model") or ""), provider_override if isinstance(provider_override, str) else None
):
return None
try:
model_info: Final = self.litellm_router_instance.get_router_model_info(
deployment=cast(dict, deployment), # cast-ok: router deployments are plain dicts
received_model_name=group,
)
window: Final = model_info.get("max_input_tokens")
except Exception: # noqa: BLE001 # best-effort: an unmappable deployment must not hide the others
return None
return window if isinstance(window, int) else None
def _group_window_facts(self, group: str) -> tuple[int | None, bool]:
"""(smallest declared context window across the group's deployments, whether any deployment
declares none). The core router picks a deployment within the group without a fit check, so
the group is only as safe as its smallest member."""
list_models: Final = getattr(self.litellm_router_instance, "get_model_list", None)
deployments: Final = list_models(model_name=group) if callable(list_models) else None
if not isinstance(deployments, list) or not deployments:
return (None, True)
windows: Final = tuple(
window for deployment in deployments if (window := self._deployment_window(group, deployment)) is not None
)
return (min(windows) if windows else None, len(windows) < len(deployments))
@staticmethod
def _out_of_band_request_text(request_kwargs: Mapping[str, object]) -> str:
"""Prompt content the resolved message list never carries: the Responses API's
`instructions`, the /v1/messages top-level `system` block, and tool definitions.
A coding agent's context is dominated by these."""
import json
instructions: Final = request_kwargs.get("instructions")
proxy_request: Final = request_kwargs.get("proxy_server_request")
body: Final = proxy_request.get("body") if isinstance(proxy_request, Mapping) else None
system: Final = body.get("system") if isinstance(body, Mapping) else None
tools: Final = (
body.get("tools") if isinstance(body, Mapping) and body.get("tools") else request_kwargs.get("tools")
)
tools_text = ""
if tools:
try:
tools_text = json.dumps(tools, default=str)
except (TypeError, ValueError):
tools_text = str(tools)
return (
(instructions if isinstance(instructions, str) else "")
+ (str(system) if system is not None else "")
+ tools_text
)
def _request_byte_upper_bound(
self, resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: Mapping[str, object]
) -> int:
"""UTF-8 byte length of all prompt content. BPE emits at least one byte per token in every
script, so the token count never exceeds this and 'bytes fit' soundly skips counting."""
content_bytes: Final = sum(len(str(m.get("content") or "").encode()) for m in resolved_messages or ())
return content_bytes + len(self._out_of_band_request_text(request_kwargs).encode())
async def _counted_request_tokens(
self, resolved_messages: Sequence[Mapping[str, object]], request_kwargs: Mapping[str, object]
) -> int | None:
"""Real-tokenizer count of the resolved messages plus the out-of-band carriers, off the
event loop; None when counting fails, and the gate then leaves the placement alone."""
import litellm
from litellm.litellm_core_utils.asyncify import asyncify
out_of_band: Final = self._out_of_band_request_text(request_kwargs)
try:
counted: Final = await asyncify(litellm.token_counter)(
messages=cast(list, resolved_messages) # cast-ok: token_counter only iterates the sequence
)
return counted + (await asyncify(litellm.token_counter)(text=out_of_band) if out_of_band else 0)
except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request
verbose_router_logger.debug("ComplexityRouter: context-window token count failed. Got - %s", e)
return None
async def _context_window_placement(
self,
tier: ComplexityTier | str,
resolved_messages: Sequence[Mapping[str, object]] | None,
request_kwargs: Mapping[str, object],
pool_override: tuple[str, ...] | None = None,
) -> _ContextWindowPlacement | None:
"""Correct a decided placement whose models provably cannot hold the prompt, or None
(the placement stands). Only a real tokenizer count ever moves a request, escalation
lands only on groups whose every deployment declares a fitting window, and a group
with no resolvable window is never moved on faith in either direction."""
if not self.config.enable_context_window_escalation or not resolved_messages:
return None
pools: Final = self._tier_pools()
pool: Final = pool_override if pool_override is not None else tuple(pools.get(_tier_name(tier), ()))
if not pool:
return None
facts: Final = MappingProxyType({group: self._group_window_facts(group) for group in pool})
known_windows: Final = tuple(window for window, _ in facts.values() if window is not None)
if not known_windows:
return None
buffer: Final = self.config.context_window_escalation_buffer
if self._request_byte_upper_bound(resolved_messages, request_kwargs) <= int(min(known_windows) * buffer):
return None
needed: Final = await self._counted_request_tokens(resolved_messages, request_kwargs)
if needed is None:
return None
return self._placement_for_tokens(tier=tier, pool=pool, pools=pools, facts=facts, needed=needed)
def _placement_for_tokens(
self,
*,
tier: ComplexityTier | str,
pool: tuple[str, ...],
pools: Mapping[str, list[str]],
facts: Mapping[str, tuple[int | None, bool]],
needed: int,
) -> _ContextWindowPlacement | None:
buffer: Final = self.config.context_window_escalation_buffer
in_tier: Final = tuple(group for group in pool if _window_can_hold(facts[group][0], needed, buffer))
if in_tier and len(in_tier) == len(pool):
return None
holdable: Final = frozenset(
group
for tier_pool in pools.values()
for group in tier_pool
if _window_can_hold(self._group_window_facts(group)[0], needed, buffer)
)
if in_tier:
return _ContextWindowPlacement(tier=tier, allowed_models=in_tier, holdable_models=holdable)
for name in self.config.tier_names()[self._active_tier_severity(tier) + 1 :]:
proven = tuple(
group
for group in pools.get(name, ())
if _group_provably_fits(self._group_window_facts(group), needed, buffer)
)
if proven:
return _ContextWindowPlacement(
tier=name if self.config.has_custom_tiers else ComplexityTier(name),
allowed_models=proven,
holdable_models=holdable,
)
return None
def _apply_plan_mode_floor(self, tier: ComplexityTier | str) -> ComplexityTier | str:
"""The higher of the decided tier and the plan-mode floor; identity when the floor is unset."""
floor: Final = self._resolve_plan_mode_floor()
@ -2381,6 +2603,26 @@ class ComplexityRouter(CustomLogger):
session_model: Final = routed_model
if plan_floored and pinned_tier is not None:
routed_model = self.get_model_for_tier(self._apply_plan_mode_floor(pinned_tier))
pin_source_tier: Final = self._tier_for_model(routed_model)
pin_placement: Final = (
await self._context_window_placement(
pin_source_tier, resolved_messages, request_kwargs, pool_override=(routed_model,)
)
if pin_source_tier is not None
else None
)
pin_context_original_tier: Final = (
pin_source_tier
if pin_placement is not None
and pin_source_tier is not None
and _tier_name(pin_placement.tier) != _tier_name(pin_source_tier)
else None
)
if pin_placement is not None and pin_context_original_tier is not None:
# The stored pin below keeps the session's own model on purpose.
routed_model = self._pick_from_tier_value(
pin_placement.allowed_models, _tier_name(pin_placement.tier)
)
# Refresh the TTL on every hit so an active session doesn't lose its
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
@ -2405,7 +2647,11 @@ class ComplexityRouter(CustomLogger):
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, routed_model=%s", cause, routed_model
)
routed_pin_tier: Final = self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier
routed_pin_tier: Final = (
pin_placement.tier
if pin_placement is not None and pin_context_original_tier is not None
else (self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier)
)
session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model)
has_original_messages: Final = messages is not None and len(messages) > 0
return self._with_session_deployment_affinity(
@ -2422,6 +2668,7 @@ class ComplexityRouter(CustomLogger):
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
context_escalation_original_tier=pin_context_original_tier,
),
)
)
@ -2616,6 +2863,8 @@ class ComplexityRouter(CustomLogger):
plan_floored: Final = tier != pre_floor_tier
if plan_floored:
signals = (*signals, "plan_mode_floor")
context_placement: Final = await self._context_window_placement(tier, resolved_messages, request_kwargs)
tier, signals, context_original_tier = _apply_context_placement(tier, signals, context_placement)
score_repr: Final = f"{score:.3f}" if score is not None else "n/a"
fallback_model: Final = self.config.default_model if not self.config.plugins else None
# A sentinel-carrying request skips the failure exit below, whether or not the floor
@ -2662,8 +2911,15 @@ class ComplexityRouter(CustomLogger):
# the cheapest tier would then contradict the floor and bound the pick below the tier
# the decision reports.
housekeeping_ceiling: Final = tier if outcome.cause == "housekeeping" else None
# A context-escalated tier becomes the hard floor: a floor the bandit can slide
# under is not a floor.
routed_model = self._soft_floor_pick(
tier, user_message, request_kwargs, hard_floor=plan_floor, hard_ceiling=housekeeping_ceiling
tier,
user_message,
request_kwargs,
hard_floor=tier if context_original_tier is not None else plan_floor,
hard_ceiling=housekeeping_ceiling,
fit_filter=context_placement.holdable_models if context_placement is not None else None,
)
adaptive: Final = self._ensure_adaptive_router()
if adaptive is not None:
@ -2680,7 +2936,13 @@ class ComplexityRouter(CustomLogger):
routed_model,
)
else:
routed_model = await self._pick_model_for_tier(tier, messages, resolved_messages, request_kwargs)
routed_model = await self._pick_model_for_tier(
tier,
messages,
resolved_messages,
request_kwargs,
allowed_models=context_placement.allowed_models if context_placement is not None else None,
)
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, tier=%s, score=%s, signals=%s, routed_model=%s",
outcome.cause,
@ -2733,5 +2995,6 @@ class ComplexityRouter(CustomLogger):
classifier_model=classifier_model,
classifier_cost=outcome.classifier_cost,
tier_litellm_params=tier_litellm_params,
context_escalation_original_tier=context_original_tier,
),
)

View file

@ -823,6 +823,32 @@ class ComplexityRouterConfig(BaseModel):
),
)
enable_context_window_escalation: bool = Field(
default=True,
description=(
"Escalate a request off a tier whose models provably cannot hold its prompt, before "
"dispatch. The classifier scores complexity and never prompt size, so a long agentic "
"session whose newest ask is trivial lands on a small-window tier and the provider "
"rejects it with a context-window 400 that nothing retries. When every model of the "
"decided tier has a declared window smaller than the estimated prompt, the request "
"moves to the lowest configured tier with a model whose declared window fits; when "
"only some of the tier's models fit, the pick is restricted to those and the tier "
"keeps the request. Models with no resolvable window are never escalated away from "
"and never escalated onto. Set false to dispatch on complexity alone, as before."
),
)
context_window_escalation_buffer: float = Field(
default=0.95,
gt=0,
le=1,
description=(
"Fraction of a model's declared context window the estimated prompt must fit within. "
"The token count is an estimate, so fitting against the full window would dispatch "
"prompts that the provider's own tokenizer then rejects; 0.95 leaves room for that "
"drift plus the response tokens."
),
)
# Semantic (embedding) matching for keyword_tier_rules instead of literal text matching
semantic_keyword_matching: bool = Field(
default=False,

View file

@ -2886,6 +2886,8 @@ class StandardLoggingRoutingDecision(TypedDict, total=False):
classifier_model: str
classifier_cost: float
escalated: bool
context_escalated: bool # writable-ok: Pydantic warns on ReadOnly TypedDict fields
context_escalation_original_tier: str # writable-ok: Pydantic warns on ReadOnly TypedDict fields
tier_boundaries: StandardLoggingRoutingDecisionTierBoundaries
reasoning_override_min_score: float # writable-ok: Pydantic warns on ReadOnly TypedDict fields
conversation_continuing: bool
@ -2912,6 +2914,8 @@ DERIVED_ROUTING_DECISION_FIELDS: Final[frozenset[str]] = frozenset(
"classifier_model",
"classifier_cost",
"escalated",
"context_escalated",
"context_escalation_original_tier",
"tier_boundaries",
"reasoning_override_min_score",
"conversation_continuing",

View file

@ -10021,3 +10021,399 @@ class TestHeuristicFirst:
)
outcome = await router.aclassify(NO_SIGNAL_PROMPT)
assert outcome.cause == "default_model_fallback"
def _windowed_router(*deployments: tuple) -> Router:
"""Real Router; each deployment is (group, provider_model, declared window or None).
None means no declared override on a model the cost map does not know: unresolvable."""
return Router(
model_list=[
{
"model_name": group,
"litellm_params": {"model": provider_model, "mock_response": "ok"},
**({"model_info": {"max_input_tokens": window}} if window is not None else {}),
}
for group, provider_model, window in deployments
]
)
_SMALL = ("small-model", "openai/gpt-3.5-turbo", 16385)
_BIG = ("big-model", "openai/gpt-4o-mini", 200000)
# A long agentic session whose newest ask is trivial: low-density filler the heuristic scores
# SIMPLE, sized well past a 16,385-token window so the fit check must move it.
_CONTEXT_FILLER = "The meeting notes were saved to the shared folder for later review this week. " * 2000
_OVERSIZED_TURNS = [
{"role": "user", "content": "Here is everything discussed so far. " + _CONTEXT_FILLER},
{"role": "assistant", "content": "Noted, I have read all of it."},
{"role": "user", "content": "ok continue"},
]
# ~40k CJK chars: chars/4 says ~10k tokens, the real tokenizer says several times that. A
# character-based shortcut would skip counting and dispatch this to a 16k window.
_CJK_TURNS = [
{"role": "user", "content": "会议记录已经保存到共享文件夹里,供大家本周晚些时候查阅和讨论使用。" * 1300},
{"role": "user", "content": "ok continue"},
]
def _tier_config(**overrides) -> Dict:
return {"tiers": {"SIMPLE": "small-model", "COMPLEX": "big-model"}, **overrides}
class TestContextWindowEscalation:
"""A tier decided on complexity alone must still hold the prompt, or the provider 400s.
The classifier never weighs prompt size (token count is a 0.10-weight scoring dimension,
below every tier boundary), so a long session ending in a trivial ask lands on the
smallest tier and dies upstream with no retry. The gate checks fit pre-dispatch, against
windows resolved through the real Router deployment chain.
"""
@pytest.mark.asyncio
async def test_an_oversized_simple_prompt_escalates_to_the_lowest_tier_that_fits(self):
"""The LIT-6503 regression: SIMPLE verdict, 17k-token prompt, 16,385-token tier model.
Unfixed, this dispatched to the small model and the provider rejected it with a
context-window 400 that neither the retry layer nor tier-keyed fallbacks catch.
"""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(),
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == "big-model"
assert result.routing_decision["context_escalated"] is True
assert result.routing_decision["context_escalation_original_tier"] == "SIMPLE"
assert result.routing_decision["tier"] == "COMPLEX"
assert "context_escalation" in result.routing_decision["signals"]
@pytest.mark.asyncio
async def test_a_prompt_that_fits_routes_exactly_as_before(self):
"""The gate must be invisible for normal traffic: same model, no escalation facts."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(),
)
result = await router.async_pre_routing_hook(
model="test-router", request_kwargs={}, messages=[{"role": "user", "content": "ok continue"}]
)
assert result is not None
assert result.model == "small-model"
assert "context_escalated" not in result.routing_decision
assert "context_escalation_original_tier" not in result.routing_decision
@pytest.mark.asyncio
async def test_the_pick_prefers_a_fitting_group_inside_the_decided_tier(self):
"""A tier holding both a small and a large group keeps the request and picks the one
that fits, which is cheaper than escalating and preserves the classifier's decision."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, ("mid-model", "openai/gpt-4o-mini", 200000), _BIG),
complexity_router_config={"tiers": {"SIMPLE": ["small-model", "mid-model"], "COMPLEX": "big-model"}},
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == "mid-model"
assert result.routing_decision["tier"] == "SIMPLE"
assert "context_escalated" not in result.routing_decision
@pytest.mark.asyncio
async def test_a_group_is_only_as_safe_as_its_smallest_deployment(self):
"""One group name can front deployments with different windows, and the core router
picks among them with no fit check, so retaining the group on its largest member
turns the pick into a coin flip against a 400. The gate judges the group by its
smallest resolvable window and escalates past it."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=Router(
model_list=[
{
"model_name": "mixed-pool",
"litellm_params": {"model": "openai/gpt-3.5-turbo", "mock_response": "ok"},
"model_info": {"max_input_tokens": 16385},
},
{
"model_name": "mixed-pool",
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "ok"},
"model_info": {"max_input_tokens": 200000},
},
{
"model_name": "big-model",
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "ok"},
"model_info": {"max_input_tokens": 200000},
},
]
),
complexity_router_config={"tiers": {"SIMPLE": "mixed-pool", "COMPLEX": "big-model"}},
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == "big-model"
assert result.routing_decision["context_escalated"] is True
@pytest.mark.asyncio
async def test_token_dense_text_cannot_slip_past_the_counting_shortcut(self):
"""CJK text runs several tokens per four characters, so a chars/4 shortcut would skip
the real count and dispatch an oversized prompt. The skip is gated on the UTF-8 byte
length, which the token count can never exceed."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(),
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_CJK_TURNS)
assert result is not None
assert result.model == "big-model"
assert result.routing_decision["context_escalated"] is True
@pytest.mark.asyncio
@pytest.mark.parametrize(
"deployments,tiers,expected_model",
[
(
(("small-model", "openai/unmapped-model-under-test", None), _BIG),
{"SIMPLE": "small-model", "COMPLEX": "big-model"},
"small-model",
),
(
(_SMALL, ("mid-model", "openai/another-unmapped-model", None), _BIG),
{"SIMPLE": "small-model", "MEDIUM": "mid-model", "COMPLEX": "big-model"},
"big-model",
),
((_SMALL,), {"SIMPLE": "small-model"}, "small-model"),
],
ids=["unknown-window-stays", "unproven-target-skipped", "nothing-fits-stays"],
)
async def test_unknown_windows_are_never_acted_on(self, deployments, tiers, expected_model):
"""No faith in either direction: a model with no resolvable window is never escalated
away from (its misfit is unprovable) and never escalated onto (its fit is unprovable);
when nothing provably fits, the classified tier stands and the client owns overflow."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(*deployments),
complexity_router_config={"tiers": tiers},
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == expected_model
@pytest.mark.asyncio
async def test_the_disabled_gate_dispatches_on_complexity_alone(self):
"""The escape hatch: enable_context_window_escalation false restores today's behavior."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(enable_context_window_escalation=False),
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == "small-model"
assert "context_escalated" not in result.routing_decision
@pytest.mark.asyncio
async def test_out_of_band_system_and_tools_count_against_the_window(self):
"""The Claude Code shape that live-testing caught: a tiny ask riding a top-level
`system` block and tool definitions that together dwarf the message list. None of
that reaches resolved messages on /v1/messages, so a gate reading only messages
dispatches a provably oversized request and the provider 400s anyway."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(),
)
result = await router.async_pre_routing_hook(
model="test-router",
request_kwargs={
"proxy_server_request": {
"body": {
"system": _CONTEXT_FILLER,
"tools": [{"name": f"tool_{i}", "description": _CONTEXT_FILLER[:500]} for i in range(20)],
}
}
},
messages=[{"role": "user", "content": "reply with exactly: rig check ok"}],
)
assert result is not None
assert result.model == "big-model"
assert result.routing_decision["context_escalated"] is True
@pytest.mark.asyncio
async def test_an_escalated_first_turn_never_becomes_the_session_pin(self):
"""Escalation describes the prompt's size, not the session: once the client compacts,
the next turn fits again, so pinning the big-window tier would hold the whole session
on it for the TTL. The escalated turn routes big, and the next fitting turn classifies
fresh instead of inheriting a pin."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(session_affinity=True),
)
session_kwargs = lambda: {"metadata": {"session_id": "s-1", "user_api_key_hash": "k-1"}} # noqa: E731
first = await router.async_pre_routing_hook(
model="test-router", request_kwargs=session_kwargs(), messages=_OVERSIZED_TURNS
)
second = await router.async_pre_routing_hook(
model="test-router", request_kwargs=session_kwargs(), messages=[{"role": "user", "content": "ok continue"}]
)
assert first is not None and first.model == "big-model"
assert second is not None and second.model == "small-model"
assert second.routing_decision["cause"] != "session_affinity_pin"
@pytest.mark.asyncio
async def test_a_pinned_session_escalates_per_request_and_keeps_its_pin(self):
"""The pin fast path skips classification, not physics: an oversized turn on a session
pinned to the small tier is served by the fitting tier, while the stored pin keeps the
session's own model so the first turn that fits again routes exactly as pinned."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=_windowed_router(_SMALL, _BIG),
complexity_router_config=_tier_config(session_affinity=True),
)
session_kwargs = lambda: {"metadata": {"session_id": "s-2", "user_api_key_hash": "k-2"}} # noqa: E731
pinned = await router.async_pre_routing_hook(
model="test-router", request_kwargs=session_kwargs(), messages=[{"role": "user", "content": "ok continue"}]
)
oversized = await router.async_pre_routing_hook(
model="test-router", request_kwargs=session_kwargs(), messages=_OVERSIZED_TURNS
)
back_to_small = await router.async_pre_routing_hook(
model="test-router", request_kwargs=session_kwargs(), messages=[{"role": "user", "content": "ok continue"}]
)
assert pinned is not None and pinned.model == "small-model"
assert oversized is not None and oversized.model == "big-model"
assert oversized.routing_decision["cause"] == "session_affinity_pin"
assert oversized.routing_decision["context_escalated"] is True
assert oversized.routing_decision["context_escalation_original_tier"] == "SIMPLE"
assert back_to_small is not None and back_to_small.model == "small-model"
assert back_to_small.routing_decision["cause"] == "session_affinity_pin"
@pytest.mark.asyncio
async def test_the_adaptive_cold_start_never_samples_a_model_that_cannot_hold_the_prompt(self):
"""The bandit's exploration is still bounded by physics: with the whole classified tier
unobserved, cold start samples only among models whose window holds the prompt."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=Router(
model_list=[
{
"model_name": "small-model",
"litellm_params": {"model": "openai/gpt-3.5-turbo", "mock_response": "ok"},
"model_info": {"max_input_tokens": 16385},
},
{
"model_name": "mid-model",
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "ok"},
"model_info": {"max_input_tokens": 200000},
},
]
),
complexity_router_config={"adaptive": True, "tiers": {"SIMPLE": ["small-model", "mid-model"]}},
)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == "mid-model"
@pytest.mark.asyncio
async def test_the_gate_never_resolves_an_authenticating_provider(self, monkeypatch, tmp_path):
"""Resolving github_copilot runs its OAuth device flow, so a window question must adopt
the declaration instead of resolving: the copilot group reads as unknown-window and the
request stays put, with zero copilot resolutions recorded."""
import json
import time
monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path))
(tmp_path / "api-key.json").write_text(json.dumps({"token": "tid=test", "expires_at": int(time.time()) + 3600}))
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=Router(
model_list=[
{"model_name": "cop-pool", "litellm_params": {"model": "github_copilot/gpt-4o"}},
{
"model_name": "big-model",
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "ok"},
"model_info": {"max_input_tokens": 200000},
},
]
),
complexity_router_config={"tiers": {"SIMPLE": "cop-pool", "COMPLEX": "big-model"}},
)
real_get_llm_provider = litellm.get_llm_provider
copilot_resolutions: List = []
def _guarded(*args, **kwargs):
target = str(kwargs.get("model") or (args[0] if args else "")) + str(kwargs.get("custom_llm_provider") or "")
if "github_copilot" in target:
copilot_resolutions.append(target)
raise RuntimeError("the gate must not resolve an authenticating provider")
return real_get_llm_provider(*args, **kwargs)
monkeypatch.setattr(litellm, "get_llm_provider", _guarded)
result = await router.async_pre_routing_hook(model="test-router", request_kwargs={}, messages=_OVERSIZED_TURNS)
assert result is not None
assert result.model == "cop-pool"
assert copilot_resolutions == []
@pytest.mark.asyncio
async def test_the_full_routing_path_serves_the_escalated_deployment(self):
"""End to end through Router.async_get_available_deployment: the auto-router alias with
an oversized prompt resolves to the big tier's deployment, and a small prompt to the
small tier's, with no mocking anywhere in the resolution chain."""
router = Router(
model_list=[
{
"model_name": "smart-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"tiers": {"SIMPLE": "small-model", "COMPLEX": "big-model"}},
},
},
{
"model_name": "small-model",
"litellm_params": {"model": "openai/gpt-3.5-turbo", "mock_response": "ok"},
"model_info": {"max_input_tokens": 16385},
},
{
"model_name": "big-model",
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "ok"},
"model_info": {"max_input_tokens": 200000},
},
]
)
oversized = await router.async_get_available_deployment(
model="smart-router", request_kwargs={}, messages=_OVERSIZED_TURNS
)
small = await router.async_get_available_deployment(
model="smart-router", request_kwargs={}, messages=[{"role": "user", "content": "ok continue"}]
)
assert oversized["model_name"] == "big-model"
assert small["model_name"] == "small-model"

View file

@ -34385,6 +34385,12 @@ export interface components {
* @description Keywords indicating code-related content
*/
code_keywords?: string[] | null;
/**
* Context Window Escalation Buffer
* @description Fraction of a model's declared context window the estimated prompt must fit within. The token count is an estimate, so fitting against the full window would dispatch prompts that the provider's own tokenizer then rejects; 0.95 leaves room for that drift plus the response tokens.
* @default 0.95
*/
context_window_escalation_buffer: number;
/**
* Custom Technical Keywords
* @description Domain-specific technical keywords appended to the effective base list (technical_keywords if set, otherwise DEFAULT_TECHNICAL_KEYWORDS). Order is preserved; duplicates are removed case-insensitively against the base list and within this list.
@ -34413,6 +34419,12 @@ export interface components {
* @description Embedding model (LiteLLM model name) used when semantic_keyword_matching is enabled
*/
embedding_model?: string | null;
/**
* Enable Context Window Escalation
* @description Escalate a request off a tier whose models provably cannot hold its prompt, before dispatch. The classifier scores complexity and never prompt size, so a long agentic session whose newest ask is trivial lands on a small-window tier and the provider rejects it with a context-window 400 that nothing retries. When every model of the decided tier has a declared window smaller than the estimated prompt, the request moves to the lowest configured tier with a model whose declared window fits; when only some of the tier's models fit, the pick is restricted to those and the tier keeps the request. Models with no resolvable window are never escalated away from and never escalated onto. Set false to dispatch on complexity alone, as before.
* @default true
*/
enable_context_window_escalation: boolean;
/**
* Escalation Keywords
* @description Case-sensitive phrases a user can include to force a bump to the next-higher complexity tier when they aren't satisfied with results (they can force a stronger model, but not choose which one). Defaults to ['LITELLM ESCALATE'] when unset; set to an empty list to disable.
@ -35602,6 +35614,10 @@ export interface components {
classifier_cost?: number;
/** Classifier Model */
classifier_model?: string;
/** Context Escalated */
context_escalated?: boolean;
/** Context Escalation Original Tier */
context_escalation_original_tier?: string;
/** Conversation Continuing */
conversation_continuing?: boolean;
/** Escalated */