fix(router): record savings metadata for adaptive strategies

This commit is contained in:
mikemikimike 2026-08-30 09:26:31 +08:00
parent 10631eb834
commit daaf871138
7 changed files with 104 additions and 2 deletions

View file

@ -8637,6 +8637,7 @@ class Router:
config=config,
model_to_prefs=model_to_prefs,
model_to_cost=model_to_cost,
litellm_router_instance=self,
)
self._register_pre_routing_strategy(
registry=self.adaptive_routers,

View file

@ -17,7 +17,7 @@ import asyncio
import time
from collections import OrderedDict
from dataclasses import asdict, dataclass
from typing import Any, Final, cast
from typing import TYPE_CHECKING, Any, Final, cast
from litellm._logging import verbose_router_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -50,8 +50,15 @@ from litellm.router_strategy.adaptive_router.signals import (
from litellm.router_strategy.adaptive_router.update_queue import (
AdaptiveRouterUpdateQueue,
)
from litellm.router_strategy.savings_baseline import (
conversation_is_continuing,
resolve_baseline,
)
from litellm.types.utils import StandardLoggingRoutingDecision
if TYPE_CHECKING:
from litellm.router import Router
# Sweep session-state cache when it exceeds this many live entries. Expired
# entries are dropped in bulk; amortizes to O(1) per insert.
_SESSION_STATE_SWEEP_THRESHOLD: Final[int] = 1024
@ -91,11 +98,13 @@ class AdaptiveRouter:
config: AdaptiveRouterConfig,
model_to_prefs: dict[str, AdaptiveRouterPreferences],
model_to_cost: dict[str, float],
litellm_router_instance: Router | None = None,
) -> None:
self.router_name = router_name
self.config = config
self.model_to_prefs = model_to_prefs
self.model_to_cost = model_to_cost
self.litellm_router_instance = litellm_router_instance
self.queue = AdaptiveRouterUpdateQueue()
self._cells: dict[tuple[RequestType, str], BanditCell] = {}
@ -203,9 +212,22 @@ class AdaptiveRouter:
routed_model=chosen_model,
cause="bandit",
request_type=request_type.value,
conversation_continuing=conversation_is_continuing(messages),
**self._savings_baseline_fields(),
),
)
def _savings_baseline_fields(self) -> dict[str, str]:
if self.litellm_router_instance is None:
return {}
baseline = resolve_baseline(self.litellm_router_instance, self.config.available_models)
if baseline is None:
return {}
fields = {"savings_baseline_model": baseline.model}
if baseline.deployment_id is not None:
fields["savings_baseline_deployment_id"] = baseline.deployment_id
return fields
# ---- Pick model ------------------------------------------------------
async def pick_model(

View file

@ -1751,6 +1751,7 @@ class ComplexityRouter(CustomLogger):
),
model_to_prefs=model_to_prefs,
model_to_cost=model_to_cost,
litellm_router_instance=self.litellm_router_instance,
)
self._adaptive_chosen_model_key = ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY
return self.adaptive_router

View file

@ -23,6 +23,10 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.router_strategy.complexity_router.complexity_router import (
ComplexityRouter,
)
from litellm.router_strategy.savings_baseline import (
conversation_is_continuing,
resolve_baseline,
)
from litellm.types.utils import StandardLoggingRoutingDecision
from .config import QualityRouterConfig, RoutingPreferences
@ -318,6 +322,15 @@ class QualityRouter(CustomLogger):
if isinstance(metadata, dict):
metadata["quality_router_decision"] = decision
def _savings_fields(self) -> dict[str, str]:
baseline = resolve_baseline(self.litellm_router_instance, self.config.available_models)
if baseline is None:
return {}
fields = {"savings_baseline_model": baseline.model}
if baseline.deployment_id is not None:
fields["savings_baseline_deployment_id"] = baseline.deployment_id
return fields
async def async_pre_routing_hook(
self,
model: str,
@ -364,6 +377,8 @@ class QualityRouter(CustomLogger):
router_type="quality",
routed_model=self.config.default_model,
cause="default_fallback",
conversation_continuing=conversation_is_continuing(messages),
**self._savings_fields(),
),
)
@ -387,6 +402,8 @@ class QualityRouter(CustomLogger):
"matched_keyword": matched_keyword,
"quality_tier": self._model_quality.get(routed_model),
"complexity_tier": None,
"conversation_continuing": conversation_is_continuing(messages),
**self._savings_fields(),
},
)
routing_decision: Final = StandardLoggingRoutingDecision(
@ -395,6 +412,8 @@ class QualityRouter(CustomLogger):
routed_model=routed_model,
cause="keyword",
matched_keyword=matched_keyword,
conversation_continuing=conversation_is_continuing(messages),
**self._savings_fields(),
)
keyword_quality_tier: Final = self._model_quality.get(routed_model)
if keyword_quality_tier is not None:
@ -433,6 +452,8 @@ class QualityRouter(CustomLogger):
"matched_keyword": None,
"quality_tier": int(quality_tier),
"complexity_tier": complexity_name,
"conversation_continuing": conversation_is_continuing(messages),
**self._savings_fields(),
},
)
@ -447,5 +468,7 @@ class QualityRouter(CustomLogger):
tier=str(int(quality_tier)),
score=score,
signals=list(signals),
conversation_continuing=conversation_is_continuing(messages),
**self._savings_fields(),
),
)

