Compare commits

...
Sign in to create a new pull request.

33 commits

Author SHA1 Message Date
danielhanchen
35e18e7680 Quote-aware Gemma strip, symmetric unstarted cleanup, ReDoS anchor
Address review findings on the tool-strip and streaming paths:

- strip_tool_call_markup stripped Gemma-native spans with a plain regex that
  stops at the first <tool_call|>, so a literal close marker inside a
  <|"|>-quoted argument truncated the span and leaked its suffix into visible
  text. A brace/quote-aware _strip_gemma_native_spans now removes complete
  spans (keeping an incomplete one unless final), matching the parser's own
  balance logic.

- The Gemma close pattern this PR added (<\|tool_call>.*?<tool_call\|>) had no
  \Z fallback, so a run of unclosed markers backtracked from every open
  position (quadratic, and the streaming stripper re-scans per token). It is
  now anchored to (?:<tool_call|>|\Z) like routes/inference.py's _TOOL_XML_RE,
  linear with identical output on well-formed input.

- _SameTaskStreamingResponse added unstarted_cleanup for the OpenAI passthrough,
  but the local GGUF/safetensors streams that enter _TrackedCancel before
  returning only unregister in the generator finally, which never runs if the
  client disconnects before the body iterator starts, leaking cancel-registry
  entries. Each such stream now passes unstarted_cleanup to exit its tracker.

- __call__ reads _unstarted_cleanup via getattr so a response built through
  __new__ (the cancel-timing test) without __init__ does not raise
  AttributeError; the test also sets the attribute explicitly.

- Document that the verbatim /v1/chat/completions passthrough delegates
  <think>/<|tool_call> splitting to llama-server (--jinja, --reasoning-format
  auto) and is intentionally not re-parsed locally, noting the llama.cpp
  dependency.

Adds a regression test for the close-marker-inside-quoted-argument strip.
2026-06-23 13:28:49 +00:00
Daniel Han
2d72bc0494
Merge branch 'main' into fix-gemma4-openai-api-streams 2026-06-23 05:51:32 -07:00
pre-commit-ci[bot]
75b65b4834 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-22 16:24:48 +00:00
danielhanchen
b0dbe43867 Guard nested markers, reset on disconnect, clean unstarted streams
Three follow-ups on the tool-parse and streaming paths:

- parse_tool_calls_from_text only skipped markers that fell inside a span it
  had already parsed successfully, so when an unquoted Gemma argument contained
  a literal marker (code:<|tool_call>call:terminal{...}<tool_call|>) the outer
  object failed to normalize, its span was never recorded, and the inner marker
  was promoted to a standalone terminal call. Candidates nested inside any other
  candidate's brace span are now skipped regardless of whether the enclosing
  candidate parsed, so a marker in malformed outer data is never executed.

- /generate/stream skipped backend.reset_generation_state() when the disconnect
  watcher set cancel_event between chunks: the loop broke and the finally's reset
  is guarded on cancel_event being unset. A subprocess backend kept decoding
  after the client left. The cancel-break path now resets the backend.

- _SameTaskStreamingResponse threw CancelledError / called aclose() on the body
  iterator on a send-side disconnect, but neither runs the try/finally of a
  generator that never started (early disconnect on http.response.start), so the
  passthrough's eagerly-opened upstream httpx stream and cancel-registry entry
  leaked. It now tracks whether the body started and, when it did not, runs an
  optional unstarted_cleanup hook; the OpenAI passthrough wires it to close the
  upstream resp/client and exit the cancel tracker.

