diff --git a/.github/workflows/consolidated-tests-ci.yml b/.github/workflows/consolidated-tests-ci.yml
index b56a6c2615..7978a200c0 100644
--- a/.github/workflows/consolidated-tests-ci.yml
+++ b/.github/workflows/consolidated-tests-ci.yml
@@ -209,7 +209,7 @@ jobs:
'peft>=0.18,<0.20' 'accelerate>=0.34,<2' \
ipython
# torchvision: unsloth_zoo.vision_utils imports it at module scope.
- pip install --index-url https://download.pytorch.org/whl/cpu \
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple \
'torch>=2.4,<2.11' 'torchvision<0.26'
# transformers + trl from the matrix combo.
pip install "$RESOLVED_TRANSFORMERS_SPEC"
@@ -268,6 +268,10 @@ jobs:
tests/saving/test_save_shell_injection.py \
tests/saving/test_patch_saving_none_tokenizer.py \
tests/saving/test_fix_sentencepiece_gguf_robustness.py \
+ tests/saving/test_compressed_export_schemes.py \
+ tests/saving/test_export_api_surface.py \
+ tests/saving/test_export_dispatch.py \
+ tests/saving/test_imatrix_export.py \
tests/utils/test_attention_masks.py \
tests/utils/test_trunc_normal_patch.py \
tests/python/test_fast_language_model_text_only.py
@@ -353,6 +357,10 @@ jobs:
tests/saving/test_save_shell_injection.py \
tests/saving/test_patch_saving_none_tokenizer.py \
tests/saving/test_fix_sentencepiece_gguf_robustness.py \
+ tests/saving/test_compressed_export_schemes.py \
+ tests/saving/test_export_api_surface.py \
+ tests/saving/test_export_dispatch.py \
+ tests/saving/test_imatrix_export.py \
tests/utils/test_attention_masks.py \
tests/utils/test_trunc_normal_patch.py \
tests/python/test_fast_language_model_text_only.py \
@@ -2166,7 +2174,7 @@ jobs:
python -m pip install --upgrade pip
# Match the matrix job's torch path so unsloth_zoo's
# `import torch` resolves to the same CPU build.
- pip install --index-url https://download.pytorch.org/whl/cpu \
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple \
'torch>=2.4,<2.11' 'torchvision<0.26'
pip install \
'numpy<3' protobuf sentencepiece \
diff --git a/.github/workflows/mlx-ci.yml b/.github/workflows/mlx-ci.yml
index 864630f9f0..a2f716a93c 100644
--- a/.github/workflows/mlx-ci.yml
+++ b/.github/workflows/mlx-ci.yml
@@ -163,7 +163,7 @@ jobs:
'pytest==9.0.3' \
'pytest-asyncio==1.3.0' \
'httpx==0.28.1'
- pip install --index-url https://download.pytorch.org/whl/cpu \
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple \
'torch==2.10.0'
# github.com occasionally 500s on the git fetch; retry the
# zoo install so a single upstream blip does not fail CI.
@@ -231,99 +231,6 @@ jobs:
tests/studio/test_is_mlx_dispatch_gate.py \
tests/studio/test_mlx_training_worker_behaviors.py
- # Studio prebuilt llama.cpp install + GGUF inference. Mirrors the
- # path Studio's setup.sh takes on macOS since #5963: plan against
- # the unslothai/llama.cpp fork's latest release, which ships the
- # bin-macos-arm64 bundle plus the llama-prebuilt-manifest.json the
- # default policy reads. After install, downloads a small published
- # GGUF (unsloth/gemma-3-270m-it-GGUF, Q4_K_M) and validates
- # llama-server /completion end to end. An install failure or a
- # non-zero binary exit is an Unsloth/Studio bug.
- - name: Studio prebuilt llama.cpp install + GGUF inference (Mac M1)
- env:
- # Withheld on PR: this step runs checked-out PR code; public GGUF still downloads.
- HF_TOKEN: ${{ github.event_name != 'pull_request' && secrets.HF_TOKEN || '' }}
- # install_llama_prebuilt.py hits the GitHub releases API to
- # resolve the asset URL. Anonymous calls share the runner-IP
- # rate-limit bucket and 403 quickly -- pass the workflow's
- # automatic GITHUB_TOKEN to bump us to the 5000/hr authenticated
- # bucket.
- GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
- GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
- run: |
- set -euo pipefail
- INSTALL_DIR="$HOME/.unsloth-studio-prebuilt-test/llama.cpp"
- rm -rf "$INSTALL_DIR"
- # Mirror studio/setup.sh on macOS (the install.sh user path):
- # it plans against the unslothai/llama.cpp fork's latest
- # release with no policy or tag flags.
- python studio/install_llama_prebuilt.py \
- --install-dir "$INSTALL_DIR" \
- --published-repo unslothai/llama.cpp
-
- # Studio bundles only llama-server + llama-quantize from the
- # prebuilt (not llama-cli) -- inference goes through
- # llama-server's HTTP /completion endpoint. Validate both:
- # llama-quantize --help proves the dynamic libs link, then
- # spin up llama-server and POST a /completion request on a
- # tiny published GGUF.
- LLAMA_SERVER="$INSTALL_DIR/build/bin/llama-server"
- LLAMA_QUANT="$INSTALL_DIR/build/bin/llama-quantize"
- [ -x "$LLAMA_SERVER" ] || { echo "::error::llama-server missing at $LLAMA_SERVER"; find "$INSTALL_DIR/build" -type f | head -40; exit 1; }
- [ -x "$LLAMA_QUANT" ] || { echo "::error::llama-quantize missing at $LLAMA_QUANT"; exit 1; }
- echo "llama-server : $LLAMA_SERVER"
- echo "llama-quantize: $LLAMA_QUANT"
- "$LLAMA_QUANT" --help >/dev/null && echo " llama-quantize loads OK"
-
- mkdir -p /tmp/ggufs
- bash .github/scripts/hf-download-with-retry.sh \
- 'unsloth/gemma-3-270m-it-GGUF' \
- 'gemma-3-270m-it-Q4_K_M.gguf' \
- /tmp/ggufs
-
- PORT=18080
- echo "=== starting llama-server on 127.0.0.1:$PORT ==="
- "$LLAMA_SERVER" \
- -m /tmp/ggufs/gemma-3-270m-it-Q4_K_M.gguf \
- --host 127.0.0.1 \
- --port "$PORT" \
- -c 256 \
- -n 16 \
- --no-warmup \
- > /tmp/llama-server.log 2>&1 &
- SERVER_PID=$!
- trap 'kill "$SERVER_PID" 2>/dev/null || true' EXIT
-
- # Wait for /health to come up
- for i in $(seq 1 30); do
- if curl -sf "http://127.0.0.1:$PORT/health" >/dev/null 2>&1; then
- echo " server up after ${i}s"
- break
- fi
- sleep 1
- done
- if ! curl -sf "http://127.0.0.1:$PORT/health" >/dev/null 2>&1; then
- echo "::error::llama-server never became healthy"
- tail -40 /tmp/llama-server.log
- exit 1
- fi
-
- PROMPT="Hello, my name is"
- echo "=== POST /completion ==="
- RESP=$(curl -sf -X POST "http://127.0.0.1:$PORT/completion" \
- -H 'Content-Type: application/json' \
- -d "{\"prompt\":\"$PROMPT\",\"n_predict\":16,\"temperature\":0,\"seed\":3407}")
- echo "raw response (head): $(echo "$RESP" | head -c 600)"
- CONTENT=$(echo "$RESP" | python -c "import json,sys; print(json.loads(sys.stdin.read()).get('content',''))")
- echo "completion content: $CONTENT"
-
- if [ -z "$CONTENT" ]; then
- echo "::error::llama-server /completion returned empty content"
- tail -40 /tmp/llama-server.log
- exit 1
- fi
- echo "OK: Studio prebuilt llama.cpp on Mac M1 + GGUF /completion works"
-
# Real MLX training + inference smoke test. Trains
# unsloth/gemma-3-270m-it for 7 deterministic LoRA steps
# (batch_size=2, gradient_accumulation_steps=3) on a single
@@ -338,6 +245,9 @@ jobs:
UNSLOTH_COMPILE_DISABLE: '1'
run: |
mkdir -p mlx_workdir
+ # Authenticate llama.cpp's release-API lookup (anonymous 403s on rate-limit);
+ # read-only GITHUB_TOKEN scoped here only, never to steps that run binaries.
+ GH_TOKEN="${{ secrets.GITHUB_TOKEN }}" GITHUB_TOKEN="${{ secrets.GITHUB_TOKEN }}" \
python tests/studio/run_real_mlx_smoke.py train \
--workdir "$PWD/mlx_workdir"
@@ -406,3 +316,88 @@ jobs:
cat "$f" 2>/dev/null || echo "(missing)"
echo
done
+
+ # Validates the macOS prebuilt path Studio's setup.sh uses (#5963): install the
+ # unslothai/llama.cpp fork's latest release, download a small public GGUF, and
+ # check llama-server /completion end to end. Split and placed last so the
+ # untrusted binary runs only in the final smoke step, after every HF_TOKEN step,
+ # leaving no token-bearing step or shared workspace for a tampered prebuilt to
+ # corrupt. GH_TOKEN: releases API; HF_TOKEN (withheld on PR): probe + GGUF fetch.
+ - name: Studio prebuilt llama.cpp install + GGUF download (Mac M1)
+ env:
+ GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
+ GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
+ HF_TOKEN: ${{ github.event_name != 'pull_request' && secrets.HF_TOKEN || '' }}
+ run: |
+ set -euo pipefail
+ INSTALL_DIR="$HOME/.unsloth-studio-prebuilt-test/llama.cpp"
+ rm -rf "$INSTALL_DIR"
+ # Download only -- no llama-quantize / llama-server launch in this step.
+ python studio/install_llama_prebuilt.py \
+ --install-dir "$INSTALL_DIR" \
+ --published-repo unslothai/llama.cpp
+ mkdir -p /tmp/ggufs
+ bash .github/scripts/hf-download-with-retry.sh \
+ 'unsloth/gemma-3-270m-it-GGUF' \
+ 'gemma-3-270m-it-Q4_K_M.gguf' \
+ /tmp/ggufs
+
+ # Final step: runs the downloaded binaries with no secrets present, and clears
+ # the GitHub Actions command files so a tampered prebuilt cannot influence the job.
+ - name: Studio prebuilt llama.cpp GGUF inference smoke (Mac M1)
+ run: |
+ set -euo pipefail
+ unset GITHUB_ENV GITHUB_PATH GITHUB_OUTPUT GITHUB_STEP_SUMMARY
+ INSTALL_DIR="$HOME/.unsloth-studio-prebuilt-test/llama.cpp"
+ # Studio bundles only llama-server + llama-quantize (not llama-cli);
+ # inference goes through llama-server's HTTP /completion endpoint.
+ LLAMA_SERVER="$INSTALL_DIR/build/bin/llama-server"
+ LLAMA_QUANT="$INSTALL_DIR/build/bin/llama-quantize"
+ [ -x "$LLAMA_SERVER" ] || { echo "::error::llama-server missing at $LLAMA_SERVER"; find "$INSTALL_DIR/build" -type f | head -40; exit 1; }
+ [ -x "$LLAMA_QUANT" ] || { echo "::error::llama-quantize missing at $LLAMA_QUANT"; exit 1; }
+ echo "llama-server : $LLAMA_SERVER"
+ echo "llama-quantize: $LLAMA_QUANT"
+ "$LLAMA_QUANT" --help >/dev/null && echo " llama-quantize loads OK"
+
+ PORT=18080
+ echo "=== starting llama-server on 127.0.0.1:$PORT ==="
+ "$LLAMA_SERVER" \
+ -m /tmp/ggufs/gemma-3-270m-it-Q4_K_M.gguf \
+ --host 127.0.0.1 \
+ --port "$PORT" \
+ -c 256 \
+ -n 16 \
+ --no-warmup \
+ > /tmp/llama-server.log 2>&1 &
+ SERVER_PID=$!
+ trap 'kill "$SERVER_PID" 2>/dev/null || true' EXIT
+
+ # Wait for /health to come up
+ for i in $(seq 1 30); do
+ if curl -sf "http://127.0.0.1:$PORT/health" >/dev/null 2>&1; then
+ echo " server up after ${i}s"
+ break
+ fi
+ sleep 1
+ done
+ if ! curl -sf "http://127.0.0.1:$PORT/health" >/dev/null 2>&1; then
+ echo "::error::llama-server never became healthy"
+ tail -40 /tmp/llama-server.log
+ exit 1
+ fi
+
+ PROMPT="Hello, my name is"
+ echo "=== POST /completion ==="
+ RESP=$(curl -sf -X POST "http://127.0.0.1:$PORT/completion" \
+ -H 'Content-Type: application/json' \
+ -d "{\"prompt\":\"$PROMPT\",\"n_predict\":16,\"temperature\":0,\"seed\":3407}")
+ echo "raw response (head): $(echo "$RESP" | head -c 600)"
+ CONTENT=$(echo "$RESP" | python -c "import json,sys; print(json.loads(sys.stdin.read()).get('content',''))")
+ echo "completion content: $CONTENT"
+
+ if [ -z "$CONTENT" ]; then
+ echo "::error::llama-server /completion returned empty content"
+ tail -40 /tmp/llama-server.log
+ exit 1
+ fi
+ echo "OK: Studio prebuilt llama.cpp on Mac M1 + GGUF /completion works"
diff --git a/.github/workflows/notebooks-ci.yml b/.github/workflows/notebooks-ci.yml
index 2edcae8ab2..0e0b35dd4d 100644
--- a/.github/workflows/notebooks-ci.yml
+++ b/.github/workflows/notebooks-ci.yml
@@ -263,7 +263,7 @@ jobs:
# unsloth_zoo.vision_utils imports PIL at module top, and the
# easiest way to get a torch-compatible PIL on a CPU runner is
# to let torchvision pull the right Pillow version.
- pip install --index-url https://download.pytorch.org/whl/cpu \
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple \
'torch>=2.8,<2.11' 'torchvision<0.26'
# Pin to the same versions update_all_notebooks.py installs in
# generated notebooks. Keep these in lockstep with PIN_TRL /
diff --git a/.github/workflows/studio-backend-ci.yml b/.github/workflows/studio-backend-ci.yml
index ea60252cf6..bce355458a 100644
--- a/.github/workflows/studio-backend-ci.yml
+++ b/.github/workflows/studio-backend-ci.yml
@@ -76,7 +76,7 @@ jobs:
# Torch CPU + transformers are required by a chunk of the backend test
# suite (gpu_selection, kv_cache_estimation, utils). CPU-only torch
# keeps the install ~250 MB / ~1 min on a clean runner.
- pip install --index-url https://download.pytorch.org/whl/cpu 'torch>=2.4,<2.11'
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple 'torch>=2.4,<2.11'
pip install 'transformers>=4.51,<5.5'
- name: Backend tests
@@ -137,7 +137,7 @@ jobs:
pyyaml jinja2 mammoth unpdf requests typer \
'numpy<3' pytest pytest-asyncio httpx
# torchvision: unsloth_zoo.vision_utils imports it at module scope.
- pip install --index-url https://download.pytorch.org/whl/cpu \
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple \
'torch>=2.4,<2.11' 'torchvision<0.26'
pip install 'transformers>=4.51,<5.5'
# bitsandbytes: hard import in unsloth/models/_utils.py. Recent
diff --git a/.github/workflows/studio-mac-ui-smoke.yml b/.github/workflows/studio-mac-ui-smoke.yml
index 512af54d53..20ca247b9f 100644
--- a/.github/workflows/studio-mac-ui-smoke.yml
+++ b/.github/workflows/studio-mac-ui-smoke.yml
@@ -185,13 +185,14 @@ jobs:
# Retry up to 3 times to absorb known macos-14 free-runner
# flakes: (1) Playwright Node 24 pipeTransport.js 'Unexpected
# end of JSON input' crash when the Chromium browser process
- # dies mid-test, and (2) Chromium net::ERR_NO_BUFFER_SPACE
- # when the runner's kernel briefly runs out of socket buffers.
- # The retry FULLY resets Studio (kill, reset-password, reboot,
- # wait /api/health, re-export bootstrap pw) before re-running
- # the script. A real test failure (assertion / timeout) does
- # NOT match either pattern so it bypasses retry and surfaces
- # immediately.
+ # dies mid-test, (2) Chromium net::ERR_NO_BUFFER_SPACE when the
+ # runner's kernel briefly runs out of socket buffers, and (3) a
+ # goto 'interrupted by another navigation' when the SPA auth
+ # guard redirects mid-navigation. The retry FULLY resets Studio
+ # (kill, reset-password, reboot, wait /api/health, re-export
+ # bootstrap pw) before re-running the script. A real test failure
+ # (assertion / timeout) does NOT match any pattern so it bypasses
+ # retry and surfaces immediately.
run: |
mkdir -p logs/playwright
attempt=1
@@ -204,8 +205,9 @@ jobs:
if [ "$rc" -eq 0 ]; then
break
fi
- if { grep -q "Unexpected end of JSON input" logs/playwright_attempt_${attempt}.log \
- || grep -q "ERR_NO_BUFFER_SPACE" logs/playwright_attempt_${attempt}.log; } \
+ if { grep -q "Unexpected end of JSON input" logs/playwright_attempt_${attempt}.log \
+ || grep -q "ERR_NO_BUFFER_SPACE" logs/playwright_attempt_${attempt}.log \
+ || grep -q "interrupted by another navigation" logs/playwright_attempt_${attempt}.log; } \
&& [ "$attempt" -lt "$max_attempts" ]; then
echo "::warning::Playwright flake on attempt ${attempt}; resetting Studio and retrying..."
kill "${STUDIO_PID}" 2>/dev/null || true
@@ -280,8 +282,8 @@ jobs:
STUDIO_UI_TURN_TIMEOUT_MS: '540000'
GGUF_REPO: ${{ env.GGUF_REPO }}
GGUF_VARIANT: ${{ env.GGUF_VARIANT }}
- # Same flake-retry shape as "Drive the chat UI with Playwright"
- # -- catches pipeTransport JSON crash and ERR_NO_BUFFER_SPACE.
+ # Same flake-retry shape as "Drive the chat UI with Playwright" -- catches
+ # pipeTransport JSON crash, ERR_NO_BUFFER_SPACE, and nav interrupts.
run: |
mkdir -p logs/playwright_extra
attempt=1
@@ -294,8 +296,9 @@ jobs:
if [ "$rc" -eq 0 ]; then
break
fi
- if { grep -q "Unexpected end of JSON input" logs/playwright_extra_attempt_${attempt}.log \
- || grep -q "ERR_NO_BUFFER_SPACE" logs/playwright_extra_attempt_${attempt}.log; } \
+ if { grep -q "Unexpected end of JSON input" logs/playwright_extra_attempt_${attempt}.log \
+ || grep -q "ERR_NO_BUFFER_SPACE" logs/playwright_extra_attempt_${attempt}.log \
+ || grep -q "interrupted by another navigation" logs/playwright_extra_attempt_${attempt}.log; } \
&& [ "$attempt" -lt "$max_attempts" ]; then
echo "::warning::Playwright flake on attempt ${attempt}; resetting Studio and retrying..."
kill "${STUDIO_EXTRA_PID}" 2>/dev/null || true
diff --git a/.github/workflows/studio-windows-inference-smoke.yml b/.github/workflows/studio-windows-inference-smoke.yml
index c44c68278d..8186c07211 100644
--- a/.github/workflows/studio-windows-inference-smoke.yml
+++ b/.github/workflows/studio-windows-inference-smoke.yml
@@ -1338,11 +1338,19 @@ jobs:
shell: pwsh
run: |
$ErrorActionPreference = 'Stop'
+ # A Program Files dir can hold a transient handle (Defender / MSBuild node)
+ # so Rename-Item intermittently fails with "Access is denied"; retry to ride it out.
+ function Rename-WithRetry($Path, $NewName) {
+ for ($i = 1; $i -le 6; $i++) {
+ try { Rename-Item -LiteralPath $Path -NewName $NewName -ErrorAction Stop; return }
+ catch { if ($i -eq 6) { throw }; Start-Sleep -Seconds 3 }
+ }
+ }
# Rename the Visual Studio install roots (incl. the Installer that holds
# vswhere.exe) so Find-VsBuildTools' vswhere + filesystem scan both miss.
foreach ($d in @("$env:ProgramFiles\Microsoft Visual Studio", "${env:ProgramFiles(x86)}\Microsoft Visual Studio")) {
if (Test-Path -LiteralPath $d) {
- Rename-Item -LiteralPath $d -NewName ((Split-Path $d -Leaf) + '.vsoff')
+ Rename-WithRetry $d ((Split-Path $d -Leaf) + '.vsoff')
Write-Host "Hid VS: $d"
}
}
@@ -1351,7 +1359,7 @@ jobs:
$hidden = @()
foreach ($c in (Get-Command cmake -All -ErrorAction SilentlyContinue)) {
if ($c.Source -and (Test-Path -LiteralPath $c.Source)) {
- Rename-Item -LiteralPath $c.Source -NewName ((Split-Path $c.Source -Leaf) + '.off')
+ Rename-WithRetry $c.Source ((Split-Path $c.Source -Leaf) + '.off')
$hidden += $c.Source
Write-Host "Hid cmake: $($c.Source)"
}
@@ -1376,7 +1384,7 @@ jobs:
- name: PyTorch CPU wheel installs and imports (no Visual Studio)
run: |
python -m pip install --upgrade pip
- python -m pip install torch --index-url https://download.pytorch.org/whl/cpu
+ python -m pip install torch --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple
python -c "import torch; print('torch', torch.__version__, 'cuda?', torch.cuda.is_available())"
- name: Install Studio (--local, --no-torch) with no build tools present
@@ -1536,8 +1544,16 @@ jobs:
shell: pwsh
run: |
$ErrorActionPreference = 'Stop'
+ # Retry the rename: a Program Files dir can hold a transient handle that
+ # makes Rename-Item intermittently fail with "Access is denied".
+ function Rename-WithRetry($Path, $NewName) {
+ for ($i = 1; $i -le 6; $i++) {
+ try { Rename-Item -LiteralPath $Path -NewName $NewName -ErrorAction Stop; return }
+ catch { if ($i -eq 6) { throw }; Start-Sleep -Seconds 3 }
+ }
+ }
foreach ($d in @("$env:ProgramFiles\Microsoft Visual Studio", "${env:ProgramFiles(x86)}\Microsoft Visual Studio")) {
- if (Test-Path -LiteralPath $d) { Rename-Item -LiteralPath $d -NewName ((Split-Path $d -Leaf) + '.vsoff'); Write-Host "Hid VS: $d" }
+ if (Test-Path -LiteralPath $d) { Rename-WithRetry $d ((Split-Path $d -Leaf) + '.vsoff'); Write-Host "Hid VS: $d" }
}
- name: Windows CUDA and ROCm prebuilts exist in unslothai/llama.cpp (what GPU users download, no VS)
diff --git a/.github/workflows/version-compat-ci.yml b/.github/workflows/version-compat-ci.yml
index 599b53df1d..e492d21e99 100644
--- a/.github/workflows/version-compat-ci.yml
+++ b/.github/workflows/version-compat-ci.yml
@@ -242,7 +242,7 @@ jobs:
run: |
python -m pip install --upgrade pip
# CPU torch (vllm/peft/st all depend on it).
- pip install --index-url https://download.pytorch.org/whl/cpu \
+ pip install --index-url https://download.pytorch.org/whl/cpu --extra-index-url https://pypi.org/simple \
'torch>=2.4,<2.11' 'torchvision<0.26' 'torchcodec<0.10'
# torchcodec is a hard requirement on transformers 5.x:
# transformers/audio_utils.py:55 does
diff --git a/pyproject.toml b/pyproject.toml
index 13c421d8ea..844ead2454 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -255,10 +255,6 @@ cu118onlytorch270 = [
"xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp310-cp310-manylinux_2_28_x86_64.whl ; python_version=='3.10' and ('linux' in sys_platform)",
"xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp311-cp311-manylinux_2_28_x86_64.whl ; python_version=='3.11' and ('linux' in sys_platform)",
"xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp312-cp312-manylinux_2_28_x86_64.whl ; python_version=='3.12' and ('linux' in sys_platform)",
- "xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp39-cp39-win_amd64.whl ; python_version=='3.9' and (sys_platform == 'win32')",
- "xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp310-cp310-win_amd64.whl ; python_version=='3.10' and (sys_platform == 'win32')",
- "xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp311-cp311-win_amd64.whl ; python_version=='3.11' and (sys_platform == 'win32')",
- "xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.30-cp312-cp312-win_amd64.whl ; python_version=='3.12' and (sys_platform == 'win32')",
]
cu126onlytorch270 = [
"xformers @ https://download.pytorch.org/whl/cu126/xformers-0.0.30-cp39-cp39-manylinux_2_28_x86_64.whl ; python_version=='3.9' and ('linux' in sys_platform)",
@@ -282,7 +278,6 @@ cu128onlytorch270 = [
]
cu118onlytorch271 = [
"xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.31.post1-cp39-abi3-manylinux_2_28_x86_64.whl ; ('linux' in sys_platform)",
- "xformers @ https://download.pytorch.org/whl/cu118/xformers-0.0.31.post1-cp39-abi3-win_amd64.whl ; (sys_platform == 'win32')",
]
cu126onlytorch271 = [
"xformers @ https://download.pytorch.org/whl/cu126/xformers-0.0.31.post1-cp39-abi3-manylinux_2_28_x86_64.whl ; ('linux' in sys_platform)",
@@ -879,14 +874,12 @@ flashattentiontorch240abiFALSEcu12x = [
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp310-cp310-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.10'",
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp311-cp311-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.11'",
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp312-cp312-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.12'",
- "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp313-cp313-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.13'",
]
flashattentiontorch240abiTRUEcu12x = [
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp39-cp39-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.9'",
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp310-cp310-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.10'",
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp311-cp311-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.11'",
"flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp312-cp312-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.12'",
- "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp313-cp313-linux_x86_64.whl ; ('linux' in sys_platform) and python_version == '3.13'",
]
intelgputorch260 = [
"unsloth_zoo[intelgpu]",
@@ -1174,14 +1167,14 @@ intelgputorch2120 = [
"unsloth_zoo[intelgpu]",
"unsloth[huggingfacenotorch]",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=844d981cb1b3948085e8cfa62c74de9f100259f6131959aa70be49123b88ae81 ; platform_system == 'Linux' and python_version == '3.10' and platform_machine == 'x86_64'",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=a16b1d00e94ad87d62af3512e390348b8656419598004100c56028bf494f086b ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=4e46e71e077cf483404a4c17ce40d71c5f0e13a81459139d4346ca427b1dd455 ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=4fdaed1bafc51d3a2834656a3420a6686a74ea226508765a49bf15d58ff3a930 ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp310-cp310-win_amd64.whl#sha256=2778b46b22e9fa0916398db299a125027a1b2331c1173b3dd2b9e2cab6263a31 ; sys_platform == 'win32' and python_version == '3.10' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp311-cp311-win_amd64.whl#sha256=ad5b147d04ee0d40f3d4d32f85f5aa3a3beb6cd5799ca026d3d7f4afa3d9e24f ; sys_platform == 'win32' and python_version == '3.11' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp312-cp312-win_amd64.whl#sha256=d9482063af2a308543f23333e32edd738ea87cbb33ade68afda9ae0fd704ccd9 ; sys_platform == 'win32' and python_version == '3.12' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
- "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp313-cp313-win_amd64.whl#sha256=5d4d67f0deb1e851c01b293e602b8dcddad26ca2be61221cee3dc0e1aa0cdefd ; sys_platform == 'win32' and python_version == '3.13' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=81ff0eb0c4fc8e19d2510b28c3e1d9382a3c7d6fdaf6a9f9631a93a030d841cf ; platform_system == 'Linux' and python_version == '3.10' and platform_machine == 'x86_64'",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=55574a68d275b85cd4d5cbf185084bae019ebf09c3f43b0bd2831b14935ec8e7 ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=a31c058c5c2e78ebe490a2e69f2f50caec6b1307ac096e944f116fdc06819d9a ; platform_system == 'Linux' and python_version == '3.12' and platform_machine == 'x86_64'",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl#sha256=e701a31efa0334775f357c98716f3821775aa944219f7888e13c2dfe2daabe2a ; platform_system == 'Linux' and python_version == '3.13' and platform_machine == 'x86_64'",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp310-cp310-win_amd64.whl#sha256=0d7730651c3e52fbf3a430cc201455f0c6600dc72e681aec495f131ea44f341a ; sys_platform == 'win32' and python_version == '3.10' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp311-cp311-win_amd64.whl#sha256=8f4a63de73e3d632098f93c8f0bd77244958a47d7c5f728b8ff35f8a91fdb983 ; sys_platform == 'win32' and python_version == '3.11' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp312-cp312-win_amd64.whl#sha256=6589ece3adc2b1ab88d90ff1267afc25df5c7b868f0b633e732cac70df36cbde ; sys_platform == 'win32' and python_version == '3.12' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
+ "triton-xpu @ https://download.pytorch.org/whl/triton_xpu-3.7.1-cp313-cp313-win_amd64.whl#sha256=2fdf001a9b0575e8b1827127259bb9b13bf36e659882be74c2dfab46597d3e7a ; sys_platform == 'win32' and python_version == '3.13' and (platform_machine == 'AMD64' or platform_machine == 'x86_64')",
"torch @ https://download.pytorch.org/whl/xpu/torch-2.12.0%2Bxpu-cp310-cp310-linux_x86_64.whl#sha256=e8923cd1fe560472904b1461b745d2f1826bb9c1bc0808225d5f28a450e4d553 ; platform_system == 'Linux' and python_version == '3.10' and platform_machine == 'x86_64'",
"torch @ https://download.pytorch.org/whl/xpu/torch-2.12.0%2Bxpu-cp311-cp311-linux_x86_64.whl#sha256=f7c082b2fc9b61def594d30ea57762dc4a8bc7111a9a9593953ed948de242e28 ; platform_system == 'Linux' and python_version == '3.11' and platform_machine == 'x86_64'",
diff --git a/scripts/scan_packages.py b/scripts/scan_packages.py
index 4be9fc5efb..ce9763e235 100644
--- a/scripts/scan_packages.py
+++ b/scripts/scan_packages.py
@@ -1208,9 +1208,10 @@ def check_js_file(content: str, filename: str, package: str) -> list[Finding]:
HIGH,
package,
filename,
- f"Python wheel ships large ({len(content) // 1024} KB) JS bundle "
- "(uncommon; manually review)",
- "",
+ # Size stays in evidence, not the check label, so the baseline key
+ # does not drift when a wheel's bundle grows by a few KB.
+ "Python wheel ships large JS bundle (uncommon; manually review)",
+ f"{len(content) // 1024} KB JS bundle",
)
)
return findings
diff --git a/scripts/scan_packages_baseline.json b/scripts/scan_packages_baseline.json
index f953f4d206..0c10ae3222 100644
--- a/scripts/scan_packages_baseline.json
+++ b/scripts/scan_packages_baseline.json
@@ -1181,7 +1181,7 @@
{
"package": "tensorboard",
"file": "tensorboard/plugins/projector/tf_projector_plugin/projector_binary.js",
- "check": "Python wheel ships large (1918 KB) JS bundle (uncommon; manually review)",
+ "check": "Python wheel ships large JS bundle (uncommon; manually review)",
"severity": "HIGH",
"evidence": ""
},
diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py
index f243f5b65a..d0461dae95 100644
--- a/studio/backend/core/export/export.py
+++ b/studio/backend/core/export/export.py
@@ -38,6 +38,26 @@ logger = get_logger(__name__)
_LLAMA_CPP_SCRIPTS_WARNING_EMITTED = False
+def _supports_kwarg(fn, name):
+ """True if `fn` accepts keyword `name` directly or via **kwargs."""
+ import inspect
+
+ try:
+ params = inspect.signature(fn).parameters
+ except (TypeError, ValueError):
+ return False
+ return name in params or any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values())
+
+
+def _compressed_export_supported():
+ """True if the installed unsloth build can do FP8/NVFP4 compressed-tensors export."""
+ try:
+ import unsloth.save as _us
+ return hasattr(_us, "_normalize_compressed_method")
+ except Exception:
+ return False
+
+
def _hf_offline(timeout = 3):
"""True if export should avoid the Hub: honors the HF offline env vars, else does one
cheap TCP reachability probe so a network-down load uses local files / the HF cache
@@ -400,16 +420,33 @@ class ExportBackend:
)
output_path: Optional[str] = None
+ # compressed-tensors formats run save_pretrained_merged with an FP8/FP4 save_method and
+ # write to a sibling "
-" directory (for vLLM).
+ _COMPRESSED = {
+ "FP8 (compressed-tensors)": ("fp8", "fp8"),
+ "NVFP4 (compressed-tensors)": ("nvfp4", "nvfp4"),
+ }
+ is_compressed = format_type in _COMPRESSED
try:
if _IS_MLX:
+ if is_compressed:
+ return False, "Compressed-tensors export is not supported on macOS/MLX.", None
mlx_save_method = "merged_4bit" if format_type == "4-bit (FP4)" else "merged_16bit"
+ elif is_compressed:
+ if not _compressed_export_supported():
+ return (
+ False,
+ "Compressed-tensors (FP8/NVFP4) export requires an Unsloth build with "
+ "compressed-tensors support. Upgrade unsloth, or choose 16-bit.",
+ None,
+ )
+ save_method = _COMPRESSED[format_type][0]
+ elif format_type == "4-bit (FP4)":
+ save_method = "merged_4bit_forced"
+ elif self._audio_type == "whisper":
+ save_method = None
else:
- if format_type == "4-bit (FP4)":
- save_method = "merged_4bit_forced"
- elif self._audio_type == "whisper":
- save_method = None
- else:
- save_method = "merged_16bit"
+ save_method = "merged_16bit"
if save_directory:
save_directory = str(resolve_export_write_dir(save_directory))
@@ -427,9 +464,15 @@ class ExportBackend:
save_directory, self.current_tokenizer, save_method = save_method
)
- self._write_export_metadata(save_directory)
- logger.info(f"Model saved successfully to {save_directory}")
- output_path = str(Path(save_directory).resolve())
+ # Compressed export writes to the "-" sibling; report that as output.
+ final_dir = (
+ f"{save_directory}-{_COMPRESSED[format_type][1]}"
+ if is_compressed
+ else save_directory
+ )
+ self._write_export_metadata(final_dir)
+ logger.info(f"Model saved successfully to {final_dir}")
+ output_path = str(Path(final_dir).resolve())
if push_to_hub:
if not repo_id or not hf_token:
@@ -464,6 +507,32 @@ class ExportBackend:
token = hf_token,
private = private,
)
+ elif is_compressed and output_path and Path(output_path).is_dir():
+ # The compressed model was already built locally in output_path; upload it
+ # directly so we do not re-run the (expensive, OOM-prone) compression that
+ # push_to_hub_merged(save_method=fp8/nvfp4) would otherwise do a second time.
+ hf_api = HfApi(token = hf_token)
+ repo_id = PushToHubMixin._create_repo(
+ PushToHubMixin,
+ repo_id = repo_id,
+ private = private,
+ token = hf_token,
+ )
+ content = MODEL_CARD.format(
+ username = repo_id.split("/")[0],
+ base_model = getattr(self.current_model.config, "_name_or_path", "unknown"),
+ model_type = getattr(self.current_model.config, "model_type", "llm"),
+ method = format_type,
+ extra = "unsloth",
+ )
+ ModelCard(content).push_to_hub(
+ repo_id, token = hf_token, commit_message = "Unsloth Model Card"
+ )
+ hf_api.upload_folder(
+ folder_path = output_path,
+ repo_id = repo_id,
+ repo_type = "model",
+ )
else:
hub_save_method = save_method if save_method is not None else "merged_16bit"
self.current_model.push_to_hub_merged(
@@ -621,6 +690,7 @@ class ExportBackend:
push_to_hub: bool = False,
repo_id: Optional[str] = None,
hf_token: Optional[str] = None,
+ imatrix_file = None,
) -> Tuple[bool, str, Optional[str]]:
"""
Export model in GGUF format.
@@ -638,6 +708,19 @@ class ExportBackend:
if not self.current_model or not self.current_tokenizer:
return False, "No model loaded. Please select a checkpoint first.", None
+ # Only forward imatrix_file to an unsloth build that accepts it; otherwise even a plain
+ # no-imatrix export would fail with an unexpected-keyword error against an older unsloth.
+ if imatrix_file is not None and not _supports_kwarg(
+ self.current_model.save_pretrained_gguf, "imatrix_file"
+ ):
+ return (
+ False,
+ "This Unsloth build does not support GGUF imatrix export. "
+ "Upgrade unsloth and unsloth_zoo, or disable the imatrix option.",
+ None,
+ )
+ imatrix_kw = {"imatrix_file": imatrix_file} if imatrix_file is not None else {}
+
output_path: Optional[str] = None
model_tmp_to_cleanup: Optional[str] = None
try:
@@ -691,6 +774,7 @@ class ExportBackend:
_model_tmp,
self.current_tokenizer,
quantization_method = quant_method,
+ **imatrix_kw,
)
# Relocate the .gguf that convert_to_gguf wrote to cwd (repo root).
@@ -757,6 +841,7 @@ class ExportBackend:
self.current_tokenizer,
quantization_method = quant_method,
token = hf_token,
+ **imatrix_kw,
)
logger.info(f"GGUF model pushed successfully to {repo_id}")
diff --git a/studio/backend/core/export/orchestrator.py b/studio/backend/core/export/orchestrator.py
index 478624b48e..052a47dd80 100644
--- a/studio/backend/core/export/orchestrator.py
+++ b/studio/backend/core/export/orchestrator.py
@@ -499,6 +499,7 @@ class ExportOrchestrator:
push_to_hub: bool = False,
repo_id: Optional[str] = None,
hf_token: Optional[str] = None,
+ imatrix_file = None,
) -> Tuple[bool, str, Optional[str]]:
"""Export model in GGUF format."""
return self._run_export(
@@ -509,6 +510,7 @@ class ExportOrchestrator:
"push_to_hub": push_to_hub,
"repo_id": repo_id,
"hf_token": hf_token,
+ "imatrix_file": imatrix_file,
},
)
diff --git a/studio/backend/core/export/worker.py b/studio/backend/core/export/worker.py
index 71a603a857..d473dcb54f 100644
--- a/studio/backend/core/export/worker.py
+++ b/studio/backend/core/export/worker.py
@@ -414,6 +414,7 @@ def _handle_export(backend, cmd: dict, resp_queue: Any) -> None:
push_to_hub = cmd.get("push_to_hub", False),
repo_id = cmd.get("repo_id"),
hf_token = cmd.get("hf_token"),
+ imatrix_file = cmd.get("imatrix_file"),
)
elif export_type == "lora":
success, message, output_path = backend.export_lora_adapter(
diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py
index 31969afb14..af62a29679 100644
--- a/studio/backend/core/inference/llama_cpp.py
+++ b/studio/backend/core/inference/llama_cpp.py
@@ -1271,6 +1271,9 @@ class LlamaCppBackend:
self._cache_type_kv: Optional[str] = None
# Whether --split-mode tensor was applied on the active load.
self._tensor_parallel: bool = False
+ # 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
self._reasoning_default: bool = True
self._speculative_type: Optional[str] = None
# Canonical UI-facing mode the user requested
@@ -1643,6 +1646,11 @@ class LlamaCppBackend:
"""Whether --split-mode tensor is active on the loaded server."""
return self._tensor_parallel
+ @property
+ def layer_preserves_tensor_intent(self) -> bool:
+ """True when a downgraded tensor request kept this layer load multi-GPU."""
+ return self._layer_preserves_tensor_intent
+
@property
def speculative_type(self) -> Optional[str]:
return self._speculative_type
@@ -2430,6 +2438,37 @@ class LlamaCppBackend:
# aborts a --split-mode tensor load, so it's dropped for the tensor attempt.
_TENSOR_PARALLEL_KV_TYPES = frozenset({"f16", "bf16", "f32"})
+ # (binary, mtime, model) that aborted on --split-mode tensor this process (#6415
+ # geometry limit, e.g. MQA n_head_kv=1). Model-keyed so one model's abort doesn't
+ # skip tensor for others; tensor is tried by default, recorded only on a real abort.
+ _tensor_split_abort_keys: set[tuple[str, int, str]] = set()
+
+ @classmethod
+ def _tensor_split_cache_key(
+ cls, binary: Optional[str], model: Optional[str]
+ ) -> Optional[tuple[str, int, str]]:
+ """(path, mtime_ns, model) key; ns mtime re-probes a same-second binary swap."""
+ if not binary or not model:
+ return None
+ try:
+ mtime = Path(binary).stat().st_mtime_ns
+ except OSError:
+ mtime = 0
+ return (binary, mtime, model)
+
+ @classmethod
+ def _tensor_split_aborts(cls, binary: Optional[str], model: Optional[str]) -> bool:
+ """True if (binary, model) aborted on --split-mode tensor this session."""
+ key = cls._tensor_split_cache_key(binary, model)
+ return key is not None and key in cls._tensor_split_abort_keys
+
+ @classmethod
+ def _record_tensor_split_abort(cls, binary: Optional[str], model: Optional[str]) -> None:
+ """Remember a (binary, model) that aborts on --split-mode tensor."""
+ key = cls._tensor_split_cache_key(binary, model)
+ if key is not None:
+ cls._tensor_split_abort_keys.add(key)
+
@staticmethod
def _windows_pip_nvidia_dll_dirs(prefix: str) -> list[str]:
"""Return DLL dirs from pip-installed CUDA wheels under
@@ -2569,9 +2608,13 @@ class LlamaCppBackend:
usable_fraction: Optional[float] = None,
total_by_idx: Optional[dict[int, int]] = None,
per_device_overhead_bytes: int = 0,
+ min_gpus: int = 1,
) -> tuple[Optional[list[int]], bool]:
"""Pick GPU(s) for a model from estimated VRAM and free memory.
+ ``min_gpus`` (default 1, capped at ``len(gpus)``) keeps a downgraded
+ tensor/multi-GPU request spread instead of collapsing to one card.
+
``model_size_bytes`` should include weights and estimated KV cache.
``usable_fraction`` (default ``_GPU_PIN_VRAM_FRACTION``) provides
headroom for compute buffers, CUDA context, and other runtime
@@ -2590,9 +2633,11 @@ class LlamaCppBackend:
if not gpus:
return None, True
+ min_gpus = max(1, min(min_gpus, len(gpus)))
model_size_mib = model_size_bytes / (1024 * 1024)
if usable_fraction is None:
usable_fraction = LlamaCppBackend._GPU_PIN_VRAM_FRACTION
+ overhead_mib = per_device_overhead_bytes / (1024 * 1024)
# Per-GPU usable budget: free - (1-frac)*total when total is known, else
# the legacy free*frac (also covers a total-0 two-column probe).
@@ -2606,19 +2651,26 @@ class LlamaCppBackend:
# card can have less usable room than a less-used small one.
ranked = sorted(gpus, key = lambda g: _usable(g[0], g[1]), reverse = True)
- # Try 1 GPU at the usable-VRAM threshold.
- if _usable(ranked[0][0], ranked[0][1]) >= model_size_mib:
+ # Cap a downgraded multi-GPU request to the usable count so it doesn't pull
+ # in a near-full card to hit min_gpus. No-op for the default min_gpus == 1.
+ usable_count = sum(1 for idx, free_mib in ranked if _usable(idx, free_mib) > overhead_mib)
+ min_gpus = max(1, min(min_gpus, usable_count or 1))
+
+ # Try 1 GPU at the usable-VRAM threshold (only when one device is allowed).
+ if min_gpus <= 1 and _usable(ranked[0][0], ranked[0][1]) >= model_size_mib:
return [ranked[0][0]], False
- # Try N GPUs (accumulate usable memory from most-free). Each GPU past the
- # first adds a fixed per-device overhead the pool must hold.
- overhead_mib = per_device_overhead_bytes / (1024 * 1024)
+ # Try N GPUs (most-free first); each past the first adds per-device overhead.
+ # Require at least min_gpus devices before accepting a fit.
cumulative = 0.0
selected = []
for idx, free_mib in ranked:
selected.append(idx)
cumulative += _usable(idx, free_mib)
- if cumulative >= model_size_mib + (len(selected) - 1) * overhead_mib:
+ if (
+ len(selected) >= min_gpus
+ and cumulative >= model_size_mib + (len(selected) - 1) * overhead_mib
+ ):
return sorted(selected), False
# Too large even for all GPUs; let --fit handle it
@@ -3645,12 +3697,22 @@ class LlamaCppBackend:
hf_repo: str,
hf_variant: Optional[str] = None,
hf_token: Optional[str] = None,
+ force: bool = False,
+ allow_smaller_fallback: bool = True,
+ cancel_event: Optional[threading.Event] = None,
) -> str:
"""Download GGUF file(s) from HuggingFace. Returns local path.
Runs WITHOUT self._lock so unload_model() can set _cancel_event at
any time; checks it between each shard download.
+
+ ``force`` re-fetches even when a (possibly stale) blob is cached.
+ ``allow_smaller_fallback=False`` raises on low disk instead of silently
+ switching to a smaller quant. ``cancel_event`` overrides
+ ``self._cancel_event`` so an update can use a private event without
+ touching the shared one; defaults to the shared event.
"""
+ cancel_event = cancel_event if cancel_event is not None else self._cancel_event
try:
import huggingface_hub # noqa: F401 -- presence check only
except ImportError:
@@ -3716,21 +3778,22 @@ class LlamaCppBackend:
# cold whenever free disk is below the full weight footprint,
# even though nothing needs downloading.
already_cached_bytes = 0
- for p in path_infos:
- if not p.size:
- continue
- try:
- cached_path = try_to_load_from_cache(hf_repo, p.path)
- except Exception:
- cached_path = None
- if isinstance(cached_path, str) and os.path.exists(cached_path):
+ if not force:
+ for p in path_infos:
+ if not p.size:
+ continue
try:
- on_disk = os.path.getsize(cached_path)
- except OSError:
- on_disk = 0
- # Satisfied only when the full blob is present.
- if on_disk >= p.size:
- already_cached_bytes += p.size
+ cached_path = try_to_load_from_cache(hf_repo, p.path)
+ except Exception:
+ cached_path = None
+ if isinstance(cached_path, str) and os.path.exists(cached_path):
+ try:
+ on_disk = os.path.getsize(cached_path)
+ except OSError:
+ on_disk = 0
+ # Satisfied only when the full blob is present.
+ if on_disk >= p.size:
+ already_cached_bytes += p.size
total_download_bytes = max(0, total_bytes - already_cached_bytes)
@@ -3753,6 +3816,13 @@ class LlamaCppBackend:
)
if total_download_bytes > free_bytes:
+ if not allow_smaller_fallback:
+ # Update path: never silently switch to a smaller quant;
+ # surface the disk shortfall for the requested variant.
+ raise RuntimeError(
+ f"Not enough disk space to download {gguf_filename}. "
+ f"Only {free_gb:.1f} GB free in {cache_dir}"
+ )
smaller = self._find_smallest_fitting_variant(
hf_repo,
free_bytes,
@@ -3793,7 +3863,7 @@ class LlamaCppBackend:
)
logger.info(f"Resolving GGUF: {gguf_label}")
try:
- if self._cancel_event.is_set():
+ if cancel_event.is_set():
raise RuntimeError("Cancelled")
dl_start = time.monotonic()
# Xet primary, HTTP fallback on stall; per-file so finished shards stay cached.
@@ -3801,18 +3871,20 @@ class LlamaCppBackend:
hf_repo,
gguf_filename,
hf_token,
- cancel_event = self._cancel_event,
+ cancel_event = cancel_event,
on_status = lambda m: logger.info(m),
+ force_download = force,
)
for shard in gguf_extra_shards:
- if self._cancel_event.is_set():
+ if cancel_event.is_set():
raise RuntimeError("Cancelled")
logger.info(f"Resolving GGUF shard: {shard}")
hf_hub_download_with_xet_fallback(
hf_repo,
shard,
hf_token,
- cancel_event = self._cancel_event,
+ cancel_event = cancel_event,
+ force_download = force,
)
except Exception as e:
if isinstance(e, RuntimeError) and "Cancelled" in str(e):
@@ -3835,6 +3907,7 @@ class LlamaCppBackend:
hf_token: Optional[str],
pick: Callable[[list[str]], Optional[str]],
label: str,
+ cancel_event: Optional[threading.Event] = None,
) -> Optional[str]:
"""Resolve and fetch a companion GGUF (mmproj / MTP drafter) by name.
@@ -3842,8 +3915,10 @@ class LlamaCppBackend:
(offline, same fallback as _download_gguf), then hf_hub_download.
Runs WITHOUT self._lock (like _download_gguf); honors _cancel_event so
an /unload between the main download and here skips the fetch.
+ ``cancel_event`` overrides ``self._cancel_event`` (defaults to it).
"""
- if self._cancel_event.is_set():
+ cancel_event = cancel_event if cancel_event is not None else self._cancel_event
+ if cancel_event.is_set():
return None
target: Optional[str] = None
@@ -3852,7 +3927,7 @@ class LlamaCppBackend:
# Retry a transient listing blip; permanent repo/auth errors and offline
# mode are not retried (offline raises at once -> fall through to cache).
for attempt in range(3):
- if self._cancel_event.is_set():
+ if cancel_event.is_set():
return None
try:
target = pick(list_repo_files(hf_repo, token = hf_token))
@@ -3868,10 +3943,10 @@ class LlamaCppBackend:
logger.debug(f"Could not list repo files for {label}: {e}")
break
logger.debug(
- f"Could not list repo files for {label} " f"(attempt {attempt + 1}/3): {e}"
+ f"Could not list repo files for {label} (attempt {attempt + 1}/3): {e}"
)
if attempt < 2:
- self._cancel_event.wait(2**attempt)
+ cancel_event.wait(2**attempt)
if target is None:
try:
@@ -3885,7 +3960,7 @@ class LlamaCppBackend:
except Exception as e:
logger.debug(f"Offline cache lookup for {label} failed: {e}")
- if target is None or self._cancel_event.is_set():
+ if target is None or cancel_event.is_set():
return None
try:
@@ -3895,7 +3970,7 @@ class LlamaCppBackend:
hf_repo,
target,
hf_token,
- cancel_event = self._cancel_event,
+ cancel_event = cancel_event,
)
except Exception as e:
logger.warning(f"Could not download {label}: {e}")
@@ -3906,11 +3981,13 @@ class LlamaCppBackend:
*,
hf_repo: str,
hf_token: Optional[str] = None,
+ cancel_event: Optional[threading.Event] = None,
) -> Optional[str]:
"""Download the mmproj (vision projection) file from a GGUF repo.
Prefers mmproj-F16.gguf, else any mmproj*.gguf. Returns the local
- path, or None if none exists.
+ path, or None if none exists. ``cancel_event`` overrides
+ ``self._cancel_event`` (defaults to it).
"""
def _pick_mmproj(candidates: list[str]) -> Optional[str]:
@@ -3931,6 +4008,7 @@ class LlamaCppBackend:
hf_token = hf_token,
pick = _pick_mmproj,
label = "mmproj",
+ cancel_event = cancel_event,
)
def _download_mtp(
@@ -4332,6 +4410,17 @@ class LlamaCppBackend:
)
)
+ @staticmethod
+ def _is_tensor_split_assert(output: str) -> bool:
+ """True only for the #6415 split-axis warmup assert (GGML_BACKEND_SPLIT_AXIS_*),
+ not any ggml assert/abort, so an unrelated invariant isn't cached. stderr is
+ merged into output."""
+ text = (output or "").lower()
+ if "ggml_assert" not in text and "ggml_abort" not in text:
+ return False
+ # the split-axis enum token, unique to this assert (not the source file).
+ return "split_axis" in text
+
@staticmethod
def _is_signal_crash(returncode: Optional[int]) -> bool:
"""True only on a hard fault (SIGSEGV/SIGABRT/SIGILL/SIGFPE/SIGBUS or a
@@ -4344,6 +4433,20 @@ class LlamaCppBackend:
return True
return -returncode in (4, 6, 7, 8, 11) # SIGILL SIGABRT SIGBUS SIGFPE SIGSEGV
+ @staticmethod
+ def _is_abort_exit(returncode: Optional[int]) -> bool:
+ """Windows CRT abort() exit code (3) from GGML_ASSERT on MSVC -- not a POSIX
+ signal or 0xC0000000+ NTSTATUS."""
+ return returncode == 3
+
+ @classmethod
+ def _should_record_tensor_split_abort(cls, returncode: Optional[int], output: str) -> bool:
+ """The #6415 split-axis abort: the marker plus a hard crash (POSIX signal or
+ Windows abort exit). Marker required so a generic crash isn't cached."""
+ return cls._is_tensor_split_assert(output) and (
+ cls._is_signal_crash(returncode) or cls._is_abort_exit(returncode)
+ )
+
@staticmethod
def _with_flash_attn_off(cmd: list[str]) -> Optional[list[str]]:
"""Return cmd with flash attention forced off, or None when its effective
@@ -4488,6 +4591,8 @@ class LlamaCppBackend:
n_gpu_layers: Optional[int] = None, # caller compat, unused
n_parallel: int = 1,
extra_args: Optional[List[str]] = None,
+ # Route-level tensor->layer fallback retry: keep the layer split multi-GPU.
+ preserve_multi_gpu_on_layer: bool = False,
) -> bool:
"""Start llama-server with a GGUF model.
@@ -4518,6 +4623,8 @@ class LlamaCppBackend:
"n_gpu_layers": n_gpu_layers,
"n_parallel": n_parallel,
"extra_args": list(extra_args) if extra_args is not None else None,
+ # Replayed by _respawn_if_dead so a downgraded model stays multi-GPU.
+ "preserve_multi_gpu_on_layer": preserve_multi_gpu_on_layer,
}
# Serialise the whole load so concurrent /load calls never leave two
# llama-server processes alive (#5401 / #5161). Doesn't block /unload.
@@ -4541,6 +4648,7 @@ class LlamaCppBackend:
chat_template_override = chat_template_override,
extra_args = extra_args,
is_vision = is_vision,
+ preserve_multi_gpu_on_layer = preserve_multi_gpu_on_layer,
):
logger.info(
f"load_model: backend already in target state for "
@@ -4626,6 +4734,9 @@ 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:
+ # 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
with self._lock:
if self._cancel_event.is_set():
logger.info("Load cancelled before diffusion server start")
@@ -4780,6 +4891,9 @@ class LlamaCppBackend:
"image input will be disabled for this session"
)
model_size = None # set in the fit try; used by the APU RAM guard
+ # 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
try:
gguf_size = self._get_gguf_size_bytes(model_path)
# Include GPU-loaded mmproj in the fit budget (#5825).
@@ -5064,10 +5178,8 @@ class LlamaCppBackend:
_apple_budget_mib = self._apple_metal_memory_budget_bytes() // (1024 * 1024)
def _restore_after_tensor_downgrade():
- # Tensor mode dropped a quantized KV and stripped the cache
- # extras (it rejects quantized); layer split supports them, so
- # restore the original type + extras (minus --split-mode) and
- # clear the env flag so the layer launch re-emits them.
+ # Restore the quantized KV + extras tensor dropped (layer
+ # split supports them), minus --split-mode.
nonlocal cache_type_kv, _cache_type_from_env, extra_args
if _tensor_dropped_cache_type_kv is not None:
cache_type_kv = _tensor_dropped_cache_type_kv
@@ -5078,13 +5190,22 @@ class LlamaCppBackend:
else extra_args
)
- if tensor_parallel and effective_is_vision:
+ # The route fallback retry is tensor-off; keep it multi-GPU.
+ if preserve_multi_gpu_on_layer:
+ _layer_min_gpus = max(_layer_min_gpus, len(gpus))
+
+ if tensor_parallel and self._tensor_split_aborts(binary, model_identifier):
+ # Aborted on tensor for this model this session (#6415); skip
+ # tensor upfront, layer split serves it.
logger.info(
- "Tensor parallelism skipped for vision model: "
- "--split-mode tensor is incompatible with --mmproj "
- "in the current llama.cpp build; using layer split."
+ "Tensor parallelism skipped: this llama.cpp build aborted "
+ "on --split-mode tensor for this model earlier this "
+ "session; using layer split across %d GPU(s).",
+ len(gpus),
)
tensor_parallel = False
+ # Keep the multi-GPU request (gated on it, not the cache).
+ _layer_min_gpus = max(_layer_min_gpus, len(gpus))
_restore_after_tensor_downgrade()
# Tensor mode replicates a compute buffer on every GPU, so drop
@@ -5124,6 +5245,11 @@ class LlamaCppBackend:
len(gpus),
)
tensor_parallel = False
+ # GPUs below tensor's compute-buffer reserve can still do layer
+ # split, so keep multi-GPU (mirrors the budget/geometry drops);
+ # _select_gpus caps unusable cards.
+ if len(gpus) >= 2:
+ _layer_min_gpus = max(_layer_min_gpus, len(gpus))
# Layer split supports a quantized KV the tensor attempt
# dropped; restore the original cache type + extras (minus
# --split-mode) so the layer launch re-emits them.
@@ -5160,8 +5286,12 @@ class LlamaCppBackend:
"per-device compute buffers; falling back to layer split."
)
tensor_parallel = False
- # Restore the dropped quantized KV + original cache extras
- # (minus --split-mode); layer split supports them.
+ # Weights needed >1 card, so keep multi-GPU across the
+ # usable tensor GPUs.
+ if len(tp_gpus) >= 2:
+ _layer_min_gpus = max(_layer_min_gpus, len(tp_gpus))
+ # Restore the dropped quantized KV + cache extras (minus
+ # --split-mode); layer split supports them.
_restore_after_tensor_downgrade()
if tensor_parallel and tp_gpus:
@@ -5263,6 +5393,7 @@ class LlamaCppBackend:
usable_fraction = _pin_fraction,
total_by_idx = total_by_idx,
per_device_overhead_bytes = _pipeline_overhead_bytes,
+ min_gpus = _layer_min_gpus,
)
# No silent shrink: effective_ctx stays == requested_ctx.
else:
@@ -5273,7 +5404,22 @@ class LlamaCppBackend:
ranked = sorted(
gpus, key = lambda g: _gpu_usable(g, pin_fraction), reverse = True
)
- for n_gpus in range(1, len(ranked) + 1):
+ # Skips _select_gpus, so apply its cap: count only cards
+ # whose usable VRAM clears the per-device layer overhead.
+ _pipeline_overhead_mib = _pipeline_overhead_bytes / (1024 * 1024)
+ _auto_min_gpus = max(
+ 1,
+ min(
+ _layer_min_gpus,
+ sum(
+ 1
+ for g in ranked
+ if _gpu_usable(g, pin_fraction) > _pipeline_overhead_mib
+ )
+ or 1,
+ ),
+ )
+ for n_gpus in range(_auto_min_gpus, len(ranked) + 1):
subset = ranked[:n_gpus]
pool_budget = _pool_budget_mib(subset, pin_fraction)
_ms = _subset_model_size(n_gpus)
@@ -5303,7 +5449,7 @@ class LlamaCppBackend:
# at 131k may pin fine with a 4096 KV (#5106).
effective_ctx = min(4096, effective_ctx)
if effective_ctx > 0:
- for n_gpus in range(1, len(ranked) + 1):
+ for n_gpus in range(_auto_min_gpus, len(ranked) + 1):
subset = ranked[:n_gpus]
kv = self._estimate_kv_cache_bytes(
effective_ctx,
@@ -5339,6 +5485,7 @@ class LlamaCppBackend:
usable_fraction = _pin_fraction,
total_by_idx = total_by_idx,
per_device_overhead_bytes = _pipeline_overhead_bytes,
+ min_gpus = _layer_min_gpus,
)
if use_fit and not explicit_ctx:
# Weights don't fit on any subset; default UI to 4096
@@ -5578,12 +5725,15 @@ class LlamaCppBackend:
]
)
self._tensor_parallel = True
+ self._layer_preserves_tensor_intent = False
logger.info(
"Tensor parallelism: --split-mode tensor, --tensor-split %s",
tp_tensor_split,
)
else:
self._tensor_parallel = False
+ # > 1 only when a tensor request was downgraded but kept multi-GPU.
+ self._layer_preserves_tensor_intent = _layer_min_gpus > 1
# Speculative decoding. See _build_speculative_flags for the
# mode resolution, benchmarks, and llama.cpp references.
@@ -5867,7 +6017,17 @@ class LlamaCppBackend:
_startup_crashed = (
self._process.poll() is not None and self._process.returncode != 0
)
- if _spawn_attempt == 0 and _fit_retry_allowed and _startup_crashed:
+ # A split-axis abort (#6415) is fit-independent: skip the
+ # --fit off retry and let the caller latch it.
+ _split_axis_crash = self._is_tensor_split_assert(
+ "\n".join(self._stdout_lines[-50:])
+ )
+ if (
+ _spawn_attempt == 0
+ and _fit_retry_allowed
+ and _startup_crashed
+ and not _split_axis_crash
+ ):
logger.warning(
"llama-server crashed during startup (exit code %s) "
"with the default memory-fit step enabled; Studio "
@@ -5913,6 +6073,21 @@ class LlamaCppBackend:
)
healthy = _spawn_and_wait(cmd)
+ # #6415 split-mode tensor warmup abort. Latch it on THIS first spawn:
+ # the flash-attn-off retry below can't run tensor (needs flash_attn),
+ # so its output drops the marker and recording later would miss it,
+ # looping every load. Record and raise to the route's layer fallback,
+ # skipping the futile flash-attn/MTP retries.
+ if not healthy and self._tensor_parallel and not self._cancel_event.is_set():
+ _ts_out = "\n".join(self._stdout_lines[-50:])
+ _ts_rc = self._process.poll() if self._process is not None else None
+ if self._should_record_tensor_split_abort(_ts_rc, _ts_out):
+ LlamaCppBackend._record_tensor_split_abort(binary, model_identifier)
+ self._kill_process()
+ raise RuntimeError(
+ "llama-server aborted on --split-mode tensor "
+ "(split-axis geometry); retrying with layer split."
+ )
# Flash-attention kernels hard-crash at startup on some ROCm/GPU
# builds (frequently inside the vision tower). Disabling FA keeps
# both vision and MTP, so retry that way before dropping either.
@@ -6057,6 +6232,7 @@ class LlamaCppBackend:
# Read the crash code before _kill_process() clears _process.
_crash_rc = self._process.poll() if self._process is not None else None
self._kill_process()
+ # The #6415 split-axis abort is latched earlier (first spawn).
# Skip if a cancel/unload is pending (mirrors the MTP guard).
if (
launched_with_mmproj
@@ -6488,6 +6664,7 @@ class LlamaCppBackend:
spec_draft_n_max: Optional[int] = None,
tensor_parallel: bool = False,
mtp_draft_path: Optional[str] = None,
+ preserve_multi_gpu_on_layer: bool = False,
) -> bool:
"""True iff the live server already satisfies these load kwargs.
@@ -6530,6 +6707,17 @@ class LlamaCppBackend:
# server. An identical request would downgrade the same way.
if not _tensor_parallel_matches_loaded(extra_args, tensor_parallel, self._tensor_parallel):
return False
+ # Preserved tensor->layer fallback + an EXPLICIT tensor drop: reload so
+ # placement re-selects instead of keeping the all-GPU mask (mirrors the route,
+ # #6659). preserve_multi_gpu_on_layer carries the route's carry-forward decision
+ # (True for an implicit same-settings reload), so those still dedupe -- the HF
+ # auto-pick / local-dir flows skip the route guard and only reach here.
+ if (
+ self._layer_preserves_tensor_intent
+ and not _effective_tensor_parallel(extra_args, tensor_parallel)
+ and not preserve_multi_gpu_on_layer
+ ):
+ return False
# Compare on the canonical requested mode. With --spec-type in
# extra_args the backend stores None; mirror that here.
@@ -6641,6 +6829,7 @@ class LlamaCppBackend:
self._supports_tools = False
self._cache_type_kv = None
self._tensor_parallel = False
+ self._layer_preserves_tensor_intent = False
self._speculative_type = None
self._requested_spec_mode = None
self._spec_draft_n_max = None
@@ -7110,7 +7299,13 @@ class LlamaCppBackend:
url = f"{self.base_url}/completion"
payload = {"prompt": "Hi", "n_predict": 4, "temperature": 0.0, "stream": False}
try:
- resp = httpx.post(url, json = payload, timeout = timeout, headers = self._auth_headers)
+ resp = httpx.post(
+ url,
+ json = payload,
+ timeout = timeout,
+ headers = self._auth_headers,
+ trust_env = False,
+ )
except Exception as e:
logger.debug(f"MTP decode probe failed: {e}")
return False
@@ -7262,7 +7457,9 @@ class LlamaCppBackend:
return False
try:
- resp = httpx.get(url, timeout = 2.0)
+ # trust_env=False: skip ambient HTTP(S)_PROXY, which if it 503s
+ # for 127.0.0.1 loops the probe until timeout and hangs load.
+ resp = httpx.get(url, timeout = 2.0, trust_env = False)
if resp.status_code == 200:
return True
except (
@@ -7309,7 +7506,7 @@ class LlamaCppBackend:
"""
url = f"{self.base_url}/props"
try:
- resp = httpx.get(url, timeout = 5.0)
+ resp = httpx.get(url, timeout = 5.0, trust_env = False)
if resp.status_code != 200:
return None
settings = resp.json().get("default_generation_settings") or {}
@@ -7389,7 +7586,9 @@ class LlamaCppBackend:
which differ only in how they parse the SSE body."""
stream_timeout = httpx.Timeout(connect = 10, read = 0.5, write = 10, pool = 10)
with httpx.Client(
- timeout = stream_timeout, limits = httpx.Limits(max_keepalive_connections = 0)
+ timeout = stream_timeout,
+ limits = httpx.Limits(max_keepalive_connections = 0),
+ trust_env = False,
) as client:
first_token_deadline = time.monotonic() + _DEFAULT_FIRST_TOKEN_TIMEOUT_S
with self._stream_with_retry(
@@ -8881,7 +9080,7 @@ class LlamaCppBackend:
system_text = _block_text(system)
try:
- with httpx.Client(timeout = 10, headers = self._auth_headers) as client:
+ with httpx.Client(timeout = 10, headers = self._auth_headers, trust_env = False) as client:
def _tokenize(text: str) -> int:
r = client.post(
@@ -8997,7 +9196,7 @@ class LlamaCppBackend:
"""Codec name on match, None on non-audio, raises on transport/JSON errors."""
if not self.is_loaded:
return None
- with httpx.Client(timeout = 10, headers = self._auth_headers) as client:
+ with httpx.Client(timeout = 10, headers = self._auth_headers, trust_env = False) as client:
def _detok(tid: int) -> str:
# Non-200 means "marker not in vocab" -- keep probing.
@@ -9112,7 +9311,9 @@ class LlamaCppBackend:
payload["n_probs"] = 1
with httpx.Client(
- timeout = httpx.Timeout(300, connect = 10), headers = self._auth_headers
+ timeout = httpx.Timeout(300, connect = 10),
+ headers = self._auth_headers,
+ trust_env = False,
) as client:
resp = client.post(f"{self.base_url}/completion", json = payload)
if resp.status_code != 200:
diff --git a/studio/backend/core/inference/llama_http.py b/studio/backend/core/inference/llama_http.py
index b554949c3e..8aa072e35b 100644
--- a/studio/backend/core/inference/llama_http.py
+++ b/studio/backend/core/inference/llama_http.py
@@ -22,11 +22,7 @@ _LIMITS = httpx.Limits(max_connections = 64, max_keepalive_connections = 32)
def _new_client() -> httpx.AsyncClient:
- try:
- return httpx.AsyncClient(limits = _LIMITS)
- except Exception:
- # Mirror external_provider: an unsupported env proxy scheme can raise.
- return httpx.AsyncClient(limits = _LIMITS, trust_env = False)
+ return httpx.AsyncClient(limits = _LIMITS, trust_env = False)
# One client per running event loop: an httpx client binds its transport to the
diff --git a/studio/backend/core/inference/tools.py b/studio/backend/core/inference/tools.py
index a5c193ff39..82c50933fc 100644
--- a/studio/backend/core/inference/tools.py
+++ b/studio/backend/core/inference/tools.py
@@ -1121,6 +1121,61 @@ def _autoinject_top_k() -> int:
return _AUTOINJECT_DEFAULT_TOP_K
+def _thread_whole_doc_enabled(scope: dict) -> bool:
+ """Whether a thread-attached file should be injected in full rather than
+ retrieved top-K. ``rag_scope.whole_doc=False`` disables it for this request."""
+ override = scope.get("whole_doc")
+ if override is False:
+ return False
+ try:
+ from core.rag import config as _rag_config
+ except Exception: # noqa: BLE001
+ return True
+ return _rag_config.THREAD_WHOLE_DOC
+
+
+_IMAGE_PART_TOKEN_ESTIMATE = 1024
+
+
+def _message_token_estimate(conversation: list[dict]) -> int:
+ """Cheap prompt-size estimate for budget guards; exact tokenization happens later."""
+ total = 0
+ for msg in conversation:
+ content = msg.get("content")
+ if isinstance(content, str):
+ total += max(1, len(content) // 4)
+ elif isinstance(content, list):
+ for part in content:
+ if isinstance(part, dict):
+ if part.get("type") in ("image_url", "input_image"):
+ total += _IMAGE_PART_TOKEN_ESTIMATE
+ else:
+ total += max(1, len(str(part.get("text") or "")) // 4)
+ total += 4 # chat-template role / separator overhead estimate
+ return total
+
+
+def _whole_doc_budget(scope: dict | None = None, conversation: list[dict] | None = None) -> int:
+ try:
+ from core.rag import config as _rag_config
+ except Exception: # noqa: BLE001
+ budget = 6000
+ else:
+ budget = _rag_config.WHOLE_DOC_MAX_TOKENS
+ if not scope:
+ return budget
+ context = _opt_int(scope.get("context_length") or scope.get("max_context_tokens"))
+ if context is None or context <= 0:
+ return budget
+ headroom = _opt_int(scope.get("response_headroom"))
+ if headroom is None:
+ headroom = max(1024, context // 4)
+ used = _message_token_estimate(conversation or [])
+ # Leave room for tool XML wrappers, citation metadata, and chat-template overhead.
+ available = context - headroom - used - 512
+ return min(budget, max(0, available))
+
+
def _last_user_text(conversation: list[dict]) -> str:
"""Plain text of the most recent user turn (text parts only)."""
for msg in reversed(conversation):
@@ -1154,7 +1209,11 @@ def build_rag_autoinject(conversation: list[dict], rag_scope: dict | None) -> di
enabled = rag_scope.get("autoinject")
if enabled is None:
enabled = _autoinject_enabled()
- if not enabled:
+ thread_id = rag_scope.get("thread_id")
+ whole_doc_requested = (
+ bool(thread_id) and not rag_scope.get("kb_id") and _thread_whole_doc_enabled(rag_scope)
+ )
+ if not enabled and not whole_doc_requested:
return None
query = _last_user_text(conversation)
if not query:
@@ -1163,35 +1222,81 @@ def build_rag_autoinject(conversation: list[dict], rag_scope: dict | None) -> di
from storage import rag_db
if not rag_db.RAG_AVAILABLE:
return None
- from core.rag.tool import search_for_autoinject
+ from core.rag.tool import render_sources, search_for_autoinject, whole_document_context
except Exception as exc: # noqa: BLE001
logger.warning("RAG auto-inject unavailable: %s", exc)
return None
+ text: str | None = None
+ sources: list[dict] = []
+
floor_override = rag_scope.get("autoinject_min_score")
floor = float(floor_override) if floor_override is not None else _autoinject_floor()
# Cap at the lean top_k, but honor a lower user setting.
lean_k = _autoinject_top_k()
sidebar_k = _opt_int(rag_scope.get("default_top_k"))
top_k = min(sidebar_k, lean_k) if sidebar_k is not None else lean_k
- try:
- found = search_for_autoinject(
- query = query,
- scope_kb_id = rag_scope.get("kb_id"),
- scope_thread_id = rag_scope.get("thread_id"),
- scope_project_id = rag_scope.get("project_id"),
- top_k = top_k,
- min_dense_score = floor,
- **_scope_retrieval_kwargs(rag_scope),
- )
- except Exception as exc: # noqa: BLE001
- logger.warning("RAG auto-inject retrieval failed: %s", exc)
- return None
- if not found:
- logger.info("RAG auto-inject: no passage >= %.2f; skipping", floor)
+
+ # Whole-document mode: a thread-attached file under budget is injected in full so
+ # the model reads everything. A KB selection is exclusive, so whole-doc never
+ # preempts it; in a project chat the project sources are still retrieved top-K and
+ # appended under one citation numbering. Oversized files (or no thread doc) fall
+ # through to the combined top-K retrieval below.
+ if whole_doc_requested:
+ try:
+ budget = _whole_doc_budget(rag_scope, conversation)
+
+ whole = whole_document_context(
+ scope_thread_id = thread_id,
+ max_tokens = budget,
+ )
+ except Exception as exc: # noqa: BLE001
+ logger.warning("RAG whole-document context failed: %s", exc)
+ whole = None
+ if whole is not None:
+ text, sources = whole
+ project_id = rag_scope.get("project_id")
+ if project_id:
+ try:
+ proj = search_for_autoinject(
+ query = query,
+ scope_project_id = project_id,
+ top_k = top_k,
+ min_dense_score = floor,
+ **_scope_retrieval_kwargs(rag_scope),
+ )
+ except Exception as exc: # noqa: BLE001
+ logger.warning("RAG project retrieval (whole-doc companion) failed: %s", exc)
+ proj = None
+ if proj is not None:
+ merged = sources + proj[1]
+ merged_text = render_sources(merged)
+ if max(1, len(merged_text) // 4) <= budget:
+ sources = merged
+ text = merged_text
+ logger.info("RAG auto-inject: whole-document context (%d chunk(s))", len(sources))
+
+ if text is None and enabled:
+ try:
+ found = search_for_autoinject(
+ query = query,
+ scope_kb_id = rag_scope.get("kb_id"),
+ scope_thread_id = rag_scope.get("thread_id"),
+ scope_project_id = rag_scope.get("project_id"),
+ top_k = top_k,
+ min_dense_score = floor,
+ **_scope_retrieval_kwargs(rag_scope),
+ )
+ except Exception as exc: # noqa: BLE001
+ logger.warning("RAG auto-inject retrieval failed: %s", exc)
+ return None
+ if not found:
+ logger.info("RAG auto-inject: no passage >= %.2f; skipping", floor)
+ return None
+ text, sources = found
+ if text is None:
return None
- text, sources = found
import json as _json
import uuid as _uuid
@@ -1236,7 +1341,7 @@ def build_rag_autoinject(conversation: list[dict], rag_scope: dict | None) -> di
"content": text,
},
]
- logger.info("RAG auto-inject: %d passage(s) >= %.2f for %r", len(sources), floor, query[:80])
+ logger.info("RAG auto-inject: %d passage(s) for %r", len(sources), query[:80])
return {"events": events, "messages": messages}
diff --git a/studio/backend/core/rag/captioner.py b/studio/backend/core/rag/captioner.py
index be8e341064..e29c9c9a7a 100644
--- a/studio/backend/core/rag/captioner.py
+++ b/studio/backend/core/rag/captioner.py
@@ -1,9 +1,12 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
-"""Caption figures with the loaded vision model and splice the text into the page
-so images are searchable via the normal FTS5 + dense path. No-op (never raises)
-without a vision model or on failure; gated by ``config.CAPTION_IMAGES``."""
+"""Vision-model helpers for ingestion: figure captioning and scanned-page OCR.
+
+Both turn pixels into indexable text and are a no-op (never raise) without a loaded
+vision model. They reuse the chat model's vision endpoint, so it must be served with
+``--ubatch-size`` >= one image's tokens (some encoders, e.g. Gemma, attend
+non-causally and abort otherwise); Studio's vision chat already requires this."""
from __future__ import annotations
@@ -15,11 +18,54 @@ from . import config
logger = logging.getLogger(__name__)
_CAPTION_PROMPT = (
- "Describe this figure or image from a document in one or two concise "
- "sentences, for search indexing. State what it depicts (e.g. a diagram, "
- "chart, table or photo) and its key content. Do not add commentary."
+ "Read this figure or image from a document for search indexing.\n"
+ "First, on a line 'TEXT:', transcribe every piece of visible text exactly as "
+ "written, in reading order: the title, axis labels and units, legend and series "
+ "names, EVERY box / node / arrow label, table headers and cells, equations, and "
+ "footnotes. List each distinct label even if it is small.\n"
+ "Then, on a line 'SUMMARY:', add one or two sentences on what it shows (chart "
+ "type and trend, diagram subject, table topic, or photo content).\n"
+ "Report only what is visible. Transcribe exactly; do not invent or guess any "
+ "text, label, or number."
)
+_OCR_PROMPT = (
+ "Transcribe all text on this document page exactly as it appears, in reading "
+ "order, including any text inside figures, diagrams, charts, and tables (keep "
+ "table rows readable). Output only the transcribed text, with no commentary or "
+ "code fences. Preserve headings, lists, and line breaks. If the page has no "
+ "readable text, output nothing."
+)
+
+
+def _collapse_runaway(
+ text: str,
+ max_repeat: int = 3,
+ max_total: int = 8,
+) -> str:
+ """Cap runaway repetition: vision models sometimes loop a line many times. Keep
+ each distinct line to ``max_repeat`` in a row and ``max_total`` total, and collapse
+ blank-line floods, so a degenerate page cannot flood the index."""
+ out: list[str] = []
+ seen: dict[str, int] = {}
+ prev: str | None = None
+ run = 0
+ for line in text.splitlines():
+ key = line.strip()
+ if not key:
+ if prev == "": # collapse runs of blank lines to a single separator
+ continue
+ prev = ""
+ out.append("")
+ continue
+ run = run + 1 if key == prev else 1
+ prev = key
+ seen[key] = seen.get(key, 0) + 1
+ if run > max_repeat or seen[key] > max_total:
+ continue
+ out.append(line)
+ return "\n".join(out)
+
def vision_endpoint() -> tuple[str, str] | None:
"""``(base_url, model)`` for a loaded vision GGUF model, else None."""
@@ -33,7 +79,28 @@ def vision_endpoint() -> tuple[str, str] | None:
return None
-def _caption_one(base_url: str, model: str, image_bytes: bytes, timeout: float) -> str | None:
+def _vision_auth_headers() -> dict | None:
+ """Bearer header for the backend's API, or None. Vision calls share the chat
+ endpoint, so they need the same key under direct-stream (``--api-key``) mode."""
+ try:
+ from routes.inference import get_llama_cpp_backend
+ return get_llama_cpp_backend()._auth_headers or None
+ except Exception: # noqa: BLE001 - auth discovery must never break ingestion
+ return None
+
+
+def _vision_complete(
+ base_url: str,
+ model: str,
+ image_bytes: bytes,
+ *,
+ prompt: str,
+ timeout: float,
+ max_tokens: int,
+ temperature: float = 0.0,
+) -> str | None:
+ """One image-in / text-out call to the loaded vision model's OpenAI-compatible
+ endpoint. Returns the stripped text or ``None`` on empty/failure (non-fatal)."""
import httpx
data_url = "data:image/png;base64," + base64.b64encode(image_bytes).decode("ascii")
@@ -43,33 +110,62 @@ def _caption_one(base_url: str, model: str, image_bytes: bytes, timeout: float)
{
"role": "user",
"content": [
- {"type": "text", "text": _CAPTION_PROMPT},
+ {"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": data_url}},
],
}
],
- "max_tokens": 200,
- "temperature": 0.2,
+ "max_tokens": max_tokens,
+ # Deterministic by default: transcription must not randomly drop labels.
+ "temperature": temperature,
"stream": False,
# Off: thinking models would spend the budget reasoning, returning "".
"chat_template_kwargs": {"enable_thinking": False},
}
try:
- r = httpx.post(f"{base_url}/v1/chat/completions", json = payload, timeout = timeout)
+ r = httpx.post(
+ f"{base_url}/v1/chat/completions",
+ json = payload,
+ timeout = timeout,
+ headers = _vision_auth_headers(),
+ )
r.raise_for_status()
text = r.json()["choices"][0]["message"]["content"]
return text.strip() or None
- except Exception: # noqa: BLE001 - a failed caption is non-fatal
- logger.debug("caption request failed", exc_info = True)
+ except Exception: # noqa: BLE001 - a failed vision call is non-fatal
+ logger.debug("vision request failed", exc_info = True)
return None
+def _caption_one(base_url: str, model: str, image_bytes: bytes, timeout: float) -> str | None:
+ return _vision_complete(
+ base_url,
+ model,
+ image_bytes,
+ prompt = _CAPTION_PROMPT,
+ timeout = timeout,
+ max_tokens = config.CAPTION_MAX_TOKENS,
+ )
+
+
+def _ocr_one(base_url: str, model: str, image_bytes: bytes, timeout: float) -> str | None:
+ return _vision_complete(
+ base_url,
+ model,
+ image_bytes,
+ prompt = _OCR_PROMPT,
+ timeout = timeout,
+ max_tokens = config.OCR_MAX_TOKENS,
+ )
+
+
def caption_images(
images: list, *, endpoint: tuple[str, str] | None = None
) -> dict[int, list[str]]:
- """Caption ``ParsedImage`` objects, keyed by 1-based page number; ``{}`` when
- disabled, no vision model, or no images. Bounded by ``CAPTION_MAX_IMAGES``."""
- if not config.CAPTION_IMAGES or not images:
+ """Caption ``ParsedImage`` objects, keyed by 1-based page number; ``{}`` when there
+ are no images or no vision model. The caller (`ingestion._run`) owns the on/off
+ policy. Bounded by ``CAPTION_MAX_IMAGES``; each caption passes ``_collapse_runaway``."""
+ if not images:
return {}
ep = endpoint or vision_endpoint()
if ep is None:
@@ -84,7 +180,50 @@ def caption_images(
caption = _caption_one(base_url, model, image_bytes, config.CAPTION_TIMEOUT_S)
if caption:
page = getattr(img, "page_number", None) or 0
- out.setdefault(int(page), []).append(caption)
+ out.setdefault(int(page), []).append(_collapse_runaway(caption))
+ return out
+
+
+def ocr_pages(
+ page_pngs: dict[int, bytes], *, endpoint: tuple[str, str] | None = None
+) -> dict[int, str]:
+ """OCR rendered page PNGs (keyed by 1-based page number) to text; ``{}`` when there
+ is no vision model or no pages. The caller (`ingestion._ocr_scanned_pages`) owns the
+ on/off policy. Bounded by ``OCR_MAX_PAGES``."""
+ if not page_pngs:
+ return {}
+ ep = endpoint or vision_endpoint()
+ if ep is None:
+ return {}
+ base_url, model = ep
+
+ out: dict[int, str] = {}
+ for page_num in sorted(page_pngs)[: config.OCR_MAX_PAGES]:
+ text = _ocr_one(base_url, model, page_pngs[page_num], config.OCR_TIMEOUT_S)
+ if text:
+ out[int(page_num)] = _collapse_runaway(text)
+ return out
+
+
+def merge_page_captions(captions: dict[int, list[str]]) -> dict[int, list[str]]:
+ """Merge a page's per-tile captions into one deduped block: drop lines repeated
+ across overlapping tiles (first kept, order preserved), then ``_collapse_runaway``,
+ so ``splice_captions`` adds a single figure block per page."""
+ out: dict[int, list[str]] = {}
+ for page, caps in captions.items():
+ seen: set[str] = set()
+ lines: list[str] = []
+ for cap in caps:
+ for line in (cap or "").splitlines():
+ stripped = line.strip()
+ key = stripped.lower()
+ if not stripped or key in seen:
+ continue
+ seen.add(key)
+ lines.append(stripped)
+ merged = _collapse_runaway("\n".join(lines))
+ if merged.strip():
+ out[page] = [merged]
return out
diff --git a/studio/backend/core/rag/config.py b/studio/backend/core/rag/config.py
index 993423683c..54a224d081 100644
--- a/studio/backend/core/rag/config.py
+++ b/studio/backend/core/rag/config.py
@@ -17,13 +17,50 @@ TOP_K_DENSE = int(os.environ.get("RAG_TOP_K_DENSE", "30"))
TOP_K_HYBRID = int(os.environ.get("RAG_TOP_K_HYBRID", "10"))
RRF_K = int(os.environ.get("RAG_RRF_K", "60"))
-UPLOAD_EXTS = {".pdf", ".txt", ".md", ".markdown", ".docx", ".html", ".htm"}
+# Whole-document context: a thread-attached file under the token budget is injected
+# in full (every chunk, in order) instead of top-K retrieval; above it, use retrieval.
+THREAD_WHOLE_DOC = os.environ.get("RAG_THREAD_WHOLE_DOC", "1") == "1"
+WHOLE_DOC_MAX_TOKENS = int(os.environ.get("RAG_WHOLE_DOC_MAX_TOKENS", "6000"))
-# Figure captioning via the loaded vision model; off by default since each caption
-# is a model call. MAX_IMAGES bounds per-doc cost.
-CAPTION_IMAGES = os.environ.get("RAG_CAPTION_IMAGES", "0") == "1"
-CAPTION_MAX_IMAGES = int(os.environ.get("RAG_CAPTION_MAX_IMAGES", "8"))
-CAPTION_TIMEOUT_S = float(os.environ.get("RAG_CAPTION_TIMEOUT_S", "30"))
+UPLOAD_EXTS = {".pdf", ".txt", ".md", ".markdown", ".docx", ".html", ".htm"}
+# Reject uploads larger than this, so one pathological file can't drive unbounded parse
+# + vision work at ingest. 0 disables the cap. Default 200 MB.
+MAX_UPLOAD_BYTES = int(os.environ.get("RAG_MAX_UPLOAD_BYTES", str(200 * 1024 * 1024)))
+
+# Extract PDF text as layout-aware Markdown (pymupdf4llm) instead of flat text, so
+# tables, headings and lists survive into chunks and retrieval. Falls back to plain
+# PyMuPDF text when off, when pymupdf4llm is missing, or when extraction fails.
+PDF_MARKDOWN = os.environ.get("RAG_PDF_MARKDOWN", "1") == "1"
+
+# Figure captioning via the loaded vision model: detected figures are transcribed +
+# described so they become searchable. On by default, a no-op without a vision model;
+# the chat's "Describe figures & charts" toggle overrides it per upload.
+CAPTION_IMAGES = os.environ.get("RAG_CAPTION_IMAGES", "1") == "1"
+# Total per-document tile budget (figure-bearing pages are tiled, see below).
+CAPTION_MAX_IMAGES = int(os.environ.get("RAG_CAPTION_MAX_IMAGES", "24"))
+CAPTION_TIMEOUT_S = float(os.environ.get("RAG_CAPTION_TIMEOUT_S", "60"))
+# Larger than a one-line caption since captions transcribe every label. FIGURE_DPI is
+# high enough to keep small box/axis labels legible when tiles are rendered.
+CAPTION_MAX_TOKENS = int(os.environ.get("RAG_CAPTION_MAX_TOKENS", "768"))
+FIGURE_DPI = int(os.environ.get("RAG_FIGURE_DPI", "200"))
+# Figure pages are tiled into an overlapping ROWS x COLS grid of high-DPI tiles (plus
+# an optional full page), so small labels and every sub-figure are covered without
+# exact region detection. MAX_PAGES bounds figure pages; MAX_IMAGES bounds total tiles.
+FIGURE_TILE_ROWS = int(os.environ.get("RAG_FIGURE_TILE_ROWS", "2"))
+FIGURE_TILE_COLS = int(os.environ.get("RAG_FIGURE_TILE_COLS", "2"))
+FIGURE_TILE_OVERLAP = float(os.environ.get("RAG_FIGURE_TILE_OVERLAP", "0.12"))
+FIGURE_FULLPAGE = os.environ.get("RAG_FIGURE_FULLPAGE", "1") == "1"
+CAPTION_MAX_PAGES = int(os.environ.get("RAG_CAPTION_MAX_PAGES", "4"))
+
+# Scanned-PDF OCR: a page with little extractable text is rendered and transcribed by
+# the vision model so it becomes searchable. Needs a vision model, else skipped (page
+# stays empty). MIN_CHARS is the text length below which a page is treated as scanned.
+OCR_SCANNED = os.environ.get("RAG_OCR_SCANNED", "1") == "1"
+OCR_MIN_CHARS = int(os.environ.get("RAG_OCR_MIN_CHARS", "16"))
+OCR_MAX_PAGES = int(os.environ.get("RAG_OCR_MAX_PAGES", "20"))
+OCR_DPI = int(os.environ.get("RAG_OCR_DPI", "150"))
+OCR_TIMEOUT_S = float(os.environ.get("RAG_OCR_TIMEOUT_S", "60"))
+OCR_MAX_TOKENS = int(os.environ.get("RAG_OCR_MAX_TOKENS", "2048"))
# Embedder backend. "auto": sentence-transformers on a CUDA/ROCm GPU (torch fp16
# wins bulk indexing), else torch-free GGUF llama-server. Switching backends changes
diff --git a/studio/backend/core/rag/ingestion.py b/studio/backend/core/rag/ingestion.py
index 77bbdc92f5..04365ab76b 100644
--- a/studio/backend/core/rag/ingestion.py
+++ b/studio/backend/core/rag/ingestion.py
@@ -99,25 +99,108 @@ def _embed_all(texts: list[str], model_name: str | None):
return vectors
+def _ocr_scanned_pages(
+ pages: list,
+ stored_path: str,
+ conn,
+ job_id: str,
+ ocr: bool | None = None,
+) -> tuple[list, set[int]]:
+ """Replace text on near-empty (scanned/image-only) PDF pages with vision-model OCR
+ so image PDFs become searchable. ``ocr`` overrides ``config.OCR_SCANNED`` per upload
+ (``None`` = config default); no-op without scanned pages or a vision model. OCR'd
+ pages have no text layer, so no preview highlight regions, but stay searchable.
+ Returns ``(pages, ocred)``: new ``Page`` objects for OCR'd pages (originals
+ otherwise) and the set of page numbers actually transcribed."""
+ if not (config.OCR_SCANNED if ocr is None else ocr):
+ return pages, set()
+ scanned = [
+ p.page_number
+ for p in pages
+ if p.page_number is not None and len((p.text or "").strip()) < config.OCR_MIN_CHARS
+ ]
+ if not scanned or captioner.vision_endpoint() is None:
+ return pages, set()
+ if len(scanned) > config.OCR_MAX_PAGES:
+ logger.warning(
+ "OCR: %d scanned pages exceed OCR_MAX_PAGES=%d; pages past the cap stay "
+ "untranscribed (raise RAG_OCR_MAX_PAGES to cover them)",
+ len(scanned),
+ config.OCR_MAX_PAGES,
+ )
+ scanned = scanned[: config.OCR_MAX_PAGES]
+ _progress(conn, job_id, "ocr", 0.25)
+ page_pngs = parsers.render_pdf_pages(stored_path, scanned, dpi = config.OCR_DPI)
+ texts = captioner.ocr_pages(page_pngs)
+ if not texts:
+ return pages, set()
+
+ from .parsers import Page
+
+ out: list = []
+ ocred: set[int] = set()
+ for page in pages:
+ text = texts.get(page.page_number)
+ if text:
+ original = (page.text or "").strip()
+ merged = text if not original or original in text else f"{original}\n\n{text}"
+ out.append(Page(text = merged, page_number = page.page_number, char_count = len(merged)))
+ ocred.add(page.page_number)
+ else:
+ out.append(page)
+ return out, ocred
+
+
def _run(
- job_id: str, document_id: str, scope: str, stored_path: str, model_name: str | None
+ job_id: str,
+ document_id: str,
+ scope: str,
+ stored_path: str,
+ model_name: str | None,
+ ocr: bool | None = None,
+ caption: bool | None = None,
) -> None:
conn = rag_db.get_connection()
try:
_progress(conn, job_id, "parsing", 0.1)
pages = parsers.parse(stored_path)
- if config.CAPTION_IMAGES and stored_path.lower().endswith(".pdf"):
- # Caption figures, splice into page text (no-op without a vision model).
+ is_pdf = stored_path.lower().endswith(".pdf")
+ ocred: set[int] = set()
+ if is_pdf:
+ pages, ocred = _ocr_scanned_pages(pages, stored_path, conn, job_id, ocr = ocr)
+ caption_on = config.CAPTION_IMAGES if caption is None else caption
+ # Skip all figure work (PDF rasterization included) without a vision model.
+ if caption_on and is_pdf and captioner.vision_endpoint() is not None:
+ # Tile figure pages, transcribe+describe each tile, then merge/dedup/splice
+ # into the page text so small labels and every sub-figure are captured.
try:
- figures = parsers.render_pdf_figures(
- stored_path, max_figures = config.CAPTION_MAX_IMAGES
+ fig_pages = parsers.pages_with_figures(
+ stored_path,
+ max_pages = config.CAPTION_MAX_PAGES,
+ # Skip only pages OCR actually transcribed (it covers them whole); a
+ # scanned figure page past the OCR cap or with empty OCR still tiles.
+ exclude_pages = ocred,
+ )
+ tiles = (
+ parsers.render_pdf_figure_tiles(
+ stored_path,
+ fig_pages,
+ dpi = config.FIGURE_DPI,
+ rows = config.FIGURE_TILE_ROWS,
+ cols = config.FIGURE_TILE_COLS,
+ overlap = config.FIGURE_TILE_OVERLAP,
+ fullpage = config.FIGURE_FULLPAGE,
+ max_tiles = config.CAPTION_MAX_IMAGES,
+ )
+ if fig_pages
+ else []
)
except Exception:
- logger.warning("figure rendering failed for job %s", job_id, exc_info = True)
- figures = []
- if figures:
- _progress(conn, job_id, "captioning", 0.2)
- captions = captioner.caption_images(figures)
+ logger.warning("figure tiling failed for job %s", job_id, exc_info = True)
+ tiles = []
+ if tiles:
+ _progress(conn, job_id, "captioning", 0.28)
+ captions = captioner.merge_page_captions(captioner.caption_images(tiles))
pages = captioner.splice_captions(pages, captions)
_progress(conn, job_id, "chunking", 0.3)
@@ -175,6 +258,8 @@ def start_ingestion(
*,
project_id: str | None = None,
model_name: str | None = None,
+ ocr: bool | None = None,
+ caption: bool | None = None,
) -> tuple[str, str]:
"""Create the document + job rows and spawn the worker, returning
``(document_id, job_id)``. A duplicate content hash in this scope returns the
@@ -191,13 +276,26 @@ def start_ingestion(
try:
existing = store.document_by_hash(conn, scope, sha)
if existing is not None:
- job_id = _new_job(conn, existing, scope, status = "completed", progress = 1.0)
- _remove_upload(stored_path)
- with _jobs_lock:
- _jobs[job_id] = queue.Queue()
- _emit(job_id, {"type": "complete", "num_chunks": 0, "deduped": True})
- _emit(job_id, None)
- return existing, job_id
+ doc = store.get_document(conn, existing)
+ empty_completed = (
+ doc is not None and doc.get("status") == "completed" and not doc.get("num_chunks")
+ )
+ if empty_completed:
+ # A prior ingest of identical bytes yielded zero chunks (e.g. a scanned
+ # PDF uploaded before a vision model loaded). Re-ingest, don't dedupe.
+ store.delete_document(conn, existing)
+ _remove_upload(doc.get("stored_path"), keep_path = stored_path)
+ else:
+ job_id = _new_job(conn, existing, scope, status = "completed", progress = 1.0)
+ _remove_upload(stored_path)
+ with _jobs_lock:
+ _jobs[job_id] = queue.Queue()
+ _emit(
+ job_id,
+ {"type": "complete", "num_chunks": doc.get("num_chunks") or 0, "deduped": True},
+ )
+ _emit(job_id, None)
+ return existing, job_id
for failed in store.failed_documents_by_hash(conn, scope, sha):
store.delete_document(conn, failed["id"])
_remove_upload(failed.get("stored_path"), keep_path = stored_path)
@@ -221,7 +319,7 @@ def start_ingestion(
_jobs[job_id] = queue.Queue()
threading.Thread(
target = _run,
- args = (job_id, document_id, scope, stored_path, model_name),
+ args = (job_id, document_id, scope, stored_path, model_name, ocr, caption),
daemon = True,
).start()
return document_id, job_id
@@ -339,10 +437,16 @@ def job_events(job_id: str):
def get_job_status(job_id: str) -> dict | None:
- """Read the persisted ingestion job row (status / stage / progress / error)."""
+ """Read the persisted ingestion job row (status / stage / progress / error), plus
+ the document's ``num_chunks`` so a client polling to completion learns the chunk
+ count (the SSE ``complete`` frame carries it, but the poll/reconcile path does not)."""
conn = rag_db.get_connection()
try:
- row = conn.execute("SELECT * FROM ingestion_jobs WHERE id=?", (job_id,)).fetchone()
+ row = conn.execute(
+ "SELECT j.*, d.num_chunks AS num_chunks FROM ingestion_jobs j "
+ "LEFT JOIN documents d ON d.id = j.document_id WHERE j.id=?",
+ (job_id,),
+ ).fetchone()
return dict(row) if row else None
finally:
conn.close()
diff --git a/studio/backend/core/rag/locators.py b/studio/backend/core/rag/locators.py
index 57c0487486..9331bb15ac 100644
--- a/studio/backend/core/rag/locators.py
+++ b/studio/backend/core/rag/locators.py
@@ -39,9 +39,11 @@ def _norm_token(token: str) -> str:
def _anchor_tokens(page_text: str, match: LocatorMatch) -> list[str]:
"""Normalized anchor tokens from the chunk's leading span. Drops first and last
- token (boundaries often slice mid-word) when long enough."""
+ token (boundaries often slice mid-word) when long enough. Pipes are split out so
+ Markdown table cells (``|Q1|$1.2M|``) become individual words that match the PDF
+ word stream."""
segment = page_text[match.start : match.end]
- raw = segment.split()
+ raw = segment.replace("|", " ").split()
if len(raw) >= MIN_ANCHOR_WORDS + 2:
raw = raw[1:-1]
tokens = [t for t in (_norm_token(w) for w in raw) if t]
diff --git a/studio/backend/core/rag/parsers.py b/studio/backend/core/rag/parsers.py
index 84da941762..ba248cf9a6 100644
--- a/studio/backend/core/rag/parsers.py
+++ b/studio/backend/core/rag/parsers.py
@@ -15,6 +15,8 @@ import os
from dataclasses import dataclass
from html.parser import HTMLParser
+from . import config
+
logger = logging.getLogger(__name__)
@@ -67,6 +69,28 @@ def _html(raw: str) -> list[Page]:
return [_page("\n".join(parser.out), 1)]
+def _pdf_markdown(doc) -> list[str] | None:
+ """Per-page layout-aware Markdown (tables, headings, lists) via pymupdf4llm; index
+ i maps to page i+1. Returns None when the lib is missing, extraction fails, or the
+ page count does not line up, so the caller falls back to plain PyMuPDF text."""
+ try:
+ import pymupdf4llm
+ except Exception:
+ return None
+ try:
+ chunks = pymupdf4llm.to_markdown(
+ doc,
+ page_chunks = True,
+ show_progress = False,
+ )
+ except Exception: # noqa: BLE001 - never let Markdown extraction break ingestion
+ logger.warning("pymupdf4llm extraction failed; using plain text", exc_info = True)
+ return None
+ if not isinstance(chunks, list) or len(chunks) != doc.page_count:
+ return None
+ return [str(c.get("text") or "") for c in chunks]
+
+
def _pdf(path: str, want_images: bool) -> tuple[list[Page], list[ParsedImage]]:
import fitz # PyMuPDF
@@ -74,8 +98,11 @@ def _pdf(path: str, want_images: bool) -> tuple[list[Page], list[ParsedImage]]:
images: list[ParsedImage] = []
doc = fitz.open(path)
try:
+ md = _pdf_markdown(doc) if config.PDF_MARKDOWN else None
for i, page in enumerate(doc):
- text = page.get_text("text") or ""
+ # Prefer layout-aware Markdown (keeps tables/headings legible for retrieval);
+ # fall back to plain text when Markdown is off, unavailable, or empty here.
+ text = (md[i] if md else "") or page.get_text("text") or ""
pages.append(_page(text, i + 1))
if want_images:
for img in page.get_images(full = True):
@@ -118,63 +145,164 @@ def _merge_rects(boxes: list) -> list:
return merged
-def render_pdf_figures(
- path: str,
+def _figure_boxes(
+ page,
*,
- dpi: int = 130,
min_area_frac: float = 0.04,
min_side: float = 40.0,
- max_figures: int = 8,
-) -> list[ParsedImage]:
- """Detect figure regions and render each to a PNG for captioning.
+) -> list:
+ """Qualifying figure-region rectangles on a page: cluster vector drawings + raster
+ placements, merge overlaps, keep the page-spanning ones (area/side filtered)."""
+ boxes: list = []
+ try:
+ boxes.extend(info["bbox"] for info in page.get_image_info())
+ except Exception:
+ pass
+ try:
+ boxes.extend(page.cluster_drawings())
+ except Exception:
+ pass
+ if not boxes:
+ return []
+ page_area = page.rect.width * page.rect.height
+ keep: list = []
+ for box in _merge_rects(boxes):
+ if (
+ box.get_area() >= min_area_frac * page_area
+ and box.width >= min_side
+ and box.height >= min_side
+ ):
+ keep.append(box)
+ return keep
- Academic figures are vector, so raster extraction yields fragments; instead
- cluster vector drawings + raster placements into boxes, keep the page-spanning
- ones, and render them. Any failure yields [], never an exception.
- """
+
+def pages_with_figures(
+ path: str,
+ *,
+ max_pages: int = 4,
+ min_area_frac: float = 0.04,
+ min_side: float = 40.0,
+ exclude_pages: set[int] | None = None,
+) -> list[int]:
+ """1-based page numbers with a qualifying figure region, capped at ``max_pages``;
+ drives figure tiling. ``exclude_pages`` (1-based) are skipped: those are the pages
+ OCR already transcribed whole, so tiling them would duplicate the vision work. Any
+ failure yields []."""
+ exclude = exclude_pages or set()
try:
import pymupdf
except Exception:
return []
-
- out: list[ParsedImage] = []
try:
doc = pymupdf.open(path)
except Exception:
return []
+ pages: list[int] = []
try:
for i, page in enumerate(doc):
- boxes: list = []
- try:
- boxes.extend(info["bbox"] for info in page.get_image_info())
- except Exception:
- pass
- try:
- boxes.extend(page.cluster_drawings())
- except Exception:
- pass
- if not boxes:
+ if (i + 1) in exclude:
continue
- page_area = page.rect.width * page.rect.height
- for box in _merge_rects(boxes):
- if (
- box.get_area() >= min_area_frac * page_area
- and box.width >= min_side
- and box.height >= min_side
- ):
- try:
- pix = page.get_pixmap(dpi = dpi, clip = box)
- out.append(
- ParsedImage(
- image_bytes = pix.tobytes("png"),
- page_number = i + 1,
- xref = 0,
- )
+ if _figure_boxes(page, min_area_frac = min_area_frac, min_side = min_side):
+ pages.append(i + 1)
+ if len(pages) >= max_pages:
+ break
+ return pages
+ finally:
+ doc.close()
+
+
+def render_pdf_figure_tiles(
+ path: str,
+ page_numbers,
+ *,
+ dpi: int = 200,
+ rows: int = 2,
+ cols: int = 2,
+ overlap: float = 0.12,
+ fullpage: bool = True,
+ max_tiles: int = 24,
+) -> list[ParsedImage]:
+ """Render figure-bearing pages as overlapping high-DPI tiles (plus an optional full
+ page), each a ``ParsedImage`` keyed by page number. Tiling keeps small labels legible
+ and covers every sub-figure without exact region detection. Any failure yields []."""
+ wanted = [int(n) for n in page_numbers]
+ if not wanted:
+ return []
+ rows, cols = max(1, int(rows)), max(1, int(cols)) # never divide by zero
+ try:
+ import pymupdf
+ except Exception:
+ return []
+ try:
+ doc = pymupdf.open(path)
+ except Exception:
+ return []
+ out: list[ParsedImage] = []
+ try:
+ for num in wanted:
+ if num < 1 or num > doc.page_count:
+ continue
+ page = doc[num - 1]
+ rect = page.rect
+ clips: list = [rect] if fullpage else []
+ cw, ch = rect.width / cols, rect.height / rows
+ ox, oy = cw * overlap, ch * overlap
+ for r in range(rows):
+ for c in range(cols):
+ clips.append(
+ pymupdf.Rect(
+ rect.x0 + c * cw - ox,
+ rect.y0 + r * ch - oy,
+ rect.x0 + (c + 1) * cw + ox,
+ rect.y0 + (r + 1) * ch + oy,
)
- except Exception:
- continue
- if len(out) >= max_figures:
- return out
+ & rect
+ )
+ for clip in clips:
+ try:
+ pix = page.get_pixmap(dpi = dpi, clip = clip)
+ out.append(ParsedImage(image_bytes = pix.tobytes("png"), page_number = num, xref = 0))
+ except Exception:
+ continue
+ if len(out) >= max_tiles:
+ return out
+ return out
+ finally:
+ doc.close()
+
+
+def render_pdf_pages(
+ path: str,
+ page_numbers,
+ *,
+ dpi: int = 150,
+) -> dict[int, bytes]:
+ """Render whole PDF pages (given as 1-based numbers) to PNG bytes, keyed by
+ page number. Backs scanned-page OCR. Any failure yields ``{}`` (or skips that
+ page), never an exception.
+ """
+ wanted = {int(n) for n in page_numbers}
+ if not wanted:
+ return {}
+ try:
+ import pymupdf
+ except Exception:
+ return {}
+ try:
+ doc = pymupdf.open(path)
+ except Exception:
+ return {}
+ out: dict[int, bytes] = {}
+ try:
+ for i, page in enumerate(doc):
+ num = i + 1
+ if num not in wanted:
+ continue
+ try:
+ pix = page.get_pixmap(dpi = dpi)
+ out[num] = pix.tobytes("png")
+ except Exception:
+ continue
return out
finally:
doc.close()
diff --git a/studio/backend/core/rag/store.py b/studio/backend/core/rag/store.py
index 7d58931e53..8e59c5fbf6 100644
--- a/studio/backend/core/rag/store.py
+++ b/studio/backend/core/rag/store.py
@@ -292,3 +292,40 @@ def chunks_by_id(conn: sqlite3.Connection, ids) -> dict:
list(ids),
).fetchall()
return {r["id"]: r for r in rows}
+
+
+def all_chunks_for_scope(conn: sqlite3.Connection, scope) -> list[dict]:
+ """Every completed-document chunk for a scope, ordered document-then-index and
+ joined with the document filename. Backs whole-document context injection, so
+ it does no retrieval or embedding."""
+ scopes = _scopes(scope)
+ if not scopes:
+ return []
+ placeholders = ",".join("?" * len(scopes))
+ rows = conn.execute(
+ f"SELECT c.id, c.text, c.document_id, c.chunk_index, c.page_number, "
+ f"c.token_count, d.filename, d.created_at "
+ f"FROM chunks c JOIN documents d ON d.id=c.document_id "
+ f"WHERE c.scope IN ({placeholders}) AND d.status='completed' "
+ f"ORDER BY d.created_at, c.document_id, c.chunk_index",
+ list(scopes),
+ ).fetchall()
+ return [dict(r) for r in rows]
+
+
+def scope_token_estimate(conn: sqlite3.Connection, scope) -> int:
+ """Upper-bound token total for a scope's completed chunks without hydrating text.
+ Mirrors ``all_chunks_for_scope`` + the ``tool._row_token_count`` fallback (stored
+ count, else length/4), so the whole-doc budget can be checked before loading text."""
+ scopes = _scopes(scope)
+ if not scopes:
+ return 0
+ placeholders = ",".join("?" * len(scopes))
+ row = conn.execute(
+ f"SELECT COALESCE(SUM(CASE WHEN c.token_count > 0 THEN c.token_count "
+ f"ELSE MAX(1, length(COALESCE(c.text, '')) / 4) END), 0) AS total "
+ f"FROM chunks c JOIN documents d ON d.id=c.document_id "
+ f"WHERE c.scope IN ({placeholders}) AND d.status='completed'",
+ list(scopes),
+ ).fetchone()
+ return int(row["total"] or 0)
diff --git a/studio/backend/core/rag/tool.py b/studio/backend/core/rag/tool.py
index ccb1b47e63..b05f8dd3a3 100644
--- a/studio/backend/core/rag/tool.py
+++ b/studio/backend/core/rag/tool.py
@@ -16,7 +16,13 @@ from xml.sax.saxutils import quoteattr
from storage import rag_db
from . import config, retrieval
-from .store import kb_scope, project_scope, thread_scope
+from .store import (
+ all_chunks_for_scope,
+ kb_scope,
+ project_scope,
+ scope_token_estimate,
+ thread_scope,
+)
SEARCH_KNOWLEDGE_BASE_TOOL = {
"type": "function",
@@ -90,6 +96,30 @@ def _format(rows, hits) -> tuple[str, list[dict]]:
return "\n\n".join(blocks), sources
+def render_sources(sources: list[dict]) -> str:
+ """Render a citation-source list to sequentially-numbered ```` blocks,
+ rewriting each source's ``citationId`` to match its 1-based position. Lets
+ independently-built source lists (a whole-document thread attachment plus
+ retrieved project passages) be merged under one citation numbering."""
+ blocks: list[str] = []
+ for i, s in enumerate(sources, 1):
+ s["citationId"] = i
+ src = quoteattr(s.get("filename") or "unknown")
+ page = s.get("page")
+ page_attr = f" page={quoteattr(str(page))}" if page else ""
+ blocks.append(f'\n{s.get("text") or ""}\n')
+ return "\n\n".join(blocks)
+
+
+def _row_token_count(row) -> int:
+ """Chunk token count for budgeting, falling back to a length estimate when the
+ stored count is missing or zero, so a malformed chunk cannot bypass the budget."""
+ tc = row["token_count"]
+ if tc:
+ return int(tc)
+ return max(1, len(row["text"] or "") // 4)
+
+
def search_knowledge_base_with_sources(
*,
query: str,
@@ -186,6 +216,55 @@ def search_for_autoinject(
return (text, sources) if sources else None
+def whole_document_context(
+ *, scope_thread_id: str | None = None, max_tokens: int
+) -> tuple[str, list[dict]] | None:
+ """Render EVERY chunk of the THREAD's attached documents (in order) as the same
+ ```` blocks + citation source-map as retrieval, so the model reads the whole
+ file rather than top-K passages. Thread-attached files only: KB and project corpora
+ are search corpora, never whole-document, so this resolves the thread scope alone.
+ ``None`` (caller falls back to retrieval) when there is no thread scope, no completed
+ chunks, or the total exceeds ``max_tokens``."""
+ if not scope_thread_id:
+ return None
+ # A non-positive budget means "never inject" (disable whole-doc via
+ # RAG_THREAD_WHOLE_DOC=0), not "inject the whole corpus unbounded".
+ if max_tokens <= 0:
+ return None
+ scope = thread_scope(scope_thread_id)
+ conn = rag_db.get_connection()
+ try:
+ # Cheap budget pre-check (SUM, no text hydration): reject an oversized attachment
+ # before loading the whole corpus; all_chunks_for_scope runs only once it fits.
+ if scope_token_estimate(conn, scope) > max_tokens:
+ return None
+ rows = all_chunks_for_scope(conn, scope)
+ finally:
+ conn.close()
+ if not rows:
+ return None
+ total = sum(_row_token_count(r) for r in rows)
+ if total > max_tokens:
+ return None
+
+ sources: list[dict] = [
+ {
+ "citationId": i,
+ "chunkId": r["id"],
+ "documentId": r["document_id"],
+ "filename": r["filename"] or "unknown",
+ "page": r["page_number"],
+ "text": r["text"] or "",
+ "score": None,
+ }
+ for i, r in enumerate(rows, 1)
+ ]
+ rendered = render_sources(sources)
+ if max(1, len(rendered) // 4) > max_tokens:
+ return None
+ return rendered, sources
+
+
def search_knowledge_base(
*,
query: str,
diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py
index 9d991c6512..f4233fcf04 100644
--- a/studio/backend/core/training/training.py
+++ b/studio/backend/core/training/training.py
@@ -299,6 +299,7 @@ class TrainingBackend:
# Build config dict for the subprocess
config = {
"model_name": kwargs["model_name"],
+ "project_name": kwargs.get("project_name"),
"training_type": kwargs.get("training_type", "LoRA/QLoRA"),
"hf_token": kwargs.get("hf_token", ""),
"load_in_4bit": kwargs.get("load_in_4bit", True),
diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py
index 3f020c8abc..610af2472e 100644
--- a/studio/backend/core/training/worker.py
+++ b/studio/backend/core/training/worker.py
@@ -44,6 +44,7 @@ if sys.platform.startswith("linux") and "HSA_ENABLE_DXG_DETECTION" not in os.env
logger = get_logger(__name__)
from utils.hardware import apply_gpu_ids
+from utils.training_runs import build_default_output_dir_name
from utils.wheel_utils import (
direct_wheel_url,
flash_attn_wheel_url,
@@ -1787,11 +1788,14 @@ def _run_mlx_training(event_queue, stop_queue, config):
# ── 5. Build output dir ──
# Resolve to ~/.unsloth/studio/outputs/ so the export page finds it
- from utils.paths import resolve_output_dir, ensure_dir, default_run_dir_name
+ from utils.paths import resolve_output_dir, ensure_dir
output_dir = config.get("output_dir", "")
if not output_dir:
- output_dir = f"{default_run_dir_name(model_name)}_{int(time.time())}"
+ output_dir = build_default_output_dir_name(
+ model_name,
+ config.get("project_name"),
+ )
output_dir = str(resolve_output_dir(output_dir))
ensure_dir(Path(output_dir))
@@ -3019,7 +3023,10 @@ def run_training_process(*, event_queue: Any, stop_queue: Any, config: dict) ->
resume_from_checkpoint
)
if not output_dir:
- output_dir = f"{default_run_dir_name(model_name)}_{int(time.time())}"
+ output_dir = build_default_output_dir_name(
+ model_name,
+ config.get("project_name"),
+ )
output_dir = str(resolve_output_dir(output_dir))
ensure_dir(Path(output_dir))
@@ -3500,7 +3507,10 @@ def _run_embedding_training(event_queue: Any, stop_queue: Any, config: dict) ->
resume_from_checkpoint
)
if not output_dir:
- output_dir = f"{default_run_dir_name(model_name)}_{int(time.time())}"
+ output_dir = build_default_output_dir_name(
+ model_name,
+ config.get("project_name"),
+ )
output_dir = str(resolve_output_dir(output_dir))
num_epochs = config.get("num_epochs", 2)
diff --git a/studio/backend/hub/schemas/inventory.py b/studio/backend/hub/schemas/inventory.py
index 44ff545e76..ef95efe2f2 100644
--- a/studio/backend/hub/schemas/inventory.py
+++ b/studio/backend/hub/schemas/inventory.py
@@ -27,6 +27,9 @@ class GgufVariantDetail(BaseModel):
downloaded: bool = Field(
False, description = "Whether this variant is already in the local HF cache"
)
+ update_available: bool = Field(
+ False, description = "Whether a newer main GGUF blob is available on Hugging Face"
+ )
partial: bool = Field(
False,
description = "Whether this variant has an in-progress (.incomplete) blob in cache",
diff --git a/studio/backend/hub/services/models/cache_inventory.py b/studio/backend/hub/services/models/cache_inventory.py
index a961a6ae9d..a27e4860e6 100644
--- a/studio/backend/hub/services/models/cache_inventory.py
+++ b/studio/backend/hub/services/models/cache_inventory.py
@@ -39,8 +39,10 @@ from hub.services.models.common import (
logger = get_logger(__name__)
-_repo_size_cache: "OrderedDict[tuple[str, str], tuple[int, frozenset[str], float]]" = OrderedDict()
-_repo_size_neg_cache: "OrderedDict[tuple[str, str], float]" = OrderedDict()
+_repo_size_cache: "OrderedDict[tuple[str, str, str], tuple[int, frozenset[str], float]]" = (
+ OrderedDict()
+)
+_repo_size_neg_cache: "OrderedDict[tuple[str, str, str], float]" = OrderedDict()
_REPO_SIZE_CACHE_MAX = 256
_REPO_SIZE_POS_TTL = 60.0
_REPO_SIZE_NEG_TTL = 60.0
@@ -52,7 +54,7 @@ def get_repo_snapshot_metadata_cached(
repo_id: str, hf_token: Optional[str] = None
) -> tuple[int, frozenset[str]]:
token_fp = hf_cache_scan.token_fingerprint(hf_token)
- cache_key = (repo_id, token_fp)
+ cache_key = (repo_id, token_fp, "snapshot")
with _repo_size_cache_lock:
cached = _repo_size_cache.get(cache_key)
if cached is not None:
@@ -119,6 +121,52 @@ def _repo_has_gguf_files(repo_info) -> bool:
return _repo_gguf_size_bytes(repo_info) > 0
+def _cached_repo_file_name(file_obj) -> str:
+ file_path = getattr(file_obj, "file_path", None)
+ if file_path:
+ try:
+ path = Path(file_path)
+ parts = path.parts
+ snapshots_idx = max(i for i, part in enumerate(parts) if part == "snapshots")
+ if len(parts) > snapshots_idx + 2:
+ return Path(*parts[snapshots_idx + 2 :]).as_posix()
+ except Exception:
+ pass
+ return str(getattr(file_obj, "file_name", "")).replace("\\", "/")
+
+
+def _repo_gguf_blob_map(repo_info, *, include_companions: bool = False) -> dict[str, set[str]]:
+ """Map each cached GGUF file's repo-relative name to the SET of its local
+ blob hashes across all cached revisions.
+
+ HF names each local cache blob FILE by the file's etag (lfs.sha256 else
+ blob_id), so a local file's blob hash == ``Path(blob_path).name``. An updated
+ repo keeps BOTH the old and new revision snapshots until HF garbage-collects
+ them, so the same file resolves to several blobs; collecting them ALL (not
+ just the first one seen, since ``repo_info.revisions`` is a frozenset and
+ yields them in arbitrary order) lets the remote-vs-local diff treat the file
+ as current when the remote (``main``) blob is present in any cached revision.
+ Mirrors the ``cached_blob_ids`` membership test in routes/models.py.
+
+ By default this keeps the historical MAIN-GGUF-only behavior. GGUF update
+ checks opt into companions so a shared mmproj/MTP blob can be compared too.
+ """
+ blob_map: dict[str, set[str]] = {}
+ for revision in repo_info.revisions:
+ for f in revision.files:
+ if include_companions:
+ if not _is_gguf_filename(f.file_name):
+ continue
+ elif not _is_main_gguf_filename(f.file_name):
+ continue
+ blob_path = getattr(f, "blob_path", None)
+ if not blob_path:
+ continue
+ name = _cached_repo_file_name(f)
+ blob_map.setdefault(name, set()).add(Path(blob_path).name)
+ return blob_map
+
+
def _prefer_cache_row(candidate: dict, existing: Optional[dict]) -> bool:
if existing is None:
return True
diff --git a/studio/backend/hub/services/models/deletion.py b/studio/backend/hub/services/models/deletion.py
index 916bb9d4f8..e7c54fc75b 100644
--- a/studio/backend/hub/services/models/deletion.py
+++ b/studio/backend/hub/services/models/deletion.py
@@ -153,6 +153,30 @@ def _remove_empty_variant_dirs(target_repos: list, variant: str) -> tuple[int, l
return removed, failures
+def _remove_empty_snapshot_dirs(target_repos: list) -> tuple[int, list[str]]:
+ removed = 0
+ failures: list[str] = []
+ for target_repo in target_repos:
+ repo_path = getattr(target_repo, "repo_path", None)
+ if not repo_path:
+ continue
+ snapshots = Path(repo_path) / "snapshots"
+ if not snapshots.is_dir():
+ continue
+ try:
+ snap_dirs = [s for s in snapshots.iterdir() if s.is_dir() and not s.is_symlink()]
+ except OSError:
+ continue
+ for snap in snap_dirs:
+ try:
+ snap.rmdir()
+ removed += 1
+ except OSError as e:
+ if e.errno != errno.ENOTEMPTY:
+ failures.append(f"{snap.name}: {e}")
+ return removed, failures
+
+
def _delete_gguf_variant_from_repos(
repo_id: str,
variant: str,
@@ -255,6 +279,9 @@ def _delete_gguf_variant_from_repos(
state_purged = download_manifest.purge_state("model", repo_id, variant)
# Reclaim the empty quant folder so it stops 404ing on delete.
removed_dirs, dir_failures = _remove_empty_variant_dirs(target_repos, variant)
+ removed_snap_dirs, snap_dir_failures = _remove_empty_snapshot_dirs(target_repos)
+ removed_dirs += removed_snap_dirs
+ dir_failures.extend(snap_dir_failures)
if dir_failures:
raise HTTPException(
status_code = 409,
@@ -284,6 +311,181 @@ def _delete_gguf_variant_from_repos(
return {"status": "deleted", "repo_id": repo_id, "variant": variant}
+def reclaim_replaced_gguf_variant(
+ repo_id: str,
+ variant: str,
+ keep_main_hashes: frozenset[str],
+ hf_token: Optional[str] = None,
+) -> dict:
+ """Prune stale main-GGUF files for a variant after a replacement verified.
+
+ This is intentionally narrower than user-driven delete: it removes only
+ same-variant main files whose local blob hash is not in *keep_main_hashes*,
+ then unlinks their blobs only if no remaining snapshot references them.
+ Shared companions and sibling variants are left intact.
+ """
+ if not keep_main_hashes:
+ logger.info(
+ "Skipping stale GGUF reclaim for %s [%s]: current main hashes unresolved",
+ repo_id,
+ variant,
+ )
+ return {
+ "status": "skipped",
+ "repo_id": repo_id,
+ "variant": variant,
+ "reason": "unresolved_hashes",
+ }
+ if not _is_valid_repo_id(repo_id) or not _is_valid_gguf_variant(variant):
+ return {
+ "status": "skipped",
+ "repo_id": repo_id,
+ "variant": variant,
+ "reason": "invalid_target",
+ }
+
+ failures: list[str] = []
+ removed_snapshots = 0
+ deleted_blobs = 0
+ deleted_bytes = 0
+ variant_key = variant.lower()
+
+ try:
+ cache_scans = cache_inventory.all_hf_cache_scans()
+ except Exception as e:
+ logger.warning(
+ "Skipping stale GGUF reclaim for %s [%s]: cache scan failed: %s",
+ repo_id,
+ variant,
+ download_registry.scrub_secrets(str(e), hf_token = hf_token),
+ )
+ return {
+ "status": "skipped",
+ "repo_id": repo_id,
+ "variant": variant,
+ "reason": "scan_failed",
+ }
+
+ candidate_repos = [
+ repo_info
+ for hf_cache in cache_scans
+ for repo_info in hf_cache.repos
+ if str(getattr(repo_info, "repo_type", "")) == "model"
+ and str(getattr(repo_info, "repo_id", "")).lower() == repo_id.lower()
+ ]
+ try:
+ matched_repo_ids = resolve_destructive_repo_ids(
+ repo_id,
+ [str(getattr(repo_info, "repo_id", "")) for repo_info in candidate_repos],
+ noun = "models",
+ )
+ except HTTPException as e:
+ detail = getattr(e, "detail", str(e))
+ logger.warning(
+ "Skipping stale GGUF reclaim for %s [%s]: %s",
+ repo_id,
+ variant,
+ download_registry.scrub_secrets(str(detail), hf_token = hf_token),
+ )
+ return {
+ "status": "skipped",
+ "repo_id": repo_id,
+ "variant": variant,
+ "reason": "ambiguous_repo",
+ }
+ target_repos = [
+ repo_info
+ for repo_info in candidate_repos
+ if str(getattr(repo_info, "repo_id", "")) in matched_repo_ids
+ ]
+
+ for target_repo in target_repos:
+ repo_dir = Path(target_repo.repo_path) if getattr(target_repo, "repo_path", None) else None
+ stale_matches: list[tuple[Path, Optional[Path], str]] = []
+ matches = _repo_file_matches(
+ target_repo,
+ lambda name: _is_main_gguf_filename(name)
+ and extract_quant_label(name).lower() == variant_key,
+ )
+ for snap, blob, name in matches:
+ blob_hash = _blob_hash_from_path(blob) if blob is not None else None
+ if blob_hash is None or blob_hash in keep_main_hashes:
+ continue
+ stale_matches.append((snap, blob, name))
+
+ if not stale_matches:
+ continue
+
+ for snap, _blob, name in stale_matches:
+ try:
+ if _path_exists_or_symlink(snap):
+ snap.unlink()
+ removed_snapshots += 1
+ except OSError as e:
+ failures.append(f"{name}: {e}")
+
+ ref_counts = _snapshot_blob_reference_counts(repo_dir)
+ seen_blobs: set[Path] = set()
+ for _snap, blob, name in stale_matches:
+ if blob is None:
+ continue
+ try:
+ blob_key = blob.resolve()
+ except OSError:
+ blob_key = blob
+ if blob_key in seen_blobs:
+ continue
+ seen_blobs.add(blob_key)
+ if ref_counts.get(blob_key, 0) > 0:
+ continue
+ try:
+ if blob.exists():
+ deleted_bytes += blob.stat().st_size
+ blob.unlink()
+ deleted_blobs += 1
+ except OSError as e:
+ failures.append(f"{name}: {e}")
+
+ removed_dirs = 0
+ dir_failures: list[str] = []
+ if target_repos:
+ removed_dirs, dir_failures = _remove_empty_variant_dirs(target_repos, variant)
+ removed_snap_dirs, snap_dir_failures = _remove_empty_snapshot_dirs(target_repos)
+ removed_dirs += removed_snap_dirs
+ dir_failures.extend(snap_dir_failures)
+ failures.extend(dir_failures)
+
+ if failures:
+ logger.warning(
+ "Stale GGUF reclaim for %s [%s] left %d failure(s): %s",
+ repo_id,
+ variant,
+ len(failures),
+ "; ".join(failures[:3]),
+ )
+
+ if removed_snapshots or deleted_blobs or removed_dirs:
+ cache_inventory.invalidate_hf_cache_scans()
+ logger.info(
+ "Reclaimed stale GGUF %s [%s]: snapshots=%d blobs=%d dirs=%d freed=%.1f MB",
+ repo_id,
+ variant,
+ removed_snapshots,
+ deleted_blobs,
+ removed_dirs,
+ deleted_bytes / (1024 * 1024),
+ )
+
+ return {
+ "status": "reclaimed",
+ "repo_id": repo_id,
+ "variant": variant,
+ "removed_snapshots": removed_snapshots,
+ "deleted_blobs": deleted_blobs,
+ "removed_dirs": removed_dirs,
+ }
+
+
def _loaded_id_matches_repo(loaded_id: str, repo_id: str) -> bool:
"""True when *loaded_id* is *repo_id* or a file within it; ``/``-boundary aware so ``org/model`` doesn't match sibling ``org/model-v2``."""
rid = repo_id.lower()
diff --git a/studio/backend/hub/services/models/gguf_variants.py b/studio/backend/hub/services/models/gguf_variants.py
index 0c0e6f9254..0147bba19a 100644
--- a/studio/backend/hub/services/models/gguf_variants.py
+++ b/studio/backend/hub/services/models/gguf_variants.py
@@ -291,6 +291,75 @@ def _partial_transport_for_variant(repo_id: str, variant: str) -> Optional[str]:
return hf_cache_scan.partial_transport_for("model", repo_id, variant)
+def _local_main_gguf_blobs_by_quant(repo_id: str) -> dict[str, dict[str, set[str]]]:
+ """Map quant -> repo-relative expected GGUF filename -> cached blob hashes.
+
+ Shared companions are copied into each main-quant bucket so update checks can
+ detect mmproj/MTP-only upstream changes without a separate remote call.
+ """
+ result: dict[str, dict[str, set[str]]] = {}
+ companion_blobs: dict[str, set[str]] = {}
+ try:
+ from hub.services.models import cache_inventory
+ scans = cache_inventory.all_hf_cache_scans()
+ except Exception as e:
+ logger.warning("Failed to scan local GGUF blobs for %s: %s", repo_id, e)
+ return result
+
+ target_lower = repo_id.lower()
+ for hf_cache in scans:
+ for repo_info in hf_cache.repos:
+ if str(getattr(repo_info, "repo_type", "")) != "model":
+ continue
+ if str(getattr(repo_info, "repo_id", "")).lower() != target_lower:
+ continue
+ for path, hashes in cache_inventory._repo_gguf_blob_map(
+ repo_info,
+ include_companions = True,
+ ).items():
+ normalized = str(path).replace("\\", "/")
+ if not hashes:
+ continue
+ if _is_mmproj_filename(normalized) or _is_mtp_drafter_path(normalized):
+ companion_blobs.setdefault(normalized, set()).update(
+ str(blob) for blob in hashes if blob
+ )
+ continue
+ quant = extract_quant_label(normalized).lower()
+ if is_big_endian_gguf_path(normalized, quant):
+ continue
+ bucket = result.setdefault(quant, {}).setdefault(normalized, set())
+ bucket.update(str(blob) for blob in hashes if blob)
+ if companion_blobs:
+ for local_blobs in result.values():
+ for path, hashes in companion_blobs.items():
+ local_blobs.setdefault(path, set()).update(hashes)
+ return result
+
+
+def _variant_update_available_from_requirement(
+ local_blobs: dict[str, set[str]], requirement: Optional[_GgufVariantRequirement], variant: str
+) -> bool:
+ if requirement is None or not local_blobs:
+ return False
+ local_by_posix = {path.replace("\\", "/"): blobs for path, blobs in local_blobs.items()}
+ for expected in requirement.expected_files:
+ path = str(expected.path).replace("\\", "/")
+ if not (
+ is_main_gguf_variant_path(path, variant)
+ or _is_mmproj_filename(path)
+ or _is_mtp_drafter_path(path)
+ ):
+ continue
+ remote_blob = expected.sha256
+ if not remote_blob:
+ continue
+ local_set = local_by_posix.get(path)
+ if not local_set or remote_blob not in local_set:
+ return True
+ return False
+
+
def delete_variant_incomplete_blobs_result(
repo_id: str,
variant: str,
@@ -657,9 +726,12 @@ async def get_gguf_variants_response(
_partial_transport_for_variant(repo_id, variant.quant),
)
+ local_blobs_by_quant = _local_main_gguf_blobs_by_quant(repo_id)
+
def _variant_detail(v) -> GgufVariantDetail:
is_partial = v.quant in partial_quants
requirement = requirements_by_quant.get(v.quant.lower())
+ downloaded = _is_fully_downloaded(v) and not is_partial
return GgufVariantDetail(
filename = v.filename,
quant = v.quant,
@@ -668,7 +740,13 @@ async def get_gguf_variants_response(
download_size_bytes = (
requirement.download_size_bytes if requirement is not None else v.size_bytes
),
- downloaded = _is_fully_downloaded(v) and not is_partial,
+ downloaded = downloaded,
+ update_available = downloaded
+ and _variant_update_available_from_requirement(
+ local_blobs_by_quant.get(v.quant.lower(), {}),
+ requirement,
+ v.quant,
+ ),
partial = is_partial,
partial_transport = (partial_quant_transports.get(v.quant) if is_partial else None),
)
diff --git a/studio/backend/hub/tests/test_model_services.py b/studio/backend/hub/tests/test_model_services.py
index 1eb7042e4e..e22aaba282 100644
--- a/studio/backend/hub/tests/test_model_services.py
+++ b/studio/backend/hub/tests/test_model_services.py
@@ -1632,6 +1632,34 @@ def test_variant_partial_accepts_variant_filtered_legacy_hashes(monkeypatch, tmp
)
+def test_variant_partial_accepts_completed_variant_in_non_latest_snapshot(monkeypatch, tmp_path):
+ """A verified GGUF update can prune an older snapshot and make that old
+ directory the newest by mtime. The variant is still complete when another
+ snapshot satisfies its manifest."""
+ monkeypatch.setattr(state_dir, "cache_root", lambda: tmp_path / "state")
+ repo_dir = tmp_path / "cache" / "models--Org--Repo"
+ old_snapshot = repo_dir / "snapshots" / "old"
+ new_snapshot = repo_dir / "snapshots" / "new"
+ old_snapshot.mkdir(parents = True)
+ new_snapshot.mkdir(parents = True)
+ (old_snapshot / "model-Q8_0.gguf").write_bytes(b"sibling")
+ (new_snapshot / "model-Q4_K_M.gguf").write_bytes(b"new")
+ assert download_manifest.write_manifest(
+ "model",
+ "Org/Repo",
+ "Q4_K_M",
+ [download_manifest.ExpectedFile(path = "model-Q4_K_M.gguf", size = 3)],
+ "http",
+ )
+
+ assert not inventory_scan.is_variant_partial(
+ "Org/Repo",
+ "Q4_K_M",
+ snapshot_dir = old_snapshot,
+ repo_cache_dir = repo_dir,
+ )
+
+
def test_gguf_variants_partial_marker_overrides_size_only_downloaded(monkeypatch, tmp_path):
async def _run_inline(fn, *args, **kwargs):
return fn(*args, **kwargs)
diff --git a/studio/backend/hub/utils/gguf_plan.py b/studio/backend/hub/utils/gguf_plan.py
index 40d5a32549..2abdb0fb79 100644
--- a/studio/backend/hub/utils/gguf_plan.py
+++ b/studio/backend/hub/utils/gguf_plan.py
@@ -36,7 +36,10 @@ def sibling_sha256(sibling) -> Optional[str]:
value = lfs.get("sha256")
else:
value = getattr(lfs, "sha256", None)
- return value if isinstance(value, str) and value else None
+ if isinstance(value, str) and value:
+ return value
+ blob_id = getattr(sibling, "blob_id", None)
+ return blob_id if isinstance(blob_id, str) and blob_id else None
def sibling_size(sibling) -> int:
diff --git a/studio/backend/hub/utils/inventory_scan.py b/studio/backend/hub/utils/inventory_scan.py
index 0f7ce6fe34..57ad7f6655 100644
--- a/studio/backend/hub/utils/inventory_scan.py
+++ b/studio/backend/hub/utils/inventory_scan.py
@@ -387,9 +387,55 @@ def _manifest_partial(
)
if resolved is None:
return True
+ if repo_type == "model" and variant is not None:
+ if download_manifest.verify_against_disk(manifest, resolved).ok:
+ return False
+ for candidate in _manifest_snapshot_dirs(repo_type, repo_id, repo_cache_dir):
+ if candidate == resolved:
+ continue
+ if download_manifest.verify_against_disk(manifest, candidate).ok:
+ return False
+ return True
return not download_manifest.verify_against_disk(manifest, resolved).ok
+def _manifest_snapshot_dirs(
+ repo_type: RepoType,
+ repo_id: str,
+ repo_cache_dir: Optional[Path] = None,
+) -> list[Path]:
+ repo_dirs = (
+ [repo_cache_dir]
+ if repo_cache_dir is not None
+ else list(iter_repo_cache_dirs(repo_type, repo_id))
+ )
+ snapshots: list[Path] = []
+ seen: set[str] = set()
+ for repo_dir in repo_dirs:
+ if repo_dir is None:
+ continue
+ snapshots_dir = repo_dir / "snapshots"
+ try:
+ if not snapshots_dir.is_dir():
+ continue
+ entries = list(snapshots_dir.iterdir())
+ except OSError:
+ continue
+ for entry in entries:
+ try:
+ if not entry.is_dir():
+ continue
+ resolved = entry.resolve()
+ except OSError:
+ continue
+ key = str(resolved)
+ if key in seen:
+ continue
+ seen.add(key)
+ snapshots.append(resolved)
+ return snapshots
+
+
def is_snapshot_partial(
repo_type: RepoType,
repo_id: str,
diff --git a/studio/backend/hub/workers/hf_download.py b/studio/backend/hub/workers/hf_download.py
index 42a8ca52b3..e45357d311 100644
--- a/studio/backend/hub/workers/hf_download.py
+++ b/studio/backend/hub/workers/hf_download.py
@@ -653,6 +653,21 @@ def _download_gguf_variant(repo_id: str, variant: str, hf_token: str | None, mod
snapshot_path,
metadata_unavailable = metadata_unavailable,
)
+ if plan is not None:
+ try:
+ from hub.services.models.deletion import reclaim_replaced_gguf_variant
+ reclaim_replaced_gguf_variant(
+ repo_id,
+ variant,
+ plan.main_hashes,
+ hf_token,
+ )
+ except Exception as e:
+ print(
+ f"Verified GGUF update for {repo_id} [{variant}], but stale-cache "
+ f"reclaim failed ({type(e).__name__}: {e})",
+ file = sys.stderr,
+ )
def _download_dataset(repo_id: str, hf_token: str | None, mode: str) -> None:
diff --git a/studio/backend/main.py b/studio/backend/main.py
index 731625ca74..0a5b775775 100644
--- a/studio/backend/main.py
+++ b/studio/backend/main.py
@@ -441,9 +441,30 @@ def _start_llama_cpp_probes_if_enabled(app: FastAPI) -> None:
).start()
+def _warm_rag_embedder() -> None:
+ """Warm RAG embeddings without blocking backend readiness."""
+ try:
+ from storage import rag_db
+
+ if not rag_db.RAG_AVAILABLE:
+ return
+ from core.rag import embeddings
+
+ embeddings.warm()
+ except Exception:
+ pass
+
+
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Startup: detect hardware, seed default admin if needed. Shutdown: clean up compiled cache."""
+
+ import time as _time
+
+ _lifespan_started = _time.perf_counter()
+ import structlog as _structlog
+
+ _lifespan_log = _structlog.get_logger(__name__)
clear_unsloth_compiled_cache()
# Remove stale .venv_overlay from old versions; switching now uses .venv_t5/.
@@ -454,6 +475,11 @@ async def lifespan(app: FastAPI):
# Detect hardware first — sets the DEVICE global used everywhere.
detect_hardware()
+ _lifespan_log.info(
+ "lifespan hardware detection completed in %.1fms",
+ (_time.perf_counter() - _lifespan_started) * 1000,
+ )
+
# Apple Silicon with MLX missing => Train/Export are greyed out (chat-only).
# Reinstall mlx by name on a background thread (off the critical path) and
# re-detect, so a reinstall/update that dropped mlx self-heals. No-op
@@ -465,7 +491,13 @@ async def lifespan(app: FastAPI):
import structlog as _structlog
_structlog.get_logger(__name__).debug("mlx autorepair skipped: %s", _mlx_exc)
- # Reap download workers orphaned by a previous crash before new downloads start.
+ # Reap workers/runs orphaned by a previous crash before new work starts.
+ try:
+ from storage.studio_db import cleanup_orphaned_runs
+ cleanup_orphaned_runs()
+ except Exception as exc:
+ _lifespan_log.warning("cleanup_orphaned_runs failed at startup: %s", exc)
+
reap_hub_orphan_workers()
# llama.cpp probes: capability (MTP support) + freshness (release age).
@@ -479,45 +511,23 @@ async def lifespan(app: FastAPI):
app.state.llama_cpp_freshness = None
_start_llama_cpp_probes_if_enabled(app)
- from storage.studio_db import cleanup_orphaned_runs
-
- try:
- cleanup_orphaned_runs()
- except Exception as exc:
- import structlog
- structlog.get_logger(__name__).warning("cleanup_orphaned_runs failed at startup: %s", exc)
-
- # Same for RAG: fail ingestion jobs stranded mid-ingest by a crash.
try:
from storage.rag_db import reconcile_orphaned_ingestion_jobs
reconcile_orphaned_ingestion_jobs()
except Exception as exc:
- import structlog
- structlog.get_logger(__name__).warning(
- "reconcile_orphaned_ingestion_jobs failed at startup: %s", exc
- )
+ _lifespan_log.warning("reconcile_orphaned_ingestion_jobs failed at startup: %s", exc)
_start_helper_precache_if_enabled()
+ threading.Thread(target = _warm_rag_embedder, daemon = True, name = "rag-embedder-warm").start()
- # Warm the RAG embedder so the first upload skips the cold load. Non-fatal.
- def _warm_rag_embedder():
- try:
- from storage import rag_db
-
- if not rag_db.RAG_AVAILABLE:
- return
- from core.rag import embeddings
-
- embeddings.warm()
- except Exception:
- pass
-
- threading.Thread(target = _warm_rag_embedder, daemon = True).start()
-
- # Initialize RSA key pair for API key encryption (external providers)
+ # Initialize RSA key pair for API key encryption (external providers).
from core.inference.key_exchange import init_key_pair
init_key_pair()
+ _lifespan_log.info(
+ "lifespan pre-auth setup completed in %.1fms",
+ (_time.perf_counter() - _lifespan_started) * 1000,
+ )
if storage.ensure_default_admin():
bootstrap_pw = storage.get_bootstrap_password()
@@ -532,6 +542,11 @@ async def lifespan(app: FastAPI):
print("=" * 60 + "\n")
else:
app.state.bootstrap_password = storage.get_bootstrap_password()
+
+ _lifespan_log.info(
+ "lifespan startup completed in %.1fms",
+ (_time.perf_counter() - _lifespan_started) * 1000,
+ )
yield
from core.inference.llama_http import aclose as _close_llama_http
@@ -919,6 +934,21 @@ install_api_error_handlers(app)
# ============ Health and System Endpoints ============
+@app.get("/api/liveness")
+async def liveness_check():
+ """Cheap process liveness for desktop port validation."""
+ return {
+ "status": "alive",
+ "service": "Unsloth UI Backend",
+ "desktop_protocol_version": 1,
+ "desktop_manageability_version": 1,
+ "supports_desktop_auth": True,
+ "supports_desktop_backend_ownership": True,
+ "studio_root_id": _studio_root_id(),
+ **({"desktop_owner": owner} if (owner := _desktop_owner()) else {}),
+ }
+
+
@app.get("/api/health")
async def health_check(request: Request):
"""Liveness plus launcher capability bits; host fingerprint gated on a bearer.
diff --git a/studio/backend/models/export.py b/studio/backend/models/export.py
index 1e8e3c4792..7e05373f11 100644
--- a/studio/backend/models/export.py
+++ b/studio/backend/models/export.py
@@ -158,9 +158,15 @@ class ExportCommonOptions(BaseModel):
class ExportMergedModelRequest(ExportCommonOptions):
"""Request for exporting a merged PEFT model."""
- format_type: Literal["16-bit (FP16)", "4-bit (FP4)"] = Field(
+ format_type: Literal[
"16-bit (FP16)",
- description = "Export precision / format for the merged model",
+ "4-bit (FP4)",
+ "FP8 (compressed-tensors)",
+ "NVFP4 (compressed-tensors)",
+ ] = Field(
+ "16-bit (FP16)",
+ description = "Export precision / format for the merged model. The compressed-tensors "
+ "options run llm-compressor for vLLM (FP8 is data-free; NVFP4 calibrates).",
)
@@ -199,6 +205,15 @@ class ExportGGUFRequest(BaseModel):
None,
description = "Hugging Face token for GGUF upload",
)
+ imatrix: bool = Field(
+ False,
+ description = "Use an importance matrix (auto-downloads the upstream unsloth GGUF "
+ "imatrix). Required for the IQ low-bit quants such as iq2_xxs / iq4_xs.",
+ )
+ imatrix_path: Optional[str] = Field(
+ None,
+ description = "Path to a custom imatrix file; overrides the auto-download when set.",
+ )
class ExportLoRAAdapterRequest(ExportCommonOptions):
diff --git a/studio/backend/models/models.py b/studio/backend/models/models.py
index 1c7ce86c2c..29083cfa3a 100644
--- a/studio/backend/models/models.py
+++ b/studio/backend/models/models.py
@@ -136,9 +136,13 @@ class GgufVariantDetail(BaseModel):
filename: str = Field(..., description = "GGUF filename (e.g., 'gemma-3-4b-it-Q4_K_M.gguf')")
quant: str = Field(..., description = "Quantization label (e.g., 'Q4_K_M')")
size_bytes: int = Field(0, description = "File size in bytes")
+ download_size_bytes: int = Field(0, description = "Total bytes needed to download this variant")
downloaded: bool = Field(
False, description = "Whether this variant is already in the local HF cache"
)
+ update_available: bool = Field(
+ False, description = "Whether a newer version of this variant is available on HF"
+ )
class GgufVariantsResponse(BaseModel):
diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py
index e64b6f731a..ff815a2fa9 100644
--- a/studio/backend/models/training.py
+++ b/studio/backend/models/training.py
@@ -9,6 +9,8 @@ import re
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from typing import Any, Optional, List, Dict, Literal
+from utils.training_runs import normalize_project_name
+
# ASCII integer, optional single sign. Rejects "++512" and Unicode digits
# ("512") that slip through str.isdigit() + int().
@@ -97,6 +99,11 @@ class TrainingStartRequest(BaseModel):
model_name: str = Field(
..., description = "Model identifier (e.g., 'unsloth/llama-3-8b-bnb-4bit')"
)
+ project_name: Optional[str] = Field(
+ None,
+ max_length = 80,
+ description = "Optional user-defined project name appended to run folders and shown in history",
+ )
training_type: Literal["LoRA/QLoRA", "Full Finetuning", "Continued Pretraining"] = Field(
...,
description = "Training type: 'LoRA/QLoRA', 'Full Finetuning', or 'Continued Pretraining'",
@@ -155,6 +162,11 @@ class TrainingStartRequest(BaseModel):
values.setdefault("train_split", values.pop("split"))
return values
+ @field_validator("project_name")
+ @classmethod
+ def _normalize_project_name(cls, value: Optional[str]) -> Optional[str]:
+ return normalize_project_name(value)
+
# NOTE: pydantic runs all `mode="after"` validators in definition order. A
# second one, `_check_steps_or_epochs`, is defined lower in this class; keep
# these cross-field checks order-independent so the two stay decoupled.
@@ -588,6 +600,7 @@ class TrainingRunSummary(BaseModel):
id: str
status: Literal["running", "completed", "stopped", "error"]
model_name: str
+ project_name: Optional[str] = None
dataset_name: str
display_name: Optional[str] = None
started_at: str
diff --git a/studio/backend/requirements/no-torch-runtime.txt b/studio/backend/requirements/no-torch-runtime.txt
index 9796fc3a50..de321f80ed 100644
--- a/studio/backend/requirements/no-torch-runtime.txt
+++ b/studio/backend/requirements/no-torch-runtime.txt
@@ -73,4 +73,9 @@ pillow
# this file installs --no-deps; without them Studio runs with RAG disabled.
sqlite-vec==0.1.9
pymupdf==1.27.2.3
+# 0.3.x keeps pymupdf-layout (which pulls onnxruntime) an optional extra; the
+# lockstep 1.27.x line makes it a hard dep we do not need for to_markdown().
+pymupdf4llm==0.3.4
python-docx==1.2.0
+
+lxml==6.0.2
diff --git a/studio/backend/requirements/studio.txt b/studio/backend/requirements/studio.txt
index b75f114c21..2a7bc9f47a 100644
--- a/studio/backend/requirements/studio.txt
+++ b/studio/backend/requirements/studio.txt
@@ -26,4 +26,7 @@ gguf
# extras-no-deps.txt; these add the lexical+dense store and document parsing.
sqlite-vec==0.1.9
pymupdf==1.27.2.3
+# 0.3.x keeps pymupdf-layout (which pulls onnxruntime) an optional extra; the
+# lockstep 1.27.x line makes it a hard dep we do not need for to_markdown().
+pymupdf4llm==0.3.4
python-docx==1.2.0
diff --git a/studio/backend/routes/data_recipe/seed.py b/studio/backend/routes/data_recipe/seed.py
index 8fb034ea4e..57a291291e 100644
--- a/studio/backend/routes/data_recipe/seed.py
+++ b/studio/backend/routes/data_recipe/seed.py
@@ -481,6 +481,37 @@ async def upload_unstructured_file(
error = "No extractable text found in file",
)
extracted_path.write_text(extracted_text, encoding = "utf-8")
+ except ImportError as e:
+ raw_path.unlink(missing_ok = True)
+ extracted_path.unlink(missing_ok = True)
+ missing = getattr(e, "name", None)
+ expected_missing = {".pdf": "pymupdf4llm", ".docx": "mammoth"}.get(ext)
+ if isinstance(e, ModuleNotFoundError) and missing == expected_missing:
+ logger.error(
+ "data_recipe.seed.text_extraction_dependency_missing",
+ error = str(e),
+ missing = missing,
+ exc_info = True,
+ )
+ return UnstructuredFileUploadResponse(
+ file_id = file_id,
+ filename = original_filename,
+ size_bytes = size_bytes,
+ status = "error",
+ error = f"Cannot read {ext} files: the '{missing}' package is not installed.",
+ )
+ logger.error(
+ "data_recipe.seed.text_extraction_failed",
+ error = str(e),
+ exc_info = True,
+ )
+ return UnstructuredFileUploadResponse(
+ file_id = file_id,
+ filename = original_filename,
+ size_bytes = size_bytes,
+ status = "error",
+ error = "Text extraction failed.",
+ )
except Exception as e:
raw_path.unlink(missing_ok = True)
extracted_path.unlink(missing_ok = True)
diff --git a/studio/backend/routes/export.py b/studio/backend/routes/export.py
index 78c1e59d2e..cf2cb2fa70 100644
--- a/studio/backend/routes/export.py
+++ b/studio/backend/routes/export.py
@@ -343,6 +343,8 @@ async def export_gguf(
"""
try:
backend = get_export_backend()
+ # A custom path wins; otherwise the imatrix toggle requests the upstream auto-download.
+ imatrix_file = request.imatrix_path or (True if request.imatrix else None)
success, message, output_path = await asyncio.to_thread(
backend.export_gguf,
save_directory = request.save_directory,
@@ -350,6 +352,7 @@ async def export_gguf(
push_to_hub = request.push_to_hub,
repo_id = request.repo_id,
hf_token = request.hf_token,
+ imatrix_file = imatrix_file,
)
if not success:
diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py
index b7202ed346..ed6819cdea 100644
--- a/studio/backend/routes/inference.py
+++ b/studio/backend/routes/inference.py
@@ -683,7 +683,9 @@ try:
detect_reasoning_flags,
)
from core.inference.llama_server_args import (
+ _effective_tensor_parallel,
_tensor_parallel_matches_loaded,
+ parse_split_mode_override,
resolve_tensor_parallel,
strip_shadowing_flags,
validate_extra_args,
@@ -718,7 +720,9 @@ except ImportError:
detect_reasoning_flags,
)
from core.inference.llama_server_args import (
+ _effective_tensor_parallel,
_tensor_parallel_matches_loaded,
+ parse_split_mode_override,
resolve_tensor_parallel,
strip_shadowing_flags,
validate_extra_args,
@@ -2086,6 +2090,32 @@ def _should_strip_split_mode(request: LoadRequest, backend_extra: Optional[list[
)
+def _carry_preserved_tensor_intent(
+ *, preserved: bool, same_model: bool, explicit_drop: bool
+) -> bool:
+ """Carry a preserved multi-GPU layer fallback forward only for a reload of the
+ SAME loaded model that doesn't explicitly drop tensor intent, so a fitting model
+ isn't collapsed to one GPU on a ctx-only change -- but an unrelated model switch
+ (without /unload) or an explicit tensor-off doesn't inherit it (#6659)."""
+ return preserved and same_model and not explicit_drop
+
+
+def _is_explicit_tensor_drop(request: LoadRequest) -> bool:
+ """True only when the request explicitly selects a non-tensor --split-mode (e.g.
+ layer/row/none), a deliberate departure from a preserved tensor->layer fallback.
+
+ A bare tensor_parallel field is NOT a drop: the Studio UI always sends it and echoes
+ the /load response's resolved value back, so after a fallback every reload carries
+ tensor_parallel=false even though the user never changed it -- treating that as a drop
+ would collapse the preserved multi-GPU placement on the next ctx/settings reload. An
+ empty clear is not a drop either (a fallback always stores --split-mode layer, never a
+ tensor split mode, so a clear never wipes tensor intent), nor is an unrelated extra
+ (--top-k) or inherit (None). tensor_parallel=true / --split-mode tensor re-engage
+ tensor. Shared by the already-loaded dedup and the load carry-forward (#6659)."""
+ override = parse_split_mode_override(request.llama_extra_args)
+ return override is not None and override.strip().lower() != "tensor"
+
+
def _request_matches_loaded_settings(
request: LoadRequest,
llama_backend: LlamaCppBackend,
@@ -2124,6 +2154,13 @@ def _request_matches_loaded_settings(
effective_extra, request.tensor_parallel, llama_backend.tensor_parallel
):
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
+ # re-selects instead of keeping the all-GPU mask (#6659). The effective check
+ # includes the env, so an env-only tensor (LLAMA_ARG_SPLIT_MODE=tensor) that
+ # can't actually be dropped falls through to the env-downgrade match, not a loop.
+ if llama_backend.layer_preserves_tensor_intent and _is_explicit_tensor_drop(request):
+ return False
# Spec decoding works on vision models too (MTP is mmproj-compatible,
# llama.cpp #22673; the old ``not is_vision`` gate is gone), so compare
# the real requested mode -- coercing vision to ``off`` here used to
@@ -2826,6 +2863,48 @@ async def load_model(
hf_variant = config.gguf_variant,
)
+ # Tensor intent for this load: the request itself, or a preserved
+ # multi-GPU layer fallback carried across a reload of the SAME model that
+ # doesn't drop it (e.g. a ctx-only change), so a fitting model doesn't
+ # silently collapse to one GPU. Only an explicit non-tensor --split-mode
+ # override counts as the drop -- the tensor field echo / unrelated extras keep
+ # the preserved placement; the same-model guard stops a switch-without-unload
+ # inheriting the prior model's intent.
+ _explicit_tensor_drop = _is_explicit_tensor_drop(request)
+ # Compare the resolved config.identifier (what load_model stores), not the
+ # raw request id: from_identifier normalizes shorthands (adds unsloth/, fixes
+ # case), so a reload with the shorthand would otherwise miss the match and
+ # drop the carry-forward. #6659
+ _same_model_loaded = (
+ llama_backend.is_loaded
+ and (llama_backend.model_identifier or "").lower()
+ == (config.identifier or "").lower()
+ )
+ # model_identifier is variant-agnostic for HF repos and dir-level for a
+ # local multi-variant directory, so also require the loaded quant to match
+ # (path else variant, mirroring _already_in_target_state) -- otherwise a
+ # different variant inherits the prior one's preserved intent. #6659
+ if _same_model_loaded:
+ if config.gguf_file and llama_backend.gguf_path:
+ try:
+ _same_model_loaded = (
+ Path(llama_backend.gguf_path).resolve()
+ == Path(config.gguf_file).resolve()
+ )
+ except OSError:
+ _same_model_loaded = False
+ else:
+ _same_model_loaded = (llama_backend.hf_variant or "").lower() == (
+ config.gguf_variant or ""
+ ).lower()
+ _tensor_intent_overall = _effective_tensor_parallel(
+ extra_llama_args, request.tensor_parallel
+ ) or _carry_preserved_tensor_intent(
+ preserved = llama_backend.layer_preserves_tensor_intent,
+ same_model = _same_model_loaded,
+ explicit_drop = _explicit_tensor_drop,
+ )
+
# Run a single load attempt with the given tensor flag + extras.
async def _attempt_gguf_load(
tensor_parallel: bool, attempt_extra_args: Optional[list[str]]
@@ -2839,6 +2918,12 @@ async def load_model(
**_source_load_kwargs,
**attempt_kwargs,
tensor_parallel = tensor_parallel,
+ # True on the layer fallback retry (tensor wanted overall but not on
+ # this attempt): keep multi-GPU. Mirrors the fallback's key.
+ preserve_multi_gpu_on_layer = bool(
+ _tensor_intent_overall
+ and not _effective_tensor_parallel(attempt_extra_args, tensor_parallel)
+ ),
)
# Tensor parallelism is arch-gated in llama.cpp and crashes some loads
@@ -6654,7 +6739,10 @@ async def openai_completions(request: Request, current_subject: str = Depends(ge
# honor stream_options.include_usage per event, while keeping SSE
# framing and token bytes intact.
_include_usage = bool((body.get("stream_options") or {}).get("include_usage"))
- client = httpx.AsyncClient(timeout = _llama_streaming_generation_timeout())
+ client = httpx.AsyncClient(
+ timeout = _llama_streaming_generation_timeout(),
+ trust_env = False,
+ )
resp = None
bytes_iter = None
disconnect_event = threading.Event()
@@ -7729,7 +7817,10 @@ async def _responses_stream(
# `async with`, explicit aclose of lines_iter BEFORE resp / client so
# the innermost httpcore byte stream is finalised in this task (not via
# the asyncgen GC in a sibling task).
- client = httpx.AsyncClient(timeout = _llama_streaming_generation_timeout())
+ client = httpx.AsyncClient(
+ timeout = _llama_streaming_generation_timeout(),
+ trust_env = False,
+ )
resp = None
lines_iter = None
disconnect_watcher = None
@@ -9239,6 +9330,7 @@ async def _anthropic_passthrough_stream(
client = httpx.AsyncClient(
timeout = _llama_streaming_generation_timeout(),
limits = httpx.Limits(max_keepalive_connections = 0),
+ trust_env = False,
)
resp = None
lines_iter = None
@@ -9765,6 +9857,7 @@ async def _openai_passthrough_stream(
client = httpx.AsyncClient(
timeout = _llama_streaming_generation_timeout(),
limits = httpx.Limits(max_keepalive_connections = 0),
+ trust_env = False,
)
resp = None
_truncate_budget = (
diff --git a/studio/backend/routes/models.py b/studio/backend/routes/models.py
index c9e4bf2230..321f87072f 100644
--- a/studio/backend/routes/models.py
+++ b/studio/backend/routes/models.py
@@ -12,6 +12,7 @@ import sys
import uuid
from pathlib import Path
from fastapi import APIRouter, Body, Depends, Header, HTTPException, Query
+from pydantic import BaseModel
from typing import List, Optional
import structlog
from loggers import get_logger
@@ -22,10 +23,27 @@ import re as _re
_VALID_REPO_ID = _re.compile(r"^[A-Za-z0-9._-]+/[A-Za-z0-9._-]+$")
+class CachedModelRepo(BaseModel):
+ repo_id: str
+ size_bytes: int
+ last_modified: Optional[float] = None
+
+
+class CachedModelsResponse(BaseModel):
+ cached: List[CachedModelRepo]
+
+
def _is_valid_repo_id(repo_id: str) -> bool:
return bool(_VALID_REPO_ID.fullmatch(repo_id))
+def _normalize_hf_token(hf_token) -> Optional[str]:
+ if not isinstance(hf_token, str):
+ return None
+ token = hf_token.strip()
+ return token or None
+
+
def _safe_is_dir(path) -> bool:
"""``Path.is_dir()`` returning ``False`` instead of raising.
@@ -74,6 +92,7 @@ if str(backend_path) not in sys.path:
sys.path.insert(0, str(backend_path))
from auth.authentication import get_current_subject
+from hub.dependencies import get_hf_token
try:
from utils.models import (
@@ -2598,109 +2617,41 @@ async def get_gguf_variants(
..., description = "HuggingFace repo ID (e.g. 'unsloth/gemma-3-4b-it-GGUF')"
),
hf_token: Optional[str] = Query(None, description = "HuggingFace token for private repos"),
+ hf_token_header: Optional[str] = Depends(get_hf_token),
current_subject: str = Depends(get_current_subject),
):
- """List GGUF quantization variants for a HF repo or local directory.
-
- Returns all variants with file sizes, vision support, and the
- recommended default.
- """
+ """List GGUF quantization variants for a HF repo or local directory."""
try:
- from utils.models.model_config import is_local_path, list_local_gguf_variants
+ hf_token = _normalize_hf_token(hf_token_header) or _normalize_hf_token(hf_token)
+ from hub.services.models import gguf_variants as hub_gguf_variants
- # Local directory path — scan filesystem.
- if is_local_path(repo_id):
- variants, has_vision = list_local_gguf_variants(repo_id)
-
- filenames = [v.filename for v in variants]
- best = _pick_best_gguf(filenames)
- default_variant = _extract_quant_label(best) if best else None
-
- return GgufVariantsResponse(
- repo_id = repo_id,
- variants = [
- GgufVariantDetail(
- filename = v.filename,
- quant = v.quant,
- size_bytes = v.size_bytes,
- downloaded = True, # all local variants are downloaded
- )
- for v in variants
- ],
- has_vision = has_vision,
- default_variant = default_variant,
- context_length = _read_native_context_length(repo_id, is_local = True),
- )
-
- # Remote HuggingFace repo — query HF API.
- variants, has_vision = list_gguf_variants(repo_id, hf_token = hf_token)
-
- filenames = [v.filename for v in variants]
- best = _pick_best_gguf(filenames)
- default_variant = _extract_quant_label(best) if best else None
-
- # Per-snapshot so a split GGUF's shards must all sit in one snapshot;
- # mmproj adapters are excluded so they can't inflate a quant's bytes.
- cached_bytes_by_quant_per_snapshot: list[dict[str, int]] = []
- try:
- from huggingface_hub import constants as hf_constants
-
- if not _is_valid_repo_id(repo_id):
- raise ValueError(f"Invalid repo_id format: {repo_id}")
-
- cache_dir = Path(hf_constants.HF_HUB_CACHE)
- target = f"models--{repo_id.replace('/', '--')}".lower()
- for entry in cache_dir.iterdir():
- if entry.name.lower() == target:
- snapshots = entry / "snapshots"
- if snapshots.is_dir():
- for snap in snapshots.iterdir():
- by_quant: dict[str, int] = {}
- for f in _iter_gguf_paths(snap):
- if _is_mmproj_filename(f.name):
- continue
- try:
- size = f.stat().st_size
- except OSError:
- continue # broken symlink / unreadable: skip
- rel = f.relative_to(snap).as_posix()
- q = _extract_quant_label(rel)
- if _is_big_endian_gguf_path(rel, q):
- continue
- q = q.lower()
- by_quant[q] = by_quant.get(q, 0) + size
- if by_quant:
- cached_bytes_by_quant_per_snapshot.append(by_quant)
- break
- except Exception:
- pass
-
- def _is_fully_downloaded(variant) -> bool:
- if variant.size_bytes == 0:
- return False
- # Complete within one snapshot (tolerance for symlink size jitter).
- quant = variant.quant.lower()
- return any(
- by_quant.get(quant, 0) >= variant.size_bytes * 0.99
- for by_quant in cached_bytes_by_quant_per_snapshot
- )
+ response = await hub_gguf_variants.get_gguf_variants_response(
+ repo_id,
+ hf_token = hf_token,
+ )
+ local = is_local_path(repo_id)
return GgufVariantsResponse(
- repo_id = repo_id,
+ repo_id = response.repo_id,
variants = [
GgufVariantDetail(
filename = v.filename,
quant = v.quant,
size_bytes = v.size_bytes,
- downloaded = _is_fully_downloaded(v),
+ download_size_bytes = int(
+ getattr(v, "download_size_bytes", v.size_bytes) or v.size_bytes
+ ),
+ downloaded = bool(v.downloaded),
+ update_available = bool(getattr(v, "update_available", False)),
)
- for v in variants
+ for v in response.variants
],
- has_vision = has_vision,
- default_variant = default_variant,
- context_length = _read_native_context_length(repo_id, is_local = False),
+ has_vision = response.has_vision,
+ default_variant = response.default_variant,
+ context_length = _read_native_context_length(repo_id, is_local = local),
)
-
+ except HTTPException:
+ raise
except Exception as e:
logger.error(f"Error listing GGUF variants for '{repo_id}': {e}", exc_info = True)
raise HTTPException(
@@ -3246,10 +3197,14 @@ def _repo_is_diffusers(repo_info) -> bool:
return False
-@router.get("/cached-models")
-async def list_cached_models(current_subject: str = Depends(get_current_subject)):
+@router.get("/cached-models", response_model = CachedModelsResponse)
+async def list_cached_models(
+ current_subject: str = Depends(get_current_subject),
+ hf_token: Optional[str] = Depends(get_hf_token),
+):
"""List non-GGUF model repos downloaded to HF cache, legacy Unsloth cache, and HF default cache."""
_WEIGHT_EXTENSIONS = (".safetensors", ".bin")
+ hf_token = _normalize_hf_token(hf_token)
try:
cache_scans = _all_hf_cache_scans()
@@ -3270,20 +3225,16 @@ async def list_cached_models(current_subject: str = Depends(get_current_subject)
)
if total_size == 0:
continue
- has_weights = any(
- f.file_name.endswith(_WEIGHT_EXTENSIONS)
+ weight_files = [
+ f
for rev in repo_info.revisions
for f in rev.files
- )
- if not has_weights:
+ if f.file_name.endswith(_WEIGHT_EXTENSIONS)
+ ]
+ if not weight_files:
continue
last_modified = max(
- (
- _blob_mtime(f)
- for rev in repo_info.revisions
- for f in rev.files
- if f.file_name.endswith(_WEIGHT_EXTENSIONS)
- ),
+ (_blob_mtime(f) for f in weight_files),
default = 0.0,
)
key = repo_id.lower()
@@ -3306,9 +3257,12 @@ async def list_cached_models(current_subject: str = Depends(get_current_subject)
repo_label = getattr(repo_info, "repo_id", "")
logger.warning(f"Skipping cached model repo {repo_label}: {e}")
continue
- # Newest download first; stable repo_id tie-break for equal/missing mtimes.
+
+ rows = list(seen_lower.values())
+ # Local-only list path: update checks are GGUF-only and happen lazily
+ # when a repo's variants are viewed.
cached = sorted(
- seen_lower.values(),
+ rows,
key = lambda c: (-(c.get("last_modified") or 0.0), c["repo_id"].lower()),
)
return {"cached": cached}
diff --git a/studio/backend/routes/rag.py b/studio/backend/routes/rag.py
index 8d23240fd5..4e35fce3c2 100644
--- a/studio/backend/routes/rag.py
+++ b/studio/backend/routes/rag.py
@@ -19,7 +19,7 @@ import secrets
import time
import uuid
-from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
+from fastapi import APIRouter, Depends, File, Form, HTTPException, Query, UploadFile
from fastapi.responses import FileResponse, StreamingResponse
from pydantic import BaseModel, Field
@@ -62,13 +62,24 @@ def _save_upload(file: UploadFile) -> tuple[str, str]:
uploads = ensure_dir(rag_uploads_root())
stored_path = str(uploads / f"{uuid.uuid4().hex}{ext}")
size = 0
+ cap = config.MAX_UPLOAD_BYTES
+ too_big = False
with open(stored_path, "wb") as out:
while True:
block = file.file.read(1 << 20)
if not block:
break
size += len(block)
+ if cap and size > cap:
+ too_big = True
+ break
out.write(block)
+ if too_big:
+ os.remove(stored_path)
+ raise HTTPException(
+ status_code = 413,
+ detail = f"File exceeds the {cap // (1024 * 1024)} MB upload limit.",
+ )
if size == 0:
os.remove(stored_path)
raise HTTPException(status_code = 400, detail = "Uploaded file is empty.")
@@ -207,6 +218,8 @@ def delete_knowledge_base(kb_id: str, subject: str = Depends(get_current_subject
async def upload_kb_document(
kb_id: str,
file: UploadFile = File(...),
+ ocr: bool | None = Form(None),
+ caption: bool | None = Form(None),
subject: str = Depends(get_current_subject),
) -> dict:
_require_rag()
@@ -218,7 +231,7 @@ async def upload_kb_document(
conn.close()
stored_path, filename = _save_upload(file)
document_id, job_id = ingestion.start_ingestion(
- store.kb_scope(kb_id), kb_id, None, filename, stored_path
+ store.kb_scope(kb_id), kb_id, None, filename, stored_path, ocr = ocr, caption = caption
)
return {"documentId": document_id, "jobId": job_id, "filename": filename}
@@ -238,12 +251,20 @@ def list_kb_documents(kb_id: str, subject: str = Depends(get_current_subject)) -
async def upload_thread_document(
thread_id: str,
file: UploadFile = File(...),
+ ocr: bool | None = Form(None),
+ caption: bool | None = Form(None),
subject: str = Depends(get_current_subject),
) -> dict:
_require_rag()
stored_path, filename = _save_upload(file)
document_id, job_id = ingestion.start_ingestion(
- store.thread_scope(thread_id), None, thread_id, filename, stored_path
+ store.thread_scope(thread_id),
+ None,
+ thread_id,
+ filename,
+ stored_path,
+ ocr = ocr,
+ caption = caption,
)
return {"documentId": document_id, "jobId": job_id, "filename": filename}
@@ -263,6 +284,8 @@ def list_thread_documents(thread_id: str, subject: str = Depends(get_current_sub
async def upload_project_document(
project_id: str,
file: UploadFile = File(...),
+ ocr: bool | None = Form(None),
+ caption: bool | None = Form(None),
subject: str = Depends(get_current_subject),
) -> dict:
_require_rag()
@@ -278,6 +301,8 @@ async def upload_project_document(
filename,
stored_path,
project_id = project_id,
+ ocr = ocr,
+ caption = caption,
)
return {"documentId": document_id, "jobId": job_id, "filename": filename}
@@ -321,6 +346,7 @@ def job_status(job_id: str, subject: str = Depends(get_current_subject)) -> dict
"stage": row.get("stage"),
"progress": row.get("progress") or 0.0,
"error": row.get("error"),
+ "numChunks": row.get("num_chunks") or 0,
}
diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py
index b232af82f5..d5ab4ae29a 100644
--- a/studio/backend/routes/training.py
+++ b/studio/backend/routes/training.py
@@ -255,6 +255,7 @@ async def start_training(
# Convert request to backend kwargs.
training_kwargs = {
"model_name": request.model_name,
+ "project_name": request.project_name,
"training_type": request.training_type,
"hf_token": request.hf_token or "",
"load_in_4bit": request.load_in_4bit,
diff --git a/studio/backend/run.py b/studio/backend/run.py
index d4cbc26b41..ccd0113972 100644
--- a/studio/backend/run.py
+++ b/studio/backend/run.py
@@ -253,12 +253,13 @@ def _verify_global_reachability(display_host: str, port: int) -> None:
local_url_c = "\033[38;5;108;1m" if use_color else "" # matches banner's URL color
reset = "\033[0m" if use_color else ""
- url = f"http://{display_host}:{port}"
+ url = f"http://{_url_host(display_host)}:{port}"
# Private/loopback/link-local addresses aren't globally routable.
try:
addr = ipaddress.ip_address(display_host)
if addr.is_loopback or addr.is_private or addr.is_link_local:
+ _public_reachable = False
print(
f"{dim} Note: {display_host} is a private/LAN address -- "
f"reachable on this network only, not from the public internet."
@@ -380,6 +381,20 @@ def _verify_global_reachability(display_host: str, port: int) -> None:
pass
+def _display_host_for_bind(host: str) -> str:
+ return _resolve_external_ip() if host in ("0.0.0.0", "::") else host
+
+
+def _loopback_bind_host_for(host: str) -> str:
+ return "::1" if host == "::" else "127.0.0.1"
+
+
+def _url_host(host: str) -> str:
+ return (
+ f"[{host}]" if ":" in host and not (host.startswith("[") and host.endswith("]")) else host
+ )
+
+
def _tool_policy_notice(host: str, secure: bool, enable_tools: "Optional[bool]") -> str:
"""One-line tool-policy summary for the plain-server startup banner, so a
network-reachable launch is never silent about code execution."""
@@ -416,7 +431,7 @@ def _emit_secure_startup_output(port: int, enable_tools: "Optional[bool]" = None
print("")
print("🦥 Unsloth Studio is running (secure)")
print("─" * 52)
- _print_cloudflare_line()
+ _print_cloudflare_line(secure = True)
print(f" On this machine only: http://127.0.0.1:{port}/")
print("─" * 52)
_emit_tool_policy_notice("127.0.0.1", True, enable_tools)
@@ -447,30 +462,108 @@ def _emit_startup_output(
_print_localhost_ipv6_mismatch_warning(localhost_mismatch_url, port)
elif wildcard_bind:
_verify_global_reachability(display_host, port)
- _print_cloudflare_line()
+ _print_cloudflare_line(loopback_host = _loopback_bind_host_for(host))
_emit_tool_policy_notice(host, False, enable_tools)
print_studio_stop_hint()
-def _print_cloudflare_line() -> None:
- """Print the Cloudflare quick-tunnel URL for 0.0.0.0 binds, if one is up.
-
- Reads the module-level URL set by ``run_server``. Prints nothing when the
- tunnel is disabled or failed -- failures are silently ignored. When the public
- reachability probe just failed (``_public_reachable is False``) but the tunnel
- is up, reword to point the user at the Cloudflare link as the way in.
- """
- if not _cloudflare_url:
- return
+def _print_cloudflare_line(secure: bool = False, loopback_host: str = "127.0.0.1") -> None:
+ """Print Cloudflare tunnel state for startup banners."""
from startup_banner import stdout_supports_color
accent = "\033[38;5;150;1m"
+ warn = "\033[38;5;215;1m"
reset = "\033[0m"
- if _public_reachable is False:
- line = f" Use the secure link access via Cloudflare instead: {_cloudflare_url}"
- else:
- line = f" Secure link access via Cloudflare: {_cloudflare_url}"
- print(f"{accent}{line}{reset}" if stdout_supports_color() else line)
+ color = stdout_supports_color()
+
+ def _emit(text: str, style: str = "") -> None:
+ print(f"{style}{text}{reset}" if (color and style) else text)
+
+ if _cloudflare_url:
+ if _public_reachable is False:
+ _emit(f" Use the secure link access via Cloudflare instead: {_cloudflare_url}", accent)
+ else:
+ _emit(f" Secure link access via Cloudflare: {_cloudflare_url}", accent)
+ if not secure:
+ if _public_reachable is True:
+ _emit(
+ " Cloudflare tunnel: ON. This Cloudflare URL is PUBLIC, and the "
+ "raw port is also publicly reachable. --no-cloudflare disables "
+ f"only the Cloudflare URL; bind {loopback_host} or close firewall "
+ "access to keep Studio private.",
+ warn,
+ )
+ else:
+ _emit(
+ " Cloudflare tunnel: ON. This is a PUBLIC internet URL: anyone "
+ "who has it can reach this Studio. Relaunch with --no-cloudflare "
+ f"to disable the Cloudflare URL; bind {loopback_host} or close "
+ "firewall access to keep Studio private.",
+ warn,
+ )
+ return
+ if _cloudflare_requested:
+ if _public_reachable is True:
+ _emit(
+ " Cloudflare tunnel: requested but failed to start. The raw port is "
+ "still reachable from the public internet (see the reachability check "
+ "above): anyone who can reach it can access this Studio.",
+ warn,
+ )
+ elif _public_reachable is False:
+ _emit(
+ " Cloudflare tunnel: requested but failed to start. Studio is reachable "
+ "on your local network only (no public link).",
+ warn,
+ )
+ else:
+ _emit(
+ " Cloudflare tunnel: requested but failed to start. There is no "
+ "Cloudflare public link. Raw port reachability was not verified; "
+ f"bind {loopback_host} or close firewall access to keep Studio private.",
+ warn,
+ )
+ elif _cloudflare_flag:
+ if _public_reachable is True:
+ _emit(
+ " Cloudflare tunnel: OFF for this mode. The raw port is still "
+ "reachable from the public internet (see the reachability check above): "
+ "anyone who can reach it can access this Studio.",
+ warn,
+ )
+ elif _public_reachable is False:
+ _emit(
+ " Cloudflare tunnel: OFF for this mode. Studio is reachable on your "
+ "local network only (no public link)."
+ )
+ else:
+ _emit(
+ " Cloudflare tunnel: OFF for this mode. There is no Cloudflare public "
+ "link. Raw port reachability was not verified; "
+ f"bind {loopback_host} or close firewall access to keep Studio private.",
+ warn,
+ )
+ elif not _cloudflare_flag:
+ if _public_reachable is True:
+ _emit(
+ " Cloudflare tunnel: OFF (--no-cloudflare). The raw port is still "
+ "reachable from the public internet (see the reachability check above): "
+ "--no-cloudflare disables only the Cloudflare link, not the public bind.",
+ warn,
+ )
+ elif _public_reachable is False:
+ _emit(
+ " Cloudflare tunnel: OFF (--no-cloudflare). Studio is reachable on your "
+ "local network only. Omit --no-cloudflare to expose a public "
+ "Cloudflare HTTPS link."
+ )
+ else:
+ _emit(
+ " Cloudflare tunnel: OFF (--no-cloudflare). There is no Cloudflare "
+ "public link. Raw port reachability was not verified; "
+ f"bind {loopback_host} or close firewall access to keep Studio private.",
+ warn,
+ )
def _get_pid_on_port(port: int) -> "tuple[int, str] | None":
@@ -697,7 +790,7 @@ _server_thread = None
# Shutdown event -- wakes the main loop on signal.
_shutdown_event = None
-# trycloudflare.com URL for 0.0.0.0 binds (set by run_server, read by the banner);
+# trycloudflare.com URL for wildcard binds (set by run_server, read by the banner);
# None when there is no tunnel (loopback, disabled, or a silently-ignored failure).
_cloudflare_url = None
@@ -707,6 +800,9 @@ _cloudflare_url = None
# not decide (timeout, blocked, private address).
_public_reachable = None
+_cloudflare_requested = False
+_cloudflare_flag = True
+
_DEFAULT_FRONTEND_PATH = Path(__file__).resolve().parent.parent / "frontend" / "dist"
@@ -880,12 +976,12 @@ def _cloudflare_tunnel_should_start(
) -> bool:
"""Whether to start the Cloudflare tunnel. --secure exposes only the tunnel
(loopback bind), so it tunnels even api-only (headless secure API serving);
- otherwise tunnel only a 0.0.0.0 bind, never api-only (Tauri) or Colab."""
+ otherwise tunnel wildcard binds, never api-only (Tauri) or Colab."""
if is_colab or not cloudflare:
return False
if secure:
return True
- return host == "0.0.0.0" and not api_only
+ return host in ("0.0.0.0", "::") and not api_only
def _apply_cli_tool_policy(enable_tools: "Optional[bool]") -> None:
@@ -933,6 +1029,9 @@ def run_server(
"""
global _server, _server_thread, _shutdown_event
+ boot_started = time.perf_counter()
+ logger.info("run_server startup begin api_only=%s host=%s port=%s", api_only, host, port)
+
# Reap every child if the parent dies abnormally (terminal close, Task
# Manager kill, SIGKILL); must run before any child can spawn.
from utils.process_lifetime import initialize_parent_lifetime
@@ -984,7 +1083,14 @@ def run_server(
from threading import Thread, Event
import uvicorn
+ import_started = time.perf_counter()
+
from main import app, setup_frontend, _IS_COLAB
+
+ logger.info(
+ "Imported FastAPI app in %.1fms",
+ (time.perf_counter() - import_started) * 1000,
+ )
from utils.paths import ensure_studio_directories
# Allow local stdio MCP servers on a loopback bind (the user's own machine),
@@ -997,6 +1103,11 @@ def run_server(
# Create all standard directories on startup.
ensure_studio_directories()
+ logger.info(
+ "Ensured Studio directories in %.1fms",
+ (time.perf_counter() - boot_started) * 1000,
+ )
+
# Auto-find a free port if the requested one is in use.
if not _is_port_free(host, port):
original_port = port
@@ -1057,9 +1168,14 @@ def run_server(
)
# Resolve once; shared by the log rewrite and banner.
- display_host = _resolve_external_ip() if host == "0.0.0.0" else host
+ display_host = _display_host_for_bind(host)
_install_uvicorn_startup_log_rewrite(host, display_host)
+ logger.info(
+ "run_server pre-uvicorn setup completed in %.1fms",
+ (time.perf_counter() - boot_started) * 1000,
+ )
+
ready_event = Event()
startup_failed = Event()
startup_errors = []
@@ -1068,6 +1184,10 @@ def run_server(
async def startup(self, *args, **kwargs):
await super().startup(*args, **kwargs)
if getattr(self, "started", False) and not self.should_exit:
+ logger.info(
+ "Uvicorn startup hook completed in %.1fms",
+ (time.perf_counter() - boot_started) * 1000,
+ )
ready_event.set()
# server_header=False suppresses uvicorn's "Server: uvicorn"; SecurityHeadersMiddleware sets its own.
@@ -1093,13 +1213,10 @@ def run_server(
# backend, not whatever a proxy/tunnel exposed. For ephemeral binds (port==0)
# leave it unset so handlers fall back to the request scope / base_url.
app.state.server_port = port if port and port > 0 else None
- # Direct (non-tunnel) base for the API panel; resolve 0.0.0.0 to the LAN IP.
+ # Direct (non-tunnel) base for the API panel; resolve wildcard binds to the LAN IP.
if port and port > 0:
- _direct_host = _resolve_external_ip() if host in ("0.0.0.0", "::") else host
- # Bracket IPv6 literals so the URL is valid (http://[2405:...]:port).
- if ":" in _direct_host and not _direct_host.startswith("["):
- _direct_host = f"[{_direct_host}]"
- app.state.server_url = f"http://{_direct_host}:{port}"
+ _direct_host = _display_host_for_bind(host)
+ app.state.server_url = f"http://{_url_host(_direct_host)}:{port}"
else:
app.state.server_url = None
app.state.secure = secure
@@ -1150,6 +1267,11 @@ def run_server(
_shutdown_event.set()
raise
+ logger.info(
+ "run_server uvicorn ready after %.1fms",
+ (time.perf_counter() - boot_started) * 1000,
+ )
+
_write_pid_file()
import atexit
@@ -1163,11 +1285,12 @@ def run_server(
if api_only and emit_tauri_port:
print(f"TAURI_PORT={port}", flush = True)
- # Free trycloudflare.com tunnel for 0.0.0.0 binds (the raw ip:port is often
+ # Free trycloudflare.com tunnel for wildcard binds (the raw ip:port is often
# unreachable). Started pre-banner and even when silent so the CLI banner can
# read app.state.cloudflare_url; torn down by _graceful_shutdown.
- global _cloudflare_url
+ global _cloudflare_url, _cloudflare_requested, _cloudflare_flag
_cloudflare_url = None
+ _cloudflare_flag = cloudflare
app.state.cloudflare_url = None
_cloudflare_enabled = _cloudflare_tunnel_should_start(
cloudflare = cloudflare,
@@ -1176,6 +1299,7 @@ def run_server(
api_only = api_only,
is_colab = _IS_COLAB,
)
+ _cloudflare_requested = _cloudflare_enabled
if _cloudflare_enabled:
try: # best-effort: any failure must not block startup
from cloudflare_tunnel import start_studio_tunnel, stop_studio_tunnel
@@ -1280,8 +1404,10 @@ def _build_arg_parser():
"--cloudflare",
action = argparse.BooleanOptionalAction,
default = True,
- help = "Auto-create a free Cloudflare HTTPS tunnel when bound to 0.0.0.0 "
- "(default on; --no-cloudflare to disable)",
+ help = "Auto-create a free Cloudflare HTTPS tunnel for non-api-only wildcard "
+ "binds (0.0.0.0 or ::), exposing Studio on a PUBLIC internet URL (default on). "
+ "Pass --no-cloudflare to disable that Cloudflare URL; it does not change a "
+ "public wildcard bind. --api-only keeps it off unless paired with --secure.",
)
parser.add_argument(
"--secure",
diff --git a/studio/backend/storage/rag_db.py b/studio/backend/storage/rag_db.py
index 4da600d768..ce27326562 100644
--- a/studio/backend/storage/rag_db.py
+++ b/studio/backend/storage/rag_db.py
@@ -119,6 +119,10 @@ def get_connection() -> sqlite3.Connection:
ensure_dir(db_path.parent)
conn = sqlite3.connect(str(db_path))
conn.row_factory = sqlite3.Row
+ # Wait for a lock instead of erroring immediately: a figure/scan-heavy ingest can
+ # hold its connection across many seconds of vision calls, and a concurrent ingest
+ # or autoinject read would otherwise hit "database is locked".
+ conn.execute("PRAGMA busy_timeout = 5000")
try:
conn.enable_load_extension(True)
sqlite_vec.load(conn)
diff --git a/studio/backend/storage/studio_db.py b/studio/backend/storage/studio_db.py
index 7421b42b2f..23b90d7002 100644
--- a/studio/backend/storage/studio_db.py
+++ b/studio/backend/storage/studio_db.py
@@ -23,6 +23,16 @@ from typing import Any, Iterable, Optional
from utils.paths import project_workspaces_root, studio_db_path, ensure_dir
+from utils.training_runs import extract_project_name
+
+
+def _extract_project_name_from_config_json(config_json: Optional[str]) -> Optional[str]:
+ if not config_json:
+ return None
+ try:
+ return extract_project_name(json.loads(config_json))
+ except (json.JSONDecodeError, TypeError):
+ return None
def _denied_path_prefixes() -> list[str]:
@@ -680,6 +690,7 @@ def list_runs(limit: int = 50, offset: int = 0) -> dict:
runs = []
for row in rows:
run = dict(row)
+ run["project_name"] = _extract_project_name_from_config_json(run.get("config_json"))
sparkline = run.get("loss_sparkline")
if sparkline:
try:
@@ -719,6 +730,7 @@ def get_run(id: str) -> Optional[dict]:
if row is None:
return None
run = dict(row)
+ run["project_name"] = _extract_project_name_from_config_json(run.get("config_json"))
sparkline = run.get("loss_sparkline")
if sparkline:
try:
diff --git a/studio/backend/tests/test_api_perf_serialization.py b/studio/backend/tests/test_api_perf_serialization.py
index f5ad53306d..348e09104c 100644
--- a/studio/backend/tests/test_api_perf_serialization.py
+++ b/studio/backend/tests/test_api_perf_serialization.py
@@ -57,6 +57,15 @@ def test_media_type_and_status():
assert err.status_code == 503
+def test_pooled_client_disables_proxy_env():
+ async def _scenario():
+ client = llama_http.nonstreaming_client()
+ assert client.trust_env is False
+ await llama_http.aclose()
+
+ asyncio.run(_scenario())
+
+
def test_pooled_client_reused_within_loop_and_recreated_after_close():
async def _scenario():
a = llama_http.nonstreaming_client()
diff --git a/studio/backend/tests/test_cached_gguf_routes.py b/studio/backend/tests/test_cached_gguf_routes.py
index 58ab41aca7..365e7624c4 100644
--- a/studio/backend/tests/test_cached_gguf_routes.py
+++ b/studio/backend/tests/test_cached_gguf_routes.py
@@ -20,6 +20,7 @@ if "structlog" not in sys.modules:
)
import routes.models as models_route
+from hub.services.models import gguf_variants as GV
def _repo(
@@ -564,21 +565,32 @@ def test_gguf_variants_mmproj_does_not_mark_quant_downloaded(monkeypatch, tmp_pa
"""The per-quant 'downloaded' flag is driven by the real weight file in a
single snapshot; an mmproj vision adapter (matching a quant label) must
not make that quant appear downloaded."""
- import huggingface_hub.constants as hf_constants
-
variants = [
- SimpleNamespace(filename = "model-Q4_K_M.gguf", quant = "Q4_K_M", size_bytes = 10_000),
- SimpleNamespace(filename = "model-F16.gguf", quant = "F16", size_bytes = 20_000),
+ SimpleNamespace(
+ filename = "model-Q4_K_M.gguf",
+ quant = "Q4_K_M",
+ display_label = None,
+ size_bytes = 10_000,
+ ),
+ SimpleNamespace(
+ filename = "model-F16.gguf",
+ quant = "F16",
+ display_label = None,
+ size_bytes = 20_000,
+ ),
]
monkeypatch.setattr(
- models_route, "list_gguf_variants", lambda repo_id, hf_token = None: (variants, True)
+ GV,
+ "list_gguf_variants",
+ lambda repo_id, hf_token = None: (variants, True, []),
)
- monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
+ monkeypatch.setattr(GV, "_local_main_gguf_blobs_by_quant", lambda _repo_id: {})
snap = tmp_path / "models--org--repo" / "snapshots" / "rev"
snap.mkdir(parents = True)
(snap / "model-Q4_K_M.gguf").write_bytes(b"x" * 10_000) # real weight, fully present
(snap / "mmproj-F16.gguf").write_bytes(b"y" * 20_000) # mmproj adapter, label "F16"
+ monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap])
result = asyncio.run(
models_route.get_gguf_variants(
@@ -592,21 +604,32 @@ def test_gguf_variants_mmproj_does_not_mark_quant_downloaded(monkeypatch, tmp_pa
def test_gguf_variants_ignore_big_endian_siblings(monkeypatch, tmp_path):
- import huggingface_hub.constants as hf_constants
-
siblings = [
SimpleNamespace(rfilename = "model-Q4_K_M-be.gguf", size = 100),
SimpleNamespace(rfilename = "model-Q4_K_M.gguf", size = 10),
]
monkeypatch.setattr(
- "huggingface_hub.model_info",
- lambda *_args, **_kwargs: SimpleNamespace(siblings = siblings),
+ GV,
+ "list_gguf_variants",
+ lambda repo_id, hf_token = None: (
+ [
+ SimpleNamespace(
+ filename = "model-Q4_K_M.gguf",
+ quant = "Q4_K_M",
+ display_label = None,
+ size_bytes = 10,
+ )
+ ],
+ False,
+ siblings,
+ ),
)
- monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
+ monkeypatch.setattr(GV, "_local_main_gguf_blobs_by_quant", lambda _repo_id: {})
snap = tmp_path / "models--org--repo" / "snapshots" / "rev"
snap.mkdir(parents = True)
(snap / "model-Q4_K_M.gguf").write_bytes(b"x" * 10)
+ monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap])
result = asyncio.run(
models_route.get_gguf_variants(
@@ -620,19 +643,25 @@ def test_gguf_variants_ignore_big_endian_siblings(monkeypatch, tmp_path):
def test_gguf_variants_cached_big_endian_does_not_satisfy_variant(monkeypatch, tmp_path):
- import huggingface_hub.constants as hf_constants
-
variants = [
- SimpleNamespace(filename = "model-Q4_K_M.gguf", quant = "Q4_K_M", size_bytes = 10),
+ SimpleNamespace(
+ filename = "model-Q4_K_M.gguf",
+ quant = "Q4_K_M",
+ display_label = None,
+ size_bytes = 10,
+ ),
]
monkeypatch.setattr(
- models_route, "list_gguf_variants", lambda repo_id, hf_token = None: (variants, False)
+ GV,
+ "list_gguf_variants",
+ lambda repo_id, hf_token = None: (variants, False, []),
)
- monkeypatch.setattr(hf_constants, "HF_HUB_CACHE", str(tmp_path))
+ monkeypatch.setattr(GV, "_local_main_gguf_blobs_by_quant", lambda _repo_id: {})
snap = tmp_path / "models--org--repo" / "snapshots" / "rev"
snap.mkdir(parents = True)
(snap / "model-Q4_K_M-be.gguf").write_bytes(b"x" * 10)
+ monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap])
result = asyncio.run(
models_route.get_gguf_variants(
diff --git a/studio/backend/tests/test_checkpoints_scan.py b/studio/backend/tests/test_checkpoints_scan.py
new file mode 100644
index 0000000000..6d473146f5
--- /dev/null
+++ b/studio/backend/tests/test_checkpoints_scan.py
@@ -0,0 +1,256 @@
+# 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 json
+import sqlite3
+import sys
+import types as _types
+from pathlib import Path
+
+_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
+if _BACKEND_DIR not in sys.path:
+ sys.path.insert(0, _BACKEND_DIR)
+
+_loggers_stub = _types.ModuleType("loggers")
+_loggers_stub.get_logger = lambda name: __import__("logging").getLogger(name)
+sys.modules.setdefault("loggers", _loggers_stub)
+sys.modules.setdefault("structlog", _types.ModuleType("structlog"))
+
+from utils.models import checkpoints as checkpoints_module
+from utils.training_runs import build_default_output_dir_name
+
+
+def _make_history_connection(db_path: Path) -> sqlite3.Connection:
+ conn = sqlite3.connect(str(db_path))
+ conn.row_factory = sqlite3.Row
+ return conn
+
+
+def _setup_training_runs_table(db_path: Path) -> None:
+ conn = _make_history_connection(db_path)
+ try:
+ conn.execute(
+ """
+ CREATE TABLE training_runs (
+ id TEXT PRIMARY KEY,
+ model_name TEXT NOT NULL,
+ config_json TEXT NOT NULL,
+ output_dir TEXT,
+ started_at TEXT NOT NULL
+ )
+ """
+ )
+ conn.commit()
+ finally:
+ conn.close()
+
+
+def _make_outputs_dir(tmp_path, monkeypatch) -> Path:
+ studio_home = tmp_path / "studio-home"
+ outputs_dir = studio_home / "outputs"
+ outputs_dir.mkdir(parents = True)
+ monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(studio_home))
+ return outputs_dir
+
+
+def test_scan_checkpoints_uses_output_dir_history_for_base_model(tmp_path, monkeypatch):
+ outputs_dir = _make_outputs_dir(tmp_path, monkeypatch)
+ run_dir = outputs_dir / "custom-run"
+ run_dir.mkdir()
+ (run_dir / "config.json").write_text("{}")
+
+ db_path = tmp_path / "studio.db"
+ _setup_training_runs_table(db_path)
+ conn = _make_history_connection(db_path)
+ try:
+ conn.execute(
+ """
+ INSERT INTO training_runs (id, model_name, config_json, output_dir, started_at)
+ VALUES (?, ?, ?, ?, ?)
+ """,
+ (
+ "run-1",
+ "unsloth/Llama-3.2-3B-Instruct",
+ "{}",
+ str(run_dir.resolve()),
+ "2026-04-09T00:00:00Z",
+ ),
+ )
+ conn.commit()
+ finally:
+ conn.close()
+
+ monkeypatch.setattr(
+ checkpoints_module,
+ "get_connection",
+ lambda: _make_history_connection(db_path),
+ )
+
+ models = checkpoints_module.scan_checkpoints(outputs_dir = str(outputs_dir))
+
+ assert models[0][2]["base_model"] == "unsloth/Llama-3.2-3B-Instruct"
+
+
+def test_scan_checkpoints_matches_project_suffixed_default_dir_against_history(
+ tmp_path, monkeypatch
+):
+ outputs_dir = _make_outputs_dir(tmp_path, monkeypatch)
+ run_name = build_default_output_dir_name(
+ "unsloth/Llama-3.2-3B-Instruct",
+ "Customer Support",
+ timestamp = 1771227800,
+ )
+ run_dir = outputs_dir / run_name
+ run_dir.mkdir()
+ (run_dir / "config.json").write_text("{}")
+
+ db_path = tmp_path / "studio.db"
+ _setup_training_runs_table(db_path)
+ conn = _make_history_connection(db_path)
+ try:
+ conn.execute(
+ """
+ INSERT INTO training_runs (id, model_name, config_json, output_dir, started_at)
+ VALUES (?, ?, ?, ?, ?)
+ """,
+ (
+ "run-2",
+ "unsloth/Llama-3.2-3B-Instruct",
+ json.dumps({"project_name": "Customer Support"}),
+ None,
+ "2026-04-09T00:00:00Z",
+ ),
+ )
+ conn.commit()
+ finally:
+ conn.close()
+
+ monkeypatch.setattr(
+ checkpoints_module,
+ "get_connection",
+ lambda: _make_history_connection(db_path),
+ )
+
+ models = checkpoints_module.scan_checkpoints(outputs_dir = str(outputs_dir))
+
+ assert models[0][2]["base_model"] == "unsloth/Llama-3.2-3B-Instruct"
+
+
+def test_scan_checkpoints_strips_project_suffix_without_history(tmp_path, monkeypatch):
+ outputs_dir = _make_outputs_dir(tmp_path, monkeypatch)
+ run_name = build_default_output_dir_name(
+ "unsloth/Llama-3.2-3B-Instruct",
+ "Customer Support",
+ timestamp = 1771227800,
+ )
+ run_dir = outputs_dir / run_name
+ run_dir.mkdir()
+ (run_dir / "config.json").write_text("{}")
+
+ db_path = tmp_path / "studio.db"
+ _setup_training_runs_table(db_path)
+ monkeypatch.setattr(
+ checkpoints_module,
+ "get_connection",
+ lambda: _make_history_connection(db_path),
+ )
+
+ models = checkpoints_module.scan_checkpoints(outputs_dir = str(outputs_dir))
+
+ assert models[0][2]["base_model"] == "unsloth/Llama-3.2-3B-Instruct"
+
+
+def test_scan_checkpoints_preserves_project_marker_in_model_without_history(tmp_path, monkeypatch):
+ outputs_dir = _make_outputs_dir(tmp_path, monkeypatch)
+ run_name = build_default_output_dir_name(
+ "org/foo__project-bar",
+ timestamp = 1771227800,
+ )
+ run_dir = outputs_dir / run_name
+ run_dir.mkdir()
+ (run_dir / "config.json").write_text("{}")
+
+ db_path = tmp_path / "studio.db"
+ _setup_training_runs_table(db_path)
+ monkeypatch.setattr(
+ checkpoints_module,
+ "get_connection",
+ lambda: _make_history_connection(db_path),
+ )
+
+ models = checkpoints_module.scan_checkpoints(outputs_dir = str(outputs_dir))
+
+ assert models[0][2]["base_model"] == "org/foo__project-bar"
+
+
+def test_scan_checkpoints_preserves_legacy_folder_name_fallback(tmp_path, monkeypatch):
+ outputs_dir = _make_outputs_dir(tmp_path, monkeypatch)
+ run_dir = outputs_dir / "unsloth_Llama-3.2-3B-Instruct_1771227800"
+ run_dir.mkdir()
+ (run_dir / "config.json").write_text("{}")
+
+ db_path = tmp_path / "studio.db"
+ _setup_training_runs_table(db_path)
+ monkeypatch.setattr(
+ checkpoints_module,
+ "get_connection",
+ lambda: _make_history_connection(db_path),
+ )
+
+ models = checkpoints_module.scan_checkpoints(outputs_dir = str(outputs_dir))
+
+ assert models[0][2]["base_model"] == "unsloth/Llama-3.2-3B-Instruct"
+
+
+def test_scan_checkpoints_prefers_exact_history_match_over_newer_suffix(tmp_path, monkeypatch):
+ outputs_dir = _make_outputs_dir(tmp_path, monkeypatch)
+ run_dir = outputs_dir / "unsloth_Test_1771227800"
+ run_dir.mkdir()
+ (run_dir / "config.json").write_text("{}")
+
+ copied_dir = tmp_path / "copied" / run_dir.name
+ copied_dir.mkdir(parents = True)
+
+ db_path = tmp_path / "studio.db"
+ _setup_training_runs_table(db_path)
+ conn = _make_history_connection(db_path)
+ try:
+ conn.execute(
+ """
+ INSERT INTO training_runs (id, model_name, config_json, output_dir, started_at)
+ VALUES (?, ?, ?, ?, ?)
+ """,
+ (
+ "run-exact",
+ "correct/base",
+ "{}",
+ str(run_dir.resolve()),
+ "2026-04-09T00:00:00Z",
+ ),
+ )
+ conn.execute(
+ """
+ INSERT INTO training_runs (id, model_name, config_json, output_dir, started_at)
+ VALUES (?, ?, ?, ?, ?)
+ """,
+ (
+ "run-suffix",
+ "wrong/base",
+ "{}",
+ str(copied_dir.resolve()),
+ "2026-04-10T00:00:00Z",
+ ),
+ )
+ conn.commit()
+ finally:
+ conn.close()
+
+ monkeypatch.setattr(
+ checkpoints_module,
+ "get_connection",
+ lambda: _make_history_connection(db_path),
+ )
+
+ models = checkpoints_module.scan_checkpoints(outputs_dir = str(outputs_dir))
+
+ assert models[0][2]["base_model"] == "correct/base"
diff --git a/studio/backend/tests/test_cloudflare_tunnel.py b/studio/backend/tests/test_cloudflare_tunnel.py
index 8042240b64..7904c70a7b 100644
--- a/studio/backend/tests/test_cloudflare_tunnel.py
+++ b/studio/backend/tests/test_cloudflare_tunnel.py
@@ -11,6 +11,7 @@ checked by AST so we never import its heavy deps (uvicorn/structlog).
import ast
import importlib.util
import io
+import os
import sys
import tarfile
import types
@@ -136,7 +137,9 @@ def test_ensure_downloads_and_chmods_when_missing(monkeypatch, tmp_path):
path = ct.ensure_cloudflared()
assert path == str(cached)
assert cached.exists()
- assert cached.stat().st_mode & 0o111 # executable bit set
+ # Host OS, not monkeypatched ct.sys.platform.
+ if os.name != "nt":
+ assert cached.stat().st_mode & 0o111
def test_ensure_returns_none_on_download_failure(monkeypatch, tmp_path):
@@ -238,7 +241,8 @@ def test_ensure_macos_extracts_tgz_and_chmods(monkeypatch, tmp_path):
path = ct.ensure_cloudflared()
assert path == str(cached)
assert cached.read_bytes() == b"mach-o"
- assert cached.stat().st_mode & 0o111 # chmod applied on posix
+ if os.name != "nt":
+ assert cached.stat().st_mode & 0o111
assert not cached.with_suffix(".tgz").exists() # temp archive cleaned up
@@ -696,6 +700,28 @@ def test_argparse_cloudflare_default_true():
assert _argparse_default(_RUN_PY.read_text(), "--cloudflare") is True
+def test_verify_global_reachability_marks_private_address_unreachable():
+ src = _RUN_PY.read_text()
+ tree = ast.parse(src)
+ func_src = next(
+ ast.get_source_segment(src, n)
+ for n in ast.walk(tree)
+ if isinstance(n, ast.FunctionDef) and n.name == "_verify_global_reachability"
+ )
+ captured = []
+ ns = {
+ "_public_reachable": None,
+ "_stdout_color_ok": lambda: False,
+ "_url_host": lambda host: host,
+ "print": lambda *a, **k: captured.append(" ".join(str(x) for x in a)),
+ }
+ exec(compile(func_src, "", "exec"), ns)
+ ns["_verify_global_reachability"]("192.168.1.10", 8888)
+
+ assert ns["_public_reachable"] is False
+ assert "private/LAN address" in "\n".join(captured)
+
+
def test_run_server_registers_tunnel_atexit_backstop():
# An abnormal exit (exception after startup -> sys.exit) bypasses
# _graceful_shutdown; an atexit backstop must still stop the tunnel.
@@ -703,16 +729,18 @@ def test_run_server_registers_tunnel_atexit_backstop():
assert "atexit.register(stop_studio_tunnel)" in src
-def test_run_server_gates_tunnel_on_wildcard():
- # Guard against accidentally widening the trigger beyond 0.0.0.0.
- source = _RUN_PY.read_text()
- assert "_cloudflare_enabled" in source
- assert 'host == "0.0.0.0"' in source
-
-
-def _run_print_cloudflare_line(monkeypatch, *, cloudflare_url, public_reachable):
- """Exec the real _print_cloudflare_line source in isolation (run.py has heavy
- deps), with the two module globals injected and startup_banner stubbed."""
+def _run_print_cloudflare_line(
+ monkeypatch,
+ *,
+ cloudflare_url,
+ public_reachable,
+ cloudflare_requested = False,
+ cloudflare_flag = True,
+ secure = False,
+ loopback_host = "127.0.0.1",
+ color = False,
+):
+ """Exec _print_cloudflare_line without importing run.py's heavy deps."""
src = _RUN_PY.read_text()
tree = ast.parse(src)
func_src = next(
@@ -721,16 +749,18 @@ def _run_print_cloudflare_line(monkeypatch, *, cloudflare_url, public_reachable)
if isinstance(n, ast.FunctionDef) and n.name == "_print_cloudflare_line"
)
stub = types.ModuleType("startup_banner")
- stub.stdout_supports_color = lambda: False
+ stub.stdout_supports_color = lambda: color
monkeypatch.setitem(sys.modules, "startup_banner", stub)
captured: list[str] = []
ns = {
"_cloudflare_url": cloudflare_url,
"_public_reachable": public_reachable,
+ "_cloudflare_requested": cloudflare_requested,
+ "_cloudflare_flag": cloudflare_flag,
"print": lambda *a, **k: captured.append(" ".join(str(x) for x in a)),
}
exec(compile(func_src, "", "exec"), ns)
- ns["_print_cloudflare_line"]()
+ ns["_print_cloudflare_line"](secure = secure, loopback_host = loopback_host)
return "\n".join(captured)
@@ -750,7 +780,6 @@ def test_cloudflare_line_default_wording_when_reachable(monkeypatch):
def test_cloudflare_line_default_wording_when_unknown(monkeypatch):
- # Probe did not run / could not decide -> keep the existing wording.
out = _run_print_cloudflare_line(
monkeypatch, cloudflare_url = "https://x.trycloudflare.com", public_reachable = None
)
@@ -758,6 +787,136 @@ def test_cloudflare_line_default_wording_when_unknown(monkeypatch):
assert "Use the secure link" not in out
-def test_cloudflare_line_prints_nothing_without_tunnel(monkeypatch):
+def test_cloudflare_line_states_inactive_when_enabled_but_not_requested(monkeypatch):
out = _run_print_cloudflare_line(monkeypatch, cloudflare_url = None, public_reachable = False)
- assert out == ""
+ assert "Cloudflare tunnel: OFF for this mode" in out
+ assert "local network only" in out
+
+
+def test_cloudflare_line_warns_when_public_url_up(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = "https://x.trycloudflare.com",
+ public_reachable = True,
+ cloudflare_requested = True,
+ )
+ assert "Secure link access via Cloudflare: https://x.trycloudflare.com" in out
+ assert "Cloudflare tunnel: ON" in out
+ assert "PUBLIC" in out
+ assert "--no-cloudflare" in out
+ assert "raw port is also publicly reachable" in out
+ assert "local network only" not in out
+
+
+def test_cloudflare_line_secure_mode_suppresses_public_warning(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = "https://x.trycloudflare.com",
+ public_reachable = True,
+ cloudflare_requested = True,
+ secure = True,
+ )
+ assert "Secure link access via Cloudflare: https://x.trycloudflare.com" in out
+ assert "Cloudflare tunnel: ON" not in out
+
+
+def test_cloudflare_line_states_disabled_when_off(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = False,
+ cloudflare_requested = False,
+ cloudflare_flag = False,
+ )
+ assert "Cloudflare tunnel: OFF" in out
+ assert "local network only" in out
+
+
+def test_cloudflare_line_states_failed_when_requested_but_no_url(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = False,
+ cloudflare_requested = True,
+ cloudflare_flag = True,
+ )
+ assert "requested but failed to start" in out
+ assert "local network only" in out
+
+
+def test_cloudflare_line_off_does_not_claim_local_only_when_unknown(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = None,
+ cloudflare_requested = False,
+ cloudflare_flag = False,
+ )
+ assert "Cloudflare tunnel: OFF" in out
+ assert "Raw port reachability was not verified" in out
+ assert "local network only" not in out
+
+
+def test_cloudflare_line_failed_does_not_claim_local_only_when_unknown(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = None,
+ cloudflare_requested = True,
+ cloudflare_flag = True,
+ )
+ assert "requested but failed to start" in out
+ assert "Raw port reachability was not verified" in out
+ assert "local network only" not in out
+
+
+@pytest.mark.parametrize(
+ "cloudflare_requested,cloudflare_flag,expected",
+ [
+ (True, True, "requested but failed to start"),
+ (False, True, "Cloudflare tunnel: OFF for this mode"),
+ (False, False, "Cloudflare tunnel: OFF"),
+ ],
+)
+def test_cloudflare_line_unknown_warns_with_loopback_host(
+ monkeypatch, cloudflare_requested, cloudflare_flag, expected
+):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = None,
+ cloudflare_requested = cloudflare_requested,
+ cloudflare_flag = cloudflare_flag,
+ loopback_host = "::1",
+ color = True,
+ )
+ assert expected in out
+ assert "bind ::1" in out
+ assert "bind 127.0.0.1" not in out
+ assert "\033[38;5;215;1m" in out
+
+
+def test_cloudflare_line_off_does_not_claim_local_only_when_publicly_reachable(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = True,
+ cloudflare_requested = False,
+ cloudflare_flag = False,
+ )
+ assert "Cloudflare tunnel: OFF" in out
+ assert "reachable from the public internet" in out
+ assert "local network only" not in out
+
+
+def test_cloudflare_line_failed_does_not_claim_local_only_when_publicly_reachable(monkeypatch):
+ out = _run_print_cloudflare_line(
+ monkeypatch,
+ cloudflare_url = None,
+ public_reachable = True,
+ cloudflare_requested = True,
+ cloudflare_flag = True,
+ )
+ assert "requested but failed to start" in out
+ assert "reachable from the public internet" in out
+ assert "local network only" not in out
diff --git a/studio/backend/tests/test_data_recipe_seed.py b/studio/backend/tests/test_data_recipe_seed.py
index 601df8bbfe..09e22116ed 100644
--- a/studio/backend/tests/test_data_recipe_seed.py
+++ b/studio/backend/tests/test_data_recipe_seed.py
@@ -1,12 +1,126 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+import asyncio
+import importlib.util
from pathlib import Path
+import pytest
-def test_seed_inspect_load_kwargs_disables_remote_code_execution():
- seed_route = (
+
+def _seed_route_source() -> str:
+ return (
Path(__file__).resolve().parent.parent / "routes" / "data_recipe" / "seed.py"
).read_text()
- assert '"trust_remote_code": False' in seed_route
+
+def test_seed_inspect_load_kwargs_disables_remote_code_execution():
+ assert '"trust_remote_code": False' in _seed_route_source()
+
+
+class _FakeUpload:
+ def __init__(self, filename: str, content: bytes):
+ self.filename = filename
+ self._content = content
+
+ async def read(self) -> bytes:
+ return self._content
+
+
+def _load_seed_route(monkeypatch: pytest.MonkeyPatch, tmp_path: Path):
+ pytest.importorskip("fastapi")
+ pytest.importorskip("multipart")
+ pytest.importorskip("structlog")
+
+ backend_root = Path(__file__).resolve().parent.parent
+ monkeypatch.syspath_prepend(str(backend_root))
+ route_path = backend_root / "routes" / "data_recipe" / "seed.py"
+ spec = importlib.util.spec_from_file_location("seed_under_test", route_path)
+ assert spec is not None and spec.loader is not None
+ seed_route = importlib.util.module_from_spec(spec)
+ spec.loader.exec_module(seed_route)
+ seed_route.UNSTRUCTURED_UPLOAD_ROOT = tmp_path / "unstructured-uploads"
+ return seed_route
+
+
+def _run_upload(
+ seed_route,
+ filename: str,
+ content: bytes,
+ block_id: str = "block",
+):
+ return asyncio.run(
+ seed_route.upload_unstructured_file(_FakeUpload(filename, content), block_id)
+ )
+
+
+def _block_files(seed_route, block_id: str = "block") -> list[str]:
+ block_dir = seed_route.UNSTRUCTURED_UPLOAD_ROOT / block_id
+ if not block_dir.exists():
+ return []
+ return sorted(path.name for path in block_dir.iterdir())
+
+
+def _raise(exc: BaseException):
+ def raise_exc(*args, **kwargs):
+ raise exc
+
+ return raise_exc
+
+
+@pytest.mark.parametrize(
+ ("filename", "package"),
+ [
+ ("paper.pdf", "pymupdf4llm"),
+ ("notes.docx", "mammoth"),
+ ],
+)
+def test_unstructured_upload_names_missing_extractor_dependency(
+ monkeypatch, tmp_path, filename, package
+):
+ seed_route = _load_seed_route(monkeypatch, tmp_path)
+ monkeypatch.setattr(
+ seed_route,
+ "_extract_text_from_file",
+ _raise(ModuleNotFoundError(f"No module named {package!r}", name = package)),
+ )
+
+ result = _run_upload(seed_route, filename, b"%PDF-1.7")
+
+ assert result.status == "error"
+ assert (
+ result.error
+ == f"Cannot read {Path(filename).suffix} files: the '{package}' package is not installed."
+ )
+ assert _block_files(seed_route) == []
+
+
+def test_unstructured_upload_keeps_txt_path_working(monkeypatch, tmp_path):
+ seed_route = _load_seed_route(monkeypatch, tmp_path)
+
+ result = _run_upload(seed_route, "notes.txt", b"hello")
+
+ assert result.status == "ok"
+ assert result.error is None
+ assert any(name.endswith(".txt") for name in _block_files(seed_route))
+ assert any(name.endswith(".extracted.txt") for name in _block_files(seed_route))
+
+
+@pytest.mark.parametrize(
+ "exc",
+ [
+ ImportError("cannot import internal symbol"),
+ ModuleNotFoundError(
+ "No module named 'missing_transitive_pkg'",
+ name = "missing_transitive_pkg",
+ ),
+ ],
+)
+def test_unstructured_upload_import_errors_stay_generic(monkeypatch, tmp_path, exc):
+ seed_route = _load_seed_route(monkeypatch, tmp_path)
+ monkeypatch.setattr(seed_route, "_extract_text_from_file", _raise(exc))
+ result = _run_upload(seed_route, "paper.pdf", b"%PDF-1.7")
+
+ assert result.status == "error"
+ assert result.error == "Text extraction failed."
+ assert _block_files(seed_route) == []
diff --git a/studio/backend/tests/test_export_imatrix_compressed.py b/studio/backend/tests/test_export_imatrix_compressed.py
new file mode 100644
index 0000000000..d914ff8651
--- /dev/null
+++ b/studio/backend/tests/test_export_imatrix_compressed.py
@@ -0,0 +1,116 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Tests for the GGUF imatrix option and compressed-tensors merged export wiring.
+
+Schema checks use the real Pydantic models; the cross-layer threading is verified with ast so it
+runs on CPU with no GPU, no model, and no llama.cpp.
+"""
+
+import ast
+from pathlib import Path
+
+import pytest
+from pydantic import ValidationError
+
+from models.export import ExportGGUFRequest, ExportMergedModelRequest
+
+_BACKEND = Path(__file__).resolve().parent.parent
+
+
+def _src(rel):
+ return (_BACKEND / rel).read_text(encoding = "utf-8")
+
+
+def _func_src(rel, name):
+ src = _src(rel)
+ node = next(
+ n for n in ast.walk(ast.parse(src)) if isinstance(n, ast.FunctionDef) and n.name == name
+ )
+ return ast.get_source_segment(src, node)
+
+
+# -- schema -------------------------------------------------------------------------------------
+
+
+def test_gguf_request_imatrix_defaults_and_set():
+ assert ExportGGUFRequest(save_directory = "/tmp/x").imatrix is False
+ assert ExportGGUFRequest(save_directory = "/tmp/x").imatrix_path is None
+ r = ExportGGUFRequest(save_directory = "/tmp/x", imatrix = True, imatrix_path = "/i.dat")
+ assert r.imatrix is True and r.imatrix_path == "/i.dat"
+
+
+def test_merged_request_accepts_compressed_formats():
+ for fmt in ("16-bit (FP16)", "FP8 (compressed-tensors)", "NVFP4 (compressed-tensors)"):
+ assert ExportMergedModelRequest(save_directory = "/tmp/x", format_type = fmt).format_type == fmt
+
+
+def test_merged_request_rejects_unknown_format():
+ with pytest.raises(ValidationError):
+ ExportMergedModelRequest(save_directory = "/tmp/x", format_type = "bogus")
+
+
+# -- threading (ast) ----------------------------------------------------------------------------
+
+
+def test_export_gguf_threads_imatrix_to_save_and_push():
+ # imatrix_file must reach both save_pretrained_gguf and push_to_hub_gguf, but only via the
+ # conditional **imatrix_kw so a no-imatrix export never sends an unsupported keyword.
+ g = _func_src("core/export/export.py", "export_gguf")
+ assert g.count("**imatrix_kw") >= 2
+ assert 'imatrix_kw = {"imatrix_file": imatrix_file} if imatrix_file is not None else {}' in g
+ # Unconditional pass-through (the old wiring) must be gone.
+ assert "imatrix_file = imatrix_file" not in g
+
+
+def test_export_gguf_guards_unsupported_imatrix_build():
+ # An older unsloth without imatrix_file support gets a clean error, not a TypeError.
+ g = _func_src("core/export/export.py", "export_gguf")
+ assert "_supports_kwarg(" in g and '"imatrix_file"' in g
+
+
+def test_export_merged_guards_unsupported_compressed_build():
+ m = _func_src("core/export/export.py", "export_merged_model")
+ assert "_compressed_export_supported()" in m
+
+
+def test_supports_kwarg_helper():
+ # exec just the helper source so the test stays free of export.py's heavy import chain.
+ ns = {}
+ exec(_func_src("core/export/export.py", "_supports_kwarg"), ns)
+ supports = ns["_supports_kwarg"]
+
+ def has_it(a, imatrix_file = None):
+ pass
+
+ def lacks_it(a):
+ pass
+
+ def via_kwargs(a, **kw):
+ pass
+
+ assert supports(has_it, "imatrix_file") is True
+ assert supports(lacks_it, "imatrix_file") is False
+ assert supports(via_kwargs, "imatrix_file") is True
+
+
+def test_orchestrator_and_worker_pass_imatrix():
+ assert "imatrix_file" in _func_src("core/export/orchestrator.py", "export_gguf")
+ assert 'imatrix_file = cmd.get("imatrix_file")' in _src("core/export/worker.py")
+
+
+def test_route_resolves_imatrix_file():
+ assert "request.imatrix_path or (True if request.imatrix else None)" in _src("routes/export.py")
+
+
+def test_export_merged_maps_compressed_to_save_method():
+ m = _func_src("core/export/export.py", "export_merged_model")
+ assert "is_compressed" in m and '"fp8"' in m and '"nvfp4"' in m
+
+
+def test_compressed_hub_push_uploads_local_dir_without_recompressing():
+ # A compressed Hub push must upload the already-built output_path, not re-run compression
+ # via push_to_hub_merged (which would compress a second time).
+ m = _func_src("core/export/export.py", "export_merged_model")
+ assert "elif is_compressed and output_path and Path(output_path).is_dir():" in m
+ assert "hf_api.upload_folder(" in m and "folder_path = output_path" in m
diff --git a/studio/backend/tests/test_hf_xet_fallback.py b/studio/backend/tests/test_hf_xet_fallback.py
index 39ecebd328..9e40fbf508 100644
--- a/studio/backend/tests/test_hf_xet_fallback.py
+++ b/studio/backend/tests/test_hf_xet_fallback.py
@@ -206,6 +206,7 @@ class _FakeAttempt:
interval,
grace_period,
on_status,
+ force_download = False,
):
self.calls.append(
_types.SimpleNamespace(
diff --git a/studio/backend/tests/test_inference_dispatcher_resilience.py b/studio/backend/tests/test_inference_dispatcher_resilience.py
index 2de833db0e..60d7503485 100644
--- a/studio/backend/tests/test_inference_dispatcher_resilience.py
+++ b/studio/backend/tests/test_inference_dispatcher_resilience.py
@@ -10,6 +10,7 @@ must be logged and skipped, not fatal. Fakes only.
from __future__ import annotations
+import ast
import queue
import sys
import threading
@@ -89,3 +90,31 @@ def test_dispatcher_survives_mailbox_put_error():
o._dispatcher_stop.set()
t.join(timeout = 5)
assert not t.is_alive()
+
+
+def test_route_llama_streaming_async_clients_disable_proxy_env():
+ """Local llama-server streaming proxies must ignore ambient HTTP_PROXY."""
+ source = (Path(__file__).resolve().parent.parent / "routes" / "inference.py").read_text(
+ encoding = "utf-8"
+ )
+ tree = ast.parse(source)
+ calls = []
+ for node in ast.walk(tree):
+ if not isinstance(node, ast.Call):
+ continue
+ func = node.func
+ if not (
+ isinstance(func, ast.Attribute)
+ and func.attr == "AsyncClient"
+ and isinstance(func.value, ast.Name)
+ and func.value.id == "httpx"
+ ):
+ continue
+ calls.append(node)
+
+ assert len(calls) == 4
+ for call in calls:
+ assert any(
+ kw.arg == "trust_env" and isinstance(kw.value, ast.Constant) and kw.value.value is False
+ for kw in call.keywords
+ ), f"httpx.AsyncClient at line {call.lineno} must set trust_env=False"
diff --git a/studio/backend/tests/test_llama_cpp_props_readback.py b/studio/backend/tests/test_llama_cpp_props_readback.py
index d87c05f2c6..316956325f 100644
--- a/studio/backend/tests/test_llama_cpp_props_readback.py
+++ b/studio/backend/tests/test_llama_cpp_props_readback.py
@@ -113,8 +113,14 @@ def _stub_props(
body = None,
exc = None,
):
- def fake_get(url, timeout = None):
+ def fake_get(
+ url,
+ timeout = None,
+ trust_env = None,
+ ):
assert url.endswith("/props")
+
+ assert trust_env is False
if exc is not None:
raise exc
return _FakeResponse(status_code, body)
diff --git a/studio/backend/tests/test_model_update_robustness.py b/studio/backend/tests/test_model_update_robustness.py
new file mode 100644
index 0000000000..9cf2a62c39
--- /dev/null
+++ b/studio/backend/tests/test_model_update_robustness.py
@@ -0,0 +1,483 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+"""Tests for model-update detection and the GGUF force-download helper.
+
+Covers:
+ * GGUF variant listing computes update_available from the already-fetched
+ sibling metadata instead of a second Hub call.
+ * hf_hub_download_with_xet_fallback(force_download=True) bypasses the
+ try_to_load_from_cache cache-first early-return.
+
+The cache "Update" action now runs through the download manager as a normal
+managed download (so it shows in the Downloads panel with progress + cancel),
+so the old POST /api/models/update endpoint and its tests are gone. Update
+*detection* — the "Update available" cue — is still exercised here.
+"""
+
+import asyncio
+import sys
+import types
+from types import SimpleNamespace
+
+if "structlog" not in sys.modules:
+
+ class _DummyLogger:
+ def __getattr__(self, _name):
+ return lambda *a, **k: None
+
+ sys.modules["structlog"] = types.SimpleNamespace(
+ BoundLogger = _DummyLogger, get_logger = lambda *a, **k: _DummyLogger()
+ )
+
+import pytest
+from hub.services.models import cache_inventory as CI
+from hub.services.models import deletion as D
+from hub.services.models import gguf_variants as GV
+
+
+def _variants():
+ return [
+ SimpleNamespace(
+ filename = "model-Q4_K_M.gguf",
+ quant = "Q4_K_M",
+ display_label = None,
+ size_bytes = 1000,
+ ),
+ SimpleNamespace(
+ filename = "model-Q8_0.gguf",
+ quant = "Q8_0",
+ display_label = None,
+ size_bytes = 2000,
+ ),
+ ]
+
+
+def _seed_cache(tmp_path, repo_id, blob_ids, gguf_files):
+ repo = tmp_path / f"models--{repo_id.replace('/', '--')}"
+ snap = repo / "snapshots" / ("a" * 40)
+ snap.mkdir(parents = True, exist_ok = True)
+ for name, size in gguf_files.items():
+ (snap / name).write_bytes(b"\0" * size)
+ blobs = repo / "blobs"
+ blobs.mkdir(exist_ok = True)
+ for b in blob_ids:
+ (blobs / b).write_bytes(b"x")
+ return repo, snap, blobs
+
+
+@pytest.fixture
+def patch_hub_gguf(monkeypatch):
+ """Patch GGUF listing and cache scans for sibling-derived update checks."""
+
+ def _sibling(
+ path: str,
+ size: int,
+ sha = None,
+ *,
+ lfs_dict = False,
+ blob_id = None,
+ ):
+ if lfs_dict:
+ lfs = {"sha256": sha} if sha else {}
+ else:
+ lfs = SimpleNamespace(sha256 = sha) if sha else None
+ return SimpleNamespace(rfilename = path, size = size, lfs = lfs, blob_id = blob_id)
+
+ def _repo_info(repo_id: str, repo_path, files: list[tuple[str, str]]):
+ return SimpleNamespace(
+ repo_id = repo_id,
+ repo_type = "model",
+ repo_path = repo_path,
+ revisions = [
+ SimpleNamespace(
+ files = [
+ SimpleNamespace(
+ file_name = name,
+ blob_path = str(repo_path / "blobs" / blob),
+ )
+ for name, blob in files
+ ]
+ )
+ ],
+ )
+
+ def _apply(tmp_path, repo_id: str, *, local_blob: str, remote_sibling):
+ with GV._VARIANT_HASH_LOCK:
+ GV._VARIANT_HASH_CACHE.clear()
+ GV._VARIANT_REQUIREMENT_CACHE.clear()
+ GV._VARIANT_REQUIREMENT_NEG_CACHE.clear()
+ repo, snap, _blobs = _seed_cache(
+ tmp_path,
+ repo_id,
+ blob_ids = [local_blob],
+ gguf_files = {"model-Q4_K_M.gguf": 1000},
+ )
+ monkeypatch.setattr(
+ GV,
+ "list_gguf_variants",
+ lambda r, hf_token = None: (_variants(), False, [remote_sibling]),
+ raising = True,
+ )
+ monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap])
+ monkeypatch.setattr(
+ CI,
+ "all_hf_cache_scans",
+ lambda: [
+ SimpleNamespace(
+ repos = [
+ _repo_info(
+ repo_id,
+ repo,
+ [("model-Q4_K_M.gguf", local_blob)],
+ )
+ ]
+ )
+ ],
+ )
+
+ return SimpleNamespace(apply = _apply, sibling = _sibling)
+
+
+def _call(coro):
+ loop = asyncio.new_event_loop()
+ try:
+ return loop.run_until_complete(coro)
+ finally:
+ loop.close()
+
+
+# ── GGUF variant update detection ───────────────────────────────
+
+
+def test_variant_update_check_missing_remote_blob_id_is_not_phantom_update(
+ tmp_path, patch_hub_gguf
+):
+ """Missing sha/blob metadata is unknown, not update_available=True."""
+ repo = "unsloth/gemma-3-4b-it-GGUF"
+ patch_hub_gguf.apply(
+ tmp_path,
+ repo,
+ local_blob = "oldsha",
+ remote_sibling = patch_hub_gguf.sibling("model-Q4_K_M.gguf", 1000, None),
+ )
+ resp = _call(GV.get_gguf_variants_response(repo))
+ assert len(resp.variants) == 2
+ q4 = next(v for v in resp.variants if v.quant == "Q4_K_M")
+ assert q4.downloaded is True
+ assert q4.update_available is False
+
+
+def test_variant_update_check_detects_update_from_existing_siblings(tmp_path, patch_hub_gguf):
+ repo = "unsloth/gemma-3-4b-it-GGUF"
+ patch_hub_gguf.apply(
+ tmp_path,
+ repo,
+ local_blob = "oldsha",
+ remote_sibling = patch_hub_gguf.sibling("model-Q4_K_M.gguf", 1000, "NEWsha"),
+ )
+ resp = _call(GV.get_gguf_variants_response(repo))
+ q4 = next(v for v in resp.variants if v.quant == "Q4_K_M")
+ assert q4.update_available is True
+
+
+def test_variant_update_check_no_update_when_blob_matches(tmp_path, patch_hub_gguf):
+ repo = "unsloth/gemma-3-4b-it-GGUF"
+ patch_hub_gguf.apply(
+ tmp_path,
+ repo,
+ local_blob = "samesha",
+ remote_sibling = patch_hub_gguf.sibling("model-Q4_K_M.gguf", 1000, "samesha"),
+ )
+ resp = _call(GV.get_gguf_variants_response(repo))
+ q4 = next(v for v in resp.variants if v.quant == "Q4_K_M")
+ assert q4.update_available is False
+
+
+@pytest.mark.parametrize(
+ ("companion_path", "has_vision"),
+ [
+ ("mmproj-F16.gguf", True),
+ ("mtp-drafter-Q8_0.gguf", False),
+ ],
+)
+def test_variant_update_check_detects_companion_only_update(
+ monkeypatch, tmp_path, patch_hub_gguf, companion_path, has_vision
+):
+ repo_id = "unsloth/gemma-4-GGUF"
+ with GV._VARIANT_HASH_LOCK:
+ GV._VARIANT_HASH_CACHE.clear()
+ GV._VARIANT_REQUIREMENT_CACHE.clear()
+ GV._VARIANT_REQUIREMENT_NEG_CACHE.clear()
+ repo, snap, _blobs = _seed_cache(
+ tmp_path,
+ repo_id,
+ blob_ids = ["mainsha", "old-companion"],
+ gguf_files = {
+ "model-Q4_K_M.gguf": 1000,
+ companion_path: 100,
+ },
+ )
+ siblings = [
+ patch_hub_gguf.sibling("model-Q4_K_M.gguf", 1000, "mainsha"),
+ patch_hub_gguf.sibling(companion_path, 100, "new-companion"),
+ ]
+ monkeypatch.setattr(
+ GV,
+ "list_gguf_variants",
+ lambda r, hf_token = None: (_variants(), has_vision, siblings),
+ raising = True,
+ )
+ monkeypatch.setattr(GV, "iter_hf_cache_snapshots", lambda _repo_id: [snap])
+ monkeypatch.setattr(
+ CI,
+ "all_hf_cache_scans",
+ lambda: [
+ SimpleNamespace(
+ repos = [
+ SimpleNamespace(
+ repo_id = repo_id,
+ repo_type = "model",
+ repo_path = repo,
+ revisions = [
+ SimpleNamespace(
+ files = [
+ SimpleNamespace(
+ file_name = "model-Q4_K_M.gguf",
+ blob_path = str(repo / "blobs" / "mainsha"),
+ ),
+ SimpleNamespace(
+ file_name = companion_path,
+ blob_path = str(repo / "blobs" / "old-companion"),
+ ),
+ ]
+ )
+ ],
+ )
+ ]
+ )
+ ],
+ )
+
+ resp = _call(GV.get_gguf_variants_response(repo_id))
+ q4 = next(v for v in resp.variants if v.quant == "Q4_K_M")
+
+ assert q4.downloaded is True
+ assert q4.update_available is True
+
+
+def test_variant_update_check_accepts_lfs_dict_and_blob_id_fallback(tmp_path, patch_hub_gguf):
+ repo = "unsloth/gemma-3-4b-it-GGUF"
+ patch_hub_gguf.apply(
+ tmp_path,
+ repo,
+ local_blob = "dictsha",
+ remote_sibling = patch_hub_gguf.sibling(
+ "model-Q4_K_M.gguf",
+ 1000,
+ "dictsha",
+ lfs_dict = True,
+ ),
+ )
+ resp = _call(GV.get_gguf_variants_response(repo))
+ assert next(v for v in resp.variants if v.quant == "Q4_K_M").update_available is False
+
+ patch_hub_gguf.apply(
+ tmp_path,
+ repo,
+ local_blob = "blobid",
+ remote_sibling = patch_hub_gguf.sibling(
+ "model-Q4_K_M.gguf",
+ 1000,
+ None,
+ blob_id = "blobid",
+ ),
+ )
+ resp = _call(GV.get_gguf_variants_response(repo))
+ assert next(v for v in resp.variants if v.quant == "Q4_K_M").update_available is False
+
+
+def test_cached_model_scan_keeps_local_safetensors_repo(monkeypatch, tmp_path):
+ repo_path = tmp_path / "models--Org--SafeTensorRepo"
+ repo = SimpleNamespace(
+ repo_id = "Org/SafeTensorRepo",
+ repo_type = "model",
+ repo_path = repo_path,
+ revisions = [
+ SimpleNamespace(
+ files = [
+ SimpleNamespace(
+ file_name = "config.json",
+ size_on_disk = 10,
+ blob_path = None,
+ ),
+ SimpleNamespace(
+ file_name = "model.safetensors",
+ size_on_disk = 100,
+ blob_path = str(repo_path / "blobs" / "modelsha"),
+ ),
+ ]
+ )
+ ],
+ )
+ monkeypatch.setattr(
+ CI,
+ "all_hf_cache_scans",
+ lambda: [SimpleNamespace(repos = [repo])],
+ )
+ monkeypatch.setattr(
+ CI.hf_cache_scan,
+ "is_snapshot_partial",
+ lambda *args, **kwargs: False,
+ )
+
+ rows = CI._scan_cached_models()
+
+ assert len(rows) == 1
+ assert rows[0]["repo_id"] == "Org/SafeTensorRepo"
+ assert rows[0]["model_format"] == "safetensors"
+ assert rows[0]["size_bytes"] == 100
+
+
+# ── hf_hub_download_with_xet_fallback force_download bypass (X2/F2) ───
+
+
+def test_force_download_bypasses_cache_first_early_return(monkeypatch):
+ """force_download=True skips the try_to_load_from_cache early-return and
+ proceeds to the real download path; force_download=False returns the cached
+ path without ever attempting a download (X2/F2)."""
+ import huggingface_hub as hf
+ import utils.hf_xet_fallback as X
+
+ cached_path = "/cache/blob/cached.gguf"
+
+ # Pretend the blob IS cached on disk (try_to_load_from_cache is imported
+ # inside the function from huggingface_hub, and os.path.exists must agree).
+ monkeypatch.setattr(hf, "try_to_load_from_cache", lambda *a, **k: cached_path, raising = False)
+ monkeypatch.setattr(X.os.path, "exists", lambda p: True, raising = False)
+
+ attempts = []
+
+ def fake_attempt(repo_id, filename, token, **kwargs):
+ attempts.append(
+ {"repo_id": repo_id, "filename": filename, "force": kwargs.get("force_download")}
+ )
+ return ("ok", "/freshly/downloaded/path")
+
+ monkeypatch.setattr(X, "_run_download_attempt", fake_attempt, raising = True)
+
+ # force_download=False: cache-first early-return, no download attempt.
+ out = X.hf_hub_download_with_xet_fallback(
+ "unsloth/repo", "model.gguf", token = None, force_download = False
+ )
+ assert out == cached_path
+ assert attempts == [] # never reached the real download
+
+ # force_download=True: bypass the early-return, run the real download.
+ out2 = X.hf_hub_download_with_xet_fallback(
+ "unsloth/repo", "model.gguf", token = None, force_download = True
+ )
+ assert out2 == "/freshly/downloaded/path"
+ assert len(attempts) == 1
+ assert attempts[0]["force"] is True
+
+
+# ── multi-revision GGUF blob comparison and update reclaim ──
+#
+# Regression for the phantom "Update available" cue that lingered AFTER a model
+# was already updated. A re-download leaves BOTH the old and new revision
+# snapshots in the HF cache, so the same gguf file resolves to several blobs.
+# The local collection must keep ALL of them (a set per file), and stale hashes
+# must be pruned only after the replacement revision verifies.
+
+
+def _rev(*files):
+ return SimpleNamespace(
+ files = [SimpleNamespace(file_name = name, blob_path = f"/blobs/{blob}") for name, blob in files]
+ )
+
+
+def test_repo_gguf_blob_map_collects_all_revision_blobs():
+ """Every cached revision's blob for a gguf file is kept as a set, not
+ collapsed to one arbitrary blob."""
+ repo_info = SimpleNamespace(
+ revisions = [
+ _rev(("lfm2-350m-q4_k_m.gguf", "OLDsha")),
+ _rev(("lfm2-350m-q4_k_m.gguf", "NEWsha")),
+ ]
+ )
+ assert CI._repo_gguf_blob_map(repo_info) == {"lfm2-350m-q4_k_m.gguf": {"OLDsha", "NEWsha"}}
+
+
+def test_reclaim_replaced_gguf_variant_prunes_old_revision_only(monkeypatch, tmp_path):
+ """After a verified update, stale same-variant files/blobs are removed while
+ the freshly downloaded hash and sibling variants remain cached."""
+ repo_id = "org/repo-GGUF"
+ repo_path = tmp_path / "models--org--repo-GGUF"
+ old_snap = repo_path / "snapshots" / ("a" * 40) / "model-Q4_K_M.gguf"
+ new_snap = repo_path / "snapshots" / ("b" * 40) / "model-Q4_K_M.gguf"
+ sibling_snap = repo_path / "snapshots" / ("b" * 40) / "model-Q8_0.gguf"
+ old_blob = repo_path / "blobs" / "OLDsha"
+ new_blob = repo_path / "blobs" / "NEWsha"
+ sibling_blob = repo_path / "blobs" / "Q8sha"
+ for path, payload in (
+ (old_snap, b"old"),
+ (new_snap, b"new"),
+ (sibling_snap, b"sibling"),
+ (old_blob, b"old-blob"),
+ (new_blob, b"new-blob"),
+ (sibling_blob, b"sibling-blob"),
+ ):
+ path.parent.mkdir(parents = True, exist_ok = True)
+ path.write_bytes(payload)
+
+ repo_info = SimpleNamespace(
+ repo_id = repo_id,
+ repo_type = "model",
+ repo_path = repo_path,
+ revisions = [
+ SimpleNamespace(
+ files = [
+ SimpleNamespace(
+ file_name = "model-Q4_K_M.gguf",
+ file_path = str(old_snap),
+ blob_path = str(old_blob),
+ )
+ ]
+ ),
+ SimpleNamespace(
+ files = [
+ SimpleNamespace(
+ file_name = "model-Q4_K_M.gguf",
+ file_path = str(new_snap),
+ blob_path = str(new_blob),
+ ),
+ SimpleNamespace(
+ file_name = "model-Q8_0.gguf",
+ file_path = str(sibling_snap),
+ blob_path = str(sibling_blob),
+ ),
+ ]
+ ),
+ ],
+ )
+ monkeypatch.setattr(
+ CI,
+ "all_hf_cache_scans",
+ lambda: [SimpleNamespace(repos = [repo_info])],
+ )
+ invalidated = []
+ monkeypatch.setattr(CI, "invalidate_hf_cache_scans", lambda: invalidated.append(True))
+
+ result = D.reclaim_replaced_gguf_variant(repo_id, "Q4_K_M", frozenset({"NEWsha"}))
+
+ assert result["removed_snapshots"] == 1
+ assert result["deleted_blobs"] == 1
+ assert result["removed_dirs"] == 1
+ assert old_snap.exists() is False
+ assert old_snap.parent.exists() is False
+ assert old_blob.exists() is False
+ assert new_snap.exists() is True
+ assert new_blob.exists() is True
+ assert sibling_snap.exists() is True
+ assert sibling_blob.exists() is True
+ assert invalidated == [True]
diff --git a/studio/backend/tests/test_offline_gguf_cache_fallback.py b/studio/backend/tests/test_offline_gguf_cache_fallback.py
index a2a505f479..4499881c4d 100644
--- a/studio/backend/tests/test_offline_gguf_cache_fallback.py
+++ b/studio/backend/tests/test_offline_gguf_cache_fallback.py
@@ -951,7 +951,11 @@ class TestWaitForHealthRetriesOnReadError:
calls = {"n": 0}
- def fake_get(url, timeout = None):
+ def fake_get(
+ url,
+ timeout = None,
+ trust_env = None,
+ ):
calls["n"] += 1
if calls["n"] == 1:
raise httpx.ReadError("WinError 10054")
diff --git a/studio/backend/tests/test_rag_captioning.py b/studio/backend/tests/test_rag_captioning.py
index 5d83a7d38d..f475c9e374 100644
--- a/studio/backend/tests/test_rag_captioning.py
+++ b/studio/backend/tests/test_rag_captioning.py
@@ -13,13 +13,15 @@ def _img(page):
return ParsedImage(image_bytes = b"\x89PNG fake", page_number = page, xref = page)
-def test_caption_images_disabled_by_default(monkeypatch):
+def test_caption_images_runs_when_images_present(monkeypatch):
+ # Policy lives in ingestion (_run); caption_images captions given images + endpoint.
monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", False)
- assert captioner.caption_images([_img(1)], endpoint = ("http://x", "local")) == {}
+ monkeypatch.setattr(captioner, "_caption_one", lambda *a: "a chart")
+ out = captioner.caption_images([_img(1)], endpoint = ("http://x", "local"))
+ assert out == {1: ["a chart"]}
def test_caption_images_groups_by_page(monkeypatch):
- monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
monkeypatch.setattr(captioner.config, "CAPTION_MAX_IMAGES", 8)
monkeypatch.setattr(captioner, "_caption_one", lambda base, model, b, t: "a chart of results")
out = captioner.caption_images([_img(1), _img(1), _img(3)], endpoint = ("http://x", "local"))
@@ -27,7 +29,6 @@ def test_caption_images_groups_by_page(monkeypatch):
def test_caption_images_respects_cap(monkeypatch):
- monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
monkeypatch.setattr(captioner.config, "CAPTION_MAX_IMAGES", 2)
calls = []
monkeypatch.setattr(captioner, "_caption_one", lambda *a: (calls.append(1) or "cap"))
@@ -36,11 +37,183 @@ def test_caption_images_respects_cap(monkeypatch):
def test_caption_images_no_endpoint(monkeypatch):
- monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
monkeypatch.setattr(captioner, "vision_endpoint", lambda: None)
assert captioner.caption_images([_img(1)]) == {}
+def test_caption_runaway_guard_applied(monkeypatch):
+ # A looping vision model must not flood the index; captions pass _collapse_runaway.
+ monkeypatch.setattr(captioner, "_caption_one", lambda *a: "\n".join(["LOOP"] * 40))
+ out = captioner.caption_images([_img(1)], endpoint = ("http://x", "local"))
+ assert out[1][0].splitlines().count("LOOP") == 3 # 40 -> 3
+
+
+def test_caption_prompt_and_token_budget(monkeypatch):
+ # Caption and OCR keep separate prompts + token caps over the shared _vision_complete.
+ captured: dict = {}
+
+ def fake_vision_complete(base_url, model, image_bytes, *, prompt, timeout, max_tokens):
+ captured.update(prompt = prompt, timeout = timeout, max_tokens = max_tokens)
+ return "ok"
+
+ monkeypatch.setattr(captioner, "_vision_complete", fake_vision_complete)
+ monkeypatch.setattr(captioner.config, "CAPTION_MAX_TOKENS", 277)
+
+ captioner._caption_one("http://x", "local", b"img", 12.0)
+ prompt = captured["prompt"].lower()
+ # Unified prompt: transcribe every label (recall) + axis/legend coverage + describe.
+ assert "transcribe" in prompt
+ assert ("axis" in prompt or "axes" in prompt) and "legend" in prompt
+ assert "do not invent" in prompt
+ assert captured["max_tokens"] == 277
+ assert captured["timeout"] == 12.0
+
+ captured.clear()
+ monkeypatch.setattr(captioner.config, "OCR_MAX_TOKENS", 999)
+ captioner._ocr_one("http://x", "local", b"img", 5.0)
+ assert captured["max_tokens"] == 999
+ assert "transcribe" in captured["prompt"].lower()
+
+
+def test_pages_with_figures_and_tiles(tmp_path):
+ from core.rag import parsers
+
+ pdf = tmp_path / "fig.pdf"
+ _figure_pdf(pdf)
+ pgs = parsers.pages_with_figures(str(pdf), max_pages = 4)
+ assert pgs == [1]
+ tiles = parsers.render_pdf_figure_tiles(str(pdf), pgs, rows = 2, cols = 2, fullpage = True)
+ assert len(tiles) == 5 # full page + 2x2 grid
+ assert all(t.image_bytes[:8] == b"\x89PNG\r\n\x1a\n" and t.page_number == 1 for t in tiles)
+ capped = parsers.render_pdf_figure_tiles(
+ str(pdf), pgs, rows = 2, cols = 2, fullpage = True, max_tiles = 3
+ )
+ assert len(capped) == 3 # max_tiles budget honored
+
+
+def test_render_pdf_figure_tiles_zero_grid_no_crash(tmp_path):
+ # A misconfigured rows/cols=0 must clamp to 1, not raise ZeroDivisionError.
+ import pymupdf
+
+ from core.rag import parsers
+
+ pdf = tmp_path / "blank.pdf"
+ doc = pymupdf.open()
+ doc.new_page()
+ doc.save(str(pdf))
+ doc.close()
+
+ out = parsers.render_pdf_figure_tiles(str(pdf), [1], rows = 0, cols = 0, fullpage = True)
+ assert len(out) == 2 # full page + a single 1x1 tile, no crash
+
+
+def test_pages_with_figures_excludes_given_pages(tmp_path):
+ # Pages OCR already transcribed (passed as exclude_pages) are skipped; every other
+ # figure page is still returned for tiling.
+ import pymupdf
+
+ from core.rag import parsers
+
+ def _draw_chart(page):
+ shape = page.new_shape()
+ shape.draw_rect(pymupdf.Rect(60, 140, 540, 520))
+ for i in range(8):
+ shape.draw_line((80, 160 + i * 40), (520, 160 + i * 40))
+ shape.finish(color = (0, 0, 0), fill = (0.8, 0.8, 0.9))
+ shape.commit()
+
+ pdf = tmp_path / "charts.pdf"
+ doc = pymupdf.open()
+ _draw_chart(doc.new_page())
+ _draw_chart(doc.new_page())
+ doc.save(str(pdf))
+ doc.close()
+
+ assert parsers.pages_with_figures(str(pdf), max_pages = 4) == [1, 2]
+ assert parsers.pages_with_figures(str(pdf), max_pages = 4, exclude_pages = {1}) == [2]
+ assert parsers.pages_with_figures(str(pdf), max_pages = 4, exclude_pages = {2}) == [1]
+
+
+def test_run_skips_figure_work_without_vision_model(
+ rag_conn, stub_embeddings, monkeypatch, tmp_path
+):
+ # No vision model -> the whole figure pass (detection + rasterization) is skipped.
+ from core.rag import parsers
+
+ monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: None)
+ touched: list[str] = []
+ monkeypatch.setattr(
+ parsers, "pages_with_figures", lambda *a, **k: touched.append("detect") or []
+ )
+ monkeypatch.setattr(
+ parsers, "render_pdf_figure_tiles", lambda *a, **k: touched.append("render") or []
+ )
+
+ pdf = tmp_path / "fig.pdf"
+ _figure_pdf(pdf)
+ _ingest_with_caption(rag_conn, "t1", pdf, None) # follow config (ON), but no model
+ assert touched == [] # neither figure detection nor tiling ran
+
+
+def test_vision_complete_sends_auth_header(monkeypatch):
+ # Direct-stream serves llama-server with --api-key; vision calls must send the bearer.
+ import httpx
+
+ monkeypatch.setattr(
+ captioner, "_vision_auth_headers", lambda: {"Authorization": "Bearer secret"}
+ )
+ captured: dict = {}
+
+ class _Resp:
+ def raise_for_status(self):
+ pass
+
+ def json(self):
+ return {"choices": [{"message": {"content": "ok"}}]}
+
+ def fake_post(url, *, json, timeout, headers):
+ captured.update(url = url, headers = headers)
+ return _Resp()
+
+ monkeypatch.setattr(httpx, "post", fake_post)
+ out = captioner._vision_complete(
+ "http://x", "local", b"img", prompt = "p", timeout = 5.0, max_tokens = 8
+ )
+ assert out == "ok"
+ assert captured["headers"] == {"Authorization": "Bearer secret"}
+
+
+def test_vision_complete_omits_header_when_unauthenticated(monkeypatch):
+ # No api-key configured -> no spurious Authorization header on plain llama-server.
+ import httpx
+
+ monkeypatch.setattr(captioner, "_vision_auth_headers", lambda: None)
+ captured: dict = {}
+
+ class _Resp:
+ def raise_for_status(self):
+ pass
+
+ def json(self):
+ return {"choices": [{"message": {"content": "ok"}}]}
+
+ def fake_post(url, *, json, timeout, headers):
+ captured["headers"] = headers
+ return _Resp()
+
+ monkeypatch.setattr(httpx, "post", fake_post)
+ captioner._vision_complete("http://x", "local", b"i", prompt = "p", timeout = 5.0, max_tokens = 8)
+ assert captured["headers"] is None
+
+
+def test_merge_page_captions_dedups():
+ out = captioner.merge_page_captions({1: ["MatMul\nScale", "Scale\nSoftMax"]})
+ text = out[1][0]
+ assert text.lower().count("scale") == 1 # repeated label from overlapping tiles dropped
+ assert "MatMul" in text and "SoftMax" in text
+
+
def test_splice_captions_appends_to_right_page():
pages = [Page("body one", 1, 8), Page("body two", 2, 8)]
out = captioner.splice_captions(pages, {2: ["a diagram of X"]})
@@ -55,29 +228,6 @@ def test_splice_captions_noop_when_empty():
assert captioner.splice_captions(pages, {}) is pages
-def test_render_pdf_figures_detects_drawing(tmp_path):
- import pymupdf
-
- from core.rag.parsers import render_pdf_figures
-
- pdf = tmp_path / "fig.pdf"
- doc = pymupdf.open()
- page = doc.new_page()
- shape = page.new_shape()
- shape.draw_rect(pymupdf.Rect(60, 60, 540, 460))
- for i in range(8):
- shape.draw_line((80, 80 + i * 40), (520, 80 + i * 40))
- shape.finish(color = (0, 0, 0), fill = (0.8, 0.8, 0.9))
- shape.commit()
- doc.save(str(pdf))
- doc.close()
-
- figs = render_pdf_figures(str(pdf))
- assert figs, "expected at least one rendered figure region"
- assert figs[0].image_bytes[:8] == b"\x89PNG\r\n\x1a\n"
- assert figs[0].page_number == 1
-
-
def test_captioned_text_is_searchable(rag_home, stub_embeddings, monkeypatch):
from core.rag import retrieval, store
from storage import rag_db
@@ -103,3 +253,100 @@ def test_captioned_text_is_searchable(rag_home, stub_embeddings, monkeypatch):
finally:
conn.close()
assert hits, "spliced caption text should be retrievable via lexical search"
+
+
+# ── per-upload caption override (parallels test_rag_ocr_fallback.py) ──
+
+
+def _figure_pdf(path):
+ """A born-digital PDF: a page with real text (so it is not treated as scanned)
+ plus a vector drawing region that figure detection picks up as a figure."""
+ import pymupdf
+
+ doc = pymupdf.open()
+ page = doc.new_page()
+ page.insert_textbox(
+ pymupdf.Rect(40, 40, 550, 120),
+ "Quarterly revenue report. The chart below shows the trend.",
+ fontsize = 11,
+ )
+ shape = page.new_shape()
+ shape.draw_rect(pymupdf.Rect(60, 140, 540, 520))
+ for i in range(8):
+ shape.draw_line((80, 160 + i * 40), (520, 160 + i * 40))
+ shape.finish(color = (0, 0, 0), fill = (0.8, 0.8, 0.9))
+ shape.commit()
+ doc.save(str(path))
+ doc.close()
+
+
+def _ingest_with_caption(rag_conn, thread_id, path, caption):
+ from core.rag import ingestion, store
+
+ scope = store.thread_scope(thread_id)
+ document_id = store.create_document(
+ rag_conn,
+ scope = scope,
+ filename = "fig.pdf",
+ sha256 = str(path) + str(caption),
+ thread_id = thread_id,
+ status = "pending",
+ stored_path = str(path),
+ )
+ job_id = ingestion._new_job(rag_conn, document_id, scope)
+ # _run(job_id, document_id, scope, stored_path, model_name, ocr, caption)
+ ingestion._run(job_id, document_id, scope, str(path), None, None, caption)
+ return store.get_document(rag_conn, document_id)
+
+
+def test_caption_override_true_runs_when_config_off(
+ rag_conn, stub_embeddings, monkeypatch, tmp_path
+):
+ # Config default OFF, but the per-upload toggle (caption=True) forces captioning.
+ from core.rag import tool
+
+ monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", False)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ monkeypatch.setattr(captioner, "_caption_one", lambda *a: "bar chart of revenue wombat-7")
+
+ pdf = tmp_path / "fig.pdf"
+ _figure_pdf(pdf)
+ _ingest_with_caption(rag_conn, "t1", pdf, True)
+
+ text, _ = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "wombat-7" in text # the spliced figure caption reached the index
+
+
+def test_caption_override_false_skips_when_config_on(
+ rag_conn, stub_embeddings, monkeypatch, tmp_path
+):
+ # Config default ON, but the per-upload toggle (caption=False) skips captioning.
+ monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ called = []
+ monkeypatch.setattr(captioner, "_caption_one", lambda *a: called.append(1) or "should not run")
+
+ pdf = tmp_path / "fig.pdf"
+ _figure_pdf(pdf)
+ _ingest_with_caption(rag_conn, "t1", pdf, False)
+
+ assert called == [] # no vision caption calls despite config ON
+
+
+def test_caption_none_follows_config(rag_conn, stub_embeddings, monkeypatch, tmp_path):
+ # Omitted override (None) falls back to config.CAPTION_IMAGES.
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ seen = []
+ monkeypatch.setattr(captioner, "_caption_one", lambda *a: seen.append(1) or "chart caption")
+
+ monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", False)
+ pdf_off = tmp_path / "off.pdf"
+ _figure_pdf(pdf_off)
+ _ingest_with_caption(rag_conn, "t1", pdf_off, None)
+ assert seen == [] # config OFF + no override -> no captioning
+
+ monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
+ pdf_on = tmp_path / "on.pdf"
+ _figure_pdf(pdf_on)
+ _ingest_with_caption(rag_conn, "t2", pdf_on, None)
+ assert seen # config ON + no override -> captioning runs
diff --git a/studio/backend/tests/test_rag_ingestion.py b/studio/backend/tests/test_rag_ingestion.py
index f0b71bc23b..7e9e803687 100644
--- a/studio/backend/tests/test_rag_ingestion.py
+++ b/studio/backend/tests/test_rag_ingestion.py
@@ -83,6 +83,34 @@ def test_ingestion_dedupe_by_hash(rag_home, stub_embeddings, tmp_path):
conn.close()
+def test_ingestion_reingests_when_existing_has_zero_chunks(rag_home, stub_embeddings, tmp_path):
+ # A prior ingest of identical bytes that yielded no chunks (e.g. a scanned PDF
+ # before a vision model loaded) must re-ingest, not dedupe to the empty record.
+ path = _write(tmp_path, "doc.txt", "alpha bravo charlie " * 50)
+ sha = ingestion._sha256_file(path)
+ scope = store.kb_scope("K1")
+ conn = rag_db.get_connection()
+ try:
+ empty_id = store.create_document(conn, scope = scope, filename = "old.txt", sha256 = sha)
+ store.set_document_status(conn, empty_id, "completed", num_chunks = 0)
+ finally:
+ conn.close()
+
+ doc_id, job_id = ingestion.start_ingestion(scope, "K1", None, "doc.txt", path)
+ events = _drain(job_id)
+ _wait_completed(job_id)
+
+ assert not any(e.get("deduped") for e in events) # not a dedupe -> real ingest
+ assert doc_id != empty_id
+ conn = rag_db.get_connection()
+ try:
+ docs = store.list_documents(conn, scope)
+ assert len(docs) == 1 # the empty record was removed, replaced by the new one
+ assert docs[0]["num_chunks"] > 0
+ finally:
+ conn.close()
+
+
def test_ingestion_dedupe_removes_duplicate_upload(rag_home, stub_embeddings):
from utils.paths import ensure_dir, rag_uploads_root
@@ -210,6 +238,41 @@ def test_delete_document_route_removes_stored_upload(rag_home):
conn.close()
+def test_get_job_status_includes_num_chunks(rag_home, stub_embeddings, tmp_path):
+ # The poll/reconcile path reads num_chunks from get_job_status (the SSE complete
+ # frame carries it, but a client that falls back to polling needs it here too).
+ path = _write(tmp_path, "doc.txt", "alpha bravo charlie " * 50)
+ scope = store.kb_scope("K1")
+ _doc_id, job_id = ingestion.start_ingestion(scope, "K1", None, "doc.txt", path)
+ _drain(job_id)
+ _wait_completed(job_id)
+ status = ingestion.get_job_status(job_id)
+ assert status["status"] == "completed"
+ assert status["num_chunks"] and status["num_chunks"] > 0
+
+
+def test_save_upload_rejects_oversize_file(rag_home, monkeypatch):
+ # A file over the cap is rejected (413) and its partial bytes are cleaned up.
+ import io
+
+ from fastapi import HTTPException
+
+ from core.rag import config
+ from routes import rag as rag_routes
+ from utils.paths import rag_uploads_root
+
+ monkeypatch.setattr(config, "MAX_UPLOAD_BYTES", 1024)
+
+ class _Up:
+ filename = "big.txt"
+ file = io.BytesIO(b"x" * 4096)
+
+ with pytest.raises(HTTPException) as ei:
+ rag_routes._save_upload(_Up())
+ assert ei.value.status_code == 413
+ assert list(rag_uploads_root().glob("*.txt")) == [] # partial upload removed
+
+
def test_ingestion_delete_removes_all_rows(rag_home, stub_embeddings, tmp_path):
path = _write(tmp_path, "doc.txt", "alpha bravo charlie delta")
scope = store.kb_scope("K1")
diff --git a/studio/backend/tests/test_rag_ocr_fallback.py b/studio/backend/tests/test_rag_ocr_fallback.py
new file mode 100644
index 0000000000..c7be1fe60b
--- /dev/null
+++ b/studio/backend/tests/test_rag_ocr_fallback.py
@@ -0,0 +1,259 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Scanned-PDF OCR fallback: a PDF page with no text layer is rendered and transcribed
+by the vision model during ingestion, so image-only PDFs become searchable. The vision
+call is stubbed, so no model is needed."""
+
+import pymupdf
+
+from core.rag import captioner, ingestion, parsers, store, tool
+
+
+def _image_only_pdf(path, *, pages = 1):
+ """A PDF whose pages carry only a raster image, so get_text returns ''."""
+ doc = pymupdf.open()
+ pix = pymupdf.Pixmap(pymupdf.csRGB, pymupdf.IRect(0, 0, 120, 120))
+ pix.clear_with(220)
+ for _ in range(pages):
+ page = doc.new_page()
+ page.insert_image(page.rect, pixmap = pix)
+ doc.save(str(path))
+ doc.close()
+
+
+def _text_pdf(path, body):
+ doc = pymupdf.open()
+ page = doc.new_page()
+ page.insert_textbox(pymupdf.Rect(40, 40, 550, 800), body, fontsize = 11)
+ doc.save(str(path))
+ doc.close()
+
+
+def _ingest(rag_conn, thread_id, filename, path):
+ """Drive the real ingestion worker synchronously and return the document row."""
+ scope = store.thread_scope(thread_id)
+ document_id = store.create_document(
+ rag_conn,
+ scope = scope,
+ filename = filename,
+ sha256 = filename,
+ thread_id = thread_id,
+ status = "pending",
+ stored_path = str(path),
+ )
+ job_id = ingestion._new_job(rag_conn, document_id, scope)
+ ingestion._run(job_id, document_id, scope, str(path), None)
+ return store.get_document(rag_conn, document_id)
+
+
+# ── parsers.render_pdf_pages ─────────────────────────────────────────
+
+
+def test_render_pdf_pages_returns_png_per_page(tmp_path):
+ pdf = tmp_path / "two.pdf"
+ _image_only_pdf(pdf, pages = 2)
+ out = parsers.render_pdf_pages(str(pdf), [1, 2], dpi = 72)
+ assert set(out) == {1, 2}
+ assert all(b.startswith(b"\x89PNG") for b in out.values())
+
+
+def test_render_pdf_pages_excludes_unwanted(tmp_path):
+ pdf = tmp_path / "three.pdf"
+ _image_only_pdf(pdf, pages = 3)
+ out = parsers.render_pdf_pages(str(pdf), [2], dpi = 72)
+ assert set(out) == {2}
+
+
+def test_render_pdf_pages_empty_request(tmp_path):
+ pdf = tmp_path / "one.pdf"
+ _image_only_pdf(pdf, pages = 1)
+ assert parsers.render_pdf_pages(str(pdf), [], dpi = 72) == {}
+
+
+# ── captioner.ocr_pages gating ───────────────────────────────────────
+
+
+def test_ocr_pages_no_endpoint(monkeypatch):
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: None)
+ assert captioner.ocr_pages({1: b"x"}) == {}
+
+
+def test_collapse_runaway_caps_repeated_lines():
+ # A looping model repeats a line hundreds of times; the guard caps it, keeps repeats.
+ text = "\n".join(["TITLE"] * 200 + ["body"] + ["Add & Norm"] * 3)
+ out = captioner._collapse_runaway(text)
+ lines = out.splitlines()
+ assert lines.count("TITLE") == 3 # 200 -> 3
+ assert lines.count("Add & Norm") == 3 # legitimate triple survives
+ assert "body" in lines
+
+
+def test_collapse_runaway_caps_interleaved_repeats():
+ # Models also loop non-consecutively; the global per-line cap bounds those too.
+ text = "\n".join(["Llion Vaswani Google", "Niki Parmar Google"] * 40)
+ out = captioner._collapse_runaway(text)
+ lines = [ln for ln in out.splitlines() if ln.strip()]
+ assert lines.count("Llion Vaswani Google") <= 8
+ assert lines.count("Niki Parmar Google") <= 8
+
+
+def test_collapse_runaway_noop_on_normal_text():
+ text = "Heading\n\nFirst paragraph.\nSecond paragraph.\n\nFooter"
+ assert captioner._collapse_runaway(text) == text
+
+
+def test_ocr_pages_applies_runaway_guard(monkeypatch):
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner, "_ocr_one", lambda *a: "\n".join(["X"] * 50))
+ out = captioner.ocr_pages({1: b"img"}, endpoint = ("http://x", "local"))
+ assert out[1].splitlines().count("X") == 3 # guard applied to stored text
+
+
+def test_ocr_pages_transcribes_and_caps(monkeypatch):
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner.config, "OCR_MAX_PAGES", 1)
+ calls = []
+ monkeypatch.setattr(
+ captioner,
+ "_ocr_one",
+ lambda base, model, b, t: (calls.append(1) or "transcribed text"),
+ )
+ out = captioner.ocr_pages({1: b"a", 2: b"b"}, endpoint = ("http://x", "local"))
+ assert out == {1: "transcribed text"} # page 2 dropped by the cap
+ assert len(calls) == 1
+
+
+def test_ocr_scanned_pages_merges_short_text_layer(rag_conn, monkeypatch):
+ # Near-empty pages can still have meaningful extractable text; OCR augments it
+ # rather than replacing it with a fallible vision transcription.
+ scope = store.thread_scope("t1")
+ document_id = store.create_document(rag_conn, scope = scope, filename = "scan.pdf", sha256 = "h")
+ job_id = ingestion._new_job(rag_conn, document_id, scope)
+ pages = [parsers.Page("ID-42", 1, 5)]
+
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner.config, "OCR_MIN_CHARS", 16)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ monkeypatch.setattr(parsers, "render_pdf_pages", lambda *a, **k: {1: b"png"})
+ monkeypatch.setattr(captioner, "ocr_pages", lambda page_pngs: {1: "OCR body text"})
+
+ out, ocred = ingestion._ocr_scanned_pages(pages, "scan.pdf", rag_conn, job_id)
+ assert ocred == {1}
+ assert out[0].text == "ID-42\n\nOCR body text"
+
+
+# ── end-to-end ingestion ─────────────────────────────────────────────
+
+
+def test_scanned_pdf_is_ocred_into_chunks(rag_conn, stub_embeddings, monkeypatch, tmp_path):
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ monkeypatch.setattr(
+ captioner, "_ocr_one", lambda base, model, b, t: "Invoice total is zebra-42 due Friday"
+ )
+
+ pdf = tmp_path / "scan.pdf"
+ _image_only_pdf(pdf, pages = 1)
+ doc = _ingest(rag_conn, "t1", "scan.pdf", pdf)
+
+ assert doc["status"] == "completed"
+ assert doc["num_chunks"] >= 1
+ # The OCR'd text is now indexed and reaches whole-document injection.
+ text, _sources = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "zebra-42" in text
+
+
+def test_scanned_page_past_ocr_cap_is_still_captioned(
+ rag_conn, stub_embeddings, monkeypatch, tmp_path
+):
+ # OCR is capped to one page, so page 2 is scanned but never transcribed. Figure
+ # captioning must still cover it (we exclude only the pages OCR actually handled),
+ # so a chart on an un-OCR'd scanned page is not silently dropped.
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner.config, "OCR_MAX_PAGES", 1)
+ monkeypatch.setattr(captioner.config, "CAPTION_IMAGES", True)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ monkeypatch.setattr(captioner, "_ocr_one", lambda *a: "scanned page alpha")
+ monkeypatch.setattr(captioner, "_caption_one", lambda *a: "figure caption bravo")
+
+ pdf = tmp_path / "scan2.pdf"
+ _image_only_pdf(pdf, pages = 2)
+ doc = _ingest(rag_conn, "t1", "scan2.pdf", pdf)
+
+ assert doc["status"] == "completed"
+ text, _ = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "scanned page alpha" in text # page 1 OCR'd, within the cap
+ assert "figure caption bravo" in text # page 2 past the cap -> captioned, not dropped
+
+
+def test_born_digital_pdf_skips_ocr(rag_conn, stub_embeddings, monkeypatch, tmp_path):
+ called = []
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner, "_ocr_one", lambda *a: called.append(1) or "should not run")
+
+ pdf = tmp_path / "digital.pdf"
+ _text_pdf(pdf, "Real born digital body text. " * 30 + "marker-quokka")
+ doc = _ingest(rag_conn, "t1", "digital.pdf", pdf)
+
+ assert doc["status"] == "completed"
+ assert called == [] # page had real text -> never considered scanned
+ text, _sources = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "marker-quokka" in text
+
+
+def _ingest_with_ocr(rag_conn, thread_id, path, ocr):
+ scope = store.thread_scope(thread_id)
+ document_id = store.create_document(
+ rag_conn,
+ scope = scope,
+ filename = "scan.pdf",
+ sha256 = str(path) + str(ocr),
+ thread_id = thread_id,
+ status = "pending",
+ stored_path = str(path),
+ )
+ job_id = ingestion._new_job(rag_conn, document_id, scope)
+ ingestion._run(job_id, document_id, scope, str(path), None, ocr = ocr)
+ return store.get_document(rag_conn, document_id)
+
+
+def test_ocr_override_false_skips_ocr_when_config_on(
+ rag_conn, stub_embeddings, monkeypatch, tmp_path
+):
+ # Config default ON, but the per-upload toggle (ocr=False) skips OCR.
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", True)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ monkeypatch.setattr(captioner, "_ocr_one", lambda *a: "should not run")
+ pdf = tmp_path / "scan.pdf"
+ _image_only_pdf(pdf, pages = 1)
+ doc = _ingest_with_ocr(rag_conn, "t1", pdf, ocr = False)
+ assert doc["num_chunks"] == 0 # scanned page left empty
+
+
+def test_ocr_override_true_runs_ocr_when_config_off(
+ rag_conn, stub_embeddings, monkeypatch, tmp_path
+):
+ # Config default OFF, but the per-upload toggle (ocr=True) forces OCR on.
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", False)
+ monkeypatch.setattr(captioner, "vision_endpoint", lambda: ("http://x", "local"))
+ monkeypatch.setattr(captioner, "_ocr_one", lambda *a: "forced ocr text quokka")
+ pdf = tmp_path / "scan.pdf"
+ _image_only_pdf(pdf, pages = 1)
+ doc = _ingest_with_ocr(rag_conn, "t1", pdf, ocr = True)
+ assert doc["num_chunks"] >= 1
+ text, _ = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "quokka" in text
+
+
+def test_ocr_disabled_leaves_scanned_pdf_empty(rag_conn, stub_embeddings, monkeypatch, tmp_path):
+ monkeypatch.setattr(captioner.config, "OCR_SCANNED", False)
+
+ pdf = tmp_path / "scan.pdf"
+ _image_only_pdf(pdf, pages = 1)
+ doc = _ingest(rag_conn, "t1", "scan.pdf", pdf)
+
+ # With OCR off, a text-less scanned page yields no chunks (prior behavior).
+ assert doc["status"] == "completed"
+ assert doc["num_chunks"] == 0
+ assert tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000) is None
diff --git a/studio/backend/tests/test_rag_parsing.py b/studio/backend/tests/test_rag_parsing.py
new file mode 100644
index 0000000000..4c46f49495
--- /dev/null
+++ b/studio/backend/tests/test_rag_parsing.py
@@ -0,0 +1,88 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""PDF text extraction: layout-aware Markdown (pymupdf4llm) with plain-text fallback."""
+
+from __future__ import annotations
+
+import pytest
+
+pytest.importorskip("pymupdf")
+
+
+def _table_pdf(path):
+ import pymupdf
+
+ doc = pymupdf.open()
+ page = doc.new_page()
+ page.insert_textbox(pymupdf.Rect(40, 40, 550, 70), "Quarterly Results", fontsize = 16)
+ rows = [("Quarter", "Revenue", "Growth"), ("Q1", "$1.2M", "12%"), ("Q2", "$1.5M", "25%")]
+ y = 90
+ for r in rows:
+ page.insert_textbox(pymupdf.Rect(40, y, 250, y + 20), r[0], fontsize = 11)
+ page.insert_textbox(pymupdf.Rect(250, y, 400, y + 20), r[1], fontsize = 11)
+ page.insert_textbox(pymupdf.Rect(400, y, 540, y + 20), r[2], fontsize = 11)
+ y += 24
+ doc.save(str(path))
+ doc.close()
+
+
+def test_pdf_extracts_markdown_table(tmp_path, monkeypatch):
+ # With Markdown on, the layout is emitted as Markdown markup (heading, and a pipe table
+ # where the extractor detects one) that flat get_text never produces.
+ pytest.importorskip("pymupdf4llm")
+ from core.rag import config, parsers
+
+ monkeypatch.setattr(config, "PDF_MARKDOWN", True)
+ pdf = tmp_path / "table.pdf"
+ _table_pdf(pdf)
+ text = "\n".join(p.text for p in parsers.parse(str(pdf)))
+ assert "Q2" in text and "$1.5M" in text # cell values preserved
+ assert "#" in text or "|" in text # Markdown markup (heading or table pipes)
+
+
+def test_pdf_markdown_off_uses_plain_text(tmp_path, monkeypatch):
+ # The toggle (RAG_PDF_MARKDOWN=0) falls back to flat PyMuPDF text: content is still
+ # there, but with no Markdown markup.
+ from core.rag import config, parsers
+
+ monkeypatch.setattr(config, "PDF_MARKDOWN", False)
+ pdf = tmp_path / "table.pdf"
+ _table_pdf(pdf)
+ text = "\n".join(p.text for p in parsers.parse(str(pdf)))
+ assert "Q2" in text and "$1.5M" in text
+ assert "#" not in text and "|" not in text # plain text path emits no Markdown markup
+
+
+def test_pdf_markdown_passes_only_supported_legacy_kwargs(monkeypatch):
+ # The pinned PyMuPDF4LLM legacy path ignores unknown kwargs; do not pass the
+ # newer layout-only OCR knobs or Markdown extraction silently loses policy control.
+ from core.rag import parsers
+
+ captured = {}
+
+ class _FakePymupdf4llm:
+ @staticmethod
+ def to_markdown(doc, **kwargs):
+ captured.update(kwargs)
+ return [{"text": "plain markdown"}]
+
+ class _Doc:
+ page_count = 1
+
+ monkeypatch.setitem(__import__("sys").modules, "pymupdf4llm", _FakePymupdf4llm)
+ assert parsers._pdf_markdown(_Doc()) == ["plain markdown"]
+ assert captured == {"page_chunks": True, "show_progress": False}
+
+
+def test_pdf_markdown_falls_back_when_lib_missing(tmp_path, monkeypatch):
+ # If pymupdf4llm extraction returns None (missing/failed), parsing still yields the
+ # plain-text pages rather than raising.
+ from core.rag import config, parsers
+
+ monkeypatch.setattr(config, "PDF_MARKDOWN", True)
+ monkeypatch.setattr(parsers, "_pdf_markdown", lambda doc: None)
+ pdf = tmp_path / "table.pdf"
+ _table_pdf(pdf)
+ pages = parsers.parse(str(pdf))
+ assert pages and "Quarter" in pages[0].text
diff --git a/studio/backend/tests/test_rag_preview.py b/studio/backend/tests/test_rag_preview.py
index e7f2a39792..0ff27897bd 100644
--- a/studio/backend/tests/test_rag_preview.py
+++ b/studio/backend/tests/test_rag_preview.py
@@ -165,6 +165,25 @@ def test_locator_handles_midword_anchor_and_locates_line():
assert r["width"] > 0 and r["height"] > 0
+def test_locator_anchors_through_markdown_table_pipes():
+ # Markdown table cells are pipe-joined with no spaces; the locator splits on pipes
+ # so a table-row chunk still anchors to the raw PDF word stream.
+ import pymupdf
+
+ from core.rag.locators import LocatorMatch, _regions_for_match
+
+ doc = pymupdf.open()
+ page = doc.new_page()
+ page.insert_text((72, 200), "Quarter Revenue Growth Q1 sales strong here", fontsize = 12)
+ # What the Markdown parser stores for the row (cells joined by pipes, no spaces).
+ page_text = "|Quarter|Revenue|Growth|Q1|sales|strong|here|"
+ match = LocatorMatch(page_index = 0, page_number = 1, start = 0, end = len(page_text))
+ rects = _regions_for_match(doc, page_text, match)
+ doc.close()
+
+ assert rects, "a Markdown table row should still anchor to the page words"
+
+
def test_sign_verify_roundtrip(rag_home):
from routes import rag as rag_routes
diff --git a/studio/backend/tests/test_rag_whole_document.py b/studio/backend/tests/test_rag_whole_document.py
new file mode 100644
index 0000000000..545d731fd2
--- /dev/null
+++ b/studio/backend/tests/test_rag_whole_document.py
@@ -0,0 +1,520 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Whole-document context mode: a thread-attached file small enough to fit is
+injected in full (every chunk, in order) instead of top-K retrieval. Covers the
+new store query, the tool-level renderer, and the auto-inject wiring + fallback.
+No embedder is needed - the whole-doc path does no query embedding."""
+
+import json
+
+from core.rag import store, tool
+from core.rag.chunking import Chunk
+from core.inference import tools as inf_tools
+
+# A vector per chunk just to satisfy add_chunks (the whole-doc path never reads
+# vectors); dimension is arbitrary but must be consistent within a connection.
+_VEC = [0.1, 0.2, 0.3, 0.4]
+
+
+def _chunk(
+ text,
+ index = 0,
+ page = None,
+ tokens = None,
+):
+ return Chunk(
+ text = text,
+ token_count = tokens if tokens is not None else len(text.split()),
+ page_number = page,
+ source_page_index = 0,
+ chunk_index = index,
+ page_char_start = 0,
+ page_char_end = len(text),
+ )
+
+
+def _add_doc(
+ conn,
+ scope,
+ doc_id,
+ filename,
+ sha,
+ texts,
+ *,
+ status = "completed",
+ tokens = None,
+ pages = None,
+):
+ chunks = [
+ _chunk(
+ t,
+ i,
+ page = (pages[i] if pages else None),
+ tokens = (tokens[i] if tokens else None),
+ )
+ for i, t in enumerate(texts)
+ ]
+ vectors = [list(_VEC) for _ in texts]
+ store.create_document(conn, scope = scope, filename = filename, sha256 = sha, document_id = doc_id)
+ store.add_chunks(conn, scope, doc_id, chunks, vectors)
+ store.set_document_status(conn, doc_id, status, num_chunks = len(texts))
+
+
+def _injected_text(result) -> str:
+ """The text spliced into the conversation as the synthetic tool result."""
+ tool_msg = next(m for m in result["messages"] if m.get("role") == "tool")
+ return tool_msg["content"]
+
+
+# ── store.all_chunks_for_scope ───────────────────────────────────────
+
+
+def test_all_chunks_for_scope_orders_by_document_then_index(rag_conn):
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "d1", "first.pdf", "h1", ["a", "b", "c"])
+ _add_doc(rag_conn, scope, "d2", "second.pdf", "h2", ["x", "y"])
+ rows = store.all_chunks_for_scope(rag_conn, scope)
+ assert [r["id"] for r in rows] == ["d1:0", "d1:1", "d1:2", "d2:0", "d2:1"]
+ assert rows[0]["filename"] == "first.pdf"
+ assert rows[-1]["filename"] == "second.pdf"
+ assert rows[0]["text"] == "a"
+
+
+def test_all_chunks_for_scope_excludes_non_completed(rag_conn):
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "done", "done.pdf", "h1", ["ready"])
+ _add_doc(rag_conn, scope, "pend", "pend.pdf", "h2", ["indexing"], status = "pending")
+ rows = store.all_chunks_for_scope(rag_conn, scope)
+ assert [r["id"] for r in rows] == ["done:0"]
+
+
+def test_all_chunks_for_scope_empty_scope(rag_conn):
+ assert store.all_chunks_for_scope(rag_conn, store.thread_scope("nope")) == []
+
+
+def test_all_chunks_for_scope_isolates_scopes(rag_conn):
+ _add_doc(rag_conn, store.thread_scope("t1"), "d1", "f", "h1", ["mine"])
+ _add_doc(rag_conn, store.thread_scope("t2"), "d2", "f", "h2", ["theirs"])
+ rows = store.all_chunks_for_scope(rag_conn, store.thread_scope("t1"))
+ assert [r["text"] for r in rows] == ["mine"]
+
+
+# ── store.scope_token_estimate (cheap whole-doc budget pre-check) ─────
+
+
+def test_scope_token_estimate_sums_without_hydrating(rag_conn):
+ # Stored counts sum directly; zero/missing falls back to length/4; non-completed out.
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "d1", "a.pdf", "h1", ["alpha", "bravo"], tokens = [10, 20])
+ # token_count 0 -> length/4 fallback: a 40-char chunk estimates to 10 tokens.
+ _add_doc(rag_conn, scope, "d2", "b.pdf", "h2", ["x" * 40], tokens = [0])
+ _add_doc(rag_conn, scope, "d3", "c.pdf", "h3", ["pending"], status = "pending", tokens = [99])
+ assert store.scope_token_estimate(rag_conn, scope) == 10 + 20 + 10
+ assert store.scope_token_estimate(rag_conn, store.thread_scope("none")) == 0
+
+
+def test_scope_token_estimate_matches_row_sum(rag_conn):
+ # Must agree with the exact per-row sum it short-circuits (one stored count, one
+ # length/4 fallback), so the pre-check never disagrees with the full path.
+ from core.rag.tool import _row_token_count
+
+ scope = store.thread_scope("t1")
+ _add_doc(
+ rag_conn, scope, "d1", "a.pdf", "h1", ["a long-ish chunk body here", "tail"], tokens = [0, 5]
+ )
+ rows = store.all_chunks_for_scope(rag_conn, scope)
+ assert store.scope_token_estimate(rag_conn, scope) == sum(_row_token_count(r) for r in rows)
+
+
+# ── tool.whole_document_context ──────────────────────────────────────
+
+
+def test_whole_document_context_returns_full_text_and_sources(rag_conn):
+ scope = store.thread_scope("t1")
+ _add_doc(
+ rag_conn,
+ scope,
+ "d1",
+ "report.pdf",
+ "h1",
+ ["chapter one body", "chapter two body"],
+ pages = [1, 2],
+ )
+ result = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert result is not None
+ text, sources = result
+ # Every chunk is present, in order, as blocks.
+ assert "chapter one body" in text
+ assert "chapter two body" in text
+ assert ' None (whole-doc is thread-attachment only).
+ assert tool.whole_document_context(max_tokens = 6000) is None
+
+
+def test_whole_document_context_null_token_count_enforces_budget(rag_conn):
+ # A missing token_count must not bypass the budget; fall back to a length estimate.
+ big = "word " * 20_000 # ~20k tokens by length estimate
+ _add_doc(rag_conn, store.thread_scope("t1"), "d1", "big.pdf", "h1", [big], tokens = [None])
+ assert tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000) is None
+ assert tool.whole_document_context(scope_thread_id = "t1", max_tokens = 1_000_000) is not None
+
+
+def test_whole_document_context_spans_multiple_docs(rag_conn):
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "d1", "a.pdf", "h1", ["alpha text"])
+ _add_doc(rag_conn, scope, "d2", "b.pdf", "h2", ["bravo text"])
+ text, sources = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "alpha text" in text and "bravo text" in text
+ assert {s["filename"] for s in sources} == {"a.pdf", "b.pdf"}
+
+
+# ── build_rag_autoinject wiring ──────────────────────────────────────
+
+
+def _convo(text = "summarize the whole document"):
+ return [{"role": "user", "content": text}]
+
+
+def test_build_rag_autoinject_uses_whole_doc(rag_conn):
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "d1", "doc.pdf", "h1", ["whole alpha part", "whole bravo part"])
+ result = inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1"})
+ assert result is not None
+ injected = _injected_text(result)
+ # Both chunks present -> the model receives the entire file, not top-K.
+ assert "whole alpha part" in injected
+ assert "whole bravo part" in injected
+ # Tool-message content is chunk text only; the citation JSON tail is internal.
+ assert inf_tools.RAG_SOURCES_SENTINEL not in injected
+
+
+def test_build_rag_autoinject_whole_doc_runs_when_autoinject_false(rag_conn, monkeypatch):
+ # Large-model Auto sets autoinject=False, but whole-doc is a separate thread-doc
+ # context mode and should still inject a fitting attachment.
+ _add_doc(rag_conn, store.thread_scope("t1"), "d1", "doc.pdf", "h1", ["entire file body"])
+ monkeypatch.setattr(
+ tool,
+ "search_for_autoinject",
+ lambda **kw: (_ for _ in ()).throw(AssertionError("retrieval should not run")),
+ )
+ result = inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1", "autoinject": False})
+ assert result is not None
+ assert "entire file body" in _injected_text(result)
+
+
+def test_build_rag_autoinject_explicit_off_disables_whole_doc(rag_conn, monkeypatch):
+ # The UI Off switch sends both autoinject=False and whole_doc=False.
+ _add_doc(rag_conn, store.thread_scope("t1"), "d1", "doc.pdf", "h1", ["small body"])
+ monkeypatch.setattr(
+ tool,
+ "search_for_autoinject",
+ lambda **kw: (_ for _ in ()).throw(AssertionError("retrieval should not run")),
+ )
+ assert (
+ inf_tools.build_rag_autoinject(
+ _convo(), {"thread_id": "t1", "autoinject": False, "whole_doc": False}
+ )
+ is None
+ )
+
+
+def test_build_rag_autoinject_falls_back_over_budget(rag_conn, monkeypatch):
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "d1", "big.pdf", "h1", ["overflow"], tokens = [50_000])
+
+ sentinel = ("TOPK_FALLBACK_TEXT", [{"citationId": 1, "filename": "big.pdf", "text": "x"}])
+ monkeypatch.setattr(tool, "search_for_autoinject", lambda **kw: sentinel)
+
+ result = inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1"})
+ assert result is not None
+ assert _injected_text(result) == "TOPK_FALLBACK_TEXT"
+
+
+def test_build_rag_autoinject_context_budget_falls_back(rag_conn, monkeypatch):
+ # Runtime context can be smaller than RAG_WHOLE_DOC_MAX_TOKENS; cap whole-doc to
+ # the active context and fall back to retrieval when it would overflow.
+ _add_doc(
+ rag_conn, store.thread_scope("t1"), "d1", "small.pdf", "h1", ["fits global"], tokens = [900]
+ )
+ sentinel = ("TOPK_CONTEXT_FALLBACK", [{"citationId": 1, "filename": "small.pdf", "text": "x"}])
+ monkeypatch.setattr(tool, "search_for_autoinject", lambda **kw: sentinel)
+ result = inf_tools.build_rag_autoinject(
+ _convo(), {"thread_id": "t1", "context_length": 1200, "whole_doc": True}
+ )
+ assert result is not None
+ assert _injected_text(result) == "TOPK_CONTEXT_FALLBACK"
+
+
+def test_whole_doc_budget_reserves_image_parts(monkeypatch):
+ from core.rag import config
+
+ monkeypatch.setattr(config, "WHOLE_DOC_MAX_TOKENS", 10_000)
+ scope = {"context_length": 7000, "response_headroom": 1000}
+ text_only = [{"role": "user", "content": [{"type": "text", "text": "summarize"}]}]
+ with_image = [
+ {
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "summarize"},
+ {"type": "image_url", "image_url": {"url": "data:image/png;base64,abc"}},
+ ],
+ }
+ ]
+
+ assert (
+ inf_tools._whole_doc_budget(scope, text_only)
+ - inf_tools._whole_doc_budget(scope, with_image)
+ == inf_tools._IMAGE_PART_TOKEN_ESTIMATE
+ )
+
+
+def test_build_rag_autoinject_server_kill_switch_blocks_whole_doc(rag_conn, monkeypatch):
+ # RAG_THREAD_WHOLE_DOC=0 stays authoritative; browser requests should not
+ # turn it back on by default.
+ from core.rag import config
+
+ monkeypatch.setattr(config, "THREAD_WHOLE_DOC", False)
+ _add_doc(rag_conn, store.thread_scope("t1"), "d1", "doc.pdf", "h1", ["small body"])
+ monkeypatch.setattr(
+ tool,
+ "search_for_autoinject",
+ lambda **kw: (_ for _ in ()).throw(AssertionError("retrieval should not run")),
+ )
+ assert (
+ inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1", "autoinject": False}) is None
+ )
+
+
+def test_whole_document_context_budgets_rendered_wrappers(rag_conn):
+ # Many tiny chunks add wrapper overhead beyond raw chunk token counts; budget
+ # the rendered prompt, not just stored text.
+ texts = ["x" for _ in range(120)]
+ _add_doc(
+ rag_conn,
+ store.thread_scope("t1"),
+ "d1",
+ "many-pages.pdf",
+ "h1",
+ texts,
+ tokens = [1 for _ in texts],
+ )
+ assert tool.whole_document_context(scope_thread_id = "t1", max_tokens = 500) is None
+
+
+def test_build_rag_autoinject_whole_doc_disabled_via_override(rag_conn, monkeypatch):
+ scope = store.thread_scope("t1")
+ _add_doc(rag_conn, scope, "d1", "doc.pdf", "h1", ["small body"])
+
+ sentinel = ("TOPK_TEXT", [{"citationId": 1, "filename": "doc.pdf", "text": "x"}])
+ monkeypatch.setattr(tool, "search_for_autoinject", lambda **kw: sentinel)
+
+ # whole_doc=False forces retrieval even though the doc fits.
+ result = inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1", "whole_doc": False})
+ assert result is not None
+ assert _injected_text(result) == "TOPK_TEXT"
+
+
+def test_build_rag_autoinject_kb_scope_never_whole_doc(rag_conn, monkeypatch):
+ # A KB-only scope (no thread) goes through retrieval, never whole-doc.
+ kb_scope = store.kb_scope("K1")
+ _add_doc(rag_conn, kb_scope, "d1", "kb.pdf", "h1", ["kb body one", "kb body two"])
+
+ sentinel = ("KB_RETRIEVAL_TEXT", [{"citationId": 1, "filename": "kb.pdf", "text": "x"}])
+ monkeypatch.setattr(tool, "search_for_autoinject", lambda **kw: sentinel)
+
+ result = inf_tools.build_rag_autoinject(_convo(), {"kb_id": "K1"})
+ assert result is not None
+ assert _injected_text(result) == "KB_RETRIEVAL_TEXT"
+
+
+def test_whole_document_context_thread_scope_only(rag_conn):
+ # A project corpus chunk is never whole-doc injected, even with a thread attachment.
+ _add_doc(rag_conn, store.thread_scope("t1"), "td", "thread.txt", "h1", ["thread attachment"])
+ _add_doc(rag_conn, store.project_scope("p1"), "pd", "project.txt", "h2", ["project corpus"])
+ text, sources = tool.whole_document_context(scope_thread_id = "t1", max_tokens = 6000)
+ assert "thread attachment" in text
+ assert "project corpus" not in text
+ assert {s["filename"] for s in sources} == {"thread.txt"}
+
+
+def test_build_rag_autoinject_appends_project_retrieval(rag_conn, monkeypatch):
+ # Project chat: thread attachment whole-doc'd AND project sources retrieved, merged.
+ _add_doc(
+ rag_conn,
+ store.thread_scope("t1"),
+ "td",
+ "thread.txt",
+ "h1",
+ ["thread chunk one", "thread chunk two"],
+ )
+ proj = (
+ "PROJ",
+ [
+ {
+ "citationId": 1,
+ "chunkId": "pj:0",
+ "documentId": "pj",
+ "filename": "project.txt",
+ "page": None,
+ "text": "project passage zeta",
+ "score": 0.91,
+ }
+ ],
+ )
+ captured = {}
+
+ def fake_search(**kw):
+ captured.update(kw)
+ return proj
+
+ monkeypatch.setattr(tool, "search_for_autoinject", fake_search)
+ result = inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1", "project_id": "p1"})
+ injected = _injected_text(result)
+ # Whole thread attachment AND the project passage are both injected.
+ assert "thread chunk one" in injected
+ assert "thread chunk two" in injected
+ assert "project passage zeta" in injected
+ # The companion retrieval was scoped to the project only (not thread or KB).
+ assert captured.get("scope_project_id") == "p1"
+ assert captured.get("scope_thread_id") is None
+ assert captured.get("scope_kb_id") is None
+ # Citation ids are sequential across the merged set: thread 1,2 then project 3.
+ assert ' whole-doc injection ────────
+
+
+def test_real_ingestion_feeds_whole_document(rag_conn, stub_embeddings, tmp_path):
+ """Drive the real ingestion worker on a multi-paragraph file, then confirm whole-doc
+ injection splices the entire document, not just retrieved chunks."""
+ from core.rag import ingestion
+
+ scope = store.thread_scope("t1")
+ body = (
+ "# Quarterly Report\n\n"
+ + ("Revenue rose across every region this period. " * 40)
+ + "\n\nThe unique closing marker is xyzzy-sentinel for the final page. " * 40
+ )
+ src = tmp_path / "report.md"
+ src.write_text(body, encoding = "utf-8")
+
+ document_id = store.create_document(
+ rag_conn,
+ scope = scope,
+ filename = "report.md",
+ sha256 = "sha-e2e",
+ thread_id = "t1",
+ status = "pending",
+ stored_path = str(src),
+ )
+ job_id = ingestion._new_job(rag_conn, document_id, scope)
+ ingestion._run(job_id, document_id, scope, str(src), None)
+
+ doc = store.get_document(rag_conn, document_id)
+ assert doc["status"] == "completed"
+ assert doc["num_chunks"] >= 2 # the doc chunked into multiple pieces
+
+ result = inf_tools.build_rag_autoinject(_convo(), {"thread_id": "t1"})
+ assert result is not None
+ injected = _injected_text(result)
+ # Opening and ending both present -> the whole file reached the model.
+ assert "Revenue rose" in injected
+ assert "xyzzy-sentinel" in injected
+ # Every stored chunk is represented as a numbered block.
+ assert injected.count(" ast.FunctionDef:
+ """Parse load_model into an AST FunctionDef (no import side effects)."""
+ src = textwrap.dedent(inspect.getsource(LlamaCppBackend.load_model))
+ return ast.parse(src).body[0]
+
+
+def _tensor_parallel_false_drop_guards() -> list[str]:
+ """Source of the guard expression for every `if ...: tensor_parallel = False`
+ (the LOCAL variable, not self._tensor_parallel) inside load_model."""
+ fn = _load_model_ast()
+
+ def _body_drops_tp(body) -> bool:
+ for n in body:
+ if (
+ isinstance(n, ast.Assign)
+ and any(isinstance(t, ast.Name) and t.id == "tensor_parallel" for t in n.targets)
+ and isinstance(n.value, ast.Constant)
+ and n.value.value is False
+ ):
+ return True
+ return False
+
+ return [
+ ast.unparse(node.test)
+ for node in ast.walk(fn)
+ if isinstance(node, ast.If) and _body_drops_tp(node.body)
+ ]
+
+
+# Every condition that may flip a requested tensor_parallel back to False. Adding
+# one must be conscious: update this allowlist and keep multi-GPU where possible.
+_ALLOWED_TP_DROP_GUARDS = {
+ # Capability: --split-mode tensor aborted for this (binary, model) (#6415).
+ # Self-healing -- tried by default, skipped only after a real abort (vs #6416).
+ "tensor_parallel and self._tensor_split_aborts(binary, model_identifier)",
+ # Capacity: tensor needs >= 2 GPUs clearing the compute-buffer reserve.
+ "tensor_parallel and len(tp_gpus) < 2",
+ # Capacity: pooled usable VRAM can't hold weights + MTP reserve -> layer split.
+ "_tp_weight_budget_mib <= _tp_required_mib",
+}
+
+
+def test_tensor_parallel_drop_sites_match_allowlist():
+ """The set of reasons a requested TP can be dropped is fixed and reviewed: a new
+ drop site fails this set-equality until consciously allowlisted (would catch #6416)."""
+ found = set(_tensor_parallel_false_drop_guards())
+ assert found == _ALLOWED_TP_DROP_GUARDS, (
+ "tensor_parallel drop sites changed.\n"
+ f" unexpected (new) : {sorted(found - _ALLOWED_TP_DROP_GUARDS)}\n"
+ f" missing (removed): {sorted(_ALLOWED_TP_DROP_GUARDS - found)}\n"
+ "A new drop means a user's TP request is ignored for a new reason -- "
+ "review it, keep multi-GPU where possible, surface it, then update "
+ "_ALLOWED_TP_DROP_GUARDS."
+ )
+
+
+def test_every_tp_drop_is_logged_not_silent():
+ """Each tensor_parallel downgrade must log why, so it never disappears silently."""
+ fn = _load_model_ast()
+
+ def _body_drops_tp(body):
+ return any(
+ isinstance(n, ast.Assign)
+ and any(isinstance(t, ast.Name) and t.id == "tensor_parallel" for t in n.targets)
+ and isinstance(n.value, ast.Constant)
+ and n.value.value is False
+ for n in body
+ )
+
+ def _body_logs(body) -> bool:
+ for n in ast.walk(ast.Module(body = list(body), type_ignores = [])):
+ if (
+ isinstance(n, ast.Call)
+ and isinstance(n.func, ast.Attribute)
+ and isinstance(n.func.value, ast.Name)
+ and n.func.value.id == "logger"
+ ):
+ return True
+ return False
+
+ for node in ast.walk(fn):
+ if isinstance(node, ast.If) and _body_drops_tp(node.body):
+ assert _body_logs(node.body), (
+ f"TP drop under `{ast.unparse(node.test)}` has no logger call -- "
+ "downgrades must explain themselves."
+ )
+
+
+def test_tensor_split_gate_is_self_healing_not_blanket():
+ """Skip is conditional on a recorded (binary, model) abort, not a blanket
+ is_vision disable (the #6416 regression)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ assert "self._tensor_split_aborts(binary, model_identifier)" in src
+ assert "if tensor_parallel and is_vision:" not in src
+ assert "if tensor_parallel and effective_is_vision:" not in src
+
+
+def test_tensor_split_skip_documents_layer_split_fallback():
+ """When the skip fires (known-bad binary+model), it states the fallback."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ gate = src.find("self._tensor_split_aborts(binary, model_identifier)")
+ assert gate != -1
+ block = src[gate : gate + 600]
+ assert "layer split" in block, "the skip should state it falls back to layer split"
+
+
+def test_tensor_split_abort_recorded_early_on_first_spawn():
+ """Recorded on the first spawn showing the marker, before the flash-attn-off
+ retry (which can't run tensor so drops the marker) -- else it loops (oobabooga, #6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ idx = src.find("_record_tensor_split_abort(binary, model_identifier)")
+ assert idx != -1, "load_model must record a (binary, model) tensor-split abort"
+ guard = src[max(0, idx - 600) : idx]
+ assert "self._tensor_parallel" in guard
+ assert (
+ "_should_record_tensor_split_abort" in guard
+ ), "record must be gated on the marker-plus-hard-crash decision helper"
+ # Recorded before the flash-attn-off retry, not after the full ladder.
+ fa_off = src.find("_with_flash_attn_off")
+ assert 0 <= idx < fa_off, "recording must latch on the first spawn, before flash-off"
+
+
+def test_vision_downgrade_preserves_multi_gpu_intent():
+ """The vision downgrade raises _layer_min_gpus and threads it into both the
+ _select_gpus and auto-context layer paths, so a fitting model still spreads."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ assert "_layer_min_gpus = max(_layer_min_gpus, len(gpus))" in src
+ assert src.count("min_gpus = _layer_min_gpus") >= 2
+ assert "range(_auto_min_gpus, len(ranked) + 1)" in src
+ auto = src.find("_auto_min_gpus = max(")
+ assert auto != -1 and "_layer_min_gpus" in src[auto : auto + 200]
+
+
+# ── per-binary capability cache (pure) ───────────────────────────────
+
+
+def test_tensor_attempted_by_default_for_unknown_binary():
+ """A (binary, model) not seen to abort -> tensor is attempted (not skipped)."""
+ assert LlamaCppBackend._tensor_split_aborts("/never/seen/llama-server", "m") is False
+ assert LlamaCppBackend._tensor_split_aborts(None, "m") is False
+ assert LlamaCppBackend._tensor_split_aborts("/x", None) is False
+
+
+def test_recorded_tensor_abort_is_per_model():
+ """A recorded (binary, model) abort trips the gate for that model only -- a
+ different model on the same binary still attempts tensor (oobabooga, #6659)."""
+ b = f"/tmp/llama-server-{id(object())}"
+ try:
+ assert LlamaCppBackend._tensor_split_aborts(b, "model-a") is False
+ LlamaCppBackend._record_tensor_split_abort(b, "model-a")
+ assert LlamaCppBackend._tensor_split_aborts(b, "model-a") is True
+ # a different model on the same binary is unaffected
+ assert LlamaCppBackend._tensor_split_aborts(b, "model-b") is False
+ finally:
+ LlamaCppBackend._tensor_split_abort_keys.discard(
+ LlamaCppBackend._tensor_split_cache_key(b, "model-a")
+ )
+
+
+# ── _select_gpus: single-GPU collapse vs honored multi-GPU intent (pure) ──
+
+
+def test_select_gpus_collapses_to_single_gpu_when_model_fits():
+ """Default (min_gpus=1): a 39 GB model on four 183 GB GPUs pins ONE GPU -- the
+ 'single GPU' symptom once TP drops, and why the downgrade needs min_gpus."""
+ gpus = [(0, 180000), (1, 180000), (2, 180000), (3, 180000)] # (idx, free MiB)
+ gpu_indices, _use_fit = LlamaCppBackend._select_gpus(int(39 * _GB), gpus)
+ assert gpu_indices is not None and len(gpu_indices) == 1
+
+
+def test_select_gpus_min_gpus_keeps_multi_gpu_for_fitting_model():
+ """min_gpus>=2 must NOT collapse to one GPU for a model that fits on one."""
+ gpus = [(0, 180000), (1, 180000), (2, 180000), (3, 180000)]
+ gpu_indices, _ = LlamaCppBackend._select_gpus(int(39 * _GB), gpus, min_gpus = 2)
+ assert gpu_indices is not None and len(gpu_indices) >= 2
+
+
+def test_select_gpus_min_gpus_capped_to_available():
+ """min_gpus larger than the GPU count is capped, not an error."""
+ gpus = [(0, 180000), (1, 180000)]
+ gi, _ = LlamaCppBackend._select_gpus(int(10 * _GB), gpus, min_gpus = 8)
+ assert gi is not None and len(gi) == 2
+
+
+def test_select_gpus_uses_multiple_gpus_when_model_does_not_fit():
+ """Sanity: selection spreads across GPUs when one card can't hold the model."""
+ gpus = [(0, 40000), (1, 40000), (2, 40000), (3, 40000)] # 40 GB free each
+ gpu_indices, _use_fit = LlamaCppBackend._select_gpus(int(120 * _GB), gpus)
+ assert gpu_indices is not None and len(gpu_indices) >= 2
+
+
+def test_select_gpus_min_gpus_excludes_unusable_gpu():
+ """min_gpus caps to usable cards: 2 free + 1 nearly-full -> 2-GPU split, not
+ forcing the full card (OOM) or tripping --fit (#6659)."""
+ gpus = [(0, 180000), (1, 180000), (2, 500)] # GPU 2 is nearly full
+ total = {0: 180000, 1: 180000, 2: 180000}
+ gi, _ = LlamaCppBackend._select_gpus(
+ int(39 * _GB),
+ gpus,
+ min_gpus = 3,
+ total_by_idx = total,
+ per_device_overhead_bytes = int(1 * _GB),
+ )
+ assert gi is not None
+ assert 2 not in gi, "a nearly-full GPU must not be forced in to satisfy min_gpus"
+ assert len(gi) == 2
+
+
+def test_tensor_abort_cache_invalidated_on_binary_mtime_change(tmp_path):
+ """Cache keys on (path, mtime, model), so a binary swapped in place (in-app
+ update, no restart) is re-probed instead of inheriting the old abort (#6659)."""
+ binp = tmp_path / "llama-server"
+ binp.write_text("v1")
+ p = str(binp)
+ try:
+ LlamaCppBackend._record_tensor_split_abort(p, "m")
+ assert LlamaCppBackend._tensor_split_aborts(p, "m") is True
+ # Simulate an in-place update bumping the binary's mtime.
+ st = binp.stat()
+ os.utime(p, (st.st_atime, st.st_mtime + 10))
+ assert (
+ LlamaCppBackend._tensor_split_aborts(p, "m") is False
+ ), "a binary swapped in place (new mtime) must be re-probed"
+ # A same-second replacement (sub-second mtime bump) must also re-probe:
+ # second-resolution mtime would inherit the stale abort (reviewer.py P2).
+ sec_ns = (binp.stat().st_mtime_ns // 1_000_000_000) * 1_000_000_000
+ os.utime(p, ns = (sec_ns, sec_ns))
+ LlamaCppBackend._record_tensor_split_abort(p, "m")
+ binp.write_text("v2")
+ os.utime(p, ns = (sec_ns, sec_ns + 1))
+ assert (
+ LlamaCppBackend._tensor_split_aborts(p, "m") is False
+ ), "a same-second in-place swap (ns mtime bump) must be re-probed"
+ finally:
+ for key in list(LlamaCppBackend._tensor_split_abort_keys):
+ if key and key[0] == p:
+ LlamaCppBackend._tensor_split_abort_keys.discard(key)
+
+
+def test_tensor_split_abort_raises_early_to_layer_fallback():
+ """The first-spawn abort raises to the route's layer fallback (not the text-only
+ mmproj strip), before the flash-attn-off retry, preserving the projector (#6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ raise_idx = src.find("(split-axis geometry); retrying with layer split")
+ assert raise_idx != -1, "the split-axis abort must raise to trigger a layer retry"
+ # raises before both the flash-attn-off retry and the text-only mmproj strip
+ assert raise_idx < src.find("_with_flash_attn_off")
+ assert raise_idx < src.find("_strip_mmproj_args(_last_spawn_cmd)")
+ # gated on the marker-plus-crash helper, which also drives the record just above
+ guard = src[max(0, raise_idx - 600) : raise_idx]
+ assert "_should_record_tensor_split_abort" in guard
+ rec_idx = src.find("_record_tensor_split_abort(binary, model_identifier)")
+ assert rec_idx != -1 and rec_idx < raise_idx
+
+
+def test_budget_downgrade_preserves_multi_gpu_intent():
+ """The pooled-VRAM downgrade raises _layer_min_gpus from the usable tensor GPUs
+ too, symmetric with the vision downgrade (reviewer.py asymmetric fix, #6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ budget = src.find("_tp_weight_budget_mib <= _tp_required_mib")
+ assert budget != -1
+ block = src[budget : budget + 1000]
+ assert "tensor_parallel = False" in block
+ assert (
+ "_layer_min_gpus = max(_layer_min_gpus, len(tp_gpus))" in block
+ ), "the budget downgrade must preserve multi-GPU intent like the vision gate"
+
+
+def test_compute_buffer_downgrade_preserves_multi_gpu_intent():
+ """The len(tp_gpus) < 2 compute-buffer downgrade raises _layer_min_gpus from the
+ full GPU set too, so it is symmetric with the budget/geometry downgrades and
+ doesn't collapse a multi-GPU layer load to one card (reviewer.py P1 on #6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ gate = src.find("tensor_parallel and len(tp_gpus) < 2")
+ assert gate != -1
+ # Bound to exactly this block: from its gate to the next (budget) downgrade.
+ nxt = src.find("_tp_weight_budget_mib <= _tp_required_mib", gate)
+ assert nxt != -1
+ block = src[gate:nxt]
+ assert "tensor_parallel = False" in block
+ assert (
+ "_layer_min_gpus = max(_layer_min_gpus, len(gpus))" in block
+ ), "the compute-buffer downgrade must preserve multi-GPU intent like the others"
+
+
+def test_tensor_split_layer_min_gpus_bump_requires_tensor_request():
+ """Every guard that bumps _layer_min_gpus off the abort cache also tests
+ tensor_parallel, so a non-tensor load on a known-bad binary doesn't grab every
+ GPU for a fitting model (#6659)."""
+ fn = _load_model_ast()
+ checked = 0
+ for node in ast.walk(fn):
+ if isinstance(node, ast.If):
+ test_src = ast.unparse(node.test)
+ if "self._tensor_split_aborts(binary, model_identifier)" not in test_src:
+ continue
+ body = "\n".join(ast.unparse(n) for n in node.body)
+ if "_layer_min_gpus" in body:
+ checked += 1
+ assert "tensor_parallel" in test_src, (
+ "the cached _layer_min_gpus bump must require a current tensor "
+ f"request, but fires under `{test_src}`"
+ )
+ assert checked >= 1, "expected an abort-cache guard that bumps _layer_min_gpus"
+
+
+# ── round-2 follow-up: route-fallback retry + auto-context cap + assert marker ──
+
+
+def test_layer_fallback_retry_preserves_multi_gpu_intent():
+ """load_model takes a preserve_multi_gpu_on_layer hint and raises _layer_min_gpus
+ for it, so the tensor-off fallback retry still spreads a fitting model (#6659)."""
+ sig = inspect.signature(LlamaCppBackend.load_model)
+ assert "preserve_multi_gpu_on_layer" in sig.parameters
+ assert sig.parameters["preserve_multi_gpu_on_layer"].default is False
+ fn = _load_model_ast()
+ found = any(
+ isinstance(n, ast.If)
+ and "preserve_multi_gpu_on_layer" in ast.unparse(n.test)
+ and "_layer_min_gpus" in "\n".join(ast.unparse(b) for b in n.body)
+ for n in ast.walk(fn)
+ )
+ assert found, "preserve_multi_gpu_on_layer must raise _layer_min_gpus"
+
+
+def test_auto_context_layer_loops_capped_to_usable_gpus():
+ """The auto-context loops bypass _select_gpus, so they apply its cap: a card
+ counts only if usable VRAM clears the per-device layer overhead (#6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ assert (
+ "range(max(1, _layer_min_gpus), len(ranked) + 1)" not in src
+ ), "auto-context loops must cap _layer_min_gpus to usable GPUs, not use it raw"
+ assert "_auto_min_gpus" in src
+ assert "range(_auto_min_gpus, len(ranked) + 1)" in src
+ # the eligibility threshold is the per-device layer overhead, not bare > 0
+ auto = src.find("_auto_min_gpus = max(")
+ assert auto != -1
+ block = src[auto : auto + 400]
+ assert "_pipeline_overhead_mib" in block, (
+ "a card must clear the per-device layer overhead to count, mirroring "
+ "_select_gpus, so a nearly-full GPU is not exposed and OOMs"
+ )
+
+
+def test_fallback_hint_uses_effective_tensor_request_not_just_toggle():
+ """Tensor intent keys off _effective_tensor_parallel (toggle + extras + env), not
+ just the toggle, so extra/env-driven tensor users keep multi-GPU (#6659)."""
+ route = Path(_BACKEND_DIR) / "routes" / "inference.py"
+ src = route.read_text()
+ idx = src.find("_tensor_intent_overall = _effective_tensor_parallel(")
+ assert idx != -1, "the GGUF load closure must compute tensor intent"
+ block = src[idx : idx + 300]
+ assert "extra_llama_args, request.tensor_parallel" in block
+ pres = src.find("preserve_multi_gpu_on_layer = bool(")
+ assert (
+ "_effective_tensor_parallel(attempt_extra_args, tensor_parallel)" in src[pres : pres + 200]
+ )
+ # not the toggle-only form this replaced
+ assert (
+ "bool(\n request.tensor_parallel and not tensor_parallel" not in src
+ )
+
+
+def test_carry_preserved_tensor_intent_truth_table():
+ """Behavioral check of the carry-forward decision: carried only for the SAME
+ model, preserved, and not an explicit drop. Catches a `not` inversion (ctx-only
+ collapse) and a missing same-model guard (cross-model leak) (#6659)."""
+ inference_routes = _load_inference_routes_module()
+ f = inference_routes._carry_preserved_tensor_intent
+ assert f(preserved = True, same_model = True, explicit_drop = False) is True
+ assert f(preserved = True, same_model = True, explicit_drop = True) is False # explicit drop
+ assert f(preserved = True, same_model = False, explicit_drop = False) is False # model switch
+ assert f(preserved = False, same_model = True, explicit_drop = False) is False # not a fallback
+
+
+def test_preserved_fallback_carried_across_non_drop_reload():
+ """The hint carries the preserved fallback via _carry_preserved_tensor_intent,
+ gated on the same model loaded, so a ctx-only reload keeps multi-GPU but a model
+ switch / explicit drop doesn't inherit it (#6659)."""
+ route = Path(_BACKEND_DIR) / "routes" / "inference.py"
+ src = route.read_text()
+ idx = src.find("_tensor_intent_overall = _effective_tensor_parallel(")
+ assert idx != -1
+ block = src[idx : idx + 400]
+ assert "_carry_preserved_tensor_intent(" in block
+ assert "preserved = llama_backend.layer_preserves_tensor_intent" in block
+ assert "same_model = _same_model_loaded" in block
+ assert "explicit_drop = _explicit_tensor_drop" in block
+
+
+def test_same_model_guard_checks_path_and_variant():
+ """The same-model guard matches the resolved config.identifier (what load_model
+ stores, after from_identifier normalizes shorthands) -- not the raw request id --
+ and also matches the loaded quant by path (local multi-variant dir) else variant (HF
+ repo), so a reload keeps the carry-forward and a different variant doesn't inherit
+ the prior one's preserved tensor intent (#6659)."""
+ route = Path(_BACKEND_DIR) / "routes" / "inference.py"
+ src = route.read_text()
+ idx = src.find("_same_model_loaded = (")
+ assert idx != -1
+ block = src[idx : idx + 1300]
+ # Identity compares the normalized config.identifier, not the raw model_identifier.
+ head = src[idx : idx + 200]
+ assert "config.identifier" in head and "== (model_identifier" not in head
+ assert "llama_backend.gguf_path" in block and "config.gguf_file" in block
+ assert "llama_backend.hf_variant" in block and "config.gguf_variant" in block
+
+
+def test_diffusion_load_clears_preserved_tensor_flag():
+ """The diffusion early-return path (skips the command builder) clears the
+ preserved-fallback flag, so a prior tensor fallback doesn't churn it (#6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ diff = src.find("if self._is_diffusion:")
+ assert diff != -1
+ start = src.find("return self._start_diffusion_server", diff)
+ assert start != -1
+ assert "self._layer_preserves_tensor_intent = False" in src[diff:start]
+
+
+def test_is_tensor_split_assert_marker():
+ """Matches the specific #6415 split-axis assert, not any ggml assert/abort, so
+ an unrelated invariant a corrupt GGUF/projector trips isn't cached (#6659)."""
+ f = LlamaCppBackend._is_tensor_split_assert
+ # the real #6415 warmup assert (split-axis enum, in ggml-backend-meta)
+ assert (
+ f(
+ "ggml-backend-meta.cpp:541: GGML_ASSERT(src_ss[0].axis != "
+ "GGML_BACKEND_SPLIT_AXIS_0) failed"
+ )
+ is True
+ )
+ # the split-axis token alone (file path elided / reworded) still matches
+ assert f("GGML_ASSERT(x.axis != GGML_BACKEND_SPLIT_AXIS_1) failed") is True
+ # UNRELATED asserts must NOT match -- including a different invariant from the
+ # same multi-assert source file (matched on the token, not the file name).
+ assert f("ggml-backend-meta.cpp:99: GGML_ASSERT(buf != NULL) failed") is False
+ assert f("/x/ggml.c:1234: GGML_ASSERT(ne == 1) failed") is False
+ assert f("ggml_abort: something else entirely") is False
+ assert f("Segmentation fault (core dumped)") is False
+ assert f("") is False
+ assert f(None) is False
+
+
+def test_layer_preserve_hint_replayed_on_respawn():
+ """The preserve hint is in the replay snapshot (_pending_load_kwargs), so a
+ respawn keeps the downgraded model multi-GPU (Codex review on #6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ pend = src.find("_pending_load_kwargs = {")
+ assert pend != -1
+ block = src[pend : src.find("}", pend) + 1]
+ assert '"preserve_multi_gpu_on_layer": preserve_multi_gpu_on_layer' in block, (
+ "the layer-preserve hint must be in the replay snapshot so _respawn_if_dead "
+ "keeps the multi-GPU placement"
+ )
+
+
+def test_should_record_tensor_split_abort_decision():
+ """Behavioral check of marker AND (signal crash OR Windows abort), so an
+ or->and typo or caching a generic crash fails here, not just the source pins."""
+ f = LlamaCppBackend._should_record_tensor_split_abort
+ marker = "ggml-backend-meta.cpp:541: GGML_ASSERT(x.axis != GGML_BACKEND_SPLIT_AXIS_0) failed"
+ # marker + a hard crash records, across every platform's abort encoding
+ assert f(-6, marker) is True # POSIX SIGABRT
+ assert f(-11, marker) is True # POSIX SIGSEGV
+ assert f(3, marker) is True # Windows CRT abort() exit (not a signal)
+ assert f(0xC0000005, marker) is True # Windows NTSTATUS access violation
+ # marker present but no hard crash -> not recorded
+ assert f(0, marker) is False # clean exit
+ assert f(-9, marker) is False # SIGKILL (OOM / unload), not a fault
+ assert f(None, marker) is False # still running
+ # hard crash but not the split-axis marker -> not recorded (no over-caching)
+ assert f(3, "some other failure") is False
+ assert f(-6, "GGML_ASSERT(buf != NULL) failed") is False
+ assert f(0xC0000005, "") is False
+
+
+def test_fit_off_retry_skipped_on_split_axis_abort():
+ """The fit-independent --fit off retry is skipped on the split-axis marker, else
+ the model crashes a second time before the latch records it (reviewer.py, #6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ retry = src.find('run_cmd = [*run_cmd, "--fit", "off"]')
+ assert retry != -1
+ guard = src[max(0, retry - 1000) : retry]
+ assert "_fit_retry_allowed" in guard and "_startup_crashed" in guard
+ assert (
+ "not _split_axis_crash" in guard
+ ), "the fit-off retry must be skipped when the crash is a split-axis abort"
+
+
+def test_is_abort_exit_recognizes_windows_crt_abort():
+ """exit code 3 (MSVC abort()) counts as a crash; signals / clean exits do not."""
+ f = LlamaCppBackend._is_abort_exit
+ assert f(3) is True
+ assert f(0) is False
+ assert f(-6) is False # POSIX SIGABRT is handled by _is_signal_crash, not here
+ assert f(None) is False
+
+
+# ── tensor-off after a multi-GPU fallback forces a reload (route dedup) ─
+
+
+class _NoopProcess:
+ """Stand-in for Popen so is_loaded is True and atexit cleanup doesn't crash."""
+
+ def terminate(self):
+ pass
+
+ def wait(self, timeout = None):
+ return 0
+
+ def kill(self):
+ pass
+
+ def poll(self):
+ return 0
+
+
+def _fallback_loaded_backend(layer_preserves_tensor_intent: bool) -> LlamaCppBackend:
+ """A loaded backend in the tensor->layer fallback state (tensor off, --split-mode
+ layer stored), differing only in the preserved-multi-GPU flag."""
+ b = LlamaCppBackend()
+ b._model_identifier = "owner/repo"
+ b._requested_n_ctx = 0
+ b._cache_type_kv = None
+ b._tensor_parallel = False
+ b._layer_preserves_tensor_intent = layer_preserves_tensor_intent
+ b._extra_args = ["--split-mode", "layer"]
+ b._requested_spec_mode = "auto"
+ b._chat_template_override = None
+ b._gguf_path = None
+ return b
+
+
+def test_tensor_off_echo_preserves_multi_gpu_fallback():
+ """The Studio UI always sends tensor_parallel and echoes the /load response's
+ resolved value, so after a fallback a ctx/settings reload carries tensor_parallel=
+ false even though the user never changed it. That echo must NOT collapse the
+ preserved multi-GPU placement -- it dedupes (Codex #6659)."""
+ from models.inference import LoadRequest
+
+ inference_routes = _load_inference_routes_module()
+
+ req = LoadRequest(model_path = "owner/repo", tensor_parallel = False)
+ assert "tensor_parallel" in req.model_fields_set, "the UI always sends the field"
+
+ # Preserved fallback + bare tensor=false echo: dedupe, keep multi-GPU (no collapse).
+ assert (
+ inference_routes._request_matches_loaded_settings(
+ req, _fallback_loaded_backend(layer_preserves_tensor_intent = True)
+ )
+ is True
+ )
+ # A genuine layer load (no preserved intent): tensor-off also dedupes, no churn.
+ assert (
+ inference_routes._request_matches_loaded_settings(
+ req, _fallback_loaded_backend(layer_preserves_tensor_intent = False)
+ )
+ is True
+ )
+
+
+def test_explicit_split_mode_layer_extras_reloads_after_multi_gpu_fallback():
+ """Tensor intent can be dropped via extras too: an explicit --split-mode layer
+ matches the stored fallback extras but must still reload (reviewer.py P1, #6659)."""
+ from models.inference import LoadRequest
+
+ inference_routes = _load_inference_routes_module()
+
+ req = LoadRequest(model_path = "owner/repo", llama_extra_args = ["--split-mode", "layer"])
+ assert "llama_extra_args" in req.model_fields_set
+ assert (
+ inference_routes._request_matches_loaded_settings(
+ req, _fallback_loaded_backend(layer_preserves_tensor_intent = True)
+ )
+ is False
+ )
+
+
+def test_tensor_off_reload_requires_explicit_toggle():
+ """An Apply that doesn't touch the toggle (e.g. a context change) isn't churned
+ by the preserved-fallback reload -- the working server is kept (Codex #6659)."""
+ from models.inference import LoadRequest
+
+ inference_routes = _load_inference_routes_module()
+
+ req = LoadRequest(model_path = "owner/repo") # tensor_parallel left unset
+ assert "tensor_parallel" not in req.model_fields_set
+ assert (
+ inference_routes._request_matches_loaded_settings(
+ req, _fallback_loaded_backend(layer_preserves_tensor_intent = True)
+ )
+ is True
+ )
+
+
+def test_tensor_off_under_env_tensor_does_not_reload_loop(monkeypatch):
+ """With LLAMA_ARG_SPLIT_MODE=tensor set, a tensor-off request can't drop tensor
+ intent, so the env-aware guard dedupes instead of reload-looping (Codex #6659)."""
+ from models.inference import LoadRequest
+
+ inference_routes = _load_inference_routes_module()
+ monkeypatch.setenv("LLAMA_ARG_SPLIT_MODE", "tensor")
+
+ req = LoadRequest(model_path = "owner/repo", tensor_parallel = False)
+ assert "tensor_parallel" in req.model_fields_set
+ # env still forces tensor -> not a real drop -> dedupe (no reload loop).
+ assert (
+ inference_routes._request_matches_loaded_settings(
+ req, _fallback_loaded_backend(layer_preserves_tensor_intent = True)
+ )
+ is True
+ )
+
+
+def test_is_explicit_tensor_drop_truth_table():
+ """Only an explicit non-tensor --split-mode override is a drop. A bare
+ tensor_parallel field (the UI always sends it and echoes the fallback's false), an
+ empty clear, an unrelated extra (--top-k), or inherit (None) must NOT collapse a
+ preserved fallback; --split-mode tensor / tensor_parallel=true re-engage (Codex
+ #6659)."""
+ from models.inference import LoadRequest
+
+ f = _load_inference_routes_module()._is_explicit_tensor_drop
+ # A non-tensor split-mode override is the one deliberate departure -> drop.
+ assert (
+ f(LoadRequest(model_path = "owner/repo", llama_extra_args = ["--split-mode", "layer"])) is True
+ )
+ # tensor / retry re-engages, never a drop.
+ assert (
+ f(LoadRequest(model_path = "owner/repo", llama_extra_args = ["--split-mode", "tensor"]))
+ is False
+ )
+ # A bare tensor_parallel field is the UI echo, not a drop (would collapse on reload).
+ assert f(LoadRequest(model_path = "owner/repo", tensor_parallel = False)) is False
+ assert f(LoadRequest(model_path = "owner/repo", tensor_parallel = True)) is False
+ # Unrelated extra / empty clear / inherit all keep the preserved placement.
+ assert f(LoadRequest(model_path = "owner/repo", llama_extra_args = ["--top-k", "20"])) is False
+ assert f(LoadRequest(model_path = "owner/repo", llama_extra_args = [])) is False
+ assert f(LoadRequest(model_path = "owner/repo")) is False
+
+
+def test_explicit_tensor_drop_uses_shared_helper_in_both_readers():
+ """Both the already-loaded dedup and the load carry-forward derive the drop from
+ _is_explicit_tensor_drop, so they agree on what counts as a drop -- a reload for
+ an unrelated extra still carries the preserved intent rather than collapsing to one
+ GPU (Codex #6659)."""
+ src = (Path(_BACKEND_DIR) / "routes" / "inference.py").read_text()
+ # Dedup reader (the preserved-fallback reload guard).
+ assert "layer_preserves_tensor_intent and _is_explicit_tensor_drop(request)" in src
+ # Load carry-forward reader feeds the same decision into the carry-forward.
+ assert "_explicit_tensor_drop = _is_explicit_tensor_drop(request)" in src
+
+
+def test_layer_preserves_tensor_intent_set_only_on_preserved_downgrade():
+ """load_model latches the flag from _layer_min_gpus (raised only when a tensor
+ request is downgraded but kept multi-GPU), and clears it when tensor stays on."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ on = src.find("self._tensor_parallel = True")
+ off = src.find("self._tensor_parallel = False")
+ assert 0 <= on and 0 <= off
+ assert "self._layer_preserves_tensor_intent = False" in src[on : on + 120]
+ assert "self._layer_preserves_tensor_intent = _layer_min_gpus > 1" in src[off : off + 400]
+
+
+def test_layer_min_gpus_bound_before_gpu_selection_try():
+ """_layer_min_gpus is bound before the GPU-selection try, so the --fit-on except
+ path can't UnboundLocalError when the command builder reads it (Codex #6659)."""
+ src = inspect.getsource(LlamaCppBackend.load_model)
+ assert src.count("_layer_min_gpus = 1") == 1, "exactly one init, before the try"
+ init = src.find("_layer_min_gpus = 1")
+ try_body = src.find("gguf_size = self._get_gguf_size_bytes")
+ fit_except = src.find("GPU selection failed")
+ use_after = src.find("self._layer_preserves_tensor_intent = _layer_min_gpus > 1")
+ assert (
+ -1 < init < try_body < fit_except < use_after
+ ), "the init must precede the try body, the except, and the command-builder use"
+
+
+def test_already_in_target_state_reloads_on_tensor_off_after_fallback():
+ """The backend fast path mirrors the route dedup: a preserved fallback reloads on
+ an EXPLICIT tensor-off request, but an implicit same-settings reload (carry-forward
+ preserve_multi_gpu_on_layer=True) still dedupes (Codex #6659)."""
+
+ def _backend(layer_preserves: bool) -> LlamaCppBackend:
+ b = _fallback_loaded_backend(layer_preserves_tensor_intent = layer_preserves)
+ b._process = _NoopProcess()
+ b._healthy = True
+ return b
+
+ kwargs = dict(
+ gguf_path = None,
+ mtp_draft_path = None,
+ model_identifier = "owner/repo",
+ hf_variant = None,
+ n_ctx = 0,
+ cache_type_kv = None,
+ speculative_type = None,
+ spec_draft_n_max = None,
+ tensor_parallel = False,
+ chat_template_override = None,
+ extra_args = ["--split-mode", "layer"],
+ is_vision = False,
+ )
+ # Preserved fallback + EXPLICIT tensor drop -> reload (not already in target state).
+ assert _backend(True)._already_in_target_state(**kwargs) is False
+ # Same preserved fallback but an implicit reload that carries the intent forward
+ # (HF auto-pick / local-dir flows skip the route guard and reach here) -> dedupe.
+ assert (
+ _backend(True)._already_in_target_state(**kwargs, preserve_multi_gpu_on_layer = True) is True
+ )
+ # A genuine layer load (no preserved intent) -> dedupe, no churn.
+ assert _backend(False)._already_in_target_state(**kwargs) is True
diff --git a/studio/backend/tests/test_training_runs.py b/studio/backend/tests/test_training_runs.py
new file mode 100644
index 0000000000..fd0d6d380f
--- /dev/null
+++ b/studio/backend/tests/test_training_runs.py
@@ -0,0 +1,103 @@
+# 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 json
+
+from storage.studio_db import _extract_project_name_from_config_json
+from utils.training_runs import (
+ build_default_output_dir_name,
+ model_segment_from_default_output_dir_name,
+ normalize_project_name,
+ slugify_project_name,
+)
+
+
+def test_normalize_project_name_trims_and_collapses_whitespace():
+ assert normalize_project_name(" Customer Support LoRA ") == "Customer Support LoRA"
+
+
+def test_normalize_project_name_returns_none_for_empty_or_invalid_values():
+ assert normalize_project_name(" ") is None
+ assert normalize_project_name(None) is None
+
+
+def test_slugify_project_name_makes_safe_suffix():
+ assert slugify_project_name("Customer Support / LoRA v2") == "customer-support-lora-v2"
+
+
+def test_slugify_project_name_rejects_path_only_or_separator_only_values():
+ assert slugify_project_name("..") is None
+ assert slugify_project_name("///") is None
+
+
+def test_build_default_output_dir_name_appends_project_slug():
+ output_dir = build_default_output_dir_name(
+ "unsloth/Llama-3.2-3B-Instruct",
+ "Customer Support",
+ timestamp = 1771227800,
+ )
+
+ assert output_dir == "unsloth_Llama-3.2-3B-Instruct__project-customer-support_1771227800"
+
+
+def test_build_default_output_dir_name_caps_final_component(tmp_path):
+ output_dir = build_default_output_dir_name(
+ "a" * 240,
+ "b" * 80,
+ timestamp = 1771227800,
+ )
+
+ assert len(output_dir.encode()) <= 255
+ (tmp_path / output_dir).mkdir()
+
+
+def test_build_default_output_dir_name_skips_invalid_project_slug():
+ output_dir = build_default_output_dir_name(
+ "unsloth/Llama-3.2-3B-Instruct",
+ "..",
+ timestamp = 1771227800,
+ )
+
+ assert output_dir == "unsloth_Llama-3.2-3B-Instruct_1771227800"
+
+
+def test_model_segment_from_default_output_dir_name_strips_project_slug():
+ assert (
+ model_segment_from_default_output_dir_name(
+ "unsloth_Llama-3.2-3B-Instruct__project-customer-support_1771227800"
+ )
+ == "unsloth_Llama-3.2-3B-Instruct"
+ )
+
+
+def test_model_segment_preserves_project_marker_text_in_model_name():
+ output_dir = build_default_output_dir_name(
+ "org/foo__project-bar",
+ timestamp = 1771227800,
+ )
+
+ assert output_dir == "org_foo__project--bar_1771227800"
+ assert model_segment_from_default_output_dir_name(output_dir) == "org_foo__project-bar"
+
+
+def test_model_segment_strips_project_slug_after_escaped_model_marker():
+ output_dir = build_default_output_dir_name(
+ "org/foo__project-bar",
+ "Customer Support",
+ timestamp = 1771227800,
+ )
+
+ assert output_dir == "org_foo__project--bar__project-customer-support_1771227800"
+ assert model_segment_from_default_output_dir_name(output_dir) == "org_foo__project-bar"
+
+
+def test_extract_project_name_from_config_json_returns_normalized_name():
+ config_json = json.dumps({"project_name": " Sales Assistant "})
+
+ assert _extract_project_name_from_config_json(config_json) == "Sales Assistant"
+
+
+def test_extract_project_name_from_config_json_handles_missing_or_invalid_payload():
+ assert _extract_project_name_from_config_json(None) is None
+ assert _extract_project_name_from_config_json("not-json") is None
+ assert _extract_project_name_from_config_json(json.dumps({"project_name": " "})) is None
diff --git a/studio/backend/tests/test_training_streaming.py b/studio/backend/tests/test_training_streaming.py
index 8ff016d3bf..70b2d6fdcc 100644
--- a/studio/backend/tests/test_training_streaming.py
+++ b/studio/backend/tests/test_training_streaming.py
@@ -195,6 +195,16 @@ def test_hf_dataset_rejects_unsafe_values(bad_hf_dataset):
)
+def test_project_name_rejects_values_over_ui_limit():
+ with pytest.raises(ValidationError):
+ TrainingStartRequest(
+ model_name = "unsloth/test",
+ project_name = "x" * 81,
+ training_type = "LoRA/QLoRA",
+ format_type = "alpaca",
+ )
+
+
# --- Start-route streaming compatibility guards ---
diff --git a/studio/backend/utils/hf_xet_fallback.py b/studio/backend/utils/hf_xet_fallback.py
index a6ba69fffc..15961ac03a 100644
--- a/studio/backend/utils/hf_xet_fallback.py
+++ b/studio/backend/utils/hf_xet_fallback.py
@@ -161,6 +161,7 @@ def _download_child_entry(
repo_type: str,
disable_xet: bool,
result_queue: Any,
+ force_download: bool = False,
) -> None:
"""Spawn-child entrypoint: download one file and report the result.
@@ -211,6 +212,7 @@ def _download_child_entry(
filename = filename,
repo_type = repo_type,
token = token,
+ force_download = force_download,
)
result_queue.put({"ok": True, "path": path})
except BaseException as e: # noqa: BLE001 - report every failure to the parent
@@ -264,6 +266,7 @@ def _run_download_attempt(
interval: float,
grace_period: float,
on_status: Optional[Callable[[str], None]],
+ force_download: bool = False,
) -> tuple[str, Optional[str]]:
"""Run one download in a spawn child supervised by the no-progress watchdog.
@@ -280,6 +283,7 @@ def _run_download_attempt(
repo_type = repo_type,
disable_xet = disable_xet,
result_queue = result_queue,
+ force_download = force_download,
),
daemon = True,
)
@@ -345,21 +349,28 @@ def hf_hub_download_with_xet_fallback(
interval: float = DEFAULT_HEARTBEAT_INTERVAL,
grace_period: float = DEFAULT_GRACE_PERIOD,
on_status: Optional[Callable[[str], None]] = None,
+ force_download: bool = False,
) -> str:
"""Download a single file with Xet primary and HTTP as a stall-only fallback.
Returns the local cache path. Raises ``RuntimeError("Cancelled")`` if
*cancel_event* is set, re-raises a deterministic child error unchanged (no
fallback), and raises ``DownloadStallError`` only if BOTH transports stall.
+
+ When *force_download* is True the cache-first early-return is skipped and the
+ flag is threaded to ``hf_hub_download`` so a newer remote blob is re-fetched
+ even if an older blob is already cached.
"""
# Finalized blob already cached: return it with no child and no network.
- try:
- from huggingface_hub import try_to_load_from_cache
- cached = try_to_load_from_cache(repo_id, filename, repo_type = repo_type)
- if isinstance(cached, str) and os.path.exists(cached):
- return cached
- except Exception as e:
- logger.debug("Cached probe failed for %s/%s: %s", repo_id, filename, e)
+ # Skipped when force_download is set so an update re-fetches a newer blob.
+ if not force_download:
+ try:
+ from huggingface_hub import try_to_load_from_cache
+ cached = try_to_load_from_cache(repo_id, filename, repo_type = repo_type)
+ if isinstance(cached, str) and os.path.exists(cached):
+ return cached
+ except Exception as e:
+ logger.debug("Cached probe failed for %s/%s: %s", repo_id, filename, e)
if cancel_event is not None and cancel_event.is_set():
raise RuntimeError("Cancelled")
@@ -386,6 +397,7 @@ def hf_hub_download_with_xet_fallback(
interval = interval,
grace_period = grace_period,
on_status = on_status,
+ force_download = force_download,
)
if kind == "ok":
diff --git a/studio/backend/utils/models/checkpoints.py b/studio/backend/utils/models/checkpoints.py
index d174f6677b..90e26d45d0 100644
--- a/studio/backend/utils/models/checkpoints.py
+++ b/studio/backend/utils/models/checkpoints.py
@@ -9,6 +9,12 @@ import structlog
from loggers import get_logger
from pathlib import Path
from typing import List, Optional, Tuple
+from storage.studio_db import get_connection
+from utils.training_runs import (
+ build_default_output_dir_name,
+ extract_project_name,
+ model_segment_from_default_output_dir_name,
+)
from utils.paths import outputs_root, resolve_output_dir
logger = get_logger(__name__)
@@ -30,6 +36,93 @@ def _checkpoint_sort_key(checkpoint_path: Path) -> tuple[int, int, str]:
return (1, 0, str(checkpoint_path))
+def _infer_base_model_from_history(checkpoint_dir: Path) -> Optional[str]:
+ """Best-effort base-model lookup using persisted Studio run metadata."""
+ checkpoint_name = checkpoint_dir.name
+ resolved_checkpoint_dir = str(checkpoint_dir.resolve())
+
+ try:
+ conn = get_connection()
+ except Exception:
+ return None
+
+ try:
+ exact_rows = conn.execute(
+ """
+ SELECT model_name
+ FROM training_runs
+ WHERE output_dir IN (?, ?)
+ ORDER BY started_at DESC
+ """,
+ (
+ resolved_checkpoint_dir,
+ str(checkpoint_dir),
+ ),
+ ).fetchall()
+ for row in exact_rows:
+ model_name = row["model_name"]
+ if model_name:
+ return model_name
+
+ suffix_rows = conn.execute(
+ """
+ SELECT model_name, output_dir
+ FROM training_runs
+ WHERE output_dir IS NOT NULL
+ ORDER BY started_at DESC
+ """
+ ).fetchall()
+ for row in suffix_rows:
+ output_dir = str(row["output_dir"] or "").rstrip("/\\")
+ if not (
+ output_dir.endswith(f"/{checkpoint_name}")
+ or output_dir.endswith(f"\\{checkpoint_name}")
+ ):
+ continue
+ model_name = row["model_name"]
+ if model_name:
+ return model_name
+
+ parts = checkpoint_name.rsplit("_", 1)
+ if len(parts) != 2 or not parts[1].isdigit():
+ return None
+
+ timestamp = int(parts[1])
+ generated_rows = conn.execute(
+ """
+ SELECT model_name, config_json
+ FROM training_runs
+ ORDER BY started_at DESC
+ """
+ ).fetchall()
+ for row in generated_rows:
+ model_name = row["model_name"]
+ if not model_name:
+ continue
+
+ project_name = None
+ config_json = row["config_json"]
+ if config_json:
+ try:
+ project_name = extract_project_name(json.loads(config_json))
+ except (TypeError, json.JSONDecodeError):
+ project_name = None
+
+ expected_dir_name = build_default_output_dir_name(
+ model_name,
+ project_name,
+ timestamp = timestamp,
+ )
+ if expected_dir_name == checkpoint_name:
+ return model_name
+ except Exception:
+ return None
+ finally:
+ conn.close()
+
+ return None
+
+
def _read_checkpoint_loss(checkpoint_path: Path) -> Optional[float]:
"""Read loss from the last log_history entry of trainer_state.json, or None."""
trainer_state = checkpoint_path / "trainer_state.json"
@@ -106,9 +199,11 @@ def scan_checkpoints(
# Fallback: extract base model name from the folder name, e.g.
# "unsloth_Llama-3.2-3B-Instruct_1771227800" → "unsloth/Llama-3.2-3B-Instruct"
if not metadata.get("base_model"):
- parts = item.name.rsplit("_", 1)
- if len(parts) == 2 and parts[1].isdigit():
- name_part = parts[0]
+ metadata["base_model"] = _infer_base_model_from_history(item)
+
+ if not metadata.get("base_model"):
+ name_part = model_segment_from_default_output_dir_name(item.name)
+ if name_part:
idx = name_part.find("_")
if idx > 0:
metadata["base_model"] = name_part[:idx] + "/" + name_part[idx + 1 :]
diff --git a/studio/backend/utils/training_runs.py b/studio/backend/utils/training_runs.py
new file mode 100644
index 0000000000..dc2535e570
--- /dev/null
+++ b/studio/backend/utils/training_runs.py
@@ -0,0 +1,104 @@
+# SPDX-License-Identifier: AGPL-3.0-only
+# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
+
+"""Helpers for naming and describing Studio training runs."""
+
+from __future__ import annotations
+
+import re
+import time
+from typing import Any, Optional
+
+_INVALID_SEGMENT_CHARS = re.compile(r"[^A-Za-z0-9._-]+")
+_MAX_RUN_DIR_NAME_CHARS = 255
+_PROJECT_MARKER = "__project-"
+_PROJECT_MARKER_ESCAPE = f"{_PROJECT_MARKER}-"
+
+
+def _trim_segment(segment: str, max_chars: int) -> str:
+ if max_chars <= 0:
+ return ""
+ return segment[:max_chars].strip("._-")
+
+
+def _escape_project_marker(segment: str) -> str:
+ return segment.replace(_PROJECT_MARKER, _PROJECT_MARKER_ESCAPE)
+
+
+def _unescape_project_marker(segment: str) -> str:
+ return segment.replace(_PROJECT_MARKER_ESCAPE, _PROJECT_MARKER)
+
+
+def _appended_project_marker_index(segment: str) -> int:
+ marker_index = segment.rfind(_PROJECT_MARKER)
+ while marker_index >= 0 and segment.startswith(_PROJECT_MARKER_ESCAPE, marker_index):
+ marker_index = segment.rfind(_PROJECT_MARKER, 0, marker_index)
+ return marker_index
+
+
+def normalize_project_name(project_name: Any) -> Optional[str]:
+ """Return a trimmed project name, or None when empty/invalid."""
+ if not isinstance(project_name, str):
+ return None
+ normalized = " ".join(project_name.strip().split())
+ return normalized or None
+
+
+def slugify_project_name(project_name: Any) -> Optional[str]:
+ """Convert a project name into a filesystem-safe suffix."""
+ normalized = normalize_project_name(project_name)
+ if normalized is None:
+ return None
+
+ slug = _INVALID_SEGMENT_CHARS.sub("-", normalized).strip("-._")
+ if not slug:
+ return None
+ return slug.lower()
+
+
+def build_default_output_dir_name(
+ model_name: str,
+ project_name: Any = None,
+ *,
+ timestamp: Optional[int] = None,
+) -> str:
+ """Build the default training output folder name."""
+ from utils.paths import default_run_dir_name
+
+ timestamp_part = str(int(time.time() if timestamp is None else timestamp))
+ timestamp_suffix = f"_{timestamp_part}"
+ model_segment = _escape_project_marker(default_run_dir_name(model_name))
+ project_slug = slugify_project_name(project_name)
+ if not project_slug:
+ max_model_chars = _MAX_RUN_DIR_NAME_CHARS - len(timestamp_suffix)
+ model_segment = _trim_segment(model_segment, max_model_chars) or "model"
+ return f"{model_segment}{timestamp_suffix}"
+
+ max_project_chars = (
+ _MAX_RUN_DIR_NAME_CHARS - len("model") - len(_PROJECT_MARKER) - len(timestamp_suffix)
+ )
+ project_slug = _trim_segment(project_slug, max_project_chars) or "project"
+ project_suffix = f"{_PROJECT_MARKER}{project_slug}{timestamp_suffix}"
+ max_model_chars = _MAX_RUN_DIR_NAME_CHARS - len(project_suffix)
+ model_segment = _trim_segment(model_segment, max_model_chars) or "model"
+ return f"{model_segment}{project_suffix}"
+
+
+def model_segment_from_default_output_dir_name(output_dir_name: str) -> Optional[str]:
+ """Return the encoded model segment from a default run folder name."""
+ parts = str(output_dir_name or "").rsplit("_", 1)
+ if len(parts) != 2 or not parts[1].isdigit():
+ return None
+ model_segment = parts[0]
+ marker_index = _appended_project_marker_index(model_segment)
+ if marker_index >= 0:
+ model_segment = model_segment[:marker_index]
+ model_segment = _unescape_project_marker(model_segment)
+ return model_segment or None
+
+
+def extract_project_name(config: Any) -> Optional[str]:
+ """Read and normalize a project name from a stored config dict."""
+ if not isinstance(config, dict):
+ return None
+ return normalize_project_name(config.get("project_name"))
diff --git a/studio/frontend/package-lock.json b/studio/frontend/package-lock.json
index 6202d3ce22..80db64553a 100644
--- a/studio/frontend/package-lock.json
+++ b/studio/frontend/package-lock.json
@@ -1704,6 +1704,7 @@
"os": [
"android"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1724,6 +1725,7 @@
"os": [
"darwin"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1744,6 +1746,7 @@
"os": [
"darwin"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1764,6 +1767,7 @@
"os": [
"linux"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1784,6 +1788,7 @@
"os": [
"linux"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1804,6 +1809,7 @@
"os": [
"linux"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1824,6 +1830,7 @@
"os": [
"linux"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1844,6 +1851,7 @@
"os": [
"linux"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1864,6 +1872,7 @@
"os": [
"linux"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1884,6 +1893,7 @@
"os": [
"win32"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -1904,6 +1914,7 @@
"os": [
"win32"
],
+ "peer": true,
"engines": {
"node": ">= 10"
},
@@ -5669,9 +5680,6 @@
"cpu": [
"arm64"
],
- "libc": [
- "glibc"
- ],
"license": "MIT",
"optional": true,
"os": [
@@ -5688,9 +5696,6 @@
"cpu": [
"arm64"
],
- "libc": [
- "musl"
- ],
"license": "MIT",
"optional": true,
"os": [
@@ -5707,9 +5712,6 @@
"cpu": [
"ppc64"
],
- "libc": [
- "glibc"
- ],
"license": "MIT",
"optional": true,
"os": [
@@ -5726,9 +5728,6 @@
"cpu": [
"s390x"
],
- "libc": [
- "glibc"
- ],
"license": "MIT",
"optional": true,
"os": [
@@ -5745,9 +5744,6 @@
"cpu": [
"x64"
],
- "libc": [
- "glibc"
- ],
"license": "MIT",
"optional": true,
"os": [
@@ -5764,9 +5760,6 @@
"cpu": [
"x64"
],
- "libc": [
- "musl"
- ],
"license": "MIT",
"optional": true,
"os": [
@@ -10282,9 +10275,9 @@
}
},
"node_modules/hono": {
- "version": "4.12.21",
- "resolved": "https://registry.npmjs.org/hono/-/hono-4.12.21.tgz",
- "integrity": "sha512-uV63apnb0kyPtAUwoWgaGh9HyIFcv8lgmzPZSiTBQAFOFGIzka5EZ1dZocmGnn0XdX0+XTqJ6Tqv7selMuGLRQ==",
+ "version": "4.12.25",
+ "resolved": "https://registry.npmjs.org/hono/-/hono-4.12.25.tgz",
+ "integrity": "sha512-2NFaIyNVgJmBs/ecmtGzlmluTFs5cHEWGTdu0t1HBwYzoGXOL5nUQBRMXsXWla5i4KkG//QMzVP88m1+I3fdAQ==",
"license": "MIT",
"engines": {
"node": ">=16.9.0"
diff --git a/studio/frontend/package.json b/studio/frontend/package.json
index 0956c710f5..a2eddecda3 100644
--- a/studio/frontend/package.json
+++ b/studio/frontend/package.json
@@ -86,7 +86,7 @@
"@tanstack/router-core": "1.169.2",
"@tanstack/history": "1.161.6",
"mermaid": "11.15.0",
- "hono": "4.12.21",
+ "hono": "4.12.25",
"qs": "6.15.2",
"ip-address": "10.1.1",
"brace-expansion@5.0.5": "5.0.6"
diff --git a/studio/frontend/src/app/provider.tsx b/studio/frontend/src/app/provider.tsx
index 914abbbf1d..176665769d 100644
--- a/studio/frontend/src/app/provider.tsx
+++ b/studio/frontend/src/app/provider.tsx
@@ -354,12 +354,11 @@ function TauriWrapper({ children }: { children: ReactNode }) {
);
}
- const showApp = status === "running" && desktopAuthReady;
+ const showApp = status === "running";
+ const desktopBooting = status === "running" && !desktopAuthReady;
+ const showInteractiveApp = showApp && desktopAuthReady;
const startupStatus = status === "running" ? "starting" : status;
- const startupProgressDetail =
- status === "running" && !desktopAuthReady
- ? "Signing in to desktop session..."
- : progressDetail;
+ const startupProgressDetail = progressDetail;
const usesCustomTitlebar = shouldUseCustomWindowTitlebar();
const usesNativeMacTitlebar = shouldUseNativeMacWindowTitlebar();
const hidesTitlebarSidebar = HIDDEN_TITLEBAR_SIDEBAR_ROUTES.has(pathname);
@@ -369,12 +368,23 @@ function TauriWrapper({ children }: { children: ReactNode }) {
-
+ {showInteractiveApp ? : null}
-
- {children}
+ {showInteractiveApp ? : null}
+ {showInteractiveApp ? children : null}
+ {desktopBooting ? (
+
+
+
Preparing Studio
+
The local backend is ready. Signing in to your desktop session before loading chats.
+
+
+ Signing in to desktop session...
+
+
+ ) : null}
>
) : (
{
@@ -593,7 +578,7 @@ export function AppSidebar() {
setRenamingTarget({ kind: "chat", item, current: item.title });
}
function openRenameRun(run: TrainingRunSummary) {
- const current = run.display_name ?? run.model_name;
+ const current = getTrainingRunDisplayTitle(run);
setRenameDraft(current);
setRenamingTarget({ kind: "run", run, current });
}
@@ -1390,7 +1375,7 @@ export function AppSidebar() {
aria-hidden
/>
- {run.display_name ?? run.model_name}
+ {getTrainingRunDisplayTitle(run)}
{formatRelativeShort(run.started_at)}
@@ -1666,8 +1651,7 @@ export function AppSidebar() {
renderEmphasizedTranslation(
t,
"shell.dialog.deleteRun.description",
- confirmingDelete.run.display_name ??
- confirmingDelete.run.model_name,
+ getTrainingRunDisplayTitle(confirmingDelete.run),
)
) : confirmingDelete?.kind === "chat" ? (
renderEmphasizedTranslation(
diff --git a/studio/frontend/src/components/assistant-ui/model-selector.tsx b/studio/frontend/src/components/assistant-ui/model-selector.tsx
index 71d75b68f5..65073258b9 100644
--- a/studio/frontend/src/components/assistant-ui/model-selector.tsx
+++ b/studio/frontend/src/components/assistant-ui/model-selector.tsx
@@ -14,6 +14,7 @@ import { isCustomProviderType } from "@/features/chat/external-providers";
import { ChevronDownStandardIcon } from "@/lib/chevron-icons";
import { cn } from "@/lib/utils";
import {
+ CheckmarkCircle02Icon,
CloudIcon,
DashboardSquare01Icon,
Download01Icon,
@@ -149,6 +150,7 @@ function ModelSelectorTrigger({
size = "default",
className,
dataTour,
+ onEject,
}: {
currentModel?: ModelOption;
isLoaded: boolean;
@@ -157,6 +159,7 @@ function ModelSelectorTrigger({
size?: "sm" | "default" | "lg";
className?: string;
dataTour?: string;
+ onEject?: () => void;
}) {
return (
@@ -164,12 +167,15 @@ function ModelSelectorTrigger({
type="button"
data-tour={dataTour}
className={cn(
- "unsloth-model-selector-trigger flex min-w-0 items-center gap-2 transition-colors",
+ "unsloth-model-selector-trigger group/trigger flex min-w-0 items-center gap-2 transition-colors",
+ // Suppress the pill's hover background while the eject hit area is
+ // hovered, so only the dot's own circle reacts.
variant === "outline" &&
- "rounded-full border border-border/60 hover:bg-[#ececec] dark:hover:bg-[#2d2e32]",
+ "rounded-full border border-border/60 hover:bg-[#ececec] has-[[data-eject-hit]:hover]:!bg-transparent dark:hover:bg-[#2d2e32]",
variant === "ghost" &&
- "rounded-full hover:bg-[#ececec] dark:hover:bg-[#2d2e32]",
- variant === "muted" && "rounded-full bg-muted hover:bg-muted/80",
+ "rounded-full hover:bg-[#ececec] has-[[data-eject-hit]:hover]:!bg-transparent dark:hover:bg-[#2d2e32]",
+ variant === "muted" &&
+ "rounded-full bg-muted hover:bg-muted/80 has-[[data-eject-hit]:hover]:!bg-muted",
// More left padding than right; the chevron is pulled close to the
// label (below) so the trigger reads balanced around the text.
size === "sm" && "h-8 pl-3 pr-1.5 text-xs",
@@ -178,9 +184,44 @@ function ModelSelectorTrigger({
className,
)}
>
- {isLoaded && (
-
- )}
+ {isLoaded &&
+ (onEject ? (
+ // Loaded status doubles as a mouse eject shortcut: green checkmark
+ // at rest, red eject icon on pill hover, click to eject. A plain
+ // span (no role/tabIndex) keeps it out of the trigger button's
+ // content model, which forbids focusable descendants. Keyboard and
+ // screen-reader users eject via the picker's "Eject model" button.
+ // aria-hidden marks it decorative; stopPropagation stops the
+ // popover from toggling. On touch (no hover) the eject icon and
+ // tooltip never reveal, so pointer-events-none disables the
+ // shortcut there and taps open the picker instead of ejecting.
+ event.stopPropagation()}
+ onClick={(event) => {
+ event.stopPropagation();
+ onEject();
+ }}
+ // Hit area larger than the icon, with a hover circle. Negative
+ // margin keeps the icon in the dot's original spot.
+ className="-m-1 flex size-5 shrink-0 cursor-pointer items-center justify-center rounded-full transition-colors hover:bg-black/10 dark:hover:bg-white/10 [@media(hover:none)]:pointer-events-none"
+ >
+
+
+
+ ) : (
+
+ ))}
{currentModel?.icon ? (
{currentModel.icon}
@@ -651,6 +692,7 @@ export function ModelSelector({
size={size}
className={className}
dataTour={triggerDataTour}
+ onEject={onEject ? handleEject : undefined}
/>
- {
if (!nextOpen && deleting) return;
setOpen(nextOpen);
}}
- >
-
-
- {title}
- {description}
-
-
- No
- {
- e.preventDefault();
- handleConfirm();
- }}
- >
- {deleting ? loadingLabel : "Yes"}
-
-
-
-
+ title={title}
+ description={description}
+ deleting={deleting}
+ onConfirm={() => void handleConfirm()}
+ />
>
);
}
diff --git a/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx b/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx
new file mode 100644
index 0000000000..d00c812325
--- /dev/null
+++ b/studio/frontend/src/components/assistant-ui/model-selector/model-update-action.tsx
@@ -0,0 +1,106 @@
+// 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 { subscribeJobListeners } from "@/features/hub/download-manager";
+import { UpdateConfirmDialog } from "@/features/hub/catalog/download-card";
+import { ggufVariantsMatch } from "@/features/hub/lib/model-identity";
+import { cn } from "@/lib/utils";
+import { RefreshCw } from "lucide-react";
+import { useCallback, useEffect, useRef, useState, type ReactNode } from "react";
+import { toast } from "sonner";
+
+interface ModelUpdateActionProps {
+ ariaLabel: string;
+ title: string;
+ description: ReactNode;
+ /** Repo + variant the update targets, so this action can refresh its caller
+ * (clearing the "update available" cue) when the matching managed download
+ * completes. `variant` is null for full-model (safetensors / MLX) rows. */
+ repoId: string;
+ variant?: string | null;
+ buttonClassName?: string;
+ iconClassName?: string;
+ disabled?: boolean;
+ /** Starts the update — which now runs as a managed download. Resolves once the
+ * download has been handed to the download manager, NOT when it finishes. */
+ onConfirm: () => Promise | void;
+ /** Fired when THIS repo+variant's managed update actually completes. */
+ onUpdated?: () => void;
+}
+
+export function ModelUpdateAction({
+ ariaLabel,
+ title,
+ description,
+ repoId,
+ variant = null,
+ buttonClassName,
+ iconClassName,
+ disabled = false,
+ onConfirm,
+ onUpdated,
+}: ModelUpdateActionProps) {
+ const [open, setOpen] = useState(false);
+
+ // The update is a managed download (it surfaces in the global Downloads panel
+ // with progress + cancel). When this exact repo+variant finishes, refresh the
+ // caller so the "update available" cue clears once the new revision is on
+ // disk. A ref keeps the subscription stable across renders without resubscribing.
+ const onUpdatedRef = useRef(onUpdated);
+ onUpdatedRef.current = onUpdated;
+ useEffect(() => {
+ return subscribeJobListeners("model", repoId, {
+ onComplete: (completedVariant) => {
+ const matches = variant
+ ? ggufVariantsMatch(completedVariant, variant)
+ : !completedVariant;
+ if (matches) onUpdatedRef.current?.();
+ },
+ });
+ }, [repoId, variant]);
+
+ const handleConfirm = useCallback(() => {
+ // Start the background re-download and close the dialog immediately; the
+ // Downloads panel owns progress + cancel from here. Only a failure to START
+ // surfaces a toast — a failed download reports itself in the panel.
+ void Promise.resolve()
+ .then(onConfirm)
+ .catch((err) => {
+ toast.error(
+ err instanceof Error ? err.message : "Failed to start update",
+ );
+ });
+ setOpen(false);
+ }, [onConfirm]);
+
+ return (
+ <>
+
+
+
+ >
+ );
+}
diff --git a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx
index cc54fd6810..622dd39d69 100644
--- a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx
+++ b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx
@@ -23,18 +23,19 @@ import {
listScanFolders,
removeScanFolder,
} from "@/features/chat/api/chat-api";
+import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store";
import type {
CachedGgufRepo,
CachedModelRepo,
LocalModelInfo,
} from "@/features/chat/api/chat-api";
-import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store";
import type { GgufVariantDetail } from "@/features/chat/types/api";
import { DotTag } from "@/features/hub/catalog/dot-tag";
import {
type HubOption,
HubOptionMenu,
} from "@/features/hub/catalog/hub-option-menu";
+import { TransportConflictDialog } from "@/features/hub/catalog/transport-conflict-dialog";
import { TrainIcon } from "@/features/hub/components/train-icon";
import { useHubInfiniteScroll } from "@/features/hub/hooks/use-hub-infinite-scroll";
import {
@@ -47,6 +48,11 @@ import { useOnlineStatus } from "@/features/hub/hooks/use-online-status";
import { isHiddenModelId } from "@/features/hub/lib/hidden-models";
import { classifyUnslothSupport } from "@/features/hub/lib/unsloth-support";
import { useHfTokenStore } from "@/features/hub/stores/hf-token-store";
+import {
+ downloadManager,
+ jobKeyOf,
+ useDownloadManagerStore,
+} from "@/features/hub/download-manager";
import { useDebouncedValue, useGpuInfo } from "@/hooks";
import { extractParamLabel } from "@/lib/model-size";
import { toast } from "@/lib/toast";
@@ -86,6 +92,7 @@ import {
hasAnyCapability,
} from "./model-capabilities";
import { ModelDeleteAction } from "./model-delete-action";
+import { ModelUpdateAction } from "./model-update-action";
import { ModelLoadSettingsAction } from "./model-load-settings-action";
import {
type ModelLoadTimes,
@@ -617,20 +624,26 @@ function normalizeGgufVariantsResponse(res: {
};
}
+function ggufVariantExpectedBytes(variant: GgufVariantDetail): number {
+ const downloadBytes = variant.download_size_bytes;
+ return typeof downloadBytes === "number" &&
+ Number.isFinite(downloadBytes) &&
+ downloadBytes > 0
+ ? downloadBytes
+ : variant.size_bytes;
+}
+
function GgufVariantExpander({
repoId,
onSelect,
gpuGb,
systemRamGb,
+ hfToken,
parentOptionKey,
onNavigatePastStart,
onNavigatePastEnd,
- onDeleteVariant,
sourceOverride,
- deleteVariantTitle = "Delete cached model?",
- renderDeleteVariantDescription,
- getDeleteVariantSuccessMessage,
- deleteDisabled = false,
+ variantActions,
onDevice = false,
onHasVision,
}: {
@@ -638,21 +651,42 @@ function GgufVariantExpander({
onSelect: (id: string, meta: ModelSelectorChangeMeta) => void;
gpuGb?: number;
systemRamGb?: number;
+ /** HF token threaded into the variant fetch so private/gated repos resolve
+ * their GGUF variants (and update badges). */
+ hfToken?: string;
parentOptionKey?: string;
onNavigatePastStart?: () => void;
onNavigatePastEnd?: () => void;
- onDeleteVariant?: (quant: string) => Promise | void;
sourceOverride?: ModelSelectorChangeMeta["source"];
- deleteVariantTitle?: string;
- renderDeleteVariantDescription?: (quant: string) => ReactNode;
- getDeleteVariantSuccessMessage?: (quant: string) => string;
- deleteDisabled?: boolean;
+ /** Update/delete actions for cached variant rows. Omitted by browse-only
+ * expanders (Recommended, etc.) that don't manage on-disk variants. */
+ variantActions?: {
+ onUpdate?: (quant: string, expectedBytes: number) => Promise | void;
+ updateTitle?: string;
+ renderUpdateDescription?: (quant: string) => ReactNode;
+ getUpdateSuccessMessage?: (quant: string) => string;
+ updateDisabled?: boolean;
+ onDelete?: (quant: string) => Promise | void;
+ deleteTitle?: string;
+ renderDeleteDescription?: (quant: string) => ReactNode;
+ getDeleteSuccessMessage?: (quant: string) => string;
+ deleteDisabled?: boolean;
+ };
/** On Device rows honor the Show all quantizations setting; Recommended and
* other browse lists always show every quant. */
onDevice?: boolean;
/** Report GGUF vision support up so the parent row can badge it. */
onHasVision?: (hasVision: boolean) => void;
}) {
+ const onUpdateVariant = variantActions?.onUpdate;
+ const updateVariantTitle = variantActions?.updateTitle ?? "Update cached model?";
+ const renderUpdateVariantDescription = variantActions?.renderUpdateDescription;
+ const updateDisabled = variantActions?.updateDisabled ?? false;
+ const onDeleteVariant = variantActions?.onDelete;
+ const deleteVariantTitle = variantActions?.deleteTitle ?? "Delete cached model?";
+ const renderDeleteVariantDescription = variantActions?.renderDeleteDescription;
+ const getDeleteVariantSuccessMessage = variantActions?.getDeleteSuccessMessage;
+ const deleteDisabled = variantActions?.deleteDisabled ?? false;
const [variants, setVariants] = useState(null);
const [defaultVariant, setDefaultVariant] = useState(null);
const [hasVision, setHasVision] = useState(false);
@@ -660,13 +694,14 @@ function GgufVariantExpander({
const [nativeContext, setNativeContext] = useState(null);
const [loading, setLoading] = useState(true);
const [error, setError] = useState(null);
+ const [refreshKey, setRefreshKey] = useState(0);
useEffect(() => {
let canceled = false;
setLoading(true);
setError(null);
- listGgufVariants(repoId)
+ listGgufVariants(repoId, hfToken)
.then((res) => {
if (canceled) return;
const normalized = normalizeGgufVariantsResponse(res);
@@ -689,7 +724,7 @@ function GgufVariantExpander({
return () => {
canceled = true;
};
- }, [repoId]);
+ }, [repoId, refreshKey, hfToken]);
// Covers Unix absolute (/), Windows drive (C:\, D:/), UNC (\\server), relative (./, ../), tilde (~/)
const isLocalPath = /^(\/|\.{1,2}[\\\/]|~[\\\/]|[A-Za-z]:[\\\/]|\\\\)/.test(
@@ -872,6 +907,7 @@ function GgufVariantExpander({
const fit = getGgufFit(v.size_bytes);
const oom = fit === "oom";
const tight = fit === "tight";
+ const expectedBytes = ggufVariantExpectedBytes(v);
const keyBase = `${repoId}:${v.filename}`;
const variantOptionKey = makeModelOptionKey("gguf-variant", keyBase);
return (
@@ -880,7 +916,7 @@ function GgufVariantExpander({
type="button"
{...variantList.getOptionProps(variantOptionKey, false)}
onClick={() =>
- handleVariantClick(v.quant, v.filename, v.downloaded, v.size_bytes)
+ handleVariantClick(v.quant, v.filename, v.downloaded, expectedBytes)
}
className={cn(
"flex min-w-0 flex-1 items-center justify-between gap-2 rounded-full px-2 py-1 text-left text-sm transition-colors hover:bg-[#ececec] focus-visible:bg-[#ececec] focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring/45 dark:hover:bg-[var(--sidebar-accent)] dark:focus-visible:bg-[var(--sidebar-accent)]",
@@ -893,9 +929,16 @@ function GgufVariantExpander({
{v.quant}
{v.downloaded ? (
-
- downloaded
-
+ <>
+
+ downloaded
+
+ {v.update_available ? (
+
+ update available
+
+ ): null}
+ >
) : v.quant === effectiveRecommended ? (
recommended
@@ -918,6 +961,29 @@ function GgufVariantExpander({
+ {v.downloaded && v.update_available && onUpdateVariant && (
+
+ This will update{" "}
+
+ {repoId} ({v.quant})
+ {"."}
+ >
+ )
+ }
+ repoId={repoId}
+ variant={v.quant}
+ buttonClassName="p-1"
+ iconClassName="size-3"
+ disabled={updateDisabled}
+ onConfirm={() => onUpdateVariant(v.quant, expectedBytes)}
+ onUpdated={() => setRefreshKey((key) => key + 1)}
+ />
+ )}
{v.downloaded && (
+ {name}
+
+ {path}
+
+ >
+ );
+}
+
/** Whether a local model is an MLX build (name hint). MLX runs on Mac only, so
* callers gate visibility on the host being a Mac. */
function localModelIsMlx(m: LocalModelInfo): boolean {
@@ -1229,18 +1306,25 @@ export function HubModelPicker({
task?: HfTaskFilter;
}) {
const gpu = useGpuInfo();
+ // The currently-loaded/running model id. We read params.checkpoint from the
+ // runtime store (backend-mirrored from /api/inference/status.active_model, see
+ // chat-runtime-store) rather than the dropdown `isSelected` highlight (which is
+ // just `value === repo_id` and can reflect a staged, not-yet-loaded pick). Used
+ // to disable the cached-row update action for the model that's live in memory.
+ const loadedModelId = useChatRuntimeStore((s) => s.params.checkpoint);
// Last-loaded timestamps power the "Recent" sort (vs "Downloaded" = file date).
const loadTimes = useModelLoadTimes(value);
// Fade the list's top edge once scrolled, and its bottom edge while more
// rows sit below the fold.
const [listScrolled, setListScrolled] = useState(false);
const [listMoreBelow, setListMoreBelow] = useState(false);
+ const hfToken = useHfTokenStore((s) => s.token);
const [query, setQuery] = useState("");
const debouncedQuery = useDebouncedValue(query);
// Shared Hub search stack (the same hooks the Hub page uses) so the picker
// and Hub run one implementation. Scoped to unsloth like the old listing.
const online = useOnlineStatus();
- const accessToken = useHfTokenStore((s) => s.token) || undefined;
+ const accessToken = hfToken || undefined;
// Recommended section: a live unsloth listing sorted by the dropdown. The
// same sort drives the search results so the dropdown works while searching.
const [recommendedSort, setRecommendedSort] =
@@ -1410,6 +1494,28 @@ export function HubModelPicker({
const alreadyCached =
_cachedGgufCache.length > 0 || _cachedModelsCache.length > 0;
const [cachedReady, setCachedReady] = useState(alreadyCached);
+ const [updateConflictKey, setUpdateConflictKey] = useState(
+ null,
+ );
+ const updateTransportConflict = useDownloadManagerStore((state) =>
+ updateConflictKey
+ ? (state.conflicts[updateConflictKey]?.info ?? null)
+ : null,
+ );
+ const cancelUpdateConflict = useCallback(() => {
+ if (updateConflictKey) downloadManager.cancelConflict(updateConflictKey);
+ setUpdateConflictKey(null);
+ }, [updateConflictKey]);
+ const resumeUpdateConflict = useCallback(() => {
+ if (!updateConflictKey) return;
+ downloadManager.resumeConflict(updateConflictKey);
+ setUpdateConflictKey(null);
+ }, [updateConflictKey]);
+ const restartUpdateConflict = useCallback(() => {
+ if (!updateConflictKey) return;
+ downloadManager.restartConflict(updateConflictKey);
+ setUpdateConflictKey(null);
+ }, [updateConflictKey]);
// LM Studio local models -- module-level cache, same pattern as above.
const [lmStudioModels, setLmStudioModels] =
@@ -1532,14 +1638,42 @@ export function HubModelPicker({
setCachedGguf(v);
})
.catch(() => {});
- listCachedModels()
+ listCachedModels(hfToken || undefined)
.then((v) => {
_cachedModelsCache = v;
setCachedModels(v);
})
.catch(() => {});
refreshLocalModelsList();
- }, [refreshLocalModelsList]);
+ }, [hfToken, refreshLocalModelsList]);
+
+ // Updates run as MANAGED downloads (they show in the global Downloads panel
+ // with manifest-based progress + a working Cancel), instead of a blocking
+ // call. The worker re-resolves `main` and pulls only changed blobs, so the
+ // cached copy stays usable until the new revision lands. The row's
+ // ModelUpdateAction refreshes the list when this repo+variant completes.
+ const startManagedUpdate = useCallback((repoId: string, variant: string, expectedBytes: number) => {
+ return downloadManager
+ .requestStart({
+ kind: "model",
+ repoId,
+ variant,
+ expectedBytes,
+ })
+ .then((outcome) => {
+ if (outcome === "conflict") {
+ setUpdateConflictKey(jobKeyOf("model", repoId, variant));
+ } else if (outcome === "error") {
+ throw new Error("Failed to start update");
+ }
+ });
+ }, []);
+
+ const updateGgufVariant = useCallback(
+ (repoId: string, quant: string, expectedBytes: number) =>
+ startManagedUpdate(repoId, quant, expectedBytes),
+ [startManagedUpdate],
+ );
useEffect(() => {
// Always refresh LM Studio + custom folder models (not gated by alreadyCached).
@@ -1564,14 +1698,14 @@ export function HubModelPicker({
})
.catch(() => {})
.finally(check);
- listCachedModels()
+ listCachedModels(hfToken || undefined)
.then((v) => {
_cachedModelsCache = v;
setCachedModels(v);
})
.catch(() => {})
.finally(check);
- }, [refreshLocalModelsList, refreshScanFolders]);
+ }, [hfToken, refreshLocalModelsList, refreshScanFolders]);
// Hide downloaded models from the recommended list. Case-insensitive
// since the HF cache lowercases repo IDs.
@@ -2417,14 +2551,21 @@ export function HubModelPicker({
onDevice={true}
onHasVision={(v) => reportVision(c.repo_id, v)}
onSelect={onSelect}
+ hfToken={hfToken || undefined}
parentOptionKey={optionKey}
onNavigatePastStart={() => hubModelList.focusOption(optionKey)}
onNavigatePastEnd={() => hubModelList.moveFocus(optionKey, "next")}
gpuGb={gpu.available ? gpu.memoryTotalGb : undefined}
systemRamGb={gpu.systemRamAvailableGb || undefined}
- onDeleteVariant={async (quant) => {
- await deleteCachedModel(c.repo_id, quant);
- refreshCachedLists();
+ variantActions={{
+ onUpdate: (quant, expectedBytes) =>
+ updateGgufVariant(c.repo_id, quant, expectedBytes),
+ // Can't update the model that's live in memory under itself.
+ updateDisabled: loadedModelId === c.repo_id,
+ onDelete: async (quant) => {
+ await deleteCachedModel(c.repo_id, quant);
+ refreshCachedLists();
+ },
}}
/>
)}
@@ -2483,7 +2624,8 @@ export function HubModelPicker({
};
return (
-
+ <>
+
{/* A small right inset shortens the search bar so Search Hub lands on the
last dropdown's right edge (none on the wider Connected box). */}
setQuery(event.target.value)}
- placeholder="Search Unsloth models"
+ placeholder={
+ section === "downloaded"
+ ? "Search local models"
+ : "Search Unsloth models"
+ }
data-model-picker-search-input={true}
className="field-soft h-9 border-0 pl-8 pr-8"
/>
@@ -3042,6 +3188,10 @@ export function HubModelPicker({
hubModelList.focusOption(optionKey)
@@ -3337,9 +3496,11 @@ export function HubModelPicker({
gpu.available ? gpu.memoryTotalGb : undefined
}
systemRamGb={gpu.systemRamAvailableGb || undefined}
- onDeleteVariant={async (quant) => {
- await deleteCachedModel(id, quant);
- refreshCachedLists();
+ variantActions={{
+ onDelete: async (quant) => {
+ await deleteCachedModel(id, quant);
+ refreshCachedLists();
+ },
}}
/>
)}
@@ -3411,6 +3572,7 @@ export function HubModelPicker({
hubModelList.focusOption(optionKey)
@@ -3422,9 +3584,11 @@ export function HubModelPicker({
gpu.available ? gpu.memoryTotalGb : undefined
}
systemRamGb={gpu.systemRamAvailableGb || undefined}
- onDeleteVariant={async (quant) => {
- await deleteCachedModel(id, quant);
- refreshCachedLists();
+ variantActions={{
+ onDelete: async (quant) => {
+ await deleteCachedModel(id, quant);
+ refreshCachedLists();
+ },
}}
/>
)}
@@ -3498,6 +3662,7 @@ export function HubModelPicker({
hubModelList.focusOption(optionKey)
@@ -3509,9 +3674,11 @@ export function HubModelPicker({
gpu.available ? gpu.memoryTotalGb : undefined
}
systemRamGb={gpu.systemRamAvailableGb || undefined}
- onDeleteVariant={async (quant) => {
- await deleteCachedModel(id, quant);
- refreshCachedLists();
+ variantActions={{
+ onDelete: async (quant) => {
+ await deleteCachedModel(id, quant);
+ refreshCachedLists();
+ },
}}
/>
)}
@@ -3538,7 +3705,7 @@ export function HubModelPicker({
) : null}
-
+
+
+ >
);
}
@@ -3698,22 +3872,21 @@ function FineTunedRows({
gpuGb={gpu.available ? gpu.memoryTotalGb : undefined}
systemRamGb={gpu.systemRamAvailableGb || undefined}
sourceOverride={isExportedGguf ? "exported" : undefined}
- deleteVariantTitle="Delete exported GGUF variant?"
- renderDeleteVariantDescription={(quant) => (
- <>
- This will remove{" "}
-
- {adapter.name} ({quant})
- {" "}
- from disk. This cannot be undone.
- >
- )}
- getDeleteVariantSuccessMessage={(quant) =>
- `Deleted ${adapter.name} ${quant}`
- }
- deleteDisabled={deleteDisabled}
- onDeleteVariant={
- isExportedGguf
+ variantActions={{
+ deleteTitle: "Delete exported GGUF variant?",
+ renderDeleteDescription: (quant) => (
+ <>
+ This will remove{" "}
+
+ {adapter.name} ({quant})
+ {" "}
+ from disk. This cannot be undone.
+ >
+ ),
+ getDeleteSuccessMessage: (quant) =>
+ `Deleted ${adapter.name} ${quant}`,
+ deleteDisabled: deleteDisabled,
+ onDelete: isExportedGguf
? async (quant) => {
await deleteFineTunedModel({
modelPath: adapter.id,
@@ -3726,8 +3899,8 @@ function FineTunedRows({
ggufVariant: quant,
});
}
- : undefined
- }
+ : undefined,
+ }}
/>
)}
diff --git a/studio/frontend/src/components/assistant-ui/model-selector/recommended-fit.ts b/studio/frontend/src/components/assistant-ui/model-selector/recommended-fit.ts
index cc96665df4..24f0edc784 100644
--- a/studio/frontend/src/components/assistant-ui/model-selector/recommended-fit.ts
+++ b/studio/frontend/src/components/assistant-ui/model-selector/recommended-fit.ts
@@ -64,38 +64,19 @@ export function matchesFormatFilter(
}
}
-// Model-size extraction from repo id, matching the backend's 3-regex priority:
-// active params (MoE "A3B") > effective params (Gemma "E4B") > total ("8B").
-// Bounded by separators so we never read "16" from "bf16" or "2" from "Kimi-K2".
-// Examples: "Qwen3.5-35B-A3B" -> 3, "gemma-4-E4B" -> 4, "Llama-3-8B" -> 8.
-const ACTIVE_PARAM_RE = /(?:^|[-_/. ])a(\d+(?:\.\d+)?)\s*[bB](?=$|[-_/. ])/i;
-const EFFECTIVE_PARAM_RE = /(?:^|[-_/. ])e(\d+(?:\.\d+)?)\s*[bB](?=$|[-_/. ])/i;
-const TOTAL_PARAM_RE = /(?:^|[-_/. ])(\d+(?:\.\d+)?)\s*[bB](?=$|[-_/. ])/;
-
-function paramsFromMatch(match: RegExpExecArray | null): number | undefined {
- if (!match) return undefined;
- const billions = parseFloat(match[1]);
- return Number.isFinite(billions) && billions > 0
- ? billions * 1e9
- : undefined;
-}
-
-/** Active/effective parameter count parsed from a repo id, if it uses explicit
- * MoE/Gemma-style notation such as A3B or E4B. */
-export function activeOrEffectiveParamsFromId(id: string): number | undefined {
- return (
- paramsFromMatch(ACTIVE_PARAM_RE.exec(id)) ??
- paramsFromMatch(EFFECTIVE_PARAM_RE.exec(id))
- );
-}
+// First "B" token in a repo id, e.g. "Qwen3-4B-GGUF" -> 4, "gpt-oss-20b" ->
+// 20, "Qwen3-30B-A3B" -> 30 (MoE total), "gemma-4-E4B" -> 4 (effective-param
+// "E" series). The digits must be bounded by a separator so we never read "16"
+// from "bf16" or the "2" in "Kimi-K2".
+const PARAM_RE = /(?:^|[-_/. ])[eE]?(\d+(?:\.\d+)?)\s*[bB](?=$|[-_./ ])/;
/** Parameter count (absolute, e.g. 4e9) parsed from a repo id, or undefined
- * when the id has no size token (so callers can treat the size as unknown).
- * Prefers MoE active-param notation (A3B) over effective (E4B) over total. */
+ * when the id has no size token (so callers can treat the size as unknown). */
export function paramsFromId(id: string): number | undefined {
- return (
- activeOrEffectiveParamsFromId(id) ?? paramsFromMatch(TOTAL_PARAM_RE.exec(id))
- );
+ const match = PARAM_RE.exec(id);
+ if (!match) return undefined;
+ const billions = parseFloat(match[1]);
+ return Number.isFinite(billions) && billions > 0 ? billions * 1e9 : undefined;
}
// Smallest practical GGUF/MLX quant (~Q2_K, low-bit). The fit check asks whether
diff --git a/studio/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx
index a05910c29f..7fe98b701f 100644
--- a/studio/frontend/src/components/assistant-ui/thread.tsx
+++ b/studio/frontend/src/components/assistant-ui/thread.tsx
@@ -3555,9 +3555,9 @@ const DiffusionCanvas: FC = () => {
/**
* AssistantMessage handles the display and inline-editing of AI responses.
- *
- * It utilizes a "Tagged Text" system ( and tags) to allow users
- * to edit structured reasoning and tool outputs within a plain-text textarea
+ *
+ * It utilizes a "Tagged Text" system ( and tags) to allow users
+ * to edit structured reasoning and tool outputs within a plain-text textarea
* while preserving the underlying data schema and tool-call metadata.
*/
const AssistantMessage: FC = () => {
@@ -3565,7 +3565,7 @@ const AssistantMessage: FC = () => {
const messageId = useAuiState(({ message }) => message.id);
const messageContent = useAuiState(({ message }) => message.content);
const incognito = useChatRuntimeStore((s) => s.incognito);
-
+
// Use global store for editing state to ensure a single source of truth
const editingId = useChatRuntimeStore((s) => s.editingMessageId);
const setEditingId = useChatRuntimeStore((s) => s.setEditingMessageId);
@@ -3588,9 +3588,9 @@ const AssistantMessage: FC = () => {
const handleSave = async () => {
const finalText = textareaRef.current?.value || "";
-
+
// Prioritize the specific thread item ID, then fallback to the global active thread ID
- const remoteId = aui.threadListItem().getState().remoteId
+ const remoteId = aui.threadListItem().getState().remoteId
|| useChatRuntimeStore.getState().activeThreadId;
if (!remoteId || remoteId === "" || remoteId === "/") {
@@ -3601,9 +3601,9 @@ const AssistantMessage: FC = () => {
try {
await updateThreadMessage({
- thread: {
- export: () => aui.thread().export(),
- import: (data) => aui.thread().import(data)
+ thread: {
+ export: () => aui.thread().export(),
+ import: (data) => aui.thread().import(data)
},
messageId,
remoteId,
@@ -3626,14 +3626,14 @@ const AssistantMessage: FC = () => {
{isEditing ? (
- {row.fitLevel && (
- <>
- •
-
- {row.fitLevel === "comfortable" ? "Comfortable" : row.fitLevel === "fits" ? "Fits GPU" : "OOM"}
-
- >
- )}
diff --git a/studio/frontend/src/features/hub/catalog/models-toolbar.tsx b/studio/frontend/src/features/hub/catalog/models-toolbar.tsx
index 9cb17736af..48f7fcffaa 100644
--- a/studio/frontend/src/features/hub/catalog/models-toolbar.tsx
+++ b/studio/frontend/src/features/hub/catalog/models-toolbar.tsx
@@ -21,7 +21,6 @@ import { HugeiconsIcon } from "@hugeicons/react";
import type { HfSortKey } from "@/features/hub/hooks/use-hub-model-search";
import type {
CapabilityFilter,
- GpuFitFilter,
ModelFormatFilter,
ModelsTab,
ResourceTypeFilter,
@@ -29,7 +28,6 @@ import type {
import {
CAPABILITY_FILTER_OPTIONS,
FORMAT_FILTER_OPTIONS,
- GPU_FIT_FILTER_OPTIONS,
} from "../lib/view-models";
import { HubOptionMenu, type HubOption } from "./hub-option-menu";
import {
@@ -70,8 +68,6 @@ export const ModelsToolbar = memo(function ModelsToolbar({
onFormatFilterChange,
capabilityFilter,
onCapabilityFilterChange,
- gpuFitFilter,
- onGpuFitFilterChange,
onManageLocalFolders,
onOpenFineTune,
}: {
@@ -88,8 +84,6 @@ export const ModelsToolbar = memo(function ModelsToolbar({
onFormatFilterChange: (value: ModelFormatFilter) => void;
capabilityFilter: CapabilityFilter;
onCapabilityFilterChange: (value: CapabilityFilter) => void;
- gpuFitFilter: GpuFitFilter;
- onGpuFitFilterChange: (value: GpuFitFilter) => void;
onManageLocalFolders: () => void;
/** Opens the curated "Fine-tune ready" channel (discover only). Exposed as a
* format-dropdown option rather than a standalone feed section. */
@@ -159,14 +153,6 @@ export const ModelsToolbar = memo(function ModelsToolbar({
})),
[],
);
- const gpuFitOptions = useMemo
[]>(
- () =>
- GPU_FIT_FILTER_OPTIONS.map((option) => ({
- value: option.value,
- label: option.label,
- })),
- [],
- );
const sortOptions = useMemo[]>(
() =>
SORT_OPTIONS.map((option) => ({
@@ -357,16 +343,6 @@ export const ModelsToolbar = memo(function ModelsToolbar({
/>
)}
- {tab === "discover" && !isDataset && (
-
- )}
-
{tab === "discover" && (