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.
This commit is contained in:
Deepanshu 2026-07-25 12:39:49 -04:00
parent a1ada63f27
commit c1f126922b
5 changed files with 280 additions and 17 deletions

View file

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

View file

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

View file

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

View file

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

View file

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