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:
1 parent
af23a10f25
commit
bbd692ea2e
16 files changed
+2085
-483
No files matched your search
+399
-109
@@ -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}"
|
||||
|
||||
Reference in new issue
Block a user