feat: multi-user authentication (dual mode)

Add optional login protection for the web UI/API, intended for shared
machines (e.g. AI servers). Dual mode: with no users configured the API
and web UI are open (as before); once at least one user exists, every
/api/* and /ws/* endpoint requires a valid session.

- bcrypt password hashing: passwords stored as $2b$ hashes in
  /etc/nvcurve/users.json (0600, root-owned); plaintext never persisted.
- 24-hour sessions: HttpOnly cookie for browsers, Authorization: Bearer
  token for CLI/scripts; in-memory, invalidated on server restart.
- Multi-user: multiple named accounts (no shared-password mode).
- New CLI: nvcurve user add|list|remove|set-password (root for mutating
  ops; password always prompted, never a CLI argument).
- New endpoints: GET /api/ping (public), /api/auth/status|login|logout|users.
- Web UI: sign-in screen when auth is enabled; status bar shows the
  signed-in user with sign-out; expired sessions (401) re-show sign-in.
- Brute-force lockout: 10 failed logins/IP within 5 min -> 15 min lockout.
- New dependency: bcrypt.

Also: LSP config (pyrightconfig.json) pointing at the project .venv, and
small error-handling cleanups in daemon.py/server.py.
This commit is contained in:
ARIA committed 2026-09-02 15:21:35 +02:00
1 parent af23a10f25
commit bbd692ea2e
16 files changed
+2085 -483

No files matched your search

+289
View File
@@ -0,0 +1,289 @@
"""Authentication: bcrypt user store, 24-hour sessions, dual-mode gating.
Dual mode
---------
- If the user store file exists and contains at least one user, the API
requires authentication for all ``/api/*`` and ``/ws/*`` endpoints
(except the public auth endpoints and ``/api/ping``).
- If the user store is absent or empty, the server runs with no
authentication at all — exactly like before.
Passwords are hashed with bcrypt (``$2b$``). Plaintext passwords are never
stored. Sessions are in-memory random 256-bit tokens valid for 24 hours;
they are presented either as an HttpOnly cookie (browser) or as an
``Authorization: Bearer <token>`` header (CLI / scripts).
A simple per-IP lockout slows down brute-force attempts against login.
"""
import json
import os
import secrets
import time
from typing import Any
import bcrypt
# ── Constants ─────────────────────────────────────────────────────────────────
SESSION_TTL_S = 24 * 60 * 60 # sessions live for 24 hours
COOKIE_NAME = "nvcurve_session"
BCRYPT_ROUNDS = 12
MAX_PASSWORD_BYTES = 72 # bcrypt only uses the first 72 bytes
# Brute-force lockout: MAX_FAILURES within WINDOW_S → locked for LOCKOUT_S.
_MAX_FAILURES = 10
_WINDOW_S = 300
_LOCKOUT_S = 900
# ── User store ────────────────────────────────────────────────────────────────
# File format (JSON): {"users": {"alice": "$2b$12$...", "bob": "$2b$12$..."}}
def _valid_username(name: str) -> bool:
if not name or len(name) > 32:
return False
return all(c.isalnum() or c in "_-." for c in name)
def load_users(path: str) -> dict[str, str]:
"""Load the user store. Returns {} if the file is absent or unreadable."""
try:
with open(path) as f:
data = json.load(f)
except (FileNotFoundError, json.JSONDecodeError, OSError):
return {}
users = data.get("users", {})
if not isinstance(users, dict):
return {}
return {k: v for k, v in users.items() if isinstance(k, str) and isinstance(v, str)}
def save_users(path: str, users: dict[str, str]) -> None:
"""Write the user store with restrictive permissions (0600, root-owned)."""
directory = os.path.dirname(path)
if directory:
try:
os.makedirs(directory, exist_ok=True)
except OSError as exc:
raise ValueError(
f"cannot create user store directory {directory!r}: {exc}"
) from exc
try:
fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
except OSError as exc:
raise ValueError(f"cannot write user store {path!r}: {exc}") from exc
with os.fdopen(fd, "w") as f:
json.dump({"users": users}, f, indent=2)
f.write("\n")
def auth_enabled(path: str) -> bool:
"""True if at least one user is configured (→ authentication required)."""
return bool(load_users(path))
def list_users(path: str) -> list[str]:
return sorted(load_users(path))
def user_store_state(path: str) -> str:
"""Report the user store state without requiring read access.
Returns one of:
- "absent": no file (auth disabled)
- "empty": file exists but has no users (auth disabled)
- "unreadable": file exists but can't be read (e.g. non-root on a 0600
root-owned store) — auth is almost certainly enabled
- "enabled": file readable and has at least one user (auth enabled)
"""
if not os.path.exists(path):
return "absent"
try:
with open(path) as f:
data = json.load(f)
except (OSError, json.JSONDecodeError):
return "unreadable"
users = data.get("users", {}) if isinstance(data, dict) else {}
if not isinstance(users, dict) or not users:
return "empty"
return "enabled"
def add_user(path: str, username: str, password: str) -> None:
"""Create a new user with a bcrypt-hashed password."""
if not _valid_username(username):
raise ValueError(
f"invalid username {username!r} — use 1–32 chars: letters, digits, '_', '-', '.'"
)
if not password:
raise ValueError("password must not be empty")
if len(password.encode("utf-8")) > MAX_PASSWORD_BYTES:
raise ValueError(f"password too long (max {MAX_PASSWORD_BYTES} bytes)")
users = load_users(path)
if username in users:
raise ValueError(f"user {username!r} already exists")
users[username] = hash_password(password)
save_users(path, users)
def remove_user(path: str, username: str) -> bool:
"""Remove a user. Returns False if the user did not exist."""
users = load_users(path)
if username not in users:
return False
del users[username]
save_users(path, users)
return True
def set_password(path: str, username: str, password: str) -> None:
"""Replace the password of an existing user."""
if not password:
raise ValueError("password must not be empty")
if len(password.encode("utf-8")) > MAX_PASSWORD_BYTES:
raise ValueError(f"password too long (max {MAX_PASSWORD_BYTES} bytes)")
users = load_users(path)
if username not in users:
raise ValueError(f"user {username!r} does not exist")
users[username] = hash_password(password)
save_users(path, users)
# ── Password hashing ──────────────────────────────────────────────────────────
def _pw_bytes(plain: str) -> bytes:
"""Encode a plaintext password, truncating to bcrypt's 72-byte limit.
bcrypt only uses the first 72 bytes and raises ValueError for longer
input; truncating here keeps the check total (no 500 on an over-long
login password) and matches what bcrypt does internally.
"""
return plain.encode("utf-8")[:MAX_PASSWORD_BYTES]
def hash_password(plain: str) -> str:
"""Hash a plaintext password with bcrypt ($2b$)."""
return bcrypt.hashpw(_pw_bytes(plain), bcrypt.gensalt(rounds=BCRYPT_ROUNDS)).decode(
"ascii"
)
def verify_password(plain: str, hashed: str) -> bool:
"""Constant-time check of a plaintext password against a bcrypt hash."""
try:
return bcrypt.checkpw(_pw_bytes(plain), hashed.encode("ascii"))
except (ValueError, TypeError):
return False
_dummy_hash: bytes | None = None
def check_credentials(users: dict[str, str], username: str, password: str) -> bool:
"""Verify username+password. Unknown users burn the same bcrypt time as
known ones so response timing does not reveal which usernames exist."""
global _dummy_hash
hashed = users.get(username)
if hashed is None:
if _dummy_hash is None:
_dummy_hash = bcrypt.hashpw(
b"nvcurve-timing-equalizer", bcrypt.gensalt(rounds=BCRYPT_ROUNDS)
)
bcrypt.checkpw(_pw_bytes(password), _dummy_hash)
return False
return verify_password(password, hashed)
# ── Sessions (in-memory, 24 h) ────────────────────────────────────────────────
_sessions: dict[str, dict[str, Any]] = {} # token -> {"username", "expires_at"}
def create_session(username: str) -> tuple[str, float]:
"""Create a session. Returns (token, expires_at_unix)."""
_purge_expired()
token = secrets.token_urlsafe(32)
expires_at = time.time() + SESSION_TTL_S
_sessions[token] = {"username": username, "expires_at": expires_at}
return token, expires_at
def get_session(token: str | None) -> str | None:
"""Return the username for a valid, unexpired token — else None."""
if not token:
return None
sess = _sessions.get(token)
if sess is None:
return None
if time.time() > sess["expires_at"]:
del _sessions[token]
return None
return sess["username"]
def get_session_expires_at(token: str | None) -> float | None:
if not token:
return None
sess = _sessions.get(token)
if sess is None:
return None
return sess["expires_at"]
def destroy_session(token: str | None) -> None:
if token:
_sessions.pop(token, None)
def _purge_expired() -> None:
now = time.time()
for t in [t for t, s in _sessions.items() if now > s["expires_at"]]:
del _sessions[t]
# ── Brute-force lockout (per client IP) ───────────────────────────────────────
_failures: dict[str, list[float]] = {} # ip -> recent failure timestamps
_lockouts: dict[str, float] = {} # ip -> unix time when lockout lifts
def is_locked_out(ip: str) -> bool:
until = _lockouts.get(ip, 0.0)
if until > time.time():
return True
if until:
del _lockouts[ip]
return False
def record_failure(ip: str) -> None:
now = time.time()
recent = [t for t in _failures.get(ip, []) if now - t < _WINDOW_S]
recent.append(now)
_failures[ip] = recent
if len(recent) >= _MAX_FAILURES:
_lockouts[ip] = now + _LOCKOUT_S
del _failures[ip]
def clear_failures(ip: str) -> None:
_failures.pop(ip, None)
# ── Token extraction ──────────────────────────────────────────────────────────
def extract_token(request: Any) -> str | None:
"""Pull the session token from an Authorization header or the cookie.
Works with both starlette Request and WebSocket objects (both expose
``.headers`` / ``.cookies``).
"""
auth_header = request.headers.get("authorization", "")
if auth_header.lower().startswith("bearer "):
token = auth_header[7:].strip()
if token:
return token
return request.cookies.get(COOKIE_NAME)
+594 -211
View File
File diff suppressed because it is too large. Load diff
+43 -10
View File
@@ -1,8 +1,9 @@
"""HTTP client for communicating with a running nvcurve server."""
import httpx
from typing import Any
import httpx
DEFAULT_BASE = "http://127.0.0.1:8042"
_TIMEOUT = 5.0
@@ -19,14 +20,20 @@ class ApiError(Exception):
class NvCurveClient:
def __init__(self, base: str = DEFAULT_BASE, gpu_index: int = 0):
def __init__(
self, base: str = DEFAULT_BASE, gpu_index: int = 0, token: str | None = None
):
self._base = base.rstrip("/")
self.gpu_index = gpu_index
self.token = token
def _url(self, path: str) -> str:
sep = "&" if "?" in path else "?"
return f"{self._base}{path}{sep}gpu_index={self.gpu_index}"
def _headers(self) -> dict:
return {"Authorization": f"Bearer {self.token}"} if self.token else {}
def _raise(self, r: httpx.Response) -> None:
if r.is_error:
try:
@@ -37,35 +44,58 @@ class NvCurveClient:
def _get(self, path: str) -> Any:
try:
r = httpx.get(self._url(path), timeout=_TIMEOUT)
r = httpx.get(self._url(path), headers=self._headers(), timeout=_TIMEOUT)
except httpx.ConnectError:
raise ServerNotRunning()
raise ServerNotRunning() from None
self._raise(r)
return r.json()
def _post(self, path: str, body: Any = None) -> Any:
try:
r = httpx.post(self._url(path), json=body, timeout=_TIMEOUT)
r = httpx.post(
self._url(path), json=body, headers=self._headers(), timeout=_TIMEOUT
)
except httpx.ConnectError:
raise ServerNotRunning()
raise ServerNotRunning() from None
self._raise(r)
return r.json()
def _delete(self, path: str) -> Any:
try:
r = httpx.delete(self._url(path), timeout=_TIMEOUT)
r = httpx.delete(self._url(path), headers=self._headers(), timeout=_TIMEOUT)
except httpx.ConnectError:
raise ServerNotRunning()
raise ServerNotRunning() from None
self._raise(r)
return r.json()
def ping(self) -> bool:
try:
httpx.get(self._url("/api/gpu"), timeout=1.0)
httpx.get(self._base + "/api/ping", timeout=1.0)
return True
except Exception:
return False
# ── Auth ─────────────────────────────────────────────────────────────────
def auth_status(self) -> dict:
"""Returns {auth_required, authenticated, username, expires_at}."""
return self._get("/api/auth/status")
def login(self, username: str, password: str) -> dict:
"""Authenticate and store the session token for subsequent calls."""
try:
r = httpx.post(
f"{self._base}/api/auth/login",
json={"username": username, "password": password},
timeout=_TIMEOUT,
)
except httpx.ConnectError:
raise ServerNotRunning() from None
self._raise(r)
data = r.json()
self.token = data.get("token")
return data
def gpus(self) -> list:
return self._get("/api/gpus")
@@ -140,7 +170,10 @@ class NvCurveClient:
return self._get("/api/config")
def config_update(self, auto_load_profile: str | None, gpu_index: int = 0) -> dict:
return self._post("/api/config", {"auto_load_profile": auto_load_profile, "gpu_index": gpu_index})
return self._post(
"/api/config",
{"auto_load_profile": auto_load_profile, "gpu_index": gpu_index},
)
# ── Server control ───────────────────────────────────────────────────────
+9 -5
View File
@@ -1,18 +1,17 @@
"""User-configurable settings with sensible defaults."""
import os
from dataclasses import dataclass, field
@dataclass
class Config:
# Safety limits
max_delta_khz: int = 3000_000 # ±3000 MHz hard cap
auto_snapshot: bool = True # Save snapshot before every write
max_snapshots: int = 20 # Maximum snapshots to keep (0 = unlimited)
max_delta_khz: int = 3000_000 # ±3000 MHz hard cap
auto_snapshot: bool = True # Save snapshot before every write
max_snapshots: int = 20 # Maximum snapshots to keep (0 = unlimited)
# Monitoring
poll_interval_s: float = 1.0 # WebSocket monitor poll rate
poll_interval_s: float = 1.0 # WebSocket monitor poll rate
# API server
host: str = "127.0.0.1"
@@ -21,6 +20,11 @@ class Config:
snapshot_dir: str = "/var/cache/nvcurve/snapshots"
profile_dir: str = "/etc/nvcurve/profiles"
# Authentication: path to the bcrypt user store. If the file exists and
# contains at least one user, the API requires login (dual mode). If it is
# absent or empty, the server runs with no authentication.
users_file: str = "/etc/nvcurve/users.json"
# Per-GPU default profiles: applied automatically on server startup.
# Key = stable GPU identifier (UUID string, "pci:{bus_id}", or "idx:{n}" fallback).
# Value = profile name (str).
+67 -28
View File
@@ -15,6 +15,7 @@ Requires root.
"""
import asyncio
import contextlib
import json
import logging
import os
@@ -22,6 +23,8 @@ import signal
import subprocess
import sys
from .config import Config
log = logging.getLogger("nvcurve.daemon")
SOCKET_PATH = "/run/nvcurve-daemon.sock"
@@ -29,27 +32,45 @@ _PERSISTENT_CONFIG_FILE = "/etc/nvcurve/config.json"
# Global server subprocess — only touched from the asyncio event loop.
_server_proc: subprocess.Popen | None = None
_cfg = None # Config instance, set in run()
_cfg: Config | None = None # Config instance, set in run()
# ── Socket command handlers ────────────────────────────────────────────────────
async def _handle_serve_start(host: str, port: int) -> dict:
global _server_proc
if _server_proc is not None and _server_proc.poll() is None:
return {"ok": False, "error": "web server already running", "pid": _server_proc.pid}
return {
"ok": False,
"error": "web server already running",
"pid": _server_proc.pid,
}
cmd = [sys.executable, "-m", "nvcurve", "serve", "start",
"--host", host, "--port", str(port), "--direct"]
cmd = [
sys.executable,
"-m",
"nvcurve",
"serve",
"start",
"--host",
host,
"--port",
str(port),
"--direct",
]
log_path = "/var/log/nvcurve-server.log"
log.info("Starting web server on %s:%d (log: %s)", host, port, log_path)
with open(log_path, "a") as lf:
_server_proc = subprocess.Popen(
cmd,
stdout=lf,
stderr=lf,
env={**os.environ, "PYTHONDONTWRITEBYTECODE": "1"},
)
try:
with open(log_path, "a") as lf:
_server_proc = subprocess.Popen(
cmd,
stdout=lf,
stderr=lf,
env={**os.environ, "PYTHONDONTWRITEBYTECODE": "1"},
)
except OSError as exc:
return {"ok": False, "error": f"cannot open log file {log_path}: {exc}"}
log.info("Web server started (PID %d)", _server_proc.pid)
return {"ok": True, "pid": _server_proc.pid}
@@ -62,9 +83,7 @@ async def _handle_serve_stop() -> dict:
log.info("Stopping web server (PID %d)…", _server_proc.pid)
_server_proc.terminate()
try:
await asyncio.get_running_loop().run_in_executor(
None, _server_proc.wait, 5
)
await asyncio.get_running_loop().run_in_executor(None, _server_proc.wait, 5)
except Exception:
_server_proc.kill()
_server_proc = None
@@ -84,6 +103,8 @@ async def _dispatch(req: dict) -> dict:
if cmd == "ping":
return {"ok": True}
elif cmd == "serve_start":
if _cfg is None:
return {"ok": False, "error": "config not initialized"}
host = req.get("host", _cfg.host)
port = req.get("port", _cfg.port)
return await _handle_serve_start(host, port)
@@ -110,20 +131,17 @@ async def _handle_client(
except Exception as exc:
resp = {"ok": False, "error": str(exc)}
finally:
try:
with contextlib.suppress(Exception):
writer.write(json.dumps(resp).encode() + b"\n")
await writer.drain()
except Exception:
pass
writer.close()
try:
with contextlib.suppress(Exception):
await writer.wait_closed()
except Exception:
pass
# ── Entrypoint ─────────────────────────────────────────────────────────────────
def _load_persistent_config() -> dict:
try:
with open(_PERSISTENT_CONFIG_FILE) as f:
@@ -147,10 +165,16 @@ def run() -> None:
cfg_data = _load_persistent_config()
from .config import Config
_cfg = Config()
for key in ("max_delta_khz", "auto_snapshot", "max_snapshots",
"snapshot_dir", "profile_dir", "host", "port"):
for key in (
"max_delta_khz",
"auto_snapshot",
"max_snapshots",
"snapshot_dir",
"profile_dir",
"host",
"port",
):
if key in cfg_data:
setattr(_cfg, key, cfg_data[key])
@@ -175,15 +199,27 @@ async def _serve_socket(auto_serve: bool = False) -> None:
# Clean up stale socket from a previous (unclean) run.
if os.path.exists(SOCKET_PATH):
os.unlink(SOCKET_PATH)
try:
os.unlink(SOCKET_PATH)
except OSError as exc:
log.warning("Could not remove stale socket %s: %s", SOCKET_PATH, exc)
server = await asyncio.start_unix_server(_handle_client, path=SOCKET_PATH)
os.chmod(SOCKET_PATH, 0o666) # allow non-root CLI to connect
# The socket must be connectable by unprivileged users: the CLI runs as the
# regular user and talks to this root daemon over the socket. 0o666 is
# intentional (standard for /run daemon sockets).
# pi-lens-ignore: S103
os.chmod(
SOCKET_PATH, 0o666
) # nosemgrep: python.lang.security.audit.insecure-file-permissions.insecure-file-permissions
log.info("Daemon listening on %s", SOCKET_PATH)
if auto_serve:
log.info("auto_serve enabled — starting web server on boot")
await _handle_serve_start(_cfg.host, _cfg.port)
if _cfg is None:
log.warning("auto_serve requested but config not initialized")
else:
log.info("auto_serve enabled — starting web server on boot")
await _handle_serve_start(_cfg.host, _cfg.port)
stop_event = asyncio.Event()
loop = asyncio.get_running_loop()
@@ -206,4 +242,7 @@ async def _serve_socket(auto_serve: bool = False) -> None:
_server_proc.kill()
if os.path.exists(SOCKET_PATH):
os.unlink(SOCKET_PATH)
try:
os.unlink(SOCKET_PATH)
except OSError as exc:
log.warning("Could not remove socket %s on shutdown: %s", SOCKET_PATH, exc)
+399 -109
View File
@@ -7,17 +7,36 @@ Requires root (NvAPI needs it).
"""
import asyncio
import json
import logging
import os
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Any
from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect
from fastapi import FastAPI, HTTPException, Request, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel
from . import auth
from .config import Config, default_config
from .hal.gpu import get_gpu, discover_gpus
from .hal.fans import (
get_fan_info,
get_temp,
interpolate_fan_speed,
reset_fan,
set_fan_speed,
validate_curve,
)
from .hal.gpu import discover_gpus, get_gpu
from .hal.limits import (
get_clock_offsets,
get_mem_offset_range,
get_power_limit,
set_clock_offsets,
set_power_limit,
)
from .hal.monitoring import (
get_driver_version,
get_vram_total,
@@ -28,41 +47,29 @@ from .hal.monitoring import (
from .hal.ranges import get_clock_ranges
from .hal.snapshot import (
list_snapshots,
)
from .hal.snapshot import (
restore as snapshot_restore,
)
from .hal.snapshot import (
save as snapshot_save,
)
from .hal.limits import (
get_power_limit,
set_power_limit,
get_clock_offsets,
set_clock_offsets,
get_mem_offset_range,
)
from .hal.fans import (
get_fan_info,
set_fan_speed,
reset_fan,
get_temp,
interpolate_fan_speed,
validate_curve,
)
from .profiles.native import (
ProfileData,
save_profile,
load_profile,
list_profiles,
delete_profile,
rename_profile,
)
from .hal.vfcurve import (
read_clock_offsets,
read_curve,
read_vfp_curve,
reset_offsets,
write_global_offset,
write_offsets,
)
from .safety import validate_write, check_negative_freq_warnings
from .profiles.native import (
ProfileData,
delete_profile,
list_profiles,
load_profile,
rename_profile,
save_profile,
)
from .safety import check_negative_freq_warnings, validate_write
log = logging.getLogger("nvcurve.server")
@@ -71,6 +78,7 @@ def _open_browser_as_user(url: str) -> None:
"""Open URL as the original (non-root) user when running under sudo."""
import os
import subprocess
sudo_user = os.environ.get("SUDO_USER")
if sudo_user and os.geteuid() == 0:
try:
@@ -83,24 +91,29 @@ def _open_browser_as_user(url: str) -> None:
except Exception:
pass
import webbrowser
webbrowser.open(url)
# ── Shared app state ──────────────────────────────────────────────────────────
_state: dict[str, Any] = {
"gpus": {}, # dict[int, dict] mapping gpu_index -> gpu state
"gpus": {}, # dict[int, dict] mapping gpu_index -> gpu state
"config": default_config,
}
def _get_gpu_state(gpu_index: int) -> dict:
if gpu_index not in _state["gpus"]:
from fastapi import HTTPException
raise HTTPException(status_code=404, detail=f"GPU {gpu_index} not found")
return _state["gpus"][gpu_index]
# ── Serialization helpers ─────────────────────────────────────────────────────
def _vfpoint_dict(p) -> dict:
return {
"index": p.index,
@@ -138,8 +151,12 @@ def _sample_dict(s) -> dict:
"pstate_label": f"P{s.pstate}" if s.pstate is not None else None,
"mem_used_bytes": s.mem_used_bytes,
"mem_total_bytes": s.mem_total_bytes,
"mem_used_mib": round(s.mem_used_bytes / (1024 ** 2), 1) if s.mem_used_bytes is not None else None,
"mem_total_mib": round(s.mem_total_bytes / (1024 ** 2), 1) if s.mem_total_bytes is not None else None,
"mem_used_mib": round(s.mem_used_bytes / (1024**2), 1)
if s.mem_used_bytes is not None
else None,
"mem_total_mib": round(s.mem_total_bytes / (1024**2), 1)
if s.mem_total_bytes is not None
else None,
"gpu_util_pct": s.gpu_util_pct,
"mem_util_pct": s.mem_util_pct,
}
@@ -147,6 +164,7 @@ def _sample_dict(s) -> dict:
# ── WebSocket broadcast ───────────────────────────────────────────────────────
async def _broadcast(clients: set, payload: dict) -> None:
"""Send JSON payload to all connected WebSocket clients, evict dead ones."""
dead = set()
@@ -160,6 +178,7 @@ async def _broadcast(clients: set, payload: dict) -> None:
# ── Background monitoring poller ─────────────────────────────────────────────
async def _monitor_poller(gpu_index: int) -> None:
"""Continuously poll GPU state and push to connected monitor WebSocket clients."""
cfg: Config = _state["config"]
@@ -198,6 +217,7 @@ async def _fan_poller(gpu_index: int) -> None:
# ── Lifespan ──────────────────────────────────────────────────────────────────
@asynccontextmanager
async def lifespan(app: FastAPI):
loop = asyncio.get_running_loop()
@@ -208,7 +228,7 @@ async def lifespan(app: FastAPI):
gpu_infos = await loop.run_in_executor(None, discover_gpus)
if not gpu_infos:
log.warning("No GPUs discovered.")
poller_tasks = []
for info in gpu_infos:
@@ -228,18 +248,18 @@ async def lifespan(app: FastAPI):
"fan_poller_task": None,
}
_state["gpus"][idx] = g_state
try:
gpu, name = await loop.run_in_executor(None, get_gpu, idx)
g_state["gpu"] = gpu
g_state["gpu_name"] = name
log.info("GPU %d: %s", idx, name)
# Read initial offsets for reconciliation baseline
offsets, err = await loop.run_in_executor(None, read_clock_offsets, gpu)
if offsets:
g_state["last_offsets"] = offsets
poller_tasks.append(asyncio.create_task(_monitor_poller(idx)))
except Exception as exc:
log.error("Failed to initialize GPU %d: %s", idx, exc)
@@ -260,16 +280,28 @@ async def lifespan(app: FastAPI):
if gpu_idx is None:
log.warning("Auto-load: no GPU found with key %r — skipping", gpu_key)
continue
log.info("Auto-loading profile %r on GPU %d (%s) [compat path]",
profile_name, gpu_idx, gpu_key)
log.info(
"Auto-loading profile %r on GPU %d (%s) [compat path]",
profile_name,
gpu_idx,
gpu_key,
)
try:
await _auto_apply_profile_with_retry(profile_name, gpu_idx)
except FileNotFoundError:
log.warning("Auto-load profile %r not found in %s — skipping GPU %d",
profile_name, cfg.profile_dir, gpu_idx)
log.warning(
"Auto-load profile %r not found in %s — skipping GPU %d",
profile_name,
cfg.profile_dir,
gpu_idx,
)
except Exception as exc:
log.warning("Auto-load profile %r failed on GPU %d: %s — skipping",
profile_name, gpu_idx, exc)
log.warning(
"Auto-load profile %r failed on GPU %d: %s — skipping",
profile_name,
gpu_idx,
exc,
)
# ──────────────────────────────────────────────────────────────────────────
yield # server is running
@@ -294,10 +326,15 @@ async def lifespan(app: FastAPI):
g_state["fan_curve"] = None
try:
await loop.run_in_executor(None, reset_fan, gpu_index)
log.info("GPU %d: restored automatic fan control on shutdown", gpu_index)
log.info(
"GPU %d: restored automatic fan control on shutdown", gpu_index
)
except Exception as exc:
log.warning("GPU %d: failed to restore automatic fan control on shutdown: %s",
gpu_index, exc)
log.warning(
"GPU %d: failed to restore automatic fan control on shutdown: %s",
gpu_index,
exc,
)
await loop.run_in_executor(None, shutdown_nvml)
@@ -314,10 +351,55 @@ app.add_middleware(
)
# ── Authentication middleware (dual mode) ─────────────────────────────────────
# When the user store contains at least one user, every /api/* endpoint requires
# a valid session (cookie or Bearer token). With no users configured, the API is
# open — exactly like before. Static files (the SPA, including the login page)
# and the public auth/ping endpoints are always reachable.
PUBLIC_API_PATHS = {
"/api/ping",
"/api/auth/login",
"/api/auth/status",
"/api/auth/logout",
}
class AuthMiddleware:
"""Require a valid session for /api/* when authentication is enabled."""
def __init__(self, app):
self.app = app
async def __call__(self, scope, receive, send):
if scope["type"] == "http":
path = scope["path"]
if (
path.startswith("/api/")
and path not in PUBLIC_API_PATHS
and scope.get("method") != "OPTIONS"
):
cfg: Config = _state["config"]
if auth.auth_enabled(cfg.users_file):
token = auth.extract_token(Request(scope))
if auth.get_session(token) is None:
response = JSONResponse(
{"detail": "Authentication required"},
status_code=401,
)
await response(scope, receive, send)
return
await self.app(scope, receive, send)
app.add_middleware(AuthMiddleware)
# ── Request models ────────────────────────────────────────────────────────────
class WriteRequest(BaseModel):
deltas: dict[int, int] # {point_index: delta_kHz}
deltas: dict[int, int] # {point_index: delta_kHz}
max_delta_khz: int | None = None # per-request safety limit override
@@ -327,7 +409,7 @@ class GlobalOffsetRequest(BaseModel):
class VerifyRequest(BaseModel):
deltas: dict[int, int] # {point_index: delta_kHz} — pre-expanded by CLI
deltas: dict[int, int] # {point_index: delta_kHz} — pre-expanded by CLI
class SnapshotRestoreRequest(BaseModel):
@@ -365,13 +447,105 @@ class FanSpeedRequest(BaseModel):
fan_pct: int
class LoginRequest(BaseModel):
username: str
password: str
# ── Helper: run blocking HAL call in thread pool ──────────────────────────────
async def _run(fn, *args):
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, fn, *args)
# ── Auth endpoints ────────────────────────────────────────────────────────────
@app.get("/api/ping")
async def api_ping():
"""Public liveness probe (no auth). Used by the CLI to detect a running server."""
return {"ok": True}
@app.get("/api/auth/status")
async def api_auth_status(request: Request):
"""Report whether auth is required and whether this request is authenticated."""
cfg: Config = _state["config"]
required = auth.auth_enabled(cfg.users_file)
token = auth.extract_token(request)
username = auth.get_session(token) if required else None
expires_at = auth.get_session_expires_at(token) if username else None
return {
"auth_required": required,
"authenticated": username is not None,
"username": username,
"expires_at": expires_at,
}
@app.post("/api/auth/login")
async def api_auth_login(req: LoginRequest, request: Request):
"""Authenticate with username+password. Sets a 24-hour session cookie.
The plaintext password is checked against the stored bcrypt hash and never
persisted. On success a random session token is returned (for Bearer use)
and set as an HttpOnly cookie (for browser use).
"""
cfg: Config = _state["config"]
users = auth.load_users(cfg.users_file)
if not users:
raise HTTPException(status_code=404, detail="Authentication is not enabled")
client_ip = request.client.host if request.client else "unknown"
if auth.is_locked_out(client_ip):
raise HTTPException(
status_code=429, detail="Too many failed attempts. Try again later."
)
if auth.check_credentials(users, req.username, req.password):
auth.clear_failures(client_ip)
token, expires_at = auth.create_session(req.username)
response = JSONResponse(
{
"ok": True,
"username": req.username,
"expires_at": expires_at,
"token": token,
}
)
response.set_cookie(
auth.COOKIE_NAME,
token,
max_age=auth.SESSION_TTL_S,
httponly=True,
samesite="lax",
path="/",
)
return response
auth.record_failure(client_ip)
raise HTTPException(status_code=401, detail="Invalid username or password")
@app.post("/api/auth/logout")
async def api_auth_logout(request: Request):
"""End the current session (idempotent)."""
token = auth.extract_token(request)
auth.destroy_session(token)
response = JSONResponse({"ok": True})
response.delete_cookie(auth.COOKIE_NAME, path="/")
return response
@app.get("/api/auth/users")
async def api_auth_users():
"""List configured usernames (requires auth)."""
cfg: Config = _state["config"]
return {"users": auth.list_users(cfg.users_file)}
def _require_gpu(gpu_index: int = 0):
g_state = _get_gpu_state(gpu_index)
gpu = g_state["gpu"]
@@ -382,10 +556,12 @@ def _require_gpu(gpu_index: int = 0):
# ── REST endpoints ────────────────────────────────────────────────────────────
@app.get("/api/gpus")
async def api_gpus():
"""List all discovered GPUs."""
from .hal.gpu import discover_gpus
gpu_infos = await _run(discover_gpus)
return [
{
@@ -397,6 +573,7 @@ async def api_gpus():
for info in gpu_infos
]
@app.get("/api/gpu")
async def api_gpu(gpu_index: int = 0):
"""GPU info: name, driver version, VRAM."""
@@ -408,7 +585,7 @@ async def api_gpu(gpu_index: int = 0):
"index": gpu_index,
"driver_version": driver,
"vram_bytes": vram,
"vram_gib": round(vram / (1024 ** 3), 2) if vram else None,
"vram_gib": round(vram / (1024**3), 2) if vram else None,
}
@@ -433,7 +610,9 @@ async def api_curve_point(point: int, gpu_index: int = 0):
if state is None:
raise HTTPException(status_code=500, detail=f"Failed to read curve: {err}")
if point < 0 or point >= len(state.points):
raise HTTPException(status_code=400, detail=f"Point index must be 0–{len(state.points)-1}")
raise HTTPException(
status_code=400, detail=f"Point index must be 0–{len(state.points) - 1}"
)
return _vfpoint_dict(state.points[point])
@@ -451,6 +630,7 @@ async def api_ranges(gpu_index: int = 0):
async def api_voltage(gpu_index: int = 0):
"""Current GPU core voltage."""
from .hal.monitoring import read_voltage
gpu, g_state = _require_gpu(gpu_index)
voltage_uv, err = await _run(read_voltage, gpu)
if voltage_uv is None:
@@ -492,6 +672,7 @@ def _persist_config_field(key: str, value) -> None:
"""
import json as _json
import os as _os
config_path = "/etc/nvcurve/config.json"
if not _os.path.exists(config_path):
return
@@ -589,6 +770,7 @@ async def _auto_apply_profile_with_retry(
Logs a warning and gives up after max_retries failed attempts.
"""
import os as _os
cfg: Config = _state["config"]
g_state = _get_gpu_state(gpu_index)
gpu = g_state["gpu"]
@@ -598,43 +780,69 @@ async def _auto_apply_profile_with_retry(
profile = await _run(load_profile, filepath) # raises FileNotFoundError if missing
expected: dict[int, int] = (
{int(k): v for k, v in profile.curve_deltas.items()} if profile.curve_deltas else {}
{int(k): v for k, v in profile.curve_deltas.items()}
if profile.curve_deltas
else {}
)
for attempt in range(max_retries):
errs = await _apply_profile(name, gpu_index)
if errs:
log.warning("Auto-load attempt %d/%d had errors: %s",
attempt + 1, max_retries, "; ".join(errs))
log.warning(
"Auto-load attempt %d/%d had errors: %s",
attempt + 1,
max_retries,
"; ".join(errs),
)
elif expected:
offsets, err = await _run(read_clock_offsets, gpu)
if offsets is None:
log.warning("Auto-load attempt %d/%d: read-back failed: %s",
attempt + 1, max_retries, err)
log.warning(
"Auto-load attempt %d/%d: read-back failed: %s",
attempt + 1,
max_retries,
err,
)
else:
mismatches = [
f"pt{idx}: expected {val/1000:+.0f}MHz got {offsets[idx]/1000:+.0f}MHz"
f"pt{idx}: expected {val / 1000:+.0f}MHz got {offsets[idx] / 1000:+.0f}MHz"
for idx, val in expected.items()
if idx < len(offsets) and offsets[idx] != val
]
if not mismatches:
log.info("Auto-load profile %r verified on GPU %d (attempt %d/%d)",
name, gpu_index, attempt + 1, max_retries)
log.info(
"Auto-load profile %r verified on GPU %d (attempt %d/%d)",
name,
gpu_index,
attempt + 1,
max_retries,
)
return
log.warning("Auto-load attempt %d/%d: read-back mismatch — %s",
attempt + 1, max_retries, "; ".join(mismatches))
log.warning(
"Auto-load attempt %d/%d: read-back mismatch — %s",
attempt + 1,
max_retries,
"; ".join(mismatches),
)
else:
log.info("Auto-load profile %r applied on GPU %d (attempt %d/%d)",
name, gpu_index, attempt + 1, max_retries)
log.info(
"Auto-load profile %r applied on GPU %d (attempt %d/%d)",
name,
gpu_index,
attempt + 1,
max_retries,
)
return
if attempt < max_retries - 1:
delay = 2 ** attempt # 1 s, 2 s, 4 s
delay = 2**attempt # 1 s, 2 s, 4 s
log.info("Retrying auto-load in %ds…", delay)
await asyncio.sleep(delay)
log.warning("Auto-load profile %r failed after %d attempts — giving up", name, max_retries)
log.warning(
"Auto-load profile %r failed after %d attempts — giving up", name, max_retries
)
async def _apply_profile(name: str, gpu_index: int = 0) -> list[str]:
@@ -645,6 +853,7 @@ async def _apply_profile(name: str, gpu_index: int = 0) -> list[str]:
Sets g_state["active_profile"] on full success.
"""
import os as _os
g_state = _get_gpu_state(gpu_index)
gpu = g_state["gpu"]
cfg: Config = _state["config"]
@@ -677,7 +886,13 @@ async def _apply_profile(name: str, gpu_index: int = 0) -> list[str]:
errs.append("Curve: " + "; ".join(errors))
else:
if cfg.auto_snapshot:
await _run(snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots)
await _run(
snapshot_save,
gpu,
g_state["gpu_name"],
cfg.snapshot_dir,
cfg.max_snapshots,
)
ret, desc = await _run(write_offsets, gpu, deltas)
if ret != 0:
errs.append(f"Curve write failed ({ret}): {desc}")
@@ -748,7 +963,9 @@ async def api_profile_delete(name: str):
g_state["active_profile"] = None
changed = any(v == name for v in cfg.auto_load_profiles.values())
if changed:
cfg.auto_load_profiles = {k: v for k, v in cfg.auto_load_profiles.items() if v != name}
cfg.auto_load_profiles = {
k: v for k, v in cfg.auto_load_profiles.items() if v != name
}
_persist_auto_load_profiles(cfg.auto_load_profiles)
return {"ok": True}
@@ -805,8 +1022,8 @@ async def api_limits(gpu_index: int = 0):
mem_off_range = await _run(get_mem_offset_range, gpu_index)
return {
**power,
**offsets, # gpc_offset_mhz, mem_offset_mhz
**mem_off_range, # min_mem_offset_mhz, max_mem_offset_mhz
**offsets, # gpc_offset_mhz, mem_offset_mhz
**mem_off_range, # min_mem_offset_mhz, max_mem_offset_mhz
}
@@ -905,6 +1122,7 @@ async def api_limits_reset(gpu_index: int = 0):
# ── Fan endpoints ──────────────────────────────────────────────────────────────
@app.get("/api/fans")
async def api_fans(gpu_index: int = 0):
"""Current fan state: fan %, curve, and whether curve control is active."""
@@ -935,12 +1153,18 @@ async def api_fans_update(req: FanCurveRequest, gpu_index: int = 0):
# set to an inappropriate speed.
if curve_data:
sample = await _run(poll, g_state["gpu"], gpu_index)
test_temp = sample.temp_c if sample and sample.temp_c is not None else curve_data[0]["temp_c"]
test_temp = (
sample.temp_c
if sample and sample.temp_c is not None
else curve_data[0]["temp_c"]
)
target = interpolate_fan_speed(curve_data, test_temp)
if target is not None:
fan_ok, fan_msg = await _run(set_fan_speed, gpu_index, target)
if not fan_ok:
raise HTTPException(status_code=500, detail=f"Fan control not available: {fan_msg}")
raise HTTPException(
status_code=500, detail=f"Fan control not available: {fan_msg}"
)
# Stop existing poller if running
if g_state.get("fan_poller_task"):
@@ -994,6 +1218,7 @@ async def api_fans_speed(req: FanSpeedRequest, gpu_index: int = 0):
# ── Write endpoints ────────────────────────────────────────────────────────────
async def _reconcile_check(gpu_index: int) -> dict | None:
"""Re-read current offsets and return a warning dict if they differ from our last known state.
@@ -1035,7 +1260,9 @@ async def api_curve_write(req: WriteRequest, gpu_index: int = 0):
vfp_state, _ = await _run(read_curve, gpu, g_state["gpu_name"])
effective_limit = req.max_delta_khz if req.max_delta_khz is not None else cfg.max_delta_khz
effective_limit = (
req.max_delta_khz if req.max_delta_khz is not None else cfg.max_delta_khz
)
errors = validate_write(req.deltas, effective_limit)
if errors:
raise HTTPException(status_code=400, detail={"errors": errors})
@@ -1052,7 +1279,13 @@ async def api_curve_write(req: WriteRequest, gpu_index: int = 0):
warning = await _reconcile_check(gpu_index)
if cfg.auto_snapshot:
await _run(snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots)
await _run(
snapshot_save,
gpu,
g_state["gpu_name"],
cfg.snapshot_dir,
cfg.max_snapshots,
)
ret, desc = await _run(write_offsets, gpu, req.deltas)
if ret != 0:
@@ -1081,7 +1314,9 @@ async def api_curve_write_global(req: GlobalOffsetRequest, gpu_index: int = 0):
raise HTTPException(status_code=500, detail="Failed to read curve")
all_deltas = {p.index: req.delta_khz for p in vfp_state.points if p.domain == "gpu"}
effective_limit = req.max_delta_khz if req.max_delta_khz is not None else cfg.max_delta_khz
effective_limit = (
req.max_delta_khz if req.max_delta_khz is not None else cfg.max_delta_khz
)
errors = validate_write(all_deltas, effective_limit)
if errors:
raise HTTPException(status_code=400, detail={"errors": errors})
@@ -1097,7 +1332,13 @@ async def api_curve_write_global(req: GlobalOffsetRequest, gpu_index: int = 0):
warning = await _reconcile_check(gpu_index)
if cfg.auto_snapshot:
await _run(snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots)
await _run(
snapshot_save,
gpu,
g_state["gpu_name"],
cfg.snapshot_dir,
cfg.max_snapshots,
)
ret, desc = await _run(write_global_offset, gpu, req.delta_khz)
if ret != 0:
@@ -1124,7 +1365,13 @@ async def api_curve_reset(gpu_index: int = 0):
warning = await _reconcile_check(gpu_index)
if cfg.auto_snapshot:
await _run(snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots)
await _run(
snapshot_save,
gpu,
g_state["gpu_name"],
cfg.snapshot_dir,
cfg.max_snapshots,
)
ret, desc = await _run(reset_offsets, gpu)
if ret != 0:
@@ -1151,10 +1398,14 @@ async def api_curve_verify(req: VerifyRequest, gpu_index: int = 0):
before_offsets, err = await _run(read_clock_offsets, gpu)
if before_offsets is None:
raise HTTPException(status_code=500, detail=f"Failed to read current state: {err}")
raise HTTPException(
status_code=500, detail=f"Failed to read current state: {err}"
)
# Always snapshot before verify — it's a testing operation
await _run(snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots)
await _run(
snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots
)
async with g_state["write_lock"]:
ret, desc = await _run(write_offsets, gpu, req.deltas)
@@ -1165,7 +1416,9 @@ async def api_curve_verify(req: VerifyRequest, gpu_index: int = 0):
after_offsets, err = await _run(read_clock_offsets, gpu)
if after_offsets is None:
raise HTTPException(status_code=500, detail=f"Verification read failed: {err}")
raise HTTPException(
status_code=500, detail=f"Verification read failed: {err}"
)
g_state["active_profile"] = None
await _update_offsets_and_broadcast(gpu_index)
@@ -1177,12 +1430,14 @@ async def api_curve_verify(req: VerifyRequest, gpu_index: int = 0):
match = actual == expected
if not match:
all_matched = False
points_result.append({
"point": point,
"expected_khz": expected,
"actual_khz": actual,
"match": match,
})
points_result.append(
{
"point": point,
"expected_khz": expected,
"actual_khz": actual,
"match": match,
}
)
collateral = [
{"point": i, "before_khz": before_offsets[i], "after_khz": after_offsets[i]}
@@ -1206,6 +1461,7 @@ async def api_shutdown():
"""Gracefully shut down the server process."""
import os
import signal
loop = asyncio.get_running_loop()
loop.call_later(0.1, lambda: os.kill(os.getpid(), signal.SIGTERM))
return {"ok": True}
@@ -1216,7 +1472,9 @@ async def api_snapshot_save(gpu_index: int = 0):
"""Save a ClockBoostTable snapshot."""
gpu, g_state = _require_gpu(gpu_index)
cfg: Config = _state["config"]
path = await _run(snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots)
path = await _run(
snapshot_save, gpu, g_state["gpu_name"], cfg.snapshot_dir, cfg.max_snapshots
)
if path is None:
raise HTTPException(status_code=500, detail="Failed to save snapshot")
return {"ok": True, "filepath": path}
@@ -1241,9 +1499,27 @@ async def api_snapshot_restore(req: SnapshotRestoreRequest, gpu_index: int = 0):
# ── WebSocket endpoints ───────────────────────────────────────────────────────
def _ws_authenticated(ws: WebSocket) -> bool:
"""True if the WebSocket connection is allowed (auth disabled or valid session).
The session token is read from the Authorization header or the cookie
(browsers send the cookie automatically on the WS handshake). The token is
deliberately NOT accepted via a query string, since uvicorn's access log
records the full path including the query string.
"""
cfg: Config = _state["config"]
if not auth.auth_enabled(cfg.users_file):
return True
return auth.get_session(auth.extract_token(ws)) is not None
@app.websocket("/ws/monitor")
async def ws_monitor(ws: WebSocket):
"""Stream MonitoringSample at poll_interval_s. Clients receive JSON objects."""
if not _ws_authenticated(ws):
await ws.close(code=1008)
return
await ws.accept()
try:
data = await ws.receive_json()
@@ -1274,7 +1550,7 @@ async def ws_monitor(ws: WebSocket):
except WebSocketDisconnect:
pass
except Exception:
pass
log.debug("monitor ws client error", exc_info=True)
finally:
g_state["monitor_clients"].discard(ws)
@@ -1282,6 +1558,9 @@ async def ws_monitor(ws: WebSocket):
@app.websocket("/ws/curve")
async def ws_curve(ws: WebSocket):
"""Push CurveState whenever the curve changes (after writes)."""
if not _ws_authenticated(ws):
await ws.close(code=1008)
return
await ws.accept()
try:
data = await ws.receive_json()
@@ -1313,47 +1592,54 @@ async def ws_curve(ws: WebSocket):
except WebSocketDisconnect:
pass
except Exception:
pass
log.debug("curve ws client error", exc_info=True)
finally:
g_state["curve_clients"].discard(ws)
# ── Frontend SPA ──────────────────────────────────────────────────────────────
import os
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
from pathlib import Path
# When set, suppresses the auto-open browser behaviour so the dev can open
# the Vite dev server (pnpm dev) manually instead.
_DEV_PORT = os.environ.get("NVCURVE_DEV_PORT")
# Robust asset resolution using importlib.resources
try:
from importlib.resources import files as _resource_files
# In a packaged installation, frontend/dist is inside the package
_dist_dir = _resource_files("nvcurve") / "frontend" / "dist"
# Fallback for local development where frontend/dist might be at project root
if not _dist_dir.is_dir():
_here = Path(__file__).parent
_dist_dir = _here.parent / "frontend" / "dist"
except (ImportError, TypeError):
# Legacy fallback for older Python or environments without importlib.resources.files
_here = Path(__file__).parent
_dist_dir = _here / "frontend" / "dist"
if not _dist_dir.is_dir():
_dist_dir = _here.parent / "frontend" / "dist"
_dist_dir = str(_dist_dir)
def _resolve_dist_dir() -> str:
"""Resolve the frontend dist directory to a string path.
Prefers the packaged location (importlib.resources), then falls back to
the project-root layout used during local development.
"""
try:
from importlib.resources import files as _resource_files
candidate = _resource_files("nvcurve") / "frontend" / "dist"
if candidate.is_dir():
return str(candidate)
except (ImportError, TypeError):
pass
here = Path(__file__).parent
for base in (here, here.parent):
dist = base / "frontend" / "dist"
if dist.is_dir():
return str(dist)
return str(here / "frontend" / "dist")
_dist_dir = _resolve_dist_dir()
if os.path.isdir(os.path.join(_dist_dir, "assets")):
app.mount("/assets", StaticFiles(directory=os.path.join(_dist_dir, "assets")), name="assets")
app.mount(
"/assets",
StaticFiles(directory=os.path.join(_dist_dir, "assets")),
name="assets",
)
@app.get("/{catchall:path}")
async def serve_spa(catchall: str):
if catchall.startswith("api/") or catchall.startswith("ws/"):
if catchall.startswith(("api/", "ws/")):
raise HTTPException(status_code=404, detail="Not Found")
if not os.path.isdir(_dist_dir):
@@ -1372,6 +1658,7 @@ async def serve_spa(catchall: str):
# ── Factory for configured app ────────────────────────────────────────────────
def create_app(config: Config = default_config) -> FastAPI:
"""Create a server app with a custom config (e.g. different gpu_index)."""
_state["config"] = config
@@ -1388,6 +1675,7 @@ def run(
"""Start the uvicorn server. Blocking."""
import socket
import threading
import uvicorn
_state["config"] = config
@@ -1404,7 +1692,9 @@ def run(
s.bind((host, port))
except OSError:
print(f"Error: port {port} is already in use.")
print(f"Use --port N to specify a different port, or free port {port} first.")
print(
f"Use --port N to specify a different port, or free port {port} first."
)
return
url = f"http://{host}:{port}"