mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
fix(router): preserve bound router fallbacks for subagents
This commit is contained in:
parent
ed0d3f9442
commit
1306a4505a
3 changed files with 50 additions and 5 deletions
|
|
@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any, Final
|
|||
import litellm
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_structure
|
||||
from litellm.router_utils.add_retry_fallback_headers import (
|
||||
add_fallback_headers_to_response,
|
||||
|
|
@ -231,8 +232,6 @@ def record_pre_routing_selection(request_kwargs: Mapping[str, Any] | None, selec
|
|||
on /v1/messages the top-level ``metadata`` dict is the provider's own request field,
|
||||
so a blanket write would forward the tier stamp upstream.
|
||||
"""
|
||||
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
|
||||
|
||||
if request_kwargs is None:
|
||||
return
|
||||
bucket: Final = request_kwargs.get(get_metadata_variable_name_from_kwargs(request_kwargs))
|
||||
|
|
@ -267,10 +266,13 @@ def get_pre_routing_selection(kwargs: Mapping[str, Any]) -> str | None:
|
|||
def fallback_lookup_groups(kwargs: Mapping[str, Any], model_group: str | None) -> tuple[str, ...]:
|
||||
"""
|
||||
Ordered keys for resolving a fallback chain: the tier a pre-routing hook selected wins,
|
||||
and the requested group still resolves when no tier-keyed chain exists, so configs keyed
|
||||
on the router name (the documented contract) keep working behind auto-routers.
|
||||
then the routed group, then the requested group. The routed group differs when Claude Code
|
||||
session affinity remaps a subagent's concrete model to its bound router.
|
||||
"""
|
||||
ordered: Final = (get_pre_routing_selection(kwargs), model_group)
|
||||
metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
|
||||
routed_group_value: Final = metadata.get("model_group") if isinstance(metadata, Mapping) else None
|
||||
routed_group: Final = routed_group_value if isinstance(routed_group_value, str) else None
|
||||
ordered: Final = (get_pre_routing_selection(kwargs), routed_group, model_group)
|
||||
return tuple(dict.fromkeys(group for group in ordered if group))
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1220,6 +1220,28 @@ class TestOrderedFallbackLookupGroups:
|
|||
assert fallback_lookup_groups({}, "smart-router") == ("smart-router",)
|
||||
assert fallback_lookup_groups({}, None) == ()
|
||||
|
||||
def test_session_remap_keeps_the_bound_router_between_tier_and_requested_group(self):
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
PRE_ROUTING_SELECTED_MODEL_KEY,
|
||||
fallback_lookup_groups,
|
||||
)
|
||||
|
||||
kwargs = {
|
||||
"litellm_metadata": {
|
||||
PRE_ROUTING_SELECTED_MODEL_KEY: "tier1",
|
||||
"model_group": "smart-router",
|
||||
}
|
||||
}
|
||||
|
||||
assert fallback_lookup_groups(kwargs, "requested-model") == (
|
||||
"tier1",
|
||||
"smart-router",
|
||||
"requested-model",
|
||||
)
|
||||
assert fallback_lookup_groups({"metadata": {"model_group": []}}, "requested-model") == (
|
||||
"requested-model",
|
||||
)
|
||||
|
||||
def test_first_resolving_group_wins_and_generic_idx_survives_a_miss(self):
|
||||
from litellm.router_utils.fallback_event_handlers import (
|
||||
get_fallback_model_group_for_lookup_groups,
|
||||
|
|
|
|||
|
|
@ -8679,6 +8679,27 @@ class TestClaudeCodeSubagentSessionRouterBinding:
|
|||
assert response.choices[0].message.content == "expensive response"
|
||||
assert subagent_kwargs["metadata"]["routing_decision"]["routed_model"] == "cheap-model"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subagent_can_use_the_bound_router_name_fallback(self):
|
||||
router = self._router(
|
||||
cheap_response="litellm.RateLimitError",
|
||||
fallbacks=[{"smart-router": ["expensive-model"]}],
|
||||
)
|
||||
|
||||
await router.acompletion(
|
||||
model="smart-router",
|
||||
messages=[{"role": "user", "content": "main turn"}],
|
||||
**self._request_kwargs(),
|
||||
)
|
||||
|
||||
response = await router.acompletion(
|
||||
model="expensive-model",
|
||||
messages=[{"role": "user", "content": "subagent turn"}],
|
||||
**self._request_kwargs(agent_id="agent-1234"),
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "expensive response"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_router_binding_is_scoped_to_the_authenticated_key(self):
|
||||
router = self._router()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue