From 4f01da3061dfa56265c3ddd8ba3f204006606529 Mon Sep 17 00:00:00 2001 From: Chau Vu / CPF-FAMILY Date: Wed, 30 Sep 2026 08:54:02 +0700 Subject: [PATCH] fix(fallback_utils): combine top-level and nested fallbacks and move tests to existing unit test file --- litellm/litellm_core_utils/fallback_utils.py | 16 ++++-- .../test_fallback_utils.py | 54 ------------------- .../litellm_core_utils/test_fallback_utils.py | 42 +++++++++++++++ 3 files changed, 54 insertions(+), 58 deletions(-) delete mode 100644 tests/litellm_utils_tests/test_fallback_utils.py diff --git a/litellm/litellm_core_utils/fallback_utils.py b/litellm/litellm_core_utils/fallback_utils.py index c56ee93f132..70d5c3e6dae 100644 --- a/litellm/litellm_core_utils/fallback_utils.py +++ b/litellm/litellm_core_utils/fallback_utils.py @@ -36,10 +36,18 @@ async def async_completion_with_fallbacks(**kwargs): model = original_model top_level_fallbacks = kwargs.pop("fallbacks", []) nested_fallbacks = nested_kwargs.pop("fallbacks", []) - raw_fallbacks: Final = top_level_fallbacks or nested_fallbacks - fallbacks: Final = [original_model] + ( - raw_fallbacks if isinstance(raw_fallbacks, list) else [raw_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/litellm_utils_tests/test_fallback_utils.py b/tests/litellm_utils_tests/test_fallback_utils.py deleted file mode 100644 index c3bf97b073d..00000000000 --- a/tests/litellm_utils_tests/test_fallback_utils.py +++ /dev/null @@ -1,54 +0,0 @@ -import pytest -from unittest.mock import AsyncMock, patch -from litellm.litellm_core_utils.fallback_utils import async_completion_with_fallbacks - - -@pytest.mark.asyncio -async def test_async_completion_with_top_level_fallbacks(): - """Verify that top-level fallbacks keyword argument is properly extracted and used.""" - mock_response = AsyncMock() - mock_response.choices = [] - - with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: - # First call fails, second succeeds - mock_acompletion.side_effect = [ - Exception("Primary model failed"), - mock_response, - ] - - res = await async_completion_with_fallbacks( - model="primary-failing-model", - fallbacks=["secondary-fallback-model"], - messages=[{"role": "user", "content": "hello"}], - ) - - assert mock_acompletion.call_count == 2 - # First attempt with primary model - assert mock_acompletion.call_args_list[0].kwargs["model"] == "primary-failing-model" - # Second attempt with fallback model - assert mock_acompletion.call_args_list[1].kwargs["model"] == "secondary-fallback-model" - assert res is not None - - -@pytest.mark.asyncio -async def test_async_completion_with_nested_fallbacks(): - """Verify backwards compatibility with nested kwargs fallbacks.""" - mock_response = AsyncMock() - mock_response.choices = [] - - with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: - mock_acompletion.side_effect = [ - Exception("Primary model failed"), - mock_response, - ] - - res = await async_completion_with_fallbacks( - model="primary-failing-model", - kwargs={"fallbacks": ["nested-fallback-model"]}, - messages=[{"role": "user", "content": "hello"}], - ) - - assert mock_acompletion.call_count == 2 - assert mock_acompletion.call_args_list[0].kwargs["model"] == "primary-failing-model" - assert mock_acompletion.call_args_list[1].kwargs["model"] == "nested-fallback-model" - assert res is not None 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"] +