mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
Fix blocking sync next
This commit is contained in:
parent
3093ef844e
commit
e38a0a6c40
2 changed files with 76 additions and 1 deletions
|
|
@ -2090,7 +2090,7 @@ class CustomStreamWrapper:
|
|||
):
|
||||
chunk = self.completion_stream
|
||||
else:
|
||||
chunk = next(self.completion_stream) # type: ignore[arg-type]
|
||||
chunk = await asyncio.to_thread(next, self.completion_stream) # type: ignore[arg-type]
|
||||
if chunk is not None and chunk != b"":
|
||||
processed_chunk = self.chunk_creator(chunk=chunk)
|
||||
if processed_chunk is None:
|
||||
|
|
|
|||
|
|
@ -1679,3 +1679,78 @@ def test_tool_use_not_dropped_when_finish_reason_already_set(
|
|||
)
|
||||
assert tool_calls[0].id == "call_1"
|
||||
assert tool_calls[0].function.name == "get_weather"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_stream_wrapper_anext_does_not_block_event_loop_for_sync_iterators(
|
||||
logging_obj: Logging,
|
||||
):
|
||||
"""
|
||||
Regression test: __anext__ must not call blocking next() on a sync iterator on the
|
||||
event loop thread. This happens for some provider streams which are sync iterators
|
||||
but used in async contexts (e.g. boto3-style streaming).
|
||||
"""
|
||||
|
||||
class BlockingIterator:
|
||||
def __init__(self, chunks, delay_s: float):
|
||||
self._it = iter(chunks)
|
||||
self._delay_s = delay_s
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
time.sleep(self._delay_s) # simulate blocking I/O
|
||||
return next(self._it)
|
||||
|
||||
test_chunk = ModelResponseStream(
|
||||
id="chatcmpl-test",
|
||||
created=int(time.time()),
|
||||
model="test-model",
|
||||
object="chat.completion.chunk",
|
||||
system_fingerprint=None,
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(
|
||||
provider_specific_fields=None,
|
||||
content="hello",
|
||||
role="assistant",
|
||||
function_call=None,
|
||||
tool_calls=None,
|
||||
audio=None,
|
||||
),
|
||||
logprobs=None,
|
||||
)
|
||||
],
|
||||
provider_specific_fields={},
|
||||
usage=None,
|
||||
)
|
||||
|
||||
# Delay is intentionally > the wait_for timeout used to detect event loop blocking.
|
||||
wrapper = CustomStreamWrapper(
|
||||
completion_stream=BlockingIterator([test_chunk], delay_s=0.3),
|
||||
model="test-model",
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider="cached_response",
|
||||
)
|
||||
|
||||
tick_event = asyncio.Event()
|
||||
|
||||
async def background_tick():
|
||||
await asyncio.sleep(0.05)
|
||||
tick_event.set()
|
||||
|
||||
bg_task = asyncio.create_task(background_tick())
|
||||
anext_task = asyncio.create_task(wrapper.__anext__())
|
||||
try:
|
||||
# If the event loop is blocked by a sync next(), this will time out.
|
||||
await asyncio.wait_for(tick_event.wait(), timeout=0.15)
|
||||
|
||||
out = await asyncio.wait_for(anext_task, timeout=2.0)
|
||||
assert isinstance(out, ModelResponseStream)
|
||||
finally:
|
||||
if not anext_task.done():
|
||||
anext_task.cancel()
|
||||
await bg_task
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue