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.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from core import tool_healing as _tool_healing
|
||||
|
||||
|
||||
# _TOOL_CLOSED_PATS: closed pairs only. _TOOL_ALL_PATS: also trailing unclosed
|
||||
# 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 = [
|
||||
re.compile(r"<tool_call>.*?</tool_call>", re.DOTALL),
|
||||
re.compile(r"<function=[\w-]+>.*?</function>", re.DOTALL),
|
||||
]
|
||||
_TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
||||
re.compile(r"<tool_call>.*$", re.DOTALL),
|
||||
re.compile(r"<function=[\w-]+>.*$", re.DOTALL),
|
||||
]
|
||||
_TOOL_ALL_PATS = _tool_healing._TOOL_ALL_PATS
|
||||
|
||||
|
||||
def parse_tool_calls_from_text(
|
||||
content: str,
|
||||
*,
|
||||
id_offset: int = 0,
|
||||
allow_incomplete: bool = True,
|
||||
) -> list[dict]:
|
||||
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.
|
||||
TOOL_XML_SIGNALS = ("<tool_call>", "<function=")
|
||||
TOOL_XML_SIGNALS = ("<tool_call>", "<|tool_call>", "<function=")
|
||||
|
||||
|
||||
# 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:
|
||||
"""Return True if ``text`` contains any tool-call XML signal."""
|
||||
return any(s in text for s in TOOL_XML_SIGNALS)
|
||||
|
|
|
|||
|
|
@ -1,14 +1,10 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# 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
|
||||
inference servers can reuse the logic without importing the inference
|
||||
External inference servers import this module without pulling in the inference
|
||||
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
|
||||
|
|
@ -17,8 +13,21 @@ import re
|
|||
# Pre-compiled patterns for tool XML stripping. The hyphen in the name
|
||||
# char-class lets dashed MCP tool/parameter names (mcp__srv__list-issues,
|
||||
# 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 = [
|
||||
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),
|
||||
]
|
||||
_TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
||||
|
|
@ -28,77 +37,358 @@ _TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
|
|||
|
||||
# Pre-compiled patterns for tool-call XML parsing.
|
||||
_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_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_PARAM_START_RE = re.compile(r"<parameter=([\w-]+)>\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]:
|
||||
"""
|
||||
Parse tool calls from XML markup in content text.
|
||||
def _balanced_brace_end(
|
||||
content: str,
|
||||
brace_start: int,
|
||||
*,
|
||||
gemma_quotes: bool = False,
|
||||
) -> int:
|
||||
depth = 0
|
||||
i = brace_start
|
||||
in_string = False
|
||||
in_gemma_string = False
|
||||
while i < len(content):
|
||||
if 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:
|
||||
<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>
|
||||
Closing tags (</tool_call>, </function>, </parameter>) are all
|
||||
optional since models frequently omit them.
|
||||
"""
|
||||
tool_calls = []
|
||||
|
||||
# Pattern 1: JSON inside <tool_call> tags. Balanced-brace extraction that
|
||||
# skips braces inside JSON strings.
|
||||
tool_calls: list[dict] = []
|
||||
# Collect JSON- and Gemma-format candidates with their byte spans, then
|
||||
# accept them in document order. Both order and spans matter:
|
||||
# * 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):
|
||||
brace_start = m.end() - 1 # position of the 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 # skip escaped character
|
||||
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:
|
||||
json_str = content[brace_start : i + 1]
|
||||
try:
|
||||
obj = json.loads(json_str)
|
||||
tc = {
|
||||
"id": f"call_{len(tool_calls)}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": obj.get("name", ""),
|
||||
"arguments": obj.get("arguments", {}),
|
||||
},
|
||||
}
|
||||
if isinstance(tc["function"]["arguments"], dict):
|
||||
tc["function"]["arguments"] = json.dumps(tc["function"]["arguments"])
|
||||
tool_calls.append(tc)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
# A marker that begins inside an open <function=...><parameter=...> value
|
||||
# is that parameter's data, not its own call; skip it (same guard the
|
||||
# XML-style parser below applies to nested <function= markers).
|
||||
if _inside_open_parameter(content, m.start()):
|
||||
continue
|
||||
end = _balanced_brace_end(content, m.end() - 1)
|
||||
if end >= 0:
|
||||
candidates.append((m.start(), end, "json", m))
|
||||
for m in _TC_GEMMA_START_RE.finditer(content):
|
||||
if _inside_open_parameter(content, m.start()):
|
||||
continue
|
||||
end = _balanced_brace_end(content, m.end() - 1, gemma_quotes = True)
|
||||
if end >= 0:
|
||||
candidates.append((m.start(), end, "gemma", m))
|
||||
candidates.sort(key = lambda c: c[0])
|
||||
|
||||
spans = [(s, e) for s, e, _kind, _m in candidates]
|
||||
for idx, (start, end, kind, m) in enumerate(candidates):
|
||||
# Skip a candidate nested inside another candidate's brace span: it is
|
||||
# the enclosing call's argument data, not its own call. Checked against
|
||||
# every candidate span (not only the ones that parsed successfully), so a
|
||||
# marker inside an outer call that later fails to normalize is still
|
||||
# never promoted to its own executable tool call.
|
||||
if any(s <= start and end <= e for j, (s, e) in enumerate(spans) if j != idx):
|
||||
continue
|
||||
if not allow_incomplete:
|
||||
tail = content[end + 1 :].lstrip()
|
||||
close_re = _TC_END_TAG_RE if kind == "json" else _TC_GEMMA_END_TAG_RE
|
||||
if close_re.match(tail) is None:
|
||||
continue
|
||||
try:
|
||||
if kind == "json":
|
||||
obj = json.loads(content[m.end() - 1 : end + 1])
|
||||
name = obj.get("name", "")
|
||||
arguments = obj.get("arguments", {})
|
||||
if isinstance(arguments, dict):
|
||||
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:
|
||||
# Step 1: Find <function=name> positions and extract bodies. Use only
|
||||
# </tool_call> or the next <function= as hard boundaries (</function>
|
||||
# can appear in code values); trim a trailing </function> afterwards.
|
||||
func_starts = list(_TC_FUNC_START_RE.finditer(content))
|
||||
func_starts = [
|
||||
fm
|
||||
for fm in _TC_FUNC_START_RE.finditer(content)
|
||||
if not _inside_open_parameter(content, fm.start())
|
||||
]
|
||||
for idx, fm in enumerate(func_starts):
|
||||
func_name = fm.group(1)
|
||||
body_start = fm.end()
|
||||
# Boundaries: next <function= tag or </tool_call>
|
||||
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:
|
||||
|
|
@ -107,36 +397,52 @@ def parse_tool_calls_from_text(content: str) -> list[dict]:
|
|||
body_end = len(content)
|
||||
body_end = min(body_end, next_func)
|
||||
body = content[body_start:body_end]
|
||||
body = _TC_FUNC_CLOSE_RE.sub("", body) # trim closing </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
|
||||
# functions, use body end as the only boundary to avoid matching
|
||||
# </parameter> inside code strings.
|
||||
arguments = {}
|
||||
arguments: dict = {}
|
||||
param_starts = list(_TC_PARAM_START_RE.finditer(body))
|
||||
if len(param_starts) == 1:
|
||||
# Value is everything after the tag to end of body, less a
|
||||
# trailing </parameter>.
|
||||
pm = param_starts[0]
|
||||
val = body[pm.end() :]
|
||||
val = _TC_PARAM_CLOSE_RE.sub("", val)
|
||||
if not allow_incomplete:
|
||||
stripped_val = val.rstrip()
|
||||
if not stripped_val.endswith(_PARAM_CLOSE_TAG):
|
||||
continue
|
||||
val = stripped_val[: -len(_PARAM_CLOSE_TAG)]
|
||||
else:
|
||||
val = _TC_PARAM_CLOSE_RE.sub("", val)
|
||||
arguments[pm.group(1)] = val.strip()
|
||||
else:
|
||||
valid_params = True
|
||||
for pidx, pm in enumerate(param_starts):
|
||||
param_name = pm.group(1)
|
||||
val_start = pm.end()
|
||||
# Value ends at next <parameter= or end of body
|
||||
next_param = (
|
||||
param_starts[pidx + 1].start()
|
||||
if pidx + 1 < len(param_starts)
|
||||
else len(body)
|
||||
)
|
||||
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()
|
||||
if not valid_params:
|
||||
continue
|
||||
|
||||
tc = {
|
||||
"id": f"call_{len(tool_calls)}",
|
||||
"id": f"call_{id_offset + len(tool_calls)}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": func_name,
|
||||
|
|
@ -147,6 +453,41 @@ def parse_tool_calls_from_text(content: str) -> list[dict]:
|
|||
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:
|
||||
"""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
|
||||
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
|
||||
for pat in patterns:
|
||||
if pat is _TC_GEMMA_CLOSED_PAT:
|
||||
continue
|
||||
text = pat.sub("", text)
|
||||
return text.strip() if final else text
|
||||
|
|
|
|||
|
|
@ -1102,6 +1102,8 @@ class ChoiceDelta(BaseModel):
|
|||
|
||||
role: 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"]
|
||||
|
|
@ -1137,6 +1139,8 @@ class CompletionMessage(BaseModel):
|
|||
role: Literal["assistant"] = "assistant"
|
||||
content: str
|
||||
refusal: Optional[str] = None
|
||||
reasoning_content: Optional[str] = None
|
||||
tool_calls: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class CompletionChoice(BaseModel):
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import uuid
|
|||
from pathlib import Path
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import StreamingResponse, JSONResponse, Response
|
||||
from starlette.requests import ClientDisconnect
|
||||
from typing import Any, List, Optional, Union
|
||||
import json
|
||||
import httpx
|
||||
|
|
@ -235,8 +236,15 @@ def _sse_streaming_response(content) -> StreamingResponse:
|
|||
a one-shot connection. Two callers build their response inline instead: the
|
||||
external-provider proxy omits ``Connection: close``, and the OpenAI
|
||||
passthrough returns an empty ``keep-alive`` stream when the request is
|
||||
cancelled before the upstream response starts."""
|
||||
return StreamingResponse(
|
||||
cancelled before the upstream response starts.
|
||||
|
||||
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,
|
||||
media_type = "text/event-stream",
|
||||
headers = {
|
||||
|
|
@ -750,6 +758,139 @@ def _set_stream_response_read_timeout(
|
|||
pass
|
||||
|
||||
|
||||
_STREAM_DISCONNECT_POLL_TIMEOUT_S = 0.25
|
||||
|
||||
|
||||
class _CompatSameTaskTimeout:
|
||||
"""Same-task timeout fallback for Python versions before asyncio.timeout."""
|
||||
|
||||
def __init__(self, timeout_s: float):
|
||||
self.timeout_s = timeout_s
|
||||
self._task = None
|
||||
self._handle = None
|
||||
self._timed_out = False
|
||||
self._cancelling = 0
|
||||
|
||||
async def __aenter__(self):
|
||||
self._task = asyncio.current_task()
|
||||
if self._task is None:
|
||||
return self
|
||||
if hasattr(self._task, "cancelling"):
|
||||
self._cancelling = self._task.cancelling()
|
||||
loop = asyncio.get_running_loop()
|
||||
self._handle = loop.call_later(max(self.timeout_s, 0), self._cancel_task)
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
if self._handle is not None:
|
||||
self._handle.cancel()
|
||||
if exc_type is not None and issubclass(exc_type, asyncio.CancelledError):
|
||||
if self._timed_out:
|
||||
if self._task is not None and hasattr(self._task, "uncancel"):
|
||||
if self._task.uncancel() > self._cancelling:
|
||||
return None
|
||||
raise asyncio.TimeoutError from exc
|
||||
return None
|
||||
|
||||
def _cancel_task(self) -> None:
|
||||
self._timed_out = True
|
||||
if self._task is not None:
|
||||
self._task.cancel()
|
||||
|
||||
|
||||
def _same_task_timeout(timeout_s: float):
|
||||
timeout_ctx = getattr(asyncio, "timeout", None)
|
||||
if timeout_ctx is not None:
|
||||
return timeout_ctx(timeout_s)
|
||||
return _CompatSameTaskTimeout(timeout_s)
|
||||
|
||||
|
||||
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(
|
||||
*,
|
||||
watchers = (),
|
||||
|
|
@ -875,8 +1016,23 @@ async def _aiter_llama_stream_items(
|
|||
raise httpx.ReadTimeout("The model did not produce a first token in time.")
|
||||
if response is not None:
|
||||
_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:
|
||||
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__()
|
||||
except asyncio.TimeoutError as exc:
|
||||
if waiting_first_item:
|
||||
|
|
@ -890,6 +1046,12 @@ async def _aiter_llama_stream_items(
|
|||
if now >= first_token_deadline:
|
||||
raise
|
||||
continue
|
||||
if (
|
||||
request is not None
|
||||
and post_first_item_read_timeout_s is not None
|
||||
and now - last_item_at < post_first_item_read_timeout_s
|
||||
):
|
||||
continue
|
||||
raise httpx.ReadTimeout("The model stopped producing tokens mid-response.")
|
||||
if (
|
||||
last_item_at is None
|
||||
|
|
@ -1125,16 +1287,17 @@ def _detect_safetensors_features(backend, chat_template: Optional[str]) -> dict:
|
|||
model_identifier = model_id,
|
||||
log_source = "safetensors",
|
||||
)
|
||||
# Our safetensors loop only parses <tool_call>{json}</tool_call> and
|
||||
# <function=name>...</function>. Llama uses <|python_tag|>, Mistral uses
|
||||
# [TOOL_CALLS]; advertising tools for those enables a pill the parser
|
||||
# can't honour. GGUF is unaffected -- llama-server normalises every
|
||||
# format into structured deltas.
|
||||
# Our safetensors loop only parses <tool_call>{json}</tool_call>,
|
||||
# <function=name>...</function>, and Gemma native <|tool_call>...<tool_call|>.
|
||||
# Llama uses <|python_tag|>, Mistral uses [TOOL_CALLS]; advertising tools for
|
||||
# those enables a pill the parser can't honour. GGUF is unaffected --
|
||||
# llama-server normalises every format into structured deltas.
|
||||
if (
|
||||
flags.get("supports_tools")
|
||||
and chat_template
|
||||
and "<tool_call>" not in chat_template
|
||||
and "<function=" not in chat_template
|
||||
and "<|tool_call>" not in chat_template
|
||||
):
|
||||
logger.info(
|
||||
"safetensors: template advertises tools but uses an "
|
||||
|
|
@ -1297,6 +1460,24 @@ async def _await_disconnect_then_close(request, resp, cancel_event) -> None:
|
|||
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
|
||||
# where the canvas tool is actually present in the tool schema; otherwise
|
||||
# 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
|
||||
# (mcp__srv__list-issues) that would otherwise leak past this strip.
|
||||
r"<(?:tool_call|function=[\w-]+)>.*?(?:</(?:tool_call|function)>|\Z)"
|
||||
r"|<\|tool_call>.*?(?:<tool_call\|>|\Z)"
|
||||
r"|</(?:tool_call|function)>"
|
||||
r"|<tool_call\|>"
|
||||
r"|</parameter>\s*\Z",
|
||||
_re.DOTALL,
|
||||
)
|
||||
|
|
@ -3221,7 +3404,9 @@ async def get_api_monitor_entry(entry_id: str, current_subject: str = Depends(ge
|
|||
|
||||
@router.post("/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.
|
||||
|
|
@ -3271,6 +3456,13 @@ async def generate_stream(
|
|||
async def stream():
|
||||
gen = None
|
||||
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:
|
||||
gen = backend.generate_chat_response(
|
||||
messages = request.messages,
|
||||
|
|
@ -3285,12 +3477,22 @@ async def generate_stream(
|
|||
)
|
||||
_DONE = object()
|
||||
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)
|
||||
if chunk is _DONE:
|
||||
completed = True
|
||||
break
|
||||
yield f"data: {json.dumps({'content': chunk})}\n\n"
|
||||
completed = True
|
||||
yield "data: [DONE]\n\n"
|
||||
if completed:
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
except asyncio.CancelledError:
|
||||
cancel_event.set()
|
||||
|
|
@ -3302,6 +3504,7 @@ async def generate_stream(
|
|||
logger.error(f"Error during generation: {e}", exc_info = True)
|
||||
yield f"data: {json.dumps({'error': _friendly_error(e)})}\n\n"
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
if not completed and not cancel_event.is_set():
|
||||
cancel_event.set()
|
||||
backend.reset_generation_state()
|
||||
|
|
@ -4725,6 +4928,9 @@ async def openai_chat_completions(
|
|||
_tracker.__enter__()
|
||||
|
||||
async def audio_input_stream():
|
||||
disconnect_watcher = asyncio.create_task(
|
||||
_await_disconnect_then_cancel(request, cancel_event)
|
||||
)
|
||||
try:
|
||||
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))
|
||||
yield f"data: {json.dumps({'error': {'message': _friendly_error(e), 'type': 'server_error'}})}\n\n"
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
_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:
|
||||
try:
|
||||
full_text = "".join(audio_input_generate())
|
||||
|
|
@ -4937,6 +5153,28 @@ async def openai_chat_completions(
|
|||
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
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) ──────────────────
|
||||
# `_effective_enable_tools` lets `unsloth run --enable-tools/--disable-tools`
|
||||
# hard-override the per-request value, else falls back to
|
||||
|
|
@ -5049,6 +5287,9 @@ async def openai_chat_completions(
|
|||
|
||||
async def gguf_tool_stream():
|
||||
gen = None
|
||||
disconnect_watcher = asyncio.create_task(
|
||||
_await_disconnect_then_cancel(request, cancel_event)
|
||||
)
|
||||
try:
|
||||
yield _chat_role_chunk(completion_id, created, model_name)
|
||||
|
||||
|
|
@ -5056,9 +5297,25 @@ async def openai_chat_completions(
|
|||
# stays free for disconnect detection.
|
||||
gen = gguf_generate_with_tools()
|
||||
prev_text = ""
|
||||
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||
_stream_usage = None
|
||||
_stream_timings = None
|
||||
_stream_finish = None
|
||||
|
||||
def _flush_reasoning_extractor():
|
||||
final_reasoning, final_visible = reasoning_extractor.finish()
|
||||
chunks = []
|
||||
if final_reasoning:
|
||||
chunks.append(
|
||||
_gguf_chat_delta_line(
|
||||
ChoiceDelta(reasoning_content = final_reasoning)
|
||||
)
|
||||
)
|
||||
if final_visible:
|
||||
api_monitor.append_reply(monitor_id, final_visible)
|
||||
chunks.append(_gguf_chat_delta_line(ChoiceDelta(content = final_visible)))
|
||||
return chunks
|
||||
|
||||
while True:
|
||||
if cancel_event.is_set():
|
||||
break
|
||||
|
|
@ -5077,7 +5334,10 @@ async def openai_chat_completions(
|
|||
# cumulative cursor so the next assistant turn
|
||||
# streams cleanly.
|
||||
if not event["text"]:
|
||||
for chunk in _flush_reasoning_extractor():
|
||||
yield chunk
|
||||
prev_text = ""
|
||||
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||
# Emit tool status as a custom SSE event (including
|
||||
# empty ones to clear UI badges)
|
||||
status_data = json.dumps(
|
||||
|
|
@ -5091,7 +5351,10 @@ async def openai_chat_completions(
|
|||
|
||||
if event["type"] in ("tool_start", "tool_end"):
|
||||
if event["type"] == "tool_start":
|
||||
for chunk in _flush_reasoning_extractor():
|
||||
yield chunk
|
||||
prev_text = ""
|
||||
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||
yield f"data: {json.dumps(event)}\n\n"
|
||||
continue
|
||||
|
||||
|
|
@ -5113,15 +5376,33 @@ async def openai_chat_completions(
|
|||
prev_text = clean_cumulative
|
||||
if not new_text:
|
||||
continue
|
||||
api_monitor.append_reply(monitor_id, new_text)
|
||||
yield _chat_content_chunk(completion_id, created, model_name, new_text)
|
||||
reasoning_delta, visible_delta = reasoning_extractor.feed(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(
|
||||
completion_id,
|
||||
created,
|
||||
model_name,
|
||||
_clamp_finish_reason(_stream_finish),
|
||||
for chunk in _flush_reasoning_extractor():
|
||||
yield 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(
|
||||
payload,
|
||||
completion_id,
|
||||
|
|
@ -5150,6 +5431,7 @@ async def openai_chat_completions(
|
|||
error_chunk = _openai_stream_error_chunk(e)
|
||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
if gen is not None:
|
||||
try:
|
||||
gen.close()
|
||||
|
|
@ -5157,7 +5439,16 @@ async def openai_chat_completions(
|
|||
pass
|
||||
_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) ─────────────────────
|
||||
|
||||
|
|
@ -5193,6 +5484,9 @@ async def openai_chat_completions(
|
|||
_tracker.__enter__()
|
||||
|
||||
async def gguf_stream_chunks():
|
||||
disconnect_watcher = asyncio.create_task(
|
||||
_await_disconnect_then_cancel(request, cancel_event)
|
||||
)
|
||||
try:
|
||||
yield _chat_role_chunk(completion_id, created, model_name)
|
||||
|
||||
|
|
@ -5200,6 +5494,7 @@ async def openai_chat_completions(
|
|||
# stays free for disconnect detection.
|
||||
gen = gguf_generate()
|
||||
prev_text = ""
|
||||
reasoning_extractor = _new_chat_reasoning_extractor()
|
||||
_stream_usage = None
|
||||
_stream_timings = None
|
||||
_stream_finish = None
|
||||
|
|
@ -5233,15 +5528,38 @@ async def openai_chat_completions(
|
|||
prev_text = cumulative
|
||||
if not new_text:
|
||||
continue
|
||||
api_monitor.append_reply(monitor_id, new_text)
|
||||
yield _chat_content_chunk(completion_id, created, model_name, new_text)
|
||||
reasoning_delta, visible_delta = reasoning_extractor.feed(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(
|
||||
completion_id,
|
||||
created,
|
||||
model_name,
|
||||
_clamp_finish_reason(_stream_finish),
|
||||
final_reasoning, final_visible = reasoning_extractor.finish()
|
||||
if final_reasoning:
|
||||
yield _gguf_chat_delta_line(ChoiceDelta(reasoning_content = final_reasoning))
|
||||
if final_visible:
|
||||
api_monitor.append_reply(monitor_id, final_visible)
|
||||
yield _gguf_chat_delta_line(ChoiceDelta(content = final_visible))
|
||||
|
||||
# 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(
|
||||
payload,
|
||||
completion_id,
|
||||
|
|
@ -5268,9 +5586,19 @@ async def openai_chat_completions(
|
|||
error_chunk = _openai_stream_error_chunk(e)
|
||||
yield f"data: {json.dumps(error_chunk)}\n\n"
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
_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:
|
||||
try:
|
||||
# ``n`` requests several independent completions; the single
|
||||
|
|
@ -5297,14 +5625,24 @@ async def openai_chat_completions(
|
|||
continue
|
||||
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(
|
||||
CompletionChoice(
|
||||
index = _idx,
|
||||
message = CompletionMessage(content = full_text),
|
||||
message = CompletionMessage(**message_kwargs),
|
||||
finish_reason = _clamp_finish_reason(completion_finish),
|
||||
)
|
||||
)
|
||||
_monitor_replies.append(full_text)
|
||||
_monitor_replies.append(visible_text)
|
||||
if completion_usage:
|
||||
# The prompt is shared across all n choices, so count its
|
||||
# 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),
|
||||
),
|
||||
)
|
||||
monitor_reply = full_text
|
||||
monitor_reply = _monitor_replies[-1] if _monitor_replies else ""
|
||||
if _n > 1:
|
||||
monitor_reply = "\n\n".join(
|
||||
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():
|
||||
gen = None
|
||||
disconnect_watcher = asyncio.create_task(
|
||||
_await_disconnect_then_cancel(request, cancel_event)
|
||||
)
|
||||
try:
|
||||
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"
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
if gen is not None:
|
||||
try:
|
||||
gen.close()
|
||||
|
|
@ -5635,7 +5977,16 @@ async def openai_chat_completions(
|
|||
_sf_tracker.__exit__(None, None, None)
|
||||
|
||||
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.
|
||||
try:
|
||||
|
|
@ -5737,6 +6088,9 @@ async def openai_chat_completions(
|
|||
_tracker.__enter__()
|
||||
|
||||
async def stream_chunks():
|
||||
disconnect_watcher = asyncio.create_task(
|
||||
_await_disconnect_then_cancel(request, cancel_event)
|
||||
)
|
||||
try:
|
||||
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"
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
_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 ────────────────────────────────────
|
||||
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 getattr(llama_backend, "reasoning_always_on", False):
|
||||
return True
|
||||
if not getattr(llama_backend, "supports_reasoning", False):
|
||||
return False
|
||||
if getattr(llama_backend, "supports_reasoning", False):
|
||||
return True
|
||||
return False
|
||||
if chat_req.enable_thinking is True:
|
||||
return True
|
||||
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
|
||||
# (finish_reason="tool_calls") don't see a spurious empty message item.
|
||||
output_items: list[dict] = []
|
||||
if reasoning_text and not text and not tool_calls:
|
||||
text = reasoning_text
|
||||
if reasoning_text:
|
||||
output_items.append(_responses_reasoning_output_item(reasoning_text))
|
||||
if text:
|
||||
|
|
@ -7144,8 +7507,8 @@ async def _responses_stream(
|
|||
client = httpx.AsyncClient(timeout = _llama_streaming_generation_timeout())
|
||||
resp = None
|
||||
lines_iter = None
|
||||
disconnect_event = threading.Event()
|
||||
disconnect_watcher = None
|
||||
disconnect_event = threading.Event()
|
||||
try:
|
||||
req = client.build_request(
|
||||
"POST", target_url, json = body, headers = {"Connection": "close"}
|
||||
|
|
@ -7205,10 +7568,10 @@ async def _responses_stream(
|
|||
)
|
||||
return
|
||||
|
||||
lines_iter = resp.aiter_lines()
|
||||
disconnect_watcher = asyncio.create_task(
|
||||
_await_disconnect_then_close(request, resp, disconnect_event)
|
||||
)
|
||||
lines_iter = resp.aiter_lines()
|
||||
async for raw_line in _aiter_llama_stream_items(
|
||||
lines_iter,
|
||||
cancel_event = disconnect_event,
|
||||
|
|
@ -7328,6 +7691,7 @@ async def _responses_stream(
|
|||
|
||||
_apply_usage(chunk_data.get("usage"))
|
||||
except asyncio.CancelledError:
|
||||
disconnect_event.set()
|
||||
api_monitor.finish(monitor_id, "cancelled")
|
||||
raise
|
||||
except (httpx.RemoteProtocolError, httpx.ReadError, httpx.CloseError) as e:
|
||||
|
|
@ -7394,21 +7758,6 @@ async def _responses_stream(
|
|||
"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]]] = []
|
||||
if reasoning_state["opened"]:
|
||||
|
|
@ -7569,7 +7918,15 @@ async def _responses_stream(
|
|||
api_monitor.finish(monitor_id)
|
||||
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")
|
||||
|
|
@ -8185,9 +8542,17 @@ async def _anthropic_tool_stream(
|
|||
drop_until_tool_end = False
|
||||
|
||||
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:
|
||||
while True:
|
||||
if await request.is_disconnected():
|
||||
if cancel_event.is_set() or await request.is_disconnected():
|
||||
cancel_event.set()
|
||||
return
|
||||
event = await asyncio.to_thread(next, gen, _sentinel)
|
||||
|
|
@ -8235,6 +8600,8 @@ async def _anthropic_tool_stream(
|
|||
if _error_event is not None:
|
||||
yield _error_event
|
||||
return
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
|
||||
stop_reason = openai_finish_to_anthropic_stop(
|
||||
captured_finish_reason, had_tool_calls = ends_on_tool_use
|
||||
|
|
@ -8271,9 +8638,17 @@ async def _anthropic_plain_stream(
|
|||
captured_finish_reason = None
|
||||
|
||||
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:
|
||||
while True:
|
||||
if await request.is_disconnected():
|
||||
if cancel_event.is_set() or await request.is_disconnected():
|
||||
cancel_event.set()
|
||||
return
|
||||
cumulative = await asyncio.to_thread(next, gen, _sentinel)
|
||||
|
|
@ -8296,6 +8671,8 @@ async def _anthropic_plain_stream(
|
|||
if _error_event is not None:
|
||||
yield _error_event
|
||||
return
|
||||
finally:
|
||||
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
|
||||
|
||||
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):
|
||||
|
|
@ -9139,6 +9516,19 @@ async def _openai_passthrough_stream(
|
|||
response ``id``, ``finish_reason`` (including ``"tool_calls"``),
|
||||
``delta.tool_calls``, and any client-requested trailing ``usage`` chunk so
|
||||
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"
|
||||
body = _build_openai_passthrough_body(
|
||||
|
|
@ -9187,7 +9577,7 @@ async def _openai_passthrough_stream(
|
|||
except Exception:
|
||||
pass
|
||||
_tracker.__exit__(None, None, None)
|
||||
return StreamingResponse(
|
||||
return _SameTaskStreamingResponse(
|
||||
iter(()),
|
||||
media_type = "text/event-stream",
|
||||
headers = {
|
||||
|
|
@ -9238,6 +9628,29 @@ async def _openai_passthrough_stream(
|
|||
_await_disconnect_then_close(request, resp, cancel_event)
|
||||
)
|
||||
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:
|
||||
lines_iter = resp.aiter_lines()
|
||||
async for raw_line in _aiter_llama_stream_items(
|
||||
|
|
@ -9251,23 +9664,117 @@ async def _openai_passthrough_stream(
|
|||
continue
|
||||
if not raw_line.startswith("data: "):
|
||||
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
|
||||
# deltas with index>=1 so only the first call streams. Only
|
||||
# lines carrying tool_calls are reparsed; everything else is
|
||||
# relayed byte-for-byte.
|
||||
if payload.parallel_tool_calls is False and '"tool_calls"' in 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_id,
|
||||
raw_line,
|
||||
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,
|
||||
# finish_reason, delta.tool_calls, and usage chunks.
|
||||
yield raw_line + "\n\n"
|
||||
if monitor_event == "done" or raw_line[6:].strip() == "[DONE]":
|
||||
if monitor_event == "done":
|
||||
monitor_done = True
|
||||
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:
|
||||
api_monitor.finish(
|
||||
monitor_id,
|
||||
|
|
@ -9303,7 +9810,24 @@ async def _openai_passthrough_stream(
|
|||
)
|
||||
_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:
|
||||
_tracker.__exit__(None, None, None)
|
||||
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 sys
|
||||
import time
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
|
||||
_backend = os.path.join(os.path.dirname(__file__), "..")
|
||||
|
|
@ -40,6 +41,123 @@ def test_stream_first_item_deadline_after_headers():
|
|||
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():
|
||||
async def _run(cancel_parent):
|
||||
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"
|
||||
|
||||
|
||||
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):
|
||||
"""A tool call not in the per-request list must be refused by the GGUF
|
||||
agentic loop (mirroring the safetensors path)."""
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ from routes.inference import (
|
|||
_openai_passthrough_stream,
|
||||
_openai_stream_usage_chunk,
|
||||
_proxy_to_external_provider,
|
||||
_SameTaskStreamingResponse,
|
||||
_set_or_prepend_system_message,
|
||||
openai_completions,
|
||||
openai_embeddings,
|
||||
|
|
@ -1245,6 +1246,79 @@ class TestGgufVisionToolRouting:
|
|||
|
||||
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):
|
||||
import routes.inference as inf_mod
|
||||
|
||||
|
|
@ -1390,6 +1464,152 @@ class TestGgufVisionToolRouting:
|
|||
assert "confirm_tool_calls requires stream=true" in entry["error"]
|
||||
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):
|
||||
import routes.inference as inf_mod
|
||||
|
||||
|
|
@ -1552,6 +1772,61 @@ class TestApiMonitorProviderAndCompletionStreams:
|
|||
async def is_disconnected(self):
|
||||
return False
|
||||
|
||||
async def _run_passthrough_stream(self, monkeypatch, lines):
|
||||
import routes.inference as inf_mod
|
||||
|
||||
class Request:
|
||||
async def is_disconnected(self):
|
||||
return False
|
||||
|
||||
async def fake_send(*_args, **_kwargs):
|
||||
return httpx.Response(200, content = b"")
|
||||
|
||||
async def fake_items(*_args, **_kwargs):
|
||||
for line in lines:
|
||||
yield line
|
||||
|
||||
monitor = ApiMonitor(max_entries = 3)
|
||||
monkeypatch.setattr(inf_mod, "api_monitor", monitor)
|
||||
monkeypatch.setattr(inf_mod, "_send_stream_with_preheader_cancel", fake_send)
|
||||
monkeypatch.setattr(inf_mod, "_aiter_llama_stream_items", fake_items)
|
||||
monitor_id = monitor.start(
|
||||
endpoint = "/v1/chat/completions",
|
||||
method = "POST",
|
||||
model = "gguf",
|
||||
prompt = "hi",
|
||||
)
|
||||
payload = ChatCompletionRequest(
|
||||
model = "default",
|
||||
messages = [ChatMessage(role = "user", content = "hi")],
|
||||
stream = True,
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
response = await _openai_passthrough_stream(
|
||||
Request(),
|
||||
threading.Event(),
|
||||
SimpleNamespace(
|
||||
base_url = "http://llama.test",
|
||||
context_length = 4096,
|
||||
_request_reasoning_kwargs = lambda *_args, **_kwargs: None,
|
||||
),
|
||||
payload,
|
||||
"gguf",
|
||||
"chatcmpl-test",
|
||||
monitor_id = monitor_id,
|
||||
)
|
||||
chunks = [chunk async for chunk in response.body_iterator]
|
||||
return SimpleNamespace(chunks = chunks, body = "".join(chunks), monitor = monitor)
|
||||
|
||||
def test_external_non_streaming_json_updates_monitor(self, monkeypatch):
|
||||
async def _run():
|
||||
import routes.inference as inf_mod
|
||||
|
|
@ -1980,6 +2255,7 @@ class TestApiMonitorProviderAndCompletionStreams:
|
|||
"chatcmpl-test",
|
||||
monitor_id = monitor_id,
|
||||
)
|
||||
assert isinstance(response, _SameTaskStreamingResponse)
|
||||
iterator = response.body_iterator
|
||||
first = await anext(iterator)
|
||||
assert "hello" in first
|
||||
|
|
@ -1997,6 +2273,88 @@ class TestApiMonitorProviderAndCompletionStreams:
|
|||
|
||||
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):
|
||||
async def _run():
|
||||
import routes.inference as inf_mod
|
||||
|
|
@ -2058,65 +2416,20 @@ class TestApiMonitorProviderAndCompletionStreams:
|
|||
|
||||
def test_passthrough_clean_eof_finalizes_monitor(self, monkeypatch):
|
||||
async def _run():
|
||||
import routes.inference as inf_mod
|
||||
|
||||
class Request:
|
||||
async def is_disconnected(self):
|
||||
return False
|
||||
|
||||
async def fake_send(*_args, **_kwargs):
|
||||
return httpx.Response(200, content = b"")
|
||||
|
||||
async def fake_items(*_args, **_kwargs):
|
||||
yield 'data: {"choices":[{"delta":{"content":"hello"}}]}'
|
||||
|
||||
monitor = ApiMonitor(max_entries = 3)
|
||||
monkeypatch.setattr(inf_mod, "api_monitor", monitor)
|
||||
monkeypatch.setattr(inf_mod, "_send_stream_with_preheader_cancel", fake_send)
|
||||
monkeypatch.setattr(inf_mod, "_aiter_llama_stream_items", fake_items)
|
||||
monitor_id = monitor.start(
|
||||
endpoint = "/v1/chat/completions",
|
||||
method = "POST",
|
||||
model = "gguf",
|
||||
prompt = "hi",
|
||||
)
|
||||
payload = ChatCompletionRequest(
|
||||
model = "default",
|
||||
messages = [ChatMessage(role = "user", content = "hi")],
|
||||
stream = True,
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
],
|
||||
result = await self._run_passthrough_stream(
|
||||
monkeypatch,
|
||||
['data: {"choices":[{"delta":{"content":"hello"}}]}'],
|
||||
)
|
||||
chunks = result.chunks
|
||||
|
||||
response = await _openai_passthrough_stream(
|
||||
Request(),
|
||||
threading.Event(),
|
||||
SimpleNamespace(
|
||||
base_url = "http://llama.test",
|
||||
context_length = 4096,
|
||||
_request_reasoning_kwargs = lambda *_args, **_kwargs: None,
|
||||
),
|
||||
payload,
|
||||
"gguf",
|
||||
"chatcmpl-test",
|
||||
monitor_id = monitor_id,
|
||||
)
|
||||
chunks = []
|
||||
async for chunk in response.body_iterator:
|
||||
chunks.append(chunk)
|
||||
|
||||
assert chunks == ['data: {"choices":[{"delta":{"content":"hello"}}]}\n\n']
|
||||
[entry] = monitor.snapshot()
|
||||
assert chunks[0] == 'data: {"choices":[{"delta":{"content":"hello"}}]}\n\n'
|
||||
compact = "".join(chunks).replace(" ", "")
|
||||
assert '"finish_reason":"stop"' in compact
|
||||
assert chunks[-1] == "data: [DONE]\n\n"
|
||||
[entry] = result.monitor.snapshot()
|
||||
assert entry["status"] == "completed"
|
||||
assert entry["reply"] == "hello"
|
||||
assert monitor.active_count() == 0
|
||||
assert result.monitor.active_count() == 0
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
|
|
|||
|
|
@ -59,8 +59,10 @@ from models.inference import (
|
|||
ResponsesUsage,
|
||||
)
|
||||
from routes.inference import (
|
||||
_SameTaskStreamingResponse,
|
||||
_build_chat_request,
|
||||
_chat_tool_calls_to_responses_output,
|
||||
_extract_responses_reasoning,
|
||||
_normalise_responses_input,
|
||||
_responses_tool_output_content,
|
||||
_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"]
|
||||
|
||||
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):
|
||||
import routes.inference as inf_mod
|
||||
|
||||
|
|
@ -927,6 +938,38 @@ class TestResponsesNonStreamingAdapter:
|
|||
assert [item["type"] for item in body["output"]] == ["message"]
|
||||
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):
|
||||
body = self._run_with_message(
|
||||
monkeypatch,
|
||||
|
|
@ -949,7 +992,7 @@ class TestResponsesNonStreamingAdapter:
|
|||
assert [item["type"] for item in body["output"]] == ["message"]
|
||||
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"})
|
||||
body = self._run_with_message(
|
||||
monkeypatch,
|
||||
|
|
@ -957,9 +1000,8 @@ class TestResponsesNonStreamingAdapter:
|
|||
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"][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):
|
||||
chunks = [
|
||||
{"choices": [{"delta": {"content": "<thi"}}]},
|
||||
|
|
@ -1286,7 +1358,7 @@ class TestResponsesStreamAdapter:
|
|||
assert entry["status"] == "completed"
|
||||
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
|
||||
|
||||
class FakeExtractor:
|
||||
|
|
@ -1327,15 +1399,16 @@ class TestResponsesStreamAdapter:
|
|||
|
||||
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()
|
||||
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 = [
|
||||
{"choices": [{"delta": {"content": "show <thi"}}]},
|
||||
{"choices": [{"delta": {"content": "nk>x</think> tags"}}]},
|
||||
{"choices": [{"delta": {"content": "<thi"}}]},
|
||||
{"choices": [{"delta": {"content": "nk>plan</think>answer"}}]},
|
||||
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
||||
]
|
||||
self._install_stream_mock(monkeypatch, chunks)
|
||||
|
|
@ -1350,13 +1423,15 @@ class TestResponsesStreamAdapter:
|
|||
|
||||
reasoning_deltas = self._payloads(lines, "response.reasoning_text.delta")
|
||||
text_deltas = self._payloads(lines, "response.output_text.delta")
|
||||
assert reasoning_deltas == []
|
||||
assert "".join(event["delta"] for event in text_deltas) == "show <think>x</think> tags"
|
||||
assert "".join(event["delta"] for event in reasoning_deltas) == "plan"
|
||||
assert "".join(event["delta"] for event in text_deltas) == "answer"
|
||||
completed = self._payloads(lines, "response.completed")[0]
|
||||
assert [item["type"] for item in completed["response"]["output"]] == ["message"]
|
||||
assert completed["response"]["output"][0]["content"][0]["text"] == (
|
||||
"show <think>x</think> tags"
|
||||
)
|
||||
assert [item["type"] for item in completed["response"]["output"]] == [
|
||||
"reasoning",
|
||||
"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):
|
||||
chunks = [
|
||||
|
|
@ -1384,7 +1459,7 @@ class TestResponsesStreamAdapter:
|
|||
"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 = [
|
||||
{"choices": [{"delta": {"content": "<think>plan</think>"}}]},
|
||||
{"choices": [], "usage": {"prompt_tokens": 2, "completion_tokens": 3}},
|
||||
|
|
@ -1402,14 +1477,34 @@ class TestResponsesStreamAdapter:
|
|||
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 "".join(event["delta"] for event in text_deltas) == "plan"
|
||||
assert text_deltas == []
|
||||
completed = self._payloads(lines, "response.completed")[0]
|
||||
assert [item["type"] for item in completed["response"]["output"]] == [
|
||||
"reasoning",
|
||||
"message",
|
||||
]
|
||||
assert [item["type"] for item in completed["response"]["output"]] == ["reasoning"]
|
||||
assert completed["response"]["output"][0]["content"][0]["text"] == "plan"
|
||||
|
||||
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"][1]["content"][0]["text"] == "plan"
|
||||
|
||||
def test_structured_reasoning_content_streams_as_reasoning(self, monkeypatch):
|
||||
chunks = [
|
||||
|
|
|
|||
|
|
@ -203,6 +203,20 @@ def test_detect_safetensors_features_function_xml_format_keeps_tools_on():
|
|||
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
|
||||
# calls as ``<tool_call>\n<function=name>...``. Faithful slice so the
|
||||
# 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.
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
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, 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):
|
||||
text = "<function=python><parameter=code>print('hi')</parameter></function>"
|
||||
result = parse_tool_calls_from_text(text)
|
||||
|
|
@ -121,6 +167,7 @@ class TestParser:
|
|||
|
||||
def test_has_tool_signal(self):
|
||||
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 not has_tool_signal("hello world")
|
||||
|
||||
|
|
@ -139,6 +186,8 @@ class TestParser:
|
|||
def test_strip_markup_closed(self):
|
||||
text = "before <tool_call>{}</tool_call> 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):
|
||||
text = "before <tool_call>{partial"
|
||||
|
|
@ -146,6 +195,7 @@ class TestParser:
|
|||
assert strip_tool_markup(text, final = True) == "before"
|
||||
# Without final=True the unclosed run is preserved.
|
||||
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):
|
||||
raw = 'before <tool_call>{"name":"web_search"'
|
||||
|
|
|
|||
|
|
@ -106,6 +106,54 @@ class TestParityWithJsonStyle:
|
|||
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:
|
||||
def test_auto_heal_still_repairs_unclosed_function(self):
|
||||
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).
|
||||
|
||||
|
||||
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) ───────────────────
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -146,24 +146,50 @@ def test_async_generators_cleanup_tracker_in_finally():
|
|||
)
|
||||
|
||||
|
||||
def test_streaming_responses_have_no_background_task():
|
||||
top = None
|
||||
for n in ast.walk(_TREE):
|
||||
if isinstance(n, ast.AsyncFunctionDef) and n.name == "openai_chat_completions":
|
||||
top = n
|
||||
break
|
||||
assert top is not None
|
||||
def test_chat_completions_streams_avoid_starlette_task_group():
|
||||
top = _async_function("openai_chat_completions")
|
||||
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
|
||||
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
|
||||
kwargs = {kw.arg for kw in sub.keywords if kw.arg}
|
||||
assert "background" not in kwargs, (
|
||||
"StreamingResponse in openai_chat_completions must not pass "
|
||||
"`background=` -- cleanup now lives in the generator's finally "
|
||||
"block; a BackgroundTask would be skipped on abrupt disconnect"
|
||||
)
|
||||
if sub.func.id == "StreamingResponse":
|
||||
legacy_calls.append(sub.lineno)
|
||||
if sub.func.id == "_SameTaskStreamingResponse":
|
||||
same_task_calls += 1
|
||||
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():
|
||||
|
|
@ -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 ───────────────────────────────────────
|
||||
|
||||
_WANTED = {
|
||||
|
|
@ -222,6 +275,21 @@ def _load_registry_module():
|
|||
return mod
|
||||
|
||||
|
||||
def _load_same_task_response_module():
|
||||
for n in _TREE.body:
|
||||
if isinstance(n, ast.ClassDef) and n.name == "_SameTaskStreamingResponse":
|
||||
source = ast.get_source_segment(SRC, n)
|
||||
break
|
||||
else:
|
||||
raise AssertionError("_SameTaskStreamingResponse missing")
|
||||
mod = {}
|
||||
exec(
|
||||
"class StreamingResponse: pass\nclass ClientDisconnect(Exception): pass\n" + source,
|
||||
mod,
|
||||
)
|
||||
return mod
|
||||
|
||||
|
||||
def _make_stream(tracker, raise_exc):
|
||||
async def gen():
|
||||
try:
|
||||
|
|
@ -326,6 +394,40 @@ def test_finally_cleanup_on_aclose():
|
|||
assert "sid-abort" not in m["_CANCEL_REGISTRY"]
|
||||
|
||||
|
||||
def test_same_task_response_closes_body_iterator_on_send_disconnect():
|
||||
m = _load_same_task_response_module()
|
||||
closed = False
|
||||
|
||||
async def body():
|
||||
nonlocal closed
|
||||
try:
|
||||
yield "data: first\n\n"
|
||||
finally:
|
||||
closed = True
|
||||
|
||||
async def run():
|
||||
agen = body()
|
||||
await agen.__anext__()
|
||||
response = m["_SameTaskStreamingResponse"].__new__(m["_SameTaskStreamingResponse"])
|
||||
response.body_iterator = agen
|
||||
response.background = None
|
||||
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():
|
||||
# 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.
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue