From bccad1d3697af68e15b7f0b50df04f067ac9429c Mon Sep 17 00:00:00 2001 From: Ankit Jhalaria Date: Tue, 14 Apr 2026 14:02:29 -0700 Subject: [PATCH] fix(router): address PR review feedback on fallback model header - Use mg.get("model") for dict-type fallbacks instead of reading back from kwargs, avoiding stale model values when mg has no "model" key - Replace misplaced duplicate test with one that actually exercises the run_async_fallback code path (skipped-loop case raises original exception) Co-Authored-By: Claude Sonnet 4.6 (1M context) --- .../router_utils/fallback_event_handlers.py | 4 +- .../router/test_fallback_headers.py | 42 ++++++++++++------- 2 files changed, 29 insertions(+), 17 deletions(-) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 247d1e971d6..e014e7a2c88 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -132,8 +132,10 @@ async def run_async_fallback( kwargs.update(mg) # Capture the effective fallback model name before the recursive call # so we can stamp it on the response header regardless of further fallbacks. + # For dict fallbacks, read "model" directly from mg rather than from kwargs + # to avoid picking up the stale model that was already in kwargs. effective_fallback_model: Optional[str] = ( - mg if isinstance(mg, str) else kwargs.get("model") + mg if isinstance(mg, str) else mg.get("model") if isinstance(mg, dict) else None ) kwargs.setdefault("metadata", {}).update( {"model_group": kwargs.get("model", None)} diff --git a/tests/test_litellm/router/test_fallback_headers.py b/tests/test_litellm/router/test_fallback_headers.py index b49a63e9f09..5d3d635e2aa 100644 --- a/tests/test_litellm/router/test_fallback_headers.py +++ b/tests/test_litellm/router/test_fallback_headers.py @@ -145,25 +145,35 @@ class TestRunAsyncFallbackHeaderPropagation: assert headers.get("x-litellm-fallback-model-used") == "claude-3-haiku" @pytest.mark.asyncio - async def test_fallback_model_header_not_present_without_fallback(self): + async def test_fallback_model_header_absent_when_primary_succeeds_via_router(self): """ - When the primary model succeeds (no fallback), x-litellm-fallback-model-used - should NOT appear in the response headers. + When run_async_fallback skips all candidates because they equal the + original_model_group, no fallback fires and the response should NOT + carry x-litellm-fallback-model-used. + + This tests the router path: fallback_model_group contains only the + original model so the loop body is never entered and error_from_fallbacks + (the original exception) is raised — which means the caller never gets a + response with the fallback header at all. We verify that by confirming + the exception propagates rather than a header-bearing response being returned. """ - from pydantic import BaseModel + mock_router = MagicMock() + mock_router.log_retry = MagicMock(side_effect=lambda kwargs, e: kwargs) - class _FakeResponse(BaseModel): - model: str = "gpt-4" - _hidden_params: dict = {} - - fake_response = _FakeResponse() - result = add_fallback_headers_to_response( - response=fake_response, - attempted_fallbacks=0, - fallback_model=None, - ) - headers = result._hidden_params.get("additional_headers", {}) - assert "x-litellm-fallback-model-used" not in headers + original_exc = Exception("primary failed") + with pytest.raises(Exception, match="primary failed"): + await run_async_fallback( + litellm_router=mock_router, + # All candidates are the same as original — loop body is skipped entirely + fallback_model_group=["gpt-4"], + original_model_group="gpt-4", + original_exception=original_exc, + max_fallbacks=3, + fallback_depth=0, + model="gpt-4", + ) + # async_function_with_fallbacks was never called — no fallback header stamped + mock_router.async_function_with_fallbacks.assert_not_called() @pytest.mark.asyncio async def test_all_fallbacks_fail_raises_exception(self):