mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(router): track routed model in fallback attempts
This commit is contained in:
parent
6adc14b4b1
commit
7b919f89a8
3 changed files with 56 additions and 5 deletions
|
|
@ -470,10 +470,11 @@ async def run_async_fallback(
|
|||
attempted: Final = (
|
||||
carried_targets if isinstance(carried_targets, AttemptedFallbackTargets) else AttemptedFallbackTargets()
|
||||
)
|
||||
attempted.record(original_model_group)
|
||||
failed_model_group: Final = get_pre_routing_selection(kwargs) or original_model_group
|
||||
attempted.record(failed_model_group)
|
||||
|
||||
for mg in fallback_model_group:
|
||||
if mg == original_model_group:
|
||||
if mg == failed_model_group:
|
||||
continue
|
||||
if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group:
|
||||
verbose_router_logger.info(
|
||||
|
|
|
|||
|
|
@ -614,6 +614,27 @@ async def test_run_async_fallback_forwards_attempted_model_groups_to_nested_call
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_can_target_the_requested_group_when_a_pre_router_replaced_it():
|
||||
"""The requested group was never called when a pre-router selected a tier, so a
|
||||
tier fallback may legitimately target that originally requested group."""
|
||||
router = RecordingRouter()
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["requested-model"],
|
||||
original_model_group="requested-model",
|
||||
original_exception=RuntimeError("selected tier failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="requested-model",
|
||||
metadata={"pre_routing_selected_model": "selected-tier"},
|
||||
)
|
||||
|
||||
assert router.received_kwargs["model"] == "requested-model"
|
||||
assert router.received_kwargs["attempted_targets"].keys == frozenset({"selected-tier", "requested-model"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"entry",
|
||||
|
|
|
|||
|
|
@ -8335,20 +8335,26 @@ class TestClaudeCodeSubagentSessionRouterBinding:
|
|||
)
|
||||
|
||||
@classmethod
|
||||
def _router(cls) -> "litellm.Router":
|
||||
def _router(
|
||||
cls,
|
||||
cheap_response: str = "cheap response",
|
||||
fallbacks: list[dict[str, list[str]]] | None = None,
|
||||
) -> "litellm.Router":
|
||||
from litellm.types.router import TaggedPreRoutingStrategy
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cheap-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "cheap response"},
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": cheap_response},
|
||||
},
|
||||
{
|
||||
"model_name": "expensive-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "mock_response": "expensive response"},
|
||||
},
|
||||
]
|
||||
],
|
||||
fallbacks=fallbacks,
|
||||
num_retries=0,
|
||||
)
|
||||
router.complexity_routers = {
|
||||
"smart-router": [TaggedPreRoutingStrategy(tags=(), strategy=cls._RewriteStrategy())]
|
||||
|
|
@ -8477,6 +8483,29 @@ class TestClaudeCodeSubagentSessionRouterBinding:
|
|||
|
||||
assert response is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subagent_can_fallback_to_its_original_requested_model(self):
|
||||
router = self._router(
|
||||
cheap_response="litellm.RateLimitError",
|
||||
fallbacks=[{"cheap-model": ["expensive-model"]}],
|
||||
)
|
||||
|
||||
await router.acompletion(
|
||||
model="smart-router",
|
||||
messages=[{"role": "user", "content": "main turn"}],
|
||||
**self._request_kwargs(),
|
||||
)
|
||||
subagent_kwargs = self._request_kwargs(agent_id="agent-1234")
|
||||
|
||||
response = await router.acompletion(
|
||||
model="expensive-model",
|
||||
messages=[{"role": "user", "content": "subagent turn"}],
|
||||
**subagent_kwargs,
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "expensive response"
|
||||
assert subagent_kwargs["metadata"]["routing_decision"]["routed_model"] == "cheap-model"
|
||||
|
||||
@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