Merge remote-tracking branch 'origin/main' into feature/deep-research

# Conflicts:
#	studio/frontend/src/features/chat/chat-page.tsx
This commit is contained in:
alkinun 2026-07-19 09:27:41 +03:00
commit d65c9520cb
43 changed files with 4215 additions and 391 deletions

View file

@ -268,6 +268,7 @@ jobs:
tests/saving/test_save_shell_injection.py \
tests/saving/test_patch_saving_none_tokenizer.py \
tests/saving/test_fix_sentencepiece_gguf_robustness.py \
tests/saving/test_fix_sentencepiece_tokenizer_guard.py \
tests/saving/test_compressed_export_schemes.py \
tests/saving/test_export_api_surface.py \
tests/saving/test_export_dispatch.py \
@ -358,6 +359,7 @@ jobs:
tests/saving/test_save_shell_injection.py \
tests/saving/test_patch_saving_none_tokenizer.py \
tests/saving/test_fix_sentencepiece_gguf_robustness.py \
tests/saving/test_fix_sentencepiece_tokenizer_guard.py \
tests/saving/test_compressed_export_schemes.py \
tests/saving/test_export_api_surface.py \
tests/saving/test_export_dispatch.py \

34
studio/MCP.md Normal file
View file

@ -0,0 +1,34 @@
# Unsloth Studio MCP server
Studio can expose a local MCP server so an MCP client can inspect models and
GPU state, validate recipes, start or stop training, inspect recipe output, and
export a loaded model.
The server is disabled by default. Enable it for a local Studio process with:
```bash
UNSLOTH_STUDIO_ENABLE_MCP=1 \
UNSLOTH_STUDIO_MCP_TOKEN='use-a-local-secret' \
unsloth studio
```
The endpoint is `http://127.0.0.1:8888/mcp/` when Studio uses its default port
(a request to `/mcp` redirects to the canonical `/mcp/`). Use the actual Studio
port when it is configured differently.
The high-impact tools are:
- `studio_status` and `list_local_models` for discovery
- `get_training_status`, `start_training`, `stop_training`, and `list_training_runs`
- `validate_recipe`, `get_recipe_job_status`, and `get_recipe_job_dataset`
- `load_checkpoint` and `export_gguf`
`start_training` accepts the same fields as the Studio `TrainingStartRequest`.
The request is validated by the existing Pydantic model before a subprocess is
started. Export paths use the existing Studio validation as well.
The endpoint always requires `UNSLOTH_STUDIO_MCP_TOKEN` and checks an exact
Bearer token for both HTTP and WebSocket connections. Keep it on localhost
unless the deployment has an authenticated reverse proxy. The MCP endpoint is
intentionally opt-in because tools can consume GPU memory, write model
artifacts, and stop active work.

View file

@ -10,10 +10,242 @@ native-chat-template fallback used by the transformers and MLX backends.
import copy
import json
import logging
from dataclasses import dataclass
from typing import Optional
_THINK_OPEN = "<think>"
_THINK_CLOSE = "</think>"
_GEMMA_CHANNEL_START = "<|channel>"
_GEMMA_THOUGHT_OPEN = "<|channel>thought"
_GEMMA_THOUGHT_CLOSE = "<channel|>"
_GEMMA_TEMPLATE_OPENERS = (
_GEMMA_THOUGHT_OPEN + "\n",
_GEMMA_THOUGHT_OPEN + "\\n",
_GEMMA_THOUGHT_OPEN + _GEMMA_THOUGHT_CLOSE,
)
def _tokenizer_objects(tokenizer) -> tuple:
"""Return a processor/tokenizer and its distinct nested tokenizer."""
if tokenizer is None:
return ()
nested = getattr(tokenizer, "tokenizer", None)
return (tokenizer,) if nested is None or nested is tokenizer else (tokenizer, nested)
def _selected_template_strings_from_value(
template,
tools = None,
*,
prefer_tool_use: bool = True,
) -> tuple[str, ...]:
"""Return the named chat template matching HF's default selection rules."""
tools = tools or None
if isinstance(template, str):
return (template,)
if not isinstance(template, dict):
return ()
if prefer_tool_use and tools and isinstance(template.get("tool_use"), str):
return (template["tool_use"],)
if isinstance(template.get("default"), str):
return (template["default"],)
values = tuple(value for value in template.values() if isinstance(value, str))
return values if len(values) == 1 else ()
def _selected_chat_template_strings(tokenizer, tools = None) -> tuple[str, ...]:
"""Return the active chat template selected for this request."""
tools = tools or None
getter = getattr(tokenizer, "get_chat_template", None)
if callable(getter):
for kwargs in ({"chat_template": None, "tools": tools}, {"tools": tools}, {}):
try:
selected = getter(**kwargs)
except Exception:
continue
if isinstance(selected, str):
return (selected,)
# ProcessorMixin.apply_chat_template does not switch to "tool_use" implicitly;
# it uses "default" unless chat_template= names another template.
is_processor = getattr(tokenizer, "tokenizer", None) is not None and callable(
getattr(tokenizer, "apply_chat_template", None)
)
return _selected_template_strings_from_value(
getattr(tokenizer, "chat_template", None),
tools,
prefer_tool_use = not is_processor,
)
def _detect_reasoning_channel_markers_from_templates(
templates: tuple[str, ...],
) -> Optional[tuple[str, str]]:
"""Return Gemma native reasoning markers only when a template emits them."""
if any(opener in template for template in templates for opener in _GEMMA_TEMPLATE_OPENERS):
return _GEMMA_THOUGHT_OPEN, _GEMMA_THOUGHT_CLOSE
return None
def detect_reasoning_channel_markers(tokenizer, tools = None) -> Optional[tuple[str, str]]:
"""Return native Gemma thought-channel markers supported by a tokenizer.
Detection uses the active chat template rather than model names or vocabulary
membership. Some models expose Gemma control tokens without using the native
thought-channel response protocol, and those must keep normal
``skip_special_tokens`` streaming.
"""
for obj in _tokenizer_objects(tokenizer):
templates = _selected_chat_template_strings(obj, tools)
if templates:
return _detect_reasoning_channel_markers_from_templates(templates)
return None
def detect_reasoning_channel_markers_from_template(
template, tools = None
) -> Optional[tuple[str, str]]:
"""Return native Gemma thought-channel markers from a raw template value."""
return _detect_reasoning_channel_markers_from_templates(
_selected_template_strings_from_value(template, tools)
)
def detect_reasoning_channel_markers_from_model_info(
tokenizer,
model_info: Optional[dict] = None,
tools = None,
) -> Optional[tuple[str, str]]:
"""Return reasoning markers from the active or cached native template."""
markers = detect_reasoning_channel_markers(tokenizer, tools = tools)
if markers is not None or not isinstance(model_info, dict):
return markers
native_templates = (
model_info.get("native_chat_template"),
(model_info.get("chat_template_info") or {}).get("template"),
)
for template in native_templates:
markers = detect_reasoning_channel_markers_from_template(template, tools)
if markers is not None:
return markers
return None
@dataclass(frozen = True)
class ChatTemplateRenderResult:
"""Prompt plus response-protocol metadata selected by the renderer."""
prompt: str
reasoning_channel_markers: Optional[tuple[str, str]] = None
def _split_partial_marker(text: str, marker: str) -> tuple[str, str]:
"""Hold the longest suffix that may become ``marker`` in the next chunk."""
for length in range(min(len(text), len(marker) - 1), 0, -1):
if text.endswith(marker[:length]):
return text[:-length], text[-length:]
return text, ""
class ReasoningChannelNormalizer:
"""Incrementally convert one native reasoning channel to ``<think>``.
The parser follows mlx-vlm's streaming boundary behavior but emits Studio's
established canonical text contract. Only the configured opening and
closing markers are consumed; tool-call and other control markers remain
available to downstream parsers.
"""
def __init__(self, opening_marker: str, closing_marker: str):
self._opening_marker = opening_marker
self._closing_marker = closing_marker
self._buffer = ""
self._in_reasoning = False
self._reasoning_done = False
self._skip_opening_newline = False
def feed(self, text: str) -> str:
"""Consume a raw text delta and return the stable canonical delta."""
self._buffer += text or ""
output: list[str] = []
while self._buffer:
if self._reasoning_done:
output.append(self._buffer)
self._buffer = ""
break
if self._in_reasoning and self._skip_opening_newline:
if self._buffer.startswith("\n"):
self._buffer = self._buffer[1:]
self._skip_opening_newline = False
if not self._buffer:
break
marker = self._closing_marker if self._in_reasoning else self._opening_marker
index = self._buffer.find(marker)
if index < 0:
stable, self._buffer = _split_partial_marker(self._buffer, marker)
output.append(stable)
break
output.append(self._buffer[:index])
self._buffer = self._buffer[index + len(marker) :]
if self._in_reasoning:
output.append(_THINK_CLOSE)
self._in_reasoning = False
self._reasoning_done = True
else:
output.append(_THINK_OPEN)
self._in_reasoning = True
self._skip_opening_newline = True
return "".join(output)
def finish(self) -> str:
"""Flush a naturally completed stream and close an open think block."""
output = self.drain()
if self._in_reasoning:
output += _THINK_CLOSE
self._in_reasoning = False
self._reasoning_done = True
return output
def drain(self) -> str:
"""Flush buffered literal text without synthesizing a closing tag."""
output = self._buffer
self._buffer = ""
return output
def normalize_reasoning_snapshots(
stream,
tokenizer = None,
cancel_event = None,
markers: Optional[tuple[str, str]] = None,
tools = None,
):
"""Normalize a prefix-monotonic cumulative text stream when supported."""
markers = markers or detect_reasoning_channel_markers(tokenizer, tools = tools)
if markers is None:
yield from stream
return
normalizer = ReasoningChannelNormalizer(*markers)
raw_output = ""
normalized_output = ""
for snapshot in stream:
if not snapshot.startswith(raw_output):
raise RuntimeError("Reasoning normalization requires cumulative text snapshots")
delta = normalizer.feed(snapshot[len(raw_output) :])
raw_output = snapshot
if delta:
normalized_output += delta
yield normalized_output
cancelled = cancel_event is not None and cancel_event.is_set()
tail = normalizer.drain() if cancelled else normalizer.finish()
if tail:
normalized_output += tail
yield normalized_output
def detect_think_prefill(prompt: Optional[str], special_tokens = None) -> str:
@ -166,7 +398,8 @@ def render_native_template(
preserve_thinking: Optional[bool] = None,
apply_fn = None,
hf_token: Optional[str] = None,
) -> Optional[str]:
return_metadata: bool = False,
):
"""Render ``messages`` + ``tools`` with the model's NATIVE chat template.
Some Unsloth override templates (e.g. ``mistral``, ``gemma-4``) do not emit
@ -175,7 +408,9 @@ def render_native_template(
tool-calling syntax. It is loaded straight from the repo (bypassing any
override on the live tokenizer) and cached on ``model_info``. Returns the
rendered prompt only if the native template actually emits the tools (render
differs with vs without tools); otherwise ``None``.
differs with vs without tools); otherwise ``None``. With ``return_metadata``,
returns ``ChatTemplateRenderResult`` so callers can stream with the response
protocol selected by this request's template.
``hf_token`` is the token the model was loaded with -- passed to the repo load
so a gated/private model's native template can still be fetched (otherwise the
@ -261,7 +496,16 @@ def render_native_template(
exc,
)
return None
return with_tools if with_tools != no_tools else None
if with_tools == no_tools:
return None
if return_metadata:
return ChatTemplateRenderResult(
with_tools,
_detect_reasoning_channel_markers_from_templates(
_selected_template_strings_from_value(native_tpl, tools)
),
)
return with_tools
def render_with_native_template_fallback(
@ -277,7 +521,8 @@ def render_with_native_template_fallback(
preserve_thinking: Optional[bool] = None,
apply_fn = None,
hf_token: Optional[str] = None,
) -> str:
return_metadata: bool = False,
):
"""Return ``formatted_prompt``, swapping in a native-template render when an
override template dropped the ``tools`` schema.
@ -285,9 +530,27 @@ def render_with_native_template_fallback(
them (detected by comparison, robust against tool names in the system prompt),
re-render with the model's native template. Shared by the transformers and MLX
backends so both advertise tools consistently. ``hf_token`` is forwarded so a
gated/private model's native template can still be fetched."""
gated/private model's native template can still be fetched. With
``return_metadata``, returns the selected prompt plus reasoning-channel markers
for the exact template used by this request."""
live_markers = detect_reasoning_channel_markers(tokenizer, tools = tools)
def _result(prompt: str, markers = live_markers):
if return_metadata:
return ChatTemplateRenderResult(prompt, markers)
return prompt
if not tools:
return formatted_prompt
# Gemma 4 can emit its native reasoning protocol even when a generation-time
# Unsloth override rendered a marker-free prompt. Preserve the live-verified
# no-tools thinking behavior without letting cached native metadata describe
# unrelated tool prompts that kept the active override.
markers = live_markers
if markers is None:
markers = detect_reasoning_channel_markers_from_model_info(
tokenizer, model_info, tools = None
)
return _result(formatted_prompt, markers)
if apply_fn is None:
apply_fn = apply_chat_template_for_generation
# Probe whether the live template dropped the schema. A tools-requiring template
@ -307,9 +570,9 @@ def render_with_native_template_fallback(
active_model_name,
exc,
)
return formatted_prompt
return _result(formatted_prompt)
if formatted_prompt != probe_no_tools:
return formatted_prompt # template already emits the tools schema
return _result(formatted_prompt) # template already emits the tools schema
native_prompt = render_native_template(
model_info = model_info,
active_model_name = active_model_name,
@ -320,6 +583,7 @@ def render_with_native_template_fallback(
preserve_thinking = preserve_thinking,
apply_fn = apply_fn,
hf_token = hf_token,
return_metadata = return_metadata,
)
if native_prompt:
logger.info(
@ -328,4 +592,4 @@ def render_with_native_template_fallback(
active_model_name,
)
return native_prompt
return formatted_prompt
return _result(formatted_prompt)

View file

@ -5,7 +5,7 @@
from unsloth import FastLanguageModel, FastVisionModel
from unsloth.chat_templates import get_chat_template
from transformers import TextStreamer
from transformers import TextIteratorStreamer, TextStreamer
from peft import PeftModel, PeftModelForCausalLM
import json
@ -32,6 +32,11 @@ from core.inference.chat_eos import (
chat_eos_repair,
resolve_chat_turn_end_eos_ids_using,
)
from core.inference.chat_template_helpers import (
ReasoningChannelNormalizer,
detect_reasoning_channel_markers,
detect_think_prefill,
)
from core.inference.presence_penalty import _make_presence_penalty_processor
from io import StringIO
import structlog
@ -187,6 +192,53 @@ class HarmonyTextStreamer:
self._queue.put(new_content)
class ReasoningTextIteratorStreamer(TextIteratorStreamer):
"""TextIteratorStreamer that preserves native channel tokens until parsed."""
def __init__(
self,
tokenizer,
*,
markers: tuple[str, str],
skip_prompt: bool = True,
timeout: float = 0.2,
cancel_event = None,
**decode_kwargs,
):
decode_kwargs["skip_special_tokens"] = False
super().__init__(tokenizer, skip_prompt = skip_prompt, timeout = timeout, **decode_kwargs)
self._normalizer = ReasoningChannelNormalizer(*markers)
self._cancel_event = cancel_event
self._aborted = False
def abort(self):
"""Mark generation as failed so ``end`` drains without closing."""
self._aborted = True
def on_finalized_text(
self,
text: str,
stream_end: bool = False,
):
"""Queue canonical deltas, closing only on natural stream completion."""
delta = self._normalizer.feed(text)
if delta:
self.text_queue.put(delta, timeout = self.timeout)
if stream_end:
cancelled = self._aborted or (
self._cancel_event is not None and self._cancel_event.is_set()
)
tail = self._normalizer.drain() if cancelled else self._normalizer.finish()
if tail:
self.text_queue.put(tail, timeout = self.timeout)
self.text_queue.put(self.stop_signal, timeout = self.timeout)
class _GenerationThreadError(RuntimeError):
"""Generation worker failures that should propagate through stream routes."""
class InferenceBackend:
"""Unified inference backend supporting text, vision, and LoRA models"""
@ -836,6 +888,7 @@ class InferenceBackend:
thread_id: Optional[str] = None,
rag_scope: Optional[dict] = None,
presence_penalty: float = 0.0,
reasoning_prefilled: bool = False,
):
"""Run an agentic tool loop on top of ``generate_chat_response``.
@ -889,6 +942,7 @@ class InferenceBackend:
session_id = session_id,
thread_id = thread_id,
rag_scope = rag_scope,
reasoning_prefilled = reasoning_prefilled,
)
def generate_chat_response(
@ -960,8 +1014,7 @@ class InferenceBackend:
thread can toggle adapters under the generation lock.
"""
if not self.active_model_name:
yield "Error: No active model"
return
raise RuntimeError("No active model")
model_info = self.models[self.active_model_name]
is_vision = model_info.get("is_vision", False)
@ -1049,6 +1102,7 @@ class InferenceBackend:
template_messages = [{"role": "system", "content": system_prompt}] + messages
else:
template_messages = messages
reasoning_channel_markers_resolved = False
try:
if not (hasattr(tokenizer, "chat_template") and tokenizer.chat_template):
raise ValueError(
@ -1058,6 +1112,7 @@ class InferenceBackend:
f"Please use a model that includes a chat template, or manually set "
f"one via tokenizer.chat_template before inference."
)
reasoning_channel_markers = None
formatted_prompt = self._apply_chat_template_for_generation(
tokenizer,
template_messages,
@ -1073,7 +1128,7 @@ class InferenceBackend:
render_with_native_template_fallback,
)
formatted_prompt = render_with_native_template_fallback(
render_result = render_with_native_template_fallback(
formatted_prompt = formatted_prompt,
tokenizer = tokenizer,
model_info = model_info,
@ -1085,13 +1140,19 @@ class InferenceBackend:
preserve_thinking = preserve_thinking,
apply_fn = self._apply_chat_template_for_generation,
hf_token = model_info.get("hf_token"),
return_metadata = True,
)
formatted_prompt = render_result.prompt
reasoning_channel_markers = render_result.reasoning_channel_markers
reasoning_channel_markers_resolved = True
logger.debug(f"Formatted prompt: {formatted_prompt[:200]}...")
except Exception as e:
logger.error(f"Error applying chat template: {e}")
# Fall back to manual formatting
formatted_prompt = self.format_chat_prompt(messages, system_prompt)
reasoning_channel_markers = None
reasoning_channel_markers_resolved = True
# Step 3: generate
yield from self.generate_stream(
@ -1105,6 +1166,8 @@ class InferenceBackend:
cancel_event = cancel_event,
_adapter_state = _adapter_state,
presence_penalty = presence_penalty,
reasoning_channel_markers = reasoning_channel_markers,
reasoning_channel_markers_resolved = reasoning_channel_markers_resolved,
)
def _generate_vision_response(
@ -1190,21 +1253,27 @@ class InferenceBackend:
# Stream with TextIteratorStreamer + background thread
try:
from core.inference.chat_template_helpers import detect_think_prefill
# Re-emit an open <think> prefill swallowed by skip_prompt (see
# generate_stream).
think_prefix = detect_think_prefill(
prompt_text, getattr(raw_tokenizer, "all_special_tokens", None)
)
from transformers import TextIteratorStreamer
import threading
streamer = TextIteratorStreamer(
streamer = self._make_text_streamer(
raw_tokenizer,
protocol_source = processor,
# The text-only VLM fallback above did not render with the
# processor template, so its native markers do not describe
# this request's response protocol.
reasoning_channel_markers = detect_reasoning_channel_markers(processor)
if image
else None,
reasoning_channel_markers_resolved = True,
skip_prompt = True,
skip_special_tokens = True,
timeout = 0.2,
cancel_event = cancel_event,
use_harmony = self._is_gpt_oss_model(),
)
generation_kwargs = dict(
@ -1226,6 +1295,10 @@ class InferenceBackend:
)
if _pp is not None:
generation_kwargs["logits_processor"] = _pp
stopping_criteria = self._cancel_stopping_criteria(cancel_event)
if stopping_criteria is not None:
generation_kwargs["stopping_criteria"] = stopping_criteria
active_stop_token_ids = self._generation_stop_token_ids(model, generation_kwargs)
err: dict[str, str] = {}
@ -1235,6 +1308,8 @@ class InferenceBackend:
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
if hasattr(streamer, "abort"):
streamer.abort()
logger.error(f"Vision generation error in thread: {e}")
finally:
try:
@ -1251,12 +1326,17 @@ class InferenceBackend:
if think_prefix:
yield think_prefix
from queue import Empty
import time
generation_complete = False
cancel_deadline = None
try:
while True:
if cancel_event is not None and cancel_event.is_set():
break
if cancel_deadline is None:
cancel_deadline = time.monotonic() + 10
elif time.monotonic() >= cancel_deadline:
break
try:
new_token = next(streamer)
except StopIteration:
@ -1265,27 +1345,48 @@ class InferenceBackend:
except Empty:
if not thread.is_alive():
generation_complete = True
output = yield from self._drain_streamer_tail(
streamer, output, active_stop_token_ids
)
break
if cancel_deadline is not None:
remaining = cancel_deadline - time.monotonic()
if remaining <= 0:
break
thread.join(timeout = remaining)
if thread.is_alive():
break
generation_complete = True
output = yield from self._drain_streamer_tail(
streamer, output, active_stop_token_ids
)
break
continue
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
output, cleaned = self._append_stream_delta(
output, new_token, active_stop_token_ids
)
yield cleaned
finally:
if cancel_event is not None and not generation_complete:
cancel_event.set()
thread.join(timeout = 10)
join_timeout = 10
if cancel_deadline is not None:
join_timeout = max(0, cancel_deadline - time.monotonic())
thread.join(timeout = join_timeout)
if thread.is_alive():
logger.warning(
"Vision generation thread did not exit after cancel/join timeout"
)
if err.get("msg"):
yield f"Error: {err['msg']}"
raise _GenerationThreadError(err["msg"])
except _GenerationThreadError:
raise
except Exception as e:
logger.error(f"Vision generation error: {e}")
yield f"Error: {str(e)}"
raise
def generate_audio_input_response(
self,
@ -1410,11 +1511,13 @@ class InferenceBackend:
)
if err.get("msg"):
yield f"Error: {err['msg']}"
raise _GenerationThreadError(err["msg"])
except _GenerationThreadError:
raise
except Exception as e:
logger.error(f"Audio input generation error: {e}")
yield f"Error: {str(e)}"
raise
def generate_whisper_response(
self,
@ -1447,6 +1550,86 @@ class InferenceBackend:
from utils.datasets import is_gpt_oss_model_name
return is_gpt_oss_model_name(model_name or self.active_model_name or "")
def _make_text_streamer(
self,
tokenizer,
*,
protocol_source = None,
reasoning_channel_markers = None,
reasoning_channel_markers_resolved: bool = False,
skip_prompt: bool = True,
timeout: float = 0.2,
cancel_event = None,
use_harmony: bool = False,
):
"""Create the streamer matching this model's native response protocol."""
if use_harmony:
try:
return HarmonyTextStreamer(
tokenizer,
skip_prompt = skip_prompt,
timeout = timeout,
)
except Exception as e:
logger.warning(f"HarmonyTextStreamer init failed, falling back: {e}")
return TextIteratorStreamer(
tokenizer,
skip_prompt = skip_prompt,
skip_special_tokens = True,
timeout = timeout,
)
markers = (
reasoning_channel_markers
if reasoning_channel_markers_resolved
else reasoning_channel_markers
or detect_reasoning_channel_markers(protocol_source or tokenizer)
)
if markers is not None:
return ReasoningTextIteratorStreamer(
tokenizer,
markers = markers,
skip_prompt = skip_prompt,
timeout = timeout,
cancel_event = cancel_event,
)
return TextIteratorStreamer(
tokenizer,
skip_prompt = skip_prompt,
skip_special_tokens = True,
timeout = timeout,
)
def _append_stream_delta(
self,
output: str,
new_token: str,
stop_token_ids = None,
):
"""Append a streamer delta and apply response-boundary cleanup."""
output += new_token
return output, self._clean_generated_text(output, stop_token_ids = stop_token_ids)
def _drain_streamer_tail(
self,
streamer,
output: str,
stop_token_ids = None,
):
"""Drain queued streamer text after the producer exits."""
while True:
try:
new_token = next(streamer)
except StopIteration:
return output
except Exception:
return output
if new_token:
output, cleaned = self._append_stream_delta(
output, new_token, stop_token_ids = stop_token_ids
)
yield cleaned
def generate_stream(
self,
prompt: str,
@ -1459,6 +1642,8 @@ class InferenceBackend:
cancel_event = None,
_adapter_state = None,
presence_penalty: float = 0.0,
reasoning_channel_markers = None,
reasoning_channel_markers_resolved: bool = False,
) -> Generator[str, None, None]:
"""Generate a streaming text response (text models only).
@ -1467,8 +1652,7 @@ class InferenceBackend:
``presence_penalty`` matches the GGUF sampling path via a logits processor (0 disables it).
"""
if not self.active_model_name:
yield "Error: No active model"
return
raise RuntimeError("No active model")
model_info = self.models[self.active_model_name]
model = model_info["model"]
@ -1481,9 +1665,7 @@ class InferenceBackend:
try:
inputs = tokenizer(prompt, return_tensors = "pt").to(model.device)
from transformers import TextIteratorStreamer
import threading
from core.inference.chat_template_helpers import detect_think_prefill
# skip_prompt swallows an open <think> prefilled by the template;
# re-emit it so the frontend can render the thinking block.
@ -1494,30 +1676,16 @@ class InferenceBackend:
else detect_think_prefill(prompt, getattr(tokenizer, "all_special_tokens", None))
)
# gpt-oss models: HarmonyTextStreamer parses the multi-channel
# harmony protocol into <think> tags
if self._is_gpt_oss_model():
try:
streamer = HarmonyTextStreamer(
tokenizer,
skip_prompt = True,
timeout = 0.2,
)
except Exception as e:
logger.warning(f"HarmonyTextStreamer init failed, falling back: {e}")
streamer = TextIteratorStreamer(
tokenizer,
skip_prompt = True,
skip_special_tokens = True,
timeout = 0.2,
)
else:
streamer = TextIteratorStreamer(
tokenizer,
skip_prompt = True,
skip_special_tokens = True,
timeout = 0.2,
)
streamer = self._make_text_streamer(
tokenizer,
protocol_source = model_info.get("tokenizer"),
reasoning_channel_markers = reasoning_channel_markers,
reasoning_channel_markers_resolved = reasoning_channel_markers_resolved,
skip_prompt = True,
timeout = 0.2,
cancel_event = cancel_event,
use_harmony = self._is_gpt_oss_model(),
)
generation_kwargs = dict(
**inputs,
@ -1535,27 +1703,16 @@ class InferenceBackend:
if tokenizer.pad_token_id is None
else tokenizer.pad_token_id,
)
active_stop_token_ids = self._generation_stop_token_ids(model, generation_kwargs)
# Presence penalty (GGUF parity); prompt_len excludes prompt tokens.
_pp = _make_presence_penalty_processor(
presence_penalty, int(inputs["input_ids"].shape[1])
)
if _pp is not None:
generation_kwargs["logits_processor"] = _pp
if cancel_event is not None:
from transformers.generation.stopping_criteria import (
StoppingCriteria,
StoppingCriteriaList,
)
class _CancelCriteria(StoppingCriteria):
def __init__(self, ev):
self.ev = ev
def __call__(self, input_ids, scores, **kwargs):
return self.ev.is_set()
generation_kwargs["stopping_criteria"] = StoppingCriteriaList(
[_CancelCriteria(cancel_event)]
)
stopping_criteria = self._cancel_stopping_criteria(cancel_event)
if stopping_criteria is not None:
generation_kwargs["stopping_criteria"] = stopping_criteria
def generate_fn():
with self._generation_lock:
@ -1565,6 +1722,8 @@ class InferenceBackend:
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
if hasattr(streamer, "abort"):
streamer.abort()
logger.error(f"Generation error: {e}")
finally:
try:
@ -1582,12 +1741,17 @@ class InferenceBackend:
if think_prefix:
yield think_prefix
from queue import Empty
import time
generation_complete = False
cancel_deadline = None
try:
while True:
if cancel_event is not None and cancel_event.is_set():
break
if cancel_deadline is None:
cancel_deadline = time.monotonic() + 10
elif time.monotonic() >= cancel_deadline:
break
try:
new_token = next(streamer)
except StopIteration:
@ -1596,11 +1760,27 @@ class InferenceBackend:
except Empty:
if not thread.is_alive():
generation_complete = True
output = yield from self._drain_streamer_tail(
streamer, output, active_stop_token_ids
)
break
if cancel_deadline is not None:
remaining = cancel_deadline - time.monotonic()
if remaining <= 0:
break
thread.join(timeout = remaining)
if thread.is_alive():
break
generation_complete = True
output = yield from self._drain_streamer_tail(
streamer, output, active_stop_token_ids
)
break
continue
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
output, cleaned = self._append_stream_delta(
output, new_token, active_stop_token_ids
)
yield cleaned
finally:
# Set cancel_event only on early exit (user cancel), NOT on
@ -1609,16 +1789,21 @@ class InferenceBackend:
# disrupt the next serialized request (e.g. compare mode).
if cancel_event is not None and not generation_complete:
cancel_event.set()
thread.join(timeout = 10)
join_timeout = 10
if cancel_deadline is not None:
join_timeout = max(0, cancel_deadline - time.monotonic())
thread.join(timeout = join_timeout)
if thread.is_alive():
logger.warning("Generation thread did not exit after cancel/join timeout")
if err.get("msg"):
yield f"Error: {err['msg']}"
raise _GenerationThreadError(err["msg"])
except _GenerationThreadError:
raise
except Exception as e:
logger.error(f"Error during generation: {e}")
yield f"Error: {str(e)}"
raise
# ── Audio (TTS) Generation ────────────────────────────────────
@ -2107,8 +2292,42 @@ class InferenceBackend:
return img.resize(new_size, Image.Resampling.LANCZOS)
return img
def _clean_generated_text(self, text: str) -> str:
"""Strip leaked special tokens using the tokenizer's own token list."""
def _generation_stop_token_ids(self, model, generation_kwargs: dict):
"""Return the stop-token ids active for a ``generate`` call."""
if "eos_token_id" in generation_kwargs:
return generation_kwargs.get("eos_token_id")
generation_config = getattr(model, "generation_config", None)
eos_token_id = getattr(generation_config, "eos_token_id", None)
if eos_token_id is not None:
return eos_token_id
config = getattr(model, "config", None)
return getattr(config, "eos_token_id", None)
def _cancel_stopping_criteria(self, cancel_event):
"""Build a Transformers stopping criteria list for user cancellation."""
if cancel_event is None:
return None
from transformers.generation.stopping_criteria import (
StoppingCriteria,
StoppingCriteriaList,
)
class _CancelCriteria(StoppingCriteria):
def __init__(self, ev):
self.ev = ev
def __call__(self, input_ids, scores, **kwargs):
return self.ev.is_set()
return StoppingCriteriaList([_CancelCriteria(cancel_event)])
def _clean_generated_text(
self,
text: str,
*,
stop_token_ids = None,
) -> str:
"""Strip leaked response-boundary tokens after streaming."""
if self._is_gpt_oss_model():
# HarmonyTextStreamer emits clean <think>...</think>. Strip any
# harmony protocol tokens and other gpt-oss tokens (e.g.
@ -2118,10 +2337,28 @@ class InferenceBackend:
return text.strip()
tokenizer = self.models.get(self.active_model_name, {}).get("tokenizer")
tokenizer = getattr(tokenizer, "tokenizer", tokenizer)
if tokenizer:
for token in getattr(tokenizer, "all_special_tokens", []):
if token in text:
text = text.replace(token, "")
if stop_token_ids is None:
stop_token_ids = self.models.get(self.active_model_name, {}).get(
"chat_turn_end_eos_ids"
)
if isinstance(stop_token_ids, int):
stop_token_ids = (stop_token_ids,)
for token_id in stop_token_ids or ():
try:
token = tokenizer.convert_ids_to_tokens(int(token_id))
except Exception:
token = None
if isinstance(token, str) and token and text.endswith(token):
text = text[: -len(token)]
elif (
isinstance(token, str)
and token
and text.endswith("</think>")
and text[: -len("</think>")].endswith(token)
):
text = text[: -len("</think>") - len(token)] + "</think>"
return text.strip()
def _load_chat_template_info(self, model_name: str):

