This commit is contained in:
Chau Vu 2026-10-01 02:21:34 +08:00 • committed by GitHub
commit b9d94b1ff4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 56 additions and 1 deletions

View file

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

View file

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