diff --git a/studio/backend/core/inference/tool_call_parser.py b/studio/backend/core/inference/tool_call_parser.py index 78e8f14103..4f5310567f 100644 --- a/studio/backend/core/inference/tool_call_parser.py +++ b/studio/backend/core/inference/tool_call_parser.py @@ -7,21 +7,24 @@ Tolerates missing closing tags in either ``{json}`` or ``v...`` shape. """ -import json -import re - from core.tool_healing import ( _TC_END_TAG_RE, _TC_FUNC_CLOSE_RE, _TC_FUNC_START_RE, + _TC_GEMMA_END_TAG_RE, _TC_GEMMA_START_RE, _TC_JSON_START_RE, _TC_PARAM_CLOSE_RE, _TC_PARAM_START_RE, _TOOL_ALL_PATS, _TOOL_CLOSED_PATS, + _FUNC_CLOSE_TAG, + _PARAM_CLOSE_TAG, _balanced_brace_end, _gemma_arguments_to_json, + _inside_open_parameter, + parse_tool_calls_from_text, + strip_tool_call_markup as strip_tool_markup, ) @@ -75,201 +78,6 @@ RAG_SEARCH_CAP_NUDGE = ( ) -_TC_GEMMA_END_TAG_RE = re.compile(r"") -_PARAM_CLOSE_TAG = "" -_FUNC_CLOSE_TAG = "" - - -def _inside_open_parameter(content: str, pos: int) -> bool: - """Return True when ``pos`` falls inside an unclosed parameter value.""" - last_param_start = -1 - for match in _TC_PARAM_START_RE.finditer(content, 0, pos): - last_param_start = match.start() - if last_param_start < 0: - return False - last_param_close = content.rfind(_PARAM_CLOSE_TAG, 0, pos) - last_func_close = content.rfind(_FUNC_CLOSE_TAG, 0, pos) - return last_param_start > max(last_param_close, last_func_close) - - -def strip_tool_markup(text: str, *, final: bool = False) -> str: - """Strip tool-call XML from streamed text. - - ``final=False`` only removes closed pairs (used during streaming so - in-progress XML stays buffered). ``final=True`` also removes a - trailing unclosed run and trims the result. - """ - pats = _TOOL_ALL_PATS if final else _TOOL_CLOSED_PATS - for pat in pats: - text = pat.sub("", text) - return text.strip() if final else text - - -def parse_tool_calls_from_text( - content: str, - *, - id_offset: int = 0, - allow_incomplete: bool = True, -) -> list[dict]: - """Parse OpenAI-format ``tool_calls`` from model text. - - Returns a list of ``{"id", "type", "function": {"name", "arguments"}}`` - dicts. ``arguments`` is always a JSON string so callers can hand it - straight back into an OpenAI-style response. - - Handles three shapes: - - - JSON inside ```` tags: - ``{"name":"web_search","arguments":{"query":"..."}}`` - - Gemma 4 native call blocks: - ``<|tool_call>call:web_search{query:"..." }`` - - XML-style function blocks: - ``v`` - - ``allow_incomplete=True`` keeps the historical healing behavior for - missing closing tags. ``allow_incomplete=False`` accepts only - well-formed wrappers so disabled Auto-Heal can still parse valid - local tool protocol without repairing truncated output. - """ - tool_calls: list[dict] = [] - - # Pattern 1: {json}. Balanced-brace scan, skipping braces in - # JSON strings. - for m in _TC_JSON_START_RE.finditer(content): - brace_start = m.end() - 1 # opening { - i = _balanced_brace_end(content, brace_start) - if i < 0: - continue - if not allow_incomplete: - tail_after_json = content[i + 1 :].lstrip() - if _TC_END_TAG_RE.match(tail_after_json) is None: - continue - json_str = content[brace_start : i + 1] - try: - obj = json.loads(json_str) - tc = { - "id": f"call_{id_offset + len(tool_calls)}", - "type": "function", - "function": { - "name": obj.get("name", ""), - "arguments": obj.get("arguments", {}), - }, - } - if isinstance(tc["function"]["arguments"], dict): - tc["function"]["arguments"] = json.dumps(tc["function"]["arguments"]) - tool_calls.append(tc) - except (json.JSONDecodeError, ValueError): - pass - - # Pattern 1b: Gemma 4 native call block: - # <|tool_call>call:terminal{command:"ls"} - for m in _TC_GEMMA_START_RE.finditer(content): - brace_start = m.end() - 1 - i = _balanced_brace_end(content, brace_start) - if i < 0: - continue - if not allow_incomplete: - tail_after_json = content[i + 1 :].lstrip() - if _TC_GEMMA_END_TAG_RE.match(tail_after_json) is None: - continue - try: - tool_calls.append( - { - "id": f"call_{id_offset + len(tool_calls)}", - "type": "function", - "function": { - "name": m.group(1), - "arguments": json.dumps(_gemma_arguments_to_json(content[m.end() : i])), - }, - } - ) - except (json.JSONDecodeError, ValueError): - pass - - # Pattern 2: v... -- closing tags optional; - # isn't a body boundary since code values can contain it. - if not tool_calls: - func_starts = [ - fm - for fm in _TC_FUNC_START_RE.finditer(content) - if not _inside_open_parameter(content, fm.start()) - ] - for idx, fm in enumerate(func_starts): - func_name = fm.group(1) - body_start = fm.end() - next_func = func_starts[idx + 1].start() if idx + 1 < len(func_starts) else len(content) - end_tag = _TC_END_TAG_RE.search(content[body_start:]) - if end_tag: - body_end = body_start + end_tag.start() - else: - body_end = len(content) - body_end = min(body_end, next_func) - body = content[body_start:body_end] - if not allow_incomplete: - # Bound the body at the closing tag rather than - # the end of the response, so a complete call followed by - # trailing prose is still accepted (matching the JSON-style - # path, which already tolerates trailing text). - # rfind picks the last , so a literal - # inside a code parameter value stays in the body. - close_idx = body.rfind(_FUNC_CLOSE_TAG) - if close_idx < 0: - continue - body = body[:close_idx] - else: - body = _TC_FUNC_CLOSE_RE.sub("", body) - - arguments: dict = {} - param_starts = list(_TC_PARAM_START_RE.finditer(body)) - if len(param_starts) == 1: - # Single param: take everything to body end so an embedded - # in code strings is preserved. - pm = param_starts[0] - val = body[pm.end() :] - if not allow_incomplete: - stripped_val = val.rstrip() - if not stripped_val.endswith(_PARAM_CLOSE_TAG): - continue - val = stripped_val[: -len(_PARAM_CLOSE_TAG)] - else: - val = _TC_PARAM_CLOSE_RE.sub("", val) - arguments[pm.group(1)] = val.strip() - else: - valid_params = True - for pidx, pm in enumerate(param_starts): - param_name = pm.group(1) - val_start = pm.end() - next_param = ( - param_starts[pidx + 1].start() - if pidx + 1 < len(param_starts) - else len(body) - ) - val = body[val_start:next_param] - if not allow_incomplete: - stripped_val = val.rstrip() - if not stripped_val.endswith(_PARAM_CLOSE_TAG): - valid_params = False - break - val = stripped_val[: -len(_PARAM_CLOSE_TAG)] - else: - val = _TC_PARAM_CLOSE_RE.sub("", val) - arguments[param_name] = val.strip() - if not valid_params: - continue - - tc = { - "id": f"call_{id_offset + len(tool_calls)}", - "type": "function", - "function": { - "name": func_name, - "arguments": json.dumps(arguments), - }, - } - tool_calls.append(tc) - - return tool_calls - - def has_tool_signal(text: str) -> bool: """Return True if ``text`` contains any tool-call XML signal.""" return any(s in text for s in TOOL_XML_SIGNALS) diff --git a/studio/backend/core/tool_healing.py b/studio/backend/core/tool_healing.py index e13c56a6c3..855755e483 100644 --- a/studio/backend/core/tool_healing.py +++ b/studio/backend/core/tool_healing.py @@ -29,10 +29,13 @@ _TC_JSON_START_RE = re.compile(r"\s*\{") _TC_GEMMA_START_RE = re.compile(r"<\|tool_call>call:([\w-]+)\s*\{") _TC_FUNC_START_RE = re.compile(r"\s*") _TC_END_TAG_RE = re.compile(r"") +_TC_GEMMA_END_TAG_RE = re.compile(r"") _TC_FUNC_CLOSE_RE = re.compile(r"\s*\s*$") _TC_PARAM_START_RE = re.compile(r"\s*") _TC_PARAM_CLOSE_RE = re.compile(r"\s*\s*$") _GEMMA_QUOTE = '<|"|>' +_PARAM_CLOSE_TAG = "" +_FUNC_CLOSE_TAG = "" def _balanced_brace_end(content: str, brace_start: int) -> int: @@ -145,52 +148,72 @@ def _gemma_arguments_to_json(args_src: str) -> dict: return json.loads(src) -def parse_tool_calls_from_text(content: str) -> list[dict]: - """ - Parse tool calls from XML markup in content text. +def _inside_open_parameter(content: str, pos: int) -> bool: + """Return True when ``pos`` falls inside an unclosed parameter value.""" + last_param_start = -1 + for match in _TC_PARAM_START_RE.finditer(content, 0, pos): + last_param_start = match.start() + if last_param_start < 0: + return False + last_param_close = content.rfind(_PARAM_CLOSE_TAG, 0, pos) + last_func_close = content.rfind(_FUNC_CLOSE_TAG, 0, pos) + return last_param_start > max(last_param_close, last_func_close) + + +def parse_tool_calls_from_text( + content: str, + *, + id_offset: int = 0, + allow_incomplete: bool = True, +) -> list[dict]: + """Parse OpenAI-format tool calls from model text. Handles formats like: {"name":"web_search","arguments":{"query":"..."}} <|tool_call>call:web_search{query:"..."} ... - Closing tags (, , ) are all - optional since models frequently omit them. """ - tool_calls = [] + tool_calls: list[dict] = [] - # Pattern 1: JSON inside tags. Balanced-brace extraction that - # skips braces inside JSON strings. for m in _TC_JSON_START_RE.finditer(content): - brace_start = m.end() - 1 # position of the opening { + brace_start = m.end() - 1 i = _balanced_brace_end(content, brace_start) - if i >= 0: - json_str = content[brace_start : i + 1] - try: - obj = json.loads(json_str) - tc = { - "id": f"call_{len(tool_calls)}", - "type": "function", - "function": { - "name": obj.get("name", ""), - "arguments": obj.get("arguments", {}), - }, - } - if isinstance(tc["function"]["arguments"], dict): - tc["function"]["arguments"] = json.dumps(tc["function"]["arguments"]) - tool_calls.append(tc) - except (json.JSONDecodeError, ValueError): - pass + if i < 0: + continue + if not allow_incomplete: + tail_after_json = content[i + 1 :].lstrip() + if _TC_END_TAG_RE.match(tail_after_json) is None: + continue + json_str = content[brace_start : i + 1] + try: + obj = json.loads(json_str) + tc = { + "id": f"call_{id_offset + len(tool_calls)}", + "type": "function", + "function": { + "name": obj.get("name", ""), + "arguments": obj.get("arguments", {}), + }, + } + if isinstance(tc["function"]["arguments"], dict): + tc["function"]["arguments"] = json.dumps(tc["function"]["arguments"]) + tool_calls.append(tc) + except (json.JSONDecodeError, ValueError): + pass - # Pattern 1b: Gemma 4 native <|tool_call>call:name{key:value}. for m in _TC_GEMMA_START_RE.finditer(content): brace_start = m.end() - 1 i = _balanced_brace_end(content, brace_start) if i < 0: continue + if not allow_incomplete: + tail_after_json = content[i + 1 :].lstrip() + if _TC_GEMMA_END_TAG_RE.match(tail_after_json) is None: + continue try: tool_calls.append( { - "id": f"call_{len(tool_calls)}", + "id": f"call_{id_offset + len(tool_calls)}", "type": "function", "function": { "name": m.group(1), @@ -201,17 +224,15 @@ def parse_tool_calls_from_text(content: str) -> list[dict]: except (json.JSONDecodeError, ValueError): pass - # Pattern 2: XML-style value - # All closing tags optional; models frequently omit them. if not tool_calls: - # Step 1: Find positions and extract bodies. Use only - # or the next - # can appear in code values); trim a trailing afterwards. - func_starts = list(_TC_FUNC_START_RE.finditer(content)) + func_starts = [ + fm + for fm in _TC_FUNC_START_RE.finditer(content) + if not _inside_open_parameter(content, fm.start()) + ] for idx, fm in enumerate(func_starts): func_name = fm.group(1) body_start = fm.end() - # Boundaries: next next_func = func_starts[idx + 1].start() if idx + 1 < len(func_starts) else len(content) end_tag = _TC_END_TAG_RE.search(content[body_start:]) if end_tag: @@ -220,36 +241,52 @@ def parse_tool_calls_from_text(content: str) -> list[dict]: body_end = len(content) body_end = min(body_end, next_func) body = content[body_start:body_end] - body = _TC_FUNC_CLOSE_RE.sub("", body) # trim closing + if not allow_incomplete: + close_idx = body.rfind(_FUNC_CLOSE_TAG) + if close_idx < 0: + continue + body = body[:close_idx] + else: + body = _TC_FUNC_CLOSE_RE.sub("", body) - # Step 2: Extract parameters from body. For single-parameter - # functions, use body end as the only boundary to avoid matching - # inside code strings. - arguments = {} + arguments: dict = {} param_starts = list(_TC_PARAM_START_RE.finditer(body)) if len(param_starts) == 1: - # Value is everything after the tag to end of body, less a - # trailing . pm = param_starts[0] val = body[pm.end() :] - val = _TC_PARAM_CLOSE_RE.sub("", val) + if not allow_incomplete: + stripped_val = val.rstrip() + if not stripped_val.endswith(_PARAM_CLOSE_TAG): + continue + val = stripped_val[: -len(_PARAM_CLOSE_TAG)] + else: + val = _TC_PARAM_CLOSE_RE.sub("", val) arguments[pm.group(1)] = val.strip() else: + valid_params = True for pidx, pm in enumerate(param_starts): param_name = pm.group(1) val_start = pm.end() - # Value ends at next + if not allow_incomplete: + stripped_val = val.rstrip() + if not stripped_val.endswith(_PARAM_CLOSE_TAG): + valid_params = False + break + val = stripped_val[: -len(_PARAM_CLOSE_TAG)] + else: + val = _TC_PARAM_CLOSE_RE.sub("", val) arguments[param_name] = val.strip() + if not valid_params: + continue tc = { - "id": f"call_{len(tool_calls)}", + "id": f"call_{id_offset + len(tool_calls)}", "type": "function", "function": { "name": func_name, diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index f0cd4d3b3e..5fd91bab56 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -1313,6 +1313,14 @@ async def _await_disconnect_then_cancel(request, cancel_event) -> None: return +async def _stop_local_disconnect_cancel_watcher(watcher) -> None: + watcher.cancel() + try: + await watcher + except (asyncio.CancelledError, Exception): + pass + + # Centralized local/server tool nudge. Keep render_html guidance gated to turns # where the canvas tool is actually present in the tool schema; otherwise # small local models can hallucinate a missing tool call instead of following @@ -5245,11 +5253,7 @@ async def openai_chat_completions( error_chunk = _openai_stream_error_chunk(e) yield f"data: {json.dumps(error_chunk)}\n\n" finally: - disconnect_watcher.cancel() - try: - await disconnect_watcher - except (asyncio.CancelledError, Exception): - pass + await _stop_local_disconnect_cancel_watcher(disconnect_watcher) if gen is not None: try: gen.close() @@ -5412,11 +5416,7 @@ async def openai_chat_completions( error_chunk = _openai_stream_error_chunk(e) yield f"data: {json.dumps(error_chunk)}\n\n" finally: - disconnect_watcher.cancel() - try: - await disconnect_watcher - except (asyncio.CancelledError, Exception): - pass + await _stop_local_disconnect_cancel_watcher(disconnect_watcher) _tracker.__exit__(None, None, None) return _SameTaskStreamingResponse( @@ -5863,11 +5863,7 @@ async def openai_chat_completions( } yield f"data: {json.dumps(error_chunk)}\n\n" finally: - disconnect_watcher.cancel() - try: - await disconnect_watcher - except (asyncio.CancelledError, Exception): - pass + await _stop_local_disconnect_cancel_watcher(disconnect_watcher) if gen is not None: try: gen.close() @@ -6097,11 +6093,7 @@ async def openai_chat_completions( } yield f"data: {json.dumps(error_chunk)}\n\n" finally: - disconnect_watcher.cancel() - try: - await disconnect_watcher - except (asyncio.CancelledError, Exception): - pass + await _stop_local_disconnect_cancel_watcher(disconnect_watcher) _tracker.__exit__(None, None, None) return _SameTaskStreamingResponse( diff --git a/studio/backend/tests/test_openai_tool_passthrough.py b/studio/backend/tests/test_openai_tool_passthrough.py index 6d5d9badc1..190bfd6a36 100644 --- a/studio/backend/tests/test_openai_tool_passthrough.py +++ b/studio/backend/tests/test_openai_tool_passthrough.py @@ -1264,6 +1264,61 @@ class TestGgufVisionToolRouting: pass return payloads + def _run_gguf_case( + self, + monkeypatch, + *, + generate = None, + tool_generate = None, + payload_kwargs = None, + backend_kwargs = None, + ): + import routes.inference as inf_mod + + reset_tool_policy() + + def _plain(**_kwargs): + raise AssertionError("plain GGUF path should not be used") + + backend_data = { + "is_loaded": True, + "is_vision": False, + "supports_tools": tool_generate is not None, + "supports_reasoning": True, + "reasoning_always_on": True, + "_is_audio": False, + "model_identifier": "test-gguf", + "context_length": 4096, + "generate_chat_completion": generate or _plain, + } + if tool_generate is not None: + backend_data["generate_chat_completion_with_tools"] = tool_generate + if backend_kwargs: + backend_data.update(backend_kwargs) + backend = SimpleNamespace(**backend_data) + + monitor = ApiMonitor(max_entries = 3) + monkeypatch.setattr(inf_mod, "api_monitor", monitor) + monkeypatch.setattr(inf_mod, "get_llama_cpp_backend", lambda: backend) + + request_data = { + "model": "default", + "messages": [{"role": "user", "content": "hi"}], + } + if payload_kwargs: + request_data.update(payload_kwargs) + payload = ChatCompletionRequest(**request_data) + response = self._drive( + openai_chat_completions(payload, request = self._Request(), current_subject = "test") + ) + result = SimpleNamespace(response = response, monitor = monitor, backend = backend) + if request_data.get("stream"): + result.chunks = self._consume_response(response) + result.payloads = self._sse_payloads(result.chunks) + else: + result.body = json.loads(response.body) + return result + def test_image_request_with_enabled_tools_enters_gguf_tool_loop(self, monkeypatch): import routes.inference as inf_mod @@ -1410,10 +1465,6 @@ class TestGgufVisionToolRouting: assert monitor.active_count() == 0 def test_standard_gguf_stream_splits_reasoning_content(self, monkeypatch): - import routes.inference as inf_mod - - reset_tool_policy() - def _generate(**_kwargs): yield "plan" @@ -1425,44 +1476,20 @@ class TestGgufVisionToolRouting: "finish_reason": "stop", } - backend = SimpleNamespace( - is_loaded = True, - is_vision = False, - supports_tools = False, - supports_reasoning = True, - reasoning_always_on = True, - _is_audio = False, - model_identifier = "test-gguf", - context_length = 4096, - generate_chat_completion = _generate, + result = self._run_gguf_case( + monkeypatch, + generate = _generate, + payload_kwargs = {"stream": True}, ) - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "get_llama_cpp_backend", lambda: backend) - - payload = ChatCompletionRequest( - model = "default", - stream = True, - messages = [{"role": "user", "content": "hi"}], - ) - - response = self._drive( - openai_chat_completions(payload, request = self._Request(), current_subject = "test") - ) - payloads = self._sse_payloads(self._consume_response(response)) - deltas = [p["choices"][0].get("delta", {}) for p in payloads if p.get("choices")] + deltas = [p["choices"][0].get("delta", {}) for p in result.payloads if p.get("choices")] assert "".join(d.get("reasoning_content", "") for d in deltas) == "plan" assert "".join(d.get("content", "") for d in deltas) == "visible" assert all("" not in d.get("content", "") for d in deltas) - [entry] = monitor.snapshot() + [entry] = result.monitor.snapshot() assert entry["reply"] == "visible" def test_reasoning_capable_gguf_stream_splits_reasoning_by_default(self, monkeypatch): - import routes.inference as inf_mod - - reset_tool_policy() - def _generate(**_kwargs): yield "planvisible" yield { @@ -1471,43 +1498,20 @@ class TestGgufVisionToolRouting: "finish_reason": "stop", } - backend = SimpleNamespace( - is_loaded = True, - is_vision = False, - supports_tools = False, - supports_reasoning = True, - reasoning_always_on = False, - _is_audio = False, - model_identifier = "test-gguf", - context_length = 4096, - generate_chat_completion = _generate, + result = self._run_gguf_case( + monkeypatch, + generate = _generate, + payload_kwargs = {"stream": True}, + backend_kwargs = {"reasoning_always_on": False}, ) - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "get_llama_cpp_backend", lambda: backend) - - payload = ChatCompletionRequest( - model = "default", - stream = True, - messages = [{"role": "user", "content": "hi"}], - ) - - response = self._drive( - openai_chat_completions(payload, request = self._Request(), current_subject = "test") - ) - payloads = self._sse_payloads(self._consume_response(response)) - deltas = [p["choices"][0].get("delta", {}) for p in payloads if p.get("choices")] + deltas = [p["choices"][0].get("delta", {}) for p in result.payloads if p.get("choices")] assert "".join(d.get("reasoning_content", "") for d in deltas) == "plan" assert "".join(d.get("content", "") for d in deltas) == "visible" - [entry] = monitor.snapshot() + [entry] = result.monitor.snapshot() assert entry["reply"] == "visible" def test_reasoning_capable_gguf_stream_sanitizes_think_tags_when_disabled(self, monkeypatch): - import routes.inference as inf_mod - - reset_tool_policy() - def _generate(**_kwargs): yield "leakedvisible" yield { @@ -1516,48 +1520,21 @@ class TestGgufVisionToolRouting: "finish_reason": "stop", } - backend = SimpleNamespace( - is_loaded = True, - is_vision = False, - supports_tools = False, - supports_reasoning = True, - reasoning_always_on = False, - _is_audio = False, - model_identifier = "test-gguf", - context_length = 4096, - generate_chat_completion = _generate, + result = self._run_gguf_case( + monkeypatch, + generate = _generate, + payload_kwargs = {"stream": True, "enable_thinking": False}, + backend_kwargs = {"reasoning_always_on": False}, ) - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "get_llama_cpp_backend", lambda: backend) - - payload = ChatCompletionRequest( - model = "default", - stream = True, - enable_thinking = False, - messages = [{"role": "user", "content": "hi"}], - ) - - response = self._drive( - openai_chat_completions(payload, request = self._Request(), current_subject = "test") - ) - payloads = self._sse_payloads(self._consume_response(response)) - deltas = [p["choices"][0].get("delta", {}) for p in payloads if p.get("choices")] + deltas = [p["choices"][0].get("delta", {}) for p in result.payloads if p.get("choices")] assert "".join(d.get("reasoning_content", "") for d in deltas) == "leaked" assert "".join(d.get("content", "") for d in deltas) == "visible" assert all("" not in d.get("content", "") for d in deltas) - [entry] = monitor.snapshot() + [entry] = result.monitor.snapshot() assert entry["reply"] == "visible" def test_gguf_tool_stream_splits_reasoning_and_strips_gemma_tool_marker(self, monkeypatch): - import routes.inference as inf_mod - - reset_tool_policy() - - def _plain(**_kwargs): - raise AssertionError("plain GGUF path should not be used") - def _tools(**_kwargs): yield { "type": "content", @@ -1569,48 +1546,26 @@ class TestGgufVisionToolRouting: "finish_reason": "stop", } - backend = SimpleNamespace( - is_loaded = True, - is_vision = False, - supports_tools = True, - supports_reasoning = True, - reasoning_always_on = True, - _is_audio = False, - model_identifier = "test-gguf", - context_length = 4096, - generate_chat_completion = _plain, - generate_chat_completion_with_tools = _tools, + result = self._run_gguf_case( + monkeypatch, + tool_generate = _tools, + payload_kwargs = { + "stream": True, + "enable_tools": True, + "enabled_tools": ["terminal"], + "messages": [{"role": "user", "content": "list files"}], + }, ) - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "get_llama_cpp_backend", lambda: backend) - - payload = ChatCompletionRequest( - model = "default", - stream = True, - enable_tools = True, - enabled_tools = ["terminal"], - messages = [{"role": "user", "content": "list files"}], - ) - - response = self._drive( - openai_chat_completions(payload, request = self._Request(), current_subject = "test") - ) - payloads = self._sse_payloads(self._consume_response(response)) - deltas = [p["choices"][0].get("delta", {}) for p in payloads if p.get("choices")] + deltas = [p["choices"][0].get("delta", {}) for p in result.payloads if p.get("choices")] assert "".join(d.get("reasoning_content", "") for d in deltas) == "plan" combined_content = "".join(d.get("content", "") for d in deltas) assert combined_content == "visible " assert "<|tool_call>" not in combined_content - [entry] = monitor.snapshot() + [entry] = result.monitor.snapshot() assert entry["reply"] == "visible " def test_non_streaming_gguf_splits_reasoning_content(self, monkeypatch): - import routes.inference as inf_mod - - reset_tool_policy() - def _generate(**_kwargs): yield "planvisible" yield { @@ -1619,35 +1574,13 @@ class TestGgufVisionToolRouting: "finish_reason": "stop", } - backend = SimpleNamespace( - is_loaded = True, - is_vision = False, - supports_tools = False, - supports_reasoning = True, - reasoning_always_on = True, - _is_audio = False, - model_identifier = "test-gguf", - context_length = 4096, - generate_chat_completion = _generate, - ) - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "get_llama_cpp_backend", lambda: backend) - - payload = ChatCompletionRequest( - model = "default", - messages = [{"role": "user", "content": "hi"}], - ) - - response = self._drive( - openai_chat_completions(payload, request = self._Request(), current_subject = "test") - ) - body = json.loads(response.body) + result = self._run_gguf_case(monkeypatch, generate = _generate) + body = result.body message = body["choices"][0]["message"] assert message["content"] == "visible" assert message["reasoning_content"] == "plan" - [entry] = monitor.snapshot() + [entry] = result.monitor.snapshot() assert entry["reply"] == "visible" def test_non_streaming_gguf_n_records_all_monitor_replies(self, monkeypatch): @@ -1812,6 +1745,61 @@ class TestApiMonitorProviderAndCompletionStreams: async def is_disconnected(self): return False + async def _run_passthrough_stream(self, monkeypatch, lines): + import routes.inference as inf_mod + + class Request: + async def is_disconnected(self): + return False + + async def fake_send(*_args, **_kwargs): + return httpx.Response(200, content = b"") + + async def fake_items(*_args, **_kwargs): + for line in lines: + yield line + + monitor = ApiMonitor(max_entries = 3) + monkeypatch.setattr(inf_mod, "api_monitor", monitor) + monkeypatch.setattr(inf_mod, "_send_stream_with_preheader_cancel", fake_send) + monkeypatch.setattr(inf_mod, "_aiter_llama_stream_items", fake_items) + monitor_id = monitor.start( + endpoint = "/v1/chat/completions", + method = "POST", + model = "gguf", + prompt = "hi", + ) + payload = ChatCompletionRequest( + model = "default", + messages = [ChatMessage(role = "user", content = "hi")], + stream = True, + tools = [ + { + "type": "function", + "function": { + "name": "lookup", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + ) + + response = await _openai_passthrough_stream( + Request(), + threading.Event(), + SimpleNamespace( + base_url = "http://llama.test", + context_length = 4096, + _request_reasoning_kwargs = lambda *_args, **_kwargs: None, + ), + payload, + "gguf", + "chatcmpl-test", + monitor_id = monitor_id, + ) + chunks = [chunk async for chunk in response.body_iterator] + return SimpleNamespace(chunks = chunks, body = "".join(chunks), monitor = monitor) + def test_external_non_streaming_json_updates_monitor(self, monkeypatch): async def _run(): import routes.inference as inf_mod @@ -2260,131 +2248,44 @@ class TestApiMonitorProviderAndCompletionStreams: def test_passthrough_stream_synthesizes_missing_finish_reason(self, monkeypatch): async def _run(): - import routes.inference as inf_mod - - class Request: - async def is_disconnected(self): - return False - - async def fake_send(*_args, **_kwargs): - return httpx.Response(200, content = b"") - - async def fake_items(*_args, **_kwargs): - yield 'data: {"id":"upstream","created":123,"model":"gguf","choices":[{"index":0,"delta":{"content":"hello"}}]}' - yield "data: [DONE]" - - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "_send_stream_with_preheader_cancel", fake_send) - monkeypatch.setattr(inf_mod, "_aiter_llama_stream_items", fake_items) - monitor_id = monitor.start( - endpoint = "/v1/chat/completions", - method = "POST", - model = "gguf", - prompt = "hi", - ) - payload = ChatCompletionRequest( - model = "default", - messages = [ChatMessage(role = "user", content = "hi")], - stream = True, - tools = [ - { - "type": "function", - "function": { - "name": "lookup", - "parameters": {"type": "object", "properties": {}}, - }, - } + result = await self._run_passthrough_stream( + monkeypatch, + [ + ( + 'data: {"id":"upstream","created":123,"model":"gguf",' + '"choices":[{"index":0,"delta":{"content":"hello"}}]}' + ), + "data: [DONE]", ], ) - - response = await _openai_passthrough_stream( - Request(), - threading.Event(), - SimpleNamespace( - base_url = "http://llama.test", - context_length = 4096, - _request_reasoning_kwargs = lambda *_args, **_kwargs: None, - ), - payload, - "gguf", - "chatcmpl-test", - monitor_id = monitor_id, - ) - chunks = [chunk async for chunk in response.body_iterator] - body = "".join(chunks) + body = result.body assert '"finish_reason":"stop"' in body.replace(" ", "") assert "data: [DONE]" in body - assert monitor.active_count() == 0 + assert result.monitor.active_count() == 0 asyncio.run(_run()) def test_passthrough_stream_synthesizes_tool_call_finish_reason(self, monkeypatch): async def _run(): - import routes.inference as inf_mod - - class Request: - async def is_disconnected(self): - return False - - async def fake_send(*_args, **_kwargs): - return httpx.Response(200, content = b"") - - async def fake_items(*_args, **_kwargs): - yield ( - 'data: {"id":"upstream","created":123,"model":"gguf",' - '"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,' - '"id":"call_1","type":"function","function":{"name":"lookup",' - '"arguments":"{}"}}]}}]}' - ) - yield "data: [DONE]" - - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "_send_stream_with_preheader_cancel", fake_send) - monkeypatch.setattr(inf_mod, "_aiter_llama_stream_items", fake_items) - monitor_id = monitor.start( - endpoint = "/v1/chat/completions", - method = "POST", - model = "gguf", - prompt = "hi", - ) - payload = ChatCompletionRequest( - model = "default", - messages = [ChatMessage(role = "user", content = "hi")], - stream = True, - tools = [ - { - "type": "function", - "function": { - "name": "lookup", - "parameters": {"type": "object", "properties": {}}, - }, - } + result = await self._run_passthrough_stream( + monkeypatch, + [ + ( + 'data: {"id":"upstream","created":123,"model":"gguf",' + '"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,' + '"id":"call_1","type":"function","function":{"name":"lookup",' + '"arguments":"{}"}}]}}]}' + ), + "data: [DONE]", ], ) - - response = await _openai_passthrough_stream( - Request(), - threading.Event(), - SimpleNamespace( - base_url = "http://llama.test", - context_length = 4096, - _request_reasoning_kwargs = lambda *_args, **_kwargs: None, - ), - payload, - "gguf", - "chatcmpl-test", - monitor_id = monitor_id, - ) - body = "".join([chunk async for chunk in response.body_iterator]) - compact = body.replace(" ", "") + compact = result.body.replace(" ", "") assert '"finish_reason":"tool_calls"' in compact assert '"finish_reason":"stop"' not in compact - assert "data: [DONE]" in body - assert monitor.active_count() == 0 + assert "data: [DONE]" in result.body + assert result.monitor.active_count() == 0 asyncio.run(_run()) @@ -2449,68 +2350,20 @@ class TestApiMonitorProviderAndCompletionStreams: def test_passthrough_clean_eof_finalizes_monitor(self, monkeypatch): async def _run(): - import routes.inference as inf_mod - - class Request: - async def is_disconnected(self): - return False - - async def fake_send(*_args, **_kwargs): - return httpx.Response(200, content = b"") - - async def fake_items(*_args, **_kwargs): - yield 'data: {"choices":[{"delta":{"content":"hello"}}]}' - - monitor = ApiMonitor(max_entries = 3) - monkeypatch.setattr(inf_mod, "api_monitor", monitor) - monkeypatch.setattr(inf_mod, "_send_stream_with_preheader_cancel", fake_send) - monkeypatch.setattr(inf_mod, "_aiter_llama_stream_items", fake_items) - monitor_id = monitor.start( - endpoint = "/v1/chat/completions", - method = "POST", - model = "gguf", - prompt = "hi", + result = await self._run_passthrough_stream( + monkeypatch, + ['data: {"choices":[{"delta":{"content":"hello"}}]}'], ) - payload = ChatCompletionRequest( - model = "default", - messages = [ChatMessage(role = "user", content = "hi")], - stream = True, - tools = [ - { - "type": "function", - "function": { - "name": "lookup", - "parameters": {"type": "object", "properties": {}}, - }, - } - ], - ) - - response = await _openai_passthrough_stream( - Request(), - threading.Event(), - SimpleNamespace( - base_url = "http://llama.test", - context_length = 4096, - _request_reasoning_kwargs = lambda *_args, **_kwargs: None, - ), - payload, - "gguf", - "chatcmpl-test", - monitor_id = monitor_id, - ) - chunks = [] - async for chunk in response.body_iterator: - chunks.append(chunk) + chunks = result.chunks assert chunks[0] == 'data: {"choices":[{"delta":{"content":"hello"}}]}\n\n' compact = "".join(chunks).replace(" ", "") assert '"finish_reason":"stop"' in compact assert chunks[-1] == "data: [DONE]\n\n" - [entry] = monitor.snapshot() + [entry] = result.monitor.snapshot() assert entry["status"] == "completed" assert entry["reply"] == "hello" - assert monitor.active_count() == 0 + assert result.monitor.active_count() == 0 asyncio.run(_run())