mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
a1ada63f27
commit
c1f126922b
5 changed files with 280 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue