fix(streaming): replace per-chunk asyncio.to_thread with queue-based async iterator

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.
This commit is contained in:
yryzhan 2026-05-20 15:26:09 +02:00
parent e59e34bed3
commit 46db4699cf
2 changed files with 138 additions and 12 deletions

View file

@ -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(

View file

@ -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()