From 2139200b3f7eca44ff8c6451e318aaf6a034195a Mon Sep 17 00:00:00 2001 From: oobabooga Date: Fri, 17 Jul 2026 19:05:46 -0300 Subject: [PATCH 01/13] Studio: don't re-download updated GGUFs on load (#7209) --- studio/backend/core/inference/llama_cpp.py | 515 +++++++++--- .../backend/hub/services/models/downloads.py | 30 + studio/backend/hub/utils/download_registry.py | 27 +- studio/backend/routes/inference.py | 50 +- .../tests/test_gguf_load_cache_reuse.py | 740 ++++++++++++++++++ .../tests/test_mtp_drafter_companion.py | 28 + .../tests/test_offline_gguf_cache_fallback.py | 40 +- 7 files changed, 1264 insertions(+), 166 deletions(-) create mode 100644 studio/backend/tests/test_gguf_load_cache_reuse.py diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 3e93a26787..fec5214cf9 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -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(): diff --git a/studio/backend/hub/services/models/downloads.py b/studio/backend/hub/services/models/downloads.py index c2ffe7bffc..862af0141a 100644 --- a/studio/backend/hub/services/models/downloads.py +++ b/studio/backend/hub/services/models/downloads.py @@ -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. diff --git a/studio/backend/hub/utils/download_registry.py b/studio/backend/hub/utils/download_registry.py index 274038b292..b6bdee3bce 100644 --- a/studio/backend/hub/utils/download_registry.py +++ b/studio/backend/hub/utils/download_registry.py @@ -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, diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 0541a81ff2..81dea96274 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -20,6 +20,7 @@ from loggers import get_logger import asyncio import threading import weakref +from contextlib import ExitStack import re as _re @@ -998,6 +999,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 +1037,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, @@ -3968,6 +3971,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 +4191,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 +4273,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 +4669,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( diff --git a/studio/backend/tests/test_gguf_load_cache_reuse.py b/studio/backend/tests/test_gguf_load_cache_reuse.py new file mode 100644 index 0000000000..15d91cd324 --- /dev/null +++ b/studio/backend/tests/test_gguf_load_cache_reuse.py @@ -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--/snapshots//`` 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 diff --git a/studio/backend/tests/test_mtp_drafter_companion.py b/studio/backend/tests/test_mtp_drafter_companion.py index 86e528ae67..02230632b6 100644 --- a/studio/backend/tests/test_mtp_drafter_companion.py +++ b/studio/backend/tests/test_mtp_drafter_companion.py @@ -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 diff --git a/studio/backend/tests/test_offline_gguf_cache_fallback.py b/studio/backend/tests/test_offline_gguf_cache_fallback.py index e24e2ca451..0b9a1e704f 100644 --- a/studio/backend/tests/test_offline_gguf_cache_fallback.py +++ b/studio/backend/tests/test_offline_gguf_cache_fallback.py @@ -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 From 5441266e3aba6a3477b1e4d908f30da0a2665d3a Mon Sep 17 00:00:00 2001 From: Long Yixing Date: Sat, 18 Jul 2026 06:30:37 +0800 Subject: [PATCH 02/13] Fix Studio reasoning channel rendering (#7121) --------- Co-authored-by: oobabooga <112222186+oobabooga@users.noreply.github.com> --- .../core/inference/chat_template_helpers.py | 282 ++++++++++++- studio/backend/core/inference/inference.py | 377 ++++++++++++++---- .../backend/core/inference/mlx_inference.py | 105 +++-- studio/backend/core/inference/orchestrator.py | 76 +++- studio/backend/routes/inference.py | 183 ++++++++- .../tests/test_mlx_inference_backend.py | 166 ++++++++ .../tests/test_openai_tool_passthrough.py | 25 +- .../test_safetensors_capability_advertise.py | 53 +++ .../test_safetensors_reasoning_stream.py | 132 ++++++ .../tests/test_sf_client_tools_passthrough.py | 55 +++ .../tests/test_think_prefill_reemit.py | 177 +++++++- 11 files changed, 1479 insertions(+), 152 deletions(-) diff --git a/studio/backend/core/inference/chat_template_helpers.py b/studio/backend/core/inference/chat_template_helpers.py index 897db8262d..5113eebb36 100644 --- a/studio/backend/core/inference/chat_template_helpers.py +++ b/studio/backend/core/inference/chat_template_helpers.py @@ -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_CLOSE = "" +_GEMMA_CHANNEL_START = "<|channel>" +_GEMMA_THOUGHT_OPEN = "<|channel>thought" +_GEMMA_THOUGHT_CLOSE = "" +_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 ````. + + 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) diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index 172e9e5546..eae705d20e 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -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""" @@ -960,8 +1012,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 +1100,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 +1110,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 +1126,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 +1138,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 +1164,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 +1251,27 @@ class InferenceBackend: # Stream with TextIteratorStreamer + background thread try: - from core.inference.chat_template_helpers import detect_think_prefill - # Re-emit an open 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 +1293,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 +1306,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 +1324,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 +1343,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 +1509,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 +1548,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 +1640,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 +1650,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 +1663,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 prefilled by the template; # re-emit it so the frontend can render the thinking block. @@ -1494,30 +1674,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 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 +1701,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 +1720,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 +1739,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 +1758,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 +1787,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 +2290,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 .... Strip any # harmony protocol tokens and other gpt-oss tokens (e.g. @@ -2118,10 +2335,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("") + and text[: -len("")].endswith(token) + ): + text = text[: -len("") - len(token)] + "" return text.strip() def _load_chat_template_info(self, model_name: str): diff --git a/studio/backend/core/inference/mlx_inference.py b/studio/backend/core/inference/mlx_inference.py index 163ade10c4..e7a90b4307 100644 --- a/studio/backend/core/inference/mlx_inference.py +++ b/studio/backend/core/inference/mlx_inference.py @@ -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 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 + # 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, diff --git a/studio/backend/core/inference/orchestrator.py b/studio/backend/core/inference/orchestrator.py index c2082bc198..ea3b164052 100644 --- a/studio/backend/core/inference/orchestrator.py +++ b/studio/backend/core/inference/orchestrator.py @@ -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: @@ -1414,12 +1445,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: @@ -1489,11 +1535,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 +1556,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 +1741,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 +1753,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()) diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index 81dea96274..d6090e2b62 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -29,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, @@ -213,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. @@ -1918,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); @@ -1941,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 " @@ -5424,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" @@ -5437,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(): @@ -7060,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( @@ -7075,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) @@ -7093,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 @@ -8637,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 ; 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 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(): @@ -8858,6 +8945,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. @@ -8945,6 +9044,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: @@ -8994,6 +9098,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", ""), @@ -9034,6 +9147,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. @@ -9120,6 +9242,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 = {} @@ -9200,6 +9328,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() @@ -9349,6 +9485,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 reasoning (GGUF parity); also covers MLX via @@ -9447,6 +9588,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) diff --git a/studio/backend/tests/test_mlx_inference_backend.py b/studio/backend/tests/test_mlx_inference_backend.py index 29fbb45158..fa50cd84d6 100644 --- a/studio/backend/tests/test_mlx_inference_backend.py +++ b/studio/backend/tests/test_mlx_inference_backend.py @@ -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", ""), + ), + 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("", 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, + ) + ) == ["", "r", "r", "ra"] + + gen = backend.generate_chat_response( + messages = [{"role": "user", "content": "ping"}], + max_new_tokens = 4, + ) + assert next(gen) == "" + 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\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", ""), + ), + 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("", 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 == [ + "\n", + "\nreason", + "\nreason", + "\nreasonanswer", + ] + 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("", 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...", + 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, + ) + ) == [ + "", + "vision", + "vision", + "vision answer", + ] diff --git a/studio/backend/tests/test_openai_tool_passthrough.py b/studio/backend/tests/test_openai_tool_passthrough.py index 5300a48557..8725b28ac8 100644 --- a/studio/backend/tests/test_openai_tool_passthrough.py +++ b/studio/backend/tests/test_openai_tool_passthrough.py @@ -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): diff --git a/studio/backend/tests/test_safetensors_capability_advertise.py b/studio/backend/tests/test_safetensors_capability_advertise.py index 0ed670ac01..bd3d8d16b9 100644 --- a/studio/backend/tests/test_safetensors_capability_advertise.py +++ b/studio/backend/tests/test_safetensors_capability_advertise.py @@ -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 , 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{% 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{% 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 %}{% 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 ``\n...``. Faithful slice so the # classifier never silently regresses for this family. diff --git a/studio/backend/tests/test_safetensors_reasoning_stream.py b/studio/backend/tests/test_safetensors_reasoning_stream.py index 4e708139b7..af5a05d266 100644 --- a/studio/backend/tests/test_safetensors_reasoning_stream.py +++ b/studio/backend/tests/test_safetensors_reasoning_stream.py @@ -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..." + all_special_tokens = [] + eos_token_id = 1 + pad_token_id = None + pieces = {10: "<|channel>thought\n", 11: "r", 12: "", 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] == "ra" + + 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..." + 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 diff --git a/studio/backend/tests/test_sf_client_tools_passthrough.py b/studio/backend/tests/test_sf_client_tools_passthrough.py index 01905b712c..f91eec9817 100644 --- a/studio/backend/tests/test_sf_client_tools_passthrough.py +++ b/studio/backend/tests/test_sf_client_tools_passthrough.py @@ -177,6 +177,15 @@ def _sse_objects(chunks): # ── Non-streaming ───────────────────────────────────────────────── +def test_non_reasoning_backend_keeps_literal_think_tags(monkeypatch): + backend = _ScriptedBackend(_fixed("show example tags")) + response = _call(_request(stream = False), monkeypatch, backend, supports_tools = False) + + message = _json_body(response)["choices"][0]["message"] + assert message["content"] == "show example 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 "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)) diff --git a/studio/backend/tests/test_think_prefill_reemit.py b/studio/backend/tests/test_think_prefill_reemit.py index 300ff92776..346399c3b2 100644 --- a/studio/backend/tests/test_think_prefill_reemit.py +++ b/studio/backend/tests/test_think_prefill_reemit.py @@ -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 ``\\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 + "\n", None) == "\n" assert detect_think_prefill(QWEN_PROMPT + "\n", []) == "\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 = "" + + class NamedTemplateProcessor: + chat_template = { + "default": "plain processor default", + "tool_use": "<|channel>thought\nprocessor tool template", + } + tokenizer = TokenMetadataOnly() + + def apply_chat_template(self, *_args, **_kwargs): + raise NotImplementedError + + expected = ("<|channel>thought", "") + 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" + + assert detect_reasoning_channel_markers( + FallbackTokenizer(), tools = [{"function": {"name": "web_search"}}] + ) == ("<|channel>thought", "") + + +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", "") + + +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"}} + + assert detect_reasoning_channel_markers_from_model_info( + SimpleNamespace(chat_template = "override has no native markers"), + model_info, + tools = None, + ) == ("<|channel>thought", "") + 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", "") + + +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"}, + "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", "") + output = "" + snapshots = [] + for chunk in ( + "<|chan", + "nel>thought", + "\nReason", + "<|tool_call>web_search", + ): + delta = parser.feed(chunk) + if delta: + output += delta + snapshots.append(output) + + assert snapshots == [ + "", + "Reason", + "Reason<|tool_call>web_search", + ] + assert snapshots[1].startswith(snapshots[0]) + compact = ReasoningChannelNormalizer("<|channel>thought", "") + assert compact.feed("<|channel>thoughtanswer") + compact.finish() == ( + "answer" + ) From 8ff2f8e70c4872ed83d2dde0a3db40ad6dcb24f8 Mon Sep 17 00:00:00 2001 From: Long Yixing Date: Sat, 18 Jul 2026 07:22:11 +0800 Subject: [PATCH 03/13] fix(studio): ignore reasoning in tool reprompts (#7134) --- studio/backend/core/inference/inference.py | 2 + studio/backend/core/inference/orchestrator.py | 2 + .../core/inference/safetensors_agentic.py | 52 ++++- studio/backend/routes/inference.py | 1 + .../tests/test_safetensors_tool_loop.py | 190 ++++++++++++++++++ 5 files changed, 243 insertions(+), 4 deletions(-) diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index eae705d20e..8d262bbb0f 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -888,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``. @@ -941,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( diff --git a/studio/backend/core/inference/orchestrator.py b/studio/backend/core/inference/orchestrator.py index ea3b164052..3afda74411 100644 --- a/studio/backend/core/inference/orchestrator.py +++ b/studio/backend/core/inference/orchestrator.py @@ -1407,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, @@ -1487,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( diff --git a/studio/backend/core/inference/safetensors_agentic.py b/studio/backend/core/inference/safetensors_agentic.py index 43b72110ff..9110315815 100644 --- a/studio/backend/core/inference/safetensors_agentic.py +++ b/studio/backend/core/inference/safetensors_agentic.py @@ -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]`` 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( { diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index d6090e2b62..52ea1f86a3 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -8886,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() diff --git a/studio/backend/tests/test_safetensors_tool_loop.py b/studio/backend/tests/test_safetensors_tool_loop.py index e3633de289..915f82ac8e 100644 --- a/studio/backend/tests/test_safetensors_tool_loop.py +++ b/studio/backend/tests/test_safetensors_tool_loop.py @@ -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 ( + "Let me prepare the requested summary carefully." + "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.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." + "more private planningThis 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" + "Let me prepare the requested summary carefully." + "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 = [ + ["Let me search for that."], + ['{"name":"web_search","arguments":{"query":"cats"}}'], + ["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.Let me search for that."], + ['{"name":"web_search","arguments":{"query":"cats"}}'], + ["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.checking details"], + ['{"name":"web_search","arguments":{"query":"cats"}}'], + ["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 "private planning detailsLet me search for that." + elif len(captured) == 2: + yield '{"name":"web_search","arguments":{"query":"cats"}}' + 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. From bf4185a2d3fa3d1c57bcfdb43ebc25912316e568 Mon Sep 17 00:00:00 2001 From: Nilay <118994073+NilayYadav@users.noreply.github.com> Date: Sat, 18 Jul 2026 04:57:28 +0530 Subject: [PATCH 04/13] Studio: don't apply nest_asyncio on plain CLI starts (breaks asyncio on Python 3.14+) (#7186) * Studio: don't apply nest_asyncio on plain CLI starts (breaks asyncio on Python 3.14+) * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Skip nest_asyncio on Python 3.14+ so notebook and embedded Studio starts also work * Tighten the nest_asyncio gate comment --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: danielhanchen --- studio/backend/run.py | 20 ++++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/studio/backend/run.py b/studio/backend/run.py index 56b9c78343..4f105b53b1 100644 --- a/studio/backend/run.py +++ b/studio/backend/run.py @@ -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 From a14b032d79ddd8f10bd874e741d3909945d1163f Mon Sep 17 00:00:00 2001 From: Vineeth Sai Date: Fri, 17 Jul 2026 16:30:45 -0700 Subject: [PATCH 05/13] Propagate fp8 block_size before the early return in get_lora_parameters_bias (#7189) * Propagate fp8 block_size before the early return in get_lora_parameters_bias get_lora_parameters_bias set the fp8 block_size on W/W_quant only after the disable_adapters/merged early return, so on the merged or disabled path (merged inference, DPO reference model) a block-fp8 weight lost its real block_size and downstream fp8 kernels fell back to [128, 128]. The non-bias sibling get_lora_parameters already sets block_size before its early return; move the block so both behave the same. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Guard the fp8 block_size against a missing quant state A decompressed compressed-tensors layer keeps quant_method == "fp8" while its weight is back to bf16, so it has no quant state and get_lora_parameters_bias must still return W_quant None for fast_linear_forward to fall back to a plain matmul. Only attach block_size when a quant state was actually found. * Guard the sibling get_lora_parameters fp8 block_size against a missing quant state Mirror the get_lora_parameters_bias guard so a decompressed compressed-tensors layer (quant_method fp8, bf16 weight, no quant state) does not raise AttributeError on the fused-LoRA path. Add a CPU-only regression test. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han --- ...get_lora_parameters_bias_fp8_block_size.py | 82 +++++++++++++++++++ ...test_get_lora_parameters_fp8_block_size.py | 81 ++++++++++++++++++ unsloth/kernels/utils.py | 18 ++-- 3 files changed, 175 insertions(+), 6 deletions(-) create mode 100644 tests/python/test_get_lora_parameters_bias_fp8_block_size.py create mode 100644 tests/python/test_get_lora_parameters_fp8_block_size.py diff --git a/tests/python/test_get_lora_parameters_bias_fp8_block_size.py b/tests/python/test_get_lora_parameters_bias_fp8_block_size.py new file mode 100644 index 0000000000..835ad19542 --- /dev/null +++ b/tests/python/test_get_lora_parameters_bias_fp8_block_size.py @@ -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] diff --git a/tests/python/test_get_lora_parameters_fp8_block_size.py b/tests/python/test_get_lora_parameters_fp8_block_size.py new file mode 100644 index 0000000000..f5f1359125 --- /dev/null +++ b/tests/python/test_get_lora_parameters_fp8_block_size.py @@ -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] diff --git a/unsloth/kernels/utils.py b/unsloth/kernels/utils.py index 1b0b5ce12e..ccfedfdef0 100644 --- a/unsloth/kernels/utils.py +++ b/unsloth/kernels/utils.py @@ -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: From a8ff8673df640710e58fabfb226df565d9751afe Mon Sep 17 00:00:00 2001 From: Ritwij Aryan Parmar <88580521+RitwijParmar@users.noreply.github.com> Date: Fri, 17 Jul 2026 19:50:39 -0400 Subject: [PATCH 06/13] feat(studio): expose an opt-in MCP control plane (#7191) * feat(studio): expose opt-in MCP control plane * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Harden Studio MCP tools: byte-safe auth, page clamping, forward export/checkpoint fields Follow-up hardening on the opt-in MCP control plane. All changes are additive and backwards compatible. - BearerTokenMiddleware now compares the Authorization header on raw bytes. A non-ASCII bearer value previously reached str-based hmac.compare_digest, which raises TypeError and surfaced as a 500 instead of a clean 401. The constructor also rejects an empty or whitespace-only token so an empty token can never match an empty "Bearer " header. - MCP tools call the route functions directly, which skips FastAPI Query validation. list_training_runs and get_recipe_job_dataset now clamp limit and offset to the same bounds the HTTP routes enforce (a negative SQLite LIMIT otherwise means "no limit"). - export_gguf forwards hf_token (the backend rejects a Hub upload without it), accepts a list of quantization methods, and exposes imatrix / imatrix_path so the IQ low-bit quants are reachable. - load_checkpoint forwards hf_token and approved_remote_code_fingerprint so gated checkpoints and the remote-code approval retry work. Its docstring is corrected: the export backend coexists with training and inference rather than freeing GPU work. - start_training passes via_api_key=False explicitly instead of relying on the unfilled Depends default. Tests: add coverage for non-ASCII and empty-token auth, the correct-token pass through, non-http scope pass through, pagination clamping, and the forwarded export/checkpoint fields. * Harden Studio MCP: cap /mcp request bodies, reject unusable tokens, fix docs Follow-up hardening from a full review pass. All changes are additive and backwards compatible. - Add "/mcp" to _BODY_PROTECTED_PREFIXES so MaxBodyMiddleware enforces the same request-body cap it already applies to every other write endpoint (/api/train, /api/export, /api/data-recipe, ...). The MCP endpoint accepts authenticated POST tool-call bodies; without this an authenticated client could send an unbounded body. The middleware only buffers the request body (not the SSE response), so streaming is unaffected, and the 500MB default cap never affects a real JSON-RPC tool call (verified live). - Reject a non-ASCII UNSLOTH_STUDIO_MCP_TOKEN at construction. HTTP header values are ASCII, so a non-ASCII token cannot be sent by a standard client and would silently lock out the endpoint; fail fast instead. - MCP.md: document the canonical /mcp/ endpoint and note that /mcp redirects to it, so clients that do not follow redirected POSTs still connect. Tests: add non-ASCII token rejection coverage. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Tighten MCP server comments --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: danielhanchen --- studio/MCP.md | 34 +++ studio/backend/main.py | 17 ++ studio/backend/mcp_server.py | 259 +++++++++++++++++++++ studio/backend/tests/test_mcp_server.py | 290 ++++++++++++++++++++++++ 4 files changed, 600 insertions(+) create mode 100644 studio/MCP.md create mode 100644 studio/backend/mcp_server.py create mode 100644 studio/backend/tests/test_mcp_server.py diff --git a/studio/MCP.md b/studio/MCP.md new file mode 100644 index 0000000000..91b39fcc69 --- /dev/null +++ b/studio/MCP.md @@ -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. \ No newline at end of file diff --git a/studio/backend/main.py b/studio/backend/main.py index e64048dc00..6e16dc00ca 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -612,6 +612,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 @@ -752,6 +768,7 @@ _BODY_PROTECTED_PREFIXES = ( "/api/settings", "/api/train", "/api/export", + "/mcp", ) _DATASET_UPLOAD_PASSTHROUGH_PREFIX = "/api/datasets/upload" _DATA_RECIPE_UNSTRUCTURED_UPLOAD_PASSTHROUGH_PREFIX = ( diff --git a/studio/backend/mcp_server.py b/studio/backend/mcp_server.py new file mode 100644 index 0000000000..f837f46425 --- /dev/null +++ b/studio/backend/mcp_server.py @@ -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 + ) diff --git a/studio/backend/tests/test_mcp_server.py b/studio/backend/tests/test_mcp_server.py new file mode 100644 index 0000000000..71792605ae --- /dev/null +++ b/studio/backend/tests/test_mcp_server.py @@ -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 From 4e09328c3bb42dcdcfe1f3de916803ad0ecc016d Mon Sep 17 00:00:00 2001 From: Andrew Chen <48723787+chuenchen309@users.noreply.github.com> Date: Sat, 18 Jul 2026 20:53:33 +0800 Subject: [PATCH 07/13] fix(tokenizer): check for tokenizer.model after saving it, not before (#7194) * fix(tokenizer): check for tokenizer.model after saving it, not before `fix_sentencepiece_tokenizer` creates its temporary directory, then returns early unless that directory already contains a tokenizer.model: if not os.path.exists(temporary_location): os.makedirs(temporary_location) # fresh, empty if not os.path.isfile(f"{temporary_location}/tokenizer.model"): return new_tokenizer # always true old_tokenizer.save_pretrained(temporary_location) # writes that file The file only appears on the line after the check, so the guard is always true and the body never runs. Nothing else writes that path either -- `convert_to_fast_tokenizer` saves into a per-name subdirectory, not `{temporary_location}/tokenizer.model`. Both call sites are in `get_chat_template` and are commented "Must fix the sentence piece tokenizer since there's no tokenizer.model file!" -- the guard defeats the exact intent the caller states. The effect is silent: the caller still gets a working `new_tokenizer`, but the sentencepiece piece rename is skipped, so the mapped token (e.g. the eos token remapped to `<|im_end|>`) is missing from tokenizer.model and GGUF/llama.cpp exports carry the old piece. `check_if_sentencepiece_model` in save.py does the same probe in the right order -- makedirs, save_pretrained, then isfile. Match it. Tests are added under tests/saving/ next to the existing sentencepiece coverage, and to the two Bucket-A lists in consolidated-tests-ci.yml, since Repo tests (CPU) --ignores tests/saving and these need protobuf. Co-Authored-By: Claude Opus 4.8 (1M context) * Clear stale tokenizer.model before the sentencepiece guard The guard now runs after old_tokenizer.save_pretrained, but the default temporary_location is a fixed reusable directory. A fast-only tokenizer writes no tokenizer.model, so a stale file from an earlier sentencepiece call could pass the guard and patch the wrong model (e.g. mixing models in one process, like a long-running server). Remove any existing tokenizer.model first, and add a regression test. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Empty the reusable sentencepiece scratch directory each call The final AutoTokenizer.from_pretrained reloads the whole temporary_location, so removing only a stale tokenizer.model still let other artifacts from a previous tokenizer (added_tokens.json, chat template, etc.) leak into the reload when the default reusable directory is used across models in one process. Recreate the directory instead, and add a regression test for the leaked-artifact case. * Clear only top-level scratch files, keep subdirectories Recreating the whole reusable directory deleted the {name} subtree that convert_to_fast_tokenizer stores a converted tokenizer's source vocab in, so old_tokenizer.save_pretrained could not copy tokenizer.model and the guard returned the tokenizer unpatched for those legacy converted tokenizers. Remove only stale top-level files (all the final reload reads) and leave subdirectories intact. Add a regression test for the converted-source subdirectory. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Keep the current tokenizer's own source vocab when clearing On a repeated get_chat_template(map_eos_token=True) call, the returned tokenizer's vocab_file points back at the top-level tokenizer.model, and the cleanup deleted that source before old_tokenizer.save_pretrained could re-emit it, so the guard returned the tokenizer unpatched. Skip removing the old tokenizer's own source vocab while still clearing stale files from a different tokenizer, and add a regression test. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Use a per-call temporary directory for the sentencepiece fix The scratch directory defaulted to a single shared path, so concurrent or repeated get_chat_template(map_eos_token=True) calls could delete or overwrite each other's tokenizer.model between save and reload (tripping the piece assertion or reloading the wrong model), and stale files from an earlier tokenizer could leak into the reload. Work in a unique per-call subdirectory instead: this isolates every call without deleting anything the caller owns, and replaces the earlier per-file cleanup. Tests updated to read the patched model from the reloaded directory and to cover isolation and source-vocab preservation. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Pass only the applied token mappings into the sentencepiece fix get_chat_template mirrors token remaps into tokenizer.model via fix_sentencepiece_tokenizer, but two caller paths passed a mapping that did not match what they wrote to the fast tokenizer JSON, so once the sentencepiece patch runs the model and JSON disagree: - the mapped-token path skipped entries whose target already existed but still passed the full mapping, renaming a piece the JSON never changed; - the EOS-swap path swapped both tokens in the JSON but passed only one direction, leaving two stop_word pieces and no old EOS piece. Pass the applied mapping (and both swap directions) instead. Add regression tests. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Tighten sentencepiece guard comments * Add SPDX license identifier to sentencepiece guard test * Reclaim the per-call sentencepiece scratch directory The per-call tempfile.mkdtemp fixed the shared-directory race but never cleaned up, so a long-running process leaked one scratch dir per call. The dir cannot be deleted eagerly for sentencepiece tokenizers because the returned tokenizer's vocab_file points into it (a later save_pretrained copies the patched tokenizer.model from there). Reclaim it correctly instead: remove the dir right away on the fast-only path (the returned tokenizer never references it), and attach a weakref.finalize so the sentencepiece dir is removed once its tokenizer is garbage collected. Add regression tests for both. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Tighten the scratch-dir reclaim comment --------- Co-authored-by: Claude Opus 4.8 (1M context) Co-authored-by: danielhanchen Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: danielhanchen --- .github/workflows/consolidated-tests-ci.yml | 2 + .../test_fix_sentencepiece_tokenizer_guard.py | 307 ++++++++++++++++++ unsloth/chat_templates.py | 12 +- unsloth/tokenizer_utils.py | 18 +- 4 files changed, 333 insertions(+), 6 deletions(-) create mode 100644 tests/saving/test_fix_sentencepiece_tokenizer_guard.py diff --git a/.github/workflows/consolidated-tests-ci.yml b/.github/workflows/consolidated-tests-ci.yml index fa84471d36..d1bea819eb 100644 --- a/.github/workflows/consolidated-tests-ci.yml +++ b/.github/workflows/consolidated-tests-ci.yml @@ -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 \ diff --git a/tests/saving/test_fix_sentencepiece_tokenizer_guard.py b/tests/saving/test_fix_sentencepiece_tokenizer_guard.py new file mode 100644 index 0000000000..1ee523d57b --- /dev/null +++ b/tests/saving/test_fix_sentencepiece_tokenizer_guard.py @@ -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 = "" + self.pad_token = "" + 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 = [("", 0.0, CONTROL), ("a", -1.0, NORMAL), ("", 0.0, CONTROL)] + old = _FakeTokenizer("old", spm_bytes = _spm_bytes(pieces), vocab = {"": 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, {"": "<|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, {"": "<|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, {"": "<|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, {"": "<|im_end|>"}, temporary_location = location + ) + tok2 = fix_sentencepiece_tokenizer( + old2, new2, {"": "<|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, {"": "<|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 = "" + self.pad_token = "" + 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 = [("", 0.0, CONTROL), ("a", -1.0, NORMAL), ("", 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, {"": "<|im_end|>"}, temporary_location = location) + + assert _read_pieces(source_model) == [ + "", + "a", + "", + ], "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 = [("", 0.0, CONTROL), ("<|im_end|>", -1.0, NORMAL), ("", 0.0, CONTROL)] + old = _FakeTokenizer("old", spm_bytes = _spm_bytes(pieces), vocab = {"": 2, "<|im_end|>": 1}) + new = _FakeTokenizer("new") + + tok = fix_sentencepiece_tokenizer( + old, new, {"": "<|im_end|>", "<|im_end|>": ""}, temporary_location = location + ) + + result = _read_pieces(f"{loaded[-1]}/tokenizer.model") + assert result.count("<|im_end|>") == 1 and result.count("") == 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 = [ + ("", 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 diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py index 2d3674fb04..bd612db1d5 100644 --- a/unsloth/chat_templates.py +++ b/unsloth/chat_templates.py @@ -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: diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py index 3a91ef188d..d6e8247ec4 100644 --- a/unsloth/tokenizer_utils.py +++ b/unsloth/tokenizer_utils.py @@ -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 From 9db639f708f52906ecc7ada0939063d6e5b9b166 Mon Sep 17 00:00:00 2001 From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Date: Sat, 18 Jul 2026 05:54:00 -0700 Subject: [PATCH 08/13] Stabilize Studio regression tests (#7192) * Stabilize Studio regression tests Rebuild on current main. Restore the set-membership sidebar account-block matcher (#6647, which fixed the same order-sensitive regex, was reverted on main, so the guard is failing on main again) and keep the watchdog replacement-race fix, whose blocked-watchdog stub now waits without a timeout so a superseded watchdog stays alive until cleanup regardless of scheduler load. * Tighten the blocked-watchdog stub comment --------- Co-authored-by: Daniel Han --- .../tests/test_training_stop_watchdog.py | 22 ++++++++++++---- .../test_studio_text_descender_clipping.py | 25 +++++++++++-------- 2 files changed, 32 insertions(+), 15 deletions(-) diff --git a/studio/backend/tests/test_training_stop_watchdog.py b/studio/backend/tests/test_training_stop_watchdog.py index 457dfc8ea2..0cd702bce2 100644 --- a/studio/backend/tests/test_training_stop_watchdog.py +++ b/studio/backend/tests/test_training_stop_watchdog.py @@ -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) diff --git a/tests/studio/test_studio_text_descender_clipping.py b/tests/studio/test_studio_text_descender_clipping.py index 7cad6cacfe..73d1244ee3 100644 --- a/tests/studio/test_studio_text_descender_clipping.py +++ b/tests/studio/test_studio_text_descender_clipping.py @@ -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'', - ) - matches = pattern.findall(src) + class_names = re.findall(r' Date: Sat, 18 Jul 2026 20:54:50 +0800 Subject: [PATCH 09/13] fix(dataprep): skip .jsonl lines that are valid JSON but not objects (#7195) * fix(dataprep): skip .jsonl lines that are valid JSON but not objects `_read_file_by_format` json.loads each line and hands the result to `_extract_text_from_json`, which assumes a dict: for field in self._TEXT_FIELDS: if field in data and isinstance(data[field], str): A JSON line does not have to be an object -- `"context"`, `["text"]` and `42` are all valid JSON. For those, `field in data` stops being a key lookup and becomes a substring/membership test, so `data[field]` raises: "context" -> "text" in "context" is True (substring!) -> TypeError: string indices must be integers ["text", "foo"] -> TypeError: list indices must be integers 42 -> TypeError: argument of type 'int' is not iterable The TypeError escapes past `except json.JSONDecodeError: continue`, so the whole load dies on one odd line. That except clause is also the tell: a *malformed* line is already skipped gracefully. A *well-formed* line that happens not to be an object should be too -- it carries no text either way. This makes the two agree. Reachable from `unsloth-cli.py:253` (`--dataset foo.jsonl` auto-detect) and `RawTextDataLoader` is exported from `unsloth/__init__.py`. Co-Authored-By: Claude Opus 4.8 (1M context) * Slim the non-object jsonl regression test and shorten the guard comment --------- Co-authored-by: Claude Opus 4.8 (1M context) Co-authored-by: Daniel Han --- tests/test_raw_text.py | 18 ++++++++++++++++++ unsloth/dataprep/raw_text.py | 4 ++++ 2 files changed, 22 insertions(+) diff --git a/tests/test_raw_text.py b/tests/test_raw_text.py index ba16e0cfc4..18549adfe8 100644 --- a/tests/test_raw_text.py +++ b/tests/test_raw_text.py @@ -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) diff --git a/unsloth/dataprep/raw_text.py b/unsloth/dataprep/raw_text.py index 128d966ecd..8623285a25 100644 --- a/unsloth/dataprep/raw_text.py +++ b/unsloth/dataprep/raw_text.py @@ -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] From 9073f07705488601ab4ff59e3828fe698083894e Mon Sep 17 00:00:00 2001 From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Date: Sat, 18 Jul 2026 21:08:55 -0700 Subject: [PATCH 10/13] fix(studio): equal padding in the dataset source segmented control (#7230) * fix(studio): equal padding in dataset source segmented control * fix(studio): scope dataset source pill layoutId per component instance --- .../studio/sections/dataset-section.tsx | 43 ++++++++++++------- 1 file changed, 28 insertions(+), 15 deletions(-) diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index e7bab1f47f..6aa9329609 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -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 (
-
From d8aa0df66e728cacd677773adb9c4a9ce66cd13a Mon Sep 17 00:00:00 2001 From: Nilay <118994073+NilayYadav@users.noreply.github.com> Date: Sun, 19 Jul 2026 11:09:49 +0530 Subject: [PATCH 11/13] Studio: keep stale canvas from surviving into a new chat (#7229) --- .../src/features/chat/artifacts/artifact-card.tsx | 7 +++++-- studio/frontend/src/features/chat/chat-page.tsx | 10 ++++------ 2 files changed, 9 insertions(+), 8 deletions(-) diff --git a/studio/frontend/src/features/chat/artifacts/artifact-card.tsx b/studio/frontend/src/features/chat/artifacts/artifact-card.tsx index ee8c26abf1..0345dc6e2a 100644 --- a/studio/frontend/src/features/chat/artifacts/artifact-card.tsx +++ b/studio/frontend/src/features/chat/artifacts/artifact-card.tsx @@ -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; } diff --git a/studio/frontend/src/features/chat/chat-page.tsx b/studio/frontend/src/features/chat/chat-page.tsx index 380ce0e0ab..ec0ad977bf 100644 --- a/studio/frontend/src/features/chat/chat-page.tsx +++ b/studio/frontend/src/features/chat/chat-page.tsx @@ -247,13 +247,13 @@ const SingleContent = memo(function SingleContent({ useState(false); const [isArtifactSurfaceVisible, setIsArtifactSurfaceVisible] = useState(false); + // Without a URL threadId the artifact must belong to the active thread. const showArtifactPanel = 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 artifactLayoutActive = showArtifactPanel || isArtifactPanelLayoutActive; @@ -1771,10 +1771,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 From 95fa3fbe30198917bdcce4000881530b9adab079 Mon Sep 17 00:00:00 2001 From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Date: Sat, 18 Jul 2026 22:47:00 -0700 Subject: [PATCH 12/13] Allow API key for Ollama connections (#7173) The Connections form hid the API key field for the Ollama preset, which blocked Ollama cloud (it requires a key). Show the optional field for Ollama; the backend already sends Authorization: Bearer when a key is set and omits the header when empty, so local keyless servers are unaffected. Fixes #7163 --- studio/backend/core/inference/providers.py | 5 +++-- studio/frontend/src/features/chat/chat-providers-dialog.tsx | 4 ++-- studio/frontend/src/features/chat/external-providers.ts | 5 +++-- 3 files changed, 8 insertions(+), 6 deletions(-) diff --git a/studio/backend/core/inference/providers.py b/studio/backend/core/inference/providers.py index 5b72373c03..d3bffc2f3d 100644 --- a/studio/backend/core/inference/providers.py +++ b/studio/backend/core/inference/providers.py @@ -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, }, diff --git a/studio/frontend/src/features/chat/chat-providers-dialog.tsx b/studio/frontend/src/features/chat/chat-providers-dialog.tsx index 95e5cfbd79..e39955e576 100644 --- a/studio/frontend/src/features/chat/chat-providers-dialog.tsx +++ b/studio/frontend/src/features/chat/chat-providers-dialog.tsx @@ -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); diff --git a/studio/frontend/src/features/chat/external-providers.ts b/studio/frontend/src/features/chat/external-providers.ts index bc718abbba..eb9e4656b0 100644 --- a/studio/frontend/src/features/chat/external-providers.ts +++ b/studio/frontend/src/features/chat/external-providers.ts @@ -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. */ From c2cf2b4a1e023f3e9de80b8889971f4d04ded9c7 Mon Sep 17 00:00:00 2001 From: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Date: Sat, 18 Jul 2026 23:15:13 -0700 Subject: [PATCH 13/13] Studio: keep the permission pill label when composer pills collapse (#7231) With 5 or more pills active the composer collapses every pill to an icon, which hid the Bypass permissions label behind a small glyph. Exempt the permission pill via data-keep-label so it always shows its label, with the collapsed icons lining up to its right. Since the pill is never icon-only now, drop the compact-mode fallthrough in the glyph off switch so it works while the other pills are collapsed. --- .../src/features/chat/permission-mode-select.tsx | 13 ++++--------- studio/frontend/src/index.css | 4 +++- 2 files changed, 7 insertions(+), 10 deletions(-) diff --git a/studio/frontend/src/features/chat/permission-mode-select.tsx b/studio/frontend/src/features/chat/permission-mode-select.tsx index 4277c1bfcf..a9cb8ce5d1 100644 --- a/studio/frontend/src/features/chat/permission-mode-select.tsx +++ b/studio/frontend/src/features/chat/permission-mode-select.tsx @@ -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. */} { - 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"); }} diff --git a/studio/frontend/src/index.css b/studio/frontend/src/index.css index 333ed70f6a..6d4c21eec8 100644 --- a/studio/frontend/src/index.css +++ b/studio/frontend/src/index.css @@ -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) {