Merge remote-tracking branch 'origin/main' into woa-nvidia-wsl-fallback

This commit is contained in:
Daniel Han 2026-06-21 02:00:04 -07:00
commit 0a2b94de71
36 changed files with 1816 additions and 150 deletions

View file

@ -86,6 +86,8 @@ unsloth studio -p 8888
```
For cloud or global access, add `-H 0.0.0.0`. By default, Unsloth is accessible only locally.
For a secure HTTPS link instead of a raw network port, use `unsloth studio --secure`. Studio stays bound to localhost and is served only through a free Cloudflare HTTPS tunnel (it fails closed if the tunnel can't start, so the raw port is never exposed).
#### Docker
Use our [Docker image](https://hub.docker.com/r/unsloth/unsloth) ```unsloth/unsloth``` container. Run:
```bash
@ -162,13 +164,19 @@ Read our [guide](https://unsloth.ai/docs/get-started/fine-tuning-llms-guide). Ad
## 📥 Advanced Installation
The below advanced instructions are for Unsloth Studio. For Unsloth Core advanced installation, [view our docs](https://unsloth.ai/docs/get-started/install/pip-install#advanced-pip-installation).
#### Developer installs: macOS, Linux, WSL:
#### Developer / Nightly / Experimental installs: macOS, Linux, WSL:
The developer install builds from the `main` branch, which is the latest (nightly) source.
```bash
git clone https://github.com/unslothai/unsloth
cd unsloth
./install.sh --local
unsloth studio -p 8888
```
To install into an isolated location (its own virtual env, `auth/`, `studio.db`, cache and llama.cpp build), set `UNSLOTH_STUDIO_HOME` and pass it again at launch:
```bash
UNSLOTH_STUDIO_HOME="$PWD/.studio" ./install.sh --local
UNSLOTH_STUDIO_HOME="$PWD/.studio" unsloth studio -p 8888
```
Then to update :
```bash
cd unsloth && git pull
@ -176,7 +184,8 @@ cd unsloth && git pull
unsloth studio -p 8888
```
#### Developer installs: Windows PowerShell:
#### Developer / Nightly / Experimental installs: Windows PowerShell:
The developer install builds from the `main` branch, which is the latest (nightly) source.
```powershell
git clone https://github.com/unslothai/unsloth.git
cd unsloth
@ -184,40 +193,31 @@ Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass
.\install.ps1 --local
unsloth studio -p 8888
```
To install into an isolated location (its own virtual env, `auth/`, `studio.db`, cache and llama.cpp build), set `UNSLOTH_STUDIO_HOME` and pass it again at launch:
```powershell
$env:UNSLOTH_STUDIO_HOME="$PWD\.studio"; .\install.ps1 --local
$env:UNSLOTH_STUDIO_HOME="$PWD\.studio"; unsloth studio -p 8888
```
Then to update :
```bash
cd unsloth && git pull
./install.sh --local
unsloth studio -p 8888
```
#### Nightly: MacOS, Linux, WSL:
```bash
git clone https://github.com/unslothai/unsloth
cd unsloth
git checkout nightly
./install.sh --local
unsloth studio -p 8888
```
Then to launch every time:
```bash
unsloth studio -p 8888
```
#### Nightly: Windows:
Run in Windows Powershell:
```powershell
git clone https://github.com/unslothai/unsloth.git
cd unsloth
git checkout nightly
Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass
cd unsloth; git pull
.\install.ps1 --local
unsloth studio -p 8888
```
Then to launch every time:
#### Remote access: `--secure` (HTTPS tunnel) vs raw port
By default `unsloth studio` binds to `127.0.0.1` (this machine only). To reach it from another device, pick one of:
- `--secure` (recommended): serve **only** through a free Cloudflare HTTPS link. Studio stays bound to localhost and the tunnel provides the public URL; it fails closed (does not start) if the tunnel can't come up, so the raw port is never exposed.
```bash
unsloth studio -p 8888
unsloth studio --secure -p 8888
```
- `-H 0.0.0.0`: bind the raw port on all network interfaces, reachable from anywhere on the network. Only use this on a trusted network.
```bash
unsloth studio -H 0.0.0.0 -p 8888
```
Server-side tools (web search, Python and terminal code execution) run as your user and are on by default. Anyone who can reach the server with the API key can run code on this machine, so keep your API key private and pass `--disable-tools` when exposing Studio.
#### Advanced launch options
Installer options can be passed as environment variables. On macOS, Linux and WSL place the variable after the pipe so the shell passes it to `sh`; on Windows set it with `$env:` before piping to `iex`.

View file

@ -2462,6 +2462,9 @@ elif case "$TORCH_INDEX_URL" in */rocm*|*/gfx*) true ;; *) false ;; esac; then
substep "ROCm: $_rocm_root"
[ -n "$_gpu_rocm_ver" ] && substep "hipconfig: $_gpu_rocm_ver"
[ -n "$_gpu_disp_mkt" ] && [ -n "$_gpu_disp_gfx" ] && substep "GPU: $_gpu_disp_mkt"
elif [ "$OS" = "macos" ] && [ "$_ARCH" = "arm64" ]; then
# Apple Silicon: PyTorch gets Metal (MPS) acceleration over unified memory, so not CPU-only.
step "gpu" "Apple Silicon (Metal, unified memory)"
else
step "gpu" "none (CPU-only)" "$C_WARN"
fi

View file

