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
This commit is contained in:
Chau Vu / CPF-FAMILY 2026-09-30 08:46:10 +07:00
parent d098b02ed9
commit ecf5be1729
2 changed files with 60 additions and 1 deletions

View file

@ -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}

View file

@ -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