diff --git a/.github/workflows/release-desktop.yml b/.github/workflows/release-desktop.yml new file mode 100644 index 0000000000..2d1c6f51f2 --- /dev/null +++ b/.github/workflows/release-desktop.yml @@ -0,0 +1,187 @@ +name: Release Desktop App + +on: + workflow_dispatch: + inputs: + draft: + description: 'Create as draft release' + type: boolean + default: true + +permissions: + contents: write + +jobs: + build: + strategy: + fail-fast: false + max-parallel: 1 + matrix: + include: + - platform: macos-latest + args: '--target aarch64-apple-darwin' + label: macOS (Apple Silicon) + # - platform: macos-latest + # args: '--target x86_64-apple-darwin' + # label: macOS (Intel) + - platform: ubuntu-22.04 + args: '' + label: Linux (x64) + - platform: windows-latest + args: '' + label: Windows (x64) + + name: Build ${{ matrix.label }} + runs-on: ${{ matrix.platform }} + + env: + FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true + + + steps: + - uses: actions/checkout@34e114876b0b11c390a56381ad16ebd13914f8d5 + + # ── Linux dependencies ── + - name: Install Linux dependencies + if: matrix.platform == 'ubuntu-22.04' + run: | + sudo apt-get update + sudo apt-get install -y libwebkit2gtk-4.1-dev libayatana-appindicator3-dev librsvg2-dev libxdo-dev libssl-dev patchelf + + # ── Node.js ── + - name: Setup Node.js + uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 + with: + node-version: 24 + + - name: Install frontend dependencies + working-directory: studio/frontend + run: npm install + + # ── Rust ── + - name: Install Rust stable + uses: dtolnay/rust-toolchain@stable + with: + targets: ${{ matrix.platform == 'macos-latest' && 'aarch64-apple-darwin,x86_64-apple-darwin' || '' }} + + - name: Rust cache + uses: swatinem/rust-cache@42dc69e1aa15d09112580998cf2ef0119e2e91ae + with: + workspaces: 'studio/src-tauri -> target' + + # ── macOS: import signing certificate ── + - name: Import Apple certificate + if: matrix.platform == 'macos-latest' + env: + APPLE_CERTIFICATE: ${{ secrets.APPLE_CERTIFICATE }} + APPLE_CERTIFICATE_PASSWORD: ${{ secrets.APPLE_CERTIFICATE_PASSWORD }} + KEYCHAIN_PASSWORD: ${{ secrets.KEYCHAIN_PASSWORD }} + run: | + echo $APPLE_CERTIFICATE | base64 --decode > certificate.p12 + security create-keychain -p "$KEYCHAIN_PASSWORD" build.keychain + security default-keychain -s build.keychain + security unlock-keychain -p "$KEYCHAIN_PASSWORD" build.keychain + security set-keychain-settings -t 3600 -u build.keychain + security import certificate.p12 -k build.keychain -P "$APPLE_CERTIFICATE_PASSWORD" -T /usr/bin/codesign + security set-key-partition-list -S apple-tool:,apple:,codesign: -s -k "$KEYCHAIN_PASSWORD" build.keychain + security find-identity -v -p codesigning build.keychain + rm -f certificate.p12 + + # ── Windows: install Azure Trusted Signing CLI ── + - name: Install trusted-signing-cli + if: matrix.platform == 'windows-latest' + run: | + cargo install trusted-signing-cli --version 0.9.0 --locked + echo "$env:USERPROFILE\.cargo\bin" | Out-File -FilePath $env:GITHUB_PATH -Encoding utf8 -Append + + # ── Windows: verify signing CLI is accessible ── + - name: Verify trusted-signing-cli + if: matrix.platform == 'windows-latest' + run: | + Write-Output "PATH: $env:PATH" + Get-Command trusted-signing-cli -ErrorAction SilentlyContinue || Write-Output "trusted-signing-cli NOT in PATH" + trusted-signing-cli --version || Write-Output "trusted-signing-cli failed to run" + + # ── Linux: build + sign + upload ── + - name: Build Linux app + if: matrix.platform == 'ubuntu-22.04' + uses: tauri-apps/tauri-action@fce9c6108b31ea247710505d3aaaa893ee6768d4 + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }} + TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY_PASSWORD }} + with: + projectPath: studio + tagName: desktop-v__VERSION__ + releaseName: 'Unsloth Studio (Desktop) v__VERSION__' + releaseBody: | + Desktop app for Unsloth Studio. + + **macOS**: Download the Apple Silicon `.dmg`. + **Windows**: Download the `-setup.exe` installer. + **Linux**: Download `.deb` (Ubuntu/Debian) or `.AppImage` (universal). + + > Linux in-app updates are AppImage-oriented. Package installs should update by downloading a new package. + > Linux AppImage on Ubuntu 24.04+ may require: `sudo apt install libfuse2t64` + releaseDraft: ${{ inputs.draft }} + prerelease: false + args: -v ${{ matrix.args }} + + # ── macOS: build + sign + notarize + upload ── + - name: Build macOS app + if: matrix.platform == 'macos-latest' + uses: tauri-apps/tauri-action@fce9c6108b31ea247710505d3aaaa893ee6768d4 + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }} + TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY_PASSWORD }} + APPLE_SIGNING_IDENTITY: ${{ secrets.APPLE_SIGNING_IDENTITY }} + APPLE_ID: ${{ secrets.APPLE_ID }} + APPLE_PASSWORD: ${{ secrets.APPLE_PASSWORD }} + APPLE_TEAM_ID: ${{ secrets.APPLE_TEAM_ID }} + with: + projectPath: studio + tagName: desktop-v__VERSION__ + releaseName: 'Unsloth Studio (Desktop) v__VERSION__' + releaseBody: | + Desktop app for Unsloth Studio. + + **macOS**: Download the Apple Silicon `.dmg`. + **Windows**: Download the `-setup.exe` installer. + **Linux**: Download `.deb` (Ubuntu/Debian) or `.AppImage` (universal). + + > Linux in-app updates are AppImage-oriented. Package installs should update by downloading a new package. + > Linux AppImage on Ubuntu 24.04+ may require: `sudo apt install libfuse2t64` + releaseDraft: ${{ inputs.draft }} + prerelease: false + args: -v ${{ matrix.args }} + + # ── Windows: build + sign + upload ── + - name: Build Windows app + if: matrix.platform == 'windows-latest' + uses: tauri-apps/tauri-action@fce9c6108b31ea247710505d3aaaa893ee6768d4 + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }} + TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY_PASSWORD }} + AZURE_CLIENT_ID: ${{ secrets.AZURE_CLIENT_ID }} + AZURE_CLIENT_SECRET: ${{ secrets.AZURE_CLIENT_SECRET }} + AZURE_TENANT_ID: ${{ secrets.AZURE_TENANT_ID }} + AZURE_TRUSTED_SIGNING_ACCOUNT_NAME: ${{ secrets.AZURE_TRUSTED_SIGNING_ACCOUNT_NAME }} + AZURE_CERTIFICATE_PROFILE_NAME: ${{ secrets.AZURE_CERTIFICATE_PROFILE_NAME }} + with: + projectPath: studio + tagName: desktop-v__VERSION__ + releaseName: 'Unsloth Studio (Desktop) v__VERSION__' + releaseBody: | + Desktop app for Unsloth Studio. + + **macOS**: Download the Apple Silicon `.dmg`. + **Windows**: Download the `-setup.exe` installer. + **Linux**: Download `.deb` (Ubuntu/Debian) or `.AppImage` (universal). + + > Linux in-app updates are AppImage-oriented. Package installs should update by downloading a new package. + > Linux AppImage on Ubuntu 24.04+ may require: `sudo apt install libfuse2t64` + releaseDraft: ${{ inputs.draft }} + prerelease: false + args: -v ${{ matrix.args }} diff --git a/.gitignore b/.gitignore index 7a24d07c6f..b6786ee655 100644 --- a/.gitignore +++ b/.gitignore @@ -204,6 +204,18 @@ tmp/ **/node_modules/ auth.db +# Tauri local build/generated output +studio/src-tauri/target/ +studio/src-tauri/gen/ +studio/src-tauri/artifacts/ +studio/src-tauri/icons/android/ +studio/src-tauri/icons/ios/ +studio/src-tauri/icons/128x128@2x.png +studio/src-tauri/icons/64x64.png +studio/src-tauri/icons/Square*Logo.png +studio/src-tauri/icons/StoreLogo.png +studio/src-tauri/icons/squarehq.png + # Local working docs **/CLAUDE.md **/claude.md diff --git a/install.ps1 b/install.ps1 index 3fc9ac4690..44464101f3 100644 --- a/install.ps1 +++ b/install.ps1 @@ -12,11 +12,13 @@ function Install-UnslothStudio { $StudioLocalInstall = $false $PackageName = "unsloth" $RepoRoot = "" + $TauriMode = $false $SkipTorch = $false $argList = $args for ($i = 0; $i -lt $argList.Count; $i++) { switch ($argList[$i]) { "--local" { $StudioLocalInstall = $true } + "--tauri" { $TauriMode = $true } "--no-torch" { $SkipTorch = $true } "--verbose" { $script:UnslothVerbose = $true } "-v" { $script:UnslothVerbose = $true } @@ -44,6 +46,20 @@ function Install-UnslothStudio { } } + # Validate --package to prevent injection into shell/Python commands + if ($PackageName -notmatch '^[a-zA-Z0-9][a-zA-Z0-9._-]*$') { + Write-Host "[ERROR] --package name contains invalid characters (allowed: a-z A-Z 0-9 . _ -)" -ForegroundColor Red + return + } + + # ── Tauri structured output ── + function Write-TauriLog { + param([string]$Tag, [string]$Message) + if ($TauriMode) { + Write-Host "[TAURI:$Tag] $Message" + } + } + $PythonVersion = "3.13" $StudioHome = Join-Path $env:USERPROFILE ".unsloth\studio" $VenvDir = Join-Path $StudioHome "unsloth_studio" @@ -609,6 +625,7 @@ shell.Run cmd, 0, False } # ── Check winget ── + Write-TauriLog "STEP" "Checking system dependencies" if (-not (Get-Command winget -ErrorAction SilentlyContinue)) { step "winget" "not available" "Red" substep "Install it from https://aka.ms/getwinget" "Yellow" @@ -688,6 +705,7 @@ shell.Run cmd, 0, False # ── Install Python if no compatible version (3.11-3.13) found ── # Find-CompatiblePython returns @{ Version = "3.13"; Path = "C:\...\python.exe" } or $null. + Write-TauriLog "STEP" "Installing Python" $DetectedPython = Find-CompatiblePython if ($DetectedPython) { step "python" "Python $($DetectedPython.Version) already installed" @@ -736,6 +754,7 @@ shell.Run cmd, 0, False } # ── Install uv if not present ── + Write-TauriLog "STEP" "Installing uv package manager" if (-not (Get-Command uv -ErrorAction SilentlyContinue)) { substep "installing uv package manager..." $prevEAP = $ErrorActionPreference @@ -746,7 +765,7 @@ shell.Run cmd, 0, False # Fallback: if winget didn't put uv on PATH, try the PowerShell installer if (-not (Get-Command uv -ErrorAction SilentlyContinue)) { substep "trying alternative uv installer..." "Yellow" - powershell -ExecutionPolicy ByPass -c "irm https://astral.sh/uv/install.ps1 | iex" + Invoke-Expression (Invoke-RestMethod -Uri "https://astral.sh/uv/install.ps1") Refresh-SessionPath } } @@ -760,6 +779,7 @@ shell.Run cmd, 0, False # ── Create venv (migrate old layout if possible, otherwise fresh) ── # Pass the resolved executable path to uv so it does not re-resolve # a version string back to a conda interpreter. + Write-TauriLog "STEP" "Creating virtual environment" if (-not (Test-Path $StudioHome)) { New-Item -ItemType Directory -Path $StudioHome -Force | Out-Null } @@ -806,6 +826,7 @@ shell.Run cmd, 0, False substep "$VenvDir" $venvExit = Invoke-InstallCommand { uv venv $VenvDir --python "$($DetectedPython.Path)" } if ($venvExit -ne 0) { + Write-TauriLog "ERROR" "Failed to create virtual environment (exit code $venvExit)" Write-Host "[ERROR] Failed to create virtual environment (exit code $venvExit)" -ForegroundColor Red return } @@ -908,11 +929,12 @@ shell.Run cmd, 0, False if ($_Migrated) { # Migrated env: force-reinstall unsloth+unsloth-zoo to ensure clean state # in the new venv location, while preserving existing torch/CUDA + Write-TauriLog "STEP" "Installing unsloth" substep "upgrading unsloth in migrated environment..." if ($SkipTorch) { # No-torch: install unsloth + unsloth-zoo with --no-deps, then # runtime deps (typer, safetensors, transformers, etc.) with --no-deps. - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.5" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.7" unsloth-zoo } if ($baseInstallExit -eq 0) { $NoTorchReq = Find-NoTorchRuntimeFile if ($NoTorchReq) { @@ -920,7 +942,7 @@ shell.Run cmd, 0, False } } } else { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.5" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --reinstall-package unsloth --reinstall-package unsloth-zoo "unsloth>=2026.4.7" unsloth-zoo } } if ($baseInstallExit -ne 0) { Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red @@ -938,19 +960,22 @@ shell.Run cmd, 0, False if ($SkipTorch) { substep "skipping PyTorch (--no-torch flag set)." "Yellow" } else { + Write-TauriLog "STEP" "Installing PyTorch" substep "installing PyTorch ($TorchIndexUrl)..." $torchInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython "torch>=2.4,<2.11.0" torchvision torchaudio --index-url $TorchIndexUrl } if ($torchInstallExit -ne 0) { + Write-TauriLog "ERROR" "Failed to install PyTorch (exit code $torchInstallExit)" Write-Host "[ERROR] Failed to install PyTorch (exit code $torchInstallExit)" -ForegroundColor Red return } } + Write-TauriLog "STEP" "Installing unsloth" substep "installing unsloth (this may take a few minutes)..." if ($SkipTorch) { # No-torch: install unsloth + unsloth-zoo with --no-deps, then # runtime deps (typer, safetensors, transformers, etc.) with --no-deps. - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --upgrade-package unsloth --upgrade-package unsloth-zoo "unsloth>=2026.4.5" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --no-deps --upgrade-package unsloth --upgrade-package unsloth-zoo "unsloth>=2026.4.7" unsloth-zoo } if ($baseInstallExit -eq 0) { $NoTorchReq = Find-NoTorchRuntimeFile if ($NoTorchReq) { @@ -958,11 +983,12 @@ shell.Run cmd, 0, False } } } elseif ($StudioLocalInstall) { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "unsloth>=2026.4.5" unsloth-zoo } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "unsloth>=2026.4.7" unsloth-zoo } } else { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth "$PackageName" } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --upgrade-package unsloth -- "$PackageName" } } if ($baseInstallExit -ne 0) { + Write-TauriLog "ERROR" "Failed to install unsloth (exit code $baseInstallExit)" Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red return } @@ -977,9 +1003,10 @@ shell.Run cmd, 0, False } } else { # Fallback: GPU detection failed to produce a URL -- let uv resolve torch + Write-TauriLog "STEP" "Installing unsloth" substep "installing unsloth (this may take a few minutes)..." if ($StudioLocalInstall) { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython unsloth-zoo "unsloth>=2026.4.5" --torch-backend=auto } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython unsloth-zoo "unsloth>=2026.4.7" --torch-backend=auto } if ($baseInstallExit -ne 0) { Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red return @@ -991,20 +1018,52 @@ shell.Run cmd, 0, False return } } else { - $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython "$PackageName" --torch-backend=auto } + $baseInstallExit = Invoke-InstallCommand { uv pip install --python $VenvPython --torch-backend=auto -- "$PackageName" } if ($baseInstallExit -ne 0) { + Write-TauriLog "ERROR" "Failed to install unsloth (exit code $baseInstallExit)" Write-Host "[ERROR] Failed to install unsloth (exit code $baseInstallExit)" -ForegroundColor Red return } } } + # Hotfix: patch install_python_stack.py for Windows GUI stdout + # The PyPI version crashes with OSError when stdout is piped from a GUI app. + # Copy our fixed version (bundled by Tauri) over the installed one. + # Remove this block once PyPI ships the fix from commit 18c5aae7. + if ($TauriMode) { + $rawPath = if ($PSCommandPath) { $PSCommandPath } else { $MyInvocation.ScriptName } + $scriptDir = Split-Path -Parent ($rawPath -replace '^\\\\\?\\', '') + $fixedPy = Join-Path $scriptDir "install_python_stack.py" + $target = Join-Path $VenvDir "Lib\site-packages\studio\install_python_stack.py" + $sentinel = "# UNSLOTH_DESKTOP_HOTFIX_APPLIED_v1" + $sentinelPattern = [regex]::Escape($sentinel) + if ((Test-Path $fixedPy) -and (Test-Path $target)) { + $installed = Get-Content $target -Raw + if ($installed -notmatch $sentinelPattern) { + Copy-Item $fixedPy $target -Force + Add-Content -Path $target -Value "`n$sentinel" + substep "patched install_python_stack.py (stdout fix)" + } else { + substep "install_python_stack.py already has stdout fix" + } + } elseif ((Test-Path $fixedPy) -and (Test-Path (Split-Path $target))) { + Copy-Item $fixedPy $target -Force + Add-Content -Path $target -Value "`n$sentinel" + substep "patched install_python_stack.py (stdout fix)" + } else { + Write-Host "[WARN] Could not patch install_python_stack.py (bundled file or target dir missing)" -ForegroundColor Yellow + } + } + # ── Run studio setup ── # setup.ps1 will handle installing Git, CMake, Visual Studio Build Tools, # CUDA Toolkit, Node.js, and other dependencies automatically via winget. + Write-TauriLog "STEP" "Running studio setup" step "setup" "running unsloth studio setup..." $UnslothExe = Join-Path $VenvDir "Scripts\unsloth.exe" if (-not (Test-Path $UnslothExe)) { + Write-TauriLog "ERROR" "unsloth CLI was not installed correctly" Write-Host "[ERROR] unsloth CLI was not installed correctly." -ForegroundColor Red Write-Host " Expected: $UnslothExe" -ForegroundColor Yellow Write-Host " This usually means an older unsloth version was installed that does not include the Studio CLI." -ForegroundColor Yellow @@ -1015,6 +1074,8 @@ shell.Run cmd, 0, False $env:SKIP_STUDIO_BASE = "1" $env:STUDIO_PACKAGE_NAME = $PackageName $env:UNSLOTH_NO_TORCH = if ($SkipTorch) { "true" } else { "false" } + # Tauri desktop app bundles its own frontend — skip Node/npm/frontend build + $env:SKIP_STUDIO_FRONTEND = if ($TauriMode) { "1" } else { "0" } # Always set STUDIO_LOCAL_INSTALL explicitly to avoid stale values from # a previous --local run in the same PowerShell session. if ($StudioLocalInstall) { @@ -1032,12 +1093,11 @@ shell.Run cmd, 0, False & $UnslothExe @studioArgs $setupExit = $LASTEXITCODE if ($setupExit -ne 0) { + Write-TauriLog "ERROR" "unsloth studio setup failed (exit code $setupExit)" Write-Host "[ERROR] unsloth studio setup failed (exit code $setupExit)" -ForegroundColor Red return } - New-StudioShortcuts -UnslothExePath $UnslothExe - # ── Expose `unsloth` via a shim dir containing only unsloth.exe ── # We do NOT add the venv Scripts dir to PATH (it also holds python.exe # and pip.exe, which would hijack the user's system interpreter). @@ -1109,6 +1169,14 @@ shell.Run cmd, 0, False } Refresh-SessionPath # sync current session with registry + # ── Tauri mode: done, skip shortcuts and auto-launch ── + if ($TauriMode) { + Write-TauriLog "DONE" "" + return + } + + New-StudioShortcuts -UnslothExePath $UnslothExe + # Launch studio automatically in interactive terminals; # in non-interactive environments (CI, Docker) just print instructions. $IsInteractive = [Environment]::UserInteractive -and (-not [Console]::IsInputRedirected) diff --git a/install.sh b/install.sh index 6915893ddf..07473e441d 100755 --- a/install.sh +++ b/install.sh @@ -35,6 +35,7 @@ substep() { printf " ${C_DIM}%-15s${2:-$C_DIM}%s${C_RST}\n" "" "$1"; } # ── Parse flags ── STUDIO_LOCAL_INSTALL=false PACKAGE_NAME="unsloth" +TAURI_MODE=false _USER_PYTHON="" _NO_TORCH_FLAG=false _VERBOSE=false @@ -54,6 +55,7 @@ for arg in "$@"; do case "$arg" in --local) STUDIO_LOCAL_INSTALL=true ;; --package) _next_is_package=true ;; + --tauri) TAURI_MODE=true ;; --python) _next_is_python=true ;; --no-torch) _NO_TORCH_FLAG=true ;; --verbose|-v) _VERBOSE=true ;; @@ -142,6 +144,24 @@ if [ "$_next_is_python" = true ]; then exit 1 fi +# Validate --package to prevent injection into shell/Python commands. +# Must start with a letter/digit (rejects leading dashes that uv would parse as flags). +case "$PACKAGE_NAME" in + [!a-zA-Z0-9]*) + echo "❌ ERROR: --package name must start with a letter or digit." >&2 + exit 1 ;; + *[!a-zA-Z0-9._-]*) + echo "❌ ERROR: --package name contains invalid characters (allowed: a-z A-Z 0-9 . _ -)" >&2 + exit 1 ;; +esac + +# ── Tauri structured output ── +tauri_log() { + if [ "$TAURI_MODE" = true ]; then + echo "[TAURI:$1] $2" + fi +} + PYTHON_VERSION="" # resolved after platform detection STUDIO_HOME="$HOME/.unsloth/studio" VENV_DIR="$STUDIO_HOME/unsloth_studio" @@ -192,6 +212,12 @@ _smart_apt_install() { return 0 fi + # In Tauri mode, report needed packages and exit — Rust handles elevation + if [ "$TAURI_MODE" = true ]; then + tauri_log "NEED_SUDO" "$_STILL_MISSING" + exit 2 + fi + # Step 3: Escalate -- need elevated permissions for remaining packages if command -v sudo >/dev/null 2>&1; then echo "" @@ -752,6 +778,7 @@ printf " ${C_DIM}%s${C_RST}\n" "$RULE" echo "" # ── Detect platform ── +tauri_log "STEP" "Detecting platform" OS="linux" if [ "$(uname)" = "Darwin" ]; then OS="macos" @@ -804,6 +831,7 @@ fi # ── Check system dependencies ── # cmake and git are needed by unsloth studio setup to build the GGUF inference # engine (llama.cpp). build-essential and libcurl-dev are also needed on Linux. +tauri_log "STEP" "Checking system dependencies" MISSING="" command -v cmake >/dev/null 2>&1 || MISSING="$MISSING cmake" @@ -868,6 +896,7 @@ else fi # ── Install uv ── +tauri_log "STEP" "Installing uv package manager" UV_MIN_VERSION="0.7.14" version_ge() { @@ -922,6 +951,7 @@ if ! command -v uv >/dev/null 2>&1 || ! _uv_version_ok uv; then fi # ── Create venv (migrate old layout if possible, otherwise fresh) ── +tauri_log "STEP" "Creating virtual environment" mkdir -p "$STUDIO_HOME" _MIGRATED=false @@ -1304,6 +1334,7 @@ case "$TORCH_INDEX_URL" in esac # ── Install unsloth directly into the venv (no activation needed) ── +tauri_log "STEP" "Installing PyTorch" _VENV_PY="$VENV_DIR/bin/python" if [ "$_MIGRATED" = true ]; then # Migrated env: force-reinstall unsloth+unsloth-zoo to ensure clean state @@ -1316,7 +1347,7 @@ if [ "$_MIGRATED" = true ]; then # to prevent transitive torch resolution. run_install_cmd "install unsloth (migrated no-torch)" uv pip install --python "$_VENV_PY" --no-deps \ --reinstall-package unsloth --reinstall-package unsloth-zoo \ - "unsloth>=2026.4.5" unsloth-zoo + "unsloth>=2026.4.7" unsloth-zoo _NO_TORCH_RT="$(_find_no_torch_runtime)" if [ -n "$_NO_TORCH_RT" ]; then run_install_cmd "install no-torch runtime deps" uv pip install --python "$_VENV_PY" --no-deps -r "$_NO_TORCH_RT" @@ -1324,7 +1355,7 @@ if [ "$_MIGRATED" = true ]; then else run_install_cmd "install unsloth (migrated)" uv pip install --python "$_VENV_PY" \ --reinstall-package unsloth --reinstall-package unsloth-zoo \ - "unsloth>=2026.4.5" unsloth-zoo + "unsloth>=2026.4.7" unsloth-zoo fi if [ "$STUDIO_LOCAL_INSTALL" = true ]; then substep "overlaying local repo (editable)..." @@ -1481,13 +1512,14 @@ elif [ -n "$TORCH_INDEX_URL" ]; then esac fi # Fresh: Step 2 - install unsloth, preserving pre-installed torch + tauri_log "STEP" "Installing Unsloth" substep "installing unsloth (this may take a few minutes)..." if [ "$SKIP_TORCH" = true ]; then # No-torch: install unsloth + unsloth-zoo with --no-deps, then # runtime deps (typer, safetensors, transformers, etc.) with --no-deps. run_install_cmd "install unsloth (no-torch)" uv pip install --python "$_VENV_PY" --no-deps \ --upgrade-package unsloth --upgrade-package unsloth-zoo \ - "unsloth>=2026.4.5" unsloth-zoo + "unsloth>=2026.4.7" unsloth-zoo _NO_TORCH_RT="$(_find_no_torch_runtime)" if [ -n "$_NO_TORCH_RT" ]; then run_install_cmd "install no-torch runtime deps" uv pip install --python "$_VENV_PY" --no-deps -r "$_NO_TORCH_RT" @@ -1498,12 +1530,12 @@ elif [ -n "$TORCH_INDEX_URL" ]; then fi elif [ "$STUDIO_LOCAL_INSTALL" = true ]; then run_install_cmd "install unsloth (local)" uv pip install --python "$_VENV_PY" \ - --upgrade-package unsloth "unsloth>=2026.4.5" unsloth-zoo + --upgrade-package unsloth "unsloth>=2026.4.7" unsloth-zoo substep "overlaying local repo (editable)..." run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps else run_install_cmd "install unsloth" uv pip install --python "$_VENV_PY" \ - --upgrade-package unsloth "$PACKAGE_NAME" + --upgrade-package unsloth -- "$PACKAGE_NAME" fi # AMD ROCm: repair torch if the unsloth/unsloth-zoo install pulled in # CUDA torch from PyPI, overwriting the ROCm wheels installed in Step 1. @@ -1523,17 +1555,19 @@ elif [ -n "$TORCH_INDEX_URL" ]; then fi else # Fallback: GPU detection failed to produce a URL -- let uv resolve torch + tauri_log "STEP" "Installing Unsloth" substep "installing unsloth (this may take a few minutes)..." if [ "$STUDIO_LOCAL_INSTALL" = true ]; then - run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" unsloth-zoo "unsloth>=2026.4.5" --torch-backend=auto + run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" unsloth-zoo "unsloth>=2026.4.7" --torch-backend=auto substep "overlaying local repo (editable)..." run_install_cmd "overlay local repo" uv pip install --python "$_VENV_PY" -e "$_REPO_ROOT" --no-deps else - run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" "$PACKAGE_NAME" --torch-backend=auto + run_install_cmd "install unsloth (auto torch backend)" uv pip install --python "$_VENV_PY" --torch-backend=auto -- "$PACKAGE_NAME" fi fi # ── Run studio setup ── +tauri_log "STEP" "Running Studio setup" # When --local, use the repo's own setup.sh directly. # Otherwise, find it inside the installed package. SETUP_SH="" @@ -1554,6 +1588,7 @@ if [ -z "$SETUP_SH" ] || [ ! -f "$SETUP_SH" ]; then fi if [ -z "$SETUP_SH" ] || [ ! -f "$SETUP_SH" ]; then + tauri_log "ERROR" "Could not find studio/setup.sh in the installed package" echo "❌ ERROR: Could not find studio/setup.sh in the installed package." exit 1 fi @@ -1571,24 +1606,32 @@ if ! command -v bash >/dev/null 2>&1; then fi step "setup" "running unsloth studio update..." -# install.sh already installs base packages (unsloth + unsloth-zoo) and -# no-torch-runtime.txt above, so tell install_python_stack.py to skip -# the base step to avoid redundant reinstallation. _SKIP_BASE=1 -# Run setup.sh outside set -e so that a llama.cpp build failure (exit 1) -# does not skip PATH setup, shortcuts, and launch below. We capture the -# exit code and propagate it after post-install steps finish. _SETUP_EXIT=0 +# Tauri desktop app bundles its own frontend — skip Node/npm/frontend build +_SKIP_FRONTEND=0 +if [ "$TAURI_MODE" = true ]; then + _SKIP_FRONTEND=1 +fi if [ "$STUDIO_LOCAL_INSTALL" = true ]; then SKIP_STUDIO_BASE="$_SKIP_BASE" \ + SKIP_STUDIO_FRONTEND="$_SKIP_FRONTEND" \ STUDIO_PACKAGE_NAME="$PACKAGE_NAME" \ STUDIO_LOCAL_INSTALL=1 \ STUDIO_LOCAL_REPO="$_REPO_ROOT" \ UNSLOTH_NO_TORCH="$SKIP_TORCH" \ bash "$SETUP_SH" Optional[str]: def create_access_token( subject: str, expires_delta: Optional[timedelta] = None, + *, + desktop: bool = False, ) -> str: """ Create a signed JWT for the given subject (e.g. username). @@ -59,6 +61,8 @@ def create_access_token( Tokens are valid across restarts because the signing secret is stored in SQLite. """ to_encode = {"sub": subject} + if desktop: + to_encode["desktop"] = True expire = datetime.now(timezone.utc) + ( expires_delta or timedelta(minutes = ACCESS_TOKEN_EXPIRE_MINUTES) ) @@ -70,7 +74,29 @@ def create_access_token( ) -def create_refresh_token(subject: str) -> str: +def is_desktop_access_token(token: str) -> bool: + """Return true only for a valid desktop-issued JWT access token.""" + if token.startswith(API_KEY_PREFIX): + return False + + subject = _decode_subject_without_verification(token) + if subject is None: + return False + + record = get_user_and_secret(subject) + if record is None: + return False + + _salt, _pwd_hash, jwt_secret, _must_change_password = record + try: + payload = jwt.decode(token, jwt_secret, algorithms = [ALGORITHM]) + except jwt.InvalidTokenError: + return False + + return payload.get("sub") == subject and payload.get("desktop") is True + + +def create_refresh_token(subject: str, *, desktop: bool = False) -> str: """ Create a random refresh token, store its hash in SQLite, and return it. @@ -78,21 +104,28 @@ def create_refresh_token(subject: str) -> str: """ token = secrets.token_urlsafe(48) expires_at = datetime.now(timezone.utc) + timedelta(days = REFRESH_TOKEN_EXPIRE_DAYS) - save_refresh_token(token, subject, expires_at.isoformat()) + save_refresh_token(token, subject, expires_at.isoformat(), is_desktop = desktop) return token -def refresh_access_token(refresh_token: str) -> Tuple[Optional[str], Optional[str]]: +def refresh_access_token( + refresh_token: str, +) -> Tuple[Optional[str], Optional[str], bool]: """ Validate a refresh token and issue a new access token. The refresh token itself is NOT consumed — it stays valid until expiry. Returns a new access_token or None if the refresh token is invalid/expired. """ - username = verify_refresh_token(refresh_token) - if username is None: - return None, None - return create_access_token(subject = username), username + verified = verify_refresh_token(refresh_token) + if verified is None: + return None, None, False + username, is_desktop = verified + return ( + create_access_token(subject = username, desktop = is_desktop), + username, + is_desktop, + ) def reload_secret() -> None: @@ -173,7 +206,8 @@ async def _get_current_subject( status_code = status.HTTP_401_UNAUTHORIZED, detail = "Invalid token payload", ) - if must_change_password and not allow_password_change: + is_desktop = payload.get("desktop") is True + if must_change_password and not allow_password_change and not is_desktop: raise HTTPException( status_code = status.HTTP_403_FORBIDDEN, detail = "Password change required", diff --git a/studio/backend/auth/storage.py b/studio/backend/auth/storage.py index 7d55a2dc59..2b0e359d39 100644 --- a/studio/backend/auth/storage.py +++ b/studio/backend/auth/storage.py @@ -6,6 +6,7 @@ SQLite storage for authentication data (user credentials + JWT secret). """ import hashlib +import os import secrets import sqlite3 from datetime import datetime, timezone @@ -54,6 +55,10 @@ def generate_bootstrap_password() -> str: # before the user changes the password. ensure_dir(_BOOTSTRAP_PW_PATH.parent) _BOOTSTRAP_PW_PATH.write_text(_bootstrap_password) + try: + os.chmod(_BOOTSTRAP_PW_PATH, 0o600) + except OSError: + pass return _bootstrap_password @@ -63,6 +68,17 @@ def get_bootstrap_password() -> Optional[str]: return _bootstrap_password +def _load_bootstrap_password() -> Optional[str]: + """Load an existing bootstrap password without creating one.""" + global _bootstrap_password + _bootstrap_password = None + if _BOOTSTRAP_PW_PATH.is_file(): + bootstrap_password = _BOOTSTRAP_PW_PATH.read_text().strip() + if bootstrap_password: + _bootstrap_password = bootstrap_password + return _bootstrap_password + + def clear_bootstrap_password() -> None: """Delete the persisted bootstrap password file (called after password change).""" global _bootstrap_password @@ -114,7 +130,8 @@ def get_connection() -> sqlite3.Connection: id INTEGER PRIMARY KEY, token_hash TEXT NOT NULL, username TEXT NOT NULL, - expires_at TEXT NOT NULL + expires_at TEXT NOT NULL, + is_desktop INTEGER NOT NULL DEFAULT 0 ); """ ) @@ -146,6 +163,13 @@ def get_connection() -> sqlite3.Connection: conn.execute( "ALTER TABLE auth_user ADD COLUMN must_change_password INTEGER NOT NULL DEFAULT 0" ) + refresh_columns = { + row["name"] for row in conn.execute("PRAGMA table_info(refresh_tokens)") + } + if "is_desktop" not in refresh_columns: + conn.execute( + "ALTER TABLE refresh_tokens ADD COLUMN is_desktop INTEGER NOT NULL DEFAULT 0" + ) conn.commit() return conn @@ -201,6 +225,9 @@ def _get_or_create_api_key_pbkdf2_salt() -> bytes: _API_KEY_PBKDF2_ITERATIONS = 100_000 +DESKTOP_SECRET_PREFIX = "desktop-" +_DESKTOP_SECRET_HASH_KEY = "desktop_secret_hash" +_DESKTOP_SECRET_CREATED_AT_KEY = "desktop_secret_created_at" def _pbkdf2_api_key(raw_key: str) -> str: @@ -233,6 +260,10 @@ def _pbkdf2_api_key(raw_key: str) -> str: return dk.hex() +def _pbkdf2_desktop_secret(raw_secret: str) -> str: + return _pbkdf2_api_key(raw_secret) + + def is_initialized() -> bool: """Check if auth is ready for login (at least one user exists in DB).""" conn = get_connection() @@ -374,6 +405,10 @@ def ensure_default_admin() -> bool: Uses a randomly generated diceware passphrase as the bootstrap password. Returns True when the default admin was created in this call. """ + if get_user_and_secret(DEFAULT_ADMIN_USERNAME) is not None: + _load_bootstrap_password() + return False + bootstrap_pw = generate_bootstrap_password() try: create_initial_user( @@ -406,12 +441,19 @@ def update_password(username: str, new_password: str) -> bool: conn.commit() if cursor.rowcount > 0: clear_bootstrap_password() + clear_desktop_secret() return cursor.rowcount > 0 finally: conn.close() -def save_refresh_token(token: str, username: str, expires_at: str) -> None: +def save_refresh_token( + token: str, + username: str, + expires_at: str, + *, + is_desktop: bool = False, +) -> None: """ Store a hashed refresh token with its associated username and expiry. """ @@ -420,21 +462,21 @@ def save_refresh_token(token: str, username: str, expires_at: str) -> None: try: conn.execute( """ - INSERT INTO refresh_tokens (token_hash, username, expires_at) - VALUES (?, ?, ?) + INSERT INTO refresh_tokens (token_hash, username, expires_at, is_desktop) + VALUES (?, ?, ?, ?) """, - (token_hash, username, expires_at), + (token_hash, username, expires_at, int(is_desktop)), ) conn.commit() finally: conn.close() -def verify_refresh_token(token: str) -> Optional[str]: +def verify_refresh_token(token: str) -> Optional[Tuple[str, bool]]: """ - Verify a refresh token and return the username. + Verify a refresh token and return the username plus desktop marker. - Returns the username if valid and not expired, None otherwise. + Returns the username and desktop marker if valid and not expired, None otherwise. The token is NOT consumed — it stays valid until it expires. """ token_hash = _hash_token(token) @@ -449,7 +491,7 @@ def verify_refresh_token(token: str) -> Optional[str]: cur = conn.execute( """ - SELECT id, username, expires_at FROM refresh_tokens + SELECT id, username, expires_at, is_desktop FROM refresh_tokens WHERE token_hash = ? """, (token_hash,), @@ -465,7 +507,7 @@ def verify_refresh_token(token: str) -> Optional[str]: conn.commit() return None - return row["username"] + return row["username"], bool(row["is_desktop"]) finally: conn.close() @@ -480,6 +522,65 @@ def revoke_user_refresh_tokens(username: str) -> None: conn.close() +def create_desktop_secret() -> str: + """Create/rotate the local desktop credential and return it once.""" + ensure_default_admin() + raw_secret = DESKTOP_SECRET_PREFIX + secrets.token_urlsafe(48) + secret_hash = _pbkdf2_desktop_secret(raw_secret) + now = datetime.now(timezone.utc).isoformat() + conn = get_connection() + try: + conn.execute( + "INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?, ?)", + (_DESKTOP_SECRET_HASH_KEY, secret_hash), + ) + conn.execute( + "INSERT OR REPLACE INTO app_secrets (key, value) VALUES (?, ?)", + (_DESKTOP_SECRET_CREATED_AT_KEY, now), + ) + conn.commit() + return raw_secret + finally: + conn.close() + + +def validate_desktop_secret(raw_secret: str) -> Optional[str]: + """Return the real admin username when the desktop secret matches.""" + if not raw_secret.startswith(DESKTOP_SECRET_PREFIX): + return None + if get_user_and_secret(DEFAULT_ADMIN_USERNAME) is None: + return None + + secret_hash = _pbkdf2_desktop_secret(raw_secret) + conn = get_connection() + try: + cur = conn.execute( + "SELECT value FROM app_secrets WHERE key = ?", + (_DESKTOP_SECRET_HASH_KEY,), + ) + row = cur.fetchone() + if row is None: + return None + if not secrets.compare_digest(row["value"], secret_hash): + return None + return DEFAULT_ADMIN_USERNAME + finally: + conn.close() + + +def clear_desktop_secret() -> None: + """Remove backend-side desktop auth state.""" + conn = get_connection() + try: + conn.execute( + "DELETE FROM app_secrets WHERE key IN (?, ?)", + (_DESKTOP_SECRET_HASH_KEY, _DESKTOP_SECRET_CREATED_AT_KEY), + ) + conn.commit() + finally: + conn.close() + + # --------------------------------------------------------------------------- # API key management # --------------------------------------------------------------------------- diff --git a/studio/backend/core/data_recipe/local_callable_validators.py b/studio/backend/core/data_recipe/local_callable_validators.py index c32b2fccaf..afd10b02d1 100644 --- a/studio/backend/core/data_recipe/local_callable_validators.py +++ b/studio/backend/core/data_recipe/local_callable_validators.py @@ -33,6 +33,11 @@ _OXC_TOOL_DIR = Path(__file__).resolve().parent / "oxc-validator" _OXC_RUNNER_PATH = _OXC_TOOL_DIR / "validate.mjs" +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + + @dataclass(frozen = True) class OxcLocalCallableValidatorSpec: name: str @@ -256,6 +261,7 @@ def _run_oxc_batch( capture_output = True, check = False, env = env, + **_windows_hidden_subprocess_kwargs(), ) except (OSError, ValueError) as exc: logger.warning("OXC subprocess launch failed: %s", exc) diff --git a/studio/backend/core/inference/audio_codecs.py b/studio/backend/core/inference/audio_codecs.py index bcf3ec2937..895b112e85 100644 --- a/studio/backend/core/inference/audio_codecs.py +++ b/studio/backend/core/inference/audio_codecs.py @@ -8,6 +8,7 @@ Supports: SNAC (Orpheus), CSM (Sesame), BiCodec (Spark), DAC (OuteTTS) import io import re +import subprocess import wave import structlog from loggers import get_logger @@ -16,6 +17,10 @@ from typing import Optional, Tuple import numpy as np import torch +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + logger = get_logger(__name__) @@ -81,7 +86,6 @@ class AudioCodecManager: return import os import sys - import subprocess # Clone SparkAudio/Spark-TTS GitHub repo for the sparktts Python package # (same approach as training — the HF model repos don't contain the package) @@ -101,6 +105,7 @@ class AudioCodecManager: spark_code_dir, ], check = True, + **_windows_hidden_subprocess_kwargs(), ) if spark_code_dir not in sys.path: @@ -119,7 +124,6 @@ class AudioCodecManager: return import os import sys - import subprocess # Clone OuteTTS repo (same pattern as Spark-TTS / BiCodec) # The pip package has problematic dependencies; the notebook clones and @@ -139,6 +143,7 @@ class AudioCodecManager: outetts_code_dir, ], check = True, + **_windows_hidden_subprocess_kwargs(), ) # Remove files that pull in heavy / incompatible dependencies # (matches notebook: gguf_model.py is under models/, others under outetts/) diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 2e26995309..c320f03b2c 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -18,6 +18,7 @@ from loggers import get_logger import shutil import socket import subprocess +import sys import threading import time from pathlib import Path @@ -26,8 +27,13 @@ from urllib.parse import urlparse import httpx +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + logger = get_logger(__name__) + # ── Pre-compiled patterns for plan-without-action re-prompt ── # Forward-looking intent signals that indicate the model is # describing what it *will* do rather than giving a final answer. @@ -81,6 +87,84 @@ _TC_PARAM_START_RE = re.compile(r"\s*") _TC_PARAM_CLOSE_RE = re.compile(r"\s*\s*$") +_TOOL_TEMPLATE_MARKERS = ( + "{%- if tools %}", + "{%- if tools -%}", + "{% if tools %}", + "{% if tools -%}", + '"role" == "tool"', + "'role' == 'tool'", + 'message.role == "tool"', + "message.role == 'tool'", +) + + +def detect_reasoning_flags( + chat_template: Optional[str], + model_identifier: Optional[str] = None, + *, + log_source: Optional[str] = None, +) -> dict: + """Classify a chat template's reasoning and tool-calling capabilities. + + Returns a dict with the same five keys populated by the GGUF sniffer: + ``supports_reasoning``, ``reasoning_style`` + (``"enable_thinking"`` | ``"reasoning_effort"``), + ``reasoning_always_on``, ``supports_preserve_thinking``, and + ``supports_tools``. Used by both the llama-server backend at load + time and the safetensors/transformers paths in ``routes/inference`` + so the two agree on what the frontend will see. + """ + flags = { + "supports_reasoning": False, + "reasoning_style": "enable_thinking", + "reasoning_always_on": False, + "supports_preserve_thinking": False, + "supports_tools": False, + } + if not chat_template: + return flags + tpl = chat_template + prefix = f"{log_source}: " if log_source else "" + + if "enable_thinking" in tpl: + flags["supports_reasoning"] = True + flags["reasoning_style"] = "enable_thinking" + logger.info(f"{prefix}model supports reasoning (enable_thinking)") + elif "reasoning_effort" in tpl: + # gpt-oss / Harmony templates use reasoning_effort + # ("low" | "medium" | "high") instead of a boolean. + flags["supports_reasoning"] = True + flags["reasoning_style"] = "reasoning_effort" + logger.info(f"{prefix}model supports reasoning (reasoning_effort)") + elif "thinking" in tpl: + # DeepSeek uses 'thinking' instead of 'enable_thinking' + normalized_id = (model_identifier or "").lower() + if "deepseek" in normalized_id: + flags["supports_reasoning"] = True + logger.info(f"{prefix}model supports reasoning (DeepSeek thinking)") + + # Hardcoded tags or reasoning_content in the template mean + # thinking is always on (no toggle to disable it). + if not flags["supports_reasoning"]: + if ("" in tpl and "" in tpl) or "reasoning_content" in tpl: + flags["supports_reasoning"] = True + flags["reasoning_always_on"] = True + logger.info(f"{prefix}model always reasons ( tags in template)") + + # preserve_thinking is an independent kwarg on some Qwen templates + # that keeps historical blocks in prior assistant turns. + if "preserve_thinking" in tpl: + flags["supports_preserve_thinking"] = True + logger.info(f"{prefix}model supports preserve_thinking") + + if any(marker in tpl for marker in _TOOL_TEMPLATE_MARKERS): + flags["supports_tools"] = True + logger.info(f"{prefix}model supports tool calling") + + return flags + + class LlamaCppBackend: """ Manages a llama-server subprocess for GGUF model inference. @@ -106,6 +190,8 @@ class LlamaCppBackend: self._chat_template: Optional[str] = None self._supports_reasoning: bool = False self._reasoning_always_on: bool = False + self._reasoning_style: str = "enable_thinking" + self._supports_preserve_thinking: bool = False self._supports_tools: bool = False self._cache_type_kv: Optional[str] = None self._reasoning_default: bool = True @@ -287,10 +373,51 @@ class LlamaCppBackend: def reasoning_always_on(self) -> bool: return self._reasoning_always_on + @property + def reasoning_style(self) -> str: + return self._reasoning_style + + @property + def supports_preserve_thinking(self) -> bool: + return self._supports_preserve_thinking + @property def reasoning_default(self) -> bool: return self._reasoning_default + def _reasoning_kwargs(self, enable_thinking: bool) -> dict: + if self._reasoning_style == "reasoning_effort": + return {"reasoning_effort": "high" if enable_thinking else "low"} + return {"enable_thinking": enable_thinking} + + def _request_reasoning_kwargs( + self, + enable_thinking: Optional[bool], + reasoning_effort: Optional[str] = None, + preserve_thinking: Optional[bool] = None, + ) -> Optional[dict]: + """Build chat_template_kwargs from per-request reasoning fields. + + Produces a merged dict covering the active model's reasoning style + (``enable_thinking`` or ``reasoning_effort``) plus the independent + ``preserve_thinking`` kwarg when the template supports it. + """ + kwargs: dict = {} + # Always-on reasoning models hardcode tags in their template + # and do not consume enable_thinking / reasoning_effort -- skip. + if self._supports_reasoning and not self._reasoning_always_on: + if self._reasoning_style == "reasoning_effort": + if reasoning_effort in ("low", "medium", "high"): + kwargs["reasoning_effort"] = reasoning_effort + elif enable_thinking is not None: + kwargs["reasoning_effort"] = "high" if enable_thinking else "low" + else: + if enable_thinking is not None: + kwargs["enable_thinking"] = enable_thinking + if self._supports_preserve_thinking and preserve_thinking is not None: + kwargs["preserve_thinking"] = preserve_thinking + return kwargs or None + @property def supports_tools(self) -> bool: return self._supports_tools @@ -440,6 +567,7 @@ class LlamaCppBackend: capture_output = True, text = True, timeout = 10, + **_windows_hidden_subprocess_kwargs(), ) if result.returncode != 0: return [] @@ -811,6 +939,9 @@ class LlamaCppBackend: self._chat_template = None self._supports_reasoning = False self._reasoning_always_on = False + self._reasoning_style = "enable_thinking" + self._reasoning_default = True + self._supports_preserve_thinking = False self._supports_tools = False self._n_layers = None self._n_kv_heads = None @@ -891,48 +1022,16 @@ class LlamaCppBackend: f"GGUF metadata: chat_template={len(self._chat_template)} chars" ) # Detect thinking/reasoning support from chat template - tpl = self._chat_template - if "enable_thinking" in tpl: - self._supports_reasoning = True - logger.info( - "GGUF metadata: model supports reasoning (enable_thinking)" - ) - elif "thinking" in tpl: - # DeepSeek uses 'thinking' instead of 'enable_thinking' - normalized_id = (self._model_identifier or "").lower() - if "deepseek" in normalized_id: - self._supports_reasoning = True - logger.info( - "GGUF metadata: model supports reasoning (DeepSeek thinking)" - ) - # Models with hardcoded tags or reasoning_content - # in their chat template always produce thinking output - # (no toggle to disable it). - if not self._supports_reasoning: - if ( - "" in tpl - and "" in tpl - or "reasoning_content" in tpl - ): - self._supports_reasoning = True - self._reasoning_always_on = True - logger.info( - "GGUF metadata: model always reasons ( tags in template)" - ) - # Detect tool calling support from chat template - tool_markers = [ - "{%- if tools %}", - "{%- if tools -%}", - "{% if tools %}", - "{% if tools -%}", - '"role" == "tool"', - "'role' == 'tool'", - 'message.role == "tool"', - "message.role == 'tool'", - ] - if any(marker in tpl for marker in tool_markers): - self._supports_tools = True - logger.info("GGUF metadata: model supports tool calling") + flags = detect_reasoning_flags( + self._chat_template, + self._model_identifier, + log_source = "GGUF metadata", + ) + self._supports_reasoning = flags["supports_reasoning"] + self._reasoning_style = flags["reasoning_style"] + self._reasoning_always_on = flags["reasoning_always_on"] + self._supports_preserve_thinking = flags["supports_preserve_thinking"] + self._supports_tools = flags["supports_tools"] except Exception as e: logger.warning(f"Failed to read GGUF metadata: {e}") @@ -1516,7 +1615,8 @@ class LlamaCppBackend: # For reasoning models, set default thinking mode. # Qwen3.5/3.6 models below 9B (0.8B, 2B, 4B) disable thinking by default. # Only 9B and larger enable thinking. - if self._supports_reasoning: + # Always-on templates ignore the kwarg entirely, so skip. + if self._supports_reasoning and not self._reasoning_always_on: thinking_default = True mid = (model_identifier or "").lower() if "qwen3.5" in mid or "qwen3.6" in mid: @@ -1524,15 +1624,14 @@ class LlamaCppBackend: if size_val is not None and size_val < 9: thinking_default = False self._reasoning_default = thinking_default + reasoning_kw = self._reasoning_kwargs(thinking_default) cmd.extend( [ "--chat-template-kwargs", - json.dumps({"enable_thinking": thinking_default}), + json.dumps(reasoning_kw), ] ) - logger.info( - f"Reasoning model: enable_thinking={thinking_default} by default" - ) + logger.info(f"Reasoning model: {reasoning_kw} by default") if mmproj_path: if not Path(mmproj_path).is_file(): @@ -1658,6 +1757,7 @@ class LlamaCppBackend: stderr = subprocess.STDOUT, text = True, env = env, + **_windows_hidden_subprocess_kwargs(), ) # Start background thread to drain stdout and prevent pipe deadlock @@ -1759,6 +1859,9 @@ class LlamaCppBackend: self._chat_template = None self._supports_reasoning = False self._reasoning_always_on = False + self._reasoning_style = "enable_thinking" + self._reasoning_default = True + self._supports_preserve_thinking = False self._supports_tools = False self._cache_type_kv = None self._speculative_type = None @@ -2309,6 +2412,8 @@ class LlamaCppBackend: stop: Optional[list[str]] = None, cancel_event: Optional[threading.Event] = None, enable_thinking: Optional[bool] = None, + reasoning_effort: Optional[str] = None, + preserve_thinking: Optional[bool] = None, ) -> Generator[str | dict, None, None]: """ Send a chat completion request to llama-server and stream tokens back. @@ -2333,9 +2438,12 @@ class LlamaCppBackend: "repeat_penalty": repetition_penalty, "presence_penalty": presence_penalty, } - # Pass enable_thinking per-request for reasoning models - if self._supports_reasoning and enable_thinking is not None: - payload["chat_template_kwargs"] = {"enable_thinking": enable_thinking} + # Pass enable_thinking / reasoning_effort / preserve_thinking per-request + _reasoning_kw = self._request_reasoning_kwargs( + enable_thinking, reasoning_effort, preserve_thinking + ) + if _reasoning_kw is not None: + payload["chat_template_kwargs"] = _reasoning_kw if max_tokens is not None: payload["max_tokens"] = max_tokens if stop: @@ -2471,6 +2579,8 @@ class LlamaCppBackend: stop: Optional[list[str]] = None, cancel_event: Optional[threading.Event] = None, enable_thinking: Optional[bool] = None, + reasoning_effort: Optional[str] = None, + preserve_thinking: Optional[bool] = None, max_tool_iterations: int = 25, auto_heal_tool_calls: bool = True, tool_call_timeout: int = 300, @@ -2549,8 +2659,11 @@ class LlamaCppBackend: "tools": tools, "tool_choice": "auto", } - if self._supports_reasoning and enable_thinking is not None: - payload["chat_template_kwargs"] = {"enable_thinking": enable_thinking} + _reasoning_kw = self._request_reasoning_kwargs( + enable_thinking, reasoning_effort, preserve_thinking + ) + if _reasoning_kw is not None: + payload["chat_template_kwargs"] = _reasoning_kw if max_tokens is not None: payload["max_tokens"] = max_tokens if stop: @@ -3199,10 +3312,11 @@ class LlamaCppBackend: "repeat_penalty": repetition_penalty, "presence_penalty": presence_penalty, } - if self._supports_reasoning and enable_thinking is not None: - stream_payload["chat_template_kwargs"] = { - "enable_thinking": enable_thinking - } + _reasoning_kw = self._request_reasoning_kwargs( + enable_thinking, reasoning_effort, preserve_thinking + ) + if _reasoning_kw is not None: + stream_payload["chat_template_kwargs"] = _reasoning_kw if max_tokens is not None: stream_payload["max_tokens"] = max_tokens if stop: diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 77cbda6b45..c1e2ac4a85 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -49,6 +49,7 @@ from unsloth.chat_templates import get_chat_template import json import threading import math +import subprocess import structlog from loggers import get_logger import time @@ -69,6 +70,10 @@ from utils.paths import ( ) from trl import SFTTrainer, SFTConfig +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + logger = get_logger(__name__) @@ -1765,6 +1770,7 @@ class UnslothTrainer: spark_code_dir, ], check = True, + **_windows_hidden_subprocess_kwargs(), ) if spark_code_dir not in sys.path: @@ -1982,8 +1988,6 @@ class UnslothTrainer: device = "cuda" if torch.cuda.is_available() else "cpu" # Clone OuteTTS repo (same as audio_codecs._load_dac) - import subprocess - base_dir = os.path.dirname(os.path.abspath(__file__)) outetts_code_dir = os.path.join(base_dir, "inference", "OuteTTS") outetts_pkg = os.path.join(outetts_code_dir, "outetts") @@ -2000,6 +2004,7 @@ class UnslothTrainer: outetts_code_dir, ], check = True, + **_windows_hidden_subprocess_kwargs(), ) for fpath in [ os.path.join(outetts_pkg, "models", "gguf_model.py"), diff --git a/studio/backend/loggers/config.py b/studio/backend/loggers/config.py index 0d32a64657..4c0d8ade28 100644 --- a/studio/backend/loggers/config.py +++ b/studio/backend/loggers/config.py @@ -44,6 +44,14 @@ class LogConfig: # Fallback to INFO if an invalid level is provided log_level = getattr(logging, log_level_name, logging.INFO) + if sys.platform == "win32": + for stream in (sys.stdout, sys.stderr): + if hasattr(stream, "reconfigure"): + try: + stream.reconfigure(encoding = "utf-8", errors = "replace") + except Exception: + pass + structlog.configure( processors = [ # Reorder processors to control field order diff --git a/studio/backend/main.py b/studio/backend/main.py index d146a8ef12..05adcaa2ea 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -179,9 +179,22 @@ logger = LogConfig.setup_logging( app.add_middleware(LoggingMiddleware) # CORS middleware +_api_only = os.environ.get("UNSLOTH_API_ONLY") == "1" +_cors_origins = ["*"] +if _api_only: + _cors_origins = [ + "tauri://localhost", # Linux/macOS Tauri webview + "http://tauri.localhost", # Windows Tauri webview + "http://localhost", # dev fallback + ] + _cors_origin_regex = None +else: + _cors_origin_regex = None + app.add_middleware( CORSMiddleware, - allow_origins = ["*"], # In production, specify allowed origins + allow_origins = _cors_origins, + allow_origin_regex = _cors_origin_regex, allow_credentials = True, allow_methods = ["*"], allow_headers = ["*"], @@ -223,6 +236,8 @@ async def health_check(): "version": UNSLOTH_VERSION, "device_type": device_type, "chat_only": _hw_module.CHAT_ONLY, + "desktop_protocol_version": 1, + "supports_desktop_auth": True, } diff --git a/studio/backend/models/auth.py b/studio/backend/models/auth.py index c55e646508..23eb0ac4c0 100644 --- a/studio/backend/models/auth.py +++ b/studio/backend/models/auth.py @@ -17,6 +17,12 @@ class AuthLoginRequest(BaseModel): password: str = Field(..., description = "Password") +class DesktopLoginRequest(BaseModel): + """Desktop-only local secret exchange payload.""" + + secret: str = Field(..., description = "Desktop local auth secret") + + class RefreshTokenRequest(BaseModel): """Refresh token payload to obtain new access + refresh tokens.""" diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index 324002ddbf..e5b037755d 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -157,12 +157,20 @@ class LoadResponse(BaseModel): ) supports_reasoning: bool = Field( False, - description = "Whether model supports thinking/reasoning mode (enable_thinking)", + description = "Whether model supports thinking/reasoning mode (enable_thinking or reasoning_effort)", + ) + reasoning_style: Literal["enable_thinking", "reasoning_effort"] = Field( + "enable_thinking", + description = "Reasoning control style: 'enable_thinking' (boolean) or 'reasoning_effort' (low|medium|high)", ) reasoning_always_on: bool = Field( False, description = "Whether reasoning is always on (hardcoded tags, not toggleable)", ) + supports_preserve_thinking: bool = Field( + False, + description = "Whether the template understands the optional preserve_thinking kwarg (Qwen3.6-style)", + ) supports_tools: bool = Field( False, description = "Whether model supports tool calling (web search, etc.)", @@ -261,9 +269,17 @@ class InferenceStatusResponse(BaseModel): supports_reasoning: bool = Field( False, description = "Whether the active model supports reasoning/thinking mode" ) + reasoning_style: Literal["enable_thinking", "reasoning_effort"] = Field( + "enable_thinking", + description = "Reasoning control style: 'enable_thinking' (boolean) or 'reasoning_effort' (low|medium|high)", + ) reasoning_always_on: bool = Field( False, description = "Whether reasoning is always on (not toggleable)" ) + supports_preserve_thinking: bool = Field( + False, + description = "Whether the active model's template understands the optional preserve_thinking kwarg", + ) supports_tools: bool = Field( False, description = "Whether the active model supports tool calling" ) @@ -481,6 +497,14 @@ class ChatCompletionRequest(BaseModel): None, description = "[x-unsloth] Enable/disable thinking/reasoning mode for supported models", ) + reasoning_effort: Optional[Literal["low", "medium", "high"]] = Field( + None, + description = "[x-unsloth] Reasoning effort level ('low'|'medium'|'high') for Harmony-style reasoning models (e.g. gpt-oss). Overrides enable_thinking when the active model uses reasoning_effort style.", + ) + preserve_thinking: Optional[bool] = Field( + None, + description = "[x-unsloth] When true, keep historical blocks from past assistant turns in the prompt (Qwen3.6 templates). Independent of enable_thinking / reasoning_effort.", + ) enable_tools: Optional[bool] = Field( None, description = "[x-unsloth] Enable tool calling for supported models", diff --git a/studio/backend/routes/auth.py b/studio/backend/routes/auth.py index 5cd23bd450..3deeb6793b 100644 --- a/studio/backend/routes/auth.py +++ b/studio/backend/routes/auth.py @@ -17,6 +17,7 @@ from models.auth import ( ChangePasswordRequest, CreateApiKeyRequest, CreateApiKeyResponse, + DesktopLoginRequest, RefreshTokenRequest, ) from models.users import Token @@ -80,6 +81,24 @@ async def login(payload: AuthLoginRequest) -> Token: ) +@router.post("/desktop-login", response_model = Token) +async def desktop_login(payload: DesktopLoginRequest) -> Token: + """Exchange a local desktop secret for normal admin-subject tokens.""" + username = storage.validate_desktop_secret(payload.secret) + if username is None: + raise HTTPException( + status_code = status.HTTP_401_UNAUTHORIZED, + detail = "Desktop authentication failed", + ) + + return Token( + access_token = create_access_token(subject = username, desktop = True), + refresh_token = create_refresh_token(subject = username, desktop = True), + token_type = "bearer", + must_change_password = False, + ) + + @router.post("/refresh", response_model = Token) async def refresh(payload: RefreshTokenRequest) -> Token: """ @@ -87,7 +106,7 @@ async def refresh(payload: RefreshTokenRequest) -> Token: The refresh token itself is reusable until it expires (7 days). """ - new_access_token, username = refresh_access_token(payload.refresh_token) + new_access_token, username, is_desktop = refresh_access_token(payload.refresh_token) if new_access_token is None or username is None: raise HTTPException( status_code = status.HTTP_401_UNAUTHORIZED, @@ -98,7 +117,9 @@ async def refresh(payload: RefreshTokenRequest) -> Token: access_token = new_access_token, refresh_token = payload.refresh_token, token_type = "bearer", - must_change_password = storage.requires_password_change(username), + must_change_password = False + if is_desktop + else storage.requires_password_change(username), ) diff --git a/studio/backend/routes/data_recipe/jobs.py b/studio/backend/routes/data_recipe/jobs.py index 00546b47a4..606ef1832c 100644 --- a/studio/backend/routes/data_recipe/jobs.py +++ b/studio/backend/routes/data_recipe/jobs.py @@ -57,6 +57,20 @@ def _resolve_local_v1_endpoint(request: Request) -> str: return f"http://127.0.0.1:{int(port)}/v1" +def _request_has_desktop_access_token(request: Request) -> bool: + auth_header = request.headers.get("authorization") + if not auth_header: + return False + + parts = auth_header.split(None, 1) + if len(parts) != 2 or parts[0].lower() != "bearer": + return False + + from auth.authentication import is_desktop_access_token + + return is_desktop_access_token(parts[1]) + + def _used_llm_model_aliases(recipe: dict[str, Any]) -> set[str]: """Return the set of model_aliases that are actually referenced by an LLM column. Used to narrow the "Chat model loaded" gate so that orphan @@ -154,6 +168,7 @@ def _inject_local_providers(recipe: dict[str, Any], request: Request) -> None: token = create_access_token( subject = "unsloth", expires_delta = timedelta(hours = 24), + desktop = _request_has_desktop_access_token(request), ) # Defensively strip any stale "external"-only fields the frontend may diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index e17f8f5882..ed331a5660 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -113,7 +113,7 @@ if str(backend_path) not in sys.path: # Import backend functions try: from core.inference import get_inference_backend - from core.inference.llama_cpp import LlamaCppBackend + from core.inference.llama_cpp import LlamaCppBackend, detect_reasoning_flags from utils.models import ModelConfig from utils.inference import load_inference_config from utils.models.model_config import load_model_defaults @@ -122,7 +122,7 @@ except ImportError: if str(parent_backend) not in sys.path: sys.path.insert(0, str(parent_backend)) from core.inference import get_inference_backend - from core.inference.llama_cpp import LlamaCppBackend + from core.inference.llama_cpp import LlamaCppBackend, detect_reasoning_flags from utils.models import ModelConfig from utils.inference import load_inference_config from utils.models.model_config import load_model_defaults @@ -273,7 +273,9 @@ async def load_model( max_context_length = llama_backend.max_context_length, native_context_length = llama_backend.native_context_length, supports_reasoning = llama_backend.supports_reasoning, + reasoning_style = llama_backend.reasoning_style, reasoning_always_on = llama_backend.reasoning_always_on, + supports_preserve_thinking = llama_backend.supports_preserve_thinking, chat_template = llama_backend.chat_template, speculative_type = llama_backend.speculative_type, ) @@ -295,6 +297,21 @@ async def load_model( logger.warning( f"Could not retrieve chat template for {backend.active_model_name}: {e}" ) + # Non-GGUF: only advertise reasoning for gpt-oss Harmony, + # which emits reasoning via channels at the tokenizer level. + # Template-level chat_template_kwargs (enable_thinking / + # preserve_thinking / tools) are not yet forwarded through + # the transformers generation path, so avoid advertising + # controls the server cannot honour outside GGUF. + _sf_supports_reasoning = False + _sf_reasoning_style = "enable_thinking" + if hasattr(backend, "_is_gpt_oss_model"): + try: + if backend._is_gpt_oss_model(): + _sf_supports_reasoning = True + _sf_reasoning_style = "reasoning_effort" + except Exception: + pass return LoadResponse( status = "already_loaded", model = backend.active_model_name, @@ -309,6 +326,11 @@ async def load_model( requires_trust_remote_code = bool( inference_config.get("trust_remote_code", False) ), + supports_reasoning = _sf_supports_reasoning, + reasoning_style = _sf_reasoning_style, + reasoning_always_on = False, + supports_preserve_thinking = False, + supports_tools = False, chat_template = _chat_template, ) @@ -422,7 +444,9 @@ async def load_model( max_context_length = llama_backend.max_context_length, native_context_length = llama_backend.native_context_length, supports_reasoning = llama_backend.supports_reasoning, + reasoning_style = llama_backend.reasoning_style, reasoning_always_on = llama_backend.reasoning_always_on, + supports_preserve_thinking = llama_backend.supports_preserve_thinking, supports_tools = llama_backend.supports_tools, cache_type_kv = llama_backend.cache_type_kv, chat_template = llama_backend.chat_template, @@ -545,6 +569,20 @@ async def load_model( except Exception: pass + # Non-GGUF: gpt-oss Harmony surfaces reasoning via tokenizer-level + # channels; other safetensors reasoning/tools/preserve-thinking + # knobs are not forwarded to tokenizer.apply_chat_template yet, so + # we only advertise support for the Harmony case here. + _sf_supports_reasoning = False + _sf_reasoning_style = "enable_thinking" + if hasattr(backend, "_is_gpt_oss_model"): + try: + if backend._is_gpt_oss_model(): + _sf_supports_reasoning = True + _sf_reasoning_style = "reasoning_effort" + except Exception: + pass + return LoadResponse( status = "loaded", model = config.identifier, @@ -559,6 +597,11 @@ async def load_model( requires_trust_remote_code = bool( inference_config.get("trust_remote_code", False) ), + supports_reasoning = _sf_supports_reasoning, + reasoning_style = _sf_reasoning_style, + reasoning_always_on = False, + supports_preserve_thinking = False, + supports_tools = False, chat_template = _chat_template, ) @@ -766,7 +809,9 @@ async def get_status( (_inference_cfg or {}).get("trust_remote_code", False) ), supports_reasoning = llama_backend.supports_reasoning, + reasoning_style = llama_backend.reasoning_style, reasoning_always_on = llama_backend.reasoning_always_on, + supports_preserve_thinking = llama_backend.supports_preserve_thinking, supports_tools = llama_backend.supports_tools, context_length = llama_backend.context_length, max_context_length = llama_backend.max_context_length, @@ -788,10 +833,18 @@ async def get_status( audio_type = model_info.get("audio_type") has_audio_input = model_info.get("has_audio_input", False) - # gpt-oss safetensors models support reasoning via harmony channels + # Non-GGUF: only gpt-oss Harmony is wired through the transformers + # generation path. Other template-level reasoning / tool kwargs + # are not yet forwarded, so we do not advertise them here. supports_reasoning = False + reasoning_style = "enable_thinking" if backend.active_model_name and hasattr(backend, "_is_gpt_oss_model"): - supports_reasoning = backend._is_gpt_oss_model() + try: + if backend._is_gpt_oss_model(): + supports_reasoning = True + reasoning_style = "reasoning_effort" + except Exception: + pass inference_config = ( load_inference_config(backend.active_model_name) if backend.active_model_name @@ -812,6 +865,10 @@ async def get_status( (inference_config or {}).get("trust_remote_code", False) ), supports_reasoning = supports_reasoning, + reasoning_style = reasoning_style, + reasoning_always_on = False, + supports_preserve_thinking = False, + supports_tools = False, ) except Exception as e: @@ -1393,6 +1450,8 @@ async def openai_chat_completions( presence_penalty = payload.presence_penalty, cancel_event = cancel_event, enable_thinking = payload.enable_thinking, + reasoning_effort = payload.reasoning_effort, + preserve_thinking = payload.preserve_thinking, auto_heal_tool_calls = payload.auto_heal_tool_calls if payload.auto_heal_tool_calls is not None else True, @@ -1562,6 +1621,8 @@ async def openai_chat_completions( presence_penalty = payload.presence_penalty, cancel_event = cancel_event, enable_thinking = payload.enable_thinking, + reasoning_effort = payload.reasoning_effort, + preserve_thinking = payload.preserve_thinking, ) _gguf_sentinel = object() diff --git a/studio/backend/run.py b/studio/backend/run.py index 9675b9ea4c..7590ef1067 100644 --- a/studio/backend/run.py +++ b/studio/backend/run.py @@ -248,6 +248,7 @@ def run_server( port: int = 8888, frontend_path: Path = Path(__file__).resolve().parent.parent / "frontend" / "dist", silent: bool = False, + api_only: bool = False, llama_parallel_slots: int = 1, ): """ @@ -258,6 +259,7 @@ def run_server( port: Port to bind to (auto-increments if in use) frontend_path: Path to frontend build directory (optional) silent: Suppress startup messages + api_only: Run API server only, no frontend serving (for Tauri desktop app) llama_parallel_slots: Number of parallel slots for llama-server Note: @@ -275,6 +277,10 @@ def run_server( except Exception: pass + # Set env var BEFORE importing main so CORS middleware picks it up + if api_only: + os.environ["UNSLOTH_API_ONLY"] = "1" + import nest_asyncio nest_asyncio.apply() @@ -310,8 +316,12 @@ def run_server( print("=" * 50) print("") - # Setup frontend if path provided - if frontend_path: + # Output port for Tauri to parse when in api-only mode + if api_only: + print(f"TAURI_PORT={port}", flush = True) + + # Setup frontend if path provided (skip in api-only mode) + if frontend_path and not api_only: if setup_frontend(app, frontend_path): if not silent: print(f"[OK] Frontend loaded from {frontend_path}") @@ -391,10 +401,17 @@ if __name__ == "__main__": help = "Path to frontend build", ) parser.add_argument("--silent", action = "store_true", help = "Suppress output") + parser.add_argument( + "--api-only", + action = "store_true", + help = "API server only, no frontend (for Tauri)", + ) args = parser.parse_args() - kwargs = dict(host = args.host, port = args.port, silent = args.silent) + kwargs = dict( + host = args.host, port = args.port, silent = args.silent, api_only = args.api_only + ) if args.frontend is not None: kwargs["frontend_path"] = Path(args.frontend) diff --git a/studio/backend/tests/test_desktop_auth.py b/studio/backend/tests/test_desktop_auth.py new file mode 100644 index 0000000000..c8cf1c7081 --- /dev/null +++ b/studio/backend/tests/test_desktop_auth.py @@ -0,0 +1,597 @@ +import importlib.util +import asyncio +import hashlib +import json +import os +import platform +import secrets +import sqlite3 +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace + +import jwt +import pytest +from fastapi import APIRouter, FastAPI +from fastapi.security import HTTPAuthorizationCredentials +from fastapi.testclient import TestClient + +from auth import storage + + +@pytest.fixture(autouse = True) +def isolated_auth_db(tmp_path, monkeypatch): + monkeypatch.setattr(storage, "DB_PATH", tmp_path / "auth.db") + monkeypatch.setattr(storage, "_BOOTSTRAP_PW_PATH", tmp_path / ".bootstrap_password") + monkeypatch.setattr(storage, "_bootstrap_password", None) + monkeypatch.setattr(storage, "_api_key_pbkdf2_salt_cache", None) + yield + + +def seed_user(*, must_change_password = False): + storage.create_initial_user( + username = storage.DEFAULT_ADMIN_USERNAME, + password = "human-password-123", + jwt_secret = secrets.token_urlsafe(64), + must_change_password = must_change_password, + ) + + +def auth_client(): + route_path = Path(__file__).resolve().parents[1] / "routes" / "auth.py" + spec = importlib.util.spec_from_file_location("_desktop_auth_route", route_path) + auth_route = importlib.util.module_from_spec(spec) + assert spec.loader is not None + spec.loader.exec_module(auth_route) + + app = FastAPI() + app.include_router(auth_route.router, prefix = "/api/auth") + return TestClient(app) + + +def data_recipe_jobs_module(): + route_path = ( + Path(__file__).resolve().parents[1] / "routes" / "data_recipe" / "jobs.py" + ) + spec = importlib.util.spec_from_file_location( + "_desktop_data_recipe_jobs", route_path + ) + jobs_route = importlib.util.module_from_spec(spec) + assert spec.loader is not None + spec.loader.exec_module(jobs_route) + return jobs_route + + +def local_recipe(): + return { + "model_providers": [{"name": "local", "is_local": True}], + "model_configs": [{"alias": "local-model", "provider": "local"}], + "columns": [{"column_type": "llm-text", "model_alias": "local-model"}], + } + + +def local_recipe_request(token): + return SimpleNamespace( + headers = {"authorization": f"Bearer {token}"}, + app = SimpleNamespace(state = SimpleNamespace(server_port = 8888)), + scope = {}, + base_url = "http://testserver/", + ) + + +@pytest.fixture +def loaded_local_model(monkeypatch): + inference_module = SimpleNamespace( + get_llama_cpp_backend = lambda: SimpleNamespace(is_loaded = True), + ) + monkeypatch.setitem(sys.modules, "routes.inference", inference_module) + + +def test_desktop_secret_round_trip_uses_real_admin_subject(): + seed_user() + raw = storage.create_desktop_secret() + + assert raw.startswith("desktop-") + assert storage.validate_desktop_secret(raw) == storage.DEFAULT_ADMIN_USERNAME + assert storage.validate_desktop_secret(raw + "x") is None + + +def test_create_desktop_secret_rotates_old_secret(): + seed_user() + old = storage.create_desktop_secret() + new = storage.create_desktop_secret() + + assert old != new + assert storage.validate_desktop_secret(old) is None + assert storage.validate_desktop_secret(new) == storage.DEFAULT_ADMIN_USERNAME + + +def test_clear_desktop_secret_invalidates_secret(): + seed_user() + raw = storage.create_desktop_secret() + + storage.clear_desktop_secret() + + assert storage.validate_desktop_secret(raw) is None + + +def test_ensure_default_admin_does_not_recreate_bootstrap_for_existing_admin(): + seed_user() + + created = storage.ensure_default_admin() + + assert created is False + assert not storage._BOOTSTRAP_PW_PATH.exists() + + +def test_ensure_default_admin_loads_existing_bootstrap_after_restart(monkeypatch): + created = storage.ensure_default_admin() + bootstrap_pw = storage._BOOTSTRAP_PW_PATH.read_text().strip() + + monkeypatch.setattr(storage, "_bootstrap_password", None) + created_again = storage.ensure_default_admin() + + assert created is True + assert storage._BOOTSTRAP_PW_PATH.exists() + assert created_again is False + assert storage.get_bootstrap_password() == bootstrap_pw + + +def test_ensure_default_admin_does_not_generate_for_empty_existing_bootstrap(): + seed_user() + storage._BOOTSTRAP_PW_PATH.write_text(" \n") + + created = storage.ensure_default_admin() + + assert created is False + assert storage._BOOTSTRAP_PW_PATH.read_text() == " \n" + assert storage.get_bootstrap_password() is None + + +def test_web_login_token_has_no_desktop_marker_and_keeps_password_gate(): + seed_user(must_change_password = True) + client = auth_client() + + response = client.post( + "/api/auth/login", + json = { + "username": storage.DEFAULT_ADMIN_USERNAME, + "password": "human-password-123", + }, + ) + + assert response.status_code == 200 + body = response.json() + assert body["must_change_password"] is True + payload = jwt.decode( + body["access_token"], + storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), + algorithms = ["HS256"], + ) + assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME + assert "desktop" not in payload + + gated = client.post( + "/api/auth/api-keys", + headers = {"Authorization": f"Bearer {body['access_token']}"}, + json = {"name": "web"}, + ) + assert gated.status_code == 403 + + +def test_desktop_login_mints_admin_token_without_clearing_web_password_change(): + seed_user(must_change_password = True) + raw = storage.create_desktop_secret() + client = auth_client() + + response = client.post("/api/auth/desktop-login", json = {"secret": raw}) + + assert response.status_code == 200 + body = response.json() + assert body["access_token"] + assert body["refresh_token"] + assert body["token_type"] == "bearer" + assert body["must_change_password"] is False + assert storage.requires_password_change(storage.DEFAULT_ADMIN_USERNAME) is True + + payload = jwt.decode( + body["access_token"], + storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), + algorithms = ["HS256"], + ) + assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME + assert payload["desktop"] is True + + +def test_desktop_refresh_preserves_desktop_marker(): + seed_user(must_change_password = True) + raw = storage.create_desktop_secret() + client = auth_client() + login_body = client.post("/api/auth/desktop-login", json = {"secret": raw}).json() + + response = client.post( + "/api/auth/refresh", + json = {"refresh_token": login_body["refresh_token"]}, + ) + + assert response.status_code == 200 + body = response.json() + assert body["must_change_password"] is False + payload = jwt.decode( + body["access_token"], + storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), + algorithms = ["HS256"], + ) + assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME + assert payload["desktop"] is True + + +def test_desktop_session_uses_real_admin_identity_for_api_keys(): + seed_user(must_change_password = True) + raw = storage.create_desktop_secret() + client = auth_client() + token = client.post("/api/auth/desktop-login", json = {"secret": raw}).json()[ + "access_token" + ] + + response = client.post( + "/api/auth/api-keys", + headers = {"Authorization": f"Bearer {token}"}, + json = {"name": "desktop"}, + ) + + assert response.status_code == 200 + rows = storage.list_api_keys(storage.DEFAULT_ADMIN_USERNAME) + assert [row["name"] for row in rows] == ["desktop"] + + +def test_local_recipe_token_preserves_desktop_marker(loaded_local_model): + from auth.authentication import create_access_token, get_current_subject + + seed_user(must_change_password = True) + jobs_route = data_recipe_jobs_module() + incoming_token = create_access_token( + subject = storage.DEFAULT_ADMIN_USERNAME, + desktop = True, + ) + recipe = local_recipe() + + jobs_route._inject_local_providers(recipe, local_recipe_request(incoming_token)) + + local_token = recipe["model_providers"][0]["api_key"] + payload = jwt.decode( + local_token, + storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), + algorithms = ["HS256"], + ) + assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME + assert payload["desktop"] is True + credentials = HTTPAuthorizationCredentials( + scheme = "Bearer", + credentials = local_token, + ) + assert ( + asyncio.run(get_current_subject(credentials)) == storage.DEFAULT_ADMIN_USERNAME + ) + + +def test_local_recipe_token_keeps_web_marker_absent(loaded_local_model): + from auth.authentication import create_access_token + + seed_user(must_change_password = False) + jobs_route = data_recipe_jobs_module() + incoming_token = create_access_token(subject = storage.DEFAULT_ADMIN_USERNAME) + recipe = local_recipe() + + jobs_route._inject_local_providers(recipe, local_recipe_request(incoming_token)) + + local_token = recipe["model_providers"][0]["api_key"] + payload = jwt.decode( + local_token, + storage.get_jwt_secret(storage.DEFAULT_ADMIN_USERNAME), + algorithms = ["HS256"], + ) + assert payload["sub"] == storage.DEFAULT_ADMIN_USERNAME + assert "desktop" not in payload + + +def test_desktop_login_rejects_invalid_secret(): + seed_user(must_change_password = False) + client = auth_client() + + response = client.post( + "/api/auth/desktop-login", + json = {"secret": "desktop-invalid"}, + ) + + assert response.status_code == 401 + + +def test_write_desktop_secret_file_is_0600_on_unix(tmp_path): + from unsloth_cli.commands import studio as studio_cli + + path = tmp_path / ".desktop_secret" + if platform.system() != "Windows": + path.write_text("old-secret") + os.chmod(path, 0o644) + + studio_cli._write_auth_secret(path, "desktop-secret") + + assert path.read_text() == "desktop-secret" + if platform.system() != "Windows": + assert oct(path.stat().st_mode & 0o777) == "0o600" + + +def test_reset_password_removes_desktop_secret_files(tmp_path, monkeypatch): + from typer.testing import CliRunner + from unsloth_cli.commands import studio as studio_cli + + auth_dir = tmp_path / "auth" + auth_dir.mkdir() + (auth_dir / "auth.db").write_text("db") + (auth_dir / ".bootstrap_password").write_text("boot") + (auth_dir / ".desktop_secret").write_text("new") + monkeypatch.setattr(studio_cli, "STUDIO_HOME", tmp_path) + + result = CliRunner().invoke(studio_cli.studio_app, ["reset-password"]) + + assert result.exit_code == 0 + assert not (auth_dir / "auth.db").exists() + assert not (auth_dir / ".bootstrap_password").exists() + assert not (auth_dir / ".desktop_secret").exists() + + +def test_reset_password_removes_desktop_secret_files_without_db(tmp_path, monkeypatch): + from typer.testing import CliRunner + from unsloth_cli.commands import studio as studio_cli + + auth_dir = tmp_path / "auth" + auth_dir.mkdir() + (auth_dir / ".desktop_secret").write_text("new") + monkeypatch.setattr(studio_cli, "STUDIO_HOME", tmp_path) + + result = CliRunner().invoke(studio_cli.studio_app, ["reset-password"]) + + assert result.exit_code == 0 + assert not (auth_dir / ".desktop_secret").exists() + + +def test_desktop_capabilities_json_reports_rollout_safe_flags(): + from typer.testing import CliRunner + import unsloth_cli.commands.studio as studio_cli + + result = CliRunner().invoke( + studio_cli.studio_app, + ["desktop-capabilities", "--json"], + ) + + assert result.exit_code == 0 + body = json.loads(result.output) + assert body["desktop_protocol_version"] == 1 + assert body["supports_provision_desktop_auth"] is True + assert body["supports_api_only"] is True + assert isinstance(body["version"], str) + + +def test_health_response_reports_desktop_capability_fields(monkeypatch): + router_stub = SimpleNamespace( + auth_router = APIRouter(), + data_recipe_router = APIRouter(), + datasets_router = APIRouter(), + export_router = APIRouter(), + inference_router = APIRouter(), + models_router = APIRouter(), + training_history_router = APIRouter(), + training_router = APIRouter(), + ) + monkeypatch.setitem(sys.modules, "routes", router_stub) + + import studio.backend.main as backend_main + + monkeypatch.setattr(backend_main._hw_module, "CHAT_ONLY", False) + + body = asyncio.run(backend_main.health_check()) + + assert body["desktop_protocol_version"] == 1 + assert body["supports_desktop_auth"] is True + + +def test_provision_desktop_auth_writes_secret_and_creates_db_without_backend_deps( + tmp_path, + monkeypatch, +): + auth_dir = tmp_path / "auth" + auth_dir.mkdir() + + code = """ +import builtins +import sys +from pathlib import Path +from typer.testing import CliRunner + +studio_home = Path(sys.argv[1]) +real_import = builtins.__import__ + +def guarded_import(name, *args, **kwargs): + blocked = ("auth", "fastapi", "structlog", "utils") + if name in blocked or name.startswith(("auth.", "utils.")): + raise ModuleNotFoundError(name) + return real_import(name, *args, **kwargs) + +builtins.__import__ = guarded_import +from unsloth_cli.commands import studio as studio_cli + +studio_cli.STUDIO_HOME = studio_home +result = CliRunner().invoke(studio_cli.studio_app, ["provision-desktop-auth"]) +if result.exit_code != 0: + print(result.output) + if result.exception is not None: + raise result.exception + raise SystemExit(result.exit_code) +""" + result = subprocess.run( + [sys.executable, "-c", code, str(tmp_path)], + cwd = Path(__file__).resolve().parents[3], + env = {**os.environ, "PYTHONPATH": "."}, + text = True, + capture_output = True, + ) + assert result.returncode == 0, result.stderr + result.stdout + secret = (auth_dir / ".desktop_secret").read_text() + assert secret.startswith("desktop-") + + conn = sqlite3.connect(auth_dir / "auth.db") + conn.row_factory = sqlite3.Row + try: + user = conn.execute( + """ + SELECT username, password_salt, password_hash, must_change_password + FROM auth_user + """ + ).fetchone() + app_secrets = { + row["key"]: row["value"] + for row in conn.execute("SELECT key, value FROM app_secrets") + } + refresh_columns = { + row["name"] for row in conn.execute("PRAGMA table_info(refresh_tokens)") + } + finally: + conn.close() + + bootstrap_password = (auth_dir / ".bootstrap_password").read_text().strip() + bootstrap_hash = hashlib.pbkdf2_hmac( + "sha256", + bootstrap_password.encode("utf-8"), + user["password_salt"].encode("utf-8"), + 100_000, + ).hex() + + assert bootstrap_password + assert user["username"] == "unsloth" + assert user["must_change_password"] == 1 + assert bootstrap_hash == user["password_hash"] + assert len(app_secrets["api_key_pbkdf2_salt"]) == 64 + assert len(app_secrets["desktop_secret_hash"]) == 64 + assert app_secrets["desktop_secret_created_at"] + assert "is_desktop" in refresh_columns + + monkeypatch.setattr(storage, "DB_PATH", auth_dir / "auth.db") + monkeypatch.setattr(storage, "_api_key_pbkdf2_salt_cache", None) + assert storage.validate_desktop_secret(secret) == storage.DEFAULT_ADMIN_USERNAME + assert storage.requires_password_change(storage.DEFAULT_ADMIN_USERNAME) is True + + +def test_provision_desktop_auth_keeps_existing_admin_password(tmp_path, monkeypatch): + from typer.testing import CliRunner + from unsloth_cli.commands import studio as studio_cli + + auth_dir = tmp_path / "auth" + auth_dir.mkdir() + monkeypatch.setattr(studio_cli, "STUDIO_HOME", tmp_path) + + conn = sqlite3.connect(auth_dir / "auth.db") + try: + conn.execute( + """ + CREATE TABLE auth_user ( + id INTEGER PRIMARY KEY, + username TEXT UNIQUE NOT NULL, + password_salt TEXT NOT NULL, + password_hash TEXT NOT NULL, + jwt_secret TEXT NOT NULL, + must_change_password INTEGER NOT NULL DEFAULT 0 + ) + """ + ) + conn.execute( + """ + INSERT INTO auth_user ( + username, password_salt, password_hash, jwt_secret, must_change_password + ) + VALUES (?, ?, ?, ?, ?) + """, + ("unsloth", "existing-salt", "existing-hash", "existing-jwt", 0), + ) + conn.commit() + finally: + conn.close() + + result = CliRunner().invoke(studio_cli.studio_app, ["provision-desktop-auth"]) + + assert result.exit_code == 0 + assert not (auth_dir / ".bootstrap_password").exists() + conn = sqlite3.connect(auth_dir / "auth.db") + conn.row_factory = sqlite3.Row + try: + user = conn.execute( + """ + SELECT password_salt, password_hash, jwt_secret, must_change_password + FROM auth_user WHERE username = ? + """, + ("unsloth",), + ).fetchone() + finally: + conn.close() + + assert dict(user) == { + "password_salt": "existing-salt", + "password_hash": "existing-hash", + "jwt_secret": "existing-jwt", + "must_change_password": 0, + } + + +def test_update_password_clears_desktop_secret(): + seed_user() + raw = storage.create_desktop_secret() + assert storage.validate_desktop_secret(raw) == storage.DEFAULT_ADMIN_USERNAME + + changed = storage.update_password( + storage.DEFAULT_ADMIN_USERNAME, "new-admin-password" + ) + assert changed is True + assert storage.validate_desktop_secret(raw) is None + + +def test_update_password_on_unknown_user_leaves_desktop_secret_intact(): + seed_user() + raw = storage.create_desktop_secret() + + changed = storage.update_password("not-a-user", "irrelevant") + assert changed is False + assert storage.validate_desktop_secret(raw) == storage.DEFAULT_ADMIN_USERNAME + + +def test_desktop_auth_provision_has_bounded_timeout(): + rs_path = ( + Path(__file__).resolve().parents[3] + / "studio" + / "src-tauri" + / "src" + / "desktop_auth.rs" + ) + src = rs_path.read_text() + start = src.index("async fn provision_desktop_auth(") + depth = 0 + body_start = src.index("{", start) + body_end = None + for i in range(body_start, len(src)): + c = src[i] + if c == "{": + depth += 1 + elif c == "}": + depth -= 1 + if depth == 0: + body_end = i + 1 + break + assert body_end is not None + body = src[start:body_end] + assert "tokio::time::timeout" in body + import re + + m = re.search(r"Duration::from_secs\(\s*(\d+)\s*\)", body) + assert m is not None + seconds = int(m.group(1)) + assert 5 <= seconds <= 120 diff --git a/studio/backend/tests/test_vision_cache.py b/studio/backend/tests/test_vision_cache.py index fae1e95311..9e7bbdd1fb 100644 --- a/studio/backend/tests/test_vision_cache.py +++ b/studio/backend/tests/test_vision_cache.py @@ -124,23 +124,50 @@ class TestVisionCacheSubprocessPath: class TestVisionCacheOnException: - """When detection raises an exception, _is_vision_model_uncached catches - it and returns False. That False must be cached so subsequent calls don't - retry and fail again.""" + """When detection raises an exception, _is_vision_model_uncached + distinguishes permanent failures (cached as False) from transient + failures (returned as None, not cached so the next call can retry). + Verify both contracts.""" + + @patch( + "utils.models.model_config.load_model_config", + side_effect = ValueError("bad config"), + ) + @patch("utils.transformers_version.needs_transformers_5", return_value = False) + def test_permanent_exception_result_cached(self, mock_needs_t5, mock_load_config): + """A permanent failure (ValueError / RepositoryNotFoundError / + GatedRepoError / JSONDecodeError) should be caught, return False, + and that False should be cached so subsequent calls don't retry. + + ValueError is used here because it's the simplest of the + code-path's cacheable exception types and does not require an + import of huggingface_hub errors (whose module path varies + across versions).""" + # First call: load_model_config raises -> except branch -> False. + assert is_vision_model("broken/model") is False + # Second call: cache hit, load_model_config not called again. + assert is_vision_model("broken/model") is False + mock_load_config.assert_called_once() @patch( "utils.models.model_config.load_model_config", side_effect = OSError("network down"), ) @patch("utils.transformers_version.needs_transformers_5", return_value = False) - def test_exception_result_cached(self, mock_needs_t5, mock_load_config): - """A real exception inside _is_vision_model_uncached should be caught, - return False, and that False should be cached for subsequent calls.""" - # First call: load_model_config raises → except branch → False + def test_transient_exception_not_cached(self, mock_needs_t5, mock_load_config): + """A transient failure (OSError, timeouts) should return None from + _is_vision_model_uncached, surface as False to the caller, and + NOT be cached, so the next call retries detection. This matches + the documented behaviour on _vision_detection_cache: + 'transient failures (network errors, timeouts) are NOT cached so + they can be retried.'""" + # First call: load_model_config raises OSError -> uncached None + # -> caller returns False without caching. assert is_vision_model("broken/model") is False - # Second call: cache hit, load_model_config not called again + # Second call: cache miss again, load_model_config called a + # second time. assert is_vision_model("broken/model") is False - mock_load_config.assert_called_once() + assert mock_load_config.call_count == 2 # --------------------------------------------------------------------------- diff --git a/studio/backend/utils/hardware/nvidia.py b/studio/backend/utils/hardware/nvidia.py index dc5295c302..274d9beb48 100644 --- a/studio/backend/utils/hardware/nvidia.py +++ b/studio/backend/utils/hardware/nvidia.py @@ -6,6 +6,10 @@ from typing import Any, Optional from loggers import get_logger +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + logger = get_logger(__name__) @@ -61,6 +65,7 @@ def get_physical_gpu_count() -> Optional[int]: capture_output = True, text = True, timeout = 5, + **_windows_hidden_subprocess_kwargs(), ) if result.returncode == 0 and result.stdout.strip(): return len(result.stdout.strip().splitlines()) @@ -85,6 +90,7 @@ def get_primary_gpu_utilization() -> dict[str, Any]: capture_output = True, text = True, timeout = 5, + **_windows_hidden_subprocess_kwargs(), ) except (OSError, subprocess.TimeoutExpired) as e: logger.warning("nvidia-smi query failed in get_primary_gpu_utilization: %s", e) @@ -135,6 +141,7 @@ def get_visible_gpu_utilization( capture_output = True, text = True, timeout = 5, + **_windows_hidden_subprocess_kwargs(), ) except (OSError, subprocess.TimeoutExpired) as e: logger.warning("nvidia-smi query failed in get_visible_gpu_utilization: %s", e) @@ -220,6 +227,7 @@ def get_backend_visible_gpu_info( capture_output = True, text = True, timeout = 10, + **_windows_hidden_subprocess_kwargs(), ) except (OSError, subprocess.TimeoutExpired) as e: logger.warning("nvidia-smi query failed in get_backend_visible_gpu_info: %s", e) diff --git a/studio/backend/utils/models/model_config.py b/studio/backend/utils/models/model_config.py index a2d48cf009..a2b0c90e59 100644 --- a/studio/backend/utils/models/model_config.py +++ b/studio/backend/utils/models/model_config.py @@ -32,6 +32,10 @@ import threading import yaml +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + logger = get_logger(__name__) # ── Model size extraction ──────────────────────────────────── @@ -579,6 +583,7 @@ def _is_vision_model_subprocess( capture_output = True, text = True, timeout = 60, + **_windows_hidden_subprocess_kwargs(), ) if result.returncode != 0: diff --git a/studio/backend/utils/subprocess_compat.py b/studio/backend/utils/subprocess_compat.py new file mode 100644 index 0000000000..bedf8cf2e6 --- /dev/null +++ b/studio/backend/utils/subprocess_compat.py @@ -0,0 +1,34 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Cross-platform subprocess helpers for the Unsloth Studio backend.""" + +import subprocess +import sys + + +def windows_hidden_subprocess_kwargs() -> dict[str, object]: + """Return Windows-only subprocess kwargs that suppress console windows. + + On non-Windows platforms returns an empty dict so callers can always + unpack the result into ``subprocess.run`` / ``subprocess.Popen`` via + ``**windows_hidden_subprocess_kwargs()``. + """ + if sys.platform != "win32": + return {} + + kwargs: dict[str, object] = {} + create_no_window = getattr(subprocess, "CREATE_NO_WINDOW", 0) + if create_no_window: + kwargs["creationflags"] = create_no_window + + startupinfo_factory = getattr(subprocess, "STARTUPINFO", None) + startf_use_showwindow = getattr(subprocess, "STARTF_USESHOWWINDOW", 0) + sw_hide = getattr(subprocess, "SW_HIDE", 0) + if startupinfo_factory is not None and startf_use_showwindow: + startupinfo = startupinfo_factory() + startupinfo.dwFlags |= startf_use_showwindow + startupinfo.wShowWindow = sw_hide + kwargs["startupinfo"] = startupinfo + + return kwargs diff --git a/studio/backend/utils/transformers_version.py b/studio/backend/utils/transformers_version.py index 36c3a4c22d..f36bdcd6e8 100644 --- a/studio/backend/utils/transformers_version.py +++ b/studio/backend/utils/transformers_version.py @@ -36,6 +36,10 @@ import subprocess import sys from pathlib import Path +from utils.subprocess_compat import ( + windows_hidden_subprocess_kwargs as _windows_hidden_subprocess_kwargs, +) + logger = get_logger(__name__) @@ -499,6 +503,7 @@ def _install_to_dir(pkg: str, target_dir: str) -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, + **_windows_hidden_subprocess_kwargs(), ) if result.returncode == 0: return True @@ -520,6 +525,7 @@ def _install_to_dir(pkg: str, target_dir: str) -> bool: stdout = subprocess.PIPE, stderr = subprocess.STDOUT, text = True, + **_windows_hidden_subprocess_kwargs(), ) if result.returncode != 0: logger.error("install failed:\n%s", result.stdout) diff --git a/studio/frontend/package.json b/studio/frontend/package.json index a2eebd5cb5..c5cb949ccd 100644 --- a/studio/frontend/package.json +++ b/studio/frontend/package.json @@ -41,6 +41,10 @@ "@tailwindcss/vite": "^4.2.2", "@tanstack/react-router": "^1.159.10", "@tanstack/react-table": "^8.21.3", + "@tauri-apps/api": "^2.10.1", + "@tauri-apps/plugin-opener": "^2.5.3", + "@tauri-apps/plugin-process": "^2.3.1", + "@tauri-apps/plugin-updater": "^2.10.1", "@toolwind/corner-shape": "^0.0.8-3", "@types/canvas-confetti": "^1.9.0", "@xyflow/react": "^12.10.0", diff --git a/studio/frontend/public/studio.png b/studio/frontend/public/studio.png new file mode 100644 index 0000000000..4e531499b7 Binary files /dev/null and b/studio/frontend/public/studio.png differ diff --git a/studio/frontend/src/app/auth-guards.ts b/studio/frontend/src/app/auth-guards.ts index 1dcdfcb143..509b8f61af 100644 --- a/studio/frontend/src/app/auth-guards.ts +++ b/studio/frontend/src/app/auth-guards.ts @@ -2,12 +2,14 @@ // Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 import { redirect } from "@tanstack/react-router"; +import { apiUrl, isTauri } from "@/lib/api-base"; import { getPostAuthRoute, hasAuthToken, hasRefreshToken, mustChangePassword, refreshSession, + tauriAutoAuth, } from "@/features/auth"; async function hasActiveSession(): Promise { @@ -16,55 +18,64 @@ async function hasActiveSession(): Promise { return refreshSession(); } -async function checkAuthInitialized(): Promise { +interface AuthStatus { + initialized: boolean; + requires_password_change: boolean; +} + +async function fetchAuthStatus(): Promise { try { - const res = await fetch("/api/auth/status"); - if (!res.ok) return true; // fallback to login on error - const data = (await res.json()) as { initialized: boolean }; - return data.initialized; + const res = await fetch(apiUrl("/api/auth/status")); + if (!res.ok) return { initialized: true, requires_password_change: mustChangePassword() }; + return (await res.json()) as AuthStatus; } catch { - return true; // fallback to login on error + return { initialized: true, requires_password_change: mustChangePassword() }; } } -async function checkPasswordChangeRequired(): Promise { - try { - const res = await fetch("/api/auth/status"); - if (!res.ok) return mustChangePassword(); - const data = (await res.json()) as { requires_password_change: boolean }; - return data.requires_password_change || mustChangePassword(); - } catch { - return mustChangePassword(); - } +function authRedirect(to: "/login" | "/change-password"): never { + throw redirect({ to }); } export async function requireAuth(): Promise { + if (isTauri) { + await tauriAutoAuth(); + return; + } + if (await hasActiveSession()) { - if (await checkPasswordChangeRequired()) { - throw redirect({ to: "/change-password" }); + const { requires_password_change } = await fetchAuthStatus(); + if (requires_password_change || mustChangePassword()) { + authRedirect("/change-password"); } return; } - const requiresPasswordChange = await checkPasswordChangeRequired(); - if (requiresPasswordChange) throw redirect({ to: "/change-password" }); - const initialized = await checkAuthInitialized(); - throw redirect({ to: initialized ? "/login" : "/change-password" }); + const status = await fetchAuthStatus(); + if (status.requires_password_change || mustChangePassword()) { + authRedirect("/change-password"); + } + authRedirect(status.initialized ? "/login" : "/change-password"); } export async function requireGuest(): Promise { + if (isTauri) { + await tauriAutoAuth(); + throw redirect({ to: "/chat" }); + } if (!(await hasActiveSession())) return; throw redirect({ to: getPostAuthRoute() }); } export async function requirePasswordChangeFlow(): Promise { - const requiresPasswordChange = await checkPasswordChangeRequired(); - - if (requiresPasswordChange) return; + if (isTauri) { + await tauriAutoAuth(); + throw redirect({ to: "/chat" }); + } + const status = await fetchAuthStatus(); + if (status.requires_password_change || mustChangePassword()) return; if (await hasActiveSession()) { throw redirect({ to: getPostAuthRoute() }); } - - const initialized = await checkAuthInitialized(); - throw redirect({ to: initialized ? "/login" : "/change-password" }); + authRedirect(status.initialized ? "/login" : "/change-password"); } diff --git a/studio/frontend/src/app/provider.tsx b/studio/frontend/src/app/provider.tsx index 68ce3061bd..b75998a169 100644 --- a/studio/frontend/src/app/provider.tsx +++ b/studio/frontend/src/app/provider.tsx @@ -1,18 +1,181 @@ // 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 { StartupScreen } from "@/components/tauri/startup-screen"; +import { UpdateBanner } from "@/components/tauri/update-banner"; +import { UpdateScreen } from "@/components/tauri/update-screen"; import { Toaster } from "@/components/ui/sonner"; +import { useTauriBackend } from "@/hooks/use-tauri-backend"; +import { useTauriUpdate } from "@/hooks/use-tauri-update"; +import { isTauri } from "@/lib/api-base"; import { ThemeProvider } from "next-themes"; -import type { ReactNode } from "react"; +import { useEffect, useRef, type ReactNode } from "react"; interface AppProviderProps { children: ReactNode; } +// --------------------------------------------------------------------------- +// Tauri window helpers (only imported in Tauri mode) +// --------------------------------------------------------------------------- + +async function showWindow(): Promise { + const { getCurrentWindow } = await import("@tauri-apps/api/window"); + await getCurrentWindow().show(); +} + +function easeOutQuart(t: number): number { + return 1 - (1 - t) ** 4; +} + +async function animateToGoldenRatio(abortRef: { current: boolean }): Promise { + const { getCurrentWindow, currentMonitor, LogicalSize } = await import("@tauri-apps/api/window"); + const win = getCurrentWindow(); + + // Ensure window is visible before resizing + await win.show(); + + const monitor = await currentMonitor(); + if (!monitor) return; + + // Convert physical pixels to logical using scale factor + const scale = monitor.scaleFactor; + const screenW = monitor.size.width / scale; + const screenH = monitor.size.height / scale; + + // Target: 75% of screen width, golden ratio height, capped at min 900x600 + const targetW = Math.max(900, Math.round(screenW * 0.75)); + const targetH = Math.max(600, Math.round(targetW / 1.618)); + // Don't exceed screen height + const finalH = Math.min(targetH, Math.round(screenH * 0.85)); + const finalW = targetW; + + // Check reduced motion preference + const prefersReducedMotion = window.matchMedia("(prefers-reduced-motion: reduce)").matches; + + if (prefersReducedMotion) { + await win.setSize(new LogicalSize(finalW, finalH)); + } else { + // Read current size instead of hardcoding — stays correct if tauri.conf.json changes + const inner = await win.innerSize(); + const factor = await win.scaleFactor(); + const startW = Math.round(inner.width / factor); + const startH = Math.round(inner.height / factor); + const steps = 15; + const stepDuration = 23; // ~350ms total + + for (let i = 1; i <= steps; i++) { + if (abortRef.current) return; + const t = easeOutQuart(i / steps); + const w = Math.round(startW + (finalW - startW) * t); + const h = Math.round(startH + (finalH - startH) * t); + await win.setSize(new LogicalSize(w, h)); + await new Promise((r) => setTimeout(r, stepDuration)); + } + } + + if (abortRef.current) return; + + // Apply constraints and finalize + await win.setResizable(true); + await win.setSizeConstraints({ minWidth: 900, minHeight: 600 }); + await win.center(); +} + +// --------------------------------------------------------------------------- +// TauriWrapper +// --------------------------------------------------------------------------- + +function TauriUpdateLayer({ isExternalServer }: { isExternalServer: boolean }) { + const update = useTauriUpdate(isExternalServer); + const isUpdating = + update.status === "updating-backend" || + update.status === "downloading" || + update.status === "installing" || + (update.status === "error" && !update.dismissed); + + if (isUpdating) { + return ( + + ); + } + + return ( + + ); +} + +function TauriWrapper({ children }: { children: ReactNode }) { + const { + status, logs, error, isExternalServer, + currentStepIndex, progressDetail, elevationPackages, + startInstall, retry, retryInstall, approveElevation, + } = useTauriBackend(); + + const hasResized = useRef(false); + const abortRef = useRef(false); + + // Show the window once the frontend mounts (for pre-running states) + useEffect(() => { + if (isTauri) void showWindow(); + }, []); + + // Animate resize when backend becomes ready + useEffect(() => { + if (status === "running" && !hasResized.current) { + hasResized.current = true; + abortRef.current = false; + animateToGoldenRatio(abortRef).catch(async () => { + // On failure, at minimum make the window resizable so user can fix manually + try { + const { getCurrentWindow } = await import("@tauri-apps/api/window"); + await getCurrentWindow().setResizable(true); + } catch { /* swallow — window may still be functional */ } + }); + } + return () => { abortRef.current = true; }; + }, [status]); + + if (!isTauri) return <>{children}; + if (status === "running") return <>{children}; + + return ( + + ); +} + export function AppProvider({ children }: AppProviderProps) { return ( - {children} + + {children} + ); diff --git a/studio/frontend/src/app/routes/__root.tsx b/studio/frontend/src/app/routes/__root.tsx index 69aeb11748..19b5763557 100644 --- a/studio/frontend/src/app/routes/__root.tsx +++ b/studio/frontend/src/app/routes/__root.tsx @@ -3,8 +3,8 @@ import { AppSidebar } from "@/components/app-sidebar"; import { Navbar } from "@/components/navbar"; +import { fetchDeviceType, usePlatformStore } from "@/config/env"; import { SidebarInset, SidebarProvider } from "@/components/ui/sidebar"; -import { usePlatformStore } from "@/config/env"; import { SettingsDialog, useSettingsDialogStore } from "@/features/settings"; import { useTrainingUnloadGuard } from "@/features/training/hooks/use-training-unload-guard"; import { useSidebarPin } from "@/hooks/use-sidebar-pin"; @@ -33,7 +33,10 @@ function isChatOnlyAllowed(pathname: string): boolean { } export const Route = createRootRoute({ - beforeLoad: ({ location }) => { + beforeLoad: async ({ location }) => { + // Ensure platform info is fetched before checking chat-only guard. + // fetchDeviceType caches after first call, so subsequent navigations are instant. + await fetchDeviceType(); const chatOnly = usePlatformStore.getState().isChatOnly(); if (chatOnly && !isChatOnlyAllowed(location.pathname)) { throw redirect({ to: "/chat" }); diff --git a/studio/frontend/src/components/assistant-ui/markdown-text.tsx b/studio/frontend/src/components/assistant-ui/markdown-text.tsx index 2a0517a44a..7eb4b21ba7 100644 --- a/studio/frontend/src/components/assistant-ui/markdown-text.tsx +++ b/studio/frontend/src/components/assistant-ui/markdown-text.tsx @@ -5,6 +5,7 @@ import { copyToClipboard } from "@/lib/copy-to-clipboard"; import { preprocessLaTeX } from "@/lib/latex"; +import { openLink } from "@/lib/open-link"; import { INTERNAL, useMessagePartText } from "@assistant-ui/react"; import { Copy02Icon, Tick02Icon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; @@ -32,9 +33,13 @@ const STREAMDOWN_COMPONENTS = { }: React.ComponentProps<"a">) => ( { + if (href && openLink(href)) { + e.preventDefault(); + } + }} {...props} > {children} diff --git a/studio/frontend/src/components/assistant-ui/sources.tsx b/studio/frontend/src/components/assistant-ui/sources.tsx index da97ff66b5..81c8b0c213 100644 --- a/studio/frontend/src/components/assistant-ui/sources.tsx +++ b/studio/frontend/src/components/assistant-ui/sources.tsx @@ -1,5 +1,6 @@ "use client"; +import { openLink } from "@/lib/open-link"; import { memo, useState, @@ -93,8 +94,8 @@ function Source({ variant, size, asChild = false, - target = "_blank", - rel = "noopener noreferrer", + href, + onClick, ...props }: SourceProps) { return ( @@ -109,8 +110,14 @@ function Source({ > { + if (href && openLink(href)) { + e.preventDefault(); + } + onClick?.(e); + }} {...(props as ComponentProps<"a">)} /> diff --git a/studio/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx index 0b97d98dfd..3f67b11b51 100644 --- a/studio/frontend/src/components/assistant-ui/thread.tsx +++ b/studio/frontend/src/components/assistant-ui/thread.tsx @@ -24,8 +24,15 @@ import { useScrollThreadToBottom, } from "@/components/assistant-ui/use-intent-aware-autoscroll"; import { Button } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; import { sentAudioNames } from "@/features/chat/api/chat-adapter"; import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store"; +import { applyQwenThinkingParams } from "@/features/chat/utils/qwen-params"; import { deleteThreadMessage } from "@/features/chat/utils/delete-thread-message"; import { AUDIO_ACCEPT, MAX_AUDIO_SIZE, fileToBase64 } from "@/lib/audio-utils"; import { copyToClipboard } from "@/lib/copy-to-clipboard"; @@ -373,18 +380,6 @@ const ComposerAudioUpload: FC = () => { ); }; -/** Qwen3/3.5 recommended params differ between thinking on/off. */ -function applyQwenThinkingParams(thinkingOn: boolean): void { - const store = useChatRuntimeStore.getState(); - const checkpoint = store.params.checkpoint?.toLowerCase() ?? ""; - if (!checkpoint.includes("qwen3")) { - return; - } - const params = thinkingOn - ? { temperature: 0.6, topP: 0.95, topK: 20, minP: 0.0 } - : { temperature: 0.7, topP: 0.8, topK: 20, minP: 0.0 }; - store.setParams({ ...store.params, ...params }); -} const ReasoningToggle: FC = () => { const modelLoaded = useChatRuntimeStore( @@ -393,8 +388,49 @@ const ReasoningToggle: FC = () => { const supportsReasoning = useChatRuntimeStore((s) => s.supportsReasoning); const reasoningEnabled = useChatRuntimeStore((s) => s.reasoningEnabled); const setReasoningEnabled = useChatRuntimeStore((s) => s.setReasoningEnabled); + const reasoningStyle = useChatRuntimeStore((s) => s.reasoningStyle); + const reasoningEffort = useChatRuntimeStore((s) => s.reasoningEffort); + const setReasoningEffort = useChatRuntimeStore((s) => s.setReasoningEffort); const disabled = !(modelLoaded && supportsReasoning); + if (reasoningStyle === "reasoning_effort") { + return ( + + + + + + {(["low", "medium", "high"] as const).map((level) => ( + setReasoningEffort(level)} + > + {level.charAt(0).toUpperCase() + level.slice(1)} + {reasoningEffort === level ? " \u2713" : ""} + + ))} + + + ); + } + return ( + ); +}; + const WebSearchToggle: FC = () => { const modelLoaded = useChatRuntimeStore( (s) => !!s.params.checkpoint && !s.modelLoading, @@ -551,6 +625,7 @@ const ComposerAction: FC<{ disabled?: boolean }> = ({ disabled }) => { + diff --git a/studio/frontend/src/components/tauri/startup-screen.tsx b/studio/frontend/src/components/tauri/startup-screen.tsx new file mode 100644 index 0000000000..af5f241416 --- /dev/null +++ b/studio/frontend/src/components/tauri/startup-screen.tsx @@ -0,0 +1,401 @@ +// 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 { ShimmerButton } from "@/components/ui/shimmer-button"; +import type { BackendStatus } from "@/hooks/use-tauri-backend"; +import { AnimatePresence, motion } from "motion/react"; + +interface StartupScreenProps { + status: BackendStatus; + logs: string[]; + error: string | null; + currentStepIndex: number; + progressDetail: string | null; + elevationPackages: string[]; + onInstall: () => void; + onRetry: () => void; + onRetryInstall: () => void; + onApproveElevation: () => void; + onStartServer: () => void; +} + +// --------------------------------------------------------------------------- +// Constants +// --------------------------------------------------------------------------- + +const INSTALL_STEPS = [ + "Detecting your system", + "Checking dependencies", + "Setting up package manager", + "Creating Python environment", + "Installing ML framework", + "Installing Unsloth", + "Finalizing setup", +] as const; + +const EASE_OUT_QUART: [number, number, number, number] = [0.165, 0.84, 0.44, 1]; + +// --------------------------------------------------------------------------- +// Sub-components +// --------------------------------------------------------------------------- + +function TealSpinner({ size = 24 }: { size?: number }) { + return ( + + ); +} + +function Logo() { + return ( +
+ Unsloth + Unsloth Studio +
+ ); +} + +function ActionButton({ + onClick, + variant = "primary", + children, +}: { + onClick: () => void; + variant?: "primary" | "secondary"; + children: React.ReactNode; +}) { + const base = "rounded-lg px-5 py-2.5 text-sm font-medium cursor-pointer transition-colors"; + const styles = + variant === "primary" + ? `${base} bg-primary text-primary-foreground hover:bg-primary/80` + : `${base} bg-muted text-foreground hover:bg-muted/80`; + return ( + + ); +} + +// --------------------------------------------------------------------------- +// Per-status renderers +// --------------------------------------------------------------------------- + +function CheckingContent() { + return ( +
+
+ +
+
+ +

