This commit is contained in:
mikemikimike 2026-09-23 14:51:13 +00:00 • committed by GitHub
commit 84ac3d5ea9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 213 additions and 28 deletions

View file

@ -9279,6 +9279,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

@ -18,7 +18,7 @@ import time
from collections import OrderedDict
from collections.abc import Mapping
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 (
@ -51,8 +51,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
@ -92,11 +99,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] = {}
@ -205,18 +214,35 @@ class AdaptiveRouter:
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = chosen_model
routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.router_name,
router_type="adaptive",
routed_model=chosen_model,
cause="bandit",
request_type=request_type.value,
conversation_continuing=conversation_is_continuing(messages),
)
routing_decision.update(self._savings_baseline_fields())
return PreRoutingHookResponse(
model=chosen_model,
messages=messages,
routing_decision=StandardLoggingRoutingDecision(
router_model_name=self.router_name,
router_type="adaptive",
routed_model=chosen_model,
cause="bandit",
request_type=request_type.value,
),
routing_decision=routing_decision,
)
def _savings_baseline_fields(self) -> StandardLoggingRoutingDecision:
if self.litellm_router_instance is None:
return {} # mutable-ok: immutable empty result for absent router
baseline = resolve_baseline(self.litellm_router_instance, self.config.available_models)
if baseline is None:
return {} # mutable-ok: immutable empty result for unresolved baseline
fields: Final[StandardLoggingRoutingDecision] = {
"savings_baseline_model": baseline.model,
**(
{"savings_baseline_deployment_id": baseline.deployment_id} if baseline.deployment_id is not None else {}
),
}
return fields
# ---- Pick model ------------------------------------------------------
async def pick_model(

View file

@ -2945,6 +2945,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

@ -24,6 +24,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
@ -319,6 +323,18 @@ class QualityRouter(CustomLogger):
if isinstance(metadata, dict):
metadata["quality_router_decision"] = decision
def _savings_fields(self) -> StandardLoggingRoutingDecision:
baseline = resolve_baseline(self.litellm_router_instance, self.config.available_models)
if baseline is None:
return {} # mutable-ok: immutable empty result for unresolved baseline
fields: Final[StandardLoggingRoutingDecision] = {
"savings_baseline_model": baseline.model,
**(
{"savings_baseline_deployment_id": baseline.deployment_id} if baseline.deployment_id is not None else {}
),
}
return fields
async def async_pre_routing_hook(
self,
model: str,
@ -334,6 +350,9 @@ class QualityRouter(CustomLogger):
verbose_router_logger.debug("QualityRouter: No messages provided, skipping routing")
return None
conversation_continuing: Final = conversation_is_continuing(messages)
savings_fields: Final = self._savings_fields()
# Extract last user message and last system prompt — same rules as
# ComplexityRouter.async_pre_routing_hook.
user_message: str | None = None
@ -357,15 +376,18 @@ class QualityRouter(CustomLogger):
verbose_router_logger.debug("QualityRouter: No user message found, routing to default model")
if not self.config.default_model:
raise ValueError("QualityRouter: no user message and no default_model configured")
default_routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=self.config.default_model,
cause="default_fallback",
conversation_continuing=conversation_continuing,
)
default_routing_decision.update(savings_fields)
return PreRoutingHookResponse(
model=self.config.default_model,
messages=messages,
routing_decision=StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=self.config.default_model,
cause="default_fallback",
),
routing_decision=default_routing_decision,
)
# Try keyword override first — it short-circuits complexity classification.
@ -388,22 +410,25 @@ class QualityRouter(CustomLogger):
"matched_keyword": matched_keyword,
"quality_tier": self._model_quality.get(routed_model),
"complexity_tier": None,
"conversation_continuing": conversation_continuing,
},
)
routing_decision: Final = StandardLoggingRoutingDecision(
keyword_routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=routed_model,
cause="keyword",
matched_keyword=matched_keyword,
conversation_continuing=conversation_continuing,
)
keyword_routing_decision.update(savings_fields)
keyword_quality_tier: Final = self._model_quality.get(routed_model)
if keyword_quality_tier is not None:
routing_decision["tier"] = str(keyword_quality_tier)
keyword_routing_decision["tier"] = str(keyword_quality_tier)
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
routing_decision=routing_decision,
routing_decision=keyword_routing_decision,
)
# No keyword match → complexity classification flow.
@ -434,19 +459,23 @@ class QualityRouter(CustomLogger):
"matched_keyword": None,
"quality_tier": int(quality_tier),
"complexity_tier": complexity_name,
"conversation_continuing": conversation_continuing,
},
)
quality_routing_decision: Final = StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=routed_model,
cause="quality_tier",
tier=str(int(quality_tier)),
score=score,
signals=list(signals), # mutable-ok: routing decision metadata uses a JSON list
conversation_continuing=conversation_continuing,
)
quality_routing_decision.update(savings_fields)
return PreRoutingHookResponse(
model=routed_model,
messages=messages,
routing_decision=StandardLoggingRoutingDecision(
router_model_name=self.model_name,
router_type="quality",
routed_model=routed_model,
cause="quality_tier",
tier=str(int(quality_tier)),
score=score,
signals=list(signals),
),
routing_decision=quality_routing_decision,
)

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 not messages or any(message.get("role") == "assistant" for message in messages)
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

