mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
e59e34bed3
commit
46db4699cf
2 changed files with 138 additions and 12 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue