From c1f126922b6924222169ce4efccc679281d235e7 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Sat, 25 Jul 2026 12:39:49 -0400 Subject: [PATCH] fix(router): re-raise mid-stream fallback on any generated content, not just text The re-raise guard added for MidStreamFallbackError only checked generated_content, which tracks text deltas alone. A stream that emitted a tool-call or reasoning-only chunk before failing had generated_content="" despite already streaming to the client, so the router silently retried and the client saw duplicated/inconsistent output. The guard now also inspects the wrapper's raw chunks for tool_calls/reasoning_content. Also moves the deferred-stream HTTP-framing-header stripping out of Router._acompletion into the proxy's _handle_llm_api_exception: Router is used directly as an SDK as well as by the proxy, and stripping headers there dropped legitimate provider metadata (content-type, proxy-authenticate) for direct SDK callers who never see the proxy's own response construction. schema.d.ts regenerated via make pre-commit; unrelated to this change. --- litellm/proxy/common_request_processing.py | 3 +- litellm/router.py | 26 ++- .../proxy/test_common_request_processing.py | 45 +++++ .../test_redact_string_in_error_paths.py | 51 ++++++ tests/test_litellm/test_router.py | 172 ++++++++++++++++-- 5 files changed, 280 insertions(+), 17 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f9cad283166..8b62f4d30c7 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -52,7 +52,7 @@ from litellm.proxy.common_utils.callback_utils import ( from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails -from litellm.router import Router +from litellm.router import _HTTP_FRAMING_HEADERS, Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.types.guardrails import GuardrailEventHooks @@ -2693,6 +2693,7 @@ class ProxyBaseLLMRequestProcessing: _response_headers = getattr(_response, "headers", None) if _response_headers: headers = get_response_headers(dict(_response_headers)) + headers = {k: v for k, v in headers.items() if k.lower() not in _HTTP_FRAMING_HEADERS} headers.update(custom_headers) # Call response headers hook for failure diff --git a/litellm/router.py b/litellm/router.py index 4cc70e1e3a0..ee2a94d9b84 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -327,6 +327,21 @@ def _strip_http_framing_headers(exc: BaseException) -> None: setattr(exc, "headers", {k: v for k, v in headers.items() if k.lower() not in _HTTP_FRAMING_HEADERS}) +def _stream_chunks_have_generated_content(chunks: List[ModelResponseStream]) -> bool: + for chunk in chunks: + if not chunk.choices: + continue + delta = chunk.choices[0].delta + if ( + delta.get("content") + or delta.get("tool_calls") + or delta.get("function_call") + or delta.get("reasoning_content") + ): + return True + return False + + class RoutingArgs(enum.Enum): ttl = 60 # 1min (RPM/TPM expire key) @@ -2114,7 +2129,9 @@ class Router: async for item in model_response: yield item except MidStreamFallbackError as e: - if not e.is_pre_first_chunk and e.generated_content: + if not e.is_pre_first_chunk and ( + e.generated_content or _stream_chunks_have_generated_content(model_response.chunks) + ): raise from litellm.main import stream_chunk_builder @@ -2655,7 +2672,9 @@ class Router: for item in model_response: yield item except MidStreamFallbackError as e: - if not e.is_pre_first_chunk and e.generated_content: + if not e.is_pre_first_chunk and ( + e.generated_content or _stream_chunks_have_generated_content(model_response.chunks) + ): raise from litellm.main import stream_chunk_builder @@ -2888,10 +2907,9 @@ class Router: if response.completion_stream is None and response.make_call is not None: try: await response.fetch_stream() - except Exception as fetch_err: + except Exception: if model_name is not None: self.success_calls[model_name] -= 1 - _strip_http_framing_headers(fetch_err) raise return await self._acompletion_streaming_iterator( model_response=response, diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 4d98a05da8d..40f10f22c21 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2952,6 +2952,51 @@ class TestHandleLLMApiExceptionRetryAfter: assert proxy_exc.headers["x-custom"] == "1" +class TestHandleLLMApiExceptionFramingHeaders: + """HTTP-framing headers on the provider exception must be stripped before the + proxy builds its own response, or they conflict with the framing the proxy + itself sets. Non-framing headers must survive unchanged.""" + + async def _invoke(self, exc: Exception): + from litellm.proxy._types import ProxyException, UserAPIKeyAuth + + processor = ProxyBaseLLMRequestProcessing(data={}) + user_api_key_dict = UserAPIKeyAuth(api_key="sk-test") + proxy_logging_obj = MagicMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) + + try: + await processor._handle_llm_api_exception( + e=exc, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + ) + except ProxyException as raised: + return raised + raise AssertionError("ProxyException was not raised") + + async def test_strips_framing_headers_preserves_others(self): + exc = litellm.RateLimitError( + message="Resource exhausted", + llm_provider="vertex_ai", + model="gemini-2.0-flash", + ) + exc.headers = { + "content-length": "42", + "transfer-encoding": "chunked", + "content-encoding": "gzip", + "content-type": "application/json", + "x-request-id": "abc-123", + } + proxy_exc = await self._invoke(exc) + assert "content-length" not in proxy_exc.headers + assert "transfer-encoding" not in proxy_exc.headers + assert "content-encoding" not in proxy_exc.headers + assert "content-type" not in proxy_exc.headers + assert proxy_exc.headers["x-request-id"] == "abc-123" + + class TestAsyncStreamingDataGeneratorFastPath: """Fast/slow path branching in async_streaming_data_generator.""" diff --git a/tests/test_litellm/test_redact_string_in_error_paths.py b/tests/test_litellm/test_redact_string_in_error_paths.py index acedb285dd4..4a624017ea2 100644 --- a/tests/test_litellm/test_redact_string_in_error_paths.py +++ b/tests/test_litellm/test_redact_string_in_error_paths.py @@ -5,8 +5,10 @@ Covers actual execution of redaction in: - WebSocket close reasons in realtime handlers (openai, azure, bedrock) - Gemini RAG ingestion x-goog-api-key header usage - Traceback redaction pattern used in proxy streaming +- Router fallback-failure traceback redaction """ +import logging import os import sys import traceback @@ -190,6 +192,55 @@ class TestProxyStreamingDataGeneratorRedaction: assert "RuntimeError" in redacted_tb +class TestRouterFallbackFailureTracebackRedaction: + """Test the fallback-failure error log in router.py's + async_function_with_fallbacks_common_utils. A prior version passed exc_info=True + alongside an already-redacted message, which bypasses redact_string() entirely + since the stdlib logging module renders exc_info separately from the message.""" + + @pytest.mark.asyncio + async def test_fallback_failure_does_not_leak_secret_via_exc_info(self, caplog): + import litellm + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"}, + }, + { + "model_name": "claude-3-haiku", + "litellm_params": {"model": "anthropic/claude-3-haiku-20240307", "api_key": "fake-key"}, + }, + ], + ) + + secret = "sk-testsecretvalue1234567890abcdef" + + with patch( + "litellm.router.run_async_fallback", + new=AsyncMock(side_effect=RuntimeError(f"boom api_key={secret}")), + ): + with caplog.at_level(logging.ERROR, logger="LiteLLM Router"): + with pytest.raises(Exception): + await router.async_function_with_fallbacks_common_utils( + e=Exception("original failure"), + disable_fallbacks=False, + fallbacks=[{"gpt-3.5-turbo": ["claude-3-haiku"]}], + context_window_fallbacks=None, + content_policy_fallbacks=None, + model_group="gpt-3.5-turbo", + args=(), + kwargs={"model": "gpt-3.5-turbo"}, + ) + + error_records = [r for r in caplog.records if r.levelno == logging.ERROR] + assert error_records, "expected an error log for the fallback failure" + for record in error_records: + assert secret not in record.getMessage() + assert secret not in (record.exc_text or "") + + def _make_mock_ingest_options(): mock = MagicMock() mock.vector_store_config = {} diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 56057817847..eee302e3165 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -2132,6 +2132,79 @@ def test_completion_streaming_iterator_reraises_mid_chunk_error(): list(result) +def test_completion_streaming_iterator_reraises_mid_chunk_error_with_no_text_content(): + """Sync: a reasoning-only chunk sets is_pre_first_chunk=False without populating + generated_content (which only tracks text deltas). The re-raise guard must still + detect this via the raw chunks on the wrapper, or the router silently retries and + the client receives duplicated/inconsistent output.""" + from unittest.mock import MagicMock + + from litellm.exceptions import MidStreamFallbackError + from litellm.types.utils import Delta, StreamingChoices + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, + } + ], + ) + + messages = [{"role": "user", "content": "Test"}] + initial_kwargs = {"model": "gpt-4", "stream": True} + + mid_chunk_error = MidStreamFallbackError( + message="Connection reset", + model="gpt-4", + llm_provider="openai", + generated_content="", + is_pre_first_chunk=False, + ) + + reasoning_chunk = litellm.ModelResponseStream( + id="chatcmpl-partial-1", + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(reasoning_content="Thinking about the answer", role="assistant"), + ) + ], + ) + + class SyncIteratorNoTextChunkError: + def __init__(self): + self.model = "gpt-4" + self.custom_llm_provider = "openai" + self.logging_obj = MagicMock() + self.chunks = [reasoning_chunk] + + def __iter__(self): + return self + + def __next__(self): + raise mid_chunk_error + + mock_response = SyncIteratorNoTextChunkError() + + with patch.object(router, "function_with_fallbacks") as mock_fallback: + result = router._completion_streaming_iterator( + model_response=mock_response, + messages=messages, + initial_kwargs=initial_kwargs, + ) + + with pytest.raises(MidStreamFallbackError): + list(result) + + assert not mock_fallback.called, ( + "fallback must not be attempted once any content, text or non-text, has already streamed" + ) + + @pytest.mark.asyncio async def test_acompletion_streaming_iterator_pre_first_chunk_skips_continuation(): """When MidStreamFallbackError has is_pre_first_chunk=True, use original messages.""" @@ -2200,6 +2273,81 @@ async def test_acompletion_streaming_iterator_pre_first_chunk_skips_continuation assert fallback_kwargs["messages"] == messages +@pytest.mark.asyncio +async def test_acompletion_streaming_iterator_reraises_mid_chunk_error_with_no_text_content(): + """Async: a reasoning-only chunk sets is_pre_first_chunk=False without populating + generated_content (which only tracks text deltas). The re-raise guard must still + detect this via the raw chunks on the wrapper, or the router silently retries and + the client receives duplicated/inconsistent output.""" + from unittest.mock import MagicMock + + from litellm.exceptions import MidStreamFallbackError + from litellm.types.utils import Delta, StreamingChoices + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, + } + ], + ) + + messages = [{"role": "user", "content": "Test"}] + initial_kwargs = {"model": "gpt-4", "stream": True} + + mid_chunk_error = MidStreamFallbackError( + message="Connection reset", + model="gpt-4", + llm_provider="openai", + generated_content="", + is_pre_first_chunk=False, + ) + + reasoning_chunk = litellm.ModelResponseStream( + id="chatcmpl-partial-1", + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(reasoning_content="Thinking about the answer", role="assistant"), + ) + ], + ) + + class AsyncIteratorNoTextChunkError: + def __init__(self): + self.model = "gpt-4" + self.custom_llm_provider = "openai" + self.logging_obj = MagicMock() + self.chunks = [reasoning_chunk] + + def __aiter__(self): + return self + + async def __anext__(self): + raise mid_chunk_error + + mock_response = AsyncIteratorNoTextChunkError() + + with patch.object(router, "async_function_with_fallbacks_common_utils") as mock_fallback_utils: + iterator = await router._acompletion_streaming_iterator( + model_response=mock_response, + messages=messages, + initial_kwargs=initial_kwargs, + ) + + with pytest.raises(MidStreamFallbackError): + async for _ in iterator: + pass + + assert not mock_fallback_utils.called, ( + "fallback must not be attempted once any content, text or non-text, has already streamed" + ) + + # --------------------------------------------------------------------------- # Shared helpers for the _aresponses_streaming_iterator test suite. # --------------------------------------------------------------------------- @@ -5984,13 +6132,13 @@ async def test_acompletion_deferred_stream_error_propagates_through_acompletion( @pytest.mark.asyncio -async def test_acompletion_deferred_stream_strips_framing_headers_on_error(): - """Content-Length / Transfer-Encoding / Content-Encoding / Content-Type from a - provider error response are stripped before the exception propagates, preventing - HTTP framing mismatches when LiteLLM builds its own error body. - - Non-framing headers (e.g. x-request-id) must be preserved. - """ +async def test_acompletion_deferred_stream_preserves_original_headers_on_error(): + """Router is used both by the proxy and directly as an SDK. HTTP-framing headers + (Content-Length, Transfer-Encoding, ...) must NOT be stripped at this layer, or + direct SDK callers lose legitimate provider metadata (e.g. content-type, + proxy-authenticate) that only the proxy's own response construction needs to + worry about. Stripping happens in the proxy layer instead + (_handle_llm_api_exception).""" import litellm as _litellm err = _litellm.RateLimitError( @@ -6027,11 +6175,11 @@ async def test_acompletion_deferred_stream_strips_framing_headers_on_error(): raised = exc_info.value headers = getattr(raised, "headers", {}) - assert "content-length" not in headers, "content-length must be stripped" - assert "transfer-encoding" not in headers, "transfer-encoding must be stripped" - assert "content-encoding" not in headers, "content-encoding must be stripped" - assert "content-type" not in headers, "content-type must be stripped" - assert headers.get("x-request-id") == "abc-123", "x-request-id must be preserved" + assert headers.get("content-length") == "42" + assert headers.get("transfer-encoding") == "chunked" + assert headers.get("content-encoding") == "gzip" + assert headers.get("content-type") == "application/json" + assert headers.get("x-request-id") == "abc-123" @pytest.mark.asyncio