Checking...

+
+
+ ); +} + +function NotInstalledContent({ onInstall }: { onInstall: () => void }) { + return ( +
+
+ +

+ To install Unsloth, click Get Started. +

+
+
+ + Get Started + +
+
+ ); +} + +function InstallingContent({ + currentStepIndex, + progressDetail, +}: { + currentStepIndex: number; + progressDetail: string | null; +}) { + const stepNum = Math.max(0, currentStepIndex) + 1; + const stepLabel = INSTALL_STEPS[Math.min(currentStepIndex, INSTALL_STEPS.length - 1)]; + + return ( +
+
+ +
+
+ +

Installing...

+

+ Please wait a few mins, then you can start training. +

+ {currentStepIndex >= 0 && ( +

+ Step {stepNum} of {INSTALL_STEPS.length}: {stepLabel} +

+ )} + {progressDetail && ( +

{progressDetail}

+ )} +
+
+ ); +} + +function RepairingContent({ + logs, + progressDetail, +}: { + logs: string[]; + progressDetail: string | null; +}) { + const latest = progressDetail ?? logs.at(-1); + + return ( +
+
+ +
+
+ +

Updating existing Studio install...

+ {latest && ( +

{latest}

+ )} +
+
+ ); +} + +function InstallErrorContent({ + error, + logs, + onRetryInstall, +}: { + error: string | null; + logs: string[]; + onRetryInstall: () => void; +}) { + return ( + <> + +
+

Setup ran into a problem

+ {error && ( +

{error}

+ )} +
+ void navigator.clipboard.writeText(logs.join("\n"))} + > + Copy Logs + + Try Again +
+
+ + ); +} + +function RepairErrorContent({ + error, + logs, + onRetry, +}: { + error: string | null; + logs: string[]; + onRetry: () => void; +}) { + return ( + <> + +
+

Update failed

+ {error && ( +

{error}

+ )} +
+ void navigator.clipboard.writeText(logs.join("\n"))} + > + Copy Logs + + Retry +
+
+ + ); +} + +function NeedsElevationContent({ + elevationPackages, + onApproveElevation, + onRetryInstall, +}: { + elevationPackages: string[]; + onApproveElevation: () => void; + onRetryInstall: () => void; +}) { + return ( + <> + +
+

Permission needed

+

+ The following system packages need to be installed: +

+
+ {elevationPackages.map((pkg) => ( +
{pkg}
+ ))} +
+
+ Cancel + Allow +
+
+ + ); +} + +function StartingContent() { + return ( +
+
+ +
+
+ +

Starting server...

+
+
+ ); +} + +function StoppedContent({ onStartServer }: { onStartServer: () => void }) { + return ( + <> + +
+

Server stopped

+
+ Start Server +
+
+ + ); +} + +function ErrorContent({ + error, + logs, + onRetry, +}: { + error: string | null; + logs: string[]; + onRetry: () => void; +}) { + return ( + <> + +
+

Something went wrong

+ {error && ( +

{error}

+ )} +
+ void navigator.clipboard.writeText(logs.join("\n"))} + > + Copy Logs + + Retry +
+
+ + ); +} + +// --------------------------------------------------------------------------- +// Main component +// --------------------------------------------------------------------------- + +export function StartupScreen({ + status, + logs, + error, + currentStepIndex, + progressDetail, + elevationPackages, + onInstall, + onRetry, + onRetryInstall, + onApproveElevation, + onStartServer, +}: StartupScreenProps) { + function renderContent() { + switch (status) { + case "checking": + return ; + case "not-installed": + return ; + case "installing": + return ; + case "install-error": + return ; + case "repairing": + return ; + case "repair-error": + return ; + case "needs-elevation": + return ( + + ); + case "starting": + return ; + case "running": + return null; + case "stopped": + return ; + case "error": + return ; + } + } + + return ( +
+
+ + + {renderContent()} + + +
+
+ ); +} diff --git a/studio/frontend/src/components/tauri/update-banner.tsx b/studio/frontend/src/components/tauri/update-banner.tsx new file mode 100644 index 0000000000..98652c4db1 --- /dev/null +++ b/studio/frontend/src/components/tauri/update-banner.tsx @@ -0,0 +1,84 @@ +// 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 { Button } from "@/components/ui/button"; +import type { UpdateInfo, UpdateStatus } from "@/hooks/use-tauri-update"; +import { AnimatePresence, motion } from "motion/react"; + +interface UpdateBannerProps { + status: UpdateStatus; + info: UpdateInfo | null; + dismissed: boolean; + isExternalServer?: boolean; + onInstall: () => void; + onDismiss: () => void; +} + +const EASE_OUT_QUART: [number, number, number, number] = [0.165, 0.84, 0.44, 1]; + +export function UpdateBanner({ + status, + info, + dismissed, + isExternalServer = false, + onInstall, + onDismiss, +}: UpdateBannerProps) { + const visible = status === "available"; + const show = visible && !dismissed; + + return ( + + {show && info && ( + +
+ {/* Close button */} + + + {/* Header */} +
+ 🦥 +
+

