From 27e0228fdfff5ae26cde9664d306e3914b2da52b Mon Sep 17 00:00:00 2001 From: wasimysaid Date: Fri, 19 Jun 2026 20:40:29 +0200 Subject: [PATCH] Fix Studio passthrough cold stream timeout --- studio/backend/routes/inference.py | 10 +--- .../tests/test_llama_route_timeouts.py | 51 +++++++++++++------ 2 files changed, 38 insertions(+), 23 deletions(-) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 5fd91bab56..da8f98b1f5 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -845,11 +845,8 @@ async def _aiter_llama_stream_items( remaining_s = first_token_deadline - time.monotonic() if remaining_s <= 0: raise httpx.ReadTimeout("The model did not produce a first token in time.") - read_timeout_s = remaining_s - if request is not None: - read_timeout_s = min(read_timeout_s, _STREAM_DISCONNECT_POLL_TIMEOUT_S) if response is not None: - _set_stream_response_read_timeout(response, read_timeout_s) + _set_stream_response_read_timeout(response, remaining_s) # Keep httpx/httpcore's AnyIO cancel scope in this task. # asyncio.wait_for would drive __anext__ in a child task. async with _same_task_timeout(remaining_s): @@ -866,10 +863,7 @@ async def _aiter_llama_stream_items( ) if stall_remaining_s <= 0: raise httpx.ReadTimeout("The model stopped producing tokens mid-response.") - _set_stream_response_read_timeout( - response, - min(stall_remaining_s, _STREAM_DISCONNECT_POLL_TIMEOUT_S), - ) + _set_stream_response_read_timeout(response, stall_remaining_s) item = await async_iter.__anext__() except asyncio.TimeoutError as exc: if waiting_first_item: diff --git a/studio/backend/tests/test_llama_route_timeouts.py b/studio/backend/tests/test_llama_route_timeouts.py index b0954619dc..24866bd03e 100644 --- a/studio/backend/tests/test_llama_route_timeouts.py +++ b/studio/backend/tests/test_llama_route_timeouts.py @@ -101,38 +101,59 @@ def test_stream_first_item_deadline_uses_compat_timeout_without_task_hop(monkeyp asyncio.run(_run()) -def test_stream_wait_polls_disconnect_without_background_watcher(): +def test_stream_wait_stops_on_known_disconnect_before_read(): async def _run(): state = SimpleNamespace(disconnect_checks = 0) cancel_event = threading.Event() - response = SimpleNamespace(request = SimpleNamespace(extensions = {"timeout": {}})) class _Request: async def is_disconnected(self): state.disconnect_checks += 1 - return state.disconnect_checks >= 2 + return True - class _SlowFirstItem: + class _Unread: async def __anext__(self): - await asyncio.sleep(0.02) - raise inf_mod.httpx.ReadTimeout("poll") + raise AssertionError("stream should stop before reading upstream") - started = time.monotonic() async for _ in inf_mod._aiter_llama_stream_items( - _SlowFirstItem(), + _Unread(), cancel_event = cancel_event, request = _Request(), - response = response, - first_token_deadline = started + 1, + first_token_deadline = time.monotonic() + 1, ): raise AssertionError("stream should stop after disconnect") assert cancel_event.is_set() - assert state.disconnect_checks >= 2 - assert time.monotonic() - started < 0.5 - assert response.request.extensions["timeout"]["read"] <= ( - inf_mod._STREAM_DISCONNECT_POLL_TIMEOUT_S - ) + assert state.disconnect_checks == 1 + + asyncio.run(_run()) + + +def test_stream_wait_does_not_shorten_upstream_read_for_disconnect_poll(): + async def _run(): + response = SimpleNamespace(request = SimpleNamespace(extensions = {"timeout": {}})) + seen_read_timeouts = [] + + class _Request: + async def is_disconnected(self): + return False + + class _NoItem: + async def __anext__(self): + seen_read_timeouts.append(response.request.extensions["timeout"]["read"]) + raise StopAsyncIteration + + async for _ in inf_mod._aiter_llama_stream_items( + _NoItem(), + cancel_event = threading.Event(), + request = _Request(), + response = response, + first_token_deadline = time.monotonic() + 1, + ): + raise AssertionError("stream should end") + + assert seen_read_timeouts + assert seen_read_timeouts[0] > inf_mod._STREAM_DISCONNECT_POLL_TIMEOUT_S asyncio.run(_run())