Compare commits
33 commits
main
...
fix-gemma4
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
35e18e7680 | ||
|
|
2d72bc0494 |
||
|
|
75b65b4834 | ||
|
|
b0dbe43867 | ||
|
|
d376756e9c | ||
|
|
e1fdb140b4 | ||
|
|
6618213da5 | ||
|
|
1e58c3707d | ||
|
|
520df9fe9d | ||
|
|
11564256b9 | ||
|
|
eca61f7f5e | ||
|
|
0083826dde | ||
|
|
b3e244d658 | ||
|
|
491586c556 | ||
|
|
21edbb94a3 | ||
|
|
4b68cac412 | ||
|
|
2b7360176b | ||
|
|
3a2c1c2b22 |
||
|
|
67ccd78440 | ||
|
|
db8d03927b | ||
|
|
6a57d3795a | ||
|
|
3ca278cb79 | ||
|
|
5c4e7b5365 | ||
|
|
27e0228fdf | ||
|
|
5ba19c2977 | ||
|
|
b53015be85 | ||
|
|
bd1e5eabe5 | ||
|
|
ede6a2bcee | ||
|
|
475ff786d8 | ||
|
|
4c877cc7d8 | ||
|
|
6705053e9e | ||
|
|
05518f2769 | ||
|
|
165f4838e3 |
14 changed files with 2081 additions and 430 deletions
|
|
@ -7,25 +7,31 @@ Tolerates missing closing tags in either ``<tool_call>{json}</tool_call>``
|
||||||
or ``<function=name><parameter=k>v...`` shape.
|
or ``<function=name><parameter=k>v...`` shape.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
from core import tool_healing as _tool_healing
|
||||||
import re
|
|
||||||
|
|
||||||
|
|
||||||
# _TOOL_CLOSED_PATS: closed pairs only. _TOOL_ALL_PATS: also trailing unclosed
|
_TOOL_ALL_PATS = _tool_healing._TOOL_ALL_PATS
|
||||||
# runs so truncated tails don't leak markup. The [\w-] name set matches OpenAI's
|
|
||||||
# so hyphenated MCP tool names (mcp__srv__list-issues) parse like built-ins.
|
|
||||||
_TOOL_CLOSED_PATS = [
|
def parse_tool_calls_from_text(
|
||||||
re.compile(r"<tool_call>.*?</tool_call>", re.DOTALL),
|
content: str,
|
||||||
re.compile(r"<function=[\w-]+>.*?</function>", re.DOTALL),
|
*,
|
||||||
]
|
id_offset: int = 0,
|
||||||
_TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
allow_incomplete: bool = True,
|
||||||
re.compile(r"<tool_call>.*$", re.DOTALL),
|
) -> list[dict]:
|
||||||
re.compile(r"<function=[\w-]+>.*$", re.DOTALL),
|
return _tool_healing.parse_tool_calls_from_text(
|
||||||
]
|
content,
|
||||||
|
id_offset = id_offset,
|
||||||
|
allow_incomplete = allow_incomplete,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def strip_tool_markup(text: str, *, final: bool = False) -> str:
|
||||||
|
return _tool_healing.strip_tool_call_markup(text, final = final)
|
||||||
|
|
||||||
|
|
||||||
# Prefixes the streaming buffer watches for to gate in-progress text.
|
# Prefixes the streaming buffer watches for to gate in-progress text.
|
||||||
TOOL_XML_SIGNALS = ("<tool_call>", "<function=")
|
TOOL_XML_SIGNALS = ("<tool_call>", "<|tool_call>", "<function=")
|
||||||
|
|
||||||
|
|
||||||
# Nudges + error prefixes shared by the GGUF and safetensors loops.
|
# Nudges + error prefixes shared by the GGUF and safetensors loops.
|
||||||
|
|
@ -74,199 +80,6 @@ RAG_SEARCH_CAP_NUDGE = (
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Pre-compiled patterns reused by ``parse_tool_calls_from_text``.
|
|
||||||
_TC_JSON_START_RE = re.compile(r"<tool_call>\s*\{")
|
|
||||||
_TC_FUNC_START_RE = re.compile(r"<function=([\w-]+)>\s*")
|
|
||||||
_TC_END_TAG_RE = re.compile(r"</tool_call>")
|
|
||||||
_TC_FUNC_CLOSE_RE = re.compile(r"\s*</function>\s*$")
|
|
||||||
# [\w-] so hyphenated MCP param names (issue-number) aren't dropped.
|
|
||||||
_TC_PARAM_START_RE = re.compile(r"<parameter=([\w-]+)>\s*")
|
|
||||||
_TC_PARAM_CLOSE_RE = re.compile(r"\s*</parameter>\s*$")
|
|
||||||
_PARAM_CLOSE_TAG = "</parameter>"
|
|
||||||
_FUNC_CLOSE_TAG = "</function>"
|
|
||||||
|
|
||||||
|
|
||||||
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 two shapes:
|
|
||||||
|
|
||||||
- JSON inside ``<tool_call>`` tags:
|
|
||||||
``<tool_call>{"name":"web_search","arguments":{"query":"..."}}</tool_call>``
|
|
||||||
- XML-style function blocks:
|
|
||||||
``<function=name><parameter=k>v</parameter></function>``
|
|
||||||
|
|
||||||
``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: <tool_call>{json}. Balanced-brace scan, skipping braces in
|
|
||||||
# JSON strings.
|
|
||||||
for m in _TC_JSON_START_RE.finditer(content):
|
|
||||||
brace_start = m.end() - 1 # opening {
|
|
||||||
depth, i = 0, brace_start
|
|
||||||
in_string = False
|
|
||||||
while i < len(content):
|
|
||||||
ch = content[i]
|
|
||||||
if in_string:
|
|
||||||
if ch == "\\" and i + 1 < len(content):
|
|
||||||
i += 2
|
|
||||||
continue
|
|
||||||
if ch == '"':
|
|
||||||
in_string = False
|
|
||||||
elif ch == '"':
|
|
||||||
in_string = True
|
|
||||||
elif ch == "{":
|
|
||||||
depth += 1
|
|
||||||
elif ch == "}":
|
|
||||||
depth -= 1
|
|
||||||
if depth == 0:
|
|
||||||
break
|
|
||||||
i += 1
|
|
||||||
if depth != 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 2: <function=name><parameter=k>v... -- closing tags optional;
|
|
||||||
# </function> 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 </function> tag rather than
|
|
||||||
# the end of the response, so a complete call followed by
|
|
||||||
# trailing prose is still accepted (matching the JSON-style
|
|
||||||
# <tool_call> path, which already tolerates trailing text).
|
|
||||||
# rfind picks the last </function>, so a literal </function>
|
|
||||||
# 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
|
|
||||||
# </parameter> 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:
|
def has_tool_signal(text: str) -> bool:
|
||||||
"""Return True if ``text`` contains any tool-call XML signal."""
|
"""Return True if ``text`` contains any tool-call XML signal."""
|
||||||
return any(s in text for s in TOOL_XML_SIGNALS)
|
return any(s in text for s in TOOL_XML_SIGNALS)
|
||||||
|
|
|
||||||
|
|
@ -1,14 +1,10 @@
|
||||||
# SPDX-License-Identifier: AGPL-3.0-only
|
# SPDX-License-Identifier: AGPL-3.0-only
|
||||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||||
|
|
||||||
"""Tool-call XML parsing and stripping helpers.
|
"""Lightweight tool-call XML parsing and stripping helpers.
|
||||||
|
|
||||||
Extracted verbatim from studio/backend/core/inference/llama_cpp.py so external
|
External inference servers import this module without pulling in the inference
|
||||||
inference servers can reuse the logic without importing the inference
|
|
||||||
orchestrator, structlog, httpx, or the rest of the studio backend.
|
orchestrator, structlog, httpx, or the rest of the studio backend.
|
||||||
|
|
||||||
Regexes and bodies are byte-for-byte identical to the original; any change must
|
|
||||||
preserve that. test_tool_healing_extraction_is_exact.py verifies via AST.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
|
@ -17,8 +13,21 @@ import re
|
||||||
# Pre-compiled patterns for tool XML stripping. The hyphen in the name
|
# Pre-compiled patterns for tool XML stripping. The hyphen in the name
|
||||||
# char-class lets dashed MCP tool/parameter names (mcp__srv__list-issues,
|
# char-class lets dashed MCP tool/parameter names (mcp__srv__list-issues,
|
||||||
# issue-number) parse alongside the built-ins.
|
# issue-number) parse alongside the built-ins.
|
||||||
|
#
|
||||||
|
# The Gemma close marker is anchored to ``(?:<tool_call|>|\Z)`` (the safe form
|
||||||
|
# routes/inference.py's _TOOL_XML_RE uses): the plain ``<\|tool_call>.*?<tool_call\|>``
|
||||||
|
# this PR introduced backtracks from every open position on a run of unclosed
|
||||||
|
# markers (quadratic, and strip_tool_markup_streaming re-scans the cumulative
|
||||||
|
# buffer per token), whereas the ``\Z`` alternative lets the first open consume
|
||||||
|
# to EOF in one linear pass. strip_tool_call_markup additionally strips Gemma
|
||||||
|
# spans via the brace/quote-aware _strip_gemma_native_spans, so a literal close
|
||||||
|
# marker inside a <|"|>-quoted argument cannot truncate the span and leak its
|
||||||
|
# suffix; the regex below is the streaming-stripper fallback.
|
||||||
|
_TC_GEMMA_CLOSED_PAT = re.compile(r"<\|tool_call>.*?(?:<tool_call\|>|\Z)", re.DOTALL)
|
||||||
_TOOL_CLOSED_PATS = [
|
_TOOL_CLOSED_PATS = [
|
||||||
re.compile(r"<tool_call>.*?</tool_call>", re.DOTALL),
|
re.compile(r"<tool_call>.*?</tool_call>", re.DOTALL),
|
||||||
|
_TC_GEMMA_CLOSED_PAT,
|
||||||
|
re.compile(r"<tool_call\|>"),
|
||||||
re.compile(r"<function=[\w-]+>.*?</function>", re.DOTALL),
|
re.compile(r"<function=[\w-]+>.*?</function>", re.DOTALL),
|
||||||
]
|
]
|
||||||
_TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
_TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
||||||
|
|
@ -28,77 +37,358 @@ _TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
||||||
|
|
||||||
# Pre-compiled patterns for tool-call XML parsing.
|
# Pre-compiled patterns for tool-call XML parsing.
|
||||||
_TC_JSON_START_RE = re.compile(r"<tool_call>\s*\{")
|
_TC_JSON_START_RE = re.compile(r"<tool_call>\s*\{")
|
||||||
|
_TC_GEMMA_START_RE = re.compile(r"<\|tool_call>call:([\w-]+)\s*\{")
|
||||||
_TC_FUNC_START_RE = re.compile(r"<function=([\w-]+)>\s*")
|
_TC_FUNC_START_RE = re.compile(r"<function=([\w-]+)>\s*")
|
||||||
_TC_END_TAG_RE = re.compile(r"</tool_call>")
|
_TC_END_TAG_RE = re.compile(r"</tool_call>")
|
||||||
|
_TC_GEMMA_END_TAG_RE = re.compile(r"<tool_call\|>")
|
||||||
_TC_FUNC_CLOSE_RE = re.compile(r"\s*</function>\s*$")
|
_TC_FUNC_CLOSE_RE = re.compile(r"\s*</function>\s*$")
|
||||||
_TC_PARAM_START_RE = re.compile(r"<parameter=([\w-]+)>\s*")
|
_TC_PARAM_START_RE = re.compile(r"<parameter=([\w-]+)>\s*")
|
||||||
_TC_PARAM_CLOSE_RE = re.compile(r"\s*</parameter>\s*$")
|
_TC_PARAM_CLOSE_RE = re.compile(r"\s*</parameter>\s*$")
|
||||||
|
_GEMMA_QUOTE = '<|"|>'
|
||||||
|
_PARAM_CLOSE_TAG = "</parameter>"
|
||||||
|
_FUNC_CLOSE_TAG = "</function>"
|
||||||
|
# A bare (unquoted) Gemma value ends at `}` or at a comma that begins the next
|
||||||
|
# `key:` pair. A comma NOT followed by a key token is part of the value (e.g.
|
||||||
|
# `location:New York, NY`), so it must not terminate the value. The key token
|
||||||
|
# must be identifier-shaped (start with a letter or underscore); a comma
|
||||||
|
# followed by digits-then-colon is value text such as a timestamp or ratio
|
||||||
|
# (`meet at 10:00, 11:00 tomorrow`), not a new key.
|
||||||
|
_GEMMA_NEXT_KEY_RE = re.compile(r"\s*[A-Za-z_][\w-]*\s*:")
|
||||||
|
|
||||||
|
|
||||||
def parse_tool_calls_from_text(content: str) -> list[dict]:
|
def _balanced_brace_end(
|
||||||
"""
|
content: str,
|
||||||
Parse tool calls from XML markup in content text.
|
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 gemma_quotes and not in_string and content.startswith(_GEMMA_QUOTE, i):
|
||||||
|
in_gemma_string = not in_gemma_string
|
||||||
|
i += len(_GEMMA_QUOTE)
|
||||||
|
continue
|
||||||
|
ch = content[i]
|
||||||
|
if in_gemma_string:
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if in_string:
|
||||||
|
if ch == "\\" and i + 1 < len(content):
|
||||||
|
i += 2
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
in_string = False
|
||||||
|
elif ch == '"':
|
||||||
|
in_string = True
|
||||||
|
elif ch == "{":
|
||||||
|
depth += 1
|
||||||
|
elif ch == "}":
|
||||||
|
depth -= 1
|
||||||
|
if depth == 0:
|
||||||
|
return i
|
||||||
|
i += 1
|
||||||
|
return -1
|
||||||
|
|
||||||
|
|
||||||
|
def _balanced_bracket_end(src: str, start: int) -> int:
|
||||||
|
"""Index of the ``]`` matching the ``[`` at ``start``, or -1. Tracks nested
|
||||||
|
``[]``/``{}`` and double-quoted strings."""
|
||||||
|
depth = 0
|
||||||
|
i = start
|
||||||
|
in_string = False
|
||||||
|
while i < len(src):
|
||||||
|
ch = src[i]
|
||||||
|
if in_string:
|
||||||
|
if ch == "\\" and i + 1 < len(src):
|
||||||
|
i += 2
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
in_string = False
|
||||||
|
elif ch == '"':
|
||||||
|
in_string = True
|
||||||
|
elif ch in "[{":
|
||||||
|
depth += 1
|
||||||
|
elif ch in "]}":
|
||||||
|
depth -= 1
|
||||||
|
if depth == 0:
|
||||||
|
return i
|
||||||
|
i += 1
|
||||||
|
return -1
|
||||||
|
|
||||||
|
|
||||||
|
def _split_top_level_commas(src: str) -> list:
|
||||||
|
"""Split on commas that are not inside a nested ``[]``/``{}`` or a string."""
|
||||||
|
parts: list[str] = []
|
||||||
|
depth = 0
|
||||||
|
in_string = False
|
||||||
|
start = 0
|
||||||
|
i = 0
|
||||||
|
while i < len(src):
|
||||||
|
ch = src[i]
|
||||||
|
if in_string:
|
||||||
|
if ch == "\\" and i + 1 < len(src):
|
||||||
|
i += 2
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
in_string = False
|
||||||
|
elif ch == '"':
|
||||||
|
in_string = True
|
||||||
|
elif ch in "[{":
|
||||||
|
depth += 1
|
||||||
|
elif ch in "]}":
|
||||||
|
depth -= 1
|
||||||
|
elif ch == "," and depth == 0:
|
||||||
|
parts.append(src[start:i])
|
||||||
|
start = i + 1
|
||||||
|
i += 1
|
||||||
|
parts.append(src[start:])
|
||||||
|
return parts
|
||||||
|
|
||||||
|
|
||||||
|
def _quote_gemma_array_elements(body: str) -> str:
|
||||||
|
"""Normalise the elements of a Gemma array value so json.loads succeeds.
|
||||||
|
|
||||||
|
Gemma may emit ``labels:[bug,ui]`` without per-element quotes, or arrays of
|
||||||
|
objects (``items:[{path:a}]``) whose keys/values also lack quotes; left
|
||||||
|
as-is json.loads fails and the whole call is dropped. Bare string elements
|
||||||
|
are quoted, object and nested-array elements are normalised recursively, and
|
||||||
|
quoted strings (already normalised from ``<|"|>``), numbers, and JSON
|
||||||
|
literals are preserved."""
|
||||||
|
out: list[str] = []
|
||||||
|
for element in _split_top_level_commas(body):
|
||||||
|
stripped = element.strip()
|
||||||
|
if not stripped or stripped[0] == '"':
|
||||||
|
out.append(element)
|
||||||
|
continue
|
||||||
|
if stripped[0] == "{":
|
||||||
|
# Object element: quote its keys/bare values like a top-level object.
|
||||||
|
out.append(_quote_gemma_object_keys(stripped))
|
||||||
|
continue
|
||||||
|
if stripped[0] == "[":
|
||||||
|
# Nested array: normalise its elements too.
|
||||||
|
inner_end = _balanced_bracket_end(stripped, 0)
|
||||||
|
if inner_end == len(stripped) - 1:
|
||||||
|
out.append("[" + _quote_gemma_array_elements(stripped[1:inner_end]) + "]")
|
||||||
|
else:
|
||||||
|
out.append(element)
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
json.loads(stripped)
|
||||||
|
out.append(element)
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
out.append(json.dumps(stripped))
|
||||||
|
return ",".join(out)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalise_gemma_quoted_strings(src: str) -> str:
|
||||||
|
parts: list[str] = []
|
||||||
|
i = 0
|
||||||
|
while i < len(src):
|
||||||
|
if not src.startswith(_GEMMA_QUOTE, i):
|
||||||
|
parts.append(src[i])
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
end = src.find(_GEMMA_QUOTE, i + len(_GEMMA_QUOTE))
|
||||||
|
if end < 0:
|
||||||
|
parts.append(src[i:])
|
||||||
|
break
|
||||||
|
raw_value = src[i + len(_GEMMA_QUOTE) : end]
|
||||||
|
parts.append(json.dumps(raw_value))
|
||||||
|
i = end + len(_GEMMA_QUOTE)
|
||||||
|
return "".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _quote_gemma_object_keys(src: str) -> str:
|
||||||
|
parts: list[str] = []
|
||||||
|
i = 0
|
||||||
|
in_string = False
|
||||||
|
while i < len(src):
|
||||||
|
ch = src[i]
|
||||||
|
if in_string:
|
||||||
|
parts.append(ch)
|
||||||
|
if ch == "\\" and i + 1 < len(src):
|
||||||
|
parts.append(src[i + 1])
|
||||||
|
i += 2
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
in_string = False
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if ch == '"':
|
||||||
|
in_string = True
|
||||||
|
parts.append(ch)
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
if ch not in "{,":
|
||||||
|
parts.append(ch)
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
parts.append(ch)
|
||||||
|
i += 1
|
||||||
|
key_start = i
|
||||||
|
while i < len(src) and src[i].isspace():
|
||||||
|
i += 1
|
||||||
|
key_name_start = i
|
||||||
|
while i < len(src) and (src[i].isalnum() or src[i] in "_-"):
|
||||||
|
i += 1
|
||||||
|
key_name = src[key_name_start:i]
|
||||||
|
colon_pos = i
|
||||||
|
while colon_pos < len(src) and src[colon_pos].isspace():
|
||||||
|
colon_pos += 1
|
||||||
|
if key_name and colon_pos < len(src) and src[colon_pos] == ":":
|
||||||
|
parts.append(src[key_start:key_name_start])
|
||||||
|
parts.append(json.dumps(key_name))
|
||||||
|
parts.append(src[i:colon_pos])
|
||||||
|
parts.append(":")
|
||||||
|
i = colon_pos + 1
|
||||||
|
# Gemma may emit bare string values ({unit:celsius}); quote them so
|
||||||
|
# json.loads succeeds. JSON scalars/objects/arrays/quoted stay as-is.
|
||||||
|
ws = i
|
||||||
|
while i < len(src) and src[i].isspace():
|
||||||
|
i += 1
|
||||||
|
parts.append(src[ws:i])
|
||||||
|
if i < len(src) and src[i] == "[":
|
||||||
|
# Array value: quote bare string elements (e.g. labels:[bug,ui])
|
||||||
|
# so json.loads succeeds instead of dropping the call.
|
||||||
|
arr_end = _balanced_bracket_end(src, i)
|
||||||
|
if arr_end < 0:
|
||||||
|
parts.append(src[i:])
|
||||||
|
i = len(src)
|
||||||
|
else:
|
||||||
|
parts.append("[" + _quote_gemma_array_elements(src[i + 1 : arr_end]) + "]")
|
||||||
|
i = arr_end + 1
|
||||||
|
elif i < len(src) and src[i] not in '"{':
|
||||||
|
v_start = i
|
||||||
|
# Consume the bare value up to `}` or a comma that starts the
|
||||||
|
# next key:value pair; a comma inside the value (e.g.
|
||||||
|
# `New York, NY`) does not terminate it.
|
||||||
|
while i < len(src):
|
||||||
|
if src[i] == "}":
|
||||||
|
break
|
||||||
|
if src[i] == "," and _GEMMA_NEXT_KEY_RE.match(src, i + 1):
|
||||||
|
break
|
||||||
|
i += 1
|
||||||
|
raw = src[v_start:i]
|
||||||
|
try:
|
||||||
|
json.loads(raw.strip())
|
||||||
|
parts.append(raw)
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
parts.append(json.dumps(raw.strip()) if raw.strip() else raw)
|
||||||
|
else:
|
||||||
|
parts.append(src[key_start:i])
|
||||||
|
return "".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
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 {}
|
||||||
|
src = _normalise_gemma_quoted_strings(args_src)
|
||||||
|
src = "{" + src + "}"
|
||||||
|
src = _quote_gemma_object_keys(src)
|
||||||
|
return json.loads(src)
|
||||||
|
|
||||||
|
|
||||||
|
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:
|
Handles formats like:
|
||||||
<tool_call>{"name":"web_search","arguments":{"query":"..."}}</tool_call>
|
<tool_call>{"name":"web_search","arguments":{"query":"..."}}</tool_call>
|
||||||
|
<|tool_call>call:web_search{query:"..."}<tool_call|>
|
||||||
<tool_call><function=web_search><parameter=query>...</parameter></function></tool_call>
|
<tool_call><function=web_search><parameter=query>...</parameter></function></tool_call>
|
||||||
Closing tags (</tool_call>, </function>, </parameter>) are all
|
|
||||||
optional since models frequently omit them.
|
|
||||||
"""
|
"""
|
||||||
tool_calls = []
|
tool_calls: list[dict] = []
|
||||||
|
# Collect JSON- and Gemma-format candidates with their byte spans, then
|
||||||
# Pattern 1: JSON inside <tool_call> tags. Balanced-brace extraction that
|
# accept them in document order. Both order and spans matter:
|
||||||
# skips braces inside JSON strings.
|
# * tools execute in returned order, so a call appearing earlier in the
|
||||||
|
# text must be emitted first even across the two formats;
|
||||||
|
# * a tool-call marker INSIDE another call's argument string is data, not a
|
||||||
|
# call, so a candidate starting within an already accepted span is
|
||||||
|
# skipped (covers a JSON marker nested in a Gemma arg and a Gemma marker
|
||||||
|
# nested in a JSON arg alike, regardless of which format is outer).
|
||||||
|
candidates = [] # (start, brace_end, kind, match)
|
||||||
for m in _TC_JSON_START_RE.finditer(content):
|
for m in _TC_JSON_START_RE.finditer(content):
|
||||||
brace_start = m.end() - 1 # position of the opening {
|
# A marker that begins inside an open <function=...><parameter=...> value
|
||||||
depth, i = 0, brace_start
|
# is that parameter's data, not its own call; skip it (same guard the
|
||||||
in_string = False
|
# XML-style parser below applies to nested <function= markers).
|
||||||
while i < len(content):
|
if _inside_open_parameter(content, m.start()):
|
||||||
ch = content[i]
|
continue
|
||||||
if in_string:
|
end = _balanced_brace_end(content, m.end() - 1)
|
||||||
if ch == "\\" and i + 1 < len(content):
|
if end >= 0:
|
||||||
i += 2 # skip escaped character
|
candidates.append((m.start(), end, "json", m))
|
||||||
continue
|
for m in _TC_GEMMA_START_RE.finditer(content):
|
||||||
if ch == '"':
|
if _inside_open_parameter(content, m.start()):
|
||||||
in_string = False
|
continue
|
||||||
elif ch == '"':
|
end = _balanced_brace_end(content, m.end() - 1, gemma_quotes = True)
|
||||||
in_string = True
|
if end >= 0:
|
||||||
elif ch == "{":
|
candidates.append((m.start(), end, "gemma", m))
|
||||||
depth += 1
|
candidates.sort(key = lambda c: c[0])
|
||||||
elif ch == "}":
|
|
||||||
depth -= 1
|
spans = [(s, e) for s, e, _kind, _m in candidates]
|
||||||
if depth == 0:
|
for idx, (start, end, kind, m) in enumerate(candidates):
|
||||||
break
|
# Skip a candidate nested inside another candidate's brace span: it is
|
||||||
i += 1
|
# the enclosing call's argument data, not its own call. Checked against
|
||||||
if depth == 0:
|
# every candidate span (not only the ones that parsed successfully), so a
|
||||||
json_str = content[brace_start : i + 1]
|
# marker inside an outer call that later fails to normalize is still
|
||||||
try:
|
# never promoted to its own executable tool call.
|
||||||
obj = json.loads(json_str)
|
if any(s <= start and end <= e for j, (s, e) in enumerate(spans) if j != idx):
|
||||||
tc = {
|
continue
|
||||||
"id": f"call_{len(tool_calls)}",
|
if not allow_incomplete:
|
||||||
"type": "function",
|
tail = content[end + 1 :].lstrip()
|
||||||
"function": {
|
close_re = _TC_END_TAG_RE if kind == "json" else _TC_GEMMA_END_TAG_RE
|
||||||
"name": obj.get("name", ""),
|
if close_re.match(tail) is None:
|
||||||
"arguments": obj.get("arguments", {}),
|
continue
|
||||||
},
|
try:
|
||||||
}
|
if kind == "json":
|
||||||
if isinstance(tc["function"]["arguments"], dict):
|
obj = json.loads(content[m.end() - 1 : end + 1])
|
||||||
tc["function"]["arguments"] = json.dumps(tc["function"]["arguments"])
|
name = obj.get("name", "")
|
||||||
tool_calls.append(tc)
|
arguments = obj.get("arguments", {})
|
||||||
except (json.JSONDecodeError, ValueError):
|
if isinstance(arguments, dict):
|
||||||
pass
|
arguments = json.dumps(arguments)
|
||||||
|
else:
|
||||||
|
name = m.group(1)
|
||||||
|
arguments = json.dumps(_gemma_arguments_to_json(content[m.end() : end]))
|
||||||
|
except (json.JSONDecodeError, ValueError):
|
||||||
|
continue
|
||||||
|
tool_calls.append(
|
||||||
|
{
|
||||||
|
"id": f"call_{id_offset + len(tool_calls)}",
|
||||||
|
"type": "function",
|
||||||
|
"function": {"name": name, "arguments": arguments},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
# Pattern 2: XML-style <function=name><parameter=key>value</parameter></function>
|
|
||||||
# All closing tags optional; models frequently omit them.
|
|
||||||
if not tool_calls:
|
if not tool_calls:
|
||||||
# Step 1: Find <function=name> positions and extract bodies. Use only
|
func_starts = [
|
||||||
# </tool_call> or the next <function= as hard boundaries (</function>
|
fm
|
||||||
# can appear in code values); trim a trailing </function> afterwards.
|
for fm in _TC_FUNC_START_RE.finditer(content)
|
||||||
func_starts = list(_TC_FUNC_START_RE.finditer(content))
|
if not _inside_open_parameter(content, fm.start())
|
||||||
|
]
|
||||||
for idx, fm in enumerate(func_starts):
|
for idx, fm in enumerate(func_starts):
|
||||||
func_name = fm.group(1)
|
func_name = fm.group(1)
|
||||||
body_start = fm.end()
|
body_start = fm.end()
|
||||||
# Boundaries: next <function= tag or </tool_call>
|
|
||||||
next_func = func_starts[idx + 1].start() if idx + 1 < len(func_starts) else len(content)
|
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:])
|
end_tag = _TC_END_TAG_RE.search(content[body_start:])
|
||||||
if end_tag:
|
if end_tag:
|
||||||
|
|
@ -107,36 +397,52 @@ def parse_tool_calls_from_text(content: str) -> list[dict]:
|
||||||
body_end = len(content)
|
body_end = len(content)
|
||||||
body_end = min(body_end, next_func)
|
body_end = min(body_end, next_func)
|
||||||
body = content[body_start:body_end]
|
body = content[body_start:body_end]
|
||||||
body = _TC_FUNC_CLOSE_RE.sub("", body) # trim closing </function>
|
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
|
arguments: dict = {}
|
||||||
# functions, use body end as the only boundary to avoid matching
|
|
||||||
# </parameter> inside code strings.
|
|
||||||
arguments = {}
|
|
||||||
param_starts = list(_TC_PARAM_START_RE.finditer(body))
|
param_starts = list(_TC_PARAM_START_RE.finditer(body))
|
||||||
if len(param_starts) == 1:
|
if len(param_starts) == 1:
|
||||||
# Value is everything after the tag to end of body, less a
|
|
||||||
# trailing </parameter>.
|
|
||||||
pm = param_starts[0]
|
pm = param_starts[0]
|
||||||
val = body[pm.end() :]
|
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()
|
arguments[pm.group(1)] = val.strip()
|
||||||
else:
|
else:
|
||||||
|
valid_params = True
|
||||||
for pidx, pm in enumerate(param_starts):
|
for pidx, pm in enumerate(param_starts):
|
||||||
param_name = pm.group(1)
|
param_name = pm.group(1)
|
||||||
val_start = pm.end()
|
val_start = pm.end()
|
||||||
# Value ends at next <parameter= or end of body
|
|
||||||
next_param = (
|
next_param = (
|
||||||
param_starts[pidx + 1].start()
|
param_starts[pidx + 1].start()
|
||||||
if pidx + 1 < len(param_starts)
|
if pidx + 1 < len(param_starts)
|
||||||
else len(body)
|
else len(body)
|
||||||
)
|
)
|
||||||
val = body[val_start:next_param]
|
val = body[val_start:next_param]
|
||||||
val = _TC_PARAM_CLOSE_RE.sub("", val) # trim trailing </parameter>
|
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()
|
arguments[param_name] = val.strip()
|
||||||
|
if not valid_params:
|
||||||
|
continue
|
||||||
|
|
||||||
tc = {
|
tc = {
|
||||||
"id": f"call_{len(tool_calls)}",
|
"id": f"call_{id_offset + len(tool_calls)}",
|
||||||
"type": "function",
|
"type": "function",
|
||||||
"function": {
|
"function": {
|
||||||
"name": func_name,
|
"name": func_name,
|
||||||
|
|
@ -147,6 +453,41 @@ def parse_tool_calls_from_text(content: str) -> list[dict]:
|
||||||
return tool_calls
|
return tool_calls
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_gemma_native_spans(text: str, *, final: bool) -> str:
|
||||||
|
"""Remove complete Gemma-native ``<|tool_call>call:NAME{...}<tool_call|>``
|
||||||
|
spans, brace- and quote-balanced so a literal ``<tool_call|>`` inside a
|
||||||
|
``<|"|>``-quoted argument does not truncate the span and leak its suffix
|
||||||
|
(which the plain ``.*?`` regex does). A span without a balanced closing
|
||||||
|
``}`` or a trailing close marker is incomplete: dropped to EOF when
|
||||||
|
``final`` (the response is over), otherwise kept verbatim so a call that is
|
||||||
|
still streaming is not stripped mid-token.
|
||||||
|
"""
|
||||||
|
out: list[str] = []
|
||||||
|
cursor = 0
|
||||||
|
for match in _TC_GEMMA_START_RE.finditer(text):
|
||||||
|
start = match.start()
|
||||||
|
if start < cursor:
|
||||||
|
continue
|
||||||
|
brace_end = _balanced_brace_end(text, match.end() - 1, gemma_quotes = True)
|
||||||
|
if brace_end < 0:
|
||||||
|
if final:
|
||||||
|
out.append(text[cursor:start])
|
||||||
|
cursor = len(text)
|
||||||
|
continue
|
||||||
|
tail = text[brace_end + 1 :]
|
||||||
|
leading_ws = len(tail) - len(tail.lstrip())
|
||||||
|
close = _TC_GEMMA_END_TAG_RE.match(tail, leading_ws)
|
||||||
|
if close is None:
|
||||||
|
if final:
|
||||||
|
out.append(text[cursor:start])
|
||||||
|
cursor = len(text)
|
||||||
|
continue
|
||||||
|
out.append(text[cursor:start])
|
||||||
|
cursor = brace_end + 1 + close.end()
|
||||||
|
out.append(text[cursor:])
|
||||||
|
return "".join(out)
|
||||||
|
|
||||||
|
|
||||||
def strip_tool_call_markup(text: str, *, final: bool = False) -> str:
|
def strip_tool_call_markup(text: str, *, final: bool = False) -> str:
|
||||||
"""Strip tool-call XML markup from text.
|
"""Strip tool-call XML markup from text.
|
||||||
|
|
||||||
|
|
@ -154,7 +495,14 @@ def strip_tool_call_markup(text: str, *, final: bool = False) -> str:
|
||||||
When ``final`` is True, trailing incomplete tool-call blocks are removed
|
When ``final`` is True, trailing incomplete tool-call blocks are removed
|
||||||
too, and the result is stripped of surrounding whitespace.
|
too, and the result is stripped of surrounding whitespace.
|
||||||
"""
|
"""
|
||||||
|
# Gemma-native spans are stripped brace/quote-aware first; the regex form is
|
||||||
|
# not quote-aware and would truncate a span at a close marker inside a quoted
|
||||||
|
# argument. Skip that regex below and let the remaining patterns handle the
|
||||||
|
# JSON/XML formats and any orphan close marker.
|
||||||
|
text = _strip_gemma_native_spans(text, final = final)
|
||||||
patterns = _TOOL_ALL_PATS if final else _TOOL_CLOSED_PATS
|
patterns = _TOOL_ALL_PATS if final else _TOOL_CLOSED_PATS
|
||||||
for pat in patterns:
|
for pat in patterns:
|
||||||
|
if pat is _TC_GEMMA_CLOSED_PAT:
|
||||||
|
continue
|
||||||
text = pat.sub("", text)
|
text = pat.sub("", text)
|
||||||
return text.strip() if final else text
|
return text.strip() if final else text
|
||||||
|
|
|
||||||
|
|
@ -1102,6 +1102,8 @@ class ChoiceDelta(BaseModel):
|
||||||
|
|
||||||
role: Optional[str] = None
|
role: Optional[str] = None
|
||||||
content: Optional[str] = None
|
content: Optional[str] = None
|
||||||
|
reasoning_content: Optional[str] = None
|
||||||
|
tool_calls: Optional[list[dict]] = None
|
||||||
|
|
||||||
|
|
||||||
OpenAIFinishReason = Literal["stop", "length", "tool_calls", "content_filter", "function_call"]
|
OpenAIFinishReason = Literal["stop", "length", "tool_calls", "content_filter", "function_call"]
|
||||||
|
|
@ -1137,6 +1139,8 @@ class CompletionMessage(BaseModel):
|
||||||
role: Literal["assistant"] = "assistant"
|
role: Literal["assistant"] = "assistant"
|
||||||
content: str
|
content: str
|
||||||
refusal: Optional[str] = None
|
refusal: Optional[str] = None
|
||||||
|
reasoning_content: Optional[str] = None
|
||||||
|
tool_calls: Optional[list[dict]] = None
|
||||||
|
|
||||||
|
|
||||||
class CompletionChoice(BaseModel):
|
class CompletionChoice(BaseModel):
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ import uuid
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
from fastapi.responses import StreamingResponse, JSONResponse, Response
|
from fastapi.responses import StreamingResponse, JSONResponse, Response
|
||||||
|
from starlette.requests import ClientDisconnect
|
||||||
from typing import Any, List, Optional, Union
|
from typing import Any, List, Optional, Union
|
||||||
import json
|
import json
|
||||||
import httpx
|
import httpx
|
||||||
|
|
@ -235,8 +236,15 @@ def _sse_streaming_response(content) -> StreamingResponse:
|
||||||
a one-shot connection. Two callers build their response inline instead: the
|
a one-shot connection. Two callers build their response inline instead: the
|
||||||
external-provider proxy omits ``Connection: close``, and the OpenAI
|
external-provider proxy omits ``Connection: close``, and the OpenAI
|
||||||
passthrough returns an empty ``keep-alive`` stream when the request is
|
passthrough returns an empty ``keep-alive`` stream when the request is
|
||||||
cancelled before the upstream response starts."""
|
cancelled before the upstream response starts.
|
||||||
return StreamingResponse(
|
|
||||||
|
Built on ``_SameTaskStreamingResponse`` (not Starlette's stock
|
||||||
|
``StreamingResponse``) so the SSE generator runs in the request task. The
|
||||||
|
legacy AnyIO task-group wrapper trips "Attempted to exit a cancel scope in a
|
||||||
|
different task" on Python 3.13 + httpx, which surfaced as a mid-stream
|
||||||
|
``response.failed``. The streaming paths that take their response inline use
|
||||||
|
``_SameTaskStreamingResponse`` directly for the same reason."""
|
||||||
|
return _SameTaskStreamingResponse(
|
||||||
content,
|
content,
|
||||||
media_type = "text/event-stream",
|
media_type = "text/event-stream",
|
||||||
headers = {
|
headers = {
|
||||||
|
|
@ -750,6 +758,139 @@ def _set_stream_response_read_timeout(
|
||||||
pass
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
class _SameTaskStreamingResponse(StreamingResponse):
|
||||||
|
"""StreamingResponse without Starlette's legacy AnyIO task-group wrapper."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*args,
|
||||||
|
unstarted_cleanup = None,
|
||||||
|
**kwargs,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
# Async callable invoked when the client disconnects before the body
|
||||||
|
# iterator is ever advanced. A generator that never started cannot run
|
||||||
|
# its own try/finally, so a stream that acquires resources before its
|
||||||
|
# first yield (the passthrough opens an upstream httpx stream eagerly)
|
||||||
|
# passes this to release them.
|
||||||
|
self._unstarted_cleanup = unstarted_cleanup
|
||||||
|
|
||||||
|
async def __call__(self, scope, receive, send) -> None:
|
||||||
|
# Track whether the body iterator was ever advanced: send() only emits a
|
||||||
|
# body message after the generator yields its first chunk, so a failure
|
||||||
|
# before then means it never entered its try/finally.
|
||||||
|
body_started = False
|
||||||
|
|
||||||
|
async def _tracking_send(message) -> None:
|
||||||
|
nonlocal body_started
|
||||||
|
if message.get("type") == "http.response.body":
|
||||||
|
body_started = True
|
||||||
|
await send(message)
|
||||||
|
|
||||||
|
try:
|
||||||
|
await self.stream_response(_tracking_send)
|
||||||
|
except OSError:
|
||||||
|
# Client disconnected mid-send.
|
||||||
|
if body_started:
|
||||||
|
# The generator produced at least one chunk and is suspended in
|
||||||
|
# its try/finally. Throw CancelledError into it (not aclose's
|
||||||
|
# GeneratorExit) so its `except asyncio.CancelledError` handler
|
||||||
|
# runs and finishes any api_monitor entry; GeneratorExit would
|
||||||
|
# skip it and only run `finally`. Fall back to aclose() without
|
||||||
|
# athrow.
|
||||||
|
athrow = getattr(self.body_iterator, "athrow", None)
|
||||||
|
if athrow is not None:
|
||||||
|
try:
|
||||||
|
await athrow(asyncio.CancelledError())
|
||||||
|
except (asyncio.CancelledError, StopAsyncIteration, RuntimeError):
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
aclose = getattr(self.body_iterator, "aclose", None)
|
||||||
|
if aclose is not None:
|
||||||
|
await aclose()
|
||||||
|
else:
|
||||||
|
# http.response.start failed before the body iterator advanced,
|
||||||
|
# so its try/finally never armed and aclose()/athrow() are no-ops
|
||||||
|
# on an unstarted generator. Release any resources acquired
|
||||||
|
# before the first yield via the explicit cleanup hook.
|
||||||
|
aclose = getattr(self.body_iterator, "aclose", None)
|
||||||
|
if aclose is not None:
|
||||||
|
await aclose()
|
||||||
|
# getattr (not self._unstarted_cleanup) so a response built via
|
||||||
|
# __new__ (some tests, pickling) without __init__ does not raise
|
||||||
|
# AttributeError here.
|
||||||
|
cleanup = getattr(self, "_unstarted_cleanup", None)
|
||||||
|
if cleanup is not None:
|
||||||
|
try:
|
||||||
|
await cleanup()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
raise ClientDisconnect()
|
||||||
|
if self.background is not None:
|
||||||
|
await self.background()
|
||||||
|
|
||||||
|
|
||||||
|
def _tracked_cancel_unstarted_cleanup(tracker):
|
||||||
|
"""Build an ``unstarted_cleanup`` for a local stream that entered ``tracker``
|
||||||
|
(a ``_TrackedCancel``) before returning the response. The generator exits the
|
||||||
|
tracker in its ``finally``, but that never runs if the client disconnects
|
||||||
|
before the body iterator starts, leaking the cancel-registry entry. This
|
||||||
|
exits the tracker on that pre-start path only (mutually exclusive with the
|
||||||
|
generator's finally, so it never double-exits)."""
|
||||||
|
|
||||||
|
async def _cleanup() -> None:
|
||||||
|
tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
|
return _cleanup
|
||||||
|
|
||||||
|
|
||||||
async def _aclose_stream_resources(
|
async def _aclose_stream_resources(
|
||||||
*,
|
*,
|
||||||
watchers = (),
|
watchers = (),
|
||||||
|
|
@ -875,8 +1016,23 @@ async def _aiter_llama_stream_items(
|
||||||
raise httpx.ReadTimeout("The model did not produce a first token in time.")
|
raise httpx.ReadTimeout("The model did not produce a first token in time.")
|
||||||
if response is not None:
|
if response is not None:
|
||||||
_set_stream_response_read_timeout(response, remaining_s)
|
_set_stream_response_read_timeout(response, remaining_s)
|
||||||
item = await asyncio.wait_for(async_iter.__anext__(), timeout = 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):
|
||||||
|
item = await async_iter.__anext__()
|
||||||
else:
|
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, stall_remaining_s)
|
||||||
item = await async_iter.__anext__()
|
item = await async_iter.__anext__()
|
||||||
except asyncio.TimeoutError as exc:
|
except asyncio.TimeoutError as exc:
|
||||||
if waiting_first_item:
|
if waiting_first_item:
|
||||||
|
|
@ -890,6 +1046,12 @@ async def _aiter_llama_stream_items(
|
||||||
if now >= first_token_deadline:
|
if now >= first_token_deadline:
|
||||||
raise
|
raise
|
||||||
continue
|
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.")
|
raise httpx.ReadTimeout("The model stopped producing tokens mid-response.")
|
||||||
if (
|
if (
|
||||||
last_item_at is None
|
last_item_at is None
|
||||||
|
|
@ -1125,16 +1287,17 @@ def _detect_safetensors_features(backend, chat_template: Optional[str]) -> dict:
|
||||||
model_identifier = model_id,
|
model_identifier = model_id,
|
||||||
log_source = "safetensors",
|
log_source = "safetensors",
|
||||||
)
|
)
|
||||||
# Our safetensors loop only parses <tool_call>{json}</tool_call> and
|
# Our safetensors loop only parses <tool_call>{json}</tool_call>,
|
||||||
# <function=name>...</function>. Llama uses <|python_tag|>, Mistral uses
|
# <function=name>...</function>, and Gemma native <|tool_call>...<tool_call|>.
|
||||||
# [TOOL_CALLS]; advertising tools for those enables a pill the parser
|
# Llama uses <|python_tag|>, Mistral uses [TOOL_CALLS]; advertising tools for
|
||||||
# can't honour. GGUF is unaffected -- llama-server normalises every
|
# those enables a pill the parser can't honour. GGUF is unaffected --
|
||||||
# format into structured deltas.
|
# llama-server normalises every format into structured deltas.
|
||||||
if (
|
if (
|
||||||
flags.get("supports_tools")
|
flags.get("supports_tools")
|
||||||
and chat_template
|
and chat_template
|
||||||
and "<tool_call>" not in chat_template
|
and "<tool_call>" not in chat_template
|
||||||
and "<function=" not in chat_template
|
and "<function=" not in chat_template
|
||||||
|
and "<|tool_call>" not in chat_template
|
||||||
):
|
):
|
||||||
logger.info(
|
logger.info(
|
||||||
"safetensors: template advertises tools but uses an "
|
"safetensors: template advertises tools but uses an "
|
||||||
|
|
@ -1297,6 +1460,24 @@ async def _await_disconnect_then_close(request, resp, cancel_event) -> None:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|
||||||
|
async def _await_disconnect_then_cancel(request, cancel_event) -> None:
|
||||||
|
"""Set ``cancel_event`` when a same-task local stream disconnects."""
|
||||||
|
try:
|
||||||
|
while not await request.is_disconnected():
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
cancel_event.set()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
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
|
# Centralized local/server tool nudge. Keep render_html guidance gated to turns
|
||||||
# where the canvas tool is actually present in the tool schema; otherwise
|
# where the canvas tool is actually present in the tool schema; otherwise
|
||||||
# small local models can hallucinate a missing tool call instead of following
|
# small local models can hallucinate a missing tool call instead of following
|
||||||
|
|
@ -1418,7 +1599,9 @@ _TOOL_XML_RE = _re.compile(
|
||||||
# Hyphen in the name char-class matches MCP tool names with dashes
|
# Hyphen in the name char-class matches MCP tool names with dashes
|
||||||
# (mcp__srv__list-issues) that would otherwise leak past this strip.
|
# (mcp__srv__list-issues) that would otherwise leak past this strip.
|
||||||
r"<(?:tool_call|function=[\w-]+)>.*?(?:</(?:tool_call|function)>|\Z)"
|
r"<(?:tool_call|function=[\w-]+)>.*?(?:</(?:tool_call|function)>|\Z)"
|
||||||
|
r"|<\|tool_call>.*?(?:<tool_call\|>|\Z)"
|
||||||
r"|</(?:tool_call|function)>"
|
r"|</(?:tool_call|function)>"
|
||||||
|
r"|<tool_call\|>"
|
||||||
r"|</parameter>\s*\Z",
|
r"|</parameter>\s*\Z",
|
||||||
_re.DOTALL,
|
_re.DOTALL,
|
||||||
)
|
)
|
||||||
|
|
@ -3221,7 +3404,9 @@ async def get_api_monitor_entry(entry_id: str, current_subject: str = Depends(ge
|
||||||
|
|
||||||
@router.post("/generate/stream")
|
@router.post("/generate/stream")
|
||||||
async def generate_stream(
|
async def generate_stream(
|
||||||
request: GenerateRequest, current_subject: str = Depends(get_current_subject)
|
request: GenerateRequest,
|
||||||
|
fastapi_request: Request,
|
||||||
|
current_subject: str = Depends(get_current_subject),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Generate a chat response with Server-Sent Events (SSE) streaming.
|
Generate a chat response with Server-Sent Events (SSE) streaming.
|
||||||
|
|
@ -3271,6 +3456,13 @@ async def generate_stream(
|
||||||
async def stream():
|
async def stream():
|
||||||
gen = None
|
gen = None
|
||||||
completed = False
|
completed = False
|
||||||
|
# Cancel the generation when the client disconnects. The generator only
|
||||||
|
# awaits asyncio.to_thread(next, gen, ...), so without a concurrent
|
||||||
|
# watcher a disconnect during a long prefill/generation would go
|
||||||
|
# unnoticed until the next send and the backend would keep generating.
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(fastapi_request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
gen = backend.generate_chat_response(
|
gen = backend.generate_chat_response(
|
||||||
messages = request.messages,
|
messages = request.messages,
|
||||||
|
|
@ -3285,12 +3477,22 @@ async def generate_stream(
|
||||||
)
|
)
|
||||||
_DONE = object()
|
_DONE = object()
|
||||||
while True:
|
while True:
|
||||||
|
if cancel_event.is_set():
|
||||||
|
# The disconnect watcher set cancel_event between chunks.
|
||||||
|
# Reset the backend here: closing the Python generator does
|
||||||
|
# not signal a subprocess backend, so without this it keeps
|
||||||
|
# decoding after the client is gone. The finally's reset is
|
||||||
|
# guarded on cancel_event being unset, so it will not run
|
||||||
|
# again for this path.
|
||||||
|
backend.reset_generation_state()
|
||||||
|
break
|
||||||
chunk = await asyncio.to_thread(next, gen, _DONE)
|
chunk = await asyncio.to_thread(next, gen, _DONE)
|
||||||
if chunk is _DONE:
|
if chunk is _DONE:
|
||||||
|
completed = True
|
||||||
break
|
break
|
||||||
yield f"data: {json.dumps({'content': chunk})}\n\n"
|
yield f"data: {json.dumps({'content': chunk})}\n\n"
|
||||||
completed = True
|
if completed:
|
||||||
yield "data: [DONE]\n\n"
|
yield "data: [DONE]\n\n"
|
||||||
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
cancel_event.set()
|
cancel_event.set()
|
||||||
|
|
@ -3302,6 +3504,7 @@ async def generate_stream(
|
||||||
logger.error(f"Error during generation: {e}", exc_info = True)
|
logger.error(f"Error during generation: {e}", exc_info = True)
|
||||||
yield f"data: {json.dumps({'error': _friendly_error(e)})}\n\n"
|
yield f"data: {json.dumps({'error': _friendly_error(e)})}\n\n"
|
||||||
finally:
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
if not completed and not cancel_event.is_set():
|
if not completed and not cancel_event.is_set():
|
||||||
cancel_event.set()
|
cancel_event.set()
|
||||||
backend.reset_generation_state()
|
backend.reset_generation_state()
|
||||||
|
|
@ -4725,6 +4928,9 @@ async def openai_chat_completions(
|
||||||
_tracker.__enter__()
|
_tracker.__enter__()
|
||||||
|
|
||||||
async def audio_input_stream():
|
async def audio_input_stream():
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
yield _chat_role_chunk(completion_id, created, model_name)
|
yield _chat_role_chunk(completion_id, created, model_name)
|
||||||
|
|
||||||
|
|
@ -4760,9 +4966,19 @@ async def openai_chat_completions(
|
||||||
api_monitor.fail(monitor_id, _friendly_error(e))
|
api_monitor.fail(monitor_id, _friendly_error(e))
|
||||||
yield f"data: {json.dumps({'error': {'message': _friendly_error(e), 'type': 'server_error'}})}\n\n"
|
yield f"data: {json.dumps({'error': {'message': _friendly_error(e), 'type': 'server_error'}})}\n\n"
|
||||||
finally:
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
return _sse_streaming_response(audio_input_stream())
|
return _SameTaskStreamingResponse(
|
||||||
|
audio_input_stream(),
|
||||||
|
unstarted_cleanup = _tracked_cancel_unstarted_cleanup(_tracker),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
try:
|
try:
|
||||||
full_text = "".join(audio_input_generate())
|
full_text = "".join(audio_input_generate())
|
||||||
|
|
@ -4937,6 +5153,28 @@ async def openai_chat_completions(
|
||||||
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||||
created = int(time.time())
|
created = int(time.time())
|
||||||
|
|
||||||
|
def _new_chat_reasoning_extractor():
|
||||||
|
return _ResponsesReasoningExtractor(
|
||||||
|
parse_think_markers = _responses_should_parse_think_markers(
|
||||||
|
payload,
|
||||||
|
llama_backend,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _gguf_chat_delta_line(delta: ChoiceDelta, finish_reason = None) -> str:
|
||||||
|
chunk = ChatCompletionChunk(
|
||||||
|
id = completion_id,
|
||||||
|
created = created,
|
||||||
|
model = model_name,
|
||||||
|
choices = [
|
||||||
|
ChunkChoice(
|
||||||
|
delta = delta,
|
||||||
|
finish_reason = finish_reason,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
return f"data: {chunk.model_dump_json(exclude_none = True)}\n\n"
|
||||||
|
|
||||||
# ── Tool-calling path (agentic loop) ──────────────────
|
# ── Tool-calling path (agentic loop) ──────────────────
|
||||||
# `_effective_enable_tools` lets `unsloth run --enable-tools/--disable-tools`
|
# `_effective_enable_tools` lets `unsloth run --enable-tools/--disable-tools`
|
||||||
# hard-override the per-request value, else falls back to
|
# hard-override the per-request value, else falls back to
|
||||||
|
|
@ -5049,6 +5287,9 @@ async def openai_chat_completions(
|
||||||
|
|
||||||
async def gguf_tool_stream():
|
async def gguf_tool_stream():
|
||||||
gen = None
|
gen = None
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
yield _chat_role_chunk(completion_id, created, model_name)
|
yield _chat_role_chunk(completion_id, created, model_name)
|
||||||
|
|
||||||
|
|
@ -5056,9 +5297,25 @@ async def openai_chat_completions(
|
||||||
# stays free for disconnect detection.
|
# stays free for disconnect detection.
|
||||||
gen = gguf_generate_with_tools()
|
gen = gguf_generate_with_tools()
|
||||||
prev_text = ""
|
prev_text = ""
|
||||||
|
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||||
_stream_usage = None
|
_stream_usage = None
|
||||||
_stream_timings = None
|
_stream_timings = None
|
||||||
_stream_finish = 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:
|
while True:
|
||||||
if cancel_event.is_set():
|
if cancel_event.is_set():
|
||||||
break
|
break
|
||||||
|
|
@ -5077,7 +5334,10 @@ async def openai_chat_completions(
|
||||||
# cumulative cursor so the next assistant turn
|
# cumulative cursor so the next assistant turn
|
||||||
# streams cleanly.
|
# streams cleanly.
|
||||||
if not event["text"]:
|
if not event["text"]:
|
||||||
|
for chunk in _flush_reasoning_extractor():
|
||||||
|
yield chunk
|
||||||
prev_text = ""
|
prev_text = ""
|
||||||
|
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||||
# Emit tool status as a custom SSE event (including
|
# Emit tool status as a custom SSE event (including
|
||||||
# empty ones to clear UI badges)
|
# empty ones to clear UI badges)
|
||||||
status_data = json.dumps(
|
status_data = json.dumps(
|
||||||
|
|
@ -5091,7 +5351,10 @@ async def openai_chat_completions(
|
||||||
|
|
||||||
if event["type"] in ("tool_start", "tool_end"):
|
if event["type"] in ("tool_start", "tool_end"):
|
||||||
if event["type"] == "tool_start":
|
if event["type"] == "tool_start":
|
||||||
|
for chunk in _flush_reasoning_extractor():
|
||||||
|
yield chunk
|
||||||
prev_text = ""
|
prev_text = ""
|
||||||
|
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||||
yield f"data: {json.dumps(event)}\n\n"
|
yield f"data: {json.dumps(event)}\n\n"
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
|
@ -5113,15 +5376,33 @@ async def openai_chat_completions(
|
||||||
prev_text = clean_cumulative
|
prev_text = clean_cumulative
|
||||||
if not new_text:
|
if not new_text:
|
||||||
continue
|
continue
|
||||||
api_monitor.append_reply(monitor_id, new_text)
|
reasoning_delta, visible_delta = reasoning_extractor.feed(new_text)
|
||||||
yield _chat_content_chunk(completion_id, created, model_name, new_text)
|
if reasoning_delta:
|
||||||
|
yield _gguf_chat_delta_line(
|
||||||
|
ChoiceDelta(reasoning_content = reasoning_delta)
|
||||||
|
)
|
||||||
|
if visible_delta:
|
||||||
|
api_monitor.append_reply(monitor_id, visible_delta)
|
||||||
|
yield _gguf_chat_delta_line(ChoiceDelta(content = visible_delta))
|
||||||
|
|
||||||
yield _chat_final_chunk(
|
for chunk in _flush_reasoning_extractor():
|
||||||
completion_id,
|
yield chunk
|
||||||
created,
|
|
||||||
model_name,
|
final_chunk = ChatCompletionChunk(
|
||||||
_clamp_finish_reason(_stream_finish),
|
id = completion_id,
|
||||||
|
created = created,
|
||||||
|
model = model_name,
|
||||||
|
choices = [
|
||||||
|
ChunkChoice(
|
||||||
|
delta = ChoiceDelta(),
|
||||||
|
finish_reason = _clamp_finish_reason(_stream_finish),
|
||||||
|
)
|
||||||
|
],
|
||||||
)
|
)
|
||||||
|
# Emit the terminal chunk carrying finish_reason before the
|
||||||
|
# optional usage chunk and [DONE], so OpenAI-compatible
|
||||||
|
# clients can detect stop/length/tool_calls.
|
||||||
|
yield f"data: {final_chunk.model_dump_json(exclude_none = True)}\n\n"
|
||||||
usage_line = _openai_stream_usage_chunk(
|
usage_line = _openai_stream_usage_chunk(
|
||||||
payload,
|
payload,
|
||||||
completion_id,
|
completion_id,
|
||||||
|
|
@ -5150,6 +5431,7 @@ async def openai_chat_completions(
|
||||||
error_chunk = _openai_stream_error_chunk(e)
|
error_chunk = _openai_stream_error_chunk(e)
|
||||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||||
finally:
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
if gen is not None:
|
if gen is not None:
|
||||||
try:
|
try:
|
||||||
gen.close()
|
gen.close()
|
||||||
|
|
@ -5157,7 +5439,16 @@ async def openai_chat_completions(
|
||||||
pass
|
pass
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
return _sse_streaming_response(gguf_tool_stream())
|
return _SameTaskStreamingResponse(
|
||||||
|
gguf_tool_stream(),
|
||||||
|
unstarted_cleanup = _tracked_cancel_unstarted_cleanup(_tracker),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
# ── Standard GGUF path (no tools) ─────────────────────
|
# ── Standard GGUF path (no tools) ─────────────────────
|
||||||
|
|
||||||
|
|
@ -5193,6 +5484,9 @@ async def openai_chat_completions(
|
||||||
_tracker.__enter__()
|
_tracker.__enter__()
|
||||||
|
|
||||||
async def gguf_stream_chunks():
|
async def gguf_stream_chunks():
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
yield _chat_role_chunk(completion_id, created, model_name)
|
yield _chat_role_chunk(completion_id, created, model_name)
|
||||||
|
|
||||||
|
|
@ -5200,6 +5494,7 @@ async def openai_chat_completions(
|
||||||
# stays free for disconnect detection.
|
# stays free for disconnect detection.
|
||||||
gen = gguf_generate()
|
gen = gguf_generate()
|
||||||
prev_text = ""
|
prev_text = ""
|
||||||
|
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||||
_stream_usage = None
|
_stream_usage = None
|
||||||
_stream_timings = None
|
_stream_timings = None
|
||||||
_stream_finish = None
|
_stream_finish = None
|
||||||
|
|
@ -5233,15 +5528,38 @@ async def openai_chat_completions(
|
||||||
prev_text = cumulative
|
prev_text = cumulative
|
||||||
if not new_text:
|
if not new_text:
|
||||||
continue
|
continue
|
||||||
api_monitor.append_reply(monitor_id, new_text)
|
reasoning_delta, visible_delta = reasoning_extractor.feed(new_text)
|
||||||
yield _chat_content_chunk(completion_id, created, model_name, new_text)
|
if reasoning_delta:
|
||||||
|
yield _gguf_chat_delta_line(
|
||||||
|
ChoiceDelta(reasoning_content = reasoning_delta)
|
||||||
|
)
|
||||||
|
if visible_delta:
|
||||||
|
api_monitor.append_reply(monitor_id, visible_delta)
|
||||||
|
yield _gguf_chat_delta_line(ChoiceDelta(content = visible_delta))
|
||||||
|
|
||||||
yield _chat_final_chunk(
|
final_reasoning, final_visible = reasoning_extractor.finish()
|
||||||
completion_id,
|
if final_reasoning:
|
||||||
created,
|
yield _gguf_chat_delta_line(ChoiceDelta(reasoning_content = final_reasoning))
|
||||||
model_name,
|
if final_visible:
|
||||||
_clamp_finish_reason(_stream_finish),
|
api_monitor.append_reply(monitor_id, final_visible)
|
||||||
|
yield _gguf_chat_delta_line(ChoiceDelta(content = final_visible))
|
||||||
|
|
||||||
|
# Final chunk
|
||||||
|
final_chunk = ChatCompletionChunk(
|
||||||
|
id = completion_id,
|
||||||
|
created = created,
|
||||||
|
model = model_name,
|
||||||
|
choices = [
|
||||||
|
ChunkChoice(
|
||||||
|
delta = ChoiceDelta(),
|
||||||
|
finish_reason = _clamp_finish_reason(_stream_finish),
|
||||||
|
)
|
||||||
|
],
|
||||||
)
|
)
|
||||||
|
# Emit the terminal chunk carrying finish_reason before the
|
||||||
|
# optional usage chunk and [DONE], so OpenAI-compatible
|
||||||
|
# clients can detect stop/length/tool_calls.
|
||||||
|
yield f"data: {final_chunk.model_dump_json(exclude_none = True)}\n\n"
|
||||||
usage_line = _openai_stream_usage_chunk(
|
usage_line = _openai_stream_usage_chunk(
|
||||||
payload,
|
payload,
|
||||||
completion_id,
|
completion_id,
|
||||||
|
|
@ -5268,9 +5586,19 @@ async def openai_chat_completions(
|
||||||
error_chunk = _openai_stream_error_chunk(e)
|
error_chunk = _openai_stream_error_chunk(e)
|
||||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||||
finally:
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
return _sse_streaming_response(gguf_stream_chunks())
|
return _SameTaskStreamingResponse(
|
||||||
|
gguf_stream_chunks(),
|
||||||
|
unstarted_cleanup = _tracked_cancel_unstarted_cleanup(_tracker),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
try:
|
try:
|
||||||
# ``n`` requests several independent completions; the single
|
# ``n`` requests several independent completions; the single
|
||||||
|
|
@ -5297,14 +5625,24 @@ async def openai_chat_completions(
|
||||||
continue
|
continue
|
||||||
full_text = token
|
full_text = token
|
||||||
|
|
||||||
|
reasoning_text, visible_text = _extract_responses_reasoning(
|
||||||
|
full_text,
|
||||||
|
parse_think_markers = _responses_should_parse_think_markers(
|
||||||
|
payload,
|
||||||
|
llama_backend,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
message_kwargs = {"content": visible_text}
|
||||||
|
if reasoning_text:
|
||||||
|
message_kwargs["reasoning_content"] = reasoning_text
|
||||||
_choices.append(
|
_choices.append(
|
||||||
CompletionChoice(
|
CompletionChoice(
|
||||||
index = _idx,
|
index = _idx,
|
||||||
message = CompletionMessage(content = full_text),
|
message = CompletionMessage(**message_kwargs),
|
||||||
finish_reason = _clamp_finish_reason(completion_finish),
|
finish_reason = _clamp_finish_reason(completion_finish),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
_monitor_replies.append(full_text)
|
_monitor_replies.append(visible_text)
|
||||||
if completion_usage:
|
if completion_usage:
|
||||||
# The prompt is shared across all n choices, so count its
|
# The prompt is shared across all n choices, so count its
|
||||||
# tokens ONCE (OpenAI bills only generated tokens for each
|
# tokens ONCE (OpenAI bills only generated tokens for each
|
||||||
|
|
@ -5326,7 +5664,7 @@ async def openai_chat_completions(
|
||||||
prompt_tokens_details = _prompt_tokens_details(_prompt_details),
|
prompt_tokens_details = _prompt_tokens_details(_prompt_details),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
monitor_reply = full_text
|
monitor_reply = _monitor_replies[-1] if _monitor_replies else ""
|
||||||
if _n > 1:
|
if _n > 1:
|
||||||
monitor_reply = "\n\n".join(
|
monitor_reply = "\n\n".join(
|
||||||
f"Choice {_idx + 1}:\n{text}" for _idx, text in enumerate(_monitor_replies)
|
f"Choice {_idx + 1}:\n{text}" for _idx, text in enumerate(_monitor_replies)
|
||||||
|
|
@ -5536,6 +5874,9 @@ async def openai_chat_completions(
|
||||||
|
|
||||||
async def sf_tool_stream():
|
async def sf_tool_stream():
|
||||||
gen = None
|
gen = None
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
yield _chat_role_chunk(completion_id, created, model_name)
|
yield _chat_role_chunk(completion_id, created, model_name)
|
||||||
|
|
||||||
|
|
@ -5627,6 +5968,7 @@ async def openai_chat_completions(
|
||||||
}
|
}
|
||||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||||
finally:
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
if gen is not None:
|
if gen is not None:
|
||||||
try:
|
try:
|
||||||
gen.close()
|
gen.close()
|
||||||
|
|
@ -5635,7 +5977,16 @@ async def openai_chat_completions(
|
||||||
_sf_tracker.__exit__(None, None, None)
|
_sf_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
if payload.stream:
|
if payload.stream:
|
||||||
return _sse_streaming_response(sf_tool_stream())
|
return _SameTaskStreamingResponse(
|
||||||
|
sf_tool_stream(),
|
||||||
|
unstarted_cleanup = _tracked_cancel_unstarted_cleanup(_sf_tracker),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
# Non-streaming JSON: drain the loop, build one ChatCompletion.
|
# Non-streaming JSON: drain the loop, build one ChatCompletion.
|
||||||
try:
|
try:
|
||||||
|
|
@ -5737,6 +6088,9 @@ async def openai_chat_completions(
|
||||||
_tracker.__enter__()
|
_tracker.__enter__()
|
||||||
|
|
||||||
async def stream_chunks():
|
async def stream_chunks():
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
yield _chat_role_chunk(completion_id, created, model_name)
|
yield _chat_role_chunk(completion_id, created, model_name)
|
||||||
|
|
||||||
|
|
@ -5813,9 +6167,19 @@ async def openai_chat_completions(
|
||||||
}
|
}
|
||||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||||
finally:
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
return _sse_streaming_response(stream_chunks())
|
return _SameTaskStreamingResponse(
|
||||||
|
stream_chunks(),
|
||||||
|
unstarted_cleanup = _tracked_cancel_unstarted_cleanup(_tracker),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
# ── Non-streaming response ────────────────────────────────────
|
# ── Non-streaming response ────────────────────────────────────
|
||||||
else:
|
else:
|
||||||
|
|
@ -6533,8 +6897,9 @@ def _responses_should_parse_think_markers(
|
||||||
if llama_backend is not None and getattr(llama_backend, "is_loaded", False):
|
if llama_backend is not None and getattr(llama_backend, "is_loaded", False):
|
||||||
if getattr(llama_backend, "reasoning_always_on", False):
|
if getattr(llama_backend, "reasoning_always_on", False):
|
||||||
return True
|
return True
|
||||||
if not getattr(llama_backend, "supports_reasoning", False):
|
if getattr(llama_backend, "supports_reasoning", False):
|
||||||
return False
|
return True
|
||||||
|
return False
|
||||||
if chat_req.enable_thinking is True:
|
if chat_req.enable_thinking is True:
|
||||||
return True
|
return True
|
||||||
return chat_req.enable_thinking is None and chat_req.reasoning_effort not in (None, "none")
|
return chat_req.enable_thinking is None and chat_req.reasoning_effort not in (None, "none")
|
||||||
|
|
@ -6830,8 +7195,6 @@ async def _responses_non_streaming(
|
||||||
# the model produced content, so clients expecting a pure tool-call turn
|
# the model produced content, so clients expecting a pure tool-call turn
|
||||||
# (finish_reason="tool_calls") don't see a spurious empty message item.
|
# (finish_reason="tool_calls") don't see a spurious empty message item.
|
||||||
output_items: list[dict] = []
|
output_items: list[dict] = []
|
||||||
if reasoning_text and not text and not tool_calls:
|
|
||||||
text = reasoning_text
|
|
||||||
if reasoning_text:
|
if reasoning_text:
|
||||||
output_items.append(_responses_reasoning_output_item(reasoning_text))
|
output_items.append(_responses_reasoning_output_item(reasoning_text))
|
||||||
if text:
|
if text:
|
||||||
|
|
@ -7144,8 +7507,8 @@ async def _responses_stream(
|
||||||
client = httpx.AsyncClient(timeout = _llama_streaming_generation_timeout())
|
client = httpx.AsyncClient(timeout = _llama_streaming_generation_timeout())
|
||||||
resp = None
|
resp = None
|
||||||
lines_iter = None
|
lines_iter = None
|
||||||
disconnect_event = threading.Event()
|
|
||||||
disconnect_watcher = None
|
disconnect_watcher = None
|
||||||
|
disconnect_event = threading.Event()
|
||||||
try:
|
try:
|
||||||
req = client.build_request(
|
req = client.build_request(
|
||||||
"POST", target_url, json = body, headers = {"Connection": "close"}
|
"POST", target_url, json = body, headers = {"Connection": "close"}
|
||||||
|
|
@ -7205,10 +7568,10 @@ async def _responses_stream(
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
lines_iter = resp.aiter_lines()
|
||||||
disconnect_watcher = asyncio.create_task(
|
disconnect_watcher = asyncio.create_task(
|
||||||
_await_disconnect_then_close(request, resp, disconnect_event)
|
_await_disconnect_then_close(request, resp, disconnect_event)
|
||||||
)
|
)
|
||||||
lines_iter = resp.aiter_lines()
|
|
||||||
async for raw_line in _aiter_llama_stream_items(
|
async for raw_line in _aiter_llama_stream_items(
|
||||||
lines_iter,
|
lines_iter,
|
||||||
cancel_event = disconnect_event,
|
cancel_event = disconnect_event,
|
||||||
|
|
@ -7328,6 +7691,7 @@ async def _responses_stream(
|
||||||
|
|
||||||
_apply_usage(chunk_data.get("usage"))
|
_apply_usage(chunk_data.get("usage"))
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
|
disconnect_event.set()
|
||||||
api_monitor.finish(monitor_id, "cancelled")
|
api_monitor.finish(monitor_id, "cancelled")
|
||||||
raise
|
raise
|
||||||
except (httpx.RemoteProtocolError, httpx.ReadError, httpx.CloseError) as e:
|
except (httpx.RemoteProtocolError, httpx.ReadError, httpx.CloseError) as e:
|
||||||
|
|
@ -7394,21 +7758,6 @@ async def _responses_stream(
|
||||||
"delta": final_visible,
|
"delta": final_visible,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
if full_reasoning and not full_text and not tool_call_state:
|
|
||||||
for event in _ensure_message_open():
|
|
||||||
yield event
|
|
||||||
full_text = full_reasoning
|
|
||||||
api_monitor.set_reply(monitor_id, full_text)
|
|
||||||
yield _sse(
|
|
||||||
"response.output_text.delta",
|
|
||||||
{
|
|
||||||
"type": "response.output_text.delta",
|
|
||||||
"item_id": message_state["item_id"],
|
|
||||||
"output_index": message_state["output_index"],
|
|
||||||
"content_index": 0,
|
|
||||||
"delta": full_text,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
close_items: list[tuple[int, str, dict[str, Any]]] = []
|
close_items: list[tuple[int, str, dict[str, Any]]] = []
|
||||||
if reasoning_state["opened"]:
|
if reasoning_state["opened"]:
|
||||||
|
|
@ -7569,7 +7918,15 @@ async def _responses_stream(
|
||||||
api_monitor.finish(monitor_id)
|
api_monitor.finish(monitor_id)
|
||||||
yield _sse("response.completed", completed_response)
|
yield _sse("response.completed", completed_response)
|
||||||
|
|
||||||
return _sse_streaming_response(event_generator())
|
return _SameTaskStreamingResponse(
|
||||||
|
event_generator(),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/responses")
|
@router.post("/responses")
|
||||||
|
|
@ -8185,9 +8542,17 @@ async def _anthropic_tool_stream(
|
||||||
drop_until_tool_end = False
|
drop_until_tool_end = False
|
||||||
|
|
||||||
gen = run_gen()
|
gen = run_gen()
|
||||||
|
# Concurrent disconnect watcher: the loop only polls is_disconnected()
|
||||||
|
# between events, so a client disconnect during a long prefill or
|
||||||
|
# generation step would otherwise hold the decode slot until the next
|
||||||
|
# event or a failed send. The watcher sets cancel_event so the backend
|
||||||
|
# stops promptly.
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
if await request.is_disconnected():
|
if cancel_event.is_set() or await request.is_disconnected():
|
||||||
cancel_event.set()
|
cancel_event.set()
|
||||||
return
|
return
|
||||||
event = await asyncio.to_thread(next, gen, _sentinel)
|
event = await asyncio.to_thread(next, gen, _sentinel)
|
||||||
|
|
@ -8235,6 +8600,8 @@ async def _anthropic_tool_stream(
|
||||||
if _error_event is not None:
|
if _error_event is not None:
|
||||||
yield _error_event
|
yield _error_event
|
||||||
return
|
return
|
||||||
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
|
|
||||||
stop_reason = openai_finish_to_anthropic_stop(
|
stop_reason = openai_finish_to_anthropic_stop(
|
||||||
captured_finish_reason, had_tool_calls = ends_on_tool_use
|
captured_finish_reason, had_tool_calls = ends_on_tool_use
|
||||||
|
|
@ -8271,9 +8638,17 @@ async def _anthropic_plain_stream(
|
||||||
captured_finish_reason = None
|
captured_finish_reason = None
|
||||||
|
|
||||||
gen = run_gen()
|
gen = run_gen()
|
||||||
|
# Concurrent disconnect watcher: the loop only polls is_disconnected()
|
||||||
|
# between chunks, so a client disconnect during a long prefill or
|
||||||
|
# generation step would otherwise hold the decode slot until the next
|
||||||
|
# chunk or a failed send. The watcher sets cancel_event so the backend
|
||||||
|
# stops promptly.
|
||||||
|
disconnect_watcher = asyncio.create_task(
|
||||||
|
_await_disconnect_then_cancel(request, cancel_event)
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
if await request.is_disconnected():
|
if cancel_event.is_set() or await request.is_disconnected():
|
||||||
cancel_event.set()
|
cancel_event.set()
|
||||||
return
|
return
|
||||||
cumulative = await asyncio.to_thread(next, gen, _sentinel)
|
cumulative = await asyncio.to_thread(next, gen, _sentinel)
|
||||||
|
|
@ -8296,6 +8671,8 @@ async def _anthropic_plain_stream(
|
||||||
if _error_event is not None:
|
if _error_event is not None:
|
||||||
yield _error_event
|
yield _error_event
|
||||||
return
|
return
|
||||||
|
finally:
|
||||||
|
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||||
|
|
||||||
stop_reason = openai_finish_to_anthropic_stop(captured_finish_reason, had_tool_calls = False)
|
stop_reason = openai_finish_to_anthropic_stop(captured_finish_reason, had_tool_calls = False)
|
||||||
for line in emitter.finish(stop_reason = stop_reason, stop_sequence = None):
|
for line in emitter.finish(stop_reason = stop_reason, stop_sequence = None):
|
||||||
|
|
@ -9139,6 +9516,19 @@ async def _openai_passthrough_stream(
|
||||||
response ``id``, ``finish_reason`` (including ``"tool_calls"``),
|
response ``id``, ``finish_reason`` (including ``"tool_calls"``),
|
||||||
``delta.tool_calls``, and any client-requested trailing ``usage`` chunk so
|
``delta.tool_calls``, and any client-requested trailing ``usage`` chunk so
|
||||||
the client sees a standard OpenAI response.
|
the client sees a standard OpenAI response.
|
||||||
|
|
||||||
|
Reasoning/tool-call extraction here is delegated to llama-server: this path
|
||||||
|
forwards to its ``/v1/chat/completions`` (Studio launches with ``--jinja``
|
||||||
|
and ``--reasoning-format auto``), which parses Gemma-native ``<think>`` into
|
||||||
|
``reasoning_content`` and ``<|tool_call>`` into structured ``tool_calls``
|
||||||
|
server-side, so the relayed ``delta.content`` carries no raw markup. This is
|
||||||
|
deliberately NOT re-parsed with the local reasoning extractor / Gemma parser
|
||||||
|
(verified end to end on the current llama.cpp build), unlike Studio's own
|
||||||
|
``/completion``-level generation paths, which must parse the raw text
|
||||||
|
themselves. The dependency is on llama.cpp's chat parser: if a future build
|
||||||
|
or chat template stops splitting ``<think>``/``<|tool_call>``, raw markup
|
||||||
|
would relay into ``content`` and this path would need the local extractor as
|
||||||
|
a safety net.
|
||||||
"""
|
"""
|
||||||
target_url = f"{llama_backend.base_url}/v1/chat/completions"
|
target_url = f"{llama_backend.base_url}/v1/chat/completions"
|
||||||
body = _build_openai_passthrough_body(
|
body = _build_openai_passthrough_body(
|
||||||
|
|
@ -9187,7 +9577,7 @@ async def _openai_passthrough_stream(
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
return StreamingResponse(
|
return _SameTaskStreamingResponse(
|
||||||
iter(()),
|
iter(()),
|
||||||
media_type = "text/event-stream",
|
media_type = "text/event-stream",
|
||||||
headers = {
|
headers = {
|
||||||
|
|
@ -9238,6 +9628,29 @@ async def _openai_passthrough_stream(
|
||||||
_await_disconnect_then_close(request, resp, cancel_event)
|
_await_disconnect_then_close(request, resp, cancel_event)
|
||||||
)
|
)
|
||||||
monitor_done = False
|
monitor_done = False
|
||||||
|
saw_finish_reason = False
|
||||||
|
saw_done = False
|
||||||
|
saw_stream_error = False
|
||||||
|
saw_tool_call_delta = False
|
||||||
|
last_chunk_id = completion_id
|
||||||
|
last_chunk_model = model_name
|
||||||
|
last_chunk_created = int(time.time())
|
||||||
|
|
||||||
|
def _synthetic_finish_line() -> str:
|
||||||
|
finish_reason = "tool_calls" if saw_tool_call_delta else "stop"
|
||||||
|
chunk = ChatCompletionChunk(
|
||||||
|
id = last_chunk_id,
|
||||||
|
created = last_chunk_created,
|
||||||
|
model = last_chunk_model,
|
||||||
|
choices = [
|
||||||
|
ChunkChoice(
|
||||||
|
delta = ChoiceDelta(),
|
||||||
|
finish_reason = finish_reason,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
return f"data: {chunk.model_dump_json(exclude_none = True)}"
|
||||||
|
|
||||||
try:
|
try:
|
||||||
lines_iter = resp.aiter_lines()
|
lines_iter = resp.aiter_lines()
|
||||||
async for raw_line in _aiter_llama_stream_items(
|
async for raw_line in _aiter_llama_stream_items(
|
||||||
|
|
@ -9251,23 +9664,117 @@ async def _openai_passthrough_stream(
|
||||||
continue
|
continue
|
||||||
if not raw_line.startswith("data: "):
|
if not raw_line.startswith("data: "):
|
||||||
continue
|
continue
|
||||||
|
data_text = raw_line[6:].strip()
|
||||||
|
if data_text == "[DONE]":
|
||||||
|
saw_done = True
|
||||||
|
if (
|
||||||
|
not saw_finish_reason
|
||||||
|
and not saw_stream_error
|
||||||
|
and not cancel_event.is_set()
|
||||||
|
):
|
||||||
|
finish_line = _synthetic_finish_line()
|
||||||
|
_monitor_openai_sse_line(
|
||||||
|
monitor_id,
|
||||||
|
finish_line,
|
||||||
|
llama_backend.context_length,
|
||||||
|
)
|
||||||
|
yield finish_line + "\n\n"
|
||||||
|
saw_finish_reason = True
|
||||||
|
_monitor_openai_sse_line(
|
||||||
|
monitor_id,
|
||||||
|
raw_line,
|
||||||
|
llama_backend.context_length,
|
||||||
|
)
|
||||||
|
yield raw_line + "\n\n"
|
||||||
|
monitor_done = True
|
||||||
|
break
|
||||||
# Honor parallel_tool_calls=false (best-effort): drop tool_call
|
# Honor parallel_tool_calls=false (best-effort): drop tool_call
|
||||||
# deltas with index>=1 so only the first call streams. Only
|
# deltas with index>=1 so only the first call streams. Only
|
||||||
# lines carrying tool_calls are reparsed; everything else is
|
# lines carrying tool_calls are reparsed; everything else is
|
||||||
# relayed byte-for-byte.
|
# relayed byte-for-byte.
|
||||||
if payload.parallel_tool_calls is False and '"tool_calls"' in raw_line:
|
if payload.parallel_tool_calls is False and '"tool_calls"' in raw_line:
|
||||||
raw_line = _cap_parallel_tool_calls_sse_line(raw_line)
|
raw_line = _cap_parallel_tool_calls_sse_line(raw_line)
|
||||||
|
data_text = raw_line[6:].strip()
|
||||||
|
try:
|
||||||
|
chunk_data = json.loads(data_text)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
chunk_data = None
|
||||||
|
if isinstance(chunk_data, dict):
|
||||||
|
if isinstance(chunk_data.get("id"), str):
|
||||||
|
last_chunk_id = chunk_data["id"]
|
||||||
|
if isinstance(chunk_data.get("model"), str):
|
||||||
|
last_chunk_model = chunk_data["model"]
|
||||||
|
if isinstance(chunk_data.get("created"), int):
|
||||||
|
last_chunk_created = chunk_data["created"]
|
||||||
|
choices = chunk_data.get("choices")
|
||||||
|
if isinstance(choices, list) and choices:
|
||||||
|
choice = choices[0]
|
||||||
|
if isinstance(choice, dict):
|
||||||
|
if choice.get("finish_reason"):
|
||||||
|
saw_finish_reason = True
|
||||||
|
delta = choice.get("delta")
|
||||||
|
if isinstance(delta, dict) and delta.get("tool_calls"):
|
||||||
|
saw_tool_call_delta = True
|
||||||
|
# Detect an upstream error chunk independently of API
|
||||||
|
# monitoring: when monitor_id is None (skip_api_monitor),
|
||||||
|
# _monitor_openai_sse_line returns before inspecting the
|
||||||
|
# error, so without this the synthetic-finish guard would
|
||||||
|
# emit a successful finish_reason after a failed stream.
|
||||||
|
if _monitor_openai_error_message(chunk_data):
|
||||||
|
saw_stream_error = True
|
||||||
monitor_event = _monitor_openai_sse_line(
|
monitor_event = _monitor_openai_sse_line(
|
||||||
monitor_id,
|
monitor_id,
|
||||||
raw_line,
|
raw_line,
|
||||||
llama_backend.context_length,
|
llama_backend.context_length,
|
||||||
)
|
)
|
||||||
|
if monitor_event == "error":
|
||||||
|
saw_stream_error = True
|
||||||
|
# If a trailing usage-only chunk (include_usage) arrives before
|
||||||
|
# any finish chunk, emit the synthetic finish first so the order
|
||||||
|
# stays finish -> usage -> [DONE], matching the other streams.
|
||||||
|
if (
|
||||||
|
isinstance(chunk_data, dict)
|
||||||
|
and chunk_data.get("usage")
|
||||||
|
and not (
|
||||||
|
isinstance(chunk_data.get("choices"), list) and chunk_data["choices"]
|
||||||
|
)
|
||||||
|
and not saw_finish_reason
|
||||||
|
and not saw_stream_error
|
||||||
|
and not cancel_event.is_set()
|
||||||
|
):
|
||||||
|
finish_line = _synthetic_finish_line()
|
||||||
|
_monitor_openai_sse_line(
|
||||||
|
monitor_id, finish_line, llama_backend.context_length
|
||||||
|
)
|
||||||
|
yield finish_line + "\n\n"
|
||||||
|
saw_finish_reason = True
|
||||||
# Relay verbatim to preserve llama-server's native id,
|
# Relay verbatim to preserve llama-server's native id,
|
||||||
# finish_reason, delta.tool_calls, and usage chunks.
|
# finish_reason, delta.tool_calls, and usage chunks.
|
||||||
yield raw_line + "\n\n"
|
yield raw_line + "\n\n"
|
||||||
if monitor_event == "done" or raw_line[6:].strip() == "[DONE]":
|
if monitor_event == "done":
|
||||||
monitor_done = True
|
monitor_done = True
|
||||||
break
|
break
|
||||||
|
if not saw_done and not saw_stream_error and not cancel_event.is_set():
|
||||||
|
# Synthesize a finish chunk only if one was not already
|
||||||
|
# emitted (e.g. before a trailing usage-only chunk), but
|
||||||
|
# always close with [DONE] whenever the upstream omitted it,
|
||||||
|
# so the stream ends on the [DONE] sentinel either way.
|
||||||
|
if not saw_finish_reason:
|
||||||
|
finish_line = _synthetic_finish_line()
|
||||||
|
_monitor_openai_sse_line(
|
||||||
|
monitor_id,
|
||||||
|
finish_line,
|
||||||
|
llama_backend.context_length,
|
||||||
|
)
|
||||||
|
yield finish_line + "\n\n"
|
||||||
|
done_line = "data: [DONE]"
|
||||||
|
_monitor_openai_sse_line(
|
||||||
|
monitor_id,
|
||||||
|
done_line,
|
||||||
|
llama_backend.context_length,
|
||||||
|
)
|
||||||
|
yield done_line + "\n\n"
|
||||||
|
monitor_done = True
|
||||||
if not monitor_done:
|
if not monitor_done:
|
||||||
api_monitor.finish(
|
api_monitor.finish(
|
||||||
monitor_id,
|
monitor_id,
|
||||||
|
|
@ -9303,7 +9810,24 @@ async def _openai_passthrough_stream(
|
||||||
)
|
)
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
return _sse_streaming_response(_stream())
|
async def _unstarted_cleanup() -> None:
|
||||||
|
# Client disconnected before the body stream started, so _stream()'s
|
||||||
|
# finally never ran. Release the eagerly-opened upstream resp/client
|
||||||
|
# and the cancel-registry entry here; the watchers and line iterator
|
||||||
|
# are created inside _stream(), so there is nothing else to close.
|
||||||
|
await _aclose_stream_resources(resp = resp, client = client)
|
||||||
|
_tracker.__exit__(None, None, None)
|
||||||
|
|
||||||
|
return _SameTaskStreamingResponse(
|
||||||
|
_stream(),
|
||||||
|
media_type = "text/event-stream",
|
||||||
|
headers = {
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "close",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
unstarted_cleanup = _unstarted_cleanup,
|
||||||
|
)
|
||||||
except BaseException:
|
except BaseException:
|
||||||
_tracker.__exit__(None, None, None)
|
_tracker.__exit__(None, None, None)
|
||||||
raise
|
raise
|
||||||
|
|
|
||||||
174
studio/backend/tests/test_gemma_tool_parse_edge_cases.py
Normal file
174
studio/backend/tests/test_gemma_tool_parse_edge_cases.py
Normal file
|
|
@ -0,0 +1,174 @@
|
||||||
|
# SPDX-License-Identifier: AGPL-3.0-only
|
||||||
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||||
|
|
||||||
|
"""Edge cases in Gemma-native tool-call parsing.
|
||||||
|
|
||||||
|
Covers two failure modes:
|
||||||
|
1. A bare (unquoted) string argument that contains a comma, e.g.
|
||||||
|
``location:New York, NY`` -- the comma must not be treated as the next
|
||||||
|
key boundary, or the whole call is dropped.
|
||||||
|
2. A tool-call marker that appears INSIDE another call's argument string is
|
||||||
|
data, not a real call, so it must not be promoted to a second tool call.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
|
||||||
|
if _BACKEND_DIR not in sys.path:
|
||||||
|
sys.path.insert(0, _BACKEND_DIR)
|
||||||
|
|
||||||
|
from core.inference.tool_call_parser import parse_tool_calls_from_text
|
||||||
|
from core.tool_healing import strip_tool_call_markup
|
||||||
|
|
||||||
|
|
||||||
|
def _args(call: dict) -> dict:
|
||||||
|
return json.loads(call["function"]["arguments"])
|
||||||
|
|
||||||
|
|
||||||
|
def test_bare_string_argument_with_comma_is_kept():
|
||||||
|
calls = parse_tool_calls_from_text(
|
||||||
|
"<|tool_call>call:get_weather{location:New York, NY,unit:celsius}<tool_call|>"
|
||||||
|
)
|
||||||
|
assert len(calls) == 1, calls
|
||||||
|
assert calls[0]["function"]["name"] == "get_weather"
|
||||||
|
assert _args(calls[0]) == {"location": "New York, NY", "unit": "celsius"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_normal_multi_key_arguments_still_split():
|
||||||
|
calls = parse_tool_calls_from_text('<|tool_call>call:f{a:1,b:hello,c:"x,y"}<tool_call|>')
|
||||||
|
assert len(calls) == 1, calls
|
||||||
|
# Numbers stay numeric, bare strings get quoted, an explicit quoted comma
|
||||||
|
# stays inside its value.
|
||||||
|
assert _args(calls[0]) == {"a": 1, "b": "hello", "c": "x,y"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_bare_value_with_timestamps_after_comma_is_kept():
|
||||||
|
# A comma followed by digits-then-colon (a timestamp/ratio) is value text,
|
||||||
|
# not a new key, so the whole query must be preserved as one argument.
|
||||||
|
calls = parse_tool_calls_from_text(
|
||||||
|
"<|tool_call>call:remind{query:meet at 10:00, 11:00 tomorrow,priority:high}<tool_call|>"
|
||||||
|
)
|
||||||
|
assert len(calls) == 1, calls
|
||||||
|
assert _args(calls[0]) == {"query": "meet at 10:00, 11:00 tomorrow", "priority": "high"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_marker_inside_json_argument_is_not_a_second_call():
|
||||||
|
# A python call whose `code` argument contains a Gemma marker string. The
|
||||||
|
# marker is data and must not execute as a second `terminal` call.
|
||||||
|
content = (
|
||||||
|
'<tool_call>{"name":"python","arguments":{"code":'
|
||||||
|
'"x = 1 # <|tool_call>call:terminal{command:ls}<tool_call|>"}}</tool_call>'
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert [c["function"]["name"] for c in calls] == ["python"], calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_two_separate_gemma_calls_both_parse():
|
||||||
|
content = "<|tool_call>call:a{x:1}<tool_call|> and <|tool_call>call:b{y:2}<tool_call|>"
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert [c["function"]["name"] for c in calls] == ["a", "b"], calls
|
||||||
|
assert _args(calls[0]) == {"x": 1}
|
||||||
|
assert _args(calls[1]) == {"y": 2}
|
||||||
|
|
||||||
|
|
||||||
|
def test_mixed_format_calls_preserve_document_order():
|
||||||
|
# A Gemma-native call precedes a JSON-format call in the text; tools execute
|
||||||
|
# in returned order, so `create` must come before `read`.
|
||||||
|
content = (
|
||||||
|
"<|tool_call>call:create{path:a}<tool_call|> then "
|
||||||
|
'<tool_call>{"name":"read","arguments":{"path":"a"}}</tool_call>'
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert [c["function"]["name"] for c in calls] == ["create", "read"], calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_json_marker_inside_gemma_argument_is_not_a_second_call():
|
||||||
|
# The reverse of the JSON-outer case: a JSON-style marker inside a Gemma
|
||||||
|
# call's quoted argument is code text, not a second `terminal` call.
|
||||||
|
content = (
|
||||||
|
'<|tool_call>call:python{code:<|"|>'
|
||||||
|
'print(<tool_call>{"name":"terminal","arguments":{"command":"ls"}}</tool_call>)'
|
||||||
|
'<|"|>}<tool_call|>'
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert [c["function"]["name"] for c in calls] == ["python"], calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_nested_gemma_marker_in_unquoted_arg_does_not_run_inner_call():
|
||||||
|
# An UNQUOTED Gemma value containing a literal marker: the outer object fails
|
||||||
|
# to normalize (the inner braces/marker break the JSON), but the inner marker
|
||||||
|
# is nested in the outer candidate span, so it must not be promoted to a
|
||||||
|
# standalone `terminal` call. The safe outcome is no executed tool call.
|
||||||
|
content = "<|tool_call>call:python{code:<|tool_call>call:terminal{command:ls}<tool_call|>}<tool_call|>"
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert "terminal" not in [c["function"]["name"] for c in calls], calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_bare_string_array_argument_is_quoted():
|
||||||
|
# Gemma may emit an array of bare strings without per-element quotes; they
|
||||||
|
# must be quoted so the call is not dropped.
|
||||||
|
calls = parse_tool_calls_from_text("<|tool_call>call:label{labels:[bug,ui]}<tool_call|>")
|
||||||
|
assert len(calls) == 1, calls
|
||||||
|
assert _args(calls[0]) == {"labels": ["bug", "ui"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_array_keeps_numbers_and_quoted_elements():
|
||||||
|
calls = parse_tool_calls_from_text(
|
||||||
|
'<|tool_call>call:f{nums:[1,2],tags:[<|"|>a,b<|"|>,c]}<tool_call|>'
|
||||||
|
)
|
||||||
|
assert _args(calls[0]) == {"nums": [1, 2], "tags": ["a,b", "c"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_array_of_objects_is_normalised():
|
||||||
|
# Arrays of objects are a common tool-schema shape; their (unquoted) keys and
|
||||||
|
# bare values must be normalised too, not left verbatim, or the call drops.
|
||||||
|
calls = parse_tool_calls_from_text(
|
||||||
|
"<|tool_call>call:batch{items:[{path:a,mode:r},{path:b,mode:w}]}<tool_call|>"
|
||||||
|
)
|
||||||
|
assert len(calls) == 1, calls
|
||||||
|
assert _args(calls[0]) == {"items": [{"path": "a", "mode": "r"}, {"path": "b", "mode": "w"}]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_nested_array_elements_are_normalised():
|
||||||
|
calls = parse_tool_calls_from_text("<|tool_call>call:grid{cells:[[a,b],[c,d]]}<tool_call|>")
|
||||||
|
assert _args(calls[0]) == {"cells": [["a", "b"], ["c", "d"]]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemma_marker_inside_xml_parameter_is_not_a_second_call():
|
||||||
|
# An XML-style <function=...> call whose <parameter=code> value contains a
|
||||||
|
# Gemma marker: the marker is the parameter's data, not a separate terminal
|
||||||
|
# call, so only the python call must be returned.
|
||||||
|
content = (
|
||||||
|
"<tool_call><function=python><parameter=code>"
|
||||||
|
"x = 1 # <|tool_call>call:terminal{command:ls}<tool_call|>"
|
||||||
|
"</parameter></function></tool_call>"
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert [c["function"]["name"] for c in calls] == ["python"], calls
|
||||||
|
assert "terminal" in _args(calls[0])["code"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_json_marker_inside_xml_parameter_is_not_a_second_call():
|
||||||
|
content = (
|
||||||
|
"<tool_call><function=python><parameter=code>"
|
||||||
|
'run(<tool_call>{"name":"terminal","arguments":{"command":"ls"}}</tool_call>)'
|
||||||
|
"</parameter></function></tool_call>"
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(content)
|
||||||
|
assert [c["function"]["name"] for c in calls] == ["python"], calls
|
||||||
|
|
||||||
|
|
||||||
|
def test_gemma_close_marker_inside_quoted_arg_is_not_leaked_when_stripping():
|
||||||
|
# A literal <tool_call|> inside a <|"|>-quoted argument must not truncate the
|
||||||
|
# span: the parser keeps it as data, and stripping must remove the whole span
|
||||||
|
# (brace/quote-aware), not stop at the inner marker and leak the suffix.
|
||||||
|
text = '<|tool_call>call:python{code:<|"|>print("<tool_call|>")<|"|>}<tool_call|>'
|
||||||
|
calls = parse_tool_calls_from_text(text)
|
||||||
|
assert len(calls) == 1, calls
|
||||||
|
assert _args(calls[0]) == {"code": 'print("<tool_call|>")'}
|
||||||
|
assert strip_tool_call_markup("before " + text + " after") == "before after"
|
||||||
|
assert strip_tool_call_markup("before " + text + " after", final = True) == "before after"
|
||||||
|
|
@ -5,6 +5,7 @@ import asyncio
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
|
import threading
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
_backend = os.path.join(os.path.dirname(__file__), "..")
|
_backend = os.path.join(os.path.dirname(__file__), "..")
|
||||||
|
|
@ -40,6 +41,123 @@ def test_stream_first_item_deadline_after_headers():
|
||||||
asyncio.run(_run())
|
asyncio.run(_run())
|
||||||
|
|
||||||
|
|
||||||
|
def test_stream_first_item_deadline_does_not_hop_tasks():
|
||||||
|
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_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_stops_on_known_disconnect_before_read():
|
||||||
|
async def _run():
|
||||||
|
state = SimpleNamespace(disconnect_checks = 0)
|
||||||
|
cancel_event = threading.Event()
|
||||||
|
|
||||||
|
class _Request:
|
||||||
|
async def is_disconnected(self):
|
||||||
|
state.disconnect_checks += 1
|
||||||
|
return True
|
||||||
|
|
||||||
|
class _Unread:
|
||||||
|
async def __anext__(self):
|
||||||
|
raise AssertionError("stream should stop before reading upstream")
|
||||||
|
|
||||||
|
async for _ in inf_mod._aiter_llama_stream_items(
|
||||||
|
_Unread(),
|
||||||
|
cancel_event = cancel_event,
|
||||||
|
request = _Request(),
|
||||||
|
first_token_deadline = time.monotonic() + 1,
|
||||||
|
):
|
||||||
|
raise AssertionError("stream should stop after disconnect")
|
||||||
|
|
||||||
|
assert cancel_event.is_set()
|
||||||
|
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())
|
||||||
|
|
||||||
|
|
||||||
def test_preheader_send_cleanup_on_disconnect_and_cancel():
|
def test_preheader_send_cleanup_on_disconnect_and_cancel():
|
||||||
async def _run(cancel_parent):
|
async def _run(cancel_parent):
|
||||||
state = SimpleNamespace(disconnected = False, closed = False, cancelled = False)
|
state = SimpleNamespace(disconnected = False, closed = False, cancelled = False)
|
||||||
|
|
|
||||||
|
|
@ -423,6 +423,46 @@ def test_tool_healing_strip_handles_hyphenated_function_names():
|
||||||
assert out == "before after"
|
assert out == "before after"
|
||||||
|
|
||||||
|
|
||||||
|
def test_tool_healing_strip_handles_gemma_native_tool_call():
|
||||||
|
from core.tool_healing import strip_tool_call_markup
|
||||||
|
out = strip_tool_call_markup(
|
||||||
|
'before <|tool_call>call:mcp__srv__list-issues{repo:"octocat/hello"}<tool_call|> after'
|
||||||
|
)
|
||||||
|
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 <tool_call|> after") == "before after"
|
||||||
|
assert strip_tool_call_markup("before <tool_call|> 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
|
||||||
|
|
||||||
|
calls = parse_tool_calls_from_text(
|
||||||
|
r'<|tool_call>call:ls{path:<|"|>C:\Users\wasim\repo<|"|>}<tool_call|>'
|
||||||
|
)
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0]["function"]["name"] == "ls"
|
||||||
|
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 = (
|
||||||
|
"<tool_call>"
|
||||||
|
+ _json.dumps({"name": "python", "arguments": {"code": "print('<|\"|>')"}})
|
||||||
|
+ "</tool_call>"
|
||||||
|
)
|
||||||
|
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):
|
def test_gguf_allow_list_blocks_unadvertised_tool(monkeypatch):
|
||||||
"""A tool call not in the per-request list must be refused by the GGUF
|
"""A tool call not in the per-request list must be refused by the GGUF
|
||||||
agentic loop (mirroring the safetensors path)."""
|
agentic loop (mirroring the safetensors path)."""
|
||||||
|
|
|
||||||
|
|
@ -48,6 +48,7 @@ from routes.inference import (
|
||||||
_openai_passthrough_stream,
|
_openai_passthrough_stream,
|
||||||
_openai_stream_usage_chunk,
|
_openai_stream_usage_chunk,
|
||||||
_proxy_to_external_provider,
|
_proxy_to_external_provider,
|
||||||
|
_SameTaskStreamingResponse,
|
||||||
_set_or_prepend_system_message,
|
_set_or_prepend_system_message,
|
||||||
openai_completions,
|
openai_completions,
|
||||||
openai_embeddings,
|
openai_embeddings,
|
||||||
|
|
@ -1245,6 +1246,79 @@ class TestGgufVisionToolRouting:
|
||||||
|
|
||||||
return TestGgufVisionToolRouting._drive(_consume())
|
return TestGgufVisionToolRouting._drive(_consume())
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _sse_payloads(chunks):
|
||||||
|
payloads = []
|
||||||
|
for chunk in chunks:
|
||||||
|
if isinstance(chunk, bytes):
|
||||||
|
chunk = chunk.decode()
|
||||||
|
for line in str(chunk).splitlines():
|
||||||
|
if not line.startswith("data: "):
|
||||||
|
continue
|
||||||
|
data = line.removeprefix("data: ")
|
||||||
|
if data == "[DONE]":
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
payloads.append(json.loads(data))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
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):
|
def test_image_request_with_enabled_tools_enters_gguf_tool_loop(self, monkeypatch):
|
||||||
import routes.inference as inf_mod
|
import routes.inference as inf_mod
|
||||||
|
|
||||||
|
|
@ -1390,6 +1464,152 @@ class TestGgufVisionToolRouting:
|
||||||
assert "confirm_tool_calls requires stream=true" in entry["error"]
|
assert "confirm_tool_calls requires stream=true" in entry["error"]
|
||||||
assert monitor.active_count() == 0
|
assert monitor.active_count() == 0
|
||||||
|
|
||||||
|
def test_standard_gguf_stream_splits_reasoning_content(self, monkeypatch):
|
||||||
|
def _generate(**_kwargs):
|
||||||
|
yield "<thi"
|
||||||
|
yield "<think>plan"
|
||||||
|
yield "<think>plan</think>vis"
|
||||||
|
yield "<think>plan</think>visible"
|
||||||
|
yield {
|
||||||
|
"type": "metadata",
|
||||||
|
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
|
||||||
|
result = self._run_gguf_case(
|
||||||
|
monkeypatch,
|
||||||
|
generate = _generate,
|
||||||
|
payload_kwargs = {"stream": True},
|
||||||
|
)
|
||||||
|
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("<think>" not in d.get("content", "") for d in deltas)
|
||||||
|
[entry] = result.monitor.snapshot()
|
||||||
|
assert entry["reply"] == "visible"
|
||||||
|
|
||||||
|
def test_reasoning_capable_gguf_stream_splits_reasoning_by_default(self, monkeypatch):
|
||||||
|
def _generate(**_kwargs):
|
||||||
|
yield "<think>plan</think>visible"
|
||||||
|
yield {
|
||||||
|
"type": "metadata",
|
||||||
|
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
|
||||||
|
result = self._run_gguf_case(
|
||||||
|
monkeypatch,
|
||||||
|
generate = _generate,
|
||||||
|
payload_kwargs = {"stream": True},
|
||||||
|
backend_kwargs = {"reasoning_always_on": False},
|
||||||
|
)
|
||||||
|
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] = result.monitor.snapshot()
|
||||||
|
assert entry["reply"] == "visible"
|
||||||
|
|
||||||
|
def test_reasoning_capable_gguf_stream_sanitizes_think_tags_when_disabled(self, monkeypatch):
|
||||||
|
def _generate(**_kwargs):
|
||||||
|
yield "<think>leaked</think>visible"
|
||||||
|
yield {
|
||||||
|
"type": "metadata",
|
||||||
|
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
|
||||||
|
result = self._run_gguf_case(
|
||||||
|
monkeypatch,
|
||||||
|
generate = _generate,
|
||||||
|
payload_kwargs = {"stream": True, "enable_thinking": False},
|
||||||
|
backend_kwargs = {"reasoning_always_on": False},
|
||||||
|
)
|
||||||
|
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("<think>" not in d.get("content", "") for d in deltas)
|
||||||
|
[entry] = result.monitor.snapshot()
|
||||||
|
assert entry["reply"] == "visible"
|
||||||
|
|
||||||
|
def test_gguf_tool_stream_splits_reasoning_and_strips_gemma_tool_marker(self, monkeypatch):
|
||||||
|
def _tools(**_kwargs):
|
||||||
|
yield {
|
||||||
|
"type": "content",
|
||||||
|
"text": '<think>plan</think>visible <|tool_call>call:terminal{command:"ls"}<tool_call|>',
|
||||||
|
}
|
||||||
|
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": "list files"}],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
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] = 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 "<think>plan</think>visible"
|
||||||
|
yield {
|
||||||
|
"type": "metadata",
|
||||||
|
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
|
||||||
|
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] = result.monitor.snapshot()
|
||||||
|
assert entry["reply"] == "visible"
|
||||||
|
|
||||||
def test_non_streaming_gguf_n_records_all_monitor_replies(self, monkeypatch):
|
def test_non_streaming_gguf_n_records_all_monitor_replies(self, monkeypatch):
|
||||||
import routes.inference as inf_mod
|
import routes.inference as inf_mod
|
||||||
|
|
||||||
|
|
@ -1552,6 +1772,61 @@ class TestApiMonitorProviderAndCompletionStreams:
|
||||||
async def is_disconnected(self):
|
async def is_disconnected(self):
|
||||||
return False
|
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):
|
def test_external_non_streaming_json_updates_monitor(self, monkeypatch):
|
||||||
async def _run():
|
async def _run():
|
||||||
import routes.inference as inf_mod
|
import routes.inference as inf_mod
|
||||||
|
|
@ -1980,6 +2255,7 @@ class TestApiMonitorProviderAndCompletionStreams:
|
||||||
"chatcmpl-test",
|
"chatcmpl-test",
|
||||||
monitor_id = monitor_id,
|
monitor_id = monitor_id,
|
||||||
)
|
)
|
||||||
|
assert isinstance(response, _SameTaskStreamingResponse)
|
||||||
iterator = response.body_iterator
|
iterator = response.body_iterator
|
||||||
first = await anext(iterator)
|
first = await anext(iterator)
|
||||||
assert "hello" in first
|
assert "hello" in first
|
||||||
|
|
@ -1997,6 +2273,88 @@ class TestApiMonitorProviderAndCompletionStreams:
|
||||||
|
|
||||||
asyncio.run(_run())
|
asyncio.run(_run())
|
||||||
|
|
||||||
|
def test_passthrough_stream_synthesizes_missing_finish_reason(self, monkeypatch):
|
||||||
|
async def _run():
|
||||||
|
result = await self._run_passthrough_stream(
|
||||||
|
monkeypatch,
|
||||||
|
[
|
||||||
|
(
|
||||||
|
'data: {"id":"upstream","created":123,"model":"gguf",'
|
||||||
|
'"choices":[{"index":0,"delta":{"content":"hello"}}]}'
|
||||||
|
),
|
||||||
|
"data: [DONE]",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
body = result.body
|
||||||
|
|
||||||
|
assert '"finish_reason":"stop"' in body.replace(" ", "")
|
||||||
|
assert "data: [DONE]" in body
|
||||||
|
assert result.monitor.active_count() == 0
|
||||||
|
|
||||||
|
asyncio.run(_run())
|
||||||
|
|
||||||
|
def test_passthrough_stream_synthesizes_tool_call_finish_reason(self, monkeypatch):
|
||||||
|
async def _run():
|
||||||
|
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]",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
compact = result.body.replace(" ", "")
|
||||||
|
|
||||||
|
assert '"finish_reason":"tool_calls"' in compact
|
||||||
|
assert '"finish_reason":"stop"' not in compact
|
||||||
|
assert "data: [DONE]" in result.body
|
||||||
|
assert result.monitor.active_count() == 0
|
||||||
|
|
||||||
|
asyncio.run(_run())
|
||||||
|
|
||||||
|
def test_passthrough_stream_error_done_skips_synthetic_finish_reason(self, monkeypatch):
|
||||||
|
async def _run():
|
||||||
|
result = await self._run_passthrough_stream(
|
||||||
|
monkeypatch,
|
||||||
|
[
|
||||||
|
'data: {"error":{"message":"boom","type":"server_error"}}',
|
||||||
|
"data: [DONE]",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
compact = result.body.replace(" ", "")
|
||||||
|
|
||||||
|
assert '"error":{"message":"boom","type":"server_error"}' in compact
|
||||||
|
assert '"finish_reason"' not in compact
|
||||||
|
assert "data: [DONE]" in result.body
|
||||||
|
[entry] = result.monitor.snapshot()
|
||||||
|
assert entry["status"] == "error"
|
||||||
|
assert entry["error"] == "boom"
|
||||||
|
assert result.monitor.active_count() == 0
|
||||||
|
|
||||||
|
asyncio.run(_run())
|
||||||
|
|
||||||
|
def test_passthrough_stream_error_eof_skips_synthetic_finish_reason(self, monkeypatch):
|
||||||
|
async def _run():
|
||||||
|
result = await self._run_passthrough_stream(
|
||||||
|
monkeypatch,
|
||||||
|
['data: {"error":{"message":"boom","type":"server_error"}}'],
|
||||||
|
)
|
||||||
|
compact = result.body.replace(" ", "")
|
||||||
|
|
||||||
|
assert '"error":{"message":"boom","type":"server_error"}' in compact
|
||||||
|
assert '"finish_reason"' not in compact
|
||||||
|
assert "data: [DONE]" not in result.body
|
||||||
|
[entry] = result.monitor.snapshot()
|
||||||
|
assert entry["status"] == "error"
|
||||||
|
assert entry["error"] == "boom"
|
||||||
|
assert result.monitor.active_count() == 0
|
||||||
|
|
||||||
|
asyncio.run(_run())
|
||||||
|
|
||||||
def test_passthrough_non_streaming_cancel_finalizes_monitor(self, monkeypatch):
|
def test_passthrough_non_streaming_cancel_finalizes_monitor(self, monkeypatch):
|
||||||
async def _run():
|
async def _run():
|
||||||
import routes.inference as inf_mod
|
import routes.inference as inf_mod
|
||||||
|
|
@ -2058,65 +2416,20 @@ class TestApiMonitorProviderAndCompletionStreams:
|
||||||
|
|
||||||
def test_passthrough_clean_eof_finalizes_monitor(self, monkeypatch):
|
def test_passthrough_clean_eof_finalizes_monitor(self, monkeypatch):
|
||||||
async def _run():
|
async def _run():
|
||||||
import routes.inference as inf_mod
|
result = await self._run_passthrough_stream(
|
||||||
|
monkeypatch,
|
||||||
class Request:
|
['data: {"choices":[{"delta":{"content":"hello"}}]}'],
|
||||||
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",
|
|
||||||
)
|
|
||||||
payload = ChatCompletionRequest(
|
|
||||||
model = "default",
|
|
||||||
messages = [ChatMessage(role = "user", content = "hi")],
|
|
||||||
stream = True,
|
|
||||||
tools = [
|
|
||||||
{
|
|
||||||
"type": "function",
|
|
||||||
"function": {
|
|
||||||
"name": "lookup",
|
|
||||||
"parameters": {"type": "object", "properties": {}},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
|
chunks = result.chunks
|
||||||
|
|
||||||
response = await _openai_passthrough_stream(
|
assert chunks[0] == 'data: {"choices":[{"delta":{"content":"hello"}}]}\n\n'
|
||||||
Request(),
|
compact = "".join(chunks).replace(" ", "")
|
||||||
threading.Event(),
|
assert '"finish_reason":"stop"' in compact
|
||||||
SimpleNamespace(
|
assert chunks[-1] == "data: [DONE]\n\n"
|
||||||
base_url = "http://llama.test",
|
[entry] = result.monitor.snapshot()
|
||||||
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)
|
|
||||||
|
|
||||||
assert chunks == ['data: {"choices":[{"delta":{"content":"hello"}}]}\n\n']
|
|
||||||
[entry] = monitor.snapshot()
|
|
||||||
assert entry["status"] == "completed"
|
assert entry["status"] == "completed"
|
||||||
assert entry["reply"] == "hello"
|
assert entry["reply"] == "hello"
|
||||||
assert monitor.active_count() == 0
|
assert result.monitor.active_count() == 0
|
||||||
|
|
||||||
asyncio.run(_run())
|
asyncio.run(_run())
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -59,8 +59,10 @@ from models.inference import (
|
||||||
ResponsesUsage,
|
ResponsesUsage,
|
||||||
)
|
)
|
||||||
from routes.inference import (
|
from routes.inference import (
|
||||||
|
_SameTaskStreamingResponse,
|
||||||
_build_chat_request,
|
_build_chat_request,
|
||||||
_chat_tool_calls_to_responses_output,
|
_chat_tool_calls_to_responses_output,
|
||||||
|
_extract_responses_reasoning,
|
||||||
_normalise_responses_input,
|
_normalise_responses_input,
|
||||||
_responses_tool_output_content,
|
_responses_tool_output_content,
|
||||||
_responses_non_streaming,
|
_responses_non_streaming,
|
||||||
|
|
@ -782,6 +784,15 @@ class TestResponsesNonStreamingAdapter:
|
||||||
assert "<think>" not in body["output"][1]["content"][0]["text"]
|
assert "<think>" not in body["output"][1]["content"][0]["text"]
|
||||||
assert "</think>" not in body["output"][1]["content"][0]["text"]
|
assert "</think>" not in body["output"][1]["content"][0]["text"]
|
||||||
|
|
||||||
|
def test_unclosed_think_block_extracts_as_reasoning(self):
|
||||||
|
reasoning, visible = _extract_responses_reasoning(
|
||||||
|
"<think>partial plan",
|
||||||
|
parse_think_markers = True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert reasoning == "partial plan"
|
||||||
|
assert visible == ""
|
||||||
|
|
||||||
def test_monitor_records_translated_visible_text(self, monkeypatch):
|
def test_monitor_records_translated_visible_text(self, monkeypatch):
|
||||||
import routes.inference as inf_mod
|
import routes.inference as inf_mod
|
||||||
|
|
||||||
|
|
@ -927,6 +938,38 @@ class TestResponsesNonStreamingAdapter:
|
||||||
assert [item["type"] for item in body["output"]] == ["message"]
|
assert [item["type"] for item in body["output"]] == ["message"]
|
||||||
assert body["output"][0]["content"][0]["text"] == "show <think>x</think> tags"
|
assert body["output"][0]["content"][0]["text"] == "show <think>x</think> tags"
|
||||||
|
|
||||||
|
def test_reasoning_capable_gguf_parses_think_tags_by_default(self, monkeypatch):
|
||||||
|
body = self._run_with_message(
|
||||||
|
monkeypatch,
|
||||||
|
{"content": "<think>plan</think>answer"},
|
||||||
|
llama_backend = SimpleNamespace(
|
||||||
|
is_loaded = True,
|
||||||
|
reasoning_always_on = False,
|
||||||
|
supports_reasoning = True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert [item["type"] for item in body["output"]] == ["reasoning", "message"]
|
||||||
|
assert body["output"][0]["content"] == [{"type": "reasoning_text", "text": "plan"}]
|
||||||
|
assert body["output"][1]["content"][0]["text"] == "answer"
|
||||||
|
|
||||||
|
def test_reasoning_capable_gguf_sanitizes_think_tags_when_disabled(self, monkeypatch):
|
||||||
|
payload = ResponsesRequest(input = "hi", reasoning = {"effort": "none"})
|
||||||
|
body = self._run_with_message(
|
||||||
|
monkeypatch,
|
||||||
|
{"content": "<think>leaked</think>answer"},
|
||||||
|
payload = payload,
|
||||||
|
llama_backend = SimpleNamespace(
|
||||||
|
is_loaded = True,
|
||||||
|
reasoning_always_on = False,
|
||||||
|
supports_reasoning = True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert [item["type"] for item in body["output"]] == ["reasoning", "message"]
|
||||||
|
assert body["output"][0]["content"] == [{"type": "reasoning_text", "text": "leaked"}]
|
||||||
|
assert body["output"][1]["content"][0]["text"] == "answer"
|
||||||
|
|
||||||
def test_structured_reasoning_content_extracts_text_parts(self, monkeypatch):
|
def test_structured_reasoning_content_extracts_text_parts(self, monkeypatch):
|
||||||
body = self._run_with_message(
|
body = self._run_with_message(
|
||||||
monkeypatch,
|
monkeypatch,
|
||||||
|
|
@ -949,7 +992,7 @@ class TestResponsesNonStreamingAdapter:
|
||||||
assert [item["type"] for item in body["output"]] == ["message"]
|
assert [item["type"] for item in body["output"]] == ["message"]
|
||||||
assert body["output"][0]["content"][0]["text"] == "33"
|
assert body["output"][0]["content"][0]["text"] == "33"
|
||||||
|
|
||||||
def test_reasoning_only_is_also_visible_message_text(self, monkeypatch):
|
def test_reasoning_only_stays_out_of_visible_message_text(self, monkeypatch):
|
||||||
payload = ResponsesRequest(input = "hi", reasoning = {"effort": "high"})
|
payload = ResponsesRequest(input = "hi", reasoning = {"effort": "high"})
|
||||||
body = self._run_with_message(
|
body = self._run_with_message(
|
||||||
monkeypatch,
|
monkeypatch,
|
||||||
|
|
@ -957,9 +1000,8 @@ class TestResponsesNonStreamingAdapter:
|
||||||
payload = payload,
|
payload = payload,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert [item["type"] for item in body["output"]] == ["reasoning", "message"]
|
assert [item["type"] for item in body["output"]] == ["reasoning"]
|
||||||
assert body["output"][0]["content"][0]["text"] == "plan"
|
assert body["output"][0]["content"][0]["text"] == "plan"
|
||||||
assert body["output"][1]["content"][0]["text"] == "plan"
|
|
||||||
|
|
||||||
|
|
||||||
# =====================================================================
|
# =====================================================================
|
||||||
|
|
@ -1033,6 +1075,36 @@ class TestResponsesStreamAdapter:
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_stream_response_avoids_legacy_receive_watcher(self, monkeypatch):
|
||||||
|
self._install_stream_mock(
|
||||||
|
monkeypatch,
|
||||||
|
[{"choices": [{"delta": {"content": "33"}}]}],
|
||||||
|
)
|
||||||
|
payload = ResponsesRequest(input = "hi", stream = True)
|
||||||
|
messages = [ChatMessage(role = "user", content = "hi")]
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
response = await _responses_stream(payload, messages, self._Request())
|
||||||
|
assert isinstance(response, _SameTaskStreamingResponse)
|
||||||
|
|
||||||
|
sent = []
|
||||||
|
|
||||||
|
async def receive():
|
||||||
|
raise AssertionError("Responses streams poll disconnects in the generator")
|
||||||
|
|
||||||
|
async def send(message):
|
||||||
|
sent.append(message)
|
||||||
|
|
||||||
|
await response({"type": "http", "asgi": {"spec_version": "2.3"}}, receive, send)
|
||||||
|
return sent
|
||||||
|
|
||||||
|
sent = asyncio.run(run())
|
||||||
|
|
||||||
|
assert sent[0]["type"] == "http.response.start"
|
||||||
|
body = b"".join(message.get("body", b"") for message in sent).decode()
|
||||||
|
assert "response.output_text.delta" in body
|
||||||
|
assert '"delta":"33"' in body.replace(" ", "")
|
||||||
|
|
||||||
def test_split_think_markers_stream_as_reasoning_and_visible_text(self, monkeypatch):
|
def test_split_think_markers_stream_as_reasoning_and_visible_text(self, monkeypatch):
|
||||||
chunks = [
|
chunks = [
|
||||||
{"choices": [{"delta": {"content": "<thi"}}]},
|
{"choices": [{"delta": {"content": "<thi"}}]},
|
||||||
|
|
@ -1286,7 +1358,7 @@ class TestResponsesStreamAdapter:
|
||||||
assert entry["status"] == "completed"
|
assert entry["status"] == "completed"
|
||||||
assert entry["reply"] == "tail"
|
assert entry["reply"] == "tail"
|
||||||
|
|
||||||
def test_reasoning_only_fallback_updates_monitor(self, monkeypatch):
|
def test_reasoning_only_stream_does_not_update_visible_monitor_reply(self, monkeypatch):
|
||||||
import routes.inference as inf_mod
|
import routes.inference as inf_mod
|
||||||
|
|
||||||
class FakeExtractor:
|
class FakeExtractor:
|
||||||
|
|
@ -1327,15 +1399,16 @@ class TestResponsesStreamAdapter:
|
||||||
|
|
||||||
lines = asyncio.run(run())
|
lines = asyncio.run(run())
|
||||||
|
|
||||||
assert self._payloads(lines, "response.output_text.delta")[-1]["delta"] == "plan"
|
assert self._payloads(lines, "response.output_text.delta") == []
|
||||||
|
assert self._payloads(lines, "response.reasoning_text.delta")[-1]["delta"] == "plan"
|
||||||
[entry] = monitor.snapshot()
|
[entry] = monitor.snapshot()
|
||||||
assert entry["status"] == "completed"
|
assert entry["status"] == "completed"
|
||||||
assert entry["reply"] == "plan"
|
assert entry["reply"] == ""
|
||||||
|
|
||||||
def test_literal_think_tags_stream_as_visible_text_without_reasoning_request(self, monkeypatch):
|
def test_reasoning_capable_gguf_stream_parses_think_tags_by_default(self, monkeypatch):
|
||||||
chunks = [
|
chunks = [
|
||||||
{"choices": [{"delta": {"content": "show <thi"}}]},
|
{"choices": [{"delta": {"content": "<thi"}}]},
|
||||||
{"choices": [{"delta": {"content": "nk>x</think> tags"}}]},
|
{"choices": [{"delta": {"content": "nk>plan</think>answer"}}]},
|
||||||
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
||||||
]
|
]
|
||||||
self._install_stream_mock(monkeypatch, chunks)
|
self._install_stream_mock(monkeypatch, chunks)
|
||||||
|
|
@ -1350,13 +1423,15 @@ class TestResponsesStreamAdapter:
|
||||||
|
|
||||||
reasoning_deltas = self._payloads(lines, "response.reasoning_text.delta")
|
reasoning_deltas = self._payloads(lines, "response.reasoning_text.delta")
|
||||||
text_deltas = self._payloads(lines, "response.output_text.delta")
|
text_deltas = self._payloads(lines, "response.output_text.delta")
|
||||||
assert reasoning_deltas == []
|
assert "".join(event["delta"] for event in reasoning_deltas) == "plan"
|
||||||
assert "".join(event["delta"] for event in text_deltas) == "show <think>x</think> tags"
|
assert "".join(event["delta"] for event in text_deltas) == "answer"
|
||||||
completed = self._payloads(lines, "response.completed")[0]
|
completed = self._payloads(lines, "response.completed")[0]
|
||||||
assert [item["type"] for item in completed["response"]["output"]] == ["message"]
|
assert [item["type"] for item in completed["response"]["output"]] == [
|
||||||
assert completed["response"]["output"][0]["content"][0]["text"] == (
|
"reasoning",
|
||||||
"show <think>x</think> tags"
|
"message",
|
||||||
)
|
]
|
||||||
|
assert completed["response"]["output"][0]["content"][0]["text"] == "plan"
|
||||||
|
assert completed["response"]["output"][1]["content"][0]["text"] == "answer"
|
||||||
|
|
||||||
def test_non_reasoning_gguf_stream_keeps_literal_think_tags_visible(self, monkeypatch):
|
def test_non_reasoning_gguf_stream_keeps_literal_think_tags_visible(self, monkeypatch):
|
||||||
chunks = [
|
chunks = [
|
||||||
|
|
@ -1384,7 +1459,7 @@ class TestResponsesStreamAdapter:
|
||||||
"show <think>x</think> tags"
|
"show <think>x</think> tags"
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_reasoning_only_streams_as_visible_message_text(self, monkeypatch):
|
def test_reasoning_only_stream_stays_out_of_visible_message_text(self, monkeypatch):
|
||||||
chunks = [
|
chunks = [
|
||||||
{"choices": [{"delta": {"content": "<think>plan</think>"}}]},
|
{"choices": [{"delta": {"content": "<think>plan</think>"}}]},
|
||||||
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
||||||
|
|
@ -1402,14 +1477,34 @@ class TestResponsesStreamAdapter:
|
||||||
reasoning_deltas = self._payloads(lines, "response.reasoning_text.delta")
|
reasoning_deltas = self._payloads(lines, "response.reasoning_text.delta")
|
||||||
text_deltas = self._payloads(lines, "response.output_text.delta")
|
text_deltas = self._payloads(lines, "response.output_text.delta")
|
||||||
assert "".join(event["delta"] for event in reasoning_deltas) == "plan"
|
assert "".join(event["delta"] for event in reasoning_deltas) == "plan"
|
||||||
assert "".join(event["delta"] for event in text_deltas) == "plan"
|
assert text_deltas == []
|
||||||
completed = self._payloads(lines, "response.completed")[0]
|
completed = self._payloads(lines, "response.completed")[0]
|
||||||
assert [item["type"] for item in completed["response"]["output"]] == [
|
assert [item["type"] for item in completed["response"]["output"]] == ["reasoning"]
|
||||||
"reasoning",
|
assert completed["response"]["output"][0]["content"][0]["text"] == "plan"
|
||||||
"message",
|
|
||||||
]
|
def test_unclosed_think_stream_stays_out_of_visible_message_text(self, monkeypatch):
|
||||||
|
chunks = [
|
||||||
|
{"choices": [{"delta": {"content": "<thi"}}]},
|
||||||
|
{"choices": [{"delta": {"content": "nk>plan"}}]},
|
||||||
|
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
||||||
|
]
|
||||||
|
self._install_stream_mock(monkeypatch, chunks)
|
||||||
|
payload = ResponsesRequest(input = "hi", stream = True, reasoning = {"effort": "high"})
|
||||||
|
messages = [ChatMessage(role = "user", content = "hi")]
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
response = await _responses_stream(payload, messages, self._Request())
|
||||||
|
return await self._collect(response)
|
||||||
|
|
||||||
|
lines = asyncio.run(run())
|
||||||
|
|
||||||
|
reasoning_deltas = self._payloads(lines, "response.reasoning_text.delta")
|
||||||
|
text_deltas = self._payloads(lines, "response.output_text.delta")
|
||||||
|
assert "".join(event["delta"] for event in reasoning_deltas) == "plan"
|
||||||
|
assert text_deltas == []
|
||||||
|
completed = self._payloads(lines, "response.completed")[0]
|
||||||
|
assert [item["type"] for item in completed["response"]["output"]] == ["reasoning"]
|
||||||
assert completed["response"]["output"][0]["content"][0]["text"] == "plan"
|
assert completed["response"]["output"][0]["content"][0]["text"] == "plan"
|
||||||
assert completed["response"]["output"][1]["content"][0]["text"] == "plan"
|
|
||||||
|
|
||||||
def test_structured_reasoning_content_streams_as_reasoning(self, monkeypatch):
|
def test_structured_reasoning_content_streams_as_reasoning(self, monkeypatch):
|
||||||
chunks = [
|
chunks = [
|
||||||
|
|
|
||||||
|
|
@ -203,6 +203,20 @@ def test_detect_safetensors_features_function_xml_format_keeps_tools_on():
|
||||||
assert flags["supports_tools"] is True
|
assert flags["supports_tools"] is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_safetensors_features_gemma_native_tool_call_keeps_tools_on():
|
||||||
|
"""Gemma 4 emits <|tool_call>call:name{...}<tool_call|>, which the shared
|
||||||
|
parser now reads, so the gate must not suppress tools for it."""
|
||||||
|
from routes.inference import _detect_safetensors_features
|
||||||
|
|
||||||
|
tpl_with_gemma_native = (
|
||||||
|
"{%- if tools -%}Tool call format: "
|
||||||
|
"<|tool_call>call:name{key:value}<tool_call|>{%- endif -%}"
|
||||||
|
)
|
||||||
|
backend = SimpleNamespace(active_model_name = "unsloth/gemma-4-12b-it")
|
||||||
|
flags = _detect_safetensors_features(backend, tpl_with_gemma_native)
|
||||||
|
assert flags["supports_tools"] is True
|
||||||
|
|
||||||
|
|
||||||
# Qwen3.5 family pin: the live GGUF + safetensors templates both wrap tool
|
# Qwen3.5 family pin: the live GGUF + safetensors templates both wrap tool
|
||||||
# calls as ``<tool_call>\n<function=name>...``. Faithful slice so the
|
# calls as ``<tool_call>\n<function=name>...``. Faithful slice so the
|
||||||
# classifier never silently regresses for this family.
|
# classifier never silently regresses for this family.
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ calls, tool-result feedback, bad-JSON heal, duplicate-call short-circuit,
|
||||||
``__IMAGES__`` sentinel stripping, executor errors, cancel, and the iteration cap.
|
``__IMAGES__`` sentinel stripping, executor errors, cancel, and the iteration cap.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
import threading
|
import threading
|
||||||
from typing import cast
|
from typing import cast
|
||||||
|
|
||||||
|
|
@ -62,6 +63,51 @@ class TestParser:
|
||||||
assert parse_tool_calls_from_text(text)[0]["function"]["name"] == "python"
|
assert parse_tool_calls_from_text(text)[0]["function"]["name"] == "python"
|
||||||
assert parse_tool_calls_from_text(text, allow_incomplete = False) == []
|
assert parse_tool_calls_from_text(text, allow_incomplete = False) == []
|
||||||
|
|
||||||
|
def test_gemma_native_tool_call(self):
|
||||||
|
text = '<|tool_call>call:terminal{command:"ls -la",workdir:"."}<tool_call|>'
|
||||||
|
result = parse_tool_calls_from_text(text)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["function"]["name"] == "terminal"
|
||||||
|
args = json.loads(result[0]["function"]["arguments"])
|
||||||
|
assert args == {"command": "ls -la", "workdir": "."}
|
||||||
|
|
||||||
|
def test_gemma_native_tool_call_template_quotes(self):
|
||||||
|
text = '<|tool_call>call:web_search{query:<|"|>openai news<|"|>}<tool_call|>'
|
||||||
|
result = parse_tool_calls_from_text(text)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["function"]["name"] == "web_search"
|
||||||
|
assert json.loads(result[0]["function"]["arguments"]) == {"query": "openai news"}
|
||||||
|
|
||||||
|
def test_gemma_native_tool_call_template_quotes_escape_backslashes(self):
|
||||||
|
text = r'<|tool_call>call:ls{path:<|"|>C:\Users\wasim\repo<|"|>}<tool_call|>'
|
||||||
|
result = parse_tool_calls_from_text(text)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["function"]["name"] == "ls"
|
||||||
|
assert json.loads(result[0]["function"]["arguments"]) == {"path": r"C:\Users\wasim\repo"}
|
||||||
|
|
||||||
|
def test_gemma_native_tool_call_hyphenated_argument_name(self):
|
||||||
|
text = '<|tool_call>call:mcp__srv__create-issue{issue-title:"Bug report"}<tool_call|>'
|
||||||
|
result = parse_tool_calls_from_text(text)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["function"]["name"] == "mcp__srv__create-issue"
|
||||||
|
assert json.loads(result[0]["function"]["arguments"]) == {"issue-title": "Bug report"}
|
||||||
|
|
||||||
|
def test_gemma_native_tool_call_keeps_braces_inside_string_value(self):
|
||||||
|
text = '<|tool_call>call:terminal{command:"echo {foo:bar}"}<tool_call|>'
|
||||||
|
result = parse_tool_calls_from_text(text)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["function"]["name"] == "terminal"
|
||||||
|
assert json.loads(result[0]["function"]["arguments"]) == {"command": "echo {foo:bar}"}
|
||||||
|
|
||||||
|
def test_gemma_native_tool_call_bare_string_values(self):
|
||||||
|
text = "<|tool_call>call:get_weather{location:Tokyo,unit:celsius}<tool_call|>"
|
||||||
|
result = parse_tool_calls_from_text(text)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert json.loads(result[0]["function"]["arguments"]) == {
|
||||||
|
"location": "Tokyo",
|
||||||
|
"unit": "celsius",
|
||||||
|
}
|
||||||
|
|
||||||
def test_xml_function_call(self):
|
def test_xml_function_call(self):
|
||||||
text = "<function=python><parameter=code>print('hi')</parameter></function>"
|
text = "<function=python><parameter=code>print('hi')</parameter></function>"
|
||||||
result = parse_tool_calls_from_text(text)
|
result = parse_tool_calls_from_text(text)
|
||||||
|
|
@ -121,6 +167,7 @@ class TestParser:
|
||||||
|
|
||||||
def test_has_tool_signal(self):
|
def test_has_tool_signal(self):
|
||||||
assert has_tool_signal("blah <tool_call> x")
|
assert has_tool_signal("blah <tool_call> x")
|
||||||
|
assert has_tool_signal("blah <|tool_call>call:terminal")
|
||||||
assert has_tool_signal("hi <function=foo>...")
|
assert has_tool_signal("hi <function=foo>...")
|
||||||
assert not has_tool_signal("hello world")
|
assert not has_tool_signal("hello world")
|
||||||
|
|
||||||
|
|
@ -139,6 +186,8 @@ class TestParser:
|
||||||
def test_strip_markup_closed(self):
|
def test_strip_markup_closed(self):
|
||||||
text = "before <tool_call>{}</tool_call> after"
|
text = "before <tool_call>{}</tool_call> after"
|
||||||
assert strip_tool_markup(text) == "before after"
|
assert strip_tool_markup(text) == "before after"
|
||||||
|
text = 'before <|tool_call>call:terminal{command:"ls"}<tool_call|> after'
|
||||||
|
assert strip_tool_markup(text) == "before after"
|
||||||
|
|
||||||
def test_strip_markup_unclosed_final(self):
|
def test_strip_markup_unclosed_final(self):
|
||||||
text = "before <tool_call>{partial"
|
text = "before <tool_call>{partial"
|
||||||
|
|
@ -146,6 +195,7 @@ class TestParser:
|
||||||
assert strip_tool_markup(text, final = True) == "before"
|
assert strip_tool_markup(text, final = True) == "before"
|
||||||
# Without final=True the unclosed run is preserved.
|
# Without final=True the unclosed run is preserved.
|
||||||
assert "partial" in strip_tool_markup(text)
|
assert "partial" in strip_tool_markup(text)
|
||||||
|
assert strip_tool_markup("before <|tool_call>call:terminal{", final = True) == "before"
|
||||||
|
|
||||||
def test_streaming_strip_respects_disabled_healing(self):
|
def test_streaming_strip_respects_disabled_healing(self):
|
||||||
raw = 'before <tool_call>{"name":"web_search"'
|
raw = 'before <tool_call>{"name":"web_search"'
|
||||||
|
|
|
||||||
|
|
@ -106,6 +106,54 @@ class TestParityWithJsonStyle:
|
||||||
assert json.loads(js[0]["function"]["arguments"]) == {"query": q}
|
assert json.loads(js[0]["function"]["arguments"]) == {"query": q}
|
||||||
|
|
||||||
|
|
||||||
|
class TestGemmaNativeStyle:
|
||||||
|
def test_closed_native_call_with_trailing_prose_is_accepted(self):
|
||||||
|
text = (
|
||||||
|
'<|tool_call>call:terminal{command:"ls -la",workdir:"."}<tool_call|>' " running it now"
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(text, allow_incomplete = False)
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0]["function"]["name"] == "terminal"
|
||||||
|
assert json.loads(calls[0]["function"]["arguments"]) == {
|
||||||
|
"command": "ls -la",
|
||||||
|
"workdir": ".",
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_unclosed_native_call_requires_healing(self):
|
||||||
|
text = '<|tool_call>call:terminal{command:"ls"}'
|
||||||
|
assert parse_tool_calls_from_text(text, allow_incomplete = False) == []
|
||||||
|
calls = parse_tool_calls_from_text(text, allow_incomplete = True)
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0]["function"]["name"] == "terminal"
|
||||||
|
|
||||||
|
def test_hyphenated_native_argument_name_is_accepted(self):
|
||||||
|
text = '<|tool_call>call:mcp__srv__create-issue{issue-title:"Bug report"}<tool_call|>'
|
||||||
|
calls = parse_tool_calls_from_text(text, allow_incomplete = False)
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert calls[0]["function"]["name"] == "mcp__srv__create-issue"
|
||||||
|
assert json.loads(calls[0]["function"]["arguments"]) == {"issue-title": "Bug report"}
|
||||||
|
|
||||||
|
def test_native_template_quotes_preserve_windows_path(self):
|
||||||
|
text = r'<|tool_call>call:ls{path:<|"|>C:\Users\wasim\repo<|"|>}<tool_call|>'
|
||||||
|
calls = parse_tool_calls_from_text(text, allow_incomplete = False)
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert json.loads(calls[0]["function"]["arguments"]) == {"path": r"C:\Users\wasim\repo"}
|
||||||
|
|
||||||
|
def test_bare_unquoted_string_values_are_accepted(self):
|
||||||
|
# Gemma can emit enum/string args unquoted; bare JSON scalars stay typed.
|
||||||
|
text = (
|
||||||
|
"<|tool_call>call:get_weather{location:Tokyo,unit:celsius,days:3,live:true}<tool_call|>"
|
||||||
|
)
|
||||||
|
calls = parse_tool_calls_from_text(text, allow_incomplete = False)
|
||||||
|
assert len(calls) == 1
|
||||||
|
assert json.loads(calls[0]["function"]["arguments"]) == {
|
||||||
|
"location": "Tokyo",
|
||||||
|
"unit": "celsius",
|
||||||
|
"days": 3,
|
||||||
|
"live": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class TestHealingPathUnaffected:
|
class TestHealingPathUnaffected:
|
||||||
def test_auto_heal_still_repairs_unclosed_function(self):
|
def test_auto_heal_still_repairs_unclosed_function(self):
|
||||||
text = "<function=web_search><parameter=query>cats"
|
text = "<function=web_search><parameter=query>cats"
|
||||||
|
|
|
||||||
|
|
@ -125,6 +125,14 @@ def test_strips_orphan_closing_tag():
|
||||||
# Mid-string </parameter> intentionally preserved (see preserve test).
|
# Mid-string </parameter> intentionally preserved (see preserve test).
|
||||||
|
|
||||||
|
|
||||||
|
def test_strips_gemma_native_orphan_closing_tag():
|
||||||
|
cleaned = _TOOL_XML_RE.sub("", "Tool call drained.<tool_call|>Visible tail.")
|
||||||
|
|
||||||
|
assert "<tool_call|>" not in cleaned
|
||||||
|
assert "Tool call drained." in cleaned
|
||||||
|
assert "Visible tail." in cleaned
|
||||||
|
|
||||||
|
|
||||||
# ── Tail-only </parameter> (PR #5735 follow-up) ───────────────────
|
# ── Tail-only </parameter> (PR #5735 follow-up) ───────────────────
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -146,24 +146,50 @@ def test_async_generators_cleanup_tracker_in_finally():
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_streaming_responses_have_no_background_task():
|
def test_chat_completions_streams_avoid_starlette_task_group():
|
||||||
top = None
|
top = _async_function("openai_chat_completions")
|
||||||
for n in ast.walk(_TREE):
|
legacy_calls = []
|
||||||
if isinstance(n, ast.AsyncFunctionDef) and n.name == "openai_chat_completions":
|
same_task_calls = 0
|
||||||
top = n
|
|
||||||
break
|
|
||||||
assert top is not None
|
|
||||||
for sub in ast.walk(top):
|
for sub in ast.walk(top):
|
||||||
if not (isinstance(sub, ast.Call) and isinstance(sub.func, ast.Name)):
|
if not (isinstance(sub, ast.Call) and isinstance(sub.func, ast.Name)):
|
||||||
continue
|
continue
|
||||||
if sub.func.id != "StreamingResponse":
|
if sub.func.id == "StreamingResponse":
|
||||||
|
legacy_calls.append(sub.lineno)
|
||||||
|
if sub.func.id == "_SameTaskStreamingResponse":
|
||||||
|
same_task_calls += 1
|
||||||
|
assert not legacy_calls, (
|
||||||
|
"Streaming /v1/chat/completions must use _SameTaskStreamingResponse, "
|
||||||
|
"not Starlette's legacy task-group StreamingResponse. Lines: "
|
||||||
|
f"{legacy_calls}"
|
||||||
|
)
|
||||||
|
assert same_task_calls >= 5
|
||||||
|
|
||||||
|
|
||||||
|
def test_openai_passthrough_stream_avoids_starlette_task_group():
|
||||||
|
top = _async_function("_openai_passthrough_stream")
|
||||||
|
legacy_calls = []
|
||||||
|
same_task_calls = 0
|
||||||
|
for sub in ast.walk(top):
|
||||||
|
if not (isinstance(sub, ast.Call) and isinstance(sub.func, ast.Name)):
|
||||||
continue
|
continue
|
||||||
kwargs = {kw.arg for kw in sub.keywords if kw.arg}
|
if sub.func.id == "StreamingResponse":
|
||||||
assert "background" not in kwargs, (
|
legacy_calls.append(sub.lineno)
|
||||||
"StreamingResponse in openai_chat_completions must not pass "
|
if sub.func.id == "_SameTaskStreamingResponse":
|
||||||
"`background=` -- cleanup now lives in the generator's finally "
|
same_task_calls += 1
|
||||||
"block; a BackgroundTask would be skipped on abrupt disconnect"
|
assert not legacy_calls, (
|
||||||
)
|
"OpenAI passthrough streams must use _SameTaskStreamingResponse, "
|
||||||
|
"not Starlette's legacy task-group StreamingResponse. Lines: "
|
||||||
|
f"{legacy_calls}"
|
||||||
|
)
|
||||||
|
assert same_task_calls >= 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_chat_streams_install_same_task_disconnect_watcher():
|
||||||
|
top = _async_function("openai_chat_completions")
|
||||||
|
assert _calls_name(top, "_await_disconnect_then_cancel"), (
|
||||||
|
"Local same-task streams must watch request disconnects themselves; "
|
||||||
|
"do not restore Starlette's task-group StreamingResponse for this."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_direct_llama_server_streams_install_disconnect_watcher():
|
def test_direct_llama_server_streams_install_disconnect_watcher():
|
||||||
|
|
@ -185,6 +211,33 @@ def test_direct_llama_server_streams_install_disconnect_watcher():
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_audio_input_stream_installs_disconnect_watcher():
|
||||||
|
audio = _async_function("audio_input_stream")
|
||||||
|
has_watcher = False
|
||||||
|
has_cleanup = False
|
||||||
|
for sub in ast.walk(audio):
|
||||||
|
if isinstance(sub, ast.Call):
|
||||||
|
fn = sub.func
|
||||||
|
if (
|
||||||
|
isinstance(fn, ast.Attribute)
|
||||||
|
and fn.attr == "create_task"
|
||||||
|
and isinstance(fn.value, ast.Name)
|
||||||
|
and fn.value.id == "asyncio"
|
||||||
|
and sub.args
|
||||||
|
and isinstance(sub.args[0], ast.Call)
|
||||||
|
and isinstance(sub.args[0].func, ast.Name)
|
||||||
|
and sub.args[0].func.id == "_await_disconnect_then_cancel"
|
||||||
|
):
|
||||||
|
has_watcher = True
|
||||||
|
if isinstance(fn, ast.Name) and fn.id == "_stop_local_disconnect_cancel_watcher":
|
||||||
|
has_cleanup = True
|
||||||
|
assert has_watcher, (
|
||||||
|
"audio_input_stream must install a disconnect watcher so client "
|
||||||
|
"disconnects set cancel_event while asyncio.to_thread(next, ...) is blocked"
|
||||||
|
)
|
||||||
|
assert has_cleanup, "audio_input_stream must stop its disconnect watcher in finally"
|
||||||
|
|
||||||
|
|
||||||
# ── Behavioral helpers ───────────────────────────────────────
|
# ── Behavioral helpers ───────────────────────────────────────
|
||||||
|
|
||||||
_WANTED = {
|
_WANTED = {
|
||||||
|
|
@ -222,6 +275,21 @@ def _load_registry_module():
|
||||||
return mod
|
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\nclass ClientDisconnect(Exception): pass\n" + source,
|
||||||
|
mod,
|
||||||
|
)
|
||||||
|
return mod
|
||||||
|
|
||||||
|
|
||||||
def _make_stream(tracker, raise_exc):
|
def _make_stream(tracker, raise_exc):
|
||||||
async def gen():
|
async def gen():
|
||||||
try:
|
try:
|
||||||
|
|
@ -326,6 +394,40 @@ def test_finally_cleanup_on_aclose():
|
||||||
assert "sid-abort" not in m["_CANCEL_REGISTRY"]
|
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
|
||||||
|
response._unstarted_cleanup = 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():
|
def test_preset_cancel_event_exits_cleanly_with_done():
|
||||||
# Pending-replay: a stashed cancel pre-set cancel_event. The loop must break
|
# 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.
|
# cleanly with final_chunk + [DONE], not propagate GeneratorExit from the GGUF wrapper.
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue