From ecf5be1729d82a12cefa80ebc9c0f4f59a6d00f6 Mon Sep 17 00:00:00 2001 From: Chau Vu / CPF-FAMILY Date: Wed, 30 Sep 2026 08:46:10 +0700 Subject: [PATCH] fix(fallback_utils): extract top-level fallbacks parameter in async_completion_with_fallbacks - Ensure direct keyword argument `fallbacks` is extracted alongside nested kwargs - Add unit tests for both top-level and nested fallbacks configurations - Fixes #43794 --- litellm/litellm_core_utils/fallback_utils.py | 7 ++- .../test_fallback_utils.py | 54 +++++++++++++++++++ 2 files changed, 60 insertions(+), 1 deletion(-) create 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 cc36f47b516..c56ee93f132 100644 --- a/litellm/litellm_core_utils/fallback_utils.py +++ b/litellm/litellm_core_utils/fallback_utils.py @@ -34,7 +34,12 @@ 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", []) + raw_fallbacks: Final = top_level_fallbacks or nested_fallbacks + fallbacks: Final = [original_model] + ( + raw_fallbacks if isinstance(raw_fallbacks, list) else [raw_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 new file mode 100644 index 00000000000..c3bf97b073d --- /dev/null +++ b/tests/litellm_utils_tests/test_fallback_utils.py @@ -0,0 +1,54 @@ +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