This commit is contained in:
Arsh Goyal 2026-09-23 14:51:47 +00:00 • committed by GitHub
commit 302684c288
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 152 additions and 9 deletions

View file

@ -218,6 +218,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()
@ -1762,6 +1763,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:
@ -1971,6 +1973,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
@ -2242,6 +2245,27 @@ 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.
The traceback is dropped first: raising one retained exception object over
and over would otherwise chain a frame per read onto it, holding every
consumer frame alive and growing the traceback without bound.
"""
if self._terminal_failure is not None:
raise self._terminal_failure.with_traceback(None)
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__.
@ -2302,17 +2326,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

@ -2257,6 +2257,123 @@ 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_failed_stream_traceback_does_not_grow_per_read():
"""A consumer that keeps reading past the failure gets the same size traceback every time."""
def looping_stream():
while True:
yield _make_chunk("the model is stuck on this")
logging_obj = MagicMock()
logging_obj.completion_start_time = None
logging_obj.failure_handler = lambda exception, traceback_exception: None
wrapper = CustomStreamWrapper(
completion_stream=looping_stream(),
model="test-model",
logging_obj=logging_obj,
custom_llm_provider="openai",
)
def _traceback_length(error: BaseException) -> int:
length = 0
tb = error.__traceback__
while tb is not None:
length += 1
tb = tb.tb_next
return length
lengths = []
for _ in range(litellm.REPEATED_STREAMING_CHUNK_LIMIT + 20):
try:
next(wrapper)
except Exception as e:
lengths.append(_traceback_length(e))
assert len(lengths) == 20
assert len(set(lengths[1:])) == 1
def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj):
"""
Test that provider-reported usage from a post-finish_reason chunk