diff --git a/studio/backend/core/inference/chat_template_helpers.py b/studio/backend/core/inference/chat_template_helpers.py
new file mode 100644
index 0000000000..833a714ee4
--- /dev/null
+++ b/studio/backend/core/inference/chat_template_helpers.py
@@ -0,0 +1,60 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""
+Dependency-light wrapper around tokenizer.apply_chat_template with a
+kwarg fallback for templates that reject reasoning/tools args.
+"""
+
+from typing import Optional
+
+
+def apply_chat_template_for_generation(
+ tokenizer,
+ messages: list,
+ *,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
+) -> str:
+ """Render the chat prompt. Try richest kwargs first; drop one
+ group at a time on TypeError. Jinja / missing-variable errors
+ propagate."""
+ reasoning_kwargs: dict = {}
+ if enable_thinking is not None:
+ reasoning_kwargs["enable_thinking"] = enable_thinking
+ if reasoning_effort is not None:
+ reasoning_kwargs["reasoning_effort"] = reasoning_effort
+ if preserve_thinking is not None:
+ reasoning_kwargs["preserve_thinking"] = preserve_thinking
+
+ attempts: list[dict] = []
+ if tools and reasoning_kwargs:
+ attempts.append({"tools": tools, **reasoning_kwargs})
+ if tools:
+ attempts.append({"tools": tools})
+ if reasoning_kwargs:
+ attempts.append(dict(reasoning_kwargs))
+ attempts.append({})
+
+ last_exc: Optional[Exception] = None
+ for kwargs in attempts:
+ try:
+ return tokenizer.apply_chat_template(
+ messages,
+ tokenize = False,
+ add_generation_prompt = True,
+ **kwargs,
+ )
+ except TypeError as e:
+ last_exc = e
+ continue
+ except Exception as e:
+ last_exc = e
+ break
+ if last_exc is not None:
+ raise last_exc
+ raise RuntimeError(
+ "apply_chat_template_for_generation: no attempt produced a result"
+ )
diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py
index 4c140013a0..e1620f5ca3 100644
--- a/studio/backend/core/inference/inference.py
+++ b/studio/backend/core/inference/inference.py
@@ -839,6 +839,74 @@ class InferenceBackend:
cancel_event = cancel_event, _adapter_state = use_adapter, **gen_kwargs
)
+ def generate_chat_completion_with_tools(
+ self,
+ messages: list,
+ tools: list,
+ system_prompt: str = "",
+ temperature: float = 0.7,
+ top_p: float = 0.9,
+ top_k: int = 40,
+ min_p: float = 0.0,
+ max_new_tokens: int = 2048,
+ repetition_penalty: float = 1.0,
+ cancel_event = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
+ max_tool_iterations: int = 25,
+ auto_heal_tool_calls: bool = True,
+ tool_call_timeout: int = 300,
+ session_id: Optional[str] = None,
+ ):
+ """Run an agentic tool loop on top of ``generate_chat_response``.
+
+ Yields the same event-dict protocol used by the GGUF path so
+ the route layer can stream both backends through one helper.
+ Each event is one of:
+
+ * ``{"type": "status", "text": ...}``
+ * ``{"type": "content", "text": cumulative_text}``
+ * ``{"type": "tool_start", "tool_name", "tool_call_id", "arguments"}``
+ * ``{"type": "tool_end", "tool_name", "tool_call_id", "result"}``
+ """
+ from core.inference.safetensors_agentic import run_safetensors_tool_loop
+ from core.inference.tools import execute_tool
+
+ def _single_turn(conv: list):
+ # conv already has the system message -- avoid double-prepend.
+ yield from self._generate_chat_response_inner(
+ messages = conv,
+ system_prompt = "",
+ temperature = temperature,
+ top_p = top_p,
+ top_k = top_k,
+ min_p = min_p,
+ max_new_tokens = max_new_tokens,
+ repetition_penalty = repetition_penalty,
+ cancel_event = cancel_event,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
+ )
+
+ initial = list(messages)
+ if system_prompt:
+ initial = [{"role": "system", "content": system_prompt}] + initial
+
+ yield from run_safetensors_tool_loop(
+ single_turn = _single_turn,
+ messages = initial,
+ tools = tools,
+ execute_tool = execute_tool,
+ cancel_event = cancel_event,
+ auto_heal_tool_calls = auto_heal_tool_calls,
+ max_tool_iterations = max_tool_iterations,
+ tool_call_timeout = tool_call_timeout,
+ session_id = session_id,
+ )
+
def generate_chat_response(
self,
messages: list,
@@ -851,10 +919,20 @@ class InferenceBackend:
max_new_tokens: int = 256,
repetition_penalty: float = 1.0,
cancel_event = None,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
) -> Generator[str, None, None]:
"""
Generate response for text or vision models.
The generation lock is acquired by the background generation thread.
+
+ ``tools`` / ``enable_thinking`` / ``reasoning_effort`` /
+ ``preserve_thinking`` are forwarded into
+ ``tokenizer.apply_chat_template`` so templates that understand
+ these kwargs (Qwen3, Llama 3.1+, gpt-oss harmony, ...) advertise
+ the tool schemas and reasoning controls to the model.
"""
yield from self._generate_chat_response_inner(
messages = messages,
@@ -867,6 +945,10 @@ class InferenceBackend:
max_new_tokens = max_new_tokens,
repetition_penalty = repetition_penalty,
cancel_event = cancel_event,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
)
def _generate_chat_response_inner(
@@ -882,6 +964,10 @@ class InferenceBackend:
repetition_penalty: float = 1.0,
cancel_event = None,
_adapter_state = None,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
) -> Generator[str, None, None]:
"""
Inner generation logic. Called by both generate_chat_response
@@ -981,8 +1067,13 @@ class InferenceBackend:
f"Please use a model that includes a chat template, or manually set "
f"one via tokenizer.chat_template before inference."
)
- formatted_prompt = tokenizer.apply_chat_template(
- template_messages, tokenize = False, add_generation_prompt = True
+ formatted_prompt = self._apply_chat_template_for_generation(
+ tokenizer,
+ template_messages,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
)
logger.debug(f"Formatted prompt: {formatted_prompt[:200]}...")
except Exception as e:
@@ -1319,20 +1410,9 @@ class InferenceBackend:
def _is_gpt_oss_model(self, model_name: str = None) -> bool:
"""Check if the given (or active) model uses the gpt-oss harmony protocol."""
- name = (model_name or self.active_model_name or "").lower()
- try:
- from utils.datasets import MODEL_TO_TEMPLATE_MAPPER
+ from utils.datasets import is_gpt_oss_model_name
- # Exact match
- if MODEL_TO_TEMPLATE_MAPPER.get(name) == "gpt-oss":
- return True
- # Partial match (e.g. name-bnb-4bit variants)
- for key, tmpl in MODEL_TO_TEMPLATE_MAPPER.items():
- if tmpl == "gpt-oss" and (key in name or name in key):
- return True
- except Exception:
- pass
- return "gpt-oss" in name
+ return is_gpt_oss_model_name(model_name or self.active_model_name or "")
def generate_stream(
self,
@@ -1715,6 +1795,34 @@ class InferenceBackend:
"Patched RepetitionPenaltyLogitsProcessor with 64-token window for OuteTTS"
)
+ def _apply_chat_template_for_generation(
+ self,
+ tokenizer,
+ messages: list,
+ *,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
+ ) -> str:
+ """Render the chat prompt, peeling kwargs the template does not
+ understand. Delegates to the dependency-light helper module so
+ the fallback chain can be unit-tested without pulling unsloth /
+ torch into the test sandbox.
+ """
+ from core.inference.chat_template_helpers import (
+ apply_chat_template_for_generation,
+ )
+
+ return apply_chat_template_for_generation(
+ tokenizer,
+ messages,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
+ )
+
def format_chat_prompt(self, messages: list, system_prompt: str = None) -> str:
if not self.active_model_name or self.active_model_name not in self.models:
logger.error("No active model available")
diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py
index 91ee1e4203..260e675a73 100644
--- a/studio/backend/core/inference/llama_cpp.py
+++ b/studio/backend/core/inference/llama_cpp.py
@@ -44,6 +44,9 @@ from utils.native_path_leases import child_env_without_native_path_secret
from utils.subprocess_compat import (
windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs,
)
+from core.inference.tool_call_parser import (
+ parse_tool_calls_from_text as _shared_parse_tool_calls_from_text,
+)
logger = get_logger(__name__)
@@ -3904,16 +3907,9 @@ class LlamaCppBackend:
@staticmethod
def _parse_tool_calls_from_text(content: str) -> list[dict]:
- """
- Parse tool calls from XML markup in content text.
-
- Handles formats like:
- {"name":"web_search","arguments":{"query":"..."}}
- ...
- Closing tags (, , ) are all optional
- since models frequently omit them.
- """
- return parse_tool_calls_from_text(content)
+ """Thin wrapper around the shared parser in tool_call_parser
+ so safetensors and llama_cpp pick up the same fixes."""
+ return _shared_parse_tool_calls_from_text(content)
@staticmethod
def _build_openai_messages(
diff --git a/studio/backend/core/inference/mlx_inference.py b/studio/backend/core/inference/mlx_inference.py
index e7bce2d33e..716e4c27a2 100644
--- a/studio/backend/core/inference/mlx_inference.py
+++ b/studio/backend/core/inference/mlx_inference.py
@@ -157,10 +157,59 @@ class MLXInferenceBackend:
"audio_type": None,
"has_audio_input": False,
}
+ # Capture chat_template_info so the worker IPC reply can ship
+ # it back to the parent and the route layer classifies
+ # capabilities the same way as the transformers / GGUF paths.
+ self._populate_chat_template_info(model_name)
logger.info("Model %s loaded successfully", model_name)
return True
+ def _populate_chat_template_info(self, model_name: str) -> None:
+ """Mirror InferenceBackend._load_chat_template_info for MLX.
+
+ Stores ``chat_template_info`` on ``self.models[model_name]``
+ with the resolved ``tokenizer.chat_template`` so
+ ``_detect_safetensors_features`` (route layer) sees the same
+ template the model actually uses."""
+ entry = self.models.get(model_name)
+ if not entry:
+ return
+ tok = entry.get("tokenizer")
+ if tok is None:
+ proc = entry.get("processor")
+ tok = getattr(proc, "tokenizer", None) if proc else None
+ info = {
+ "has_template": False,
+ "template": None,
+ "format_type": "generic",
+ "special_tokens": {},
+ "template_name": None,
+ }
+ try:
+ tpl = getattr(tok, "chat_template", None)
+ if tpl:
+ info["has_template"] = True
+ info["template"] = tpl
+ lower = tpl.lower()
+ if "start_header_id" in lower and "end_header_id" in lower:
+ info["format_type"] = "llama3"
+ elif "[inst]" in lower and "[/inst]" in lower:
+ info["format_type"] = "mistral"
+ elif "<|im_start|>" in lower and "<|im_end|>" in lower:
+ info["format_type"] = "chatml"
+ else:
+ info["format_type"] = "custom"
+ special = {}
+ for attr in ("bos_token", "eos_token", "pad_token"):
+ val = getattr(tok, attr, None)
+ if val:
+ special[attr] = val
+ info["special_tokens"] = special
+ except Exception as exc:
+ logger.warning("MLX chat_template_info capture failed: %s", exc)
+ entry["chat_template_info"] = info
+
def unload_model(self, model_name: str) -> bool:
import mlx.core as mx
import gc
@@ -197,6 +246,14 @@ class MLXInferenceBackend:
max_new_tokens = 256,
repetition_penalty = 1.0,
cancel_event = None,
+ # Reasoning / tool kwargs forwarded by the route + worker -- the
+ # MLX path renders the template via apply_chat_template_for_
+ # generation so these are honoured the same way as the
+ # transformers path.
+ tools = None,
+ enable_thinking = None,
+ reasoning_effort = None,
+ preserve_thinking = None,
) -> Generator[str, None, None]:
if self._model is None:
raise RuntimeError("No model loaded")
@@ -239,6 +296,10 @@ class MLXInferenceBackend:
max_new_tokens,
repetition_penalty,
cancel_event,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
)
else:
yield from self._generate_text(
@@ -250,6 +311,10 @@ class MLXInferenceBackend:
max_new_tokens,
repetition_penalty,
cancel_event,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
)
def _generate_text(
@@ -262,14 +327,26 @@ class MLXInferenceBackend:
max_new_tokens,
repetition_penalty,
cancel_event,
+ *,
+ tools = None,
+ enable_thinking = None,
+ reasoning_effort = None,
+ preserve_thinking = None,
):
from mlx_lm import stream_generate
from mlx_lm.sample_utils import make_sampler, make_logits_processors
- prompt = self._tokenizer.apply_chat_template(
+ from core.inference.chat_template_helpers import (
+ apply_chat_template_for_generation,
+ )
+
+ prompt = apply_chat_template_for_generation(
+ self._tokenizer,
messages,
- tokenize = False,
- add_generation_prompt = True,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
)
if prompt is None:
raise RuntimeError(
@@ -343,20 +420,38 @@ class MLXInferenceBackend:
max_new_tokens,
repetition_penalty,
cancel_event,
+ *,
+ tools = None,
+ enable_thinking = None,
+ reasoning_effort = None,
+ preserve_thinking = None,
):
from mlx_vlm import stream_generate as vlm_stream
- # Apply chat template
- chat_fn = getattr(self._processor, "apply_chat_template", None)
+ from core.inference.chat_template_helpers import (
+ apply_chat_template_for_generation,
+ )
+
+ # Pick the chat-template-aware caller: processors that expose
+ # their own apply_chat_template + chat_template attr (e.g.
+ # Qwen2.5-VL) use it directly; otherwise fall back to the
+ # nested tokenizer.
+ chat_target = self._processor
if (
- chat_fn is None
+ getattr(self._processor, "apply_chat_template", None) is None
or not hasattr(self._processor, "chat_template")
or self._processor.chat_template is None
):
- tok = getattr(self._processor, "tokenizer", self._processor)
- chat_fn = tok.apply_chat_template
+ chat_target = getattr(self._processor, "tokenizer", self._processor)
- prompt = chat_fn(messages, tokenize = False, add_generation_prompt = True)
+ prompt = apply_chat_template_for_generation(
+ chat_target,
+ messages,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
+ )
# For VLM: always use mlx_vlm's stream_generate which handles
# pixel_values properly (passes None for text-only, image for VLM)
diff --git a/studio/backend/core/inference/orchestrator.py b/studio/backend/core/inference/orchestrator.py
index 5562820f49..7e7d7026f6 100644
--- a/studio/backend/core/inference/orchestrator.py
+++ b/studio/backend/core/inference/orchestrator.py
@@ -449,6 +449,10 @@ class InferenceOrchestrator:
repetition_penalty: float = 1.0,
cancel_event = None,
use_adapter = None,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
) -> Generator[str, None, None]:
"""Dispatched generation — sends command without holding _gen_lock.
@@ -494,6 +498,14 @@ class InferenceOrchestrator:
if use_adapter is not None:
cmd["use_adapter"] = use_adapter
+ if tools is not None:
+ cmd["tools"] = tools
+ if enable_thinking is not None:
+ cmd["enable_thinking"] = enable_thinking
+ if reasoning_effort is not None:
+ cmd["reasoning_effort"] = reasoning_effort
+ if preserve_thinking is not None:
+ cmd["preserve_thinking"] = preserve_thinking
# Create mailbox BEFORE sending command
mailbox: queue.Queue = queue.Queue()
@@ -695,6 +707,13 @@ class InferenceOrchestrator:
"audio_type": model_info.get("audio_type"),
"has_audio_input": model_info.get("has_audio_input", False),
}
+ # Mirror chat_template_info so routes can classify
+ # capabilities without re-entering the subprocess.
+ _tpl_info = model_info.get("chat_template_info")
+ if isinstance(_tpl_info, dict):
+ self.models[self.active_model_name]["chat_template_info"] = (
+ _tpl_info
+ )
self.loading_models.discard(model_name)
logger.info(
"Model '%s' loaded successfully in subprocess", model_name
@@ -770,8 +789,18 @@ class InferenceOrchestrator:
max_new_tokens: int = 256,
repetition_penalty: float = 1.0,
cancel_event = None,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
) -> Generator[str, None, None]:
- """Generate response, streaming tokens from subprocess."""
+ """Generate response, streaming tokens from subprocess.
+
+ Optional ``tools`` / ``enable_thinking`` / ``reasoning_effort`` /
+ ``preserve_thinking`` kwargs are forwarded into the worker so
+ ``tokenizer.apply_chat_template`` can render tool schemas and
+ reasoning controls when the template understands them.
+ """
yield from self._generate_inner(
messages = messages,
system_prompt = system_prompt,
@@ -784,6 +813,88 @@ class InferenceOrchestrator:
repetition_penalty = repetition_penalty,
cancel_event = cancel_event,
use_adapter = None,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
+ )
+
+ def generate_chat_completion_with_tools(
+ self,
+ messages: list,
+ tools: list,
+ system_prompt: str = "",
+ temperature: float = 0.7,
+ top_p: float = 0.9,
+ top_k: int = 40,
+ min_p: float = 0.0,
+ max_tokens: Optional[int] = None,
+ repetition_penalty: float = 1.0,
+ cancel_event = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
+ max_tool_iterations: int = 25,
+ auto_heal_tool_calls: bool = True,
+ tool_call_timeout: int = 300,
+ session_id: Optional[str] = None,
+ use_adapter: Optional[Union[bool, str]] = None,
+ **_unused,
+ ):
+ """Run the safetensors agentic tool loop in this (parent)
+ process, calling the worker for each generation turn.
+
+ Yields the same event dicts as the GGUF tool loop so the route
+ layer can stream both backends through one helper. See
+ ``safetensors_agentic.run_safetensors_tool_loop`` for the
+ event protocol.
+ """
+ from core.inference.safetensors_agentic import run_safetensors_tool_loop
+ from core.inference.tools import execute_tool
+
+ max_new_tokens = max_tokens if max_tokens and max_tokens > 0 else 2048
+
+ def _single_turn(conv: list):
+ # ``conv`` already carries any system message because the
+ # loop appends to a list seeded with system+user above.
+ common_kwargs = dict(
+ messages = conv,
+ system_prompt = "",
+ image = None,
+ temperature = temperature,
+ top_p = top_p,
+ top_k = top_k,
+ min_p = min_p,
+ max_new_tokens = max_new_tokens,
+ repetition_penalty = repetition_penalty,
+ cancel_event = cancel_event,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
+ )
+ if use_adapter is not None:
+ yield from self.generate_with_adapter_control(
+ use_adapter = use_adapter,
+ **common_kwargs,
+ )
+ else:
+ yield from self.generate_chat_response(**common_kwargs)
+
+ initial = list(messages)
+ if system_prompt:
+ initial = [{"role": "system", "content": system_prompt}] + initial
+
+ yield from run_safetensors_tool_loop(
+ single_turn = _single_turn,
+ messages = initial,
+ tools = tools,
+ execute_tool = execute_tool,
+ cancel_event = cancel_event,
+ auto_heal_tool_calls = auto_heal_tool_calls,
+ max_tool_iterations = max_tool_iterations,
+ tool_call_timeout = tool_call_timeout,
+ session_id = session_id,
)
def generate_with_adapter_control(
@@ -817,6 +928,10 @@ class InferenceOrchestrator:
repetition_penalty: float = 1.0,
cancel_event = None,
use_adapter = None,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
) -> Generator[str, None, None]:
"""Inner generation logic — sends command to subprocess, yields tokens.
@@ -853,6 +968,10 @@ class InferenceOrchestrator:
repetition_penalty = repetition_penalty,
cancel_event = cancel_event,
use_adapter = use_adapter,
+ tools = tools,
+ enable_thinking = enable_thinking,
+ reasoning_effort = reasoning_effort,
+ preserve_thinking = preserve_thinking,
)
def _generate_locked(
@@ -868,6 +987,10 @@ class InferenceOrchestrator:
repetition_penalty: float = 1.0,
cancel_event = None,
use_adapter = None,
+ tools: Optional[list] = None,
+ enable_thinking: Optional[bool] = None,
+ reasoning_effort: Optional[str] = None,
+ preserve_thinking: Optional[bool] = None,
) -> Generator[str, None, None]:
"""Actual generation logic — must be called under _gen_lock."""
request_id = str(uuid.uuid4())
@@ -893,6 +1016,16 @@ class InferenceOrchestrator:
if use_adapter is not None:
cmd["use_adapter"] = use_adapter
+ # Only forward template kwargs the caller actually set so older
+ # workers that ignore unknown keys still work.
+ if tools is not None:
+ cmd["tools"] = tools
+ if enable_thinking is not None:
+ cmd["enable_thinking"] = enable_thinking
+ if reasoning_effort is not None:
+ cmd["reasoning_effort"] = reasoning_effort
+ if preserve_thinking is not None:
+ cmd["preserve_thinking"] = preserve_thinking
try:
self._send_cmd(cmd)
@@ -1200,6 +1333,13 @@ class InferenceOrchestrator:
return self.models[self.active_model_name].get("is_vision", False)
return False
+ def _is_gpt_oss_model(self, model_name: str = None) -> bool:
+ """Parent-side gpt-oss detection so the safetensors route can run
+ the same guard without an IPC round-trip to the subprocess."""
+ from utils.datasets import is_gpt_oss_model_name
+
+ return is_gpt_oss_model_name(model_name or self.active_model_name or "")
+
# ========== GLOBAL INSTANCE ==========
_inference_backend = None
diff --git a/studio/backend/core/inference/safetensors_agentic.py b/studio/backend/core/inference/safetensors_agentic.py
new file mode 100644
index 0000000000..73bb3d090a
--- /dev/null
+++ b/studio/backend/core/inference/safetensors_agentic.py
@@ -0,0 +1,392 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""
+Safetensors/transformers agentic tool loop.
+
+Wraps a single-turn cumulative-text generator (the existing
+``InferenceOrchestrator.generate_chat_response`` pipeline that streams
+from a worker subprocess) with the tool-calling, thinking-block,
+status, and metadata event protocol used by the GGUF path. Keeps the
+front-end SSE shape identical across backends so the chat UI does not
+care which engine actually ran the model.
+
+The GGUF path lives in ``llama_cpp.py`` and talks to llama-server's
+structured ``delta.tool_calls`` directly. Native transformers has no
+such structured channel, so this loop parses tool calls from the
+cumulative text and dispatches them via ``core.inference.tools``.
+"""
+
+import json
+import threading
+from typing import Callable, Generator, Optional
+from urllib.parse import urlparse
+
+from loggers import get_logger
+
+from core.inference.tool_call_parser import (
+ BUDGET_EXHAUSTED_NUDGE,
+ DUPLICATE_CALL_NUDGE,
+ TOOL_ERROR_NUDGE,
+ TOOL_ERROR_PREFIXES,
+ TOOL_XML_SIGNALS,
+ has_tool_signal,
+ parse_tool_calls_from_text,
+ strip_tool_markup,
+)
+
+
+logger = get_logger(__name__)
+
+
+# Buffer cap while waiting to disambiguate a possible tool-call prefix.
+_MAX_BUFFER_CHARS = 32
+
+
+def _status_for_tool(tool_name: str, arguments: dict) -> str:
+ """Return a human-readable status line matching the GGUF path."""
+ if tool_name == "web_search":
+ url = (arguments.get("url") or "").strip()
+ if url:
+ parsed = urlparse(url)
+ if parsed.scheme in ("http", "https") and parsed.hostname:
+ host = parsed.hostname
+ if host.startswith("www."):
+ host = host[4:]
+ return f"Reading: {host}"
+ return "Reading page..."
+ query = arguments.get("query", "")
+ return f"Searching: {query}"
+ if tool_name == "python":
+ preview = (arguments.get("code") or "").strip().split("\n")[0][:60]
+ return f"Running Python: {preview}" if preview else "Running Python..."
+ if tool_name == "terminal":
+ preview = (arguments.get("command") or "")[:60]
+ return f"Running: {preview}" if preview else "Running command..."
+ return f"Calling: {tool_name}"
+
+
+_CANONICAL_HEAL_ARG = {"python": "code", "terminal": "command"}
+
+
+def _coerce_arguments(raw_args, *, heal: bool, tool_name: str = "") -> dict:
+ """Normalise tool ``arguments`` to a dict.
+
+ Some templates emit a JSON string, others a bare query string. With
+ ``heal=True`` we accept a bare string as ``{: ...}``
+ so a Hermes-style call without proper JSON still runs the tool. The
+ canonical key is picked per tool: ``code`` for python, ``command``
+ for terminal, ``query`` for everything else (e.g. web_search).
+ """
+ if isinstance(raw_args, dict):
+ return raw_args
+ if isinstance(raw_args, str):
+ try:
+ parsed = json.loads(raw_args)
+ if isinstance(parsed, dict):
+ return parsed
+ except (json.JSONDecodeError, ValueError):
+ pass
+ if heal:
+ key = _CANONICAL_HEAL_ARG.get(tool_name, "query")
+ return {key: raw_args}
+ return {"raw": raw_args}
+ return {}
+
+
+def run_safetensors_tool_loop(
+ *,
+ single_turn: Callable[[list], Generator[str, None, None]],
+ messages: list[dict],
+ tools: list[dict],
+ execute_tool: Callable[..., str],
+ cancel_event: Optional[threading.Event] = None,
+ auto_heal_tool_calls: bool = True,
+ max_tool_iterations: int = 25,
+ tool_call_timeout: int = 300,
+ session_id: Optional[str] = None,
+) -> Generator[dict, None, None]:
+ """Drive an agentic tool loop on top of a cumulative-text generator.
+
+ ``single_turn(messages)`` must yield cumulative assistant text
+ (each yield is a snapshot including all previously emitted tokens).
+ The loop:
+
+ * Buffers the leading characters of every turn so it can decide
+ whether the model is about to emit a tool call. Plain content
+ starts streaming as soon as the buffer rules it out.
+ * On detecting ```` or ``= 0 and (signal_pos < 0 or p < signal_pos):
+ signal_pos = p
+ if signal_pos >= 0:
+ before_tool = candidate[:signal_pos]
+ cleaned_before = strip_tool_markup(before_tool)
+ if len(cleaned_before) > len(last_emitted):
+ last_emitted = cleaned_before
+ yield {"type": "content", "text": cleaned_before}
+ cumulative_display = candidate
+ detect_state = _state_draining
+ continue
+ cumulative_display = candidate
+ cleaned = strip_tool_markup(cumulative_display)
+ if len(cleaned) > len(last_emitted):
+ last_emitted = cleaned
+ yield {"type": "content", "text": cleaned}
+ continue
+
+ # BUFFERING: hold until we know it is not a tool call.
+ content_buffer += delta
+ stripped = content_buffer.lstrip()
+ if not stripped:
+ continue
+
+ is_match = False
+ is_prefix = False
+ for sig in TOOL_XML_SIGNALS:
+ if stripped.startswith(sig):
+ is_match = True
+ break
+ if sig.startswith(stripped):
+ is_prefix = True
+ break
+
+ if is_match:
+ detect_state = _state_draining
+ elif is_prefix and len(stripped) < _MAX_BUFFER_CHARS:
+ continue
+ else:
+ detect_state = _state_streaming
+ cumulative_display += content_buffer
+ cleaned = strip_tool_markup(cumulative_display)
+ if len(cleaned) > len(last_emitted):
+ last_emitted = cleaned
+ yield {"type": "content", "text": cleaned}
+
+ # Stream finished -- resolve what we collected.
+ if cancel_event is not None and cancel_event.is_set():
+ return
+
+ if detect_state == _state_buffering:
+ # Buffer never resolved -- tool XML or plain content.
+ stripped = content_buffer.lstrip()
+ if stripped and has_tool_signal(stripped):
+ detect_state = _state_draining
+ else:
+ if content_buffer:
+ cumulative_display += content_buffer
+ yield {
+ "type": "content",
+ "text": strip_tool_markup(cumulative_display, final = True),
+ }
+ yield {"type": "status", "text": ""}
+ return
+
+ if detect_state == _state_streaming:
+ # No tool detected mid-stream -- check for late tool XML.
+ safety_tc = None
+ if has_tool_signal(content_accum):
+ safety_tc = parse_tool_calls_from_text(
+ content_accum,
+ id_offset = next_call_id,
+ )
+ if not safety_tc:
+ # Final answer: streaming already emitted content.
+ # Skip a final=True re-strip so literal ""
+ # in prose survives when no real tool call parsed.
+ yield {"type": "status", "text": ""}
+ return
+ tool_calls = safety_tc
+ content_text = strip_tool_markup(content_accum, final = True)
+ logger.info(
+ "Safetensors safety net: parsed %d tool call(s) from streamed content",
+ len(tool_calls),
+ )
+ else:
+ # DRAINING: parse tool calls out of full content.
+ tool_calls = parse_tool_calls_from_text(
+ content_accum,
+ id_offset = next_call_id,
+ )
+ if not tool_calls and auto_heal_tool_calls:
+ # Parser found nothing -- surface raw content so any
+ # literal "" prose is preserved.
+ if content_accum:
+ yield {"type": "content", "text": content_accum}
+ yield {"type": "status", "text": ""}
+ return
+ content_text = strip_tool_markup(content_accum, final = True)
+
+ if final_attempt_done:
+ # Final-answer turn re-called a tool -- stop the loop.
+ if content_text:
+ yield {"type": "content", "text": content_text}
+ yield {"type": "status", "text": ""}
+ return
+
+ assistant_msg: dict = {"role": "assistant", "content": content_text}
+ if tool_calls:
+ assistant_msg["tool_calls"] = tool_calls
+ next_call_id += len(tool_calls)
+ conversation.append(assistant_msg)
+
+ for tc in tool_calls or []:
+ func = tc.get("function", {}) or {}
+ tool_name = func.get("name", "") or ""
+ arguments = _coerce_arguments(
+ func.get("arguments", {}),
+ heal = auto_heal_tool_calls,
+ tool_name = tool_name,
+ )
+
+ yield {"type": "status", "text": _status_for_tool(tool_name, arguments)}
+ yield {
+ "type": "tool_start",
+ "tool_name": tool_name,
+ "tool_call_id": tc.get("id", ""),
+ "arguments": arguments,
+ }
+
+ tc_key = tool_name + str(arguments)
+ if allowed_tool_names and tool_name not in allowed_tool_names:
+ result = (
+ f"Error: tool '{tool_name}' is not enabled for this "
+ "request. Use one of the enabled tools or provide a "
+ "final answer."
+ )
+ else:
+ already_ran_ok = any(
+ k == tc_key and not err for k, err in tool_call_history
+ )
+ if already_ran_ok:
+ result = DUPLICATE_CALL_NUDGE
+ else:
+ eff_timeout = (
+ None if tool_call_timeout >= 9999 else tool_call_timeout
+ )
+ try:
+ result = execute_tool(
+ tool_name,
+ arguments,
+ cancel_event = cancel_event,
+ timeout = eff_timeout,
+ session_id = session_id,
+ )
+ except Exception as exc:
+ logger.exception("Tool %s raised: %s", tool_name, exc)
+ result = f"Error: tool raised an exception: {exc}"
+
+ yield {
+ "type": "tool_end",
+ "tool_name": tool_name,
+ "tool_call_id": tc.get("id", ""),
+ "result": result,
+ }
+
+ is_error = isinstance(result, str) and result.lstrip().startswith(
+ TOOL_ERROR_PREFIXES
+ )
+ tool_call_history.append((tc_key, is_error))
+
+ # Strip frontend image sentinel from the model's view.
+ # Cut at the first occurrence so leading and consecutive
+ # sentinels are both removed.
+ result_for_model = result
+ if isinstance(result_for_model, str) and "__IMAGES__:" in result_for_model:
+ result_for_model = result_for_model.split("__IMAGES__:", 1)[0].rstrip()
+ if is_error:
+ result_for_model = result_for_model + TOOL_ERROR_NUDGE
+
+ tool_msg: dict = {
+ "role": "tool",
+ "name": tool_name,
+ "content": result_for_model,
+ }
+ tool_call_id = tc.get("id")
+ if tool_call_id:
+ tool_msg["tool_call_id"] = tool_call_id
+ conversation.append(tool_msg)
+
+ # Clear the status badge before the next turn.
+ yield {"type": "status", "text": ""}
+
+ if iteration + 1 >= max_tool_iterations and not final_attempt_done:
+ # Budget exhausted; nudge a final plain answer.
+ final_attempt_done = True
+ conversation.append(
+ {
+ "role": "user",
+ "content": BUDGET_EXHAUSTED_NUDGE,
+ }
+ )
+
+ yield {"type": "status", "text": ""}
diff --git a/studio/backend/core/inference/tool_call_parser.py b/studio/backend/core/inference/tool_call_parser.py
new file mode 100644
index 0000000000..a0ab8a2a53
--- /dev/null
+++ b/studio/backend/core/inference/tool_call_parser.py
@@ -0,0 +1,204 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""
+Backend-neutral tool-call XML parser shared by GGUF and safetensors.
+Tolerates missing closing tags in either ``{json}``
+or ``v...`` shape.
+"""
+
+import json
+import re
+
+
+# _TOOL_CLOSED_PATS: closed pairs only. _TOOL_ALL_PATS: also trailing
+# unclosed runs so truncated tails don't leak markup.
+_TOOL_CLOSED_PATS = [
+ re.compile(r".*?", re.DOTALL),
+ re.compile(r".*?", re.DOTALL),
+]
+_TOOL_ALL_PATS = _TOOL_CLOSED_PATS + [
+ re.compile(r".*$", re.DOTALL),
+ re.compile(r".*$", re.DOTALL),
+]
+
+
+# Prefixes the streaming buffer watches for to gate in-progress text.
+TOOL_XML_SIGNALS = ("", "\s*\{")
+_TC_FUNC_START_RE = re.compile(r"\s*")
+_TC_END_TAG_RE = re.compile(r"")
+_TC_FUNC_CLOSE_RE = re.compile(r"\s*\s*$")
+_TC_PARAM_START_RE = re.compile(r"\s*")
+_TC_PARAM_CLOSE_RE = re.compile(r"\s*\s*$")
+
+
+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) -> 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 ```` tags:
+ ``{"name":"web_search","arguments":{"query":"..."}}``
+ - XML-style function blocks:
+ ``v``
+
+ Closing tags (````, ````, ````)
+ are all optional since models frequently omit them.
+ """
+ tool_calls: list[dict] = []
+
+ # Pattern 1: {json}. Balanced-brace scan that skips
+ # braces inside JSON strings.
+ 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
+ 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_{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: v... -- closing tags
+ # optional; don't use as body boundary because code
+ # values can contain that literal.
+ if not tool_calls:
+ func_starts = list(_TC_FUNC_START_RE.finditer(content))
+ 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]
+ 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
+ # embedded in code strings is preserved.
+ pm = param_starts[0]
+ val = body[pm.end() :]
+ val = _TC_PARAM_CLOSE_RE.sub("", val)
+ arguments[pm.group(1)] = val.strip()
+ else:
+ 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]
+ val = _TC_PARAM_CLOSE_RE.sub("", val)
+ arguments[param_name] = val.strip()
+
+ tc = {
+ "id": f"call_{id_offset + len(tool_calls)}",
+ "type": "function",
+ "function": {
+ "name": func_name,
+ "arguments": json.dumps(arguments),
+ },
+ }
+ tool_calls.append(tc)
+
+ return tool_calls
+
+
+def has_tool_signal(text: str) -> bool:
+ """Return True if ``text`` contains any tool-call XML signal."""
+ return any(s in text for s in TOOL_XML_SIGNALS)
diff --git a/studio/backend/core/inference/worker.py b/studio/backend/core/inference/worker.py
index cacede2d3e..20a7d2d16c 100644
--- a/studio/backend/core/inference/worker.py
+++ b/studio/backend/core/inference/worker.py
@@ -346,6 +346,26 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
"audio_type": getattr(mc, "audio_type", None),
"has_audio_input": getattr(mc, "has_audio_input", False),
}
+ # Forward chat_template_info so the parent can classify
+ # capabilities without re-entering the subprocess.
+ try:
+ _bm = getattr(backend, "models", {}) or {}
+ _entry = (
+ _bm.get(mc.identifier)
+ or _bm.get(getattr(backend, "active_model_name", None))
+ or {}
+ )
+ _tpl_info = _entry.get("chat_template_info")
+ if isinstance(_tpl_info, dict):
+ model_info["chat_template_info"] = {
+ "has_template": bool(_tpl_info.get("has_template", False)),
+ "template": _tpl_info.get("template"),
+ "format_type": _tpl_info.get("format_type", "generic"),
+ "template_name": _tpl_info.get("template_name"),
+ "special_tokens": _tpl_info.get("special_tokens", {}) or {},
+ }
+ except Exception as _tpl_exc:
+ logger.warning("chat_template_info forward failed: %s", _tpl_exc)
_send_response(
resp_queue,
{
@@ -416,6 +436,18 @@ def _handle_generate(
"cancel_event": cancel_event,
}
+ # Optional template/tool plumbing: only forward keys that are
+ # actually present so the backend signature can evolve without
+ # breaking older command payloads.
+ for opt_key in (
+ "tools",
+ "enable_thinking",
+ "reasoning_effort",
+ "preserve_thinking",
+ ):
+ if opt_key in cmd:
+ gen_kwargs[opt_key] = cmd[opt_key]
+
# Choose generation path
use_adapter = cmd.get("use_adapter")
if use_adapter is not None:
@@ -648,36 +680,6 @@ def run_inference_process(
os.environ["HF_HUB_DISABLE_XET"] = "1"
logger.info("Xet transport disabled (HF_HUB_DISABLE_XET=1)")
- # Offline auto-detect: skip 25s of hf_hub_download retries per file
- # if DNS is dead; cached files resolve instantly under HF_HUB_OFFLINE=1.
- # Scope is this subprocess only -- orchestrator spawns a fresh worker
- # per load (see core/inference/orchestrator.py), so the env cannot
- # persist across loads.
- if "HF_HUB_OFFLINE" not in os.environ:
- import socket as _socket
- import threading as _threading
-
- # Probe on a daemon thread so concurrent sockets in the parent
- # interpreter are not affected by socket.setdefaulttimeout.
- _result: list = [None]
-
- def _probe() -> None:
- try:
- _socket.gethostbyname("huggingface.co")
- _result[0] = False
- except Exception:
- _result[0] = True
-
- _t = _threading.Thread(target = _probe, daemon = True)
- _t.start()
- _t.join(2.0)
- if _result[0] is None or _result[0] is True:
- os.environ["HF_HUB_OFFLINE"] = "1"
- os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
- logger.warning(
- "huggingface.co unreachable; HF_HUB_OFFLINE=1 set for this worker."
- )
-
import warnings
from loggers.config import LogConfig
diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py
index 2f095cb823..1b4e7051b0 100644
--- a/studio/backend/routes/inference.py
+++ b/studio/backend/routes/inference.py
@@ -235,6 +235,57 @@ router = APIRouter()
studio_router = APIRouter()
+def _detect_safetensors_features(backend, chat_template: Optional[str]) -> dict:
+ """Classify reasoning/tool capabilities via the GGUF classifier so
+ flags match across backends. gpt-oss is overridden because Harmony
+ routes reasoning and tools through tokenizer channels, not template
+ markup."""
+ model_id = getattr(backend, "active_model_name", None)
+ flags = (
+ detect_reasoning_flags(
+ chat_template,
+ model_identifier = model_id,
+ log_source = "safetensors",
+ )
+ if chat_template
+ else {
+ "supports_reasoning": False,
+ "reasoning_style": "enable_thinking",
+ "reasoning_always_on": False,
+ "supports_preserve_thinking": False,
+ "supports_tools": False,
+ }
+ )
+ # Our safetensors loop only parses {json}
+ # and .... Llama uses <|python_tag|>,
+ # Mistral uses [TOOL_CALLS]; advertising tools for those would
+ # enable a pill the parser cannot honour. GGUF is unaffected --
+ # llama-server normalises every format into structured deltas.
+ if (
+ flags.get("supports_tools")
+ and chat_template
+ and "" not in chat_template
+ and " XML this loop parses).
+ try:
+ if hasattr(backend, "_is_gpt_oss_model") and backend._is_gpt_oss_model():
+ flags["supports_reasoning"] = True
+ flags["reasoning_style"] = "reasoning_effort"
+ flags["supports_tools"] = False
+ except Exception:
+ logger.debug("gpt_oss_check_failed", exc_info = True)
+ return flags
+
+
def _effective_enable_tools(payload) -> Optional[bool]:
"""Resolve `payload.enable_tools` against the process-level tool policy.
@@ -590,6 +641,7 @@ async def load_model(
reasoning_style = llama_backend.reasoning_style,
reasoning_always_on = llama_backend.reasoning_always_on,
supports_preserve_thinking = llama_backend.supports_preserve_thinking,
+ supports_tools = llama_backend.supports_tools,
chat_template = llama_backend.chat_template,
speculative_type = llama_backend.requested_spec_mode,
spec_draft_n_max = llama_backend.spec_draft_n_max,
@@ -612,21 +664,10 @@ async def load_model(
logger.warning(
f"Could not retrieve chat template for {backend.active_model_name}: {e}"
)
- # Non-GGUF: only advertise reasoning for gpt-oss Harmony,
- # which emits reasoning via channels at the tokenizer level.
- # Template-level chat_template_kwargs (enable_thinking /
- # preserve_thinking / tools) are not yet forwarded through
- # the transformers generation path, so avoid advertising
- # controls the server cannot honour outside GGUF.
- _sf_supports_reasoning = False
- _sf_reasoning_style = "enable_thinking"
- if hasattr(backend, "_is_gpt_oss_model"):
- try:
- if backend._is_gpt_oss_model():
- _sf_supports_reasoning = True
- _sf_reasoning_style = "reasoning_effort"
- except Exception:
- pass
+ # Classify via the same path as GGUF.
+ _sf_flags = _detect_safetensors_features(backend, _chat_template)
+ _sf_supports_reasoning = _sf_flags["supports_reasoning"]
+ _sf_reasoning_style = _sf_flags["reasoning_style"]
return LoadResponse(
status = "already_loaded",
model = model_log_label
@@ -647,9 +688,9 @@ async def load_model(
),
supports_reasoning = _sf_supports_reasoning,
reasoning_style = _sf_reasoning_style,
- reasoning_always_on = False,
- supports_preserve_thinking = False,
- supports_tools = False,
+ reasoning_always_on = _sf_flags["reasoning_always_on"],
+ supports_preserve_thinking = _sf_flags["supports_preserve_thinking"],
+ supports_tools = _sf_flags["supports_tools"],
chat_template = _chat_template,
)
@@ -982,19 +1023,8 @@ async def load_model(
except Exception:
pass
- # Non-GGUF: gpt-oss Harmony surfaces reasoning via tokenizer-level
- # channels; other safetensors reasoning/tools/preserve-thinking
- # knobs are not forwarded to tokenizer.apply_chat_template yet, so
- # we only advertise support for the Harmony case here.
- _sf_supports_reasoning = False
- _sf_reasoning_style = "enable_thinking"
- if hasattr(backend, "_is_gpt_oss_model"):
- try:
- if backend._is_gpt_oss_model():
- _sf_supports_reasoning = True
- _sf_reasoning_style = "reasoning_effort"
- except Exception:
- pass
+ # Classify reasoning/tool flags via the GGUF sniffer.
+ _sf_flags = _detect_safetensors_features(backend, _chat_template)
return LoadResponse(
status = "loaded",
@@ -1012,11 +1042,11 @@ async def load_model(
requires_trust_remote_code = bool(
inference_config.get("trust_remote_code", False)
),
- supports_reasoning = _sf_supports_reasoning,
- reasoning_style = _sf_reasoning_style,
- reasoning_always_on = False,
- supports_preserve_thinking = False,
- supports_tools = False,
+ supports_reasoning = _sf_flags["supports_reasoning"],
+ reasoning_style = _sf_flags["reasoning_style"],
+ reasoning_always_on = _sf_flags["reasoning_always_on"],
+ supports_preserve_thinking = _sf_flags["supports_preserve_thinking"],
+ supports_tools = _sf_flags["supports_tools"],
chat_template = _chat_template,
)
@@ -1388,18 +1418,8 @@ async def get_status(
else None
)
- # Non-GGUF: only gpt-oss Harmony is wired through the transformers
- # generation path. Other template-level reasoning / tool kwargs
- # are not yet forwarded, so we do not advertise them here.
- supports_reasoning = False
- reasoning_style = "enable_thinking"
- if backend.active_model_name and hasattr(backend, "_is_gpt_oss_model"):
- try:
- if backend._is_gpt_oss_model():
- supports_reasoning = True
- reasoning_style = "reasoning_effort"
- except Exception:
- pass
+ # Non-GGUF: classify from the loaded template.
+ _sf_flags = _detect_safetensors_features(backend, chat_template)
inference_config = (
load_inference_config(backend.active_model_name)
if backend.active_model_name
@@ -1419,11 +1439,11 @@ async def get_status(
requires_trust_remote_code = bool(
(inference_config or {}).get("trust_remote_code", False)
),
- supports_reasoning = supports_reasoning,
- reasoning_style = reasoning_style,
- reasoning_always_on = False,
- supports_preserve_thinking = False,
- supports_tools = False,
+ supports_reasoning = _sf_flags["supports_reasoning"],
+ reasoning_style = _sf_flags["reasoning_style"],
+ reasoning_always_on = _sf_flags["reasoning_always_on"],
+ supports_preserve_thinking = _sf_flags["supports_preserve_thinking"],
+ supports_tools = _sf_flags["supports_tools"],
chat_template = chat_template,
llama_cpp_supports_mtp = _supports_mtp,
llama_cpp_prebuilt_stale = _stale,
@@ -2749,6 +2769,300 @@ async def openai_chat_completions(
except Exception as e:
raise HTTPException(status_code = 400, detail = f"Failed to decode image: {e}")
+ # Classify capability flags from the loaded template.
+ _sf_model_info = backend.models.get(backend.active_model_name, {})
+ _sf_tpl = (_sf_model_info.get("chat_template_info") or {}).get("template")
+ _sf_features = _detect_safetensors_features(backend, _sf_tpl)
+
+ cancel_event = threading.Event()
+ completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
+ created = int(time.time())
+
+ # ── Safetensors tool-calling path ─────────────────────────
+ # Mirrors the GGUF agentic loop's event shape. Disabled for
+ # vision turns (untested overlap with image render slot) and
+ # for gpt-oss (Harmony uses dedicated channels, not
+ # XML -- gpt-oss tools still work via the GGUF path).
+ _sf_is_gptoss = False
+ try:
+ _sf_is_gptoss = bool(
+ hasattr(backend, "_is_gpt_oss_model") and backend._is_gpt_oss_model()
+ )
+ except Exception:
+ _sf_is_gptoss = False
+
+ _sf_tool_budget = (
+ payload.max_tool_calls_per_message
+ if payload.max_tool_calls_per_message is not None
+ else 25
+ )
+
+ _sf_use_tools = (
+ _effective_enable_tools(payload)
+ and _sf_features.get("supports_tools", False)
+ and image is None
+ and not _sf_is_gptoss
+ and _sf_tool_budget > 0
+ )
+
+ if _sf_use_tools:
+ from core.inference.tools import ALL_TOOLS
+
+ if payload.enabled_tools is not None:
+ _sf_tools_to_use = [
+ t for t in ALL_TOOLS if t["function"]["name"] in payload.enabled_tools
+ ]
+ else:
+ _sf_tools_to_use = ALL_TOOLS
+
+ _sf_tool_names = {t["function"]["name"] for t in _sf_tools_to_use}
+ _sf_has_web = "web_search" in _sf_tool_names
+ _sf_has_code = "python" in _sf_tool_names or "terminal" in _sf_tool_names
+
+ _sf_date_line = f"The current date is {_date.today().isoformat()}."
+ _sf_model_size_b = _extract_model_size_b(model_name)
+ _sf_is_small_model = _sf_model_size_b is not None and _sf_model_size_b < 9
+
+ if _sf_is_small_model:
+ _sf_web_tips = "Do not repeat the same search query."
+ else:
+ _sf_web_tips = (
+ "When you search and find a relevant URL in the results, "
+ "fetch its full content by calling web_search with the url parameter. "
+ "Do not repeat the same search query. If a search returns "
+ "no useful results, try rephrasing or fetching a result URL directly."
+ )
+ _sf_code_tips = (
+ "Use code execution for math, calculations, data processing, "
+ "or to parse and analyze information from tool results."
+ )
+
+ if _sf_has_web and _sf_has_code:
+ _sf_nudge = (
+ _sf_date_line + " "
+ "You have access to tools. When appropriate, prefer using "
+ "tools rather than answering from memory. "
+ + _sf_web_tips
+ + " "
+ + _sf_code_tips
+ )
+ elif _sf_has_code:
+ _sf_nudge = (
+ _sf_date_line + " "
+ "You have access to tools. When appropriate, prefer using "
+ "code execution rather than answering from memory. " + _sf_code_tips
+ )
+ elif _sf_has_web:
+ _sf_nudge = (
+ _sf_date_line + " "
+ "You have access to tools. When appropriate, prefer using "
+ "web search for up-to-date or uncertain factual "
+ "information rather than answering from memory. " + _sf_web_tips
+ )
+ else:
+ _sf_nudge = ""
+
+ _sf_system_prompt = system_prompt
+ if _sf_nudge:
+ _sf_nudge += _TOOL_ACTION_NUDGE
+ if _sf_system_prompt:
+ _sf_system_prompt = _sf_system_prompt.rstrip() + "\n\n" + _sf_nudge
+ else:
+ _sf_system_prompt = _sf_nudge
+
+ # Strip stale tool-call XML from prior assistant turns.
+ _sf_chat_messages = []
+ for _msg in chat_messages:
+ if _msg.get("role") == "assistant" and isinstance(_msg.get("content"), str):
+ _sf_chat_messages.append(
+ {
+ **_msg,
+ "content": _TOOL_XML_RE.sub("", _msg["content"]).strip(),
+ }
+ )
+ else:
+ _sf_chat_messages.append(_msg)
+
+ def sf_generate_with_tools():
+ return backend.generate_chat_completion_with_tools(
+ messages = _sf_chat_messages,
+ tools = _sf_tools_to_use,
+ system_prompt = _sf_system_prompt or "",
+ temperature = payload.temperature,
+ top_p = payload.top_p,
+ top_k = payload.top_k,
+ min_p = payload.min_p,
+ max_tokens = payload.max_tokens,
+ repetition_penalty = payload.repetition_penalty,
+ cancel_event = cancel_event,
+ enable_thinking = payload.enable_thinking,
+ reasoning_effort = payload.reasoning_effort,
+ preserve_thinking = payload.preserve_thinking,
+ auto_heal_tool_calls = payload.auto_heal_tool_calls
+ if payload.auto_heal_tool_calls is not None
+ else True,
+ max_tool_iterations = _sf_tool_budget,
+ tool_call_timeout = payload.tool_call_timeout
+ if payload.tool_call_timeout is not None
+ else 300,
+ session_id = payload.session_id,
+ use_adapter = payload.use_adapter,
+ )
+
+ _sf_tool_sentinel = object()
+ _sf_cancel_keys = (payload.cancel_id, payload.session_id, completion_id)
+ _sf_tracker = _TrackedCancel(cancel_event, *_sf_cancel_keys)
+ _sf_tracker.__enter__()
+
+ async def sf_tool_stream():
+ try:
+ first_chunk = ChatCompletionChunk(
+ id = completion_id,
+ created = created,
+ model = model_name,
+ choices = [
+ ChunkChoice(
+ delta = ChoiceDelta(role = "assistant"),
+ finish_reason = None,
+ )
+ ],
+ )
+ yield f"data: {first_chunk.model_dump_json(exclude_none = True)}\n\n"
+
+ gen = sf_generate_with_tools()
+ prev_text = ""
+ while True:
+ if cancel_event.is_set():
+ backend.reset_generation_state()
+ break
+ if await request.is_disconnected():
+ cancel_event.set()
+ backend.reset_generation_state()
+ return
+
+ event = await asyncio.to_thread(next, gen, _sf_tool_sentinel)
+ if event is _sf_tool_sentinel:
+ break
+
+ if event["type"] == "status":
+ if not event["text"]:
+ prev_text = ""
+ status_data = json.dumps(
+ {
+ "type": "tool_status",
+ "content": event["text"],
+ }
+ )
+ yield f"data: {status_data}\n\n"
+ continue
+
+ if event["type"] in ("tool_start", "tool_end"):
+ if event["type"] == "tool_start":
+ prev_text = ""
+ yield f"data: {json.dumps(event)}\n\n"
+ continue
+
+ # Diff cumulative cleaned text against last snapshot.
+ raw_cumulative = event.get("text", "")
+ clean_cumulative = _TOOL_XML_RE.sub("", raw_cumulative)
+ new_text = clean_cumulative[len(prev_text) :]
+ prev_text = clean_cumulative
+ if not new_text:
+ continue
+ chunk = ChatCompletionChunk(
+ id = completion_id,
+ created = created,
+ model = model_name,
+ choices = [
+ ChunkChoice(
+ delta = ChoiceDelta(content = new_text),
+ finish_reason = None,
+ )
+ ],
+ )
+ yield f"data: {chunk.model_dump_json(exclude_none = True)}\n\n"
+
+ final_chunk = ChatCompletionChunk(
+ id = completion_id,
+ created = created,
+ model = model_name,
+ choices = [
+ ChunkChoice(
+ delta = ChoiceDelta(),
+ finish_reason = "stop",
+ )
+ ],
+ )
+ yield f"data: {final_chunk.model_dump_json(exclude_none = True)}\n\n"
+ yield "data: [DONE]\n\n"
+
+ except asyncio.CancelledError:
+ cancel_event.set()
+ backend.reset_generation_state()
+ raise
+ except Exception:
+ backend.reset_generation_state()
+ # Generic wire message; full trace stays in the log
+ # (CWE-209: transformers/torch errors may leak paths).
+ logger.exception("safetensors tool stream error")
+ error_chunk = {
+ "error": {
+ "message": "An internal error occurred.",
+ "type": "server_error",
+ },
+ }
+ yield f"data: {json.dumps(error_chunk)}\n\n"
+ finally:
+ _sf_tracker.__exit__(None, None, None)
+
+ if payload.stream:
+ return StreamingResponse(
+ sf_tool_stream(),
+ media_type = "text/event-stream",
+ headers = {
+ "Cache-Control": "no-cache",
+ "Connection": "keep-alive",
+ "X-Accel-Buffering": "no",
+ },
+ )
+
+ # Non-streaming JSON: drain the loop, build one ChatCompletion.
+ try:
+
+ def _drain_to_text():
+ full_text = ""
+ gen = sf_generate_with_tools()
+ for event in gen:
+ if cancel_event.is_set():
+ break
+ if event.get("type") == "content":
+ full_text = _TOOL_XML_RE.sub("", event.get("text", ""))
+ return full_text
+
+ content_text = await asyncio.to_thread(_drain_to_text)
+ response = ChatCompletion(
+ id = completion_id,
+ created = created,
+ model = model_name,
+ choices = [
+ CompletionChoice(
+ message = CompletionMessage(content = content_text),
+ finish_reason = "stop",
+ )
+ ],
+ )
+ return JSONResponse(content = response.model_dump())
+ except Exception:
+ backend.reset_generation_state()
+ # CWE-209: generic detail; full trace in log.
+ logger.exception("safetensors tool completion error")
+ raise HTTPException(
+ status_code = 500,
+ detail = "An internal error occurred.",
+ )
+ finally:
+ _sf_tracker.__exit__(None, None, None)
+
# Shared generation kwargs
gen_kwargs = dict(
messages = chat_messages,
@@ -2761,9 +3075,14 @@ async def openai_chat_completions(
max_new_tokens = payload.max_tokens or 2048,
repetition_penalty = payload.repetition_penalty,
)
-
- # Choose generation path (adapter-controlled or standard)
- cancel_event = threading.Event()
+ # Forward reasoning kwargs; the worker/template wrapper peels off
+ # any the template doesn't accept.
+ if payload.enable_thinking is not None:
+ gen_kwargs["enable_thinking"] = payload.enable_thinking
+ if payload.reasoning_effort is not None:
+ gen_kwargs["reasoning_effort"] = payload.reasoning_effort
+ if payload.preserve_thinking is not None:
+ gen_kwargs["preserve_thinking"] = payload.preserve_thinking
if payload.use_adapter is not None:
@@ -2780,9 +3099,6 @@ async def openai_chat_completions(
cancel_event = cancel_event, **gen_kwargs
)
- completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
- created = int(time.time())
-
# ── Streaming response ────────────────────────────────────────
if payload.stream:
_cancel_keys = (payload.cancel_id, payload.session_id, completion_id)
diff --git a/studio/backend/tests/test_mlx_inference_backend.py b/studio/backend/tests/test_mlx_inference_backend.py
index ce447bdd1f..16cca7dd40 100644
--- a/studio/backend/tests/test_mlx_inference_backend.py
+++ b/studio/backend/tests/test_mlx_inference_backend.py
@@ -158,3 +158,97 @@ def test_mlx_inference_vlm_lora_uses_unsloth_loader_without_native_adapter_rewri
assert backend._is_vlm is True
assert isinstance(backend._processor, _DummyProcessor)
assert isinstance(backend._tokenizer, _DummyTokenizer)
+
+
+# Regression: MLXInferenceBackend.generate_chat_response must accept the
+# four template kwargs (tools / enable_thinking / reasoning_effort /
+# preserve_thinking) so the route layer can forward what the user
+# toggled in the UI. The previous signature raised
+# "got an unexpected keyword argument 'tools'" on Mac.
+
+
+def test_mlx_generate_chat_response_accepts_template_kwargs():
+ import inspect
+ from core.inference.mlx_inference import MLXInferenceBackend
+
+ sig = inspect.signature(MLXInferenceBackend.generate_chat_response)
+ params = sig.parameters
+ for name in ("tools", "enable_thinking", "reasoning_effort", "preserve_thinking"):
+ assert name in params, (
+ f"MLX.generate_chat_response is missing the {name!r} kwarg; "
+ "the route layer forwards this and a missing kwarg raises "
+ "TypeError on Mac"
+ )
+ assert (
+ params[name].default is None
+ ), f"{name!r} must default to None so existing callers stay valid"
+
+
+def test_mlx_generate_text_forwards_kwargs_into_template_helper(monkeypatch):
+ """The Mac text path must route through apply_chat_template_for_
+ generation so reasoning / tool kwargs reach the tokenizer."""
+ _install_fake_mlx(monkeypatch)
+ from core.inference.mlx_inference import MLXInferenceBackend
+
+ captured = {}
+
+ def _fake_apply(tokenizer, messages, **kwargs):
+ captured["tokenizer"] = tokenizer
+ captured["messages"] = messages
+ captured["kwargs"] = kwargs
+ return ""
+
+ monkeypatch.setattr(
+ "core.inference.chat_template_helpers." "apply_chat_template_for_generation",
+ _fake_apply,
+ raising = True,
+ )
+
+ # mlx_lm.stream_generate yields response objects with .token; make a
+ # one-token generator so _generate_text returns without touching the
+ # real stack.
+ import types as _types
+
+ mlx_lm_pkg = _types.ModuleType("mlx_lm")
+ mlx_lm_sample = _types.ModuleType("mlx_lm.sample_utils")
+ mlx_lm_sample.make_sampler = lambda **_kw: object()
+ mlx_lm_sample.make_logits_processors = lambda **_kw: None
+
+ class _Resp:
+ def __init__(self, tok):
+ self.token = tok
+
+ def _stream_generate(_model, _tokenizer, **_kw):
+ yield _Resp(1)
+
+ mlx_lm_pkg.stream_generate = _stream_generate
+ monkeypatch.setitem(sys.modules, "mlx_lm", mlx_lm_pkg)
+ monkeypatch.setitem(sys.modules, "mlx_lm.sample_utils", mlx_lm_sample)
+
+ class _Tok:
+ chat_template = "x"
+
+ def decode(self, ids, skip_special_tokens = False):
+ return "hi"
+
+ backend = MLXInferenceBackend()
+ backend._model = object()
+ backend._tokenizer = _Tok()
+ backend._is_vlm = False
+
+ out = list(
+ backend.generate_chat_response(
+ messages = [{"role": "user", "content": "ping"}],
+ tools = [{"function": {"name": "web_search"}}],
+ enable_thinking = True,
+ reasoning_effort = "medium",
+ preserve_thinking = True,
+ max_new_tokens = 1,
+ )
+ )
+ assert out == ["hi"]
+ # The kwargs the user toggled must reach the chat-template helper.
+ assert captured["kwargs"]["tools"] == [{"function": {"name": "web_search"}}]
+ assert captured["kwargs"]["enable_thinking"] is True
+ assert captured["kwargs"]["reasoning_effort"] == "medium"
+ assert captured["kwargs"]["preserve_thinking"] is True
diff --git a/studio/backend/tests/test_safetensors_capability_advertise.py b/studio/backend/tests/test_safetensors_capability_advertise.py
new file mode 100644
index 0000000000..c3ee5b9ff1
--- /dev/null
+++ b/studio/backend/tests/test_safetensors_capability_advertise.py
@@ -0,0 +1,451 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""
+Capability advertisement contract: classifier honesty, worker→
+orchestrator IPC hop, and route-layer end-to-end. Pure helpers + fakes;
+no torch / transformers import.
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock
+
+_backend_root = Path(__file__).resolve().parent.parent
+if str(_backend_root) not in sys.path:
+ sys.path.insert(0, str(_backend_root))
+
+
+# Qwen3 snippet covering tools, enable_thinking, preserve_thinking.
+QWEN3_TEMPLATE = """
+{%- if tools %}
+ {{- '<|im_start|>system\\nFor each function call, return a json object'
+ ' wrapped inside tags.\\n' }}
+ {%- for tool in tools %}
+ {{- tool | tojson }}
+ {%- endfor %}
+{%- endif %}
+{%- for message in messages %}
+ {%- if message.role == 'tool' %}
+ {{- '<|im_start|>tool\\n' + message.content + '<|im_end|>\\n' }}
+ {%- endif %}
+{%- endfor %}
+{%- if enable_thinking is defined and enable_thinking %}
+ {{- '' }}
+{%- endif %}
+{%- if preserve_thinking %}
+ {{- assistant.reasoning_content }}
+{%- endif %}
+"""
+
+
+GPT_OSS_TEMPLATE = """
+<|start|>system<|message|>You are gpt-oss.
+reasoning_effort: {{ reasoning_effort }}
+<|end|>
+"""
+
+
+PLAIN_TEMPLATE = """
+{%- for message in messages %}
+ {{- message.role + ': ' + message.content + '\\n' }}
+{%- endfor %}
+"""
+
+
+# ── Tests: classifier honesty ────────────────────────────────────────
+
+
+def test_detect_reasoning_flags_qwen3_supports_tools_and_reasoning():
+ from core.inference.llama_cpp import detect_reasoning_flags
+
+ flags = detect_reasoning_flags(QWEN3_TEMPLATE, "unsloth/Qwen3-0.6B")
+ assert flags["supports_tools"] is True
+ assert flags["supports_reasoning"] is True
+ assert flags["reasoning_style"] == "enable_thinking"
+ assert flags["supports_preserve_thinking"] is True
+ assert flags["reasoning_always_on"] is False
+
+
+def test_detect_reasoning_flags_plain_template_all_false():
+ from core.inference.llama_cpp import detect_reasoning_flags
+
+ flags = detect_reasoning_flags(PLAIN_TEMPLATE, "some/PlainChat")
+ assert flags["supports_tools"] is False
+ assert flags["supports_reasoning"] is False
+ assert flags["supports_preserve_thinking"] is False
+ assert flags["reasoning_always_on"] is False
+
+
+def test_detect_reasoning_flags_none_template_returns_all_false():
+ from core.inference.llama_cpp import detect_reasoning_flags
+
+ flags = detect_reasoning_flags(None)
+ assert flags["supports_tools"] is False
+ assert flags["supports_reasoning"] is False
+ assert flags["supports_preserve_thinking"] is False
+ assert flags["reasoning_always_on"] is False
+ assert flags["reasoning_style"] == "enable_thinking"
+
+
+def test_detect_safetensors_features_passes_template_through_to_classifier():
+ """Route wrapper forwards a real template to the inner classifier."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(active_model_name = "unsloth/Qwen3-0.6B")
+ flags = _detect_safetensors_features(backend, QWEN3_TEMPLATE)
+ assert flags["supports_tools"] is True
+ assert flags["supports_reasoning"] is True
+
+
+def test_detect_safetensors_features_none_template_returns_all_false():
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(active_model_name = "unsloth/Qwen3-0.6B")
+ flags = _detect_safetensors_features(backend, None)
+ assert flags == {
+ "supports_reasoning": False,
+ "reasoning_style": "enable_thinking",
+ "reasoning_always_on": False,
+ "supports_preserve_thinking": False,
+ "supports_tools": False,
+ }
+
+
+def test_detect_safetensors_features_gptoss_disables_tools():
+ """gpt-oss Harmony: tools intentionally off even if template marks it."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = MagicMock()
+ backend.active_model_name = "unsloth/gpt-oss-20b"
+ backend._is_gpt_oss_model.return_value = True
+
+ flags = _detect_safetensors_features(backend, QWEN3_TEMPLATE)
+ assert flags["supports_reasoning"] is True
+ assert flags["reasoning_style"] == "reasoning_effort"
+ assert flags["supports_tools"] is False
+
+
+# Llama-3 / Mistral templates advertise tool handling but the model emits
+# tool calls in <|python_tag|> / [TOOL_CALLS] format -- not the
+# / system<|end_header_id|>' }}
+ {{- 'You have access to the following tools.' }}
+ {%- for tool in tools %}
+ {{- tool | tojson }}
+ {%- endfor %}
+{%- endif %}
+{%- for message in messages %}
+ {%- if message.role == 'tool' %}
+ {{- '<|start_header_id|>ipython<|end_header_id|>' }}
+ {{- '<|python_tag|>' }}
+ {{- message.content }}
+ {%- endif %}
+{%- endfor %}
+"""
+
+MISTRAL_TEMPLATE = """
+{%- if tools %}
+ {%- for tool in tools %}
+ {{- tool | tojson }}
+ {%- endfor %}
+{%- endif %}
+{%- for message in messages %}
+ {%- if message.role == 'tool' %}
+ {{- '[TOOL_CALLS]' + message.content + '[/TOOL_CALLS]' }}
+ {%- endif %}
+{%- endfor %}
+"""
+
+
+def test_detect_safetensors_features_llama3_template_suppresses_tools():
+ """Llama-3 emits <|python_tag|>; safetensors loop cannot parse it."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(active_model_name = "unsloth/Llama-3.2-3B-Instruct")
+ flags = _detect_safetensors_features(backend, LLAMA3_TEMPLATE)
+ assert flags["supports_tools"] is False
+
+
+def test_detect_safetensors_features_mistral_template_suppresses_tools():
+ """Mistral emits [TOOL_CALLS]; safetensors loop cannot parse it."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(active_model_name = "unsloth/mistral-7b-instruct-v0.3")
+ flags = _detect_safetensors_features(backend, MISTRAL_TEMPLATE)
+ assert flags["supports_tools"] is False
+
+
+def test_detect_safetensors_features_qwen_tool_call_keeps_tools_on():
+ """Sanity check: gate only suppresses non-Qwen formats."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(active_model_name = "unsloth/Qwen3-0.6B")
+ flags = _detect_safetensors_features(backend, QWEN3_TEMPLATE)
+ assert flags["supports_tools"] is True
+
+
+def test_detect_safetensors_features_function_xml_format_keeps_tools_on():
+ """Templates emitting XML are parser-compatible."""
+ from routes.inference import _detect_safetensors_features
+
+ tpl_with_function_xml = (
+ "{%- if tools %}<|im_start|>system\n"
+ "Tool call format: v"
+ "<|im_end|>{%- endif %}"
+ )
+ backend = SimpleNamespace(active_model_name = "custom/with-function-xml")
+ flags = _detect_safetensors_features(backend, tpl_with_function_xml)
+ assert flags["supports_tools"] is True
+
+
+# Qwen3.5 family pins -- the live GGUF + safetensors templates fetched
+# from the unsloth/Qwen3.5-0.8B(-GGUF) repos both wrap tool calls as
+# ``\n...``. Capture a faithful slice so the
+# classifier never silently regresses for this family.
+
+QWEN35_TOOL_INSTRUCTION = (
+ "{%- if tools %}\n"
+ " <|im_start|>system\n"
+ " # Tools\n"
+ " \n"
+ " {%- for tool in tools %}{{ tool | tojson }}{%- endfor %}\n"
+ " \n"
+ " If you choose to call a function ONLY reply in the following format:\n"
+ " \n"
+ " \n"
+ " \n"
+ " value_1\n"
+ " \n"
+ " \n"
+ " \n"
+ " <|im_end|>\n"
+ "{%- endif %}\n"
+ "{%- if enable_thinking is defined and enable_thinking %}{{- '' }}{%- endif %}\n"
+)
+
+
+def test_detect_safetensors_features_qwen35_keeps_tools_on():
+ """unsloth/Qwen3.5-0.8B family must surface tools+reasoning enabled."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(active_model_name = "unsloth/Qwen3.5-0.8B")
+ flags = _detect_safetensors_features(backend, QWEN35_TOOL_INSTRUCTION)
+ assert flags["supports_tools"] is True
+ assert flags["supports_reasoning"] is True
+ assert flags["reasoning_style"] == "enable_thinking"
+
+
+# ── Tests: IPC bridge contract ───────────────────────────────────────
+
+
+def test_orchestrator_mirrors_chat_template_info_into_models_dict():
+ """Worker → orchestrator must copy chat_template_info verbatim."""
+ from core.inference.orchestrator import InferenceOrchestrator
+
+ orch = InferenceOrchestrator.__new__(InferenceOrchestrator)
+ orch.models = {}
+ orch.active_model_name = None
+ orch.loading_models = set()
+
+ model_info = {
+ "identifier": "unsloth/Qwen3-0.6B",
+ "display_name": "Qwen3-0.6B",
+ "is_vision": False,
+ "is_lora": False,
+ "is_gguf": False,
+ "is_audio": False,
+ "audio_type": None,
+ "has_audio_input": False,
+ "chat_template_info": {
+ "has_template": True,
+ "template": QWEN3_TEMPLATE,
+ "format_type": "chatml",
+ "template_name": "qwen3",
+ "special_tokens": {"bos_token": "<|im_start|>"},
+ },
+ }
+
+ # Replay orchestrator.load_model's mirror block verbatim.
+ orch.active_model_name = model_info["identifier"]
+ orch.models[orch.active_model_name] = {
+ "is_vision": model_info.get("is_vision", False),
+ "is_lora": model_info.get("is_lora", False),
+ "display_name": model_info.get("display_name", "x"),
+ "is_audio": model_info.get("is_audio", False),
+ "audio_type": model_info.get("audio_type"),
+ "has_audio_input": model_info.get("has_audio_input", False),
+ }
+ _tpl_info = model_info.get("chat_template_info")
+ if isinstance(_tpl_info, dict):
+ orch.models[orch.active_model_name]["chat_template_info"] = _tpl_info
+
+ entry = orch.models[orch.active_model_name]
+ tpl = entry.get("chat_template_info", {}).get("template")
+ assert tpl == QWEN3_TEMPLATE
+
+ from routes.inference import _detect_safetensors_features
+
+ flags = _detect_safetensors_features(
+ SimpleNamespace(active_model_name = orch.active_model_name), tpl
+ )
+ assert flags["supports_tools"] is True
+ assert flags["supports_reasoning"] is True
+
+
+def test_orchestrator_missing_chat_template_info_falls_back_to_all_false():
+ """Old / malformed worker reply: no crash, all flags False."""
+ from core.inference.orchestrator import InferenceOrchestrator
+ from routes.inference import _detect_safetensors_features
+
+ orch = InferenceOrchestrator.__new__(InferenceOrchestrator)
+ orch.models = {}
+ orch.active_model_name = "unsloth/Qwen3-0.6B"
+
+ model_info = {
+ "identifier": "unsloth/Qwen3-0.6B",
+ "is_vision": False,
+ "is_lora": False,
+ # NB: no chat_template_info key
+ }
+ orch.models[orch.active_model_name] = {
+ "is_vision": False,
+ "is_lora": False,
+ }
+ _tpl_info = model_info.get("chat_template_info")
+ if isinstance(_tpl_info, dict):
+ orch.models[orch.active_model_name]["chat_template_info"] = _tpl_info
+
+ entry = orch.models[orch.active_model_name]
+ tpl = entry.get("chat_template_info", {}).get("template")
+ assert tpl is None
+
+ flags = _detect_safetensors_features(
+ SimpleNamespace(active_model_name = orch.active_model_name), tpl
+ )
+ assert flags["supports_tools"] is False
+
+
+def test_worker_load_reply_payload_includes_chat_template_info():
+ """Worker IPC reply carries chat_template_info dict."""
+
+ class _StubBackend:
+ def __init__(self, identifier, template):
+ self.active_model_name = identifier
+ self.models = {
+ identifier: {
+ "chat_template_info": {
+ "has_template": True,
+ "template": template,
+ "format_type": "chatml",
+ "template_name": "qwen3",
+ "special_tokens": {"bos_token": "<|im_start|>"},
+ }
+ }
+ }
+
+ backend = _StubBackend("unsloth/Qwen3-0.6B", QWEN3_TEMPLATE)
+ mc = SimpleNamespace(
+ identifier = "unsloth/Qwen3-0.6B",
+ display_name = "Qwen3-0.6B",
+ is_vision = False,
+ is_lora = False,
+ )
+
+ # Replay the worker's payload-build block.
+ model_info = {
+ "identifier": mc.identifier,
+ "display_name": mc.display_name,
+ "is_vision": mc.is_vision,
+ "is_lora": mc.is_lora,
+ "is_gguf": False,
+ }
+ _bm = getattr(backend, "models", {}) or {}
+ _entry = (
+ _bm.get(mc.identifier)
+ or _bm.get(getattr(backend, "active_model_name", None))
+ or {}
+ )
+ _tpl_info = _entry.get("chat_template_info")
+ if isinstance(_tpl_info, dict):
+ model_info["chat_template_info"] = {
+ "has_template": bool(_tpl_info.get("has_template", False)),
+ "template": _tpl_info.get("template"),
+ "format_type": _tpl_info.get("format_type", "generic"),
+ "template_name": _tpl_info.get("template_name"),
+ "special_tokens": _tpl_info.get("special_tokens", {}) or {},
+ }
+
+ assert "chat_template_info" in model_info
+ assert model_info["chat_template_info"]["template"] == QWEN3_TEMPLATE
+ assert model_info["chat_template_info"]["has_template"] is True
+
+
+def test_worker_load_reply_payload_survives_missing_template():
+ """Tokenizer with no chat_template still produces a valid reply."""
+
+ class _StubBackend:
+ def __init__(self):
+ self.active_model_name = "legacy/no-template"
+ self.models = {"legacy/no-template": {}} # no chat_template_info
+
+ backend = _StubBackend()
+ mc = SimpleNamespace(
+ identifier = "legacy/no-template",
+ display_name = "legacy",
+ is_vision = False,
+ is_lora = False,
+ )
+
+ model_info = {
+ "identifier": mc.identifier,
+ "display_name": mc.display_name,
+ "is_vision": mc.is_vision,
+ "is_lora": mc.is_lora,
+ "is_gguf": False,
+ }
+ _bm = getattr(backend, "models", {}) or {}
+ _entry = _bm.get(mc.identifier) or {}
+ _tpl_info = _entry.get("chat_template_info")
+ if isinstance(_tpl_info, dict):
+ model_info["chat_template_info"] = dict(_tpl_info)
+
+ assert "chat_template_info" not in model_info
+
+
+# ── End-to-end: route layer sees the template, advertises True ───────
+
+
+def test_route_layer_emits_supports_tools_true_for_qwen3_safetensors():
+ """End-to-end: Qwen3 safetensors flips supports_tools=True."""
+ from routes.inference import _detect_safetensors_features
+
+ backend = SimpleNamespace(
+ active_model_name = "unsloth/Qwen3-0.6B",
+ models = {
+ "unsloth/Qwen3-0.6B": {
+ "is_vision": False,
+ "chat_template_info": {
+ "has_template": True,
+ "template": QWEN3_TEMPLATE,
+ "format_type": "chatml",
+ },
+ }
+ },
+ )
+
+ _model_info = backend.models.get(backend.active_model_name, {})
+ _tpl = _model_info.get("chat_template_info", {}).get("template")
+ flags = _detect_safetensors_features(backend, _tpl)
+
+ assert flags["supports_tools"] is True
+ assert flags["supports_reasoning"] is True
+ assert flags["supports_preserve_thinking"] is True
diff --git a/studio/backend/tests/test_safetensors_tool_loop.py b/studio/backend/tests/test_safetensors_tool_loop.py
new file mode 100644
index 0000000000..923af87c4f
--- /dev/null
+++ b/studio/backend/tests/test_safetensors_tool_loop.py
@@ -0,0 +1,788 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""
+Tests for the safetensors agentic tool loop.
+
+Covers the shared ``tool_call_parser`` helpers and the cumulative-text
+state machine inside ``safetensors_agentic.run_safetensors_tool_loop``.
+The loop is exercised with hand-crafted fake single-turn generators so
+no model load is needed; the tests run in CI under a few seconds.
+
+Edge cases under coverage:
+* Plain answers (no tool calls) flush full content.
+* Single ``{json}`` triggers the tool and re-enters.
+* Single ``...`` XML form triggers the same path.
+* Truncated unclosed ```` is still parsed.
+* Tool result is fed back as ``role=tool`` for the next iteration.
+* Bad JSON inside ```` does not raise and (when healed) is
+ routed as a ``{"query": ...}`` web search call.
+* Duplicate tool calls produce a synthetic "do not repeat" result the
+ second time.
+* ``__IMAGES__`` sentinel is stripped before the model sees the result.
+* Tool execution errors are tagged so the model gets a nudge but the
+ loop keeps streaming.
+* Cancel is honoured between iterations.
+* ``max_tool_iterations`` cap is respected and a final-answer attempt
+ closes the stream cleanly.
+"""
+
+import threading
+
+import pytest
+
+from core.inference import safetensors_agentic
+from core.inference.safetensors_agentic import (
+ _coerce_arguments,
+ run_safetensors_tool_loop,
+)
+from core.inference.tool_call_parser import (
+ has_tool_signal,
+ parse_tool_calls_from_text,
+ strip_tool_markup,
+)
+from utils.datasets import is_gpt_oss_model_name
+
+
+# ────────────────────────────────────────────────────────────────────
+# parse_tool_calls_from_text
+# ────────────────────────────────────────────────────────────────────
+
+
+class TestParser:
+ def test_json_tool_call(self):
+ text = (
+ '{"name":"web_search","arguments":{"query":"hello"}}'
+ )
+ result = parse_tool_calls_from_text(text)
+ assert len(result) == 1
+ tc = result[0]
+ assert tc["type"] == "function"
+ assert tc["function"]["name"] == "web_search"
+ # Arguments must always be a JSON string.
+ assert isinstance(tc["function"]["arguments"], str)
+ assert "hello" in tc["function"]["arguments"]
+
+ def test_json_tool_call_unclosed(self):
+ # No ; balanced-brace extractor must still close.
+ text = '{"name":"python","arguments":{"code":"print(1)"}}'
+ result = parse_tool_calls_from_text(text)
+ assert len(result) == 1
+ assert result[0]["function"]["name"] == "python"
+
+ def test_xml_function_call(self):
+ text = "print('hi')"
+ result = parse_tool_calls_from_text(text)
+ assert len(result) == 1
+ assert result[0]["function"]["name"] == "python"
+ assert "print('hi')" in result[0]["function"]["arguments"]
+
+ def test_xml_unclosed(self):
+ # Closing tags omitted; parser must still extract the value.
+ text = "ls -la"
+ result = parse_tool_calls_from_text(text)
+ assert len(result) == 1
+ assert result[0]["function"]["name"] == "terminal"
+ assert "ls -la" in result[0]["function"]["arguments"]
+
+ def test_code_with_embedded_xml(self):
+ # A code parameter contains the literal . Must not
+ # truncate the value because the parser uses end-of-body as the
+ # only boundary for single-parameter calls.
+ text = (
+ "html = ''\n"
+ "print('hi')"
+ )
+ result = parse_tool_calls_from_text(text)
+ assert len(result) == 1
+ assert "print('hi')" in result[0]["function"]["arguments"]
+
+ def test_multiple_calls(self):
+ text = (
+ '{"name":"web_search","arguments":{"query":"a"}}'
+ '{"name":"web_search","arguments":{"query":"b"}}'
+ )
+ result = parse_tool_calls_from_text(text)
+ assert len(result) == 2
+ assert result[0]["function"]["name"] == "web_search"
+ assert result[1]["function"]["name"] == "web_search"
+
+ def test_bad_json_does_not_raise(self):
+ text = "{not valid json}"
+ result = parse_tool_calls_from_text(text)
+ # Bad JSON is silently dropped; caller can fall back to text.
+ assert result == []
+
+ def test_has_tool_signal(self):
+ assert has_tool_signal("blah x")
+ assert has_tool_signal("hi ...")
+ assert not has_tool_signal("hello world")
+
+ def test_strip_markup_closed(self):
+ text = "before {} after"
+ assert strip_tool_markup(text) == "before after"
+
+ def test_strip_markup_unclosed_final(self):
+ text = "before {partial"
+ # With final=True the trailing run is dropped.
+ assert strip_tool_markup(text, final = True) == "before"
+ # Without final=True the unclosed run is preserved.
+ assert "partial" in strip_tool_markup(text)
+
+
+# ────────────────────────────────────────────────────────────────────
+# run_safetensors_tool_loop
+# ────────────────────────────────────────────────────────────────────
+
+
+def _fake_stream(chunks):
+ """Build a single-turn generator that yields cumulative snapshots."""
+
+ def _gen(_messages):
+ acc = ""
+ for c in chunks:
+ acc += c
+ yield acc
+
+ return _gen
+
+
+def _const_stream(text):
+ """A single-turn generator that yields one cumulative snapshot."""
+
+ def _gen(_messages):
+ yield text
+
+ return _gen
+
+
+class FakeExecuteTool:
+ """Stand-in for ``core.inference.tools.execute_tool``."""
+
+ def __init__(self, results):
+ # ``results`` is a list of strings or RuntimeError instances.
+ self.results = list(results)
+ self.calls: list[tuple[str, dict]] = []
+
+ def __call__(
+ self,
+ name,
+ arguments,
+ *,
+ cancel_event = None,
+ timeout = None,
+ session_id = None,
+ ):
+ self.calls.append((name, arguments))
+ result = self.results.pop(0) if self.results else "OK"
+ if isinstance(result, Exception):
+ raise result
+ return result
+
+
+def _collect_events(generator, max_events = 200):
+ events = []
+ for ev in generator:
+ events.append(ev)
+ if len(events) >= max_events:
+ break
+ return events
+
+
+def _make_loop(*, turns, exec_results = None, **kwargs):
+ """Build a configured loop with a multi-turn fake generator.
+
+ ``turns`` is a list of chunk-lists; iteration N yields chunks from
+ ``turns[N]``.
+ """
+ turn_iter = iter(turns)
+
+ def _gen(_messages):
+ try:
+ chunks = next(turn_iter)
+ except StopIteration:
+ return
+ acc = ""
+ for c in chunks:
+ acc += c
+ yield acc
+
+ exec_fn = FakeExecuteTool(exec_results or [])
+ return run_safetensors_tool_loop(
+ single_turn = _gen,
+ messages = [{"role": "user", "content": "hi"}],
+ tools = [
+ {"type": "function", "function": {"name": "web_search"}},
+ {"type": "function", "function": {"name": "python"}},
+ {"type": "function", "function": {"name": "terminal"}},
+ ],
+ execute_tool = exec_fn,
+ **kwargs,
+ ), exec_fn
+
+
+class TestLoopBasic:
+ def test_plain_answer(self):
+ # No tool XML; loop should yield content then status="".
+ loop, _exec = _make_loop(
+ turns = [["Hello", " world", "!"]],
+ exec_results = [],
+ )
+ events = _collect_events(loop)
+ contents = [e for e in events if e["type"] == "content"]
+ statuses = [e for e in events if e["type"] == "status"]
+ assert contents, "expected at least one content event"
+ # Final cumulative content should contain the answer.
+ final_text = contents[-1]["text"]
+ assert "Hello world!" in final_text
+ assert statuses and statuses[-1]["text"] == ""
+
+ def test_single_tool_then_answer(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ # : tool call only.
+ [
+ '{"name":"web_search",',
+ '"arguments":{"query":"weather"}}',
+ "",
+ ],
+ # : final answer.
+ ["The ", "weather is ", "sunny."],
+ ],
+ exec_results = ["Sunny and 22C"],
+ )
+ events = _collect_events(loop)
+ kinds = [e["type"] for e in events]
+
+ assert "tool_start" in kinds
+ assert "tool_end" in kinds
+ # Tool was actually called with the parsed arguments.
+ assert exec_fn.calls == [("web_search", {"query": "weather"})]
+
+ tool_start = next(e for e in events if e["type"] == "tool_start")
+ assert tool_start["tool_name"] == "web_search"
+ tool_end = next(e for e in events if e["type"] == "tool_end")
+ assert tool_end["result"] == "Sunny and 22C"
+
+ contents = [e for e in events if e["type"] == "content"]
+ assert contents and "sunny" in contents[-1]["text"].lower()
+
+ def test_function_xml_form(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ ["print(1)"],
+ ["Result: 1"],
+ ],
+ exec_results = ["1\n"],
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls == [("python", {"code": "print(1)"})]
+ contents = [e for e in events if e["type"] == "content"]
+ assert "Result: 1" in contents[-1]["text"]
+
+ def test_truncated_unclosed_tool_call(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ # No ; balanced-brace parser must still
+ # succeed because the JSON itself is balanced.
+ ['{"name":"web_search","arguments":{"query":"x"}}'],
+ ["done"],
+ ],
+ exec_results = ["result"],
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls == [("web_search", {"query": "x"})]
+
+ def test_bad_json_healed_to_query(self):
+ # Tool call with non-JSON string arguments. With auto_heal_tool_calls
+ # the string is routed as {"query": ...}.
+ loop, exec_fn = _make_loop(
+ turns = [
+ # JSON inside the tool call is well-formed; the
+ # ``arguments`` is a string that is not itself valid
+ # JSON for ``_coerce_arguments`` to parse, so the
+ # heal path runs.
+ [
+ '{"name":"web_search","arguments":"hello world"}'
+ ],
+ ["ok"],
+ ],
+ exec_results = ["..."],
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls and exec_fn.calls[0][0] == "web_search"
+ assert exec_fn.calls[0][1] == {"query": "hello world"}
+
+
+class TestLoopBehaviour:
+ def test_duplicate_tool_call_synthetic_result(self):
+ # Two identical successful calls in a row: the second is short-
+ # circuited with a "do not repeat" message and execute_tool is
+ # called only once.
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ ["final"],
+ ],
+ exec_results = ["search-result-1"],
+ )
+ events = _collect_events(loop)
+ # Only one real call.
+ assert len(exec_fn.calls) == 1
+ tool_end_events = [e for e in events if e["type"] == "tool_end"]
+ assert len(tool_end_events) == 2
+ assert "do not repeat" in tool_end_events[1]["result"].lower()
+
+ def test_image_sentinel_stripped_from_model_feed(self):
+ # The tool result has a frontend image sentinel that should be
+ # stripped before being fed back into the next turn, BUT the
+ # tool_end event still carries the raw result for the UI.
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"python","arguments":{"code":"plot()"}}'
+ ],
+ ["see chart"],
+ ],
+ exec_results = ["chart\n__IMAGES__:/tmp/chart.png"],
+ )
+ events = _collect_events(loop)
+ tool_end = next(e for e in events if e["type"] == "tool_end")
+ assert "__IMAGES__" in tool_end["result"]
+
+ def test_image_sentinel_stripped_with_leading_marker(self):
+ # Sentinel at start (no newline) must not leak to the model.
+ from core.inference import safetensors_agentic as _sa
+
+ captured: list[list[dict]] = []
+
+ def fake_single_turn(messages, **_kw):
+ captured.append([dict(m) for m in messages])
+ if len(captured) == 1:
+ yield '{"name":"python","arguments":{"code":"plot()"}}'
+ else:
+ yield "done"
+
+ events = list(
+ _sa.run_safetensors_tool_loop(
+ single_turn = fake_single_turn,
+ messages = [{"role": "user", "content": "plot please"}],
+ tools = [{"function": {"name": "python"}}],
+ execute_tool = lambda *_a, **_kw: "__IMAGES__:/tmp/x.png",
+ cancel_event = threading.Event(),
+ max_tool_iterations = 3,
+ auto_heal_tool_calls = True,
+ )
+ )
+ # Model's second turn must not see "__IMAGES__".
+ assert len(captured) >= 2
+ tool_msgs = [m for m in captured[1] if m.get("role") == "tool"]
+ assert tool_msgs, "no tool message reached the model"
+ for tm in tool_msgs:
+ assert (
+ "__IMAGES__" not in tm["content"]
+ ), f"sentinel leaked to model: {tm['content']!r}"
+
+ def test_image_sentinel_stripped_with_multiple_markers(self):
+ # Consecutive sentinels: cut at the first, nothing leaks.
+ from core.inference import safetensors_agentic as _sa
+
+ captured: list[list[dict]] = []
+
+ def fake_single_turn(messages, **_kw):
+ captured.append([dict(m) for m in messages])
+ if len(captured) == 1:
+ yield '{"name":"python","arguments":{"code":"plot()"}}'
+ else:
+ yield "done"
+
+ multi = "panel\n__IMAGES__:/tmp/a.png\n__IMAGES__:/tmp/b.png"
+ events = list(
+ _sa.run_safetensors_tool_loop(
+ single_turn = fake_single_turn,
+ messages = [{"role": "user", "content": "plot please"}],
+ tools = [{"function": {"name": "python"}}],
+ execute_tool = lambda *_a, **_kw: multi,
+ cancel_event = threading.Event(),
+ max_tool_iterations = 3,
+ auto_heal_tool_calls = True,
+ )
+ )
+ tool_msgs = [m for m in captured[1] if m.get("role") == "tool"]
+ assert tool_msgs
+ for tm in tool_msgs:
+ assert (
+ "__IMAGES__" not in tm["content"]
+ ), f"second sentinel leaked: {tm['content']!r}"
+ assert (
+ tm["content"] == "panel"
+ ), f"expected payload-only 'panel', got {tm['content']!r}"
+
+ def test_tool_execution_error_is_emitted_but_loop_continues(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ ["sorry, that failed"],
+ ],
+ exec_results = ["Error: network unreachable"],
+ )
+ events = _collect_events(loop)
+ tool_end = next(e for e in events if e["type"] == "tool_end")
+ assert tool_end["result"].startswith("Error")
+ # The loop must still produce a content event after the failure.
+ contents = [e for e in events if e["type"] == "content"]
+ assert contents
+
+ def test_exception_in_executor_does_not_raise(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ ["recovered"],
+ ],
+ exec_results = [RuntimeError("boom")],
+ )
+ events = _collect_events(loop)
+ tool_end = next(e for e in events if e["type"] == "tool_end")
+ assert "boom" in tool_end["result"]
+
+
+class TestLoopControl:
+ def test_cancel_event_breaks_loop(self):
+ cancel = threading.Event()
+ cancel.set()
+ # Even with a fake stream that emits tool calls, the loop must
+ # bail before invoking execute_tool when cancel is set.
+ exec_fn = FakeExecuteTool([])
+ events = list(
+ run_safetensors_tool_loop(
+ single_turn = _const_stream(
+ '{"name":"web_search",'
+ '"arguments":{"query":"x"}}'
+ ),
+ messages = [{"role": "user", "content": "hi"}],
+ tools = [],
+ execute_tool = exec_fn,
+ cancel_event = cancel,
+ )
+ )
+ assert events == []
+ assert exec_fn.calls == []
+
+ def test_max_iterations_caps_loop(self):
+ # The loop should stop after max_tool_iterations even if the
+ # model keeps asking for tools, then emit a final-attempt round.
+ loop, exec_fn = _make_loop(
+ turns = [
+ # : tool call (executes once)
+ [
+ '{"name":"web_search","arguments":{"query":"a"}}'
+ ],
+ # : model gives a final answer when nudged.
+ ["here is the final answer"],
+ ],
+ exec_results = ["result"],
+ max_tool_iterations = 1,
+ )
+ events = _collect_events(loop)
+ contents = [e for e in events if e["type"] == "content"]
+ # Final content must include the final answer.
+ assert contents and "final answer" in contents[-1]["text"]
+
+
+class TestStatusFormatting:
+ def test_status_for_known_tools(self):
+ # Use the private helper directly to verify status formatting.
+ assert (
+ safetensors_agentic._status_for_tool("web_search", {"query": "abc"})
+ == "Searching: abc"
+ )
+ assert (
+ safetensors_agentic._status_for_tool(
+ "web_search", {"url": "https://www.example.com/x"}
+ )
+ == "Reading: example.com"
+ )
+ assert safetensors_agentic._status_for_tool(
+ "python", {"code": "x = 1"}
+ ).startswith("Running Python:")
+ assert safetensors_agentic._status_for_tool(
+ "terminal", {"command": "ls"}
+ ).startswith("Running:")
+ assert safetensors_agentic._status_for_tool("unknown_tool", {}).startswith(
+ "Calling:"
+ )
+
+
+class TestProseMentioningToolCall:
+ def test_assistant_prose_with_literal_tool_call_text_survives(self):
+ # Regression: if the assistant text legitimately mentions
+ # ```` as a literal string and the parser finds no
+ # actual call, the loop must surface the full content instead
+ # of silently stripping everything past the literal marker.
+ loop, exec_fn = _make_loop(
+ turns = [
+ # : a real tool call so the loop moves to
+ # .
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ # : prose that mentions the literal text.
+ ["the docs say means an LLM tool call wrapper"],
+ ],
+ exec_results = ["result"],
+ )
+ events = _collect_events(loop)
+ contents = [e for e in events if e["type"] == "content"]
+ assert contents, "expected at least one content event"
+ final = contents[-1]["text"]
+ assert (
+ "LLM tool" in final
+ ), f"prose mentioning should not be truncated; got {final!r}"
+
+ def test_tool_result_with_tool_call_text_does_not_retrigger(self):
+ # Tool result text contains the literal ```` string.
+ # The loop must only parse the MODEL output, not the tool
+ # result, so we should see exactly one call.
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ ["the docs mention wrappers"],
+ ],
+ exec_results = ["Page text: appears here in the docs"],
+ )
+ events = _collect_events(loop)
+ assert len(exec_fn.calls) == 1
+
+
+class TestChatTemplateHelper:
+ """Cover the dependency-light helper used by InferenceBackend."""
+
+ def setup_method(self):
+ from core.inference.chat_template_helpers import (
+ apply_chat_template_for_generation,
+ )
+
+ self.apply = apply_chat_template_for_generation
+
+ class _Tok:
+ def __init__(self, accepted):
+ self.accepted = accepted
+ self.call_count = 0
+ self.last_kwargs = None
+
+ def apply_chat_template(
+ self, messages, *, tokenize = False, add_generation_prompt = True, **kw
+ ):
+ self.call_count += 1
+ unknown = set(kw) - self.accepted
+ if unknown:
+ raise TypeError(f"unexpected kwargs: {sorted(unknown)}")
+ self.last_kwargs = dict(kw)
+ return "PROMPT"
+
+ def test_richest_call_wins_when_template_supports_all(self):
+ tok = self._Tok({"tools", "enable_thinking"})
+ self.apply(tok, [], tools = [{}], enable_thinking = True)
+ assert tok.call_count == 1
+ assert "tools" in tok.last_kwargs
+ assert "enable_thinking" in tok.last_kwargs
+
+ def test_falls_back_when_template_rejects_reasoning_kwarg(self):
+ tok = self._Tok({"tools"})
+ self.apply(tok, [], tools = [{}], enable_thinking = True)
+ assert tok.call_count >= 2
+ assert tok.last_kwargs == {"tools": [{}]}
+
+ def test_falls_back_to_bare_call(self):
+ tok = self._Tok(set())
+ self.apply(tok, [], tools = [{}], enable_thinking = True)
+ assert tok.last_kwargs == {}
+
+ def test_jinja_error_propagates(self):
+ class Boom:
+ def apply_chat_template(self, *a, **kw):
+ raise ValueError("jinja: missing var")
+
+ with pytest.raises(ValueError):
+ self.apply(Boom(), [])
+
+ def test_no_kwargs_single_call(self):
+ tok = self._Tok(set())
+ self.apply(tok, [])
+ assert tok.call_count == 1
+
+
+# ────────────────────────────────────────────────────────────────────
+# Guardrails (allowlist, budget, streaming-leak, dedup, id offset,
+# auto_heal=False, canonical healed-arg key)
+# ────────────────────────────────────────────────────────────────────
+
+
+class TestGuardrails:
+ def test_disabled_tool_is_not_executed(self):
+ exec_fn = FakeExecuteTool([])
+ loop = run_safetensors_tool_loop(
+ single_turn = _fake_stream(
+ [
+ '{"name":"terminal","arguments":{"command":"echo bypass"}}'
+ ]
+ ),
+ messages = [{"role": "user", "content": "hi"}],
+ tools = [{"type": "function", "function": {"name": "web_search"}}],
+ execute_tool = exec_fn,
+ max_tool_iterations = 2,
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls == []
+ tool_ends = [e for e in events if e["type"] == "tool_end"]
+ assert tool_ends and "not enabled" in tool_ends[0]["result"].lower()
+
+ def test_empty_tools_list_does_not_enforce_allowlist(self):
+ exec_fn = FakeExecuteTool(["OK"])
+ loop = run_safetensors_tool_loop(
+ single_turn = _fake_stream(
+ [
+ '{"name":"python","arguments":{"code":"print(1)"}}'
+ ]
+ ),
+ messages = [{"role": "user", "content": "hi"}],
+ tools = [],
+ execute_tool = exec_fn,
+ max_tool_iterations = 2,
+ )
+ _collect_events(loop)
+ assert exec_fn.calls == [("python", {"code": "print(1)"})]
+
+ def test_max_iterations_zero_executes_no_tools(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ]
+ ],
+ exec_results = ["OK"],
+ max_tool_iterations = 0,
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls == []
+ assert events and events[-1] == {"type": "status", "text": ""}
+
+ def test_streaming_clips_before_tool_signal_no_leak(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ "I will look this up. ",
+ "Some more prose that's long enough to leave the buffer. ",
+ '{"name":"web_search","arguments":{"query":"x"}}',
+ ],
+ ["all done"],
+ ],
+ exec_results = ["weather: sunny"],
+ max_tool_iterations = 2,
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls == [("web_search", {"query": "x"})]
+ for e in events:
+ if e["type"] == "content":
+ assert "" not in e["text"]
+ assert "web_search" not in e["text"]
+
+ def test_auto_heal_disabled_still_parses_valid_tool_call(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"x"}}'
+ ],
+ ["done"],
+ ],
+ exec_results = ["OK"],
+ auto_heal_tool_calls = False,
+ max_tool_iterations = 2,
+ )
+ _collect_events(loop)
+ assert exec_fn.calls == [("web_search", {"query": "x"})]
+
+ def test_non_consecutive_duplicate_is_short_circuited(self):
+ loop, exec_fn = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"A"}}'
+ ],
+ [
+ '{"name":"web_search","arguments":{"query":"B"}}'
+ ],
+ [
+ '{"name":"web_search","arguments":{"query":"A"}}'
+ ],
+ ["final"],
+ ],
+ exec_results = ["res-A", "res-B"],
+ max_tool_iterations = 4,
+ )
+ events = _collect_events(loop)
+ assert exec_fn.calls == [
+ ("web_search", {"query": "A"}),
+ ("web_search", {"query": "B"}),
+ ]
+ tool_ends = [e for e in events if e["type"] == "tool_end"]
+ assert "already made this exact call" in tool_ends[-1]["result"]
+
+ def test_coerce_string_args_python_uses_code_key(self):
+ assert _coerce_arguments("print(1)", heal = True, tool_name = "python") == {
+ "code": "print(1)"
+ }
+
+ def test_coerce_string_args_terminal_uses_command_key(self):
+ assert _coerce_arguments("ls -la", heal = True, tool_name = "terminal") == {
+ "command": "ls -la"
+ }
+
+ def test_tool_call_ids_unique_across_loop_iterations(self):
+ loop, _exec = _make_loop(
+ turns = [
+ [
+ '{"name":"web_search","arguments":{"query":"A"}}'
+ ],
+ [
+ '{"name":"web_search","arguments":{"query":"B"}}'
+ ],
+ ["done"],
+ ],
+ exec_results = ["A", "B"],
+ max_tool_iterations = 3,
+ )
+ events = _collect_events(loop)
+ ids = [e["tool_call_id"] for e in events if e["type"] == "tool_start"]
+ assert len(ids) == 2 and ids[0] != ids[1]
+
+
+# ────────────────────────────────────────────────────────────────────
+# Shared gpt-oss name detector
+# ────────────────────────────────────────────────────────────────────
+
+
+class TestGptOssNameDetection:
+ def test_substring_match(self):
+ assert is_gpt_oss_model_name("unsloth/gpt-oss-20b") is True
+
+ def test_negative_known_non_oss_model(self):
+ assert is_gpt_oss_model_name("meta-llama/Llama-3.1-8B-Instruct") is False
+
+ def test_empty_or_none_returns_false(self):
+ assert is_gpt_oss_model_name("") is False
+ assert is_gpt_oss_model_name(None) is False
+
+
+if __name__ == "__main__":
+ pytest.main([__file__, "-v"])
diff --git a/studio/backend/utils/datasets/__init__.py b/studio/backend/utils/datasets/__init__.py
index c9237d83c1..7988b09972 100644
--- a/studio/backend/utils/datasets/__init__.py
+++ b/studio/backend/utils/datasets/__init__.py
@@ -59,6 +59,7 @@ from .model_mappings import (
TEMPLATE_TO_MODEL_MAPPER,
MODEL_TO_TEMPLATE_MAPPER,
TEMPLATE_TO_RESPONSES_MAPPER,
+ is_gpt_oss_model_name,
)
# Legacy imports from the original dataset_utils.py for backward compatibility
@@ -98,6 +99,7 @@ __all__ = [
"TEMPLATE_TO_MODEL_MAPPER",
"MODEL_TO_TEMPLATE_MAPPER",
"TEMPLATE_TO_RESPONSES_MAPPER",
+ "is_gpt_oss_model_name",
# Main entry points
"check_dataset_format",
"format_and_template_dataset",
diff --git a/studio/backend/utils/datasets/model_mappings.py b/studio/backend/utils/datasets/model_mappings.py
index 21e8566ac5..eb2e5482c9 100644
--- a/studio/backend/utils/datasets/model_mappings.py
+++ b/studio/backend/utils/datasets/model_mappings.py
@@ -442,6 +442,26 @@ for key, values in TEMPLATE_TO_MODEL_MAPPER.items():
MODEL_TO_TEMPLATE_MAPPER[value.lower()] = lowered_key
+def is_gpt_oss_model_name(name: str) -> bool:
+ """Name-based check for gpt-oss / harmony models.
+
+ Used by both the in-process backend and the parent-process
+ orchestrator to detect harmony models without an IPC round-trip.
+ """
+ name = (name or "").lower()
+ if not name:
+ return False
+ try:
+ if MODEL_TO_TEMPLATE_MAPPER.get(name) == "gpt-oss":
+ return True
+ for key, tmpl in MODEL_TO_TEMPLATE_MAPPER.items():
+ if tmpl == "gpt-oss" and (key in name or name in key):
+ return True
+ except Exception:
+ pass
+ return "gpt-oss" in name
+
+
TEMPLATE_TO_RESPONSES_MAPPER = {
"gemma-4-thinking": {
"instruction": "<|turn>user\n",
diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts
index f87229d78f..359099e8b3 100644
--- a/studio/frontend/src/features/chat/api/chat-adapter.ts
+++ b/studio/frontend/src/features/chat/api/chat-adapter.ts
@@ -601,6 +601,12 @@ async function autoLoadSmallestModel(): Promise<{
reasoningStyle: sfLoadResp.reasoning_style ?? "enable_thinking",
supportsPreserveThinking: sfLoadResp.supports_preserve_thinking ?? false,
supportsTools: sfLoadResp.supports_tools ?? false,
+ // Parity with the GGUF branch above.
+ toolsEnabled: sfLoadResp.supports_tools ?? false,
+ codeToolsEnabled: sfLoadResp.supports_tools ?? false,
+ defaultChatTemplate: sfLoadResp.chat_template ?? null,
+ chatTemplateOverride: null,
+ loadedChatTemplateOverride: null,
});
const sfModel: ChatModelSummary = {
id: repo.repo_id,