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

+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}"