+ New version: v{info.version} +

+

+ {isExternalServer + ? "Run `unsloth studio update` from your terminal" + : "A new app update is available"} +

+
+
+ + {/* Actions */} +
+ + + +
+
+
+ )} +
+ ); +} diff --git a/studio/frontend/src/components/tauri/update-screen.tsx b/studio/frontend/src/components/tauri/update-screen.tsx new file mode 100644 index 0000000000..a8b77ab1fa --- /dev/null +++ b/studio/frontend/src/components/tauri/update-screen.tsx @@ -0,0 +1,173 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +import type { UpdateStatus } from "@/hooks/use-tauri-update"; +import { AnimatePresence, motion } from "motion/react"; +import { useEffect, useRef } from "react"; + +interface UpdateScreenProps { + status: UpdateStatus; + logs: string[]; + progress: number; + error: string | null; + onRetry: () => void; + onSkipRestart: () => void; +} + +const EASE_OUT_QUART: [number, number, number, number] = [0.165, 0.84, 0.44, 1]; + +function Spinner({ size = 24 }: { size?: number }) { + return ( + + ); +} + +function Logo() { + return ( +
+ Unsloth + Unsloth Studio +
+ ); +} + +function statusLabel(status: UpdateStatus): string { + switch (status) { + case "updating-backend": + return "Updating backend..."; + case "downloading": + return "Downloading app update..."; + case "installing": + return "Installing update..."; + case "error": + return "Update failed"; + default: + return "Updating..."; + } +} + +function statusSubtext(status: UpdateStatus, progress: number): string { + switch (status) { + case "updating-backend": + return "This may take a few minutes. Do not close the app."; + case "downloading": + return `${progress}% downloaded`; + case "installing": + return "The app will restart shortly."; + case "error": + return "Something went wrong during the update."; + default: + return ""; + } +} + +function LogViewer({ logs }: { logs: string[] }) { + const scrollRef = useRef(null); + + useEffect(() => { + if (scrollRef.current) { + scrollRef.current.scrollTop = scrollRef.current.scrollHeight; + } + }, [logs]); + + if (logs.length === 0) return null; + + return ( +
+ {logs.map((line, i) => ( +
+ {line} +
+ ))} +
+ ); +} + +export function UpdateScreen({ + status, + logs, + progress, + error, + onRetry, + onSkipRestart, +}: UpdateScreenProps) { + const isError = status === "error"; + + return ( +
+ + + +
+ {!isError && } +

+ {statusLabel(status)} +

+

+ {statusSubtext(status, progress)} +

+
+ + {/* Download progress bar */} + {status === "downloading" && ( +
+ +
+ )} + + {/* Error display */} + + {isError && error && ( + +

{error}

+
+ )} +
+ + {/* Error actions */} + {isError && ( +
+ + +
+ )} + + {/* Log viewer */} + +
+
+ ); +} diff --git a/studio/frontend/src/components/ui/shimmer-button.tsx b/studio/frontend/src/components/ui/shimmer-button.tsx new file mode 100644 index 0000000000..d675cc0979 --- /dev/null +++ b/studio/frontend/src/components/ui/shimmer-button.tsx @@ -0,0 +1,96 @@ +import React, { type ComponentPropsWithoutRef, type CSSProperties } from "react" + +import { cn } from "@/lib/utils" + +export interface ShimmerButtonProps extends ComponentPropsWithoutRef<"button"> { + shimmerColor?: string + shimmerSize?: string + borderRadius?: string + shimmerDuration?: string + background?: string + className?: string + children?: React.ReactNode +} + +export const ShimmerButton = React.forwardRef< + HTMLButtonElement, + ShimmerButtonProps +>( + ( + { + shimmerColor = "#ffffff", + shimmerSize = "0.05em", + shimmerDuration = "3s", + borderRadius = "100px", + background = "rgba(0, 0, 0, 1)", + className, + children, + ...props + }, + ref + ) => { + return ( + + ) + } +) + +ShimmerButton.displayName = "ShimmerButton" diff --git a/studio/frontend/src/config/env.ts b/studio/frontend/src/config/env.ts index 91e17f6bb9..72bb3fa815 100644 --- a/studio/frontend/src/config/env.ts +++ b/studio/frontend/src/config/env.ts @@ -1,6 +1,7 @@ // 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 { apiUrl } from "@/lib/api-base"; import { create } from "zustand"; export const env = { @@ -21,9 +22,21 @@ interface PlatformState { isChatOnly: () => boolean; } +// Client-side platform detection as fallback when backend isn't ready yet. +function detectLocalPlatform(): DeviceType { + if (typeof navigator === "undefined") return "linux"; + const platform = navigator.platform.toLowerCase(); + const ua = navigator.userAgent.toLowerCase(); + if (platform.includes("mac") || ua.includes("mac")) return "mac"; + if (platform.includes("win") || ua.includes("win")) return "windows"; + return "linux"; +} + +const localDeviceType = detectLocalPlatform(); + export const usePlatformStore = create()((_, get) => ({ - deviceType: "linux", - chatOnly: false, + deviceType: localDeviceType, + chatOnly: localDeviceType === "mac", fetched: false, isChatOnly: () => get().chatOnly, })); @@ -33,16 +46,22 @@ export async function fetchDeviceType(): Promise { if (fetched) return usePlatformStore.getState().deviceType; try { - const res = await fetch("/api/health"); + const res = await fetch(apiUrl("/api/health")); if (res.ok) { const data = (await res.json()) as { device_type?: string; chat_only?: boolean }; - const deviceType = data.device_type ?? "linux"; + const deviceType = data.device_type ?? detectLocalPlatform(); const chatOnly = data.chat_only ?? deviceType === "mac"; usePlatformStore.setState({ deviceType, chatOnly, fetched: true }); return deviceType; } - } catch (err) { - console.warn("[platform] Failed to fetch device type, will retry", err); + } catch { + // Backend not ready — use client-side detection so chat-only guard + // still works on initial load (important for macOS). Keep fetched=false + // so a later call retries against the backend. + const deviceType = detectLocalPlatform(); + const chatOnly = deviceType === "mac"; + usePlatformStore.setState({ deviceType, chatOnly, fetched: false }); + return deviceType; } return usePlatformStore.getState().deviceType; diff --git a/studio/frontend/src/features/auth/api.ts b/studio/frontend/src/features/auth/api.ts index 3bb2c9139c..97d815950d 100644 --- a/studio/frontend/src/features/auth/api.ts +++ b/studio/frontend/src/features/auth/api.ts @@ -1,6 +1,7 @@ // 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 { apiUrl, isTauri } from "@/lib/api-base"; import { clearAuthTokens, getAuthToken, @@ -34,7 +35,7 @@ async function redirectToAuth(): Promise { let target = "/login"; try { - const res = await fetch("/api/auth/status"); + const res = await fetch(apiUrl("/api/auth/status")); if (res.ok) { const data = (await res.json()) as { requires_password_change: boolean }; if (data.requires_password_change || mustChangePassword()) target = "/change-password"; @@ -46,12 +47,32 @@ async function redirectToAuth(): Promise { window.location.href = target; } +async function retryWithCurrentToken( + input: RequestInfo | URL, + init?: RequestInit, +): Promise { + const retryHeaders = new Headers(init?.headers); + const token = getAuthToken(); + if (token) retryHeaders.set("Authorization", `Bearer ${token}`); + return fetch(input, { ...init, headers: retryHeaders }); +} + +async function retryWithTauriAutoAuth( + input: RequestInfo | URL, + init?: RequestInit, +): Promise { + clearAuthTokens(); + const { tauriAutoAuth } = await import("./tauri-auto-auth"); + if (await tauriAutoAuth()) return retryWithCurrentToken(input, init); + return null; +} + export async function refreshSession(): Promise { const refreshToken = getRefreshToken(); if (!refreshToken) return false; try { - const response = await fetch("/api/auth/refresh", { + const response = await fetch(apiUrl("/api/auth/refresh"), { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ refresh_token: refreshToken }), @@ -78,6 +99,7 @@ export async function authFetch( input: RequestInfo | URL, init?: RequestInit, ): Promise { + const resolvedInput = typeof input === 'string' ? apiUrl(input) : input; const headers = new Headers(init?.headers); const accessToken = getAuthToken(); if (accessToken) { @@ -86,14 +108,18 @@ export async function authFetch( let response: Response; try { - response = await fetch(input, { ...init, headers }); + response = await fetch(resolvedInput, { ...init, headers }); } catch (err) { if (err instanceof TypeError) { throw new Error("Studio isn't running -- please relaunch it."); } throw err; } + if (await isPasswordChangeRequiredResponse(response)) { + if (isTauri) { + return (await retryWithTauriAutoAuth(resolvedInput, init)) ?? response; + } void redirectToAuth(); return response; } @@ -101,25 +127,24 @@ export async function authFetch( const refreshed = await refreshSession(); if (!refreshed) { + if (isTauri) { + return (await retryWithTauriAutoAuth(resolvedInput, init)) ?? response; + } clearAuthTokens(); void redirectToAuth(); return response; } if (mustChangePassword()) { + if (isTauri) { + return (await retryWithTauriAutoAuth(resolvedInput, init)) ?? response; + } void redirectToAuth(); return response; } - const retryHeaders = new Headers(init?.headers); - const newToken = getAuthToken(); - if (newToken) { - retryHeaders.set("Authorization", `Bearer ${newToken}`); - } else { - clearAuthTokens(); - } - - return fetch(input, { ...init, headers: retryHeaders }); + if (!getAuthToken()) clearAuthTokens(); + return retryWithCurrentToken(resolvedInput, init); } export function logout(): void { diff --git a/studio/frontend/src/features/auth/components/auth-form.tsx b/studio/frontend/src/features/auth/components/auth-form.tsx index d9190429bd..090a3081a4 100644 --- a/studio/frontend/src/features/auth/components/auth-form.tsx +++ b/studio/frontend/src/features/auth/components/auth-form.tsx @@ -1,6 +1,7 @@ // 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 { apiUrl } from "@/lib/api-base"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; @@ -49,7 +50,7 @@ async function loginWithPassword( username: string, password: string, ): Promise { - const response = await fetch("/api/auth/login", { + const response = await fetch(apiUrl("/api/auth/login"), { method: "POST", headers: { "Content-Type": "application/json", @@ -96,7 +97,7 @@ export function AuthForm({ mode }: AuthFormProps): ReactElement | null { // (e.g. tokens from a previous install attempt). The server's // /api/auth/status is the source of truth for requires_password_change. try { - const response = await fetch("/api/auth/status"); + const response = await fetch(apiUrl("/api/auth/status")); if (!response.ok) throw new Error("Failed to load auth status."); const result = (await response.json()) as AuthStatusResponse; if (!canceled) { @@ -147,14 +148,15 @@ export function AuthForm({ mode }: AuthFormProps): ReactElement | null { }; }, [navigate]); - // Seed password from bootstrap credentials injected into HTML + // Seed password from bootstrap credentials injected into HTML by web CLI. useEffect(() => { - const bootstrap = window.__UNSLOTH_BOOTSTRAP__; - if (bootstrap) { - if (!isLoginMode && !password) { + function loadBootstrap() { + const bootstrap = window.__UNSLOTH_BOOTSTRAP__; + if (bootstrap && !isLoginMode && !password) { setPassword(bootstrap.password); } } + loadBootstrap(); }, []); const blockedByState = @@ -241,7 +243,7 @@ export function AuthForm({ mode }: AuthFormProps): ReactElement | null { accessToken = bootstrapToken.access_token; } - const response = await fetch("/api/auth/change-password", { + const response = await fetch(apiUrl("/api/auth/change-password"), { method: "POST", headers: { "Content-Type": "application/json", diff --git a/studio/frontend/src/features/auth/index.ts b/studio/frontend/src/features/auth/index.ts index af629cfb9a..9cc1599195 100644 --- a/studio/frontend/src/features/auth/index.ts +++ b/studio/frontend/src/features/auth/index.ts @@ -15,3 +15,8 @@ export { resetOnboardingDone, setMustChangePassword, } from "./session"; +export { + clearTauriAuthFailure, + getTauriAuthFailure, + tauriAutoAuth, +} from "./tauri-auto-auth"; diff --git a/studio/frontend/src/features/auth/session.ts b/studio/frontend/src/features/auth/session.ts index 6012174077..49a2722bdf 100644 --- a/studio/frontend/src/features/auth/session.ts +++ b/studio/frontend/src/features/auth/session.ts @@ -2,6 +2,7 @@ // Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 import { usePlatformStore } from "@/config/env"; +import { isTauri } from "@/lib/api-base"; export const AUTH_TOKEN_KEY = "unsloth_auth_token"; export const AUTH_REFRESH_TOKEN_KEY = "unsloth_auth_refresh_token"; @@ -78,6 +79,7 @@ export function resetOnboardingDone(): void { } export function getPostAuthRoute(): PostAuthRoute { + if (isTauri) return "/chat"; if (mustChangePassword()) return "/change-password"; if (usePlatformStore.getState().isChatOnly()) return "/chat"; return "/chat"; diff --git a/studio/frontend/src/features/auth/tauri-auto-auth.ts b/studio/frontend/src/features/auth/tauri-auto-auth.ts new file mode 100644 index 0000000000..d67730d2d6 --- /dev/null +++ b/studio/frontend/src/features/auth/tauri-auto-auth.ts @@ -0,0 +1,95 @@ +// 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 { isTauri } from "@/lib/api-base"; +import { + hasAuthToken, + hasRefreshToken, + mustChangePassword, + storeAuthTokens, +} from "./session"; +import { refreshSession } from "./api"; + +type DesktopAuthResponse = { + access_token: string; + refresh_token: string; +}; + +// Concurrency guard: multiple route guards can call tauriAutoAuth simultaneously. +// Without this, the first-launch password-change could race with itself. +let pending: Promise | null = null; +let lastTauriAuthFailure: string | null = null; + +const TAURI_AUTH_FAILURE_FALLBACK = + "Desktop authentication failed. Update or repair the managed Studio install, then restart Studio."; +const BACKEND_NOT_READY_MESSAGE = "Backend is not ready"; + +function authFailureMessage(error: unknown): string { + if (typeof error === "string" && error) return error; + if (error instanceof Error && error.message) return error.message; + return TAURI_AUTH_FAILURE_FALLBACK; +} + +export function getTauriAuthFailure(): string | null { + return lastTauriAuthFailure; +} + +export function clearTauriAuthFailure(): void { + lastTauriAuthFailure = null; +} + +function setTauriAuthFailure(error: unknown): void { + lastTauriAuthFailure = authFailureMessage(error); + window.dispatchEvent( + new CustomEvent("tauri-auth-failed", { detail: lastTauriAuthFailure }), + ); +} + +function isBackendNotReady(error: unknown): boolean { + return authFailureMessage(error).includes(BACKEND_NOT_READY_MESSAGE); +} + +async function doTauriAutoAuth(): Promise { + // Desktop must handle password-change state internally in Rust. + if (hasAuthToken() && !mustChangePassword()) { + clearTauriAuthFailure(); + return true; + } + + // Try refreshing existing session + if (hasRefreshToken()) { + const refreshed = await refreshSession(); + if (refreshed && hasAuthToken() && !mustChangePassword()) { + clearTauriAuthFailure(); + return true; + } + } + + try { + const { invoke } = await import("@tauri-apps/api/core"); + const tokens = await invoke("desktop_auth"); + storeAuthTokens(tokens.access_token, tokens.refresh_token, false); + clearTauriAuthFailure(); + return true; + } catch (error) { + if (isBackendNotReady(error)) return false; + setTauriAuthFailure(error); + return false; + } +} + +/** + * Silently authenticate in Tauri desktop mode. + * + * Delegates bootstrap/password handling to Rust and only stores returned tokens. + * + * Returns true if authentication succeeded. + * Concurrent calls are coalesced into a single in-flight attempt. + */ +export function tauriAutoAuth(): Promise { + if (!isTauri) return Promise.resolve(false); + if (!pending) { + pending = doTauriAutoAuth().finally(() => { pending = null; }); + } + return pending; +} diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index 53a935cf45..a93120397e 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -385,6 +385,8 @@ async function autoLoadSmallestModel(): Promise<{ supportsReasoning: loadResp.supports_reasoning ?? false, reasoningAlwaysOn: loadResp.reasoning_always_on ?? false, reasoningEnabled: loadResp.supports_reasoning ?? false, + reasoningStyle: loadResp.reasoning_style ?? "enable_thinking", + supportsPreserveThinking: loadResp.supports_preserve_thinking ?? false, supportsTools: loadResp.supports_tools ?? false, toolsEnabled: loadResp.supports_tools ?? false, codeToolsEnabled: loadResp.supports_tools ?? false, @@ -433,6 +435,14 @@ async function autoLoadSmallestModel(): Promise<{ sfLoadResp.requires_trust_remote_code ?? false, ); store.setParams({ ...store.params, maxTokens: 4096 }); + useChatRuntimeStore.setState({ + supportsReasoning: sfLoadResp.supports_reasoning ?? false, + reasoningAlwaysOn: sfLoadResp.reasoning_always_on ?? false, + reasoningEnabled: sfLoadResp.supports_reasoning ?? false, + reasoningStyle: sfLoadResp.reasoning_style ?? "enable_thinking", + supportsPreserveThinking: sfLoadResp.supports_preserve_thinking ?? false, + supportsTools: sfLoadResp.supports_tools ?? false, + }); const sfModel: ChatModelSummary = { id: repo.repo_id, name: sfLoadResp.display_name ?? repo.repo_id, @@ -501,6 +511,8 @@ async function autoLoadSmallestModel(): Promise<{ supportsReasoning: loadResp.supports_reasoning ?? false, reasoningAlwaysOn: loadResp.reasoning_always_on ?? false, reasoningEnabled: loadResp.supports_reasoning ?? false, + reasoningStyle: loadResp.reasoning_style ?? "enable_thinking", + supportsPreserveThinking: loadResp.supports_preserve_thinking ?? false, supportsTools: loadResp.supports_tools ?? false, toolsEnabled: loadResp.supports_tools ?? false, codeToolsEnabled: loadResp.supports_tools ?? false, @@ -696,7 +708,14 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { let serverMetadata: { usage?: ServerUsage; timings?: ServerTimings } | null = null; try { - const { supportsReasoning, reasoningEnabled } = runtime; + const { + supportsReasoning, + reasoningEnabled, + reasoningStyle, + reasoningEffort, + supportsPreserveThinking, + preserveThinking, + } = runtime; const stream = streamChatCompletions( { model: params.checkpoint, @@ -712,7 +731,12 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { image_base64: imageBase64, audio_base64: audioBase64, ...(useAdapter === undefined ? {} : { use_adapter: useAdapter }), - ...(supportsReasoning ? { enable_thinking: reasoningEnabled } : {}), + ...(supportsReasoning + ? reasoningStyle === "reasoning_effort" + ? { reasoning_effort: reasoningEffort } + : { enable_thinking: reasoningEnabled } + : {}), + ...(supportsPreserveThinking ? { preserve_thinking: preserveThinking } : {}), ...(supportsTools && (toolsEnabled || codeToolsEnabled) ? { enable_tools: true, diff --git a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts index 6c870d74a9..c9dcd911a2 100644 --- a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts +++ b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts @@ -253,7 +253,9 @@ export function useChatModelRuntime() { max_context_length: statusRes.max_context_length, native_context_length: statusRes.native_context_length, supports_reasoning: statusRes.supports_reasoning, + reasoning_style: statusRes.reasoning_style, reasoning_always_on: statusRes.reasoning_always_on, + supports_preserve_thinking: statusRes.supports_preserve_thinking, supports_tools: statusRes.supports_tools, speculative_type: statusRes.speculative_type, }; @@ -265,6 +267,8 @@ export function useChatModelRuntime() { // Restore reasoning/tools support flags and context length const supportsReasoning = statusRes.supports_reasoning ?? false; const reasoningAlwaysOn = statusRes.reasoning_always_on ?? false; + const reasoningStyle = statusRes.reasoning_style ?? "enable_thinking"; + const supportsPreserveThinking = statusRes.supports_preserve_thinking ?? false; const supportsTools = statusRes.supports_tools ?? false; const currentGgufContextLength = statusRes.is_gguf ? (statusRes.context_length ?? null) @@ -279,7 +283,14 @@ export function useChatModelRuntime() { useChatRuntimeStore.setState({ supportsReasoning, reasoningAlwaysOn, + reasoningStyle, + supportsPreserveThinking, supportsTools, + // Reset per-turn reasoning flag so models that do not support + // reasoning do not inherit a stale off state from a prior model. + reasoningEnabled: supportsReasoning + ? useChatRuntimeStore.getState().reasoningEnabled + : true, ggufContextLength: currentGgufContextLength, ggufMaxContextLength, ggufNativeContextLength, @@ -498,6 +509,8 @@ export function useChatModelRuntime() { supportsReasoning: loadResponse.supports_reasoning ?? false, reasoningAlwaysOn, reasoningEnabled: reasoningAlwaysOn ? true : reasoningDefault, + reasoningStyle: loadResponse.reasoning_style ?? "enable_thinking", + supportsPreserveThinking: loadResponse.supports_preserve_thinking ?? false, supportsTools: loadResponse.supports_tools ?? false, toolsEnabled: loadResponse.supports_tools ?? false, codeToolsEnabled: loadResponse.supports_tools ?? false, diff --git a/studio/frontend/src/features/chat/shared-composer.tsx b/studio/frontend/src/features/chat/shared-composer.tsx index e133669a1c..3ba2f8b6f5 100644 --- a/studio/frontend/src/features/chat/shared-composer.tsx +++ b/studio/frontend/src/features/chat/shared-composer.tsx @@ -4,6 +4,13 @@ import { TooltipIconButton } from "@/components/assistant-ui/tooltip-icon-button"; import { CodeToggleIcon } from "@/components/assistant-ui/code-toggle-icon"; import { Button } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { applyQwenThinkingParams } from "@/features/chat/utils/qwen-params"; import { AUDIO_ACCEPT, MAX_AUDIO_SIZE, fileToBase64 } from "@/lib/audio-utils"; import { useAui } from "@assistant-ui/react"; import { cn } from "@/lib/utils"; @@ -245,6 +252,12 @@ export function SharedComposer({ const reasoningAlwaysOn = useChatRuntimeStore((s) => s.reasoningAlwaysOn); const reasoningEnabled = useChatRuntimeStore((s) => s.reasoningEnabled); const setReasoningEnabled = useChatRuntimeStore((s) => s.setReasoningEnabled); + const reasoningStyle = useChatRuntimeStore((s) => s.reasoningStyle); + const reasoningEffort = useChatRuntimeStore((s) => s.reasoningEffort); + const setReasoningEffort = useChatRuntimeStore((s) => s.setReasoningEffort); + const supportsPreserveThinking = useChatRuntimeStore((s) => s.supportsPreserveThinking); + const preserveThinking = useChatRuntimeStore((s) => s.preserveThinking); + const setPreserveThinking = useChatRuntimeStore((s) => s.setPreserveThinking); const supportsTools = useChatRuntimeStore((s) => s.supportsTools); const toolsEnabled = useChatRuntimeStore((s) => s.toolsEnabled); const setToolsEnabled = useChatRuntimeStore((s) => s.setToolsEnabled); @@ -391,6 +404,13 @@ export function SharedComposer({ store.setModelRequiresTrustRemoteCode( resp.requires_trust_remote_code ?? false, ); + useChatRuntimeStore.setState({ + supportsReasoning: resp.supports_reasoning ?? false, + reasoningAlwaysOn: resp.reasoning_always_on ?? false, + reasoningStyle: resp.reasoning_style ?? "enable_thinking", + supportsPreserveThinking: resp.supports_preserve_thinking ?? false, + supportsTools: resp.supports_tools ?? false, + }); return resp.status; } @@ -566,41 +586,93 @@ export function SharedComposer({ )} - + + + {(["low", "medium", "high"] as const).map((level) => ( + setReasoningEffort(level)} + > + {level.charAt(0).toUpperCase() + level.slice(1)} + {reasoningEffort === level ? " \u2713" : ""} + + ))} + + + ) : ( + + )} + {supportsPreserveThinking && ( + + > + {preserveThinking && modelLoaded ? ( + + ) : ( + + )} + Preserve Thinking + + )}