fix(streaming): bounded queue, proper cleanup, get_running_loop

- Replace unbounded asyncio.Queue() with maxsize=64 + backpressure
- Add _QueueError class instead of fragile tuple-matching sentinel
- Add _queue_wrapper cleanup in httpx.TimeoutException and generic
  Exception handlers to prevent zombie threads
- Replace asyncio.get_event_loop() with asyncio.get_running_loop()
- Add proper Optional type annotation for _queue_wrapper in __init__
- Replace hasattr checks with direct None checks
- Add test_queue_wrapper_empty_stream test case
- Fix tests to use asyncio.get_running_loop()
This commit is contained in:
yryzhan 2026-05-20 16:14:04 +02:00
parent 46db4699cf
commit caf9b501c9
2 changed files with 52 additions and 19 deletions

View file

@ -75,7 +75,14 @@ def _next_sync_or_exhausted(it: Any) -> Any:
_QUEUE_EXHAUSTED = object()
_QUEUE_ERROR = object()
_QUEUE_MAX_SIZE = 64
class _QueueError:
__slots__ = ("exc",)
def __init__(self, exc: BaseException) -> None:
self.exc = exc
class _SyncIteratorToQueue:
@ -89,31 +96,41 @@ class _SyncIteratorToQueue:
def __init__(self, sync_iterator: Any, loop: asyncio.AbstractEventLoop) -> None:
self._iterator = sync_iterator
self._loop = loop
self._queue: asyncio.Queue = asyncio.Queue()
self._queue: asyncio.Queue = asyncio.Queue(maxsize=_QUEUE_MAX_SIZE)
self._stop_event = threading.Event()
self._thread = threading.Thread(target=self._producer, daemon=True)
self._thread.start()
def _put(self, item: Any) -> None:
"""Enqueue with backpressure — blocks producer until space is available."""
while not self._stop_event.is_set():
fut = asyncio.run_coroutine_threadsafe(self._queue.put(item), self._loop)
try:
fut.result(timeout=0.5)
return
except TimeoutError:
continue
except Exception:
return
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)
self._put(item)
except StopIteration:
self._loop.call_soon_threadsafe(
self._queue.put_nowait, _QUEUE_EXHAUSTED
)
self._put(_QUEUE_EXHAUSTED)
return
except Exception as exc:
self._loop.call_soon_threadsafe(self._queue.put_nowait, (_QUEUE_ERROR, exc))
self._put(_QueueError(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]
if isinstance(item, _QueueError):
raise item.exc
return item
def close(self) -> None:
@ -226,6 +243,7 @@ class CustomStreamWrapper:
self.is_function_call = self.check_is_function_call(logging_obj=logging_obj)
self.created: Optional[int] = None
self._last_returned_hidden_params: Optional[dict] = None
self._queue_wrapper: Optional[_SyncIteratorToQueue] = None
def _check_max_streaming_duration(self) -> None:
"""Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS."""
@ -2172,12 +2190,10 @@ class CustomStreamWrapper:
):
chunk = self.completion_stream
else:
if (
not hasattr(self, "_queue_wrapper")
or self._queue_wrapper is None
):
if self._queue_wrapper is None:
self._queue_wrapper = _SyncIteratorToQueue(
self.completion_stream, asyncio.get_event_loop()
self.completion_stream,
asyncio.get_running_loop(),
)
try:
chunk = await self._queue_wrapper.get()
@ -2202,7 +2218,7 @@ 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:
if self._queue_wrapper is not None:
self._queue_wrapper.close()
self._queue_wrapper = None
if self.sent_last_chunk is True:
@ -2286,6 +2302,9 @@ class CustomStreamWrapper:
processed_chunk = self.finish_reason_handler()
return processed_chunk
except httpx.TimeoutException as e: # if httpx read timeout error occues
if self._queue_wrapper is not None:
self._queue_wrapper.close()
self._queue_wrapper = None
traceback_exception = traceback.format_exc()
## ADD DEBUG INFORMATION - E.G. LITELLM REQUEST TIMEOUT
traceback_exception += "\nLiteLLM Default Request Timeout - {}".format(
@ -2303,6 +2322,9 @@ class CustomStreamWrapper:
)
self._handle_stream_fallback_error(e)
except Exception as e:
if self._queue_wrapper is not None:
self._queue_wrapper.close()
self._queue_wrapper = None
traceback_exception = traceback.format_exc()
if self.logging_obj is not None:
## LOGGING

View file

@ -2133,7 +2133,7 @@ from litellm.litellm_core_utils.streaming_handler import _SyncIteratorToQueue
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()
loop = asyncio.get_running_loop()
wrapper = _SyncIteratorToQueue(iter(items), loop)
results = []
@ -2153,7 +2153,7 @@ async def test_queue_wrapper_propagates_exception():
yield "chunk2"
raise ValueError("intentional error")
loop = asyncio.get_event_loop()
loop = asyncio.get_running_loop()
wrapper = _SyncIteratorToQueue(failing_iter(), loop)
assert await wrapper.get() == "chunk1"
@ -2174,7 +2174,7 @@ async def test_queue_wrapper_close_stops_producer():
call_count += 1
yield call_count
loop = asyncio.get_event_loop()
loop = asyncio.get_running_loop()
wrapper = _SyncIteratorToQueue(counting_iter(), loop)
await wrapper.get()
@ -2185,10 +2185,21 @@ async def test_queue_wrapper_close_stops_producer():
assert call_count == count_after_close or call_count <= count_after_close + 1
@pytest.mark.asyncio
async def test_queue_wrapper_empty_stream():
"""An empty iterator raises StopAsyncIteration immediately."""
loop = asyncio.get_running_loop()
wrapper = _SyncIteratorToQueue(iter([]), loop)
with pytest.raises(StopAsyncIteration):
await wrapper.get()
wrapper.close()
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)
wrapper = _SyncIteratorToQueue(iter([1]), loop)
assert not hasattr(wrapper, "__aiter__")
assert not hasattr(wrapper, "__anext__")
wrapper.close()