feat(router): opt in to prompt-cache cost routing (#43232)

* feat(router): opt in to prompt-cache cost routing

* fix(router): address prompt-cache routing review

* ci: include cache-routing regressions in coverage
This commit is contained in:
tin-berri 2026-09-26 18:13:13 -07:00 • committed by GitHub
parent 501ef23f4a
commit de06c93767
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1337 additions and 124 deletions

View file

@ -46,6 +46,7 @@ legacy_paths() {
echo tests/unit/google_genai
echo tests/unit/router_strategy
echo tests/unit/router_utils
echo tests/unit/proxy/common_utils/test_cache_aware_routing.py
echo tests/unit/enterprise/enterprise_callbacks/send_emails
echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py
echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py

View file

@ -0,0 +1,236 @@
from __future__ import annotations
import time
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from urllib.parse import urlparse
from pydantic import BaseModel, JsonValue, TypeAdapter
import litellm
from litellm._internal_context import current_billing_time, pinned_billing_time
from litellm.caching.dual_cache import DualCache
from litellm.llms.anthropic.prompt_cache_prediction import (
NativePredictionTarget,
PromptPrefix,
TokenCounter,
UnsupportedPredictionTarget,
cache_scope,
count_prompt_tokens,
parse_prompt,
resolve_prediction_target,
supported_prediction_headers,
)
from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens
from litellm.proxy.hooks.prompt_cache_prediction import lookup
from litellm.types.management_endpoints.prompt_cache_prediction import (
CacheCostScenario,
CacheEvidence,
CachePredictionArm,
CacheTokenBuckets,
)
from litellm.types.router import Deployment
from litellm.utils import get_prompt_cache_min_tokens
__all__: Final = ("AnthropicCacheRouting", "TokenCounter", "predict_arm")
_JSON: Final = TypeAdapter(Mapping[str, JsonValue])
_NATIVE_OPTIONS: Final = frozenset(
(
"max_tokens",
"system",
"tools",
"tool_choice",
"thinking",
"output_config",
"cache_control",
"speed",
"service_tier",
"temperature",
"top_p",
"top_k",
"stop_sequences",
"stream",
)
)
class _ModelLimits(BaseModel):
max_input_tokens: int | None = None
max_output_tokens: int | None = None
@dataclass(frozen=True, slots=True)
class AnthropicCacheRouting:
body: Mapping[str, JsonValue]
prefix: PromptPrefix
requested_output_limit: int
@staticmethod
def request_body(
url: str,
headers: Mapping[str, str],
body: Mapping[str, JsonValue],
request_kwargs: Mapping[str, object],
messages: Sequence[Mapping[str, object]] | None,
) -> Mapping[str, JsonValue] | None:
if not urlparse(url).path.endswith("/v1/messages") or not supported_prediction_headers(headers):
return None
return _JSON.validate_python(
MappingProxyType(
{
**body,
**MappingProxyType({key: request_kwargs[key] for key in _NATIVE_OPTIONS if key in request_kwargs}),
"messages": messages,
}
)
)
@classmethod
def from_body(cls, body: Mapping[str, JsonValue]) -> AnthropicCacheRouting | None:
prefix: Final = parse_prompt(body)
limit: Final = body.get("max_tokens")
if prefix is None or not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0:
return None
return cls(body, prefix, limit)
@staticmethod
def supports(deployment: Deployment) -> bool:
return isinstance(resolve_prediction_target(deployment.litellm_params), NativePredictionTarget)
async def is_warm(self, deployment: Deployment, caller: str, cache: DualCache, now: float) -> bool:
target: Final = resolve_prediction_target(deployment.litellm_params)
if not isinstance(target, NativePredictionTarget):
return False
scope: Final = cache_scope(caller, deployment.model_info.id or "", target.api_key, target.model)
observation: Final = await lookup(cache, scope, self.prefix, now=now)
return observation is not None and observation.expires_at > now
@staticmethod
def fits(deployment: Deployment, input_tokens: int, output_tokens: int) -> bool:
target: Final = resolve_prediction_target(deployment.litellm_params)
if not isinstance(target, NativePredictionTarget):
return False
limits: Final = _ModelLimits.model_validate(
MappingProxyType(
{
**litellm.get_model_info(target.model, custom_llm_provider="anthropic"),
**deployment.model_info.model_dump(exclude_none=True),
}
)
)
return (
limits.max_input_tokens is not None
and input_tokens + output_tokens <= limits.max_input_tokens
and limits.max_output_tokens is not None
and output_tokens <= limits.max_output_tokens
)
async def predict(
self,
deployment: Deployment,
caller: str,
cache: DualCache,
counter: TokenCounter,
now: float | None,
) -> CachePredictionArm:
return await predict_arm(deployment, self.body, self.prefix, caller, cache, counter, now=now)
@staticmethod
def cost(arm: CachePredictionArm, output_tokens: int) -> float | None:
return (
price_cache_tokens(arm.model or "", arm.deployment_id, arm.estimate.tokens, output_tokens)
if arm.estimate is not None
else None
)
@staticmethod
async def count_tokens(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None:
return await count_prompt_tokens(model, api_key, body)
def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets:
return CacheTokenBuckets(
uncached_input_tokens=suffix_tokens,
cache_read_input_tokens=read_tokens,
cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0,
cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0,
)
def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None:
cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens)
return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None
async def predict_arm(
deployment: Deployment,
body: Mapping[str, JsonValue],
prefix: PromptPrefix,
caller_key_hash: str,
cache: DualCache,
token_counter: TokenCounter,
now: float | None = None,
) -> CachePredictionArm:
deployment_id: Final = deployment.model_info.id or ""
params: Final = deployment.litellm_params
unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model)
if deployment.model_info.blocked:
return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"}))
target: Final = resolve_prediction_target(params)
if isinstance(target, UnsupportedPredictionTarget):
return unknown.model_copy(update=MappingProxyType({"reason": target.reason}))
model: Final = target.model
api_key: Final = target.api_key
total_count: Final = await token_counter(model, api_key, body)
prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body)
if total_count is None or prefix_count is None or total_count < prefix_count:
return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"}))
scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model)
checked_at: Final = time.time() if now is None else now
observation: Final = await lookup(cache, scope, prefix, now=checked_at)
exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint
cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count
if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable):
return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"}))
suffix: Final = total_count - cacheable
evidence: Final = (
CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at)
if observation is not None
else None
)
if cacheable < get_prompt_cache_min_tokens(params.model):
disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count))
if disabled is None:
return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"}))
return CachePredictionArm(
deployment_id=deployment_id,
model=model,
cache_state="disabled",
reason="below_cache_minimum",
estimate=disabled,
cold=disabled,
warm=disabled,
token_count_source="anthropic_count_tokens",
)
fresh: Final = observation is not None and observation.expires_at > checked_at
read: Final = observation.cached_tokens if fresh and observation is not None else 0
with pinned_billing_time(current_billing_time()):
cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds))
warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds))
estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds))
if cold is None or warm is None or estimate is None:
return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"}))
return CachePredictionArm(
deployment_id=deployment_id,
model=model,
cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown",
reason=None if fresh else "observation_expired" if observation else "no_compatible_observation",
estimate=estimate,
cold=cold,
warm=warm,
evidence=evidence,
token_count_source="anthropic_count_tokens",
)