View file

@ -14,7 +14,7 @@ bare string with no provider beside them; an operator who writes ``deepseek-r1``
Azure would otherwise be priced against whoever else owns that name.
"""
from collections.abc import Iterable
from collections.abc import Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Final, NamedTuple
from litellm._logging import verbose_router_logger
@ -45,6 +45,11 @@ class Baseline(NamedTuple):
deployment_id: str | None = None
def conversation_is_continuing(messages: Sequence[Mapping[str, object]] | None) -> bool:
"""Return whether the request contains evidence of an earlier assistant turn."""
return any(message.get("role") == "assistant" for message in messages or ())
def canonical_model(model: str, custom_llm_provider: str | None = None) -> str | None:
"""``provider/model``, or ``None`` when the pair names no known provider.

View file

@ -12,6 +12,7 @@ from unittest.mock import AsyncMock
import pytest
from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter
from litellm.router_strategy.savings_baseline import Baseline
from litellm.types.router import (
AdaptiveRouterConfig,
PreRoutingHookResponse,
@ -43,6 +44,32 @@ async def test_returns_pre_routing_hook_response_with_chosen_model():
assert response.model == "smart"
@pytest.mark.asyncio
async def test_records_savings_baseline_and_conversation_shape(monkeypatch):
router = _make_router()
router.litellm_router_instance = object()
router.pick_model = AsyncMock(return_value="smart") # type: ignore[method-assign]
monkeypatch.setattr(
"litellm.router_strategy.adaptive_router.adaptive_router.resolve_baseline",
lambda router, models: Baseline("openai/smart", "deployment-id"),
)
response = await router.async_pre_routing_hook(
model="smart-cheap-router",
request_kwargs={},
messages=[
{"role": "user", "content": "first"},
{"role": "assistant", "content": "answer"},
{"role": "user", "content": "next"},
],
)
assert response is not None
assert response.routing_decision["conversation_continuing"] is True
assert response.routing_decision["savings_baseline_model"] == "openai/smart"
assert response.routing_decision["savings_baseline_deployment_id"] == "deployment-id"
@pytest.mark.asyncio
async def test_classifies_last_user_message_for_request_type():
r = _make_router()

View file

@ -19,6 +19,7 @@ from litellm.router_strategy.quality_router.config import (
DEFAULT_COMPLEXITY_TO_QUALITY,
)
from litellm.router_strategy.quality_router.quality_router import QualityRouter
from litellm.router_strategy.savings_baseline import Baseline
def _make_model_list(spec: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
@ -822,6 +823,28 @@ class TestKeywordOverride:
class TestDecisionMetadata:
@pytest.mark.asyncio
async def test_decision_includes_savings_baseline_and_conversation_shape(self, quality_router, monkeypatch):
monkeypatch.setattr(
"litellm.router_strategy.quality_router.quality_router.resolve_baseline",
lambda router, models: Baseline("openai/opus-next", "id-opus-next"),
)
request_kwargs: Dict[str, Any] = {}
response = await quality_router.async_pre_routing_hook(
model="quality-router-test",
request_kwargs=request_kwargs,
messages=[
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
{"role": "user", "content": "continue"},
],
)
assert response is not None
assert response.routing_decision["conversation_continuing"] is True
assert response.routing_decision["savings_baseline_model"] == "openai/opus-next"
assert response.routing_decision["savings_baseline_deployment_id"] == "id-opus-next"
@pytest.mark.asyncio
async def test_hook_stashes_decision_in_request_kwargs_metadata(
self, quality_router