diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index b6c82d184d2..3fc1b2a8e1e 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -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 ------------------------------------------------------ diff --git a/litellm/router_strategy/quality_router/quality_router.py b/litellm/router_strategy/quality_router/quality_router.py index aba274de68b..59af232b498 100644 --- a/litellm/router_strategy/quality_router/quality_router.py +++ b/litellm/router_strategy/quality_router/quality_router.py @@ -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, diff --git a/tests/test_litellm/router_strategy/test_quality_router.py b/tests/test_litellm/router_strategy/test_quality_router.py index 98ba9915342..be984bcb5a1 100644 --- a/tests/test_litellm/router_strategy/test_quality_router.py +++ b/tests/test_litellm/router_strategy/test_quality_router.py @@ -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 = { diff --git a/tests/test_litellm/router_strategy/test_savings_baseline.py b/tests/test_litellm/router_strategy/test_savings_baseline.py index 5efc73d2dcb..859482edc08 100644 --- a/tests/test_litellm/router_strategy/test_savings_baseline.py +++ b/tests/test_litellm/router_strategy/test_savings_baseline.py @@ -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"