View file

@ -0,0 +1,343 @@
from __future__ import annotations
import asyncio
import time
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
from litellm.llms.anthropic.cache_aware_routing import AnthropicCacheRouting, TokenCounter
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
from litellm.types.router import Deployment, PreRoutingHookResponse
from litellm.types.utils import StandardLoggingRoutingDecision
if TYPE_CHECKING:
from litellm.router import Router
_MESSAGES: Final = TypeAdapter(list[Mapping[str, object]])
_MAPPING: Final = TypeAdapter(Mapping[str, object])
_DEPLOYMENTS: Final[TypeAdapter[tuple[Deployment, ...] | Deployment]] = TypeAdapter(tuple[Deployment, ...] | Deployment)
_MARKER_OPTIONS: Final = frozenset(
("model", "complexity_router_config", "rpm", "tpm", "tags", "timeout", "stream_timeout", "num_retries")
)
_CLASSIFIED_CAUSES: Final = frozenset(
{
"heuristic_scorer",
"heuristic_v2",
"reasoning_override",
"llm_classifier",
"llm_v2_classifier",
"jev_classifier",
"capability_classifier",
"heuristic_first_short_circuit",
"hybrid_short_circuit",
"classifier_plugin",
}
)
class _ProxyRequest(BaseModel):
model_config = ConfigDict(strict=True)
url: str
body: Mapping[str, JsonValue]
headers: Mapping[str, str]
class _CallerSettings(BaseModel):
config: Mapping[str, object] | None = None
@dataclass(frozen=True, slots=True)
class CacheAwareChoice:
model: str
tier: str
deployment_id: str
original_cost: float
estimated_cost: float
@dataclass(frozen=True, slots=True)
class _Candidate:
model: str
tier: str
deployment: Deployment
def eligible_models(
config: ComplexityRouterConfig, decision: StandardLoggingRoutingDecision
) -> tuple[tuple[str, str], ...]:
tier: Final = decision.get("tier")
order: Final = config.tier_names()
tier_entries: Final = chain.from_iterable(config.tier_model_configs.values())
if (
tier is None
or tier not in order
or decision.get("cause") not in _CLASSIFIED_CAUSES
or config.has_custom_tiers
or config.plugins
or config.adaptive
or config.session_affinity
or config.classification_mode != "every_request"
or any(entry.litellm_params for entry in tier_entries)
or any(not isinstance(model, str) for model in config.tiers.values())
):
return ()
floor: Final = order.index(tier)
eligible: Final = tuple(
(name, model) for name, model in config.tiers.items() if isinstance(model, str) and name in order[floor:]
)
return tuple(entry for index, entry in enumerate(eligible) if entry[1] not in tuple(m for _, m in eligible[:index]))
def _candidate(router: Router, tier: str, model: str, request_kwargs: Mapping[str, object]) -> _Candidate | None:
deployments: Final = router.deployments_for_request(model, request_kwargs)
if len(deployments) != 1:
return None
deployment: Final = Deployment.model_validate(deployments[0])
if deployment.model_info.blocked or not deployment.model_info.id or not AnthropicCacheRouting.supports(deployment):
return None
return _Candidate(model, tier, deployment)
async def _available(
candidate: _Candidate,
router: Router,
caller: UserAPIKeyAuth,
request_kwargs: Mapping[str, object],
messages: Sequence[Mapping[str, object]] | None,
) -> bool:
try:
await can_key_call_resolved_model(
model=candidate.model, llm_model_list=router.get_model_list(), valid_token=caller, llm_router=router
)
healthy: Final = _DEPLOYMENTS.validate_python(
await router.async_get_healthy_deployments( # pyright: ignore[reportUnknownMemberType] # legacy router results are validated at this boundary
model=candidate.model,
messages=_MESSAGES.validate_python(messages) if messages else None, # pyright: ignore[reportArgumentType] # router annotations predate structured native messages
request_kwargs=dict(request_kwargs), # mutable-ok: Router's filtering API accepts a request dictionary
)
)
except Exception: # noqa: BLE001 # an unavailable optional candidate must not fail the originally selected route
return False
available: Final = (healthy,) if isinstance(healthy, Deployment) else healthy
return any(entry.model_info.id == candidate.deployment.model_info.id for entry in available)
def supported_router_marker(router: Router, alias: str, request_kwargs: Mapping[str, object]) -> bool:
markers: Final = tuple(
Deployment.model_validate(entry) for entry in router.deployments_for_request(alias, request_kwargs)
)
return bool(markers) and all(
marker.litellm_params.model == "auto_router/complexity_router"
and not frozenset(marker.litellm_params.model_dump(exclude_defaults=True, exclude_none=True)) - _MARKER_OPTIONS
for marker in markers
)
async def select_cached_model(
*,
router: Router,
config: ComplexityRouterConfig,
params_for_model: Callable[[str, str], Mapping[str, object]],
response: PreRoutingHookResponse,
body: Mapping[str, JsonValue],
request_kwargs: Mapping[str, object],
messages: Sequence[Mapping[str, object]] | None,
caller: UserAPIKeyAuth,
cache: DualCache,
counter_for_model: Callable[[str], TokenCounter],
now: float | None = None,
) -> CacheAwareChoice | None:
checked_at: Final = time.time() if now is None else now
decision: Final = response.routing_decision
provider: Final = AnthropicCacheRouting.from_body(body)
if not config.cache_aware_routing or decision is None or provider is None or not caller.api_key:
return None
names: Final = eligible_models(config, decision)
if not names or response.model not in tuple(model for _, model in names):
return None
candidates: Final = tuple(
candidate for tier, model in names if (candidate := _candidate(router, tier, model, request_kwargs)) is not None
)
original: Final = next((candidate for candidate in candidates if candidate.model == response.model), None)
if original is None:
return None
alternatives: Final = tuple(candidate for candidate in candidates if candidate.model != original.model)
warm_flags: Final = await asyncio.gather(
*(provider.is_warm(candidate.deployment, caller.api_key, cache, checked_at) for candidate in alternatives)
)
warm: Final = tuple(candidate for candidate, fresh in zip(alternatives, warm_flags) if fresh)
if not warm:
return None
considered: Final = (original, *warm)
availability: Final = await asyncio.gather(
*(_available(candidate, router, caller, request_kwargs, messages) for candidate in considered)
)
authorized: Final = tuple(candidate for candidate, available in zip(warm, availability[1:]) if available)
if not availability[0] or not authorized:
return None
compared: Final = (original, *authorized)
output_limits: Final = tuple(
params_for_model(candidate.tier, candidate.model).get("max_tokens", provider.requested_output_limit)
for candidate in compared
)
if any(not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0 for limit in output_limits):
return None
limits: Final = tuple(limit for limit in output_limits if isinstance(limit, int))
arms: Final = await asyncio.gather(
*(
provider.predict(
candidate.deployment,
caller.api_key,
cache,
counter_for_model(candidate.model),
now=now,
)
for candidate in compared
)
)
costs: Final = tuple(
provider.cost(arm, min(config.cache_aware_routing_output_tokens, limit)) for arm, limit in zip(arms, limits)
)
original_cost: Final = costs[0]
if original_cost is None:
return None
finished_at: Final = time.time() if now is None else now
qualifying: Final = tuple(
CacheAwareChoice(candidate.model, candidate.tier, arm.deployment_id, original_cost, cost)
for candidate, arm, cost, limit in zip(authorized, arms[1:], costs[1:], limits[1:])
if cost is not None
and cost < original_cost
and arm.cache_state in ("warm", "partial")
and arm.evidence is not None
and arm.evidence.expires_at > finished_at
and arm.estimate is not None
and provider.fits(candidate.deployment, arm.estimate.tokens.total_tokens, limit)
)
return min(qualifying, key=lambda choice: choice.estimated_cost, default=None)
async def _choose_cached_model(
*,
router: Router,
config: ComplexityRouterConfig,
params_for_model: Callable[[str, str], Mapping[str, object]],
response: PreRoutingHookResponse | None,
request_kwargs: Mapping[str, object],
messages: Sequence[Mapping[str, object]] | None,
) -> CacheAwareChoice | None:
if not config.cache_aware_routing or response is None or response.routing_decision is None:
return None
if not eligible_models(config, response.routing_decision):
return None
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner
)
from litellm.router_strategy.complexity_router.context_compaction import compaction_pending
if (
proxy_server.llm_router is not router
or router.routing_plugins
or has_request_transforms()
or compaction_pending(request_kwargs)
or not supported_router_marker(router, response.routing_decision.get("router_model_name") or "", request_kwargs)
):
return None
metadata: Final = _MAPPING.validate_python(
request_kwargs.get(get_metadata_variable_name_from_kwargs(request_kwargs)) or MappingProxyType({})
)
caller: Final = metadata.get("user_api_key_auth")
if not isinstance(caller, UserAPIKeyAuth):
return None
settings: Final = _CallerSettings.model_validate(caller, from_attributes=True)
if settings.config:
return None
try:
incoming: Final = _ProxyRequest.model_validate(request_kwargs.get("proxy_server_request"))
except ValidationError:
return None
if any(
request_kwargs.get(key)
for key in (
"guardrails",
"cache_control_injection_points",
"api_key",
"api_base",
"extra_headers",
"prompt_id",
"mock_response",
"model_info",
"custom_llm_provider",
)
):
return None
body: Final = AnthropicCacheRouting.request_body(
incoming.url, incoming.headers, incoming.body, request_kwargs, messages
)
if body is None:
return None
limiter: Final = proxy_server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter")
if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3):
return None
def counter_for_model(model_name: str) -> TokenCounter:
async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None:
try:
async with limiter.request_capacity(caller, model_name, request_data=request_kwargs):
return await AnthropicCacheRouting.count_tokens(model, api_key, body)
except Exception: # noqa: BLE001 # an optional prediction denied capacity is an unavailable estimate
return None
return count
return await select_cached_model(
router=router,
config=config,
params_for_model=params_for_model,
response=response,
body=body,
request_kwargs=request_kwargs,
messages=messages,
caller=caller,
cache=proxy_server.proxy_logging_obj.internal_usage_cache.dual_cache,
counter_for_model=counter_for_model,
)
async def choose_cached_model(
*,
router: Router,
config: ComplexityRouterConfig,
params_for_model: Callable[[str, str], Mapping[str, object]],
response: PreRoutingHookResponse | None,
request_kwargs: Mapping[str, object],
messages: Sequence[Mapping[str, object]] | None,
) -> CacheAwareChoice | None:
if not config.cache_aware_routing:
return None
try:
return await asyncio.wait_for(
_choose_cached_model(
router=router,
config=config,
params_for_model=params_for_model,
response=response,
request_kwargs=request_kwargs,
messages=messages,
),
timeout=config.cache_aware_routing_timeout_ms / 1000,
)
except Exception: # noqa: BLE001 # cache prediction is optional and must preserve normal routing on failure
verbose_router_logger.debug("Cache-aware routing unavailable; keeping the classified model")
return None

