diff --git a/studio/backend/core/tool_healing.py b/studio/backend/core/tool_healing.py index 855755e483..b0e74f20fa 100644 --- a/studio/backend/core/tool_healing.py +++ b/studio/backend/core/tool_healing.py @@ -16,6 +16,7 @@ import re _TOOL_CLOSED_PATS = [ re.compile(r".*?", re.DOTALL), re.compile(r"<\|tool_call>.*?", re.DOTALL), + re.compile(r""), re.compile(r".*?", re.DOTALL), ] _TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [ @@ -38,13 +39,13 @@ _PARAM_CLOSE_TAG = "" _FUNC_CLOSE_TAG = "" -def _balanced_brace_end(content: str, brace_start: int) -> int: +def _balanced_brace_end(content: str, brace_start: int, *, gemma_quotes: bool = False) -> int: depth = 0 i = brace_start in_string = False in_gemma_string = False while i < len(content): - if content.startswith(_GEMMA_QUOTE, i): + if gemma_quotes and not in_string and content.startswith(_GEMMA_QUOTE, i): in_gemma_string = not in_gemma_string i += len(_GEMMA_QUOTE) continue @@ -203,7 +204,7 @@ def parse_tool_calls_from_text( for m in _TC_GEMMA_START_RE.finditer(content): brace_start = m.end() - 1 - i = _balanced_brace_end(content, brace_start) + i = _balanced_brace_end(content, brace_start, gemma_quotes = True) if i < 0: continue if not allow_incomplete: diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 365b1c3f3a..f6feac72c5 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -752,6 +752,9 @@ class _SameTaskStreamingResponse(StreamingResponse): try: await self.stream_response(send) except OSError: + aclose = getattr(self.body_iterator, "aclose", None) + if aclose is not None: + await aclose() raise ClientDisconnect() if self.background is not None: await self.background() @@ -5136,6 +5139,21 @@ async def openai_chat_completions( _stream_usage = None _stream_timings = None _stream_finish = None + + def _flush_reasoning_extractor(): + final_reasoning, final_visible = reasoning_extractor.finish() + chunks = [] + if final_reasoning: + chunks.append( + _gguf_chat_delta_line( + ChoiceDelta(reasoning_content = final_reasoning) + ) + ) + if final_visible: + api_monitor.append_reply(monitor_id, final_visible) + chunks.append(_gguf_chat_delta_line(ChoiceDelta(content = final_visible))) + return chunks + while True: if cancel_event.is_set(): break @@ -5154,6 +5172,8 @@ async def openai_chat_completions( # cumulative cursor so the next assistant turn # streams cleanly. if not event["text"]: + for chunk in _flush_reasoning_extractor(): + yield chunk prev_text = "" reasoning_extractor = _new_chat_reasoning_extractor() # Emit tool status as a custom SSE event (including @@ -5169,6 +5189,8 @@ async def openai_chat_completions( if event["type"] in ("tool_start", "tool_end"): if event["type"] == "tool_start": + for chunk in _flush_reasoning_extractor(): + yield chunk prev_text = "" reasoning_extractor = _new_chat_reasoning_extractor() yield f"data: {json.dumps(event)}\n\n" @@ -5201,12 +5223,8 @@ async def openai_chat_completions( api_monitor.append_reply(monitor_id, visible_delta) yield _gguf_chat_delta_line(ChoiceDelta(content = visible_delta)) - final_reasoning, final_visible = reasoning_extractor.finish() - if final_reasoning: - yield _gguf_chat_delta_line(ChoiceDelta(reasoning_content = final_reasoning)) - if final_visible: - api_monitor.append_reply(monitor_id, final_visible) - yield _gguf_chat_delta_line(ChoiceDelta(content = final_visible)) + for chunk in _flush_reasoning_extractor(): + yield chunk final_chunk = ChatCompletionChunk( id = completion_id, diff --git a/studio/backend/tests/test_mcp_servers.py b/studio/backend/tests/test_mcp_servers.py index d1a9662b1d..81153e26ba 100644 --- a/studio/backend/tests/test_mcp_servers.py +++ b/studio/backend/tests/test_mcp_servers.py @@ -431,6 +431,13 @@ def test_tool_healing_strip_handles_gemma_native_tool_call(): assert out == "before after" +def test_tool_healing_strip_handles_gemma_close_only_marker(): + from core.tool_healing import strip_tool_call_markup + + assert strip_tool_call_markup("before after") == "before after" + assert strip_tool_call_markup("before after", final = True) == "before after" + + def test_tool_healing_parser_handles_gemma_native_windows_path(): from core.tool_healing import parse_tool_calls_from_text import json as _json @@ -443,6 +450,18 @@ def test_tool_healing_parser_handles_gemma_native_windows_path(): assert _json.loads(calls[0]["function"]["arguments"]) == {"path": r"C:\Users\wasim\repo"} +def test_tool_healing_json_parser_preserves_literal_gemma_quote_token(): + from core.tool_healing import parse_tool_calls_from_text + import json as _json + + text = "" + _json.dumps( + {"name": "python", "arguments": {"code": "print('<|\"|>')"}} + ) + "" + calls = parse_tool_calls_from_text(text, allow_incomplete = False) + assert len(calls) == 1 + assert _json.loads(calls[0]["function"]["arguments"]) == {"code": "print('<|\"|>')"} + + def test_gguf_allow_list_blocks_unadvertised_tool(monkeypatch): """A tool call not in the per-request list must be refused by the GGUF agentic loop (mirroring the safetensors path).""" diff --git a/studio/backend/tests/test_openai_tool_passthrough.py b/studio/backend/tests/test_openai_tool_passthrough.py index 1a649f09fe..aaef9e4dcc 100644 --- a/studio/backend/tests/test_openai_tool_passthrough.py +++ b/studio/backend/tests/test_openai_tool_passthrough.py @@ -1565,6 +1565,33 @@ class TestGgufVisionToolRouting: [entry] = result.monitor.snapshot() assert entry["reply"] == "visible " + def test_gguf_tool_stream_flushes_held_text_before_status_reset(self, monkeypatch): + def _tools(**_kwargs): + yield {"type": "content", "text": "answer <"} + yield {"type": "status", "text": ""} + yield { + "type": "metadata", + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + "finish_reason": "stop", + } + + result = self._run_gguf_case( + monkeypatch, + tool_generate = _tools, + payload_kwargs = { + "stream": True, + "enable_tools": True, + "enabled_tools": ["terminal"], + "messages": [{"role": "user", "content": "say literal"}], + }, + ) + deltas = [p["choices"][0].get("delta", {}) for p in result.payloads if p.get("choices")] + + combined_content = "".join(d.get("content", "") for d in deltas) + assert combined_content == "answer <" + [entry] = result.monitor.snapshot() + assert entry["reply"] == "answer <" + def test_non_streaming_gguf_splits_reasoning_content(self, monkeypatch): def _generate(**_kwargs): yield "planvisible" diff --git a/tests/studio/test_stream_cancel_registration_timing.py b/tests/studio/test_stream_cancel_registration_timing.py index bc1688fef4..09e30bb730 100644 --- a/tests/studio/test_stream_cancel_registration_timing.py +++ b/tests/studio/test_stream_cancel_registration_timing.py @@ -275,6 +275,23 @@ def _load_registry_module(): return mod +def _load_same_task_response_module(): + for n in _TREE.body: + if isinstance(n, ast.ClassDef) and n.name == "_SameTaskStreamingResponse": + source = ast.get_source_segment(SRC, n) + break + else: + raise AssertionError("_SameTaskStreamingResponse missing") + mod = {} + exec( + "class StreamingResponse: pass\n" + "class ClientDisconnect(Exception): pass\n" + + source, + mod, + ) + return mod + + def _make_stream(tracker, raise_exc): async def gen(): try: @@ -379,6 +396,39 @@ def test_finally_cleanup_on_aclose(): assert "sid-abort" not in m["_CANCEL_REGISTRY"] +def test_same_task_response_closes_body_iterator_on_send_disconnect(): + m = _load_same_task_response_module() + closed = False + + async def body(): + nonlocal closed + try: + yield "data: first\n\n" + finally: + closed = True + + async def run(): + agen = body() + await agen.__anext__() + response = m["_SameTaskStreamingResponse"].__new__(m["_SameTaskStreamingResponse"]) + response.body_iterator = agen + response.background = None + + async def stream_response(_send): + raise OSError("client disconnected") + + response.stream_response = stream_response + try: + await response({}, None, lambda _message: None) + except m["ClientDisconnect"]: + pass + else: + raise AssertionError("expected ClientDisconnect") + + asyncio.run(run()) + assert closed + + def test_preset_cancel_event_exits_cleanly_with_done(): # Pending-replay: a stashed cancel pre-set cancel_event. The loop must break # cleanly with final_chunk + [DONE], not propagate GeneratorExit from the GGUF wrapper.