"""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())