* fix(studio/colab): merge iframe+keepalive into start(), add proxy_headers to uvicorn - Move serve_kernel_port_as_iframe and keepalive loop into colab.start() so both run in the same cell execution context, eliminating the race where the proxy URL was shown before the iframe cell had a chance to run - Add a 2s sleep after run_server() before show_link() to give Colab's proxy infrastructure time to register the bound port - Add proxy_headers=True and forwarded_allow_ips="*" to uvicorn Config so X-Forwarded-Proto/Host from Colab's reverse proxy are trusted - Simplify notebook start cell (no more separate iframe cell needed) * fix(studio/colab): fix iframe blocking and server thread crash in Colab Two root causes for the long-standing proxy/iframe breakage: 1. SecurityHeadersMiddleware set X-Frame-Options: DENY and frame-ancestors 'none' unconditionally, blocking serve_kernel_port_as_iframe regardless of server health. Fix: detect Colab via COLAB_BACKEND_URL/COLAB_GPU env vars, relax frame-ancestors to *.prod.colab.dev and omit X-Frame-Options. 2. asyncio.run() in the daemon thread conflicted with nest_asyncio's global patches applied on the main thread, causing the server to crash silently after ready_event fired. Fix: use explicit new_event_loop() + run_until_complete() in the daemon thread to bypass nest_asyncio's asyncio.run patch. Also replace blind time.sleep(2) with a health endpoint poll so the link and iframe are only shown once the server is truly reachable. * fix(studio/colab): use reliable /content + google.colab path for Colab detection COLAB_BACKEND_URL and COLAB_GPU env vars aren't consistently set across all Colab runtime versions. Use /content dir + google.colab package path as a more reliable signal, computed once at module load. * fix(studio/colab): fix port mismatch, health-check silence, and CSP framing Four bugs causing the iframe and URL button to always fail: 1. Port not propagated back: run_server auto-increments when 8888 is taken, but start() kept using the original port for show_link() and serve_kernel_port_as_iframe() — now reads app.state.server_port. 2. Silent health-check failure: the poll loop never checked whether any attempt succeeded; on all-fail it continued and showed a dead link — now exits early with a clear error message. 3. CSP frame-ancestors too narrow: '*.prod.colab.dev' only matches one subdomain level; actual Colab proxy URLs are two levels deep (e.g. foo.region.prod.colab.dev), and the parent frame may also be colab.research.google.com or a sandboxed null-origin output iframe — changed to '*' in Colab mode (single-user sandbox, no security loss). 4. _IS_COLAB detection hardcoded python3.10/3.11 paths: Python 3.12+ Colab runtimes wouldn't match when env vars aren't set — replaced with a glob over python3.*/dist-packages/google/colab. * fix(studio/colab): harden Colab startup against every known failure mode colab.py: - get_colab_url: retry eval_js up to 3x (10s timeout each), validate that result is a real https:// URL containing the port before accepting it; log a clear warning when falling back to localhost - show_link: safe short_url truncation (try/except around str.index so an unexpected URL shape never blocks the link card from rendering); also emit the URL via logger so it's visible in cell text output even if HTML display is suppressed - start: detect "already running" at entry — on cell re-run Studio is still healthy on port 8888; skip re-launch and go straight to show+iframe so the user never ends up with mismatched port state - start: wrap run_server in try/except (SystemExit + Exception) so startup errors surface as readable messages rather than cell crashes - start: check frontend_path/index.html exists, not just the directory - start: remove unused `import sys` - start / keepalive: catch KeyboardInterrupt so interrupting the cell prints a clean "stopped" message instead of a raw traceback - extract _is_studio_healthy() and _show_and_embed() helpers to deduplicate the fast-path and normal-path logic main.py: - _build_csp: in Colab mode, extend script-src to include *.prod.colab.dev and *.googleusercontent.com (Colab injects scripts from these origins into the output iframe scaffolding) - _build_csp: in Colab mode, extend connect-src with blob:, data:, wss://*.prod.colab.dev, and wss://*.googleusercontent.com so WebSocket streams and Colab kernel traffic are not blocked by CSP * fix(studio/colab): fix iframe width responsiveness and height sizing Replace serve_kernel_port_as_iframe with a raw CSS iframe for two reasons: 1. Width responsiveness: serve_kernel_port_as_iframe sets the width as an HTML attribute (width="100%") which Colab's output machinery can bake into a fixed pixel value on first render, causing the Studio to stop following the notebook panel width when it opens/closes or the window resizes. A CSS style property (style="width:100%") participates in normal reflow and always tracks the parent container width. 2. Height sizing: the hardcoded height=1200 was too tall on short monitors (forced outer-page scroll) and wasted space on tall ones. A small JS snippet reads screen.availHeight and sets height to ~82% of the screen, clamped to [600, 1100]px, with a resize listener that re-fits on zoom changes and panel open/close events. Also eliminate the double eval_js call: _show_and_embed now fetches the Colab proxy URL once and passes it to show_link via the new _url kwarg, so google.colab.kernel.proxyPort is only called once per invocation. Falls back to serve_kernel_port_as_iframe if IPython.display.HTML is unavailable for any reason. * fix(studio/colab): fix link button + add fullscreen hover button to iframe Link button: target="_blank" is blocked by Colab's output sandbox. Switch to onclick="window.open(url,'_blank')" which the sandbox allows. Fullscreen: add a small button that appears on hover in the top-right corner of the iframe. Clicking it calls requestFullscreen() on the wrapper div and stretches the iframe to 100vh/100vw. Exits back to normal on fullscreen change. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * revert(studio/colab): remove fullscreen button * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * fix(studio/colab): address review feedback - Wrap both urlopen calls in with statements to prevent socket/fd leaks - Replace JS resize listener with CSS height:82vh — simpler, responsive, and no risk of leaked window listeners on cell re-runs - Use importlib.util.find_spec("google.colab") instead of a glob path to detect Colab; more robust across Python versions and venv layouts * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * fix(studio/colab): fall back to href navigation when window.open is blocked window.open from a cross-origin sandboxed Colab output iframe can be silently blocked by the browser (returns null, no exception). The old code returned false unconditionally, so a blocked popup left the button doing nothing. Now: if window.open succeeds the new tab opens and the href is suppressed; if it returns null the browser follows the href, navigating the output cell to Studio — always does something useful. * fix(studio/colab): remove button, give iframe a branded header bar The "Open Unsloth Studio" button was unreliable in Colab's sandboxed output context regardless of how window.open was called. Since the iframe already loads Studio inline, the button added no value and confused users with a URL that 404s outside the output cell. Replace the separate link card + bare iframe with a single block: a slim black header bar (Unsloth logo + truncated URL) flush on top of the full-height responsive iframe. Cleaner and removes the broken button entirely. * studio: gate uvicorn proxy_headers/forwarded_allow_ips behind _IS_COLAB forwarded_allow_ips="*" was applied unconditionally, so every Studio deployment trusted X-Forwarded-* headers from any client. Only Colab needs that, because its reverse proxy fronts the kernel. For a normal local/standalone Studio this is an unwanted relaxation, especially when bound to 0.0.0.0. Now proxy_headers/forwarded_allow_ips are only set when _IS_COLAB. Standalone runs fall back to uvicorn's defaults (proxy_headers honored from loopback only), restoring the prior security posture, while Colab keeps the wide trust its proxy requires. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han <danielhanchen@gmail.com>
931 lines
33 KiB
Python
931 lines
33 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""
|
|
Main FastAPI application for Unsloth UI Backend
|
|
"""
|
|
|
|
import os
|
|
import sys
|
|
from pathlib import Path as _Path
|
|
|
|
# Suppress annoying C-level dependency warnings globally
|
|
os.environ["PYTHONWARNINGS"] = "ignore"
|
|
|
|
# Ensure backend dir is on sys.path so _platform_compat is importable when
|
|
# main.py is launched directly (e.g. `uvicorn main:app`).
|
|
_backend_dir = str(_Path(__file__).parent)
|
|
if _backend_dir not in sys.path:
|
|
sys.path.insert(0, _backend_dir)
|
|
|
|
# `uvicorn main:app` bypasses run.py; seed thread caps here too.
|
|
from utils.cpu_threads import configure_cpu_threads
|
|
|
|
try:
|
|
configure_cpu_threads()
|
|
except ValueError as exc:
|
|
_raw = os.environ.get("UNSLOTH_CPU_THREADS")
|
|
raise SystemExit(
|
|
f"Error: Invalid UNSLOTH_CPU_THREADS value {_raw!r}: {exc}"
|
|
) from None
|
|
|
|
# Fix for Anaconda/conda-forge Python: seed platform._sys_version_cache before
|
|
# any library imports that trigger attrs -> rich -> structlog -> platform crash.
|
|
# See: https://github.com/python/cpython/issues/102396
|
|
import _platform_compat # noqa: F401
|
|
|
|
# Direct `uvicorn main:app` launches bypass run.py, so re-export here too
|
|
# (mirrors run.py). Required BEFORE the unsloth-zoo import below, since
|
|
# its LLAMA_CPP_DEFAULT_DIR binding is import-time.
|
|
from utils.paths.storage_roots import studio_root as _studio_root
|
|
|
|
try:
|
|
_LEGACY_STUDIO_ROOT = (_Path.home() / ".unsloth" / "studio").resolve()
|
|
except (OSError, ValueError):
|
|
_LEGACY_STUDIO_ROOT = _Path.home() / ".unsloth" / "studio"
|
|
try:
|
|
_STUDIO_ROOT_RESOLVED = _studio_root().resolve()
|
|
except (OSError, ValueError):
|
|
_STUDIO_ROOT_RESOLVED = _studio_root()
|
|
if _STUDIO_ROOT_RESOLVED != _LEGACY_STUDIO_ROOT:
|
|
if not os.environ.get("UNSLOTH_STUDIO_HOME"):
|
|
os.environ["UNSLOTH_STUDIO_HOME"] = str(_STUDIO_ROOT_RESOLVED)
|
|
if not os.environ.get("UNSLOTH_LLAMA_CPP_PATH"):
|
|
os.environ["UNSLOTH_LLAMA_CPP_PATH"] = str(_STUDIO_ROOT_RESOLVED / "llama.cpp")
|
|
|
|
import hashlib
|
|
import mimetypes
|
|
import re as _re
|
|
import shutil
|
|
import warnings
|
|
from contextlib import asynccontextmanager
|
|
from importlib.metadata import PackageNotFoundError, version as package_version
|
|
from typing import Optional
|
|
from urllib.parse import urlparse
|
|
|
|
|
|
_STUDIO_INSTALL_ID_RE = _re.compile(r"^[0-9a-f]{64}$")
|
|
|
|
|
|
def _read_studio_install_id() -> str:
|
|
"""Per-install opaque id written by install.sh / install.ps1 at
|
|
$STUDIO_HOME/share/studio_install_id. Returns "" when the file is
|
|
absent (pre-PR install, fresh tree never run through the installer)
|
|
or contains anything other than a 64-char lowercase-hex token --
|
|
in which case /api/health emits "" and the launcher's _check_health
|
|
falls back to the existing "no baked id, accept any healthy
|
|
Unsloth backend" path. This intentionally replaces a previous
|
|
sha256(resolved_install_path) so the field carries no install-path
|
|
information for callers reaching /api/health (relevant when Studio
|
|
is run with -H 0.0.0.0)."""
|
|
try:
|
|
token = (
|
|
(_STUDIO_ROOT_RESOLVED / "share" / "studio_install_id").read_text().strip()
|
|
)
|
|
except (OSError, ValueError):
|
|
return ""
|
|
return token if _STUDIO_INSTALL_ID_RE.fullmatch(token) else ""
|
|
|
|
|
|
_STUDIO_ROOT_ID_CACHE: str = _read_studio_install_id()
|
|
|
|
|
|
def _studio_root_id() -> str:
|
|
"""Same-install discriminator for /api/health: a per-install opaque
|
|
token written once by the installer and read once at module import.
|
|
Empty when no installer-written token is present; the launcher
|
|
contract treats "" as "no baked id, accept any healthy backend"."""
|
|
return _STUDIO_ROOT_ID_CACHE
|
|
|
|
|
|
# Fix broken Windows registry MIME types. Some Windows installs map .js to
|
|
# "text/plain" in the registry (HKCR\.js\Content Type). Python's mimetypes
|
|
# module reads from the registry, and FastAPI/Starlette's StaticFiles uses
|
|
# mimetypes.guess_type() to set Content-Type headers. Browsers enforce strict
|
|
# MIME checking for ES module scripts (<script type="module">) and will refuse
|
|
# to execute .js files served as text/plain — resulting in a blank page.
|
|
# Calling add_type() *before* StaticFiles is instantiated ensures the correct
|
|
# types are used regardless of the OS registry.
|
|
if sys.platform == "win32":
|
|
mimetypes.add_type("application/javascript", ".js")
|
|
mimetypes.add_type("text/css", ".css")
|
|
|
|
# Suppress annoying dependency warnings in production
|
|
if os.getenv("ENVIRONMENT_TYPE", "production") == "production":
|
|
warnings.filterwarnings("ignore")
|
|
# Alternatively, you can be more specific:
|
|
# warnings.filterwarnings("ignore", category=DeprecationWarning)
|
|
# warnings.filterwarnings("ignore", module="triton.*")
|
|
|
|
from fastapi import Depends, FastAPI, HTTPException, Request
|
|
from fastapi.middleware.cors import CORSMiddleware
|
|
from fastapi.staticfiles import StaticFiles
|
|
from fastapi.responses import FileResponse, HTMLResponse, Response
|
|
from pathlib import Path
|
|
from datetime import datetime
|
|
|
|
# Import routers
|
|
from routes import (
|
|
auth_router,
|
|
chat_history_router,
|
|
data_recipe_router,
|
|
datasets_router,
|
|
export_router,
|
|
inference_router,
|
|
inference_studio_router,
|
|
mcp_servers_router,
|
|
models_router,
|
|
providers_router,
|
|
training_history_router,
|
|
training_router,
|
|
)
|
|
from auth import storage
|
|
from auth.authentication import get_current_subject
|
|
from utils.hardware import (
|
|
detect_hardware,
|
|
get_device,
|
|
DeviceType,
|
|
get_backend_visible_gpu_info,
|
|
)
|
|
import utils.hardware.hardware as _hw_module
|
|
|
|
from utils.cache_cleanup import clear_unsloth_compiled_cache
|
|
from utils.native_path_leases import native_path_leases_supported
|
|
from utils.update_status import (
|
|
get_studio_install_source_status,
|
|
get_studio_update_status,
|
|
)
|
|
from utils.studio_version import get_studio_version
|
|
|
|
|
|
def get_unsloth_version() -> str:
|
|
try:
|
|
return package_version("unsloth")
|
|
except PackageNotFoundError:
|
|
pass
|
|
|
|
version_file = (
|
|
_Path(__file__).resolve().parents[2] / "unsloth" / "models" / "_utils.py"
|
|
)
|
|
try:
|
|
for line in version_file.read_text(encoding = "utf-8").splitlines():
|
|
if line.startswith("__version__ = "):
|
|
return line.split("=", 1)[1].strip().strip('"').strip("'")
|
|
except OSError:
|
|
pass
|
|
return "dev"
|
|
|
|
|
|
UNSLOTH_VERSION = get_unsloth_version()
|
|
STUDIO_VERSION = get_studio_version()
|
|
|
|
|
|
def _load_desktop_owner() -> dict[str, str] | None:
|
|
token = os.environ.pop("UNSLOTH_STUDIO_DESKTOP_OWNER_TOKEN", "")
|
|
kind = os.environ.pop("UNSLOTH_STUDIO_DESKTOP_OWNER_KIND", "")
|
|
if kind != "tauri" or not token:
|
|
return None
|
|
return {
|
|
"kind": "tauri",
|
|
"token_sha256": hashlib.sha256(token.encode("utf-8")).hexdigest(),
|
|
}
|
|
|
|
|
|
_DESKTOP_OWNER = _load_desktop_owner()
|
|
|
|
|
|
def _desktop_owner() -> dict[str, str] | None:
|
|
return _DESKTOP_OWNER
|
|
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app: FastAPI):
|
|
"""Startup: detect hardware, seed default admin if needed. Shutdown: clean up compiled cache."""
|
|
# Clean up any stale compiled cache from previous runs
|
|
clear_unsloth_compiled_cache()
|
|
|
|
# Remove stale .venv_overlay from previous versions — no longer used.
|
|
# Version switching now uses .venv_t5/ (pre-installed by setup.sh).
|
|
overlay_dir = Path(__file__).resolve().parent.parent.parent / ".venv_overlay"
|
|
if overlay_dir.is_dir():
|
|
shutil.rmtree(overlay_dir, ignore_errors = True)
|
|
|
|
# Detect hardware first — sets DEVICE global used everywhere
|
|
detect_hardware()
|
|
|
|
# llama.cpp probes: capability (MTP support) + freshness (release age).
|
|
# Both cached; freshness has a 24h disk TTL.
|
|
try:
|
|
from core.inference.llama_cpp import LlamaCppBackend
|
|
from utils.llama_cpp_freshness import (
|
|
check_prebuilt_freshness,
|
|
format_stale_warning,
|
|
)
|
|
|
|
_bin = LlamaCppBackend._find_llama_server_binary()
|
|
_caps = LlamaCppBackend.probe_server_capabilities(_bin)
|
|
app.state.llama_cpp_capabilities = _caps
|
|
_freshness = check_prebuilt_freshness(_bin)
|
|
app.state.llama_cpp_freshness = _freshness
|
|
|
|
import structlog as _structlog
|
|
|
|
_log = _structlog.get_logger(__name__)
|
|
if _caps.get("found") and not _caps.get("supports_mtp"):
|
|
_msg = (
|
|
"llama.cpp prebuilt lacks MTP support "
|
|
"(--spec-type mtp/draft-mtp). Run `unsloth studio update`. "
|
|
"MTP GGUFs will load without speculative decoding."
|
|
)
|
|
_log.warning(_msg)
|
|
print(f"WARNING: {_msg}", flush = True)
|
|
if _freshness.get("stale"):
|
|
_msg = format_stale_warning(_freshness)
|
|
_log.warning(_msg)
|
|
print(f"WARNING: {_msg}", flush = True)
|
|
except Exception as _probe_exc:
|
|
import structlog as _structlog
|
|
|
|
_structlog.get_logger(__name__).debug(
|
|
"llama.cpp startup probes failed: %s", _probe_exc
|
|
)
|
|
|
|
from storage.studio_db import cleanup_orphaned_runs
|
|
|
|
try:
|
|
cleanup_orphaned_runs()
|
|
except Exception as exc:
|
|
import structlog
|
|
|
|
structlog.get_logger(__name__).warning(
|
|
"cleanup_orphaned_runs failed at startup: %s", exc
|
|
)
|
|
|
|
# Pre-cache the helper GGUF model for LLM-assisted dataset detection.
|
|
# Runs in a background thread so it doesn't block server startup.
|
|
import threading
|
|
|
|
def _precache():
|
|
try:
|
|
from utils.datasets.llm_assist import precache_helper_gguf
|
|
|
|
precache_helper_gguf()
|
|
except Exception:
|
|
pass # non-critical
|
|
|
|
threading.Thread(target = _precache, daemon = True).start()
|
|
|
|
# Initialize RSA key pair for API key encryption (external providers)
|
|
from core.inference.key_exchange import init_key_pair
|
|
|
|
init_key_pair()
|
|
|
|
if storage.ensure_default_admin():
|
|
bootstrap_pw = storage.get_bootstrap_password()
|
|
app.state.bootstrap_password = bootstrap_pw
|
|
|
|
bootstrap_path = storage.DB_PATH.parent / ".bootstrap_password"
|
|
print("\n" + "=" * 60)
|
|
print("DEFAULT ADMIN ACCOUNT CREATED")
|
|
print(f" username: {storage.DEFAULT_ADMIN_USERNAME}")
|
|
print(f" password saved to: {bootstrap_path}")
|
|
print(" Open the Studio UI to sign in and change it.")
|
|
print("=" * 60 + "\n")
|
|
else:
|
|
app.state.bootstrap_password = storage.get_bootstrap_password()
|
|
yield
|
|
# Cleanup
|
|
_hw_module.DEVICE = None
|
|
clear_unsloth_compiled_cache()
|
|
|
|
|
|
# Create FastAPI app
|
|
app = FastAPI(
|
|
title = "Unsloth UI Backend",
|
|
version = UNSLOTH_VERSION,
|
|
description = "Backend API for Unsloth UI - Training and Model Management",
|
|
lifespan = lifespan,
|
|
)
|
|
|
|
# Initialize structured logging
|
|
from loggers.config import LogConfig
|
|
from loggers.handlers import LoggingMiddleware
|
|
|
|
logger = LogConfig.setup_logging(
|
|
service_name = "unsloth-studio-backend",
|
|
env = os.getenv("ENVIRONMENT_TYPE", "production"),
|
|
)
|
|
|
|
app.add_middleware(LoggingMiddleware)
|
|
|
|
|
|
# Citation favicons load from www.google.com/s2/favicons; *.gstatic.com is
|
|
# kept for legacy web-search faviconV2 paths. Everything else is same-origin.
|
|
from starlette.middleware.base import BaseHTTPMiddleware # noqa: E402
|
|
from starlette.requests import Request as _StarletteRequest # noqa: E402
|
|
|
|
|
|
_CSP_SCRIPT_NONCE_HEADER = "x-internal-script-nonce"
|
|
|
|
|
|
# /content is Colab's working directory — more reliable than env vars which
|
|
# aren't always set depending on Colab runtime version.
|
|
import importlib.util as _importlib_util
|
|
|
|
_IS_COLAB = os.path.isdir("/content") and (
|
|
bool(os.environ.get("COLAB_BACKEND_URL"))
|
|
or bool(os.environ.get("COLAB_JUPYTER_IP"))
|
|
or _importlib_util.find_spec("google.colab") is not None
|
|
)
|
|
|
|
|
|
def _build_csp(script_nonce: "str | None" = None) -> str:
|
|
script_src = "script-src 'self'"
|
|
if script_nonce:
|
|
script_src += f" 'nonce-{script_nonce}'"
|
|
# In Colab the parent frame can be colab.research.google.com, a multi-level
|
|
# *.prod.colab.dev subdomain (e.g. foo.region.prod.colab.dev — note: CSP
|
|
# wildcards only match one level, so *.prod.colab.dev misses these), or a
|
|
# sandboxed null-origin output iframe. Use '*' so any ancestor is allowed;
|
|
# Colab is already a sandboxed single-user environment.
|
|
frame_ancestors = "*" if _IS_COLAB else "'none'"
|
|
|
|
# In Colab the frontend is served over the Colab reverse-proxy at an HTTPS
|
|
# *.prod.colab.dev URL. Colab's kernel communication layer and the output
|
|
# iframe scaffolding inject scripts from *.prod.colab.dev and
|
|
# *.googleusercontent.com, and make fetch/WebSocket connections to those
|
|
# same origins. Widen script-src and connect-src in Colab mode so those
|
|
# requests are not blocked. 'unsafe-inline' for scripts is still omitted;
|
|
# our own inline script uses a nonce.
|
|
if _IS_COLAB:
|
|
script_src += " https://*.prod.colab.dev https://*.googleusercontent.com"
|
|
connect_src = (
|
|
"'self' blob: data: "
|
|
"https://huggingface.co https://datasets-server.huggingface.co "
|
|
"https://*.prod.colab.dev wss://*.prod.colab.dev "
|
|
"https://*.googleusercontent.com wss://*.googleusercontent.com"
|
|
)
|
|
else:
|
|
connect_src = (
|
|
"'self' https://huggingface.co https://datasets-server.huggingface.co"
|
|
)
|
|
|
|
return (
|
|
"default-src 'self'; "
|
|
"img-src 'self' data: blob: https://t0.gstatic.com "
|
|
"https://t1.gstatic.com https://t2.gstatic.com "
|
|
"https://t3.gstatic.com https://www.google.com; "
|
|
f"connect-src {connect_src}; "
|
|
"style-src 'self' 'unsafe-inline'; "
|
|
f"{script_src}; "
|
|
"font-src 'self' data:; "
|
|
f"frame-ancestors {frame_ancestors}; "
|
|
"form-action 'self'; "
|
|
"base-uri 'self'"
|
|
)
|
|
|
|
|
|
class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
|
"""Set baseline security headers; splice per-response inline-script nonces into CSP."""
|
|
|
|
async def dispatch(self, request: _StarletteRequest, call_next):
|
|
response = await call_next(request)
|
|
# Strip the internal nonce hand-off header so it never reaches the client.
|
|
nonce = response.headers.get(_CSP_SCRIPT_NONCE_HEADER)
|
|
if nonce is not None:
|
|
del response.headers[_CSP_SCRIPT_NONCE_HEADER]
|
|
response.headers.setdefault("Content-Security-Policy", _build_csp(nonce))
|
|
# Omit X-Frame-Options in Colab — CSP frame-ancestors handles it, and
|
|
# DENY would block serve_kernel_port_as_iframe regardless of CSP.
|
|
if not _IS_COLAB:
|
|
response.headers.setdefault("X-Frame-Options", "DENY")
|
|
response.headers.setdefault("X-Content-Type-Options", "nosniff")
|
|
response.headers.setdefault("Referrer-Policy", "no-referrer")
|
|
response.headers.setdefault(
|
|
"Permissions-Policy",
|
|
"camera=(), microphone=(), geolocation=(), interest-cohort=()",
|
|
)
|
|
response.headers["server"] = "unsloth-studio"
|
|
return response
|
|
|
|
|
|
app.add_middleware(SecurityHeadersMiddleware)
|
|
|
|
|
|
# Cap upload body on protected POSTs; default 500 MB, env-tunable.
|
|
import json as _json_for_413 # noqa: E402
|
|
|
|
|
|
_MAX_BODY_BYTES = int(os.environ.get("UNSLOTH_STUDIO_MAX_BODY_MB", "500")) * 1024 * 1024
|
|
_BODY_PROTECTED_PREFIXES = (
|
|
"/v1/chat/completions",
|
|
"/v1/completions",
|
|
"/api/inference",
|
|
"/api/data-recipe",
|
|
"/api/datasets",
|
|
"/api/chat",
|
|
"/api/train",
|
|
"/api/export",
|
|
)
|
|
|
|
|
|
async def _send_413(send, total_bytes: int) -> None:
|
|
payload = _json_for_413.dumps(
|
|
{
|
|
"detail": (
|
|
f"Request body too large "
|
|
f"({total_bytes:,} bytes; max {_MAX_BODY_BYTES:,})."
|
|
)
|
|
},
|
|
).encode("utf-8")
|
|
await send(
|
|
{
|
|
"type": "http.response.start",
|
|
"status": 413,
|
|
"headers": [
|
|
(b"content-type", b"application/json"),
|
|
(b"content-length", str(len(payload)).encode("ascii")),
|
|
],
|
|
}
|
|
)
|
|
await send({"type": "http.response.body", "body": payload, "more_body": False})
|
|
|
|
|
|
class MaxBodyMiddleware:
|
|
"""Reject oversized bodies on protected POST/PUT/PATCH; raw ASGI so chunked uploads cannot bypass the cap."""
|
|
|
|
def __init__(self, app, max_bytes: int, protected_prefixes: tuple):
|
|
self.app = app
|
|
self.max_bytes = max_bytes
|
|
self.protected_prefixes = protected_prefixes
|
|
|
|
async def __call__(self, scope, receive, send):
|
|
if scope["type"] != "http":
|
|
await self.app(scope, receive, send)
|
|
return
|
|
method = scope.get("method", "").upper()
|
|
path = scope.get("path", "")
|
|
if method not in ("POST", "PUT", "PATCH") or not any(
|
|
path.startswith(p) for p in self.protected_prefixes
|
|
):
|
|
await self.app(scope, receive, send)
|
|
return
|
|
|
|
declared = None
|
|
for name, value in scope.get("headers", []):
|
|
if name == b"content-length":
|
|
try:
|
|
declared = int(value.decode("latin-1"))
|
|
except (ValueError, UnicodeDecodeError):
|
|
declared = None
|
|
break
|
|
if declared is not None and declared > self.max_bytes:
|
|
await _send_413(send, declared)
|
|
return
|
|
|
|
chunks: list = []
|
|
total = 0
|
|
while True:
|
|
msg = await receive()
|
|
mtype = msg.get("type")
|
|
if mtype == "http.disconnect":
|
|
return
|
|
if mtype != "http.request":
|
|
# Mid-stream unexpected frame: forwarding would corrupt downstream.
|
|
return
|
|
body = msg.get("body", b"") or b""
|
|
if body:
|
|
total += len(body)
|
|
if total > self.max_bytes:
|
|
await _send_413(send, total)
|
|
return
|
|
chunks.append(body)
|
|
if not msg.get("more_body", False):
|
|
break
|
|
|
|
replayed = {"sent": False}
|
|
|
|
async def replay_receive():
|
|
if not replayed["sent"]:
|
|
replayed["sent"] = True
|
|
return {
|
|
"type": "http.request",
|
|
"body": b"".join(chunks),
|
|
"more_body": False,
|
|
}
|
|
# After replay, fall through so http.disconnect still propagates.
|
|
return await receive()
|
|
|
|
await self.app(scope, replay_receive, send)
|
|
|
|
|
|
app.add_middleware(
|
|
MaxBodyMiddleware,
|
|
max_bytes = _MAX_BODY_BYTES,
|
|
protected_prefixes = _BODY_PROTECTED_PREFIXES,
|
|
)
|
|
|
|
|
|
from starlette.responses import RedirectResponse as _RedirectResponse # noqa: E402
|
|
|
|
|
|
@app.get("/recipes", include_in_schema = False)
|
|
@app.get("/recipes/{rest:path}", include_in_schema = False)
|
|
async def _recipes_redirect(rest: str = ""):
|
|
target = "/data-recipes" + (("/" + rest) if rest else "")
|
|
return _RedirectResponse(url = target, status_code = 308)
|
|
|
|
|
|
# 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
|
|
"http://localhost:5173", # Tauri dev/Vite
|
|
"http://127.0.0.1:5173", # Tauri dev/Vite fallback
|
|
]
|
|
_cors_origin_regex = None
|
|
else:
|
|
_cors_origin_regex = None
|
|
|
|
app.add_middleware(
|
|
CORSMiddleware,
|
|
allow_origins = _cors_origins,
|
|
allow_origin_regex = _cors_origin_regex,
|
|
allow_credentials = True,
|
|
allow_methods = ["*"],
|
|
allow_headers = ["*"],
|
|
)
|
|
|
|
# ============ Register API Routes ============
|
|
|
|
# Register routers
|
|
app.include_router(auth_router, prefix = "/api/auth", tags = ["auth"])
|
|
app.include_router(training_router, prefix = "/api/train", tags = ["training"])
|
|
app.include_router(models_router, prefix = "/api/models", tags = ["models"])
|
|
app.include_router(chat_history_router, prefix = "/api/chat", tags = ["chat"])
|
|
app.include_router(inference_router, prefix = "/api/inference", tags = ["inference"])
|
|
# Studio-only inference endpoints (cancel, etc.) are intentionally NOT
|
|
# exposed on the /v1 OpenAI-compat prefix below.
|
|
app.include_router(inference_studio_router, prefix = "/api/inference", tags = ["inference"])
|
|
|
|
# OpenAI-compatible endpoints: mount the same inference router at /v1
|
|
# so external tools (Open WebUI, SillyTavern, etc.) can use the
|
|
# standard /v1/chat/completions path.
|
|
app.include_router(inference_router, prefix = "/v1", tags = ["openai-compat"])
|
|
app.include_router(providers_router, prefix = "/api/providers", tags = ["providers"])
|
|
app.include_router(mcp_servers_router, prefix = "/api/mcp/servers", tags = ["mcp"])
|
|
app.include_router(datasets_router, prefix = "/api/datasets", tags = ["datasets"])
|
|
app.include_router(data_recipe_router, prefix = "/api/data-recipe", tags = ["data-recipe"])
|
|
app.include_router(export_router, prefix = "/api/export", tags = ["export"])
|
|
app.include_router(
|
|
training_history_router, prefix = "/api/train", tags = ["training-history"]
|
|
)
|
|
|
|
|
|
# ============ Health and System Endpoints ============
|
|
|
|
|
|
@app.get("/api/health")
|
|
async def health_check(request: Request):
|
|
"""Liveness plus launcher capability bits; install fingerprint gated on a valid bearer.
|
|
|
|
Unauthenticated callers (Tauri watchdog, frontend bootstrap polls) need
|
|
``service`` / ``studio_root_id`` / ``chat_only`` / ``desktop_*`` / ``native_path_leases_supported``
|
|
to (a) re-adopt a sibling backend across restarts and (b) gate UI surfaces
|
|
before any token is available. None of those leak install path or version.
|
|
``version`` / ``studio_version`` / ``device_type`` still require a bearer
|
|
because they fingerprint the host.
|
|
"""
|
|
base = {
|
|
"status": "healthy",
|
|
"timestamp": datetime.now().isoformat(),
|
|
"service": "Unsloth UI Backend",
|
|
"chat_only": _hw_module.CHAT_ONLY,
|
|
"desktop_protocol_version": 1,
|
|
"desktop_manageability_version": 1,
|
|
"supports_desktop_auth": True,
|
|
"supports_desktop_backend_ownership": True,
|
|
# Opaque per-install id; launchers reject sibling Studios on the same port.
|
|
"studio_root_id": _studio_root_id(),
|
|
"native_path_leases_supported": native_path_leases_supported(),
|
|
**({"desktop_owner": owner} if (owner := _desktop_owner()) else {}),
|
|
}
|
|
auth = request.headers.get("authorization", "")
|
|
if not auth.lower().startswith("bearer "):
|
|
return base
|
|
try:
|
|
from auth.authentication import get_current_subject as _gcs
|
|
from fastapi.security import HTTPAuthorizationCredentials
|
|
|
|
creds = HTTPAuthorizationCredentials(
|
|
scheme = "Bearer", credentials = auth.split(" ", 1)[1]
|
|
)
|
|
# Must await: a bare coroutine is truthy and would skip the auth check.
|
|
subject = await _gcs(creds)
|
|
except HTTPException:
|
|
return base
|
|
except Exception:
|
|
return base
|
|
if not subject:
|
|
return base
|
|
|
|
platform_map = {"darwin": "mac", "win32": "windows", "linux": "linux"}
|
|
device_type = platform_map.get(sys.platform, sys.platform)
|
|
return {
|
|
**base,
|
|
"version": UNSLOTH_VERSION,
|
|
"studio_version": STUDIO_VERSION,
|
|
"device_type": device_type,
|
|
}
|
|
|
|
|
|
@app.get("/api/studio/install-source")
|
|
def studio_install_source(_current_subject: str = Depends(get_current_subject)):
|
|
"""Return source-aware install metadata without remote update checks."""
|
|
return get_studio_install_source_status(UNSLOTH_VERSION)
|
|
|
|
|
|
@app.get("/api/studio/update-status")
|
|
def studio_update_status(_current_subject: str = Depends(get_current_subject)):
|
|
"""Return source-aware manual update status for browser-served Studio."""
|
|
return get_studio_update_status(UNSLOTH_VERSION)
|
|
|
|
|
|
@app.post("/api/shutdown")
|
|
async def shutdown_server(
|
|
request: Request,
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""Gracefully shut down the Unsloth Studio server.
|
|
|
|
Called by the frontend quit dialog so users can stop the server from the UI
|
|
without needing to use the CLI or kill the process manually.
|
|
"""
|
|
import asyncio
|
|
|
|
async def _delayed_shutdown():
|
|
await asyncio.sleep(0.2) # Let the HTTP response return first
|
|
trigger = getattr(request.app.state, "trigger_shutdown", None)
|
|
if trigger is not None:
|
|
trigger()
|
|
else:
|
|
# Fallback when not launched via run_server() (e.g. direct uvicorn)
|
|
import signal
|
|
import os
|
|
|
|
os.kill(os.getpid(), signal.SIGTERM)
|
|
|
|
request.app.state._shutdown_task = asyncio.create_task(_delayed_shutdown())
|
|
return {"status": "shutting_down"}
|
|
|
|
|
|
@app.get("/api/system")
|
|
async def get_system_info(
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""Get system information.
|
|
|
|
Gated behind auth: the response includes platform, Python version,
|
|
GPU name, memory total, and ML package set -- enough to fingerprint
|
|
a host. Studio's chat-only-mode design assumes only the local user
|
|
reaches /api/system; in -H 0.0.0.0 / Colab / Tauri-relayed setups
|
|
that assumption breaks unless we require a bearer.
|
|
"""
|
|
import platform
|
|
import psutil
|
|
from utils.hardware import get_device
|
|
from utils.hardware.hardware import _backend_label
|
|
|
|
visibility_info = get_backend_visible_gpu_info()
|
|
gpu_info = {
|
|
"available": visibility_info["available"],
|
|
"devices": visibility_info["devices"],
|
|
}
|
|
|
|
# CPU & Memory
|
|
memory = psutil.virtual_memory()
|
|
|
|
return {
|
|
"platform": platform.platform(),
|
|
"python_version": platform.python_version(),
|
|
# Use the centralized _backend_label helper so the /api/system
|
|
# endpoint reports "rocm" on AMD hosts instead of "cuda", matching
|
|
# the /api/hardware and /api/gpu-visibility endpoints.
|
|
"device_backend": _backend_label(get_device()),
|
|
"cpu_count": psutil.cpu_count(),
|
|
"memory": {
|
|
"total_gb": round(memory.total / 1e9, 2),
|
|
"available_gb": round(memory.available / 1e9, 2),
|
|
"percent_used": memory.percent,
|
|
},
|
|
"gpu": gpu_info,
|
|
}
|
|
|
|
|
|
@app.get("/api/system/gpu-visibility")
|
|
async def get_gpu_visibility(
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
return get_backend_visible_gpu_info()
|
|
|
|
|
|
@app.get("/api/system/hardware")
|
|
async def get_hardware_info(
|
|
current_subject: str = Depends(get_current_subject),
|
|
):
|
|
"""Return GPU name, total VRAM, and key ML package versions.
|
|
|
|
Gated behind auth alongside /api/system -- same fingerprinting
|
|
concern. /api/system/gpu-visibility is also auth-gated already.
|
|
"""
|
|
from utils.hardware import get_gpu_summary, get_package_versions
|
|
|
|
return {
|
|
"gpu": get_gpu_summary(),
|
|
"versions": get_package_versions(),
|
|
}
|
|
|
|
|
|
# ============ Serve Frontend (Optional) ============
|
|
|
|
|
|
def _strip_crossorigin(html_bytes: bytes) -> bytes:
|
|
"""Remove ``crossorigin`` attributes from script/link tags.
|
|
|
|
Vite adds ``crossorigin`` by default which forces CORS mode on font
|
|
subresource loads. When Studio is served over plain HTTP, Firefox
|
|
HTTPS-Only Mode does not exempt CORS font requests -- causing all
|
|
@font-face downloads to fail silently. Stripping the attribute
|
|
makes them regular same-origin fetches that work on any protocol.
|
|
"""
|
|
import re as _re
|
|
|
|
html = html_bytes.decode("utf-8")
|
|
html = _re.sub(r'\s+crossorigin(?:="[^"]*")?', "", html)
|
|
return html.encode("utf-8")
|
|
|
|
|
|
def _inject_bootstrap(html_bytes: bytes, app: FastAPI):
|
|
"""Inject bootstrap credentials when password change is pending.
|
|
Returns ``(html_bytes, script_nonce_or_None)``; callers forward the
|
|
nonce via ``_CSP_SCRIPT_NONCE_HEADER`` so CSP allows the inline script.
|
|
"""
|
|
import json as _json
|
|
import secrets as _secrets
|
|
|
|
if not storage.requires_password_change(storage.DEFAULT_ADMIN_USERNAME):
|
|
return html_bytes, None
|
|
|
|
bootstrap_pw = getattr(app.state, "bootstrap_password", None)
|
|
if not bootstrap_pw:
|
|
return html_bytes, None
|
|
|
|
payload = _json.dumps(
|
|
{
|
|
"username": storage.DEFAULT_ADMIN_USERNAME,
|
|
"password": bootstrap_pw,
|
|
}
|
|
)
|
|
nonce = _secrets.token_urlsafe(16)
|
|
tag = f'<script nonce="{nonce}">window.__UNSLOTH_BOOTSTRAP__={payload}</script>'
|
|
html = html_bytes.decode("utf-8")
|
|
html = html.replace("</head>", f"{tag}</head>", 1)
|
|
return html.encode("utf-8"), nonce
|
|
|
|
|
|
_DEFAULT_PORTS = {"http": 80, "https": 443, "ws": 80, "wss": 443}
|
|
|
|
|
|
def _canonical_origin(scheme: str, netloc: str) -> Optional[tuple[str, str, int]]:
|
|
"""Canonicalise an Origin to ``(scheme, host, port)`` for equality.
|
|
Browsers strip default ports (RFC 6454 sec 6.1) and scheme/host are
|
|
case-insensitive (RFC 3986), so bare string compare misclassifies
|
|
same-origin requests as cross-origin. Returns ``None`` on unparseable
|
|
input so callers fall to the safer cross-origin default.
|
|
"""
|
|
scheme = (scheme or "").strip().lower()
|
|
if not scheme or not netloc:
|
|
return None
|
|
# Strip userinfo (RFC 3986); Origin never carries credentials.
|
|
if "@" in netloc:
|
|
netloc = netloc.rsplit("@", 1)[1]
|
|
# IPv6 hosts use brackets (RFC 3986 sec 3.2.2): ``[::1]:8902``. Bare
|
|
# ``partition(":")`` mis-parses these and breaks ``unsloth studio -H ::1``.
|
|
if netloc.startswith("["):
|
|
close = netloc.find("]")
|
|
if close == -1:
|
|
return None
|
|
host = netloc[1:close]
|
|
rest = netloc[close + 1 :]
|
|
if rest.startswith(":"):
|
|
port_str = rest[1:]
|
|
elif rest == "":
|
|
port_str = ""
|
|
else:
|
|
return None
|
|
else:
|
|
host, _, port_str = netloc.partition(":")
|
|
host = host.strip().lower()
|
|
if not host:
|
|
return None
|
|
if port_str:
|
|
try:
|
|
port = int(port_str)
|
|
except ValueError:
|
|
return None
|
|
else:
|
|
port = _DEFAULT_PORTS.get(scheme, 0)
|
|
return (scheme, host, port)
|
|
|
|
|
|
def _is_same_origin_request(request: Request) -> bool:
|
|
"""True when Origin is missing or matches request's scheme://host:port.
|
|
Top-level same-document GETs omit Origin, so missing counts as same-origin.
|
|
Callers must also emit ``Vary: Origin``. Both sides are canonicalised via
|
|
:func:`_canonical_origin` so default-port stripping and scheme/host case
|
|
do not misclassify same-origin requests as cross-origin.
|
|
"""
|
|
origin = request.headers.get("origin")
|
|
if origin is None:
|
|
# Missing header: top-level same-document GETs omit Origin.
|
|
return True
|
|
# Empty string is not a valid serialised origin (RFC 6454 sec 6.1).
|
|
if not origin:
|
|
return False
|
|
# "null" token (sandboxed iframes, file:// pages) is never same-origin.
|
|
if origin == "null":
|
|
return False
|
|
# ``urlparse`` raises ``ValueError`` on malformed IPv6 brackets; swallow
|
|
# so a garbage Origin doesn't 500 the SPA handler.
|
|
try:
|
|
parsed = urlparse(origin)
|
|
except ValueError:
|
|
return False
|
|
origin_canon = _canonical_origin(parsed.scheme, parsed.netloc)
|
|
if origin_canon is None:
|
|
return False
|
|
try:
|
|
self_canon = _canonical_origin(request.url.scheme, request.url.netloc)
|
|
except ValueError:
|
|
return False
|
|
if self_canon is None:
|
|
return False
|
|
return origin_canon == self_canon
|
|
|
|
|
|
def setup_frontend(app: FastAPI, build_path: Path):
|
|
"""Mount frontend static files (optional)"""
|
|
if not build_path.exists():
|
|
return False
|
|
|
|
# Mount assets
|
|
assets_dir = build_path / "assets"
|
|
if assets_dir.exists():
|
|
app.mount("/assets", StaticFiles(directory = assets_dir), name = "assets")
|
|
|
|
def _build_index_response(request: Request) -> Response:
|
|
content = (build_path / "index.html").read_bytes()
|
|
content = _strip_crossorigin(content)
|
|
# Bootstrap pw is same-origin only; Vary: Origin keeps caches honest.
|
|
if _is_same_origin_request(request):
|
|
content, nonce = _inject_bootstrap(content, app)
|
|
else:
|
|
nonce = None
|
|
headers = {
|
|
"Cache-Control": "no-cache, no-store, must-revalidate",
|
|
"Vary": "Origin",
|
|
}
|
|
if nonce:
|
|
headers[_CSP_SCRIPT_NONCE_HEADER] = nonce
|
|
return Response(
|
|
content = content,
|
|
media_type = "text/html",
|
|
headers = headers,
|
|
)
|
|
|
|
@app.get("/")
|
|
async def serve_root(request: Request):
|
|
return _build_index_response(request)
|
|
|
|
@app.get("/{full_path:path}")
|
|
async def serve_frontend(request: Request, full_path: str):
|
|
if full_path in {"api", "v1"} or full_path.startswith(("api/", "v1/")):
|
|
return {"error": "API endpoint not found"}
|
|
|
|
file_path = (build_path / full_path).resolve()
|
|
|
|
# Block path traversal — ensure resolved path stays inside build_path
|
|
if not file_path.is_relative_to(build_path.resolve()):
|
|
return Response(status_code = 403)
|
|
|
|
if file_path.is_file():
|
|
return FileResponse(file_path)
|
|
|
|
# Serve index.html as bytes — avoids Content-Length mismatch
|
|
return _build_index_response(request)
|
|
|
|
return True
|