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())