From caf9b501c93c556769c13eebbf35928d20c28903 Mon Sep 17 00:00:00 2001 From: yryzhan Date: Wed, 20 May 2026 16:14:04 +0200 Subject: [PATCH] fix(streaming): bounded queue, proper cleanup, get_running_loop - Replace unbounded asyncio.Queue() with maxsize=64 + backpressure - Add _QueueError class instead of fragile tuple-matching sentinel - Add _queue_wrapper cleanup in httpx.TimeoutException and generic Exception handlers to prevent zombie threads - Replace asyncio.get_event_loop() with asyncio.get_running_loop() - Add proper Optional type annotation for _queue_wrapper in __init__ - Replace hasattr checks with direct None checks - Add test_queue_wrapper_empty_stream test case - Fix tests to use asyncio.get_running_loop() --- .../litellm_core_utils/streaming_handler.py | 52 +++++++++++++------ .../test_streaming_handler.py | 19 +++++-- 2 files changed, 52 insertions(+), 19 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 529507ed6f8..d2ab43fdfaf 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -75,7 +75,14 @@ def _next_sync_or_exhausted(it: Any) -> Any: _QUEUE_EXHAUSTED = object() -_QUEUE_ERROR = object() +_QUEUE_MAX_SIZE = 64 + + +class _QueueError: + __slots__ = ("exc",) + + def __init__(self, exc: BaseException) -> None: + self.exc = exc class _SyncIteratorToQueue: @@ -89,31 +96,41 @@ class _SyncIteratorToQueue: def __init__(self, sync_iterator: Any, loop: asyncio.AbstractEventLoop) -> None: self._iterator = sync_iterator self._loop = loop - self._queue: asyncio.Queue = asyncio.Queue() + self._queue: asyncio.Queue = asyncio.Queue(maxsize=_QUEUE_MAX_SIZE) self._stop_event = threading.Event() self._thread = threading.Thread(target=self._producer, daemon=True) self._thread.start() + def _put(self, item: Any) -> None: + """Enqueue with backpressure — blocks producer until space is available.""" + while not self._stop_event.is_set(): + fut = asyncio.run_coroutine_threadsafe(self._queue.put(item), self._loop) + try: + fut.result(timeout=0.5) + return + except TimeoutError: + continue + except Exception: + return + def _producer(self) -> None: try: while not self._stop_event.is_set(): try: item = next(self._iterator) - self._loop.call_soon_threadsafe(self._queue.put_nowait, item) + self._put(item) except StopIteration: - self._loop.call_soon_threadsafe( - self._queue.put_nowait, _QUEUE_EXHAUSTED - ) + self._put(_QUEUE_EXHAUSTED) return except Exception as exc: - self._loop.call_soon_threadsafe(self._queue.put_nowait, (_QUEUE_ERROR, exc)) + self._put(_QueueError(exc)) async def get(self) -> Any: item = await self._queue.get() if item is _QUEUE_EXHAUSTED: raise StopAsyncIteration - if isinstance(item, tuple) and len(item) == 2 and item[0] is _QUEUE_ERROR: - raise item[1] + if isinstance(item, _QueueError): + raise item.exc return item def close(self) -> None: @@ -226,6 +243,7 @@ class CustomStreamWrapper: self.is_function_call = self.check_is_function_call(logging_obj=logging_obj) self.created: Optional[int] = None self._last_returned_hidden_params: Optional[dict] = None + self._queue_wrapper: Optional[_SyncIteratorToQueue] = None def _check_max_streaming_duration(self) -> None: """Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS.""" @@ -2172,12 +2190,10 @@ class CustomStreamWrapper: ): chunk = self.completion_stream else: - if ( - not hasattr(self, "_queue_wrapper") - or self._queue_wrapper is None - ): + if self._queue_wrapper is None: self._queue_wrapper = _SyncIteratorToQueue( - self.completion_stream, asyncio.get_event_loop() + self.completion_stream, + asyncio.get_running_loop(), ) try: chunk = await self._queue_wrapper.get() @@ -2202,7 +2218,7 @@ class CustomStreamWrapper: self.chunks.append(processed_chunk) return processed_chunk except (StopAsyncIteration, StopIteration): - if hasattr(self, "_queue_wrapper") and self._queue_wrapper is not None: + if self._queue_wrapper is not None: self._queue_wrapper.close() self._queue_wrapper = None if self.sent_last_chunk is True: @@ -2286,6 +2302,9 @@ class CustomStreamWrapper: processed_chunk = self.finish_reason_handler() return processed_chunk except httpx.TimeoutException as e: # if httpx read timeout error occues + if self._queue_wrapper is not None: + self._queue_wrapper.close() + self._queue_wrapper = None traceback_exception = traceback.format_exc() ## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT traceback_exception += "\nLiteLLM Default Request Timeout - {}".format( @@ -2303,6 +2322,9 @@ class CustomStreamWrapper: ) self._handle_stream_fallback_error(e) except Exception as e: + if self._queue_wrapper is not None: + self._queue_wrapper.close() + self._queue_wrapper = None traceback_exception = traceback.format_exc() if self.logging_obj is not None: ## LOGGING diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index b5afa988eb7..1184f7c62d9 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -2133,7 +2133,7 @@ from litellm.litellm_core_utils.streaming_handler import _SyncIteratorToQueue async def test_queue_wrapper_delivers_chunks_in_order(): """All chunks from the sync iterator arrive in order via the queue.""" items = list(range(10)) - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() wrapper = _SyncIteratorToQueue(iter(items), loop) results = [] @@ -2153,7 +2153,7 @@ async def test_queue_wrapper_propagates_exception(): yield "chunk2" raise ValueError("intentional error") - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() wrapper = _SyncIteratorToQueue(failing_iter(), loop) assert await wrapper.get() == "chunk1" @@ -2174,7 +2174,7 @@ async def test_queue_wrapper_close_stops_producer(): call_count += 1 yield call_count - loop = asyncio.get_event_loop() + loop = asyncio.get_running_loop() wrapper = _SyncIteratorToQueue(counting_iter(), loop) await wrapper.get() @@ -2185,10 +2185,21 @@ async def test_queue_wrapper_close_stops_producer(): assert call_count == count_after_close or call_count <= count_after_close + 1 +@pytest.mark.asyncio +async def test_queue_wrapper_empty_stream(): + """An empty iterator raises StopAsyncIteration immediately.""" + loop = asyncio.get_running_loop() + wrapper = _SyncIteratorToQueue(iter([]), loop) + + with pytest.raises(StopAsyncIteration): + await wrapper.get() + wrapper.close() + + def test_queue_wrapper_no_aiter(): """Queue wrapper must NOT expose __aiter__ to avoid is_async_iterable() rerouting.""" loop = asyncio.new_event_loop() - wrapper = _SyncIteratorToQueue(iter([]), loop) + wrapper = _SyncIteratorToQueue(iter([1]), loop) assert not hasattr(wrapper, "__aiter__") assert not hasattr(wrapper, "__anext__") wrapper.close()