View file

@ -0,0 +1,20 @@
from typing import Final
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.anthropic.cache_aware_routing import predict_arm
__all__: Final = ("has_request_transforms", "predict_arm")
def has_request_transforms() -> bool:
from litellm.proxy.hooks import PROXY_HOOKS
builtins: Final = frozenset(PROXY_HOOKS.values())
hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook")
callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger)
return any(
type(callback) not in builtins
and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks)
for callback in callbacks
)

View file

@ -1,4 +1,5 @@
from collections.abc import Mapping
from datetime import datetime, timezone
from math import isfinite
from typing import Final
@ -20,12 +21,13 @@ def _valid_price(value: object) -> bool:
return isinstance(value, (int, float)) and not isinstance(value, bool) and isfinite(value) and value >= 0
def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets) -> bool:
def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets, completion_tokens: int = 0) -> bool:
required: Final = (
("input_cost_per_token", True),
("cache_read_input_token_cost", tokens.cache_read_input_tokens > 0),
("cache_creation_input_token_cost", tokens.cache_creation_5m_input_tokens > 0),
("cache_creation_input_token_cost_above_1hr", tokens.cache_creation_1h_input_tokens > 0),
("output_cost_per_token", completion_tokens > 0),
)
if any(needed and not _valid_price(prices.get(key)) for key, needed in required):
return False
@ -36,7 +38,9 @@ def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets
)
def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> float | None:
def price_cache_tokens(
model: str, deployment_id: str, tokens: CacheTokenBuckets, completion_tokens: int = 0
) -> float | None:
try:
selected_model: Final = _select_model_name_for_cost_calc(
model=model,
@ -53,12 +57,12 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets
if price_entry is None:
return None
prices: Final = _PRICE_ENTRY.validate_python(price_entry)
if not _has_required_prices(prices, tokens):
if completion_tokens < 0 or not _has_required_prices(prices, tokens, completion_tokens):
return None
usage: Final = Usage(
prompt_tokens=tokens.total_tokens,
completion_tokens=0,
total_tokens=tokens.total_tokens,
completion_tokens=completion_tokens,
total_tokens=tokens.total_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=tokens.cache_read_input_tokens,
cache_creation_tokens=tokens.cache_creation_5m_input_tokens + tokens.cache_creation_1h_input_tokens,
@ -73,7 +77,7 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets
messages=[], # mutable-ok: Logging requires a list
stream=False,
call_type="completion",
start_time=None,
start_time=datetime.now(timezone.utc),
litellm_call_id="prompt-cache-prediction",
function_id="prompt-cache-prediction",
)
@ -85,7 +89,12 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets
router_model_id=deployment_id,
litellm_logging_obj=logging_obj,
)
cost: Final = logging_obj.cost_breakdown.get("input_cost") if logging_obj.cost_breakdown is not None else None
return cost if cost is not None and _valid_price(cost) else None
breakdown: Final = logging_obj.cost_breakdown
input_cost: Final = breakdown.get("input_cost") if breakdown is not None else None
output_cost: Final = breakdown.get("output_cost") if breakdown is not None else None
if input_cost is None or output_cost is None:
return None
cost: Final = input_cost + output_cost
return cost if _valid_price(cost) else None
except Exception: # noqa: BLE001 # the shared pricing owners raise plain Exception for unpriceable models
return None

