fix(audit): security, correctness, refactors, docs (no wire change) #306

Merged
TapTap merged 44 commits from fix/audit-cycle into dev 2026-09-21 22:23:33 +02:00
2 changed files with 105 additions and 23 deletions
Showing only changes of commit 379f127859 - Show all commits
+55 -12
View File
@@ -20,6 +20,12 @@ CLIENT_CMD = [os.path.join(BUILD_DIR, "client")]
_WORKER = os.environ.get("PYTEST_XDIST_WORKER") _WORKER = os.environ.get("PYTEST_XDIST_WORKER")
TEST_DATA_DIR = os.path.join(PROJECT_ROOT, f"test_data-{_WORKER}" if _WORKER else "test_data") TEST_DATA_DIR = os.path.join(PROJECT_ROOT, f"test_data-{_WORKER}" if _WORKER else "test_data")
# Default wall-clock budget for a short-lived client invocation. Every client
# is expected to finish well within this; the bound exists so a hung client
# fails the test instead of stalling the whole CI run indefinitely. Callers
# that legitimately need longer can pass an explicit ``timeout``.
CLIENT_TIMEOUT = 180
class ServerManager: class ServerManager:
"""Manages a long-lived server process. Reuses across test cases.""" """Manages a long-lived server process. Reuses across test cases."""
@@ -137,7 +143,40 @@ class CountingProxy:
return result return result
def run_client(source_dir, dest_dir, flags=None, port=None, extra_args=None): def _run_client_cmd(cmd, timeout):
"""Run one client command, returning ``(result, duration)``.
On timeout the client is killed and a result-like ``CompletedProcess`` with
a non-zero returncode is returned instead of raising, so callers keep the
established ``(result, duration)`` contract and the failure carries the
command plus whatever output was captured for diagnosis.
"""
start = time.monotonic()
try:
result = subprocess.run(cmd, text=True, capture_output=True, timeout=timeout)
except subprocess.TimeoutExpired as exc:
duration = time.monotonic() - start
stdout = exc.stdout or ""
stderr = exc.stderr or ""
if isinstance(stdout, bytes):
stdout = stdout.decode(errors="replace")
if isinstance(stderr, bytes):
stderr = stderr.decode(errors="replace")
diagnostic = (
f"client timed out after {timeout}s\n"
f"command: {cmd!r}\n"
f"--- captured stdout ---\n{stdout}\n"
f"--- captured stderr ---\n{stderr}"
)
result = subprocess.CompletedProcess(cmd, returncode=-1,
stdout=stdout, stderr=diagnostic)
return result, duration
duration = time.monotonic() - start
return result, duration
def run_client(source_dir, dest_dir, flags=None, port=None, extra_args=None,
timeout=CLIENT_TIMEOUT):
"""Run the client and return (result, duration).""" """Run the client and return (result, duration)."""
cmd = CLIENT_CMD + ["--source-dir", source_dir, "--dest-dir", dest_dir, "--save-to-disk"] cmd = CLIENT_CMD + ["--source-dir", source_dir, "--dest-dir", dest_dir, "--save-to-disk"]
if port: if port:
@@ -146,23 +185,18 @@ def run_client(source_dir, dest_dir, flags=None, port=None, extra_args=None):
cmd += flags cmd += flags
if extra_args: if extra_args:
cmd += extra_args cmd += extra_args
start = time.monotonic() return _run_client_cmd(cmd, timeout)
result = subprocess.run(cmd, text=True, capture_output=True)
duration = time.monotonic() - start
return result, duration
def run_client_posix(source_dir, dest_dir, flags=None, port=None): def run_client_posix(source_dir, dest_dir, flags=None, port=None,
timeout=CLIENT_TIMEOUT):
"""Run the client with positional args (rsync-style).""" """Run the client with positional args (rsync-style)."""
cmd = CLIENT_CMD + [source_dir, dest_dir, "--save-to-disk"] cmd = CLIENT_CMD + [source_dir, dest_dir, "--save-to-disk"]
if port: if port:
cmd += ["--server-port", str(port)] cmd += ["--server-port", str(port)]
if flags: if flags:
cmd += flags cmd += flags
start = time.monotonic() return _run_client_cmd(cmd, timeout)
result = subprocess.run(cmd, text=True, capture_output=True)
duration = time.monotonic() - start
return result, duration
def generate_test_files(source_dir, full=False): def generate_test_files(source_dir, full=False):
@@ -238,8 +272,17 @@ def make_result(name, success, duration=None, error=""):
def get_dest_received_dir(dest_dir, source_dir): def get_dest_received_dir(dest_dir, source_dir):
"""Get the path where received files land inside dest_dir.""" """Get the path where received files land inside dest_dir.
return os.path.join(dest_dir, os.path.abspath(source_dir).lstrip(os.sep))
FastSync mirrors the absolute source path below the receive root with the
leading root separator removed. Strip that separator explicitly rather
than with ``str.lstrip(os.sep)``: ``lstrip`` removes a *set* of characters
rather than a path prefix, which is not the same operation.
"""
abs_source = os.path.abspath(source_dir)
if abs_source.startswith(os.sep):
abs_source = abs_source[len(os.sep):]
return os.path.join(dest_dir, abs_source)
def _find_free_port(): def _find_free_port():
+49 -10
View File
@@ -1,23 +1,57 @@
"""CLI validation and preflight checks.""" """CLI validation and preflight checks."""
import socket
import subprocess import subprocess
import sys import sys
import os import os
import shutil import shutil
import time
import pytest import pytest
sys.path.insert(0, os.path.dirname(__file__)) sys.path.insert(0, os.path.dirname(__file__))
from common import BUILD_DIR, CLIENT_CMD, SERVER_CMD, TEST_DATA_DIR, run_client, verify_transfer from common import (
BUILD_DIR,
CLIENT_CMD,
CLIENT_TIMEOUT,
SERVER_CMD,
TEST_DATA_DIR,
get_dest_received_dir,
run_client,
verify_transfer,
)
DEFAULT_PORT = 8080
def _port_is_listening(host, port, timeout=0.3):
"""True if something accepts a TCP connection on host:port right now."""
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
def _wait_for_listener(host, port, timeout=5.0):
"""Poll host:port until a listener accepts, or the deadline passes."""
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
if _port_is_listening(host, port):
return True
time.sleep(0.05)
return False
class TestHelp: class TestHelp:
def test_client_help(self): def test_client_help(self):
r = subprocess.run(CLIENT_CMD + ["--help"], capture_output=True, text=True) r = subprocess.run(CLIENT_CMD + ["--help"], capture_output=True,
text=True, timeout=CLIENT_TIMEOUT)
assert r.returncode == 0 assert r.returncode == 0
assert "Usage:" in r.stdout assert "Usage:" in r.stdout
assert "SSH transport" in r.stdout assert "SSH transport" in r.stdout
def test_server_help(self): def test_server_help(self):
r = subprocess.run(SERVER_CMD + ["--help"], capture_output=True, text=True) r = subprocess.run(SERVER_CMD + ["--help"], capture_output=True,
text=True, timeout=CLIENT_TIMEOUT)
assert r.returncode == 0 assert r.returncode == 0
assert "Usage:" in r.stdout assert "Usage:" in r.stdout
@@ -67,19 +101,24 @@ class TestServerPort:
def test_default_port(self): def test_default_port(self):
"""Server should start on default port 8080.""" """Server should start on default port 8080."""
if _port_is_listening("127.0.0.1", DEFAULT_PORT):
pytest.skip(f"port {DEFAULT_PORT} already in use by another process")
proc = subprocess.Popen( proc = subprocess.Popen(
SERVER_CMD, stdout=subprocess.DEVNULL, stderr=None, SERVER_CMD, stdout=subprocess.DEVNULL, stderr=None,
) )
try: try:
import socket, time if not _wait_for_listener("127.0.0.1", DEFAULT_PORT, timeout=5.0):
time.sleep(0.5) if proc.poll() is not None and _port_is_listening("127.0.0.1", DEFAULT_PORT):
with socket.create_connection(("127.0.0.1", 8080), timeout=2): pytest.skip(f"port {DEFAULT_PORT} was taken by another process")
pass # Port is listening pytest.fail(f"Server not listening on default port {DEFAULT_PORT}")
except (ConnectionRefusedError, OSError):
pytest.fail("Server not listening on default port 8080")
finally: finally:
proc.terminate() proc.terminate()
try:
proc.wait(timeout=5) proc.wait(timeout=5)
except subprocess.TimeoutExpired:
proc.kill()
proc.wait()
def _seed_protocol_source(source): def _seed_protocol_source(source):
@@ -105,7 +144,7 @@ class TestProtocol:
port=shared_server.port) port=shared_server.port)
assert result.returncode == 0, \ assert result.returncode == 0, \
f"--protocol current run failed: {(result.stderr or result.stdout)[:400]}" f"--protocol current run failed: {(result.stderr or result.stdout)[:400]}"
received = os.path.join(dest, os.path.abspath(source).lstrip(os.sep)) received = get_dest_received_dir(dest, source)
mismatches, missing = verify_transfer(source, received) mismatches, missing = verify_transfer(source, received)
assert not mismatches and not missing, \ assert not mismatches and not missing, \
f"transfer mismatch: missing={missing} mismatches={mismatches}" f"transfer mismatch: missing={missing} mismatches={mismatches}"