From 475ff786d8ec1df3f7ec8416ca667891fe254fc0 Mon Sep 17 00:00:00 2001 From: wasimysaid Date: Fri, 19 Jun 2026 17:07:47 +0200 Subject: [PATCH] Address Gemma stream review comments --- studio/backend/core/tool_healing.py | 1 + studio/backend/routes/inference.py | 77 ++++++++++++++++++- .../tests/test_llama_route_timeouts.py | 68 ++++++++++++++++ studio/backend/tests/test_tool_xml_strip.py | 8 ++ 4 files changed, 152 insertions(+), 2 deletions(-) diff --git a/studio/backend/core/tool_healing.py b/studio/backend/core/tool_healing.py index c419d8e21c..e13c56a6c3 100644 --- a/studio/backend/core/tool_healing.py +++ b/studio/backend/core/tool_healing.py @@ -135,6 +135,7 @@ def _quote_gemma_object_keys(src: str) -> str: def _gemma_arguments_to_json(args_src: str) -> dict: + """Parse Gemma 4's native call:name{key:value} argument object.""" args_src = args_src.strip() if not args_src: return {} diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index f40348a9ae..9c0937907d 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -697,6 +697,53 @@ def _set_stream_response_read_timeout( pass +_STREAM_DISCONNECT_POLL_TIMEOUT_S = 0.25 + + +class _CompatSameTaskTimeout: + """Same-task timeout fallback for Python versions before asyncio.timeout.""" + + def __init__(self, timeout_s: float): + self.timeout_s = timeout_s + self._task = None + self._handle = None + self._timed_out = False + self._cancelling = 0 + + async def __aenter__(self): + self._task = asyncio.current_task() + if self._task is None: + return self + if hasattr(self._task, "cancelling"): + self._cancelling = self._task.cancelling() + loop = asyncio.get_running_loop() + self._handle = loop.call_later(max(self.timeout_s, 0), self._cancel_task) + return self + + async def __aexit__(self, exc_type, exc, tb): + if self._handle is not None: + self._handle.cancel() + if exc_type is not None and issubclass(exc_type, asyncio.CancelledError): + if self._timed_out: + if self._task is not None and hasattr(self._task, "uncancel"): + if self._task.uncancel() > self._cancelling: + return None + raise asyncio.TimeoutError from exc + return None + + def _cancel_task(self) -> None: + self._timed_out = True + if self._task is not None: + self._task.cancel() + + +def _same_task_timeout(timeout_s: float): + timeout_ctx = getattr(asyncio, "timeout", None) + if timeout_ctx is not None: + return timeout_ctx(timeout_s) + return _CompatSameTaskTimeout(timeout_s) + + async def _preheader_cancelled(cancel_event = None, request: Optional[Request] = None) -> bool: if cancel_event is not None and cancel_event.is_set(): return True @@ -785,13 +832,31 @@ 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, remaining_s) + _set_stream_response_read_timeout(response, read_timeout_s) # Keep httpx/httpcore's AnyIO cancel scope in this task. # asyncio.wait_for would drive __anext__ in a child task. - async with asyncio.timeout(remaining_s): + async with _same_task_timeout(remaining_s): item = await async_iter.__anext__() else: + if ( + request is not None + and response is not None + and post_first_item_read_timeout_s is not None + and last_item_at is not None + ): + stall_remaining_s = post_first_item_read_timeout_s - ( + time.monotonic() - last_item_at + ) + 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), + ) item = await async_iter.__anext__() except asyncio.TimeoutError as exc: if waiting_first_item: @@ -805,6 +870,12 @@ async def _aiter_llama_stream_items( if now >= first_token_deadline: raise continue + if ( + request is not None + and post_first_item_read_timeout_s is not None + and now - last_item_at < post_first_item_read_timeout_s + ): + continue raise httpx.ReadTimeout("The model stopped producing tokens mid-response.") if ( last_item_at is None @@ -1291,6 +1362,7 @@ _TOOL_XML_RE = _re.compile( r"<(?:tool_call|function=[\w-]+)>.*?(?:|\Z)" r"|<\|tool_call>.*?(?:|\Z)" r"|" + r"|" r"|\s*\Z", _re.DOTALL, ) @@ -7400,6 +7472,7 @@ async def _responses_stream( async for raw_line in _aiter_llama_stream_items( lines_iter, cancel_event = disconnect_event, + request = request, first_token_deadline = first_token_deadline, response = resp, ): diff --git a/studio/backend/tests/test_llama_route_timeouts.py b/studio/backend/tests/test_llama_route_timeouts.py index dc6e13a3de..b0954619dc 100644 --- a/studio/backend/tests/test_llama_route_timeouts.py +++ b/studio/backend/tests/test_llama_route_timeouts.py @@ -5,6 +5,7 @@ import asyncio import os import sys import time +import threading from types import SimpleNamespace _backend = os.path.join(os.path.dirname(__file__), "..") @@ -69,6 +70,73 @@ def test_stream_first_item_deadline_does_not_hop_tasks(): asyncio.run(_run()) +def test_stream_first_item_deadline_uses_compat_timeout_without_task_hop(monkeypatch): + monkeypatch.setattr(inf_mod.asyncio, "timeout", None, raising = False) + + async def _run(): + outer_task = asyncio.current_task() + seen_tasks = [] + + class _One: + def __init__(self): + self.done = False + + async def __anext__(self): + seen_tasks.append(asyncio.current_task()) + if self.done: + raise StopAsyncIteration + self.done = True + return "data: {}" + + out = [] + async for item in inf_mod._aiter_llama_stream_items( + _One(), + first_token_deadline = time.monotonic() + 1, + ): + out.append(item) + + assert out == ["data: {}"] + assert seen_tasks == [outer_task, outer_task] + + asyncio.run(_run()) + + +def test_stream_wait_polls_disconnect_without_background_watcher(): + 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 + + class _SlowFirstItem: + async def __anext__(self): + await asyncio.sleep(0.02) + raise inf_mod.httpx.ReadTimeout("poll") + + started = time.monotonic() + async for _ in inf_mod._aiter_llama_stream_items( + _SlowFirstItem(), + cancel_event = cancel_event, + request = _Request(), + response = response, + first_token_deadline = started + 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 + ) + + asyncio.run(_run()) + + def test_preheader_send_cleanup_on_disconnect_and_cancel(): async def _run(cancel_parent): state = SimpleNamespace(disconnected = False, closed = False, cancelled = False) diff --git a/studio/backend/tests/test_tool_xml_strip.py b/studio/backend/tests/test_tool_xml_strip.py index 2ba3310fbe..c2dc1fe8db 100644 --- a/studio/backend/tests/test_tool_xml_strip.py +++ b/studio/backend/tests/test_tool_xml_strip.py @@ -125,6 +125,14 @@ def test_strips_orphan_closing_tag(): # Mid-string intentionally preserved (see preserve test). +def test_strips_gemma_native_orphan_closing_tag(): + cleaned = _TOOL_XML_RE.sub("", "Tool call drained.Visible tail.") + + assert "" not in cleaned + assert "Tool call drained." in cleaned + assert "Visible tail." in cleaned + + # ── Tail-only (PR #5735 follow-up) ───────────────────