View file

@ -1,4 +1,3 @@
import time
from collections.abc import Mapping
from types import MappingProxyType
from typing import Annotated, Final
@ -6,18 +5,10 @@ from typing import Annotated, Final
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import BaseModel, JsonValue, TypeAdapter
import litellm
from litellm._internal_context import current_billing_time, pinned_billing_time
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.anthropic.prompt_cache_prediction import (
PromptPrefix,
TokenCounter,
UnsupportedPredictionTarget,
cache_scope,
count_prompt_tokens,
parse_prompt,
resolve_prediction_target,
supported_prediction_headers,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -27,22 +18,16 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary
)
from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens
from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms, predict_arm
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner
)
from litellm.proxy.hooks.prompt_cache_prediction import lookup
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.management_endpoints.prompt_cache_prediction import (
CacheCostScenario,
CacheEvidence,
CachePredictionArm,
CachePredictionRequest,
CachePredictionResponse,
CacheTokenBuckets,
)
from litellm.types.router import Deployment
from litellm.utils import get_prompt_cache_min_tokens
router: Final = APIRouter()
_REQUEST_DATA: Final = TypeAdapter(Mapping[str, object])
@ -52,33 +37,6 @@ class _CallerSettings(BaseModel):
config: Mapping[str, object] | None = None
def has_request_transforms() -> bool:
from litellm.proxy.hooks import PROXY_HOOKS
builtins: Final = frozenset(PROXY_HOOKS.values())
hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook")
callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger)
return any(
type(callback) not in builtins
and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks)
for callback in callbacks
)
def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets:
return CacheTokenBuckets(
uncached_input_tokens=suffix_tokens,
cache_read_input_tokens=read_tokens,
cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0,
cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0,
)
def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None:
cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens)
return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None
def _capacity_counter(
limiter: _PROXY_MaxParallelRequestsHandler_v3,
caller: UserAPIKeyAuth,
@ -103,75 +61,6 @@ def _capacity_request_data(
return MappingProxyType(data)
async def predict_arm(
deployment: Deployment,
body: Mapping[str, JsonValue],
prefix: PromptPrefix,
caller_key_hash: str,
cache: DualCache,
token_counter: TokenCounter,
) -> CachePredictionArm:
deployment_id: Final = deployment.model_info.id or ""
params: Final = deployment.litellm_params
unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model)
if deployment.model_info.blocked:
return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"}))
target: Final = resolve_prediction_target(params)
if isinstance(target, UnsupportedPredictionTarget):
return unknown.model_copy(update=MappingProxyType({"reason": target.reason}))
model: Final = target.model
api_key: Final = target.api_key
total_count: Final = await token_counter(model, api_key, body)
prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body)
if total_count is None or prefix_count is None or total_count < prefix_count:
return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"}))
scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model)
observation: Final = await lookup(cache, scope, prefix)
exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint
cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count
if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable):
return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"}))
suffix: Final = total_count - cacheable
evidence: Final = (
CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at)
if observation is not None
else None
)
if cacheable < get_prompt_cache_min_tokens(params.model):
disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count))
if disabled is None:
return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"}))
return CachePredictionArm(
deployment_id=deployment_id,
model=model,
cache_state="disabled",
reason="below_cache_minimum",
estimate=disabled,
cold=disabled,
warm=disabled,
token_count_source="anthropic_count_tokens",
)
fresh: Final = observation is not None and observation.expires_at > time.time()
read: Final = observation.cached_tokens if fresh and observation is not None else 0
with pinned_billing_time(current_billing_time()):
cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds))
warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds))
estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds))
if cold is None or warm is None or estimate is None:
return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"}))
return CachePredictionArm(
deployment_id=deployment_id,
model=model,
cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown",
reason=None if fresh else "observation_expired" if observation else "no_compatible_observation",
estimate=estimate,
cold=cold,
warm=warm,
evidence=evidence,
token_count_source="anthropic_count_tokens",
)
@router.post(
"/cost/predict-cache",
tags=["Cost Tracking"], # mutable-ok: FastAPI requires a list for OpenAPI tags

View file

@ -1,6 +1,6 @@
# Complexity Router
A rule-based routing strategy that classifies requests by complexity and routes them to appropriate models - with zero API calls and sub-millisecond latency.
A routing strategy that classifies requests by complexity and routes them to appropriate models. The default rule-based classifier scores requests locally. Optional classifiers and cache-aware routing can make provider calls
## Overview
@ -68,6 +68,45 @@ still resolve to a deployment in `model_list`; this configuration does not creat
- abc
```
### Opt in to prompt-cache costs
Set `cache_aware_routing: true` to consider observed prompt-cache savings after classification. This is disabled by default. A warm model in the same or a higher tier can replace the classified model when its estimated input and output cost is strictly lower. Cache savings never lower the required tier
```yaml
model_list:
- model_name: smart-router
litellm_params:
model: auto_router/complexity_router
complexity_router_config:
cache_aware_routing: true
cache_aware_routing_output_tokens: 1024
cache_aware_routing_timeout_ms: 2000
context_compaction: false
tiers:
SIMPLE: haiku
COMPLEX: sonnet
- model_name: haiku
litellm_params:
model: anthropic/claude-haiku-4-5
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: sonnet
litellm_params:
model: anthropic/claude-sonnet-5
api_key: os.environ/ANTHROPIC_API_KEY
```
This first version supports the proxy's native `POST /v1/messages` endpoint with Anthropic, text and client tools, and one explicit message-content `cache_control` breakpoint. Each tier must name one model group with one deployment. The default v3 rate limiter must be enabled. It uses the same observations and token counting as `/cost/predict-cache`; it does not prewarm caches or enable provider caching on the application's behalf
The proxy must have observed a successful cache read or write for the candidate's matching prefix, under the same caller key, deployment, provider key and model. A fresh observation allows a cache discount; missing or expired evidence does not. Provider eviction can still turn an expected hit into a miss
The comparison includes uncached input, cache writes at the requested TTL, cache reads and expected output tokens. Set `cache_aware_routing_output_tokens` to your workload's expected response length; it defaults to 1024 and is capped separately by each model's effective output limit. With `max_tokens_from_tier_model: true` (the default), this is the model's known output ceiling; when disabled or unknown, the caller's `max_tokens` applies. The full effective output limit, together with the counted input, must fit the candidate's known limits. Custom deployment prices are respected
Prediction makes up to two token-count requests per compared model. These use rate and concurrency capacity and add latency. The default total timeout is two seconds; timeout, missing counts or prices, and prediction failures preserve the classified route. No provider count requests run when there is no warm eligible alternative
Session affinity, user-turn classification, adaptive routing, routing plugins, custom tier ladders, tier pools and per-tier parameter overrides keep their existing behavior without a cache adjustment. The same applies to unsupported providers or prompt shapes, beta headers, custom provider endpoints, request transforms, and pending context compaction. Disable context compaction as in the example so it cannot rewrite the predicted prompt. Alias markers should contain only routing configuration and rate, timeout or tag settings
When cache costs change the model, the routing decision reports `cause: prompt_cache_cost`. Its signals include the original model, classification cause and both estimated costs
### Capability forecasting
Set `classifier_type: capability` to use

View file

@ -4390,10 +4390,13 @@ class ComplexityRouter(CustomLogger):
resolved_messages=resolved_messages,
context_fit=context_fit,
)
cache_adjusted_response: Final = await self._apply_prompt_cache_routing(
routed_response, messages, request_kwargs, context_fit
)
response: Final = (
await self._gate_response_health(
await self._gate_response_modality(
routed_response, messages, resolved_messages, request_kwargs, context_fit
cache_adjusted_response, messages, resolved_messages, request_kwargs, context_fit
),
messages,
input,
@ -4401,7 +4404,7 @@ class ComplexityRouter(CustomLogger):
request_kwargs,
context_fit,
)
if routed_response is not None
if cache_adjusted_response is not None
else None
)
# Sentinel presence, not the plan_mode cause, gates the pin write: a plan-mode turn
@ -4425,6 +4428,53 @@ class ComplexityRouter(CustomLogger):
)
return self._with_session_deployment_affinity(response)
async def _apply_prompt_cache_routing(
self,
response: PreRoutingHookResponse | None,
messages: Sequence[Mapping[str, object]] | None,
request_kwargs: Mapping[str, object],
context_fit: _RequestContextFit,
) -> PreRoutingHookResponse | None:
if not self.config.cache_aware_routing or response is None or response.routing_decision is None:
return response
from litellm.proxy.common_utils.cache_aware_routing import choose_cached_model
choice: Final = await choose_cached_model(
router=self.litellm_router_instance,
config=self.config,
params_for_model=self._litellm_params_for_model,
response=response,
request_kwargs=request_kwargs,
messages=messages,
)
if choice is None or not context_fit.accepts(choice.model):
return response
params: Final = self._litellm_params_for_model(choice.tier, choice.model)
decision: Final[StandardLoggingRoutingDecision] = {
**response.routing_decision,
"routed_model": choice.model,
"cause": "prompt_cache_cost",
"tier": choice.tier,
"tier_label": (self.config.tier_labels or {}).get(choice.tier, choice.tier),
"tier_litellm_params": params,
"signals": (
*(response.routing_decision.get("signals") or ()),
f"cache-aware:classified-model={response.model}",
f"cache-aware:classification-cause={response.routing_decision.get('cause')}",
f"cache-aware:estimated-cost={choice.estimated_cost:.8f};original-cost={choice.original_cost:.8f}",
),
}
verbose_router_logger.info(
"ComplexityRouter: cache-aware choice model=%s original=%s estimated_cost=%s original_cost=%s",
choice.model,
response.model,
choice.estimated_cost,
choice.original_cost,
)
return response.model_copy(
update=MappingProxyType({"model": choice.model, "litellm_params": params, "routing_decision": decision})
)
async def _classify_and_route(
self,
model: str,

View file

@ -1422,6 +1422,25 @@ class ComplexityRouterConfig(BaseModel):
),
)
cache_aware_routing: bool = Field(
default=False,
description=(
"Opt in to comparing prompt-cache costs after classification. On supported native Anthropic proxy requests, "
"an already warm model in the same or a higher tier may replace the classified model when its estimated "
"input and output cost is lower. Unsupported requests and unavailable estimates keep ordinary routing."
),
)
cache_aware_routing_output_tokens: int = Field(
default=1024,
ge=0,
description="Expected output tokens used in cache-aware cost comparisons; capped by each model's effective output limit.",
)
cache_aware_routing_timeout_ms: int = Field(
default=2000,
gt=0,
description="Total time budget for cache-aware predictions; expiry preserves the original routing decision.",
)
# Session affinity: pin the first turn's routed model for the rest of the session
session_affinity: bool = Field(
default=False,

View file

@ -2970,6 +2970,7 @@ class StandardLoggingRoutingDecisionTierBoundaries(TypedDict):
RoutingDecisionCause = Literal[
"prompt_cache_cost",
"heuristic_scorer",
"heuristic_v2",
# The scorer found 2+ reasoning markers and forced REASONING regardless of score.

View file

@ -0,0 +1,588 @@
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Final
import pytest
from pydantic import JsonValue
from litellm import Router
from litellm.caching.dual_cache import DualCache
from litellm.llms.anthropic.prompt_cache_prediction import TokenCounter, cache_scope, parse_prompt
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.cache_aware_routing import (
CacheAwareChoice,
choose_cached_model,
eligible_models,
select_cached_model,
)
from litellm.proxy.hooks.prompt_cache_prediction import CacheObservation, _cache_key
from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
from litellm.types.router import PreRoutingHookResponse
_CALLER: Final = "test-cache-aware-caller"
_PROVIDER_KEY: Final = "test-cache-aware-provider"
_NOW: Final = 1000.0
@dataclass(frozen=True, slots=True)
class _Counts:
total: int | None = 51000
prefix: int | None = 50000
async def __call__(self, model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None:
return self.total if "max_tokens" in body else self.prefix
def _counter_for_model(model: str) -> TokenCounter:
return _Counts()
def _forbidden_counter(model: str) -> TokenCounter:
raise AssertionError("No provider counts should run without a warm eligible alternative")
def _body(text: str = "Stable cached context") -> dict[str, JsonValue]:
return {
"max_tokens": 20000,
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": text, "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": "What is 2 + 2?"},
],
}
],
}
def _router(
strong_output_rate: float = 0.000015, *, free: bool = False, cheap_limit: int = 30000, strong_limit: int = 30000
) -> Router:
return Router(
model_list=[
{
"model_name": "cheap",
"litellm_params": {
"model": "anthropic/claude-haiku-4-5",
"api_key": _PROVIDER_KEY,
"input_cost_per_token": 0 if free else 0.000001,
"output_cost_per_token": 0 if free else 0.000004,
"cache_read_input_token_cost": 0 if free else 0.0000001,
"cache_creation_input_token_cost": 0 if free else 0.00000125,
},
"model_info": {"id": "test-cache-cheap", "max_input_tokens": 100000, "max_output_tokens": cheap_limit},
},
{
"model_name": "strong",
"litellm_params": {
"model": "anthropic/claude-sonnet-5",
"api_key": _PROVIDER_KEY,
"input_cost_per_token": 0 if free else 0.000003,
"output_cost_per_token": 0 if free else strong_output_rate,
"cache_read_input_token_cost": 0 if free else 0.0000003,
"cache_creation_input_token_cost": 0 if free else 0.00000375,
},
"model_info": {
"id": "test-cache-strong",
"max_input_tokens": 100000,
"max_output_tokens": strong_limit,
},
},
]
)
def _config(**overrides: object) -> ComplexityRouterConfig:
return ComplexityRouterConfig.model_validate(
{
"tiers": {"SIMPLE": "cheap", "COMPLEX": "strong"},
"cache_aware_routing": True,
**overrides,
}
)
def _response(tier: str = "SIMPLE", model: str = "cheap") -> PreRoutingHookResponse:
return PreRoutingHookResponse(
model=model,
messages=None,
routing_decision={
"router_model_name": "smart",
"router_type": "complexity",
"routed_model": model,
"tier": tier,
"cause": "heuristic_scorer",
},
)
async def _observed(cache: DualCache, *, caller: str = _CALLER, expires_at: float = 1290.0) -> None:
prefix: Final = parse_prompt(_body())
assert prefix is not None
scope: Final = cache_scope(caller, "test-cache-strong", _PROVIDER_KEY, "claude-sonnet-5")
observation: Final = CacheObservation(
fingerprint=prefix.fingerprint,
cached_tokens=50000,
observed_at=990.0,
expires_at=expires_at,
)
await cache.async_set_cache(_cache_key(scope, prefix.fingerprint), observation.model_dump_json(), ttl=3600)
async def _select(
*,
router: Router,
config: ComplexityRouterConfig,
response: PreRoutingHookResponse,
body: Mapping[str, JsonValue],
request_kwargs: Mapping[str, object],
messages: Sequence[Mapping[str, object]] | None,
caller: UserAPIKeyAuth,
cache: DualCache,
counter_for_model: Callable[[str], TokenCounter],
now: float,
) -> CacheAwareChoice | None:
complexity: Final = ComplexityRouter("smart", router, config.model_dump())
return await select_cached_model(
router=router,
config=config,
params_for_model=complexity._litellm_params_for_model,
response=response,
body=body,
request_kwargs=request_kwargs,
messages=messages,
caller=caller,
cache=cache,
counter_for_model=counter_for_model,
now=now,
)
@pytest.mark.asyncio
async def test_warm_stronger_model_wins_after_counting_input_and_output_cost() -> None:
cache: Final = DualCache()
await _observed(cache)
choice: Final = await _select(
router=_router(),
config=_config(),
response=_response(),
body=_body(),
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=_counter_for_model,
now=_NOW,
)
assert choice is not None
assert (choice.model, choice.tier, choice.deployment_id) == ("strong", "COMPLEX", "test-cache-strong")
assert choice.original_cost == pytest.approx(50000 * 0.00000125 + 1000 * 0.000001 + 1024 * 0.000004)
assert choice.estimated_cost == pytest.approx(50000 * 0.0000003 + 1000 * 0.000003 + 1024 * 0.000015)
@pytest.mark.asyncio
async def test_output_price_can_outweigh_the_cache_saving() -> None:
cache: Final = DualCache()
await _observed(cache)
choice: Final = await _select(
router=_router(strong_output_rate=0.001),
config=_config(),
response=_response(),
body=_body(),
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=_counter_for_model,
now=_NOW,
)
assert choice is None
@pytest.mark.asyncio
@pytest.mark.parametrize("case", ["missing", "expired", "different_caller", "changed_prefix", "unauthorized"])
async def test_no_cache_discount_without_fresh_authorized_matching_evidence(case: str) -> None:
cache: Final = DualCache()
if case != "missing":
await _observed(
cache,
caller="someone-else" if case == "different_caller" else _CALLER,
expires_at=999.0 if case == "expired" else 1290.0,
)
choice: Final = await _select(
router=_router(),
config=_config(),
response=_response(),
body=_body("Changed context" if case == "changed_prefix" else "Stable cached context"),
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap"] if case == "unauthorized" else ["cheap", "strong"]),
cache=cache,
counter_for_model=_forbidden_counter,
now=_NOW,
)
assert choice is None
@pytest.mark.asyncio
async def test_disabled_setting_does_not_access_prediction_services() -> None:
config: Final = ComplexityRouterConfig(tiers={"SIMPLE": "cheap"})
assert config.cache_aware_routing is False
assert (
await choose_cached_model(
router=_router(),
config=config,
params_for_model=ComplexityRouter("smart", _router(), config.model_dump())._litellm_params_for_model,
response=_response(),
request_kwargs={},
messages=None,
)
is None
)
def test_cache_prices_cannot_add_a_model_below_the_classified_tier() -> None:
response: Final = _response("COMPLEX", "strong")
assert response.routing_decision is not None
assert eligible_models(_config(), response.routing_decision) == (("COMPLEX", "strong"),)
@pytest.mark.parametrize(
"overrides", [{"adaptive": True}, {"session_affinity": True}, {"classification_mode": "user_turn"}]
)
def test_existing_pinned_or_adaptive_policies_are_preserved(overrides: Mapping[str, object]) -> None:
response: Final = _response()
assert response.routing_decision is not None
assert eligible_models(_config(**overrides), response.routing_decision) == ()
@pytest.mark.parametrize("total,prefix", [(None, 50000), (51000, None), (1000, 50000)])
@pytest.mark.asyncio
async def test_unavailable_or_inconsistent_counts_keep_the_classified_model(
total: int | None, prefix: int | None
) -> None:
cache: Final = DualCache()
await _observed(cache)
assert (
await _select(
router=_router(),
config=_config(),
response=_response(),
body=_body(),
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=lambda _: _Counts(total, prefix),
now=_NOW,
)
is None
)
@pytest.mark.asyncio
async def test_output_estimate_is_capped_by_the_requested_limit() -> None:
cache: Final = DualCache()
await _observed(cache)
choice: Final = await _select(
router=_router(),
config=_config(cache_aware_routing_output_tokens=100000, max_tokens_from_tier_model=False),
response=_response(),
body={**_body(), "max_tokens": 1},
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=_counter_for_model,
now=_NOW,
)
assert choice is not None
assert choice.estimated_cost == pytest.approx(50000 * 0.0000003 + 1000 * 0.000003 + 0.000015)
@pytest.mark.asyncio
async def test_warm_model_that_cannot_fit_the_request_is_not_selected() -> None:
cache: Final = DualCache()
await _observed(cache)
assert (
await _select(
router=_router(),
config=_config(max_tokens_from_tier_model=False),
response=_response(),
body={**_body(), "max_tokens": 100000000},
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=_counter_for_model,
now=_NOW,
)
is None
)
def test_repeated_model_in_multiple_tiers_is_only_considered_once() -> None:
decision: Final = _response().routing_decision
assert decision is not None
assert eligible_models(_config(tiers={"SIMPLE": "cheap", "MEDIUM": "strong", "COMPLEX": "strong"}), decision) == (
("SIMPLE", "cheap"),
("MEDIUM", "strong"),
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"enabled,behavior,expected",
[
(False, "success", "cheap"),
(True, "success", "strong"),
(True, "tier_cost", "cheap"),
(True, "tier_context", "cheap"),
(True, "error", "cheap"),
(True, "deadline", "cheap"),
(True, "cancel", None),
(True, "transformed", "cheap"),
(True, "unsupported_shape", "cheap"),
(True, "custom_endpoint", "cheap"),
(True, "compaction", "cheap"),
(True, "guardrail", "cheap"),
],
)
async def test_router_applies_opt_in_and_preserves_failure_semantics(
monkeypatch: pytest.MonkeyPatch, enabled: bool, behavior: str, expected: str | None
) -> None:
import asyncio
import json
import httpx
import litellm
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from litellm.proxy.utils import ProxyLogging
from litellm.router_strategy.complexity_router.context_compaction import initialize_compaction_state
config: Final = _config(
cache_aware_routing=enabled, cache_aware_routing_timeout_ms=1 if behavior == "deadline" else 2000
)
models: Final = _router(
strong_output_rate=0.000048 if behavior == "tier_cost" else 0.000015,
cheap_limit=50 if behavior == "tier_cost" else 30000,
strong_limit=60000 if behavior == "tier_context" else 30000,
).get_model_list()
assert models is not None
router: Final = Router(
model_list=[
*models,
{
"model_name": "smart",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": config.model_dump(),
**({"temperature": 0.1} if behavior == "transformed" else {}),
},
},
]
)
logging: Final = ProxyLogging(UserApiKeyCache())
logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3(
logging.internal_usage_cache
)
await _observed(logging.internal_usage_cache.dual_cache, expires_at=1e100)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging)
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
requests: Final = asyncio.Queue[httpx.Request]()
async def count(request: httpx.Request) -> httpx.Response:
requests.put_nowait(request)
assert enabled and behavior not in ("transformed", "unsupported_shape", "custom_endpoint")
if behavior == "error":
return httpx.Response(503, json={"error": "Provider unavailable"})
if behavior == "cancel":
raise asyncio.CancelledError()
if behavior == "deadline":
await asyncio.Future()
payload: Final = json.loads(request.content)
assert request.url == "https://api.anthropic.com/v1/messages/count_tokens"
assert request.headers["x-api-key"] == _PROVIDER_KEY
return httpx.Response(200, json={"input_tokens": 51000 if "What is 2 + 2?" in str(payload) else 50000})
async with httpx.AsyncClient(transport=httpx.MockTransport(count)) as client:
handler: Final = AsyncHTTPHandler()
await handler.client.aclose()
handler.client = client
litellm.in_memory_llm_clients_cache.set_cache("async_httpx_clientanthropic", handler)
body: Final = {
**_body(),
**({"max_tokens": 1000} if behavior in ("tier_cost", "tier_context") else {}),
**({"thinking": {"type": "enabled", "budget_tokens": 10000}} if behavior == "unsupported_shape" else {}),
}
kwargs: Final = {
"litellm_metadata": {
"user_api_key_auth": UserAPIKeyAuth(api_key=_CALLER, models=["smart", "cheap", "strong"])
},
"proxy_server_request": {"url": "http://localhost/v1/messages", "body": body, "headers": {}},
**({"api_base": "https://custom.example"} if behavior == "custom_endpoint" else {}),
**(
{"_context_compaction_state": initialize_compaction_state({}, "messages")}
if behavior == "compaction"
else {}
),
**({"guardrails": ["test-guardrail"]} if behavior == "guardrail" else {}),
}
if expected is None:
with pytest.raises(asyncio.CancelledError):
await router.async_pre_routing_hook(model="smart", request_kwargs=kwargs, messages=body["messages"])
return
response: Final = await router.async_pre_routing_hook(
model="smart", request_kwargs=kwargs, messages=body["messages"]
)
assert response is not None
assert response.model == expected
if not enabled or behavior in (
"transformed",
"unsupported_shape",
"custom_endpoint",
"compaction",
"guardrail",
):
assert requests.qsize() == 0
if enabled and behavior == "success":
assert requests.qsize() == 4
assert response.routing_decision is not None
assert response.routing_decision["cause"] == (
"prompt_cache_cost" if expected == "strong" else "heuristic_scorer"
)
@pytest.mark.asyncio
async def test_equal_costs_keep_the_classified_model() -> None:
cache: Final = DualCache()
await _observed(cache)
assert (
await _select(
router=_router(free=True),
config=_config(),
response=_response(),
body=_body(),
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=_counter_for_model,
now=_NOW,
)
is None
)
@pytest.mark.parametrize(
"cause", ["llm_v2_classifier", "capability_classifier", "heuristic_first_short_circuit", "hybrid_short_circuit"]
)
def test_successful_classifiers_can_consider_cache_costs(cause: str) -> None:
response: Final = PreRoutingHookResponse.model_validate(
{
"model": "cheap",
"messages": None,
"routing_decision": {"tier": "SIMPLE", "cause": cause},
}
)
assert response.routing_decision is not None
assert eligible_models(_config(), response.routing_decision) == (("SIMPLE", "cheap"), ("COMPLEX", "strong"))
@pytest.mark.asyncio
@pytest.mark.parametrize(
"cheap_limit,strong_limit,requested,from_tier,output_rate,expected_limits",
[
(50, 1000, 1000, True, 0.000048, None),
(30000, 30000, 1, True, 0.000049, None),
(30000, 60000, 1, True, 0.000015, None),
(50, 100, 20000, True, 0.000048, (50, 100)),
(30000, 30000, 1, False, 0.000048, (1, 1)),
],
)
async def test_each_candidate_uses_its_effective_routed_output_limit(
cheap_limit: int,
strong_limit: int,
requested: int,
from_tier: bool,
output_rate: float,
expected_limits: tuple[int, int] | None,
) -> None:
cache: Final = DualCache()
await _observed(cache)
choice: Final = await _select(
router=_router(strong_output_rate=output_rate, cheap_limit=cheap_limit, strong_limit=strong_limit),
config=_config(max_tokens_from_tier_model=from_tier),
response=_response(),
body={**_body(), "max_tokens": requested},
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]),
cache=cache,
counter_for_model=_counter_for_model,
now=_NOW,
)
if expected_limits is None:
assert choice is None
return
assert choice is not None
assert choice.original_cost == pytest.approx(50000 * 0.00000125 + 1000 * 0.000001 + expected_limits[0] * 0.000004)
assert choice.estimated_cost == pytest.approx(
50000 * 0.0000003 + 1000 * 0.000003 + expected_limits[1] * output_rate
)
@pytest.mark.asyncio
@pytest.mark.parametrize("warm,authorized", [(False, True), (True, True), (True, False)])
async def test_authorization_only_runs_for_original_and_warm_alternatives_before_provider_counts(
monkeypatch: pytest.MonkeyPatch, warm: bool, authorized: bool
) -> None:
from unittest.mock import AsyncMock
from litellm.proxy.common_utils import cache_aware_routing
cache: Final = DualCache()
if warm:
await _observed(cache)
models: Final = _router().get_model_list()
assert models is not None
router: Final = Router(
model_list=[
*models,
{**models[0], "model_name": "cold", "model_info": {"id": "test-cache-cold"}},
]
)
authorization: Final = AsyncMock(wraps=cache_aware_routing.can_key_call_resolved_model)
monkeypatch.setattr(cache_aware_routing, "can_key_call_resolved_model", authorization)
def counter_for_model(model: str) -> TokenCounter:
assert warm and authorized
assert authorization.await_count == 2
return _Counts()
choice: Final = await _select(
router=router,
config=_config(tiers={"SIMPLE": "cheap", "MEDIUM": "cold", "COMPLEX": "strong"}),
response=_response(),
body=_body(),
request_kwargs={},
messages=None,
caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong", "cold"] if authorized else ["cheap", "cold"]),
cache=cache,
counter_for_model=counter_for_model,
now=_NOW,
)
assert (choice is not None) == (warm and authorized)
assert tuple(call.kwargs["model"] for call in authorization.await_args_list) == (
("cheap", "strong") if warm else ()
)

View file

@ -39483,6 +39483,24 @@ export interface components {
adaptive_eligible: "all" | "classified_tier";
/** @description Quality vs cost weights for adaptive selection (used when adaptive=True) */
adaptive_weights?: components["schemas"]["AdaptiveRouterWeights"];
/**
* Cache Aware Routing
* @description Opt in to comparing prompt-cache costs after classification. On supported native Anthropic proxy requests, an already warm model in the same or a higher tier may replace the classified model when its estimated input and output cost is lower. Unsupported requests and unavailable estimates keep ordinary routing.
* @default false
*/
cache_aware_routing: boolean;
/**
* Cache Aware Routing Output Tokens
* @description Expected output tokens used in cache-aware cost comparisons; capped by each model's effective output limit.
* @default 1024
*/
cache_aware_routing_output_tokens: number;
/**
* Cache Aware Routing Timeout Ms
* @description Total time budget for cache-aware predictions; expiry preserves the original routing decision.
* @default 2000
*/
cache_aware_routing_timeout_ms: number;
/** @description Probability threshold policy required when classifier_type is 'capability'. The classifier forecasts p_solve for efficient_tier, adjusts base_threshold using the capability-card boundary, and otherwise routes to capable_tier */
capability_classifier_config?: components["schemas"]["CapabilityClassifierConfig"] | null;
/**
@ -42605,7 +42623,7 @@ export interface components {
* Cause
* @enum {string}
*/
cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "jev_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
cause?: "prompt_cache_cost" | "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "jev_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
/** Classifier Calibrated Capable P Solve */
classifier_calibrated_capable_p_solve?: number;
/** Classifier Calibrated Efficient P Solve */