fix(streaming): stop re-raising and re-logging a failed stream on every next()

A stream that raised keeps its provider iterator alive, so a consumer that iterates past the error re-detected, re-raised and re-logged the same failure on every call, flooding failure callbacks for one request. Latch the failure and re-raise it without logging again.
This commit is contained in:
Business.arshgoyal 2026-08-29 10:02:09 +00:00
parent 3993829a21
commit 2571ffe9f7
2 changed files with 112 additions and 9 deletions

View file

@ -241,6 +241,7 @@ class CustomStreamWrapper:
self.custom_llm_provider = custom_llm_provider
self.logging_obj: LiteLLMLoggingObject = logging_obj
self.completion_stream = completion_stream
self._terminal_failure: Exception | None = None
self.sent_first_chunk = False
self.sent_last_chunk = False
self._stream_created_time: float = time.time()
@ -1915,6 +1916,7 @@ class CustomStreamWrapper:
cache_hit = False
if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response":
cache_hit = True
self._raise_if_stream_already_failed()
self._check_max_streaming_duration()
try:
if self.completion_stream is None:
@ -2123,6 +2125,7 @@ class CustomStreamWrapper:
cache_hit = False
if self.custom_llm_provider is not None and self.custom_llm_provider == "cached_response":
cache_hit = True
self._raise_if_stream_already_failed()
try:
# Inside the try (not before it) so a raised litellm.Timeout flows
# through the same except Exception -> _handle_stream_fallback_error
@ -2389,6 +2392,24 @@ class CustomStreamWrapper:
recover_error,
)
def _raise_if_stream_already_failed(self) -> None:
"""
A stream that already raised is finished, but its underlying iterator can
still hand out chunks - a consumer that keeps iterating (or a wrapper that
swallows the error and asks for the next chunk) re-runs detection, then
re-raises and re-logs the same failure on every call, flooding the failure
callbacks with thousands of copies of one request
(https://github.com/BerriAI/litellm/issues/13786).
Re-raise the failure the stream died of, without logging it a second time.
"""
if self._terminal_failure is not None:
raise self._terminal_failure
def _raise_terminal_failure(self, exc: Exception) -> "NoReturn":
self._terminal_failure = exc
raise exc
def _handle_stream_fallback_error(self, e: Exception) -> "NoReturn":
"""
Common error handling for both __next__ and __anext__.
@ -2449,17 +2470,19 @@ class CustomStreamWrapper:
# Exception: 429 (rate-limit) IS retriable/transient — allow it
# through so the Router can switch to a different model group.
if mapped_status_code is not None and 400 <= mapped_status_code < 500 and mapped_status_code != 429:
raise mapped_exception
self._raise_terminal_failure(mapped_exception)
if original_status_code is not None and 400 <= original_status_code < 500 and original_status_code != 429:
raise mapped_exception
self._raise_terminal_failure(mapped_exception)
raise MidStreamFallbackError(
message=str(mapped_exception),
model=self.model,
llm_provider=self.custom_llm_provider or "anthropic",
original_exception=mapped_exception,
generated_content=self.response_uptil_now,
is_pre_first_chunk=not self.sent_first_chunk,
self._raise_terminal_failure(
MidStreamFallbackError(
message=str(mapped_exception),
model=self.model,
llm_provider=self.custom_llm_provider or "anthropic",
original_exception=mapped_exception,
generated_content=self.response_uptil_now,
is_pre_first_chunk=not self.sent_first_chunk,
)
)
@staticmethod

View file

@ -2190,6 +2190,86 @@ def test_raise_on_model_repetition_tolerates_empty_choices(
wrapper.raise_on_model_repetition()
@pytest.mark.asyncio
async def test_stream_stays_failed_after_raising_and_logs_failure_once():
"""
Regression test for https://github.com/BerriAI/litellm/issues/13786
A model stuck repeating one chunk makes the wrapper raise, but the provider
iterator keeps yielding, so a consumer that iterates past the error used to
get a freshly detected, freshly logged failure on every __anext__ call -
thousands of failure callbacks for a single request.
"""
async def looping_stream():
while True:
yield _make_chunk("the model is stuck on this")
logging_obj = MagicMock()
logging_obj.completion_start_time = None
dispatched = []
async def _record_failure(exception, traceback_exception, **kwargs):
dispatched.append(exception)
logging_obj.dispatch_failure_handlers = _record_failure
wrapper = CustomStreamWrapper(
completion_stream=looping_stream(),
model="test-model",
logging_obj=logging_obj,
custom_llm_provider="openai",
)
raised = []
for _ in range(litellm.REPEATED_STREAMING_CHUNK_LIMIT + 5):
try:
await wrapper.__anext__()
except Exception as e:
raised.append(e)
await asyncio.sleep(0)
assert len(raised) == 5
assert all(error is raised[0] for error in raised[1:])
assert len(dispatched) == 1
def test_sync_stream_stays_failed_after_raising_and_logs_failure_once():
"""Sync counterpart of the async test above (see issue #13786)."""
def looping_stream():
while True:
yield _make_chunk("the model is stuck on this")
logging_obj = MagicMock()
logging_obj.completion_start_time = None
failures = []
logging_obj.failure_handler = lambda exception, traceback_exception: failures.append(
exception
)
wrapper = CustomStreamWrapper(
completion_stream=looping_stream(),
model="test-model",
logging_obj=logging_obj,
custom_llm_provider="openai",
)
raised = []
for _ in range(litellm.REPEATED_STREAMING_CHUNK_LIMIT + 5):
try:
next(wrapper)
except Exception as e:
raised.append(e)
time.sleep(0.1)
assert len(raised) == 5
assert all(error is raised[0] for error in raised[1:])
assert len(failures) == 1
def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj):
"""
Test that provider-reported usage from a post-finish_reason chunk