feat(proxy): enhance fallback handling in streaming mode

This commit is contained in:
lorenzbaraldi 2026-05-16 18:46:36 +02:00
parent ec2f3aadb8
commit 8665cddf44
5 changed files with 86 additions and 2 deletions

View file

@ -1255,6 +1255,16 @@ class ProxyBaseLLMRequestProcessing:
# streaming generators. Pre-call processing can rewrite `self.data["model"]` for
# aliasing/routing, but the OpenAI-compatible response `model` field should reflect
# what the client sent.
attempted_fallbacks = additional_headers.get(
"x-litellm-attempted-fallbacks", 0
)
try:
attempted_fallbacks = int(attempted_fallbacks or 0)
except (TypeError, ValueError):
attempted_fallbacks = 0
if attempted_fallbacks > 0:
self.data["_litellm_attempted_fallbacks"] = attempted_fallbacks
if requested_model_from_client:
self.data["_litellm_client_requested_model"] = (
requested_model_from_client

View file

@ -6709,6 +6709,22 @@ def _restamp_streaming_chunk_model(
if request_data.get("fastest_response", False):
return chunk, model_mismatch_logged
attempted_fallbacks = request_data.get("_litellm_attempted_fallbacks", 0)
if attempted_fallbacks is None:
attempted_fallbacks = 0
hidden_params = getattr(chunk, "_hidden_params", {}) or {}
if isinstance(hidden_params, dict):
additional_headers = hidden_params.get("additional_headers", {}) or {}
attempted_fallbacks = additional_headers.get(
"x-litellm-attempted-fallbacks", attempted_fallbacks
)
try:
attempted_fallbacks = int(attempted_fallbacks or 0)
except (TypeError, ValueError):
attempted_fallbacks = 0
if attempted_fallbacks > 0:
return chunk, model_mismatch_logged
downstream_model = (
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
)

View file

@ -2062,6 +2062,35 @@ class Router:
)
setattr(fallback_item, "usage", combined_usage)
@staticmethod
def _copy_stream_wrapper_additional_headers_to_chunk(
stream_response: Any,
stream_chunk: Any,
) -> Any:
"""
Copy wrapper-level headers onto streamed chunks.
Fallback headers are attached to the stream wrapper after a successful
fallback. The proxy only sees individual chunks during SSE serialization,
so the chunk needs the same hidden header metadata.
"""
response_hidden_params = getattr(stream_response, "_hidden_params", {}) or {}
if not isinstance(response_hidden_params, dict):
return stream_chunk
response_headers = response_hidden_params.get("additional_headers", {}) or {}
if not response_headers or not hasattr(stream_chunk, "_hidden_params"):
return stream_chunk
chunk_hidden_params = getattr(stream_chunk, "_hidden_params", {}) or {}
if not isinstance(chunk_hidden_params, dict):
chunk_hidden_params = {}
chunk_hidden_params.setdefault("additional_headers", {})
chunk_hidden_params["additional_headers"].update(response_headers)
setattr(stream_chunk, "_hidden_params", chunk_hidden_params)
return stream_chunk
async def _acompletion_streaming_iterator(
self,
model_response: CustomStreamWrapper,
@ -2160,6 +2189,12 @@ class Router:
# If fallback returns a streaming response, iterate over it
if hasattr(fallback_response, "__aiter__"):
async for fallback_item in fallback_response: # type: ignore
fallback_item = (
self._copy_stream_wrapper_additional_headers_to_chunk(
stream_response=fallback_response,
stream_chunk=fallback_item,
)
)
if (
fallback_item
and isinstance(fallback_item, ModelResponseStream)
@ -2296,6 +2331,10 @@ class Router:
if hasattr(fallback_response, "__iter__"):
for fallback_item in fallback_response:
fallback_item = router_self._copy_stream_wrapper_additional_headers_to_chunk(
stream_response=fallback_response,
stream_chunk=fallback_item,
)
if (
fallback_item
and isinstance(fallback_item, ModelResponseStream)
@ -8844,7 +8883,7 @@ class Router:
# - else return the model's rate limit headers
"""
if (
isinstance(response, BaseModel)
(isinstance(response, BaseModel) or hasattr(response, "_hidden_params"))
and hasattr(response, "_hidden_params")
and isinstance(response._hidden_params, dict) # type: ignore
):

View file

@ -9,7 +9,9 @@ def _add_headers_to_response(response: Any, headers: dict) -> Any:
"""
Helper function to add headers to a response's hidden params
"""
if response is None or not isinstance(response, BaseModel):
if response is None or not (
isinstance(response, BaseModel) or hasattr(response, "_hidden_params")
):
return response
hidden_params: Optional[Union[dict, HiddenParams]] = getattr(

View file

@ -86,6 +86,23 @@ def test_restamp_streaming_chunk_skips_matching_model():
assert result.model == "client-model"
assert model_mismatch_logged is False
def test_restamp_streaming_chunk_preserves_fallback_model():
from litellm.proxy.proxy_server import _restamp_streaming_chunk_model
chunk = _make_model_response_stream_chunk("fallback-model")
chunk._hidden_params.setdefault("additional_headers", {})
chunk._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] = 1
result, model_mismatch_logged = _restamp_streaming_chunk_model(
chunk=chunk,
requested_model_from_client="client-model",
request_data={"litellm_call_id": "test-call-id"},
model_mismatch_logged=False,
)
assert result is chunk
assert result.model == "fallback-model"
assert model_mismatch_logged is False
def test_fast_serialize_simple_streaming_chunk_matches_model_dump_json():
from litellm.proxy.proxy_server import _serialize_streaming_chunk