Files
FastSync/tests/integration/common.py
T

318 lines
11 KiB
Python

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):
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
self._proc = subprocess.Popen(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
_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()