From 46db4699cf64c6b37461358fa11348b7333ac556 Mon Sep 17 00:00:00 2001 From: yryzhan Date: Wed, 20 May 2026 15:26:09 +0200 Subject: [PATCH] fix(streaming): replace per-chunk asyncio.to_thread with queue-based async iterator MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Each streaming chunk previously spawned a new thread via asyncio.to_thread, adding ~200µs overhead per chunk. Replace with a single background thread pushing items through an asyncio.Queue — the consumer awaits queue.get() with near-zero latency. The wrapper intentionally does NOT implement __aiter__/__anext__ to avoid is_async_iterable() rerouting the stream into the wrong code path. The str/bytes fast-path is preserved unchanged. --- .../litellm_core_utils/streaming_handler.py | 63 +++++++++++++- .../test_streaming_handler.py | 87 +++++++++++++++++-- 2 files changed, 138 insertions(+), 12 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index fa7faf3035d..529507ed6f8 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -74,6 +74,52 @@ def _next_sync_or_exhausted(it: Any) -> Any: return _SYNC_ITER_EXHAUSTED +_QUEUE_EXHAUSTED = object() +_QUEUE_ERROR = object() + + +class _SyncIteratorToQueue: + """ + Bridges a sync iterator to an async consumer via a single background thread + and an asyncio.Queue — avoiding per-chunk asyncio.to_thread overhead. + + Does NOT implement __aiter__/__anext__ to avoid is_async_iterable() rerouting. + """ + + def __init__(self, sync_iterator: Any, loop: asyncio.AbstractEventLoop) -> None: + self._iterator = sync_iterator + self._loop = loop + self._queue: asyncio.Queue = asyncio.Queue() + self._stop_event = threading.Event() + self._thread = threading.Thread(target=self._producer, daemon=True) + self._thread.start() + + 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) + except StopIteration: + self._loop.call_soon_threadsafe( + self._queue.put_nowait, _QUEUE_EXHAUSTED + ) + return + except Exception as exc: + self._loop.call_soon_threadsafe(self._queue.put_nowait, (_QUEUE_ERROR, 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] + return item + + def close(self) -> None: + self._stop_event.set() + + def is_async_iterable(obj: Any) -> bool: """ Check if an object is an async iterable (can be used with 'async for'). @@ -2126,9 +2172,17 @@ class CustomStreamWrapper: ): chunk = self.completion_stream else: - chunk = await asyncio.to_thread(_next_sync_or_exhausted, self.completion_stream) # type: ignore[arg-type] - if chunk is _SYNC_ITER_EXHAUSTED: - raise StopAsyncIteration + if ( + not hasattr(self, "_queue_wrapper") + or self._queue_wrapper is None + ): + self._queue_wrapper = _SyncIteratorToQueue( + self.completion_stream, asyncio.get_event_loop() + ) + try: + chunk = await self._queue_wrapper.get() + except StopAsyncIteration: + raise if chunk is not None and chunk != b"": processed_chunk = self.chunk_creator(chunk=chunk) if processed_chunk is None: @@ -2148,6 +2202,9 @@ 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: + self._queue_wrapper.close() + self._queue_wrapper = None if self.sent_last_chunk is True: # log the final chunk with accurate streaming values complete_streaming_response = litellm.stream_chunk_builder( 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 49d3c51e340..b5afa988eb7 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -2036,23 +2036,19 @@ async def test_azure_streaming_role_preserved_with_include_usage(sync_mode: bool chunks.append(chunk) # The prompt_filter chunk should be forwarded with choices=[] - assert len(chunks[0].choices) == 0, ( - f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices" - ) + assert ( + len(chunks[0].choices) == 0 + ), f"Expected prompt_filter chunk with choices=[], got {len(chunks[0].choices)} choices" # At least one chunk must have role='assistant' in its delta has_role = any( - len(c.choices) > 0 - and getattr(c.choices[0].delta, "role", None) == "assistant" + len(c.choices) > 0 and getattr(c.choices[0].delta, "role", None) == "assistant" for c in chunks ) assert has_role, ( "No chunk contained role='assistant' in delta (issue #24221). " "Chunk deltas: " - + str([ - c.choices[0].delta if c.choices else "no choices" - for c in chunks - ]) + + str([c.choices[0].delta if c.choices else "no choices" for c in chunks]) ) @@ -2124,3 +2120,76 @@ def test_gemini_legacy_vertex_tool_calls_finish_reason_with_stop_enum(): f"Expected 'tool_calls' but got {final.choices[0].finish_reason!r}. " "STOP enum was not normalised through map_finish_reason()." ) + + +# ============================================================================ +# _SyncIteratorToQueue tests +# ============================================================================ + +from litellm.litellm_core_utils.streaming_handler import _SyncIteratorToQueue + + +@pytest.mark.asyncio +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() + wrapper = _SyncIteratorToQueue(iter(items), loop) + + results = [] + for _ in range(10): + results.append(await wrapper.get()) + + assert results == items + wrapper.close() + + +@pytest.mark.asyncio +async def test_queue_wrapper_propagates_exception(): + """Exceptions from the sync iterator propagate to the async consumer.""" + + def failing_iter(): + yield "chunk1" + yield "chunk2" + raise ValueError("intentional error") + + loop = asyncio.get_event_loop() + wrapper = _SyncIteratorToQueue(failing_iter(), loop) + + assert await wrapper.get() == "chunk1" + assert await wrapper.get() == "chunk2" + with pytest.raises(ValueError, match="intentional error"): + await wrapper.get() + wrapper.close() + + +@pytest.mark.asyncio +async def test_queue_wrapper_close_stops_producer(): + """Calling close() stops the producer thread from calling next().""" + call_count = 0 + + def counting_iter(): + nonlocal call_count + while True: + call_count += 1 + yield call_count + + loop = asyncio.get_event_loop() + wrapper = _SyncIteratorToQueue(counting_iter(), loop) + + await wrapper.get() + wrapper.close() + await asyncio.sleep(0.05) + count_after_close = call_count + await asyncio.sleep(0.05) + assert call_count == count_after_close or call_count <= count_after_close + 1 + + +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) + assert not hasattr(wrapper, "__aiter__") + assert not hasattr(wrapper, "__anext__") + wrapper.close() + loop.close()