fix(streaming): drain sync (boto3) streams in one thread to stop bursty delivery

This commit is contained in:
Devin AI 2026-07-24 20:17:21 +00:00
parent 35dc982692
commit 7d21ac9300
2 changed files with 161 additions and 12 deletions

View file

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

View file

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