@ -7,7 +7,7 @@ model on metadata, and return a PreRoutingHookResponse.
Routing is stateless per-turn — `pick_model` does not take a session id.
"""
from unittest.mock import AsyncMock
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -43,6 +43,51 @@ 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():
router = _make_router()
router_instance = MagicMock()
router_instance.model_name_to_deployment_indices = {"fast": [0], "smart": [1]}
router_instance.model_list = [
{"litellm_params": {"model": "openai/fast"}, "model_info": {"id": "fast-id"}},
{"litellm_params": {"model": "openai/smart"}, "model_info": {"id": "smart-id"}},
]
router_instance.get_deployment_model_info.side_effect = lambda deployment_id, model: {
"input_cost_per_token": 0.00000015 if deployment_id == "fast-id" else 0.000005,
"output_cost_per_token": 0.00000015 if deployment_id == "fast-id" else 0.000005,
}
router.litellm_router_instance = router_instance
router.pick_model = AsyncMock(return_value="smart") # type: ignore[method-assign]
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"] == "smart-id"
@pytest.mark.asyncio
async def test_empty_messages_use_conservative_continuing_shape():
router = _make_router()
router.pick_model = AsyncMock(return_value="smart") # type: ignore[method-assign]
response = await router.async_pre_routing_hook(
model="smart-cheap-router", request_kwargs={}, messages=None
)
assert response is not None
assert response.routing_decision["conversation_continuing"] is True
@pytest.mark.asyncio
async def test_classifies_last_user_message_for_request_type():
r = _make_router()

View file

@ -822,6 +822,70 @@ class TestKeywordOverride:
class TestDecisionMetadata:
@pytest.mark.asyncio
@pytest.mark.parametrize("content, cause", [("hi", "quality_tier"), ("write python", "keyword")])
async def test_conversation_shape_is_read_once_for_metadata_and_savings(self, keyword_router, content, cause):
class CountedMessage(dict):
role_reads = 0
def get(self, key, default=None):
if key == "role":
self.role_reads += 1
return super().get(key, default)
message = CountedMessage(role="user", content=content)
kwargs: Dict[str, Any] = {}
response = await keyword_router.async_pre_routing_hook("qr", kwargs, [message])
assert response is not None
assert response.routing_decision["cause"] == cause
assert response.routing_decision["conversation_continuing"] is False
assert kwargs["metadata"]["quality_router_decision"]["conversation_continuing"] is False
assert message.role_reads == 2
@pytest.mark.asyncio
async def test_decision_includes_savings_baseline_and_conversation_shape(self, quality_router):
quality_router.litellm_router_instance.model_name_to_deployment_indices = {
"haiku": [0],
"sonnet": [1],
"opus": [2],
"opus-next": [3],
}
quality_router.litellm_router_instance.get_deployment_model_info.side_effect = (
lambda deployment_id, model: {
"input_cost_per_token": {
"id-haiku": 0.000001,
"id-sonnet": 0.000002,
"id-opus": 0.000003,
"id-opus-next": 0.000004,
}[deployment_id],
"output_cost_per_token": 0.000001,
}
)
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/sonnet"
assert response.routing_decision["savings_baseline_deployment_id"] == "id-sonnet"
first_turn = await quality_router.async_pre_routing_hook(
model="quality-router-test",
request_kwargs={},
messages=[{"role": "user", "content": "first request"}],
)
assert first_turn is not None
assert first_turn.routing_decision["conversation_continuing"] is False
@pytest.mark.asyncio
async def test_hook_stashes_decision_in_request_kwargs_metadata(
self, quality_router

View file

@ -4,6 +4,7 @@ from litellm.router import Router
from litellm.router_strategy.savings_baseline import (
Baseline,
canonical_model,
conversation_is_continuing,
_models_in,
_most_expensive,
resolve_baseline,
@ -22,6 +23,19 @@ def parent() -> Router:
)
@pytest.mark.parametrize(
"messages, expected",
[
(None, True),
([], True),
([{"role": "user", "content": "hello"}], False),
([{"role": "user"}, {"role": "assistant"}, {"role": "user"}], True),
],
)
def test_conversation_shape_for_savings(messages, expected):
assert conversation_is_continuing(messages) is expected
class TestCanonicalModel:
def test_qualifies_a_bare_name_with_the_provider_that_owns_it(self):
assert canonical_model("claude-opus-5") == "anthropic/claude-opus-5"