fix(router): reuse quality savings metadata

This commit is contained in:
mikemikimike 2026-09-13 15:48:25 +08:00
parent 00f2ce757e
commit d0562c5f44
4 changed files with 60 additions and 18 deletions

View file

@ -224,12 +224,14 @@ class AdaptiveRouter:
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: StandardLoggingRoutingDecision = {
return {
"savings_baseline_model": baseline.model,
} # mutable-ok: assemble TypedDict kwargs
if baseline.deployment_id is not None:
fields["savings_baseline_deployment_id"] = baseline.deployment_id
return fields
**(
{"savings_baseline_deployment_id": baseline.deployment_id}
if baseline.deployment_id is not None
else {}
),
}
# ---- Pick model ------------------------------------------------------

View file

@ -326,12 +326,14 @@ class QualityRouter(CustomLogger):
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: StandardLoggingRoutingDecision = {
return {
"savings_baseline_model": baseline.model,
} # mutable-ok: assemble TypedDict kwargs
if baseline.deployment_id is not None:
fields["savings_baseline_deployment_id"] = baseline.deployment_id
return fields
**(
{"savings_baseline_deployment_id": baseline.deployment_id}
if baseline.deployment_id is not None
else {}
),
}
async def async_pre_routing_hook(
self,
@ -348,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
@ -376,9 +381,9 @@ class QualityRouter(CustomLogger):
router_type="quality",
routed_model=self.config.default_model,
cause="default_fallback",
conversation_continuing=conversation_is_continuing(messages),
conversation_continuing=conversation_continuing,
)
default_routing_decision.update(self._savings_fields())
default_routing_decision.update(savings_fields)
return PreRoutingHookResponse(
model=self.config.default_model,
messages=messages,
@ -405,7 +410,7 @@ class QualityRouter(CustomLogger):
"matched_keyword": matched_keyword,
"quality_tier": self._model_quality.get(routed_model),
"complexity_tier": None,
"conversation_continuing": conversation_is_continuing(messages),
"conversation_continuing": conversation_continuing,
},
)
keyword_routing_decision: Final = StandardLoggingRoutingDecision(
@ -414,9 +419,9 @@ class QualityRouter(CustomLogger):
routed_model=routed_model,
cause="keyword",
matched_keyword=matched_keyword,
conversation_continuing=conversation_is_continuing(messages),
conversation_continuing=conversation_continuing,
)
keyword_routing_decision.update(self._savings_fields())
keyword_routing_decision.update(savings_fields)
keyword_quality_tier: Final = self._model_quality.get(routed_model)
if keyword_quality_tier is not None:
keyword_routing_decision["tier"] = str(keyword_quality_tier)
@ -454,7 +459,7 @@ class QualityRouter(CustomLogger):
"matched_keyword": None,
"quality_tier": int(quality_tier),
"complexity_tier": complexity_name,
"conversation_continuing": conversation_is_continuing(messages),
"conversation_continuing": conversation_continuing,
},
)
@ -466,9 +471,9 @@ class QualityRouter(CustomLogger):
tier=str(int(quality_tier)),
score=score,
signals=list(signals),
conversation_continuing=conversation_is_continuing(messages),
conversation_continuing=conversation_continuing,
)
quality_routing_decision.update(self._savings_fields())
quality_routing_decision.update(savings_fields)
return PreRoutingHookResponse(
model=routed_model,
messages=messages,

View file

@ -822,6 +822,27 @@ 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 = {

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"