Adds a nested-unquoted-marker regression test.
2026-06-22 16:24:09 +00:00
pre-commit-ci[bot]
d376756e9c [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-22 15:37:37 +00:00
danielhanchen
e1fdb140b4 Watch disconnects on Anthropic streams; keep timestamps in Gemma values
Two follow-ups on the streaming and tool-parse paths:

- _anthropic_tool_stream and _anthropic_plain_stream drove generation through
  asyncio.to_thread(next, gen, ...) and only polled is_disconnected() between
  events, so a client disconnect during prefill or a long generation/tool step
  held the decode slot until the next event or a failed send. Both now run the
  _await_disconnect_then_cancel watcher used by the other local streams, stop it
  in finally, and break promptly when cancel_event is set.

- _GEMMA_NEXT_KEY_RE treated any comma followed by word-chars-then-colon as the
  next key, so a bare value such as "meet at 10:00, 11:00 tomorrow" was split
  into bogus keys. The next-key token must now be identifier-shaped (start with
  a letter or underscore), so a comma before a timestamp, ratio, or other
  numeric-then-colon text stays part of the value.

Adds a timestamp-in-bare-value regression test.
2026-06-22 15:37:00 +00:00
pre-commit-ci[bot]
6618213da5 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-22 15:16:46 +00:00
danielhanchen
1e58c3707d Harden Gemma array parsing, XML-parameter guard, and stream teardown
Address five review findings on the Gemma tool-call and OpenAI passthrough
streaming paths:

- parse_tool_calls_from_text collected JSON and Gemma markers without the
  _inside_open_parameter guard, so a marker embedded in an existing
  <function=...><parameter=...> value was promoted to a separate tool call.
  Candidates that start inside an open XML parameter are now skipped, matching
  the guard the XML-style parser already applies.

- _quote_gemma_array_elements preserved array elements starting with { or [
  verbatim, so an array of objects (items:[{path:a}]) or a nested array failed
  json.loads and the whole call was dropped. Object and nested-array elements
  are now normalised recursively.

- _openai_passthrough_stream synthesized a finish chunk before a trailing
  usage-only chunk and set saw_finish_reason, which made the EOF guard skip the
  [DONE] sentinel. The EOF path now emits [DONE] whenever the upstream omitted
  it, even after a finish chunk was already synthesized.

- /generate/stream drove generation through asyncio.to_thread with no
  disconnect watcher, so a client disconnect during a long generation went
  unnoticed until the next send. It now runs _await_disconnect_then_cancel
  against the request, matching the other local streaming endpoints.

- _SameTaskStreamingResponse closed the body iterator with aclose() on a
  send-side disconnect, raising GeneratorExit so the generators' cancellation
  handlers (which finish the api_monitor entry) never ran. It now throws
  CancelledError, falling back to aclose() when athrow is unavailable.

Extends tests/test_gemma_tool_parse_edge_cases.py with array-of-objects,
nested-array, and marker-inside-XML-parameter cases.
2026-06-22 15:14:42 +00:00
danielhanchen
520df9fe9d Quote bare Gemma array elements; order finish before trailing usage
- _quote_gemma_object_keys skipped array values, so a Gemma call with a
  bare-string array argument like labels:[bug,ui] produced invalid JSON and
  the whole tool call was dropped. Array values are now scanned and bare
  string elements quoted, while numbers, quoted strings, and JSON literals
  are preserved.

- In the OpenAI passthrough stream, a trailing usage-only chunk
  (stream_options.include_usage) that arrived before any finish chunk was
  relayed before the synthetic finish, producing usage -> finish -> [DONE].
  Emit the synthetic finish before that usage chunk so the order matches the
  other streams (finish -> usage -> [DONE]).

Extends tests/test_gemma_tool_parse_edge_cases.py with the bare-array cases.
2026-06-22 12:59:00 +00:00
pre-commit-ci[bot]
11564256b9 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-22 12:26:05 +00:00
danielhanchen
eca61f7f5e Parse tool calls in document order and skip nested markers both ways
Unify the JSON- and Gemma-format tool-call passes into a single
position-ordered scan:

- Calls are now emitted in byte order across both formats, so a mixed
  output like `<|tool_call>call:create{...}<tool_call|> ... <tool_call>
  {"name":"read",...}</tool_call>` executes create before read, matching
  the order they appear in (tools run in returned order).

- A candidate that starts inside an already-accepted call's span is
  skipped, in both directions: a JSON marker inside a Gemma argument and a
  Gemma marker inside a JSON argument are treated as data, not promoted to
  a second executable tool call.

Extends tests/test_gemma_tool_parse_edge_cases.py with the ordering and
JSON-in-Gemma nesting cases.
2026-06-22 12:24:41 +00:00
danielhanchen
0083826dde Emit the terminal finish_reason chunk in GGUF streams
The OpenAI chat-completions GGUF tool stream and plain stream both built a
final ChatCompletionChunk carrying finish_reason but never yielded it, so
clients received the optional usage chunk and [DONE] with no chunk carrying
finish_reason. OpenAI-compatible consumers rely on that terminal choice to
distinguish stop/length/tool_calls. Yield it before the usage chunk and
[DONE], matching the other streaming paths.
2026-06-22 11:49:46 +00:00
danielhanchen
b3e244d658 Harden Gemma tool-call parsing and stream-error detection
Address three issues in the Gemma-native tool-call path:

- _quote_gemma_object_keys stopped a bare (unquoted) string value at the
  first comma, so an argument like `location:New York, NY` was split
  mid-value and the synthesized JSON failed to parse, dropping the whole
  tool call. A bare value now ends only at `}` or a comma that begins the
  next `key:` pair.

- parse_tool_calls_from_text scanned the entire response for Gemma markers
  even inside a tool call already parsed from a `<tool_call>{...}` JSON
  block, so a marker-like string inside an argument (data) was promoted to
  a second, unintended tool call. Matches inside an already-consumed call
  span are now skipped.

- _openai_passthrough_stream relied on _monitor_openai_sse_line to flag a
  stream error, which returns early when monitor_id is None
  (skip_api_monitor), so an upstream error chunk left saw_stream_error
  unset and the synthetic-finish guard emitted a successful finish_reason
  after a failed stream. Error chunks are now detected independently of API
  monitoring.

Adds tests/test_gemma_tool_parse_edge_cases.py covering the comma and
marker-injection cases.
2026-06-22 11:03:21 +00:00
pre-commit-ci[bot]
491586c556 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-22 04:59:32 +00:00
danielhanchen
21edbb94a3 Merge branch 'main' into fix-gemma4-openai-api-streams
Resolve studio/backend/routes/inference.py:
- keep the reasoning_content split (the #6471 fix) on the GGUF chat streams
- keep each streaming path's direct _SameTaskStreamingResponse return (the
  cancel-scope fix) and route main's new _sse_streaming_response helper through
  _SameTaskStreamingResponse too, so every SSE path stays in the request task
- keep both the new _SameTaskStreamingResponse class and main's
  _aclose_stream_resources teardown helper

All 466 tests across the 9 changed test files pass.
2026-06-22 04:58:22 +00:00
pre-commit-ci[bot]
4b68cac412 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-20 07:15:30 +00:00
Daniel Han
2b7360176b Gemma 4: parse bare-string tool args and keep safetensors tools for native <|tool_call>
- Quote bare unquoted string values in Gemma native tool-call args (e.g.
  {location:Tokyo,unit:celsius}) so they parse; JSON scalars stay typed.
- Stop _detect_safetensors_features from suppressing supports_tools for
  templates that emit Gemma native <|tool_call>, which the shared parser
  now reads.
- Add tests for both.
2026-06-20 07:14:44 +00:00
Daniel Han
3a2c1c2b22
Merge branch 'main' into fix-gemma4-openai-api-streams 2026-06-20 00:00:27 -07:00
pre-commit-ci[bot]
67ccd78440 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-19 20:49:41 +00:00
wasimysaid
db8d03927b Address stream cleanup and Gemma parser reviews 2026-06-19 22:49:06 +02:00
wasimysaid
6a57d3795a Avoid synthetic finish after passthrough errors 2026-06-19 21:54:50 +02:00
wasimysaid
3ca278cb79 Preserve audio stream disconnect cancellation 2026-06-19 21:16:29 +02:00
wasimysaid
5c4e7b5365 Fix tool parser compatibility exports lint 2026-06-19 20:48:10 +02:00
wasimysaid
27e0228fdf Fix Studio passthrough cold stream timeout 2026-06-19 20:40:29 +02:00
wasimysaid
5ba19c2977 Clean up Studio OpenAI stream helpers 2026-06-19 20:09:13 +02:00
wasimysaid
b53015be85 Address OpenAI stream review issues 2026-06-19 19:40:29 +02:00
wasimysaid
bd1e5eabe5 Harden OpenAI chat completion streams 2026-06-19 19:03:09 +02:00
wasimysaid
ede6a2bcee Avoid Responses stream task-group cleanup 2026-06-19 17:23:02 +02:00
wasimysaid
475ff786d8 Address Gemma stream review comments 2026-06-19 17:07:47 +02:00
wasimysaid
4c877cc7d8 Keep reasoning-only Responses output hidden 2026-06-19 16:51:40 +02:00
wasimysaid
6705053e9e Avoid duplicate Responses stream disconnect watcher 2026-06-19 16:33:32 +02:00
pre-commit-ci[bot]
05518f2769 [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
2026-06-19 14:29:10 +00:00
wasimysaid
165f4838e3 Fix Gemma 4 GGUF OpenAI API streams 2026-06-19 16:27:33 +02:00
14 changed files with 2081 additions and 430 deletions

View file

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

View file

@ -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

View file

@ -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):

View file

@ -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

View 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"

View file

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

View file

@ -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)."""

View file

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

View file

@ -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 = [

View file

@ -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.

View file

@ -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"'

View file

@ -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"

View file

@ -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) ───────────────────

View file

@ -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.