fix(fallbacks): preserve fallback model in response when using SDK-level fallbacks

This commit is contained in:
Varshith 2026-05-19 10:46:54 -05:00
parent ec9353cb69
commit 366d8d1db9
3 changed files with 100 additions and 3 deletions

View file

@ -256,6 +256,12 @@ def process_response_headers(response_headers: Union[httpx.Headers, dict]) -> di
"llm_provider-"
): # return raw provider headers (incl. openai-compatible ones)
processed_headers[k] = v
elif k.startswith("x-litellm-"):
# LiteLLM's own internal headers (e.g. x-litellm-attempted-fallbacks,
# x-litellm-model-group) are not LLM provider headers and must not be
# prefixed. Downstream consumers (proxy override, callers checking
# whether a fallback happened) look up the bare key.
processed_headers[k] = v
else:
additional_headers["{}-{}".format("llm_provider", k)] = v

View file

@ -7,6 +7,9 @@ from litellm.litellm_core_utils.core_helpers import (
safe_deep_copy,
filter_internal_params,
)
from litellm.router_utils.add_retry_fallback_headers import (
add_fallback_headers_to_response,
)
from .asyncify import run_async_function
@ -42,7 +45,7 @@ async def async_completion_with_fallbacks(**kwargs):
# Try each fallback model
most_recent_exception_str: Optional[str] = None
for fallback in fallbacks:
for attempted_fallbacks, fallback in enumerate(fallbacks):
try:
completion_kwargs = safe_deep_copy(base_kwargs)
# Handle dictionary fallback configurations
@ -63,7 +66,10 @@ async def async_completion_with_fallbacks(**kwargs):
)
if response is not None:
return response
return add_fallback_headers_to_response(
response=response,
attempted_fallbacks=attempted_fallbacks,
)
except Exception as e:
verbose_logger.exception(

View file

@ -1,7 +1,12 @@
"""Tests for litellm.litellm_core_utils.fallback_utils."""
import pytest
import litellm
from litellm.litellm_core_utils.fallback_utils import async_completion_with_fallbacks
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.fallback_utils import (
async_completion_with_fallbacks,
)
@pytest.mark.asyncio
@ -41,3 +46,83 @@ async def test_fallback_dict_not_mutated(monkeypatch):
"primary-model",
"fallback-model",
]
@pytest.mark.asyncio
async def test_async_completion_with_fallbacks_sets_attempted_fallbacks_header():
"""
When a fallback succeeds, the response must carry the
`x-litellm-attempted-fallbacks` header so the proxy and other callers can
detect that a fallback occurred. Without it,
`_override_openai_response_model` stamps the requested model back over the
fallback model used. See issue #28241.
"""
response = await async_completion_with_fallbacks(
model="openai/primary-llm",
messages=[{"role": "user", "content": "hi"}],
api_key="fake-key",
mock_response=Exception("forced failure"),
kwargs={
"fallbacks": [
{
"model": "openai/backup-llm",
"api_key": "fake-key",
"mock_response": "backup-resp",
}
]
},
)
hidden_params = getattr(response, "_hidden_params", None)
assert isinstance(hidden_params, dict)
headers = hidden_params.get("additional_headers") or {}
assert headers.get("x-litellm-attempted-fallbacks") == 1
@pytest.mark.asyncio
async def test_async_completion_with_fallbacks_header_is_zero_when_primary_succeeds():
"""
When the primary model succeeds on the first attempt, the header should be
`0` (no fallback was used). This mirrors the existing router-level
semantics in `async_function_with_fallbacks`.
"""
response = await async_completion_with_fallbacks(
model="openai/primary-llm",
messages=[{"role": "user", "content": "hi"}],
api_key="fake-key",
mock_response="primary-resp",
kwargs={
"fallbacks": [
{
"model": "openai/backup-llm",
"api_key": "fake-key",
"mock_response": "backup-resp",
}
]
},
)
hidden_params = getattr(response, "_hidden_params", None)
assert isinstance(hidden_params, dict)
headers = hidden_params.get("additional_headers") or {}
assert headers.get("x-litellm-attempted-fallbacks") == 0
assert response.choices[0].message.content == "primary-resp"
def test_process_response_headers_preserves_x_litellm_headers():
"""
`process_response_headers` must not add the `llm_provider-` prefix to
LiteLLM's own internal headers (anything starting with `x-litellm-`).
These are markers set by LiteLLM (e.g. fallback / retry headers); the
proxy and other callers look up the bare key.
"""
result = process_response_headers(
{
"x-litellm-attempted-fallbacks": 1,
"x-litellm-model-group": "gpt-4",
"x-stainless-arch": "arm64",
}
)
assert result["x-litellm-attempted-fallbacks"] == 1
assert result["x-litellm-model-group"] == "gpt-4"
assert result["llm_provider-x-stainless-arch"] == "arm64"