@ -657,11 +657,32 @@ def detect_reasoning_flags(
return flags
# Gemma 4 ships MTP as a separate drafter (no "-mtp" in the name). Gemma 3n
# ships no drafter, so it is excluded -- it takes the normal non-MTP path.
_GEMMA_MTP_FAMILY_RE = re.compile(r"gemma[-_]?4[-_]", re.IGNORECASE)
def _is_gemma_mtp_family(name: Optional[str]) -> bool:
"""Match Gemma 4 by name."""
return bool(name) and bool(_GEMMA_MTP_FAMILY_RE.search(name))
def _is_gemma_mtp_name(model_identifier: Optional[str], gguf_path: Optional[str] = None) -> bool:
"""Match Gemma 4 by id or GGUF filename."""
return _is_gemma_mtp_family(model_identifier) or _is_gemma_mtp_family(
Path(gguf_path).name if gguf_path else None
)
def _is_mtp_model_name(model_identifier: Optional[str], gguf_path: Optional[str] = None) -> bool:
"""Name-based MTP detector. Fallback for the metadata signal."""
for cand in (model_identifier, Path(gguf_path).name if gguf_path else None):
if cand and "-mtp" in cand.lower():
return True
# Recognise Gemma 4 too, so a failed drafter download surfaces a
# fallback reason instead of silently defaulting.
if cand and _is_gemma_mtp_family(cand):
return True
return False
@ -851,6 +872,23 @@ def _auto_mode_drops_mtp(
return req_mode == "auto" and size_b is not None and size_b < _MTP_MIN_SIZE_B
def _mla_mtp_auto_enabled() -> bool:
"""Whether Auto may pick embedded MTP for an MLA model (GLM-5.2/DeepSeek/Kimi).
Off by default: llama.cpp's MLA/DSA MTP path keeps a duplicated full target-KV
context and recomputes the sparse-attention indexer every draft step, so it runs
~2x slower than no speculation (GLM-5.2 bench: 27 vs 45 tok/s, flat across draft
depth and 96-100% acceptance) -- the opposite of the vLLM/SGLang speedup on the
same model. Set UNSLOTH_MLA_MTP_ENABLED=1 to let Auto promote MLA MTP again once
that path is optimized upstream. Forced mtp / mtp+ngram ignore this gate."""
return os.environ.get("UNSLOTH_MLA_MTP_ENABLED", "0").strip().lower() in (
"1",
"true",
"yes",
"on",
)
def _extra_args_set_spec_type(extra_args: Optional[Iterable[str]]) -> bool:
"""User passed --spec-type / --spec-default? llama-server takes one
--spec-type (comma-separated to chain), so suppress auto-emit."""
@ -1357,6 +1395,11 @@ class LlamaCppBackend:
def gguf_path(self) -> Optional[str]:
return self._gguf_path
@property
def hf_repo(self) -> Optional[str]:
"""HF repo of the loaded model, or None for local/native file loads."""
return self._hf_repo
@property
def mtp_draft_path(self) -> Optional[str]:
return self._mtp_draft_path
@ -2823,12 +2866,16 @@ class LlamaCppBackend:
drafter_path: Optional[str] = None,
draft_weights_bytes: int = 0,
n_parallel: int = 1,
mtp_keeps_target_ctx: bool = True,
) -> Optional[int]:
"""MTP draft reserve at ``n_ctx`` = draft KV (grows with ctx) + separate-
drafter weights + (MLA only) a duplicated target KV context. The verify
buffer rides in the ctx-fit headroom (no tuned constant). None when the
draft KV can't be sized (caller keeps the flat fallback).
``draft_weights_bytes`` is the drafter file size (0 for embedded)."""
drafter weights + (MTP + MLA only) a duplicated target KV context. The
verify buffer rides in the ctx-fit headroom (no tuned constant). None when
the draft KV can't be sized (caller keeps the flat fallback).
``draft_weights_bytes`` is the drafter file size (0 for embedded).
``mtp_keeps_target_ctx`` is True for MTP draft modes (which keep the
duplicated target context) and False for separate-drafter spec modes
(draft-simple/draft-eagle3), which do not."""
draft_kv = self._mtp_draft_kv_bytes(
n_ctx,
drafter_path = drafter_path,
@ -2837,16 +2884,19 @@ class LlamaCppBackend:
n_parallel = n_parallel,
)
weights = max(0, draft_weights_bytes)
# MLA models (GLM-5.x, DeepSeek, Kimi-K2) keep a *second* full copy of the
# target model's KV context for MTP draft verification -- llama.cpp's
# MLA models (GLM-5.x, DeepSeek, Kimi-K2) under MTP keep a *second* full copy
# of the target model's KV context for draft verification -- llama.cpp's
# `ctx_tgt=yes` -- allocated at f16 regardless of the main cache type. It is
# ~the main KV again and dwarfs the embedded draft head (GLM-5.2 @ 1M ctx:
# a ~2 GiB head next to a ~89 GiB target copy), so omitting it lets auto-fit
# pick a context that fits on paper but OOMs cublasCreate at the first
# decode. Non-MLA MTP (Qwen/Gemma) keeps no such copy, so this is gated
# strictly on MLA (kv_lora_rank present) and leaves those models unchanged.
# decode. Gated on both MLA (kv_lora_rank present) and the engaged mode
# actually being MTP: non-MLA MTP (Qwen/Gemma) keeps no such copy, and the
# separate-drafter spec modes (draft-simple/draft-eagle3) load a small
# distinct drafter with its own KV -- already counted in draft_kv/weights --
# rather than duplicating the target, so they must not be charged for it.
target_ctx_copy = 0
if self._kv_lora_rank is not None:
if mtp_keeps_target_ctx and self._kv_lora_rank is not None:
target_ctx_copy = self._estimate_kv_cache_bytes(n_ctx, "f16", n_parallel = n_parallel)
if draft_kv is None:
# KV unsized (exotic/remote drafter): still reserve known weights + any
@ -3789,11 +3839,31 @@ class LlamaCppBackend:
return None
target: Optional[str] = None
try:
from huggingface_hub import list_repo_files
target = pick(list_repo_files(hf_repo, token = hf_token))
except Exception as e:
logger.debug(f"Could not list repo files for {label}: {e}")
from huggingface_hub import list_repo_files
# Retry a transient listing blip; permanent repo/auth errors and offline
# mode are not retried (offline raises at once -> fall through to cache).
for attempt in range(3):
if self._cancel_event.is_set():
return None
try:
target = pick(list_repo_files(hf_repo, token = hf_token))
break
except Exception as e:
if type(e).__name__ in (
"RepositoryNotFoundError",
"GatedRepoError",
"RevisionNotFoundError",
"EntryNotFoundError",
"OfflineModeIsEnabled",
):
logger.debug(f"Could not list repo files for {label}: {e}")
break
logger.debug(
f"Could not list repo files for {label} " f"(attempt {attempt + 1}/3): {e}"
)
if attempt < 2:
self._cancel_event.wait(2**attempt)
if target is None:
try:
@ -4807,6 +4877,12 @@ class LlamaCppBackend:
self._nextn_predict_layers
or _is_mtp_model_name(model_identifier, model_path)
or bool(mtp_draft_path)
) and not (
# Drafterless Gemma falls back to ngram-mod; reserve no
# drafter VRAM for it (mirrors the launch resolver).
_is_gemma_mtp_name(model_identifier, model_path)
and not mtp_draft_path
and not self._nextn_predict_layers
)
_mtp_binary_ok = True
_mtp_probe_raised = False
@ -4818,25 +4894,31 @@ class LlamaCppBackend:
except Exception:
_mtp_binary_ok = False
_mtp_probe_raised = True
_mtp_will_engage = bool(
_user_mtp_via_extras
or _user_draft_via_extras
or (
not _extra_args_set_spec_type(extra_args)
and _mtp_model_for_fit
and (
_mtp_effective in ("mtp", "mtp+ngram")
or (_mtp_effective == "auto" and not _mtp_sub_3b_for_fit)
)
and (
_mtp_binary_ok
# Reserve on a raised (uncached) probe too: it re-probes in
# _build_speculative_flags and may still engage MTP (embedded
# head or separate drafter -- _mtp_model_for_fit covers both).
or _mtp_probe_raised
)
_auto_studio_mtp = (
not _extra_args_set_spec_type(extra_args)
and _mtp_model_for_fit
and (
_mtp_effective in ("mtp", "mtp+ngram")
or (_mtp_effective == "auto" and not _mtp_sub_3b_for_fit)
)
and (
_mtp_binary_ok
# Reserve on a raised (uncached) probe too: it re-probes in
# _build_speculative_flags and may still engage MTP (embedded
# head or separate drafter -- _mtp_model_for_fit covers both).
or _mtp_probe_raised
)
)
_mtp_will_engage = bool(
_user_mtp_via_extras or _user_draft_via_extras or _auto_studio_mtp
)
# The duplicated full target-KV copy (ctx_tgt) is an MTP-only
# cost: the MTP head runs a second context over the target
# model's own KV geometry. The separate-drafter spec modes
# (draft-simple/draft-eagle3, reached via _user_draft_via_extras)
# load a small distinct drafter with its own KV and keep no such
# copy, so only charge it when the engaged mode is truly MTP.
_engaged_is_mtp = bool(_user_mtp_via_extras or _auto_studio_mtp)
# Effective draft depth: extras win (last-wins at launch), else
# the field, else the platform default (2 GPU / 3 CPU).
@ -4905,6 +4987,7 @@ class LlamaCppBackend:
drafter_path = _mtp_draft_for_budget,
draft_weights_bytes = _mtp_draft_weights,
n_parallel = n_parallel,
mtp_keeps_target_ctx = _engaged_is_mtp,
)
is not None
):
@ -4920,6 +5003,7 @@ class LlamaCppBackend:
_dp: Optional[str] = _mtp_draft_for_budget,
_w: int = _mtp_draft_weights,
_np: int = n_parallel,
_mtp: bool = _engaged_is_mtp,
) -> int:
v = self._estimate_mtp_overhead_bytes(
ctx,
@ -4929,6 +5013,7 @@ class LlamaCppBackend:
drafter_path = _dp,
draft_weights_bytes = _w,
n_parallel = _np,
mtp_keeps_target_ctx = _mtp,
)
return v if v is not None else 0
@ -6169,6 +6254,24 @@ class LlamaCppBackend:
_mtp_too_small = (
_mtp_size_b is not None and _mtp_size_b < _MTP_MIN_SIZE_B and not bool(mtp_draft_path)
)
# Drafterless Gemma (name-only MTP, no embedded head): emitting MTP
# would abort llama-server, so every mode below falls back instead.
_mtp_drafter_missing = (
_is_gemma_mtp_name(model_identifier, model_path)
and not mtp_draft_path
and not self._nextn_predict_layers
)
# Embedded MTP head on an MLA model (GLM-5.2/DeepSeek/Kimi, detected by
# kv_lora_rank): llama.cpp's MLA/DSA MTP path is ~2x slower than no spec,
# so Auto drops it (override via the Settings dropdown / forced mtp, or
# UNSLOTH_MLA_MTP_ENABLED=1). Separate drafters (Gemma, mtp_draft_path) and
# non-MLA embedded heads (Qwen, no kv_lora_rank) are unaffected.
_auto_mla_embedded_mtp = (
bool(self._nextn_predict_layers)
and self._kv_lora_rank is not None
and not bool(mtp_draft_path)
and not _mla_mtp_auto_enabled()
)
if user_owns_spec_type:
# User --spec-type wins outright; suppress auto-emit to avoid a
@ -6263,6 +6366,20 @@ class LlamaCppBackend:
logger.info("Spec decoding: ngram-mod")
return True
def _fallback_drafter_not_found() -> None:
"""Drafterless Gemma: use ngram-mod (or spec-default) and record why."""
logger.warning(
"Model %s is MTP-capable but no drafter or head was found; "
"falling back. Check network or run `unsloth studio update`.",
model_identifier,
)
if self.probe_server_capabilities(binary).get("supports_ngram_mod"):
_emit_ngram_mod()
else:
flags.append("--spec-default")
self._speculative_type = "default"
self._spec_fallback_reason = "drafter_not_found"
if effective_mode == "off":
return flags # nothing to emit
if effective_mode == "ngram-simple":
@ -6283,6 +6400,10 @@ class LlamaCppBackend:
flags.append("--spec-default")
self._speculative_type = "default"
return flags
if _mtp_drafter_missing:
# Drafterless: draft-mtp would abort llama-server, so fall back.
_fallback_drafter_not_found()
return flags
if _mtp_too_small:
logger.warning(
f"Forcing MTP on a {_mtp_size_b:.1f}B model; "
@ -6301,6 +6422,10 @@ class LlamaCppBackend:
)
_emit_ngram_mod()
return flags
if _mtp_drafter_missing:
# No head/drafter: keep ngram-mod, drop the draft-mtp chain.
_fallback_drafter_not_found()
return flags
if _mtp_too_small:
logger.warning(
f"Forcing MTP+Ngram on a {_mtp_size_b:.1f}B model; "
@ -6312,14 +6437,42 @@ class LlamaCppBackend:
# effective_mode == "auto": the promotion path. llama.cpp #22673:
# MTP is compatible with mmproj, so there's no vision gate.
if is_mtp_model and not _mtp_too_small:
# GPU: MTP-only. CPU/Mac: chain ngram-mod + MTP.
_emit_mtp(chain_ngram = not gpus)
if _auto_mla_embedded_mtp:
# MLA embedded-MTP (GLM-5.2 et al.): the MTP path regresses vs spec-off
# on llama.cpp today, so Auto drops it and falls back to ngram-mod (or
# spec-off if unsupported), mirroring the sub-3B branch. Forced mtp /
# mtp+ngram (handled above) still engage; UNSLOTH_MLA_MTP_ENABLED=1
# re-enables this promotion once upstream optimizes the path.
self._spec_fallback_reason = "mla_mtp_disabled"
_mla_caps = self.probe_server_capabilities(binary)
if _mla_caps.get("supports_ngram_mod"):
logger.info(
"Auto: MLA embedded-MTP model detected; llama.cpp's MLA/DSA "
"MTP path is slower than no speculation, so using ngram-mod "
"instead. Override via the Studio Speculative Decoding "
"dropdown or UNSLOTH_MLA_MTP_ENABLED=1."
)
_emit_ngram_mod()
else:
logger.info(
"Auto: MLA embedded-MTP model detected; disabling speculative "
"decoding (this llama-server does not advertise ngram-mod). "
"Override via the dropdown or UNSLOTH_MLA_MTP_ENABLED=1."
)
# spec-off: emit nothing, mirroring the sub-3B no-ngram path.
elif is_mtp_model and not _mtp_too_small:
if _mtp_drafter_missing:
# Name-only MTP, drafter did not resolve (download failed/absent).
_fallback_drafter_not_found()
else:
# GPU: MTP-only. CPU/Mac: chain ngram-mod + MTP.
_emit_mtp(chain_ngram = not gpus)
elif is_mtp_model and _mtp_too_small:
# Sub-3B fallback: drop the MTP draft head, keep ngram-mod when
# the binary supports it.
_small_caps = self.probe_server_capabilities(binary)
if _small_caps.get("supports_ngram_mod"):
if _mtp_drafter_missing:
_fallback_drafter_not_found()
elif self.probe_server_capabilities(binary).get("supports_ngram_mod"):
logger.info(
f"MTP GGUF detected but model size {_mtp_size_b:.1f}B "
"is below the 3B speedup threshold; using ngram-mod "
@ -6409,6 +6562,16 @@ class LlamaCppBackend:
if req_mode != backend_mode:
return False
# Prior HF load fell back with drafter_not_found; a same-settings reload
# must retry the download in load_model, not dedupe to the stale fallback
# (HF loads resolve the drafter there, so gguf_path is None here).
if (
self._spec_fallback_reason == "drafter_not_found"
and gguf_path is None
and req_mode in ("auto", "mtp", "mtp+ngram")
):
return False
# spec_draft_n_max only matters when an MTP variant is engaged. Compare
# on the resolved spec so an Auto request promoted to draft-mtp still
# bounces a reload when n_max changes.

View file

@ -29,6 +29,7 @@ from hub.schemas.inventory import (
DeleteCachedModelResponse,
GgufVariantsResponse,
LocalModelListResponse,
ModelsFolderResponse,
RecommendedFoldersResponse,
RemoveScanFolderResponse,
ScanFolderInfo,
@ -91,6 +92,11 @@ def browse_folders(
return folder_browser.browse_folders_response(path, show_hidden)
@router.get("/models-folder", response_model = ModelsFolderResponse)
def get_models_folder(current_subject: str = Depends(get_current_subject)):
return local_inventory.get_models_folder_response()
@router.get("/gguf-variants", response_model = GgufVariantsResponse)
async def get_gguf_variants(
repo_id: str = Query(

View file

@ -284,3 +284,13 @@ class BrowseFoldersResponse(BaseModel):
"they contain only files, no subdirectories)."
),
)
class ModelsFolderResponse(BaseModel):
"""The directory where downloaded models are stored (the active HF hub
cache, honoring ``HF_HOME`` / ``HF_HUB_CACHE``)."""
path: str = Field(
...,
description = "Path to the model download directory.",
)

View file

@ -670,6 +670,31 @@ async def list_local_models_response(models_dir: str = "./models") -> LocalModel
)
def get_models_folder_response() -> dict:
"""Return the directory where downloaded models are stored.
This is the active HF hub cache (honors ``HF_HOME`` / ``HF_HUB_CACHE``);
the desktop app reveals it in the OS file manager.
"""
path = _resolve_hf_cache_dir()
# Create it if missing so "Open folder" works before the first download:
# HF builds the cache lazily, and studio only pre-creates the *default*
# dir, not a user's explicit HF_HOME / HF_HUB_CACHE.
try:
path.mkdir(parents = True, exist_ok = True)
except OSError as e:
raise HTTPException(
status_code = 500,
detail = f"Failed to create models folder: {path}: {e}",
) from e
if not path.is_dir():
raise HTTPException(
status_code = 500,
detail = f"Models folder path is not a directory: {path}",
)
return {"path": str(path)}
def get_scan_folders_response() -> dict:
return {"folders": list_scan_folders()}

View file

@ -181,6 +181,55 @@ def test_browse_folders_hides_sensitive_dirs(monkeypatch, tmp_path):
assert ".ssh" not in names
def test_get_models_folder_response_creates_and_returns_dir(monkeypatch, tmp_path):
# The endpoint creates the cache dir on demand so the desktop "Open folder"
# action works even before the first download.
target = tmp_path / "hub"
monkeypatch.setattr(local_inventory, "_resolve_hf_cache_dir", lambda: target)
response = local_inventory.get_models_folder_response()
assert response == {"path": str(target)}
assert target.is_dir()
def test_get_models_folder_response_reports_create_failure(monkeypatch, tmp_path):
target = tmp_path / "hub"
target.write_text("not a directory")
monkeypatch.setattr(local_inventory, "_resolve_hf_cache_dir", lambda: target)
with pytest.raises(HTTPException) as exc_info:
local_inventory.get_models_folder_response()
assert exc_info.value.status_code == 500
assert "Failed to create models folder" in exc_info.value.detail
def test_get_models_folder_response_requires_directory(monkeypatch, tmp_path):
class MissingPath:
def __init__(self, value: Path):
self.value = value
def mkdir(self, *, parents: bool, exist_ok: bool):
assert parents is True
assert exist_ok is True
def is_dir(self):
return False
def __str__(self):
return str(self.value)
target = MissingPath(tmp_path / "hub")
monkeypatch.setattr(local_inventory, "_resolve_hf_cache_dir", lambda: target)
with pytest.raises(HTTPException) as exc_info:
local_inventory.get_models_folder_response()
assert exc_info.value.status_code == 500
assert "not a directory" in exc_info.value.detail
def test_contained_link_path_confines_to_link_dir(tmp_path):
link_dir = tmp_path / "ollama" / ".studio_links" / "abc123"

View file

@ -412,8 +412,13 @@ class InferenceStatusResponse(BaseModel):
"(auto on an MTP model, or forced mtp / mtp+ngram). "
"'binary_no_mtp' / 'binary_outdated' -> a newer prebuilt would "
"re-enable it (show the update affordance); 'runtime_error' -> the "
"current build could not run it. None when MTP engaged or was not "
"requested."
"current build could not run it; 'drafter_not_found' -> the model's "
"separate MTP drafter could not be resolved; 'mla_mtp_disabled' -> "
"an Auto-mode policy downgrade: the model is MLA (GLM-5.2 et al.) "
"whose llama.cpp MTP path runs slower than no speculation, so Auto "
"used ngram-mod or spec-off instead -- updating won't help; choose "
"MTP in Settings (or set UNSLOTH_MLA_MTP_ENABLED=1) to force it. "
"None when MTP engaged or was not requested."
),
)
llama_cpp_prebuilt_stale: bool = Field(

View file

@ -1829,6 +1829,16 @@ def _request_matches_loaded_settings(
backend_mode = llama_backend.requested_spec_mode or "auto"
if req_mode != backend_mode:
return False
# Prior HF load fell back with drafter_not_found: a same-settings reload must
# retry the download, not dedupe to the stale fallback. HF only (hf_repo set);
# local/native loads have no download to retry (handled by the path compare).
if (
llama_backend.hf_repo
and llama_backend.spec_fallback_reason == "drafter_not_found"
and req_mode in ("auto", "mtp", "mtp+ngram")
and not _extra_args_set_spec_type(effective_extra)
):
return False
# spec_draft_n_max only matters with an MTP variant; None means "platform
# default" and matches whatever the backend chose.
if backend_mode in ("mtp", "mtp+ngram") and request.spec_draft_n_max is not None:
@ -2206,7 +2216,9 @@ async def load_model(
and llama_backend.model_identifier.lower() == model_identifier.lower()
# Match runtime settings so Apply isn't dropped (#5401).
and _request_matches_loaded_settings(
request, llama_backend, effective_chat_template_override
request,
llama_backend,
effective_chat_template_override,
)
# Skip if a prior audio probe failed -- let load_model retry.
and getattr(llama_backend, "_audio_probed", True)
@ -3127,9 +3139,13 @@ async def generate_stream(
log = logger,
)
cancel_event = threading.Event()
async def stream():
gen = None
completed = False
try:
for chunk in backend.generate_chat_response(
gen = backend.generate_chat_response(
messages = request.messages,
system_prompt = request.system_prompt,
image = image,
@ -3138,14 +3154,35 @@ async def generate_stream(
top_k = request.top_k,
max_new_tokens = request.max_new_tokens,
repetition_penalty = request.repetition_penalty,
):
cancel_event = cancel_event,
)
_DONE = object()
while True:
chunk = await asyncio.to_thread(next, gen, _DONE)
if chunk is _DONE:
break
yield f"data: {json.dumps({'content': chunk})}\n\n"
completed = True
yield "data: [DONE]\n\n"
except asyncio.CancelledError:
cancel_event.set()
backend.reset_generation_state()
raise
except Exception as e:
cancel_event.set()
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"
finally:
if not completed and not cancel_event.is_set():
cancel_event.set()
backend.reset_generation_state()
if gen is not None:
try:
await asyncio.to_thread(gen.close)
except (RuntimeError, ValueError):
pass
return StreamingResponse(
stream(),

View file

@ -62,6 +62,7 @@ from core.inference.llama_cpp import (
_extra_args_set_any_flag,
_extra_args_set_spec_type,
_is_mtp_model_name,
_mla_mtp_auto_enabled,
)
@ -1329,6 +1330,282 @@ def test_forced_mtp_ngram_on_non_mtp_model_keeps_ngram(monkeypatch):
assert backend.requested_spec_mode == "mtp+ngram"
# ── Auto drops embedded MTP for MLA models (GLM-5.2 et al.) ───────────
#
# llama.cpp's MLA/DSA MTP path runs ~2x slower than no speculation (GLM-5.2
# bench), so Auto downgrades it to ngram-mod (or spec-off). The clean
# metadata separator from non-MLA MTP (Qwen, kept on draft-mtp) is
# self._kv_lora_rank. Forced mtp / mtp+ngram and separate drafters (Gemma)
# stay on draft-mtp; UNSLOTH_MLA_MTP_ENABLED=1 re-enables Auto promotion.
# GLM-5.2's repo name has no "MTP" marker, so its MTP signal is metadata-only
# (nextn_predict_layers) -- exactly the embedded-MLA case we gate.
_GLM_MLA_MODEL = "unsloth/GLM-5.2-GGUF"
def _mla_resolver_backend(
monkeypatch,
*,
ngram_supported = True,
kv_lora_rank = 512,
nextn = 1,
):
"""Resolver backend posing as an embedded-MTP MLA model (kv_lora_rank set)."""
backend = _resolver_backend(monkeypatch, ngram_supported = ngram_supported)
backend._nextn_predict_layers = nextn
backend._kv_lora_rank = kv_lora_rank
return backend
@pytest.mark.parametrize("gpus", [True, False])
def test_auto_mla_embedded_mtp_falls_back_to_ngram(monkeypatch, gpus):
# Auto + MLA embedded MTP + ngram supported -> ngram-mod on BOTH platforms
# (the CPU chain ngram-mod,draft-mtp is dropped: no draft-mtp for MLA).
backend = _mla_resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = _GLM_MLA_MODEL,
model_path = None,
gpus = gpus,
binary = "/fake/llama-server",
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "ngram-mod"
assert "--spec-draft-n-max" not in parsed
assert "--spec-ngram-mod-n-match" in parsed
assert backend.speculative_type == "ngram-mod"
assert backend.requested_spec_mode == "auto"
assert backend.spec_fallback_reason == "mla_mtp_disabled"
assert backend.spec_draft_n_max is None
def test_auto_mla_embedded_mtp_no_ngram_disables_spec(monkeypatch):
# Auto + MLA embedded MTP + no ngram-mod support -> emit nothing (spec-off),
# mirroring the sub-3B no-ngram path. Still flagged as a policy downgrade.
backend = _mla_resolver_backend(monkeypatch, ngram_supported = False)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = _GLM_MLA_MODEL,
model_path = None,
gpus = True,
binary = "/fake/llama-server",
)
assert "--spec-type" not in flags
assert backend.speculative_type is None
assert backend.requested_spec_mode == "auto"
assert backend.spec_fallback_reason == "mla_mtp_disabled"
def test_auto_non_mla_embedded_mtp_keeps_draft_mtp(monkeypatch):
# Auto + embedded MTP + NON-MLA (kv_lora_rank None, e.g. Qwen) -> unchanged:
# still draft-mtp at the platform default. No policy downgrade.
backend = _mla_resolver_backend(monkeypatch, kv_lora_rank = None)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = _MTP_MODEL,
model_path = None,
gpus = True,
binary = "/fake/llama-server",
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "draft-mtp"
assert parsed.get("--spec-draft-n-max") == "2"
assert backend.speculative_type == "draft-mtp"
assert backend.spec_fallback_reason is None
def test_auto_mla_separate_drafter_keeps_mtp(monkeypatch):
# Auto + MLA + a separate drafter (mtp_draft_path) -> the drafter exemption
# wins over the MLA gate: still draft-mtp (Gemma-style external drafter is
# not the slow embedded MLA/DSA path).
backend = _mla_resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = _GLM_MLA_MODEL,
model_path = None,
gpus = True,
binary = "/fake/llama-server",
mtp_draft_path = "/fake/mtp-draft.gguf",
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "draft-mtp"
assert backend.speculative_type == "draft-mtp"
assert backend.spec_fallback_reason is None
def test_auto_non_mtp_mla_model_unaffected(monkeypatch):
# Auto + MLA but NO embedded MTP head (kv_lora_rank set, nextn None, e.g.
# GLM-4.7-Flash) -> non-MTP default; no accidental ngram drop.
backend = _mla_resolver_backend(monkeypatch, nextn = None)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = "unsloth/GLM-4.7-Flash-GGUF",
model_path = None,
gpus = True,
binary = "/fake/llama-server",
)
assert "--spec-default" in flags
assert "ngram-mod" not in flags
assert backend.speculative_type == "default"
assert backend.spec_fallback_reason is None
@pytest.mark.parametrize(
"mode, expect_spec_type, expect_n_max",
[
("mtp", "draft-mtp", "2"),
("mtp+ngram", "ngram-mod,draft-mtp", "2"),
],
)
def test_forced_mtp_on_mla_still_engages(monkeypatch, mode, expect_spec_type, expect_n_max):
# Explicit override engages the deliberately-slower MTP route on MLA models,
# regardless of the Auto gate. No policy downgrade reason.
backend = _mla_resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = mode,
spec_draft_n_max = None,
extra_args = None,
model_identifier = _GLM_MLA_MODEL,
model_path = None,
gpus = True,
binary = "/fake/llama-server",
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == expect_spec_type
assert parsed.get("--spec-draft-n-max") == expect_n_max
assert backend.speculative_type == "draft-mtp"
assert backend.requested_spec_mode == mode
assert backend.spec_fallback_reason is None
def test_env_flag_reenables_auto_mla_mtp(monkeypatch):
# UNSLOTH_MLA_MTP_ENABLED=1 -> Auto promotes MLA embedded MTP to draft-mtp
# again (the forward hook for when llama.cpp optimizes the path).
monkeypatch.setenv("UNSLOTH_MLA_MTP_ENABLED", "1")
backend = _mla_resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = _GLM_MLA_MODEL,
model_path = None,
gpus = True,
binary = "/fake/llama-server",
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "draft-mtp"
assert backend.speculative_type == "draft-mtp"
assert backend.spec_fallback_reason is None
@pytest.mark.parametrize("value", ["1", "true", "yes", "on", "TRUE", "On"])
def test_mla_mtp_auto_enabled_truthy_values(monkeypatch, value):
monkeypatch.setenv("UNSLOTH_MLA_MTP_ENABLED", value)
assert _mla_mtp_auto_enabled() is True
@pytest.mark.parametrize("value", ["0", "false", "no", "off", "", " ", "bogus"])
def test_mla_mtp_auto_disabled_default_and_falsy(monkeypatch, value):
monkeypatch.setenv("UNSLOTH_MLA_MTP_ENABLED", value)
assert _mla_mtp_auto_enabled() is False
def test_mla_mtp_auto_disabled_when_unset(monkeypatch):
monkeypatch.delenv("UNSLOTH_MLA_MTP_ENABLED", raising = False)
assert _mla_mtp_auto_enabled() is False
def test_read_gguf_metadata_captures_kv_lora_rank(tmp_path):
# GLM-5.2-style header: MLA (kv_lora_rank) + embedded MTP (nextn) populate
# both fields, so the Auto gate sees an MLA embedded-MTP model.
gguf = _write_minimal_gguf(
tmp_path / "model.gguf",
arch = "glm-dsa",
nextn = 1,
extra_uint32 = {
"glm-dsa.block_count": 4,
"glm-dsa.attention.kv_lora_rank": 512,
},
)
backend = LlamaCppBackend()
backend._read_gguf_metadata(str(gguf))
assert backend._nextn_predict_layers == 1
assert backend._kv_lora_rank == 512
def test_read_gguf_metadata_qwen_mtp_has_no_kv_lora_rank(tmp_path):
# Qwen MTP header: embedded MTP but non-MLA, so kv_lora_rank stays None and
# Auto keeps it on draft-mtp.
gguf = _write_minimal_gguf(
tmp_path / "model.gguf",
arch = "qwen35moe",
nextn = 1,
extra_uint32 = {"qwen35moe.block_count": 4},
)
backend = LlamaCppBackend()
backend._read_gguf_metadata(str(gguf))
assert backend._nextn_predict_layers == 1
assert backend._kv_lora_rank is None
def test_reload_skip_auto_mla_ngram_is_idempotent():
# A GLM model resolved to ngram-mod under Auto must not churn: a duplicate
# Auto /load at the same settings is already-satisfied.
backend = _mtp_backend(
_model_identifier = _GLM_MLA_MODEL,
_speculative_type = "ngram-mod",
_requested_spec_mode = "auto",
)
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = _GLM_MLA_MODEL,
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = "auto",
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is True
)
def test_reload_forced_mtp_bounces_auto_mla():
# Overriding Auto (ngram-mod) with a forced mtp request must reload (to the
# slower draft-mtp route), not dedup against the running ngram-mod server.
backend = _mtp_backend(
_model_identifier = _GLM_MLA_MODEL,
_speculative_type = "ngram-mod",
_requested_spec_mode = "auto",
)
assert (
backend._already_in_target_state(
gguf_path = None,
model_identifier = _GLM_MLA_MODEL,
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = "mtp",
chat_template_override = None,
extra_args = None,
is_vision = False,
)
is False
)
# ── Full named-repo resolver matrix (the shipping Studio families) ─────
#
# Locks auto / off / forced-mtp routing for every Qwen3.5 (MTP + plain) and
@ -1376,6 +1653,8 @@ _REAL_REPO_MATRIX = [
def _resolve_real(monkeypatch, repo, drafter, mode):
backend = _resolver_backend(monkeypatch)
if "qwen" in repo.lower() and "-mtp" in repo.lower():
backend._nextn_predict_layers = 1
flags = backend._build_speculative_flags(
speculative_type = mode,
spec_draft_n_max = None,
@ -1566,3 +1845,126 @@ def test_spec_fallback_reason_reset_on_off(monkeypatch):
binary = "/fake/llama-server",
)
assert backend.spec_fallback_reason is None
def test_is_gemma_mtp_family():
from core.inference.llama_cpp import _is_gemma_mtp_family
assert _is_gemma_mtp_family("unsloth/gemma-4-E4B-it-GGUF") is True
assert _is_gemma_mtp_family("unsloth/gemma-4-12b-it-GGUF") is True
# gemma-3n ships no separate drafter, so it is not a drafter family.
assert _is_gemma_mtp_family("unsloth/gemma-3n-E2B-it-GGUF") is False
assert _is_gemma_mtp_family("unsloth/Qwen3.5-35B-A3B-MTP-GGUF") is False
assert _is_gemma_mtp_family("unsloth/llama-3-8b") is False
def test_gemma_3n_without_drafter_is_not_mtp(monkeypatch):
# gemma-3n ships no drafter; it must take the normal non-MTP path, not
# drafter_not_found (which would make every reload retry a missing drafter).
backend = _resolver_backend(monkeypatch)
backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = "unsloth/gemma-3n-E4B-it-GGUF",
model_path = None,
gpus = True,
binary = "/fake/llama-server",
mtp_draft_path = None,
)
assert backend.spec_fallback_reason is None
def test_spec_fallback_reason_drafter_not_found(monkeypatch):
# Drafterless Gemma should fall back to ngram-mod + drafter_not_found.
backend = _resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = "unsloth/gemma-4-E4B-it-GGUF",
model_path = None,
gpus = True,
binary = "/fake/llama-server",
mtp_draft_path = None, # Drafter download failed
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "ngram-mod"
assert backend.speculative_type == "ngram-mod"
assert backend.spec_fallback_reason == "drafter_not_found"
def test_is_gemma_mtp_name_none_safe():
# model_identifier=None (local load) must not raise; recognise via filename.
from core.inference.llama_cpp import _is_gemma_mtp_family, _is_gemma_mtp_name
assert _is_gemma_mtp_family(None) is False
assert _is_gemma_mtp_name(None, "/models/gemma-4-E4B-it-Q4_K_M.gguf") is True
assert _is_gemma_mtp_name("unsloth/Qwen3.5-4B-MTP-GGUF", None) is False
@pytest.mark.parametrize("mode", ["mtp", "mtp+ngram"])
def test_forced_mtp_gemma_without_drafter_falls_back(monkeypatch, mode):
# Forced MTP on a drafterless Gemma must fall back, not emit draft-mtp.
backend = _resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = mode,
spec_draft_n_max = None,
extra_args = None,
model_identifier = "unsloth/gemma-4-E4B-it-GGUF",
model_path = None,
gpus = True,
binary = "/fake/llama-server",
mtp_draft_path = None,
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "ngram-mod"
assert "--model-draft" not in parsed
assert backend.spec_fallback_reason == "drafter_not_found"
def test_local_gemma_gguf_without_identifier_falls_back(monkeypatch):
# Local Gemma GGUF (family only in filename) must not crash; falls back.
backend = _resolver_backend(monkeypatch)
flags = backend._build_speculative_flags(
speculative_type = "auto",
spec_draft_n_max = None,
extra_args = None,
model_identifier = None,
model_path = "/models/gemma-4-E4B-it-Q4_K_M.gguf",
gpus = True,
binary = "/fake/llama-server",
mtp_draft_path = None,
)
parsed = _flags_dict(flags)
assert parsed.get("--spec-type") == "ngram-mod"
assert backend.spec_fallback_reason == "drafter_not_found"
def _drafter_not_found_kwargs():
return dict(
model_identifier = "unsloth/gemma-4-E4B-it-GGUF",
hf_variant = "Q4_K_M",
n_ctx = 8192,
cache_type_kv = None,
speculative_type = "auto",
chat_template_override = None,
extra_args = None,
is_vision = False,
gguf_path = None, # HF load: drafter resolves inside load_model
)
def test_already_in_target_state_retries_after_hf_drafter_not_found():
# Recoverable drafter_not_found must not dedupe; reload re-attempts download.
backend = _mtp_backend(
_model_identifier = "unsloth/gemma-4-E4B-it-GGUF",
_speculative_type = "ngram-mod",
_spec_fallback_reason = "drafter_not_found",
_mtp_draft_path = None,
_gguf_path = None,
)
assert backend._already_in_target_state(**_drafter_not_found_kwargs()) is False
# Sanity: with no fallback reason the same request still dedupes (matches).
ok = _mtp_backend(_model_identifier = "unsloth/gemma-4-E4B-it-GGUF", _gguf_path = None)
assert ok._already_in_target_state(**_drafter_not_found_kwargs()) is True

View file

@ -172,6 +172,22 @@ class TestMlaTargetCtxReserve:
ctx = 131072
assert mla._estimate_mtp_overhead_bytes(ctx) > non._estimate_mtp_overhead_bytes(ctx)
def test_separate_drafter_mode_drops_target_copy(self):
# The duplicated target context is MTP-only. draft-simple / draft-eagle3
# load a small separate drafter with its own KV (counted in the draft KV)
# and keep no target copy, so even on an MLA model the reserve must drop
# the f16 copy when mtp_keeps_target_ctx=False -- which is what the loader
# threads for those modes. The default (True) keeps the MTP copy.
b = _make_mla_backend()
ctx = 262144
mtp = b._estimate_mtp_overhead_bytes(ctx) # default True == MTP draft
separate = b._estimate_mtp_overhead_bytes(ctx, mtp_keeps_target_ctx = False)
# Separate-drafter overhead is exactly the draft KV (no target copy)...
assert separate == b._mtp_draft_kv_bytes(ctx)
# ...and the MTP reserve is that plus the full f16 target copy.
assert mtp == separate + b._estimate_kv_cache_bytes(ctx, "f16")
assert mtp > separate
class TestMlaFitPreventsOom:
"""The corrected reserve must actually lower the auto-fit context so the
@ -200,7 +216,7 @@ class TestMlaFitPreventsOom:
self.MODEL_BYTES,
mtp_engaged = True,
total_mib = self.TOTAL_MIB,
mtp_overhead_fn = lambda c: (b._mtp_draft_kv_bytes(c) or 0),
mtp_overhead_fn = lambda c: b._mtp_draft_kv_bytes(c) or 0,
)
assert draft_only == self.REQ_CTX # reproduces the over-advertised context
assert with_copy < self.REQ_CTX # corrected reserve backs the context off

