mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(router): reuse quality savings metadata
This commit is contained in:
parent
00f2ce757e
commit
d0562c5f44
4 changed files with 60 additions and 18 deletions
|
|
@ -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 ------------------------------------------------------
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue