mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(fallback_utils): combine top-level and nested fallbacks and move tests to existing unit test file
This commit is contained in:
parent
ecf5be1729
commit
4f01da3061
3 changed files with 54 additions and 58 deletions
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue