import filecmp import os import random import shutil import socket import subprocess import sys import tempfile import threading import time PROJECT_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) BUILD_DIR = os.path.join(PROJECT_ROOT, "build") SERVER_CMD = [os.path.join(BUILD_DIR, "server")] CLIENT_CMD = [os.path.join(BUILD_DIR, "client")] # Under pytest-xdist each worker process gets its own PYTEST_XDIST_WORKER id # ('gw0', 'gw1', ...). Worker-key the transient working dir so concurrent # workers on the shared filesystem never collide on fixtures. Outside xdist # (or with -n1) this stays the historical 'test_data' path. _WORKER = os.environ.get("PYTEST_XDIST_WORKER") 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: """Manages a long-lived server process. Reuses across test cases.""" def __init__(self): self._proc = None self._port = None def start(self, extra_args=None, env=None): self.stop() self._port = _find_free_port() # Plain TCP is intentionally explicit in the server; integration tests # exercise that opt-in mode rather than relying on the secure default. cmd = SERVER_CMD + ["-p", str(self._port), "--allow-unauthenticated"] if extra_args: cmd += extra_args proc_env = dict(os.environ) if env: proc_env.update(env) self._proc = subprocess.Popen(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, env=proc_env) _wait_for_port(self._port, timeout=5) def stop(self): if self._proc: _wait_proc(self._proc) self._proc = None @property def port(self): return self._port def __enter__(self): self.start() return self def __exit__(self, *args): self.stop() def __del__(self): self.stop() class CountingProxy: """One-shot TCP forwarder that counts the bytes flowing in each direction between one client and the real server. Client output and --stats report SOURCE lengths, so a delta/fuzzy transfer that moves only a few percent of the file is invisible in normal output. Routing the client through this proxy makes the actual wire usage observable: client_to_server counts every byte the client sent (config, paths, and file/delta payloads), server_to_client counts the reply bytes (including the receiver's delta signatures). """ def __init__(self, target_port): self.target_port = target_port self._listener = socket.socket() self._listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) self._listener.bind(("127.0.0.1", 0)) self._listener.listen(1) self._listener.settimeout(30) self.port = self._listener.getsockname()[1] self.client_to_server = 0 self.server_to_client = 0 @staticmethod def _pump(src, dst, counter): while True: try: data = src.recv(65536) except OSError: return if not data: try: dst.shutdown(socket.SHUT_WR) except OSError: pass return try: dst.sendall(data) except OSError: return counter[0] += len(data) def run(self, cmd, join_timeout=20): """Forward one client run (the full command list) to the real server and return the CompletedProcess after the counts have settled. ``join_timeout`` bounds how long to wait for the forwarding threads. The client->server count is published as soon as the client side reaches EOF (i.e. once the client process has exited), so callers that only need that count can pass a small value instead of waiting for the server to close its idle socket.""" def serve(): try: client_sock, _ = self._listener.accept() server_sock = socket.create_connection(("127.0.0.1", self.target_port), timeout=10) except OSError: self._listener.close() return c2s, s2c = [0], [0] a = threading.Thread(target=self._pump, args=(client_sock, server_sock, c2s)) b = threading.Thread(target=self._pump, args=(server_sock, client_sock, s2c)) a.start() b.start() a.join() self.client_to_server = c2s[0] b.join() self.server_to_client = s2c[0] self._listener.close() thread = threading.Thread(target=serve, daemon=True) thread.start() result = subprocess.run(cmd, capture_output=True, text=True, timeout=180) thread.join(join_timeout) return result 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).""" cmd = CLIENT_CMD + ["--source-dir", source_dir, "--dest-dir", dest_dir, "--save-to-disk"] if port: cmd += ["--server-port", str(port)] if flags: cmd += flags if extra_args: cmd += extra_args return _run_client_cmd(cmd, timeout) def run_client_posix(source_dir, dest_dir, flags=None, port=None, timeout=CLIENT_TIMEOUT): """Run the client with positional args (rsync-style).""" cmd = CLIENT_CMD + [source_dir, dest_dir, "--save-to-disk"] if port: cmd += ["--server-port", str(port)] if flags: cmd += flags return _run_client_cmd(cmd, timeout) def generate_test_files(source_dir, full=False): """Generate structured test data. Returns total bytes written.""" if os.path.exists(source_dir): shutil.rmtree(source_dir) os.makedirs(source_dir) target_total = 25 * 1024 * 1024 if full else 0 written = 0 files = { "small.txt": b"hello world\n", "medium.txt": b"the quick brown fox jumps over the lazy dog\n" * 5000, "binary.bin": bytes(range(256)) * 1000, "nested/subdir/deep.txt": b"deeply nested file\n", "nested/another.txt": b"another nested file\n" * 50, } for rel_path, content in files.items(): full_path = os.path.join(source_dir, rel_path) os.makedirs(os.path.dirname(full_path), exist_ok=True) with open(full_path, "wb") as f: f.write(content) written += len(content) if full: os.makedirs(os.path.join(source_dir, "bulk"), exist_ok=True) i = 0 while written < target_total: chunk_size = min(5 * 1024 * 1024, target_total - written) with open(os.path.join(source_dir, f"bulk/file_{i}.dat"), "wb") as f: f.write(random.randbytes(chunk_size)) written += chunk_size i += 1 return written def verify_transfer(source_dir, received_dir): """Verify all files from source exist in received_dir and match. Returns (mismatches, missing).""" source_dir = os.path.abspath(source_dir) received_dir = os.path.abspath(received_dir) if not os.path.exists(received_dir): return [], ["no received files found"] mismatches, missing = [], [] for root, dirs, files in os.walk(source_dir): for f in files: src_path = os.path.join(root, f) rel = os.path.relpath(src_path, source_dir) dst_path = os.path.join(received_dir, rel) if not os.path.exists(dst_path): missing.append(rel) elif not filecmp.cmp(src_path, dst_path, shallow=False): mismatches.append(rel) return mismatches, missing def clean_dir(path): """Remove and recreate a directory.""" if os.path.exists(path): shutil.rmtree(path) os.makedirs(path, exist_ok=True) def make_result(name, success, duration=None, error=""): """Create a standardized result dict.""" return { "name": name, "status": "Success" if success else "Failed", "time": f"{duration:.4f}s" if duration is not None else "N/A", "error": error, } def get_dest_received_dir(dest_dir, source_dir): """Get the path where received files land inside dest_dir. 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(): with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: s.bind(("", 0)) return s.getsockname()[1] def _wait_for_port(port, timeout=5): deadline = time.monotonic() + timeout while time.monotonic() < deadline: try: with socket.create_connection(("127.0.0.1", port), timeout=0.3): return except (ConnectionRefusedError, OSError): time.sleep(0.05) raise RuntimeError(f"Server port {port} not ready after {timeout}s") def _wait_proc(proc, timeout=5): """Stop a long-lived subprocess promptly. The server installs a SIGTERM handler, so signal first and only escalate to SIGKILL if it does not exit; waiting without signalling would burn the full timeout on every stop.""" if proc.poll() is not None: proc.wait() return proc.terminate() try: proc.wait(timeout=timeout) except subprocess.TimeoutExpired: proc.kill() proc.wait()