Merge remote-tracking branch 'upstream/main' into fix-oom-labels-vulkan-6414

# Conflicts:
#	studio/backend/routes/inference.py
#	studio/backend/routes/training_vram.py
#	studio/backend/tests/test_gpu_memory_mode.py
This commit is contained in:
LeoBorcherding 2026-07-24 00:16:40 -05:00
commit aa79356c27
63 changed files with 5679 additions and 840 deletions

View file

@ -228,6 +228,7 @@ jobs:
tests/sh/test_system_node_readonly.sh \
tests/sh/test_nvcc_meets_llama_minimum.sh \
tests/sh/test_resolve_cuda_archs.sh \
tests/sh/test_staged_validation_enabled.sh \
tests/sh/test_tauri_install_exit_order.sh \
tests/sh/test_torch_constraint.sh \
tests/sh/test_torch_flavor.sh \

View file

@ -625,6 +625,36 @@ _is_pkg_installed() {
esac
}
# ── Helper: human-readable apt distro label for the sudo package prompt (#6207) ──
# Reads /etc/os-release so the Accept? prompt can say which distro we detected and
# that packages come from that distro's official apt repos (not a tarball).
_apt_distro_description() {
# Plain ( ... ) subshell — not $() — so case/;; stays bash-3.2-safe on macOS.
# Bash 3.2 misparses case arms inside command substitution and errors on `;;`.
(
if [ ! -r /etc/os-release ]; then
printf 'a debian-like system'
exit 0
fi
# shellcheck disable=SC1091
. /etc/os-release 2>/dev/null || true
if [ -n "${NAME:-}" ] && [ -n "${VERSION_ID:-}" ]; then
_ad_label="$NAME $VERSION_ID"
elif [ -n "${PRETTY_NAME:-}" ]; then
_ad_label="$PRETTY_NAME"
elif [ -n "${NAME:-}" ]; then
_ad_label="$NAME"
else
printf 'a debian-like system'
exit 0
fi
case " ${ID:-} ${ID_LIKE:-} " in
*" debian "*|*" ubuntu "*) _ad_label="${_ad_label} (debian-like)" ;;
esac
printf '%s' "$_ad_label"
)
}
# ── Helper: install packages via apt, escalating to sudo only if needed ──
# Usage: _smart_apt_install pkg1 pkg2 pkg3 ...
_smart_apt_install() {
@ -655,11 +685,14 @@ _smart_apt_install() {
# Step 3: Escalate -- need elevated permissions for remaining packages
if command -v sudo >/dev/null 2>&1; then
_ad_desc="$(_apt_distro_description)"
echo ""
echo " !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!"
echo " WARNING: We require sudo elevated permissions to install:"
echo " $_STILL_MISSING"
echo " If you accept, we'll run sudo now, and it'll prompt your password."
echo " Detected ${_ad_desc}."
echo " If you accept, we'll run sudo apt-get to install these packages"
echo " from your distro's official repositories (not a third-party tarball)."
echo " !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!"
echo ""
printf " Accept? [Y/n] "

View file

@ -93,6 +93,17 @@ huggingfacenotorch = [
"trl>=0.18.2,!=0.19.0,<=0.24.0",
"sentence-transformers",
]
# torchcodec backend for Gemma audio / datasets>=4 (#7225).
# Pick the audio-torch* pin matching your torch minor (see TORCH_TORCHCODEC).
audio-torch210 = [
"torchcodec>=0.10.0,<0.11.0 ; python_version >= '3.10'",
]
audio-torch290 = [
"torchcodec>=0.8.0,<0.10.0 ; python_version >= '3.10'",
]
audio-torch280 = [
"torchcodec>=0.6.0,<0.8.0 ; python_version >= '3.9'",
]
huggingface = [
"unsloth[huggingfacenotorch]",
"unsloth_zoo>=2026.7.6",
@ -532,16 +543,19 @@ cu126-torch2100 = [
"unsloth[huggingface]",
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0",
"unsloth[cu126onlytorch2100]",
"unsloth[audio-torch210]",
]
cu128-torch2100 = [
"unsloth[huggingface]",
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0",
"unsloth[cu128onlytorch2100]",
"unsloth[audio-torch210]",
]
cu130-torch2100 = [
"unsloth[huggingface]",
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0",
"unsloth[cu130onlytorch2100]",
"unsloth[audio-torch210]",
]
kaggle = [
"unsloth[huggingface]",
@ -831,16 +845,19 @@ cu126-ampere-torch2100 = [
"unsloth[huggingface]",
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0",
"unsloth[cu126onlytorch2100]",
"unsloth[audio-torch210]",
]
cu128-ampere-torch2100 = [
"unsloth[huggingface]",
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0",
"unsloth[cu128onlytorch2100]",
"unsloth[audio-torch210]",
]
cu130-ampere-torch2100 = [
"unsloth[huggingface]",
"bitsandbytes>=0.45.5,!=0.46.0,!=0.48.0",
"unsloth[cu130onlytorch2100]",
"unsloth[audio-torch210]",
]
flashattentiontorch260abiFALSEcu12x = [
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp39-cp39-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.9'",
@ -1125,7 +1142,8 @@ intelgputorch210 = [
"torchvision @ https://download.pytorch.org/whl/xpu/torchvision-0.25.0%2Bxpu-cp313-cp313-win_amd64.whl#sha256=1c4b44b36a557f7381e3076fb8843366742238648441d607c8d049c6da0f8886 ; sys_platform == 'win32' and python_version == '3.13' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
]
intel-gpu-torch210 = [
"unsloth[intelgputorch210]"
"unsloth[intelgputorch210]",
"unsloth[audio-torch210]",
]
intelgputorch2110 = [
"unsloth_zoo[intelgpu]",
@ -1279,6 +1297,7 @@ rocm72-torch2100 = [
"torchvision @ https://repo.radeon.com/rocm/manylinux/rocm-rel-7.2/torchvision-0.25.0%2Brocm7.2.0.git82df5f59-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
"torchvision @ https://repo.radeon.com/rocm/manylinux/rocm-rel-7.2/torchvision-0.25.0%2Brocm7.2.0.git82df5f59-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
"torchvision @ https://repo.radeon.com/rocm/manylinux/rocm-rel-7.2/torchvision-0.25.0%2Brocm7.2.0.git82df5f59-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
"unsloth[audio-torch210]",
]
rocm711-torch2100 = [
"unsloth[amd]",
@ -1297,6 +1316,7 @@ rocm711-torch2100 = [
"torchvision @ https://repo.radeon.com/rocm/manylinux/rocm-rel-7.1.1/torchvision-0.25.0%2Brocm7.1.1.git82df5f59-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
"torchvision @ https://repo.radeon.com/rocm/manylinux/rocm-rel-7.1.1/torchvision-0.25.0%2Brocm7.1.1.git82df5f59-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
"torchvision @ https://repo.radeon.com/rocm/manylinux/rocm-rel-7.1.1/torchvision-0.25.0%2Brocm7.1.1.git82df5f59-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
"unsloth[audio-torch210]",
]
[project.urls]

View file

@ -95,8 +95,8 @@ COLAB_ORACLE_BASE_URL = "https://raw.githubusercontent.com/googlecolab/backend-i
# Source: pytorch/torchcodec compatibility matrix on its README.
TORCH_TORCHCODEC: dict[str, set[str]] = {
"2.10": {"0.10"},
"2.9": {"0.7", "0.8", "0.9"},
"2.8": {"0.6"},
"2.9": {"0.8", "0.9"},
"2.8": {"0.6", "0.7"},
"2.7": {"0.3", "0.4", "0.5"},
"2.6": {"0.2", "0.3"},
"2.5": {"0.1", "0.2"},

View file

@ -1,5 +1,5 @@
{
"_comment": "scan_packages.py allowlist. Each entry is a CRITICAL/HIGH finding manually judged benign. Matched on (package, package-relative file, check, evidence_hash); evidence_hash is over the matched code with L<NN>: markers stripped, so version bumps and line shifts do not reopen an entry but changed code does. severity and evidence are for review only. Regenerate with --write-baseline AFTER reviewing every line.",
"_comment": "scan_packages.py allowlist (reviewed). Each entry is a CRITICAL/HIGH finding manually judged benign. Matched on (package, package-relative file, check, evidence_hash); evidence_hash is over the matched code with L<NN>: markers stripped, so version bumps and line shifts do not reopen an entry but changed code does. severity and evidence are for review only. Regenerate with --write-baseline AFTER reviewing every line.",
"version": 1,
"entries": [
{
@ -303,8 +303,8 @@
"file": "openai/_base_client.py",
"check": "C2 polling/beaconing loop detected",
"severity": "CRITICAL",
"evidence": "L264: while True: sha256:95ca67e46d42354ae650abbdc5b0d97df8b0ed43187800bf40f5690c3901b94b",
"evidence_hash": "a57d8d15fed0bf04f9967dcc18a18b80bb19f4095675bccbb78ac0450d7fce14"
"evidence": "L274: while True: sha256:90a38e5c1e26893c7c273354143612640e9a9c0f079d3e2b60612d79f24e80a6",
"evidence_hash": "1022e8e8649436ec64a98a9d9141d085452c49549fd2157b0278fc369a83ac66"
},
{
"package": "openai",
@ -319,8 +319,8 @@
"file": "openai/auth/_workload.py",
"check": "Accesses cloud metadata/IMDS AND makes network calls",
"severity": "CRITICAL",
"evidence": "IMDS: L96: url = \"http://169.254.169.254/metadata/identity/oauth2/token\" | L149: url = \"http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity\"\nNetwork: L77: http_client: httpx.Client | None = None, | L108: with httpx.Client() as client: | L133: http_client: httpx.Client | None = None, | L155: with httpx.Client() as client: | L248: with httpx.Client() as client:",
"evidence_hash": "1581d9f4a23393e9af23fbe5ef9f66807b22c5b5a3f1fe167254c9ebee108567"
"evidence": "IMDS: L97: url = \"http://169.254.169.254/metadata/identity/oauth2/token\" | L150: url = \"http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity\"\nNetwork: L78: http_client: httpx.Client | None = None, | L109: with httpx.Client() as client: | L134: http_client: httpx.Client | None = None, | L156: with httpx.Client() as client: | L251: exchange_client = DefaultHttpx2Client(follow_redirects=False) if self._use_httpx2 else httpx.Client()",
"evidence_hash": "9717e51cb961dc14c458955d91a1e48e3753997346ecea0106bded3a8d64bfe0"
},
{
"package": "openai",
@ -343,8 +343,8 @@
"file": "openai/resources/beta/responses/responses.py",
"check": "C2 polling/beaconing loop detected",
"severity": "CRITICAL",
"evidence": "L3999: while True: sha256:df298b6eaf3416589b79f4ef283f8fb76e54d505bfda8840673f8e6419117e2e",
"evidence_hash": "10ce5cb5a7097fcff4042ddcfb4802edda60aa4b7b113c8b926a52ddb76f78c2"
"evidence": "L4000: while True: sha256:f8ab538118daba9ec06e27399dbdc90a4521c3390e6a47a6348a1f180a83effd",
"evidence_hash": "31481ea83c687acc27144d72d3832d4fb98dd1c79fb5e0ddd85080de95997b9f"
},
{
"package": "openai",
@ -359,16 +359,16 @@
"file": "openai/resources/realtime/realtime.py",
"check": "C2 polling/beaconing loop detected",
"severity": "CRITICAL",
"evidence": "L310: while True: sha256:458198ff3d3f05870bf98c9564cbfd68c739e57b9bbe4120ed81e3eb6af74a05",
"evidence_hash": "a3165d21e46b3ce553795daeae53e8f80e8e89c5cb228e68e6dcaff54bca5a89"
"evidence": "L311: while True: sha256:5b63313072aae9ca28677e03426513ccf12221e4f4e0ea6c31efbe09790633b5",
"evidence_hash": "05e1af469d651b51673763a7c4cdf759af9472fb627b7b470adc28cc237bd650"
},
{
"package": "openai",
"file": "openai/resources/responses/responses.py",
"check": "C2 polling/beaconing loop detected",
"severity": "CRITICAL",
"evidence": "L3950: while True: sha256:1ce0b5a388c747945cdfda1a71b77afdfd03ae840d7aa9fa62f02eb00aa5e29f",
"evidence_hash": "6de300ebb5e6e17cb51c89cbcdf08515a44655182f0776f0908a9d1043ebbcd7"
"evidence": "L3951: while True: sha256:d68ef896bf0743ca430cfacb9a3353da1f3b9c51c3a21b6450a07a32b55aa2ac",
"evidence_hash": "160eecdd79b521bffbe8476f782b69a0724c35d1b19376a7600807165fd54f9f"
},
{
"package": "openai",

View file

@ -2023,6 +2023,10 @@ class LlamaCppBackend:
self._tensor_split: Optional[List[float]] = None
# User-picked physical GPU indices (None = automatic selection).
self._gpu_ids: Optional[List[int]] = None
# RAW requested GPU pin, before the fit narrowed it. self._gpu_ids records the
# EFFECTIVE (fit-narrowed) pin for /status; dedupe compares this raw value so a
# [0, 1] narrowed to [0] and re-sent as [0, 1] still matches (#7239).
self._requested_gpu_ids: Optional[List[int]] = None
# Layer load kept multi-GPU only to honor a downgraded tensor request, so a
# later explicit tensor-off reloads instead of deduping to it (#6659).
self._layer_preserves_tensor_intent: bool = False
@ -2494,6 +2498,46 @@ class LlamaCppBackend:
"""User-picked physical GPU indices, or None for automatic selection."""
return self._gpu_ids
@property
def requested_gpu_ids(self) -> Optional[List[int]]:
"""RAW requested GPU pin (before the fit narrowed it), or None for auto.
gpu_ids echoes the EFFECTIVE pin for /status."""
return self._requested_gpu_ids
def matches_gpu_ids(self, gpu_ids: Optional[List[int]]) -> bool:
"""Whether a requested pin is already satisfied by the active runner.
A regular GGUF load may narrow the requested placement pool to the
smallest fitting subset. Accept both the original request and the
effective status-echoed subset so either can round-trip without a
needless reload. Diffusion drives one device and keeps its existing
lowest-device normalization.
"""
if self._is_diffusion:
requested = [sorted(int(x) for x in gpu_ids)[0]] if gpu_ids else None
return requested == (self._gpu_ids or None)
requested = sorted(int(x) for x in gpu_ids) if gpu_ids else None
raw = self._requested_gpu_ids or None
effective = self._gpu_ids or None
return requested == raw or requested == effective
def _record_matching_gpu_request(self, gpu_ids: Optional[List[int]]) -> None:
"""Adopt the caller's explicit pool after a full already-loaded match.
Matching an effective subset avoids a reload, but the incoming request
is still the user's latest placement intent. Record it so status and a
later reload do not restore GPUs the user just removed.
"""
if self._is_diffusion:
self._requested_gpu_ids = [sorted(int(x) for x in gpu_ids)[0]] if gpu_ids else None
else:
self._requested_gpu_ids = sorted(int(x) for x in gpu_ids) if gpu_ids else None
if self._last_load_kwargs is not None:
self._last_load_kwargs["gpu_ids"] = (
list(self._requested_gpu_ids) if self._requested_gpu_ids else None
)
@property
def n_layers(self) -> Optional[int]:
"""Model layer count (GGUF block_count), or None if unknown."""
@ -4644,6 +4688,14 @@ class LlamaCppBackend:
LlamaCppBackend._gguf_skip_value(f, atype)
return None
@classmethod
def _gguf_path_is_diffusion(cls, gguf_path: str, model_identifier: str) -> bool:
"""Classify a downloaded GGUF without mutating the active backend."""
probe = object.__new__(cls)
probe._model_identifier = model_identifier
probe._read_gguf_metadata(gguf_path)
return probe._is_diffusion
def _read_gguf_metadata(self, gguf_path: str) -> None:
"""Read context_length, architecture params, and chat_template from a GGUF header.
@ -5095,11 +5147,14 @@ class LlamaCppBackend:
# the unload reset) so /status doesn't misreport TP and an identical
# re-Apply doesn't reload against stale tensor-parallel state.
self._tensor_parallel = False
# Record only the single device the runner actually uses (the lowest
# selected GPU, chosen above) -- not the whole pick. The diffusion runner
# is single-device, so echoing a multi-GPU list would misreport placement
# in /status and let a re-Apply dedup against GPUs the runner never used.
# The single-device runner records only the lowest selected GPU (chosen
# above), not the whole pick, and clears any explicit pin from a prior
# chat load; a multi-GPU list would misreport placement and mis-dedup.
self._gpu_ids = [sorted(gpu_ids)[0]] if gpu_ids else None
# The frontend prefers requested_gpu_ids when hydrating the picker.
# Diffusion uses only one device, so echo the collapsed effective pin,
# not unused members of the original request.
self._requested_gpu_ids = list(self._gpu_ids) if self._gpu_ids else None
if hf_variant:
self._hf_variant = hf_variant
elif gguf_path:
@ -6224,6 +6279,8 @@ class LlamaCppBackend:
gpu_layers: int = -1,
n_cpu_moe: int = 0,
tensor_split: Optional[List[float]] = None,
# Explicit GPU placement pool (issue #7164). None/[] = auto-select;
# the fitter may pin the smallest subset of this pool that fits.
gpu_ids: Optional[List[int]] = None,
n_threads: Optional[int] = None,
n_gpu_layers: Optional[int] = None, # caller compat, unused
@ -6321,15 +6378,63 @@ class LlamaCppBackend:
self._cancel_event.clear()
# ── Phase 1: kill old process (under lock, fast) ──────────
with self._lock:
self._kill_process()
# Resolve llama-server now but defer a not-found error: a block-diffusion
# GGUF uses the diffusion runner, and its arch is only known after the header.
binary = self._find_llama_server_binary()
is_vulkan_backend = self._is_vulkan_backend(binary)
# ── Vulkan-ordinal preflight (BEFORE the Phase 1 kill) ────────
# An explicit Vulkan pin the ggml probe never enumerated cannot be honored.
# Validate it ABOVE the kill so an invalid selection leaves the live model
# untouched: CUDA ids are range-checked at the route, but Vulkan ordinals are
# not, so a stale gpu_ids=[99] used to kill the server then 400, leaving
# nothing running (#7239). _get_gpu_memory needs only the binary (safe pre-
# download) and reuses the later fit's issubset logic. Guarded on a found
# Vulkan build + a pin so a deferred not-found stays deferred for diffusion.
if is_vulkan_backend and gpu_ids and binary:
_pf_wanted = {int(x) for x in gpu_ids}
_pf_probed = {g[0] for g in self._get_gpu_memory(binary)}
if not _pf_wanted.issubset(_pf_probed):
raise ValueError(
f"Requested Vulkan GPU ordinal(s) {sorted(_pf_wanted)} not "
f"present. Available Vulkan devices: {sorted(_pf_probed)}."
)
# A remote uncached GGUF may only reveal that it needs the
# single-device diffusion runner after download. On Vulkan, an
# explicit gpu_ids request cannot be mapped from ggml ordinals to
# that runner's CUDA physical index. Download and classify the main
# file before killing the healthy server so this late rejection is
# non-destructive. The Phase 2 call below reuses this cached path.
_preflight_model_path = None
if is_vulkan_backend and gpu_ids and hf_repo:
_resolved_repo = _resolve_repo_id_casing(hf_repo)
if _resolved_repo != hf_repo:
logger.info(
"Using cached repo_id casing '%s' for requested '%s'",
_resolved_repo,
hf_repo,
)
hf_repo = _resolved_repo
with _hf_offline_if_dns_dead():
_preflight_model_path = self._download_gguf(
hf_repo = hf_repo,
hf_variant = hf_variant,
hf_token = hf_token,
)
if self._gguf_path_is_diffusion(_preflight_model_path, model_identifier):
raise ValueError(
"GPU selection (gpu_ids) is not supported for a DiffusionGemma "
"GGUF on a Vulkan llama.cpp build: the diffusion runner selects "
"its device by CUDA physical index, which has no defined mapping "
"to ggml Vulkan device ordinals. Omit gpu_ids to use the default "
"device."
)
# ── Phase 1: kill old process (under lock, fast) ──────────
with self._lock:
self._kill_process()
# ── Phase 2: download (NO lock held, so cancel can proceed) ──
# mtp_draft_path arrives set for local Gemma loads (detected
# sibling); for -hf loads it's None here and resolved just below.
@ -6351,7 +6456,7 @@ class LlamaCppBackend:
)
hf_repo = _resolved_repo
with _hf_offline_if_dns_dead():
model_path = self._download_gguf(
model_path = _preflight_model_path or self._download_gguf(
hf_repo = hf_repo,
hf_variant = hf_variant,
hf_token = hf_token,
@ -6401,6 +6506,18 @@ class LlamaCppBackend:
# Block-diffusion GGUFs (DiffusionGemma) cannot run on llama-server;
# serve them with the diffusion runner (same OpenAI-compat interface).
if self._is_diffusion:
# The diffusion runner pins its child by CUDA visibility mask, so a
# ggml Vulkan ordinal cannot be honored (wrong GPU / CPU fallback).
# Route and remote-download preflights reject before teardown; keep
# this as a final defense if classification ever disagrees.
if is_vulkan_backend and gpu_ids:
raise ValueError(
"GPU selection (gpu_ids) is not supported for a DiffusionGemma "
"GGUF on a Vulkan llama.cpp build: the diffusion runner selects "
"its device by CUDA physical index, which has no defined mapping "
"to ggml Vulkan device ordinals. Omit gpu_ids to use the default "
"device."
)
# Not a tensor/layer GGUF: clear any preserved-fallback flag from a
# prior load (this path skips the command builder that clears it).
self._layer_preserves_tensor_intent = False
@ -6637,6 +6754,12 @@ class LlamaCppBackend:
# Layer-fallback min GPUs; raised below on a tensor downgrade. Bound
# before the try so the --fit-on except path still has it (no UnboundLocal).
_layer_min_gpus = 1
# An explicit Vulkan ordinal absent from the ggml probe cannot be
# honored; flag it in the fit and reject after the try (raising inside
# would be swallowed into the --fit-on fallback). Bound before the try.
_vulkan_explicit_unmatched = False
_vulkan_requested_ids: list[int] = []
_vulkan_available_ordinals: list[int] = []
try:
gguf_size = self._get_gguf_size_bytes(model_path)
# Include GPU-loaded mmproj in the fit budget (#5825).
@ -6649,6 +6772,28 @@ class LlamaCppBackend:
# Pass binary so a Vulkan build probes ggml's Vulkan ordinals.
_gpu_mem = self._get_gpu_memory(binary)
gpus = [(idx, free) for idx, free, _t in _gpu_mem]
# Restrict the fit (and thus the layer plan + pin env) to the
# selected GPUs; fail-open if none match so a stale UI choice
# can't strand the load on CPU (issue #7164).
if gpu_ids:
# A Vulkan build indexes by ggml ordinal. An explicit ordinal
# absent from the probe can't be pinned, so reject after the try
# rather than fail-open onto a device the user didn't pick.
_wanted_ids = {int(x) for x in gpu_ids}
# Reject if ANY requested ordinal is absent, not only when none
# match: [0, 99] against {0, 1} silently drops 99. Comparing the
# full requested set (before filter narrows) still lets the fitter
# pick a valid subset later -- that is narrowing, not absence.
_probed_ordinals = {g[0] for g in gpus}
if is_vulkan_backend and not _wanted_ids.issubset(_probed_ordinals):
_vulkan_explicit_unmatched = True
_vulkan_requested_ids = sorted(_wanted_ids)
_vulkan_available_ordinals = sorted(_probed_ordinals)
# Restrict the probed pool to the selection; fail-open (keep the
# full pool) if none match so a stale UI choice can't strand the
# load on CPU (issue #7164).
_sel_gpus = [g for g in gpus if g[0] in _wanted_ids]
gpus = _sel_gpus if _sel_gpus else gpus
total_by_idx = {idx: total for idx, _f, total in _gpu_mem}
# GPU picker: restrict every mode to the chosen devices, so
# auto selection only considers them and manual mask to
@ -7475,6 +7620,17 @@ class LlamaCppBackend:
tp_tensor_split = None
effective_ctx = requested_ctx # fall back to original
# An unenumerated explicit Vulkan ordinal can't be pinned; fail loudly
# instead of fitting onto an unselected device. Clear the raw selection
# the early state-publish recorded so it never leaks into gpu_ids (#7239).
if _vulkan_explicit_unmatched:
self._gpu_ids = None
self._requested_gpu_ids = None
raise ValueError(
f"Requested Vulkan GPU ordinal(s) {_vulkan_requested_ids} not "
f"present. Available Vulkan devices: {_vulkan_available_ordinals}."
)
# GPU picker: when no narrower subset was chosen (manual, or
# a failed/file-size selection), pin the whole picked set so the
# model can't spill onto an unpicked GPU.
@ -7838,11 +7994,45 @@ class LlamaCppBackend:
", ".join(unsupported_cache_flags),
)
# Vulkan pins via --device (a cmd arg, unlike the env-based
# CUDA/ROCm pin below), emitted BEFORE user extras so llama.cpp's
# last-wins parsing lets a user --device override Unsloth's pick.
if is_vulkan_backend and gpu_indices is not None:
cmd += LlamaCppBackend._vulkan_pin_args(gpu_indices)
# Vulkan pins via --device (a cmd arg), before user extras so a user
# --device wins. Fall back to raw ids when the fit did not narrow.
_vulkan_pin_ids = gpu_indices if gpu_indices is not None else (gpu_ids or None)
# Record the pin actually applied (fit-narrowed gpu_indices, else the raw
# request) for the keep-warm loop, dedupe, and /status, so an explicit
# [0, 1] narrowed to [0] records [0] and /status never echoes an ordinal
# the child never saw. Auto selection (no gpu_ids) stays None (#7239).
if is_vulkan_backend:
# Only record an EXPLICIT Vulkan pin: an auto pick still narrows +
# pins below, but recording it would misreport an explicit pin and
# make dedupe miss the loaded server; mirrors the CUDA/ROCm branch.
self._gpu_ids = (
sorted(int(x) for x in _vulkan_pin_ids)
if (gpu_ids and _vulkan_pin_ids)
else None
)
elif gpu_ids:
# Physical pin: the fit-selected subset when the fit ran, else the raw
# user selection so an explicit choice is honoured even when the fit
# could not size the model.
_effective_pin_ids = (
[int(x) for x in gpu_indices]
if gpu_indices is not None
else [int(x) for x in gpu_ids]
)
self._gpu_ids = (
sorted(int(x) for x in _effective_pin_ids) if _effective_pin_ids else None
)
else:
self._gpu_ids = None
# Also record the RAW requested pin (before the fit narrowed it). Load
# dedupe compares this so a [0, 1] narrowed to [0] and re-sent as [0, 1]
# still matches, while /status keeps echoing the effective pin (#7239).
self._requested_gpu_ids = sorted(int(x) for x in gpu_ids) if gpu_ids else None
if is_vulkan_backend and _vulkan_pin_ids is not None:
cmd += LlamaCppBackend._vulkan_pin_args(_vulkan_pin_ids)
# User pass-through args go last so llama.cpp's last-wins parsing
# lets the user override Unsloth's auto-set flags. Already
@ -7911,10 +8101,10 @@ class LlamaCppBackend:
f"Data-center GPU detected: applied DC llama.cpp env tuning (multi_gpu={multi_gpu})"
)
# Pin to selected GPU(s). On ROCm, narrowing only
# CUDA_VISIBLE_DEVICES leaves an AMD child seeing the full set, so
# set HIP_VISIBLE_DEVICES too. Vulkan is pinned via --device
# (above), not here.
# Pin to selected GPU(s) (issue #7164; resolved above into gpu_indices).
# On ROCm, narrowing only CUDA_VISIBLE_DEVICES leaves the AMD child
# seeing the full set, so set HIP_VISIBLE_DEVICES too. Vulkan is pinned
# via --device (above), not here.
# A deliberate zero-offload load with no GPU companions runs
# entirely on CPU, yet a visible CUDA device still costs the child
# ~0.5 GB (context + compute scratch) that the CPU-only
@ -8835,16 +9025,10 @@ class LlamaCppBackend:
)
):
return False
# A changed GPU pick must reload (compare order-insensitively; None/[]
# both mean automatic). The diffusion runner collapses a multi-GPU pick
# to its single lowest device, so self._gpu_ids holds just that device;
# normalize the request the same way, or a multi-GPU pick that resolves
# to the same device needlessly reloads.
if self._is_diffusion:
requested_gpu_pick = [sorted(gpu_ids)[0]] if gpu_ids else None
else:
requested_gpu_pick = sorted(gpu_ids) if gpu_ids else None
if (self._gpu_ids or None) != requested_gpu_pick:
# A changed GPU pick must reload. Regular GGUF accepts either the raw
# requested placement pool or the effective status-echoed subset;
# diffusion compares its normalized single-device pick.
if not self.matches_gpu_ids(gpu_ids):
return False
# Compare on the canonical requested mode. With --spec-type in
@ -8902,6 +9086,7 @@ class LlamaCppBackend:
current = list(self._extra_args) if self._extra_args is not None else []
if list(extra_args) != current:
return False
self._record_matching_gpu_request(gpu_ids)
return True
def _classify_gpu_offload(
@ -9033,12 +9218,15 @@ class LlamaCppBackend:
self._supports_preserve_thinking = False
self._supports_tools = False
self._cache_type_kv = None
# GPU-pin state describes the active runner only; clear it so an explicit
# pin never leaks into the next (or diffusion) runner.
self._gpu_ids = None
self._requested_gpu_ids = None
self._tensor_parallel = False
self._gpu_memory_mode = "auto"
self._gpu_layers = -1
self._n_cpu_moe = 0
self._tensor_split = None
self._gpu_ids = None
self._layer_preserves_tensor_intent = False
self._speculative_type = None
self._requested_spec_mode = None

View file

@ -18,6 +18,7 @@ import queue
import random
import re
import shlex
import shutil
import ssl
import subprocess
import sys
@ -328,6 +329,7 @@ def _find_blocked_commands(command: str) -> set[str]:
# Directory holding the sandbox ``sitecustomize.py`` shim (code-interpreter
# path remap); placed on the sandboxed child's PYTHONPATH in _build_safe_env.
_SANDBOX_SITE_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "sandbox_site")
# ── "Approve for me" (permission_mode="auto") safety detection ──────────────
# Auto mode pauses only calls classified here as potentially unsafe. The sandbox
# and hard blocks (blocklist, rlimits) still apply at run time; this gate only
@ -2491,15 +2493,124 @@ def is_potentially_unsafe_tool_call(name: str, arguments: dict) -> bool:
return True
def _canon_win_path(p: str) -> str:
"""Canonical form for trust comparison: realpath (expands 8.3 aliases and
resolves junctions/symlinks) + normcase/normpath."""
return os.path.normcase(os.path.normpath(os.path.realpath(p)))
def _augment_native_program_roots(roots: list[str]) -> list[str]:
"""Add the native Program Files sibling for any x86 root by stripping the
`` (x86)`` suffix, so a 32-bit process (whose known-folder ids map only to
the x86 root) still trusts a 64-bit Git install."""
out = list(roots)
for root in roots:
base = root.rstrip("\\/")
if base.lower().endswith(" (x86)"):
native = base[: -len(" (x86)")]
if native and native not in out:
out.append(native)
return out
def _windows_program_roots() -> list[str]:
"""Program Files install roots, resolved ONLY from the Windows known-folder
API (SHGetKnownFolderPath). Fails closed (returns ``[]``) if the API is
unavailable: env vars (%ProgramFiles%, even %SystemDrive%) are caller-
overrideable and could relocate the trust boundary, so we never derive a
trusted root from them. On any real Windows host shell32 is present, so
this only returns empty in a broken/non-Windows environment where the
sandbox git-PATH feature is not needed anyway (#7317).
"""
roots: list[str] = []
try:
import ctypes
from ctypes import wintypes
# FOLDERID_ProgramFiles, _ProgramFilesX86, _ProgramFilesX64. The X64
# id (Win10 1703+) yields the native root even from a 32-bit process,
# where the first two both map to Program Files (x86).
folder_ids = (
"{905e63b6-c1bf-494e-b29c-65b732d3d21a}",
"{7C5A40EF-A0FB-4BFC-874A-C0F2E0B9FA8E}",
"{6D809377-6AF0-444b-8957-A3773F02200E}",
)
_SHGet = ctypes.windll.shell32.SHGetKnownFolderPath
_CoTaskMemFree = ctypes.windll.ole32.CoTaskMemFree
for fid in folder_ids:
guid = ctypes.create_string_buffer(16)
ctypes.windll.ole32.CLSIDFromString(wintypes.LPCWSTR(fid), ctypes.byref(guid))
ptr = ctypes.c_wchar_p()
if _SHGet(ctypes.byref(guid), 0, None, ctypes.byref(ptr)) == 0:
if ptr.value:
roots.append(ptr.value)
_CoTaskMemFree(ptr)
except Exception:
return []
return _augment_native_program_roots(roots)
def _resolve_trusted_windows_git() -> tuple[str, str]:
"""Find a git launcher in a TRUSTED Program Files dir. Returns
``(canonical_dir, ext)`` or ``("", "")``.
``shutil.which`` returns only the first PATH match, which may be an
untrusted user shim; scan the remaining PATH entries for a later trusted
Git so bare ``git`` still resolves (#7317).
"""
exts = [e for e in (os.environ.get("PATHEXT") or ".EXE;.CMD;.BAT;.COM").split(os.pathsep)]
candidates: list[str] = []
primary = shutil.which("git")
if primary:
candidates.append(primary)
for entry in (os.environ.get("PATH") or "").split(os.pathsep):
entry = entry.strip().strip('"')
if not entry or not os.path.isabs(entry):
continue
for ext in exts:
cand = os.path.join(entry, "git" + ext)
if os.path.isfile(cand):
candidates.append(cand)
for git_exe in candidates:
git_dir = os.path.dirname(git_exe)
if os.path.isabs(git_dir) and _is_trusted_windows_program_dir(git_dir):
return os.path.realpath(git_dir), os.path.splitext(git_exe)[1].upper()
return "", ""
def _is_trusted_windows_program_dir(path: str) -> bool:
"""True when ``path`` sits under a system-managed Program Files root.
Only the Program Files roots are trusted (admin-writable only), resolved
via the known-folder API so an overridden env var cannot relocate them,
never ``%SystemRoot%`` (Git does not install there and it holds
world-writable subdirs like ``Windows\\Temp``). Per-user managers
(Scoop/Choco shims under the profile) are refused. Paths are canonicalized
so 8.3 aliases and junctions still resolve to their real root (#7317).
"""
norm = _canon_win_path(path)
for root in _windows_program_roots():
root_norm = _canon_win_path(root)
if norm == root_norm or norm.startswith(root_norm + os.sep):
return True
return False
def _build_safe_env(workdir: str) -> dict[str, str]:
"""Build a minimal, credential-free environment for sandboxed subprocesses.
Whitelist-built from scratch (parent env NOT inherited): only PATH/HOME/
TMPDIR/LANG/TERM/PYTHONIOENCODING/PYTHONPATH (+VIRTUAL_ENV or Windows
SystemRoot) reach the child; all credential vars (HF_TOKEN, AWS_*, etc.)
are absent. HOME points at the sandbox workdir so SDKs can't read the
SystemRoot and a minimal PATHEXT) reach the child; all credential vars
(HF_TOKEN, AWS_*, etc.) are absent. HOME points at the sandbox workdir so SDKs can't read the
operator's cached creds. PYTHONPATH carries only the sandbox sitecustomize
shim directory.
PATH starts with the Studio interpreter / venv and OS system dirs so
``python``/``pip`` stay pinned. On Windows only, Git-for-Windows install
dirs from the host PATH are appended so bare ``git`` resolves (#7317).
User-writable host PATH entries (venv, ``node_modules/.bin``, etc.) are
never inherited they could shadow auto-safe terminal commands.
"""
# Start from the running interpreter's dir so 'python'/'pip' resolve to the
# same environment the Unsloth server runs in.
@ -2519,6 +2630,20 @@ def _build_safe_env(workdir: str) -> dict[str, str]:
else:
path_entries.extend(["/usr/local/bin", "/usr/bin", "/bin"])
# Windows Git installs live outside System32; inherit the dir of the git
# the HOST shell resolves, but ONLY when it sits under a system install
# root (Program Files, windir). A user-writable dir (Scoop/Choco shims)
# is refused: it would let an attacker drop rg.exe/jq.exe beside git and
# have an auto-approved bare command execute it (#7317).
git_ext = ""
if sys.platform == "win32":
# Append the CANONICAL (realpath) trusted git dir, scanning past any
# untrusted user shim that sorts first on PATH; the canonical path
# cannot be retargeted via a junction after the trust check.
_trusted_git_dir, git_ext = _resolve_trusted_windows_git()
if _trusted_git_dir:
path_entries.append(_trusted_git_dir)
# Deduplicate, preserving order.
deduped = list(dict.fromkeys(p for p in path_entries if p))
@ -2538,6 +2663,15 @@ def _build_safe_env(workdir: str) -> dict[str, str]:
# Windows needs SystemRoot for Python/subprocess to work.
if sys.platform == "win32":
env["SystemRoot"] = os.environ.get("SystemRoot", r"C:\Windows")
# Restrict PATHEXT so cwd .BAT/.CMD cannot hijack bare names (#7317).
pathext = ".EXE;.COM"
if git_ext and git_ext not in (".EXE", ".COM"):
# Keep the host git launcher (e.g. a .CMD shim) resolvable.
pathext += ";" + git_ext
env["PATHEXT"] = pathext
# cmd/CreateProcess search cwd before PATH for bare names; disable so
# a workdir rg.exe/git.exe cannot shadow auto-approved commands.
env["NoDefaultCurrentDirectoryInExePath"] = "1"
return env

View file

@ -74,7 +74,7 @@ class LoadRequest(BaseModel):
)
gpu_ids: Optional[List[int]] = Field(
None,
description = "Physical GPU indices to use, for example [0, 1]. Omit or pass [] to use automatic selection. Explicit gpu_ids are unsupported when the parent CUDA_VISIBLE_DEVICES uses UUID/MIG entries. For GGUF models the picked devices are pinned via CUDA/HIP_VISIBLE_DEVICES.",
description = "GPU placement pool, for example [0, 1]. Omit or pass [] to use automatic selection. CUDA/ROCm values are physical GPU indices and are unsupported when the parent CUDA_VISIBLE_DEVICES uses UUID/MIG entries; Vulkan values are ggml device ordinals. For GGUF models the fitter may pin the smallest subset of this pool that fits.",
)
speculative_type: Optional[str] = Field(
None,
@ -485,7 +485,14 @@ class LoadResponse(BaseModel):
)
gpu_ids: Optional[List[int]] = Field(
None,
description = "Physical GPU indices the model is pinned to, or None for automatic selection.",
description = "Effective GPU indices the model is using after fit-time narrowing, or None for automatic selection.",
)
requested_gpu_ids: Optional[List[int]] = Field(
None,
description = (
"GPU placement pool requested by the user before fit-time narrowing, "
"or None for automatic selection."
),
)
@ -649,7 +656,14 @@ class InferenceStatusResponse(BaseModel):
)
gpu_ids: Optional[List[int]] = Field(
None,
description = "Physical GPU indices the model is pinned to, or None for automatic selection.",
description = "Effective GPU indices the model is using after fit-time narrowing, or None for automatic selection.",
)
requested_gpu_ids: Optional[List[int]] = Field(
None,
description = (
"GPU placement pool requested by the user before fit-time narrowing, "
"or None for automatic selection."
),
)
llama_cpp_supports_mtp: bool = Field(
True,

View file

@ -47,6 +47,14 @@ class ProviderCreate(BaseModel):
None,
description = "Custom base URL (overrides registry default). Omit to use the default.",
)
models: list[str] = Field(
default_factory = list,
description = "Enabled model IDs for this connection",
)
available_models: list[str] = Field(
default_factory = list,
description = "Discovered catalog model IDs last fetched for this connection",
)
class ProviderUpdate(BaseModel):
@ -55,6 +63,11 @@ class ProviderUpdate(BaseModel):
display_name: Optional[str] = Field(None, description = "New display name")
base_url: Optional[str] = Field(None, description = "New base URL")
is_enabled: Optional[bool] = Field(None, description = "Enable or disable this provider")
models: Optional[list[str]] = Field(None, description = "Enabled model IDs for this connection")
available_models: Optional[list[str]] = Field(
None,
description = "Discovered catalog model IDs last fetched for this connection",
)
class ProviderResponse(BaseModel):
@ -65,6 +78,14 @@ class ProviderResponse(BaseModel):
display_name: str = Field(..., description = "User-chosen label")
base_url: str = Field(..., description = "API base URL")
is_enabled: bool = Field(True, description = "Whether this provider is enabled")
models: list[str] = Field(
default_factory = list,
description = "Enabled model IDs for this connection",
)
available_models: list[str] = Field(
default_factory = list,
description = "Discovered catalog model IDs last fetched for this connection",
)
created_at: str = Field(..., description = "ISO 8601 creation timestamp")
updated_at: str = Field(..., description = "ISO 8601 last-update timestamp")

View file

@ -3241,15 +3241,10 @@ def _request_matches_loaded_settings(
)
):
return False
# A changed GPU pick must reload. The diffusion runner collapses a multi-GPU
# request to its single lowest device (it drives one device only), so the
# backend records just that device; compare the request the same way, or a
# multi-GPU pick that resolves to the same device needlessly reloads.
if llama_backend.is_diffusion:
_req_gpu_ids = [sorted(request.gpu_ids)[0]] if request.gpu_ids else None
else:
_req_gpu_ids = sorted(request.gpu_ids) if request.gpu_ids else None
if _req_gpu_ids != llama_backend.gpu_ids:
# A regular GGUF may narrow the requested placement pool. Accept either the
# original request or the effective status-echoed subset; diffusion keeps
# its single-device normalization.
if not llama_backend.matches_gpu_ids(request.gpu_ids):
return False
# Preserved tensor->layer fallback (both report tensor=off, so the check above
# matches): if the user now explicitly drops tensor intent, reload so placement
@ -3897,15 +3892,19 @@ def _classify_diffusion_gguf(config: ModelConfig) -> Optional[bool]:
"""Classify a GGUF as diffusion, normal, or unknown before it is loaded.
``None`` is important here: a remote GGUF whose header is not cached can
still be routed to the single-GPU diffusion runner after download. Treating
that case as normal would let Manual mode skip the training guard even
though the runner ignores Manual's llama-server placement controls.
still be routed to the single-GPU diffusion runner after download. Default
placement keeps that unknown case guarded until the header is available.
"""
identity = " ".join(
str(getattr(config, attr, "") or "") for attr in ("identifier", "gguf_hf_repo", "gguf_file")
).lower()
if "diffusion" in identity:
return True
# Name-only hint, used ONLY as a pre-download fallback, scoped to the
# DiffusionGemma runner family: a bare "diffusion" substring is common in
# ordinary text-model names/paths (e.g. "stable-diffusion-prompt"), and treating
# those as diffusion falsely rejects a valid Vulkan+gpu_ids GGUF (#7239). Normalize
# non-alphanumerics so "DiffusionGemma"/"diffusion-gemma" collapse to one token.
# The local header below stays authoritative.
name_says_diffusion = "diffusiongemma" in _re.sub(r"[^a-z0-9]+", "", identity)
try:
main = getattr(config, "gguf_file", None)
@ -3915,23 +3914,86 @@ def _classify_diffusion_gguf(config: ModelConfig) -> Optional[bool]:
if repo and variant:
from hub.utils.gguf import resolve_local_gguf_path
main = resolve_local_gguf_path(repo, variant)
if not main or not Path(main).is_file():
return None
probe = LlamaCppBackend()
probe._read_gguf_metadata(str(main))
if probe.is_diffusion:
return True
# A successfully decoded architecture proves that this is a normal
# llama-server GGUF. No architecture means the lightweight probe could
# not establish the routing decision, so preserve the unknown state.
if getattr(probe, "_architecture", None):
return False
return None
if main and Path(main).is_file():
# The local GGUF header is authoritative (same probe the loader uses), so
# it can't be fooled by a "diffusion"-flavored name/path.
probe = LlamaCppBackend()
probe._read_gguf_metadata(str(main))
if probe.is_diffusion:
return True
# A decoded architecture proves a normal llama-server GGUF; no architecture
# means the probe was inconclusive, so fall through to the name hint below.
if getattr(probe, "_architecture", None):
return False
except Exception as e:
logger.debug("Could not identify diffusion GGUF for training guard: %s", e)
# Header unavailable (remote uncached) or inconclusive: True only for the
# DiffusionGemma name family; otherwise None keeps an unknown remote GGUF guarded
# as potentially diffusion until its header proves otherwise.
return True if name_says_diffusion else None
async def _resolve_gguf_gpu_ids_for_request(
config: ModelConfig, gpu_ids: Optional[List[int]]
) -> Optional[List[int]]:
"""Resolve and fully validate an explicit GGUF GPU placement pool.
CUDA and ROCm use physical IDs. Vulkan uses ggml ordinals, so its device
existence check comes from the same ggml probe used by the loader. Both
/load and /validate call this before their training guard or any teardown.
"""
if not gpu_ids:
return None
from utils.hardware import DeviceType, get_device
from utils.hardware.hardware import resolve_requested_gpu_ids
is_vulkan = LlamaCppBackend._is_vulkan_backend()
if get_device() == DeviceType.XPU and not is_vulkan:
raise HTTPException(
status_code = 400,
detail = (
"GPU selection (gpu_ids) is not supported on Intel XPU. "
"Omit gpu_ids to use all devices."
),
)
if is_vulkan and _classify_diffusion_gguf(config) is True:
raise HTTPException(
status_code = 400,
detail = (
"GPU selection (gpu_ids) is not supported for a DiffusionGemma "
"GGUF on a Vulkan llama.cpp build: the diffusion runner selects "
"its device by CUDA physical index, which has no defined mapping "
"to ggml Vulkan device ordinals. Omit gpu_ids to use the default "
"device."
),
)
try:
resolved = resolve_requested_gpu_ids(gpu_ids, is_vulkan = is_vulkan)
except ValueError as exc:
raise HTTPException(status_code = 400, detail = str(exc)) from exc
if is_vulkan and resolved:
binary = LlamaCppBackend._find_llama_server_binary()
if binary:
probed = {
gpu[0] for gpu in await asyncio.to_thread(LlamaCppBackend._get_gpu_memory, binary)
}
wanted = {int(gpu_id) for gpu_id in resolved}
if not wanted.issubset(probed):
raise HTTPException(
status_code = 400,
detail = (
f"Requested Vulkan GPU ordinal(s) {sorted(wanted)} not "
f"present. Available Vulkan devices: {sorted(probed)}."
),
)
return resolved
def _guard_chat_load_against_training(
config: ModelConfig,
@ -3971,23 +4033,20 @@ def _guard_chat_load_against_training(
if is_gguf and gpu_memory_mode == "manual" and diffusion_kind is False:
return
# A Vulkan build's gpu_ids are ggml Vulkan ordinals with no defined mapping
# to the physical index space this guard sizes against (the free-VRAM rows
# of get_visible_gpu_utilization, _diffusion_gpu_arg's device token). Don't
# resolve them as physical ids, but keep the pick COUNT: dropping to the
# whole-pool estimate could OK a load that lands on a busy selected card and
# OOMs training, so size against the N most-constrained visible cards.
worst_case_gpu_count = None
if is_gguf and requested_gpu_ids:
# Vulkan GGUF pins are ggml ordinals, not CUDA physical IDs. Detect this
# before deriving a possible diffusion fallback device so an unknown remote
# GGUF never sends its ordinal through the CUDA single-device path; the
# can_load guard then sizes a Vulkan pick against the N most-constrained
# visible cards (via requested_gpu_ids count) instead of a physical index.
is_vulkan = False
if is_gguf:
try:
if LlamaCppBackend._is_vulkan_backend():
worst_case_gpu_count = len(set(requested_gpu_ids))
requested_gpu_ids = None
is_vulkan = LlamaCppBackend._is_vulkan_backend()
except Exception as e:
logger.debug("Vulkan backend check failed in chat-load guard: %s", e)
logger.warning("Could not detect Vulkan backend for chat-load guard: %s", e)
diffusion_gpu = None
if is_gguf and diffusion_kind is not False:
if is_gguf and diffusion_kind is not False and not (is_vulkan and requested_gpu_ids):
# Use the same token selection as the runner: an explicit pick wins,
# followed by DG_GPU, the first parent-visible token, then GPU 0.
diffusion_gpu = LlamaCppBackend._diffusion_gpu_arg(
@ -4014,9 +4073,9 @@ def _guard_chat_load_against_training(
max_seq_length = max_seq_length,
requested_gpu_ids = requested_gpu_ids,
is_gguf = is_gguf,
is_vulkan = is_vulkan,
required_override_gb = required_override_gb,
single_device_gpu = diffusion_gpu,
worst_case_gpu_count = worst_case_gpu_count,
)
if ok:
return
@ -4321,6 +4380,7 @@ async def _load_model_impl(
# Skip if a prior audio probe failed -- let load_model retry.
and getattr(llama_backend, "_audio_probed", True)
):
llama_backend._record_matching_gpu_request(request.gpu_ids)
logger.info(
"Model already loaded (GGUF): "
f"{model_log_label} variant={request.gguf_variant or llama_backend.hf_variant}, skipping reload"
@ -4367,6 +4427,7 @@ async def _load_model_impl(
n_layers = llama_backend.n_layers,
n_moe_layers = llama_backend.n_moe_layers,
gpu_ids = llama_backend.gpu_ids,
requested_gpu_ids = llama_backend.requested_gpu_ids,
)
else:
if (
@ -4433,64 +4494,14 @@ async def _load_model_impl(
# Normalize gpu_ids: empty list means auto-selection, same as None
effective_gpu_ids = request.gpu_ids if request.gpu_ids else None
# GGUF supports gpu_ids: validate the pick up front (before the training
# guard) so a bad pick is a clean 400, not masked by a VRAM 409. Rejects
# negative / out-of-range / duplicate ids and UUID/MIG parents. XPU hosts
# are rejected outright: the picker's indices are torch-xpu ordinals neither
# applicator speaks (CUDA/HIP masks don't apply, the Vulkan --device pin
# uses ggml's own Vulkan ordinals), so a pick could land on the wrong device.
if config.is_gguf and effective_gpu_ids is not None:
from utils.hardware import DeviceType, get_device
from utils.hardware.hardware import resolve_requested_gpu_ids
# A Vulkan build validates against ggml's own Vulkan ordinals:
# /api/system reports gguf_devices in that space, load_model pins the
# pick with --device Vulkan<i>, so probe, picker, and pin share one
# index space (physical/torch ids are never involved). Check it
# BEFORE the XPU ban -- a Vulkan pick on an Intel/XPU host does not
# rely on torch-xpu ordinals, so the ban must not reject it.
if LlamaCppBackend._is_vulkan_backend():
# Diffusion GGUFs bypass llama-server: the diffusion runner
# forwards gpu_ids[0] as a CUDA/DG_GPU device token, NOT
# --device Vulkan<i>, so a Vulkan ordinal would target the wrong
# card. Reject only a CONFIRMED diffusion GGUF (`is True`) here:
# `None` is the ordinary first-load case for an uncached Hub GGUF
# (no local header to classify yet), and rejecting it would make
# the picker unusable for remote GGUFs. An uncached model that
# turns out to be diffusion is caught post-download by the
# spawn-time Vulkan backstop in load_model.
if _classify_diffusion_gguf(config) is True:
raise HTTPException(
status_code = 400,
detail = (
"GPU selection (gpu_ids) is not supported for diffusion "
"GGUF models on a Vulkan llama.cpp build: the diffusion "
"runner cannot map ggml Vulkan ordinals. Omit gpu_ids."
),
)
# validate_vulkan_gpu_ids may spawn the Vulkan device probe
# (blocking subprocess.run). Run it off the event loop so a
# stalled driver/probe can't freeze status/progress/unload for
# up to the probe timeout.
try:
await asyncio.to_thread(
LlamaCppBackend.validate_vulkan_gpu_ids, effective_gpu_ids
)
except ValueError as exc:
raise HTTPException(status_code = 400, detail = str(exc)) from exc
elif get_device() == DeviceType.XPU:
raise HTTPException(
status_code = 400,
detail = (
"GPU selection (gpu_ids) is not supported on Intel XPU. "
"Omit gpu_ids to use all devices."
),
)
else:
try:
resolve_requested_gpu_ids(effective_gpu_ids)
except ValueError as exc:
raise HTTPException(status_code = 400, detail = str(exc)) from exc
# Validate the full GGUF placement pool before the training guard so an
# invalid physical ID or Vulkan ordinal is a clean 400, not a masked VRAM
# 409. The shared helper rejects XPU picks (unless Vulkan), rejects a
# diffusion GGUF pick on Vulkan, and runs the Vulkan device probe off the
# event loop. The same helper is used by /validate.
gguf_gpu_ids: Optional[List[int]] = None
if config.is_gguf:
gguf_gpu_ids = await _resolve_gguf_gpu_ids_for_request(config, effective_gpu_ids)
if not config.is_gguf and _mlx_distributed_launch_detected():
raise HTTPException(
status_code = 400,
@ -4614,8 +4625,9 @@ async def _load_model_impl(
gpu_layers = request.gpu_layers,
n_cpu_moe = request.n_cpu_moe,
tensor_split = request.tensor_split,
gpu_ids = effective_gpu_ids,
n_parallel = _n_parallel,
# Issue #7164: explicit GPU pin resolved to physical ids above.
gpu_ids = gguf_gpu_ids,
)
if config.gguf_hf_repo:
# HF mode: download via huggingface_hub then start llama-server
@ -4789,6 +4801,7 @@ async def _load_model_impl(
n_layers = llama_backend.n_layers,
n_moe_layers = llama_backend.n_moe_layers,
gpu_ids = llama_backend.gpu_ids,
requested_gpu_ids = llama_backend.requested_gpu_ids,
)
# ── Standard path: load via Unsloth/transformers ──────────
@ -5082,54 +5095,12 @@ async def validate_model(
# Apply the same training coexistence policy as /load before the frontend
# unloads the current model.
effective_gpu_ids = request.gpu_ids if request.gpu_ids else None
# Mirror /load: GGUF supports gpu_ids, so validate the pick (a bad one is
# a clean 400) before the guard sizes the model against training VRAM.
# XPU-host picks are rejected like /load (no defined mapping from the
# picker's torch-xpu ordinals to the launcher's device spaces).
if config.is_gguf and effective_gpu_ids is not None:
from utils.hardware import DeviceType, get_device
from utils.hardware.hardware import resolve_requested_gpu_ids
# Mirror /load: a Vulkan build validates the pick in ggml's own
# Vulkan ordinal space (the space the --device pin uses), and rejects
# picks for CONFIRMED diffusion GGUFs (their runner takes a CUDA/DG_GPU
# token, not --device Vulkan<i>, so an ordinal targets the wrong card).
# `None` (uncached, unclassifiable) is allowed through so first-time
# remote GGUF loads still work; the spawn-time backstop catches an
# uncached model that turns out to be diffusion after download. Check
# the Vulkan path BEFORE the XPU ban: a Vulkan pick on an XPU host
# uses ggml ordinals, not torch-xpu ones, so the ban must not hide it.
if LlamaCppBackend._is_vulkan_backend():
if _classify_diffusion_gguf(config) is True:
raise HTTPException(
status_code = 400,
detail = (
"GPU selection (gpu_ids) is not supported for diffusion "
"GGUF models on a Vulkan llama.cpp build: the diffusion "
"runner cannot map ggml Vulkan ordinals. Omit gpu_ids."
),
)
# Off-loop: validate_vulkan_gpu_ids may spawn the blocking Vulkan
# probe subprocess (see /load).
try:
await asyncio.to_thread(
LlamaCppBackend.validate_vulkan_gpu_ids, effective_gpu_ids
)
except ValueError as exc:
raise HTTPException(status_code = 400, detail = str(exc)) from exc
elif get_device() == DeviceType.XPU:
raise HTTPException(
status_code = 400,
detail = (
"GPU selection (gpu_ids) is not supported on Intel XPU. "
"Omit gpu_ids to use all devices."
),
)
else:
try:
resolve_requested_gpu_ids(effective_gpu_ids)
except ValueError as exc:
raise HTTPException(status_code = 400, detail = str(exc)) from exc
# Mirror /load: the shared helper validates the GGUF pick (a bad one is a
# clean 400) before the guard sizes against training VRAM -- rejecting
# XPU picks unless Vulkan, rejecting a diffusion GGUF pick on Vulkan, and
# running the Vulkan device probe off the event loop.
if config.is_gguf:
await _resolve_gguf_gpu_ids_for_request(config, effective_gpu_ids)
effective_load_in_4bit = _effective_load_in_4bit(config, request.load_in_4bit)
# Both checks cover the [adapter, base] set (matching the scan route and workers):
@ -5954,6 +5925,7 @@ async def get_status(current_subject: str = Depends(get_current_subject)):
n_layers = llama_backend.n_layers,
n_moe_layers = llama_backend.n_moe_layers,
gpu_ids = llama_backend.gpu_ids,
requested_gpu_ids = llama_backend.requested_gpu_ids,
llama_cpp_supports_mtp = _supports_mtp,
spec_fallback_reason = llama_backend.spec_fallback_reason,
llama_cpp_prebuilt_stale = _stale,
@ -6118,6 +6090,7 @@ async def generate_audio(
# Advertised repo id after an auto-switch load, else a clean public id,
# never the absolute .gguf path.
model_name = _llama_public_model_id(llama_backend)
_audio_model_id = getattr(llama_backend, "model_identifier", None) or model_name
gen = lambda: llama_backend.generate_audio_response(
text = text,
audio_type = llama_backend._audio_type,
@ -6136,6 +6109,7 @@ async def generate_audio(
if not model_info.get("is_audio"):
raise HTTPException(status_code = 400, detail = "Active model is not an audio model.")
model_name = public_model_id(backend.active_model_name)
_audio_model_id = getattr(backend, "active_model_name", None) or model_name
gen = lambda: backend.generate_audio_response(
text = text,
temperature = payload.temperature,
@ -6147,6 +6121,13 @@ async def generate_audio(
use_adapter = payload.use_adapter,
)
# Apply per-model recommended sampling + any operator UNSLOTH_SAMPLING_* pin before
# generating, so `unsloth run --temperature` (and the other pins) and per-model
# recommendations reach audio (TTS) generation too, not just chat. The gen lambdas read
# payload.* lazily at call time, so filling here takes effect; this covers both the direct
# /audio/generate route and the chat-completions audio branches that delegate here.
_fill_recommended_sampling_openai(payload, _audio_model_id)
try:
wav_bytes, sample_rate = await asyncio.to_thread(gen)
except Exception as e:
@ -7438,6 +7419,51 @@ async def delete_openai_container(
await client.close()
def _fill_recommended_sampling_openai(payload, model_id) -> None:
"""Apply per-model recommended sampling (and any operator UNSLOTH_SAMPLING_* pin) to a
ChatCompletionRequest in place.
Only the sampling fields the client did NOT explicitly send (tracked via
``model_fields_set``) are overwritten, so a client that sets a field stays byte-identical
unless an operator pins it. Fields with neither a recommendation nor a pin keep their
existing (schema-default) value.
"""
from utils.inference.inference_config import resolve_effective_sampling, SAMPLING_FIELD_NAMES
explicit = {
f: (getattr(payload, f) if f in payload.model_fields_set else None)
for f in SAMPLING_FIELD_NAMES
}
effective = resolve_effective_sampling(model_id, explicit)
for field, value in effective.items():
setattr(payload, field, value)
# /v1/completions is proxied to llama-server verbatim; its repetition knob is "repeat_penalty",
# and every other sampling field keeps its name (mirrors _build_passthrough_payload).
_COMPLETIONS_SAMPLING_BODY_KEY = {"repetition_penalty": "repeat_penalty"}
def _fill_recommended_sampling_completions(body: dict, model_id) -> None:
"""Apply per-model recommended sampling (and any operator UNSLOTH_SAMPLING_* pin) to a raw
``/v1/completions`` body in place, so the legacy (non-chat) endpoint honors the same pins as
``/v1/chat/completions``.
Unlike :func:`_fill_recommended_sampling_openai`, which fills a ChatCompletionRequest whose
schema already carries per-field defaults, this body is proxied to llama-server as-is. A field
with no operator pin, client value, or per-model recommendation is therefore left untouched
(``fill_defaults = False``) so llama-server keeps its own default rather than being forced onto
this schema's value. llama-server names the repetition knob ``repeat_penalty``, so read and
write that alias for the client-sent value and any pin.
"""
from utils.inference.inference_config import resolve_effective_sampling, SAMPLING_FIELD_NAMES
explicit = {f: body.get(_COMPLETIONS_SAMPLING_BODY_KEY.get(f, f)) for f in SAMPLING_FIELD_NAMES}
effective = resolve_effective_sampling(model_id, explicit, fill_defaults = False)
for field, value in effective.items():
body[_COMPLETIONS_SAMPLING_BODY_KEY.get(field, field)] = value
@router.post("/chat/completions")
async def openai_chat_completions(
payload: ChatCompletionRequest,
@ -7767,6 +7793,13 @@ async def openai_chat_completions(
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
created = int(time.time())
# Apply recommended sampling + operator pins to the omitted fields before generating,
# so audio-input (non-whisper) generation honors `unsloth run --temperature` and
# per-model recommendations like chat does. Whisper (ASR) ignores these fields.
_fill_recommended_sampling_openai(
payload, getattr(backend, "active_model_name", None) or model_name
)
def audio_input_generate():
if model_info.get("audio_type") == "whisper":
return backend.generate_whisper_response(
@ -7910,6 +7943,18 @@ async def openai_chat_completions(
),
)
# Apply per-model recommended sampling (and any operator UNSLOTH_SAMPLING_* pin) to the
# fields the client omitted, so agents and API clients get the model's tuned defaults
# unless they set the field explicitly. Placed after external-provider routing (which
# returned above) so only local llama-server / transformers requests are touched, and it
# covers both the passthrough and non-passthrough branches below since both read payload.*.
_reco_model_id = (
getattr(llama_backend, "model_identifier", None)
if using_gguf
else getattr(backend, "active_model_name", None)
) or model_name
_fill_recommended_sampling_openai(payload, _reco_model_id)
# ── Standard OpenAI function-calling pass-through (GGUF only) ────
# When a client (opencode / Claude Code via OpenAI compat / Cursor /
# Continue / ...) sends standard OpenAI `tools` without Unsloth's
@ -10687,6 +10732,10 @@ async def openai_completions(request: Request, current_subject: str = Depends(ge
if _resolved_max_tokens is not None
else (llama_backend.context_length or _DEFAULT_MAX_TOKENS_FLOOR)
)
# Apply per-model recommended sampling and any operator UNSLOTH_SAMPLING_* pin to the raw
# body so /v1/completions honors the same pins as /v1/chat/completions; it is otherwise a
# verbatim proxy that would keep llama-server's defaults for every omitted sampling field.
_fill_recommended_sampling_completions(body, getattr(llama_backend, "model_identifier", None))
target_url = f"{llama_backend.base_url}/v1/completions"
is_stream = body.get("stream", False)
prompt_text = _flatten_monitor_prompt(body.get("prompt", ""))
@ -11645,6 +11694,9 @@ async def _responses_stream(
detail = "Image provided but current GGUF model does not support vision.",
)
# Streaming /v1/responses builds the passthrough body directly (bypassing
# openai_chat_completions), so apply recommended sampling here too.
_fill_recommended_sampling_openai(chat_req, getattr(llama_backend, "model_identifier", None))
body = _build_openai_passthrough_body(
chat_req, backend_ctx = llama_backend.context_length, llama_backend = llama_backend
)
@ -13076,14 +13128,28 @@ async def anthropic_messages(
# endpoint matches /v1/chat/completions.
_has_image = _normalize_anthropic_openai_images(openai_messages, llama_backend.is_vision)
temperature = payload.temperature if payload.temperature is not None else 0.6
top_p = payload.top_p if payload.top_p is not None else 0.95
top_k = payload.top_k if payload.top_k is not None else 20
min_p = payload.min_p if payload.min_p is not None else 0.01
repetition_penalty = (
payload.repetition_penalty if payload.repetition_penalty is not None else 1.0
# Fill omitted sampling fields with the per-model recommendation (or an operator
# UNSLOTH_SAMPLING_* pin); an explicit client value wins unless the operator pinned it.
# Anthropic sampling fields are Optional, so None already marks "client omitted".
from utils.inference.inference_config import resolve_effective_sampling
_anthropic_sampling = resolve_effective_sampling(
getattr(llama_backend, "model_identifier", None) or model_name,
{
"temperature": payload.temperature,
"top_p": payload.top_p,
"top_k": payload.top_k,
"min_p": payload.min_p,
"repetition_penalty": payload.repetition_penalty,
"presence_penalty": payload.presence_penalty,
},
)
presence_penalty = payload.presence_penalty if payload.presence_penalty is not None else 0.0
temperature = _anthropic_sampling["temperature"]
top_p = _anthropic_sampling["top_p"]
top_k = _anthropic_sampling["top_k"]
min_p = _anthropic_sampling["min_p"]
repetition_penalty = _anthropic_sampling["repetition_penalty"]
presence_penalty = _anthropic_sampling["presence_penalty"]
stop = payload.stop_sequences or None
# Translate Anthropic tool_choice to OpenAI format for llama-server. Falls

View file

@ -47,6 +47,20 @@ logger = structlog.get_logger(__name__)
router = APIRouter()
def _provider_response(row: dict) -> ProviderResponse:
return ProviderResponse(
id = row["id"],
provider_type = row["provider_type"],
display_name = row["display_name"],
base_url = row["base_url"],
is_enabled = bool(row["is_enabled"]),
models = row.get("models") or [],
available_models = row.get("available_models") or [],
created_at = row["created_at"],
updated_at = row["updated_at"],
)
# ── Public key for API key encryption ─────────────────────────────
@ -89,18 +103,7 @@ async def get_pricing_snapshot(current_subject: str = Depends(get_current_subjec
async def list_provider_configs(current_subject: str = Depends(get_current_subject)):
"""List all saved provider configurations."""
rows = providers_db.list_providers()
return [
ProviderResponse(
id = row["id"],
provider_type = row["provider_type"],
display_name = row["display_name"],
base_url = row["base_url"],
is_enabled = bool(row["is_enabled"]),
created_at = row["created_at"],
updated_at = row["updated_at"],
)
for row in rows
]
return [_provider_response(row) for row in rows]
@router.post("/", response_model = ProviderResponse, status_code = 201)
@ -124,18 +127,12 @@ async def create_provider_config(
provider_type = payload.provider_type,
display_name = payload.display_name,
base_url = base_url,
models = payload.models,
available_models = payload.available_models,
)
row = providers_db.get_provider(provider_id)
return ProviderResponse(
id = row["id"],
provider_type = row["provider_type"],
display_name = row["display_name"],
base_url = row["base_url"],
is_enabled = bool(row["is_enabled"]),
created_at = row["created_at"],
updated_at = row["updated_at"],
)
return _provider_response(row)
@router.put("/{provider_id}", response_model = ProviderResponse)
@ -154,20 +151,14 @@ async def update_provider_config(
display_name = payload.display_name,
base_url = payload.base_url,
is_enabled = payload.is_enabled,
models = payload.models,
available_models = payload.available_models,
)
if not updated:
raise HTTPException(status_code = 400, detail = "No fields to update")
row = providers_db.get_provider(provider_id)
return ProviderResponse(
id = row["id"],
provider_type = row["provider_type"],
display_name = row["display_name"],
base_url = row["base_url"],
is_enabled = bool(row["is_enabled"]),
created_at = row["created_at"],
updated_at = row["updated_at"],
)
return _provider_response(row)
@router.delete("/{provider_id}", status_code = 204)

View file

@ -225,24 +225,24 @@ def can_load_chat_during_training(
max_seq_length: int,
requested_gpu_ids: Optional[List[int]],
is_gguf: bool = False,
is_vulkan: bool = False,
required_override_gb: Optional[float] = None,
single_device_gpu: Optional[str] = None,
worst_case_gpu_count: Optional[int] = None,
) -> Tuple[bool, Dict[str, Any]]:
"""Decide if a NEW chat model can load without OOMing active training (inverse
of can_keep_chat_during_training: training is already resident, so size the
chat model against the free VRAM that remains). Sizes/places it the same way
the loader will: HF auto reuses auto_select_gpu_ids; HF explicit requires an
even-share per-GPU floor for device_map="balanced"; GGUF sizes from
required_override_gb over the visible pool. ``single_device_gpu`` is the
exact physical device token selected by a single-device runner.
``worst_case_gpu_count`` sizes against the N most-constrained visible cards
without a physical index: a Vulkan-build pick selects that many ggml
ordinals whose physical mapping is unknown, so assume the load lands on the
busiest N rather than the whole pool.
`load_in_4bit` must be effective (LoRA can flip 4-bit -> 16-bit). Non-CUDA
allows the load; default-deny on any CUDA case it can't size, so a load never
OOMs training."""
required_override_gb over the visible pool. A Vulkan GGUF selection picks by ggml
Vulkan ordinal (separate index space from CUDA ids), so its requested_gpu_ids is
NOT resolved against the CUDA set (which would raise -> invalid_gpu_ids -> bypass
the OOM check); conservatively size an N-device request against the least-free
N visible GPUs instead.
``single_device_gpu`` is the exact physical device token selected by a
single-device runner. `load_in_4bit` must be effective (LoRA can flip 4-bit
-> 16-bit). Non-CUDA allows the load; default-deny on any CUDA case it can't
size, so a load never OOMs training."""
try:
from utils.hardware import (
DeviceType,
@ -263,6 +263,11 @@ def can_load_chat_during_training(
max_seq_length = max_seq_length or 2048,
)
# A Vulkan GGUF selection uses ggml Vulkan ordinals, not CUDA physical ids;
# size it against the full visible pool (GGUF self-placement) rather than
# resolving ordinals against the CUDA parent-visible set.
vulkan_gguf = is_gguf and is_vulkan
# HF auto: reuse the loader's selector; fits iff its pick clears the margin.
if not requested_gpu_ids and not is_gguf:
_selected, meta = auto_select_gpu_ids(model_name, **est_kwargs)
@ -288,7 +293,9 @@ def can_load_chat_during_training(
}
# Explicit GPUs, or GGUF: size directly and check live free VRAM.
if single_device_gpu is not None:
if requested_gpu_ids and vulkan_gguf:
mode = "gguf_vulkan"
elif single_device_gpu is not None:
mode = "single_device"
elif is_gguf:
mode = "gguf"
@ -301,7 +308,17 @@ def can_load_chat_during_training(
return False, {"mode": mode, "reason": "estimate_unavailable"}
free_by_index = _free_vram_by_index(get_visible_gpu_utilization().get("devices", []))
if single_device_gpu is not None:
if requested_gpu_ids and vulkan_gguf:
# Vulkan ordinals cannot be mapped to CUDA physical indices. Budget
# the least-free N visible cards for an N-device request. If that
# conservative subset fits, any physical mapping of the ordinals
# fits, without collapsing a multi-GPU request to one card.
visible_free = list(free_by_index.values())
if not visible_free:
return False, {"mode": "gguf_vulkan", "reason": "no_visible_gpus"}
n_pins = min(len(requested_gpu_ids), len(visible_free))
free_vals = sorted(visible_free)[:n_pins]
elif single_device_gpu is not None:
token = str(single_device_gpu).strip()
if not token:
# Empty token = a CPU-only single-device runner (e.g. a CPU
@ -328,15 +345,9 @@ def can_load_chat_during_training(
except ValueError:
return True, {"mode": mode, "reason": "invalid_gpu_ids"}
free_vals = [free_by_index.get(i, 0.0) for i in resolved]
elif worst_case_gpu_count:
# Vulkan-build pick: the ggml ordinals have no physical mapping, so
# size against the N most-constrained visible cards (worst case)
# rather than the whole pool -- else a pool that fits could still OK
# a load that lands on a busy selected card and OOMs training.
ranked_asc = sorted(free_by_index.values())
free_vals = ranked_asc[: max(1, worst_case_gpu_count)]
else:
# GGUF: llama.cpp picks the GPU(s); any visible GPU is a candidate.
# GGUF self-placement / auto Vulkan (no requested ids): llama.cpp picks
# the GPU(s), so any visible GPU is a candidate -> size the whole pool.
free_vals = list(free_by_index.values())
if not free_vals:

View file

@ -6,8 +6,12 @@
Same pattern as studio_db.py (module-level functions, raw sqlite3, WAL,
per-function connections). API keys are NOT stored here: they live only in
the browser (localStorage) and are sent encrypted per-request.
Enabled model selections and discovered catalog IDs are stored server-side so
remote Studio clients see the same connection state (#7281).
"""
import json
import logging
import sqlite3
import threading
@ -22,6 +26,33 @@ _schema_lock = threading.Lock()
_schema_ready = False
def _encode_models_json(models: Optional[list[str]]) -> str:
if not models:
return "[]"
return json.dumps([str(model).strip() for model in models if str(model).strip()])
def _decode_models_json(raw: Optional[str]) -> list[str]:
if not raw:
return []
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
return []
if not isinstance(parsed, list):
return []
return [str(model).strip() for model in parsed if str(model).strip()]
def _row_models(row: sqlite3.Row) -> tuple[list[str], list[str]]:
return (
_decode_models_json(row["models_json"] if "models_json" in row.keys() else None),
_decode_models_json(
row["available_models_json"] if "available_models_json" in row.keys() else None
),
)
def _ensure_schema(conn: sqlite3.Connection) -> None:
"""Create the llm_providers table if absent. Called once per process."""
conn.execute("PRAGMA journal_mode=WAL")
@ -38,6 +69,13 @@ def _ensure_schema(conn: sqlite3.Connection) -> None:
)
"""
)
existing_cols = {row[1] for row in conn.execute("PRAGMA table_info(llm_providers)").fetchall()}
if "models_json" not in existing_cols:
conn.execute("ALTER TABLE llm_providers ADD COLUMN models_json TEXT NOT NULL DEFAULT '[]'")
if "available_models_json" not in existing_cols:
conn.execute(
"ALTER TABLE llm_providers ADD COLUMN available_models_json TEXT NOT NULL DEFAULT '[]'"
)
def get_connection() -> sqlite3.Connection:
@ -59,17 +97,37 @@ def get_connection() -> sqlite3.Connection:
return conn
def create_provider(id: str, provider_type: str, display_name: str, base_url: str) -> None:
def create_provider(
id: str,
provider_type: str,
display_name: str,
base_url: str,
models: Optional[list[str]] = None,
available_models: Optional[list[str]] = None,
) -> None:
"""Insert a new provider configuration."""
now = datetime.now(timezone.utc).isoformat()
conn = get_connection()
try:
conn.execute(
"""
INSERT INTO llm_providers (id, provider_type, display_name, base_url, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?)
INSERT INTO llm_providers (
id, provider_type, display_name, base_url,
models_json, available_models_json,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(id, provider_type, display_name, base_url, now, now),
(
id,
provider_type,
display_name,
base_url,
_encode_models_json(models),
_encode_models_json(available_models),
now,
now,
),
)
conn.commit()
finally:
@ -81,6 +139,8 @@ def update_provider(
display_name: Optional[str] = None,
base_url: Optional[str] = None,
is_enabled: Optional[bool] = None,
models: Optional[list[str]] = None,
available_models: Optional[list[str]] = None,
) -> bool:
"""Update fields on an existing provider. Returns True if a row was updated."""
updates = []
@ -94,6 +154,12 @@ def update_provider(
if is_enabled is not None:
updates.append("is_enabled = ?")
params.append(1 if is_enabled else 0)
if models is not None:
updates.append("models_json = ?")
params.append(_encode_models_json(models))
if available_models is not None:
updates.append("available_models_json = ?")
params.append(_encode_models_json(available_models))
if not updates:
return False
updates.append("updated_at = ?")
@ -128,7 +194,13 @@ def get_provider(id: str) -> Optional[dict]:
conn = get_connection()
try:
row = conn.execute("SELECT * FROM llm_providers WHERE id = ?", (id,)).fetchone()
return dict(row) if row else None
if not row:
return None
data = dict(row)
models, available_models = _row_models(row)
data["models"] = models
data["available_models"] = available_models
return data
finally:
conn.close()
@ -138,6 +210,13 @@ def list_providers() -> list[dict]:
conn = get_connection()
try:
rows = conn.execute("SELECT * FROM llm_providers ORDER BY created_at").fetchall()
return [dict(row) for row in rows]
providers: list[dict] = []
for row in rows:
data = dict(row)
models, available_models = _row_models(row)
data["models"] = models
data["available_models"] = available_models
providers.append(data)
return providers
finally:
conn.close()

View file

@ -0,0 +1,90 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Audio (TTS) generation applies recommended sampling + operator pins, like chat.
Regression guard for the fix that moved the sampling fill ahead of the audio generators: a
prior version resolved sampling only after the audio branches returned, so `unsloth run
--temperature` (UNSLOTH_SAMPLING_*) and per-model recommendations never reached audio
generation. These exercise the transformers TTS path of ``generate_audio`` (the direct
``/audio/generate`` route, which the chat-completions audio branches also delegate to).
"""
import asyncio
import pytest
import routes.inference as inference_route
from models.inference import ChatCompletionRequest
from utils.inference import inference_config as ic
class _FakeLlama:
# is_loaded False forces the transformers (non-GGUF) TTS branch in generate_audio.
is_loaded = False
_is_audio = False
class _FakeTransformersBackend:
def __init__(self):
self.active_model_name = "some/custom-tts"
self.models = {"some/custom-tts": {"is_audio": True}}
self.captured = {}
def generate_audio_response(self, **kwargs):
self.captured.update(kwargs)
return (b"RIFFfake", 24000)
@pytest.fixture(autouse = True)
def _isolate(monkeypatch):
ic._recommended_sampling.cache_clear()
for field in ic.SAMPLING_FIELD_NAMES:
monkeypatch.delenv(ic._SAMPLING_FIELDS[field][0], raising = False)
yield
ic._recommended_sampling.cache_clear()
def _run_generate_audio(
monkeypatch,
*,
recommended = None,
temperature = None,
):
backend = _FakeTransformersBackend()
monkeypatch.setattr(inference_route, "get_llama_cpp_backend", lambda: _FakeLlama())
monkeypatch.setattr(inference_route, "get_inference_backend", lambda: backend)
async def _noop_switch(*a, **k):
return None
monkeypatch.setattr(inference_route, "_maybe_auto_switch_model", _noop_switch)
# Recommendation source == the Chat UI's .inference block.
monkeypatch.setattr(ic, "load_inference_config", lambda mid: dict(recommended or {}))
ic._recommended_sampling.cache_clear()
kwargs = {"model": "some/custom-tts", "messages": [{"role": "user", "content": "hi"}]}
if temperature is not None:
kwargs["temperature"] = temperature
payload = ChatCompletionRequest(**kwargs)
asyncio.run(inference_route.generate_audio(payload, request = None, current_subject = "t"))
return backend.captured
def test_audio_uses_recommended_sampling_when_omitted(monkeypatch):
captured = _run_generate_audio(monkeypatch, recommended = {"temperature": 1.0, "top_k": 64})
assert captured["temperature"] == 1.0
assert captured["top_k"] == 64
def test_audio_operator_pin_overrides_client(monkeypatch):
monkeypatch.setenv("UNSLOTH_SAMPLING_TEMPERATURE", "0.9")
captured = _run_generate_audio(monkeypatch, recommended = {"temperature": 1.0}, temperature = 0.2)
assert captured["temperature"] == 0.9 # operator pin wins even over an explicit client value
def test_audio_client_explicit_preserved(monkeypatch):
captured = _run_generate_audio(monkeypatch, recommended = {"temperature": 1.0}, temperature = 0.2)
assert captured["temperature"] == 0.2 # explicit client value preserved over recommendation

View file

@ -170,6 +170,7 @@ class TestCanLoadGGUF(_GpuCacheResetMixin, unittest.TestCase):
estimate = None,
single_device_gpu = None,
gpu_ids = None,
is_vulkan = False,
):
with (
patch("utils.hardware.get_device", return_value = DeviceType.CUDA),
@ -185,6 +186,7 @@ class TestCanLoadGGUF(_GpuCacheResetMixin, unittest.TestCase):
max_seq_length = 0,
requested_gpu_ids = gpu_ids,
is_gguf = True,
is_vulkan = is_vulkan,
required_override_gb = required_override,
single_device_gpu = single_device_gpu,
)
@ -234,6 +236,35 @@ class TestCanLoadGGUF(_GpuCacheResetMixin, unittest.TestCase):
self.assertFalse(blocked)
self.assertEqual(blocked_info["usable_gb"], 10.0)
def test_vulkan_pin_takes_precedence_over_unknown_diffusion_fallback(self):
# An uncached GGUF can carry a speculative single-device fallback while
# its explicit pin is actually a ggml Vulkan ordinal. Never interpret
# that ordinal as the same-numbered CUDA physical device.
ok, info, _ = self._run(
devices = _devices((0, 80, 0), (1, 80, 78)),
required_override = 20.0,
single_device_gpu = "0",
gpu_ids = [0],
is_vulkan = True,
)
self.assertFalse(ok)
self.assertEqual(info["mode"], "gguf_vulkan")
self.assertEqual(info["usable_gb"], 2.0)
def test_vulkan_multi_gpu_guard_counts_requested_devices(self):
# The ordinal mapping is unknown, so use the least-free two visible
# cards for a two-device request. Their aggregate capacity is still
# available instead of collapsing the request to one card.
ok, info, _ = self._run(
devices = _devices((0, 80, 70), (1, 80, 70), (2, 80, 0)),
required_override = 10.0,
gpu_ids = [0, 1],
is_vulkan = True,
)
self.assertTrue(ok)
self.assertEqual(info["mode"], "gguf_vulkan")
self.assertEqual(info["usable_gb"], 18.5)
def test_single_device_unresolved_token_sizes_against_worst_device(self):
# A non-numeric device token (a CUDA UUID / MIG handle) can't map to a
# free-VRAM index. The runner still drives ONE device, so size against the
@ -421,15 +452,16 @@ class TestChatLoadGuardRoute(unittest.TestCase):
def test_allows_when_fits(self):
self._guard(training_active = True, decision = (True, {"mode": "auto"}))
def test_vulkan_build_drops_gpu_ids_before_physical_sizing(self):
# A Vulkan build's gpu_ids are ggml Vulkan ordinals; forwarding them to
# can_load_chat_during_training would resolve them in physical index
# space (the wrong card's free-VRAM row, or a ValueError that skips the
# protection entirely). The guard must size such loads as unpinned.
def test_vulkan_build_flags_gpu_ids_for_ordinal_sizing(self):
# A Vulkan build's gpu_ids are ggml Vulkan ordinals. The guard must pass
# is_vulkan so the sizer treats the pick as an N-device request in ggml
# ordinal space (worst-case least-free N cards) instead of resolving the
# ordinals as physical ids, and must NOT derive a single-device physical
# fallback for the unclassified case.
captured = []
config = SimpleNamespace(is_gguf = True, is_lora = False, path = None)
with (
patch.object(self.route, "_classify_diffusion_gguf", lambda c: False),
patch.object(self.route, "_classify_diffusion_gguf", lambda c: None),
patch.object(self.route, "_estimate_gguf_required_gb", lambda *a, **k: 2.0),
patch.object(
self.route.LlamaCppBackend,
@ -441,10 +473,12 @@ class TestChatLoadGuardRoute(unittest.TestCase):
config = config,
captured = captured,
training_active = True,
decision = (True, {"mode": "gguf"}),
decision = (True, {"mode": "gguf_vulkan"}),
requested_gpu_ids = [1],
)
self.assertIsNone(captured[0]["requested_gpu_ids"])
self.assertEqual(captured[0]["requested_gpu_ids"], [1])
self.assertTrue(captured[0]["is_vulkan"])
self.assertIsNone(captured[0]["single_device_gpu"])
def test_non_vulkan_build_keeps_gpu_ids_for_sizing(self):
captured = []
@ -524,58 +558,19 @@ class TestChatLoadGuardRoute(unittest.TestCase):
def test_manual_known_normal_gguf_bypasses_training_estimate(self):
captured = []
config = SimpleNamespace(is_gguf = True)
with patch.object(self.route, "_classify_diffusion_gguf", return_value = False):
with patch.object(self.route, "_classify_diffusion_gguf", return_value = False) as classify:
self._guard(
config = config,
captured = captured,
training_active = True,
decision = (False, {"reason": "must not run"}),
gpu_memory_mode = "manual",
requested_gpu_ids = [1, 3],
)
classify.assert_called_once_with(config)
self.assertEqual(captured, [])
def test_manual_unknown_gguf_keeps_single_device_training_guard(self):
captured = []
config = SimpleNamespace(is_gguf = True)
with (
patch.object(self.route, "_classify_diffusion_gguf", return_value = None),
patch.object(self.route, "_estimate_gguf_required_gb", return_value = 12.5),
patch.object(
self.route.LlamaCppBackend,
"_diffusion_gpu_arg",
return_value = "2",
),
):
self._guard(
config = config,
captured = captured,
training_active = True,
decision = (True, {"mode": "single_device"}),
gpu_memory_mode = "manual",
)
self.assertEqual(len(captured), 1)
self.assertEqual(captured[0]["single_device_gpu"], "2")
def test_manual_diffusion_uses_single_device_guard(self):
captured = []
config = SimpleNamespace(is_gguf = True)
with (
patch.object(self.route, "_classify_diffusion_gguf", return_value = True),
patch.object(self.route, "_estimate_gguf_required_gb", return_value = 12.5),
):
self._guard(
config = config,
captured = captured,
training_active = True,
decision = (True, {"mode": "gguf"}),
gpu_memory_mode = "manual",
requested_gpu_ids = [3, 1],
)
self.assertEqual(len(captured), 1)
self.assertEqual(captured[0]["single_device_gpu"], "1")
self.assertEqual(captured[0]["requested_gpu_ids"], [3, 1])
def test_unpinned_diffusion_uses_runner_default_gpu(self):
def test_manual_diffusion_keeps_single_device_training_guard(self):
captured = []
config = SimpleNamespace(is_gguf = True)
with (
@ -586,11 +581,6 @@ class TestChatLoadGuardRoute(unittest.TestCase):
"_effective_gpu_count",
return_value = 2,
),
patch.object(
self.route.LlamaCppBackend,
"_diffusion_gpu_arg",
return_value = "3",
) as gpu_arg,
):
self._guard(
config = config,
@ -598,9 +588,11 @@ class TestChatLoadGuardRoute(unittest.TestCase):
training_active = True,
decision = (True, {"mode": "single_device"}),
gpu_memory_mode = "manual",
requested_gpu_ids = [3, 1],
)
gpu_arg.assert_called_once_with(None, cpu_only = False)
self.assertEqual(captured[0]["single_device_gpu"], "3")
self.assertEqual(len(captured), 1)
self.assertEqual(captured[0]["single_device_gpu"], "1")
self.assertEqual(captured[0]["requested_gpu_ids"], [3, 1])
def test_refuses_with_headroom_number(self):
info = {"required_gb": 30.0, "usable_gb": 6.0, "needed_gb": 39.0, "mode": "auto"}

View file

@ -591,10 +591,23 @@ def test_load_request_accepts_gpu_ids():
@pytest.mark.parametrize("model_cls", [LoadResponse, InferenceStatusResponse])
def test_response_models_emit_gpu_ids(model_cls):
if model_cls is LoadResponse:
obj = model_cls(status = "loaded", model = "m", display_name = "m", inference = {}, gpu_ids = [1])
obj = model_cls(
status = "loaded",
model = "m",
display_name = "m",
inference = {},
gpu_ids = [1],
requested_gpu_ids = [1, 2],
)
else:
obj = model_cls(gpu_ids = [1])
obj = model_cls(gpu_ids = [1], requested_gpu_ids = [1, 2])
assert obj.model_dump()["gpu_ids"] == [1]
assert obj.model_dump()["requested_gpu_ids"] == [1, 2]
def test_gguf_load_and_status_responses_include_requested_gpu_pool():
route_src = (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
assert route_src.count("requested_gpu_ids = llama_backend.requested_gpu_ids") == 3
def test_gpu_ids_property_default_and_reset():
@ -625,6 +638,10 @@ def _target_state_gpu_ids(backend, gpu_ids):
def test_gpu_ids_reload_detection_is_order_insensitive():
backend = _loaded_backend("auto")
backend._gpu_ids = [0, 1]
# A real non-narrowed load records the raw request too; the non-diffusion
# dedupe now compares that raw pin (#7239). Set it to match the effective pin
# (no narrowing) so this exercises the order-insensitive comparison.
backend._requested_gpu_ids = [0, 1]
# Same set, different order -> no reload.
assert _target_state_gpu_ids(backend, [1, 0]) is True
# Different set -> reload.
@ -633,6 +650,26 @@ def test_gpu_ids_reload_detection_is_order_insensitive():
assert _target_state_gpu_ids(backend, None) is False
def test_gpu_ids_reload_detection_accepts_raw_and_effective_pin():
backend = _loaded_backend("auto")
backend._requested_gpu_ids = [0, 1]
backend._gpu_ids = [0]
backend._last_load_kwargs = {"gpu_ids": [0, 1], "model_identifier": "owner/repo"}
# The original request still matches after the fitter narrows it.
assert _target_state_gpu_ids(backend, [1, 0]) is True
assert backend.requested_gpu_ids == [0, 1]
# The status response echoes the effective pin, which must also round-trip.
# Treat the incoming subset as the latest intent so status and a future
# reload do not restore GPU 1 after the user removed it.
assert _target_state_gpu_ids(backend, [0]) is True
assert backend.requested_gpu_ids == [0]
assert backend._last_load_kwargs == {"gpu_ids": [0], "model_identifier": "owner/repo"}
# A genuinely different placement pool still reloads.
assert _target_state_gpu_ids(backend, [1]) is False
assert _target_state_gpu_ids(backend, None) is False
def test_gpu_ids_reload_detection_collapses_diffusion_to_single_device():
# The diffusion runner drives only its single lowest device, so the backend
# records [lowest]. A later multi-GPU request that still resolves to that
@ -642,6 +679,7 @@ def test_gpu_ids_reload_detection_collapses_diffusion_to_single_device():
backend._is_diffusion = True
backend._gpu_ids = [1] # loaded on the lowest of an earlier [3, 1] pick
assert _target_state_gpu_ids(backend, [3, 1]) is True
assert backend.requested_gpu_ids == [1]
assert _target_state_gpu_ids(backend, [1]) is True
# Lowest device changes (2, not 1) -> reload.
assert _target_state_gpu_ids(backend, [3, 2]) is False
@ -649,6 +687,56 @@ def test_gpu_ids_reload_detection_collapses_diffusion_to_single_device():
assert _target_state_gpu_ids(backend, None) is False
def test_remote_vulkan_diffusion_preflight_runs_before_teardown(monkeypatch):
def _mark_diffusion(probe, path):
assert path == "/cache/model.gguf"
probe._is_diffusion = True
monkeypatch.setattr(LlamaCppBackend, "_read_gguf_metadata", _mark_diffusion)
assert LlamaCppBackend._gguf_path_is_diffusion("/cache/model.gguf", "owner/model") is True
src = inspect.getsource(llama_cpp_module.LlamaCppBackend.load_model)
preflight = src.index("_preflight_model_path = self._download_gguf(")
teardown = src.index("# ── Phase 1: kill old process")
assert preflight < teardown
assert "model_path = _preflight_model_path or self._download_gguf(" in src
def test_remote_vulkan_diffusion_rejection_keeps_active_server(monkeypatch):
backend = LlamaCppBackend()
killed = []
monkeypatch.setattr(backend, "_find_llama_server_binary", lambda **_kwargs: "/bin/llama")
monkeypatch.setattr(backend, "_is_vulkan_backend", lambda _binary = None: True)
monkeypatch.setattr(backend, "_get_gpu_memory", lambda _binary = None: [(0, 1024, 2048)])
monkeypatch.setattr(
backend,
"_download_gguf",
lambda **_kwargs: "/cache/diffusion.gguf",
)
monkeypatch.setattr(backend, "_gguf_path_is_diffusion", lambda *_args: True)
monkeypatch.setattr(backend, "_kill_process", lambda: killed.append(True))
monkeypatch.setattr(
llama_cpp_module,
"_resolve_repo_id_casing",
lambda repo: repo,
)
monkeypatch.setattr(
llama_cpp_module,
"_hf_offline_if_dns_dead",
lambda: __import__("contextlib").nullcontext(),
)
with pytest.raises(ValueError, match = "DiffusionGemma"):
backend.load_model(
hf_repo = "owner/model",
hf_variant = "Q4_K_M",
model_identifier = "owner/model",
gpu_ids = [0],
)
assert killed == []
def test_start_diffusion_server_resets_tensor_parallel():
# A prior tensor-parallel chat load leaves self._tensor_parallel True (load_model
# phase 1 only kills the process, it skips the unload reset). Diffusion is never
@ -656,19 +744,19 @@ def test_start_diffusion_server_resets_tensor_parallel():
# diffusion re-Apply reloads against stale tensor-parallel state.
src = inspect.getsource(llama_cpp_module.LlamaCppBackend._start_diffusion_server)
assert "self._tensor_parallel = False" in src
assert "self._requested_gpu_ids = list(self._gpu_ids) if self._gpu_ids else None" in src
def test_vulkan_gpu_gate_allows_unclassified_gguf():
# The pre-download /load + /validate Vulkan gates must reject only a
# CONFIRMED-diffusion pick (`is True`). `None` -- the ordinary first-load
# case for an uncached Hub GGUF with no local header -- has to pass, or the
# GPU picker is unusable for first-time remote GGUF loads (Codex #7356).
# The shared GGUF gpu_ids validator must reject only a CONFIRMED-diffusion
# pick on Vulkan (`is True`). `None` -- the ordinary first-load case for an
# uncached Hub GGUF with no local header -- has to pass, or the GPU picker is
# unusable for first-time remote GGUF loads (Codex #7356).
route_src = (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
# Both Vulkan validation sites (/load, /validate) gate on `is True`...
assert route_src.count("_classify_diffusion_gguf(config) is True") == 2
# ...and no diffusion *rejection* keys off the old over-broad `is not False`
# (which also caught the unclassifiable None). The training guard keeps its
# own conservative `diffusion_kind is not False` sizing -- a different name.
assert "if is_vulkan and _classify_diffusion_gguf(config) is True:" in route_src
# No diffusion *rejection* keys off the old over-broad `is not False` (which
# also caught the unclassifiable None). The training guard keeps its own
# conservative `diffusion_kind is not False` sizing -- a different name.
assert "_classify_diffusion_gguf(config) is not False" not in route_src
@ -685,16 +773,13 @@ def test_diffusion_vulkan_load_drops_unmappable_gpu_pin():
assert guard < drop < spawn
def test_route_matches_loaded_settings_collapses_diffusion_gpu_ids():
# The route-level reload dedupe mirrors the backend: for a loaded diffusion
# model it compares the request against the single recorded device, not the
# full requested list, or a same-device multi-GPU pick reloads needlessly.
def test_route_matches_loaded_settings_uses_shared_gpu_pin_matcher():
# Route-level and backend race dedupe must share one normalization path so
# raw, effective, and diffusion pins cannot drift apart.
route_src = (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text(encoding = "utf-8")
match_impl = route_src[route_src.index("def _request_matches_loaded_settings") :]
guard = match_impl.index("if llama_backend.is_diffusion:")
collapse = match_impl.index("[sorted(request.gpu_ids)[0]] if request.gpu_ids else None")
compare = match_impl.index("if _req_gpu_ids != llama_backend.gpu_ids:")
assert guard < collapse < compare
assert "if not llama_backend.matches_gpu_ids(request.gpu_ids):" in match_impl
assert "llama_backend._record_matching_gpu_request(request.gpu_ids)" in match_impl
# ── Manual tensor split: child enumeration pinned to the picker's order ──────

View file

@ -130,6 +130,26 @@ class TestResolveRequestedGpuIds(_GpuCacheResetMixin, unittest.TestCase):
):
self.assertEqual(resolve_requested_gpu_ids([]), [1, 3])
def test_vulkan_ordinals_bypass_cuda_parent_visible_validation(self):
# Vulkan build on a CPU-only torch host: no CUDA parent-visible set and a
# zero physical count, yet a valid Vulkan ordinal must not be rejected as
# a CUDA physical id (issue #7239).
with (
patch.dict(os.environ, {}, clear = True),
patch("utils.hardware.hardware.get_physical_gpu_count", return_value = 0),
):
# As a CUDA physical id, [0] is outside the empty parent-visible set.
with self.assertRaises(ValueError):
resolve_requested_gpu_ids([0])
# As Vulkan ordinals, [0] and [0, 1] pass through unchanged.
self.assertEqual(resolve_requested_gpu_ids([0], is_vulkan = True), [0])
self.assertEqual(resolve_requested_gpu_ids([0, 1], is_vulkan = True), [0, 1])
# Malformed ordinals are still rejected.
with self.assertRaisesRegex(ValueError, "duplicate GPU IDs"):
resolve_requested_gpu_ids([0, 0], is_vulkan = True)
with self.assertRaisesRegex(ValueError, "non-negative"):
resolve_requested_gpu_ids([-1], is_vulkan = True)
def test_apply_gpu_ids_only_updates_cuda_visible_devices(self):
with patch.dict(
os.environ,
@ -853,6 +873,171 @@ class TestRouteErrors(unittest.TestCase):
self.assertIn("only supported on CUDA devices", str(exc_info.exception))
def test_inference_route_resolves_gguf_gpu_ids(self):
# GGUF gpu_ids are now supported: /load routes them through the same
# resolution as non-GGUF loads (rejecting only genuinely invalid ids with
# the resolver's actionable message) rather than a blanket "not supported"
# reject, so /validate can stay consistent with /load (#7239).
import utils.hardware.hardware as hardware_mod
inference_route = _load_route_module(
"inference_route_module_for_gguf_gpu_ids_test",
"routes/inference.py",
)
request = LoadRequest(model_path = "unsloth/test.gguf", gpu_ids = [0, 1])
model_config = SimpleNamespace(
is_gguf = True,
is_lora = False,
gguf_hf_repo = None,
gguf_file = "/tmp/test.gguf",
gguf_mmproj_file = None,
gguf_variant = None,
identifier = "unsloth/test.gguf",
display_name = "unsloth/test.gguf",
is_vision = False,
is_audio = False,
audio_type = None,
has_audio_input = False,
)
def _fake_resolve(ids, is_vulkan = False):
raise ValueError("SENTINEL requested GPUs are outside the parent-visible set")
with (
patch.object(
inference_route,
"ModelConfig",
SimpleNamespace(from_identifier = lambda **_kwargs: model_config),
),
# Patch both the package re-export and the defining module so the stub
# fires no matter which import path the route uses.
patch("utils.hardware.resolve_requested_gpu_ids", _fake_resolve),
patch.object(hardware_mod, "resolve_requested_gpu_ids", _fake_resolve),
patch.object(
inference_route,
"_guard_chat_load_against_training",
return_value = None,
),
patch.object(inference_route.asyncio, "to_thread", new = _inline_to_thread),
patch.object(inference_route, "_hf_offline_if_dns_dead", nullcontext),
):
with self.assertRaises(HTTPException) as exc_info:
asyncio.run(
inference_route._load_model_impl(
request,
SimpleNamespace(
app = SimpleNamespace(
state = SimpleNamespace(llama_parallel_slots = 1),
),
),
current_subject = "test-user",
)
)
# The selection was routed through resolution (not the old blanket reject).
self.assertEqual(exc_info.exception.status_code, 400)
self.assertIn("SENTINEL", exc_info.exception.detail)
self.assertNotIn("not supported for GGUF", exc_info.exception.detail)
def test_load_rejects_unavailable_vulkan_ordinal_before_training_guard(self):
inference_route = _load_route_module(
"inference_route_module_for_vulkan_preflight_test",
"routes/inference.py",
)
request = LoadRequest(model_path = "unsloth/test.gguf", gpu_ids = [99])
model_config = SimpleNamespace(
is_gguf = True,
is_lora = False,
gguf_hf_repo = None,
gguf_file = "/tmp/test.gguf",
gguf_mmproj_file = None,
gguf_variant = None,
identifier = "unsloth/test.gguf",
display_name = "unsloth/test.gguf",
is_vision = False,
is_audio = False,
audio_type = None,
has_audio_input = False,
)
with (
patch.object(
inference_route,
"ModelConfig",
SimpleNamespace(from_identifier = lambda **_kwargs: model_config),
),
patch("utils.hardware.get_device", return_value = DeviceType.CUDA),
patch.object(inference_route, "_classify_diffusion_gguf", return_value = None),
patch.object(
inference_route.LlamaCppBackend,
"_is_vulkan_backend",
return_value = True,
),
patch.object(
inference_route.LlamaCppBackend,
"_find_llama_server_binary",
return_value = "/tmp/llama-server",
),
patch.object(
inference_route.LlamaCppBackend,
"_get_gpu_memory",
return_value = [(0, 8 * 1024**3, 16 * 1024**3)],
),
patch.object(
inference_route,
"_guard_chat_load_against_training",
return_value = None,
) as training_guard,
patch.object(inference_route.asyncio, "to_thread", new = _inline_to_thread),
patch.object(inference_route, "_hf_offline_if_dns_dead", nullcontext),
):
with self.assertRaises(HTTPException) as exc_info:
asyncio.run(
inference_route._load_model_impl(
request,
SimpleNamespace(
app = SimpleNamespace(
state = SimpleNamespace(llama_parallel_slots = 1),
),
),
current_subject = "test-user",
)
)
self.assertEqual(exc_info.exception.status_code, 400)
self.assertIn("Vulkan GPU ordinal(s) [99]", exc_info.exception.detail)
training_guard.assert_not_called()
def test_vulkan_ordinals_are_allowed_on_xpu_hosts(self):
import utils.hardware.hardware as hardware_mod
inference_route = _load_route_module(
"inference_route_module_for_xpu_vulkan_test",
"routes/inference.py",
)
config = SimpleNamespace(is_gguf = True)
with (
patch("utils.hardware.get_device", return_value = DeviceType.XPU),
patch.object(
inference_route.LlamaCppBackend,
"_is_vulkan_backend",
return_value = True,
),
patch.object(inference_route, "_classify_diffusion_gguf", return_value = False),
patch.object(hardware_mod, "resolve_requested_gpu_ids", return_value = [0, 1]),
patch.object(
inference_route.LlamaCppBackend,
"_find_llama_server_binary",
return_value = None,
),
):
resolved = asyncio.run(
inference_route._resolve_gguf_gpu_ids_for_request(config, [1, 0])
)
self.assertEqual(resolved, [0, 1])
def test_inference_route_validates_gpu_ids_for_gguf(self):
# gpu_ids is now SUPPORTED for GGUF (the GPU picker), but still
# validated: a rejected pick surfaces as a clean 400, not the old
@ -861,7 +1046,7 @@ class TestRouteErrors(unittest.TestCase):
import utils.hardware.hardware as hardware_mod
inference_route = _load_route_module(
"inference_route_module_for_gguf_gpu_ids_test",
"inference_route_module_for_gguf_gpu_ids_test2",
"routes/inference.py",
)
request = LoadRequest(model_path = "unsloth/test.gguf", gpu_ids = [0, 1])
@ -886,6 +1071,17 @@ class TestRouteErrors(unittest.TestCase):
"ModelConfig",
SimpleNamespace(from_identifier = lambda **_kwargs: model_config),
),
# Patch both the package re-export and the defining module so the stub
# fires no matter which import path the route uses.
patch(
"utils.hardware.resolve_requested_gpu_ids",
side_effect = ValueError("Invalid gpu_ids [0, 1]: rejected by test"),
),
patch.object(
hardware_mod,
"resolve_requested_gpu_ids",
side_effect = ValueError("Invalid gpu_ids [0, 1]: rejected by test"),
),
patch.object(
inference_route,
"_guard_chat_load_against_training",
@ -893,11 +1089,6 @@ class TestRouteErrors(unittest.TestCase):
),
patch.object(inference_route.asyncio, "to_thread", new = _inline_to_thread),
patch.object(inference_route, "_hf_offline_if_dns_dead", nullcontext),
patch.object(
hardware_mod,
"resolve_requested_gpu_ids",
side_effect = ValueError("Invalid gpu_ids [0, 1]: rejected by test"),
),
):
with self.assertRaises(HTTPException) as exc_info:
asyncio.run(

View file

@ -51,6 +51,27 @@ def _modules_json(*paths):
_COMMIT = "0123456789abcdef0123456789abcdef01234567"
def _fs_case_sensitive(root):
"""Whether root's filesystem is case-sensitive (Linux yes; macOS/Windows usually no). The gate
mirrors the loader, whose file lookups follow the same rule, so some cases only exist on one."""
probe = Path(root) / "_case_probe"
probe.write_text("x")
try:
return not (Path(root) / "_CASE_PROBE").exists()
finally:
probe.unlink()
def _requires_case_sensitive_fs(root):
if not _fs_case_sensitive(root):
pytest.skip("requires a case-sensitive filesystem")
def _requires_case_insensitive_fs(root):
if _fs_case_sensitive(root):
pytest.skip("requires a case-insensitive filesystem")
def _make_cache(
root,
repo_id,
@ -382,6 +403,329 @@ def test_gate_blocks_sharded_pickle(hf_cache):
assert _offline_decision("org/shard").blocked is True
def test_gate_blocks_indexed_pickle_shard_in_subdirectory(hf_cache):
# from_pretrained follows weight_map paths relative to the root index, so these nested shards
# are deserialized even though they are not direct children of the load root (iterdir misses
# them). The online gate blocks index-referenced subdir pickles; the offline gate must too.
_make_cache(
hf_cache,
"org/indexed-shard",
{
"pytorch_model.bin.index.json": (
'{"weight_map": {"layer.weight": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"shards/pytorch_model-00001-of-00001.bin": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/indexed-shard")
assert decision.blocked is True
assert any(
u["path"] == "shards/pytorch_model-00001-of-00001.bin" for u in decision.unsafe_files
)
def test_gate_blocks_indexed_pickle_shard_with_nonstandard_stem(hf_cache):
# The index tells the loader to deserialize this file, so a pickle EXTENSION is enough -- the
# shard's stem need not match the on-disk weight-name heuristic (which only guesses bare files).
_make_cache(
hf_cache,
"org/indexed-odd",
{
"pytorch_model.bin.index.json": '{"weight_map": {"w": "shards/evil-00001-of-00001.bin"}}',
"shards/evil-00001-of-00001.bin": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/indexed-odd")
assert decision.blocked is True
assert any(u["path"] == "shards/evil-00001-of-00001.bin" for u in decision.unsafe_files)
def test_gate_blocks_safetensors_index_pointing_to_pickle_shard(hf_cache):
# load_state_dict picks safetensors vs torch.load by each shard's own suffix, so a
# model.safetensors.index.json that maps a weight to a .bin shard still deserializes it. The
# index's own existence must not suppress the shard it names.
_make_cache(
hf_cache,
"org/st-index-pickle",
{
"model.safetensors.index.json": (
'{"weight_map": {"w": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"shards/pytorch_model-00001-of-00001.bin": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/st-index-pickle")
assert decision.blocked is True
assert any(
u["path"] == "shards/pytorch_model-00001-of-00001.bin" for u in decision.unsafe_files
)
def test_gate_blocks_indexed_shard_with_no_pickle_extension(hf_cache):
# Transformers torch.loads any indexed shard not ending in .safetensors, so an unconventional
# extensionless name is still a deserialization target.
_make_cache(
hf_cache,
"org/indexed-noext",
{
"pytorch_model.bin.index.json": '{"weight_map": {"w": "shards/payload"}}',
"shards/payload": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/indexed-noext")
assert decision.blocked is True
assert any(u["path"] == "shards/payload" for u in decision.unsafe_files)
_UPPER_INDEX_FILES = {
"PYTORCH_MODEL.BIN.INDEX.JSON": (
'{"weight_map": {"w": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"shards/pytorch_model-00001-of-00001.bin": "pickle",
}
def test_gate_blocks_uppercase_index_on_case_insensitive_fs(hf_cache):
# On a case-insensitive volume (Windows/macOS) from_pretrained opens an oddly-cased index when it
# requests the canonical lowercase name, so the loader-mirror lookup resolves it and blocks.
_requires_case_insensitive_fs(hf_cache)
_make_cache(hf_cache, "org/upper-index", _UPPER_INDEX_FILES)
with _no_network():
decision = _offline_decision("org/upper-index")
assert decision.blocked is True
assert any(
u["path"] == "shards/pytorch_model-00001-of-00001.bin" for u in decision.unsafe_files
)
def test_gate_allows_uppercase_index_on_case_sensitive_fs(hf_cache):
# On a case-sensitive FS from_pretrained's os.path.isfile of the canonical lowercase name misses
# the uppercase artifact and never loads its shard, so the gate must not over-block it.
_requires_case_sensitive_fs(hf_cache)
_make_cache(hf_cache, "org/upper-index", _UPPER_INDEX_FILES)
with _no_network():
assert _offline_decision("org/upper-index").blocked is False
def test_gate_blocks_indexed_shard_named_with_backslash(hf_cache):
# On POSIX a backslash is a literal filename char, so from_pretrained joins the raw weight_map
# value and deserializes a file actually named "dir\payload.bin"; the gate must probe it verbatim.
import os
if os.sep != "/":
pytest.skip("backslash is a path separator off POSIX")
_make_cache(
hf_cache,
"org/backslash",
{
"pytorch_model.bin.index.json": '{"weight_map": {"w": "dir\\\\payload.bin"}}',
"dir\\payload.bin": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/backslash")
assert decision.blocked is True
assert any(u["path"] == "dir\\payload.bin" for u in decision.unsafe_files)
def test_gate_blocks_indexed_shard_with_uppercase_safetensors_suffix(hf_cache):
# load_state_dict's endswith(".safetensors") is case-sensitive, so a shard named payload.SAFETENSORS
# falls to torch.load. The gate must classify shard suffixes case-sensitively to match it.
_make_cache(
hf_cache,
"org/upper-suffix",
{
"pytorch_model.bin.index.json": '{"weight_map": {"w": "shards/payload.SAFETENSORS"}}',
"shards/payload.SAFETENSORS": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/upper-suffix")
assert decision.blocked is True
assert any(u["path"] == "shards/payload.SAFETENSORS" for u in decision.unsafe_files)
def test_gate_allows_stale_safetensors_index_beside_direct_safetensors(hf_cache):
# A complete direct model.safetensors is selected before either index, so a stale
# model.safetensors.index.json referencing a .bin shard never deserializes -> must not block.
_make_cache(
hf_cache,
"org/direct-plus-stale-index",
{
"model.safetensors": "tensors",
"model.safetensors.index.json": (
'{"weight_map": {"w": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"shards/pytorch_model-00001-of-00001.bin": "pickle",
},
)
with _no_network():
assert _offline_decision("org/direct-plus-stale-index").blocked is False
def test_gate_blocks_pytorch_index_with_uppercase_safetensors_decoy(hf_cache):
# On a case-sensitive FS, from_pretrained asks for the canonical lowercase model.safetensors, does
# not find an uppercase decoy, and selects the pytorch index instead. The decoy must not suppress.
_requires_case_sensitive_fs(hf_cache)
_make_cache(
hf_cache,
"org/upper-decoy",
{
"MODEL.SAFETENSORS": "decoy",
"pytorch_model.bin.index.json": (
'{"weight_map": {"w": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"shards/pytorch_model-00001-of-00001.bin": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/upper-decoy")
assert decision.blocked is True
assert any(
u["path"] == "shards/pytorch_model-00001-of-00001.bin" for u in decision.unsafe_files
)
def test_gate_blocks_direct_pickle_with_uppercase_safetensors_decoy(hf_cache):
# Same decoy against a direct pytorch_model.bin: the loader selects the pickle, so the uppercase
# safetensors must not suppress it on a case-sensitive FS.
_requires_case_sensitive_fs(hf_cache)
_make_cache(
hf_cache,
"org/upper-decoy-direct",
{"MODEL.SAFETENSORS": "decoy", "pytorch_model.bin": "pickle"},
)
with _no_network():
decision = _offline_decision("org/upper-decoy-direct")
assert decision.blocked is True
assert any(u["path"] == "pytorch_model.bin" for u in decision.unsafe_files)
def test_gate_blocks_indexed_pickle_shard_in_module_subdir(hf_cache):
# A weight index inside a sentence-transformers module load root points at a nested pickle shard.
_make_cache(
hf_cache,
"org/mod-indexed",
{
"modules.json": _modules_json("0_Transformer"),
"0_Transformer/pytorch_model.bin.index.json": (
'{"weight_map": {"w": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"0_Transformer/shards/pytorch_model-00001-of-00001.bin": "pickle",
},
)
with _no_network():
decision = _offline_decision("org/mod-indexed")
assert decision.blocked is True
assert any(
u["path"] == "0_Transformer/shards/pytorch_model-00001-of-00001.bin"
for u in decision.unsafe_files
)
def test_gate_allows_indexed_pickle_shard_with_safetensors_sibling(hf_cache):
# A base model.safetensors makes the loader ignore the pickle index entirely, so it must not
# block (mirrors the direct-file safetensors-sibling suppression).
_make_cache(
hf_cache,
"org/indexed-both",
{
"pytorch_model.bin.index.json": (
'{"weight_map": {"w": "shards/pytorch_model-00001-of-00001.bin"}}'
),
"shards/pytorch_model-00001-of-00001.bin": "pickle",
"model.safetensors": "y",
},
)
with _no_network():
assert _offline_decision("org/indexed-both").blocked is False
def test_gate_allows_indexed_safetensors_shard_in_subdirectory(hf_cache):
# A safetensors index lists inert shards -- following it must never block (guards against a
# scanner that flags every indexed shard regardless of format).
_make_cache(
hf_cache,
"org/st-indexed",
{
"model.safetensors.index.json": (
'{"weight_map": {"w": "shards/model-00001-of-00001.safetensors"}}'
),
"shards/model-00001-of-00001.safetensors": "tensors",
},
)
with _no_network():
assert _offline_decision("org/st-indexed").blocked is False
def test_gate_blocks_on_index_path_traversal(hf_cache):
# A weight_map entry escaping the snapshot via ".." is abnormal/hostile -> fail closed.
_make_cache(
hf_cache,
"org/escape",
{"pytorch_model.bin.index.json": '{"weight_map": {"w": "../../../../etc/evil.bin"}}'},
)
with _no_network():
assert _offline_decision("org/escape").blocked is True
def test_gate_allows_symlinked_sharded_safetensors(tmp_path, monkeypatch):
# Real HF caches store snapshot files as symlinks into blobs/. A resolve()-based containment
# check would escape the snapshot and false-block every sharded model; the lexical gate must not.
import hashlib
import os
from huggingface_hub.file_download import repo_folder_name
root = tmp_path / "hub"
root.mkdir()
monkeypatch.setenv("HF_HOME", str(tmp_path))
monkeypatch.setenv("HF_HUB_CACHE", str(root))
monkeypatch.setattr(
"utils.hf_cache_settings.get_hf_cache_paths",
lambda: SimpleNamespace(hub_cache = root),
)
repo_dir = root / repo_folder_name(repo_id = "org/sym", repo_type = "model")
(repo_dir / "refs").mkdir(parents = True)
(repo_dir / "refs" / "main").write_text(_COMMIT)
blobs = repo_dir / "blobs"
blobs.mkdir()
snapshot = repo_dir / "snapshots" / _COMMIT
(snapshot / "shards").mkdir(parents = True)
def _blobbed(rel, content):
digest = hashlib.sha256(content.encode()).hexdigest()
(blobs / digest).write_text(content)
target = snapshot / rel
target.parent.mkdir(parents = True, exist_ok = True)
target.symlink_to(os.path.relpath(blobs / digest, target.parent))
_blobbed("config.json", "{}")
_blobbed(
"model.safetensors.index.json",
'{"weight_map": {"w": "shards/model-00001-of-00001.safetensors"}}',
)
_blobbed("shards/model-00001-of-00001.safetensors", "tensors")
with _no_network():
assert _offline_decision("org/sym").blocked is False
def test_gate_allows_index_without_weight_map(hf_cache):
# An index whose top-level JSON has no dict weight_map lets the loader resolve no shards, so it
# must not crash or block on its own (only inert safetensors are cached here).
_make_cache(
hf_cache,
"org/no-wm",
{"model.safetensors.index.json": "[]", "model.safetensors": "x"},
)
with _no_network():
assert _offline_decision("org/no-wm").blocked is False
def test_gate_allows_nothing_cached(hf_cache):
with _no_network():
assert _offline_decision("org/missing").blocked is False

View file

@ -0,0 +1,70 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Unit tests for provider model persistence (unslothai/unsloth#7281)."""
from __future__ import annotations
from pathlib import Path
import pytest
import storage.providers_db as providers_db
@pytest.fixture()
def isolated_providers_db(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
db_path = tmp_path / "studio.db"
monkeypatch.setattr(providers_db, "studio_db_path", lambda: db_path)
monkeypatch.setattr(providers_db, "ensure_dir", lambda _path: None)
providers_db._schema_ready = False
yield db_path
providers_db._schema_ready = False
def test_create_and_list_provider_models(isolated_providers_db: Path):
providers_db.create_provider(
id = "ollama1",
provider_type = "ollama",
display_name = "Home Ollama",
base_url = "http://127.0.0.1:11434",
models = ["llama3.2", "qwen2.5"],
available_models = ["llama3.2", "qwen2.5", "mistral"],
)
row = providers_db.get_provider("ollama1")
assert row is not None
assert row["models"] == ["llama3.2", "qwen2.5"]
assert row["available_models"] == ["llama3.2", "qwen2.5", "mistral"]
listed = providers_db.list_providers()
assert len(listed) == 1
assert listed[0]["models"] == ["llama3.2", "qwen2.5"]
def test_update_provider_models(isolated_providers_db: Path):
providers_db.create_provider(
id = "vllm1",
provider_type = "vllm",
display_name = "Remote vLLM",
base_url = "http://studio-host:8000/v1",
models = ["meta-llama/Llama-3.2-1B-Instruct"],
available_models = ["meta-llama/Llama-3.2-1B-Instruct"],
)
assert providers_db.update_provider(
id = "vllm1",
models = ["meta-llama/Llama-3.2-3B-Instruct"],
available_models = [
"meta-llama/Llama-3.2-1B-Instruct",
"meta-llama/Llama-3.2-3B-Instruct",
],
)
row = providers_db.get_provider("vllm1")
assert row is not None
assert row["models"] == ["meta-llama/Llama-3.2-3B-Instruct"]
assert row["available_models"] == [
"meta-llama/Llama-3.2-1B-Instruct",
"meta-llama/Llama-3.2-3B-Instruct",
]

View file

@ -0,0 +1,270 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Effective sampling resolution: per-model recommendation + operator pins.
Precedence per field: operator UNSLOTH_SAMPLING_* pin -> client explicit value ->
per-model recommendation (load_inference_config) -> static schema default.
"""
import pytest
from utils.inference.inference_config import resolve_effective_sampling, SAMPLING_FIELD_NAMES
from utils.inference import inference_config as ic
_SCHEMA_DEFAULTS = {
"temperature": 0.6,
"top_p": 0.95,
"top_k": 20,
"min_p": 0.01,
"repetition_penalty": 1.0,
"presence_penalty": 0.0,
}
@pytest.fixture(autouse = True)
def _isolate(monkeypatch):
# The recommended lookup is lru-cached; clear it so a patched config takes effect.
ic._recommended_sampling.cache_clear()
for field in SAMPLING_FIELD_NAMES:
monkeypatch.delenv(ic._SAMPLING_FIELDS[field][0], raising = False)
yield
ic._recommended_sampling.cache_clear()
def _all_omitted():
return {f: None for f in SAMPLING_FIELD_NAMES}
def _set_recommended(monkeypatch, mapping):
# _recommended_sampling sources from load_inference_config -- the exact block the Chat UI
# seeds from -- so patch that directly. Fields absent from `mapping` fall to schema defaults.
monkeypatch.setattr(ic, "load_inference_config", lambda mid: dict(mapping))
ic._recommended_sampling.cache_clear()
def test_recommended_applies_when_client_omits(monkeypatch):
_set_recommended(monkeypatch, {"temperature": 1.0, "top_k": 64, "min_p": 0.0})
eff = resolve_effective_sampling("some/model", _all_omitted())
assert eff["temperature"] == 1.0
assert eff["top_k"] == 64
assert eff["min_p"] == 0.0
# A field with no recommendation keeps the static schema default.
assert eff["top_p"] == 0.95
def test_client_explicit_beats_recommended(monkeypatch):
_set_recommended(monkeypatch, {"temperature": 1.0})
eff = resolve_effective_sampling("some/model", {**_all_omitted(), "temperature": 0.2})
assert eff["temperature"] == 0.2
def test_operator_pin_beats_client_and_recommended(monkeypatch):
_set_recommended(monkeypatch, {"temperature": 1.0})
monkeypatch.setenv("UNSLOTH_SAMPLING_TEMPERATURE", "0.9")
eff = resolve_effective_sampling("some/model", {**_all_omitted(), "temperature": 0.2})
assert eff["temperature"] == 0.9
def test_unknown_model_matches_ui_inference_block(monkeypatch):
# An unknown model gets the same values the Chat UI would seed (load_inference_config's
# default.yaml fallback: temp 0.7 / top_k -1), NOT the request schema defaults.
ui_block = {
"temperature": 0.7,
"top_p": 0.95,
"top_k": -1,
"min_p": 0.01,
"presence_penalty": 0.0,
"repetition_penalty": 1.0,
}
monkeypatch.setattr(ic, "load_inference_config", lambda mid: dict(ui_block))
ic._recommended_sampling.cache_clear()
eff = resolve_effective_sampling("some/unknown-model", _all_omitted())
assert eff["temperature"] == 0.7
assert eff["top_k"] == -1
assert eff["min_p"] == 0.01
def test_empty_recommendation_falls_back_to_schema_defaults(monkeypatch):
# If load_inference_config yields nothing usable, the resolver falls back to the request
# schema defaults.
monkeypatch.setattr(ic, "load_inference_config", lambda mid: {})
ic._recommended_sampling.cache_clear()
eff = resolve_effective_sampling("some/model", _all_omitted())
assert eff == _SCHEMA_DEFAULTS
@pytest.mark.parametrize(
"model",
["unsloth/gemma-4-E4B", "unsloth/Qwen3-4B", "unsloth/Qwen3.5-9B", "someorg/unknown-xyz"],
)
def test_recommendation_matches_ui_source(model):
# Parity guard: what the server recommends for omitted fields equals the Chat UI's source
# (load_inference_config) for every field the UI adopts (mergeBackendRecommendedInference).
ic._recommended_sampling.cache_clear()
ui = ic.load_inference_config(model)
rec = ic._recommended_sampling(model)
for f in ic._UI_RECOMMENDED_FIELDS:
cleaned = ic._clean_sampling_value(f, ui.get(f))
if cleaned is not None:
assert rec.get(f) == cleaned, f"{model}:{f} rec={rec.get(f)} ui={ui.get(f)}"
def test_repetition_penalty_not_auto_recommended(monkeypatch):
# The Chat UI's mergeBackendRecommendedInference never adopts a backend repetition_penalty
# (e.g. lfm2's family value 1.05), so the server must not auto-apply one either. It stays at
# the schema default unless the client sends it or an operator pins it.
monkeypatch.setattr(
ic, "load_inference_config", lambda mid: {"temperature": 0.7, "repetition_penalty": 1.05}
)
ic._recommended_sampling.cache_clear()
eff = resolve_effective_sampling("some/lfm2-model", _all_omitted())
assert eff["temperature"] == 0.7 # a UI-adopted field is recommended
assert eff["repetition_penalty"] == 1.0 # rep is NOT auto-recommended (matches the UI)
# An operator can still pin it explicitly.
monkeypatch.setenv("UNSLOTH_SAMPLING_REPETITION_PENALTY", "1.05")
eff2 = resolve_effective_sampling("some/lfm2-model", _all_omitted())
assert eff2["repetition_penalty"] == 1.05
@pytest.mark.parametrize(
"raw, expected",
[
("0.5", 0.5),
("abc", None), # unparseable
("9.0", None), # above temperature max (2.0)
("-1", None), # below temperature min (0.0)
(" ", None), # blank
("nan", None), # NaN would pass a naive range check
("inf", None), # non-finite
("-inf", None), # non-finite
],
)
def test_operator_override_parsing(monkeypatch, raw, expected):
monkeypatch.setenv("UNSLOTH_SAMPLING_TEMPERATURE", raw)
assert ic._operator_sampling_override("temperature") == expected
def test_out_of_range_recommendation_is_dropped(monkeypatch):
# A malformed model recommendation (out of range) is ignored, so the request keeps the
# schema default rather than forwarding a bad value to llama-server.
_set_recommended(monkeypatch, {"temperature": 5.0, "top_k": 64})
eff = resolve_effective_sampling("some/model", _all_omitted())
assert eff["temperature"] == 0.6 # 5.0 is outside [0, 2] -> schema default
assert eff["top_k"] == 64 # a valid recommendation is still applied
def test_operator_override_top_k_int_and_range(monkeypatch):
monkeypatch.setenv("UNSLOTH_SAMPLING_TOP_K", "40")
assert ic._operator_sampling_override("top_k") == 40
monkeypatch.setenv("UNSLOTH_SAMPLING_TOP_K", "200") # above max 100
assert ic._operator_sampling_override("top_k") is None
monkeypatch.setenv("UNSLOTH_SAMPLING_TOP_K", "-1") # min allowed
assert ic._operator_sampling_override("top_k") == -1
@pytest.mark.parametrize(
"field, val",
[
("top_k", 10**400), # oversized int on an int field: int() ok, but math.isfinite raises
("top_k", float("nan")), # NaN reaching an int field: int(nan) raises ValueError
("top_k", float("inf")), # inf reaching an int field: int(inf) raises OverflowError
(
"temperature",
10**400,
), # oversized int on a float field: float(huge_int) raises OverflowError
],
)
def test_clean_sampling_value_rejects_unrepresentable(field, val):
# None of these may raise; each is unusable and must be dropped to None (regression: an
# oversized value used to raise OverflowError before the range check could drop it).
assert ic._clean_sampling_value(field, val) is None
def test_oversized_operator_override_ignored(monkeypatch):
# A huge integer string parses via int() but overflows float(); math.isfinite would raise
# OverflowError and 500 the request. It must be ignored like any other bad override and the
# field must fall back to the schema default -- no exception.
monkeypatch.setenv("UNSLOTH_SAMPLING_TOP_K", "9" * 400)
assert ic._operator_sampling_override("top_k") is None
_set_recommended(monkeypatch, {}) # no per-model recommendation -> schema default applies
eff = resolve_effective_sampling("some/model", _all_omitted())
assert eff["top_k"] == 20 # schema default, resolved without raising
def test_oversized_recommendation_ignored(monkeypatch):
# A malformed per-model recommendation carrying an oversized int must not raise while
# resolving either; the field simply falls back to the schema default.
_set_recommended(monkeypatch, {"temperature": 10**400, "top_k": 64})
eff = resolve_effective_sampling("some/model", _all_omitted())
assert eff["temperature"] == 0.6 # oversized -> dropped -> schema default
assert eff["top_k"] == 64 # a valid recommendation is still applied
def test_fill_recommended_sampling_openai_payload(monkeypatch):
from models.inference import ChatCompletionRequest
from routes.inference import _fill_recommended_sampling_openai
_set_recommended(monkeypatch, {"temperature": 1.0, "top_k": 64, "min_p": 0.0})
# Client sent only temperature; top_k / min_p were omitted.
payload = ChatCompletionRequest(
model = "m", messages = [{"role": "user", "content": "hi"}], temperature = 0.2
)
_fill_recommended_sampling_openai(payload, "some/model")
assert payload.temperature == 0.2 # explicit client value preserved
assert payload.top_k == 64 # recommended fills the omitted field
assert payload.min_p == 0.0
assert payload.top_p == 0.95 # no recommendation -> schema default unchanged
def test_fill_recommended_sampling_openai_operator_pin_overrides_client(monkeypatch):
from models.inference import ChatCompletionRequest
from routes.inference import _fill_recommended_sampling_openai
monkeypatch.setattr(ic, "load_model_defaults", lambda mid: {})
monkeypatch.setattr(ic, "get_family_inference_params", lambda mid: {})
ic._recommended_sampling.cache_clear()
monkeypatch.setenv("UNSLOTH_SAMPLING_TEMPERATURE", "0.9")
payload = ChatCompletionRequest(
model = "m", messages = [{"role": "user", "content": "hi"}], temperature = 0.2
)
_fill_recommended_sampling_openai(payload, "some/model")
assert payload.temperature == 0.9 # operator pin wins even over an explicit client value
def test_fill_recommended_sampling_completions_body(monkeypatch):
# /v1/completions is a raw proxy: recommendations fill omitted fields, but a field with no
# recommendation and no pin is left absent so llama-server keeps its own default (unlike the
# chat schema, which carries per-field defaults).
from routes.inference import _fill_recommended_sampling_completions
_set_recommended(monkeypatch, {"temperature": 1.0, "top_k": 64, "min_p": 0.0})
body = {"prompt": "hi", "temperature": 0.2}
_fill_recommended_sampling_completions(body, "some/model")
assert body["temperature"] == 0.2 # explicit client value preserved
assert body["top_k"] == 64 # recommendation fills the omitted field
assert body["min_p"] == 0.0
# No recommendation and no pin -> NOT injected (llama-server keeps its default).
assert "top_p" not in body
assert "presence_penalty" not in body
assert "repeat_penalty" not in body
def test_fill_recommended_sampling_completions_operator_pin(monkeypatch):
# An operator pin overrides the client's raw-body value, and the repetition pin is written
# under llama-server's "repeat_penalty" key (the schema field is repetition_penalty).
from routes.inference import _fill_recommended_sampling_completions
monkeypatch.setattr(ic, "load_inference_config", lambda mid: {})
ic._recommended_sampling.cache_clear()
monkeypatch.setenv("UNSLOTH_SAMPLING_TEMPERATURE", "0.9")
monkeypatch.setenv("UNSLOTH_SAMPLING_REPETITION_PENALTY", "1.2")
body = {"prompt": "hi", "temperature": 0.2, "repeat_penalty": 1.05}
_fill_recommended_sampling_completions(body, "some/model")
assert body["temperature"] == 0.9 # operator pin wins over the client's explicit value
assert body["repeat_penalty"] == 1.2 # repetition pin lands on llama-server's key
assert "repetition_penalty" not in body # never leak the schema field name into the body

View file

@ -297,6 +297,8 @@ class TestSandboxEnvIsolation:
"PYTHONPATH",
"VIRTUAL_ENV",
"SystemRoot",
"PATHEXT", # Windows only; minimal list so cwd scripts cannot hijack
"NoDefaultCurrentDirectoryInExePath", # Windows only; no cwd-first lookup
}
extras = set(env.keys()) - allowed
assert not extras, f"sandbox env added unexpected keys: {extras}"
@ -305,6 +307,220 @@ class TestSandboxEnvIsolation:
assert env["PYTHONPATH"].endswith("sandbox_site")
assert "leak-me" not in env["PYTHONPATH"]
def test_host_git_dir_appended_after_curated(self, monkeypatch, tmp_path):
# #7317: Windows Git lives under Program Files, not System32. Sandbox
# PATH resolves bare `git` by appending the dir of the git the HOST
# shell resolves (shutil.which), after the curated prefix.
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
prog = tmp_path / "Program Files"
monkeypatch.setattr(tools_mod, "_windows_program_roots", lambda: [str(prog)])
git_dir = prog / "Git" / "cmd"
git_dir.mkdir(parents = True)
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: str(git_dir / "git.exe"))
env = _build_safe_env(str(tmp_path))
parts = env["PATH"].split(os.pathsep)
assert str(git_dir) in parts
# Curated prefix stays ahead of host Git so Studio python/pip win.
assert parts.index(str(git_dir)) > 0
def test_host_path_dirs_not_inherited(self, monkeypatch, tmp_path):
"""Host PATH dirs (user-writable, git-lookalike) are never inherited;
only the resolved git dir is. No git resolved -> nothing appended."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
venv_scripts = tmp_path / "venv" / "Scripts"
venv_scripts.mkdir(parents = True)
fake_git = tmp_path / "scratch" / "Git" / "cmd"
fake_git.mkdir(parents = True)
monkeypatch.setenv(
"PATH",
os.pathsep.join([str(venv_scripts), str(fake_git), os.environ.get("PATH", "")]),
)
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: None)
env = _build_safe_env(str(tmp_path))
parts = env["PATH"].split(os.pathsep)
assert str(venv_scripts) not in parts
# A git-suffixed but unresolved (user-writable) dir is NOT trusted.
assert str(fake_git) not in parts
def test_git_cmd_shim_extension_added_to_pathext(self, monkeypatch, tmp_path):
"""A host git resolved as a .cmd shim under a trusted root stays
resolvable under the restricted PATHEXT (cwd lookup disabled)."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
prog = tmp_path / "Program Files"
monkeypatch.setattr(tools_mod, "_windows_program_roots", lambda: [str(prog)])
git_dir = prog / "Git" / "cmd"
git_dir.mkdir(parents = True)
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: str(git_dir / "git.cmd"))
env = _build_safe_env(str(tmp_path))
assert str(git_dir) in env["PATH"].split(os.pathsep)
assert env["PATHEXT"] == ".EXE;.COM;.CMD"
def test_user_writable_git_dir_refused(self, monkeypatch, tmp_path):
"""Git resolved from a per-user manager (Scoop shims) is NOT trusted:
an attacker could drop rg.exe beside it and hit the auto-approve gate."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
monkeypatch.setattr(
tools_mod, "_windows_program_roots", lambda: [str(tmp_path / "Program Files")]
)
shim_dir = tmp_path / "users" / "alice" / "scoop" / "shims"
shim_dir.mkdir(parents = True)
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: str(shim_dir / "git.exe"))
env = _build_safe_env(str(tmp_path))
assert str(shim_dir) not in env["PATH"].split(os.pathsep)
# No trusted git launcher -> PATHEXT stays minimal.
assert env["PATHEXT"] == ".EXE;.COM"
def test_trust_uses_known_folder_not_env_override(self, monkeypatch, tmp_path):
"""Trust is driven by the resolved Program Files roots, so a git under
an attacker-overridden %ProgramFiles% env value is still refused."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
real_prog = tmp_path / "RealProgramFiles"
(real_prog).mkdir()
evil = tmp_path / "attacker"
(evil / "Git" / "cmd").mkdir(parents = True)
# Resolver returns the genuine root; env is overridden to the evil dir.
monkeypatch.setattr(tools_mod, "_windows_program_roots", lambda: [str(real_prog)])
monkeypatch.setenv("ProgramFiles", str(evil))
monkeypatch.setattr(
tools_mod.shutil, "which", lambda name: str(evil / "Git" / "cmd" / "git.exe")
)
env = _build_safe_env(str(tmp_path))
assert str(evil / "Git" / "cmd") not in env["PATH"].split(os.pathsep)
def test_canonical_git_dir_appended(self, monkeypatch, tmp_path):
"""The PATH entry is the realpath of the trusted dir, not a junction
alias, so it cannot be retargeted after the trust check."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
real_prog = tmp_path / "Program Files"
real_git = real_prog / "Git" / "cmd"
real_git.mkdir(parents = True)
link = tmp_path / "link"
try:
link.symlink_to(real_prog, target_is_directory = True)
except (OSError, NotImplementedError):
pytest.skip("symlink unsupported in this environment")
monkeypatch.setattr(tools_mod, "_windows_program_roots", lambda: [str(real_prog)])
monkeypatch.setattr(
tools_mod.shutil,
"which",
lambda name: str(link / "Git" / "cmd" / "git.exe"),
)
env = _build_safe_env(str(tmp_path))
parts = env["PATH"].split(os.pathsep)
assert str(real_git) in parts # canonical, not the `link/...` alias
def test_windows_temp_git_dir_refused(self, monkeypatch, tmp_path):
"""A git under a world-writable %SystemRoot% subdir (Windows\\Temp) is
NOT trusted, even though it sits under the Windows root."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
monkeypatch.setattr(
tools_mod, "_windows_program_roots", lambda: [str(tmp_path / "Program Files")]
)
temp_git = tmp_path / "Windows" / "Temp" / "Git" / "cmd"
temp_git.mkdir(parents = True)
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: str(temp_git / "git.exe"))
env = _build_safe_env(str(tmp_path))
assert str(temp_git) not in env["PATH"].split(os.pathsep)
def test_trusted_program_dir_matches_via_realpath(self, monkeypatch, tmp_path):
"""The trust check canonicalizes paths, so a symlinked/short alias of
Program Files still matches (stand-in for 8.3 PROGRA~1 on Windows)."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
real_prog = tmp_path / "Program Files"
(real_prog / "Git" / "cmd").mkdir(parents = True)
alias = tmp_path / "PROGRA~1"
try:
alias.symlink_to(real_prog, target_is_directory = True)
except (OSError, NotImplementedError):
pytest.skip("symlink unsupported in this environment")
monkeypatch.setattr(tools_mod, "_windows_program_roots", lambda: [str(real_prog)])
git_via_alias = alias / "Git" / "cmd" / "git.exe"
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: str(git_via_alias))
env = _build_safe_env(str(tmp_path))
parts = [os.path.normcase(os.path.realpath(p)) for p in env["PATH"].split(os.pathsep)]
assert os.path.normcase(str(real_prog / "Git" / "cmd")) in parts
def test_scan_past_untrusted_git_shim(self, monkeypatch, tmp_path):
"""When an untrusted shim sorts first on PATH, the scan still finds a
later trusted Program Files git."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
prog = tmp_path / "Program Files"
trusted_git = prog / "Git" / "cmd"
trusted_git.mkdir(parents = True)
(trusted_git / "git.EXE").write_text("") # match PATHEXT case on this FS
shim = tmp_path / "scoop" / "shims"
shim.mkdir(parents = True)
(shim / "git.EXE").write_text("")
monkeypatch.setattr(tools_mod, "_windows_program_roots", lambda: [str(prog)])
# shutil.which returns the untrusted shim first.
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: str(shim / "git.EXE"))
monkeypatch.setenv("PATH", os.pathsep.join([str(shim), str(trusted_git)]))
monkeypatch.setenv("PATHEXT", ".EXE")
env = _build_safe_env(str(tmp_path))
parts = env["PATH"].split(os.pathsep)
assert str(trusted_git) in parts
assert str(shim) not in parts
def test_program_roots_fails_closed_without_known_folder_api(self, monkeypatch):
"""When the known-folder API is unavailable, no roots are trusted: env
vars (even %SystemDrive%) are caller-overrideable, so we never derive a
trusted root from them."""
import core.inference.tools as tools_mod
# ctypes fails on this Linux host, so the API path raises and we fail
# closed. Any attacker override of these env vars must be irrelevant.
monkeypatch.setenv("ProgramFiles", r"D:\attacker-writable")
monkeypatch.setenv("ProgramW6432", r"D:\attacker-writable")
monkeypatch.setenv("SystemDrive", "D:")
assert tools_mod._windows_program_roots() == []
def test_augment_native_program_roots_derives_native_sibling(self):
"""A 32-bit process only sees the x86 root; the native sibling is
derived by stripping the ` (x86)` suffix."""
import core.inference.tools as tools_mod
roots = tools_mod._augment_native_program_roots([r"C:\Program Files (x86)"])
lowered = [r.lower() for r in roots]
assert r"c:\program files (x86)" in lowered
assert r"c:\program files" in lowered
def test_no_default_current_directory_in_exe_path_set_on_windows(self, monkeypatch, tmp_path):
"""cmd/CreateProcess must not search cwd for bare names in the sandbox."""
import core.inference.tools as tools_mod
from core.inference.tools import _build_safe_env
monkeypatch.setattr(sys, "platform", "win32")
monkeypatch.setattr(tools_mod.shutil, "which", lambda name: None)
env = _build_safe_env(str(tmp_path))
assert env["NoDefaultCurrentDirectoryInExePath"] == "1"
def test_home_points_at_sandbox_workdir(self, tmp_path):
from core.inference.tools import _build_safe_env

View file

@ -1639,17 +1639,34 @@ def get_parent_visible_gpu_ids() -> list[int]:
return list(parent_visible_ids) if parent_visible_ids is not None else []
def resolve_requested_gpu_ids(gpu_ids: Optional[list[int]]) -> list[int]:
def resolve_requested_gpu_ids(
gpu_ids: Optional[list[int]], *, is_vulkan: bool = False
) -> list[int]:
parent_visible_spec = _get_parent_visible_gpu_spec()
parent_visible_ids = get_parent_visible_gpu_ids()
physical_gpu_count = get_physical_gpu_count()
if gpu_ids is None:
return parent_visible_ids
return [] if is_vulkan else parent_visible_ids
requested_ids = list(gpu_ids)
if len(requested_ids) == 0:
return parent_visible_ids
return [] if is_vulkan else parent_visible_ids
if is_vulkan:
# A Vulkan build selects by ggml Vulkan ordinal (--device VulkanN), a separate
# index space from CUDA/ROCm ids that may be empty under CPU-only torch. The
# CUDA parent-visible / physical-count checks below do not apply; only reject
# malformed ordinals (issue #7239).
if len(set(requested_ids)) != len(requested_ids):
raise ValueError(f"Invalid gpu_ids {requested_ids}: duplicate GPU IDs are not allowed.")
negative_ids = [gpu_id for gpu_id in requested_ids if gpu_id < 0]
if negative_ids:
raise ValueError(
f"Invalid gpu_ids {requested_ids}: GPU IDs must be non-negative. "
f"Rejected IDs: {negative_ids}."
)
return requested_ids
if not parent_visible_spec["supports_explicit_gpu_ids"]:
raise ValueError(
@ -2193,12 +2210,13 @@ def auto_select_gpu_ids(
metadata["selection_mode"] = "auto"
metadata["selected_gpu_ids"] = selected
logger.debug(
"Selected GPUs automatically",
model_name = model_name,
selected_gpu_ids = selected,
usable_gb = metadata["usable_gb"],
required_gb = metadata.get("required_gb"),
multi_gpu_overhead = multi_gpu_overhead,
"Selected GPUs automatically: model=%s selected=%s usable_gb=%s "
"required_gb=%s multi_gpu_overhead=%s",
model_name,
selected,
metadata["usable_gb"],
metadata.get("required_gb"),
multi_gpu_overhead,
)
return selected, metadata
@ -2214,12 +2232,13 @@ def auto_select_gpu_ids(
metadata["usable_gb"] = round(fallback_usable, 3)
metadata["selected_gpu_ids"] = fallback_all
logger.warning(
"Falling back to all visible GPUs -- model may not fit",
model_name = model_name,
selected_gpu_ids = fallback_all,
usable_gb = metadata["usable_gb"],
required_gb = metadata.get("required_gb"),
multi_gpu_overhead = multi_gpu_overhead,
"Falling back to all visible GPUs; model may not fit: model=%s "
"selected=%s usable_gb=%s required_gb=%s multi_gpu_overhead=%s",
model_name,
fallback_all,
metadata["usable_gb"],
metadata.get("required_gb"),
multi_gpu_overhead,
)
return fallback_all, metadata

View file

@ -5,7 +5,10 @@
from pathlib import Path
from typing import Dict, Any, Optional
from functools import lru_cache
import json
import math
import os
import yaml
import structlog
from loggers import get_logger
@ -160,3 +163,137 @@ def load_inference_config(model_identifier: str) -> Dict[str, Any]:
}
return inference_config
# ── Effective sampling resolution for `unsloth run` / `unsloth start` ──────────
#
# Per-model recommended sampling is applied to a request only for the fields the
# client omitted; an operator can pin a field from the CLI via UNSLOTH_SAMPLING_*
# (a hard override that wins even over an explicit client value). Precedence per
# field: operator pin -> client explicit -> per-model recommendation -> the static
# schema default (mirroring ChatCompletionRequest, so behavior is unchanged when
# nothing is recommended or pinned).
# field -> (env var, static default, min, max, is_int)
_SAMPLING_FIELDS = {
"temperature": ("UNSLOTH_SAMPLING_TEMPERATURE", 0.6, 0.0, 2.0, False),
"top_p": ("UNSLOTH_SAMPLING_TOP_P", 0.95, 0.0, 1.0, False),
"top_k": ("UNSLOTH_SAMPLING_TOP_K", 20, -1, 100, True),
"min_p": ("UNSLOTH_SAMPLING_MIN_P", 0.01, 0.0, 1.0, False),
"repetition_penalty": ("UNSLOTH_SAMPLING_REPETITION_PENALTY", 1.0, 1.0, 2.0, False),
"presence_penalty": ("UNSLOTH_SAMPLING_PRESENCE_PENALTY", 0.0, 0.0, 2.0, False),
}
# Public, ordered tuple of the sampling fields callers resolve.
SAMPLING_FIELD_NAMES = tuple(_SAMPLING_FIELDS)
# Fields the Studio Chat UI adopts as *per-model recommendations* from the backend
# `.inference` block. Its frontend `mergeBackendRecommendedInference`
# (presets/preset-policy.ts) seeds exactly these five and never reads repetition_penalty,
# so the server auto-recommends the same five for request parity. repetition_penalty stays a
# manual-only knob (client-sent or an UNSLOTH_SAMPLING_REPETITION_PENALTY operator pin),
# matching the UI where it is never auto-filled per model.
_UI_RECOMMENDED_FIELDS = ("temperature", "top_p", "top_k", "min_p", "presence_penalty")
def _clean_sampling_value(field: str, val: Any):
"""Coerce ``val`` to the field's numeric type when it is a finite, in-range number, else None.
Rejects bool, non-numeric, NaN/inf, and out-of-range values so neither a bad operator env
var nor a malformed model recommendation can reach llama-server. NaN matters because
``nan < lo`` and ``nan > hi`` are both False, so a plain range check would let it through.
Coerce before the finiteness check: ``math.isfinite`` and ``float()`` raise ``OverflowError``
on an int too big for a C double (an oversized UNSLOTH_SAMPLING_TOP_K would otherwise 500 the
request), while an in-range int is range-checked exactly and ``int()`` rejects a NaN/inf that
reached an int field.
"""
if isinstance(val, bool) or not isinstance(val, (int, float)):
return None
_env, _default, lo, hi, is_int = _SAMPLING_FIELDS[field]
try:
val = int(val) if is_int else float(val)
except (ValueError, OverflowError):
# int(nan)/int(inf) and float(oversized_int) raise; treat them as unusable.
return None
# After coercion an int is always finite; only a float can still be NaN/inf.
if isinstance(val, float) and not math.isfinite(val):
return None
if val < lo or val > hi:
return None
return val
def _operator_sampling_override(field: str):
"""Operator-pinned value for a sampling field from UNSLOTH_SAMPLING_*, or None.
An unparseable, non-finite, or out-of-range value is ignored so a bad env var can never
reach llama-server; the field then falls back to the client / recommended value.
"""
_env, _default, _lo, _hi, is_int = _SAMPLING_FIELDS[field]
raw = os.environ.get(_env)
if raw is None or raw.strip() == "":
return None
try:
val = int(raw) if is_int else float(raw)
except (TypeError, ValueError):
return None
return _clean_sampling_value(field, val)
@lru_cache(maxsize = 128)
def _recommended_sampling(model_id: str) -> Dict[str, Any]:
"""Per-model recommended sampling, resolved through the SAME path the Studio Chat UI uses.
The Chat UI seeds its sampling from the ``.inference`` block of the load/status responses,
which is exactly :func:`load_inference_config` (model-specific YAML -> family defaults
(inference_defaults.json) -> default.yaml). Sourcing recommendations here keeps the values
the server applies to a request identical to what the UI shows for the same model. Only the
fields the UI actually adopts (:data:`_UI_RECOMMENDED_FIELDS`) are recommended; each value
is validated (finite + in range) before use. Cached by model id.
"""
if not model_id:
return {}
try:
cfg = load_inference_config(model_id) or {}
except Exception as e:
logger.debug(f"Could not load recommended sampling for '{model_id}': {e}")
return {}
recommended: Dict[str, Any] = {}
for field in _UI_RECOMMENDED_FIELDS:
cleaned = _clean_sampling_value(field, cfg.get(field))
if cleaned is not None:
recommended[field] = cleaned
return recommended
def resolve_effective_sampling(
model_id: Optional[str],
explicit: Dict[str, Any],
*,
fill_defaults: bool = True,
) -> Dict[str, Any]:
"""Resolve the effective sampling params for a request.
``explicit`` maps each field in :data:`SAMPLING_FIELD_NAMES` to the client-sent
value, or ``None`` when the client omitted it. Precedence (highest first): an
operator ``UNSLOTH_SAMPLING_*`` pin, then the client's explicit value, then the
per-model recommendation, then the static schema default.
When ``fill_defaults`` is False a field with no operator pin, client value, or
per-model recommendation is omitted from the result instead of set to the static
schema default, so a raw proxy body (``/v1/completions``) keeps llama-server's own
default for that field rather than being forced onto this schema's value.
"""
recommended = _recommended_sampling(model_id or "")
effective: Dict[str, Any] = {}
for field, (_env, default, _lo, _hi, _int) in _SAMPLING_FIELDS.items():
override = _operator_sampling_override(field)
if override is not None:
effective[field] = override
elif explicit.get(field) is not None:
effective[field] = explicit[field]
elif field in recommended:
effective[field] = recommended[field]
elif fill_defaults:
effective[field] = default
return effective

View file

@ -46,17 +46,6 @@ _PICKLE_WEIGHT_RE = re.compile(
r"\.(bin|pt|pth|ckpt|pkl|pickle)$",
re.IGNORECASE,
)
# Base-model safetensors set: HF names the base pickle pytorch_model.bin but the safetensors
# model.safetensors (stems differ), so a base pickle is replaced only by these, not an adapter's.
_BASE_SAFETENSORS_RE = re.compile(
r"^(model(-\d+-of-\d+)?\.safetensors|model\.safetensors\.index\.json)$",
re.IGNORECASE,
)
# Adapter (PEFT) safetensors set: adapter_model.safetensors, its shards, or index.
_ADAPTER_SAFETENSORS_RE = re.compile(
r"^(adapter_model(-\d+-of-\d+)?\.safetensors|adapter_model\.safetensors\.index\.json)$",
re.IGNORECASE,
)
# Non-blocking levels: clean or not-yet-finished. Anything else (unsafe/suspicious/
# malicious or a future label) blocks, so Hub schema drift fails CLOSED.
@ -94,6 +83,13 @@ _INERT_SUFFIXES = frozenset(
_SOURCE_SUFFIXES = frozenset({".py", ".pyc", ".pyx", ".pyi"})
# Torch-family weight indexes: from_pretrained feeds each shard they name to load_state_dict, which
# torch.load()s (pickle) any shard whose name does not end in .safetensors, whatever its stem. A
# pytorch index is superseded when a base safetensors is present (the loader prefers it); a
# safetensors index IS the chosen archive, so a non-safetensors target it names still loads. tf/flax
# indexes load via non-pickle loaders, so they are not a torch.load vector here.
_TORCH_INDEX_FILES = ("pytorch_model.bin.index.json", "model.safetensors.index.json")
# Root weight-index files. from_pretrained reads these to find sharded weights, so a
# flagged subdir pickle is a load vector iff a root index references it.
_TRANSFORMERS_INDEX_FILES = (
@ -313,13 +309,72 @@ def _st_load_roots(snapshot: Path) -> list:
return roots
def _indexed_pickle_shards(index_path: Path, root: Path, snapshot: Path) -> list:
"""Shards a torch weight index points a ``from_pretrained`` load at that load_state_dict would
torch.load (pickle): every ``weight_map`` target NOT ending in ``.safetensors``, whatever its
stem (an arbitrary name like ``shards/payload`` still deserializes). Resolved relative to the
index dir (``root``) like the loader, so a shard in a nested dir is followed (iterdir misses it).
Lexical only, never ``Path.resolve()`` (HF snapshot files symlink into ``blobs/``, so resolving
escapes the snapshot and false-blocks every shard). Raises OSError -> caller fails CLOSED on an
unreadable/invalid index or a target escaping the snapshot."""
import json
import os
try:
# JSON is UTF-8 by spec; pin it so a non-ASCII index is not misdecoded (and needlessly
# blocked) under Windows' cp1252 default.
parsed = json.loads(index_path.read_text(encoding = "utf-8"))
except (OSError, ValueError) as exc:
raise OSError(f"unreadable weight index: {index_path}") from exc
weight_map = parsed.get("weight_map") if isinstance(parsed, dict) else None
if not isinstance(weight_map, dict):
return [] # no dict weight_map -> the loader resolves no shards from this index
snapshot_norm = os.path.normpath(str(snapshot))
shards = []
for shard in weight_map.values():
raw = str(shard)
if not raw:
continue
# Join the RAW weight_map value like from_pretrained's os.path.join: on POSIX a backslash is a
# literal filename char (not a separator), so normalizing it would probe a different path than
# the loader opens. normpath + containment stay platform-aware (os.sep) to block "..".
joined = os.path.normpath(os.path.join(str(root), raw))
if joined != snapshot_norm and not joined.startswith(snapshot_norm + os.sep):
raise OSError(f"weight index escapes the snapshot: {index_path}")
shard_path = Path(joined)
# Case-SENSITIVE, mirroring load_state_dict's own endswith(".safetensors"): a shard named
# payload.SAFETENSORS is not treated as safetensors by the loader and falls to torch.load.
if not shard_path.name.endswith(".safetensors") and shard_path.is_file():
shards.append(shard_path)
return shards
def _loader_resolves(root: Path, name: str) -> bool:
"""True iff from_pretrained would open ``name`` under ``root``. ``is_file()`` honors the platform
(case-sensitive on Linux, case-insensitive on Windows/macOS), so it mirrors the loader's own
lookup: an oddly-cased decoy counts as an alternative only where the loader would truly open it.
A name-fold instead would let an uppercase MODEL.SAFETENSORS suppress the scan on Linux while the
loader, asking for the canonical lowercase name, silently falls through to a pickle index."""
return (root / name).is_file()
def _cached_pickle_weight_files(snapshot: Path) -> list:
"""Pickle weight files in snapshot's ST load roots, EXCLUDING those whose weight family also
ships an inert safetensors in the same dir (the loader prefers it): a base pickle is suppressed
only by a base model.safetensors, an adapter pickle only by adapter_model.safetensors -- an
unrelated safetensors is no substitute. Load roots only. Raises OSError if the snapshot root is
unreadable (caller blocks)."""
"""Pickle weight files a SentenceTransformer/Transformers load deserializes from snapshot's ST
load roots, EXCLUDING those whose weight family also ships an inert safetensors in the same dir
(the loader prefers it): a base pickle is suppressed only by a base model.safetensors, an adapter
pickle only by adapter_model.safetensors -- an unrelated safetensors is no substitute. Covers
both direct-child pickles AND pickle shards referenced by a local weight index (which the loader
follows into nested dirs, matching the online gate). Raises OSError -- caller fails CLOSED -- if
the snapshot root or a weight index is unreadable, or an index reference escapes the snapshot."""
blocked = []
seen = set()
def _add(path: Path):
key = str(path)
if key not in seen:
seen.add(key)
blocked.append(path)
for root in _st_load_roots(snapshot):
try:
entries = [p for p in root.iterdir() if p.is_file()]
@ -327,15 +382,35 @@ def _cached_pickle_weight_files(snapshot: Path) -> list:
if root == snapshot:
raise # top-level unreadable -> fail closed
continue # unreadable module subdir: nothing loadable to attest here
has_base_safetensors = any(_BASE_SAFETENSORS_RE.match(p.name) for p in entries)
has_adapter_safetensors = any(_ADAPTER_SAFETENSORS_RE.match(p.name) for p in entries)
# Safetensors alternatives the loader would actually resolve (never a bare name-fold, which
# fails OPEN: see _loader_resolves). A base pickle is replaced only by a base safetensors, an
# adapter pickle only by an adapter one. A single model.safetensors also outranks BOTH indexes.
has_direct_base_safetensors = _loader_resolves(root, "model.safetensors")
has_base_safetensors = has_direct_base_safetensors or _loader_resolves(
root, "model.safetensors.index.json"
)
has_adapter_safetensors = _loader_resolves(root, "adapter_model.safetensors")
for path in entries:
if not _PICKLE_WEIGHT_RE.match(path.name):
continue
is_adapter = path.name.lower().startswith("adapter_model")
has_alternative = has_adapter_safetensors if is_adapter else has_base_safetensors
if not has_alternative:
blocked.append(path)
_add(path)
# A torch weight index makes from_pretrained load nested shards iterdir never sees; the loader
# torch.loads any not ending in .safetensors. Probe the canonical index name with the loader's
# own lookup (_loader_resolves), so an oddly-cased artifact it would never open does not block.
# A direct model.safetensors wins over BOTH indexes; failing that a base safetensors still
# outranks the pytorch index, while a safetensors index is itself the chosen archive.
for index_name in _TORCH_INDEX_FILES:
if not _loader_resolves(root, index_name):
continue
if has_direct_base_safetensors:
continue
if index_name == "pytorch_model.bin.index.json" and has_base_safetensors:
continue
for shard_path in _indexed_pickle_shards(root / index_name, root, snapshot):
_add(shard_path)
return blocked

View file

@ -901,8 +901,8 @@ export function AppSidebar() {
: "group/recent-item relative";
const actionClass =
variant === "project"
? "sidebar-row-action group-hover/project-chat-item:opacity-100 group-hover/project-chat-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto"
: "sidebar-row-action group-hover/recent-item:opacity-100 group-hover/recent-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto";
? "sidebar-row-action sidebar-touch-reveal group-hover/project-chat-item:opacity-100 group-hover/project-chat-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto"
: "sidebar-row-action sidebar-touch-reveal group-hover/recent-item:opacity-100 group-hover/recent-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto";
const buttonClass = cn(
"sidebar-nav-btn h-[33px] cursor-pointer rounded-full pr-4 text-ui-14p5 leading-ui-19 tracking-nav font-medium",
// pl-3 (12px) over the content's pl-1.5 (6px) = 18px, aligning the
@ -912,13 +912,14 @@ export function AppSidebar() {
isPinned && variant !== "project" && "gap-[8.5px]",
variant === "project"
? // Room for the hover pin quick-action plus the kebab.
"group-hover/project-chat-item:pr-14 group-has-[.sidebar-row-action[data-state=open]]/project-chat-item:pr-8"
"group-hover/project-chat-item:pr-14 group-has-[.sidebar-row-action[data-state=open]]/project-chat-item:pr-8 [@media(pointer:coarse)]:pr-14"
: isPinned
? // Pinned rows show an extra unpin button on hover, so reserve more room
// (pr-8 when the menu is open keeps the unpin button clear of the title).
"group-hover/recent-item:pr-16 group-has-[.sidebar-row-action[data-state=open]]/recent-item:pr-8"
"group-hover/recent-item:pr-16 group-has-[.sidebar-row-action[data-state=open]]/recent-item:pr-8 [@media(pointer:coarse)]:pr-16"
: // Hover room for the kebab only; title keeps one more character.
"group-hover/recent-item:pr-6 group-has-[.sidebar-row-action[data-state=open]]/recent-item:pr-6",
// Touch rows clear the full always-visible kebab hit area (pr-10).
"group-hover/recent-item:pr-6 group-has-[.sidebar-row-action[data-state=open]]/recent-item:pr-6 [@media(pointer:coarse)]:pr-10",
);
const isRenamingThis =
@ -987,7 +988,7 @@ export function AppSidebar() {
togglePinnedChat(item.id);
}}
aria-label={isPinned ? "Unpin chat" : "Pin chat"}
className="sidebar-row-action is-unpin-action group-hover/project-chat-item:opacity-100 group-hover/project-chat-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto"
className="sidebar-row-action sidebar-touch-reveal is-unpin-action group-hover/project-chat-item:opacity-100 group-hover/project-chat-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto"
>
<span className="sidebar-row-action-glyph">
<HugeiconsIcon icon={isPinned ? PinOffIcon : PinIcon} strokeWidth={1.75} className="size-icon" />
@ -1002,7 +1003,7 @@ export function AppSidebar() {
togglePinnedChat(item.id);
}}
aria-label="Unpin chat"
className="sidebar-row-action is-unpin-action group-hover/recent-item:opacity-100 group-hover/recent-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto"
className="sidebar-row-action sidebar-touch-reveal is-unpin-action group-hover/recent-item:opacity-100 group-hover/recent-item:pointer-events-auto focus-visible:opacity-100 focus-visible:pointer-events-auto"
>
<span className="sidebar-row-action-glyph">
<HugeiconsIcon icon={PinOffIcon} strokeWidth={1.75} className="size-icon" />

View file

@ -1556,6 +1556,7 @@ const Composer: FC<{
const draftThreadId = referenceThreadId;
const draftKey = draftThreadId ? composerDraftKey(draftThreadId) : null;
const lastDraftKeyRef = useRef(draftKey);
const draftSaveTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null);
useEffect(() => {
const draft = draftKey ? (readComposerDraft(draftKey) ?? "") : "";
const composer = aui.composer();
@ -1574,8 +1575,25 @@ const Composer: FC<{
return;
}
const t = setTimeout(() => writeComposerDraft(draftKey, composerText), 300);
draftSaveTimerRef.current = t;
return () => clearTimeout(t);
}, [composerText, draftKey]);
// Without this the restore effect above puts the sent text back when the
// runtime rebinds on the first message.
const draftKeyRef = useRef(draftKey);
useEffect(() => {
draftKeyRef.current = draftKey;
}, [draftKey]);
const clearStoredDraft = useCallback(() => {
if (draftSaveTimerRef.current !== null) {
clearTimeout(draftSaveTimerRef.current);
draftSaveTimerRef.current = null;
}
const key = draftKeyRef.current;
if (key) {
writeComposerDraft(key, "");
}
}, []);
// react-textarea-autosize re-measures only on value change or window resize,
// not on the width swap from expanding, so it keeps the taller height and
// leaves a stray blank row. Nudge a resize whenever input width changes.
@ -1726,9 +1744,10 @@ const Composer: FC<{
setPendingSend(false);
dismissWaitToast();
if (text.trim().length > 0 || attachments.length > 0) {
clearStoredDraft();
aui.composer().send();
}
}, [pendingSend, indexingActive, aui, dismissWaitToast]);
}, [pendingSend, indexingActive, aui, clearStoredDraft, dismissWaitToast]);
// Drop any queued send + toast on unmount (e.g. thread switch).
useEffect(
@ -1771,6 +1790,7 @@ const Composer: FC<{
flushResourcesSync(() => {
aui.composer().setText("");
});
clearStoredDraft();
startPromptQueue(
[queuedPrompt],
createPromptQueueTarget(),
@ -1804,6 +1824,7 @@ const Composer: FC<{
closeOverlay();
return;
}
clearStoredDraft();
setImageToolsEnabled(true);
setPendingImageEditReference({
threadId: overlay.threadId ?? referenceThreadId,
@ -1821,11 +1842,15 @@ const Composer: FC<{
);
});
closeOverlay();
return;
}
clearStoredDraft();
},
[
aui,
canQueueCurrentPrompt,
clearStoredDraft,
closeOverlay,
composerText,
createPromptQueueTarget,
@ -1947,6 +1972,7 @@ const Composer: FC<{
flushResourcesSync(() => {
aui.composer().setText("");
});
clearStoredDraft();
startPromptQueue([queuedPrompt], createPromptQueueTarget(), true);
}}
onSendClick={interceptSend}

View file

@ -83,7 +83,7 @@ function CopyBtn({ text }: { text: string }) {
);
}
/** Save the executed script as a .py file via a client-side Blob (no server file serving). */
/** Save the script as a .py file via a client-side Blob. */
function DownloadBtn({ code, name = "script.py" }: { code: string; name?: string }) {
const download = useCallback(() => {
if (typeof document === "undefined") {
@ -229,8 +229,8 @@ const PythonToolUIImpl: ToolCallMessagePartComponent = ({
const authToken = getAuthToken();
return (
// Run status and output collapse from history, but the script source is
// rendered outside ToolFallbackContent so it stays visible on reopen (#7165).
// Status/output collapse from history; the script source renders outside
// ToolFallbackContent so it stays visible on reopen (#7165).
<ToolFallbackRoot defaultOpen={isRunning}>
<ToolFallbackTrigger
toolName={firstLine ? `Python: ${firstLine}` : "Python"}

View file

@ -92,10 +92,8 @@ export const OPTIMIZER_OPTIONS: ReadonlyArray<{ value: string; label: string }>
{ value: "adamw_torch_fused", label: "AdamW (PyTorch Fused)" },
];
// Optimizers the MLX trainer actually supports on Apple Silicon. Values must
// match SUPPORTED_MLX_OPTIMIZERS in unsloth-zoo's mlx/trainer.py; on MLX the
// bitsandbytes/torch names above have no meaning and are remapped to plain
// AdamW, so Studio offers this list instead when running on a Mac.
// MLX trainer optimizers (Apple Silicon); must match SUPPORTED_MLX_OPTIMIZERS in
// unsloth-zoo's mlx/trainer.py. The CUDA/torch names above are remapped to AdamW on MLX.
export const MLX_OPTIMIZER_OPTIONS: ReadonlyArray<{ value: string; label: string }> = [
{ value: "adamw", label: "AdamW" },
{ value: "adam", label: "Adam" },

View file

@ -23,6 +23,8 @@ export interface ProviderConfig {
display_name: string;
base_url: string;
is_enabled: boolean;
models?: string[];
available_models?: string[];
created_at: string;
updated_at: string;
}
@ -123,6 +125,8 @@ export async function createProviderConfig(payload: {
providerType: string;
displayName: string;
baseUrl?: string | null;
models?: string[];
availableModels?: string[];
}): Promise<ProviderConfig> {
const response = await authFetch("/api/providers/", {
method: "POST",
@ -131,6 +135,8 @@ export async function createProviderConfig(payload: {
provider_type: payload.providerType,
display_name: payload.displayName,
base_url: payload.baseUrl ?? null,
models: payload.models ?? [],
available_models: payload.availableModels ?? [],
}),
});
return parseJsonOrThrow<ProviderConfig>(response);
@ -158,6 +164,8 @@ export async function updateProviderConfig(
displayName?: string;
baseUrl?: string | null;
isEnabled?: boolean;
models?: string[];
availableModels?: string[];
},
): Promise<ProviderConfig> {
const response = await authFetch(`/api/providers/${providerId}`, {
@ -167,6 +175,10 @@ export async function updateProviderConfig(
...(payload.displayName === undefined ? {} : { display_name: payload.displayName }),
...(payload.baseUrl === undefined ? {} : { base_url: payload.baseUrl }),
...(payload.isEnabled === undefined ? {} : { is_enabled: payload.isEnabled }),
...(payload.models === undefined ? {} : { models: payload.models }),
...(payload.availableModels === undefined
? {}
: { available_models: payload.availableModels }),
}),
});
return parseJsonOrThrow<ProviderConfig>(response);

View file

@ -175,6 +175,7 @@ import {
} from "./stores/chat-runtime-store";
import { useChatPreferencesStore } from "./stores/chat-preferences-store";
import { useExternalProvidersStore } from "./stores/external-providers-store";
import { syncExternalProvidersFromBackend } from "./sync-external-providers";
import { buildChatTourSteps } from "./tour";
import type { ChatView, MessageRecord } from "./types";
import {
@ -1762,8 +1763,18 @@ export function ChatPage({
const externalProvidersForChat = connectionsEnabled ? externalProviders : [];
useEffect(() => {
void hydratePersistedSettings();
}, [hydratePersistedSettings]);
void (async () => {
await hydratePersistedSettings();
try {
const synced = await syncExternalProvidersFromBackend(
useExternalProvidersStore.getState().providers,
);
setExternalProviders(synced);
} catch {
// Silent on startup; Connections settings still surfaces load errors.
}
})();
}, [hydratePersistedSettings, setExternalProviders]);
useEffect(() => {
// Skip while off-route: ChatPage stays mounted, and toast+navigate here would

View file

@ -41,7 +41,6 @@ import {
type ProviderRegistryEntry,
createProviderConfig,
deleteProviderConfig,
listProviderConfigs,
listProviderModels,
listProviderRegistry,
testProviderConnection,
@ -49,7 +48,6 @@ import {
} from "./api/providers-api";
import type { ExternalProviderConfig } from "./external-providers";
import {
CUSTOM_BACKEND_PROVIDER_TYPE,
CUSTOM_PROVIDER_PRESETS,
allowsManualModelIdsWithCatalog,
customProviderBaseUrlPlaceholder,
@ -68,6 +66,10 @@ import {
toExternalBackendProviderType,
} from "./external-providers";
import { useExternalProvidersStore } from "./stores/external-providers-store";
import {
pruneProviderModelIds,
syncExternalProvidersFromBackend,
} from "./sync-external-providers";
/** Matches navbar / thread layout easing (see index.css --ease-out-quart) */
const PROVIDER_FORM_EASE: [number, number, number, number] = [
@ -76,58 +78,7 @@ const PROVIDER_FORM_EASE: [number, number, number, number] = [
const PROVIDER_FORM_DURATION = 0.2;
const CUSTOM_PROVIDER_MISSING_KEY_MESSAGE =
"No API key found. Add a valid API key for this connection.";
const ANTHROPIC_DATED_SNAPSHOT_SUFFIX = /-\d{8}$/;
const OPENAI_DEPRECATED_MODELS = new Set(["gpt-5.3"]);
const HIDDEN_PROVIDER_TYPES = new Set(["qwen"]);
const OPENROUTER_EXCLUDED_MODELS = new Set([
"google/chirp-3",
"kwaivgi/kling-v3.0-pro",
"openai/whisper-1",
"openai/gpt-4o-mini-transcribe",
"recraft/recraft-v4-pro",
]);
function normalizeUrl(input: string): string {
return input.trim().replace(/\/+$/, "");
}
function resolveUiProviderTypeFromConfig(
configProviderType: string,
configDisplayName: string | null | undefined,
configBaseUrl: string | null | undefined,
registryRows: ProviderRegistryEntry[],
existingProviderType: string | undefined,
): string {
if (existingProviderType && isCustomProviderType(existingProviderType)) {
return existingProviderType;
}
if (configProviderType !== CUSTOM_BACKEND_PROVIDER_TYPE) {
return configProviderType;
}
const displayName = (configDisplayName ?? "").trim().toLowerCase();
const matchingCustomPreset = CUSTOM_PROVIDER_PRESETS.find(
(preset) => preset.displayName.toLowerCase() === displayName,
);
if (matchingCustomPreset) {
return matchingCustomPreset.providerType;
}
const openAiRegistry = registryRows.find(
(entry) => entry.provider_type === CUSTOM_BACKEND_PROVIDER_TYPE,
);
if (!openAiRegistry) {
return configProviderType;
}
const openAiDisplayName = openAiRegistry.display_name.trim().toLowerCase();
if (displayName.length > 0 && displayName !== openAiDisplayName) {
return LEGACY_CUSTOM_PROVIDER_TYPE;
}
const configUrl = normalizeUrl(configBaseUrl ?? "");
const defaultUrl = normalizeUrl(openAiRegistry.base_url ?? "");
if (configUrl.length > 0 && configUrl !== defaultUrl) {
return LEGACY_CUSTOM_PROVIDER_TYPE;
}
return configProviderType;
}
function parseManualModelIds(text: string): string[] {
const seen = new Set<string>();
@ -182,19 +133,6 @@ function shouldAppendOpenAiVersionPath(providerType: string): boolean {
);
}
function pruneProviderModelIds(providerType: string, modelIds: string[]): string[] {
if (providerType === "anthropic") {
return modelIds.filter((id) => !ANTHROPIC_DATED_SNAPSHOT_SUFFIX.test(id));
}
if (providerType === "openai") {
return modelIds.filter((id) => !OPENAI_DEPRECATED_MODELS.has(id));
}
if (providerType === "openrouter") {
return modelIds.filter((id) => !OPENROUTER_EXCLUDED_MODELS.has(id));
}
return modelIds;
}
function formatModelSummary(models: string[]): string {
if (models.length === 0) {
return "No models enabled";
@ -360,9 +298,9 @@ export function ChatProvidersSettings({
}
let syncSucceeded = false;
try {
const [registryRows, configRows] = await Promise.all([
const [registryRows, syncedProviders] = await Promise.all([
listProviderRegistry(),
listProviderConfigs(),
syncExternalProvidersFromBackend(providersRef.current),
]);
if (!isMounted) return;
syncSucceeded = true;
@ -377,61 +315,6 @@ export function ChatProvidersSettings({
}
return registryRows[0]?.provider_type ?? "";
});
const existingById = new Map<string, ExternalProviderConfig>();
for (const provider of providersRef.current) {
existingById.set(provider.id, provider);
}
const syncedProviders: ExternalProviderConfig[] = configRows
.filter((config) => config.is_enabled)
.map((config) => {
const existing = existingById.get(config.id);
const uiProviderType = resolveUiProviderTypeFromConfig(
config.provider_type,
config.display_name,
config.base_url,
registryRows,
existing?.providerType,
);
const createdAt = Number.isFinite(Date.parse(config.created_at))
? Date.parse(config.created_at)
: Date.now();
const updatedAt = Number.isFinite(Date.parse(config.updated_at))
? Date.parse(config.updated_at)
: Date.now();
const registryEntry =
registryRows.find((entry) => entry.provider_type === uiProviderType) ??
registryRows.find((entry) => entry.provider_type === config.provider_type);
const defaultModels = pruneProviderModelIds(
uiProviderType,
registryEntry?.default_models ?? [],
);
const savedModels = existing?.models ?? [];
const savedAvailableModels = existing?.availableModels ?? [];
const existingModels = pruneProviderModelIds(
uiProviderType,
savedModels.length > 0 ? savedModels : defaultModels,
);
const existingAvailableModels = pruneProviderModelIds(
uiProviderType,
savedAvailableModels.length > 0 ? savedAvailableModels : defaultModels,
);
return {
id: config.id,
providerType: uiProviderType,
name: config.display_name,
baseUrl: config.base_url ?? "",
models: existingModels,
availableModels: existingAvailableModels,
enablePromptCaching: supportsProviderPromptCaching(uiProviderType)
? (existing?.enablePromptCaching ?? true)
: undefined,
isReasoningModel: supportsProviderReasoningToggle(uiProviderType)
? existing?.isReasoningModel === true
: undefined,
createdAt: existing?.createdAt ?? createdAt,
updatedAt,
};
});
// Trust the backend response. An empty array means every connection was
// removed (often from another tab); mirror that locally, else stale
// entries become un-removable here until localStorage is cleared.
@ -699,6 +582,10 @@ export function ChatProvidersSettings({
providerType: backendProviderType,
displayName,
baseUrl,
models: modelsToSave,
availableModels: manualOnly
? []
: pruneProviderModelIds(providerType, availableModels),
});
const createdAt = Number.isFinite(Date.parse(created.created_at))
? Date.parse(created.created_at)
@ -814,6 +701,10 @@ export function ChatProvidersSettings({
customProviderDisplayName(existing.providerType)
: existing.name,
baseUrl,
models: modelsToSave,
availableModels: manualOnly
? []
: pruneProviderModelIds(existing.providerType, availableModels),
});
if (apiKey.trim()) {
setExternalProviderApiKey(editingProviderId, apiKey.trim());

View file

@ -222,7 +222,9 @@ export function applyActiveModelStatusToStore(
incomingGpuMode === "manual" ? (status.n_cpu_moe ?? null) : null;
const incomingSplit =
incomingGpuMode === "manual" ? (status.tensor_split ?? null) : null;
const incomingGpuIds = status.is_gguf ? (status.gpu_ids ?? null) : null;
const incomingGpuIds = status.is_gguf
? (status.requested_gpu_ids ?? status.gpu_ids ?? null)
: null;
const gpuStatusChanged =
prevState.loadedGpuMemoryMode !== incomingGpuMode ||
prevState.loadedGpuLayers !== incomingGpuLayers ||

View file

@ -622,6 +622,7 @@ export function loadedGpuMemoryFields(resp: {
n_layers?: number | null;
n_moe_layers?: number;
gpu_ids?: number[] | null;
requested_gpu_ids?: number[] | null;
}) {
// GPU-memory state is meaningful only for a GGUF chat load. A non-GGUF response
// still carries gpu_memory_mode (its default "auto" is serialized), so gate on
@ -649,7 +650,9 @@ export function loadedGpuMemoryFields(resp: {
};
}
const mode = resp.gpu_memory_mode ?? "auto";
const gpuIds = resp.gpu_ids ?? null;
// Keep the user's placement pool editable across status/load hydration.
// gpu_ids remains the effective fitted subset for diagnostics.
const gpuIds = resp.requested_gpu_ids ?? resp.gpu_ids ?? null;
// Layer/MoE/split knobs apply (and are reported) only in manual mode; in auto
// the server ignores them, so don't seed the loaded baseline or the editable
// knobs with values it never applied. In manual, the server reports gpu_layers
@ -687,7 +690,7 @@ export function loadedGpuMemoryFields(resp: {
ggufLayerCount: resp.n_layers ?? null,
// MoE expert-layer count: the n_cpu_moe slider max, and 0 hides the slider.
moeLayerCount: resp.n_moe_layers ?? null,
// The picker reflects what loaded (the request sent the user's pick).
// The picker reflects the requested placement pool, not a fitted subset.
selectedGpuIds: gpuIds,
// What the running server loaded is by definition in the current backend's
// index space.

View file

@ -0,0 +1,221 @@
// 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 ProviderRegistryEntry,
listProviderConfigs,
listProviderRegistry,
updateProviderConfig,
} from "./api/providers-api";
import {
CUSTOM_BACKEND_PROVIDER_TYPE,
CUSTOM_PROVIDER_PRESETS,
type ExternalProviderConfig,
isCustomProviderType,
isPromptCacheTtl,
LEGACY_CUSTOM_PROVIDER_TYPE,
supportsProviderPromptCaching,
supportsProviderPromptCacheTtl,
supportsProviderReasoningToggle,
} from "./external-providers";
const ANTHROPIC_DATED_SNAPSHOT_SUFFIX = /-\d{8}$/;
const OPENAI_DEPRECATED_MODELS = new Set(["gpt-5.3"]);
const OPENROUTER_EXCLUDED_MODELS = new Set([
"google/chirp-3",
"kwaivgi/kling-v3.0-pro",
"openai/whisper-1",
"openai/gpt-4o-mini-transcribe",
"recraft/recraft-v4-pro",
]);
function normalizeUrl(input: string): string {
return input.trim().replace(/\/+$/, "");
}
export function resolveUiProviderTypeFromConfig(
configProviderType: string,
configDisplayName: string | null | undefined,
configBaseUrl: string | null | undefined,
registryRows: ProviderRegistryEntry[],
existingProviderType: string | undefined,
): string {
if (existingProviderType && isCustomProviderType(existingProviderType)) {
return existingProviderType;
}
if (configProviderType !== CUSTOM_BACKEND_PROVIDER_TYPE) {
return configProviderType;
}
const displayName = (configDisplayName ?? "").trim().toLowerCase();
const matchingCustomPreset = CUSTOM_PROVIDER_PRESETS.find(
(preset) => preset.displayName.toLowerCase() === displayName,
);
if (matchingCustomPreset) {
return matchingCustomPreset.providerType;
}
const openAiRegistry = registryRows.find(
(entry) => entry.provider_type === CUSTOM_BACKEND_PROVIDER_TYPE,
);
if (!openAiRegistry) {
return configProviderType;
}
const openAiDisplayName = openAiRegistry.display_name.trim().toLowerCase();
if (displayName.length > 0 && displayName !== openAiDisplayName) {
return LEGACY_CUSTOM_PROVIDER_TYPE;
}
const configUrl = normalizeUrl(configBaseUrl ?? "");
const defaultUrl = normalizeUrl(openAiRegistry.base_url ?? "");
if (configUrl.length > 0 && configUrl !== defaultUrl) {
return LEGACY_CUSTOM_PROVIDER_TYPE;
}
return configProviderType;
}
export function pruneProviderModelIds(
providerType: string,
modelIds: string[],
): string[] {
if (providerType === "anthropic") {
return modelIds.filter((id) => !ANTHROPIC_DATED_SNAPSHOT_SUFFIX.test(id));
}
if (providerType === "openai") {
return modelIds.filter((id) => !OPENAI_DEPRECATED_MODELS.has(id));
}
if (providerType === "openrouter") {
return modelIds.filter((id) => !OPENROUTER_EXCLUDED_MODELS.has(id));
}
return modelIds;
}
/** Carry browser-local provider knobs through a backend sync rebuild. */
export function mergeLocalProviderOptions(
existing: ExternalProviderConfig | undefined,
synced: ExternalProviderConfig,
): ExternalProviderConfig {
if (!existing) {
return synced;
}
const providerType = synced.providerType;
return {
...synced,
enablePromptCaching: supportsProviderPromptCaching(providerType)
? (existing.enablePromptCaching ?? synced.enablePromptCaching ?? true)
: undefined,
promptCacheTtl:
supportsProviderPromptCacheTtl(providerType) &&
isPromptCacheTtl(existing.promptCacheTtl)
? existing.promptCacheTtl
: synced.promptCacheTtl,
isReasoningModel: supportsProviderReasoningToggle(providerType)
? (existing.isReasoningModel ?? synced.isReasoningModel)
: undefined,
openaiContainerTtlMinutes:
providerType === "openai" &&
typeof existing.openaiContainerTtlMinutes === "number" &&
existing.openaiContainerTtlMinutes >= 1
? Math.min(existing.openaiContainerTtlMinutes, 20)
: synced.openaiContainerTtlMinutes,
};
}
/** Merge enabled backend provider configs with local store state. */
export async function syncExternalProvidersFromBackend(
existingProviders: ExternalProviderConfig[],
): Promise<ExternalProviderConfig[]> {
const [registryRows, configRows] = await Promise.all([
listProviderRegistry(),
listProviderConfigs(),
]);
const existingById = new Map<string, ExternalProviderConfig>();
for (const provider of existingProviders) {
existingById.set(provider.id, provider);
}
const backfillTasks: Promise<unknown>[] = [];
const syncedProviders = configRows
.filter((config) => config.is_enabled)
.map((config) => {
const existing = existingById.get(config.id);
const uiProviderType = resolveUiProviderTypeFromConfig(
config.provider_type,
config.display_name,
config.base_url,
registryRows,
existing?.providerType,
);
const createdAt = Number.isFinite(Date.parse(config.created_at))
? Date.parse(config.created_at)
: Date.now();
const updatedAt = Number.isFinite(Date.parse(config.updated_at))
? Date.parse(config.updated_at)
: Date.now();
const registryEntry =
registryRows.find((entry) => entry.provider_type === uiProviderType) ??
registryRows.find((entry) => entry.provider_type === config.provider_type);
const defaultModels = pruneProviderModelIds(
uiProviderType,
registryEntry?.default_models ?? [],
);
const serverModels = pruneProviderModelIds(
uiProviderType,
config.models ?? [],
);
const serverAvailableModels = pruneProviderModelIds(
uiProviderType,
config.available_models ?? [],
);
const savedModels = existing?.models ?? [];
const savedAvailableModels = existing?.availableModels ?? [];
const resolvedModels = pruneProviderModelIds(
uiProviderType,
serverModels.length > 0
? serverModels
: savedModels.length > 0
? savedModels
: defaultModels,
);
const resolvedAvailableModels = pruneProviderModelIds(
uiProviderType,
serverAvailableModels.length > 0
? serverAvailableModels
: savedAvailableModels.length > 0
? savedAvailableModels
: defaultModels,
);
const needsModelBackfill =
serverModels.length === 0 && savedModels.length > 0;
const needsAvailableBackfill =
serverAvailableModels.length === 0 && savedAvailableModels.length > 0;
if (needsModelBackfill || needsAvailableBackfill) {
backfillTasks.push(
updateProviderConfig(config.id, {
models: resolvedModels,
availableModels: resolvedAvailableModels,
}),
);
}
const synced: ExternalProviderConfig = {
id: config.id,
providerType: uiProviderType,
name: config.display_name,
baseUrl: config.base_url ?? "",
models: resolvedModels,
availableModels: resolvedAvailableModels,
enablePromptCaching: supportsProviderPromptCaching(uiProviderType)
? (existing?.enablePromptCaching ?? true)
: undefined,
isReasoningModel: supportsProviderReasoningToggle(uiProviderType)
? existing?.isReasoningModel === true
: undefined,
createdAt: existing?.createdAt ?? createdAt,
updatedAt,
};
return mergeLocalProviderOptions(existing, synced);
});
if (backfillTasks.length > 0) {
await Promise.allSettled(backfillTasks);
}
return syncedProviders;
}

View file

@ -188,7 +188,10 @@ export interface LoadModelResponse {
n_layers?: number | null;
/** Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not MoE. */
n_moe_layers?: number;
/** Effective GPU placement after fit-time narrowing. */
gpu_ids?: number[] | null;
/** User-requested GPU placement pool before fit-time narrowing. */
requested_gpu_ids?: number[] | null;
}
export interface UnloadModelRequest {
@ -240,7 +243,10 @@ export interface InferenceStatusResponse {
/** n_ctx the active GGUF load was invoked with (0 = Auto); re-seeds a
* Manual + Auto-layers context pin on hydration. Null for non-GGUF. */
requested_context_length?: number | null;
/** Effective GPU placement after fit-time narrowing. */
gpu_ids?: number[] | null;
/** User-requested GPU placement pool before fit-time narrowing. */
requested_gpu_ids?: number[] | null;
n_layers?: number | null;
/** Model's MoE expert-layer count (the n_cpu_moe ceiling); 0 if not MoE. */
n_moe_layers?: number;

View file

@ -205,26 +205,19 @@ export function ParamsSection(): ReactElement {
setCtxInput(String(store.contextLength));
}, [store.contextLength]);
// On Apple Silicon the MLX trainer supports a different optimizer set than
// the CUDA/bitsandbytes list, so offer the MLX names there.
// Apple Silicon (MLX) supports a different optimizer set than the CUDA list.
const isMac = platformDeviceType === "mac";
const optimizerOptions = isMac ? MLX_OPTIMIZER_OPTIONS : OPTIMIZER_OPTIONS;
// On Mac, the MLX backend normalizes every CUDA/bitsandbytes optimizer in
// OPTIMIZER_OPTIONS (including the shared default) to plain AdamW, so show
// AdamW for those to keep the control truthful and non-blank. Any other
// value -- an MLX optimizer the user picked, or an unrecognized/non-canonical
// imported one -- is shown as-is rather than mislabeled as AdamW, since the
// backend would run or reject it on its own terms. Non-Mac display unchanged.
// On Mac the MLX backend remaps CUDA optimizers to AdamW, so label those as
// AdamW; other values (MLX or imported) show as-is. Non-Mac unchanged.
const isCudaAliasOptimizer = OPTIMIZER_OPTIONS.some(
(o) => o.value === store.optimizerType,
);
const selectedOptimizer =
isMac && isCudaAliasOptimizer ? "adamw" : store.optimizerType;
// LoftQ is not supported on MLX (the backend rejects it), so clear a stale
// selection to lora on Apple Silicon -- whether persisted, applied from a
// model default, or imported -- so the backend never receives it.
// LoftQ is unsupported on MLX; clear a stale selection to lora on Apple Silicon.
const setLoraVariant = store.setLoraVariant;
useEffect(() => {
if (isMac && store.loraVariant === "loftq") {
@ -232,8 +225,7 @@ export function ParamsSection(): ReactElement {
}
}, [isMac, store.loraVariant, setLoraVariant]);
// Packing is not supported on MLX (the backend forces it off), so clear it on
// Apple Silicon -- the checkbox is disabled and the flag is never sent.
// Packing is unsupported on MLX; clear it on Apple Silicon (checkbox disabled).
const setPacking = store.setPacking;
useEffect(() => {
if (isMac && store.packing) {

View file

@ -189,9 +189,8 @@ export function ProgressSection({
const cfgLoraDropout = cfg?.loraDropout;
const cfgLoraVariant = cfg?.loraVariant;
// Mirror the training form: on Mac the CUDA/bitsandbytes optimizer names run
// as plain AdamW (the MLX backend normalizes them), so label them AdamW here
// too rather than by the requested, unnormalized name.
// Mirror the training form: on Mac the MLX backend runs CUDA optimizers as
// AdamW, so label them AdamW here too.
const effectiveOptimizer =
platformDeviceType === "mac" &&
OPTIMIZER_OPTIONS.some((o) => o.value === cfgOptimizerType)

View file

@ -132,15 +132,17 @@ function toGpuInfo(data: SystemInfoResponse | null): GpuInfo {
}
function toGpuDevices(data: SystemInfoResponse | null): SystemGpuDevice[] {
// Unpinnable configurations must hide every pick surface: XPU indices are
// torch-xpu ordinals no applicator speaks -- /load and /validate 400 picks,
// so the backend reports gpu.gguf_gpu_ids_supported and every gate keyed on
// physicalIndex (picker, persisted-pick reconcile) follows it. The device
// flavor lives on the TOP-LEVEL device_backend field; absent support info
// defaults to pinnable (older backend).
const pinnableBackend =
data?.device_backend !== "xpu" &&
data?.gpu?.gguf_gpu_ids_supported !== false;
// Unpinnable configurations must hide every pick surface: /load and /validate
// 400 picks the applicator can't place, so the backend reports
// gpu.gguf_gpu_ids_supported and every gate keyed on physicalIndex (picker,
// persisted-pick reconcile) follows it. Absent support info defaults to
// pinnable (older backend).
const picksAccepted = data?.gpu?.gguf_gpu_ids_supported !== false;
// The XPU ban is specific to torch-xpu PHYSICAL ordinals (no applicator speaks
// them). A Vulkan pick uses ggml ordinals (--device Vulkan<i>), which don't
// rely on torch-xpu, so a Vulkan build stays pinnable even on an XPU host --
// and the backend already reports gguf_gpu_ids_supported true there.
const pinnablePhysical = picksAccepted && data?.device_backend !== "xpu";
// These devices exist to drive GGUF loads, so when the backend reports the
// llama-server (Vulkan) inventory, that list is authoritative: it can see
// cards torch can't, its indices are the ggml ordinals /load pins with
@ -154,7 +156,7 @@ function toGpuDevices(data: SystemInfoResponse | null): SystemGpuDevice[] {
name: d.name ?? `GPU ${d.index}`,
memoryTotalGb: d.memory_total_gb ?? 0,
memoryFreeGb: d.vram_free_gb ?? 0,
physicalIndex: pinnableBackend && d.index_kind === "vulkan",
physicalIndex: picksAccepted && d.index_kind === "vulkan",
}));
}
return (data?.gpu?.devices ?? [])
@ -164,7 +166,7 @@ function toGpuDevices(data: SystemInfoResponse | null): SystemGpuDevice[] {
name: d.name ?? `GPU ${d.index}`,
memoryTotalGb: d.memory_total_gb ?? 0,
memoryFreeGb: d.vram_free_gb ?? 0,
physicalIndex: pinnableBackend && d.index_kind === "physical",
physicalIndex: pinnablePhysical && d.index_kind === "physical",
}));
}

View file

@ -1018,6 +1018,12 @@ html[data-chat-font] .aui-root {
.sidebar-row-action {
@apply absolute top-0 bottom-0 right-0 inline-flex cursor-pointer items-center justify-end pl-2 pr-1.5 opacity-0 pointer-events-none outline-none;
}
@media (pointer: coarse) {
/* Only chat rows reserve touch padding (#7276); other rows stay hover-revealed to avoid clipped labels. */
.sidebar-row-action.sidebar-touch-reveal {
@apply opacity-100 pointer-events-auto;
}
}
.sidebar-row-action[data-state="open"] {
@apply opacity-100 pointer-events-auto;
}

View file

@ -186,8 +186,24 @@ VALIDATION_MODEL_CACHE_FILENAME = "stories260K.gguf"
# in validate_prebuilt_choice. Disabled for now: the llama-server GPU forward pass
# JIT-compiles CUDA kernels on first load and stalls every install and update by
# minutes on Blackwell (sm_100). The check and the source-build fallback it triggers
# are kept intact -- set this to True to re-enable them.
# are kept intact -- set this to True, or set UNSLOTH_LLAMA_STAGED_VALIDATION=1, to
# re-enable them (#5854 gap 2).
_RUN_STAGED_PREBUILT_VALIDATION = False
def staged_validation_enabled() -> bool:
"""True when the expensive llama-server GPU smoke test should run.
Default off (Blackwell CUDA JIT stalls installs). Opt in via the module
constant or ``UNSLOTH_LLAMA_STAGED_VALIDATION`` (1/true/yes/on). Used by both
the prebuilt path and setup.sh's source-build post-check (#5854).
"""
if _RUN_STAGED_PREBUILT_VALIDATION:
return True
raw = os.environ.get("UNSLOTH_LLAMA_STAGED_VALIDATION", "").strip().lower()
return raw in ("1", "true", "yes", "on")
INSTALL_LOCK_TIMEOUT_SECONDS = 300
INSTALL_STAGING_ROOT_NAME = ".staging"
GITHUB_AUTH_HOSTS = {"api.github.com", "github.com"}
@ -5868,9 +5884,10 @@ def validate_prebuilt_choice(
# so they are always validated. For an approved bundle the sha256 manifest
# already proves integrity, so its runtime smoke test -- a cold CUDA-JIT pass
# costing minutes on Blackwell sm_100 -- is gated behind
# _RUN_STAGED_PREBUILT_VALIDATION, disabled for now. The check and the
# source-build fallback it triggers are kept intact; flip the flag to restore it.
if choice.expected_sha256 is None or _RUN_STAGED_PREBUILT_VALIDATION:
# staged_validation_enabled() (constant or UNSLOTH_LLAMA_STAGED_VALIDATION),
# disabled for now. The check and the source-build fallback it triggers are
# kept intact; flip the flag / env to restore it (#5854).
if choice.expected_sha256 is None or staged_validation_enabled():
validate_quantize(
quantize_path,
probe_path,
@ -5891,6 +5908,49 @@ def validate_prebuilt_choice(
return server_path, quantize_path
def validate_existing_install(
install_dir: Path,
*,
install_kind: str | None = None,
host: HostInfo | None = None,
) -> None:
"""Run the staged smoke test against an already-built llama.cpp tree (#5854).
Used by setup.sh after a GPU source build when ``UNSLOTH_LLAMA_STAGED_VALIDATION``
is set. Raises ``PrebuiltFallback`` on failure so the caller can retry CPU.
"""
host = host or detect_host()
bin_dir = install_dir / "build" / "bin"
server_name = "llama-server.exe" if host.is_windows else "llama-server"
quantize_name = "llama-quantize.exe" if host.is_windows else "llama-quantize"
server_path = bin_dir / server_name
quantize_path = bin_dir / quantize_name
if not server_path.is_file():
raise PrebuiltFallback(f"llama-server not found at {server_path}")
with tempfile.TemporaryDirectory(prefix = "unsloth-llama-source-validate-") as tmp:
work_dir = Path(tmp)
probe_path = work_dir / "stories260K.gguf"
quantized_path = work_dir / "stories260K-q4.gguf"
download_validation_model(probe_path, validation_model_cache_path(install_dir))
if quantize_path.is_file():
validate_quantize(
quantize_path,
probe_path,
quantized_path,
install_dir,
host,
)
validate_server(
server_path,
probe_path,
host,
install_dir,
install_kind = install_kind,
)
log(f"staged source-build validation succeeded for {install_dir}")
def validate_prebuilt_attempts(
attempts: Iterable[AssetChoice],
host: HostInfo,
@ -6345,6 +6405,24 @@ def parse_args() -> argparse.Namespace:
"fork). Use --output-format json."
),
)
resolve_group.add_argument(
"--validate-install",
metavar = "DIR",
help = (
"Run the staged llama-server smoke test against an existing build "
"tree (setup.sh source-build post-check, #5854). Exit 2 on failure. "
"Normally gated by UNSLOTH_LLAMA_STAGED_VALIDATION; this flag always "
"runs the check."
),
)
parser.add_argument(
"--install-kind",
default = None,
help = (
"Install kind for --validate-install GPU offload (e.g. linux-cuda, "
"linux-rocm, macos-arm64). When omitted, host detection decides."
),
)
parser.add_argument(
"--output-format",
choices = ("plain", "json"),
@ -6381,6 +6459,17 @@ def emit_resolver_output(payload: dict[str, Any], *, output_format: str) -> None
def main() -> int:
args = parse_args()
if args.validate_install is not None:
try:
validate_existing_install(
Path(args.validate_install),
install_kind = args.install_kind,
)
except PrebuiltFallback as exc:
print(str(exc), file = sys.stderr)
raise SystemExit(EXIT_FALLBACK) from exc
return EXIT_SUCCESS
if args.resolve_llama_tag is not None:
resolved = resolve_requested_llama_tag(
args.resolve_llama_tag,

View file

@ -335,6 +335,31 @@ def artifacts_for_host(
# ── Slim selection (paired with the installed llama.cpp ggml runtime) ──
def _llama_ggml_commit(tag: str) -> str | None:
"""The ggml commit a llama.cpp fork tag was built against. Fork tags are
"b<upstream_build>-mix-<ggml_commit>"; the ggml commit after "-mix-" fixes
the ggml ABI the slim whisper bundle links against, while the build number
only tracks upstream llama / fork PRs outside ggml. None when the tag has no
"-mix-" marker (then only an exact tag pairs)."""
marker = "-mix-"
idx = tag.rfind(marker)
end = idx + len(marker)
return tag[end:] if idx >= 0 and end < len(tag) else None
def llama_runtime_pairs(installed_tag: str, required_tag: Any) -> bool:
"""Whether an installed llama tag can back a slim bundle needing required_tag.
An exact tag always pairs; so does a shared ggml commit, since a newer llama
build with the same ggml ships an ABI-identical runtime. requires_ggml_sonames
stays the real per-file ABI gate."""
if not isinstance(required_tag, str):
return False
if installed_tag == required_tag:
return True
commit = _llama_ggml_commit(installed_tag)
return commit is not None and commit == _llama_ggml_commit(required_tag)
def slim_pairing_for_artifact(
artifact: dict[str, Any], host: HostInfo, backend: str
) -> tuple[Path, str] | None:
@ -348,10 +373,10 @@ def slim_pairing_for_artifact(
return None
llama_bin_dir, llama_tag, _profile = runtime
requires_tag = artifact.get("requires_llama_tag")
if not isinstance(requires_tag, str) or requires_tag != llama_tag:
if not llama_runtime_pairs(llama_tag, requires_tag):
log(
f"slim_selection: {asset} skipped: installed llama tag {llama_tag!r} "
f"!= required {requires_tag!r}"
f"does not pair with required {requires_tag!r}"
)
return None
sonames = artifact.get("requires_ggml_sonames")
@ -466,11 +491,10 @@ def _slim_release_incompatibility(manifest: dict[str, Any], host: HostInfo) -> s
for artifact in os_compatible
if isinstance(artifact.get("requires_llama_tag"), str)
}
if required_tags and installed_tag not in required_tags:
if required_tags and not any(llama_runtime_pairs(installed_tag, tag) for tag in required_tags):
required_tag = sorted(required_tags)[0]
return (
f"slim bundle requires llama.cpp {required_tag}; "
f"installed llama.cpp is {installed_tag}"
f"slim bundle requires llama.cpp {required_tag}; installed llama.cpp is {installed_tag}"
)
return None
@ -820,7 +844,7 @@ def selection_from_artifact(
# A slim selection carries its pairing so the install wiring and marker know
# which llama runtime provides the ggml libraries.
runtime = installed_llama_runtime()
if runtime is None or runtime[1] != artifact.get("requires_llama_tag"):
if runtime is None or not llama_runtime_pairs(runtime[1], artifact.get("requires_llama_tag")):
raise PrebuiltFallback(
"the paired llama.cpp runtime changed underneath the slim whisper selection"
)

View file

@ -251,6 +251,38 @@ _resolve_cuda_archs() {
printf '%s' "$_archs"
}
# Opt-in staged GPU smoke test after a source build (#5854 gap 2). Default off:
# llama-server's first GPU forward pass JIT-compiles CUDA kernels and stalls
# installs for minutes on Blackwell. Same env as install_llama_prebuilt.py.
_staged_validation_enabled() {
local _raw="${UNSLOTH_LLAMA_STAGED_VALIDATION:-}"
# Match install_llama_prebuilt.py staged_validation_enabled(): strip + lowercase.
_raw="$(printf '%s' "$_raw" | sed -e 's/^[[:space:]]*//' -e 's/[[:space:]]*$//' | tr '[:upper:]' '[:lower:]')"
case "$_raw" in
1|true|yes|on) return 0 ;;
*) return 1 ;;
esac
}
# Map the source-build GPU backend to install_llama_prebuilt --install-kind so
# validate_server enables --n-gpu-layers for the right backends.
_source_smoke_install_kind() {
if [ "${_TRY_METAL_CPU_FALLBACK:-false}" = true ]; then
printf '%s' "macos-arm64"
return 0
fi
case "${GPU_BACKEND:-}" in
cuda)
case "$(uname -m 2>/dev/null || true)" in
aarch64|arm64) printf '%s' "linux-arm64-cuda" ;;
*) printf '%s' "linux-cuda" ;;
esac
;;
rocm) printf '%s' "linux-rocm" ;;
*) printf '%s' "" ;;
esac
}
# Run a GPU probe under a 10s timeout when `timeout` is available so a wedged
# NVIDIA driver cannot hang setup; fall back to a bare call where it is not.
_setup_run_smi() {
@ -1900,6 +1932,37 @@ else
run_quiet_no_exit "build diffusion visual server" cmake --build "$_BUILD_TMP/build" --config Release --target llama-diffusion-gemma-visual-server -j"$NCPU" || true
fi
# Opt-in post-build GPU smoke test (#5854 gap 2). Default off (Blackwell
# CUDA JIT stalls). On failure, reuse the CPU fallback path so the user
# still gets a working llama-server. Runs before the install swap.
if [ "$BUILD_OK" = true ] && _staged_validation_enabled; then
_FB_LABEL="$(_gpu_fallback_label)"
_SMOKE_KIND="$(_source_smoke_install_kind)"
if [ -n "$_FB_LABEL" ]; then
_SMOKE_CMD=(
python "$SCRIPT_DIR/install_llama_prebuilt.py"
--validate-install "$_BUILD_TMP"
)
[ -n "$_SMOKE_KIND" ] && _SMOKE_CMD+=(--install-kind "$_SMOKE_KIND")
if ! run_quiet_no_exit "validate source llama.cpp" "${_SMOKE_CMD[@]}"; then
substep "$_FB_LABEL source build failed smoke test; retrying CPU build..." "$C_WARN"
_TRY_METAL_CPU_FALLBACK=false
rm -rf "$_BUILD_TMP/build"
if run_quiet_no_exit "cmake llama.cpp (cpu fallback)" cmake $CMAKE_GENERATOR_ARGS -S "$_BUILD_TMP" -B "$_BUILD_TMP/build" $CPU_FALLBACK_CMAKE_ARGS; then
_BUILD_DESC="building (CPU fallback after $_FB_LABEL smoke failed)"
GPU_BACKEND=""
run_quiet_no_exit "build llama-server (cpu fallback)" cmake --build "$_BUILD_TMP/build" --config Release --target llama-server -j"$NCPU" || BUILD_OK=false
if [ "$BUILD_OK" = true ]; then
run_quiet_no_exit "build llama-quantize (cpu fallback)" cmake --build "$_BUILD_TMP/build" --config Release --target llama-quantize -j"$NCPU" || true
run_quiet_no_exit "build diffusion visual server (cpu fallback)" cmake --build "$_BUILD_TMP/build" --config Release --target llama-diffusion-gemma-visual-server -j"$NCPU" || true
fi
else
BUILD_OK=false
fi
fi
fi
fi
# Swap only after build succeeds -- preserves existing install on failure
if [ "$BUILD_OK" = true ]; then
_assert_studio_owned_or_absent "$LLAMA_CPP_DIR" "llama.cpp install"

View file

@ -0,0 +1,129 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
"""torch / torchcodec ABI guardrails (unslothai/unsloth#7225)."""
from __future__ import annotations
import importlib.util
import re
import sys
import types
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
PYPROJECT = REPO_ROOT / "pyproject.toml"
IMPORT_FIXES_PATH = REPO_ROOT / "unsloth" / "import_fixes.py"
def _load_import_fixes_module():
spec = importlib.util.spec_from_file_location(
"unsloth_import_fixes_under_test",
IMPORT_FIXES_PATH,
)
assert spec and spec.loader
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)
return mod
def test_pyproject_declares_torch210_audio_extra_with_python_gate():
text = PYPROJECT.read_text(encoding = "utf-8")
assert "audio-torch210 = [" in text
assert "torchcodec>=0.10.0,<0.11.0" in text
assert "python_version >= '3.10'" in text
assert "audio-torch290 = [" in text
assert "audio-torch280 = [" in text
assert "\naudio = [" not in text
def _stub_torch(monkeypatch, version: str):
torch_mod = types.ModuleType("torch")
torch_mod.__version__ = version
monkeypatch.setitem(sys.modules, "torch", torch_mod)
def test_torch210_extras_bundle_audio_torch210():
text = PYPROJECT.read_text(encoding = "utf-8")
for extra in (
"cu128-torch2100",
"cu126-ampere-torch2100",
"rocm72-torch2100",
):
match = re.search(rf"^{extra} = \[(.*?)^\]", text, re.MULTILINE | re.DOTALL)
assert match is not None, extra
assert "unsloth[audio-torch210]" in match.group(1)
def test_torchcodec_matrix_matches_notebook_validator():
from scripts import notebook_validator as nv
fixes = _load_import_fixes_module()
assert fixes._TORCH_TORCHCODEC_MINORS == nv.TORCH_TORCHCODEC
def test_torchcodec_exclusive_upper_bound():
fixes = _load_import_fixes_module()
assert fixes._torchcodec_exclusive_upper("0.10") == "<0.11.0"
assert fixes._torchcodec_exclusive_upper("0.9") == "<0.10.0"
def test_torch290_rejects_torchcodec_07(monkeypatch):
import importlib.metadata
fixes = _load_import_fixes_module()
_stub_torch(monkeypatch, "2.9.0+cu128")
monkeypatch.setattr(importlib.metadata, "version", lambda _name: "0.7.0")
hint = fixes._torchcodec_version_mismatch_hint()
assert hint is not None
assert "audio-torch210" not in hint
def test_torch280_accepts_torchcodec_07(monkeypatch):
import importlib.metadata
fixes = _load_import_fixes_module()
_stub_torch(monkeypatch, "2.8.0+cu128")
monkeypatch.setattr(importlib.metadata, "version", lambda _name: "0.7.0")
assert fixes._torchcodec_version_mismatch_hint() is None
def test_torch210_rejects_torchcodec_011(monkeypatch):
import importlib.metadata
fixes = _load_import_fixes_module()
_stub_torch(monkeypatch, "2.10.0+cu128")
monkeypatch.setattr(
importlib.metadata,
"version",
lambda _name: "0.11.0",
)
hint = fixes._torchcodec_version_mismatch_hint()
assert hint is not None
assert "torchcodec 0.11.0" in hint
assert "audio-torch210" in hint
assert "<0.11.0" in hint
assert "<11.0" not in hint
def test_torch210_accepts_torchcodec_010(monkeypatch):
import importlib.metadata
fixes = _load_import_fixes_module()
_stub_torch(monkeypatch, "2.10.0+cu128")
monkeypatch.setattr(
importlib.metadata,
"version",
lambda _name: "0.10.0+cu128",
)
assert fixes._torchcodec_version_mismatch_hint() is None
def test_import_fixes_loads_on_python39_syntax():
"""Regression: module must import on 3.9 (postponed annotations for str | None)."""
fixes = _load_import_fixes_module()
assert callable(fixes._torchcodec_version_mismatch_hint)

View file

@ -12,6 +12,7 @@ sh "$TESTS_DIR/sh/test_mac_intel_compat.sh"
sh "$TESTS_DIR/sh/test_torch_constraint.sh"
sh "$TESTS_DIR/sh/test_nvcc_meets_llama_minimum.sh"
sh "$TESTS_DIR/sh/test_resolve_cuda_archs.sh"
sh "$TESTS_DIR/sh/test_staged_validation_enabled.sh"
sh "$TESTS_DIR/sh/test_strixhalo_wsl_reroute.sh"
sh "$TESTS_DIR/sh/test_uninstall_shared_icon.sh"
sh "$TESTS_DIR/sh/test_torch_flavor.sh"

View file

@ -0,0 +1,94 @@
#!/bin/bash
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
# Unit tests for install.sh's _apt_distro_description helper (#6207).
# The sudo Accept? prompt should name the detected distro and say packages come
# from official apt repos. Hermetic: extract the helper and rewrite
# /etc/os-release to per-test fixtures (same pattern as test_strixhalo_wsl_reroute.sh).
set -e
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
INSTALL_SH="$SCRIPT_DIR/../../install.sh"
PASS=0
FAIL=0
_TMP_ROOT=$(mktemp -d)
trap 'rm -rf "$_TMP_ROOT"' EXIT
assert_eq() {
_label="$1"; _expected="$2"; _actual="$3"
if [ "$_actual" = "$_expected" ]; then
echo " PASS: $_label"; PASS=$((PASS + 1))
else
echo " FAIL: $_label (expected '$_expected', got '$_actual')"; FAIL=$((FAIL + 1))
fi
}
assert_contains() {
_label="$1"; _hay="$2"; _needle="$3"
case "$_hay" in
*"$_needle"*) echo " PASS: $_label"; PASS=$((PASS + 1)) ;;
*) echo " FAIL: $_label (missing '$_needle' in: $_hay)"; FAIL=$((FAIL + 1)) ;;
esac
}
# Extract helper with /etc/os-release rewritten to $1.
build_func() {
_fix="$1"
_f=$(mktemp -p "$_TMP_ROOT")
sed -n '/^_apt_distro_description()/,/^}/p' "$INSTALL_SH" \
| sed -e "s#/etc/os-release#$_fix/os-release#g" \
> "$_f"
echo "$_f"
}
run_desc() {
_os="$1"
_d=$(mktemp -d -p "$_TMP_ROOT")
printf '%s\n' "$_os" > "$_d/os-release"
_f=$(build_func "$_d")
# shellcheck disable=SC1090
. "$_f"
_apt_distro_description
}
echo "=== _apt_distro_description ==="
assert_eq "ubuntu name+version debian-like" \
"Ubuntu 24.04 (debian-like)" \
"$(run_desc "$(printf 'NAME=\"Ubuntu\"\nVERSION_ID=\"24.04\"\nID=ubuntu\nID_LIKE=debian\n')")"
assert_eq "debian name+version debian-like" \
"Debian GNU/Linux 12 (debian-like)" \
"$(run_desc "$(printf 'NAME=\"Debian GNU/Linux\"\nVERSION_ID=\"12\"\nID=debian\n')")"
assert_eq "pretty_name fallback when name/version missing" \
"Linux Mint 22 (debian-like)" \
"$(run_desc "$(printf 'PRETTY_NAME=\"Linux Mint 22\"\nID=linuxmint\nID_LIKE=\"ubuntu debian\"\n')")"
# NAME alone (no VERSION_ID) — still prefer NAME over PRETTY_NAME.
assert_eq "name only" \
"Pop!_OS (debian-like)" \
"$(run_desc "$(printf 'NAME=\"Pop!_OS\"\nID=pop\nID_LIKE=\"ubuntu debian\"\n')")"
assert_eq "missing os-release file" \
"a debian-like system" \
"$(
_d=$(mktemp -d -p "$_TMP_ROOT")
_f=$(build_func "$_d")
# shellcheck disable=SC1090
. "$_f"
_apt_distro_description
)"
echo "=== _smart_apt_install prompt contract ==="
_smart=$(sed -n '/^_smart_apt_install()/,/^}/p' "$INSTALL_SH")
assert_contains "calls distro helper" "$_smart" '_apt_distro_description'
assert_contains "names detected distro" "$_smart" 'Detected ${_ad_desc}'
assert_contains "mentions apt-get" "$_smart" 'sudo apt-get'
assert_contains "mentions official repos" "$_smart" "official repositories"
assert_contains "rejects tarball worry" "$_smart" "not a third-party tarball"
echo ""
echo "Results: $PASS passed, $FAIL failed"
[ "$FAIL" -eq 0 ]

View file

@ -0,0 +1,116 @@
#!/bin/bash
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
# Unit tests for setup.sh staged-validation helpers (#5854 gap 2).
# Opt-in GPU smoke after a source build; default off (Blackwell JIT stall).
set -e
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
SETUP_SH="$SCRIPT_DIR/../../studio/setup.sh"
PASS=0
FAIL=0
_FUNC_FILE=$(mktemp)
{
sed -n '/^_staged_validation_enabled()/,/^}/p' "$SETUP_SH"
sed -n '/^_source_smoke_install_kind()/,/^}/p' "$SETUP_SH"
} > "$_FUNC_FILE"
# shellcheck disable=SC1090
. "$_FUNC_FILE"
rm -f "$_FUNC_FILE"
assert_eq() {
_label="$1"; _expected="$2"; _actual="$3"
if [ "$_actual" = "$_expected" ]; then
echo " PASS: $_label"; PASS=$((PASS + 1))
else
echo " FAIL: $_label (expected '$_expected', got '$_actual')"; FAIL=$((FAIL + 1))
fi
}
assert_rc() {
_label="$1"; _expected="$2"
shift 2
set +e
"$@" >/dev/null 2>&1
_rc=$?
set -e
if [ "$_rc" -eq "$_expected" ]; then
echo " PASS: $_label"; PASS=$((PASS + 1))
else
echo " FAIL: $_label (expected rc $_expected, got $_rc)"; FAIL=$((FAIL + 1))
fi
}
echo "=== _staged_validation_enabled ==="
unset UNSLOTH_LLAMA_STAGED_VALIDATION
assert_rc "default off" 1 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=0
assert_rc "0 is off" 1 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=1
assert_rc "1 is on" 0 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=true
assert_rc "true is on" 0 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=yes
assert_rc "yes is on" 0 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=on
assert_rc "on is on" 0 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=True
assert_rc "True is on" 0 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=' yes '
assert_rc "whitespace yes is on" 0 _staged_validation_enabled
UNSLOTH_LLAMA_STAGED_VALIDATION=maybe
assert_rc "maybe is off" 1 _staged_validation_enabled
unset UNSLOTH_LLAMA_STAGED_VALIDATION
echo "=== _source_smoke_install_kind ==="
_TRY_METAL_CPU_FALLBACK=true
GPU_BACKEND=""
assert_eq "metal" "macos-arm64" "$(_source_smoke_install_kind)"
_TRY_METAL_CPU_FALLBACK=false
GPU_BACKEND=cuda
_kind="$(_source_smoke_install_kind)"
case "$(uname -m)" in
aarch64|arm64) assert_eq "cuda arm" "linux-arm64-cuda" "$_kind" ;;
*) assert_eq "cuda x86" "linux-cuda" "$_kind" ;;
esac
GPU_BACKEND=rocm
assert_eq "rocm" "linux-rocm" "$(_source_smoke_install_kind)"
GPU_BACKEND=""
assert_eq "cpu empty" "" "$(_source_smoke_install_kind)"
echo "=== setup.sh source smoke contract ==="
assert_contains() {
_label="$1"; _hay="$2"; _needle="$3"
case "$_hay" in
*"$_needle"*) echo " PASS: $_label"; PASS=$((PASS + 1)) ;;
*) echo " FAIL: $_label (missing '$_needle')"; FAIL=$((FAIL + 1)) ;;
esac
}
_src=$(cat "$SETUP_SH")
assert_contains "env gate present" "$_src" "UNSLOTH_LLAMA_STAGED_VALIDATION"
assert_contains "calls validate-install" "$_src" "--validate-install"
assert_contains "smoke fail retries CPU" "$_src" "source build failed smoke test; retrying CPU build"
# Smoke must run before the install swap.
_smoke_pos=$(printf '%s' "$_src" | awk '/validate source llama.cpp/{print NR; exit}')
_swap_pos=$(printf '%s' "$_src" | awk '/mv "\$_BUILD_TMP" "\$LLAMA_CPP_DIR"/{print NR; exit}')
if [ -n "$_smoke_pos" ] && [ -n "$_swap_pos" ] && [ "$_smoke_pos" -lt "$_swap_pos" ]; then
echo " PASS: smoke before install swap"; PASS=$((PASS + 1))
else
echo " FAIL: smoke before install swap (smoke=$_smoke_pos swap=$_swap_pos)"; FAIL=$((FAIL + 1))
fi
echo ""
echo "Results: $PASS passed, $FAIL failed"
[ "$FAIL" -eq 0 ]

View file

@ -3324,6 +3324,66 @@ def test_validate_prebuilt_choice_approved_validation_runs_when_flag_enabled(tmp
assert calls == {"quantize": 1, "server": 1}
def test_staged_validation_enabled_default_off(monkeypatch):
monkeypatch.setattr(INSTALL_LLAMA_PREBUILT, "_RUN_STAGED_PREBUILT_VALIDATION", False)
monkeypatch.delenv("UNSLOTH_LLAMA_STAGED_VALIDATION", raising = False)
assert INSTALL_LLAMA_PREBUILT.staged_validation_enabled() is False
@pytest.mark.parametrize("value", ["1", "true", "YES", "on"])
def test_staged_validation_enabled_env_opt_in(monkeypatch, value):
monkeypatch.setattr(INSTALL_LLAMA_PREBUILT, "_RUN_STAGED_PREBUILT_VALIDATION", False)
monkeypatch.setenv("UNSLOTH_LLAMA_STAGED_VALIDATION", value)
assert INSTALL_LLAMA_PREBUILT.staged_validation_enabled() is True
def test_validate_prebuilt_choice_approved_validation_runs_when_env_enabled(tmp_path, monkeypatch):
monkeypatch.setattr(INSTALL_LLAMA_PREBUILT, "_RUN_STAGED_PREBUILT_VALIDATION", False)
monkeypatch.setenv("UNSLOTH_LLAMA_STAGED_VALIDATION", "1")
calls = _run_validate_prebuilt_choice(monkeypatch, tmp_path, expected_sha256 = "ab" * 32)
assert calls == {"quantize": 1, "server": 1}
def test_validate_existing_install_runs_server_smoke(tmp_path, monkeypatch):
# setup.sh --validate-install path: exercise smoke helpers without a real GPU.
install_dir = tmp_path / "llama.cpp"
bin_dir = install_dir / "build" / "bin"
bin_dir.mkdir(parents = True)
(bin_dir / "llama-server").write_text("#!/bin/sh\n", encoding = "utf-8")
(bin_dir / "llama-quantize").write_text("#!/bin/sh\n", encoding = "utf-8")
calls: dict[str, int] = {"quantize": 0, "server": 0, "download": 0}
monkeypatch.setattr(
INSTALL_LLAMA_PREBUILT,
"download_validation_model",
lambda path, cache = None: calls.__setitem__("download", calls["download"] + 1),
)
monkeypatch.setattr(
INSTALL_LLAMA_PREBUILT,
"validate_quantize",
lambda *a, **k: calls.__setitem__("quantize", calls["quantize"] + 1),
)
monkeypatch.setattr(
INSTALL_LLAMA_PREBUILT,
"validate_server",
lambda *a, **k: calls.__setitem__("server", calls["server"] + 1),
)
monkeypatch.setattr(
INSTALL_LLAMA_PREBUILT,
"detect_host",
lambda: linux_host(),
)
INSTALL_LLAMA_PREBUILT.validate_existing_install(install_dir, install_kind = "linux-cuda")
assert calls == {"quantize": 1, "server": 1, "download": 1}
def test_validate_existing_install_missing_server_raises(tmp_path, monkeypatch):
monkeypatch.setattr(INSTALL_LLAMA_PREBUILT, "detect_host", lambda: linux_host())
with pytest.raises(INSTALL_LLAMA_PREBUILT.PrebuiltFallback, match = "llama-server not found"):
INSTALL_LLAMA_PREBUILT.validate_existing_install(tmp_path / "missing")
def test_diffusion_visual_server_uses_approved_checksum_download(monkeypatch, tmp_path: Path):
asset_name = "llama-diffusion-gemma-visual-server-linux-x64"
expected_sha = "a" * 64

View file

@ -441,8 +441,9 @@ def test_main_forwards_requested_whisper_tags(tmp_path, monkeypatch):
monkeypatch.setattr(
M,
"resolve_prebuilt",
lambda host, **kwargs: seen.update(kwargs)
or {"prebuilt_available": False, "repo": "unslothai/whisper.cpp"},
lambda host, **kwargs: (
seen.update(kwargs) or {"prebuilt_available": False, "repo": "unslothai/whisper.cpp"}
),
)
assert M.main(["--resolve-prebuilt", "v1.8.0", "--output-format", "json"]) == 0
assert seen["whisper_tag"] == "v1.8.0"
@ -807,6 +808,50 @@ def test_slim_release_tag_skew_has_distinct_compatibility_error(tmp_path, monkey
M.select_artifact_with_fallback(manifest, _cuda_host(), "cuda")
# A newer llama build that keeps the same ggml commit as SLIM_LLAMA_TAG.
NEWER_LLAMA_TAG = "b10079-mix-fb3d4ca"
@pytest.mark.parametrize(
"installed,required,pairs",
[
(SLIM_LLAMA_TAG, SLIM_LLAMA_TAG, True), # exact tag
(NEWER_LLAMA_TAG, SLIM_LLAMA_TAG, True), # newer build, same ggml commit
("b10069-mix-0000000", SLIM_LLAMA_TAG, False), # same build, different ggml
(SLIM_LLAMA_TAG, None, False), # no requirement recorded
("b10069", "b10069", True), # tag without -mix-, exact only
("b10070", "b10069", False), # tag without -mix-, no shared key
],
)
def test_llama_runtime_pairs_keys_on_ggml_commit(installed, required, pairs):
assert M.llama_runtime_pairs(installed, required) is pairs
def test_slim_pairs_across_llama_build_bump_with_same_ggml(tmp_path, monkeypatch):
# The live failure: the llama installer advances to a newer build that keeps
# the same ggml commit, so the slim bundle's paired runtime is ABI-identical
# and must still select rather than degrade to CPU or report unavailable.
bin_dir = _fake_llama_bin(tmp_path)
monkeypatch.setattr(
M, "installed_llama_runtime", lambda: (bin_dir, NEWER_LLAMA_TAG, "cuda13-newer")
)
artifact, backend, used_fallback = M.select_artifact_with_fallback(
_slim_manifest(), _cuda_host(), "cuda"
)
assert artifact["asset"] == SLIM_ASSET
assert backend == "cuda" and used_fallback is False
def test_slim_build_bump_same_ggml_is_not_a_compatibility_error(tmp_path, monkeypatch):
# A same-ggml build bump must not surface as a release incompatibility (the
# update path reports that as unavailable); only a real ggml skew does.
bin_dir = _fake_llama_bin(tmp_path)
monkeypatch.setattr(
M, "installed_llama_runtime", lambda: (bin_dir, NEWER_LLAMA_TAG, "cuda13-newer")
)
assert M._slim_release_incompatibility(_slim_manifest(), _cuda_host()) is None
def test_link_ggml_runtime_hardlinks_every_ggml_library(tmp_path):
bin_dir = _fake_llama_bin(tmp_path)
whisper_bin = tmp_path / "whisper.cpp" / "build" / "bin"

View file

@ -36,6 +36,9 @@ ART_DIR = os.environ.get("PW_ART_DIR", "logs/playwright_extra")
ART = Path(ART_DIR)
ART.mkdir(parents = True, exist_ok = True)
STRICT = os.environ.get("STUDIO_UI_STRICT", "0") == "1"
# The Voice-picker media-access crash is specific to headless Chromium on macos-14; only there do we
# downgrade a renderer crash to a warning. Linux/Windows strict smoke jobs keep hard crash coverage.
MACOS_RUNNER = os.environ.get("RUNNER_OS", "").lower() == "macos" or sys.platform == "darwin"
# Longer turn timeout: gemma-3-270m CPU inference is 3-5x slower on macos-14 runners.
TURN_TIMEOUT_MS = int(os.environ.get("STUDIO_UI_TURN_TIMEOUT_MS", "180000"))
WALL_TIMEOUT_S = float(os.environ.get("STUDIO_UI_WALL_TIMEOUT_S", "720"))
@ -71,6 +74,18 @@ def runtime_warn(m: str) -> None:
info(f"WARN (runtime): {m}")
def page_crashed(pg, exc: Exception) -> bool:
"""True when the browser/page/context died (a macos-14 renderer crash) rather than a live-page
assertion failing -- so the caller can downgrade CI-environment flakiness to a runtime warning."""
try:
if pg.is_closed():
return True
except Exception:
return True
msg = str(exc).lower()
return "has been closed" in msg or "target closed" in msg or "crash" in msg
with sync_playwright() as p:
_watchdog = install_wall_clock_watchdog(
WALL_TIMEOUT_S,
@ -544,13 +559,17 @@ with sync_playwright() as p:
if voice_tab.count() == 0:
fail("Voice settings tab not found")
else:
voice_tab.click()
page.get_by_label("Dictation engine").click()
page.get_by_role("option", name = "Local transcription").click()
page.get_by_label("Speech recognition model").click()
page.get_by_placeholder("Search model").fill("whisper")
results = page.get_by_test_id("stt-model-results")
# The dictation-engine dropdown touches a media-access path that can crash headless
# Chromium on macos-14 (CheckMediaAccessPermission). A resulting TargetClosedError is CI
# flakiness there, not a product bug, so on macOS a crash is a runtime warning + page
# recovery; on Linux/Windows a crash and any live-page failure stay a hard fail.
try:
voice_tab.click()
page.get_by_label("Dictation engine").click()
page.get_by_role("option", name = "Local transcription").click()
page.get_by_label("Speech recognition model").click()
page.get_by_placeholder("Search model").fill("whisper")
results = page.get_by_test_id("stt-model-results")
page.wait_for_function(
"""() => {
const node = document.querySelector('[data-testid="stt-model-results"]');
@ -569,10 +588,23 @@ with sync_playwright() as p:
)
info("OK Voice model picker mouse wheel changed scrollTop")
except Exception as exc:
fail(f"Voice model picker did not wheel-scroll: {exc!r}")
shoot("10-settings-tabs-visited")
page.keyboard.press("Escape")
page.wait_for_timeout(300)
if page_crashed(page, exc) and MACOS_RUNNER:
runtime_warn(f"Voice model picker aborted (browser/page unstable): {exc!r}")
page = recover_or_replace_page(
page,
ctx,
default_timeout_ms = 60_000,
info = lambda m: info(f"recovery: {m}"),
)
else:
fail(f"Voice model picker did not wheel-scroll: {exc!r}")
# When the crash closed the context/browser (not just the page), recover_or_replace_page
# cannot mint a replacement and hands back the closed page; skip the cosmetic teardown rather
# than re-raise TargetClosedError on it. is_closed() is a local check and never raises.
if not page.is_closed():
shoot("10-settings-tabs-visited")
page.keyboard.press("Escape")
page.wait_for_timeout(300)
info(f"visited Settings tabs: {seen_tabs}")
if not seen_tabs:
soft_fail("no Settings tabs were visitable")
@ -591,4 +623,7 @@ with sync_playwright() as p:
sys.exit(1)
info("PASS extra UI flow")
_watchdog.cancel()
browser.close()
try:
browser.close()
except Exception:
pass # a crashed browser may already be gone; never fail teardown after PASS

View file

@ -14,6 +14,7 @@ DATA_TAB = FRONTEND / "features/settings/tabs/data-tab.tsx"
PROMPT_STORAGE = FRONTEND / "features/chat/prompt-storage/prompt-storage-dialog.tsx"
APP_SIDEBAR = FRONTEND / "components/app-sidebar.tsx"
INDEX_CSS = FRONTEND / "index.css"
THREAD = FRONTEND / "components/assistant-ui/thread.tsx"
THREAD_SIDEBAR = FRONTEND / "features/chat/thread-sidebar.tsx"
SHARED_COMPOSER = FRONTEND / "features/chat/shared-composer.tsx"
@ -129,3 +130,26 @@ def test_expanded_titlebar_button_and_corner_match_sidebar_edge():
'className="pointer-events-none absolute top-full size-3 -translate-x-px rounded-tl-[12px] border-l border-t border-sidebar-border bg-background"'
in source
)
def test_chat_sidebar_row_actions_visible_on_coarse_pointers():
"""unslothai/unsloth#7276: Recents chat kebab must be tappable on iPad."""
sidebar_source = APP_SIDEBAR.read_text(encoding = "utf-8")
css_source = INDEX_CSS.read_text(encoding = "utf-8")
assert "renderChatSidebarItem" in sidebar_source
block = sidebar_source.split("function renderChatSidebarItem", 1)[1].split("\n function ", 1)[
0
]
assert "[@media(pointer:coarse)]:pr-10" in block
assert "sidebar-touch-reveal" in block
# Coarse-pointer visibility must come after .sidebar-row-action { opacity-0 }.
coarse_idx = css_source.index("@media (pointer: coarse)")
base_idx = css_source.index(".sidebar-row-action {")
assert coarse_idx > base_idx
coarse_block = css_source[coarse_idx : coarse_idx + 280]
assert "sidebar-touch-reveal" in coarse_block
assert "opacity-100" in coarse_block
assert "pointer-events-auto" in coarse_block
# Must not reveal every sidebar-row-action (project/run/nav rows lack padding).
assert ".sidebar-row-action {\n\t\t\t@apply opacity-100" not in coarse_block
assert ".sidebar-row-action.sidebar-touch-reveal" in coarse_block

View file

@ -151,6 +151,19 @@ def test_active_model_config_round_trips_gpu_fields():
assert "export function gpuFieldsSignature" in shared
def test_gpu_picker_round_trips_requested_pool_not_fitted_subset():
"""A GGUF fit may narrow [0, 1] to [0], but load/status hydration must keep
[0, 1] as the editable pool so a later reload can grow back onto GPU 1."""
types = _read("features/chat/types/api.ts")
assert types.count("requested_gpu_ids?: number[] | null") >= 2
store = _read("features/chat/stores/chat-runtime-store.ts")
assert "resp.requested_gpu_ids ?? resp.gpu_ids ?? null" in store
status = _read("features/chat/lib/apply-inference-status-to-store.ts")
assert "status.requested_gpu_ids ?? status.gpu_ids ?? null" in status
def test_compare_load_uses_each_models_gpu_config():
src = _read("features/chat/shared-composer.tsx")
assert "ownConfig.gpuMemoryMode ?? compareLoadKnobs.gpuMemoryMode" in src

View file

@ -0,0 +1,61 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Static contracts for remote connection model persistence (#7281)."""
from pathlib import Path
REPO = Path(__file__).resolve().parents[2]
FRONTEND = REPO / "studio/frontend/src"
PROVIDERS_API = FRONTEND / "features/chat/api/providers-api.ts"
SYNC_PROVIDERS = FRONTEND / "features/chat/sync-external-providers.ts"
CHAT_PAGE = FRONTEND / "features/chat/chat-page.tsx"
PROVIDERS_DB = REPO / "studio/backend/storage/providers_db.py"
PROVIDERS_MODELS = REPO / "studio/backend/models/providers.py"
def test_providers_db_stores_model_json_columns():
source = PROVIDERS_DB.read_text(encoding = "utf-8")
assert "models_json" in source
assert "available_models_json" in source
assert "ALTER TABLE llm_providers ADD COLUMN models_json" in source
def test_provider_api_schemas_expose_models():
source = PROVIDERS_MODELS.read_text(encoding = "utf-8")
assert "models: list[str]" in source
assert "available_models: list[str]" in source
def test_frontend_sync_prefers_server_models_on_remote_clients():
source = SYNC_PROVIDERS.read_text(encoding = "utf-8")
assert "config.models" in source
assert "config.available_models" in source
assert "serverModels.length > 0" in source
def test_frontend_sync_backfills_local_models_to_backend():
source = SYNC_PROVIDERS.read_text(encoding = "utf-8")
assert "updateProviderConfig" in source
assert "needsModelBackfill" in source
assert "Promise.allSettled(backfillTasks)" in source
def test_frontend_sync_preserves_local_provider_options():
source = SYNC_PROVIDERS.read_text(encoding = "utf-8")
assert "mergeLocalProviderOptions" in source
assert "promptCacheTtl" in source
assert "openaiContainerTtlMinutes" in source
def test_chat_page_hydrates_connections_on_startup():
source = CHAT_PAGE.read_text(encoding = "utf-8")
assert "syncExternalProvidersFromBackend" in source
assert "await hydratePersistedSettings()" in source
def test_providers_api_sends_models_to_backend():
source = PROVIDERS_API.read_text(encoding = "utf-8")
assert "available_models: payload.availableModels" in source
assert "models: payload.models" in source

View file

@ -12,6 +12,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from __future__ import annotations
import os
import importlib.abc
import importlib.machinery
@ -1525,6 +1527,59 @@ def patch_torchcodec_audio_decoder():
pass
# torch.minor -> compatible torchcodec.minor strings (see notebook_validator.py).
_TORCH_TORCHCODEC_MINORS: dict[str, set[str]] = {
"2.10": {"0.10"},
"2.9": {"0.8", "0.9"},
"2.8": {"0.6", "0.7"},
"2.7": {"0.3", "0.4", "0.5"},
"2.6": {"0.2", "0.3"},
"2.5": {"0.1", "0.2"},
}
def _torchcodec_exclusive_upper(pin: str) -> str:
"""Next torchcodec minor as an exclusive pip upper bound (0.10 -> <0.11.0)."""
major, minor = pin.split(".", 1)
return f"<{major}.{int(minor) + 1}.0"
def _torchcodec_version_mismatch_hint() -> str | None:
"""Return a user-facing hint when installed torchcodec mismatches torch."""
try:
import importlib.metadata as importlib_metadata
import torch
from packaging.version import Version
torchcodec_version = importlib_metadata.version("torchcodec")
except Exception:
return None
def _minor(version: str) -> str:
parts = Version(version.split("+", 1)[0]).release
return ".".join(str(p) for p in parts[:2])
try:
torch_minor = _minor(torch.__version__)
codec_minor = _minor(torchcodec_version)
except Exception:
# Non-PEP440 version strings must never break `import unsloth`.
return None
allowed = _TORCH_TORCHCODEC_MINORS.get(torch_minor)
if allowed is None or codec_minor in allowed:
return None
pin = sorted(allowed)[-1]
upper = _torchcodec_exclusive_upper(pin)
install_hint = f"`pip install 'torchcodec>={pin},{upper}'`"
if torch_minor == "2.10":
install_hint += " or `pip install 'unsloth[audio-torch210]'`"
return (
f"torchcodec {torchcodec_version} is incompatible with torch {torch.__version__}; "
f"install a matching build with {install_hint}."
)
def disable_torchcodec_if_broken():
"""Make broken torchcodec behave as if uninstalled (#5446).
@ -1533,6 +1588,15 @@ def disable_torchcodec_if_broken():
flags and seat a sys.modules sentinel so downstream imports fall through
their existing except ImportError handlers cleanly.
"""
mismatch_hint = _torchcodec_version_mismatch_hint()
if mismatch_hint is not None:
try:
import warnings
warnings.warn(mismatch_hint, stacklevel = 2)
except Exception:
# Warning filters promoted to errors must not abort the disable
# fallback below (e.g. PYTHONWARNINGS=error, pytest -W error).
pass
try:
import importlib.util
if importlib.util.find_spec("torchcodec") is None:

View file

@ -19,6 +19,8 @@ from unsloth_cli.commands.start import (
_CLAUDE_ENV_UNSET,
_SUBAGENT_DESCRIPTION,
_SUBAGENT_INSTRUCTIONS,
_SUBAGENT_PLAN_DESCRIPTION,
_SUBAGENT_PLAN_INSTRUCTIONS,
_claude_flags,
_claude_local_env,
_wsl_shim_env,
@ -113,7 +115,11 @@ def _stop_child(process: subprocess.Popen) -> None:
pass
def run_local_agent(task: str, cancel_event: threading.Event | None = None) -> str:
def run_local_agent(
task: str,
cancel_event: threading.Event | None = None,
read_only: bool = False,
) -> str:
base = _required_env("UNSLOTH_CLAUDE_SUBAGENT_BASE_URL")
key = _required_env("UNSLOTH_CLAUDE_SUBAGENT_API_KEY")
model = _required_env("UNSLOTH_CLAUDE_SUBAGENT_MODEL")
@ -135,16 +141,20 @@ def run_local_agent(task: str, cancel_event: threading.Event | None = None) -> s
*_claude_flags(model),
"--permission-mode",
(
"bypassPermissions"
if os.environ.get("UNSLOTH_CLAUDE_SUBAGENT_BYPASS_PERMISSIONS") == "1"
else "acceptEdits"
"plan"
if read_only
else (
"bypassPermissions"
if os.environ.get("UNSLOTH_CLAUDE_SUBAGENT_BYPASS_PERMISSIONS") == "1"
else "acceptEdits"
)
),
"--print",
"--output-format",
"json",
"--no-session-persistence",
"--append-system-prompt",
_SUBAGENT_INSTRUCTIONS,
_SUBAGENT_PLAN_INSTRUCTIONS if read_only else _SUBAGENT_INSTRUCTIONS,
f"Task: {task}",
]
bridged, wsl_names = _wsl_shim_env(command, local_env, _CLAUDE_ENV_UNSET)
@ -196,7 +206,15 @@ def run_local_agent(task: str, cancel_event: threading.Event | None = None) -> s
return _result_text(stdout)
def _response(request: dict, run_agent: Callable[[str], str] = run_local_agent) -> dict | None:
def _response(
request: dict,
run_agent: Callable[[str], str] = run_local_agent,
tool_name: str = "unsloth_agent",
tool_description: str | None = None,
run_read_only_agent: Callable[[str], str] | None = None,
read_only_tool_name: str | None = None,
instructions: str | None = None,
) -> dict | None:
request_id = request.get("id")
method = request.get("method")
if request_id is None:
@ -208,40 +226,55 @@ def _response(request: dict, run_agent: Callable[[str], str] = run_local_agent)
"capabilities": {"tools": {"listChanged": False}},
"serverInfo": {"name": "unsloth-local-agent", "version": "1.0.0"},
}
if instructions:
result["instructions"] = instructions
elif method == "ping":
result = {}
elif method == "tools/list":
result = {
"tools": [
{
"name": "unsloth_agent",
"title": "Unsloth local agent",
"description": _SUBAGENT_DESCRIPTION,
"inputSchema": {
"type": "object",
"properties": {
"task": {
"type": "string",
"description": "The complete task for the local Unsloth agent.",
}
},
"required": ["task"],
"additionalProperties": False,
def tool_definition(name: str, description: str, read_only: bool) -> dict:
return {
"name": name,
"title": "Unsloth local plan agent" if read_only else "Unsloth local agent",
"description": description,
"inputSchema": {
"type": "object",
"properties": {
"task": {
"type": "string",
"description": "The complete task for the local Unsloth agent.",
}
},
"annotations": {
"readOnlyHint": False,
"destructiveHint": True,
"idempotentHint": False,
"openWorldHint": True,
},
"_meta": {"anthropic/maxResultSizeChars": _MAX_RESULT_CHARACTERS},
}
]
}
"required": ["task"],
"additionalProperties": False,
},
"annotations": {
"readOnlyHint": read_only,
"destructiveHint": not read_only,
"idempotentHint": read_only,
"openWorldHint": True,
},
"_meta": {"anthropic/maxResultSizeChars": _MAX_RESULT_CHARACTERS},
}
tools = [tool_definition(tool_name, tool_description or _SUBAGENT_DESCRIPTION, False)]
if read_only_tool_name and run_read_only_agent:
tools.append(tool_definition(read_only_tool_name, _SUBAGENT_PLAN_DESCRIPTION, True))
result = {"tools": tools}
elif method == "tools/call":
params = request.get("params") or {}
arguments = params.get("arguments") or {}
task = arguments.get("task") if params.get("name") == "unsloth_agent" else None
requested_tool = params.get("name")
selected_agent = (
run_agent
if requested_tool == tool_name
else (
run_read_only_agent
if requested_tool == read_only_tool_name and run_read_only_agent
else None
)
)
task = arguments.get("task") if selected_agent else None
if not isinstance(task, str) or not task.strip():
result = {
"content": [{"type": "text", "text": "A non-empty task is required."}],
@ -249,7 +282,7 @@ def _response(request: dict, run_agent: Callable[[str], str] = run_local_agent)
}
else:
try:
text = run_agent(task.strip())
text = selected_agent(task.strip())
result = {"content": [{"type": "text", "text": text}], "isError": False}
except Exception as exc:
result = {
@ -269,6 +302,11 @@ def serve(
stdin: Any = sys.stdin,
stdout: Any = sys.stdout,
run_agent: Callable[[str, threading.Event], str] = run_local_agent,
tool_name: str = "unsloth_agent",
tool_description: str | None = None,
run_read_only_agent: Callable[[str, threading.Event], str] | None = None,
read_only_tool_name: str | None = None,
instructions: str | None = None,
) -> None:
active: dict[object, threading.Event] = {}
workers: list[threading.Thread] = []
@ -308,6 +346,15 @@ def serve(
response = _response(
request,
run_agent = lambda task: run_agent(task, cancel_event),
tool_name = tool_name,
tool_description = tool_description,
run_read_only_agent = (
(lambda task: run_read_only_agent(task, cancel_event))
if run_read_only_agent
else None
),
read_only_tool_name = read_only_tool_name,
instructions = instructions,
)
if not cancel_event.is_set():
send(response)
@ -343,7 +390,18 @@ def serve(
worker.start()
response = None
else:
response = _response(request)
response = _response(
request,
tool_name = tool_name,
tool_description = tool_description,
run_read_only_agent = (
(lambda task: run_read_only_agent(task, threading.Event()))
if run_read_only_agent
else None
),
read_only_tool_name = read_only_tool_name,
instructions = instructions,
)
except Exception as exc:
response = {
"jsonrpc": "2.0",
@ -362,5 +420,14 @@ def serve(
signal.signal(signum, handler)
def main() -> None:
serve(
run_read_only_agent = lambda task, cancel_event: run_local_agent(
task, cancel_event, read_only = True
),
read_only_tool_name = "unsloth_plan_agent",
)
if __name__ == "__main__":
serve()
main()

View file

@ -0,0 +1,171 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Small stdio MCP bridge from cloud Codex to an explicit local Codex child."""
from __future__ import annotations
import json
import os
import shutil
import subprocess
import sys
import threading
from pathlib import Path
from typing import Any
from unsloth_cli.claude_subagent_mcp import _bounded, _stop_child, serve
from unsloth_cli.commands.start import (
_CODEX_ENV_KEY,
_CODEX_ENV_UNSET,
_CODEX_PROFILE,
_CODEX_SUBAGENT_CONFIG_ENV,
_CODEX_SUBAGENT_MCP_TOOL,
_CODEX_SUBAGENT_TOOL_DESCRIPTION,
_CODEX_SUBAGENT_ROUTING_INSTRUCTIONS,
_SUBAGENT_INSTRUCTIONS,
_merge_wslenv,
_wsl_shim_env,
)
_CANCEL_POLL_SECONDS = 0.1
_SERVER_INSTRUCTIONS = _CODEX_SUBAGENT_ROUTING_INSTRUCTIONS
def _config() -> dict:
path = os.environ.get(_CODEX_SUBAGENT_CONFIG_ENV, "").strip()
if not path:
raise RuntimeError(f"Missing {_CODEX_SUBAGENT_CONFIG_ENV}.")
try:
config = json.loads(Path(path).read_text(encoding = "utf-8"))
except (OSError, ValueError) as exc:
raise RuntimeError("Could not read the local Codex agent configuration.") from exc
if not isinstance(config, dict):
raise RuntimeError("The local Codex agent configuration must be an object.")
for name in ("api_key", "codex_home"):
if not isinstance(config.get(name), str) or not config[name].strip():
raise RuntimeError(f"The local Codex agent configuration is missing {name}.")
return config
def _result_text(stdout: str) -> str:
messages = []
errors = []
for line in stdout.splitlines():
try:
event = json.loads(line)
except ValueError:
continue
if not isinstance(event, dict):
continue
item = event.get("item")
if (
event.get("type") == "item.completed"
and isinstance(item, dict)
and item.get("type") == "agent_message"
and isinstance(item.get("text"), str)
and item["text"].strip()
):
messages.append(item["text"].strip())
if event.get("type") in ("error", "turn.failed"):
detail = event.get("message") or event.get("error")
if isinstance(detail, dict):
detail = detail.get("message") or json.dumps(detail)
if detail:
errors.append(str(detail))
if errors:
raise RuntimeError(_bounded(errors[-1]))
if messages:
return _bounded(messages[-1])
raise RuntimeError("The local Codex agent returned no readable result.")
def run_local_agent(task: str, cancel_event: threading.Event | None = None) -> str:
config = _config()
executable = shutil.which("codex")
if executable is None:
raise RuntimeError("`codex` is not installed or is not on PATH.")
cancel_event = cancel_event or threading.Event()
if cancel_event.is_set():
raise RuntimeError("The local Codex agent was cancelled.")
permissions = (
["--dangerously-bypass-approvals-and-sandbox"]
if config.get("bypass_permissions") is True
else ["--sandbox", "workspace-write", "--ask-for-approval", "never"]
)
command = [
"codex",
"--oss",
"--profile",
_CODEX_PROFILE,
*permissions,
"exec",
"--ephemeral",
"--json",
"--skip-git-repo-check",
f"{_SUBAGENT_INSTRUCTIONS}\n\nTask: {task}",
]
local_env = {
_CODEX_ENV_KEY: config["api_key"],
"CODEX_HOME": config["codex_home"],
"CODEX_SQLITE_HOME": config["codex_home"],
}
bridged, wsl_names = _wsl_shim_env(command, local_env, _CODEX_ENV_UNSET)
child_env = dict(os.environ)
if wsl_names:
bridged = {**bridged, "PWD": os.getcwd()}
child_env["WSLENV"] = _merge_wslenv(child_env.get("WSLENV", ""), wsl_names)
for name in _CODEX_ENV_UNSET:
child_env[name] = ""
else:
for name in _CODEX_ENV_UNSET:
child_env.pop(name, None)
child_env.update(bridged)
popen_kwargs: dict[str, Any] = {
"cwd": os.getcwd(),
"env": child_env,
"stdin": subprocess.DEVNULL,
"stdout": subprocess.PIPE,
"stderr": subprocess.PIPE,
"text": True,
}
if os.name == "nt":
popen_kwargs["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
else:
popen_kwargs["start_new_session"] = True
process = subprocess.Popen([executable, *command[1:]], **popen_kwargs)
try:
while True:
try:
stdout, stderr = process.communicate(timeout = _CANCEL_POLL_SECONDS)
break
except subprocess.TimeoutExpired:
if cancel_event.is_set():
_stop_child(process)
raise RuntimeError("The local Codex agent was cancelled.")
except BaseException:
if process.poll() is None:
_stop_child(process)
raise
if process.returncode != 0:
detail = stderr.strip() or stdout.strip()
raise RuntimeError(
_bounded(detail) or f"Local Codex exited with code {process.returncode}."
)
return _result_text(stdout)
def main() -> None:
if len(sys.argv) > 1:
os.environ[_CODEX_SUBAGENT_CONFIG_ENV] = sys.argv[1]
serve(
run_agent = run_local_agent,
tool_name = _CODEX_SUBAGENT_MCP_TOOL,
tool_description = _CODEX_SUBAGENT_TOOL_DESCRIPTION,
instructions = _SERVER_INSTRUCTIONS,
)
if __name__ == "__main__":
main()

View file

@ -84,8 +84,33 @@ _SUBAGENT_INSTRUCTIONS = (
"use the available tools when useful, verify your work, and return a concise result to the "
"parent agent."
)
_SUBAGENT_PLAN_DESCRIPTION = (
"Read-only local coding subagent powered by Unsloth for planning and codebase research. "
"Use this local agent when Claude is in plan mode."
)
_SUBAGENT_PLAN_INSTRUCTIONS = (
"You are a read-only local coding subagent powered by Unsloth. Investigate the assigned "
"task with read-only tools, produce a concrete plan or answer, and return a concise result "
"to the parent agent. Do not modify files."
)
_CLAUDE_SUBAGENT_MCP_MODULE = "unsloth_cli.claude_subagent_mcp"
_CLAUDE_SUBAGENT_TOOL = "mcp__plugin_unsloth-local-agent_unsloth__unsloth_agent"
_CLAUDE_SUBAGENT_PLAN_TOOL = "mcp__plugin_unsloth-local-agent_unsloth__unsloth_plan_agent"
_CODEX_SUBAGENT_MCP_MODULE = "unsloth_cli.codex_subagent_mcp"
_CODEX_SUBAGENT_MCP_SERVER = "unsloth_local_agent"
_CODEX_SUBAGENT_MCP_TOOL = "spawn_local_agent"
_CODEX_SUBAGENT_CONFIG_ENV = "UNSLOTH_CODEX_SUBAGENT_CONFIG"
_CODEX_PARENT_OVERLAY_MANIFEST = ".unsloth-parent-overlay.json"
_CODEX_SUBAGENT_TOOL_DESCRIPTION = (
f"{_SUBAGENT_DESCRIPTION} Use this tool instead of the built-in spawn_agent tool for those "
"requests. Other subagent requests may use the built-in tools normally."
)
_CODEX_SUBAGENT_ROUTING_INSTRUCTIONS = (
"When the user asks to spawn an Unsloth agent or local agent, you must call the "
"spawn_local_agent MCP tool once with the complete task. Do not answer, simulate the "
"result, call wait, or use a built-in subagent before calling the tool. Use built-in "
"subagents for other delegation requests."
)
_PI_SUBAGENT_EXTENSION = Path(__file__).parent.parent / "pi_subagent.ts"
# OpenCode selects a model by "<providerID>/<modelID>". Use a dedicated id to avoid
# colliding with a user's providers; provider filters are set in the launch-time overlay.
@ -113,12 +138,14 @@ class _PassthroughCommand(TyperCommand):
_CLAUDE_ENV_UNSET = ("ANTHROPIC_API_KEY", "CLAUDE_CODE_OAUTH_TOKEN")
_CODEX_ENV_UNSET = ("OPENAI_API_KEY", "CODEX_API_KEY", "CODEX_ACCESS_TOKEN")
# Shared by every agent command; only the config/env/command differ.
# Help is grouped into rich panels so `--help` reads as Model / Server / Session
# instead of one long unaligned list.
_PANEL_MODEL = "Model"
_PANEL_SERVER = "Server"
_PANEL_SAMPLING = "Sampling"
_PANEL_SESSION = "Agent session"
_MODEL_OPTION = typer.Option(
@ -186,6 +213,56 @@ _TOOL_CALL_NUDGING_OPTION = typer.Option(
help = "Retry once with a nudge when a non-streaming passthrough tool call can't be healed. "
"On by default; when the flag is omitted an inherited UNSLOTH_TOOL_CALL_NUDGE is kept.",
)
# Sampling overrides pin a value on the auto-started server (winning over the client and the
# per-model recommendation). Default unset -> the model's recommended sampling is used.
_TEMPERATURE_OPTION = typer.Option(
None,
"--temperature",
min = 0.0,
max = 2.0,
rich_help_panel = _PANEL_SAMPLING,
help = "Pin the sampling temperature. Default: unset (per-model recommendation).",
)
_TOP_P_OPTION = typer.Option(
None,
"--top-p",
min = 0.0,
max = 1.0,
rich_help_panel = _PANEL_SAMPLING,
help = "Pin top-p (nucleus) sampling. Default: unset (per-model recommendation).",
)
_TOP_K_OPTION = typer.Option(
None,
"--top-k",
min = -1,
max = 100,
rich_help_panel = _PANEL_SAMPLING,
help = "Pin top-k sampling. Default: unset (per-model recommendation).",
)
_MIN_P_OPTION = typer.Option(
None,
"--min-p",
min = 0.0,
max = 1.0,
rich_help_panel = _PANEL_SAMPLING,
help = "Pin min-p sampling threshold. Default: unset (per-model recommendation).",
)
_REPETITION_PENALTY_OPTION = typer.Option(
None,
"--repetition-penalty",
min = 1.0,
max = 2.0,
rich_help_panel = _PANEL_SAMPLING,
help = "Pin the repetition penalty. Default: unset (per-model recommendation).",
)
_PRESENCE_PENALTY_OPTION = typer.Option(
None,
"--presence-penalty",
min = 0.0,
max = 2.0,
rich_help_panel = _PANEL_SAMPLING,
help = "Pin the presence penalty. Default: unset (per-model recommendation).",
)
# Agent-session knobs.
_KEY_OPTION = typer.Option(
@ -389,6 +466,12 @@ class ServerOptions(NamedTuple):
enable_tools: bool = False
tool_call_healing: Optional[bool] = None
tool_call_nudging: Optional[bool] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
top_k: Optional[int] = None
min_p: Optional[float] = None
repetition_penalty: Optional[float] = None
presence_penalty: Optional[float] = None
def _split_repo_variant(model: str) -> tuple:
@ -932,6 +1015,18 @@ def _start_studio_server(
child_env["UNSLOTH_TOOL_CALL_NUDGE"] = "1" if server.tool_call_nudging else "0"
elif "UNSLOTH_TOOL_CALL_NUDGE" not in child_env:
child_env["UNSLOTH_TOOL_CALL_NUDGE"] = "1"
# Forward any sampling pin via the env; `unsloth run` reads UNSLOTH_SAMPLING_* and the
# backend resolver applies it as a hard override. Only set fields the operator specified.
for _sampling_env, _sampling_value in (
("UNSLOTH_SAMPLING_TEMPERATURE", server.temperature),
("UNSLOTH_SAMPLING_TOP_P", server.top_p),
("UNSLOTH_SAMPLING_TOP_K", server.top_k),
("UNSLOTH_SAMPLING_MIN_P", server.min_p),
("UNSLOTH_SAMPLING_REPETITION_PENALTY", server.repetition_penalty),
("UNSLOTH_SAMPLING_PRESENCE_PENALTY", server.presence_penalty),
):
if _sampling_value is not None:
child_env[_sampling_env] = str(_sampling_value)
kwargs: dict = {
"stdout": log,
"stderr": subprocess.STDOUT,
@ -1019,6 +1114,30 @@ def _require_studio(
"""Return (base, server). server is a Popen only when WE auto-started it."""
base = find_studio_server()
if base is not None:
# Attaching to a server someone else started: UNSLOTH_SAMPLING_* pins only reach the
# server process when WE launch it (via _start_studio_server), so a sampling flag on the
# attach path can't take effect. Warn instead of silently dropping it, so the operator is
# not misled into thinking generation now uses the pinned value.
_pinned = [
_flag
for _flag, _value in (
("--temperature", server_options.temperature),
("--top-p", server_options.top_p),
("--top-k", server_options.top_k),
("--min-p", server_options.min_p),
("--repetition-penalty", server_options.repetition_penalty),
("--presence-penalty", server_options.presence_penalty),
)
if _value is not None
]
if _pinned:
typer.echo(
f"Warning: an Unsloth server is already running at {base}; sampling pins "
f"({', '.join(_pinned)}) apply only when this command starts the server, so the "
"running server keeps its current sampling. Stop it with `unsloth studio stop` "
"and re-run to apply them.",
err = True,
)
return base, None
expected = os.environ.get("UNSLOTH_STUDIO_URL", "http://127.0.0.1:8888").rstrip("/")
# Auto-start a local server only for an interactive launch with a model to serve, and
@ -1093,6 +1212,13 @@ def _write_private_json(path: Path, data: dict) -> None:
handle.write(json.dumps(data, indent = 2) + "\n")
def _write_private_text(path: Path, text: str) -> None:
path.parent.mkdir(parents = True, exist_ok = True, mode = 0o700)
fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
with os.fdopen(fd, "w", encoding = "utf-8") as handle:
handle.write(text)
def _read_json_object(path: Path) -> Optional[dict]:
# {} when missing, None when it can't be parsed as an object (so the caller
# leaves a user-managed file untouched rather than clobbering it).
@ -1422,10 +1548,18 @@ _DYNAMIC_SECTIONS_FLAG = "--exclude-dynamic-system-prompt-sections"
def _claude_settings_overlay(model_id: str) -> str:
# Session-only `claude --settings` overlay (command-line tier, no ~/.claude write):
# suppress the attribution header, and pin availableModels to the served model so a
# user allowlist can't reject it. The pin must be non-empty; [] is ignored.
# suppress the attribution header, keep every subagent on the served model (a user
# CLAUDE_CODE_SUBAGENT_MODEL pin would otherwise route delegated work off the local
# endpoint), and pin availableModels to the served model so a user allowlist can't
# reject it. The pin must be non-empty; [] is ignored.
return json.dumps(
{"env": {"CLAUDE_CODE_ATTRIBUTION_HEADER": "0"}, "availableModels": [model_id]}
{
"env": {
"CLAUDE_CODE_ATTRIBUTION_HEADER": "0",
"CLAUDE_CODE_SUBAGENT_MODEL": "inherit",
},
"availableModels": [model_id],
}
)
@ -1591,62 +1725,214 @@ def write_codex_config(base: str, model: dict, home: Path) -> None:
typer.echo(f"Updated {profile}")
def write_codex_subagent_config(base: str, key: str, model: dict, home: Path) -> Path:
"""Write a session-scoped Codex custom agent without replacing the main model."""
home.mkdir(parents = True, exist_ok = True)
model_id = model["id"]
window = model.get("context_length") or model.get("max_context_length")
catalog_name = "unsloth-model-catalog.json"
text = (
f"name = {json.dumps(_SUBAGENT_NAME)}\n"
f"description = {json.dumps(_SUBAGENT_DESCRIPTION)}\n"
f"developer_instructions = {json.dumps(_SUBAGENT_INSTRUCTIONS)}\n"
f"model_provider = {json.dumps(_CODEX_PROFILE)}\n"
f"model = {json.dumps(model_id)}\n"
def write_codex_subagent_bridge(
base: str, key: str, model: dict, home: Path, *, yolo: bool
) -> Path:
"""Write private config for an explicit local Codex child launched through MCP."""
child_home = home / "child"
write_codex_config(base, model, child_home)
path = home / "subagent.json"
_write_private_json(
path,
{
"api_key": key,
"codex_home": str(child_home),
"bypass_permissions": yolo,
},
)
if _codex_supports_model_catalog() and _CODEX_FALLBACK_PROMPT.is_file():
catalog = home / catalog_name
catalog_text = json.dumps(_codex_model_catalog(model), indent = 2) + "\n"
if not catalog.exists() or catalog.read_text(encoding = "utf-8") != catalog_text:
catalog.write_text(catalog_text, encoding = "utf-8")
typer.echo(f"Updated {catalog}")
text += f"model_catalog_json = {json.dumps(catalog_name)}\n"
if window:
text += f"model_context_window = {int(window)}\n"
credential = home / "unsloth-auth.json"
_write_private_json(credential, {"token": key})
auth_command = sys.executable
auth_args = [
"-c",
"import json,sys; print(json.load(open(sys.argv[1], encoding='utf-8'))['token'])",
str(credential),
]
if _wsl_windows_executable(["codex"]):
auth_command = "wsl.exe"
auth_args = [
"-d",
os.environ["WSL_DISTRO_NAME"],
"--",
sys.executable,
*auth_args,
]
text += (
f"\n{_PROVIDER_HEADER}\n"
'name = "Unsloth Studio"\n'
f"base_url = {json.dumps(base + '/v1')}\n"
'wire_api = "responses"\n'
f"\n{_PROVIDER_HEADER[:-1]}.auth]\n"
f"command = {json.dumps(auth_command)}\n"
f"args = {json.dumps(auth_args)}\n"
"timeout_ms = 5000\n"
)
path = home / f"{_SUBAGENT_NAME}.toml"
if not path.exists() or path.read_text(encoding = "utf-8") != text:
path.write_text(text, encoding = "utf-8")
typer.echo(f"Updated {path}")
return path
def _wsl_windows_user_profile(executable: str) -> Path:
"""Return the Windows user profile as a path accessible from WSL."""
profile = os.environ.get("USERPROFILE", "").strip()
if not profile:
try:
profile = subprocess.check_output(
["cmd.exe", "/d", "/c", "echo %USERPROFILE%"],
text = True,
stderr = subprocess.DEVNULL,
cwd = str(Path(executable).parent),
).strip()
except (OSError, subprocess.CalledProcessError) as exc:
_fail(f"Could not find the Windows user profile for Codex: {exc}")
if not profile or profile == "%USERPROFILE%":
_fail("Could not find the Windows user profile for Codex.")
if profile.startswith("/"):
return Path(profile)
try:
translated = subprocess.check_output(
["wslpath", "-u", profile],
text = True,
stderr = subprocess.DEVNULL,
).strip()
except (OSError, subprocess.CalledProcessError) as exc:
_fail(f"Could not translate Windows user profile {profile}: {exc}")
if not translated:
_fail(f"Could not translate Windows user profile {profile}.")
return Path(translated)
def _codex_source_home(*, ignore_configured: bool = False) -> Path:
configured = None if ignore_configured else os.environ.get("CODEX_HOME")
if configured:
if _wsl_windows_executable(["codex"]) and _looks_like_path(configured):
if not configured.startswith("/"):
try:
configured = subprocess.check_output(
["wslpath", "-u", configured],
text = True,
stderr = subprocess.DEVNULL,
).strip()
except (OSError, subprocess.CalledProcessError) as exc:
_fail(f"Could not translate Windows CODEX_HOME {configured}: {exc}")
if not configured:
_fail("Could not translate Windows CODEX_HOME.")
return Path(configured).expanduser()
executable = _wsl_windows_executable(["codex"])
if executable:
return _wsl_windows_user_profile(executable) / ".codex"
return Path.home() / ".codex"
def _remove_overlay_entry(path: Path) -> None:
is_junction = getattr(path, "is_junction", None)
if is_junction and is_junction():
path.rmdir()
elif path.is_symlink() or path.is_file():
path.unlink()
elif path.is_dir():
shutil.rmtree(path)
elif path.exists():
path.unlink()
def _create_directory_junction(source: Path, target: Path) -> bool:
if os.name != "nt":
return False
try:
result = subprocess.run(
["cmd.exe", "/d", "/c", "mklink", "/J", str(target), str(source)],
capture_output = True,
text = True,
timeout = 30,
check = False,
)
except (OSError, subprocess.SubprocessError):
return False
return result.returncode == 0
def write_codex_parent_overlay(overlay: Path) -> Path:
"""Add local-agent routing without replacing the cloud parent's configuration."""
overlay.mkdir(parents = True, exist_ok = True, mode = 0o700)
manifest_path = overlay / _CODEX_PARENT_OVERLAY_MANIFEST
try:
manifest = json.loads(manifest_path.read_text(encoding = "utf-8"))
except (FileNotFoundError, OSError, json.JSONDecodeError):
manifest = None
source_home = _codex_source_home()
overlay_key = str(overlay.resolve(strict = False))
source_key = str(source_home.resolve(strict = False))
if source_key == overlay_key:
previous_source = manifest.get("source_home") if isinstance(manifest, dict) else None
if isinstance(previous_source, str) and previous_source:
candidate = Path(previous_source).expanduser()
if str(candidate.resolve(strict = False)) != overlay_key:
source_home = candidate
else:
source_home = _codex_source_home(ignore_configured = True)
else:
source_home = _codex_source_home(ignore_configured = True)
source_key = str(source_home.resolve(strict = False))
same_source = isinstance(manifest, dict) and manifest.get("source_home") == source_key
if same_source:
managed_entries = manifest.get("entries", [])
if not isinstance(managed_entries, list):
managed_entries = []
for name in managed_entries:
if isinstance(name, str) and name not in {"", ".", ".."} and Path(name).name == name:
_remove_overlay_entry(overlay / name)
else:
# A reused overlay must never mix credentials, config, or plugins from two
# different Codex homes. Legacy overlays have no manifest, so rebuild them once.
for target in list(overlay.iterdir()):
_remove_overlay_entry(target)
# Keep the user's auth, config, plugins, agents, skills, rules, and session state visible.
# Symlinks make this an overlay rather than a stale copy. If Windows denies them,
# use directory junctions so large runtime state remains shared without a bulk copy.
# Copy the configuration surfaces and sessions only if both link forms are unavailable.
fallback_dirs = {"agents", "skills", "rules", "plugins", "marketplaces", "sessions"}
entries = []
if source_home.is_dir():
for source in source_home.iterdir():
if source.name in {
"AGENTS.md",
"AGENTS.override.md",
_CODEX_PARENT_OVERLAY_MANIFEST,
}:
continue
target = overlay / source.name
_remove_overlay_entry(target)
try:
target.symlink_to(source, target_is_directory = source.is_dir())
entries.append(source.name)
except OSError:
if source.is_file():
shutil.copy2(source, target)
entries.append(source.name)
elif source.is_dir():
if _create_directory_junction(source, target):
entries.append(source.name)
elif source.name in fallback_dirs:
shutil.copytree(source, target)
entries.append(source.name)
_write_private_json(
manifest_path,
{"source_home": source_key, "entries": sorted(entries)},
)
inherited = ""
instruction_name = "AGENTS.md"
for candidate in (source_home / "AGENTS.override.md", source_home / "AGENTS.md"):
try:
text = candidate.read_text(encoding = "utf-8")
except FileNotFoundError:
continue
except OSError as exc:
_fail(f"Could not preserve Codex instructions from {candidate}: {exc}")
if text.strip():
inherited = text.rstrip()
instruction_name = candidate.name
break
other_name = "AGENTS.md" if instruction_name == "AGENTS.override.md" else "AGENTS.override.md"
other = overlay / other_name
if other.is_file() or other.is_symlink():
other.unlink()
routing = _CODEX_SUBAGENT_ROUTING_INSTRUCTIONS
combined = f"{inherited}\n\n{routing}\n" if inherited else f"{routing}\n"
_write_private_text(overlay / instruction_name, combined)
return overlay
@contextlib.contextmanager
def _codex_parent_overlay(session_home: Path, *, launch: bool, persist: bool):
if launch and not persist:
temp_root = _agents_config_root() / ".tmp"
temp_root.mkdir(parents = True, exist_ok = True, mode = 0o700)
overlay = Path(tempfile.mkdtemp(prefix = "codex-parent-", dir = temp_root))
try:
yield write_codex_parent_overlay(overlay)
finally:
shutil.rmtree(overlay, ignore_errors = True)
else:
yield write_codex_parent_overlay(session_home / "parent")
def _agent_config_path(path: Path, command: list) -> str:
"""Translate a generated config path when a Windows agent runs through WSL."""
return _wsl_windows_path(path) if _wsl_windows_executable(command) else str(path)
@ -1770,25 +2056,42 @@ def write_claude_subagent_plugin(path: Path, server_env: dict) -> Path:
"description: Delegate a task to the local agent powered by Unsloth. Use when the "
"user asks to spawn an Unsloth agent or local agent.\n"
"---\n\n"
"Call the Unsloth local agent tool once with the complete task. Return its result "
"to the user without claiming that the cloud parent completed the local work.\n",
"Call the Unsloth local agent tool once with the complete task. In plan mode, call "
"the read-only Unsloth plan agent instead. Return its result to the user without "
"claiming that the cloud parent completed the local work.\n",
encoding = "utf-8",
)
return plugin
def _codex_subagent_flags(path: Path) -> list[str]:
config_path = _agent_config_path(path, ["codex"])
return [
"--enable",
"multi_agent",
"-c",
"agents.max_depth=1",
"-c",
f"agents.{_SUBAGENT_NAME}.description={json.dumps(_SUBAGENT_DESCRIPTION)}",
"-c",
f"agents.{_SUBAGENT_NAME}.config_file={json.dumps(config_path)}",
]
command = sys.executable
package_root = str(Path(__file__).resolve().parents[2])
bootstrap = (
f"import sys;sys.path.insert(0,{json.dumps(package_root)});"
f"from {_CODEX_SUBAGENT_MCP_MODULE} import main;main()"
)
args = ["-c", bootstrap, str(path)]
if _wsl_windows_executable(["codex"]):
command = "wsl.exe"
args = [
"-d",
os.environ["WSL_DISTRO_NAME"],
"--",
sys.executable,
"-c",
bootstrap,
str(path),
]
server = (
"{ "
f"command = {json.dumps(command)}, "
f"args = {json.dumps(args)}, "
f"required = true, enabled_tools = [{json.dumps(_CODEX_SUBAGENT_MCP_TOOL)}], "
'default_tools_approval_mode = "approve", '
"startup_timeout_sec = 15, tool_timeout_sec = 3600 }"
)
return ["-c", f"mcp_servers.{_CODEX_SUBAGENT_MCP_SERVER}={server}"]
def _wsl_windows_executable(command: list) -> Optional[str]:
@ -2597,6 +2900,12 @@ def claude(
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
temperature: Optional[float] = _TEMPERATURE_OPTION,
top_p: Optional[float] = _TOP_P_OPTION,
top_k: Optional[int] = _TOP_K_OPTION,
min_p: Optional[float] = _MIN_P_OPTION,
repetition_penalty: Optional[float] = _REPETITION_PENALTY_OPTION,
presence_penalty: Optional[float] = _PRESENCE_PENALTY_OPTION,
serve: bool = _SERVE_OPTION,
yolo: bool = _YOLO_OPTION,
persist: bool = _PERSIST_OPTION,
@ -2611,7 +2920,17 @@ def claude(
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
serve = serve,
launch = launch,
server_options = ServerOptions(enable_tools, tool_call_healing, tool_call_nudging),
server_options = ServerOptions(
enable_tools = enable_tools,
tool_call_healing = tool_call_healing,
tool_call_nudging = tool_call_nudging,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
repetition_penalty = repetition_penalty,
presence_penalty = presence_penalty,
),
)
model_id = entry["id"]
install_hint = (
@ -2639,7 +2958,7 @@ def claude(
_agent_config_path(plugin, ["claude"]),
# Before ctx.args: a forwarded `--` would turn later flags positional.
"--allowedTools",
_CLAUDE_SUBAGENT_TOOL,
f"{_CLAUDE_SUBAGENT_TOOL},{_CLAUDE_SUBAGENT_PLAN_TOOL}",
*_yolo_command_flags("claude", yolo),
*ctx.args,
]
@ -2698,6 +3017,12 @@ def codex(
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
temperature: Optional[float] = _TEMPERATURE_OPTION,
top_p: Optional[float] = _TOP_P_OPTION,
top_k: Optional[int] = _TOP_K_OPTION,
min_p: Optional[float] = _MIN_P_OPTION,
repetition_penalty: Optional[float] = _REPETITION_PENALTY_OPTION,
presence_penalty: Optional[float] = _PRESENCE_PENALTY_OPTION,
serve: bool = _SERVE_OPTION,
yolo: bool = _YOLO_OPTION,
persist: bool = _PERSIST_OPTION,
@ -2712,7 +3037,17 @@ def codex(
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
serve = serve,
launch = launch,
server_options = ServerOptions(enable_tools, tool_call_healing, tool_call_nudging),
server_options = ServerOptions(
enable_tools = enable_tools,
tool_call_healing = tool_call_healing,
tool_call_nudging = tool_call_nudging,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
repetition_penalty = repetition_penalty,
presence_penalty = presence_penalty,
),
)
# This preflight runs after _connect may have auto-started a server but before _run
# takes over its lifecycle, so tear the server down here if it rejects the model
@ -2726,25 +3061,32 @@ def codex(
subagent_id = _subagent_model_id(base, key, entry, model, gguf_variant)
subagent_model = {**entry, "id": subagent_id}
with _session_config("codex-subagent", launch, persist = persist) as home:
agent_config = write_codex_subagent_config(base, key, subagent_model, home)
command = [
"codex",
*_codex_subagent_flags(agent_config),
*_yolo_command_flags("codex", yolo),
*ctx.args,
]
typer.echo(
"Unsloth is available as the `unsloth` local agent. "
"Ask Codex to spawn an Unsloth or local agent."
)
_run(
bridge_config = write_codex_subagent_bridge(
base,
key,
subagent_model,
{},
command,
launch = launch,
install_hint = "npm install -g @openai/codex",
home,
yolo = yolo,
)
with _codex_parent_overlay(home, launch = launch, persist = persist) as parent_home:
command = [
"codex",
*_codex_subagent_flags(bridge_config),
*_yolo_command_flags("codex", yolo),
*ctx.args,
]
typer.echo(
"Unsloth is available as a local agent. "
"Ask Codex to spawn an Unsloth or local agent."
)
_run(
base,
subagent_model,
{"CODEX_HOME": str(parent_home)},
command,
launch = launch,
install_hint = "npm install -g @openai/codex",
)
return
command = [
"codex",
@ -2773,6 +3115,12 @@ def openclaw(
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
temperature: Optional[float] = _TEMPERATURE_OPTION,
top_p: Optional[float] = _TOP_P_OPTION,
top_k: Optional[int] = _TOP_K_OPTION,
min_p: Optional[float] = _MIN_P_OPTION,
repetition_penalty: Optional[float] = _REPETITION_PENALTY_OPTION,
presence_penalty: Optional[float] = _PRESENCE_PENALTY_OPTION,
serve: bool = _SERVE_OPTION,
yolo: bool = _YOLO_OPTION,
persist: bool = _PERSIST_OPTION,
@ -2787,7 +3135,17 @@ def openclaw(
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
serve = serve,
launch = launch,
server_options = ServerOptions(enable_tools, tool_call_healing, tool_call_nudging),
server_options = ServerOptions(
enable_tools = enable_tools,
tool_call_healing = tool_call_healing,
tool_call_nudging = tool_call_nudging,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
repetition_penalty = repetition_penalty,
presence_penalty = presence_penalty,
),
)
openclaw_args = list(ctx.args)
# Default a bare `unsloth start openclaw` to the local TUI. Anything the caller
@ -2837,6 +3195,12 @@ def opencode(
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
temperature: Optional[float] = _TEMPERATURE_OPTION,
top_p: Optional[float] = _TOP_P_OPTION,
top_k: Optional[int] = _TOP_K_OPTION,
min_p: Optional[float] = _MIN_P_OPTION,
repetition_penalty: Optional[float] = _REPETITION_PENALTY_OPTION,
presence_penalty: Optional[float] = _PRESENCE_PENALTY_OPTION,
serve: bool = _SERVE_OPTION,
yolo: bool = _YOLO_OPTION,
persist: bool = _PERSIST_OPTION,
@ -2851,7 +3215,17 @@ def opencode(
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
serve = serve,
launch = launch,
server_options = ServerOptions(enable_tools, tool_call_healing, tool_call_nudging),
server_options = ServerOptions(
enable_tools = enable_tools,
tool_call_healing = tool_call_healing,
tool_call_nudging = tool_call_nudging,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
repetition_penalty = repetition_penalty,
presence_penalty = presence_penalty,
),
)
if as_subagent:
subagent_id = _subagent_model_id(base, key, entry, model, gguf_variant)
@ -2981,6 +3355,12 @@ def hermes(
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
temperature: Optional[float] = _TEMPERATURE_OPTION,
top_p: Optional[float] = _TOP_P_OPTION,
top_k: Optional[int] = _TOP_K_OPTION,
min_p: Optional[float] = _MIN_P_OPTION,
repetition_penalty: Optional[float] = _REPETITION_PENALTY_OPTION,
presence_penalty: Optional[float] = _PRESENCE_PENALTY_OPTION,
serve: bool = _SERVE_OPTION,
yolo: bool = _YOLO_OPTION,
persist: bool = _PERSIST_OPTION,
@ -2997,7 +3377,17 @@ def hermes(
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
serve = serve,
launch = launch,
server_options = ServerOptions(enable_tools, tool_call_healing, tool_call_nudging),
server_options = ServerOptions(
enable_tools = enable_tools,
tool_call_healing = tool_call_healing,
tool_call_nudging = tool_call_nudging,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
repetition_penalty = repetition_penalty,
presence_penalty = presence_penalty,
),
)
install_hint = _hermes_install_hint()
with _session_config("hermes", launch, persist = persist) as home:
@ -3021,6 +3411,12 @@ def pi(
enable_tools: bool = _ENABLE_TOOLS_OPTION,
tool_call_healing: Optional[bool] = _TOOL_CALL_HEALING_OPTION,
tool_call_nudging: Optional[bool] = _TOOL_CALL_NUDGING_OPTION,
temperature: Optional[float] = _TEMPERATURE_OPTION,
top_p: Optional[float] = _TOP_P_OPTION,
top_k: Optional[int] = _TOP_K_OPTION,
min_p: Optional[float] = _MIN_P_OPTION,
repetition_penalty: Optional[float] = _REPETITION_PENALTY_OPTION,
presence_penalty: Optional[float] = _PRESENCE_PENALTY_OPTION,
serve: bool = _SERVE_OPTION,
yolo: bool = _YOLO_OPTION,
persist: bool = _PERSIST_OPTION,
@ -3035,7 +3431,17 @@ def pi(
LoadOptions(gguf_variant, max_seq_length, load_in_4bit, tensor_parallel),
serve = serve,
launch = launch,
server_options = ServerOptions(enable_tools, tool_call_healing, tool_call_nudging),
server_options = ServerOptions(
enable_tools = enable_tools,
tool_call_healing = tool_call_healing,
tool_call_nudging = tool_call_nudging,
temperature = temperature,
top_p = top_p,
top_k = top_k,
min_p = min_p,
repetition_penalty = repetition_penalty,
presence_penalty = presence_penalty,
),
)
install_hint = "npm install -g --ignore-scripts @earendil-works/pi-coding-agent"
if as_subagent:

View file

@ -1661,6 +1661,7 @@ def _consume_legacy_short_aliases(
_RUN_PANEL_MODEL = "Model"
_RUN_PANEL_SERVER = "Server & network"
_RUN_PANEL_TOOLS = "Tool calls"
_RUN_PANEL_SAMPLING = "Sampling"
_RUN_PANEL_ADVANCED = "Advanced"
@ -1758,6 +1759,57 @@ def run(
"Default: on. No effect on streaming requests or the server-side agentic loop."
),
),
temperature: Optional[float] = typer.Option(
None,
"--temperature",
min = 0.0,
max = 2.0,
rich_help_panel = _RUN_PANEL_SAMPLING,
help = (
"Pin the sampling temperature for every request that omits it, overriding the "
"model's recommended value. Default: unset (use the per-model recommendation)."
),
),
top_p: Optional[float] = typer.Option(
None,
"--top-p",
min = 0.0,
max = 1.0,
rich_help_panel = _RUN_PANEL_SAMPLING,
help = "Pin top-p (nucleus) sampling. Default: unset (per-model recommendation).",
),
top_k: Optional[int] = typer.Option(
None,
"--top-k",
min = -1,
max = 100,
rich_help_panel = _RUN_PANEL_SAMPLING,
help = "Pin top-k sampling. Default: unset (per-model recommendation).",
),
min_p: Optional[float] = typer.Option(
None,
"--min-p",
min = 0.0,
max = 1.0,
rich_help_panel = _RUN_PANEL_SAMPLING,
help = "Pin min-p sampling threshold. Default: unset (per-model recommendation).",
),
repetition_penalty: Optional[float] = typer.Option(
None,
"--repetition-penalty",
min = 1.0,
max = 2.0,
rich_help_panel = _RUN_PANEL_SAMPLING,
help = "Pin the repetition penalty. Default: unset (per-model recommendation).",
),
presence_penalty: Optional[float] = typer.Option(
None,
"--presence-penalty",
min = 0.0,
max = 2.0,
rich_help_panel = _RUN_PANEL_SAMPLING,
help = "Pin the presence penalty. Default: unset (per-model recommendation).",
),
yes: bool = typer.Option(
False,
"--yes",
@ -1841,7 +1893,7 @@ def run(
Example:
unsloth studio run --model unsloth/Qwen3-1.7B-GGUF --gguf-variant UD-Q4_K_XL
unsloth studio run --model unsloth/Qwen3-1.7B-GGUF --top-k 20 --seed 42 --parallel 8
unsloth studio run --model unsloth/Qwen3-1.7B-GGUF --temperature 0.7 --seed 42 --parallel 8
unsloth studio run --model some-model --chat-template-file /path/to/tpl.jinja
unsloth studio run --model unsloth/Qwen3-27B-GGUF --gguf-variant Q8_0 --tensor-parallel
"""
@ -1870,6 +1922,21 @@ def run(
elif "UNSLOTH_TOOL_CALL_NUDGE" not in os.environ:
os.environ["UNSLOTH_TOOL_CALL_NUDGE"] = "1"
# Sampling overrides: the backend resolver reads UNSLOTH_SAMPLING_* to hard-pin a field
# (winning over both the client and the per-model recommendation). Only write a flag that
# was set explicitly so an omitted flag inherits any value the parent forwarded (e.g.
# `unsloth start`) and, when nothing is set, leaves the per-model recommendation in charge.
for _sampling_env, _sampling_value in (
("UNSLOTH_SAMPLING_TEMPERATURE", temperature),
("UNSLOTH_SAMPLING_TOP_P", top_p),
("UNSLOTH_SAMPLING_TOP_K", top_k),
("UNSLOTH_SAMPLING_MIN_P", min_p),
("UNSLOTH_SAMPLING_REPETITION_PENALTY", repetition_penalty),
("UNSLOTH_SAMPLING_PRESENCE_PENALTY", presence_penalty),
):
if _sampling_value is not None:
os.environ[_sampling_env] = str(_sampling_value)
# Set before any re-exec so the in-venv server inherits it via the env.
# `run --verbose` used to pass through to llama-server (its own -v); keep
# that by forwarding --log-verbose so we add Unsloth logs without dropping it.

View file

@ -7,6 +7,7 @@ import { Type } from "typebox";
const provider = "unsloth";
const maxResultCharacters = 100_000;
const maxParallelAgents = 4;
const cancelGraceMilliseconds = 2_000;
const configPath = process.env.UNSLOTH_PI_SUBAGENT_CONFIG || "";
delete process.env.UNSLOTH_PI_SUBAGENT_CONFIG;
@ -27,6 +28,8 @@ const baseUrl = typeof config.baseUrl === "string" ? config.baseUrl : "";
const apiKey = typeof config.apiKey === "string" ? config.apiKey : "";
const contextWindow = positiveInt(config.contextWindow, 32768);
const maxTokens = positiveInt(config.maxTokens, Math.min(Math.floor(contextWindow / 4), 8192));
let activeAgents = 0;
const waitingAgents: Array<() => boolean> = [];
function positiveInt(value: unknown, fallback: number): number {
const parsed = Number.parseInt(typeof value === "string" ? value : String(value || ""), 10);
@ -47,6 +50,45 @@ function boundedResult(text: string): string {
return `${text.slice(0, maxResultCharacters)}\n\n[Local agent output truncated]`;
}
function agentSlotRelease(): () => void {
let released = false;
return () => {
if (released) return;
released = true;
while (waitingAgents.length) {
if (waitingAgents.shift()!()) return;
}
activeAgents -= 1;
};
}
function acquireAgentSlot(signal: AbortSignal | undefined): Promise<() => void> {
if (signal?.aborted) return Promise.reject(new Error("The local Unsloth agent was cancelled."));
if (activeAgents < maxParallelAgents) {
activeAgents += 1;
return Promise.resolve(agentSlotRelease());
}
return new Promise((resolve, reject) => {
let waiting = true;
const grant = () => {
if (!waiting) return false;
waiting = false;
signal?.removeEventListener("abort", cancel);
resolve(agentSlotRelease());
return true;
};
const cancel = () => {
if (!waiting) return;
waiting = false;
const index = waitingAgents.indexOf(grant);
if (index >= 0) waitingAgents.splice(index, 1);
reject(new Error("The local Unsloth agent was cancelled."));
};
waitingAgents.push(grant);
signal?.addEventListener("abort", cancel, { once: true });
});
}
function piInvocation(args: string[]): { command: string; args: string[] } {
const currentScript = process.argv[1];
const bunVirtualScript = currentScript?.startsWith("/$bunfs/root/");
@ -107,6 +149,136 @@ async function stopChildTree(child: ChildProcess): Promise<void> {
signalProcessGroup(child, "SIGKILL");
}
interface LocalAgentResult {
task: string;
response: string;
transcript: any[];
error?: string;
}
async function runLocalAgent(
task: string,
cwd: string,
signal: AbortSignal | undefined,
onProgress: (result: LocalAgentResult) => void,
): Promise<LocalAgentResult> {
const extension = fileURLToPath(import.meta.url);
const args = [
"--mode",
"json",
"--print",
"--no-session",
"--provider",
provider,
"--model",
model,
"--no-extensions",
"--extension",
extension,
`Task: ${task}`,
];
const invocation = piInvocation(args);
let output = "";
let stderr = "";
let childError = "";
let aborted = false;
const result: LocalAgentResult = { task, response: "", transcript: [] };
const transcriptEntries = new Set<string>();
const appendTranscript = (messages: any[]): boolean => {
let changed = false;
for (const message of messages) {
const entry = JSON.stringify(message);
if (transcriptEntries.has(entry)) continue;
transcriptEntries.add(entry);
result.transcript.push(message);
changed = true;
}
return changed;
};
const processLine = (line: string) => {
try {
const event = JSON.parse(line);
if (event.type === "message_end" && event.message && appendTranscript([event.message])) {
onProgress(result);
}
if (
event.type === "turn_end" &&
Array.isArray(event.toolResults) &&
event.toolResults.length &&
appendTranscript(event.toolResults)
) {
onProgress(result);
}
if (event.type !== "message_end") return;
const message = event.message;
// Pi reports model/API failures as message_end events while still
// exiting 0, so the exit status alone cannot surface them.
if (message?.stopReason === "error" || message?.stopReason === "aborted") {
childError =
(typeof message.errorMessage === "string" && message.errorMessage) ||
`The local Unsloth agent stopped: ${message.stopReason}.`;
return;
}
const response = finalText(message);
if (response) {
result.response = boundedResult(response);
childError = "";
}
} catch {
// Ignore non-JSON diagnostic lines. The exit status still reports failures.
}
};
const exitCode = await new Promise<number>((resolve, reject) => {
const child = spawn(invocation.command, invocation.args, {
cwd,
detached: process.platform !== "win32",
shell: false,
stdio: ["ignore", "pipe", "pipe"],
env: {
...process.env,
UNSLOTH_PI_SUBAGENT_CHILD: "1",
UNSLOTH_PI_SUBAGENT_CONFIG: configPath,
},
});
let cleanup: Promise<void> | undefined;
const cancel = () => {
if (aborted) return;
aborted = true;
cleanup = stopChildTree(child);
};
child.on("error", (error) => {
signal?.removeEventListener("abort", cancel);
reject(error);
});
child.stdout.on("data", (chunk) => {
output += chunk.toString();
const lines = output.split("\n");
output = lines.pop() || "";
for (const line of lines) processLine(line);
});
child.stderr.on("data", (chunk) => {
stderr = (stderr + chunk.toString()).slice(-100_000);
});
child.on("close", async (code) => {
signal?.removeEventListener("abort", cancel);
await cleanup;
if (output.trim()) processLine(output);
resolve(code ?? 1);
});
signal?.addEventListener("abort", cancel, { once: true });
if (signal?.aborted) cancel();
});
if (aborted) throw new Error("The local Unsloth agent was cancelled.");
if (exitCode !== 0) {
result.error = stderr.trim() || `The local Unsloth agent exited with code ${exitCode}.`;
}
if (childError) result.error = boundedResult(childError);
if (!result.response && !result.error) result.response = "The local agent returned no text.";
return result;
}
export default function unslothSubagent(pi: ExtensionAPI): void {
if (!model || !baseUrl || !apiKey || !configPath) {
throw new Error("Unsloth subagent configuration is incomplete.");
@ -137,104 +309,97 @@ export default function unslothSubagent(pi: ExtensionAPI): void {
name: "unsloth_agent",
label: "Unsloth agent",
description:
"Local coding subagent powered by Unsloth for debugging, implementation, and codebase research. Use when the user asks to spawn an Unsloth or local agent.",
"Run local coding agents powered by Unsloth for debugging, implementation, and codebase research. Use task for one agent. To run multiple independent agents, use tasks; up to four run concurrently. The tool returns only after every requested agent finishes.",
parameters: Type.Object({
task: Type.String({ description: "The complete task for the local Unsloth agent." }),
task: Type.Optional(
Type.String({ description: "The complete task for one local Unsloth agent." }),
),
tasks: Type.Optional(
Type.Array(Type.String({ description: "A complete task for one local Unsloth agent." }), {
description: "Independent tasks to run concurrently, one local agent per task.",
minItems: 2,
maxItems: maxParallelAgents,
}),
),
}),
async execute(_toolCallId, params, signal, _onUpdate, ctx) {
const extension = fileURLToPath(import.meta.url);
const args = [
"--mode",
"json",
"--print",
"--no-session",
"--provider",
provider,
"--model",
model,
"--no-extensions",
"--extension",
extension,
`Task: ${params.task}`,
];
const invocation = piInvocation(args);
let output = "";
let stderr = "";
let lastResponse = "";
let childError = "";
let aborted = false;
const processLine = (line: string) => {
try {
const event = JSON.parse(line);
if (event.type !== "message_end") return;
const message = event.message;
// Pi reports model/API failures as message_end events while still
// exiting 0, so the exit status alone cannot surface them.
if (message?.stopReason === "error" || message?.stopReason === "aborted") {
childError =
(typeof message.errorMessage === "string" && message.errorMessage) ||
`The local Unsloth agent stopped: ${message.stopReason}.`;
return;
}
const response = finalText(message);
if (response) {
lastResponse = boundedResult(response);
childError = "";
}
} catch {
// Ignore non-JSON diagnostic lines. The exit status still reports failures.
}
};
const exitCode = await new Promise<number>((resolve, reject) => {
const child = spawn(invocation.command, invocation.args, {
cwd: ctx.cwd,
detached: process.platform !== "win32",
shell: false,
stdio: ["ignore", "pipe", "pipe"],
env: {
...process.env,
UNSLOTH_PI_SUBAGENT_CHILD: "1",
UNSLOTH_PI_SUBAGENT_CONFIG: configPath,
},
});
let cleanup: Promise<void> | undefined;
const cancel = () => {
if (aborted) return;
aborted = true;
cleanup = stopChildTree(child);
};
child.on("error", (error) => {
signal?.removeEventListener("abort", cancel);
reject(error);
});
child.stdout.on("data", (chunk) => {
output += chunk.toString();
const lines = output.split("\n");
output = lines.pop() || "";
for (const line of lines) processLine(line);
});
child.stderr.on("data", (chunk) => {
stderr = (stderr + chunk.toString()).slice(-100_000);
});
child.on("close", async (code) => {
signal?.removeEventListener("abort", cancel);
await cleanup;
if (output.trim()) processLine(output);
resolve(code ?? 1);
});
signal?.addEventListener("abort", cancel, { once: true });
if (signal?.aborted) cancel();
});
if (aborted) throw new Error("The local Unsloth agent was cancelled.");
if (exitCode !== 0) {
throw new Error(stderr.trim() || `The local Unsloth agent exited with code ${exitCode}.`);
executionMode: "parallel",
async execute(_toolCallId, params, signal, onUpdate, ctx) {
const singleTask = typeof params.task === "string" && params.task.trim() ? params.task.trim() : "";
const parallelTasks = Array.isArray(params.tasks)
? params.tasks.map((task) => task.trim()).filter(Boolean)
: [];
if (Boolean(singleTask) === Boolean(parallelTasks.length)) {
throw new Error("Provide exactly one of task or tasks.");
}
if (childError) throw new Error(boundedResult(childError));
if (parallelTasks.length > maxParallelAgents) {
throw new Error(`At most ${maxParallelAgents} local agents can run concurrently.`);
}
if (parallelTasks.length === 1) {
throw new Error("Use task for one local agent, or tasks for two to four agents.");
}
const tasks = singleTask ? [singleTask] : parallelTasks;
const results: Array<LocalAgentResult | undefined> = new Array(tasks.length);
let completed = 0;
const details = () => ({
provider,
model,
mode: tasks.length === 1 ? "single" : "parallel",
results: results.filter((result): result is LocalAgentResult => Boolean(result)),
});
const emitUpdate = () => {
onUpdate?.({
content: [
{
type: "text",
text: `Local agents: ${completed}/${tasks.length} completed`,
},
],
details: details(),
});
};
await Promise.all(
tasks.map(async (task, index) => {
let releaseAgentSlot: (() => void) | undefined;
try {
releaseAgentSlot = await acquireAgentSlot(signal);
results[index] = await runLocalAgent(task, ctx.cwd, signal, (partial) => {
results[index] = partial;
emitUpdate();
});
} catch (error) {
results[index] = {
task,
response: "",
transcript: results[index]?.transcript || [],
error: String(error),
};
} finally {
releaseAgentSlot?.();
completed += 1;
emitUpdate();
}
}),
);
if (signal?.aborted) throw new Error("The local Unsloth agent was cancelled.");
const completedResults = results.filter(
(result): result is LocalAgentResult => Boolean(result),
);
const succeeded = completedResults.filter((result) => !result.error).length;
const response =
completedResults.length === 1
? completedResults[0].error || completedResults[0].response
: [
`Parallel: ${succeeded}/${tasks.length} local agents succeeded`,
...completedResults.map(
(result, index) =>
`\n### Agent ${index + 1}${result.error ? " failed" : ""}\n\n${result.error || result.response}`,
),
].join("\n");
if (succeeded !== completedResults.length) throw new Error(response);
return {
content: [{ type: "text", text: lastResponse || "The local agent returned no text." }],
details: { provider, model },
content: [{ type: "text", text: response }],
details: details(),
};
},
});

View file

@ -43,6 +43,41 @@ def test_protocol_lists_and_calls_local_agent():
}
def test_protocol_exposes_read_only_agent_for_claude_plan_mode():
listed = bridge._response(
{"jsonrpc": "2.0", "id": 1, "method": "tools/list"},
run_read_only_agent = lambda task: task,
read_only_tool_name = "unsloth_plan_agent",
)
tools = {tool["name"]: tool for tool in listed["result"]["tools"]}
assert tools["unsloth_agent"]["annotations"]["readOnlyHint"] is False
assert tools["unsloth_plan_agent"]["annotations"] == {
"readOnlyHint": True,
"destructiveHint": False,
"idempotentHint": True,
"openWorldHint": True,
}
called = bridge._response(
{
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": {
"name": "unsloth_plan_agent",
"arguments": {"task": " inspect this "},
},
},
run_agent = lambda task: f"write: {task}",
run_read_only_agent = lambda task: f"plan: {task}",
read_only_tool_name = "unsloth_plan_agent",
)
assert called["result"] == {
"content": [{"type": "text", "text": "plan: inspect this"}],
"isError": False,
}
def test_protocol_returns_tool_errors_to_parent():
response = bridge._response(
{
@ -212,6 +247,38 @@ def test_local_child_uses_unsloth_without_overwriting_parent_auth(
assert "CLAUDE_CODE_OAUTH_TOKEN" not in child_env
def test_read_only_local_child_uses_plan_mode(monkeypatch, tmp_path):
captured = {}
monkeypatch.setenv("UNSLOTH_CLAUDE_SUBAGENT_BASE_URL", "http://127.0.0.1:8888")
monkeypatch.setenv("UNSLOTH_CLAUDE_SUBAGENT_API_KEY", "sk-unsloth-test")
monkeypatch.setenv("UNSLOTH_CLAUDE_SUBAGENT_MODEL", "unsloth/model-GGUF:Q4_K_M")
monkeypatch.setenv("UNSLOTH_CLAUDE_SUBAGENT_BYPASS_PERMISSIONS", "1")
monkeypatch.setenv("CLAUDE_PROJECT_DIR", str(tmp_path))
monkeypatch.setattr(bridge.shutil, "which", lambda _: "/usr/local/bin/claude")
monkeypatch.setattr(bridge, "_claude_flags", lambda model: [])
class Process:
pid = 1234
returncode = 0
def communicate(self, timeout):
return json.dumps({"is_error": False, "result": "PLAN_OK"}), ""
def poll(self):
return self.returncode
def popen(command, **kwargs):
captured["command"] = command
return Process()
monkeypatch.setattr(bridge.subprocess, "Popen", popen)
assert bridge.run_local_agent("plan this", read_only = True) == "PLAN_OK"
command = captured["command"]
assert command[command.index("--permission-mode") + 1] == "plan"
prompt = command[command.index("--append-system-prompt") + 1]
assert "read-only local coding subagent" in prompt
def test_local_child_process_is_stopped_on_cancellation(monkeypatch, tmp_path):
monkeypatch.setenv("UNSLOTH_CLAUDE_SUBAGENT_BASE_URL", "http://127.0.0.1:8888")
monkeypatch.setenv("UNSLOTH_CLAUDE_SUBAGENT_API_KEY", "sk-unsloth-test")

View file

@ -0,0 +1,228 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
from __future__ import annotations
import io
import json
import os
import subprocess
import pytest
import unsloth_cli.codex_subagent_mcp as bridge
def _write_config(tmp_path, *, bypass_permissions = False):
path = tmp_path / "subagent.json"
path.write_text(
json.dumps(
{
"api_key": "sk-unsloth-test",
"codex_home": str(tmp_path / "child"),
"bypass_permissions": bypass_permissions,
}
)
)
return path
def test_protocol_uses_codex_specific_tool_name():
requests = "\n".join(
[
json.dumps({"jsonrpc": "2.0", "id": 0, "method": "initialize"}),
json.dumps({"jsonrpc": "2.0", "id": 1, "method": "tools/list"}),
json.dumps(
{
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": {
"name": bridge._CODEX_SUBAGENT_MCP_TOOL,
"arguments": {"task": " inspect this "},
},
}
),
]
)
output = io.StringIO()
bridge.serve(
io.StringIO(requests),
output,
run_agent = lambda task, cancel_event: f"completed: {task}",
tool_name = bridge._CODEX_SUBAGENT_MCP_TOOL,
tool_description = bridge._CODEX_SUBAGENT_TOOL_DESCRIPTION,
instructions = bridge._SERVER_INSTRUCTIONS,
)
responses = {
response["id"]: response for response in map(json.loads, output.getvalue().splitlines())
}
assert responses[0]["result"]["instructions"] == bridge._SERVER_INSTRUCTIONS
assert len(bridge._SERVER_INSTRUCTIONS) <= 512
assert responses[1]["result"]["tools"][0]["name"] == "spawn_local_agent"
assert (
"Use this tool instead of the built-in spawn_agent tool"
in responses[1]["result"]["tools"][0]["description"]
)
assert responses[1]["result"]["tools"][0]["annotations"]["destructiveHint"] is True
assert responses[2]["result"] == {
"content": [{"type": "text", "text": "completed: inspect this"}],
"isError": False,
}
@pytest.mark.parametrize("bypass_permissions", [False, True])
@pytest.mark.parametrize("wsl_bridge", [False, True])
def test_local_child_uses_explicit_unsloth_profile(
monkeypatch, tmp_path, bypass_permissions, wsl_bridge
):
config = _write_config(tmp_path, bypass_permissions = bypass_permissions)
monkeypatch.setenv(bridge._CODEX_SUBAGENT_CONFIG_ENV, str(config))
credential_names = ("OPENAI_API_KEY", "CODEX_API_KEY", "CODEX_ACCESS_TOKEN")
for name in credential_names:
monkeypatch.setenv(name, "cloud-key")
monkeypatch.setenv("CODEX_SQLITE_HOME", str(tmp_path / "parent-sqlite"))
if wsl_bridge:
monkeypatch.setattr(
bridge,
"_wsl_shim_env",
lambda command, env, unset: (
env,
(
bridge._CODEX_ENV_KEY,
"CODEX_HOME/p",
"CODEX_SQLITE_HOME/p",
*unset,
"PWD/p",
),
),
)
monkeypatch.setattr(bridge.shutil, "which", lambda _: "/usr/local/bin/codex")
captured = {}
class Process:
pid = 1234
returncode = 0
def communicate(self, timeout):
captured["timeout"] = timeout
return (
json.dumps(
{
"type": "item.completed",
"item": {"type": "agent_message", "text": "LOCAL_OK"},
}
),
"",
)
def poll(self):
return self.returncode
def popen(command, **kwargs):
captured["command"] = command
captured.update(kwargs)
return Process()
monkeypatch.setattr(bridge.subprocess, "Popen", popen)
assert bridge.run_local_agent("reply exactly LOCAL_OK") == "LOCAL_OK"
command = captured["command"]
assert command[:4] == ["/usr/local/bin/codex", "--oss", "--profile", "unsloth_api"]
if bypass_permissions:
assert "--dangerously-bypass-approvals-and-sandbox" in command
else:
assert command[4:8] == ["--sandbox", "workspace-write", "--ask-for-approval", "never"]
assert command[command.index("exec") + 1 : command.index("exec") + 4] == [
"--ephemeral",
"--json",
"--skip-git-repo-check",
]
assert command[-1].endswith("Task: reply exactly LOCAL_OK")
assert captured["cwd"] == os.getcwd()
assert captured["stdin"] is subprocess.DEVNULL
assert captured["stdout"] is subprocess.PIPE
assert captured["stderr"] is subprocess.PIPE
if os.name == "nt":
assert captured["creationflags"] == subprocess.CREATE_NEW_PROCESS_GROUP
else:
assert captured["start_new_session"] is True
assert captured["env"]["CODEX_HOME"] == str(tmp_path / "child")
assert captured["env"]["CODEX_SQLITE_HOME"] == str(tmp_path / "child")
assert captured["env"][bridge._CODEX_ENV_KEY] == "sk-unsloth-test"
if wsl_bridge:
assert all(captured["env"][name] == "" for name in credential_names)
wslenv = captured["env"]["WSLENV"].split(":")
assert all(
name in {entry.split("/", 1)[0] for entry in wslenv} for name in bridge._CODEX_ENV_UNSET
)
assert "CODEX_SQLITE_HOME/p" in wslenv
assert "PWD/p" in wslenv
else:
assert all(name not in captured["env"] for name in credential_names)
def test_local_child_returns_last_agent_message():
output = "\n".join(
[
json.dumps(
{
"type": "item.completed",
"item": {"type": "agent_message", "text": "intermediate"},
}
),
json.dumps(
{
"type": "item.completed",
"item": {"type": "agent_message", "text": "final"},
}
),
]
)
assert bridge._result_text(output) == "final"
def test_local_child_prioritizes_failed_turn_over_progress():
output = "\n".join(
[
json.dumps(
{
"type": "item.completed",
"item": {"type": "agent_message", "text": "still working"},
}
),
json.dumps({"type": "turn.failed", "error": {"message": "local failure"}}),
]
)
with pytest.raises(RuntimeError, match = "local failure"):
bridge._result_text(output)
def test_local_child_process_is_stopped_on_cancellation(monkeypatch, tmp_path):
config = _write_config(tmp_path)
monkeypatch.setenv(bridge._CODEX_SUBAGENT_CONFIG_ENV, str(config))
monkeypatch.setattr(bridge.shutil, "which", lambda _: "/usr/local/bin/codex")
cancel_event = bridge.threading.Event()
stopped = []
class Process:
pid = 1234
returncode = None
def communicate(self, timeout):
cancel_event.set()
raise subprocess.TimeoutExpired("codex", timeout)
def poll(self):
return self.returncode
process = Process()
monkeypatch.setattr(bridge.subprocess, "Popen", lambda *args, **kwargs: process)
def stop(child):
stopped.append(child)
child.returncode = -15
monkeypatch.setattr(bridge, "_stop_child", stop)
with pytest.raises(RuntimeError, match = "cancelled"):
bridge.run_local_agent("wait", cancel_event)
assert stopped == [process]

View file

@ -64,7 +64,12 @@ import {{ existsSync }} from "node:fs";
import {{ pathToFileURL }} from "node:url";
mock.module("typebox", () => ({{
Type: {{ Object: (value) => value, String: (value) => value }},
Type: {{
Object: (value) => value,
String: (value) => value,
Optional: (value) => value,
Array: (value) => value,
}},
}}));
test("cancellation stops the Pi child process group", async () => {{
@ -138,10 +143,25 @@ def test_pi_child_error_events_fail_the_tool_call(tmp_path):
driver = tmp_path / "pi-driver.js"
driver.write_text(
"""
const event = {
type: "message_end",
message: { role: "assistant", stopReason: "error", errorMessage: "backend unreachable", content: [] },
};
const task = process.argv.at(-1).replace(/^Task: /, "");
const event = task === "pass"
? {
type: "message_end",
message: {
role: "assistant",
stopReason: "stop",
content: [{ type: "text", text: "PASS_OK" }],
},
}
: {
type: "message_end",
message: {
role: "assistant",
stopReason: "error",
errorMessage: "backend unreachable",
content: [],
},
};
console.log(JSON.stringify(event));
""",
encoding = "utf-8",
@ -154,7 +174,12 @@ import {{ expect, mock, test }} from "bun:test";
import {{ pathToFileURL }} from "node:url";
mock.module("typebox", () => ({{
Type: {{ Object: (value) => value, String: (value) => value }},
Type: {{
Object: (value) => value,
String: (value) => value,
Optional: (value) => value,
Array: (value) => value,
}},
}}));
test("child error events fail the tool call", async () => {{
@ -168,14 +193,30 @@ test("child error events fail the tool call", async () => {{
registerTool(value) {{ tool = value; }},
}});
const execution = tool.execute(
const singleExecution = tool.execute(
"call",
{{ task: "fail" }},
undefined,
undefined,
{{ cwd: {str(tmp_path)!r} }},
);
await expect(execution).rejects.toThrow("backend unreachable");
await expect(singleExecution).rejects.toThrow("backend unreachable");
const parallelExecution = tool.execute(
"call",
{{ tasks: ["pass", "fail"] }},
undefined,
undefined,
{{ cwd: {str(tmp_path)!r} }},
);
const parallelError = await parallelExecution.then(
() => "",
(error) => String(error),
);
expect(parallelError).toContain("Parallel: 1/2 local agents succeeded");
expect(parallelError).toContain("PASS_OK");
expect(parallelError).toContain("Agent 2 failed");
expect(parallelError).toContain("backend unreachable");
}}, 10_000);
""",
encoding = "utf-8",
@ -189,3 +230,247 @@ test("child error events fail the tool call", async () => {{
)
assert completed.returncode == 0, completed.stdout + completed.stderr
@pytest.mark.skipif(os.name == "nt", reason = "POSIX driver script")
def test_pi_parallel_agents_run_together_and_preserve_transcripts(tmp_path):
bun = shutil.which("bun")
if bun is None:
pytest.skip("Bun is required to execute the bundled Pi extension")
config = tmp_path / "subagent.json"
config.write_text(
json.dumps(
{
"baseUrl": "http://127.0.0.1:8000/v1",
"apiKey": "private-token",
"model": "local-model",
"contextWindow": 32768,
"maxTokens": 8192,
}
),
encoding = "utf-8",
)
starts = tmp_path / "starts"
driver = tmp_path / "pi-driver.js"
driver.write_text(
f"""
import * as fs from "node:fs";
const task = process.argv.at(-1).replace(/^Task: /, "");
fs.appendFileSync({str(starts)!r}, `${{task}}\\n`);
for (let attempt = 0; attempt < 100; attempt++) {{
const count = fs.readFileSync({str(starts)!r}, "utf8").trim().split("\\n").filter(Boolean).length;
if (count >= 2) break;
await Bun.sleep(20);
}}
const event = {{
type: "message_end",
message: {{
role: "assistant",
stopReason: "stop",
content: [{{ type: "text", text: `DONE_${{task}}` }}],
}},
}};
console.log(JSON.stringify(event));
console.log(JSON.stringify({{
type: "tool_execution_end",
toolCallId: `tool_${{task}}`,
toolName: "read",
result: {{ content: [{{ type: "text", text: `TOOL_${{task}}` }}] }},
isError: false,
}}));
const toolResult = {{
role: "toolResult",
toolCallId: `tool_${{task}}`,
toolName: "read",
content: [{{ type: "text", text: `TOOL_${{task}}` }}],
isError: false,
}};
// Current Pi emits a completed tool result both as message_end and in the
// following turn_end. Preserve it once in the transcript.
console.log(JSON.stringify({{
type: "message_end",
message: toolResult,
}}));
console.log(JSON.stringify({{
type: "turn_end",
message: event.message,
toolResults: [toolResult],
}}));
""",
encoding = "utf-8",
)
extension = Path(__file__).parents[1] / "pi_subagent.ts"
test_file = tmp_path / "pi-parallel.test.ts"
test_file.write_text(
f"""
import {{ expect, mock, test }} from "bun:test";
import {{ pathToFileURL }} from "node:url";
mock.module("typebox", () => ({{
Type: {{
Object: (value) => value,
String: (value) => value,
Optional: (value) => value,
Array: (value) => value,
}},
}}));
test("parallel tasks launch one child each and retain their transcripts", async () => {{
process.env.UNSLOTH_PI_SUBAGENT_CONFIG = {str(config)!r};
process.argv[1] = {str(driver)!r};
const loaded = await import(pathToFileURL({str(extension)!r}).href);
let tool;
loaded.default({{
registerProvider() {{}},
registerTool(value) {{ tool = value; }},
}});
expect(tool.executionMode).toBe("parallel");
const result = await tool.execute(
"call",
{{ tasks: ["ALPHA", "BETA"] }},
undefined,
undefined,
{{ cwd: {str(tmp_path)!r} }},
);
expect(result.content[0].text).toContain("Parallel: 2/2 local agents succeeded");
expect(result.content[0].text).toContain("DONE_ALPHA");
expect(result.content[0].text).toContain("DONE_BETA");
expect(result.details.mode).toBe("parallel");
expect(result.details.results).toHaveLength(2);
expect(result.details.results[0].transcript).toHaveLength(2);
expect(result.details.results[1].transcript).toHaveLength(2);
expect(result.details.results[0].transcript[0].content[0].text).toBe("DONE_ALPHA");
expect(result.details.results[0].transcript[1].content[0].text).toBe("TOOL_ALPHA");
expect(result.details.results[1].transcript[0].content[0].text).toBe("DONE_BETA");
expect(result.details.results[1].transcript[1].content[0].text).toBe("TOOL_BETA");
}}, 10_000);
""",
encoding = "utf-8",
)
completed = subprocess.run(
[bun, "test", str(test_file)],
capture_output = True,
text = True,
timeout = 15,
)
assert completed.returncode == 0, completed.stdout + completed.stderr
@pytest.mark.skipif(os.name == "nt", reason = "POSIX driver script")
def test_pi_parallel_agent_cap_spans_concurrent_tool_calls(tmp_path):
bun = shutil.which("bun")
if bun is None:
pytest.skip("Bun is required to execute the bundled Pi extension")
config = tmp_path / "subagent.json"
config.write_text(
json.dumps(
{
"baseUrl": "http://127.0.0.1:8000/v1",
"apiKey": "private-token",
"model": "local-model",
"contextWindow": 32768,
"maxTokens": 8192,
}
),
encoding = "utf-8",
)
markers = tmp_path / "active"
markers.mkdir()
peaks = tmp_path / "peaks"
driver = tmp_path / "pi-driver.js"
driver.write_text(
f"""
import * as fs from "node:fs";
const task = process.argv.at(-1).replace(/^Task: /, "");
const marker = `{str(markers)!s}/${{process.pid}}`;
fs.writeFileSync(marker, task);
await Bun.sleep(150);
fs.appendFileSync({str(peaks)!r}, `${{fs.readdirSync({str(markers)!r}).length}}\\n`);
await Bun.sleep(150);
fs.unlinkSync(marker);
console.log(JSON.stringify({{
type: "message_end",
message: {{
role: "assistant",
stopReason: "stop",
content: [{{ type: "text", text: `DONE_${{task}}` }}],
}},
}}));
""",
encoding = "utf-8",
)
extension = Path(__file__).parents[1] / "pi_subagent.ts"
test_file = tmp_path / "pi-global-cap.test.ts"
test_file.write_text(
f"""
import {{ expect, mock, test }} from "bun:test";
import {{ pathToFileURL }} from "node:url";
mock.module("typebox", () => ({{
Type: {{
Object: (value) => value,
String: (value) => value,
Optional: (value) => value,
Array: (value) => value,
}},
}}));
test("concurrent tool calls share the four-agent cap", async () => {{
process.env.UNSLOTH_PI_SUBAGENT_CONFIG = {str(config)!r};
process.argv[1] = {str(driver)!r};
const loaded = await import(pathToFileURL({str(extension)!r}).href);
let tool;
loaded.default({{
registerProvider() {{}},
registerTool(value) {{ tool = value; }},
}});
const first = tool.execute(
"call-1",
{{ tasks: ["A1", "A2", "A3", "A4"] }},
undefined,
undefined,
{{ cwd: {str(tmp_path)!r} }},
);
const second = tool.execute(
"call-2",
{{ tasks: ["B1", "B2", "B3", "B4"] }},
undefined,
undefined,
{{ cwd: {str(tmp_path)!r} }},
);
const results = await Promise.all([first, second]);
expect(results[0].content[0].text).toContain("4/4 local agents succeeded");
expect(results[1].content[0].text).toContain("4/4 local agents succeeded");
const afterQueue = await tool.execute(
"call-3",
{{ task: "C" }},
undefined,
undefined,
{{ cwd: {str(tmp_path)!r} }},
);
expect(afterQueue.content[0].text).toContain("DONE_C");
}}, 10_000);
""",
encoding = "utf-8",
)
completed = subprocess.run(
[bun, "test", str(test_file)],
capture_output = True,
text = True,
timeout = 15,
)
assert completed.returncode == 0, completed.stdout + completed.stderr
observed = [int(value) for value in peaks.read_text().splitlines()]
assert max(observed) == 4

View file

@ -127,6 +127,8 @@ def test_claude_settings_overlay_pins_served_model():
assert overlay["availableModels"] == [MODEL["id"]]
# The attribution-header suppression is preserved alongside it.
assert overlay["env"]["CLAUDE_CODE_ATTRIBUTION_HEADER"] == "0"
# Subagents fall through to the served model instead of a user's opus/sonnet pin.
assert overlay["env"]["CLAUDE_CODE_SUBAGENT_MODEL"] == "inherit"
def test_install_agent_prompts_then_installs(monkeypatch):
@ -619,51 +621,223 @@ def test_write_codex_config_omits_catalog_for_old_codex(tmp_path, monkeypatch):
assert not (tmp_path / "model-catalog.json").exists()
def test_write_codex_subagent_config_keeps_parent_model_out(tmp_path, monkeypatch):
def test_write_codex_subagent_bridge_keeps_parent_credentials_out(tmp_path, monkeypatch):
monkeypatch.setattr(start, "_codex_supports_model_catalog", lambda: True)
local = {**MODEL, "id": MODEL["id"] + ":UD-Q4_K_XL"}
path = start.write_codex_subagent_config(BASE, "private-token", local, tmp_path)
agent = _parse_toml(path.read_text())
assert agent["name"] == "unsloth"
assert "local agent" in agent["description"].lower()
assert agent["model_provider"] == start._CODEX_PROFILE
assert agent["model"] == local["id"]
assert agent["model_context_window"] == MODEL["context_length"]
assert agent["model_providers"][start._CODEX_PROFILE] == {
"name": "Unsloth Studio",
"base_url": f"{BASE}/v1",
"wire_api": "responses",
"auth": {
"command": sys.executable,
"args": [
"-c",
"import json,sys; print(json.load(open(sys.argv[1], encoding='utf-8'))['token'])",
str(tmp_path / "unsloth-auth.json"),
],
"timeout_ms": 5000,
},
path = start.write_codex_subagent_bridge(
BASE,
"private-token",
local,
tmp_path,
yolo = False,
)
assert json.loads(path.read_text()) == {
"api_key": "private-token",
"codex_home": str(tmp_path / "child"),
"bypass_permissions": False,
}
assert json.loads((tmp_path / "unsloth-auth.json").read_text()) == {"token": "private-token"}
catalog = json.loads((tmp_path / agent["model_catalog_json"]).read_text())
assert path.stat().st_mode & 0o077 == 0
profile = _parse_toml((tmp_path / "child" / "unsloth_api.config.toml").read_text())
assert profile["model"] == local["id"]
assert profile["model_provider"] == start._CODEX_PROFILE
assert profile["model_context_window"] == MODEL["context_length"]
config = _parse_toml((tmp_path / "child" / "config.toml").read_text())
assert config["model_providers"][start._CODEX_PROFILE]["base_url"] == f"{BASE}/v1"
catalog = json.loads((tmp_path / "child" / profile["model_catalog_json"]).read_text())
assert catalog["models"][0]["slug"] == local["id"]
def test_write_codex_parent_overlay_preserves_user_state_and_instructions(tmp_path, monkeypatch):
source = tmp_path / "user-codex"
source.mkdir()
(source / "config.toml").write_text('model = "cloud-model"\n')
(source / "auth.json").write_text('{"auth": "cloud"}\n')
(source / "sessions").mkdir()
(source / "AGENTS.override.md").write_text("Keep my existing instructions.\n")
monkeypatch.setenv("CODEX_HOME", str(source))
overlay = start.write_codex_parent_overlay(tmp_path / "managed" / "parent")
assert (overlay / "config.toml").read_text() == 'model = "cloud-model"\n'
assert (overlay / "auth.json").read_text() == '{"auth": "cloud"}\n'
assert (overlay / "sessions").is_dir()
instructions = (overlay / "AGENTS.override.md").read_text()
assert instructions.startswith("Keep my existing instructions.\n")
assert start._CODEX_SUBAGENT_ROUTING_INSTRUCTIONS in instructions
assert not (overlay / "AGENTS.md").exists()
assert (overlay / "AGENTS.override.md").stat().st_mode & 0o077 == 0
assert (source / "AGENTS.override.md").read_text() == "Keep my existing instructions.\n"
def test_write_codex_parent_overlay_refreshes_reused_entries(tmp_path, monkeypatch):
first = tmp_path / "first-codex"
first.mkdir()
(first / "auth.json").write_text('{"auth": "old"}\n')
(first / "old-only.toml").write_text("old\n")
second = tmp_path / "second-codex"
second.mkdir()
(second / "auth.json").write_text('{"auth": "new"}\n')
overlay_path = tmp_path / "managed" / "parent"
monkeypatch.setenv("CODEX_HOME", str(first))
overlay = start.write_codex_parent_overlay(overlay_path)
assert (overlay / "auth.json").read_text() == '{"auth": "old"}\n'
assert (overlay / "old-only.toml").exists()
monkeypatch.setenv("CODEX_HOME", str(second))
overlay = start.write_codex_parent_overlay(overlay_path)
assert (overlay / "auth.json").read_text() == '{"auth": "new"}\n'
assert not (overlay / "old-only.toml").exists()
def test_write_codex_parent_overlay_does_not_use_itself_as_source(tmp_path, monkeypatch):
source = tmp_path / "user-codex"
source.mkdir()
(source / "auth.json").write_text('{"auth": "cloud"}\n')
overlay_path = tmp_path / "managed" / "parent"
monkeypatch.setenv("CODEX_HOME", str(source))
overlay = start.write_codex_parent_overlay(overlay_path)
monkeypatch.setenv("CODEX_HOME", str(overlay))
overlay = start.write_codex_parent_overlay(overlay_path)
assert (overlay / "auth.json").read_text() == '{"auth": "cloud"}\n'
manifest = json.loads((overlay / start._CODEX_PARENT_OVERLAY_MANIFEST).read_text())
assert manifest["source_home"] == str(source)
def test_write_codex_parent_overlay_refreshes_fallback_copies(tmp_path, monkeypatch):
source = tmp_path / "user-codex"
source.mkdir()
config = source / "config.toml"
config.write_text('model = "first"\n')
sessions = source / "sessions"
sessions.mkdir()
(sessions / "existing.jsonl").write_text("existing session\n")
monkeypatch.setenv("CODEX_HOME", str(source))
def deny_symlink(*args, **kwargs):
raise OSError("symlinks unavailable")
monkeypatch.setattr(Path, "symlink_to", deny_symlink)
monkeypatch.setattr(start, "_create_directory_junction", lambda source, target: False)
overlay = start.write_codex_parent_overlay(tmp_path / "managed" / "parent")
(overlay / "history.jsonl").write_text("session state\n")
config.write_text('model = "second"\n')
overlay = start.write_codex_parent_overlay(overlay)
assert (overlay / "config.toml").read_text() == 'model = "second"\n'
assert (overlay / "sessions" / "existing.jsonl").read_text() == "existing session\n"
assert (overlay / "history.jsonl").read_text() == "session state\n"
config.unlink()
overlay = start.write_codex_parent_overlay(overlay)
assert not (overlay / "config.toml").exists()
assert (overlay / "history.jsonl").read_text() == "session state\n"
def test_create_directory_junction_uses_windows_mklink(tmp_path, monkeypatch):
captured = {}
monkeypatch.setattr(start.os, "name", "nt")
def run(command, **kwargs):
captured["command"] = command
captured["kwargs"] = kwargs
return SimpleNamespace(returncode = 0)
monkeypatch.setattr(start.subprocess, "run", run)
source = tmp_path / "source"
target = tmp_path / "target"
assert start._create_directory_junction(source, target) is True
assert captured["command"] == [
"cmd.exe",
"/d",
"/c",
"mklink",
"/J",
str(target),
str(source),
]
assert captured["kwargs"] == {
"capture_output": True,
"text": True,
"timeout": 30,
"check": False,
}
@pytest.mark.skipif(os.name == "nt", reason = "WSL scenario")
def test_codex_subagent_auth_uses_wsl_for_windows_codex(monkeypatch, tmp_path):
def test_write_codex_parent_overlay_uses_windows_home_for_windows_codex(tmp_path, monkeypatch):
windows_profile = tmp_path / "windows-profile"
source = windows_profile / ".codex"
source.mkdir(parents = True)
(source / "auth.json").write_text('{"auth": "windows"}\n')
executable = "/mnt/c/Users/x/AppData/Roaming/npm/codex"
monkeypatch.delenv("CODEX_HOME", raising = False)
monkeypatch.delenv("USERPROFILE", raising = False)
monkeypatch.setenv("WSL_DISTRO_NAME", "Ubuntu")
monkeypatch.setattr(start.shutil, "which", lambda _: executable)
def check_output(command, **kwargs):
if command[0] == "cmd.exe":
assert kwargs["cwd"] == str(Path(executable).parent)
return r"C:\Users\x" + "\n"
assert command == ["wslpath", "-u", r"C:\Users\x"]
return str(windows_profile) + "\n"
monkeypatch.setattr(start.subprocess, "check_output", check_output)
overlay = start.write_codex_parent_overlay(tmp_path / "managed" / "parent")
assert (overlay / "auth.json").read_text() == '{"auth": "windows"}\n'
def test_codex_parent_overlay_launch_uses_private_temp_root_and_cleans_up(tmp_path, monkeypatch):
source = tmp_path / "user-codex"
source.mkdir()
(source / "auth.json").write_text("{}\n")
monkeypatch.setenv("CODEX_HOME", str(source))
agents_root = tmp_path / "agents"
monkeypatch.setattr(start, "_agents_config_root", lambda: agents_root)
with start._codex_parent_overlay(tmp_path / "session", launch = True, persist = False) as overlay:
assert overlay.parent == agents_root / ".tmp"
assert start._CODEX_SUBAGENT_ROUTING_INSTRUCTIONS in (overlay / "AGENTS.md").read_text()
assert overlay.exists()
assert not overlay.exists()
@pytest.mark.skipif(os.name == "nt", reason = "WSL scenario")
def test_codex_subagent_bridge_uses_wsl_for_windows_codex(monkeypatch, tmp_path):
monkeypatch.setenv("WSL_DISTRO_NAME", "Ubuntu")
monkeypatch.setattr(start, "_codex_supports_model_catalog", lambda: False)
monkeypatch.setattr(
start.shutil,
"which",
lambda _: "/mnt/c/Users/x/AppData/Roaming/npm/codex.exe",
)
path = start.write_codex_subagent_config(BASE, "private-token", MODEL, tmp_path)
auth = _parse_toml(path.read_text())["model_providers"][start._CODEX_PROFILE]["auth"]
assert auth["command"] == "wsl.exe"
assert auth["args"][:5] == ["-d", "Ubuntu", "--", sys.executable, "-c"]
assert auth["args"][-1] == str(tmp_path / "unsloth-auth.json")
flags = start._codex_subagent_flags(tmp_path / "subagent.json")
prefix = f"mcp_servers.{start._CODEX_SUBAGENT_MCP_SERVER}="
override = next(value for value in flags if value.startswith(prefix))
server = _parse_toml("server = " + override.removeprefix(prefix))["server"]
assert server["command"] == "wsl.exe"
assert server["args"] == [
"-d",
"Ubuntu",
"--",
sys.executable,
"-c",
server["args"][5],
str(tmp_path / "subagent.json"),
]
assert "sys.path.insert" in server["args"][5]
assert f"from {start._CODEX_SUBAGENT_MCP_MODULE} import main" in server["args"][5]
assert server["required"] is True
assert server["enabled_tools"] == [start._CODEX_SUBAGENT_MCP_TOOL]
assert server["default_tools_approval_mode"] == "approve"
assert not any(value.startswith("developer_instructions=") for value in flags)
@pytest.mark.skipif(os.name == "nt", reason = "WSL scenario")
@ -784,7 +958,10 @@ def test_connect_claude_no_launch(fake_studio):
_assert_env_set(result.output, "CLAUDE_AUTOCOMPACT_PCT_OVERRIDE", "90")
assert f"claude --model {MODEL['id']} --exclude-dynamic-system-prompt-sections" in result.output
# Overlay is passed inline (session-only), not a path into the user's ~/.claude.
assert "--settings" in result.output
command = _launch_command(result.output)
settings = json.loads(command[command.index("--settings") + 1])
assert settings["env"]["CLAUDE_CODE_SUBAGENT_MODEL"] == "inherit"
assert "--plugin-dir" not in command
assert ".claude/settings.json" not in result.output
@ -808,7 +985,7 @@ def test_connect_claude_as_subagent_preserves_cloud_parent(fake_studio, tmp_path
"--plugin-dir",
str(plugin),
"--allowedTools",
start._CLAUDE_SUBAGENT_TOOL,
f"{start._CLAUDE_SUBAGENT_TOOL},{start._CLAUDE_SUBAGENT_PLAN_TOOL}",
"hello",
]
assert "--model" not in command
@ -836,6 +1013,7 @@ def test_connect_claude_as_subagent_preserves_cloud_parent(fake_studio, tmp_path
}
skill = (plugin / "skills" / "local-agent" / "SKILL.md").read_text()
assert "spawn an Unsloth agent or local agent" in skill
assert "In plan mode" in skill
assert "Ask Claude to spawn an Unsloth or local agent." in result.output
@ -1011,6 +1189,11 @@ def test_connect_codex_no_launch(fake_studio, tmp_path):
def test_connect_codex_as_subagent_preserves_cloud_parent(fake_studio, tmp_path, monkeypatch):
monkeypatch.setattr(start, "_codex_supports_model_catalog", lambda: True)
source_home = tmp_path / "user-codex"
source_home.mkdir()
(source_home / "config.toml").write_text('model = "cloud-model"\n')
(source_home / "AGENTS.md").write_text("Keep the user's guidance.\n")
monkeypatch.setenv("CODEX_HOME", str(source_home))
result = CliRunner().invoke(
start.start_app,
[
@ -1024,20 +1207,34 @@ def test_connect_codex_as_subagent_preserves_cloud_parent(fake_studio, tmp_path,
assert result.exit_code == 0, result.output
command = _launch_command(result.output)
assert command[0] == "codex"
assert command[1:3] == ["--enable", "multi_agent"]
assert "agents.max_depth=1" in command
assert "--oss" not in command
assert "--profile" not in command
assert "--model" not in command
assert "CODEX_HOME" not in result.output
parent_home = tmp_path / "agents" / "codex-subagent" / "parent"
_assert_env_set(result.output, "CODEX_HOME", str(parent_home))
assert start._CODEX_ENV_KEY not in result.output
assert "sk-unsloth-feedfacefeedface" not in result.output
home = tmp_path / "agents" / "codex-subagent"
agent_path = home / "unsloth.toml"
agent = _parse_toml(agent_path.read_text())
assert agent["model"] == MODEL["id"] + ":UD-Q4_K_XL"
assert "env_key" not in agent["model_providers"][start._CODEX_PROFILE]
assert f"agents.unsloth.config_file={json.dumps(str(agent_path))}" in command
bridge_path = home / "subagent.json"
bridge = json.loads(bridge_path.read_text())
assert bridge["api_key"] == "sk-unsloth-feedfacefeedface"
assert bridge["codex_home"] == str(home / "child")
assert bridge["bypass_permissions"] is False
profile = _parse_toml((home / "child" / "unsloth_api.config.toml").read_text())
assert profile["model"] == MODEL["id"] + ":UD-Q4_K_XL"
prefix = f"mcp_servers.{start._CODEX_SUBAGENT_MCP_SERVER}="
override = next(value for value in command if value.startswith(prefix))
assert override.startswith(prefix)
server = _parse_toml("server = " + override.removeprefix(prefix))["server"]
assert server["command"] == sys.executable
assert server["args"] == ["-c", server["args"][1], str(bridge_path)]
assert "sys.path.insert" in server["args"][1]
assert f"from {start._CODEX_SUBAGENT_MCP_MODULE} import main" in server["args"][1]
assert server["enabled_tools"] == [start._CODEX_SUBAGENT_MCP_TOOL]
assert not any(value.startswith("developer_instructions=") for value in command)
parent_instructions = (parent_home / "AGENTS.md").read_text()
assert parent_instructions.startswith("Keep the user's guidance.\n")
assert start._CODEX_SUBAGENT_ROUTING_INSTRUCTIONS in parent_instructions
assert "Ask Codex to spawn an Unsloth or local agent." in result.output
@ -1787,6 +1984,118 @@ def test_start_studio_server_respects_inherited_tool_call_env(monkeypatch):
assert env["UNSLOTH_TOOL_CALL_NUDGE"] == "1"
def test_start_studio_server_forwards_sampling_via_env(monkeypatch):
# Sampling pins ride to the child server through UNSLOTH_SAMPLING_*; unset ones stay absent
# so the backend keeps the per-model recommendation.
captured = {}
class FakePopen:
def __init__(self, command, **kwargs):
captured["kwargs"] = kwargs
self.pid = 1
def poll(self):
return None
monkeypatch.setattr(start.subprocess, "Popen", FakePopen)
monkeypatch.setattr(start, "_studio_healthy", lambda base, timeout = 3.0: True)
monkeypatch.setattr(start, "_log_tail", lambda path, lines = 20: "API Key: sk-unsloth-x")
monkeypatch.setattr(start.time, "sleep", lambda _s: None)
for _v in ("TEMPERATURE", "TOP_P", "TOP_K", "MIN_P", "REPETITION_PENALTY", "PRESENCE_PENALTY"):
monkeypatch.delenv(f"UNSLOTH_SAMPLING_{_v}", raising = False)
# No sampling flags -> nothing forwarded.
start._start_studio_server("http://127.0.0.1:8888", "unsloth/M-GGUF", start.LoadOptions())
env = captured["kwargs"]["env"]
assert not any(k.startswith("UNSLOTH_SAMPLING_") for k in env)
# Pins are forwarded; unset ones stay absent.
start._start_studio_server(
"http://127.0.0.1:8888",
"unsloth/M-GGUF",
start.LoadOptions(),
start.ServerOptions(temperature = 0.3, top_k = 40, min_p = 0.05),
)
env = captured["kwargs"]["env"]
assert env["UNSLOTH_SAMPLING_TEMPERATURE"] == "0.3"
assert env["UNSLOTH_SAMPLING_TOP_K"] == "40"
assert env["UNSLOTH_SAMPLING_MIN_P"] == "0.05"
assert "UNSLOTH_SAMPLING_TOP_P" not in env
def test_require_studio_warns_on_sampling_pin_when_reusing_server(monkeypatch, capsys):
# Attaching to an already-running server can't apply UNSLOTH_SAMPLING_* pins (only
# _start_studio_server forwards them), so a sampling flag on the attach path must warn
# instead of being silently dropped while the command "succeeds".
monkeypatch.setattr(start, "find_studio_server", lambda: BASE)
base, server = start._require_studio(
"unsloth/M-GGUF",
start.LoadOptions(),
serve = True,
launch = True,
server_options = start.ServerOptions(temperature = 0.3, top_k = 40),
)
assert base == BASE
assert server is None # attach path: we did not start the server
err = capsys.readouterr().err
assert "already running" in err
assert "--temperature" in err and "--top-k" in err
# Only the pinned fields are named; an unset one is not.
assert "--top-p" not in err
def test_require_studio_no_sampling_warning_without_pins(monkeypatch, capsys):
# Reusing a server with no sampling pins stays silent (tool flags are out of scope here).
monkeypatch.setattr(start, "find_studio_server", lambda: BASE)
base, server = start._require_studio(
"unsloth/M-GGUF",
start.LoadOptions(),
serve = True,
server_options = start.ServerOptions(enable_tools = True),
)
assert base == BASE and server is None
assert "sampling" not in capsys.readouterr().err.lower()
def test_start_claude_parses_sampling_flags(fake_studio, monkeypatch):
# `unsloth start claude ... --temperature 0.3 --top-k 40` routes the pins into ServerOptions.
monkeypatch.setenv("UNSLOTH_STUDIO_URL", "http://127.0.0.1:8888")
monkeypatch.setattr(start, "find_studio_server", lambda: None)
captured = {}
fake = SimpleNamespace(pid = 1, poll = lambda: None)
def fake_start(
base,
model,
load,
server_options = None,
):
captured["server_options"] = server_options
start._auto_served_server = fake
return fake
monkeypatch.setattr(start, "_start_studio_server", fake_start)
monkeypatch.setattr(start, "_shutdown_server", lambda server: None)
monkeypatch.setattr(start.shutil, "which", lambda _: "/usr/local/bin/claude")
monkeypatch.setattr(start.subprocess, "run", lambda command, env: SimpleNamespace(returncode = 0))
result = CliRunner().invoke(
start.start_app,
[
"claude",
"--model",
"unsloth/gemma-4-E2B-it-GGUF",
"--temperature",
"0.3",
"--top-k",
"40",
],
)
assert result.exit_code == 0, result.output
so = captured["server_options"]
assert so.temperature == 0.3 and so.top_k == 40 and so.top_p is None
def test_connect_model_bare_id_matches_loaded_without_reload(fake_studio):
# A bare `--model <loaded repo>` (no load knobs) attaches to the already-loaded model
# without touching /api/inference/load, so it can never evict another session.

View file

@ -305,6 +305,43 @@ def test_run_omitted_flag_respects_inherited_env(monkeypatch, inherited):
assert studio_mod.os.environ["UNSLOTH_TOOL_CALL_NUDGE"] == inherited
_SAMPLING_ENV_SUFFIXES = (
"TEMPERATURE",
"TOP_P",
"TOP_K",
"MIN_P",
"REPETITION_PENALTY",
"PRESENCE_PENALTY",
)
def test_run_sampling_flags_set_env(monkeypatch):
"""`--temperature`/`--top-k` write UNSLOTH_SAMPLING_* (a hard override the backend applies);
an omitted sampling flag leaves its env unset so the per-model recommendation stays."""
studio_mod = _load_run_command()
for _v in _SAMPLING_ENV_SUFFIXES:
monkeypatch.delenv(f"UNSLOTH_SAMPLING_{_v}", raising = False)
_invoke_run(monkeypatch, _BASE + ["--temperature", "0.3", "--top-k", "40"])
assert studio_mod.os.environ["UNSLOTH_SAMPLING_TEMPERATURE"] == "0.3"
assert studio_mod.os.environ["UNSLOTH_SAMPLING_TOP_K"] == "40"
assert "UNSLOTH_SAMPLING_TOP_P" not in studio_mod.os.environ
def test_run_no_sampling_flags_leaves_env_unset(monkeypatch):
"""Plain `unsloth run` writes no UNSLOTH_SAMPLING_*; the server keeps the recommendation."""
studio_mod = _load_run_command()
for _v in _SAMPLING_ENV_SUFFIXES:
monkeypatch.delenv(f"UNSLOTH_SAMPLING_{_v}", raising = False)
_invoke_run(monkeypatch, _BASE)
assert not any(k.startswith("UNSLOTH_SAMPLING_") for k in studio_mod.os.environ)
def test_run_rejects_out_of_range_sampling_flag(monkeypatch):
"""typer enforces the documented ranges before a value can reach the server."""
result, _captured = _invoke_run(monkeypatch, _BASE + ["--temperature", "9"])
assert result.exit_code != 0
@pytest.mark.parametrize("platform", ["linux", "darwin", "win32"])
def test_reexec_argv_is_consistent_across_platforms(monkeypatch, platform):
"""Linux/Darwin (execvp) and Windows (Popen) must build the same argv."""
@ -344,12 +381,14 @@ def test_reexec_mixed_parallel_with_passthrough(monkeypatch):
"""--parallel + llama-server pass-through flags must all reach the child."""
result, captured = _invoke_run(
monkeypatch,
_BASE + ["--parallel", "8", "--top-k", "20", "--temp", "0.7"],
# --top-k is now a first-class sampling flag (routed via UNSLOTH_SAMPLING_*), so use
# --seed / --temp here, which remain genuine llama-server pass-through flags.
_BASE + ["--parallel", "8", "--seed", "42", "--temp", "0.7"],
)
assert len(captured) == 1
argv = captured[0]["argv"]
assert _value_after(argv, "--parallel") == "8", argv
assert _value_after(argv, "--top-k") == "20", argv
assert _value_after(argv, "--seed") == "42", argv
assert _value_after(argv, "--temp") == "0.7", argv