diff --git a/README.md b/README.md index 650afe6..3588f5f 100644 --- a/README.md +++ b/README.md @@ -78,6 +78,20 @@ sudo nvcurve user remove alice # remove a user Adding the first user enables authentication immediately; removing the last user disables it. See the [Usage Guide](docs/Usage-Guide.md#authentication-multi-user) for details. +## TLS (HTTPS) + +The server speaks **plain HTTP by default**. For network access (e.g. behind a reverse proxy or on a LAN), you can enable TLS so the web UI, API, and WebSocket all run over HTTPS — the session cookie is then marked `Secure`. + +```bash +# One-off (this server run only) +nvcurve serve start --ssl-certfile /path/to/cert.pem --ssl-keyfile /path/to/key.pem + +# Persistent (stored in /etc/nvcurve/config.json; used by the daemon too) +sudo nvcurve service configure --ssl-certfile /path/to/cert.pem --ssl-keyfile /path/to/key.pem +``` + +With TLS enabled the UI is at `https://:8042` and the CLI switches to `https://` automatically. A self-signed certificate works for local use (the browser will warn); for multi-user setups use a certificate your browser trusts (e.g. via your internal CA or a reverse proxy). + ## Systemd Service Install the daemon for automatic profile loading on boot and optional web server auto-start: @@ -149,6 +163,10 @@ The daemon reads settings from `/etc/nvcurve/config.json`: "max_delta_khz": 3000000, "auto_snapshot": true, "max_snapshots": 20, + "ssl_certfile": null, + "ssl_keyfile": null, + "trusted_proxies": [], + "allow_api_shutdown": true, "auto_load_profiles": { "idx:0": "my_profile" } @@ -160,9 +178,12 @@ The daemon reads settings from `/etc/nvcurve/config.json`: | `host` | Web server bind address (`0.0.0.0` for network access) | | `port` | Web server port (default `8042`) | | `auto_serve` | Auto-start web server on boot | -| `max_delta_khz` | Safety cap for frequency offsets (default 3000 MHz) | +| `max_delta_khz` | Safety cap for frequency offsets (default 3000 MHz). Enforced server-side; API clients cannot raise it per request | | `auto_snapshot` | Save snapshot before every write | | `max_snapshots` | Max snapshots to keep (`0` = unlimited) | +| `ssl_certfile` / `ssl_keyfile` | TLS certificate/key — enables HTTPS when both are set (default: off) | +| `trusted_proxies` | Proxy IPs whose `X-Forwarded-For` is trusted for the login lockout (e.g. `["127.0.0.1"]` for a local reverse proxy) | +| `allow_api_shutdown` | Allow authenticated users to stop the server via `POST /api/shutdown` (set `false` on shared systems; use systemd instead) | | `auto_load_profiles` | Per-GPU profile to apply on boot (`{gpu_key: profile_name}`) | The GPU key can be a UUID, `pci:XXXX`, or `idx:N` fallback. Find your GPU key with `nvcurve gpus`. diff --git a/docs/Usage-Guide.md b/docs/Usage-Guide.md index 657a11b..8d4f998 100644 --- a/docs/Usage-Guide.md +++ b/docs/Usage-Guide.md @@ -145,6 +145,27 @@ Adding the first user **switches the server into authenticated mode immediately* - The user store file should stay root-owned and `0600` (the CLI enforces this). - The web UI and API are still only as safe as the network path to the server — bind to a trusted interface (`--host`) and/or firewall the port. Authentication protects against casual access, not a determined network attacker. - The `nvcurve user` commands and the user store require root; day-to-day sign-in does not. +- The login lockout is keyed by client IP. Behind a reverse proxy all clients share the proxy's IP — set `trusted_proxies` in `/etc/nvcurve/config.json` (e.g. `["127.0.0.1"]`) so the lockout uses the real client IP from `X-Forwarded-For`. The header is only honoured for peers you list there (it is spoofable otherwise). +- On shared systems consider setting `allow_api_shutdown: false` so users cannot stop the server via the API (manage it with systemd instead). +- The frequency safety cap (`max_delta_khz`) is enforced by the server from its config; API clients cannot raise it per request. The CLI's `--max-delta` (root-only, direct hardware path) can still override it for a single write. + +## TLS (HTTPS) + +The server speaks **plain HTTP by default**. When you expose it beyond localhost, enable TLS so credentials and session cookies are not sent in cleartext: + +```bash +# Persistent (stored in /etc/nvcurve/config.json, used by the daemon too) +sudo nvcurve service configure --ssl-certfile /path/to/cert.pem --ssl-keyfile /path/to/key.pem + +# One-off +nvcurve serve start --ssl-certfile /path/to/cert.pem --ssl-keyfile /path/to/key.pem +``` + +- Both files must be set for TLS to activate; the UI then lives at `https://:8042` and the WebSocket upgrades to `wss://` automatically. +- To disable TLS again: `sudo nvcurve service configure --no-ssl` (removes the certificate/key from the config). +- The session cookie gets the `Secure` flag, so it is only sent over HTTPS. +- A self-signed certificate is fine for a home LAN (the browser shows a warning); for multi-user setups use a certificate your browser trusts. +- The CLI detects TLS from the config/runtime info and switches to `https://` automatically. ## CLI Reference @@ -251,13 +272,14 @@ nvcurve service uninstall sudo nvcurve service configure --auto-serve sudo nvcurve service configure --no-auto-serve sudo nvcurve service configure --host 0.0.0.0 --port 8042 +sudo nvcurve service configure --ssl-certfile /path/to/cert.pem --ssl-keyfile /path/to/key.pem ``` ## Configuration Files | File | Purpose | | --- | --- | -| `/etc/nvcurve/config.json` | Persistent config (host, port, auto-serve, default profiles) | +| `/etc/nvcurve/config.json` | Persistent config (host, port, auto-serve, TLS, safety cap, default profiles) | | `/etc/nvcurve/profiles/*.json` | Saved profiles | | `/var/cache/nvcurve/snapshots/` | Auto-saved snapshots before writes | | `/run/nvcurve.json` | Runtime server info (host, port, PID) | diff --git a/nvcurve/cli.py b/nvcurve/cli.py index 09c4090..35b2aa7 100644 --- a/nvcurve/cli.py +++ b/nvcurve/cli.py @@ -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"]) diff --git a/nvcurve/client.py b/nvcurve/client.py index bad5574..1a3a3cf 100644 --- a/nvcurve/client.py +++ b/nvcurve/client.py @@ -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") diff --git a/nvcurve/config.py b/nvcurve/config.py index b148504..733afbf 100644 --- a/nvcurve/config.py +++ b/nvcurve/config.py @@ -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 [] diff --git a/nvcurve/daemon.py b/nvcurve/daemon.py index bd81388..dd748ba 100644 --- a/nvcurve/daemon.py +++ b/nvcurve/daemon.py @@ -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() diff --git a/nvcurve/hal/snapshot.py b/nvcurve/hal/snapshot.py index a58711c..632d5fe 100644 --- a/nvcurve/hal/snapshot.py +++ b/nvcurve/hal/snapshot.py @@ -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() diff --git a/nvcurve/server.py b/nvcurve/server.py index 6ff9e36..28cb99a 100644 --- a/nvcurve/server.py +++ b/nvcurve/server.py @@ -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, + ) diff --git a/tests/test_security.py b/tests/test_security.py new file mode 100644 index 0000000..6f0bd9b --- /dev/null +++ b/tests/test_security.py @@ -0,0 +1,263 @@ +"""Security regression tests for nvcurve. + +Standalone (no pytest required): + + python tests/test_security.py + +Also works under pytest if available. Covers the security-critical logic: +SPA path containment, snapshot restore containment, login lockout client-IP +derivation, TLS scheme detection, daemon socket hardening, and the +server-enforced safety cap. +""" + +import asyncio +import builtins +import io +import os +import sys +import tempfile + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +from nvcurve import ( + daemon, # noqa: E402 + server, # noqa: E402 +) +from nvcurve.cli import _discover_server_url # noqa: E402 +from nvcurve.config import ( # noqa: E402 + Config, + normalize_trusted_proxies, + tls_enabled, +) +from nvcurve.hal.snapshot import restore as snapshot_restore # noqa: E402 + +PASS = 0 +FAIL = 0 + + +def check(name: str, cond: bool) -> None: + global PASS, FAIL + if cond: + PASS += 1 + print(f" PASS {name}") + else: + FAIL += 1 + print(f" FAIL {name}") + + +def _encoded_traversal(target: str = "/etc/hostname") -> str: + """Build an encoded '..'-based traversal path deep enough to escape any dist dir.""" + dist = os.path.abspath(server._dist_dir) + depth = len(dist.rstrip("/").split("/")) + enc = lambda s: s.replace("/", "%2f") # noqa: E731 + return "/" + enc("../*" + str(depth + 2) + target) + + +def test_spa_path_containment() -> None: + """The SPA catch-all must not serve files outside the dist directory.""" + from fastapi.testclient import TestClient + + client = TestClient(server.app) + r = client.get(_encoded_traversal()) + check("SPA traversal -> 404", r.status_code == 404) + r = client.get(_encoded_traversal("/etc/passwd")) + check("SPA traversal /etc/passwd -> 404", r.status_code == 404) + # Normal static files must still be served. + r = client.get("/index.html") + check("SPA /index.html -> 200", r.status_code == 200) + r = client.get("/api/nonexistent") + check("unknown /api/ path -> 404", r.status_code == 404) + + +def test_snapshot_restore_containment() -> None: + """Snapshot restore must only read files inside the snapshot dir.""" + tmp = tempfile.mkdtemp() + snap_dir = os.path.join(tmp, "snaps") + os.makedirs(snap_dir) + outside = os.path.join(tmp, "evil.bin") + with open(outside, "wb") as f: + f.write(b"\x00" * 9248) + + check( + "restore(outside file) rejected", + snapshot_restore(None, snap_dir, outside) is False, + ) + check( + "restore(nonexistent) rejected", + snapshot_restore(None, snap_dir, "/etc/hostname") is False, + ) + link = os.path.join(snap_dir, "link.bin") + os.symlink(outside, link) + check( + "restore(symlink escape) rejected", + snapshot_restore(None, snap_dir, link) is False, + ) + + +def test_safety_cap_not_client_overridable() -> None: + """The API must not accept a per-request safety cap override.""" + check( + "WriteRequest has no max_delta_khz field", + "max_delta_khz" not in server.WriteRequest.model_fields, + ) + check( + "GlobalOffsetRequest has no max_delta_khz field", + "max_delta_khz" not in server.GlobalOffsetRequest.model_fields, + ) + + +def test_client_ip_derivation() -> None: + """X-Forwarded-For is only honoured for configured trusted proxies.""" + + class FakeReq: + def __init__(self, peer: str, headers: dict): + self.client = type("C", (), {"host": peer})() + self.headers = headers + + check( + "trusted proxy -> XFF used", + server._client_ip( + FakeReq("127.0.0.1", {"x-forwarded-for": "9.9.9.9"}), ["127.0.0.1"] + ) + == "9.9.9.9", + ) + check( + "untrusted peer -> XFF ignored", + server._client_ip( + FakeReq("8.8.8.8", {"x-forwarded-for": "9.9.9.9"}), ["127.0.0.1"] + ) + == "8.8.8.8", + ) + check( + "rightmost untrusted hop", + server._client_ip( + FakeReq("127.0.0.1", {"x-forwarded-for": "127.0.0.1, 9.9.9.9"}), + ["127.0.0.1"], + ) + == "9.9.9.9", + ) + check( + "no XFF -> peer", + server._client_ip(FakeReq("127.0.0.1", {}), ["127.0.0.1"]) == "127.0.0.1", + ) + check( + "all-trusted chain -> peer", + server._client_ip( + FakeReq("127.0.0.1", {"x-forwarded-for": "10.0.0.1, 10.0.0.2"}), + ["127.0.0.1", "10.0.0.1", "10.0.0.2"], + ) + == "127.0.0.1", + ) + + +def test_trusted_proxies_normalization() -> None: + """String values must be normalized to lists (no substring matching).""" + check("list passthrough", normalize_trusted_proxies(["1.2.3.4"]) == ["1.2.3.4"]) + check( + "comma string split", + normalize_trusted_proxies("1.2.3.4, 5.6.7.8") == ["1.2.3.4", "5.6.7.8"], + ) + check("None -> []", normalize_trusted_proxies(None) == []) + check("junk -> []", normalize_trusted_proxies(42) == []) + # The original bug: substring membership. After normalization, "127.0.0.1" + # must NOT be trusted when only "127.0.0.10" is listed. + trusted = normalize_trusted_proxies(["127.0.0.10"]) + check("no substring trust", "127.0.0.1" not in trusted) + + +def test_tls_scheme_detection() -> None: + """_discover_server_url must pick https when TLS is configured.""" + from nvcurve import cli as cli_mod + + # Hermetic: hide any live server's runtime info file and any existing + # persistent config so the defaults level of the priority chain is + # exercised. + real_info_file = cli_mod._SERVER_INFO_FILE + real_persistent_cfg = cli_mod._PERSISTENT_CONFIG_FILE + hidden = tempfile.mkdtemp() + cli_mod._SERVER_INFO_FILE = os.path.join(hidden, "nvcurve.json") + cli_mod._PERSISTENT_CONFIG_FILE = os.path.join(hidden, "config.json") + try: + cfg = Config( + host="10.0.0.5", port=9000, ssl_certfile="/x/c.pem", ssl_keyfile="/x/k.pem" + ) + check("tls_enabled true", tls_enabled(cfg) is True) + check("https url", _discover_server_url(cfg) == "https://10.0.0.5:9000") + cfg2 = Config(host="10.0.0.5", port=9000) + check("http url", _discover_server_url(cfg2) == "http://10.0.0.5:9000") + finally: + cli_mod._SERVER_INFO_FILE = real_info_file + cli_mod._PERSISTENT_CONFIG_FILE = real_persistent_cfg + + +def test_daemon_ignores_caller_host_port() -> None: + """serve_start must bind the configured host/port, never the caller's.""" + daemon._cfg = Config(host="10.1.1.1", port=9999) + captured: dict = {} + + class FakePopen: + def __init__(self, cmd, **kw): + captured["cmd"] = cmd + self.pid = 4242 + + def poll(self): + return 0 + + def terminate(self): + pass + + def wait(self, *a): + return 0 + + real_open = builtins.open + + def fake_open(path, *a, **kw): + if str(path).endswith("nvcurve-server.log"): + return io.StringIO() + return real_open(path, *a, **kw) + + orig_popen = daemon.subprocess.Popen + daemon.subprocess.Popen = FakePopen + builtins.open = fake_open + try: + loop = asyncio.new_event_loop() + resp = loop.run_until_complete( + # "0.0.0.0" is a test payload proving the daemon ignores caller + # host/port — no socket is bound here. + daemon._dispatch({"cmd": "serve_start", "host": "0.0.0.0", "port": 12345}) # noqa: S104 + ) + finally: + builtins.open = real_open + daemon.subprocess.Popen = orig_popen + + cmd = " ".join(captured.get("cmd", [])) + check("config host/port used", "10.1.1.1" in cmd and "9999" in cmd) + check("caller host/port ignored", "0.0.0.0" not in cmd and "12345" not in cmd) # noqa: S104 + check( + "response reports configured values", + resp.get("ok") is True + and resp.get("host") == "10.1.1.1" + and resp.get("port") == 9999, + ) + check("response includes tls flag", "tls" in resp) + + +def main() -> int: + tests = [ + test_spa_path_containment, + test_snapshot_restore_containment, + test_safety_cap_not_client_overridable, + test_client_ip_derivation, + test_trusted_proxies_normalization, + test_tls_scheme_detection, + test_daemon_ignores_caller_host_port, + ] + for t in tests: + print(f"== {t.__name__} ==") + t() + print(f"\n{PASS} passed, {FAIL} failed") + return 1 if FAIL else 0 + + +if __name__ == "__main__": + sys.exit(main())