diff --git a/litellm/litellm_core_utils/fallback_utils.py b/litellm/litellm_core_utils/fallback_utils.py index cc36f47b516..70d5c3e6dae 100644 --- a/litellm/litellm_core_utils/fallback_utils.py +++ b/litellm/litellm_core_utils/fallback_utils.py @@ -34,7 +34,20 @@ async def async_completion_with_fallbacks(**kwargs): nested_kwargs: Final = kwargs.pop("kwargs", {}) original_model: Final = kwargs["model"] model = original_model - fallbacks: Final = [original_model] + nested_kwargs.pop("fallbacks", []) + top_level_fallbacks = kwargs.pop("fallbacks", []) + nested_fallbacks = nested_kwargs.pop("fallbacks", []) + combined_fallbacks = [] + if isinstance(top_level_fallbacks, list): + combined_fallbacks.extend(top_level_fallbacks) + elif top_level_fallbacks: + combined_fallbacks.append(top_level_fallbacks) + + if isinstance(nested_fallbacks, list): + combined_fallbacks.extend(nested_fallbacks) + elif nested_fallbacks: + combined_fallbacks.append(nested_fallbacks) + + fallbacks: Final = [original_model] + combined_fallbacks kwargs.pop("acompletion", None) # Remove to prevent keyword conflicts litellm_call_id: Final = str(uuid.uuid4()) base_kwargs: Final = {**kwargs, **nested_kwargs, "litellm_call_id": litellm_call_id} diff --git a/tests/unit/litellm_core_utils/test_fallback_utils.py b/tests/unit/litellm_core_utils/test_fallback_utils.py index 90a61696e9d..1db4b222d9d 100644 --- a/tests/unit/litellm_core_utils/test_fallback_utils.py +++ b/tests/unit/litellm_core_utils/test_fallback_utils.py @@ -167,3 +167,45 @@ def test_process_response_headers_ignores_preserve_flag_for_httpx_headers(): result = process_response_headers(raw, preserve_litellm_internal_headers=True) assert "x-litellm-attempted-fallbacks" not in result assert result["llm_provider-x-litellm-attempted-fallbacks"] == "1" + + +@pytest.mark.asyncio +async def test_async_completion_with_top_level_fallbacks(monkeypatch): + attempted_models: list[str] = [] + + async def _fake_acompletion(*, model: str, **kwargs): + attempted_models.append(model) + if model == "primary-model": + raise Exception("primary failed") + return {"model": model} + + monkeypatch.setattr(litellm, "acompletion", _fake_acompletion) + + response = await async_completion_with_fallbacks( + model="primary-model", + fallbacks=["fallback-model-1", "fallback-model-2"], + ) + assert response["model"] == "fallback-model-1" + assert attempted_models == ["primary-model", "fallback-model-1"] + + +@pytest.mark.asyncio +async def test_async_completion_with_combined_top_level_and_nested_fallbacks(monkeypatch): + attempted_models: list[str] = [] + + async def _fake_acompletion(*, model: str, **kwargs): + attempted_models.append(model) + if model in ["primary-model", "top-fallback"]: + raise Exception(f"{model} failed") + return {"model": model} + + monkeypatch.setattr(litellm, "acompletion", _fake_acompletion) + + response = await async_completion_with_fallbacks( + model="primary-model", + fallbacks=["top-fallback"], + kwargs={"fallbacks": ["nested-fallback"]}, + ) + assert response["model"] == "nested-fallback" + assert attempted_models == ["primary-model", "top-fallback", "nested-fallback"] +