Standalone support for the ThermalGrizzly WireView Pro II without depending on the external wireview_reporter exporter: - wireview.py: serial protocol (STX/ETX + 16-bit CRC16-CCITT) with vendor-data product identification, config version, UID, build string, screen layout, and temperature/power/current sensor reads; hwmon fallback with per-channel index resolution; udev-based detection (vendor 0x2560 / product 0x0101) with fallback port probing; serial read timeout and watchdog reconnect - server.py: startup detection (root only), 1 Hz poller gated on subscribed clients, WS 'wireview' channel, GET /api/wireview, rejected-product memoization, stale-read race guard - frontend: WireView tab (visible when a device is detected) with live temperature/power/current cards, sparkline history, and device info; nullable temp channels - tests: 76 tests covering parser, CRC, fault classification, hwmon resolution, JSON safety, and the serial transport against a pty-based fake device Verified live: tab appears with the device connected, disappears when unplugged, reconnects on re-plug, no serial traffic when idle.
2216 lines
77 KiB
Python
2216 lines
77 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 .wireview import create_device, find_wireview_ports
|
||
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,
|
||
"wireview": {
|
||
"clients": set(), # connected /ws/wireview clients
|
||
"device": None, # WireViewSerialDevice | WireViewHwmonDevice | None
|
||
"info": None, # device identity (static per connection)
|
||
"last_sample": None, # most recent sensor sample
|
||
"connected": False, # last read succeeded
|
||
"failures": 0, # consecutive failed reads
|
||
"rejected_ports": set(), # ports reporting an unsupported product
|
||
},
|
||
}
|
||
|
||
|
||
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)
|
||
|
||
|
||
# ── WireView Pro II (Thermal Grizzly) ─────────────────────────────────────────
|
||
# Consecutive failed reads after which a connected device is given up. Same
|
||
# rule as the exporter: one corrupt frame or a slow read never trips it, and
|
||
# an unplug is caught at once by the port-node check.
|
||
WIREVIEW_MAX_FAILED_READS = 6
|
||
WIREVIEW_WATCHDOG_INTERVAL_S = 5.0
|
||
|
||
|
||
async def _wireview_disconnect() -> None:
|
||
"""Drop the connected WireView device and tell clients it is gone."""
|
||
wv = _state["wireview"]
|
||
device = wv["device"]
|
||
if device is None:
|
||
return
|
||
wv["device"] = None
|
||
wv["connected"] = False
|
||
wv["info"] = None
|
||
wv["last_sample"] = None
|
||
wv["failures"] = 0
|
||
await _run(device.close)
|
||
log.info("WireView disconnected")
|
||
await _broadcast(wv["clients"], {"type": "unavailable"})
|
||
|
||
|
||
async def _wireview_connect() -> None:
|
||
"""Detect and connect a WireView Pro II. No-op when none is attached or
|
||
one is already connected."""
|
||
wv = _state["wireview"]
|
||
if wv["device"] is not None:
|
||
return
|
||
# Forget rejected ports that disappeared, so a re-plug gets a fresh probe.
|
||
if wv["rejected_ports"]:
|
||
present = set(await _run(find_wireview_ports))
|
||
wv["rejected_ports"] &= present
|
||
device = await _run(create_device)
|
||
if device is None:
|
||
return
|
||
# A port that reported an unsupported product is not re-probed (and
|
||
# re-logged) on every watchdog tick.
|
||
if device.port in wv["rejected_ports"]:
|
||
return
|
||
ok = await _run(device.connect)
|
||
if not ok:
|
||
await _run(device.close)
|
||
if device.rejected:
|
||
wv["rejected_ports"].add(device.port)
|
||
return
|
||
wv["device"] = device
|
||
wv["failures"] = 0
|
||
wv["connected"] = True
|
||
wv["info"] = await _run(device.info)
|
||
# Push the first sample right away so clients don't wait for the next tick.
|
||
sample = await _run(device.read_sample)
|
||
if sample is not None:
|
||
wv["last_sample"] = sample
|
||
await _broadcast(
|
||
wv["clients"], {"type": "sample", "info": wv["info"], "sample": sample}
|
||
)
|
||
else:
|
||
wv["connected"] = False
|
||
|
||
|
||
async def _wireview_poller() -> None:
|
||
"""Read the WireView at poll_interval_s and push to connected WS clients.
|
||
|
||
Skips reads while no client is subscribed (like the monitor poller), so
|
||
idle polling never contends for the port with other tools (official GUI,
|
||
wireviewd)."""
|
||
cfg: Config = _state["config"]
|
||
while True:
|
||
try:
|
||
wv = _state["wireview"]
|
||
device = wv["device"]
|
||
if device is not None and wv["clients"]:
|
||
sample = await _run(device.read_sample)
|
||
if wv["device"] is not device:
|
||
# The device was disconnected (unplug) while the read was
|
||
# in flight — drop the stale result instead of
|
||
# resurrecting state.
|
||
pass
|
||
elif sample is not None:
|
||
wv["failures"] = 0
|
||
wv["connected"] = True
|
||
wv["last_sample"] = sample
|
||
await _broadcast(
|
||
wv["clients"],
|
||
{"type": "sample", "info": wv["info"], "sample": sample},
|
||
)
|
||
else:
|
||
wv["failures"] += 1
|
||
if (
|
||
not await _run(device.node_exists)
|
||
or wv["failures"] >= WIREVIEW_MAX_FAILED_READS
|
||
):
|
||
await _wireview_disconnect()
|
||
except asyncio.CancelledError:
|
||
return
|
||
except Exception as exc:
|
||
log.warning("WireView poller error: %s", exc)
|
||
await asyncio.sleep(cfg.poll_interval_s)
|
||
|
||
|
||
async def _wireview_watchdog() -> None:
|
||
"""Hot-plug detection: connect when a WireView appears, disconnect when
|
||
its device node disappears."""
|
||
while True:
|
||
await asyncio.sleep(WIREVIEW_WATCHDOG_INTERVAL_S)
|
||
try:
|
||
wv = _state["wireview"]
|
||
device = wv["device"]
|
||
if device is None:
|
||
await _wireview_connect()
|
||
elif not await _run(device.node_exists):
|
||
await _wireview_disconnect()
|
||
except asyncio.CancelledError:
|
||
return
|
||
except Exception as exc:
|
||
log.warning("WireView watchdog error: %s", exc)
|
||
|
||
|
||
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)
|
||
|
||
# ── WireView Pro II detection ─────────────────────────────────────────────
|
||
# Independent of GPU discovery: the tab appears whenever the connector
|
||
# monitor is attached, and the watchdog handles hot-plug afterwards.
|
||
try:
|
||
await _wireview_connect()
|
||
except Exception as exc:
|
||
log.warning("WireView initial detection failed: %s", exc)
|
||
poller_tasks.append(asyncio.create_task(_wireview_poller()))
|
||
poller_tasks.append(asyncio.create_task(_wireview_watchdog()))
|
||
|
||
# ── 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
|
||
|
||
# Release the WireView port (if connected) so other tools can use it.
|
||
wv = _state["wireview"]
|
||
if wv["device"] is not None:
|
||
with suppress(Exception):
|
||
await _run(wv["device"].close)
|
||
wv["device"] = None
|
||
wv["connected"] = False
|
||
|
||
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}
|
||
|
||
|
||
# ── WireView Pro II (Thermal Grizzly) ─────────────────────────────────────────
|
||
|
||
|
||
@app.get("/api/wireview")
|
||
async def api_wireview():
|
||
"""WireView Pro II availability and the most recent sensor sample.
|
||
|
||
available: a device is connected (the UI shows the WireView tab).
|
||
sample: None until the first successful read.
|
||
"""
|
||
wv = _state["wireview"]
|
||
return {
|
||
"available": wv["device"] is not None,
|
||
"connected": wv["connected"],
|
||
"info": wv["info"],
|
||
"sample": wv["last_sample"],
|
||
}
|
||
|
||
|
||
# ── 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)
|
||
|
||
|
||
@app.websocket("/ws/wireview")
|
||
async def ws_wireview(ws: WebSocket):
|
||
"""Stream WireView Pro II sensor samples at poll_interval_s.
|
||
|
||
Messages:
|
||
{"type": "unavailable"} — no device connected
|
||
{"type": "sample", "info": {...}, "sample": {...}} — new reading
|
||
"""
|
||
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
|
||
except WebSocketDisconnect:
|
||
return
|
||
except Exception:
|
||
await ws.close()
|
||
return
|
||
|
||
wv = _state["wireview"]
|
||
wv["clients"].add(ws)
|
||
try:
|
||
# Send the current state immediately so the client does not have to
|
||
# wait for the next poll tick.
|
||
if wv["device"] is not None and wv["last_sample"] is not None:
|
||
await ws.send_json(
|
||
{"type": "sample", "info": wv["info"], "sample": wv["last_sample"]}
|
||
)
|
||
else:
|
||
await ws.send_json({"type": "unavailable"})
|
||
|
||
while True:
|
||
await ws.receive_text()
|
||
except WebSocketDisconnect:
|
||
pass
|
||
except Exception:
|
||
log.debug("wireview ws client error", exc_info=True)
|
||
finally:
|
||
wv["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,
|
||
)
|