Adds an experimental power-cap mode using the undocumented RM ioctl interface (based on panchovix's LACT PR #1205) to set power limits below the VBIOS minimum (down to 30 W). - hal/rm_power.py: RM ioctl power-cap read/write/reset + runtime probe - limits.py: power_cap_mode (nvml/ioctl) with support detection - config.py: persist power_cap_mode per GPU - profiles: record/apply power_cap_mode - server.py: POST /api/limits validates ioctl support (409 on failure) - cli.py: profile save falls back to persisted mode - client.py: power_cap_mode in Limits - frontend: toggle + warning with panchovix attribution (LACT #1205) - tests: test_rm_power.py (unit) + integration coverage - Makefile: add test_rm_power.py to make test Also includes automated linter reformatting (prettier, ruff, shellcheck, isort, markdownlint) that the linter would apply anyway.
2004 lines
69 KiB
Python
2004 lines
69 KiB
Python
"""FastAPI API server — REST + WebSocket.
|
||
|
||
Run via: nvcurve serve [--host 127.0.0.1 --port 8042]
|
||
Or: uvicorn nvcurve.server:app
|
||
|
||
Requires root (NvAPI needs it).
|
||
"""
|
||
|
||
import asyncio
|
||
import logging
|
||
import os
|
||
from contextlib import asynccontextmanager, suppress
|
||
from pathlib import Path
|
||
from typing import Any, Protocol
|
||
|
||
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, tls_enabled
|
||
from .hal.dashboard import get_dashboard_info
|
||
from .hal.fans import (
|
||
get_fan_info,
|
||
get_num_fans,
|
||
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,
|
||
init_nvml,
|
||
poll,
|
||
shutdown_nvml,
|
||
throttle_reasons_label,
|
||
)
|
||
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.vfcurve import (
|
||
read_clock_offsets,
|
||
read_curve,
|
||
reset_offsets,
|
||
write_global_offset,
|
||
write_offsets,
|
||
)
|
||
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")
|
||
|
||
|
||
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:
|
||
subprocess.Popen(
|
||
["runuser", "-u", sudo_user, "--", "xdg-open", url],
|
||
stdout=subprocess.DEVNULL,
|
||
stderr=subprocess.DEVNULL,
|
||
)
|
||
return
|
||
except Exception as exc:
|
||
log.debug("runuser xdg-open failed, falling back to webbrowser: %s", exc)
|
||
import webbrowser
|
||
|
||
webbrowser.open(url)
|
||
|
||
|
||
# ── Shared app state ──────────────────────────────────────────────────────────
|
||
|
||
_state: dict[str, Any] = {
|
||
"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,
|
||
"freq_khz": p.freq_khz,
|
||
"freq_mhz": p.freq_mhz,
|
||
"volt_uv": p.volt_uv,
|
||
"volt_mv": p.volt_mv,
|
||
"delta_khz": p.delta_khz,
|
||
"delta_mhz": p.delta_mhz,
|
||
"effective_freq_khz": p.effective_freq_khz,
|
||
"effective_freq_mhz": p.effective_freq_mhz,
|
||
"domain": p.domain,
|
||
}
|
||
|
||
|
||
def _curve_state_dict(state) -> dict:
|
||
return {
|
||
"gpu_name": state.gpu_name,
|
||
"timestamp": state.timestamp,
|
||
"points": [_vfpoint_dict(p) for p in state.points],
|
||
}
|
||
|
||
|
||
def _sample_dict(s) -> dict:
|
||
return {
|
||
"timestamp": s.timestamp,
|
||
"voltage_uv": s.voltage_uv,
|
||
"voltage_mv": s.voltage_uv / 1000.0 if s.voltage_uv is not None else None,
|
||
"clock_mhz": s.clock_mhz,
|
||
"mem_clock_mhz": s.mem_clock_mhz,
|
||
"temp_c": s.temp_c,
|
||
"power_w": s.power_w,
|
||
"fan_pct": s.fan_pct,
|
||
"fans": s.fans,
|
||
"pstate": s.pstate,
|
||
"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,
|
||
"gpu_util_pct": s.gpu_util_pct,
|
||
"mem_util_pct": s.mem_util_pct,
|
||
"throttle_reasons": s.throttle_reasons,
|
||
"throttle_reasons_label": throttle_reasons_label(s.throttle_reasons),
|
||
"pcie_link_width": s.pcie_link_width,
|
||
"pcie_link_generation": s.pcie_link_generation,
|
||
"mem_temp_c": s.mem_temp_c,
|
||
}
|
||
|
||
|
||
def _int_key_deltas(deltas: dict) -> dict[int, int]:
|
||
"""Convert string-keyed deltas (from JSON) to int-keyed.
|
||
|
||
Raises ValueError if any key is not a valid integer (corrupted profile),
|
||
so a bad profile fails closed rather than partially applying to hardware.
|
||
"""
|
||
out: dict[int, int] = {}
|
||
for k, v in deltas.items():
|
||
try:
|
||
out[int(k)] = v
|
||
except (TypeError, ValueError) as exc:
|
||
raise ValueError(f"Invalid curve point index in profile: {k!r}") from exc
|
||
return out
|
||
|
||
|
||
# ── WebSocket broadcast ───────────────────────────────────────────────────────
|
||
|
||
|
||
async def _broadcast(clients: set, payload: dict) -> None:
|
||
"""Send JSON payload to all connected WebSocket clients, evict dead ones."""
|
||
dead = set()
|
||
for ws in list(clients):
|
||
try:
|
||
await ws.send_json(payload)
|
||
except Exception:
|
||
dead.add(ws)
|
||
clients -= dead
|
||
|
||
|
||
# ── 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"]
|
||
while True:
|
||
try:
|
||
g_state = _state["gpus"].get(gpu_index)
|
||
if g_state and g_state["gpu"] is not None and g_state["monitor_clients"]:
|
||
loop = asyncio.get_running_loop()
|
||
sample = await loop.run_in_executor(
|
||
None, poll, g_state["gpu"], gpu_index
|
||
)
|
||
await _broadcast(g_state["monitor_clients"], _sample_dict(sample))
|
||
except Exception as exc:
|
||
log.warning("Monitor poller error for GPU %d: %s", gpu_index, exc)
|
||
await asyncio.sleep(cfg.poll_interval_s)
|
||
|
||
|
||
async def _fan_poller(gpu_index: int) -> None:
|
||
"""Continuously read GPU temp, interpolate fan speed from active curve, and apply it to the target fans."""
|
||
last_write_error: str | None = None
|
||
while True:
|
||
try:
|
||
g_state = _state["gpus"].get(gpu_index)
|
||
if g_state and g_state.get("fan_curve_active") and g_state.get("fan_curve"):
|
||
temp = await _run(get_temp, gpu_index)
|
||
if temp is not None:
|
||
curve = g_state["fan_curve"]
|
||
target = interpolate_fan_speed(curve, temp)
|
||
if target is not None:
|
||
ok, msg = await _run(
|
||
set_fan_speed,
|
||
gpu_index,
|
||
target,
|
||
g_state.get("fan_targets"),
|
||
)
|
||
if not ok:
|
||
# Log once per distinct failure so a stuck target
|
||
# list doesn't spam a warning every 2 s tick.
|
||
if msg != last_write_error:
|
||
log.warning(
|
||
"Fan write failed for GPU %d (targets=%s): %s",
|
||
gpu_index,
|
||
g_state.get("fan_targets"),
|
||
msg,
|
||
)
|
||
last_write_error = msg
|
||
else:
|
||
last_write_error = None
|
||
except asyncio.CancelledError:
|
||
return
|
||
except Exception as exc:
|
||
log.warning("Fan poller error for GPU %d: %s", gpu_index, exc)
|
||
await asyncio.sleep(2.0)
|
||
|
||
|
||
async def _activate_fan_curve(
|
||
gpu_index: int, curve: list, fans: list[int] | None = None
|
||
) -> None:
|
||
"""Set the active fan curve, (re)start the poller, and persist it.
|
||
|
||
fans=None targets all fans on the device; a list targets the given
|
||
fan indices. Persistence (config.json) is what makes the curve survive
|
||
server restarts: fan control is volatile, so the driver reverts to
|
||
automatic mode on reboot and the saved curve is re-applied at the next
|
||
server start.
|
||
"""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
|
||
# Validate explicit fan targets against the hardware. Stale indices
|
||
# (e.g. a profile saved on a 2-fan GPU applied to a 1-fan GPU, or a
|
||
# persisted entry restored after a hardware change) would otherwise
|
||
# make the poller fail silently on every tick.
|
||
if fans is not None:
|
||
num_fans = await _run(get_num_fans, gpu_index)
|
||
if not fans or (num_fans > 0 and any(f < 0 or f >= num_fans for f in fans)):
|
||
log.warning(
|
||
"Fan targets %s invalid for GPU %d (%d fan(s)); falling back to all fans",
|
||
fans,
|
||
gpu_index,
|
||
num_fans,
|
||
)
|
||
fans = None
|
||
|
||
# Stop existing poller if running
|
||
if g_state.get("fan_poller_task"):
|
||
g_state["fan_poller_task"].cancel()
|
||
with suppress(asyncio.CancelledError):
|
||
await g_state["fan_poller_task"]
|
||
|
||
g_state["fan_curve"] = curve
|
||
g_state["fan_curve_active"] = True
|
||
g_state["fan_targets"] = fans
|
||
g_state["fan_poller_task"] = asyncio.create_task(_fan_poller(gpu_index))
|
||
|
||
cfg: Config = _state["config"]
|
||
cfg.fan_curves[_gpu_stable_key(gpu_index)] = {"curve": curve, "fans": fans}
|
||
_persist_fan_curves(cfg.fan_curves)
|
||
|
||
|
||
async def _deactivate_fan_curve(gpu_index: int, reset_hardware: bool = True) -> None:
|
||
"""Clear the active fan curve, stop the poller, and clear its persistence.
|
||
|
||
When reset_hardware is True the GPU is returned to automatic fan control.
|
||
"""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
|
||
if g_state.get("fan_poller_task"):
|
||
g_state["fan_poller_task"].cancel()
|
||
with suppress(asyncio.CancelledError):
|
||
await g_state["fan_poller_task"]
|
||
g_state["fan_poller_task"] = None
|
||
|
||
g_state["fan_curve"] = None
|
||
g_state["fan_curve_active"] = False
|
||
g_state["fan_targets"] = None
|
||
|
||
if reset_hardware:
|
||
ok, msg = await _run(reset_fan, gpu_index)
|
||
if not ok:
|
||
log.warning("Fan reset warning: %s", msg)
|
||
|
||
cfg: Config = _state["config"]
|
||
key = _gpu_stable_key(gpu_index)
|
||
if key in cfg.fan_curves:
|
||
del cfg.fan_curves[key]
|
||
_persist_fan_curves(cfg.fan_curves)
|
||
|
||
|
||
# ── Lifespan ──────────────────────────────────────────────────────────────────
|
||
|
||
|
||
@asynccontextmanager
|
||
async def lifespan(app: FastAPI):
|
||
loop = asyncio.get_running_loop()
|
||
|
||
# Initialize NVML (best-effort)
|
||
await loop.run_in_executor(None, init_nvml)
|
||
|
||
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:
|
||
idx = info.index
|
||
g_state = {
|
||
"gpu": None,
|
||
"gpu_name": info.name,
|
||
"uuid": info.uuid,
|
||
"pci_bus_id": info.pci_bus_id,
|
||
"write_lock": asyncio.Lock(),
|
||
"last_offsets": None,
|
||
"active_profile": None,
|
||
"monitor_clients": set(),
|
||
"curve_clients": set(),
|
||
"fan_curve": None,
|
||
"fan_curve_active": False,
|
||
"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)
|
||
|
||
# ── Backward Compatibility Bridge ──────────────────────────────────────────
|
||
# NOTE: This auto-load path is for users running the server directly (e.g.
|
||
# via an old systemd unit file that lacks the new daemon mode).
|
||
# In the future, this will be removed and auto-loading will be the
|
||
# responsibility of daemon.py only.
|
||
cfg: Config = _state["config"]
|
||
if cfg.auto_load_profiles:
|
||
# Build a reverse map: stable_key → current gpu_index
|
||
key_to_idx = {_gpu_stable_key(idx): idx for idx in _state["gpus"]}
|
||
for gpu_key, profile_name in cfg.auto_load_profiles.items():
|
||
if not profile_name:
|
||
continue
|
||
gpu_idx = key_to_idx.get(gpu_key)
|
||
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,
|
||
)
|
||
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,
|
||
)
|
||
except Exception as exc:
|
||
log.warning(
|
||
"Auto-load profile %r failed on GPU %d: %s — skipping",
|
||
profile_name,
|
||
gpu_idx,
|
||
exc,
|
||
)
|
||
# ──────────────────────────────────────────────────────────────────────────
|
||
|
||
# ── Restore persisted fan curves ──────────────────────────────────────────
|
||
# Fan control is volatile: the driver reverts to automatic mode on reboot,
|
||
# so a curve applied via the UI is persisted in config.json and re-applied
|
||
# here at startup. Runs after the auto-load profile path so the user's
|
||
# explicit fan curve setting takes precedence.
|
||
for gpu_index, g_state in _state["gpus"].items():
|
||
if g_state["gpu"] is None:
|
||
continue
|
||
if g_state.get("fan_curve_active"):
|
||
continue # already activated by the auto-load profile path
|
||
key = _gpu_stable_key(gpu_index)
|
||
entry = cfg.fan_curves.get(key)
|
||
if not entry:
|
||
continue
|
||
# Migrate the legacy format (bare curve list) to the current
|
||
# {"curve": ..., "fans": ...} shape; legacy entries targeted all fans.
|
||
if isinstance(entry, list):
|
||
entry = {"curve": entry, "fans": None}
|
||
curve = entry.get("curve") if isinstance(entry, dict) else None
|
||
fans = entry.get("fans") if isinstance(entry, dict) else None
|
||
if not curve:
|
||
continue
|
||
ok, msg = validate_curve(curve)
|
||
if not ok:
|
||
log.warning("Skipping persisted fan curve for GPU %d: %s", gpu_index, msg)
|
||
continue
|
||
log.info("Restoring persisted fan curve on GPU %d (%s)", gpu_index, key)
|
||
try:
|
||
await _activate_fan_curve(gpu_index, curve, fans)
|
||
except Exception as exc:
|
||
log.warning(
|
||
"Failed to restore persisted fan curve on GPU %d: %s",
|
||
gpu_index,
|
||
exc,
|
||
)
|
||
# ──────────────────────────────────────────────────────────────────────────
|
||
|
||
yield # server is running
|
||
|
||
for task in poller_tasks:
|
||
task.cancel()
|
||
for task in poller_tasks:
|
||
with suppress(asyncio.CancelledError):
|
||
await task
|
||
|
||
for gpu_index, g_state in _state["gpus"].items():
|
||
if g_state.get("fan_poller_task"):
|
||
g_state["fan_poller_task"].cancel()
|
||
with suppress(asyncio.CancelledError):
|
||
await g_state["fan_poller_task"]
|
||
if g_state.get("fan_curve_active"):
|
||
g_state["fan_curve_active"] = False
|
||
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
|
||
)
|
||
except Exception as 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)
|
||
|
||
|
||
# ── App ───────────────────────────────────────────────────────────────────────
|
||
|
||
app = FastAPI(title="nvcurve", version="0.5.0", lifespan=lifespan)
|
||
|
||
app.add_middleware(
|
||
CORSMiddleware,
|
||
# The SPA is served same-origin by this server, so CORS only matters for
|
||
# local development (e.g. the Vite dev server). Restrict to localhost
|
||
# origins rather than a wildcard.
|
||
allow_origin_regex=r"https?://(localhost|127\.0\.0\.1)(:\d+)?$",
|
||
allow_methods=["*"],
|
||
allow_headers=["*"],
|
||
)
|
||
|
||
|
||
# ── 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}
|
||
|
||
|
||
class GlobalOffsetRequest(BaseModel):
|
||
delta_khz: int
|
||
|
||
|
||
class VerifyRequest(BaseModel):
|
||
deltas: dict[int, int] # {point_index: delta_kHz} — pre-expanded by CLI
|
||
|
||
|
||
class SnapshotRestoreRequest(BaseModel):
|
||
filepath: str | None = None
|
||
|
||
|
||
class LimitsRequest(BaseModel):
|
||
power_limit_w: int | None = None
|
||
mem_offset_mhz: int | None = None
|
||
# "nvml" (default) or "ioctl" (experimental RM power control).
|
||
power_cap_mode: str | None = None
|
||
|
||
|
||
class ProfileSaveRequest(BaseModel):
|
||
name: str
|
||
|
||
|
||
class ProfileRenameRequest(BaseModel):
|
||
new_name: str
|
||
|
||
|
||
class ConfigUpdateRequest(BaseModel):
|
||
auto_load_profile: str | None = None
|
||
gpu_index: int = 0
|
||
|
||
|
||
class FanCurvePoint(BaseModel):
|
||
temp_c: int
|
||
fan_pct: int
|
||
|
||
|
||
class FanCurveRequest(BaseModel):
|
||
curve: list[FanCurvePoint]
|
||
# Fan indices to control (0-based); None = all fans on the device.
|
||
fans: list[int] | None = None
|
||
|
||
|
||
class FanSpeedRequest(BaseModel):
|
||
fan_pct: int
|
||
# Specific fan index to set (0-based); None = all fans on the device.
|
||
fan: int | None = None
|
||
|
||
|
||
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 = _client_ip(request, cfg.trusted_proxies)
|
||
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",
|
||
secure=tls_enabled(cfg),
|
||
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"]
|
||
if gpu is None:
|
||
raise HTTPException(status_code=503, detail=f"GPU {gpu_index} not initialized")
|
||
return gpu, g_state
|
||
|
||
|
||
class _ClientIpSource(Protocol):
|
||
"""Structural type for the request objects _client_ip accepts.
|
||
|
||
Both Starlette's Request and WebSocket expose these; tests may pass
|
||
lightweight duck types.
|
||
"""
|
||
|
||
client: Any
|
||
headers: Any
|
||
|
||
|
||
def _client_ip(request: _ClientIpSource, trusted_proxies: list[str]) -> str:
|
||
"""Best-effort client IP for the login lockout.
|
||
|
||
When the direct peer is a configured trusted proxy (e.g. a TLS
|
||
reverse proxy), use the rightmost X-Forwarded-For entry that is not
|
||
itself a trusted proxy. Otherwise use the direct peer address —
|
||
X-Forwarded-For is spoofable, so it is only honoured for peers the
|
||
operator explicitly listed in ``trusted_proxies``.
|
||
"""
|
||
peer = request.client.host if request.client else "unknown"
|
||
if not trusted_proxies or peer not in trusted_proxies:
|
||
return peer
|
||
hops = [h.strip() for h in request.headers.get("x-forwarded-for", "").split(",")]
|
||
for hop in reversed(hops):
|
||
if hop and hop not in trusted_proxies:
|
||
return hop
|
||
return peer
|
||
|
||
|
||
# ── 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 [
|
||
{
|
||
"index": info.index,
|
||
"name": info.name,
|
||
"uuid": info.uuid,
|
||
"pci_bus_id": info.pci_bus_id,
|
||
}
|
||
for info in gpu_infos
|
||
]
|
||
|
||
|
||
@app.get("/api/gpu")
|
||
async def api_gpu(gpu_index: int = 0):
|
||
"""GPU info: name, driver version, VRAM."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
driver = get_driver_version()
|
||
vram = get_vram_total(gpu_index)
|
||
return {
|
||
"name": g_state["gpu_name"],
|
||
"index": gpu_index,
|
||
"driver_version": driver,
|
||
"vram_bytes": vram,
|
||
"vram_gib": round(vram / (1024**3), 2) if vram else None,
|
||
}
|
||
|
||
|
||
@app.get("/api/dashboard")
|
||
async def api_dashboard(gpu_index: int = 0):
|
||
"""Static GPU info for the Dashboard tab (VBIOS, CUDA cores, PCIe, BAR1, etc.).
|
||
|
||
Live values (clocks, temps, power, throttle) come from the monitor WebSocket.
|
||
"""
|
||
_, g_state = _require_gpu(gpu_index)
|
||
info = await _run(get_dashboard_info, gpu_index, g_state["gpu_name"])
|
||
return info
|
||
|
||
|
||
@app.get("/api/curve")
|
||
async def api_curve(gpu_index: int = 0):
|
||
"""Full CurveState: all V/F points with base freq, voltage, delta, effective freq."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
state, err = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
if state is None:
|
||
raise HTTPException(status_code=500, detail=f"Failed to read curve: {err}")
|
||
|
||
# Update reconciliation baseline
|
||
g_state["last_offsets"] = [p.delta_khz for p in state.points]
|
||
return _curve_state_dict(state)
|
||
|
||
|
||
@app.get("/api/curve/{point}")
|
||
async def api_curve_point(point: int, gpu_index: int = 0):
|
||
"""Single V/F point detail."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
state, err = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
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}"
|
||
)
|
||
return _vfpoint_dict(state.points[point])
|
||
|
||
|
||
@app.get("/api/ranges")
|
||
async def api_ranges(gpu_index: int = 0):
|
||
"""Clock boost domain ranges (min/max offset per domain)."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
ranges, err = await _run(get_clock_ranges, gpu)
|
||
if ranges is None:
|
||
raise HTTPException(status_code=500, detail=f"Failed to read ranges: {err}")
|
||
return ranges
|
||
|
||
|
||
@app.get("/api/voltage")
|
||
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:
|
||
raise HTTPException(status_code=500, detail=f"Failed to read voltage: {err}")
|
||
return {"voltage_uv": voltage_uv, "voltage_mv": voltage_uv / 1000.0}
|
||
|
||
|
||
@app.get("/api/monitor")
|
||
async def api_monitor(gpu_index: int = 0):
|
||
"""One-shot monitoring snapshot: voltage, clock, temp, power, fan, p-state, VRAM, utilization."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
sample = await _run(poll, gpu, gpu_index)
|
||
return _sample_dict(sample)
|
||
|
||
|
||
@app.get("/api/snapshots")
|
||
async def api_snapshots():
|
||
"""List saved ClockBoostTable snapshots."""
|
||
cfg: Config = _state["config"]
|
||
snapshots = await _run(list_snapshots, cfg.snapshot_dir)
|
||
return [
|
||
{
|
||
"filepath": s.filepath,
|
||
"timestamp": s.timestamp,
|
||
"gpu": s.gpu,
|
||
"nonzero_offsets": s.nonzero_offsets,
|
||
"size": s.size,
|
||
}
|
||
for s in snapshots
|
||
]
|
||
|
||
|
||
def _persist_config_field(key: str, value) -> None:
|
||
"""Write a single key into /etc/nvcurve/config.json if the file exists.
|
||
|
||
The file is created by `service install`. If it doesn't exist (e.g. the
|
||
service was never installed), config changes are in-memory only for the
|
||
current server session. Silently ignores errors.
|
||
"""
|
||
import json as _json
|
||
import os as _os
|
||
|
||
config_path = "/etc/nvcurve/config.json"
|
||
if not _os.path.exists(config_path):
|
||
return
|
||
try:
|
||
with open(config_path) as f:
|
||
data = _json.load(f)
|
||
except Exception:
|
||
data = {}
|
||
if value is not None:
|
||
data[key] = value
|
||
else:
|
||
data.pop(key, None)
|
||
try:
|
||
with open(config_path, "w") as f:
|
||
_json.dump(data, f)
|
||
except Exception as exc:
|
||
log.warning("Failed to persist config field %r: %s", key, exc)
|
||
|
||
|
||
def _gpu_stable_key(gpu_index: int) -> str:
|
||
"""Return a stable identifier for a GPU suitable for use as a config key.
|
||
|
||
Preference order: NVML UUID → PCI bus ID → 'idx:{n}' fallback.
|
||
UUID is the most stable across reboots and GPU slot changes.
|
||
"""
|
||
g_state = _state["gpus"].get(gpu_index, {})
|
||
uuid = g_state.get("uuid")
|
||
if uuid:
|
||
return uuid
|
||
pci = g_state.get("pci_bus_id")
|
||
if pci is not None:
|
||
return f"pci:{pci:04x}"
|
||
return f"idx:{gpu_index}"
|
||
|
||
|
||
def _persist_auto_load_profiles(profiles: dict[str, str]) -> None:
|
||
"""Persist auto_load_profiles dict to config.json."""
|
||
_persist_config_field("auto_load_profiles", profiles if profiles else None)
|
||
|
||
|
||
def _persist_fan_curves(fan_curves: dict) -> None:
|
||
"""Persist the per-GPU active fan curves dict to config.json."""
|
||
_persist_config_field("fan_curves", fan_curves if fan_curves else None)
|
||
|
||
|
||
def _persist_power_cap_modes(modes: dict[str, str]) -> None:
|
||
"""Persist the per-GPU experimental power-cap mode dict to config.json."""
|
||
_persist_config_field("power_cap_modes", modes if modes else None)
|
||
|
||
|
||
def _power_cap_mode(cfg: Config, gpu_index: int) -> str:
|
||
"""Return the effective power-cap mode for a GPU ("nvml" or "ioctl")."""
|
||
mode = cfg.power_cap_modes.get(_gpu_stable_key(gpu_index), "nvml")
|
||
return mode if mode in ("nvml", "ioctl") else "nvml"
|
||
|
||
|
||
@app.get("/api/profiles")
|
||
async def api_profiles(gpu_index: int = 0):
|
||
"""List saved native profiles, the active profile name, and the auto-load profile name."""
|
||
cfg: Config = _state["config"]
|
||
profiles = await _run(list_profiles, cfg.profile_dir)
|
||
g_state = _state["gpus"].get(gpu_index)
|
||
active = g_state["active_profile"] if g_state else None
|
||
return {
|
||
"profiles": profiles,
|
||
"active": active,
|
||
"auto_load": cfg.auto_load_profiles.get(_gpu_stable_key(gpu_index)),
|
||
}
|
||
|
||
|
||
@app.post("/api/profiles")
|
||
async def api_profile_save(req: ProfileSaveRequest, gpu_index: int = 0):
|
||
"""Save current GPU state (curve deltas + limits) as a named profile."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
|
||
state, err = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
if state is None:
|
||
raise HTTPException(status_code=500, detail=f"Failed to read curve: {err}")
|
||
|
||
curve_deltas = {str(p.index): p.delta_khz for p in state.points if p.delta_khz != 0}
|
||
|
||
try:
|
||
mode = _power_cap_mode(cfg, gpu_index)
|
||
power_info = await _run(get_power_limit, gpu_index, mode)
|
||
offsets = await _run(get_clock_offsets, gpu_index)
|
||
power_limit_w = power_info.get("power_limit_w")
|
||
mem_offset_mhz = offsets.get("mem_offset_mhz")
|
||
except Exception:
|
||
power_limit_w = None
|
||
mem_offset_mhz = None
|
||
mode = "nvml"
|
||
|
||
data = ProfileData(
|
||
name=req.name,
|
||
gpu_name=g_state["gpu_name"],
|
||
curve_deltas=curve_deltas,
|
||
mem_offset_mhz=mem_offset_mhz,
|
||
power_limit_w=power_limit_w,
|
||
power_cap_mode=mode,
|
||
fan_curve=g_state.get("fan_curve") if g_state.get("fan_curve_active") else None,
|
||
fan_targets=g_state.get("fan_targets")
|
||
if g_state.get("fan_curve_active")
|
||
else None,
|
||
)
|
||
filepath = await _run(save_profile, cfg.profile_dir, data)
|
||
g_state["active_profile"] = req.name
|
||
return {"ok": True, "filepath": filepath}
|
||
|
||
|
||
async def _auto_apply_profile_with_retry(
|
||
name: str, gpu_index: int = 0, max_retries: int = 3
|
||
) -> None:
|
||
"""Apply profile on startup with read-back verification and exponential backoff retry.
|
||
|
||
Raises FileNotFoundError if the profile does not exist (no point retrying).
|
||
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"]
|
||
|
||
safe_name = "".join(c for c in name if c.isalnum() or c in " _-()").strip()
|
||
filepath = _os.path.join(cfg.profile_dir, f"{safe_name}.json")
|
||
profile = await _run(load_profile, filepath) # raises FileNotFoundError if missing
|
||
|
||
expected: dict[int, int] = (
|
||
_int_key_deltas(profile.curve_deltas) 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),
|
||
)
|
||
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,
|
||
)
|
||
else:
|
||
mismatches = [
|
||
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,
|
||
)
|
||
return
|
||
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,
|
||
)
|
||
return
|
||
|
||
if attempt < max_retries - 1:
|
||
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
|
||
)
|
||
|
||
|
||
async def _apply_profile(name: str, gpu_index: int = 0) -> list[str]:
|
||
"""Load and apply a saved profile to hardware.
|
||
|
||
Returns a list of error strings. An empty list means success.
|
||
Raises FileNotFoundError if the profile file does not exist.
|
||
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"]
|
||
|
||
safe_name = "".join(c for c in name if c.isalnum() or c in " _-()").strip()
|
||
filepath = _os.path.join(cfg.profile_dir, f"{safe_name}.json")
|
||
|
||
# Let FileNotFoundError propagate so callers can map it to 404 or a warning.
|
||
profile = await _run(load_profile, filepath)
|
||
|
||
errs: list[str] = []
|
||
|
||
# Apply mem offset first — driver may reset curve table as a side-effect.
|
||
if profile.mem_offset_mhz is not None:
|
||
ok, msg = await _run(set_clock_offsets, None, profile.mem_offset_mhz, gpu_index)
|
||
if not ok:
|
||
errs.append(f"Mem offset: {msg}")
|
||
|
||
if profile.power_limit_w is not None:
|
||
mode = profile.power_cap_mode or "nvml"
|
||
ok, msg = await _run(set_power_limit, profile.power_limit_w, gpu_index, mode)
|
||
if not ok:
|
||
errs.append(f"Power limit: {msg}")
|
||
|
||
# Apply curve deltas (after mem offset which may have wiped them).
|
||
async with g_state["write_lock"]:
|
||
if profile.curve_deltas:
|
||
deltas = _int_key_deltas(profile.curve_deltas)
|
||
errors = validate_write(deltas, cfg.max_delta_khz)
|
||
if errors:
|
||
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,
|
||
)
|
||
ret, desc = await _run(write_offsets, gpu, deltas)
|
||
if ret != 0:
|
||
errs.append(f"Curve write failed ({ret}): {desc}")
|
||
else:
|
||
await _run(reset_offsets, gpu)
|
||
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
|
||
# Apply fan curve if present in profile (persists it so it survives restarts);
|
||
# otherwise deactivate any active fan curve (and clear its persistence).
|
||
if profile.fan_curve:
|
||
ok, msg = validate_curve(profile.fan_curve)
|
||
if not ok:
|
||
errs.append(f"Fan curve: {msg}")
|
||
else:
|
||
await _activate_fan_curve(gpu_index, profile.fan_curve, profile.fan_targets)
|
||
elif g_state.get("fan_curve_active"):
|
||
await _deactivate_fan_curve(gpu_index, reset_hardware=True)
|
||
|
||
if not errs:
|
||
g_state["active_profile"] = name
|
||
return errs
|
||
|
||
|
||
@app.post("/api/profiles/{name}/apply")
|
||
async def api_profile_apply(name: str, gpu_index: int = 0):
|
||
"""Apply a saved profile to hardware (curve deltas + limits)."""
|
||
_require_gpu(gpu_index)
|
||
try:
|
||
errs = await _apply_profile(name, gpu_index)
|
||
except FileNotFoundError as err:
|
||
raise HTTPException(
|
||
status_code=404, detail=f"Profile '{name}' not found"
|
||
) from err
|
||
except Exception as e:
|
||
raise HTTPException(
|
||
status_code=500, detail=f"Failed to load profile: {e}"
|
||
) from e
|
||
if errs:
|
||
raise HTTPException(status_code=500, detail="; ".join(errs))
|
||
return {"ok": True}
|
||
|
||
|
||
@app.delete("/api/profiles/{name}")
|
||
async def api_profile_delete(name: str):
|
||
"""Delete a saved profile by name."""
|
||
cfg: Config = _state["config"]
|
||
ok = await _run(delete_profile, cfg.profile_dir, name)
|
||
if not ok:
|
||
raise HTTPException(status_code=404, detail=f"Profile '{name}' not found")
|
||
for g_state in _state["gpus"].values():
|
||
if g_state["active_profile"] == name:
|
||
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
|
||
}
|
||
_persist_auto_load_profiles(cfg.auto_load_profiles)
|
||
return {"ok": True}
|
||
|
||
|
||
@app.post("/api/profiles/{name}/rename")
|
||
async def api_profile_rename(name: str, req: ProfileRenameRequest):
|
||
"""Rename a profile."""
|
||
cfg: Config = _state["config"]
|
||
if not req.new_name.strip():
|
||
raise HTTPException(status_code=400, detail="New name cannot be empty")
|
||
ok = await _run(rename_profile, cfg.profile_dir, name, req.new_name.strip())
|
||
if not ok:
|
||
raise HTTPException(status_code=404, detail=f"Profile '{name}' not found")
|
||
for g_state in _state["gpus"].values():
|
||
if g_state["active_profile"] == name:
|
||
g_state["active_profile"] = req.new_name.strip()
|
||
changed = any(v == name for v in cfg.auto_load_profiles.values())
|
||
if changed:
|
||
cfg.auto_load_profiles = {
|
||
k: (req.new_name.strip() if v == name else v)
|
||
for k, v in cfg.auto_load_profiles.items()
|
||
}
|
||
_persist_auto_load_profiles(cfg.auto_load_profiles)
|
||
return {"ok": True}
|
||
|
||
|
||
@app.get("/api/config")
|
||
async def api_config_get(gpu_index: int = 0):
|
||
"""Get mutable server configuration for a specific GPU."""
|
||
cfg: Config = _state["config"]
|
||
return {"auto_load_profile": cfg.auto_load_profiles.get(_gpu_stable_key(gpu_index))}
|
||
|
||
|
||
@app.post("/api/config")
|
||
async def api_config_update(req: ConfigUpdateRequest):
|
||
"""Update mutable server configuration. Changes persist to /etc/nvcurve/config.json if present."""
|
||
if req.gpu_index not in _state["gpus"]:
|
||
raise HTTPException(status_code=404, detail=f"GPU {req.gpu_index} not found")
|
||
cfg: Config = _state["config"]
|
||
key = _gpu_stable_key(req.gpu_index)
|
||
if req.auto_load_profile:
|
||
cfg.auto_load_profiles[key] = req.auto_load_profile
|
||
else:
|
||
cfg.auto_load_profiles.pop(key, None)
|
||
_persist_auto_load_profiles(cfg.auto_load_profiles)
|
||
return {"ok": True, "auto_load_profile": cfg.auto_load_profiles.get(key)}
|
||
|
||
|
||
@app.get("/api/limits")
|
||
async def api_limits(gpu_index: int = 0):
|
||
"""Current performance limits: power and clock offsets."""
|
||
cfg: Config = _state["config"]
|
||
mode = _power_cap_mode(cfg, gpu_index)
|
||
power = await _run(get_power_limit, gpu_index, mode)
|
||
offsets = await _run(get_clock_offsets, gpu_index)
|
||
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
|
||
}
|
||
|
||
|
||
@app.post("/api/limits")
|
||
async def api_limits_update(req: LimitsRequest, gpu_index: int = 0):
|
||
"""Update performance limits."""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
errs = []
|
||
|
||
if req.power_cap_mode is not None:
|
||
if req.power_cap_mode not in ("nvml", "ioctl"):
|
||
raise HTTPException(
|
||
status_code=400, detail="power_cap_mode must be 'nvml' or 'ioctl'"
|
||
)
|
||
if req.power_cap_mode == "ioctl":
|
||
# Verify the GPU actually exposes the RM interface before enabling,
|
||
# so a client can't lock a GPU into a mode where every power
|
||
# operation fails (ioctl mode has no NVML fallback by design).
|
||
info = await _run(get_power_limit, gpu_index, "ioctl")
|
||
if not info.get("rm_power_supported"):
|
||
raise HTTPException(
|
||
status_code=409,
|
||
detail="Experimental RM power control is not supported "
|
||
"on this GPU/driver",
|
||
)
|
||
key = _gpu_stable_key(gpu_index)
|
||
if req.power_cap_mode == "nvml":
|
||
cfg.power_cap_modes.pop(key, None)
|
||
else:
|
||
cfg.power_cap_modes[key] = "ioctl"
|
||
_persist_power_cap_modes(cfg.power_cap_modes)
|
||
|
||
mode = _power_cap_mode(cfg, gpu_index)
|
||
|
||
if req.power_limit_w is not None:
|
||
ok, msg = await _run(set_power_limit, req.power_limit_w, gpu_index, mode)
|
||
if not ok:
|
||
errs.append(f"Power Limit: {msg}")
|
||
|
||
if req.mem_offset_mhz is not None:
|
||
ok, msg = await _run(set_clock_offsets, None, req.mem_offset_mhz, gpu_index)
|
||
if not ok:
|
||
errs.append(f"Mem Offset: {msg}")
|
||
else:
|
||
# Setting mem offset may reset the GPC/curve table as a driver side-effect.
|
||
# Re-apply the last known curve offsets to restore them.
|
||
await _reapply_curve(gpu_index)
|
||
|
||
if errs:
|
||
raise HTTPException(status_code=500, detail="; ".join(errs))
|
||
|
||
g_state["active_profile"] = None
|
||
|
||
return {"ok": True}
|
||
|
||
|
||
async def _reapply_curve(gpu_index: int) -> None:
|
||
"""Re-write the last known V/F curve offsets to hardware and notify WS clients."""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
gpu = g_state["gpu"]
|
||
last = g_state["last_offsets"]
|
||
if gpu is None or not last:
|
||
return
|
||
deltas = {i: off for i, off in enumerate(last) if off != 0}
|
||
if not deltas:
|
||
return
|
||
try:
|
||
await _run(write_offsets, gpu, deltas)
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
except Exception as exc:
|
||
log.warning("_reapply_curve: %s", exc)
|
||
|
||
|
||
async def _update_offsets_and_broadcast(gpu_index: int) -> None:
|
||
"""Re-read curve offsets, update the reconciliation baseline, and push to WS clients.
|
||
|
||
When curve WS clients are connected, a single read_curve call covers both
|
||
updating the baseline and the broadcast payload — avoiding a redundant
|
||
read_clock_offsets (ClockBoostTable) call that would otherwise happen first.
|
||
"""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
gpu = g_state["gpu"]
|
||
if gpu is None:
|
||
return
|
||
if g_state["curve_clients"]:
|
||
state, _ = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
if state:
|
||
g_state["last_offsets"] = [p.delta_khz for p in state.points]
|
||
await _broadcast(g_state["curve_clients"], _curve_state_dict(state))
|
||
else:
|
||
offsets, _ = await _run(read_clock_offsets, gpu)
|
||
g_state["last_offsets"] = offsets
|
||
|
||
|
||
@app.post("/api/limits/reset")
|
||
async def api_limits_reset(gpu_index: int = 0):
|
||
"""Reset power limit to hardware default and memory clock offset to 0."""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
errs = []
|
||
|
||
# Reset uses the GPU's current mode: in ioctl mode the default is
|
||
# restored through the RM route (which can also restore a previous
|
||
# below-VBIOS-minimum cap).
|
||
mode = _power_cap_mode(cfg, gpu_index)
|
||
power = await _run(get_power_limit, gpu_index, mode)
|
||
default_w = power.get("default_power_limit_w")
|
||
if default_w is not None:
|
||
ok, msg = await _run(set_power_limit, default_w, gpu_index, mode)
|
||
if not ok:
|
||
errs.append(f"Power Limit: {msg}")
|
||
|
||
ok, msg = await _run(set_clock_offsets, None, 0, gpu_index)
|
||
if not ok:
|
||
errs.append(f"Mem Offset: {msg}")
|
||
else:
|
||
await _reapply_curve(gpu_index)
|
||
|
||
if errs:
|
||
raise HTTPException(status_code=500, detail="; ".join(errs))
|
||
|
||
g_state["active_profile"] = None
|
||
|
||
return {"ok": True}
|
||
|
||
|
||
# ── Fan endpoints ──────────────────────────────────────────────────────────────
|
||
|
||
|
||
@app.get("/api/fans")
|
||
async def api_fans(gpu_index: int = 0):
|
||
"""Current fan state: per-fan %, curve, and whether curve control is active."""
|
||
_get_gpu_state(gpu_index)
|
||
g_state = _state["gpus"][gpu_index]
|
||
info = await _run(get_fan_info, gpu_index)
|
||
curve_active = g_state.get("fan_curve_active", False)
|
||
return {
|
||
**info,
|
||
"fan_mode": "curve" if curve_active else "auto",
|
||
"curve": g_state.get("fan_curve"),
|
||
"curve_active": curve_active,
|
||
"fan_targets": g_state.get("fan_targets"),
|
||
}
|
||
|
||
|
||
@app.post("/api/fans")
|
||
async def api_fans_update(req: FanCurveRequest, gpu_index: int = 0):
|
||
"""Set or update the fan curve. Starts the fan control poller.
|
||
|
||
req.fans selects which fans the curve drives (None = all fans).
|
||
"""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
|
||
curve_data = [{"temp_c": p.temp_c, "fan_pct": p.fan_pct} for p in req.curve]
|
||
ok, msg = validate_curve(curve_data)
|
||
if not ok:
|
||
raise HTTPException(status_code=400, detail=msg)
|
||
|
||
fans = sorted(set(req.fans)) if req.fans is not None else None
|
||
|
||
# Test that fan control is available on this GPU, probing with the
|
||
# target for the *current* temperature so the fan is never briefly
|
||
# 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"]
|
||
)
|
||
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, fans)
|
||
if not fan_ok:
|
||
raise HTTPException(
|
||
status_code=500, detail=f"Fan control not available: {fan_msg}"
|
||
)
|
||
|
||
# Activate the curve (starts the poller) and persist it so it survives restarts.
|
||
await _activate_fan_curve(gpu_index, curve_data, fans)
|
||
|
||
return {"ok": True}
|
||
|
||
|
||
@app.post("/api/fans/reset")
|
||
async def api_fans_reset(gpu_index: int = 0):
|
||
"""Deactivate fan curve control and restore automatic fan mode."""
|
||
_get_gpu_state(gpu_index)
|
||
|
||
# Stop the poller, clear state, restore automatic fan control, and clear
|
||
# the persisted curve so it is not re-applied on the next server start.
|
||
await _deactivate_fan_curve(gpu_index, reset_hardware=True)
|
||
|
||
return {"ok": True}
|
||
|
||
|
||
@app.post("/api/fans/speed")
|
||
async def api_fans_speed(req: FanSpeedRequest, gpu_index: int = 0):
|
||
"""One-shot set fan(s) to an exact percentage (bypasses curve).
|
||
|
||
req.fan selects a single fan index; None sets all fans.
|
||
"""
|
||
_get_gpu_state(gpu_index)
|
||
pct = max(0, min(100, req.fan_pct))
|
||
fans = [req.fan] if req.fan is not None else None
|
||
ok, msg = await _run(set_fan_speed, gpu_index, pct, fans)
|
||
if not ok:
|
||
raise HTTPException(status_code=500, detail=msg)
|
||
return {"ok": True}
|
||
|
||
|
||
# ── 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.
|
||
|
||
Returns None if no external change detected (or no baseline).
|
||
"""
|
||
g_state = _get_gpu_state(gpu_index)
|
||
gpu = g_state["gpu"]
|
||
last = g_state["last_offsets"]
|
||
if last is None:
|
||
return None
|
||
|
||
current, err = await _run(read_clock_offsets, gpu)
|
||
if current is None:
|
||
return None # Can't read — let the write attempt proceed
|
||
|
||
changed = [i for i, (a, b) in enumerate(zip(last, current, strict=False)) if a != b]
|
||
if not changed:
|
||
return None
|
||
|
||
# External tool changed the curve — active profile is no longer current.
|
||
g_state["active_profile"] = None
|
||
|
||
return {
|
||
"warning": "external_change_detected",
|
||
"message": (
|
||
f"{len(changed)} point(s) changed since last read "
|
||
f"(e.g. by LACT, nvidia-smi, or another tool). "
|
||
"The write will proceed using the current hardware state."
|
||
),
|
||
"changed_points": changed[:20], # cap list for readability
|
||
}
|
||
|
||
|
||
@app.post("/api/curve/write")
|
||
async def api_curve_write(req: WriteRequest, gpu_index: int = 0):
|
||
"""Write per-point frequency offsets. {deltas: {point_index: delta_kHz}}"""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
|
||
vfp_state, _ = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
|
||
# The safety cap is always the server-side config value — clients cannot
|
||
# raise it per request (shared systems must not let one user override the
|
||
# hardware safety limit). Raise it in /etc/nvcurve/config.json if needed.
|
||
errors = validate_write(req.deltas, cfg.max_delta_khz)
|
||
if errors:
|
||
raise HTTPException(status_code=400, detail={"errors": errors})
|
||
|
||
# Check for negative-freq warnings before writing (best-effort, non-blocking)
|
||
freq_warnings: list[str] = []
|
||
if vfp_state:
|
||
vfp_freqs = [p.freq_khz for p in vfp_state.points]
|
||
freq_warnings = check_negative_freq_warnings(
|
||
req.deltas, vfp_freqs, g_state["last_offsets"] or []
|
||
)
|
||
|
||
async with g_state["write_lock"]:
|
||
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,
|
||
)
|
||
|
||
ret, desc = await _run(write_offsets, gpu, req.deltas)
|
||
if ret != 0:
|
||
raise HTTPException(status_code=500, detail=f"Write failed ({ret}): {desc}")
|
||
|
||
# Update baseline and push curve update to WS clients
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
g_state["active_profile"] = None
|
||
|
||
result = {"ok": True, "return_code": ret, "description": desc}
|
||
if warning:
|
||
result["warning"] = warning
|
||
if freq_warnings:
|
||
result["freq_warnings"] = freq_warnings
|
||
return result
|
||
|
||
|
||
@app.post("/api/curve/write/global")
|
||
async def api_curve_write_global(req: GlobalOffsetRequest, gpu_index: int = 0):
|
||
"""Apply a uniform frequency offset to all curve points."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
|
||
vfp_state, _ = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
if not vfp_state:
|
||
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"}
|
||
# Safety cap is the server-side config value only (see api_curve_write).
|
||
errors = validate_write(all_deltas, cfg.max_delta_khz)
|
||
if errors:
|
||
raise HTTPException(status_code=400, detail={"errors": errors})
|
||
|
||
freq_warnings: list[str] = []
|
||
if vfp_state:
|
||
vfp_freqs = [p.freq_khz for p in vfp_state.points]
|
||
freq_warnings = check_negative_freq_warnings(
|
||
all_deltas, vfp_freqs, g_state["last_offsets"] or []
|
||
)
|
||
|
||
async with g_state["write_lock"]:
|
||
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,
|
||
)
|
||
|
||
ret, desc = await _run(write_global_offset, gpu, req.delta_khz)
|
||
if ret != 0:
|
||
raise HTTPException(status_code=500, detail=f"Write failed ({ret}): {desc}")
|
||
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
g_state["active_profile"] = None
|
||
|
||
result = {"ok": True, "return_code": ret, "description": desc}
|
||
if warning:
|
||
result["warning"] = warning
|
||
if freq_warnings:
|
||
result["freq_warnings"] = freq_warnings
|
||
return result
|
||
|
||
|
||
@app.post("/api/curve/reset")
|
||
async def api_curve_reset(gpu_index: int = 0):
|
||
"""Reset all frequency offsets to zero."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
|
||
async with g_state["write_lock"]:
|
||
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,
|
||
)
|
||
|
||
ret, desc = await _run(reset_offsets, gpu)
|
||
if ret != 0:
|
||
raise HTTPException(status_code=500, detail=f"Reset failed ({ret}): {desc}")
|
||
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
g_state["active_profile"] = None
|
||
|
||
result = {"ok": True, "return_code": ret, "description": desc}
|
||
if warning:
|
||
result["warning"] = warning
|
||
return result
|
||
|
||
|
||
@app.post("/api/curve/verify")
|
||
async def api_curve_verify(req: VerifyRequest, gpu_index: int = 0):
|
||
"""Write-verify-read cycle. Returns per-point match results and collateral changes."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
|
||
errors = validate_write(req.deltas, cfg.max_delta_khz)
|
||
if errors:
|
||
raise HTTPException(status_code=400, detail={"errors": errors})
|
||
|
||
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}"
|
||
)
|
||
|
||
# Always snapshot before verify — it's a testing operation
|
||
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)
|
||
if ret != 0:
|
||
raise HTTPException(status_code=500, detail=f"Write failed ({ret}): {desc}")
|
||
|
||
await asyncio.sleep(0.2)
|
||
|
||
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}"
|
||
)
|
||
|
||
g_state["active_profile"] = None
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
|
||
points_result = []
|
||
all_matched = True
|
||
for point, expected in sorted(req.deltas.items()):
|
||
actual = after_offsets[point]
|
||
match = actual == expected
|
||
if not match:
|
||
all_matched = False
|
||
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]}
|
||
for i in range(len(before_offsets))
|
||
if i not in req.deltas and before_offsets[i] != after_offsets[i]
|
||
]
|
||
|
||
return {
|
||
"ok": all_matched and not collateral,
|
||
"all_matched": all_matched,
|
||
"no_side_effects": not collateral,
|
||
"return_code": ret,
|
||
"description": desc,
|
||
"points": points_result,
|
||
"collateral_changes": collateral,
|
||
}
|
||
|
||
|
||
@app.post("/api/shutdown")
|
||
async def api_shutdown():
|
||
"""Gracefully shut down the server process.
|
||
|
||
Disabled when ``allow_api_shutdown`` is false in the config — on shared
|
||
systems stop the service via systemd instead.
|
||
"""
|
||
import os
|
||
import signal
|
||
|
||
cfg: Config = _state["config"]
|
||
if not cfg.allow_api_shutdown:
|
||
raise HTTPException(
|
||
status_code=403,
|
||
detail="API shutdown is disabled (allow_api_shutdown: false)",
|
||
)
|
||
|
||
loop = asyncio.get_running_loop()
|
||
loop.call_later(0.1, lambda: os.kill(os.getpid(), signal.SIGTERM))
|
||
return {"ok": True}
|
||
|
||
|
||
@app.post("/api/snapshot/save")
|
||
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
|
||
)
|
||
if path is None:
|
||
raise HTTPException(status_code=500, detail="Failed to save snapshot")
|
||
return {"ok": True, "filepath": path}
|
||
|
||
|
||
@app.post("/api/snapshot/restore")
|
||
async def api_snapshot_restore(req: SnapshotRestoreRequest, gpu_index: int = 0):
|
||
"""Restore a ClockBoostTable snapshot. Uses most recent if filepath not specified."""
|
||
gpu, g_state = _require_gpu(gpu_index)
|
||
cfg: Config = _state["config"]
|
||
|
||
async with g_state["write_lock"]:
|
||
ok = await _run(snapshot_restore, gpu, cfg.snapshot_dir, req.filepath)
|
||
if not ok:
|
||
raise HTTPException(status_code=500, detail="Failed to restore snapshot")
|
||
|
||
await _update_offsets_and_broadcast(gpu_index)
|
||
g_state["active_profile"] = None
|
||
|
||
return {"ok": True}
|
||
|
||
|
||
# ── 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()
|
||
if data.get("action") != "subscribe":
|
||
await ws.close()
|
||
return
|
||
gpu_index = data.get("gpu_index", 0)
|
||
except WebSocketDisconnect:
|
||
return
|
||
except Exception:
|
||
await ws.close()
|
||
return
|
||
|
||
g_state = _state["gpus"].get(gpu_index)
|
||
if not g_state:
|
||
await ws.close()
|
||
return
|
||
|
||
g_state["monitor_clients"].add(ws)
|
||
try:
|
||
gpu = g_state["gpu"]
|
||
if gpu is not None:
|
||
sample = await _run(poll, gpu, gpu_index)
|
||
await ws.send_json(_sample_dict(sample))
|
||
|
||
while True:
|
||
await ws.receive_text()
|
||
except WebSocketDisconnect:
|
||
pass
|
||
except Exception:
|
||
log.debug("monitor ws client error", exc_info=True)
|
||
finally:
|
||
g_state["monitor_clients"].discard(ws)
|
||
|
||
|
||
@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()
|
||
if data.get("action") != "subscribe":
|
||
await ws.close()
|
||
return
|
||
gpu_index = data.get("gpu_index", 0)
|
||
except WebSocketDisconnect:
|
||
return
|
||
except Exception:
|
||
await ws.close()
|
||
return
|
||
|
||
g_state = _state["gpus"].get(gpu_index)
|
||
if not g_state:
|
||
await ws.close()
|
||
return
|
||
|
||
g_state["curve_clients"].add(ws)
|
||
try:
|
||
gpu = g_state["gpu"]
|
||
if gpu is not None:
|
||
state, _ = await _run(read_curve, gpu, g_state["gpu_name"])
|
||
if state:
|
||
await ws.send_json(_curve_state_dict(state))
|
||
|
||
while True:
|
||
await ws.receive_text()
|
||
except WebSocketDisconnect:
|
||
pass
|
||
except Exception:
|
||
log.debug("curve ws client error", exc_info=True)
|
||
finally:
|
||
g_state["curve_clients"].discard(ws)
|
||
|
||
|
||
# ── Frontend SPA ──────────────────────────────────────────────────────────────
|
||
|
||
|
||
# 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")
|
||
|
||
|
||
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:
|
||
# Project requires Python >= 3.12, so the 3.7-compat finding is a false positive.
|
||
from importlib.resources import ( # nosemgrep: python.lang.compatibility.python37.python37-compatibility-importlib2
|
||
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.get("/{catchall:path}")
|
||
async def serve_spa(catchall: str):
|
||
if catchall.startswith(("api/", "ws/")):
|
||
raise HTTPException(status_code=404, detail="Not Found")
|
||
|
||
if not os.path.isdir(_dist_dir):
|
||
return {"error": "Frontend not built. Run pnpm build in frontend/."}
|
||
|
||
# Contain the resolved path inside the dist directory. The raw URL path
|
||
# can carry encoded ".." segments (e.g. /%2e%2e/etc/passwd) that would
|
||
# otherwise escape the dist dir via os.path.join — an unauthenticated
|
||
# arbitrary-file-read since the server runs as root.
|
||
base = os.path.realpath(_dist_dir)
|
||
path = os.path.realpath(os.path.join(_dist_dir, catchall))
|
||
if path != base and not path.startswith(base + os.sep):
|
||
raise HTTPException(status_code=404, detail="Not Found")
|
||
|
||
if os.path.isfile(path) and catchall:
|
||
return FileResponse(path)
|
||
|
||
index = os.path.join(_dist_dir, "index.html")
|
||
if os.path.isfile(index):
|
||
return FileResponse(index)
|
||
|
||
raise HTTPException(status_code=404, detail="Not Found")
|
||
|
||
|
||
# ── 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
|
||
return app
|
||
|
||
|
||
def run(
|
||
host: str = "127.0.0.1",
|
||
port: int = 8042,
|
||
gpu_index: int = 0,
|
||
config: Config = default_config,
|
||
open_browser: bool = False,
|
||
ssl_certfile: str | None = None,
|
||
ssl_keyfile: str | None = None,
|
||
) -> None:
|
||
"""Start the uvicorn server. Blocking.
|
||
|
||
When both ssl_certfile and ssl_keyfile are given (either here or in the
|
||
config), the server serves HTTPS and the session cookie is Secure.
|
||
"""
|
||
import socket
|
||
import threading
|
||
|
||
import uvicorn
|
||
|
||
_state["config"] = config
|
||
|
||
# CLI flags take precedence over config values.
|
||
certfile = ssl_certfile or config.ssl_certfile
|
||
keyfile = ssl_keyfile or config.ssl_keyfile
|
||
if certfile:
|
||
config.ssl_certfile = certfile
|
||
if keyfile:
|
||
config.ssl_keyfile = keyfile
|
||
tls = bool(certfile and keyfile)
|
||
|
||
# Fail fast on a bad TLS configuration — otherwise uvicorn dies at
|
||
# startup and (in daemon mode) the error is only visible in the server
|
||
# log while `serve status` reports "not running".
|
||
tls = False
|
||
if certfile and keyfile:
|
||
missing = [
|
||
f"{label} ({path})"
|
||
for label, path in (("certificate", certfile), ("key", keyfile))
|
||
if not os.path.isfile(path)
|
||
]
|
||
if missing:
|
||
print(f"Error: TLS file(s) not found: {', '.join(missing)}")
|
||
print(
|
||
"Fix the path (nvcurve service configure --ssl-certfile/--ssl-keyfile) "
|
||
"or disable TLS (--no-ssl)."
|
||
)
|
||
return
|
||
tls = True
|
||
|
||
# Suppress noisy websockets keepalive ping-timeout tracebacks — these are
|
||
# normal disconnection events (browser tab closed, network hiccup) and
|
||
# logging them at ERROR level creates false alarm noise.
|
||
logging.getLogger("websockets").setLevel(logging.CRITICAL)
|
||
|
||
# Fail fast if the port is already in use — silently shifting ports breaks
|
||
# client discovery. Users should configure a different port explicitly.
|
||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||
try:
|
||
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."
|
||
)
|
||
return
|
||
|
||
url = f"{'https' if tls else 'http'}://{host}:{port}"
|
||
|
||
# Print banner *before* uvicorn starts so it appears above uvicorn's own output.
|
||
# GPU name is populated by the lifespan; we omit it here since the server
|
||
# hasn't started yet, and the lifespan logs it via log.info.
|
||
print("\033[1;36m" + "─" * 60 + "\033[0m")
|
||
print("\033[1;32m" + " NVCurve".center(60) + "\033[0m")
|
||
print(f" {url}".center(60))
|
||
print("\033[1;36m" + "─" * 60 + "\033[0m")
|
||
print(" Press Ctrl+C to stop.")
|
||
print()
|
||
|
||
if open_browser and not _DEV_PORT:
|
||
threading.Timer(1.2, lambda: _open_browser_as_user(url)).start()
|
||
|
||
uvicorn.run(
|
||
app,
|
||
host=host,
|
||
port=port,
|
||
log_level="warning",
|
||
access_log=False,
|
||
ssl_certfile=certfile if tls else None,
|
||
ssl_keyfile=keyfile if tls else None,
|
||
)
|