View file

@ -829,6 +829,20 @@ class TestExtraArgsMtpDetection:
"request.tensor_parallel,llama_backend.tensor_parallel)" in body
)
def test_route_matcher_retries_after_drafter_not_found(self):
# drafter_not_found must not report "already loaded" or the reload never
# retries the download (#6459). Read source: importing routes pulls deps.
routes_src = (
Path(__file__).resolve().parent.parent / "routes" / "inference.py"
).read_text()
start = routes_src.index("def _request_matches_loaded_settings")
end = routes_src.index("\ndef ", start + 1)
body = "".join(routes_src[start:end].split())
assert 'llama_backend.spec_fallback_reason=="drafter_not_found"' in body
assert "not_extra_args_set_spec_type(effective_extra)" in body
# HF-only (hf_repo): local/native loads have no download to retry.
assert "llama_backend.hf_repo" in body
def test_extra_args_main_cache_type_heavier_axis(self):
# Asymmetric --cache-type-k/-v must budget the heavier axis (extras win
# per axis at launch), not the last-wins single type that under-reserves.

View file

@ -107,6 +107,7 @@ def test_detect_safetensors_features_none_template_returns_all_false():
"supports_reasoning": False,
"reasoning_style": "enable_thinking",
"reasoning_always_on": False,
"reasoning_effort_levels": [],
"supports_preserve_thinking": False,
"supports_tools": False,
}

