mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
501ef23f4a
commit
de06c93767
12 changed files with 1337 additions and 124 deletions
|
|
@ -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
|
||||
|
|
|
|||
236
litellm/llms/anthropic/cache_aware_routing.py
Normal file
236
litellm/llms/anthropic/cache_aware_routing.py
Normal 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",
|
||||
)
|
||||
343
litellm/proxy/common_utils/cache_aware_routing.py
Normal file
343
litellm/proxy/common_utils/cache_aware_routing.py
Normal 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
|
||||
20
litellm/proxy/common_utils/prompt_cache_prediction.py
Normal file
20
litellm/proxy/common_utils/prompt_cache_prediction.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
588
tests/unit/proxy/common_utils/test_cache_aware_routing.py
Normal file
588
tests/unit/proxy/common_utils/test_cache_aware_routing.py
Normal 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 ()
|
||||
)
|
||||
20
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
20
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue