security: harden web server, daemon socket, and write paths #9
No files matched your search
@@ -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://<host>: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`.
|
||||
|
||||
+23
-1
@@ -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://<host>: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) |
|
||||
|
||||
+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,
|
||||
)
|
||||
@@ -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())
|
||||
Reference in new issue
Block a user