diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 038d2d81277..71413668e23 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 268b244c86d..a744885b3bb 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) ) diff --git a/litellm/router.py b/litellm/router.py index fac48b45fb3..64a6b34e26a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 ): diff --git a/litellm/router_utils/add_retry_fallback_headers.py b/litellm/router_utils/add_retry_fallback_headers.py index 6b921a0db8a..840150bf335 100644 --- a/litellm/router_utils/add_retry_fallback_headers.py +++ b/litellm/router_utils/add_retry_fallback_headers.py @@ -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( diff --git a/tests/test_litellm/proxy/test_response_model_sanitization.py b/tests/test_litellm/proxy/test_response_model_sanitization.py index 621291b8331..f85224c6962 100644 --- a/tests/test_litellm/proxy/test_response_model_sanitization.py +++ b/tests/test_litellm/proxy/test_response_model_sanitization.py @@ -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