View file

@ -365,6 +365,11 @@ def test_runtime_recovery_reloads_without_mtp(monkeypatch):
while b._spec_fallback_reason != "runtime_error" and time.monotonic() < deadline:
time.sleep(0.02)
assert b._spec_fallback_reason == "runtime_error"
# The reload thread clears the single-flight flag in its finally, a beat after
# it sets the fallback reason -- wait for that instead of racing the thread.
deadline = time.monotonic() + 2
while b._mtp_runtime_fallback_in_progress and time.monotonic() < deadline:
time.sleep(0.02)
assert b._mtp_runtime_fallback_in_progress is False

View file

@ -67,6 +67,7 @@ import {
listPromptEntries,
type PromptEntry,
} from "@/features/chat/api/prompts-api";
import { useChatPreferencesStore } from "@/features/chat/stores/chat-preferences-store";
import { useChatProjects } from "@/features/chat/hooks/use-chat-projects";
import { NewProjectDialog } from "@/features/chat/components/new-project-dialog";
import { parseExternalModelId } from "@/features/chat/external-providers";
@ -135,6 +136,7 @@ import {
ChevronLeftIcon,
ChevronRightIcon,
Columns2Icon,
CornerDownRightIcon,
GitBranchIcon,
GlobeIcon,
HeadphonesIcon,
@ -178,12 +180,33 @@ type PromptQueueUIEntry = {
total: number;
};
type PromptQueueUIItemStatus = "queued" | "next" | "waiting" | "running";
type PromptQueueUIItem = {
id: string;
prompt: string;
position: number;
total: number;
status: PromptQueueUIItemStatus;
threadIds: string[];
canEdit: boolean;
canRemove: boolean;
};
interface PromptQueueUIState {
byThreadId: Record<string, PromptQueueUIEntry>;
current: number;
total: number;
items: PromptQueueUIItem[];
isRunning: boolean;
}
const usePromptQueueUI = create<PromptQueueUIState>(() => ({
byThreadId: {},
current: 0,
total: 0,
items: [],
isRunning: false,
}));
type PromptQueueTarget = {
@ -195,8 +218,10 @@ type PromptQueueTarget = {
};
type PromptQueueItem = {
id: string;
prompt: string;
target: PromptQueueTarget;
dispatched: boolean;
};
const PROMPT_QUEUE_INDEXING_RETRY_MS = 500;
@ -214,6 +239,10 @@ function compactIds(ids: Array<string | null | undefined>) {
return Array.from(new Set(ids.filter((id): id is string => Boolean(id))));
}
function createPromptQueueItemId() {
return `prompt-queue-${crypto.randomUUID()}`;
}
function stopPromptQueueSubscription({
resetRunningState = true,
}: {
@ -228,7 +257,7 @@ function stopPromptQueueSubscription({
}
}
function resetPromptQueue(showToast = false) {
function resetPromptQueue() {
promptQueueGeneration += 1;
promptQueueIsRunning = false;
promptQueueItems = [];
@ -240,16 +269,10 @@ function resetPromptQueue(showToast = false) {
}
stopPromptQueueSubscription();
syncPromptQueueUI();
if (showToast) {
toast.success("Prompt queue complete");
}
}
function queueToastDescription(prompt: string) {
return prompt.length > 80 ? `${prompt.slice(0, 80)}...` : prompt;
}
function appendQueuedPrompt(item: PromptQueueItem) {
item.dispatched = true;
syncPromptQueueUI();
item.target.append(item.prompt);
}
@ -335,8 +358,10 @@ async function dispatchQueuedPrompt(
function createQueuedPrompt(prompt: string, target: PromptQueueTarget) {
return {
id: createPromptQueueItemId(),
prompt,
target,
dispatched: false,
};
}
@ -368,13 +393,62 @@ function findPromptQueueEntry(
return null;
}
function canEditPromptQueueItem(item: PromptQueueItem) {
return !item.dispatched;
}
function canRemovePromptQueueItem(item: PromptQueueItem) {
return !item.dispatched;
}
function promptQueueItemMatchesThreadIds(
item: PromptQueueUIItem,
threadIds: string[],
) {
return item.threadIds.some((threadId) => threadIds.includes(threadId));
}
function syncPromptQueueUI() {
if (!promptQueueIsRunning || promptQueueItems.length === 0) {
usePromptQueueUI.setState({ byThreadId: {} });
usePromptQueueUI.setState({
byThreadId: {},
current: 0,
total: 0,
items: [],
isRunning: false,
});
return;
}
const activeItemIndex = Math.max(promptQueueIndex, 0);
const total = promptQueueItems.length;
const current = promptQueueIndex >= 0 ? Math.min(activeItemIndex + 1, total) : 0;
const items = promptQueueItems
.map((item, index): PromptQueueUIItem | null => {
if (index < activeItemIndex || item.dispatched) {
return null;
}
const threadIds = getPromptQueueTargetIds(item.target);
const isActive = promptQueueIndex >= 0 && index === activeItemIndex;
const status: PromptQueueUIItemStatus = item.dispatched
? "running"
: isActive
? promptQueueWaitingForTargetIdle
? "waiting"
: "next"
: "queued";
return {
id: item.id,
prompt: item.prompt,
position: index + 1,
total,
status,
threadIds,
canEdit: canEditPromptQueueItem(item),
canRemove: canRemovePromptQueueItem(item),
};
})
.filter((item): item is PromptQueueUIItem => Boolean(item));
const groups: Array<{
ids: Set<string>;
current: number;
@ -423,7 +497,80 @@ function syncPromptQueueUI() {
});
}
usePromptQueueUI.setState({ byThreadId });
usePromptQueueUI.setState({
byThreadId,
current,
total,
items,
isRunning: true,
});
}
function editPromptQueueItem(itemId: string, prompt: string) {
const nextPrompt = prompt.trim();
if (!nextPrompt) {
return false;
}
const itemIndex = promptQueueItems.findIndex(
(candidate) => candidate.id === itemId,
);
if (itemIndex < 0) {
return false;
}
const item = promptQueueItems[itemIndex];
if (!canEditPromptQueueItem(item)) {
return false;
}
item.prompt = nextPrompt;
syncPromptQueueUI();
return true;
}
function clearPromptQueueRetryTimer() {
if (!promptQueueRetryTimer) {
return;
}
clearTimeout(promptQueueRetryTimer);
promptQueueRetryTimer = null;
}
function removePromptQueueItem(itemId: string) {
const itemIndex = promptQueueItems.findIndex((item) => item.id === itemId);
if (itemIndex < 0) {
return false;
}
const item = promptQueueItems[itemIndex];
if (!canRemovePromptQueueItem(item)) {
return false;
}
const wasActive =
promptQueueIndex >= 0 && itemIndex === Math.max(promptQueueIndex, 0);
promptQueueItems.splice(itemIndex, 1);
if (promptQueueItems.length === 0) {
resetPromptQueue();
return true;
}
if (itemIndex < promptQueueIndex) {
promptQueueIndex -= 1;
}
if (wasActive && promptQueueIndex >= promptQueueItems.length) {
resetPromptQueue();
return true;
}
syncPromptQueueUI();
if (wasActive) {
clearPromptQueueRetryTimer();
promptQueueWaitingForTargetIdle = false;
promptQueuePrevStoreRunning = false;
const next = promptQueueItems[promptQueueIndex];
if (next) {
scheduleQueuedPromptDispatch(next, 50);
}
}
return true;
}
function isPromptQueueTargetRunning(
@ -456,15 +603,12 @@ function isActivePromptQueueTargetRunning(
function advancePromptQueue() {
const nextIndex = promptQueueIndex + 1;
if (nextIndex >= promptQueueItems.length) {
resetPromptQueue(true);
resetPromptQueue();
return;
}
promptQueueIndex = nextIndex;
syncPromptQueueUI();
const next = promptQueueItems[nextIndex];
toast(`Prompt ${nextIndex + 1} / ${promptQueueItems.length}`, {
description: queueToastDescription(next.prompt),
});
promptQueueWaitingForTargetIdle = false;
promptQueuePrevStoreRunning = false;
scheduleQueuedPromptDispatch(next, 100);
@ -529,9 +673,6 @@ function startPromptQueue(
...filtered.map((prompt) => createQueuedPrompt(prompt, target)),
);
syncPromptQueueUI();
toast.success("Added to prompt queue", {
description: `${filtered.length} prompt${filtered.length === 1 ? "" : "s"} queued.`,
});
return;
}
@ -547,12 +688,6 @@ function startPromptQueue(
promptQueueIsRunning = true;
promptQueuePrevStoreRunning = shouldWaitForCurrentRun;
syncPromptQueueUI();
toast(
shouldWaitForCurrentRun ? "Prompt queued" : `Prompt 1 / ${filtered.length}`,
{
description: queueToastDescription(filtered[0]),
},
);
startPromptQueueSubscription();
if (!shouldWaitForCurrentRun) {
const first = promptQueueItems[0];
@ -563,10 +698,15 @@ function startPromptQueue(
}
function stopPromptQueueRun() {
const activeTarget = promptQueueItems[Math.max(promptQueueIndex, 0)]?.target;
const activeItem = promptQueueItems[Math.max(promptQueueIndex, 0)];
const activeTarget = activeItem?.target;
const shouldCancelActiveRun = Boolean(activeItem?.dispatched);
resetPromptQueue();
if (!shouldCancelActiveRun) {
return;
}
try {
activeTarget?.cancel();
activeTarget.cancel();
} catch {
// The active run may have already ended.
}
@ -1018,6 +1158,29 @@ const ThreadComposerDock: FC<{
onHeightChange?: (height: number | null) => void;
}> = ({ disabled, threadId, onHeightChange }) => {
const { overlay } = useGeneratedImageOverlay();
const activeThreadId = useChatRuntimeStore((s) => s.activeThreadId);
const threadListItemId = useAuiState(
({ threadListItem }) => threadListItem.id,
);
const threadListItemRemoteId = useAuiState(
({ threadListItem }) => threadListItem.remoteId,
);
const promptQueueThreadIds = compactIds([
threadListItemId,
threadListItemRemoteId,
threadId,
activeThreadId,
]);
const queueVisible = usePromptQueueUI(
(s) =>
Boolean(findPromptQueueEntry(s, promptQueueThreadIds)) &&
s.items.some((item) =>
promptQueueItemMatchesThreadIds(item, promptQueueThreadIds),
),
);
const showModelDisclaimer = useChatPreferencesStore(
(s) => s.showModelDisclaimer,
);
// Report dock height so the viewport reserves matching scroll space when
// attachments or multiline input grow the composer.
@ -1046,7 +1209,12 @@ const ThreadComposerDock: FC<{
{/* Fade the top edge so scrolling text is not cut off by a hard line. */}
<div
aria-hidden={true}
className="absolute inset-x-0 bottom-0 top-[10px] bg-gradient-to-t from-background from-[calc(100%_-_28px)] to-transparent"
className={cn(
"absolute inset-x-0 bottom-0 bg-gradient-to-t from-background from-[calc(100%_-_28px)] to-transparent",
queueVisible
? "h-32 backdrop-blur-[1px] [mask-image:linear-gradient(to_top,black_0%,black_58%,transparent_100%)]"
: "top-[10px]",
)}
/>
<div className="relative px-5 pb-2">
<div className="pointer-events-auto mx-auto w-full max-w-(--thread-max-width)">
@ -1056,9 +1224,11 @@ const ThreadComposerDock: FC<{
menuSide="top"
/>
</div>
<p className="composer-footer-note">
LLMs can make mistakes. Double-check responses.
</p>
{showModelDisclaimer && (
<p className="composer-footer-note">
LLMs can make mistakes. Double-check responses.
</p>
)}
</div>
</div>
);
@ -1743,14 +1913,15 @@ const Composer: FC<{
aria-disabled={disabled}
onSubmit={handleSubmit}
>
<PromptQueueStack queueThreadIds={promptQueueThreadIds} />
{isTauri ? (
// Phase 1 native model owns Tauri local-path drops. Restore browser
// attachment drops in Tauri once Phase 1d adds token bridging.
<div className="aui-composer-attachment-dropzone unsloth-composer-surface">
<div className="aui-composer-attachment-dropzone unsloth-composer-surface relative z-10">
{composerContent}
</div>
) : (
<ComposerPrimitive.AttachmentDropzone className="group/dropzone aui-composer-attachment-dropzone unsloth-composer-surface relative">
<ComposerPrimitive.AttachmentDropzone className="group/dropzone aui-composer-attachment-dropzone unsloth-composer-surface relative z-10">
{composerContent}
{/* Gemini-style drop affordance, shown while a file is dragged over
the composer. Absolute + pointer-events-none so the outline adds
@ -2991,6 +3162,184 @@ const ComposerToolsMenu: FC<{ side?: "top" | "bottom" }> = ({
);
};
function promptQueueStatusLabel(status: PromptQueueUIItemStatus) {
switch (status) {
case "running":
return "Running now";
case "waiting":
return "Waiting";
case "next":
return "Next";
case "queued":
return "Queued";
default: {
const exhaustiveStatus: never = status;
throw new Error(`Unhandled prompt queue status: ${exhaustiveStatus}`);
}
}
}
const PromptQueueStack: FC<{ queueThreadIds: string[] }> = ({
queueThreadIds,
}) => {
const queueEntry = usePromptQueueUI((s) =>
findPromptQueueEntry(s, queueThreadIds),
);
const items = usePromptQueueUI((s) => s.items);
const [editingItemId, setEditingItemId] = useState<string | null>(null);
const [draftPrompt, setDraftPrompt] = useState("");
const editInputRef = useRef<HTMLTextAreaElement>(null);
const visibleItems = items.filter((item) =>
promptQueueItemMatchesThreadIds(item, queueThreadIds),
);
const editingItem = visibleItems.find((item) => item.id === editingItemId);
const editingItemCanEdit = editingItem?.canEdit ?? false;
const activeEditingItemId = editingItem ? editingItemId : null;
useEffect(() => {
if (!activeEditingItemId) {
return;
}
editInputRef.current?.focus();
editInputRef.current?.select();
}, [activeEditingItemId]);
useEffect(() => {
if (!editingItemId || editingItemCanEdit) {
return;
}
setEditingItemId(null);
setDraftPrompt("");
}, [editingItemCanEdit, editingItemId]);
if (!queueEntry || visibleItems.length === 0) {
return null;
}
const { current, total } = queueEntry;
const startEditing = (item: PromptQueueUIItem) => {
if (!item.canEdit) {
return;
}
setEditingItemId(item.id);
setDraftPrompt(item.prompt);
};
const saveEditing = () => {
if (!activeEditingItemId) {
return;
}
if (editPromptQueueItem(activeEditingItemId, draftPrompt)) {
setEditingItemId(null);
setDraftPrompt("");
}
};
const cancelEditing = () => {
setEditingItemId(null);
setDraftPrompt("");
};
return (
<div
className="relative z-0 mx-7 mb-[-8px] max-h-[28vh] overflow-y-auto rounded-t-[18px] rounded-b-none border border-border/45 bg-background/90 px-5 py-2 text-muted-foreground shadow-none backdrop-blur-md dark:bg-card/85"
aria-label={`Prompt queue, ${current} of ${total}`}
>
<div className="divide-y divide-border/25">
{visibleItems.map((item, visibleIndex) => {
const isEditing = item.id === activeEditingItemId;
const visiblePosition = visibleIndex + 1;
return (
<div
key={item.id}
className={cn("min-h-10", isEditing ? "h-auto" : "h-10")}
aria-label={`${promptQueueStatusLabel(item.status)} prompt ${visiblePosition} of ${visibleItems.length}: ${item.prompt}`}
>
{isEditing ? (
<div className="grid min-h-10 grid-cols-[minmax(0,1fr)_auto_auto] items-center gap-2.5 py-1">
<textarea
ref={editInputRef}
value={draftPrompt}
rows={1}
onChange={(event) =>
setDraftPrompt(event.currentTarget.value)
}
onKeyDown={(event) => {
if (
event.key === "Enter" &&
(event.metaKey || event.ctrlKey)
) {
event.preventDefault();
saveEditing();
} else if (event.key === "Escape") {
event.preventDefault();
cancelEditing();
}
}}
className="max-h-20 min-h-8 min-w-0 resize-none rounded-md border border-border/45 bg-transparent px-2 py-1.5 text-sm leading-5 text-foreground outline-none transition-colors focus-visible:border-ring focus-visible:ring-2 focus-visible:ring-ring/35"
aria-label={`Edit queued prompt ${visiblePosition}`}
/>
<Button
type="button"
variant="ghost"
size="sm"
className="h-7 px-2 text-xs text-muted-foreground"
onClick={cancelEditing}
>
Cancel
</Button>
<Button
type="button"
size="sm"
className="h-7 px-2 text-xs"
disabled={draftPrompt.trim().length === 0}
onClick={saveEditing}
>
Save
</Button>
</div>
) : (
<div className="grid h-10 grid-cols-[minmax(0,1fr)_auto_2rem] items-center gap-2.5">
<div className="flex min-w-0 items-center gap-2.5">
<CornerDownRightIcon className="size-4 shrink-0 text-muted-foreground/50" />
<div className="truncate text-sm text-muted-foreground">
{item.prompt}
</div>
</div>
{item.canEdit ? (
<Button
type="button"
variant="ghost"
size="sm"
className="h-7 w-[5.25rem] justify-center gap-1 px-0 text-sm font-normal text-muted-foreground/80 hover:text-foreground"
onClick={() => startEditing(item)}
>
<HugeiconsIcon icon={Edit03Icon} strokeWidth={2} />
Edit
</Button>
) : null}
<TooltipIconButton
tooltip="Remove from queue"
side="bottom"
type="button"
variant="ghost"
size="icon"
className="col-start-3 size-7 justify-self-center text-muted-foreground/70 hover:text-destructive"
aria-label={`Remove queued prompt ${visiblePosition}`}
disabled={!item.canRemove}
onClick={() => removePromptQueueItem(item.id)}
>
<HugeiconsIcon icon={Delete02Icon} strokeWidth={2} />
</TooltipIconButton>
</div>
)}
</div>
);
})}
</div>
</div>
);
};
const ComposerRightControls: FC<{
disabled?: boolean;
queueDisabled?: boolean;
@ -3014,8 +3363,6 @@ const ComposerRightControls: FC<{
findPromptQueueEntry(s, queueThreadIds),
);
const isQueueRunning = Boolean(queueEntry);
const queueCurrent = queueEntry?.current ?? 0;
const queueTotal = queueEntry?.total ?? 0;
return (
<div className="aui-composer-action-wrapper flex shrink-0 items-center gap-1.5">
<ReasoningToggle side={menuSide} />
@ -3043,14 +3390,6 @@ const ComposerRightControls: FC<{
</TooltipIconButton>
</ComposerPrimitive.StopDictation>
</ComposerPrimitive.If>
{isQueueRunning ? (
<span
className="ml-1 flex h-7 items-center rounded-full bg-primary/10 px-2 text-[11px] font-semibold text-primary"
aria-live="polite"
>
<span className="tabular-nums">Queue {queueCurrent}/{queueTotal}</span>
</span>
) : null}
<AuiIf condition={({ thread }) => !thread.isRunning && !isQueueRunning}>
<ComposerPrimitive.Send asChild={true}>
<TooltipIconButton

View file

@ -104,6 +104,7 @@ import {
loadOptionalBool,
useChatRuntimeStore,
} from "./stores/chat-runtime-store";
import { useChatPreferencesStore } from "./stores/chat-preferences-store";
import { useExternalProvidersStore } from "./stores/external-providers-store";
import { buildChatTourSteps } from "./tour";
import { ArtifactSurface } from "./artifacts/artifact-surface";
@ -492,6 +493,9 @@ function CompareShell({
children: ReactElement;
composer: ReactElement;
}): ReactElement {
const showModelDisclaimer = useChatPreferencesStore(
(s) => s.showModelDisclaimer,
);
return (
<CompareHandlesProvider handlesRef={handlesRef}>
<div className="flex min-h-0 min-w-0 flex-1 basis-0 flex-col">
@ -503,9 +507,11 @@ function CompareShell({
</div>
<div className="shrink-0 bg-background pl-5 pr-5 md:pr-[30px] pb-2 pt-1">
<div className="mx-auto w-full max-w-[48rem]">{composer}</div>
<p className="composer-footer-note">
LLMs can make mistakes. Double-check responses.
</p>
{showModelDisclaimer && (
<p className="composer-footer-note">
LLMs can make mistakes. Double-check responses.
</p>
)}
</div>
</div>
</CompareHandlesProvider>

View file

@ -1005,12 +1005,16 @@ export function ChatSettingsPanel({
speculativeType === "mtp+ngram") && (
<div className="rounded-lg bg-amber-500/[0.08] px-3 py-2 text-[12px] leading-[1.4] text-nav-fg/80">
<p>
{specFallbackReason === "runtime_error"
{specFallbackReason === "mla_mtp_disabled"
? "MTP is disabled by default for this model architecture because it currently runs slower than standard decoding. Select MTP above to force it."
: specFallbackReason === "runtime_error"
? "MTP could not start for this model on the installed llama.cpp build, so it is running without speculative decoding."
: "MTP is not available in the installed llama.cpp build, so this model is running without it." +
(llamaUpdateStatus?.update_available
? " Update llama.cpp to enable it."
: "")}
: specFallbackReason === "drafter_not_found"
? "This model supports MTP, but its drafter file could not be downloaded, so MTP is off and it falls back to n-gram speculative decoding where the llama.cpp build supports it. Check your network connection or Hugging Face access, then reload the model to retry the drafter."
: "MTP is not available in the installed llama.cpp build, so this model is running without it." +
(llamaUpdateStatus?.update_available
? " Update llama.cpp to enable it."
: "")}
</p>
{mtpUpdatable && llamaUpdateStatus?.update_available && (
<Button

View file

@ -6,9 +6,12 @@ import { persist } from "zustand/middleware";
// Client-side chat UI prefs kept in localStorage, not the chat DB.
// confirmDeleteChats: when off, deleting a chat skips the confirm dialog.
// showModelDisclaimer: when off, hide the "LLMs can make mistakes" footer note.
export interface ChatPreferencesState {
confirmDeleteChats: boolean;
setConfirmDeleteChats: (value: boolean) => void;
showModelDisclaimer: boolean;
setShowModelDisclaimer: (value: boolean) => void;
}
export const useChatPreferencesStore = create<ChatPreferencesState>()(
@ -17,6 +20,9 @@ export const useChatPreferencesStore = create<ChatPreferencesState>()(
confirmDeleteChats: true,
setConfirmDeleteChats: (confirmDeleteChats) =>
set({ confirmDeleteChats }),
showModelDisclaimer: true,
setShowModelDisclaimer: (showModelDisclaimer) =>
set({ showModelDisclaimer }),
}),
{
name: "unsloth_chat_preferences",
@ -25,6 +31,7 @@ export const useChatPreferencesStore = create<ChatPreferencesState>()(
return {
...current,
confirmDeleteChats: saved?.confirmDeleteChats ?? true,
showModelDisclaimer: saved?.showModelDisclaimer ?? true,
};
},
},

View file

@ -196,7 +196,10 @@ export interface InferenceStatusResponse {
/**
* Why MTP was disabled on the loaded model despite being requested.
* "binary_no_mtp" / "binary_outdated" -> updating llama.cpp would re-enable
* it; "runtime_error" -> the current build could not run it. Null otherwise.
* it; "runtime_error" -> the current build could not run it;
* "mla_mtp_disabled" -> an Auto-mode policy downgrade for MLA models
* (GLM-5.2 et al.) whose llama.cpp MTP path is slower than no speculation
* (updating won't help; choose MTP in Settings to force it). Null otherwise.
*/
spec_fallback_reason?: string | null;
}

View file

@ -30,6 +30,7 @@ import { useEffect, useRef, useState } from "react";
import { useShallow } from "zustand/react/shallow";
import { EXPORT_METHODS, type ExportMethod } from "../constants";
import type { ExportLogEntry } from "../api/export-api";
import { getExportLogLineClass } from "../lib/log-style";
import {
selectExportProgressPercent,
useExportRuntimeStore,
@ -520,22 +521,16 @@ export function ExportRunPanel(props: ExportRunPanelProps) {
</span>
</div>
) : (
<pre className="whitespace-pre-wrap break-words">
<div className="whitespace-pre-wrap break-words">
{run.logLines.map((entry, idx) => (
<div
key={idx}
className={
entry.stream === "stderr"
? "text-rose-300/90"
: entry.stream === "status"
? "text-sky-300/90"
: ""
}
className={getExportLogLineClass(entry)}
>
{formatLogLine(entry)}
</div>
))}
</pre>
</div>
)}
</div>
</div>

View file

@ -0,0 +1,40 @@
// 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 type { ExportLogEntry } from "../api/export-api";
type ExportLogTone = "stdout" | "stderr" | "status" | "warning";
const WARNING_LINE_PATTERNS = [
/Skipping import of cpp extensions due to incompatible torch version/i,
/Please see GitHub issue #2919 for more info/i,
/torch_dtype is deprecated!\s*Use dtype instead!/i,
] as const;
function isWarningLine(line: string): boolean {
return WARNING_LINE_PATTERNS.some((pattern) => pattern.test(line));
}
export function getExportLogTone(entry: ExportLogEntry): ExportLogTone {
if (entry.stream === "status") {
return "status";
}
if (isWarningLine(entry.line)) {
return "warning";
}
return entry.stream === "stderr" ? "stderr" : "stdout";
}
export function getExportLogLineClass(entry: ExportLogEntry): string {
const tone = getExportLogTone(entry);
if (tone === "stderr") {
return "text-rose-300/90";
}
if (tone === "status") {
return "text-sky-300/90";
}
if (tone === "warning") {
return "text-status-warning";
}
return "";
}

View file

@ -5,6 +5,8 @@ import { cn } from "@/lib/utils";
import { ArrowLeft01Icon, ArrowRight01Icon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import {
type MouseEvent as ReactMouseEvent,
type PointerEvent as ReactPointerEvent,
type ReactNode,
useCallback,
useEffect,
@ -102,13 +104,79 @@ export function CardCarousel<T>({
[stepPx],
);
// Click-and-drag panning (mouse only; touch/pen keep native scrolling).
const drag = useRef<{ id: number; x: number; left: number; moved: boolean } | null>(
null,
);
const suppressClick = useRef(false);
const onPointerDown = useCallback((e: ReactPointerEvent<HTMLDivElement>) => {
suppressClick.current = false;
const el = scrollerRef.current;
if (!el || e.pointerType !== "mouse" || e.button !== 0) return;
drag.current = { id: e.pointerId, x: e.clientX, left: el.scrollLeft, moved: false };
}, []);
const onPointerMove = useCallback((e: ReactPointerEvent<HTMLDivElement>) => {
const d = drag.current;
const el = scrollerRef.current;
if (!d || !el || e.pointerId !== d.id) return;
// Primary button no longer held: the press ended off the scroller, so no
// pointerup reached us. Drop the stale drag instead of scrolling on hover.
if ((e.buttons & 1) === 0) {
if (d.moved) el.style.scrollSnapType = "";
drag.current = null;
return;
}
const dx = e.clientX - d.x;
// Ignore tiny moves so plain clicks still register.
if (!d.moved && Math.abs(dx) < 5) return;
if (!d.moved) {
d.moved = true;
// Snap fights the per-frame scrollLeft writes; disable it while dragging.
el.style.scrollSnapType = "none";
el.setPointerCapture(d.id);
}
el.scrollLeft = d.left - dx;
}, []);
const endDrag = useCallback((e: ReactPointerEvent<HTMLDivElement>) => {
const d = drag.current;
if (!d || e.pointerId !== d.id) return;
if (d.moved) {
// A drag just happened: swallow the click it would fire on a card.
suppressClick.current = true;
const el = scrollerRef.current;
// Restore snap so the row settles on a card after the drag.
if (el) el.style.scrollSnapType = "";
el?.releasePointerCapture?.(d.id);
}
drag.current = null;
}, []);
const onClickCapture = useCallback((e: ReactMouseEvent<HTMLDivElement>) => {
if (!suppressClick.current) return;
suppressClick.current = false;
e.preventDefault();
e.stopPropagation();
}, []);
return (
<div className="relative">
<div
ref={scrollerRef}
onScroll={updateArrows}
onPointerDown={onPointerDown}
onPointerMove={onPointerMove}
onPointerUp={endDrag}
onPointerCancel={endDrag}
onClickCapture={onClickCapture}
// Stop the avatar image from starting a native drag during a pan.
onDragStart={(e) => e.preventDefault()}
aria-label={ariaLabel}
className="hub-carousel flex snap-x gap-4 overflow-x-auto pb-4 pt-2"
// px-2 + -mx-2 give card shadows room so the edge cards aren't clipped;
// scroll-px-2 keeps snap-start aligned with the heading.
className="hub-carousel -mx-2 flex cursor-grab snap-x scroll-px-2 gap-4 overflow-x-auto px-2 pb-4 pt-2 select-none active:cursor-grabbing"
>
{items.map((item) => (
<div

View file

@ -23,11 +23,14 @@ import {
listScanFolders,
removeScanFolder,
} from "@/features/hub/inventory";
import { openModelsDir } from "@/features/native-intents/api";
import { isTauri } from "@/lib/api-base";
import { cn } from "@/lib/utils";
import {
Delete02Icon,
FileSearchIcon,
FolderAddIcon,
FolderExportIcon,
FolderOpenIcon,
FolderSearchIcon,
PlusSignIcon,
@ -138,6 +141,16 @@ export function OnDeviceFoldersDialog({
[handleInventoryChanged, pending],
);
// Scan folders are arbitrary paths that may be moved or deleted after they
// were registered, so surface the command's failure as a toast.
const handleOpen = useCallback(async (folder: ScanFolderInfo) => {
try {
await openModelsDir(folder.path);
} catch (err) {
toast.error("Couldn't open location", { description: formatError(err) });
}
}, []);
const handleRemove = useCallback(
async (folder: ScanFolderInfo) => {
const key = `remove:${folder.id}` as const;
@ -333,6 +346,27 @@ export function OnDeviceFoldersDialog({
</TooltipContent>
</Tooltip>
</div>
{isTauri ? (
<Tooltip>
<TooltipTrigger asChild={true}>
<button
type="button"
aria-label={`Open ${folder.path}`}
onClick={() => void handleOpen(folder)}
className="inline-flex size-8 shrink-0 items-center justify-center rounded-full text-muted-foreground transition-colors hover:bg-muted hover:text-foreground"
>
<HugeiconsIcon
icon={FolderExportIcon}
strokeWidth={1.75}
className="size-4"
/>
</button>
</TooltipTrigger>
<TooltipContent side="left" className="tooltip-compact">
Open in file manager
</TooltipContent>
</Tooltip>
) : null}
<Tooltip>
<TooltipTrigger asChild={true}>
<button

View file

@ -147,7 +147,7 @@
pointer-events: none;
position: absolute;
z-index: 5;
width: 56px;
width: 44px;
opacity: 0;
transition: opacity 240ms ease;
}
@ -157,7 +157,9 @@
}
.hub-page .hub-carousel-fade-left {
left: 0;
/* -8px matches the scroller's -mx-2 bleed so the opaque edge sits on the
clip edge and no card peeks out beside the fade. */
left: -8px;
background: linear-gradient(
to right,
var(--background),
@ -167,7 +169,8 @@
}
.hub-page .hub-carousel-fade-right {
right: 0;
/* Mirror of fade-left: offset by the -mx-2 bleed to reach the clip edge. */
right: -8px;
background: linear-gradient(
to left,
var(--background),

View file

@ -44,3 +44,9 @@ export async function revealPathToken(token: string): Promise<void> {
export async function openPathToken(token: string): Promise<void> {
return invokeNative<void>("open_path_token", { token });
}
// Open a backend-resolved directory (e.g. the models/HF cache folder) in the
// OS file manager. The Tauri command validates the path is a real directory.
export async function openModelsDir(path: string): Promise<void> {
return invokeNative<void>("open_models_dir", { path });
}

View file

@ -1,7 +1,6 @@
// 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 type { ReactNode } from "react";
import {
Select,
SelectContent,
@ -10,22 +9,20 @@ import {
SelectValue,
} from "@/components/ui/select";
import { Slider } from "@/components/ui/slider";
import {
ToggleGroup,
ToggleGroupItem,
} from "@/components/ui/toggle-group";
import { ToggleGroup, ToggleGroupItem } from "@/components/ui/toggle-group";
import {
Tooltip,
TooltipContent,
TooltipTrigger,
} from "@/components/ui/tooltip";
import { InfoIcon } from "lucide-react";
import { cn } from "@/lib/utils";
import {
type RagAutoInject,
type RagMode,
useChatRuntimeStore,
} from "@/features/chat/stores/chat-runtime-store";
import { cn } from "@/lib/utils";
import { InfoIcon } from "lucide-react";
import type { ReactNode } from "react";
const MODE_LABEL: Record<RagMode, string> = {
hybrid: "Hybrid",
@ -92,6 +89,7 @@ function SliderRow({
disabled={disabled}
onValueChange={([v]) => onChange(v)}
aria-label={label}
className="panel-slider"
/>
</div>
);
@ -152,6 +150,7 @@ export function RetrievalSettingsSection() {
step={1}
onValueChange={([value]) => setRagTopK(value)}
aria-label="Number of passages to retrieve"
className="panel-slider"
/>
</div>
@ -175,7 +174,9 @@ export function RetrievalSettingsSection() {
value={ragAutoInject}
onValueChange={(value) => {
// Radix clears on re-click; ignore empty so one stays selected.
if (value) setRagAutoInject(value as RagAutoInject);
if (value) {
setRagAutoInject(value as RagAutoInject);
}
}}
className="w-full"
aria-label="Auto-retrieve documents"

View file

@ -0,0 +1,40 @@
// 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 { authFetch } from "@/features/auth";
import { readFastApiError } from "@/lib/format-fastapi-error";
export type ModelsFolder = {
path: string;
};
// The path is resolved once at backend startup and never changes, so cache it
// and dedupe concurrent loads (same shape as the sibling settings loaders).
let cachedModelsFolder: ModelsFolder | null = null;
let inFlightModelsFolder: Promise<ModelsFolder> | null = null;
async function fetchModelsFolder(): Promise<ModelsFolder> {
const res = await authFetch("/api/hub/models-folder");
if (!res.ok) {
throw new Error(
await readFastApiError(res, "Failed to load models folder"),
);
}
const data = (await res.json()) as { path: string };
return { path: data.path };
}
export async function loadModelsFolder(): Promise<ModelsFolder> {
if (cachedModelsFolder) {
return cachedModelsFolder;
}
inFlightModelsFolder ??= fetchModelsFolder()
.then((folder) => {
cachedModelsFolder = folder;
return folder;
})
.finally(() => {
inFlightModelsFolder = null;
});
return inFlightModelsFolder;
}

View file

@ -187,6 +187,12 @@ export function ChatTab() {
const setConfirmDeleteChats = useChatPreferencesStore(
(state) => state.setConfirmDeleteChats,
);
const showModelDisclaimer = useChatPreferencesStore(
(state) => state.showModelDisclaimer,
);
const setShowModelDisclaimer = useChatPreferencesStore(
(state) => state.setShowModelDisclaimer,
);
useEffect(() => {
void countAllChats().then(setCount);
@ -310,6 +316,15 @@ export function ChatTab() {
/>
</SettingsRow>
))}
<SettingsRow
label={t("settings.chat.modelDisclaimer")}
description={t("settings.chat.modelDisclaimerDescription")}
>
<Switch
checked={showModelDisclaimer}
onCheckedChange={setShowModelDisclaimer}
/>
</SettingsRow>
</SettingsSection>
<SettingsSection title={t("settings.chat.artifacts.title")}>

View file

@ -13,6 +13,11 @@ import {
import { Input } from "@/components/ui/input";
import { Switch } from "@/components/ui/switch";
import { usePlatformStore } from "@/config/env";
import { isTauri } from "@/lib/api-base";
import { openModelsDir } from "@/features/native-intents/api";
import { copyToClipboard } from "@/lib/copy-to-clipboard";
import { toast } from "@/lib/toast";
import { loadModelsFolder, type ModelsFolder } from "../api/models-folder";
import { resetOnboardingDone } from "@/features/auth";
import { useChatRuntimeStore } from "@/features/chat";
import {
@ -134,6 +139,7 @@ export function GeneralTab() {
null,
);
const [isSavingHelperPrecache, setIsSavingHelperPrecache] = useState(false);
const [modelsFolder, setModelsFolder] = useState<ModelsFolder | null>(null);
const draftRef = useRef(draftToken);
useEffect(() => {
@ -199,6 +205,43 @@ export function GeneralTab() {
};
}, [t]);
useEffect(() => {
let cancelled = false;
void loadModelsFolder()
.then((folder) => {
if (cancelled) return;
setModelsFolder(folder);
})
.catch(() => {
// Non-critical: leave the row hidden if the path can't be resolved.
});
return () => {
cancelled = true;
};
}, []);
// Desktop opens the folder in the OS file manager; the browser can't, so it
// falls back to copying the path (which is the info users actually want).
const handleModelsFolder = async () => {
const folder = modelsFolder;
if (!folder) return;
if (isTauri) {
try {
await openModelsDir(folder.path);
} catch (error) {
toast.error(t("settings.general.storage.openError"), {
description: error instanceof Error ? error.message : undefined,
});
}
return;
}
if (await copyToClipboard(folder.path)) {
toast.success(t("settings.general.storage.copied"));
} else {
toast.error(t("settings.general.storage.copyError"));
}
};
const saveHelperPrecache = async (enabled: boolean) => {
setIsSavingHelperPrecache(true);
setHelperPrecacheError(null);
@ -287,6 +330,33 @@ export function GeneralTab() {
</SettingsRow>
</SettingsSection>
{modelsFolder ? (
<SettingsSection title={t("settings.general.storage.sectionTitle")}>
<SettingsRow
label={t("settings.general.storage.modelsFolder")}
description={t("settings.general.storage.modelsFolderDescription")}
>
<div className="flex items-center gap-2">
<span
title={modelsFolder.path}
className="max-w-[280px] truncate font-mono text-xs text-muted-foreground"
>
{modelsFolder.path}
</span>
<Button
variant="outline"
size="sm"
onClick={() => void handleModelsFolder()}
>
{isTauri
? t("settings.general.storage.openAction")
: t("settings.general.storage.copyAction")}
</Button>
</div>
</SettingsRow>
</SettingsSection>
) : null}
<SettingsSection title={t("settings.general.chatDefaults")}>
<SettingsRow
label={t("settings.general.autoTitleNewChats")}

View file

@ -136,6 +136,17 @@ export const en = {
maxUploadSizeDescription:
"Default is {defaultSize} MB.",
},
storage: {
sectionTitle: "Storage",
modelsFolder: "Models folder",
modelsFolderDescription:
"Where downloaded models are stored.",
openAction: "Open",
copyAction: "Copy path",
copied: "Path copied",
openError: "Couldn't open the folder",
copyError: "Couldn't copy the path",
},
resetPreferences: {
sectionTitle: "Danger zone",
label: "Reset all local preferences",
@ -197,6 +208,9 @@ export const en = {
chat: {
title: "Chat",
description: "Manage chat history stored on this device.",
modelDisclaimer: "Show model disclaimer",
modelDisclaimerDescription:
'Show "LLMs can make mistakes" under the chat box.',
artifacts: {
title: "Canvas",
collapseHtmlBlocks: "Collapse HTML blocks",

View file

@ -20,6 +20,7 @@ import subprocess
import sys
import sysconfig
import tempfile
import textwrap
import urllib.request
from pathlib import Path
@ -1401,6 +1402,7 @@ VERBOSE: bool = os.environ.get("UNSLOTH_VERBOSE", "0") == "1"
# Update _TOTAL if you add/remove steps in install_python_stack().
_STEP: int = 0
_TOTAL: int = 0 # set at runtime in install_python_stack() based on platform
_PROGRESS_LINE_ACTIVE: bool = False
# -- Paths --------------------------------------------------------------
SCRIPT_DIR = Path(__file__).resolve().parent
@ -1486,6 +1488,7 @@ _HAS_COLOR = _stdout_supports_color()
# 2-space indent, 15-char label (dim), then value.
_LABEL = "deps"
_COL = 15
_INDENT = 2
def _green(msg: str) -> str:
@ -1517,15 +1520,38 @@ def _step(
color_fn = None,
) -> None:
"""Print a single step line in the column format."""
global _PROGRESS_LINE_ACTIVE
if color_fn is None:
color_fn = _green
padded = label[:_COL]
_safe_print(f" {_dim(padded)}{' ' * (_COL - len(padded))}{color_fn(value)}")
plain_prefix_width = _INDENT + _COL
prefix = f"{' ' * _INDENT}{_dim(padded)}{' ' * (_COL - len(padded))}"
wrap_width = max(
24,
shutil.get_terminal_size((100, 20)).columns - plain_prefix_width,
)
lines = textwrap.wrap(
value,
width = wrap_width,
break_long_words = False,
break_on_hyphens = False,
) or [""]
if _PROGRESS_LINE_ACTIVE and not VERBOSE:
try:
sys.stdout.write("\n")
sys.stdout.flush()
except OSError:
pass
_PROGRESS_LINE_ACTIVE = False
_safe_print(f"{prefix}{color_fn(lines[0])}")
continuation_prefix = " " * plain_prefix_width
for line in lines[1:]:
_safe_print(f"{continuation_prefix}{color_fn(line)}")
def _progress(label: str) -> None:
"""Print an in-place progress bar aligned to the step column layout."""
global _STEP
global _STEP, _PROGRESS_LINE_ACTIVE
_STEP += 1
if VERBOSE:
return
@ -1537,6 +1563,7 @@ def _progress(label: str) -> None:
try:
sys.stdout.write(f"\r {_dim(_LABEL)}{pad}[{bar}] {_STEP:2}/{_TOTAL} {label:<20}{end}")
sys.stdout.flush()
_PROGRESS_LINE_ACTIVE = end == ""
except OSError:
pass

View file

@ -968,6 +968,9 @@ elif [ "$_setup_amd_detected" = true ]; then
substep "ROCm: $_setup_rocm_root"
[ -n "$_setup_rocm_ver" ] && substep "hipconfig: $_setup_rocm_ver"
[ -n "$_setup_mkt" ] && [ -n "$_setup_gfx" ] && substep "GPU: $_setup_mkt"
elif [ "$(uname -s 2>/dev/null)" = "Darwin" ] && [ "$(uname -m 2>/dev/null)" = "arm64" ]; then
# Apple Silicon: llama.cpp builds with Metal over unified memory, so not a CPU-only host.
step "gpu" "Apple Silicon (Metal, unified memory)"
else
step "gpu" "none (chat-only / GGUF)" "$C_WARN"
substep "Training and GPU inference require an NVIDIA or AMD ROCm GPU."

View file

@ -321,17 +321,29 @@ pub fn get_server_logs(state: tauri::State<'_, BackendState>) -> Vec<String> {
}
}
/// Open the Unsloth Studio directory in the system file manager.
#[tauri::command]
pub fn open_logs_dir() -> Result<(), String> {
let home = dirs::home_dir().ok_or("Could not determine home directory")?;
let dir = home.join(".unsloth").join("studio");
if !dir.exists() {
/// Open an existing directory in the system file manager. Validates the path
/// up front so callers get a clean error instead of a raw OS failure.
fn open_existing_dir(dir: &std::path::Path) -> Result<(), String> {
if !dir.is_dir() {
return Err(format!("Directory does not exist: {}", dir.display()));
}
open::that(dir).map_err(|e| format!("Failed to open directory: {}", e))
}
open::that(&dir).map_err(|e| format!("Failed to open directory: {}", e))
/// Open the Unsloth Studio directory in the system file manager.
#[tauri::command]
pub fn open_logs_dir(window: tauri::WebviewWindow) -> Result<(), String> {
crate::native_intents::ensure_main_window(&window)?;
let home = dirs::home_dir().ok_or("Could not determine home directory")?;
open_existing_dir(&home.join(".unsloth").join("studio"))
}
/// Open a models directory (resolved by the backend, e.g. the HF cache) in the
/// system file manager.
#[tauri::command]
pub fn open_models_dir(window: tauri::WebviewWindow, path: String) -> Result<(), String> {
crate::native_intents::ensure_main_window(&window)?;
open_existing_dir(std::path::Path::new(&path))
}
/// Start the first-launch installation process.

View file

@ -204,6 +204,7 @@ fn main() {
commands::check_health,
commands::get_server_logs,
commands::open_logs_dir,
commands::open_models_dir,
commands::start_backend_update,
commands::start_managed_repair,
commands::cancel_pending_elevation,

View file

@ -258,7 +258,7 @@ fn prune_expired(inner: &mut NativeIntakeInner) {
.retain(|intent| intent.path.expires_at_ms > now);
}
fn ensure_main_window(window: &WebviewWindow) -> Result<(), String> {
pub(crate) fn ensure_main_window(window: &WebviewWindow) -> Result<(), String> {
if window.label() == "main" {
Ok(())
} else {

View file

@ -10,6 +10,7 @@ from __future__ import annotations
import ast
import asyncio
import json
import threading
import time
from pathlib import Path
@ -453,6 +454,147 @@ def test_audio_input_stream_offloads_blocking_next_to_thread():
)
def test_generate_stream_offloads_blocking_next_to_thread():
outer = None
for fn in ast.walk(_TREE):
if isinstance(fn, ast.AsyncFunctionDef) and fn.name == "generate_stream":
outer = fn
break
assert outer is not None, "generate_stream handler missing"
inner = None
for sub in ast.walk(outer):
if isinstance(sub, ast.AsyncFunctionDef) and sub.name == "stream":
inner = sub
break
assert inner is not None, "generate_stream inner stream() generator missing"
for sub in ast.walk(inner):
if isinstance(sub, (ast.For, ast.AsyncFor)):
it_src = ast.unparse(sub.iter)
assert "generate_chat_response" not in it_src, (
"generate_stream's inner stream() must not iterate "
"backend.generate_chat_response() directly -- that blocks the event "
"loop on every blocking subprocess read between tokens. Use "
"`await asyncio.to_thread(next, gen, _DONE)` inside a `while True` "
"loop instead"
)
found_to_thread_next = False
for sub in ast.walk(inner):
if not isinstance(sub, ast.Call):
continue
fn_expr = sub.func
if not (
isinstance(fn_expr, ast.Attribute)
and fn_expr.attr == "to_thread"
and isinstance(fn_expr.value, ast.Name)
and fn_expr.value.id == "asyncio"
):
continue
if sub.args and isinstance(sub.args[0], ast.Name) and sub.args[0].id == "next":
found_to_thread_next = True
break
assert found_to_thread_next, (
"generate_stream's inner stream() must call "
"`asyncio.to_thread(next, gen, _DONE)` to keep the event loop free while the "
"worker subprocess produces the next token"
)
def test_generate_stream_cancels_backend_on_stream_cancelled_error():
outer = None
for fn in ast.walk(_TREE):
if isinstance(fn, ast.AsyncFunctionDef) and fn.name == "generate_stream":
outer = fn
break
assert outer is not None, "generate_stream handler missing"
outer_src = ast.unparse(outer)
assert "cancel_event = threading.Event()" in outer_src
inner = None
for sub in ast.walk(outer):
if isinstance(sub, ast.AsyncFunctionDef) and sub.name == "stream":
inner = sub
break
assert inner is not None, "generate_stream inner stream() generator missing"
def _awaits_to_thread_gen_close(node: ast.AST) -> bool:
for sub in ast.walk(node):
if not isinstance(sub, ast.Await):
continue
call = sub.value
if not isinstance(call, ast.Call):
continue
fn_expr = call.func
if not (
isinstance(fn_expr, ast.Attribute)
and fn_expr.attr == "to_thread"
and isinstance(fn_expr.value, ast.Name)
and fn_expr.value.id == "asyncio"
):
continue
if not call.args:
continue
close_expr = call.args[0]
if (
isinstance(close_expr, ast.Attribute)
and close_expr.attr == "close"
and isinstance(close_expr.value, ast.Name)
and close_expr.value.id == "gen"
):
return True
return False
found_cancel_kwarg = False
found_cancel_handler = False
found_finally_cleanup = False
for sub in ast.walk(inner):
if isinstance(sub, ast.Call):
call_src = ast.unparse(sub.func)
if call_src.endswith("generate_chat_response"):
found_cancel_kwarg = any(
kw.arg == "cancel_event"
and isinstance(kw.value, ast.Name)
and kw.value.id == "cancel_event"
for kw in sub.keywords
)
if isinstance(sub, ast.ExceptHandler):
exc_src = ast.unparse(sub.type) if sub.type is not None else ""
if exc_src != "asyncio.CancelledError":
continue
body_src = "\n".join(ast.unparse(stmt) for stmt in sub.body)
found_cancel_handler = (
"cancel_event.set()" in body_src
and "backend.reset_generation_state()" in body_src
and any(isinstance(stmt, ast.Raise) and stmt.exc is None for stmt in sub.body)
)
if isinstance(sub, ast.Try) and sub.finalbody:
final_src = "\n".join(ast.unparse(stmt) for stmt in sub.finalbody)
found_finally_cleanup = (
"not completed" in final_src
and "not cancel_event.is_set()" in final_src
and "cancel_event.set()" in final_src
and "backend.reset_generation_state()" in final_src
and _awaits_to_thread_gen_close(sub)
)
assert found_cancel_kwarg, (
"generate_stream must pass cancel_event into backend.generate_chat_response "
"so cancelled streams can stop backend generation"
)
assert found_cancel_handler, (
"generate_stream must catch asyncio.CancelledError, set cancel_event, "
"reset backend state, and re-raise"
)
assert found_finally_cleanup, (
"generate_stream cleanup must cancel/reset incomplete streams and "
"offload gen.close() with asyncio.to_thread so backend joins cannot "
"block the event loop"
)
def test_stream_chunks_cancel_branch_resets_backend_state():
# The cancel branch must call backend.reset_generation_state() to flush
# GPU/KV-cache state, else cancel-via-POST leaves the subprocess dirty.
@ -542,6 +684,106 @@ def test_unsloth_stream_loop_breaks_on_external_cancel_event():
)
def test_generate_stream_stays_responsive_under_blocking_next():
# Same sync-generator shape as generate_stream, with resp_queue.get modeled
# by sleep. The output must stay unchanged while next() moves off-loop.
chunks = ["alpha", "beta", "gamma", "delta"]
def _generate_chat_response():
for chunk in chunks:
time.sleep(0.08)
yield chunk
def _sse(chunk):
return f"data: {json.dumps({'content': chunk})}\n\n"
async def _direct_loop():
out = []
for chunk in _generate_chat_response():
out.append(_sse(chunk))
out.append("data: [DONE]\n\n")
return out
async def _to_thread_loop():
_DONE = object()
gen = _generate_chat_response()
out = []
try:
while True:
chunk = await asyncio.to_thread(next, gen, _DONE)
if chunk is _DONE:
break
out.append(_sse(chunk))
out.append("data: [DONE]\n\n")
return out
finally:
try:
gen.close()
except (RuntimeError, ValueError):
pass
async def _run_with_heartbeat(loop_coro):
ticks = 0
max_gap = 0.0
async def _heartbeat():
nonlocal ticks, max_gap
last = time.monotonic()
while True:
await asyncio.sleep(0.01)
now = time.monotonic()
max_gap = max(max_gap, now - last)
last = now
ticks += 1
heartbeat = asyncio.create_task(_heartbeat())
await asyncio.sleep(0)
try:
out = await loop_coro()
finally:
heartbeat.cancel()
try:
await heartbeat
except asyncio.CancelledError:
pass
return out, ticks, max_gap
async def _main():
direct_out, direct_ticks, direct_max_gap = await _run_with_heartbeat(_direct_loop)
threaded_out, threaded_ticks, threaded_max_gap = await _run_with_heartbeat(_to_thread_loop)
return (
direct_out,
direct_ticks,
direct_max_gap,
threaded_out,
threaded_ticks,
threaded_max_gap,
)
(
direct_out,
direct_ticks,
direct_max_gap,
threaded_out,
threaded_ticks,
threaded_max_gap,
) = asyncio.run(_main())
assert threaded_out == direct_out == [_sse(chunk) for chunk in chunks] + ["data: [DONE]\n\n"]
assert direct_ticks == 0, (
f"direct generate_stream loop should block the event loop; "
f"got {direct_ticks} heartbeat ticks and max gap {direct_max_gap:.3f}s"
)
assert threaded_ticks >= 8, (
f"to_thread generate_stream loop should let the event loop run; "
f"got {threaded_ticks} heartbeat ticks"
)
assert threaded_max_gap < 0.06, (
f"to_thread generate_stream loop should avoid long heartbeat gaps; "
f"got {threaded_max_gap:.3f}s"
)
def test_audio_stream_stays_responsive_under_blocking_next():
# Assert the pre-fix `for chunk in audio_input_generate()` pattern blocks the
# event loop, then confirm the post-fix pattern exits promptly.