View file

@ -9,6 +9,7 @@ OpenAI-compatible /v1/chat/completions endpoint.
import atexit
import contextlib
import functools
import json
import os
import re
@ -999,6 +1000,283 @@ def _cached_colocated_split_main(
return None
def _cached_variant_resolution(repo_id: str, hf_variant: str) -> tuple[Optional[str], list[str]]:
"""Find a cached main GGUF and its shards for a variant."""
candidate = next(_cached_variant_candidates(repo_id, hf_variant), None)
if candidate is None:
return None, []
_, main, shards, _ = candidate
return main, shards
def _cached_variant_candidates(
repo_id: str,
hf_variant: str,
*,
require_mmproj: bool = False,
) -> Generator[tuple[str, str, list[str], Path], None, None]:
"""Yield complete cached variant copies in snapshot preference order."""
try:
from utils.models.model_config import _iter_hf_cache_snapshots
for snap in _iter_hf_cache_snapshots(repo_id):
cached_files = _gguf_snapshot_files(snap)
matches = _gguf_files_for_variant(cached_files, hf_variant)
if not matches:
continue
main = matches[0]
shards = _gguf_extra_shards(matches, main)
split = _SHARD_FULL_RE.match(main)
if split:
numbers = {
int(match.group(2))
for path in [main, *shards]
if (match := _SHARD_FULL_RE.match(path))
}
if numbers != set(range(1, int(split.group(3)) + 1)):
continue
main_path = snap.joinpath(*main.replace("\\", "/").split("/"))
if not main_path.is_file() or not _snapshot_has_all_shards(
str(main_path), main, shards, {}
):
continue
if require_mmproj and not _pick_mmproj(cached_files):
continue
yield str(main_path), main, shards, snap
except Exception as e:
logger.debug(f"Cache lookup for variant failed: {e}")
def _cached_candidate_matches_revision_size(
repo_id: str, candidate: tuple[str, str, list[str], Path], hf_token: Optional[str]
) -> bool:
"""Check cached byte sizes against the snapshot's own Hub revision.
A snapshot pointer is normally published only after its blob is complete.
When the old revision is still queryable, also compare every weight file's
size so a manually truncated cache entry is not treated as reusable. If
metadata cannot be reached, retain the cache's normal offline semantics.
"""
main_path, main, shards, snap = candidate
paths = [main, *shards]
try:
from huggingface_hub import get_paths_info
infos = list(
get_paths_info(
repo_id,
paths,
revision = snap.name,
token = hf_token,
)
)
except Exception as e:
logger.debug(
"Could not size-check cached GGUF %s at revision %s: %s",
repo_id,
snap.name,
e,
)
return True
if not infos:
# The Hub answers an unknown (e.g. force-pushed away) revision with an
# empty result, not an error; treat it like unreachable metadata.
return True
expected_sizes = {info.path: info.size for info in infos if info.size is not None}
if any(path not in expected_sizes for path in paths):
return False
try:
if os.path.getsize(main_path) < expected_sizes[main]:
return False
except OSError:
return False
return _snapshot_has_all_shards(main_path, main, shards, expected_sizes)
def _cached_complete_candidate(
repo_id: str, gguf_filename: Optional[str], shards: list[str]
) -> Optional[tuple[str, str, list[str], Path]]:
"""Return one complete exact-filename cache candidate with snapshot context."""
if not gguf_filename:
return None
if shards:
main_path = _cached_colocated_split_main(repo_id, gguf_filename, shards, {})
else:
m = _SHARD_FULL_RE.match(gguf_filename)
if m and int(m.group(3)) > 1:
return None
main_path = _cached_hf_snapshot_file(repo_id, gguf_filename)
if main_path is None:
return None
snap = _snapshot_dir_of(main_path)
if snap is None:
return None
return main_path, gguf_filename, shards, snap
def cached_gguf_for_load(
hf_repo: str,
hf_variant: Optional[str],
*,
require_mmproj: bool = False,
verify_sizes: bool = False,
hf_token: Optional[str] = None,
) -> Optional[str]:
"""Return a cached GGUF that can be loaded without downloading."""
if not hf_variant:
return None
hf_repo = _resolve_repo_id_casing(hf_repo)
for candidate in _cached_variant_candidates(
hf_repo,
hf_variant,
require_mmproj = require_mmproj,
):
if verify_sizes and not _cached_candidate_matches_revision_size(
hf_repo, candidate, hf_token
):
continue
return candidate[0]
return None
def _snapshot_dir_of(path: str) -> Optional[Path]:
"""Return the HF cache snapshot containing path, if any."""
try:
p = Path(os.path.abspath(path))
except OSError:
return None
for ancestor in p.parents:
if ancestor.parent.name == "snapshots":
return ancestor
return None
def _companion_snapshot_sibling(
near_path: str, pick: Callable[[list[str]], Optional[str]]
) -> Optional[str]:
"""Find a companion in the same snapshot as near_path."""
snap = _snapshot_dir_of(near_path)
if snap is None:
return None
try:
sibling = pick(_gguf_snapshot_files(snap))
except Exception:
return None
if not sibling:
return None
candidate = snap / sibling
return str(candidate) if candidate.is_file() else None
def _pick_mmproj(candidates: list[str]) -> Optional[str]:
mmproj_files = sorted(
f for f in candidates if f.lower().endswith(".gguf") and "mmproj" in Path(f).name.lower()
)
if not mmproj_files:
return None
return next((f for f in mmproj_files if f.lower().endswith("-f16.gguf")), mmproj_files[0])
def _hub_download_in_flight(hf_repo: str) -> bool:
try:
from hub.utils.download_registry import get_models_registry
return bool(get_models_registry().active_job_refs(hf_repo))
except Exception:
return False
def _hub_download_blocks_gguf_load(
hf_repo: str,
hf_variant: Optional[str],
*,
require_mmproj: bool = False,
hf_token: Optional[str] = None,
) -> bool:
"""Whether an active Hub job makes this GGUF load unsafe.
Same-variant jobs can reclaim the stale snapshot a load would reuse, so
they always block. Other jobs block only when this load lacks a complete
cached copy and would write to the shared cache itself.
"""
try:
from hub.utils.download_registry import get_models_registry
registry = get_models_registry()
if not registry.active_job_refs(hf_repo):
return False
if registry.has_active_variant(hf_repo, hf_variant):
return True
except Exception:
return False
return (
cached_gguf_for_load(
hf_repo,
hf_variant,
require_mmproj = require_mmproj,
verify_sizes = True,
hf_token = hf_token,
)
is None
)
# Active GGUF loads by normalized repo ID.
_LOADS_IN_FLIGHT: dict[str, int] = {}
_LOADS_IN_FLIGHT_LOCK = threading.Lock()
@contextlib.contextmanager
def gguf_load_in_flight(hf_repo: Optional[str]):
"""Track an HF GGUF load until the context exits."""
key = (hf_repo or "").strip().lower()
if not key:
yield
return
with _LOADS_IN_FLIGHT_LOCK:
_LOADS_IN_FLIGHT[key] = _LOADS_IN_FLIGHT.get(key, 0) + 1
try:
yield
finally:
with _LOADS_IN_FLIGHT_LOCK:
remaining = _LOADS_IN_FLIGHT.get(key, 1) - 1
if remaining <= 0:
_LOADS_IN_FLIGHT.pop(key, None)
else:
_LOADS_IN_FLIGHT[key] = remaining
def hf_gguf_load_in_flight(hf_repo: str) -> bool:
"""Return whether a GGUF load is active for hf_repo."""
key = (hf_repo or "").strip().lower()
if not key:
return False
with _LOADS_IN_FLIGHT_LOCK:
return _LOADS_IN_FLIGHT.get(key, 0) > 0
def _with_gguf_load_marker(load: Callable):
"""Keep an HF repo marked for the full synchronous load call."""
@functools.wraps(load)
def wrapped(self, *args, **kwargs):
hf_repo = kwargs.get("hf_repo")
with gguf_load_in_flight(hf_repo):
if hf_repo and _hub_download_blocks_gguf_load(
hf_repo,
kwargs.get("hf_variant"),
require_mmproj = bool(
kwargs.get("is_vision")
and not extra_args_disable_mmproj(kwargs.get("extra_args"))
),
hf_token = kwargs.get("hf_token"),
):
raise RuntimeError(
f"'{hf_repo}' is currently being downloaded by the download manager"
)
return load(self, *args, **kwargs)
return wrapped
def _gguf_extra_shards(files: Iterable[str], first_shard: str) -> list[str]:
m = _SHARD_FULL_RE.match(first_shard)
if not m:
@ -3806,13 +4084,13 @@ class LlamaCppBackend:
hf_repo: str,
free_bytes: int,
hf_token: Optional[str] = None,
) -> Optional[tuple[str, int]]:
) -> Optional[tuple[str, int, list[str]]]:
"""Find the smallest GGUF variant (including all shards) that fits.
Groups split shards by variant prefix and sums their sizes (e.g.
UD-Q4_K_XL with 9 shards of 50 GB each = 450 GB total).
Returns (first_shard_filename, total_size_bytes) or None.
Returns (first_shard_filename, total_size_bytes, extra_shards) or None.
"""
try:
from huggingface_hub import get_paths_info, list_repo_files
@ -3848,9 +4126,13 @@ class LlamaCppBackend:
# Smallest that fits
variant_sizes.sort(key = lambda x: x[1])
for first_file, total_size, _ in variant_sizes:
for first_file, total_size, shard_files in variant_sizes:
if total_size > 0 and total_size <= free_bytes:
return first_file, total_size
return (
first_file,
total_size,
[path for path in sorted(shard_files) if path != first_file],
)
return None
except Exception:
@ -4446,42 +4728,52 @@ class LlamaCppBackend:
except Exception as e:
logger.warning(f"Could not list repo files: {e}")
# Offline: resolve variant -> filename from the local HF cache.
# The heuristic below assumes filenames echo the repo name, which
# breaks for e.g. Qwen3.6-27B-MTP-GGUF (no "MTP" in file). Match
# against the rel path (not just basename) so subdir layouts like
# ``BF16/foo.gguf`` are findable.
# Fall back to the local cache when the repo listing is unavailable.
if not gguf_filename:
try:
from utils.models.model_config import _iter_hf_cache_snapshots
for snap in _iter_hf_cache_snapshots(hf_repo):
cached_files = _gguf_snapshot_files(snap)
matches = _gguf_files_for_variant(cached_files, hf_variant)
if not matches:
continue
gguf_filename = matches[0]
gguf_extra_shards = _gguf_extra_shards(matches, gguf_filename)
logger.info(
"Resolved variant %s -> %s from local HF cache",
hf_variant,
gguf_filename,
)
break
except Exception as e:
logger.debug(f"Offline cache lookup for variant failed: {e}")
cached_name, cached_shards = _cached_variant_resolution(hf_repo, hf_variant)
if cached_name:
gguf_filename = cached_name
gguf_extra_shards = cached_shards
logger.info(
"Resolved variant %s -> %s from local HF cache",
hf_variant,
gguf_filename,
)
if not gguf_filename:
repo_name = hf_repo.split("/")[-1].replace("-GGUF", "")
gguf_filename = f"{repo_name}-{hf_variant}.gguf"
# Prefer the existing model. Updates use force=True to fetch a new revision.
if not force:
if hf_variant:
# Resolve by variant so a newer revision's filename does not hide
# the complete older copy. Size-check against that older snapshot's
# own revision when its metadata remains available.
cached_main = cached_gguf_for_load(
hf_repo,
hf_variant,
verify_sizes = True,
hf_token = hf_token,
)
else:
candidate = _cached_complete_candidate(hf_repo, gguf_filename, gguf_extra_shards)
cached_main = (
candidate[0]
if candidate is not None
and _cached_candidate_matches_revision_size(hf_repo, candidate, hf_token)
else None
)
if cached_main is not None:
logger.info(f"Reusing cached GGUF: {cached_main}")
return cached_main
# Check disk space; fall back to a smaller variant if needed
all_gguf_files = [gguf_filename] + gguf_extra_shards
expected_sizes: dict[str, int] = {}
try:
from huggingface_hub import get_paths_info, try_to_load_from_cache
path_infos = list(get_paths_info(hf_repo, all_gguf_files, token = hf_token))
expected_sizes = {p.path: p.size for p in path_infos if p.size}
total_bytes = sum((p.size or 0) for p in path_infos)
# Subtract bytes already in the HF cache so we only preflight
@ -4490,25 +4782,10 @@ class LlamaCppBackend:
# cold whenever free disk is below the full weight footprint,
# even though nothing needs downloading.
already_cached_bytes = 0
# Cross-snapshot / case-variant cache reuse is offline-only (see the download
# path below); online, hf_hub_download fetches the current revision and
# resumes partials, so an old snapshot must not be counted as cached here or
# the preflight would under-count the download and skip the disk fallback.
# Count only files that can resume this download.
offline = _hf_env_offline()
# A split GGUF whose shards are not co-located in a single snapshot is
# refetched as a whole set later, so it must not be counted as cached here.
split_needs_refetch = False
if offline and not force and gguf_extra_shards:
# Scan all snapshots for one that holds the whole set co-located, so a
# newer snapshot with only the first shard does not mask an older
# complete one and needlessly trip the disk fallback.
if (
_cached_colocated_split_main(
hf_repo, gguf_filename, gguf_extra_shards, expected_sizes
)
is None
):
split_needs_refetch = True
# Offline split sets are reusable only when every shard shares a snapshot.
split_needs_refetch = bool(offline and not force and gguf_extra_shards)
if not force and not split_needs_refetch:
for p in path_infos:
if not p.size:
@ -4569,32 +4846,26 @@ class LlamaCppBackend:
hf_token,
)
if smaller:
fallback_file, fallback_size = smaller
fallback_file, fallback_size, fallback_shards = smaller
logger.info(
f"Selected variant too large ({total_gb:.1f} GB), "
f"falling back to {fallback_file} ({fallback_size / (1024**3):.1f} GB)"
)
gguf_filename = fallback_file
_m = _SHARD_RE.match(gguf_filename)
_prefix = _m.group(1) if _m else None
if _prefix:
prefix_lower = _prefix.lower()
gguf_extra_shards = sorted(
f
for f in all_gguf_files
if f.lower().startswith(prefix_lower)
and f != gguf_filename
and not _is_companion_gguf_path(f)
gguf_extra_shards = fallback_shards
# The selected fallback is a new load target. Apply the
# same any-revision reuse policy before starting a fetch.
fallback_candidate = _cached_complete_candidate(
hf_repo, gguf_filename, gguf_extra_shards
)
if fallback_candidate is not None and (
_cached_candidate_matches_revision_size(
hf_repo, fallback_candidate, hf_token
)
else:
gguf_extra_shards = []
# Record the fallback's size so the later cache-reuse probe can
# size-verify it; only for a single-file fallback, since
# _find_smallest_fitting_variant returns the whole-variant size
# and using that as the first shard's expected size would reject
# a valid cached first shard of a split fallback.
if not gguf_extra_shards:
expected_sizes[fallback_file] = fallback_size
):
logger.info(f"Reusing cached fallback GGUF: {fallback_candidate[0]}")
return fallback_candidate[0]
else:
raise RuntimeError(
f"Not enough disk space to download any variant. "
@ -4614,45 +4885,25 @@ class LlamaCppBackend:
raise RuntimeError("Cancelled")
dl_start = time.monotonic()
# Xet primary, HTTP fallback on stall; per-file so finished shards stay cached.
local_path = None
# Reuse a cached copy from another snapshot / case-variant repo dir only when
# offline. Online, fall through to hf_hub_download so its revision/etag check
# fetches the current file (and resumes a partial) instead of serving a stale
# same-name blob from an older revision.
if not force and _hf_env_offline():
if gguf_extra_shards:
# A split GGUF must load every shard from one snapshot; reuse only a
# snapshot that holds the whole set co-located, scanning past a newer
# snapshot that has just the first shard while an older one is complete.
local_path = _cached_colocated_split_main(
hf_repo, gguf_filename, gguf_extra_shards, expected_sizes
)
else:
local_path = _cached_hf_snapshot_file(
hf_repo,
gguf_filename,
expected_size = expected_sizes.get(gguf_filename),
)
if local_path is None:
local_path = hf_hub_download_with_xet_fallback(
local_path = hf_hub_download_with_xet_fallback(
hf_repo,
gguf_filename,
hf_token,
cancel_event = cancel_event,
on_status = lambda m: logger.info(m),
force_download = force,
)
for shard in gguf_extra_shards:
if cancel_event.is_set():
raise RuntimeError("Cancelled")
logger.info(f"Resolving GGUF shard: {shard}")
hf_hub_download_with_xet_fallback(
hf_repo,
gguf_filename,
shard,
hf_token,
cancel_event = cancel_event,
on_status = lambda m: logger.info(m),
force_download = force,
)
for shard in gguf_extra_shards:
if cancel_event.is_set():
raise RuntimeError("Cancelled")
logger.info(f"Resolving GGUF shard: {shard}")
hf_hub_download_with_xet_fallback(
hf_repo,
shard,
hf_token,
cancel_event = cancel_event,
force_download = force,
)
except Exception as e:
if isinstance(e, RuntimeError) and "Cancelled" in str(e):
raise
@ -4675,10 +4926,12 @@ class LlamaCppBackend:
pick: Callable[[list[str]], Optional[str]],
label: str,
cancel_event: Optional[threading.Event] = None,
near_path: Optional[str] = None,
) -> Optional[str]:
"""Resolve and fetch a companion GGUF (mmproj / MTP drafter) by name.
Tries the live repo file list, then the local HF cache snapshots
Prefers a companion co-located with ``near_path``'s cache snapshot,
then tries the live repo file list, then the local HF cache snapshots
(offline, same fallback as _download_gguf), then hf_hub_download.
Runs WITHOUT self._lock (like _download_gguf); honors _cancel_event so
an /unload between the main download and here skips the fetch.
@ -4688,6 +4941,17 @@ class LlamaCppBackend:
if cancel_event.is_set():
return None
# Keep companion files in the main GGUF's snapshot.
if near_path:
cached = _companion_snapshot_sibling(near_path, pick)
if cached:
logger.info("Reusing cached %s: %s", label, cached)
return cached
if _hub_download_in_flight(hf_repo):
logger.info("Skipping %s download while a hub download is active", label)
return None
target: Optional[str] = None
from huggingface_hub import list_repo_files
@ -4760,33 +5024,23 @@ class LlamaCppBackend:
hf_repo: str,
hf_token: Optional[str] = None,
cancel_event: Optional[threading.Event] = None,
near_path: Optional[str] = None,
) -> Optional[str]:
"""Download the mmproj (vision projection) file from a GGUF repo.
Prefers mmproj-F16.gguf, else any mmproj*.gguf. Returns the local
path, or None if none exists. ``cancel_event`` overrides
``self._cancel_event`` (defaults to it).
``self._cancel_event`` (defaults to it). ``near_path`` prefers a
copy co-located with the main GGUF's cache snapshot.
"""
def _pick_mmproj(candidates: list[str]) -> Optional[str]:
mmproj_files = sorted(
f
for f in candidates
if f.lower().endswith(".gguf") and "mmproj" in Path(f).name.lower()
)
if not mmproj_files:
return None
for f in mmproj_files:
if f.lower().endswith("-f16.gguf"):
return f
return mmproj_files[0]
return self._download_companion_gguf(
hf_repo = hf_repo,
hf_token = hf_token,
pick = _pick_mmproj,
label = "mmproj",
cancel_event = cancel_event,
near_path = near_path,
)
def _cached_repo_mtp_drafter(self, hf_repo: str) -> Optional[str]:
@ -4817,6 +5071,7 @@ class LlamaCppBackend:
*,
hf_repo: str,
hf_token: Optional[str] = None,
near_path: Optional[str] = None,
) -> Optional[str]:
"""Download the separate MTP drafter (speculative head) from a GGUF repo.
@ -4828,16 +5083,6 @@ class LlamaCppBackend:
are intentionally skipped. Returns the local path, or None.
"""
# Offline, reuse any drafter already on disk (a fresh copy can't be
# fetched). Online, _download_companion_gguf/hf_hub_download reuse the
# current cached file and refetch a changed one, so skip the probe here
# rather than pair new weights with a stale draft.
if _hf_env_offline():
cached = self._cached_repo_mtp_drafter(hf_repo)
if cached:
logger.info(f"Reusing cached MTP drafter (offline): {cached}")
return cached
def _pick_mtp(candidates: list[str]) -> Optional[str]:
# Root-level only: MTP/ subdir copies now share the mtp- prefix but
# are explicit-selection, not auto-fetch (they'd sort ahead of root).
@ -4850,11 +5095,28 @@ class LlamaCppBackend:
)
return mtp_files[0] if mtp_files else None
if near_path:
cached = _companion_snapshot_sibling(near_path, _pick_mtp)
if cached:
logger.info("Reusing cached MTP drafter: %s", cached)
return cached
# Offline, reuse any drafter already on disk (a fresh copy can't be
# fetched). Online, _download_companion_gguf/hf_hub_download reuse the
# current cached file and refetch a changed one, so skip the probe here
# rather than pair new weights with a stale draft.
if _hf_env_offline():
cached = self._cached_repo_mtp_drafter(hf_repo)
if cached:
logger.info(f"Reusing cached MTP drafter (offline): {cached}")
return cached
return self._download_companion_gguf(
hf_repo = hf_repo,
hf_token = hf_token,
pick = _pick_mtp,
label = "MTP drafter",
near_path = near_path,
)
def _resolve_launch_mmproj_path(
@ -5436,6 +5698,7 @@ class LlamaCppBackend:
)
self._stdout_thread.start()
@_with_gguf_load_marker
def load_model(
self,
*,
@ -5581,6 +5844,7 @@ class LlamaCppBackend:
mmproj_path = self._download_mmproj(
hf_repo = hf_repo,
hf_token = hf_token,
near_path = model_path,
)
# Auto-download the separate MTP drafter (e.g. Gemma) when
# the requested spec mode can use it. Repos with the head
@ -5598,6 +5862,7 @@ class LlamaCppBackend:
mtp_draft_path = self._download_mtp(
hf_repo = hf_repo,
hf_token = hf_token,
near_path = model_path,
)
elif gguf_path:
if not Path(gguf_path).is_file():

View file

@ -11,6 +11,10 @@ import threading
from typing import Optional, Generator
from core.inference.message_content import content_to_text
from core.inference.runtime_context import runtime_context_length
from core.inference.chat_template_helpers import (
ReasoningChannelNormalizer,
normalize_reasoning_snapshots,
)
from loggers import get_logger
logger = get_logger(__name__)
@ -533,7 +537,7 @@ class MLXInferenceBackend:
break
if self._is_vlm:
yield from self._generate_vlm(
stream = self._generate_vlm(
full_messages,
image,
temperature,
@ -550,7 +554,7 @@ class MLXInferenceBackend:
presence_penalty = presence_penalty,
)
else:
yield from self._generate_text(
stream = self._generate_text(
full_messages,
temperature,
top_p,
@ -565,6 +569,7 @@ class MLXInferenceBackend:
preserve_thinking = preserve_thinking,
presence_penalty = presence_penalty,
)
yield from stream
def _generate_text(
self,
@ -609,7 +614,7 @@ class MLXInferenceBackend:
# probe and native render share a renderer. (VLM renders via the
# processor for image tokens and is not wired here.)
model_info = self.models.get(self.active_model_name, {})
prompt = render_with_native_template_fallback(
render_result = render_with_native_template_fallback(
formatted_prompt = prompt,
tokenizer = self._tokenizer,
model_info = model_info,
@ -620,7 +625,10 @@ class MLXInferenceBackend:
reasoning_effort = reasoning_effort,
preserve_thinking = preserve_thinking,
hf_token = model_info.get("hf_token"),
return_metadata = True,
)
prompt = render_result.prompt
reasoning_channel_markers = render_result.reasoning_channel_markers
# An open <think> prefilled by the template lives in the prompt, not
# the generated tokens; re-emit it so the frontend renders the block.
@ -654,7 +662,17 @@ class MLXInferenceBackend:
if not logits_processors:
logits_processors = None
preserve_native_channels = reasoning_channel_markers is not None
token_ids = []
normalizer = (
ReasoningChannelNormalizer(*reasoning_channel_markers)
if reasoning_channel_markers is not None
else None
)
# MLX consumers diff cumulative snapshots. Keep a prompt-prefilled
# <think> prefix on every native-protocol snapshot just as the normal
# decoding path does below.
normalized_output = think_prefix
logger.info(
"Generating: prompt_len=%d, max_tokens=%d, model=%s, tokenizer=%s",
len(prompt),
@ -678,12 +696,19 @@ class MLXInferenceBackend:
**gen_kwargs,
):
final_response = response
token_ids.append(response.token)
cumulative = self._tokenizer.decode(
token_ids,
skip_special_tokens = True,
)
yield think_prefix + cumulative
if preserve_native_channels:
piece = getattr(response, "text", None) or ""
delta = normalizer.feed(piece)
if delta:
normalized_output += delta
yield normalized_output
else:
token_ids.append(response.token)
cumulative = self._tokenizer.decode(
token_ids,
skip_special_tokens = True,
)
yield think_prefix + cumulative
if cancel_event and cancel_event.is_set():
break
@ -700,6 +725,12 @@ class MLXInferenceBackend:
getattr(final_response, "generation_tokens", 0),
getattr(final_response, "generation_tps", 0.0),
)
if normalizer is not None:
cancelled = cancel_event is not None and cancel_event.is_set()
tail = normalizer.drain() if cancelled else normalizer.finish()
if tail:
normalized_output += tail
yield normalized_output
def _generate_vlm(
self,
@ -858,31 +889,37 @@ class MLXInferenceBackend:
elif _rep_active:
vlm_kwargs["repetition_penalty"] = float(repetition_penalty)
with self._generation_lock:
final_response = None
try:
for response in vlm_stream(
self._model,
self._processor,
prompt,
images,
**vlm_kwargs,
):
final_response = response
token_text = response.text if hasattr(response, "text") else str(response)
cumulative += token_text
yield cumulative
if cancel_event and cancel_event.is_set():
break
finally:
# mlx_vlm exposes the same stats fields as mlx_lm.
if final_response is not None:
self.last_generation_stats = _build_generation_stats(
getattr(final_response, "prompt_tokens", 0),
getattr(final_response, "prompt_tps", 0.0),
getattr(final_response, "generation_tokens", 0),
getattr(final_response, "generation_tps", 0.0),
)
def _stream_vlm_snapshots():
nonlocal cumulative
with self._generation_lock:
final_response = None
try:
for response in vlm_stream(
self._model,
self._processor,
prompt,
images,
**vlm_kwargs,
):
final_response = response
token_text = response.text if hasattr(response, "text") else str(response)
cumulative += token_text
yield cumulative
if cancel_event and cancel_event.is_set():
break
finally:
# mlx_vlm exposes the same stats fields as mlx_lm.
if final_response is not None:
self.last_generation_stats = _build_generation_stats(
getattr(final_response, "prompt_tokens", 0),
getattr(final_response, "prompt_tps", 0.0),
getattr(final_response, "generation_tokens", 0),
getattr(final_response, "generation_tps", 0.0),
)
yield from normalize_reasoning_snapshots(
_stream_vlm_snapshots(), chat_target, cancel_event, tools = tools
)
def generate_with_adapter_control(
self,

View file

@ -59,7 +59,32 @@ class GenStreamError(str):
"Error:" by checking isinstance(chunk, GenStreamError).
"""
__slots__ = ()
__slots__ = ("public",)
def __new__(
cls,
value,
*,
public: bool = False,
):
obj = str.__new__(cls, value)
obj.public = bool(public)
return obj
class GenStreamErrorRaised(RuntimeError):
"""Internal exception form of ``GenStreamError`` for generator boundaries."""
__slots__ = ("public",)
def __init__(
self,
value,
*,
public: bool = False,
):
super().__init__(value)
self.public = bool(public)
class InferenceOrchestrator:
@ -531,13 +556,19 @@ class InferenceOrchestrator:
initial_resp_queue = self._resp_queue
while True:
if self._proc is not initial_proc or self._resp_queue is not initial_resp_queue:
yield GenStreamError(f"Error: {self._subprocess_crash_message(crash_context)}")
yield GenStreamError(
f"Error: {self._subprocess_crash_message(crash_context)}",
public = True,
)
return
resp = read_one(read_timeout)
if resp is None:
# Check subprocess health
if not self._ensure_subprocess_alive():
yield GenStreamError(f"Error: {self._subprocess_crash_message(crash_context)}")
yield GenStreamError(
f"Error: {self._subprocess_crash_message(crash_context)}",
public = True,
)
return
continue
@ -689,11 +720,11 @@ class InferenceOrchestrator:
GPU work stays serialized; this only avoids orchestrator lock contention.
"""
if not self._ensure_subprocess_alive():
yield GenStreamError("Error: Inference subprocess is not running")
yield GenStreamError("Error: Inference subprocess is not running", public = True)
return
if not self.active_model_name:
yield GenStreamError("Error: No active model")
yield GenStreamError("Error: No active model", public = True)
return
# Latch the target model so the recheck below can detect a switch that completed
# between _start_dispatcher and mailbox registration (mirrors the locked path's
@ -704,7 +735,7 @@ class InferenceOrchestrator:
# so without this early-out a compare request would enqueue a generate on the
# outgoing model and delay the switch.
if self._unload_pending:
yield GenStreamError("Error: model is being unloaded")
yield GenStreamError("Error: model is being unloaded", public = True)
return
# Ensure the dispatcher runs. _start_dispatcher serializes concurrent starters under
@ -776,7 +807,7 @@ class InferenceOrchestrator:
# _stop_dispatcher joins the dispatcher, which itself takes that lock.
if orphaned_dispatcher:
self._stop_dispatcher()
yield GenStreamError("Error: model is being unloaded")
yield GenStreamError("Error: model is being unloaded", public = True)
return
try:
@ -1376,6 +1407,7 @@ class InferenceOrchestrator:
use_adapter: Optional[Union[bool, str]] = None,
stats_holder: Optional[dict] = None,
presence_penalty: float = 0.0,
reasoning_prefilled: bool = False,
**_unused,
):
"""Run the safetensors agentic tool loop in the parent process,
@ -1414,12 +1446,27 @@ class InferenceOrchestrator:
presence_penalty = presence_penalty,
)
if use_adapter is not None:
yield from self.generate_with_adapter_control(
stream = self.generate_with_adapter_control(
use_adapter = use_adapter,
**common_kwargs,
)
else:
yield from self.generate_chat_response(**common_kwargs)
stream = self.generate_chat_response(**common_kwargs)
close_stream = False
try:
for chunk in stream:
if isinstance(chunk, GenStreamError):
close_stream = True
raise GenStreamErrorRaised(str(chunk), public = chunk.public)
yield chunk
finally:
if close_stream:
close = getattr(stream, "close", None)
if callable(close):
try:
close()
except Exception:
logger.debug("failed to close errored generation stream", exc_info = True)
initial = list(messages)
if system_prompt:
@ -1441,6 +1488,7 @@ class InferenceOrchestrator:
confirm_tool_calls = confirm_tool_calls,
bypass_permissions = bypass_permissions,
permission_mode = permission_mode,
reasoning_prefilled = reasoning_prefilled,
)
def generate_with_adapter_control(
@ -1489,11 +1537,11 @@ class InferenceOrchestrator:
readers don't consume each other's tokens off the shared resp_queue.
"""
if not self._ensure_subprocess_alive():
yield GenStreamError("Error: Inference subprocess is not running")
yield GenStreamError("Error: Inference subprocess is not running", public = True)
return
if not self.active_model_name:
yield GenStreamError("Error: No active model")
yield GenStreamError("Error: No active model", public = True)
return
expected_model = self.active_model_name
@ -1510,7 +1558,7 @@ class InferenceOrchestrator:
# so we never generate on the wrong one.
if self._unload_pending or self.active_model_name != expected_model:
# Won the lock handoff during a switch; don't start on the outgoing model.
yield GenStreamError("Error: model is being unloaded")
yield GenStreamError("Error: model is being unloaded", public = True)
return
request_id = str(uuid.uuid4())
image_b64 = self._pil_to_base64(image) if image is not None else None
@ -1695,10 +1743,10 @@ class InferenceOrchestrator:
) -> Generator[str, None, None]:
"""Shared inner logic for audio input generation (Whisper + ASR)."""
if not self._ensure_subprocess_alive():
yield GenStreamError("Error: Inference subprocess is not running")
yield GenStreamError("Error: Inference subprocess is not running", public = True)
return
if not self.active_model_name:
yield GenStreamError("Error: No active model")
yield GenStreamError("Error: No active model", public = True)
return
expected_model = self.active_model_name
@ -1707,7 +1755,7 @@ class InferenceOrchestrator:
# cleared or swapped the model while we waited.
if self._unload_pending or self.active_model_name != expected_model:
# Won the lock handoff during a switch; don't start on the outgoing model.
yield GenStreamError("Error: model is being unloaded")
yield GenStreamError("Error: model is being unloaded", public = True)
return
request_id = str(uuid.uuid4())

View file

@ -276,8 +276,9 @@ PROVIDER_REGISTRY: dict[str, dict[str, Any]] = {
"auth_header": "Authorization",
"auth_prefix": "Bearer ",
"notes": (
"Local Ollama server. OpenAI-compatible /v1/chat/completions; "
"no API key. Surfaced via CUSTOM_PROVIDER_PRESETS in the frontend."
"Ollama server (local or cloud). OpenAI-compatible "
"/v1/chat/completions; API key optional (required by Ollama "
"cloud). Surfaced via CUSTOM_PROVIDER_PRESETS in the frontend."
),
"hidden": True,
},

View file

@ -50,6 +50,7 @@ from core.inference.tool_call_parser import (
# pattern lists, so the safetensors streaming strip stays aligned with the parser.
from core.tool_healing import (
_REHEARSAL_TAIL_STRIP_RE,
_THINK_CLOSE_RE,
_strip_bracket_tag_calls,
_think_spans_outside_tool_markup,
apply_tool_strip_patterns,
@ -304,6 +305,45 @@ def _status_for_tool(tool_name: str, arguments: dict) -> str:
return status_for_tool(tool_name, arguments)
def _reprompt_intent_text(text: str, *, reasoning_prefilled: bool = False) -> str:
"""Return visible answer text for the plan-without-action classifier.
Safetensors reasoning shares the cumulative text channel with the answer.
Forward-looking phrases inside ``<think>`` / ``[THINK]`` are private
planning, not a user-visible promise to call a tool. Match GGUF's behavior:
classify visible content when present and fall back to reasoning only for a
reasoning-only stall.
"""
prefilled_reasoning = ""
if reasoning_prefilled:
close = _THINK_CLOSE_RE.search(text)
if close is None:
return text.strip()
prefilled_reasoning = text[: close.end()].strip()
text = text[close.end() :].strip()
if not text:
return prefilled_reasoning
spans = _think_spans_outside_tool_markup(text)
if not spans:
return text.strip()
visible: list[str] = []
reasoning: list[str] = []
cursor = 0
for start, end in spans:
visible.append(text[cursor:start])
reasoning.append(text[start:end])
cursor = end
visible.append(text[cursor:])
visible_text = "".join(visible).strip()
reasoning_text = "".join(reasoning).strip()
if visible_text:
return visible_text
return "\n".join(part for part in (prefilled_reasoning, reasoning_text) if part).strip()
def _looks_like_enabled_bare_json(text: str, enabled_tool_names: Optional[set]) -> bool:
"""True when ``text`` opens with an ENABLED markerless bare-JSON call; an ordinary JSON answer returns False."""
probe = strip_llama3_leading_sentinels(text.lstrip())
@ -448,6 +488,7 @@ def run_safetensors_tool_loop(
confirm_tool_calls: bool = False,
bypass_permissions: bool = False,
permission_mode: Optional[str] = None,
reasoning_prefilled: bool = False,
) -> Generator[dict, None, None]:
"""Drive an agentic tool loop on top of a cumulative-text generator.
@ -956,7 +997,10 @@ def run_safetensors_tool_loop(
# (GGUF loop parity). The retry is gated on nudge_tool_calls so
# Studio callers (which send True) always nudge, while API callers
# who omit the flag keep today's no-reprompt behavior (opt-in).
stripped_answer = content_accum.strip()
intent_text = _reprompt_intent_text(
content_accum,
reasoning_prefilled = reasoning_prefilled,
)
if (
auto_heal_tool_calls
and nudge_tool_calls
@ -965,7 +1009,7 @@ def run_safetensors_tool_loop(
and not rag_autoinjected
and not tool_denied
and not any(record.executed for record in tool_controller.history)
and is_short_intent_without_action(stripped_answer)
and is_short_intent_without_action(intent_text)
):
reprompt_count += 1
logger.info(
@ -973,9 +1017,9 @@ def run_safetensors_tool_loop(
"calling tools (%d chars)",
reprompt_count,
MAX_ACT_REPROMPTS,
len(stripped_answer),
len(intent_text),
)
conversation.append({"role": "assistant", "content": stripped_answer})
conversation.append({"role": "assistant", "content": intent_text})
tool_hint = " or ".join(_active_tool_names(active_tools)) or "an available tool"
conversation.append(
{

View file

@ -60,6 +60,30 @@ def _job_status(
return DownloadJobStatus(state = state, error = error, generation = generation)
def _load_in_flight(repo_id: str) -> bool:
try:
from core.inference.llama_cpp import hf_gguf_load_in_flight
return hf_gguf_load_in_flight(repo_id)
except Exception:
return False
def _load_in_flight_error(repo_id: str) -> HTTPException:
return HTTPException(
status_code = 409,
detail = (
f"A model load for '{repo_id}' is in progress and may be "
"downloading it. Wait for the load to finish (or cancel it), "
"then start the download."
),
)
def _reject_if_load_in_flight(repo_id: str) -> None:
if _load_in_flight(repo_id):
raise _load_in_flight_error(repo_id)
def _spawn_download_worker(
repo_id: str,
variant: Optional[str],
@ -89,6 +113,9 @@ async def download_model_response(body: DownloadModelRequest, hf_token: Optional
# Canonicalize so two different-cased paste-ins share one job + cache dir.
repo_id = await asyncio.to_thread(resolve_cached_repo_id_case, repo_id, repo_type = "model")
# Avoid concurrent writers to the same HF cache files.
_reject_if_load_in_flight(repo_id)
variant = (body.gguf_variant or "").strip() or None
if variant is not None and not _is_valid_gguf_variant(variant):
raise HTTPException(
@ -147,9 +174,12 @@ async def download_model_response(body: DownloadModelRequest, hf_token: Optional
blob_hashes = variant_blob_hashes,
progress_blob_hashes = variant_progress_blob_hashes,
completed_baseline_bytes = completed_baseline_bytes,
admission_check = lambda: not _load_in_flight(repo_id),
)
generation = _registry.current_generation(key)
if not claimed:
if claim_state == "admission_blocked":
raise _load_in_flight_error(repo_id)
# claim_state is the blocking job's state. The client can attach only
# when the blocker is this key's own in-flight job (adoptable); a
# cross-variant conflict or in-progress delete is not accepted.

View file

@ -47,7 +47,7 @@ import time
import weakref
from dataclasses import dataclass, field, replace
from pathlib import Path
from typing import Iterator, Literal, Optional
from typing import Callable, Iterator, Literal, Optional
from loggers import get_logger
@ -1028,6 +1028,7 @@ class DownloadRegistry:
blob_hashes: Optional[frozenset[str]] = None,
progress_blob_hashes: Optional[frozenset[str]] = None,
completed_baseline_bytes: int = 0,
admission_check: Optional[Callable[[], bool]] = None,
generation: Optional[int] = None,
replace_active: bool = False,
metadata_transport: Optional[str] = None,
@ -1038,6 +1039,13 @@ class DownloadRegistry:
requested_hashes = blob_hashes or frozenset()
requested_progress_hashes = progress_blob_hashes or frozenset()
with self._lock:
# Run the final external admission check while the registry lock is
# held, immediately before inspecting and publishing active state.
# The GGUF load path establishes its marker before calling
# its active-job probe, so either this claim observes that marker
# or the load's later probe observes this claim.
if admission_check is not None and not admission_check():
return False, "admission_blocked"
deleting_scopes = self._deleting.get(repo)
if deleting_scopes is not None and (
None in deleting_scopes or variant_from_key(key) in deleting_scopes
@ -1222,6 +1230,23 @@ class DownloadRegistry:
)
return refs
def has_active_variant(self, repo_id: str, variant: Optional[str]) -> bool:
"""Whether an active model job targets this exact GGUF variant.
Scans the job table rather than only ``_repo_active`` so an XET-to-HTTP
retry handoff remains visible while it has temporarily released its
active slot.
"""
repo_key = normalize_repo_key(repo_id)
target = (variant or "").strip().lower() or None
with self._lock:
for key, job in self._jobs.items():
if _repo_of_key(key) != repo_key or job.state not in _ACTIVE_STATES:
continue
if self._active_job_variant_locked(key) == target:
return True
return False
def begin_delete(
self,
repo_id: str,

View file

@ -622,6 +622,22 @@ app = FastAPI(
lifespan = lifespan,
)
# The MCP surface is opt-in because it can start GPU jobs and write model
# artifacts. Mount it only when explicitly enabled by the Studio process.
if os.environ.get("UNSLOTH_STUDIO_ENABLE_MCP") == "1":
from fastmcp.utilities.lifespan import combine_lifespans
from mcp_server import BearerTokenMiddleware, create_studio_mcp
_studio_mcp_app = create_studio_mcp().http_app(path = "/")
_studio_mcp_lifespan = _studio_mcp_app.lifespan
_mcp_token = os.environ.get("UNSLOTH_STUDIO_MCP_TOKEN")
if not _mcp_token:
raise RuntimeError("UNSLOTH_STUDIO_MCP_TOKEN is required when MCP is enabled")
_studio_mcp_app = BearerTokenMiddleware(_studio_mcp_app, _mcp_token)
app.router.lifespan_context = combine_lifespans(lifespan, _studio_mcp_lifespan)
app.mount("/mcp", _studio_mcp_app)
from loggers.config import LogConfig
from loggers.handlers import LoggingMiddleware
@ -780,6 +796,7 @@ _BODY_PROTECTED_PREFIXES = (
"/api/settings",
"/api/train",
"/api/export",
"/mcp",
)
_DATASET_UPLOAD_PASSTHROUGH_PREFIX = "/api/datasets/upload"
_DATA_RECIPE_UNSTRUCTURED_UPLOAD_PASSTHROUGH_PREFIX = (

View file

@ -0,0 +1,259 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Curated MCP tools for driving an Unsloth Studio instance.
The MCP surface deliberately wraps the existing Studio services instead of
duplicating training or export logic. It is opt-in because several tools can
start GPU work or write model artifacts.
"""
from __future__ import annotations
import hmac
from typing import Any
from fastmcp import FastMCP
class BearerTokenMiddleware:
"""Require an exact bearer token when Studio MCP is exposed remotely."""
def __init__(self, app: Any, token: str) -> None:
if not token or not token.strip():
raise ValueError("Studio MCP bearer token must be a non-empty value")
if not token.isascii():
# A non-ASCII token cannot be sent in an HTTP header; reject it here.
raise ValueError("Studio MCP bearer token must contain ASCII characters only")
self.app = app
# Compare on raw header bytes: str hmac.compare_digest raises on non-ASCII
# input, which would surface as a 500 instead of a clean 401.
self.expected = token.encode("utf-8")
async def __call__(self, scope: dict[str, Any], receive: Any, send: Any) -> None:
scope_type = scope.get("type")
if scope_type not in ("http", "websocket"):
await self.app(scope, receive, send)
return
headers = dict(scope.get("headers", []))
raw_auth = headers.get(b"authorization", b"")
scheme, _, supplied = raw_auth.partition(b" ")
if scheme.lower() != b"bearer" or not hmac.compare_digest(supplied, self.expected):
await _send_unauthorized(send, scope_type)
return
await self.app(scope, receive, send)
async def _send_unauthorized(send: Any, scope_type: str) -> None:
if scope_type == "websocket":
await send({"type": "websocket.close", "code": 4401})
return
await send(
{
"type": "http.response.start",
"status": 401,
"headers": [(b"content-type", b"application/json"), (b"www-authenticate", b"Bearer")],
}
)
await send(
{
"type": "http.response.body",
"body": b'{"detail":"MCP bearer token required"}',
}
)
def _dump(value: Any) -> Any:
"""Convert Pydantic responses to plain JSON values for MCP clients."""
if hasattr(value, "model_dump"):
return value.model_dump(mode = "json")
return value
def _clamp(value: int, low: int, high: int) -> int:
"""Clamp an MCP-supplied integer into an inclusive range.
MCP tools call the Studio route functions directly, which skips FastAPI's
Query(ge=, le=) validation, so we re-apply the same bounds here.
"""
return max(low, min(value, high))
def create_studio_mcp() -> FastMCP:
"""Create the Studio MCP server and register the high-value tools."""
mcp = FastMCP(
"Unsloth Studio",
instructions = (
"Use read tools to inspect the local Studio state before starting GPU work. "
"Training and export tools can consume substantial VRAM and write files. "
"Never expose tokens or local paths from tool results unless the user asks."
),
)
@mcp.tool
async def studio_status() -> dict[str, Any]:
"""Return the current training, export, inference, and GPU state."""
from routes.export import get_export_status
from routes.inference import get_status as get_inference_status
from routes.training import get_training_status
from utils.hardware import get_gpu_utilization
training, export, inference = await _gather_status(
get_training_status(current_subject = "mcp"),
get_export_status(current_subject = "mcp"),
get_inference_status(current_subject = "mcp"),
)
return {
"training": _dump(training),
"export": _dump(export),
"inference": _dump(inference),
"hardware": get_gpu_utilization(),
}
@mcp.tool
async def list_local_models(models_dir: str = "./models") -> dict[str, Any]:
"""List local and cached models available to Studio."""
from routes.models import list_local_models as list_models
return _dump(await list_models(models_dir = models_dir, current_subject = "mcp"))
@mcp.tool
async def get_training_status() -> dict[str, Any]:
"""Read the active training job, phase, progress, and recent metrics."""
from routes.training import get_training_status as get_status
return _dump(await get_status(current_subject = "mcp"))
@mcp.tool
async def start_training(config: dict[str, Any]) -> dict[str, Any]:
"""Start a validated Studio training job from a TrainingStartRequest-shaped object.
The config is validated by the same Pydantic model used by the Studio UI.
Call get_training_status first and do not start work while another job runs.
"""
from models import TrainingStartRequest
from routes.training import start_training as start
request = TrainingStartRequest.model_validate(config)
# Pass via_api_key explicitly (a direct call leaves it a Depends object).
# MCP drives Studio like the UI session, so it coexists and frees VRAM.
return _dump(await start(request, current_subject = "mcp", via_api_key = False))
@mcp.tool
async def stop_training(save: bool = True) -> dict[str, Any]:
"""Ask the active training process to stop at its next safe checkpoint."""
from routes.training import TrainingStopRequest, stop_training as stop
return _dump(await stop(TrainingStopRequest(save = save), current_subject = "mcp"))
@mcp.tool
async def list_training_runs(limit: int = 50, offset: int = 0) -> dict[str, Any]:
"""List completed and stopped training runs, newest first."""
from routes.training_history import list_training_runs as list_runs
# Clamp here (direct call skips Query bounds); a negative LIMIT = no limit.
limit = _clamp(limit, 1, 200)
offset = max(0, offset)
return _dump(await list_runs(limit = limit, offset = offset, current_subject = "mcp"))
@mcp.tool
def validate_recipe(recipe: dict[str, Any]) -> dict[str, Any]:
"""Validate a Data Recipe with the same validator used by Studio."""
from models.data_recipe import RecipePayload
from routes.data_recipe.validate import validate
return _dump(validate(RecipePayload(recipe = recipe)))
@mcp.tool
def get_recipe_job_status(job_id: str) -> dict[str, Any]:
"""Read the status of a Data Recipe job."""
from routes.data_recipe.jobs import job_status
return _dump(job_status(job_id))
@mcp.tool
def get_recipe_job_dataset(
job_id: str,
limit: int = 20,
offset: int = 0,
) -> dict[str, Any]:
"""Read a bounded page of generated Data Recipe rows."""
from routes.data_recipe.jobs import job_dataset
# Clamp here (direct call skips FastAPI's Query bounds).
limit = _clamp(limit, 1, 500)
offset = max(0, offset)
return _dump(job_dataset(job_id, limit = limit, offset = offset))
@mcp.tool
async def load_checkpoint(
checkpoint_path: str,
max_seq_length: int = 2048,
load_in_4bit: bool = True,
trust_remote_code: bool = False,
approved_remote_code_fingerprint: str | None = None,
hf_token: str | None = None,
) -> dict[str, Any]:
"""Load a checkpoint into the export backend.
Export runs in its own subprocess and coexists with training and
inference; it does not unload them, so a load can fail with a clear
out-of-memory error if the GPU is already full. Pass hf_token to load a
gated checkpoint, and approved_remote_code_fingerprint to retry a
trust_remote_code load that was blocked pending review.
"""
from models import LoadCheckpointRequest
from routes.export import load_checkpoint as load
request = LoadCheckpointRequest(
checkpoint_path = checkpoint_path,
max_seq_length = max_seq_length,
load_in_4bit = load_in_4bit,
trust_remote_code = trust_remote_code,
approved_remote_code_fingerprint = approved_remote_code_fingerprint,
hf_token = hf_token,
)
return _dump(await load(request, current_subject = "mcp"))
@mcp.tool
async def export_gguf(
save_directory: str,
quantization_method: str | list[str] = "Q4_K_M",
push_to_hub: bool = False,
repo_id: str | None = None,
hf_token: str | None = None,
imatrix: bool = False,
imatrix_path: str | None = None,
) -> dict[str, Any]:
"""Export the loaded model to GGUF using Studio's existing path validation.
quantization_method may be a single method or a list to produce several
GGUFs from one load. Pass hf_token when push_to_hub is set (the backend
rejects a Hub upload without it). Set imatrix (or imatrix_path) for the
IQ low-bit quants that require an importance matrix.
"""
from models import ExportGGUFRequest
from routes.export import export_gguf as export
request = ExportGGUFRequest(
save_directory = save_directory,
quantization_method = quantization_method,
push_to_hub = push_to_hub,
repo_id = repo_id,
hf_token = hf_token,
imatrix = imatrix,
imatrix_path = imatrix_path,
)
return _dump(await export(request, current_subject = "mcp"))
return mcp
async def _gather_status(*coroutines: Any) -> tuple[Any, ...]:
"""Gather independent status calls without letting one optional backend fail all state."""
import asyncio
results = await asyncio.gather(*coroutines, return_exceptions = True)
return tuple(
{"error": str(result)} if isinstance(result, Exception) else result for result in results
)

View file

@ -20,6 +20,7 @@ from loggers import get_logger
import asyncio
import threading
import weakref
from contextlib import ExitStack
import re as _re
@ -28,6 +29,7 @@ import re as _re
from utils.models import extract_model_size_b as _extract_model_size_b
from utils.api_errors import openai_error_body, anthropic_error_body
from core.inference.orchestrator import GenStreamError, GenStreamErrorRaised
from core.inference.llama_admission import (
LlamaAdmissionCancelled,
LlamaAdmissionConfig,
@ -212,6 +214,14 @@ def _friendly_error(exc: Exception) -> str:
return "An internal error occurred"
def _friendly_gen_stream_error(value) -> str:
"""Return a client-safe message for typed local generation errors."""
text = str(value)
if getattr(value, "public", False):
return text
return safe_error_detail(RuntimeError(text), fallback = "An internal error occurred.")
def _friendly_upstream_error(text: str) -> str:
"""Rewrite a raw llama-server error body into an actionable message where we can.
@ -998,6 +1008,7 @@ try:
from core.inference.llama_server_args import (
_effective_tensor_parallel,
_tensor_parallel_matches_loaded,
extra_args_disable_mmproj,
parse_split_mode_override,
resolve_tensor_parallel,
strip_shadowing_flags,
@ -1035,6 +1046,7 @@ except ImportError:
from core.inference.llama_server_args import (
_effective_tensor_parallel,
_tensor_parallel_matches_loaded,
extra_args_disable_mmproj,
parse_split_mode_override,
resolve_tensor_parallel,
strip_shadowing_flags,
@ -1915,16 +1927,57 @@ async def artifact_preview_frame(allow_network: bool = False):
_BARE_JSON_NAME_MARKER_RE = _re.compile(r'\{\s*\\?"(?:name|function)\\?"\s*:')
def _detect_safetensors_features(backend, chat_template: Optional[str]) -> dict:
def _detect_safetensors_features(
backend,
chat_template: Optional[str],
tools = None,
) -> dict:
"""Classify reasoning/tool capabilities via the GGUF classifier so flags
match across backends. gpt-oss is overridden: Harmony routes reasoning and
tools through tokenizer channels, not template markup."""
model_id = getattr(backend, "active_model_name", None)
feature_template = chat_template
try:
from core.inference.chat_template_helpers import _selected_template_strings_from_value
selected_templates = _selected_template_strings_from_value(chat_template, tools)
if selected_templates:
feature_template = selected_templates[0]
except Exception:
logger.debug("safetensors_named_template_selection_failed", exc_info = True)
flags = detect_reasoning_flags(
chat_template,
feature_template,
model_identifier = model_id,
log_source = "safetensors",
)
if not flags.get("supports_reasoning"):
try:
from core.inference.chat_template_helpers import (
detect_reasoning_channel_markers_from_template,
)
templates = [chat_template]
models = getattr(backend, "models", None)
model_info = (
models.get(model_id, {})
if isinstance(models, dict) and model_id is not None
else {}
)
if isinstance(model_info, dict):
templates.extend(
(
model_info.get("native_chat_template"),
(model_info.get("chat_template_info") or {}).get("template"),
)
)
if any(
detect_reasoning_channel_markers_from_template(template, tools = tools) is not None
for template in templates
):
flags["supports_reasoning"] = True
flags["reasoning_always_on"] = True
logger.info("safetensors: model always reasons (native channel markers)")
except Exception:
logger.debug("safetensors_native_reasoning_marker_check_failed", exc_info = True)
# Markers any supported parser recognises (template advertises tools but
# uses none -> drop the pill). Reuse the parser's own signal list so this
# gate never drifts (a hand-maintained copy lost the DeepSeek variants);
@ -1938,9 +1991,9 @@ def _detect_safetensors_features(backend, chat_template: Optional[str]) -> dict:
)
if (
flags.get("supports_tools")
and chat_template
and not any(m in chat_template for m in _PARSER_MARKERS)
and not _BARE_JSON_NAME_MARKER_RE.search(chat_template)
and isinstance(feature_template, str)
and not any(m in feature_template for m in _PARSER_MARKERS)
and not _BARE_JSON_NAME_MARKER_RE.search(feature_template)
):
logger.info(
"safetensors: template advertises tools but uses an "
@ -3968,6 +4021,7 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
native_grant_backed = False
model_log_label = request.model_path
gguf_load_stack = ExitStack()
try:
# Validate user pass-through args up front so a managed-flag collision
# returns 400 before any model work.
@ -4187,15 +4241,9 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
llama_backend = get_llama_cpp_backend()
unsloth_backend = get_inference_backend()
# Unload any active Unsloth model to free VRAM (off the event loop:
# unload takes _gen_lock and can wait on an in-flight stream).
if unsloth_backend.active_model_name:
logger.info(
f"Unloading Unsloth model '{unsloth_backend.active_model_name}' before loading GGUF"
)
await asyncio.to_thread(
unsloth_backend.unload_model, unsloth_backend.active_model_name
)
if config.gguf_hf_repo:
from core.inference.llama_cpp import gguf_load_in_flight
gguf_load_stack.enter_context(gguf_load_in_flight(config.gguf_hf_repo))
# Inherit llama_extra_args from the previous load when the request
# omits the field (the chat-settings Apply path doesn't round-trip
@ -4275,6 +4323,38 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
extra_llama_args,
)
# Block cache writes that would race the download manager. This runs
# after pass-through argument inheritance so a carried --no-mmproj
# changes the companion requirement exactly as it does for the load.
if config.gguf_hf_repo:
from core.inference.llama_cpp import _hub_download_blocks_gguf_load
if await asyncio.to_thread(
_hub_download_blocks_gguf_load,
config.gguf_hf_repo,
config.gguf_variant,
require_mmproj = bool(
config.is_vision and not extra_args_disable_mmproj(extra_llama_args)
),
hf_token = request.hf_token,
):
raise HTTPException(
status_code = 409,
detail = (
f"'{model_log_label}' is currently being downloaded "
"by the download manager. Wait for the download to "
"finish (or cancel it), then load the model."
),
)
# Unload any active Unsloth model only after every hub conflict check.
if unsloth_backend.active_model_name:
logger.info(
f"Unloading Unsloth model '{unsloth_backend.active_model_name}' before loading GGUF"
)
await asyncio.to_thread(
unsloth_backend.unload_model, unsloth_backend.active_model_name
)
# Route to HF or local mode based on config. Run in a thread so the
# event loop stays free for progress polling and other requests
# during the (potentially long) GGUF download + llama-server start.
@ -4639,6 +4719,8 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
logger.error(f"Error loading model: {e}", exc_info = True)
msg = _maybe_unsupported_message(redacted_msg)
raise HTTPException(status_code = 500, detail = f"Failed to load model: {msg}")
finally:
gguf_load_stack.close()
def _requires_trust_remote_code_for_model(
@ -5392,6 +5474,10 @@ async def generate_stream(
if chunk is _DONE:
completed = True
break
if isinstance(chunk, GenStreamError):
yield f"data: {json.dumps({'error': _friendly_gen_stream_error(chunk)})}\n\n"
yield "data: [DONE]\n\n"
return
yield f"data: {json.dumps({'content': chunk})}\n\n"
if completed:
yield "data: [DONE]\n\n"
@ -5405,6 +5491,7 @@ async def generate_stream(
backend.reset_generation_state()
logger.error(f"Error during generation: {e}", exc_info = True)
yield f"data: {json.dumps({'error': _friendly_error(e)})}\n\n"
yield "data: [DONE]\n\n"
finally:
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
if not completed and not cancel_event.is_set():
@ -7028,6 +7115,13 @@ async def openai_chat_completions(
chunk_text = await asyncio.to_thread(next, gen, _DONE)
if chunk_text is _DONE:
break
if isinstance(chunk_text, GenStreamError):
_msg = _friendly_gen_stream_error(chunk_text)
api_monitor.fail(monitor_id, _msg)
yield _openai_stream_error_sse(
{"error": {"message": _msg, "type": "server_error"}}
)
return
if chunk_text:
api_monitor.append_reply(monitor_id, chunk_text)
yield _chat_content_chunk(
@ -7043,8 +7137,11 @@ async def openai_chat_completions(
raise
except Exception as e:
logger.error(f"Error during audio input streaming: {e}", exc_info = True)
api_monitor.fail(monitor_id, _friendly_error(e))
yield f"data: {json.dumps({'error': {'message': _friendly_error(e), 'type': 'server_error'}})}\n\n"
_msg = _friendly_error(e)
api_monitor.fail(monitor_id, _msg)
yield _openai_stream_error_sse(
{"error": {"message": _msg, "type": "server_error"}}
)
finally:
await _stop_local_disconnect_cancel_watcher(disconnect_watcher)
_tracker.__exit__(None, None, None)
@ -7061,7 +7158,15 @@ async def openai_chat_completions(
)
else:
try:
full_text = "".join(audio_input_generate())
full_text = ""
for chunk_text in audio_input_generate():
if isinstance(chunk_text, GenStreamError):
_msg = _friendly_gen_stream_error(chunk_text)
api_monitor.fail(monitor_id, _msg)
raise HTTPException(status_code = 500, detail = _msg)
full_text += chunk_text
except HTTPException:
raise
except Exception as e:
api_monitor.fail(monitor_id, _friendly_error(e))
raise
@ -8605,19 +8710,33 @@ async def openai_chat_completions(
# 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)
# GGUF parity: enable_thinking templates prefill an unclosed <think>; split into
# reasoning_content deltas so the UI renders the block for safetensors and MLX.
_sf_parse_think = bool(
_sf_features.get("supports_reasoning") or _sf_features.get("reasoning_always_on")
# Named templates may expose native reasoning only in their ``tool_use``
# branch. Use a truthy placeholder for Studio-managed tools, whose concrete
# schemas are selected below, and the request schemas for client passthrough.
_sf_server_tool_intent = bool(
_effective_enable_tools(payload) or _explicit_studio_tool_loop_requested(payload)
)
# Prefilled-open only for prefill styles with thinking on; gpt-oss uses the normal mode.
_sf_reasoning_prefilled = _sf_reasoning_prefill_mode(
_sf_features,
payload.enable_thinking,
_sf_tpl,
reasoning_effort = payload.reasoning_effort,
_sf_template_tools = payload.tools if payload.tool_choice != "none" else None
if not _sf_template_tools and _sf_server_tool_intent:
_sf_template_tools = ({},)
def _sf_response_protocol(tools = None):
features = _detect_safetensors_features(backend, _sf_tpl, tools = tools)
parse_think = bool(
features.get("supports_reasoning") or features.get("reasoning_always_on")
)
reasoning_prefilled = _sf_reasoning_prefill_mode(
features,
payload.enable_thinking,
_sf_tpl,
reasoning_effort = payload.reasoning_effort,
)
return features, parse_think, reasoning_prefilled
# GGUF parity: split canonical <think> output into reasoning_content. The
# selected template branch must match whether this request renders tools.
_sf_features, _sf_parse_think, _sf_reasoning_prefilled = _sf_response_protocol(
_sf_template_tools
)
def _new_sf_reasoning_extractor():
@ -8767,6 +8886,7 @@ async def openai_chat_completions(
permission_mode = payload.permission_mode,
use_adapter = payload.use_adapter,
stats_holder = _sf_stats_holder,
reasoning_prefilled = _sf_reasoning_prefilled,
)
_sf_tool_sentinel = object()
@ -8826,6 +8946,18 @@ async def openai_chat_completions(
_sf_next_task = None
if event is _sf_tool_sentinel:
break
if isinstance(event, GenStreamError):
backend.reset_generation_state()
_msg = _friendly_gen_stream_error(event)
api_monitor.fail(monitor_id, _msg)
yield _openai_stream_error_sse(
{"error": {"message": _msg, "type": "server_error"}}
)
return
if not isinstance(event, dict):
raise RuntimeError(
f"Invalid safetensors tool event: {type(event).__name__}"
)
if event["type"] == "heartbeat":
# Tool-execution wrapper heartbeat -> SSE keepalive.
@ -8913,6 +9045,11 @@ async def openai_chat_completions(
backend.reset_generation_state()
api_monitor.finish(monitor_id, "cancelled")
raise
except GenStreamErrorRaised as exc:
backend.reset_generation_state()
_msg = _friendly_gen_stream_error(exc)
api_monitor.fail(monitor_id, _msg)
yield _openai_stream_error_sse({"error": {"message": _msg, "type": "server_error"}})
except Exception:
backend.reset_generation_state()
# Generic wire message; full trace stays in the log (CWE-209:
@ -8962,6 +9099,15 @@ async def openai_chat_completions(
for event in gen:
if cancel_event.is_set():
break
if isinstance(event, GenStreamError):
raise HTTPException(
status_code = 500,
detail = _friendly_gen_stream_error(event),
)
if not isinstance(event, dict):
raise RuntimeError(
f"Invalid safetensors tool event: {type(event).__name__}"
)
if event.get("type") == "content":
full_text = _strip_tool_xml_for_display(
event.get("text", ""),
@ -9002,6 +9148,15 @@ async def openai_chat_completions(
backend.reset_generation_state()
api_monitor.finish(monitor_id, "cancelled")
raise
except GenStreamErrorRaised as exc:
backend.reset_generation_state()
_msg = _friendly_gen_stream_error(exc)
api_monitor.fail(monitor_id, _msg)
raise HTTPException(status_code = 500, detail = _msg)
except HTTPException as exc:
backend.reset_generation_state()
api_monitor.fail(monitor_id, str(exc.detail))
raise
except Exception:
backend.reset_generation_state()
# CWE-209: generic detail; full trace in log.
@ -9088,6 +9243,12 @@ async def openai_chat_completions(
else:
gen_kwargs["tools"] = payload.tools
# The potential tool context above is needed before server/client routing is
# known. This standard path now has the exact schemas that will be rendered,
# so resolve reasoning parsing again to keep empty registries, forced-tool
# misses, and tool_choice="none" on the marker-free template branch.
_, _sf_parse_think, _sf_reasoning_prefilled = _sf_response_protocol(gen_kwargs.get("tools"))
# Request-scoped usage/timings receptacle (filled at gen_done).
stats_holder: dict = {}
@ -9168,6 +9329,14 @@ async def openai_chat_completions(
_next_task = None
if cumulative is _DONE:
break
if isinstance(cumulative, GenStreamError):
backend.reset_generation_state()
_msg = _friendly_gen_stream_error(cumulative)
api_monitor.fail(monitor_id, _msg)
yield _openai_stream_error_sse(
{"error": {"message": _msg, "type": "server_error"}}
)
return
if await request.is_disconnected():
cancel_event.set()
backend.reset_generation_state()
@ -9317,6 +9486,11 @@ async def openai_chat_completions(
try:
full_text = ""
for token in generate():
if isinstance(token, GenStreamError):
backend.reset_generation_state()
_msg = _friendly_gen_stream_error(token)
api_monitor.fail(monitor_id, _msg)
raise HTTPException(status_code = 500, detail = _msg)
full_text = token
# Split prefilled <think> reasoning (GGUF parity); also covers MLX via
@ -9415,6 +9589,8 @@ async def openai_chat_completions(
api_monitor.finish(monitor_id)
return _model_json_response(response)
except HTTPException:
raise
except Exception as e:
backend.reset_generation_state()
logger.error(f"Error during OpenAI completion: {e}", exc_info = True)

View file

@ -1354,11 +1354,23 @@ def run_server(
if secure:
os.environ["UNSLOTH_SECURE"] = "1"
import nest_asyncio
nest_asyncio.apply()
import asyncio
# nest_asyncio is for Colab/IPython, where the main thread already runs a loop
# the blocking waits below would collide with. Apply it only with a loop running
# (a plain CLI start has nothing to nest) and only on Python <= 3.13: on 3.14+
# its global Task patch leaves asyncio.current_task() None (tracking moved into
# C), which also breaks the background uvicorn loop and 500s every request. It
# is archived upstream, so no 3.14 fix is coming; skip it there.
if sys.version_info < (3, 14):
try:
asyncio.get_running_loop()
except RuntimeError:
pass
else:
import nest_asyncio
nest_asyncio.apply()
from threading import Thread, Event
import uvicorn

View file

@ -0,0 +1,740 @@
# 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 cached GGUF reuse and load/download exclusion.
No GPU, network, or subprocesses are required.
"""
from __future__ import annotations
import asyncio
import sys
import threading
import types as _types
from pathlib import Path
from unittest.mock import patch
import pytest
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
if _BACKEND_DIR not in sys.path:
sys.path.insert(0, _BACKEND_DIR)
# Stub optional dependencies before importing the modules under test.
_loggers_stub = _types.ModuleType("loggers")
_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
sys.modules.setdefault("loggers", _loggers_stub)
_structlog_stub = _types.ModuleType("structlog")
sys.modules.setdefault("structlog", _structlog_stub)
try:
import httpx # noqa: F401
except ImportError:
_httpx_stub = _types.ModuleType("httpx")
for _exc_name in (
"ConnectError",
"TimeoutException",
"ReadTimeout",
"ReadError",
"RemoteProtocolError",
"CloseError",
"HTTPError",
"RequestError",
"HTTPStatusError",
):
setattr(_httpx_stub, _exc_name, type(_exc_name, (Exception,), {}))
_httpx_stub.Response = type("Response", (), {})
_httpx_stub.Request = type("Request", (), {})
class _FakeTimeout:
def __init__(self, *a, **kw):
pass
_httpx_stub.Timeout = _FakeTimeout
_httpx_stub.Client = type(
"Client",
(),
{
"__init__": lambda self, **kw: None,
"__enter__": lambda self: self,
"__exit__": lambda self, *a: None,
},
)
sys.modules.setdefault("httpx", _httpx_stub)
from huggingface_hub import constants as hf_constants
from core.inference.llama_cpp import (
LlamaCppBackend,
cached_gguf_for_load,
gguf_load_in_flight,
hf_gguf_load_in_flight,
)
REPO = "unsloth/gemma-test-GGUF"
VARIANT = "UD-Q4_K_XL"
MAIN = f"gemma-test-{VARIANT}.gguf"
def _build_cache(
root: Path,
repo_id: str,
files: dict[str, int],
*,
snapshot_sha: str = "a" * 40,
) -> Path:
"""Create ``$root/models--<repo>/snapshots/<sha>/<rel>`` for each entry."""
repo_dir = root / f"models--{repo_id.replace('/', '--')}"
(repo_dir / "blobs").mkdir(parents = True, exist_ok = True)
snap = repo_dir / "snapshots" / snapshot_sha
snap.mkdir(parents = True, exist_ok = True)
for rel, size in files.items():
full = snap / rel
full.parent.mkdir(parents = True, exist_ok = True)
full.write_bytes(b"\0" * size)
return snap
@pytest.fixture
def hf_cache(tmp_path, monkeypatch):
monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising = False)
return tmp_path
def _fail_download(*_args, **_kwargs):
raise AssertionError("must reuse the cached GGUF instead of downloading")
def _fail_get_paths_info(*_args, **_kwargs):
raise AssertionError("cached reuse must return before the sizing preflight")
class TestLoadReusesCachedCopy:
def test_online_reuse_after_revision_bump(self, hf_cache):
"""A new repo revision does not replace a complete cached model."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / MAIN)
def test_reuse_size_check_uses_cached_snapshot_revision(self, hf_cache):
"""Current-revision size changes do not invalidate an older complete copy."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
revisions: list[str | None] = []
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
revisions.append(revision)
size = 4 if revision == snap.name else 8
return [_types.SimpleNamespace(path = path, size = size) for path in paths]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / MAIN)
assert revisions == [snap.name]
def test_reuse_when_cached_revision_vanished_from_hub(self, hf_cache):
"""The Hub answers an unknown revision with an empty result, not an error."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", lambda *_a, **_k: []),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / MAIN)
def test_truncated_cached_file_is_not_reused(self, hf_cache):
backend = LlamaCppBackend()
_build_cache(hf_cache, REPO, {MAIN: 4})
downloaded: list[str] = []
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
return [_types.SimpleNamespace(path = path, size = 8) for path in paths]
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [MAIN]
assert out == f"/fake/{REPO}/{MAIN}"
def test_truncated_cached_split_shard_is_not_reused(self, hf_cache):
backend = LlamaCppBackend()
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
_build_cache(hf_cache, REPO, {shard1: 8, shard2: 4})
downloaded: list[str] = []
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
return [_types.SimpleNamespace(path = path, size = 8) for path in paths]
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [shard1, shard2]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [shard1, shard2]
assert out == f"/fake/{REPO}/{shard1}"
def test_online_reuse_when_reupload_renamed_the_file(self, hf_cache):
"""A renamed variant still reuses its cached file."""
backend = LlamaCppBackend()
old_name = f"gemma-test-old-{VARIANT}.gguf"
snap = _build_cache(hf_cache, REPO, {old_name: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / old_name)
def test_downloads_when_nothing_cached(self, hf_cache):
backend = LlamaCppBackend()
downloaded: list[str] = []
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 1) for p in paths if p is not None]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [MAIN]
assert out == f"/fake/{REPO}/{MAIN}"
def test_force_redownloads_despite_cache(self, hf_cache):
"""A forced download ignores a complete cached copy."""
backend = LlamaCppBackend()
_build_cache(hf_cache, REPO, {MAIN: 4})
downloaded: list[str] = []
def fake_download(
repo_id,
filename,
token = None,
**kwargs,
):
assert kwargs.get("force_download") is True
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 1) for p in paths if p is not None]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT, force = True)
assert downloaded == [MAIN]
assert out == f"/fake/{REPO}/{MAIN}"
def test_split_reused_only_when_colocated(self, hf_cache):
backend = LlamaCppBackend()
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
snap = _build_cache(hf_cache, REPO, {shard1: 4, shard2: 4})
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [shard1, shard2]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / shard1)
def test_partial_split_set_downloads(self, hf_cache):
"""A partial split set is not reused."""
backend = LlamaCppBackend()
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
_build_cache(hf_cache, REPO, {shard1: 4})
downloaded: list[str] = []
def fake_download(
repo_id,
filename,
token = None,
**_kwargs,
):
downloaded.append(filename)
return f"/fake/{repo_id}/{filename}"
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 4) for p in paths if p is not None]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [shard1, shard2]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert downloaded == [shard1, shard2]
assert out == f"/fake/{REPO}/{shard1}"
def test_reuse_prefers_newest_snapshot_after_update(self, hf_cache):
"""Loads prefer the newest complete snapshot."""
import os
backend = LlamaCppBackend()
old_snap = _build_cache(hf_cache, REPO, {MAIN: 4}, snapshot_sha = "a" * 40)
new_snap = _build_cache(hf_cache, REPO, {MAIN: 6}, snapshot_sha = "b" * 40)
os.utime(old_snap, (1_000_000, 1_000_000))
os.utime(new_snap, (2_000_000, 2_000_000))
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", _fail_get_paths_info),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(new_snap / MAIN)
def test_low_disk_fallback_reuses_cached_copy(self, hf_cache):
backend = LlamaCppBackend()
fallback = "gemma-test-Q2_K.gguf"
snap = _build_cache(hf_cache, REPO, {fallback: 4})
def fake_get_paths_info(
_repo,
paths,
*,
revision = None,
token = None,
):
size = 4 if revision == snap.name else 100
return [_types.SimpleNamespace(path = path, size = size) for path in paths]
with (
patch("huggingface_hub.list_repo_files", lambda *_a, **_k: [MAIN]),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("shutil.disk_usage", lambda *_a, **_k: _types.SimpleNamespace(free = 10)),
patch.object(
backend,
"_find_smallest_fitting_variant",
lambda *_a, **_k: (fallback, 4, []),
),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_gguf(hf_repo = REPO, hf_variant = VARIANT)
assert out == str(snap / fallback)
def test_companion_prefers_main_snapshot_sibling(self, hf_cache):
"""A cached mmproj is reused from the main model's snapshot."""
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4, "mmproj-F16.gguf": 2})
def _fail_list(*_args, **_kwargs):
raise AssertionError("snapshot sibling must resolve without a repo listing")
with patch("huggingface_hub.list_repo_files", _fail_list):
out = backend._download_mmproj(hf_repo = REPO, near_path = str(snap / MAIN))
assert out == str(snap / "mmproj-F16.gguf")
def test_companion_finds_snapshot_through_hf_symlink(self, hf_cache):
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {})
blobs = snap.parent.parent / "blobs"
main_blob = blobs / "main"
mmproj_blob = blobs / "mmproj"
main_blob.write_bytes(b"main")
mmproj_blob.write_bytes(b"mmproj")
try:
(snap / MAIN).symlink_to(main_blob)
(snap / "mmproj-F16.gguf").symlink_to(mmproj_blob)
except OSError as exc:
pytest.skip(f"symlinks unavailable: {exc}")
with patch("huggingface_hub.list_repo_files", _fail_download):
out = backend._download_mmproj(hf_repo = REPO, near_path = str(snap / MAIN))
assert out == str(snap / "mmproj-F16.gguf")
def test_companion_does_not_download_during_hub_job(self, hf_cache):
backend = LlamaCppBackend()
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
registry = _types.SimpleNamespace(active_job_refs = lambda _repo: [object()])
with (
patch("huggingface_hub.list_repo_files", _fail_download),
patch("hub.utils.download_registry.get_models_registry", lambda: registry),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", _fail_download),
):
out = backend._download_mmproj(hf_repo = REPO, near_path = str(snap / MAIN))
assert out is None
class TestCachedGgufForLoadProbe:
def test_complete_copy_found(self, hf_cache):
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
assert cached_gguf_for_load(REPO, VARIANT) == str(snap / MAIN)
def test_absent_copy_is_none(self, hf_cache):
assert cached_gguf_for_load(REPO, VARIANT) is None
def test_partial_split_is_none(self, hf_cache):
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
_build_cache(hf_cache, REPO, {shard1: 4})
assert cached_gguf_for_load(REPO, VARIANT) is None
def test_partial_new_snapshot_does_not_hide_complete_split(self, hf_cache):
import os
shard1 = f"gemma-test-{VARIANT}-00001-of-00002.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00002.gguf"
old = _build_cache(
hf_cache,
REPO,
{shard1: 4, shard2: 4},
snapshot_sha = "a" * 40,
)
new = _build_cache(hf_cache, REPO, {shard1: 4}, snapshot_sha = "b" * 40)
os.utime(old, (1_000_000, 1_000_000))
os.utime(new, (2_000_000, 2_000_000))
assert cached_gguf_for_load(REPO, VARIANT) == str(old / shard1)
def test_split_requires_every_declared_shard(self, hf_cache):
shard1 = f"gemma-test-{VARIANT}-00001-of-00003.gguf"
shard2 = f"gemma-test-{VARIANT}-00002-of-00003.gguf"
_build_cache(hf_cache, REPO, {shard1: 4, shard2: 4})
assert cached_gguf_for_load(REPO, VARIANT) is None
def test_required_mmproj_must_share_main_snapshot(self, hf_cache):
snap = _build_cache(hf_cache, REPO, {MAIN: 4})
assert cached_gguf_for_load(REPO, VARIANT) == str(snap / MAIN)
assert cached_gguf_for_load(REPO, VARIANT, require_mmproj = True) is None
(snap / "mmproj-F16.gguf").write_bytes(b"mmproj")
assert cached_gguf_for_load(REPO, VARIANT, require_mmproj = True) == str(snap / MAIN)
def test_required_mmproj_scans_past_newer_main_only_snapshot(self, hf_cache):
import os
old = _build_cache(
hf_cache,
REPO,
{MAIN: 4, "mmproj-F16.gguf": 2},
snapshot_sha = "a" * 40,
)
new = _build_cache(hf_cache, REPO, {MAIN: 4}, snapshot_sha = "b" * 40)
os.utime(old, (1_000_000, 1_000_000))
os.utime(new, (2_000_000, 2_000_000))
assert cached_gguf_for_load(REPO, VARIANT, require_mmproj = True) == str(old / MAIN)
class TestLoadHubDownloadExclusion:
def test_in_flight_marker_counts_and_normalizes_case(self):
assert not hf_gguf_load_in_flight(REPO)
with gguf_load_in_flight(REPO):
assert hf_gguf_load_in_flight(REPO.upper())
with gguf_load_in_flight(REPO.lower()):
assert hf_gguf_load_in_flight(REPO)
assert hf_gguf_load_in_flight(REPO)
assert not hf_gguf_load_in_flight(REPO)
def test_marker_noops_for_local_loads(self):
with gguf_load_in_flight(None):
assert not hf_gguf_load_in_flight("")
def test_marker_cleared_on_exception(self):
with pytest.raises(RuntimeError):
with gguf_load_in_flight(REPO):
raise RuntimeError("boom")
assert not hf_gguf_load_in_flight(REPO)
def test_hub_download_refused_while_load_in_flight(self):
from fastapi import HTTPException
from hub.schemas.downloads import DownloadModelRequest
from hub.services.models import downloads as dl
body = DownloadModelRequest(repo_id = REPO, gguf_variant = VARIANT)
with (
patch.object(dl, "resolve_cached_repo_id_case", lambda repo_id, repo_type: repo_id),
gguf_load_in_flight(REPO),
):
with pytest.raises(HTTPException) as exc_info:
asyncio.run(dl.download_model_response(body))
assert exc_info.value.status_code == 409
assert "load" in exc_info.value.detail.lower()
def test_hub_download_rechecks_marker_before_claim(self):
from fastapi import HTTPException
from hub.schemas.downloads import DownloadModelRequest
from hub.services.models import downloads as dl
scope = None
def mark_load(*_args, **_kwargs):
nonlocal scope
if scope is None:
scope = gguf_load_in_flight(REPO)
scope.__enter__()
return frozenset()
class _Registry:
def claim(self, *_args, admission_check, **_kwargs):
assert admission_check() is False
return False, "admission_blocked"
def current_generation(self, _key):
return 0
registry = _Registry()
body = DownloadModelRequest(repo_id = REPO, gguf_variant = VARIANT)
try:
with (
patch.object(dl, "resolve_cached_repo_id_case", lambda repo_id, repo_type: repo_id),
patch.object(dl.gguf_variants, "gguf_variant_blob_hashes", mark_load),
patch.object(dl, "_registry", registry),
):
with pytest.raises(HTTPException) as exc_info:
asyncio.run(dl.download_model_response(body))
finally:
if scope is not None:
scope.__exit__(None, None, None)
assert exc_info.value.status_code == 409
def test_registry_admission_check_prevents_claim(self):
from hub.utils.download_registry import DownloadRegistry, TRANSPORT_HTTP
registry = DownloadRegistry()
claimed, state = registry.claim(
f"{REPO}::{VARIANT}",
TRANSPORT_HTTP,
repo_type = "model",
repo_id = REPO,
variant = VARIANT,
admission_check = lambda: False,
)
assert claimed is False
assert state == "admission_blocked"
assert registry.active_jobs(REPO) == {}
def test_same_variant_job_stays_visible_during_retry_handoff(self):
from hub.utils.download_registry import DownloadRegistry, TRANSPORT_XET
from core.inference.llama_cpp import _hub_download_blocks_gguf_load
registry = DownloadRegistry()
key = f"{REPO}::{VARIANT}"
claimed, _ = registry.claim(
key,
TRANSPORT_XET,
repo_type = "model",
repo_id = REPO,
variant = VARIANT,
)
assert claimed is True
assert registry.has_active_variant(REPO, VARIANT.lower()) is True
registry.release_active_slot(key)
assert registry.active_jobs(REPO) == {}
assert registry.active_job_refs(REPO)
assert registry.has_active_variant(REPO, VARIANT) is True
with (
patch("hub.utils.download_registry.get_models_registry", lambda: registry),
patch(
"core.inference.llama_cpp.cached_gguf_for_load",
side_effect = AssertionError("same-variant jobs must block before cache reuse"),
),
):
assert _hub_download_blocks_gguf_load(REPO, VARIANT) is True
registry.set_job(key, "complete")
assert registry.has_active_variant(REPO, VARIANT) is False
def test_other_variant_job_still_allows_complete_cached_load(self):
from core.inference.llama_cpp import _hub_download_blocks_gguf_load
from hub.utils.download_registry import DownloadRegistry, TRANSPORT_HTTP
registry = DownloadRegistry()
registry.claim(
f"{REPO}::Q8_0",
TRANSPORT_HTTP,
repo_type = "model",
repo_id = REPO,
variant = "Q8_0",
)
with (
patch("hub.utils.download_registry.get_models_registry", lambda: registry),
patch(
"core.inference.llama_cpp.cached_gguf_for_load",
return_value = "/cached/model.gguf",
) as cached_probe,
):
assert _hub_download_blocks_gguf_load(REPO, VARIANT) is False
cached_probe.assert_called_once_with(
REPO,
VARIANT,
require_mmproj = False,
verify_sizes = True,
hf_token = None,
)
def test_cancelled_request_keeps_marker_until_load_thread_finishes(self):
from core.inference.llama_cpp import _with_gguf_load_marker
started = threading.Event()
release = threading.Event()
finished = threading.Event()
class FakeBackend:
@_with_gguf_load_marker
def load_model(self, *, hf_repo):
started.set()
release.wait(timeout = 2)
finished.set()
return True
async def scenario():
with patch(
"core.inference.llama_cpp._hub_download_blocks_gguf_load",
return_value = False,
):
task = asyncio.create_task(
asyncio.to_thread(FakeBackend().load_model, hf_repo = REPO)
)
assert await asyncio.to_thread(started.wait, 1)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert hf_gguf_load_in_flight(REPO)
release.set()
assert await asyncio.to_thread(finished.wait, 1)
for _ in range(100):
if not hf_gguf_load_in_flight(REPO):
break
await asyncio.sleep(0.001)
assert not hf_gguf_load_in_flight(REPO)
asyncio.run(scenario())
def test_load_marker_precedes_hub_guard_and_unload(self):
source = (Path(__file__).resolve().parent.parent / "routes" / "inference.py").read_text()
gguf_branch = source[source.index("if config.is_gguf:") :]
assert (
gguf_branch.index("enter_context(gguf_load_in_flight")
< gguf_branch.index("if request.llama_extra_args is None")
< gguf_branch.index("_hub_download_blocks_gguf_load")
< gguf_branch.index("unsloth_backend.unload_model")
)
llama_source = (
Path(__file__).resolve().parent.parent / "core" / "inference" / "llama_cpp.py"
).read_text()
assert "@_with_gguf_load_marker\n def load_model(" in llama_source

View file

@ -0,0 +1,290 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import asyncio
import sys
import types
import pytest
from mcp_server import BearerTokenMiddleware, _clamp, _dump, create_studio_mcp
def _get_tool(name):
tools = asyncio.run(create_studio_mcp().list_tools())
return {tool.name: tool for tool in tools}[name]
def test_studio_mcp_registers_control_plane_tools():
tools = asyncio.run(create_studio_mcp().list_tools())
assert {tool.name for tool in tools} == {
"studio_status",
"list_local_models",
"get_training_status",
"start_training",
"stop_training",
"list_training_runs",
"validate_recipe",
"get_recipe_job_status",
"get_recipe_job_dataset",
"load_checkpoint",
"export_gguf",
}
def test_dump_serializes_pydantic_values():
class Response:
def model_dump(self, *, mode):
assert mode == "json"
return {"ok": True}
assert _dump(Response()) == {"ok": True}
assert _dump({"already": "json"}) == {"already": "json"}
def test_bearer_token_middleware_rejects_wrong_token():
events = []
async def app(scope, receive, send):
events.append("app")
async def send(message):
events.append(message)
middleware = BearerTokenMiddleware(app, "secret")
asyncio.run(
middleware(
{"type": "http", "headers": [(b"authorization", b"Bearer wrong")]},
None,
send,
)
)
assert events[0]["status"] == 401
assert "app" not in events
def test_bearer_token_middleware_closes_unauthorized_websocket():
events = []
async def app(scope, receive, send):
events.append("app")
async def send(message):
events.append(message)
middleware = BearerTokenMiddleware(app, "secret")
asyncio.run(
middleware(
{"type": "websocket", "headers": []},
None,
send,
)
)
assert events == [{"type": "websocket.close", "code": 4401}]
def test_bearer_token_middleware_rejects_non_ascii_authorization():
# A non-ASCII bearer value must produce a clean 401, not a 500. Comparing on
# bytes avoids the str hmac.compare_digest TypeError on non-ASCII input.
events = []
async def app(scope, receive, send):
events.append("app")
async def send(message):
events.append(message)
middleware = BearerTokenMiddleware(app, "secret")
asyncio.run(
middleware(
{"type": "http", "headers": [(b"authorization", b"Bearer \xff\xff")]},
None,
send,
)
)
assert events[0]["status"] == 401
assert "app" not in events
def test_bearer_token_middleware_accepts_correct_token():
events = []
async def app(scope, receive, send):
events.append("app")
async def send(message):
events.append(message)
middleware = BearerTokenMiddleware(app, "secret")
asyncio.run(
middleware(
{"type": "http", "headers": [(b"authorization", b"Bearer secret")]},
None,
send,
)
)
assert events == ["app"]
def test_bearer_token_middleware_requires_non_empty_token():
async def app(scope, receive, send):
pass
for bad in ("", " "):
with pytest.raises(ValueError):
BearerTokenMiddleware(app, bad)
def test_bearer_token_middleware_rejects_non_ascii_token():
async def app(scope, receive, send):
pass
# non-ASCII tokens cannot be transmitted in an HTTP header by a standard
# client, so they are rejected at construction instead of locking out.
for bad in ("töken", "\U0001f600"):
with pytest.raises(ValueError):
BearerTokenMiddleware(app, bad)
def test_bearer_token_middleware_passes_through_non_http_scopes():
events = []
async def app(scope, receive, send):
events.append("app")
async def send(message):
events.append(message)
middleware = BearerTokenMiddleware(app, "secret")
asyncio.run(middleware({"type": "lifespan"}, None, send))
assert events == ["app"]
def test_clamp_restricts_to_inclusive_bounds():
assert _clamp(5, 1, 200) == 5
assert _clamp(-10, 1, 200) == 1
assert _clamp(10_000, 1, 200) == 200
assert _clamp(0, 1, 500) == 1
assert _clamp(1_000, 1, 500) == 500
def test_export_and_checkpoint_tools_expose_forwarded_fields():
export_props = set(_get_tool("export_gguf").parameters["properties"])
assert {"hf_token", "imatrix", "imatrix_path"} <= export_props
checkpoint_props = set(_get_tool("load_checkpoint").parameters["properties"])
assert {"hf_token", "approved_remote_code_fingerprint"} <= checkpoint_props
def _stub_module(monkeypatch, name, **attrs):
module = types.ModuleType(name)
for key, value in attrs.items():
setattr(module, key, value)
if "." in name:
module.__path__ = [] # mark package-like so submodule imports resolve
monkeypatch.setitem(sys.modules, name, module)
return module
def test_export_gguf_forwards_hf_token_and_imatrix(monkeypatch):
captured = {}
class FakeExportGGUFRequest:
def __init__(self, **kwargs):
captured.update(kwargs)
async def fake_export(request, current_subject):
return {"current_subject": current_subject}
_stub_module(monkeypatch, "models", ExportGGUFRequest = FakeExportGGUFRequest)
_stub_module(monkeypatch, "routes")
_stub_module(monkeypatch, "routes.export", export_gguf = fake_export)
tool = _get_tool("export_gguf")
result = asyncio.run(
tool.fn(
save_directory = "/tmp/out",
quantization_method = ["Q4_K_M", "Q8_0"],
push_to_hub = True,
repo_id = "me/model",
hf_token = "hf_secret",
imatrix = True,
imatrix_path = "/tmp/imatrix.dat",
)
)
assert captured["hf_token"] == "hf_secret"
assert captured["imatrix"] is True
assert captured["imatrix_path"] == "/tmp/imatrix.dat"
assert captured["quantization_method"] == ["Q4_K_M", "Q8_0"]
assert result["current_subject"] == "mcp"
def test_load_checkpoint_forwards_token_and_fingerprint(monkeypatch):
captured = {}
class FakeLoadCheckpointRequest:
def __init__(self, **kwargs):
captured.update(kwargs)
async def fake_load(request, current_subject):
return {"current_subject": current_subject}
_stub_module(monkeypatch, "models", LoadCheckpointRequest = FakeLoadCheckpointRequest)
_stub_module(monkeypatch, "routes")
_stub_module(monkeypatch, "routes.export", load_checkpoint = fake_load)
tool = _get_tool("load_checkpoint")
asyncio.run(
tool.fn(
checkpoint_path = "/tmp/ckpt",
approved_remote_code_fingerprint = "sha256:abc",
hf_token = "hf_secret",
)
)
assert captured["hf_token"] == "hf_secret"
assert captured["approved_remote_code_fingerprint"] == "sha256:abc"
def test_list_training_runs_clamps_pagination(monkeypatch):
captured = {}
async def fake_list_runs(limit, offset, current_subject):
captured["limit"] = limit
captured["offset"] = offset
return {"ok": True}
_stub_module(monkeypatch, "routes")
_stub_module(monkeypatch, "routes.training_history", list_training_runs = fake_list_runs)
tool = _get_tool("list_training_runs")
asyncio.run(tool.fn(limit = 10_000, offset = -5))
assert captured["limit"] == 200
assert captured["offset"] == 0
def test_get_recipe_job_dataset_clamps_pagination(monkeypatch):
captured = {}
def fake_job_dataset(job_id, limit, offset):
captured["limit"] = limit
captured["offset"] = offset
return {"ok": True}
_stub_module(monkeypatch, "routes")
_stub_module(monkeypatch, "routes.data_recipe")
_stub_module(monkeypatch, "routes.data_recipe.jobs", job_dataset = fake_job_dataset)
tool = _get_tool("get_recipe_job_dataset") # this tool is synchronous
tool.fn(job_id = "job-1", limit = -1, offset = -9)
assert captured["limit"] == 1
assert captured["offset"] == 0

View file

@ -523,3 +523,169 @@ def test_mlx_generate_text_forwards_kwargs_into_template_helper(monkeypatch):
assert render["kwargs"]["enable_thinking"] is True
assert render["kwargs"]["reasoning_effort"] == "medium"
assert render["kwargs"]["preserve_thinking"] is True
def test_mlx_text_normalizes_native_reasoning_and_close_releases_lock(monkeypatch):
_install_fake_mlx(monkeypatch)
from core.inference.mlx_inference import MLXInferenceBackend
monkeypatch.setattr(
"core.inference.chat_template_helpers.apply_chat_template_for_generation",
lambda *_args, **_kwargs: "prompt",
raising = True,
)
monkeypatch.setattr(
"core.inference.chat_template_helpers.render_with_native_template_fallback",
lambda formatted_prompt, **_kwargs: SimpleNamespace(
prompt = formatted_prompt,
reasoning_channel_markers = ("<|channel>thought\n", "<channel|>"),
),
raising = True,
)
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, text, tok):
self.text = text
self.token = tok
def _stream_generate(_model, _tokenizer, **_kw):
yield _Resp("<|channel>thought\n", 10)
yield _Resp("r", 11)
yield _Resp("<channel|>", 12)
yield _Resp("a", 13)
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)
backend = MLXInferenceBackend()
backend._model = object()
backend._tokenizer = SimpleNamespace(all_special_tokens = [])
backend._is_vlm = False
assert list(
backend.generate_chat_response(
messages = [{"role": "user", "content": "ping"}],
max_new_tokens = 4,
)
) == ["<think>", "<think>r", "<think>r</think>", "<think>r</think>a"]
gen = backend.generate_chat_response(
messages = [{"role": "user", "content": "ping"}],
max_new_tokens = 4,
)
assert next(gen) == "<think>"
assert backend._generation_lock.locked()
gen.close()
assert not backend._generation_lock.locked()
def test_mlx_text_native_metadata_preserves_prefilled_think_snapshots(monkeypatch):
_install_fake_mlx(monkeypatch)
from core.inference.mlx_inference import MLXInferenceBackend
monkeypatch.setattr(
"core.inference.chat_template_helpers.apply_chat_template_for_generation",
lambda *_args, **_kwargs: "prompt<think>\n",
raising = True,
)
monkeypatch.setattr(
"core.inference.chat_template_helpers.render_with_native_template_fallback",
lambda formatted_prompt, **_kwargs: SimpleNamespace(
prompt = formatted_prompt,
reasoning_channel_markers = ("<|channel>thought", "<channel|>"),
),
raising = True,
)
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, text, tok):
self.text = text
self.token = tok
def _stream_generate(_model, _tokenizer, **_kw):
yield _Resp("reason", 10)
yield _Resp("</think>", 11)
yield _Resp("answer", 12)
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)
backend = MLXInferenceBackend()
backend._model = object()
backend._tokenizer = SimpleNamespace(all_special_tokens = [])
backend._is_vlm = False
snapshots = list(
backend.generate_chat_response(
messages = [{"role": "user", "content": "ping"}],
max_new_tokens = 3,
)
)
assert snapshots == [
"<think>\n",
"<think>\nreason",
"<think>\nreason</think>",
"<think>\nreason</think>answer",
]
assert all(current.startswith(previous) for previous, current in zip(snapshots, snapshots[1:]))
def test_mlx_vlm_normalizes_native_reasoning_channels(monkeypatch):
_install_fake_mlx(monkeypatch)
from core.inference.mlx_inference import MLXInferenceBackend
monkeypatch.setattr(
"core.inference.chat_template_helpers.apply_chat_template_for_generation",
lambda *_args, **_kwargs: "prompt",
raising = True,
)
mlx_vlm_pkg = types.ModuleType("mlx_vlm")
class _Resp:
def __init__(self, text, tok):
self.text = text
self.token = tok
def _stream_generate(_model, _processor, _prompt, _images, **_kw):
yield _Resp("<|channel>thought\n", 10)
yield _Resp("vision", 11)
yield _Resp("<channel|>", 12)
yield _Resp(" answer", 13)
mlx_vlm_pkg.stream_generate = _stream_generate
monkeypatch.setitem(sys.modules, "mlx_vlm", mlx_vlm_pkg)
backend = MLXInferenceBackend()
backend._model = SimpleNamespace(config = SimpleNamespace())
backend._processor = SimpleNamespace(
chat_template = "<|channel>thought\n...<channel|>",
all_special_tokens = [],
apply_chat_template = lambda *_args, **_kwargs: "prompt",
)
backend._is_vlm = True
assert list(
backend.generate_chat_response(
messages = [{"role": "user", "content": "describe"}],
image = object(),
max_new_tokens = 4,
)
) == [
"<think>",
"<think>vision",
"<think>vision</think>",
"<think>vision</think> answer",
]

View file

@ -328,6 +328,7 @@ def test_download_mtp_prefers_root_over_new_scheme_copies(monkeypatch):
pick,
label,
cancel_event = None,
near_path = None,
):
captured["pick"] = pick
return None
@ -428,6 +429,32 @@ def test_download_mtp_reuse_follows_snapshot_order_offline(tmp_path, monkeypatch
assert got is not None and Path(got).parent.parent.name == "newest"
def test_download_mtp_prefers_main_snapshot_offline(tmp_path, monkeypatch):
import utils.models.model_config as mc
from core.inference.llama_cpp import LlamaCppBackend
monkeypatch.setenv("HF_HUB_OFFLINE", "1")
snapshots = tmp_path / "models--unsloth--gemma" / "snapshots"
old = snapshots / "old"
new = snapshots / "new"
old.mkdir(parents = True)
new.mkdir(parents = True)
main = old / "gemma-UD-Q4_K_XL.gguf"
old_drafter = old / "mtp-gemma.gguf"
new_drafter = new / "mtp-gemma.gguf"
main.write_bytes(b"main")
old_drafter.write_bytes(b"old")
new_drafter.write_bytes(b"new")
monkeypatch.setattr(mc, "_iter_hf_cache_snapshots", lambda _repo: [new, old])
got = LlamaCppBackend()._download_mtp(
hf_repo = "unsloth/gemma-GGUF",
near_path = str(main),
)
assert got == str(old_drafter)
def test_download_mtp_online_skips_cache_reuse(tmp_path, monkeypatch):
# Online, do not reuse a cached copy: go to the download path so a changed
# drafter is refetched (hf_hub_download checks the current revision).
@ -447,6 +474,7 @@ def test_download_mtp_online_skips_cache_reuse(tmp_path, monkeypatch):
pick,
label,
cancel_event = None,
near_path = None,
):
reached["hit"] = True
return None

View file

@ -244,9 +244,7 @@ class TestGgufVariantFileResolution:
def test_download_reuses_older_snapshot_when_current_ref_snapshot_is_partial(
self, monkeypatch, hf_cache
):
# Cross-snapshot reuse is an offline-resilience path: online, hf_hub_download
# resumes the partial current-ref download and revalidates the revision instead
# of serving an older snapshot's same-name blob.
# Keep coverage for offline reuse; online reuse is tested separately.
monkeypatch.setenv("HF_HUB_OFFLINE", "1")
backend = LlamaCppBackend()
repo = "unsloth/vision-GGUF"
@ -292,8 +290,7 @@ class TestGgufVariantFileResolution:
def test_download_reuses_cached_gguf_when_lowercase_partial_cache_shadows_it(
self, monkeypatch, hf_cache
):
# Case-variant cross-dir reuse is offline-only; online the canonical repo id
# resolves up front and hf_hub_download fetches the current revision.
# Keep coverage for case-insensitive offline cache lookup.
monkeypatch.setenv("HF_HUB_OFFLINE", "1")
backend = LlamaCppBackend()
canonical_repo = "unsloth/gemma-4-E2B-it-GGUF"
@ -348,45 +345,26 @@ class TestGgufVariantFileResolution:
assert out == str(snap / gguf_file)
assert seen_repos
def test_download_online_does_not_reuse_old_snapshot(self, monkeypatch, hf_cache):
# Online, an older same-name snapshot must not be served (it may be a stale
# revision); hf_hub_download is called so the current revision is fetched and
# its etag revalidated.
def test_download_online_reuses_complete_cached_snapshot(self, monkeypatch, hf_cache):
# Loads reuse complete cached models across repo revisions.
monkeypatch.delenv("HF_HUB_OFFLINE", raising = False)
backend = LlamaCppBackend()
repo = "unsloth/vision-GGUF"
_build_cache(hf_cache, repo, {"model-UD-Q4_K_XL.gguf": 4}, snapshot_sha = "a" * 40)
downloaded: list[str] = []
snap = _build_cache(hf_cache, repo, {"model-UD-Q4_K_XL.gguf": 4}, snapshot_sha = "a" * 40)
def fake_get_paths_info(
_repo_id,
paths,
token = None,
):
return [_types.SimpleNamespace(path = p, size = 4) for p in paths if p]
def fake_download(
repo_id,
filename,
token = None,
**kwargs,
):
downloaded.append(filename)
return f"/fresh/{filename}"
def fail_download(*_args, **_kwargs):
raise AssertionError("must reuse the cached GGUF instead of downloading")
with (
patch(
"huggingface_hub.list_repo_files",
lambda *_a, **_k: ["model-UD-Q4_K_XL.gguf"],
),
patch("huggingface_hub.get_paths_info", fake_get_paths_info),
patch("huggingface_hub.try_to_load_from_cache", lambda *_a, **_k: None),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fake_download),
patch("core.inference.llama_cpp.hf_hub_download_with_xet_fallback", fail_download),
):
out = backend._download_gguf(hf_repo = repo, hf_variant = "UD-Q4_K_XL")
assert downloaded == ["model-UD-Q4_K_XL.gguf"]
assert out == "/fresh/model-UD-Q4_K_XL.gguf"
assert out == str(snap / "model-UD-Q4_K_XL.gguf")
def test_download_reuses_older_snapshot_when_offline_env_is_true(self, monkeypatch, hf_cache):
# HF_HUB_OFFLINE accepts truthy spellings beyond "1" (true/yes/on); the offline

View file

@ -1054,7 +1054,7 @@ class TestChatCompletionRequestToolFields:
monkeypatch.setattr(
inference_route,
"_detect_safetensors_features",
lambda backend, chat_template: {"supports_tools": True},
lambda backend, chat_template, tools = None: {"supports_tools": True},
)
monitor = ApiMonitor(max_entries = 3)
monkeypatch.setattr(inference_route, "api_monitor", monitor)
@ -6620,6 +6620,29 @@ class TestApiMonitorAudioInput:
assert entry["reply"] == "hello world"
assert monitor.active_count() == 0
def failing_chunks():
yield "partial"
raise RuntimeError("generation failed")
self._patch_audio_backend(monkeypatch, failing_chunks())
error_monitor = ApiMonitor(max_entries = 3)
monkeypatch.setattr(inf_mod, "api_monitor", error_monitor)
error_response = await openai_chat_completions(
payload,
request = request,
current_subject = "test",
)
error_chunks = [
chunk.decode() if isinstance(chunk, bytes) else chunk
async for chunk in error_response.body_iterator
]
assert '"type": "server_error"' in error_chunks[-1]
assert error_chunks[-1].endswith("data: [DONE]\n\n")
[error_entry] = error_monitor.snapshot()
assert error_entry["status"] == "error"
assert error_monitor.active_count() == 0
asyncio.run(_run())
def test_non_gguf_tts_auto_route_records_monitor(self, monkeypatch):

View file

@ -417,6 +417,59 @@ def test_detect_safetensors_features_gemma_native_tool_call_keeps_tools_on():
assert flags["supports_tools"] is True
def test_detect_safetensors_features_gemma_native_reasoning_is_parseable_not_prefilled():
"""Native Gemma channels are normalized to <think>, then split by the route."""
from routes.inference import _detect_safetensors_features, _sf_reasoning_prefill_mode
tpl_with_gemma_native = "{% if add_generation_prompt %}<|channel>thought\n<channel|>{% endif %}"
backend = SimpleNamespace(
active_model_name = "unsloth/gemma-4-E2B-it",
models = {
"unsloth/gemma-4-E2B-it": {
"native_chat_template": tpl_with_gemma_native,
"chat_template_info": {"template": "override has no native markers"},
}
},
)
flags = _detect_safetensors_features(backend, "override has no native markers")
missing_arg_flags = _detect_safetensors_features(backend, None)
assert flags["supports_reasoning"] is True
assert flags["reasoning_always_on"] is True
assert missing_arg_flags["supports_reasoning"] is True
assert _sf_reasoning_prefill_mode(flags, None, tpl_with_gemma_native) is False
def test_detect_safetensors_features_selects_native_reasoning_from_tool_template():
"""Request tools select a marker-bearing named template without affecting default chat."""
from routes.inference import _detect_safetensors_features
named_template = {
"default": "plain default template",
"tool_use": "{% if tools %}<|channel>thought\n<channel|>{% endif %}",
}
backend = SimpleNamespace(
active_model_name = "custom/named-native-reasoning",
models = {
"custom/named-native-reasoning": {
"native_chat_template": named_template,
"chat_template_info": {"template": "{% if tools %}<tool_call>{% endif %}"},
}
},
)
default_flags = _detect_safetensors_features(backend, "plain override")
tool_flags = _detect_safetensors_features(
backend,
"plain override",
tools = [{"type": "function"}],
)
assert default_flags["supports_reasoning"] is False
assert tool_flags["supports_reasoning"] is True
assert tool_flags["reasoning_always_on"] is True
# Qwen3.5 family pin: the live GGUF + safetensors templates both wrap tool
# calls as ``<tool_call>\n<function=name>...``. Faithful slice so the
# classifier never silently regresses for this family.

View file

@ -215,3 +215,135 @@ def test_s6_reasoning_effort_none_disables_prefill_for_enable_thinking_effort():
swallowed = _replay_sf_reasoning_stream(events, prefilled = True)
assert swallowed["visible"] == ""
assert swallowed["reasoning"] == "The capital of France is Paris."
def test_native_reasoning_streamer_selected_and_errors_raise():
import threading
import pytest
torch = pytest.importorskip("torch")
inf = pytest.importorskip("core.inference.inference")
class Batch(dict):
def to(self, _device):
return self
class Tok:
chat_template = "<|channel>thought\n...<channel|>"
all_special_tokens = []
eos_token_id = 1
pad_token_id = None
pieces = {10: "<|channel>thought\n", 11: "r", 12: "<channel|>", 13: "a"}
def __call__(self, *_args, **_kwargs):
return Batch({"input_ids": torch.zeros((1, 1), dtype = torch.long)})
def decode(self, ids, **_kwargs):
return "".join(self.pieces.get(int(token_id), "") for token_id in ids)
class Model:
device = "cpu"
generation_config = type("Cfg", (), {"eos_token_id": 1})()
config = generation_config
def __init__(self, fail = False):
self.fail = fail
self.kwargs = None
def generate(self, **kwargs):
self.kwargs = kwargs
streamer = kwargs["streamer"]
streamer.put(torch.zeros((1, 1), dtype = torch.long))
for token_id in [10, 11, 12, 13]:
streamer.put(torch.tensor([token_id]))
if self.fail:
raise RuntimeError("boom")
backend = inf.InferenceBackend.__new__(inf.InferenceBackend)
backend.active_model_name = "gemma-test"
backend._generation_lock = threading.Lock()
backend.models = {"gemma-test": {"model": Model(), "tokenizer": Tok()}}
assert list(backend.generate_stream("prompt", max_new_tokens = 4))[-1] == "<think>r</think>a"
backend.models["gemma-test"]["model"] = Model(fail = True)
with pytest.raises(inf._GenerationThreadError, match = "boom"):
list(backend.generate_stream("prompt", max_new_tokens = 4))
def test_text_only_vlm_fallback_resolves_native_markers_off():
import threading
import pytest
torch = pytest.importorskip("torch")
inf = pytest.importorskip("core.inference.inference")
class Batch(dict):
def to(self, _device):
return self
class Tokenizer:
all_special_tokens = []
eos_token_id = 1
pad_token_id = None
def __call__(self, *_args, **_kwargs):
return Batch({"input_ids": torch.zeros((1, 1), dtype = torch.long)})
class Processor:
chat_template = "<|channel>thought\n...<channel|>"
tokenizer = Tokenizer()
class Model:
device = "cpu"
generation_config = type("Cfg", (), {"eos_token_id": 1})()
config = generation_config
def generate(self, **_kwargs):
return None
class EmptyStreamer:
def __next__(self):
raise StopIteration
def end(self):
return None
captured = {}
backend = inf.InferenceBackend.__new__(inf.InferenceBackend)
backend.active_model_name = "vision-test"
backend._generation_lock = threading.Lock()
backend.models = {
"vision-test": {
"model": Model(),
"processor": Processor(),
"tokenizer": Processor(),
}
}
backend.format_chat_prompt = lambda *_args, **_kwargs: "manual text-only prompt"
def make_streamer(*_args, **kwargs):
captured.update(kwargs)
return EmptyStreamer()
backend._make_text_streamer = make_streamer
assert (
list(
backend._generate_vision_response(
messages = [{"role": "user", "content": "hello"}],
system_prompt = "",
image = None,
temperature = 0.7,
top_p = 0.9,
top_k = 40,
min_p = 0.0,
max_new_tokens = 1,
repetition_penalty = 1.0,
)
)
== []
)
assert captured["reasoning_channel_markers"] is None
assert captured["reasoning_channel_markers_resolved"] is True

View file

@ -3205,6 +3205,196 @@ class TestLoopBehaviour:
class TestLoopRePrompt:
"""Plan-without-action re-prompt parity with GGUF: nudge instead of terminating, up to ``MAX_ACT_REPROMPTS`` extra slots. Studio always nudges, so these drive the loop with ``nudge_tool_calls=True``."""
def test_reasoning_intent_does_not_reprompt_a_visible_answer(self):
generations = 0
def _gen(_messages, active_tools = None):
nonlocal generations
generations += 1
yield (
"<think>Let me prepare the requested summary carefully.</think>"
"This is the final visible answer."
)
exec_fn = FakeExecuteTool([])
events = _collect_events(
run_safetensors_tool_loop(
single_turn = _gen,
messages = [{"role": "user", "content": "summarize this"}],
tools = [{"type": "function", "function": {"name": "web_search"}}],
execute_tool = exec_fn,
nudge_tool_calls = True,
)
)
assert generations == 1
assert exec_fn.calls == []
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1].endswith("This is the final visible answer.")
def test_prefilled_reasoning_intent_does_not_reprompt_a_visible_answer(self):
generations = 0
def _gen(_messages, active_tools = None):
nonlocal generations
generations += 1
yield "Let me prepare the requested summary carefully.</think>This is the final visible answer."
exec_fn = FakeExecuteTool([])
events = _collect_events(
run_safetensors_tool_loop(
single_turn = _gen,
messages = [{"role": "user", "content": "summarize this"}],
tools = [{"type": "function", "function": {"name": "web_search"}}],
execute_tool = exec_fn,
nudge_tool_calls = True,
reasoning_prefilled = True,
)
)
assert generations == 1
assert exec_fn.calls == []
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1].endswith("This is the final visible answer.")
def test_prefilled_reasoning_with_reemitted_think_does_not_reprompt(self):
generations = 0
def _gen(_messages, active_tools = None):
nonlocal generations
generations += 1
yield (
"Let me prepare the requested summary carefully."
"<think>more private planning</think>This is the final visible answer."
)
exec_fn = FakeExecuteTool([])
events = _collect_events(
run_safetensors_tool_loop(
single_turn = _gen,
messages = [{"role": "user", "content": "summarize this"}],
tools = [{"type": "function", "function": {"name": "web_search"}}],
execute_tool = exec_fn,
nudge_tool_calls = True,
reasoning_prefilled = True,
)
)
assert generations == 1
assert exec_fn.calls == []
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1].endswith("This is the final visible answer.")
def test_prefilled_reasoning_with_later_think_does_not_reprompt(self):
generations = 0
def _gen(_messages, active_tools = None):
nonlocal generations
generations += 1
yield (
"private prefilled planning</think>"
"<think>Let me prepare the requested summary carefully.</think>"
"This is the final visible answer."
)
exec_fn = FakeExecuteTool([])
events = _collect_events(
run_safetensors_tool_loop(
single_turn = _gen,
messages = [{"role": "user", "content": "summarize this"}],
tools = [{"type": "function", "function": {"name": "web_search"}}],
execute_tool = exec_fn,
nudge_tool_calls = True,
reasoning_prefilled = True,
)
)
assert generations == 1
assert exec_fn.calls == []
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1].endswith("This is the final visible answer.")
def test_reasoning_only_intent_still_reprompts_and_uses_a_tool(self):
loop, exec_fn = _make_loop(
turns = [
["<think>Let me search for that.</think>"],
['<tool_call>{"name":"web_search","arguments":{"query":"cats"}}</tool_call>'],
["Here is the answer."],
],
exec_results = ["result"],
nudge_tool_calls = True,
)
events = _collect_events(loop)
assert exec_fn.calls == [("web_search", {"query": "cats"})]
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1] == "Here is the answer."
def test_prefilled_no_close_reasoning_intent_still_reprompts(self):
loop, exec_fn = _make_loop(
turns = [
["I need more context.<think>Let me search for that."],
['<tool_call>{"name":"web_search","arguments":{"query":"cats"}}</tool_call>'],
["Here is the answer."],
],
exec_results = ["result"],
nudge_tool_calls = True,
reasoning_prefilled = True,
)
events = _collect_events(loop)
assert exec_fn.calls == [("web_search", {"query": "cats"})]
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1] == "Here is the answer."
def test_prefilled_reasoning_prefix_is_kept_for_reasoning_only_reprompt(self):
loop, exec_fn = _make_loop(
turns = [
["Let me search for that.</think><think>checking details</think>"],
['<tool_call>{"name":"web_search","arguments":{"query":"cats"}}</tool_call>'],
["Here is the answer."],
],
exec_results = ["result"],
nudge_tool_calls = True,
reasoning_prefilled = True,
)
events = _collect_events(loop)
assert exec_fn.calls == [("web_search", {"query": "cats"})]
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1] == "Here is the answer."
def test_reprompt_history_uses_visible_intent_text(self):
captured: list[list[dict]] = []
def _gen(messages, active_tools = None):
captured.append([dict(message) for message in messages])
if len(captured) == 1:
yield "<think>private planning details</think>Let me search for that."
elif len(captured) == 2:
yield '<tool_call>{"name":"web_search","arguments":{"query":"cats"}}</tool_call>'
else:
yield "Here is the answer."
exec_fn = FakeExecuteTool(["result"])
events = _collect_events(
run_safetensors_tool_loop(
single_turn = _gen,
messages = [{"role": "user", "content": "find cats"}],
tools = [{"type": "function", "function": {"name": "web_search"}}],
execute_tool = exec_fn,
nudge_tool_calls = True,
)
)
assert exec_fn.calls == [("web_search", {"query": "cats"})]
assert captured[1][1] == {"role": "assistant", "content": "Let me search for that."}
contents = [e["text"] for e in events if e["type"] == "content"]
assert contents[-1] == "Here is the answer."
def test_intent_signal_triggers_reprompt(self):
# Turn 1: intent signal, no tool call.
# Turn 2 (re-prompt): proper tool call -> executes.

View file

@ -177,6 +177,15 @@ def _sse_objects(chunks):
# ── Non-streaming ─────────────────────────────────────────────────
def test_non_reasoning_backend_keeps_literal_think_tags(monkeypatch):
backend = _ScriptedBackend(_fixed("show <think>example</think> tags"))
response = _call(_request(stream = False), monkeypatch, backend, supports_tools = False)
message = _json_body(response)["choices"][0]["message"]
assert message["content"] == "show <think>example</think> tags"
assert message["reasoning_content"] is None
def test_xml_healed_to_tool_calls_non_streaming(monkeypatch):
backend = _ScriptedBackend(_fixed(_CALL_XML))
payload = _request(tools = [LOOKUP_TOOL], stream = False)
@ -485,6 +494,52 @@ def test_streaming_no_tools_verbatim(monkeypatch):
assert finishes == ["stop"]
def test_streaming_gen_stream_error_is_not_model_text(monkeypatch):
from core.inference.orchestrator import GenStreamError
class _ErrorAfterPartial(_ScriptedBackend):
def __init__(self):
super().__init__(_fixed())
def generate_chat_response(self, **_kwargs):
yield "<think>partial"
yield GenStreamError("Error: /tmp/secret traceback")
backend = _ErrorAfterPartial()
payload = _request(stream = True)
response = _call(payload, monkeypatch, backend, supports_tools = False)
chunks = _collect_sse(response)
objs = _sse_objects(chunks)
deltas = [o.get("choices", [{}])[0].get("delta", {}) for o in objs if o.get("choices")]
assert any("partial" in json.dumps(delta) for delta in deltas)
assert not any("/tmp/secret" in json.dumps(delta) for delta in deltas)
errors = [o["error"]["message"] for o in objs if "error" in o]
assert errors == ["An internal error occurred."]
assert any(
"data: [DONE]" in (chunk.decode() if isinstance(chunk, bytes) else chunk)
for chunk in chunks
)
def test_server_tool_streaming_invalid_event_is_error(monkeypatch):
class _InvalidEventBackend(_ScriptedBackend):
def __init__(self):
super().__init__(_fixed())
def generate_chat_completion_with_tools(self, **_kwargs):
yield {"type": "content", "text": "partial"}
yield "not-an-event"
backend = _InvalidEventBackend()
payload = _request(tools = [LOOKUP_TOOL], enable_tools = True, stream = True)
response = _call(payload, monkeypatch, backend)
objs = _sse_objects(_collect_sse(response))
errors = [o["error"]["message"] for o in objs if "error" in o]
assert errors == ["An internal error occurred."]
def test_streaming_repeated_snapshot_no_duplicate_call(monkeypatch):
# Repeated then shrunk cumulative snapshots must not double-heal.
backend = _ScriptedBackend(_fixed(_CALL_XML, _CALL_XML, _CALL_XML[:5], _CALL_XML))

View file

@ -2,7 +2,7 @@
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""
Unit tests for detect_think_prefill.
Unit tests for local reasoning-stream helpers.
Reasoning templates (Qwen3.6-style) end the generation prompt with an open
``<think>\\n`` so the model starts reasoning immediately. skip_prompt
@ -16,7 +16,13 @@ import sys
_backend = os.path.join(os.path.dirname(__file__), "..")
sys.path.insert(0, _backend)
from core.inference.chat_template_helpers import detect_think_prefill
from core.inference.chat_template_helpers import (
ReasoningChannelNormalizer,
detect_reasoning_channel_markers,
detect_reasoning_channel_markers_from_model_info,
detect_think_prefill,
render_with_native_template_fallback,
)
QWEN_PROMPT = "<|im_start|>user\nHi!<|im_end|>\n<|im_start|>assistant\n"
@ -87,3 +93,170 @@ def test_guard_emits_when_think_not_special():
def test_guard_default_and_empty_keep_emitting():
assert detect_think_prefill(QWEN_PROMPT + "<think>\n", None) == "<think>\n"
assert detect_think_prefill(QWEN_PROMPT + "<think>\n", []) == "<think>\n"
def test_gemma_channel_detection_uses_active_template_not_token_metadata():
class TemplateTokenizer:
chat_template = {"default": "...<|channel>thought\\n{{ eoc_token }}"}
class NamedTemplateTokenizer:
chat_template = {
"default": "plain assistant template",
"tool_use": "...<|channel>thought\\n{{ eoc_token }}",
}
class TokenMetadataOnly:
chat_template = None
soc_token = "<|channel>"
eoc_token = "<channel|>"
class NamedTemplateProcessor:
chat_template = {
"default": "plain processor default",
"tool_use": "<|channel>thought\nprocessor tool template<channel|>",
}
tokenizer = TokenMetadataOnly()
def apply_chat_template(self, *_args, **_kwargs):
raise NotImplementedError
expected = ("<|channel>thought", "<channel|>")
assert detect_reasoning_channel_markers(TemplateTokenizer()) == expected
assert detect_reasoning_channel_markers(NamedTemplateTokenizer()) is None
assert (
detect_reasoning_channel_markers(
NamedTemplateTokenizer(), tools = [{"function": {"name": "web_search"}}]
)
== expected
)
assert detect_reasoning_channel_markers(NamedTemplateTokenizer(), tools = []) is None
assert (
detect_reasoning_channel_markers(
NamedTemplateProcessor(), tools = [{"function": {"name": "web_search"}}]
)
is None
)
assert detect_reasoning_channel_markers(TokenMetadataOnly()) is None
def test_gemma_channel_detection_tries_no_argument_getter_fallback():
class FallbackTokenizer:
chat_template = "plain fallback template"
def get_chat_template(self, **kwargs):
if kwargs:
raise ValueError("tools are not supported")
return "...<|channel>thought\n<channel|>"
assert detect_reasoning_channel_markers(
FallbackTokenizer(), tools = [{"function": {"name": "web_search"}}]
) == ("<|channel>thought", "<channel|>")
def test_native_template_fallback_returns_selected_reasoning_metadata():
from types import SimpleNamespace
messages = [{"role": "user", "content": "hi"}]
tools = [{"type": "function", "function": {"name": "web_search"}}]
def render(tokenizer, msgs, *, tools, **_kw):
body = "".join(message["content"] for message in msgs)
suffix = "|TOOLS" if tools else ""
return body + suffix if tokenizer.chat_template == "NATIVE <|channel>thought\n" else body
result = render_with_native_template_fallback(
formatted_prompt = "hi",
tokenizer = SimpleNamespace(chat_template = "OVERRIDE"),
model_info = {
"native_chat_template": "NATIVE <|channel>thought\n",
"tokenizer": SimpleNamespace(chat_template = "OVERRIDE"),
},
active_model_name = "gemma-test",
messages = messages,
tools = tools,
apply_fn = render,
return_metadata = True,
)
assert result.prompt == "hi|TOOLS"
assert result.reasoning_channel_markers == ("<|channel>thought", "<channel|>")
def test_cached_native_template_metadata_recovers_reasoning_markers_without_tools():
from types import SimpleNamespace
model_info = {"chat_template_info": {"template": "native <|channel>thought\n<channel|>"}}
assert detect_reasoning_channel_markers_from_model_info(
SimpleNamespace(chat_template = "override has no native markers"),
model_info,
tools = None,
) == ("<|channel>thought", "<channel|>")
result = render_with_native_template_fallback(
formatted_prompt = "prompt from override",
tokenizer = SimpleNamespace(chat_template = "override has no native markers"),
model_info = model_info,
active_model_name = "gemma-test",
messages = [{"role": "user", "content": "hi"}],
tools = None,
return_metadata = True,
)
assert result.prompt == "prompt from override"
assert result.reasoning_channel_markers == ("<|channel>thought", "<channel|>")
def test_cached_native_markers_do_not_describe_live_tool_template():
from types import SimpleNamespace
tools = [{"type": "function", "function": {"name": "web_search"}}]
class LiveTokenizer:
chat_template = "live tool template without native markers"
def render(_tokenizer, _messages, *, tools, **_kwargs):
return "prompt with tools" if tools else "prompt without tools"
result = render_with_native_template_fallback(
formatted_prompt = "prompt with tools",
tokenizer = LiveTokenizer(),
model_info = {
"chat_template_info": {"template": "native <|channel>thought\n<channel|>"},
"tokenizer": SimpleNamespace(),
},
active_model_name = "gemma-test",
messages = [{"role": "user", "content": "hi"}],
tools = tools,
apply_fn = render,
return_metadata = True,
)
assert result.prompt == "prompt with tools"
assert result.reasoning_channel_markers is None
def test_gemma_channel_normalization_is_prefix_monotonic_and_preserves_tools():
parser = ReasoningChannelNormalizer("<|channel>thought", "<channel|>")
output = ""
snapshots = []
for chunk in (
"<|chan",
"nel>thought",
"\nReason",
"<chan",
"nel|><|tool_call>web_search<tool_call|>",
):
delta = parser.feed(chunk)
if delta:
output += delta
snapshots.append(output)
assert snapshots == [
"<think>",
"<think>Reason",
"<think>Reason</think><|tool_call>web_search<tool_call|>",
]
assert snapshots[1].startswith(snapshots[0])
compact = ReasoningChannelNormalizer("<|channel>thought", "<channel|>")
assert compact.feed("<|channel>thought<channel|>answer") + compact.finish() == (
"<think></think>answer"
)

View file

@ -258,15 +258,27 @@ def test_watchdog_no_op_when_worker_superseded(monkeypatch):
def test_new_run_gets_its_own_watchdog(monkeypatch):
# A stale watchdog sleeping on an old proc must not stop a new run's stop from
# creating its own watcher.
monkeypatch.setitem(_G, "_STOP_GRACE_S", 100.0)
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0)
b = TrainingBackend()
_record_force_terminate(monkeypatch, b)
started = []
release = threading.Event()
def _blocked_watchdog(
target_proc,
cancel,
watched_job_id = None,
):
started.append(target_proc)
# No timeout: the finally always releases this, so a superseded watchdog stays
# alive through the assertions regardless of load; as a daemon it can't hang exit.
release.wait()
monkeypatch.setattr(b, "_stop_watchdog_loop", _blocked_watchdog)
old_proc = _FakeProc(alive = True)
b._proc = old_proc
b._start_stop_watchdog(cancel = False)
first_wd = b._stop_watchdog
assert _wait_until(lambda: started == [old_proc])
# New run: fresh worker replaces the handle; its stop must get a new watcher
# even though the old (superseded) watchdog is still alive.
@ -276,12 +288,12 @@ def test_new_run_gets_its_own_watchdog(monkeypatch):
second_wd = b._stop_watchdog
try:
assert _wait_until(lambda: started == [old_proc, new_proc])
assert first_wd.is_alive()
assert second_wd is not first_wd, "a new run must get its own watchdog"
assert b._stop_watchdog_proc is new_proc
finally:
old_proc._alive = False
new_proc._alive = False
release.set()
first_wd.join(timeout = 5)
second_wd.join(timeout = 5)

View file

@ -8,7 +8,7 @@ import { cn } from "@/lib/utils";
import { useAuiState } from "@assistant-ui/react";
import { LayoutTwoColumnIcon as Layout2ColumnIcon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { useLayoutEffect, useMemo } from "react";
import { useLayoutEffect, useMemo, useRef } from "react";
import { useChatRuntimeStore } from "../stores/chat-runtime-store";
import type { ArtifactViewMode } from "./html-frame";
import {
@ -83,15 +83,18 @@ export function ArtifactCard({
],
);
const surface = artifactThreadId ? "panel" : "overlay";
// Once per mount, so a view-change cleanup can't re-trigger a stale open.
const autoOpenAttemptedRef = useRef(false);
useLayoutEffect(() => {
if (selectedArtifactId === artifact.id) {
updateArtifact(artifact);
}
if (!autoOpen) {
if (!autoOpen || autoOpenAttemptedRef.current) {
return;
}
autoOpenAttemptedRef.current = true;
if (hasAutoOpenedArtifact(artifact.id)) {
return;
}

View file

@ -272,13 +272,13 @@ const SingleContent = memo(function SingleContent({
openResearchRun.threadId === (threadId ?? activeThreadId),
);
const showResearchPanel = researchMatchesThread && !isMobile;
// Without a URL threadId the artifact must belong to the active thread.
const showArtifactPanel = !showResearchPanel && Boolean(
artifact &&
artifactSurface === "panel" &&
(threadId
? !artifact.threadId || artifact.threadId === threadId
: Boolean(newThreadNonce) ||
Boolean(artifact.threadId && artifact.threadId === activeThreadId)),
: Boolean(artifact.threadId && artifact.threadId === activeThreadId)),
);
const showContextPanel = showResearchPanel || showArtifactPanel;
@ -1843,10 +1843,8 @@ export function ChatPage({
useEffect(() => {
if (view.mode !== "single") return;
if (view.threadId || view.newThreadNonce || !selectedArtifact) return;
// view excludes __LOCALID_ threads (they fall through to mode:"single"
// with no threadId/nonce). Don't close a canvas whose thread is the
// active local thread.
if (view.threadId || !selectedArtifact) return;
// Close any canvas that doesn't belong to the active thread.
if (
selectedArtifact.threadId &&
selectedArtifact.threadId === activeThreadId

View file

@ -244,8 +244,8 @@ export function ChatProvidersSettings({
(s) => s.setConnectionsEnabled,
);
const isCustomProvider = isCustomProviderType(providerType);
// Local presets (Ollama, llama.cpp) never use API keys — hide the field.
// vLLM may optionally use a bearer token on secured deployments.
// llama.cpp hides the key field. Ollama and vLLM show an optional key:
// Ollama cloud and secured vLLM need one; local servers leave it empty.
const showApiKeyField = !customPresetSkipsApiKeyField(providerType);
const showReasoningToggle = supportsProviderReasoningToggle(providerType);

View file

@ -184,11 +184,12 @@ export function supportsRemoteModelCatalog(
);
}
/** Presets that skip the API-key field (local servers with no auth by default). */
/** Presets that hide the API-key field. Ollama is not skipped: Ollama cloud
* requires a key; local servers leave the optional field empty. */
export function customPresetSkipsApiKeyField(
providerType: string | null | undefined,
): boolean {
return providerType === "ollama" || providerType === "llama_cpp";
return providerType === "llama_cpp";
}
/** Catalog load plus optional manual model IDs. */

View file

@ -278,28 +278,23 @@ export function PermissionModeComposerPill({
data-pill-label={active.label}
data-active={fullAccess ? "true" : "false"}
data-variant={fullAccess ? "danger" : undefined}
data-keep-label="true"
aria-label="Permission level for tool calls"
title={`${active.label}: ${active.description}`}
>
{/* The icon doubles as an off switch (mirrors the MCP pill): hover
swaps it to an X; clicking it turns bypass permissions Off (no
prompts, sandbox on) without opening the menu. In compact
icon-only mode the glyph is the whole button, so clicks fall
through and open the menu instead. */}
prompts, sandbox on) without opening the menu. data-keep-label
exempts this pill from compact icon-only mode, so the off switch
stays clickable even while the other pills are collapsed. */}
<span
role="button"
aria-label="Turn off bypass permissions"
tabIndex={-1}
onPointerDown={(e) => {
if (e.currentTarget.closest('[data-pill-compact="true"]')) {
return;
}
e.stopPropagation();
}}
onClick={(e) => {
if (e.currentTarget.closest('[data-pill-compact="true"]')) {
return;
}
e.stopPropagation();
setPermissionMode("off");
}}

View file

@ -70,11 +70,13 @@ import {
} from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { useNavigate } from "@tanstack/react-router";
import { motion, useReducedMotion } from "motion/react";
import {
type ChangeEvent,
type DragEvent,
useCallback,
useEffect,
useId,
useMemo,
useRef,
useState,
@ -153,6 +155,9 @@ function normalizeSliceInput(value: string): string | null {
export function DatasetSection() {
const t = useT();
const navigate = useNavigate();
const reducedMotion = useReducedMotion();
// Scopes the pill layoutId so multiple instances never share one.
const sourcePillLayoutId = useId();
const {
dataset,
datasetSource,
@ -686,6 +691,9 @@ export function DatasetSection() {
{(() => {
// Hub-style sliding-pill segmented control, matching the Hub tabs
// via the shared .hub-tab-toggle / .hub-tab-toggle-pill classes.
// flex-auto buttons share leftover space equally so padding stays
// equal for all labels; the pill sits inside the active button so
// it always matches its bounds.
const sourceTabs: {
value: "huggingface" | "upload" | "s3";
label: string;
@ -696,24 +704,12 @@ export function DatasetSection() {
? []
: [{ value: "s3" as const, label: "Amazon S3" }]),
];
const activeIndex = Math.max(
0,
sourceTabs.findIndex((item) => item.value === datasetSource),
);
return (
<div
role="radiogroup"
aria-label="Dataset source"
className="hub-tab-toggle relative inline-flex h-9 w-full items-center rounded-full"
className="hub-tab-toggle relative flex h-9 w-full items-center rounded-full"
>
<span
aria-hidden="true"
className="hub-tab-toggle-pill pointer-events-none absolute inset-y-0 left-0 rounded-full transition-transform duration-200 ease-out"
style={{
width: `${100 / sourceTabs.length}%`,
transform: `translateX(${activeIndex * 100}%)`,
}}
/>
{sourceTabs.map((item) => (
<button
key={item.value}
@ -732,13 +728,30 @@ export function DatasetSection() {
}
}}
className={cn(
"relative z-10 inline-flex h-9 flex-1 cursor-pointer items-center justify-center rounded-full px-3 text-[12.5px] font-medium transition-colors",
"relative inline-flex h-9 flex-auto cursor-pointer items-center justify-center rounded-full px-3 text-[12.5px] font-medium transition-colors",
datasetSource === item.value
? "text-foreground"
: "text-muted-foreground hover:text-foreground",
)}
>
{item.label}
{datasetSource === item.value && (
<motion.span
aria-hidden="true"
layoutId={sourcePillLayoutId}
className="hub-tab-toggle-pill absolute inset-0 rounded-full"
transition={
reducedMotion
? { duration: 0 }
: {
type: "spring",
stiffness: 500,
damping: 35,
mass: 0.5,
}
}
/>
)}
<span className="relative z-10">{item.label}</span>
</button>
))}
</div>

View file

@ -1467,7 +1467,9 @@ html[data-chat-font] .aui-root {
}
/* With more than 4 tools on, drop pill labels to icons only to cut clutter.
Compare keeps its label via data-keep-label. */
Compare and the bypass-permissions pill keep their labels via
data-keep-label; the permission pill sits before the collapsed icons, so
they line up to its right. */
[data-pill-compact="true"]
.composer-pill-btn:not([data-keep-label])
> span:not(.composer-pill-glyph) {

View file

@ -0,0 +1,82 @@
import ast
from pathlib import Path
def _load_function(name):
# Extract a function from kernels/utils.py without importing unsloth (which
# needs a GPU / torch / bitsandbytes). The functions exercised here only use
# getattr and the _FP8_WEIGHT_DTYPES name on the paths under test.
source = Path(__file__).parents[2] / "unsloth" / "kernels" / "utils.py"
tree = ast.parse(source.read_text(encoding = "utf-8"))
funcs = [
node for node in ast.walk(tree) if isinstance(node, ast.FunctionDef) and node.name == name
]
assert len(funcs) == 1, (name, funcs)
namespace = {"getattr": getattr, "_FP8_WEIGHT_DTYPES": ()}
module = ast.Module(body = funcs, type_ignores = [])
ast.fix_missing_locations(module)
exec(compile(module, str(source), "exec"), namespace)
return namespace[name]
class _Obj:
pass
def _make_disabled_block_fp8_proj(block_size):
# A merged/disabled projection whose base layer is a block-fp8 weight that
# ships a non-default block size on its checkpoint.
weight = _Obj()
weight.quant_state = _Obj()
base_layer = _Obj()
base_layer.weight = weight
base_layer.quant_method = "fp8"
base_layer.block_size = block_size
base_layer.bias = None
proj = _Obj()
proj.base_layer = base_layer
proj.merged = True
proj.disable_adapters = True
return proj, weight.quant_state
def test_bias_variant_propagates_fp8_block_size_on_disabled_path():
# Downstream fp8 kernels read getattr(weight_scale, "block_size", [128, 128]),
# so the checkpoint's real block size must survive the merged/disabled path,
# exactly as it does for the non-bias sibling get_lora_parameters.
get_lora_parameters_bias = _load_function("get_lora_parameters_bias")
proj, weight_scale = _make_disabled_block_fp8_proj([64, 128])
get_lora_parameters_bias(proj)
assert getattr(weight_scale, "block_size", [128, 128]) == [64, 128]
def _make_decompressed_merged_proj():
# A merged compressed-tensors layer that was decompressed back to bf16. It keeps
# quant_method == "fp8" from the checkpoint metadata, but the live weight is bf16
# so there is no quant state to attach a block size to.
weight = _Obj()
weight.dtype = "bfloat16"
base_layer = _Obj()
base_layer.weight = weight
base_layer.quant_method = "fp8"
base_layer.block_size = [128, 128]
base_layer.bias = None
proj = _Obj()
proj.base_layer = base_layer
proj.merged = True
proj.disable_adapters = True
return proj
def test_bias_variant_keeps_none_quant_state_for_decompressed_layer():
# Such a layer has no quant state, and fast_linear_forward relies on getting
# W_quant None back so it can fall back to a plain matmul, so setting the block
# size must not assume a quant state is present.
get_lora_parameters_bias = _load_function("get_lora_parameters_bias")
W, W_quant = get_lora_parameters_bias(_make_decompressed_merged_proj())[:2]
assert W_quant is None
assert getattr(W, "block_size", None) == [128, 128]

View file

@ -0,0 +1,81 @@
import ast
from pathlib import Path
def _load_function(name):
# Extract a function from kernels/utils.py without importing unsloth (which
# needs a GPU / torch / bitsandbytes). get_lora_parameters only uses getattr,
# hasattr and the _FP8_WEIGHT_DTYPES name on the paths under test.
source = Path(__file__).parents[2] / "unsloth" / "kernels" / "utils.py"
tree = ast.parse(source.read_text(encoding = "utf-8"))
funcs = [
node for node in ast.walk(tree) if isinstance(node, ast.FunctionDef) and node.name == name
]
assert len(funcs) == 1, (name, funcs)
namespace = {"getattr": getattr, "hasattr": hasattr, "_FP8_WEIGHT_DTYPES": ()}
module = ast.Module(body = funcs, type_ignores = [])
ast.fix_missing_locations(module)
exec(compile(module, str(source), "exec"), namespace)
return namespace[name]
class _Obj:
pass
def _make_disabled_block_fp8_proj(block_size):
# A merged/disabled projection whose base layer is a block-fp8 weight that
# ships a non-default block size on its checkpoint.
weight = _Obj()
weight.quant_state = _Obj()
base_layer = _Obj()
base_layer.weight = weight
base_layer.quant_method = "fp8"
base_layer.block_size = block_size
proj = _Obj()
proj.base_layer = base_layer
proj.merged = True
proj.disable_adapters = True
return proj, weight.quant_state
def test_propagates_fp8_block_size_on_disabled_path():
# get_lora_parameters already sets block_size before its early return; downstream
# fp8 kernels read getattr(weight_scale, "block_size", [128, 128]), so the
# checkpoint's real block size must survive the merged/disabled path.
get_lora_parameters = _load_function("get_lora_parameters")
proj, weight_scale = _make_disabled_block_fp8_proj([64, 128])
get_lora_parameters(proj)
assert getattr(weight_scale, "block_size", [128, 128]) == [64, 128]
def _make_decompressed_merged_proj():
# A merged compressed-tensors layer that was decompressed back to bf16. It keeps
# quant_method == "fp8" from the checkpoint metadata, but the live weight is bf16
# so there is no quant state to attach a block size to.
weight = _Obj()
weight.dtype = "bfloat16"
base_layer = _Obj()
base_layer.weight = weight
base_layer.quant_method = "fp8"
base_layer.block_size = [128, 128]
proj = _Obj()
proj.base_layer = base_layer
proj.merged = True
proj.disable_adapters = True
return proj
def test_keeps_none_quant_state_for_decompressed_layer():
# Mirrors the get_lora_parameters_bias guard: with no quant state, assigning
# W_quant.block_size must not assume one is present, or it raises AttributeError
# on None. fast_lora relies on getting W_quant None back to fall back to a plain
# matmul, so this path must stay crash-free.
get_lora_parameters = _load_function("get_lora_parameters")
W, W_quant = get_lora_parameters(_make_decompressed_merged_proj())[:2]
assert W_quant is None
assert getattr(W, "block_size", None) == [128, 128]

View file

@ -0,0 +1,307 @@
# SPDX-License-Identifier: AGPL-3.0-only
import gc
import os
os.environ.setdefault("PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION", "python")
import transformers
from transformers.utils import sentencepiece_model_pb2
from unsloth.tokenizer_utils import fix_sentencepiece_tokenizer
NORMAL, CONTROL = 1, 3
def _spm_bytes(pieces):
m = sentencepiece_model_pb2.ModelProto()
for piece, score, typ in pieces:
p = m.pieces.add()
p.piece = piece
p.score = score
p.type = typ
return m.SerializeToString()
def _read_pieces(path):
m = sentencepiece_model_pb2.ModelProto()
with open(path, "rb") as f:
m.ParseFromString(f.read())
return [p.piece for p in m.pieces]
class _FakeTokenizer:
"""Minimal stand-in for a sentencepiece-backed slow tokenizer.
``save_pretrained`` writes a tokenizer.model, which is what the real slow
tokenizers do and what fix_sentencepiece_tokenizer reads back.
"""
def __init__(
self,
name,
spm_bytes = None,
vocab = None,
):
self.name = name
self.eos_token = "</s>"
self.pad_token = "<pad>"
self._spm_bytes = spm_bytes
self._vocab = vocab or {}
self.saved_to = []
def save_pretrained(self, location):
self.saved_to.append(location)
os.makedirs(location, exist_ok = True)
if self._spm_bytes is not None:
with open(os.path.join(location, "tokenizer.model"), "wb") as f:
f.write(self._spm_bytes)
def __call__(
self,
texts,
add_special_tokens = False,
):
class _Encoded:
pass
encoded = _Encoded()
encoded.input_ids = [[self._vocab[text]] for text in texts]
return encoded
def _tokenizers():
pieces = [("<s>", 0.0, CONTROL), ("a", -1.0, NORMAL), ("</s>", 0.0, CONTROL)]
old = _FakeTokenizer("old", spm_bytes = _spm_bytes(pieces), vocab = {"</s>": 2})
new = _FakeTokenizer("new")
return old, new
class _ReloadedTokenizer:
"""Weakref-able stand-in for the tokenizer AutoTokenizer.from_pretrained returns."""
def __init__(self, location):
self.location = location
def _stub_auto_tokenizer(monkeypatch):
"""fix_sentencepiece_tokenizer reloads the patched directory through
AutoTokenizer at the end; that needs a full tokenizer on disk, which is
out of scope here. Record the reload location and hand back a sentinel.
"""
loaded = []
class _StubAutoTokenizer:
@staticmethod
def from_pretrained(location, **kwargs):
loaded.append(location)
return _ReloadedTokenizer(location)
monkeypatch.setattr(transformers, "AutoTokenizer", _StubAutoTokenizer)
return loaded
def test_old_tokenizer_is_saved_so_its_model_can_be_read(tmp_path, monkeypatch):
"""The guard must not skip the body on a fresh temporary directory.
fix_sentencepiece_tokenizer creates its scratch directory itself and then
checks for a tokenizer.model inside it, but that file only appears once
old_tokenizer.save_pretrained() has run.
"""
_stub_auto_tokenizer(monkeypatch)
old, new = _tokenizers()
location = str(tmp_path / "_unsloth_sentencepiece_temp")
fix_sentencepiece_tokenizer(old, new, {"</s>": "<|im_end|>"}, temporary_location = location)
assert old.saved_to, "old tokenizer was never saved: the body did not run"
def test_token_mapping_is_applied_to_the_sentencepiece_model(tmp_path, monkeypatch):
loaded = _stub_auto_tokenizer(monkeypatch)
old, new = _tokenizers()
location = str(tmp_path / "_unsloth_sentencepiece_temp")
# Hold the returned tokenizer so its scratch dir survives until we read it.
tok = fix_sentencepiece_tokenizer(old, new, {"</s>": "<|im_end|>"}, temporary_location = location)
assert "<|im_end|>" in _read_pieces(f"{loaded[-1]}/tokenizer.model")
assert tok is not None
def test_tokenizer_without_a_sentencepiece_model_is_returned_untouched(tmp_path, monkeypatch):
"""A fast-only tokenizer writes no tokenizer.model, so the guard still
short-circuits and the caller gets new_tokenizer back unchanged. Its scratch
dir is unreferenced and reclaimed immediately.
"""
_stub_auto_tokenizer(monkeypatch)
old = _FakeTokenizer("old", spm_bytes = None)
new = _FakeTokenizer("new")
location = str(tmp_path / "_unsloth_sentencepiece_temp")
result = fix_sentencepiece_tokenizer(
old, new, {"</s>": "<|im_end|>"}, temporary_location = location
)
assert result is new
assert not any(
name.startswith("tokenizer_") for name in os.listdir(location)
), "the fast-only scratch dir was not reclaimed"
def test_each_call_uses_a_fresh_isolated_subdirectory(tmp_path, monkeypatch):
"""Each call must work in its own unique subdirectory, so concurrent or
repeated calls never share scratch files, stale artifacts never leak into
the reload, and nothing the caller left in the scratch location is deleted.
"""
loaded = _stub_auto_tokenizer(monkeypatch)
location = str(tmp_path / "_unsloth_sentencepiece_temp")
os.makedirs(location, exist_ok = True)
# A pre-existing artifact in the shared scratch location.
marker = os.path.join(location, "leftover.json")
with open(marker, "w") as f:
f.write("{}")
old1, new1 = _tokenizers()
old2, new2 = _tokenizers()
# Hold both returned tokenizers so their scratch dirs stay alive.
tok1 = fix_sentencepiece_tokenizer(
old1, new1, {"</s>": "<|im_end|>"}, temporary_location = location
)
tok2 = fix_sentencepiece_tokenizer(
old2, new2, {"</s>": "<|im_end|>"}, temporary_location = location
)
work1, work2 = loaded[0], loaded[1]
assert work1 != work2, "two calls reused the same directory"
assert os.path.dirname(work1) == location and os.path.dirname(work2) == location
assert os.path.isdir(work1) and os.path.isdir(work2)
# Nothing the caller left behind is deleted, and it never leaks into a work dir.
assert os.path.isfile(marker), "a pre-existing scratch file was deleted"
assert not os.path.isfile(os.path.join(work1, "leftover.json"))
assert not os.path.isfile(os.path.join(work2, "leftover.json"))
assert tok1 is not None and tok2 is not None
def test_sentencepiece_scratch_dir_is_reclaimed_once_the_tokenizer_is_gone(tmp_path, monkeypatch):
"""The scratch dir must live as long as the returned tokenizer (its vocab_file
points there), then be reclaimed when the tokenizer is garbage collected.
"""
loaded = _stub_auto_tokenizer(monkeypatch)
old, new = _tokenizers()
location = str(tmp_path / "_unsloth_sentencepiece_temp")
tok = fix_sentencepiece_tokenizer(old, new, {"</s>": "<|im_end|>"}, temporary_location = location)
work = loaded[-1]
assert os.path.isdir(work), "scratch dir vanished while the tokenizer was alive"
del tok
gc.collect()
assert not os.path.isdir(work), "scratch dir was not reclaimed after the tokenizer was freed"
class _CopyFromSubdirTokenizer:
"""A slow tokenizer whose sentencepiece source lives elsewhere (like the
tokenizers convert_to_fast_tokenizer produces under {location}/{name}).
save_pretrained copies that source into the destination, as HF slow
tokenizers copy their vocab_file.
"""
def __init__(self, source_model_path):
self.eos_token = "</s>"
self.pad_token = "<pad>"
self._source_model_path = source_model_path
def save_pretrained(self, location):
os.makedirs(location, exist_ok = True)
if os.path.isfile(self._source_model_path):
with open(self._source_model_path, "rb") as src:
data = src.read()
with open(os.path.join(location, "tokenizer.model"), "wb") as dst:
dst.write(data)
def __call__(
self,
texts,
add_special_tokens = False,
):
class _Encoded:
pass
encoded = _Encoded()
encoded.input_ids = [[2] for _ in texts]
return encoded
def test_source_vocab_outside_the_work_directory_is_not_disturbed(tmp_path, monkeypatch):
"""A tokenizer whose sentencepiece source lives elsewhere (e.g. the subtree
convert_to_fast_tokenizer created) is copied into the fresh work directory
and patched there; the original source is left untouched.
"""
loaded = _stub_auto_tokenizer(monkeypatch)
location = str(tmp_path / "_unsloth_sentencepiece_temp")
subdir = os.path.join(location, "some_model")
os.makedirs(subdir, exist_ok = True)
pieces = [("<s>", 0.0, CONTROL), ("a", -1.0, NORMAL), ("</s>", 0.0, CONTROL)]
source_model = os.path.join(subdir, "tokenizer.model")
with open(source_model, "wb") as f:
f.write(_spm_bytes(pieces))
old = _CopyFromSubdirTokenizer(source_model)
new = _FakeTokenizer("new")
tok = fix_sentencepiece_tokenizer(old, new, {"</s>": "<|im_end|>"}, temporary_location = location)
assert _read_pieces(source_model) == [
"<s>",
"a",
"</s>",
], "the original source vocab was modified"
assert "<|im_end|>" in _read_pieces(f"{loaded[-1]}/tokenizer.model")
assert tok is not None
def test_swap_mapping_swaps_both_pieces_without_duplicating(tmp_path, monkeypatch):
"""When the caller swaps eos and stop_word in the fast JSON it must pass both
directions here; a one-way mapping would leave two stop_word pieces and no eos.
"""
loaded = _stub_auto_tokenizer(monkeypatch)
location = str(tmp_path / "_unsloth_sentencepiece_temp")
pieces = [("<s>", 0.0, CONTROL), ("<|im_end|>", -1.0, NORMAL), ("</s>", 0.0, CONTROL)]
old = _FakeTokenizer("old", spm_bytes = _spm_bytes(pieces), vocab = {"</s>": 2, "<|im_end|>": 1})
new = _FakeTokenizer("new")
tok = fix_sentencepiece_tokenizer(
old, new, {"</s>": "<|im_end|>", "<|im_end|>": "</s>"}, temporary_location = location
)
result = _read_pieces(f"{loaded[-1]}/tokenizer.model")
assert result.count("<|im_end|>") == 1 and result.count("</s>") == 1, result
assert tok is not None
def test_only_applied_mappings_are_patched(tmp_path, monkeypatch):
"""When the caller skips a mapping whose target already exists, it must not
pass that mapping here, or the skipped source token gets renamed anyway and
duplicates the existing target in the model.
"""
loaded = _stub_auto_tokenizer(monkeypatch)
location = str(tmp_path / "_unsloth_sentencepiece_temp")
pieces = [
("<s>", 0.0, CONTROL),
("aa", -1.0, NORMAL),
("bb", -1.0, NORMAL),
("X", -1.0, NORMAL),
]
old = _FakeTokenizer("old", spm_bytes = _spm_bytes(pieces), vocab = {"aa": 1, "bb": 2})
new = _FakeTokenizer("new")
# Caller skipped aa->X (X already exists) and applied bb->Y, so only bb->Y is passed.
tok = fix_sentencepiece_tokenizer(old, new, {"bb": "Y"}, temporary_location = location)
result = _read_pieces(f"{loaded[-1]}/tokenizer.model")
assert result.count("X") == 1 and "Y" in result and "aa" in result, result
assert tok is not None

View file

@ -34,17 +34,22 @@ def test_model_selector_trigger_label_uses_leading_tight():
def test_sidebar_account_block_uses_leading_tight():
src = _read(APP_SIDEBAR)
# Match the account-block parent div regardless of its gap utility; this
# guard is about the leading-* class, not the spacing.
pattern = re.compile(
r'<div\s+className="flex\s+flex-col\s+gap-\S+\s+(\S+)\s+group-data-\[collapsible=icon\]:hidden">',
)
matches = pattern.findall(src)
class_names = re.findall(r'<div\s+className="([^"]+)"', src)
required = {
"flex",
"flex-1",
"flex-col",
"group-data-[collapsible=icon]:hidden",
}
matches = [classes for classes in class_names if required <= set(classes.split())]
assert matches, "could not find sidebar account-block parent div"
leading_classes = [m for m in matches if m.startswith("leading-")]
assert leading_classes, f"no leading-* class on sidebar account-block parent: {matches}"
for cls in leading_classes:
assert cls == "leading-tight", f"sidebar account-block must use leading-tight, got: {cls}"
for classes in matches:
leading_classes = [cls for cls in classes.split() if cls.startswith("leading-")]
assert leading_classes, f"no leading-* class on sidebar account-block parent: {classes}"
for cls in leading_classes:
assert (
cls == "leading-tight"
), f"sidebar account-block must use leading-tight, got: {cls}"
def test_no_truncate_plus_leading_none_in_changed_files():

View file

@ -295,7 +295,25 @@ def test_smart_chunk_text_single_chunk_no_eos_returns_plain_list():
return True
def test_load_from_file_skips_non_object_json_lines():
"""Non-object .jsonl lines (valid JSON, not dicts) are skipped, not fatal."""
# "context" contains "text", ["text"] holds it, 42 isn't iterable -- each
# would reach data[field] and raise TypeError without the isinstance guard.
with tempfile.NamedTemporaryFile("w", suffix = ".jsonl", delete = False) as f:
f.write('"context"\n["text", "x"]\n42\n{"text": "keep this"}\n')
path = f.name
try:
text = RawTextDataLoader(None)._read_file_by_format(path, "json_lines")
assert text == "keep this", text
finally:
os.unlink(path)
print("test_load_from_file_skips_non_object_json_lines passed")
return True
if __name__ == "__main__":
success = test_raw_text_loader()
success = test_smart_chunk_text_single_chunk_no_eos_returns_plain_list() and success
success = test_load_from_file_skips_non_object_json_lines() and success
sys.exit(0 if success else 1)

View file

@ -1929,6 +1929,9 @@ def get_chat_template(
string_vocab = tokenizer._tokenizer.to_str()
skipped = 0
# Only mirror applied mappings into the spm model; a skipped one would
# rename a piece the JSON never changed and desync the two.
applied_mapping = {}
for old_token, new_token in token_mapping.items():
old_count = string_vocab.count(f'"{old_token}"')
new_count = string_vocab.count(f'"{new_token}"')
@ -1939,6 +1942,7 @@ def get_chat_template(
raise RuntimeError(f"{old_token} was not part of the tokenizer!")
else:
string_vocab = string_vocab.replace(f'"{old_token}"', f'"{new_token}"')
applied_mapping[old_token] = new_token
pass
pass
@ -1973,7 +1977,7 @@ def get_chat_template(
# Must fix the sentence piece tokenizer since there's no tokenizer.model file!
from .tokenizer_utils import fix_sentencepiece_tokenizer
tokenizer = fix_sentencepiece_tokenizer(tokenizer, new_tokenizer, token_mapping,)
tokenizer = fix_sentencepiece_tokenizer(tokenizer, new_tokenizer, applied_mapping,)
else:
pass
@ -1997,8 +2001,11 @@ def get_chat_template(
string_vocab = string_vocab.replace(old_eos_token, temporary_stop_token)
string_vocab = string_vocab.replace(stop_word, old_eos_token)
string_vocab = string_vocab.replace(temporary_stop_token, stop_word)
# JSON swapped both, so swap both here too; a one-way map leaves two stop_word pieces.
sentencepiece_mapping = { old_eos_token : stop_word, stop_word : old_eos_token, }
else:
string_vocab = string_vocab.replace(old_eos_token, stop_word)
sentencepiece_mapping = { old_eos_token : stop_word, }
pass
new_tokenizer = tokenizer._tokenizer.from_str(string_vocab)
@ -2017,9 +2024,8 @@ def get_chat_template(
)
# Must fix the sentence piece tokenizer since there's no tokenizer.model file!
token_mapping = { old_eos_token : stop_word, }
from .tokenizer_utils import fix_sentencepiece_tokenizer
tokenizer = fix_sentencepiece_tokenizer(tokenizer, new_tokenizer, token_mapping,)
tokenizer = fix_sentencepiece_tokenizer(tokenizer, new_tokenizer, sentencepiece_mapping,)
pass
else:

View file

@ -236,6 +236,10 @@ class RawTextDataLoader:
def _extract_text_from_json(self, data):
"""Extract text from JSON object using common field names."""
# Skip non-object lines (str/list/number): `field in data` would be a
# substring/membership test, not a key lookup, and `data[field]` raises.
if not isinstance(data, dict):
return ""
for field in self._TEXT_FIELDS:
if field in data and isinstance(data[field], str):
return data[field]

View file

@ -325,7 +325,10 @@ def get_lora_parameters(proj):
if getattr(base_layer, "quant_method", None) == "fp8":
# we need to somehow store and pass this information :)
W.block_size = getattr(base_layer, "block_size", [128, 128])
W_quant.block_size = W.block_size
# A decompressed compressed-tensors layer keeps quant_method == "fp8" while its
# weight is back to bf16, so it has no quant state to carry the block size.
if W_quant is not None:
W_quant.block_size = W.block_size
# if not hasattr(proj, "disable_adapters") or proj.disable_adapters or proj.merged:
if getattr(proj, "disable_adapters", True) or proj.merged:
@ -375,14 +378,17 @@ def get_lora_parameters_bias(proj):
if W_quant is None:
W_quant = getattr(base_layer, "weight_scale", None)
# if not hasattr(proj, "disable_adapters") or proj.disable_adapters or proj.merged:
if getattr(proj, "disable_adapters", True) or proj.merged:
return W, W_quant, None, None, None, base_layer.bias
if getattr(base_layer, "quant_method", None) == "fp8":
# we need to somehow store and pass this information :)
W.block_size = getattr(base_layer, "block_size", [128, 128])
W_quant.block_size = W.block_size
# A decompressed compressed-tensors layer keeps quant_method == "fp8" while its
# weight is back to bf16, so it has no quant state to carry the block size.
if W_quant is not None:
W_quant.block_size = W.block_size
# if not hasattr(proj, "disable_adapters") or proj.disable_adapters or proj.merged:
if getattr(proj, "disable_adapters", True) or proj.merged:
return W, W_quant, None, None, None, base_layer.bias
adapter = getattr(proj, "active_adapters", None)
if adapter is None:

View file

@ -17,6 +17,9 @@ from transformers.convert_slow_tokenizer import convert_slow_tokenizer
from transformers import PreTrainedTokenizerFast
import re
import os
import shutil
import tempfile
import weakref
from transformers.models.llama.modeling_llama import logger
from peft import PeftModelForCausalLM
import torch
@ -370,13 +373,19 @@ def fix_sentencepiece_tokenizer(
if not os.path.exists(temporary_location):
os.makedirs(temporary_location)
# Check if tokenizer.model exists
if not os.path.isfile(f"{temporary_location}/tokenizer.model"):
return new_tokenizer
# Fresh per-call subdir so concurrent/repeated calls can't clobber each other's
# tokenizer.model or leak stale files, without deleting anything the caller owns.
temporary_location = tempfile.mkdtemp(prefix = "tokenizer_", dir = temporary_location)
# First save the old tokenizer
old_tokenizer.save_pretrained(temporary_location)
# Only sentencepiece tokenizers write tokenizer.model, so check after the save.
if not os.path.isfile(f"{temporary_location}/tokenizer.model"):
# new_tokenizer was built in memory and never references this dir, so drop it.
shutil.rmtree(temporary_location, ignore_errors = True)
return new_tokenizer
tokenizer_file = sentencepiece_model_pb2.ModelProto()
tokenizer_file.ParseFromString(open(f"{temporary_location}/tokenizer.model", "rb").read())
@ -414,6 +423,9 @@ def fix_sentencepiece_tokenizer(
eos_token = new_tokenizer.eos_token,
pad_token = new_tokenizer.pad_token,
)
# vocab_file points here, so the dir must outlive the tokenizer (a later
# save_pretrained copies the patched tokenizer.model from it); reclaim it on GC.
weakref.finalize(tokenizer, shutil.rmtree, temporary_location, ignore_errors = True)
return tokenizer