mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(streaming): drain sync (boto3) streams in one thread to stop bursty delivery
This commit is contained in:
parent
35dc982692
commit
7d21ac9300
2 changed files with 161 additions and 12 deletions
|
|
@ -9,6 +9,7 @@ import traceback
|
|||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
AsyncIterator,
|
||||
Callable,
|
||||
Dict,
|
||||
|
|
@ -63,18 +64,70 @@ _SYNC_ITER_EXHAUSTED = object()
|
|||
_GCHUNK_FIELDS: frozenset = frozenset(GChunk.__annotations__)
|
||||
|
||||
|
||||
def _next_sync_or_exhausted(it: Any) -> Any:
|
||||
"""
|
||||
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _SyncStreamError:
|
||||
"""Carries an exception raised while draining a synchronous stream across the
|
||||
thread boundary so it can be re-raised on the consuming event loop."""
|
||||
|
||||
asyncio.to_thread re-raises thread exceptions inside a coroutine, where PEP 479
|
||||
converts StopIteration to RuntimeError before any except clause can catch it.
|
||||
Returning a sentinel instead keeps StopIteration out of the coroutine boundary.
|
||||
exc: BaseException
|
||||
|
||||
|
||||
async def _aiter_sync_stream_in_thread(sync_stream: Any) -> AsyncGenerator[Any, None]:
|
||||
"""
|
||||
Drain a synchronous iterator (e.g. a boto3 Bedrock event stream) in a single
|
||||
background thread, delivering items to the event loop through an asyncio.Queue.
|
||||
|
||||
Dispatching one asyncio.to_thread per chunk pays thread-pool scheduling latency
|
||||
on every chunk; during that latency the upstream socket buffers several SSE
|
||||
events that then arrive back-to-back, so clients see stalls followed by bursts.
|
||||
Iterating the whole stream in one thread lets each chunk reach the consumer as
|
||||
soon as the provider emits it, restoring the provider's native cadence.
|
||||
|
||||
A dedicated daemon thread is used rather than a pooled executor so a long-lived
|
||||
stream never occupies a slot in the shared/default thread pool that the rest of
|
||||
the event loop relies on for unrelated blocking work.
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
queue: asyncio.Queue[Any] = asyncio.Queue()
|
||||
stop = threading.Event()
|
||||
|
||||
def _deliver(item: Any) -> None:
|
||||
"""
|
||||
Hand one item to the consumer's loop. A consumer may abandon the stream
|
||||
(client disconnect) and let its event loop close while this thread is
|
||||
still blocked inside next(); call_soon_threadsafe then raises, so stop
|
||||
draining rather than crash the thread.
|
||||
"""
|
||||
if stop.is_set():
|
||||
return
|
||||
try:
|
||||
loop.call_soon_threadsafe(queue.put_nowait, item)
|
||||
except RuntimeError:
|
||||
stop.set()
|
||||
|
||||
def _drain() -> None:
|
||||
try:
|
||||
for item in sync_stream:
|
||||
if stop.is_set():
|
||||
break
|
||||
_deliver(item)
|
||||
except Exception as exc:
|
||||
_deliver(_SyncStreamError(exc))
|
||||
finally:
|
||||
_deliver(_SYNC_ITER_EXHAUSTED)
|
||||
|
||||
thread = threading.Thread(target=_drain, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
return next(it)
|
||||
except StopIteration:
|
||||
return _SYNC_ITER_EXHAUSTED
|
||||
while True:
|
||||
item = await queue.get()
|
||||
if item is _SYNC_ITER_EXHAUSTED:
|
||||
return
|
||||
if isinstance(item, _SyncStreamError):
|
||||
raise item.exc
|
||||
yield item
|
||||
finally:
|
||||
stop.set()
|
||||
|
||||
|
||||
def is_async_iterable(obj: Any) -> bool:
|
||||
|
|
@ -127,6 +180,7 @@ class CustomStreamWrapper:
|
|||
self.custom_llm_provider = custom_llm_provider
|
||||
self.logging_obj: LiteLLMLoggingObject = logging_obj
|
||||
self.completion_stream = completion_stream
|
||||
self._threaded_sync_stream: AsyncGenerator[Any, None] | None = None
|
||||
self.sent_first_chunk = False
|
||||
self.sent_last_chunk = False
|
||||
self._stream_created_time: float = time.time()
|
||||
|
|
@ -223,6 +277,8 @@ class CustomStreamWrapper:
|
|||
return self
|
||||
|
||||
async def aclose(self):
|
||||
threaded_sync_stream = self._threaded_sync_stream
|
||||
self._threaded_sync_stream = None
|
||||
if self.completion_stream is not None:
|
||||
stream_to_close = self.completion_stream
|
||||
self.completion_stream = None
|
||||
|
|
@ -230,6 +286,14 @@ class CustomStreamWrapper:
|
|||
# Without this, CancelledError is thrown into every await during
|
||||
# task group cancellation, preventing HTTP connection release.
|
||||
with anyio.CancelScope(shield=True):
|
||||
if threaded_sync_stream is not None:
|
||||
try:
|
||||
await threaded_sync_stream.aclose()
|
||||
except BaseException as e:
|
||||
verbose_logger.debug(
|
||||
"CustomStreamWrapper.aclose: error closing threaded sync stream: %s",
|
||||
e,
|
||||
)
|
||||
try:
|
||||
if hasattr(stream_to_close, "aclose"):
|
||||
await stream_to_close.aclose()
|
||||
|
|
@ -1982,9 +2046,9 @@ class CustomStreamWrapper:
|
|||
if isinstance(self.completion_stream, str) or isinstance(self.completion_stream, bytes):
|
||||
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 self._threaded_sync_stream is None:
|
||||
self._threaded_sync_stream = _aiter_sync_stream_in_thread(self.completion_stream)
|
||||
chunk = await self._threaded_sync_stream.__anext__()
|
||||
if chunk is not None and chunk != b"":
|
||||
processed_chunk = self.chunk_creator(chunk=chunk)
|
||||
if processed_chunk is None:
|
||||
|
|
|
|||
|
|
@ -2660,6 +2660,91 @@ async def test_custom_stream_wrapper_anext_exhaustion_raises_stop_async_iteratio
|
|||
pytest.fail(f"PEP 479 regression: StopIteration leaked as RuntimeError: {e}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_stream_wrapper_anext_drains_sync_iterator_concurrently(
|
||||
logging_obj: Logging,
|
||||
):
|
||||
"""
|
||||
Regression for #34502: a synchronous (boto3-style) stream must be drained in a
|
||||
single background thread so the provider keeps emitting chunks while the async
|
||||
consumer is still processing earlier ones. With the previous per-chunk
|
||||
asyncio.to_thread dispatch the sync iterator only advances when the consumer
|
||||
calls __anext__, so a slow consumer stalls the provider and chunks arrive in
|
||||
bursts. This test asserts the provider races ahead of a slow consumer.
|
||||
"""
|
||||
|
||||
def _make_chunk(content: str) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-concurrent",
|
||||
created=int(time.time()),
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
system_fingerprint=None,
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(
|
||||
provider_specific_fields=None,
|
||||
content=content,
|
||||
role="assistant",
|
||||
function_call=None,
|
||||
tool_calls=None,
|
||||
audio=None,
|
||||
),
|
||||
logprobs=None,
|
||||
)
|
||||
],
|
||||
provider_specific_fields={},
|
||||
usage=None,
|
||||
)
|
||||
|
||||
class TimestampingIterator:
|
||||
"""Fast sync producer that records when each chunk leaves the iterator."""
|
||||
|
||||
def __init__(self, chunks, produced_at: list):
|
||||
self._it = iter(chunks)
|
||||
self._produced_at = produced_at
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
chunk = next(self._it) # raises StopIteration when exhausted
|
||||
self._produced_at.append(time.monotonic())
|
||||
return chunk
|
||||
|
||||
n_chunks = 5
|
||||
produced_at: list = []
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=TimestampingIterator(
|
||||
[_make_chunk(str(i)) for i in range(n_chunks)], produced_at
|
||||
),
|
||||
model="test-model",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="cached_response",
|
||||
)
|
||||
|
||||
consumer_delay = 0.05
|
||||
consumed_at: list = []
|
||||
async for _ in wrapper:
|
||||
consumed_at.append(time.monotonic())
|
||||
await asyncio.sleep(consumer_delay) # simulate a slow downstream client
|
||||
|
||||
assert len(produced_at) == n_chunks
|
||||
assert len(consumed_at) >= n_chunks # finalize appends a trailing finish chunk
|
||||
# The single-thread pump produces every chunk while the consumer is still
|
||||
# draining the first couple, so production completes well before the consumer
|
||||
# reaches the tail. Per-chunk dispatch would interleave the two, making the
|
||||
# last production timestamp land right before each corresponding consumption.
|
||||
reference_consumption = consumed_at[n_chunks - 2]
|
||||
assert max(produced_at) < reference_consumption, (
|
||||
f"provider did not race ahead of slow consumer "
|
||||
f"(last produced at {max(produced_at):.3f}, "
|
||||
f"reference consumption at {reference_consumption:.3f})"
|
||||
)
|
||||
|
||||
|
||||
# Azure streaming chunks that reproduce issue #24221:
|
||||
# Azure sends an initial chunk with prompt_filter_results and choices=[],
|
||||
# then a chunk with role='assistant' and content='', then content chunks.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue