security: harden web server, daemon socket, and write paths
Security review findings, fixed and verified: Critical - Fix unauthenticated arbitrary file read: the SPA catch-all route joined the raw URL path onto the dist dir without containment, so encoded '..' segments (/%2e%2e/etc/passwd) leaked any file readable by the root server. Resolve with realpath and reject paths outside the dist dir (fail-closed 404). High - Daemon socket: serve_start no longer accepts caller-chosen host/port. The socket is world-connectable (unprivileged CLI users), so callers could previously rebind the root web server to 0.0.0.0. The daemon now always binds the operator-configured address and reports it in the response; the CLI warns on mismatch. Medium - Remove the per-request max_delta_khz override from the API: the server-enforced safety cap is now authoritative. CLI direct paths (write, profile apply, verify) honor the configured cap; --max-delta still overrides for explicit root use. - Snapshot restore: confine filepath to the snapshot directory (realpath containment; blocks symlink escapes). - Login lockout: honor X-Forwarded-For only for peers listed in the new trusted_proxies config (rightmost untrusted hop), so the per-IP lockout works behind a reverse proxy. Spoofed headers from untrusted peers are ignored. - /api/shutdown: new allow_api_shutdown config (default true); shared systems can disable the API shutdown path. TLS (opt-in, like auth) - New ssl_certfile/ssl_keyfile config + CLI flags (serve start, service install/configure, --no-ssl to disable). When active: HTTPS for UI/API, wss:// for WebSockets, Secure session cookie, CLI auto-switches to https://. Cert/key paths are validated up front with a clear error instead of a silent uvicorn crash. Tests & docs - tests/test_security.py: standalone regression tests (no new deps) covering SPA containment, snapshot containment, cap removal, client-IP derivation, proxy normalization, TLS scheme detection, and daemon host/port hardening. - README + Usage-Guide: TLS section, new config keys, updated security notes.
This commit is contained in:
1 parent
8956fc9d7b
commit
39701c12ff
9 files changed
+648
-55
No files matched your search
+151
-16
@@ -38,7 +38,7 @@ import sys
|
||||
import time
|
||||
|
||||
from .client import ApiError, NvCurveClient, ServerNotRunning
|
||||
from .config import Config, default_config
|
||||
from .config import Config, default_config, normalize_trusted_proxies, tls_enabled
|
||||
from .nvapi.constants import (
|
||||
CT_BASE,
|
||||
CT_POINTS,
|
||||
@@ -472,6 +472,21 @@ _PERSISTENT_CONFIG_FILE = (
|
||||
)
|
||||
_DAEMON_SOCKET_PATH = "/run/nvcurve-daemon.sock"
|
||||
|
||||
|
||||
def _configured_max_delta() -> int:
|
||||
"""Frequency safety cap from the persistent config (built-in default fallback).
|
||||
|
||||
The operator-configured cap is authoritative for all direct-hardware
|
||||
write paths (write, profile apply, verify) unless explicitly overridden
|
||||
with --max-delta.
|
||||
"""
|
||||
try:
|
||||
with open(_PERSISTENT_CONFIG_FILE) as f:
|
||||
return json.load(f).get("max_delta_khz", default_config.max_delta_khz)
|
||||
except (FileNotFoundError, json.JSONDecodeError, OSError):
|
||||
return default_config.max_delta_khz
|
||||
|
||||
|
||||
_ALLOWED_HOSTS = {"127.0.0.1", "::1", "localhost"}
|
||||
|
||||
|
||||
@@ -548,12 +563,18 @@ def _discover_server_url(cfg: Config) -> str:
|
||||
1. /run/nvcurve.json — runtime info written by the running server process
|
||||
2. /etc/nvcurve/config.json — persistent config written by `service install`
|
||||
3. Config defaults — 127.0.0.1:8042
|
||||
|
||||
The scheme is https when TLS is configured (or reported by the running
|
||||
server), http otherwise.
|
||||
"""
|
||||
scheme = "https" if tls_enabled(cfg) else "http"
|
||||
|
||||
# 1. Runtime info (most accurate — reflects the actual running port)
|
||||
info = _read_server_info()
|
||||
if info:
|
||||
host = _safe_host(info["host"], cfg)
|
||||
return f"http://{host}:{info['port']}"
|
||||
s = "https" if info.get("tls") else scheme
|
||||
return f"{s}://{host}:{info['port']}"
|
||||
|
||||
# 2. Persistent config (survives reboots; written by `service install`)
|
||||
try:
|
||||
@@ -561,12 +582,12 @@ def _discover_server_url(cfg: Config) -> str:
|
||||
data = json.load(f)
|
||||
host = _safe_host(data.get("host", cfg.host), cfg)
|
||||
port = data.get("port", cfg.port)
|
||||
return f"http://{host}:{port}"
|
||||
return f"{scheme}://{host}:{port}"
|
||||
except (FileNotFoundError, json.JSONDecodeError, KeyError):
|
||||
pass
|
||||
|
||||
# 3. Hardcoded defaults
|
||||
return f"http://{cfg.host}:{cfg.port}"
|
||||
return f"{scheme}://{cfg.host}:{cfg.port}"
|
||||
|
||||
|
||||
# ── Subcommand handlers ───────────────────────────────────────────────────────
|
||||
@@ -823,7 +844,7 @@ def cmd_write(args):
|
||||
}
|
||||
|
||||
effective_max = (
|
||||
max_delta_khz if max_delta_khz is not None else default_config.max_delta_khz
|
||||
max_delta_khz if max_delta_khz is not None else _configured_max_delta()
|
||||
)
|
||||
errors = validate_write(point_deltas, effective_max)
|
||||
if errors:
|
||||
@@ -864,6 +885,7 @@ def cmd_verify(args):
|
||||
from .hal.gpu import get_gpu
|
||||
from .hal.snapshot import save as snapshot_save
|
||||
from .hal.vfcurve import read_clock_offsets, write_offsets
|
||||
from .safety import validate_write
|
||||
|
||||
try:
|
||||
delta_khz = int(args.delta * 1000)
|
||||
@@ -881,6 +903,13 @@ def cmd_verify(args):
|
||||
|
||||
point_deltas = dict.fromkeys(points, delta_khz)
|
||||
|
||||
# Enforce the operator-configured safety cap (same as cmd_write).
|
||||
errors = validate_write(point_deltas, _configured_max_delta())
|
||||
if errors:
|
||||
for e in errors:
|
||||
print(f"Error: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
gpu, gpu_name = get_gpu(index=getattr(args, "gpu_index", 0))
|
||||
|
||||
print("=== Write-Verify Cycle ===")
|
||||
@@ -1233,7 +1262,7 @@ def cmd_profile(args):
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
errors = validate_write(deltas, default_config.max_delta_khz)
|
||||
errors = validate_write(deltas, _configured_max_delta())
|
||||
if errors:
|
||||
errs.append("Curve: " + "; ".join(errors))
|
||||
else:
|
||||
@@ -1571,6 +1600,10 @@ def cmd_service(args):
|
||||
port = getattr(args, "port", 8042)
|
||||
auto_serve = getattr(args, "auto_serve", False)
|
||||
persistent_cfg.update({"host": host, "port": port, "auto_serve": auto_serve})
|
||||
if getattr(args, "ssl_certfile", None):
|
||||
persistent_cfg["ssl_certfile"] = args.ssl_certfile
|
||||
if getattr(args, "ssl_keyfile", None):
|
||||
persistent_cfg["ssl_keyfile"] = args.ssl_keyfile
|
||||
try:
|
||||
with open(_PERSISTENT_CONFIG_FILE, "w") as f:
|
||||
json.dump(persistent_cfg, f, indent=2)
|
||||
@@ -1578,11 +1611,17 @@ def cmd_service(args):
|
||||
print(f"Failed to write {_PERSISTENT_CONFIG_FILE}: {exc}", file=sys.stderr)
|
||||
return
|
||||
print(f"Persistent config written to {_PERSISTENT_CONFIG_FILE}")
|
||||
scheme = (
|
||||
"https"
|
||||
if persistent_cfg.get("ssl_certfile") and persistent_cfg.get("ssl_keyfile")
|
||||
else "http"
|
||||
)
|
||||
if auto_serve:
|
||||
print(f" Web server will auto-start on boot at {host}:{port}")
|
||||
print(f" Web server will auto-start on boot at {scheme}://{host}:{port}")
|
||||
else:
|
||||
print(
|
||||
f" Web server default: {host}:{port} (start on demand: nvcurve serve start)"
|
||||
f" Web server default: {scheme}://{host}:{port} "
|
||||
"(start on demand: nvcurve serve start)"
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1706,12 +1745,14 @@ def cmd_service(args):
|
||||
auto_serve = pcfg.get("auto_serve", False)
|
||||
host = pcfg.get("host", "127.0.0.1")
|
||||
port = pcfg.get("port", 8042)
|
||||
tls = bool(pcfg.get("ssl_certfile") and pcfg.get("ssl_keyfile"))
|
||||
print()
|
||||
print(f"web server auto-start: {'on' if auto_serve else 'off'}")
|
||||
print(f"web server address: {host}:{port}")
|
||||
print(f"web server address: {'https' if tls else 'http'}://{host}:{port}")
|
||||
print(f"web server TLS: {'on' if tls else 'off'}")
|
||||
print()
|
||||
print(
|
||||
"Change with: nvcurve service configure [--auto-serve|--no-auto-serve] [--host H] [--port P]"
|
||||
"Change with: nvcurve service configure [--auto-serve|--no-auto-serve] [--host H] [--port P] [--ssl-certfile C --ssl-keyfile K]"
|
||||
)
|
||||
|
||||
elif action == "configure":
|
||||
@@ -1736,6 +1777,13 @@ def cmd_service(args):
|
||||
pcfg["host"] = args.host
|
||||
if hasattr(args, "port") and args.port is not None:
|
||||
pcfg["port"] = args.port
|
||||
if getattr(args, "ssl_certfile", None):
|
||||
pcfg["ssl_certfile"] = args.ssl_certfile
|
||||
if getattr(args, "ssl_keyfile", None):
|
||||
pcfg["ssl_keyfile"] = args.ssl_keyfile
|
||||
if getattr(args, "no_ssl", False):
|
||||
pcfg.pop("ssl_certfile", None)
|
||||
pcfg.pop("ssl_keyfile", None)
|
||||
|
||||
try:
|
||||
with open(_PERSISTENT_CONFIG_FILE, "w") as f:
|
||||
@@ -1747,6 +1795,9 @@ def cmd_service(args):
|
||||
print(f" auto-serve: {'on' if pcfg.get('auto_serve', False) else 'off'}")
|
||||
print(f" host: {pcfg.get('host', '127.0.0.1')}")
|
||||
print(f" port: {pcfg.get('port', 8042)}")
|
||||
print(
|
||||
f" TLS: {'on' if pcfg.get('ssl_certfile') and pcfg.get('ssl_keyfile') else 'off'}"
|
||||
)
|
||||
|
||||
if os.path.exists(unit_path):
|
||||
try:
|
||||
@@ -1766,12 +1817,28 @@ def _cmd_serve_start(args, cfg: Config, open_browser: bool = False) -> None:
|
||||
host = getattr(args, "host", cfg.host)
|
||||
port = getattr(args, "port", cfg.port)
|
||||
|
||||
# Optional TLS (CLI flags override the persistent config).
|
||||
ssl_certfile = getattr(args, "ssl_certfile", None)
|
||||
ssl_keyfile = getattr(args, "ssl_keyfile", None)
|
||||
if ssl_certfile:
|
||||
cfg.ssl_certfile = ssl_certfile
|
||||
if ssl_keyfile:
|
||||
cfg.ssl_keyfile = ssl_keyfile
|
||||
|
||||
# --direct: skip daemon round-trip (used when the daemon itself spawns us).
|
||||
if getattr(args, "direct", False):
|
||||
require_root()
|
||||
try:
|
||||
with open(_SERVER_INFO_FILE, "w") as f:
|
||||
json.dump({"pid": os.getpid(), "host": host, "port": port}, f)
|
||||
json.dump(
|
||||
{
|
||||
"pid": os.getpid(),
|
||||
"host": host,
|
||||
"port": port,
|
||||
"tls": tls_enabled(cfg),
|
||||
},
|
||||
f,
|
||||
)
|
||||
except OSError as exc:
|
||||
print(f"Failed to write {_SERVER_INFO_FILE}: {exc}", file=sys.stderr)
|
||||
return
|
||||
@@ -1791,13 +1858,34 @@ def _cmd_serve_start(args, cfg: Config, open_browser: bool = False) -> None:
|
||||
return
|
||||
|
||||
# Prefer daemon socket: no root required, daemon manages the server process.
|
||||
resp = _daemon_send({"cmd": "serve_start", "host": host, "port": port})
|
||||
# The daemon always binds the configured host/port (callers cannot choose
|
||||
# the interface), so report the address from the daemon's response.
|
||||
resp = _daemon_send({"cmd": "serve_start"})
|
||||
if resp is not None:
|
||||
if resp.get("ok"):
|
||||
print(f"Web server starting (PID {resp['pid']}) at http://{host}:{port}")
|
||||
rhost = resp.get("host", host)
|
||||
rport = resp.get("port", port)
|
||||
if (rhost, rport) != (host, port):
|
||||
print(
|
||||
f"Note: daemon uses the configured bind address {rhost}:{rport} "
|
||||
"(change with: nvcurve service configure --host/--port)",
|
||||
file=sys.stderr,
|
||||
)
|
||||
if (ssl_certfile or ssl_keyfile) and not resp.get("tls"):
|
||||
print(
|
||||
"Note: --ssl-certfile/--ssl-keyfile are ignored while the daemon "
|
||||
"manages the server — the daemon uses the TLS settings from "
|
||||
"/etc/nvcurve/config.json (set with: nvcurve service configure "
|
||||
"--ssl-certfile/--ssl-keyfile)",
|
||||
file=sys.stderr,
|
||||
)
|
||||
scheme = "https" if resp.get("tls") else "http"
|
||||
print(
|
||||
f"Web server starting (PID {resp['pid']}) at {scheme}://{rhost}:{rport}"
|
||||
)
|
||||
if open_browser:
|
||||
time.sleep(1.5)
|
||||
_open_browser_as_user(f"http://{host}:{port}")
|
||||
_open_browser_as_user(f"{scheme}://{rhost}:{rport}")
|
||||
else:
|
||||
print(f"Daemon: {resp.get('error')}", file=sys.stderr)
|
||||
return
|
||||
@@ -1807,7 +1895,8 @@ def _cmd_serve_start(args, cfg: Config, open_browser: bool = False) -> None:
|
||||
|
||||
info = _read_server_info()
|
||||
if info:
|
||||
url = f"http://{info['host']}:{info['port']}"
|
||||
scheme = "https" if info.get("tls") else "http"
|
||||
url = f"{scheme}://{info['host']}:{info['port']}"
|
||||
print(f"Server is already running (PID {info['pid']}) at {url}.")
|
||||
if open_browser:
|
||||
_open_browser_as_user(url)
|
||||
@@ -1827,6 +1916,10 @@ def _cmd_serve_start(args, cfg: Config, open_browser: bool = False) -> None:
|
||||
"--port",
|
||||
str(port),
|
||||
]
|
||||
if ssl_certfile:
|
||||
cmd += ["--ssl-certfile", ssl_certfile]
|
||||
if ssl_keyfile:
|
||||
cmd += ["--ssl-keyfile", ssl_keyfile]
|
||||
if getattr(args, "gpu_index", 0):
|
||||
cmd += ["--gpu", str(args.gpu_index)]
|
||||
log_path = _log_file()
|
||||
@@ -1846,7 +1939,15 @@ def _cmd_serve_start(args, cfg: Config, open_browser: bool = False) -> None:
|
||||
# Foreground mode — write info file so clients can discover host:port.
|
||||
try:
|
||||
with open(_SERVER_INFO_FILE, "w") as f:
|
||||
json.dump({"pid": os.getpid(), "host": host, "port": port}, f)
|
||||
json.dump(
|
||||
{
|
||||
"pid": os.getpid(),
|
||||
"host": host,
|
||||
"port": port,
|
||||
"tls": tls_enabled(cfg),
|
||||
},
|
||||
f,
|
||||
)
|
||||
except OSError as exc:
|
||||
print(f"Failed to write {_SERVER_INFO_FILE}: {exc}", file=sys.stderr)
|
||||
return
|
||||
@@ -2056,6 +2157,12 @@ Examples:
|
||||
"--host", default="127.0.0.1", help="Bind address (default 127.0.0.1)"
|
||||
)
|
||||
p_start.add_argument("--port", type=int, default=8042, help="Port (default 8042)")
|
||||
p_start.add_argument(
|
||||
"--ssl-certfile", default=None, help="TLS certificate (enables HTTPS)"
|
||||
)
|
||||
p_start.add_argument(
|
||||
"--ssl-keyfile", default=None, help="TLS private key (enables HTTPS)"
|
||||
)
|
||||
p_start.add_argument(
|
||||
"--detach", "-d", action="store_true", help="Run in background"
|
||||
)
|
||||
@@ -2090,6 +2197,16 @@ Examples:
|
||||
default=8042,
|
||||
help="Default web server port (stored in config)",
|
||||
)
|
||||
p_install.add_argument(
|
||||
"--ssl-certfile",
|
||||
default=None,
|
||||
help="TLS certificate (stored in config; enables HTTPS)",
|
||||
)
|
||||
p_install.add_argument(
|
||||
"--ssl-keyfile",
|
||||
default=None,
|
||||
help="TLS private key (stored in config; enables HTTPS)",
|
||||
)
|
||||
|
||||
p_configure = s_svc.add_parser(
|
||||
"configure", help="Update config and restart daemon (escalates to root)"
|
||||
@@ -2109,6 +2226,19 @@ Examples:
|
||||
)
|
||||
p_configure.add_argument("--host", default=None, help="Web server bind address")
|
||||
p_configure.add_argument("--port", type=int, default=None, help="Web server port")
|
||||
p_configure.add_argument(
|
||||
"--ssl-certfile",
|
||||
default=None,
|
||||
help="TLS certificate (stored in config; enables HTTPS)",
|
||||
)
|
||||
p_configure.add_argument(
|
||||
"--ssl-keyfile", default=None, help="TLS private key (stored in config)"
|
||||
)
|
||||
p_configure.add_argument(
|
||||
"--no-ssl",
|
||||
action="store_true",
|
||||
help="Disable TLS (remove certificate/key from config)",
|
||||
)
|
||||
|
||||
s_svc.add_parser("uninstall", help="Remove systemd service (escalates to root)")
|
||||
s_svc.add_parser("start", help="Start systemd service (escalates to root)")
|
||||
@@ -2149,9 +2279,14 @@ def main():
|
||||
"users_file",
|
||||
"host",
|
||||
"port",
|
||||
"ssl_certfile",
|
||||
"ssl_keyfile",
|
||||
"allow_api_shutdown",
|
||||
):
|
||||
if key in data:
|
||||
setattr(cfg, key, data[key])
|
||||
if "trusted_proxies" in data:
|
||||
cfg.trusted_proxies = normalize_trusted_proxies(data["trusted_proxies"])
|
||||
if "auto_load_profiles" in data:
|
||||
# Keys are stable GPU identifiers (UUID, "pci:XXXX", or "idx:N")
|
||||
cfg.auto_load_profiles = dict(data["auto_load_profiles"])
|
||||
|
||||
+5
-10
@@ -120,18 +120,13 @@ class NvCurveClient:
|
||||
def write_curve(
|
||||
self,
|
||||
deltas: dict[int, int],
|
||||
max_delta_khz: int | None = None,
|
||||
) -> dict:
|
||||
body: dict = {"deltas": deltas}
|
||||
if max_delta_khz is not None:
|
||||
body["max_delta_khz"] = max_delta_khz
|
||||
return self._post("/api/curve/write", body)
|
||||
# The server enforces its configured safety cap; clients cannot
|
||||
# override it per request.
|
||||
return self._post("/api/curve/write", {"deltas": deltas})
|
||||
|
||||
def write_global(self, delta_khz: int, max_delta_khz: int | None = None) -> dict:
|
||||
body: dict = {"delta_khz": delta_khz}
|
||||
if max_delta_khz is not None:
|
||||
body["max_delta_khz"] = max_delta_khz
|
||||
return self._post("/api/curve/write/global", body)
|
||||
def write_global(self, delta_khz: int) -> dict:
|
||||
return self._post("/api/curve/write/global", {"delta_khz": delta_khz})
|
||||
|
||||
def reset_curve(self) -> dict:
|
||||
return self._post("/api/curve/reset")
|
||||
|
||||
@@ -17,6 +17,21 @@ class Config:
|
||||
host: str = "127.0.0.1"
|
||||
port: int = 8042
|
||||
|
||||
# Optional TLS: when both are set, the server serves HTTPS and the
|
||||
# session cookie is marked Secure. Off by default (plain HTTP).
|
||||
ssl_certfile: str | None = None
|
||||
ssl_keyfile: str | None = None
|
||||
|
||||
# Proxy IPs (e.g. a reverse proxy on 127.0.0.1) whose X-Forwarded-For
|
||||
# header is trusted for the login brute-force lockout. JSON array in
|
||||
# config.json (a comma-separated string is also accepted and normalized).
|
||||
# Without this, all proxied clients share the proxy's IP.
|
||||
trusted_proxies: list[str] = field(default_factory=list)
|
||||
|
||||
# Allow any authenticated user to stop the server via POST /api/shutdown.
|
||||
# Set false on shared systems; manage the service via systemd instead.
|
||||
allow_api_shutdown: bool = True
|
||||
|
||||
snapshot_dir: str = "/var/cache/nvcurve/snapshots"
|
||||
profile_dir: str = "/etc/nvcurve/profiles"
|
||||
|
||||
@@ -42,3 +57,25 @@ class Config:
|
||||
|
||||
# Module-level default config instance.
|
||||
default_config = Config()
|
||||
|
||||
|
||||
def tls_enabled(cfg: Config) -> bool:
|
||||
"""True when both TLS files are configured (server serves HTTPS)."""
|
||||
return bool(cfg.ssl_certfile and cfg.ssl_keyfile)
|
||||
|
||||
|
||||
def normalize_trusted_proxies(value) -> list[str]:
|
||||
"""Normalize a trusted_proxies config value to a list of IP strings.
|
||||
|
||||
Accepts a JSON array (the documented format) or a comma-separated string
|
||||
(tolerated for convenience). Normalizing matters because the server does
|
||||
exact list membership tests — a raw string would degrade to substring
|
||||
matching (e.g. "127.0.0.1" in "127.0.0.10").
|
||||
"""
|
||||
if value is None:
|
||||
return []
|
||||
if isinstance(value, str):
|
||||
return [h.strip() for h in value.split(",") if h.strip()]
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [str(h).strip() for h in value if str(h).strip()]
|
||||
return []
|
||||
+32
-9
@@ -7,10 +7,16 @@ Protocol: newline-delimited JSON, one request → one response, connection close
|
||||
|
||||
Commands:
|
||||
{"cmd": "ping"}
|
||||
{"cmd": "serve_start", "host": "127.0.0.1", "port": 8042}
|
||||
{"cmd": "serve_start"}
|
||||
{"cmd": "serve_stop"}
|
||||
{"cmd": "serve_status"}
|
||||
|
||||
The socket is world-connectable (unprivileged users drive it via the CLI),
|
||||
so the command surface is deliberately minimal: serve_start ALWAYS binds the
|
||||
configured host/port from /etc/nvcurve/config.json — callers cannot choose
|
||||
the bind address (no ad-hoc 0.0.0.0 exposure). Changing the bind address is
|
||||
an operator action via `nvcurve service configure`.
|
||||
|
||||
Requires root.
|
||||
"""
|
||||
|
||||
@@ -23,7 +29,7 @@ import signal
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from .config import Config
|
||||
from .config import Config, normalize_trusted_proxies
|
||||
|
||||
log = logging.getLogger("nvcurve.daemon")
|
||||
|
||||
@@ -38,7 +44,13 @@ _cfg: Config | None = None # Config instance, set in run()
|
||||
# ── Socket command handlers ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def _handle_serve_start(host: str, port: int) -> dict:
|
||||
async def _handle_serve_start() -> dict:
|
||||
"""Start the web server on the *configured* host/port.
|
||||
|
||||
The bind address is taken from /etc/nvcurve/config.json only — the
|
||||
socket is reachable by unprivileged users, so callers must not be able
|
||||
to choose the interface (e.g. binding 0.0.0.0 to expose the API).
|
||||
"""
|
||||
global _server_proc
|
||||
if _server_proc is not None and _server_proc.poll() is None:
|
||||
return {
|
||||
@@ -47,6 +59,8 @@ async def _handle_serve_start(host: str, port: int) -> dict:
|
||||
"pid": _server_proc.pid,
|
||||
}
|
||||
|
||||
host = _cfg.host if _cfg is not None else "127.0.0.1"
|
||||
port = _cfg.port if _cfg is not None else 8042
|
||||
cmd = [
|
||||
sys.executable,
|
||||
"-m",
|
||||
@@ -72,7 +86,13 @@ async def _handle_serve_start(host: str, port: int) -> dict:
|
||||
except OSError as exc:
|
||||
return {"ok": False, "error": f"cannot open log file {log_path}: {exc}"}
|
||||
log.info("Web server started (PID %d)", _server_proc.pid)
|
||||
return {"ok": True, "pid": _server_proc.pid}
|
||||
return {
|
||||
"ok": True,
|
||||
"pid": _server_proc.pid,
|
||||
"host": host,
|
||||
"port": port,
|
||||
"tls": bool(_cfg and _cfg.ssl_certfile and _cfg.ssl_keyfile),
|
||||
}
|
||||
|
||||
|
||||
async def _handle_serve_stop() -> dict:
|
||||
@@ -105,9 +125,7 @@ async def _dispatch(req: dict) -> dict:
|
||||
elif cmd == "serve_start":
|
||||
if _cfg is None:
|
||||
return {"ok": False, "error": "config not initialized"}
|
||||
host = req.get("host", _cfg.host)
|
||||
port = req.get("port", _cfg.port)
|
||||
return await _handle_serve_start(host, port)
|
||||
return await _handle_serve_start()
|
||||
elif cmd == "serve_stop":
|
||||
return await _handle_serve_stop()
|
||||
elif cmd == "serve_status":
|
||||
@@ -174,9 +192,13 @@ def run() -> None:
|
||||
"profile_dir",
|
||||
"host",
|
||||
"port",
|
||||
"ssl_certfile",
|
||||
"ssl_keyfile",
|
||||
):
|
||||
if key in cfg_data:
|
||||
setattr(_cfg, key, cfg_data[key])
|
||||
if "trusted_proxies" in cfg_data:
|
||||
_cfg.trusted_proxies = normalize_trusted_proxies(cfg_data["trusted_proxies"])
|
||||
|
||||
# Apply auto-load profiles in a subprocess so the daemon process itself
|
||||
# never loads NvAPI/NVML/HAL modules — keeps steady-state RSS low.
|
||||
@@ -207,7 +229,8 @@ async def _serve_socket(auto_serve: bool = False) -> None:
|
||||
server = await asyncio.start_unix_server(_handle_client, path=SOCKET_PATH)
|
||||
# The socket must be connectable by unprivileged users: the CLI runs as the
|
||||
# regular user and talks to this root daemon over the socket. 0o666 is
|
||||
# intentional (standard for /run daemon sockets).
|
||||
# intentional — the command surface is restricted accordingly (serve_start
|
||||
# always uses the configured host/port; see module docstring).
|
||||
# pi-lens-ignore: S103
|
||||
_SOCKET_MODE = 0o666
|
||||
os.chmod(
|
||||
@@ -220,7 +243,7 @@ async def _serve_socket(auto_serve: bool = False) -> None:
|
||||
log.warning("auto_serve requested but config not initialized")
|
||||
else:
|
||||
log.info("auto_serve enabled — starting web server on boot")
|
||||
await _handle_serve_start(_cfg.host, _cfg.port)
|
||||
await _handle_serve_start()
|
||||
|
||||
stop_event = asyncio.Event()
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
@@ -121,6 +121,15 @@ def restore(gpu, snapshot_dir: str, filepath: str | None = None) -> bool:
|
||||
print(f"Snapshot file not found: {filepath}")
|
||||
return False
|
||||
|
||||
# Contain the path inside the snapshot directory — callers (in
|
||||
# particular the HTTP API) must not be able to point the restore at
|
||||
# arbitrary files on the filesystem.
|
||||
snap_dir = os.path.realpath(snapshot_dir)
|
||||
resolved = os.path.realpath(filepath)
|
||||
if not resolved.startswith(snap_dir + os.sep):
|
||||
print(f"Snapshot path outside snapshot directory: {filepath}")
|
||||
return False
|
||||
|
||||
try:
|
||||
with open(filepath, "rb") as f:
|
||||
raw = f.read()
|
||||
|
||||
+106
-18
@@ -11,7 +11,7 @@ import logging
|
||||
import os
|
||||
from contextlib import asynccontextmanager, suppress
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, WebSocket, WebSocketDisconnect
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
@@ -20,7 +20,7 @@ from fastapi.staticfiles import StaticFiles
|
||||
from pydantic import BaseModel
|
||||
|
||||
from . import auth
|
||||
from .config import Config, default_config
|
||||
from .config import Config, default_config, tls_enabled
|
||||
from .hal.dashboard import get_dashboard_info
|
||||
from .hal.fans import (
|
||||
get_fan_info,
|
||||
@@ -552,12 +552,10 @@ app.add_middleware(AuthMiddleware)
|
||||
|
||||
class WriteRequest(BaseModel):
|
||||
deltas: dict[int, int] # {point_index: delta_kHz}
|
||||
max_delta_khz: int | None = None # per-request safety limit override
|
||||
|
||||
|
||||
class GlobalOffsetRequest(BaseModel):
|
||||
delta_khz: int
|
||||
max_delta_khz: int | None = None # per-request safety limit override
|
||||
|
||||
|
||||
class VerifyRequest(BaseModel):
|
||||
@@ -654,7 +652,7 @@ async def api_auth_login(req: LoginRequest, request: Request):
|
||||
if not users:
|
||||
raise HTTPException(status_code=404, detail="Authentication is not enabled")
|
||||
|
||||
client_ip = request.client.host if request.client else "unknown"
|
||||
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."
|
||||
@@ -677,6 +675,7 @@ async def api_auth_login(req: LoginRequest, request: Request):
|
||||
max_age=auth.SESSION_TTL_S,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
secure=tls_enabled(cfg),
|
||||
path="/",
|
||||
)
|
||||
return response
|
||||
@@ -710,6 +709,36 @@ def _require_gpu(gpu_index: int = 0):
|
||||
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 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -1407,10 +1436,10 @@ async def api_curve_write(req: WriteRequest, gpu_index: int = 0):
|
||||
|
||||
vfp_state, _ = await _run(read_curve, gpu, g_state["gpu_name"])
|
||||
|
||||
effective_limit = (
|
||||
req.max_delta_khz if req.max_delta_khz is not None else cfg.max_delta_khz
|
||||
)
|
||||
errors = validate_write(req.deltas, effective_limit)
|
||||
# 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})
|
||||
|
||||
@@ -1461,10 +1490,8 @@ async def api_curve_write_global(req: GlobalOffsetRequest, gpu_index: int = 0):
|
||||
raise HTTPException(status_code=500, detail="Failed to read curve")
|
||||
|
||||
all_deltas = {p.index: req.delta_khz for p in vfp_state.points if p.domain == "gpu"}
|
||||
effective_limit = (
|
||||
req.max_delta_khz if req.max_delta_khz is not None else cfg.max_delta_khz
|
||||
)
|
||||
errors = validate_write(all_deltas, effective_limit)
|
||||
# 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})
|
||||
|
||||
@@ -1605,10 +1632,21 @@ async def api_curve_verify(req: VerifyRequest, gpu_index: int = 0):
|
||||
|
||||
@app.post("/api/shutdown")
|
||||
async def api_shutdown():
|
||||
"""Gracefully shut down the server process."""
|
||||
"""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}
|
||||
@@ -1795,7 +1833,15 @@ async def serve_spa(catchall: str):
|
||||
if not os.path.isdir(_dist_dir):
|
||||
return {"error": "Frontend not built. Run pnpm build in frontend/."}
|
||||
|
||||
path = os.path.join(_dist_dir, catchall)
|
||||
# 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)
|
||||
|
||||
@@ -1821,8 +1867,14 @@ def run(
|
||||
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."""
|
||||
"""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
|
||||
|
||||
@@ -1830,6 +1882,34 @@ def run(
|
||||
|
||||
_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.
|
||||
@@ -1847,7 +1927,7 @@ def run(
|
||||
)
|
||||
return
|
||||
|
||||
url = f"http://{host}:{port}"
|
||||
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
|
||||
@@ -1862,4 +1942,12 @@ def run(
|
||||
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)
|
||||
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,
|
||||
)
|
||||
Reference in new issue
Block a user