You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
339 lines
11 KiB
339 lines
11 KiB
#!/usr/bin/env python3 |
|
"""Burst/load peer for tcp_proxy integration test (full-duplex, byte-exact). |
|
|
|
Both sides stream a known pseudo-random byte sequence in random chunks |
|
(10B..1MB) with optional random idle (100..500ms) between chunks, over full-duplex |
|
TCP connections, and verify the received stream byte-for-byte. |
|
|
|
Because the sequence is deterministic (self-contained xorshift64* generator), |
|
every received byte is checked against a pre-known value — any corruption, |
|
cross-connection mixing or stall is detected exactly. |
|
|
|
Multiple concurrent connections are supported (--connections N); each connection |
|
carries its index in a 4-byte handshake so both peers derive distinct per-connection |
|
seeds and any cross-connection byte mixing is caught. |
|
|
|
Roles: |
|
server: burst_peer.py --server --host 127.0.0.1 --port 19193 --connections 8 |
|
client: burst_peer.py --host 10.200.100.1 --port 19193 --mark 1 --connections 8 |
|
|
|
Exit code: 0 = PASS (all connections byte-exact), 1 = FAIL. |
|
""" |
|
|
|
import argparse |
|
import math |
|
import random |
|
import select |
|
import signal |
|
import socket |
|
import struct |
|
import sys |
|
import threading |
|
import time |
|
|
|
SO_MARK = 36 |
|
MASK64 = (1 << 64) - 1 |
|
|
|
|
|
class ByteStream: |
|
"""Deterministic xorshift64* byte stream (version-independent).""" |
|
|
|
def __init__(self, seed): |
|
self.s = (seed ^ 0x9E3779B97F4A7C15) & MASK64 |
|
if self.s == 0: |
|
self.s = 0x123456789ABCDEF |
|
|
|
def next_u64(self): |
|
x = self.s |
|
x ^= (x >> 12) & MASK64 |
|
x ^= (x << 25) & MASK64 |
|
x ^= (x >> 27) & MASK64 |
|
self.s = x |
|
return (x * 0x2545F4914F6CDD1D) & MASK64 |
|
|
|
def next_bytes(self, n): |
|
out = bytearray() |
|
while len(out) < n: |
|
out += self.next_u64().to_bytes(8, "little") |
|
return bytes(out[:n]) |
|
|
|
|
|
def rand_chunk(rnd, lo, hi): |
|
"""Log-uniform chunk size in [lo, hi] — covers small (keep-alive) and large.""" |
|
if hi <= lo: |
|
return lo |
|
v = math.exp(rnd.uniform(math.log(lo), math.log(hi))) |
|
return max(lo, min(hi, int(round(v)))) |
|
|
|
|
|
class StuckError(Exception): |
|
pass |
|
|
|
|
|
def send_all(sock, data): |
|
view = memoryview(data) |
|
while len(view): |
|
n = sock.send(view) |
|
if n <= 0: |
|
raise RuntimeError("send failed (connection closed)") |
|
view = view[n:] |
|
|
|
|
|
def recv_exact(sock, n): |
|
out = bytearray() |
|
while len(out) < n: |
|
c = sock.recv(n - len(out)) |
|
if not c: |
|
raise RuntimeError("connection closed") |
|
out += c |
|
return bytes(out) |
|
|
|
|
|
def receive_session(sock, recv_seed, total, stall_ms, name): |
|
"""Read frames until END, verify byte-for-byte. Returns (bytes, frames).""" |
|
rng = ByteStream(recv_seed) |
|
received = 0 |
|
nframes = 0 |
|
last_progress = time.monotonic() |
|
last_log = time.monotonic() |
|
|
|
def read_n(n): |
|
nonlocal last_progress |
|
out = bytearray() |
|
while len(out) < n: |
|
r, _, _ = select.select([sock], [], [], 0.1) |
|
if r: |
|
c = sock.recv(n - len(out)) |
|
if not c: |
|
raise RuntimeError("connection closed at %d/%d bytes" % (received, total)) |
|
out += c |
|
last_progress = time.monotonic() |
|
else: |
|
if time.monotonic() - last_progress >= stall_ms: |
|
raise StuckError("no progress for %dms (received %d/%d)" % |
|
(stall_ms, received, total)) |
|
return bytes(out) |
|
|
|
while True: |
|
hdr = read_n(4) |
|
(length,) = struct.unpack(">I", hdr) |
|
if length == 0: |
|
if received != total: |
|
raise RuntimeError("peer ended early: %d/%d" % (received, total)) |
|
return received, nframes |
|
if received + length > total: |
|
raise RuntimeError("frame overflow: %d+%d > %d" % (received, length, total)) |
|
payload = read_n(length) |
|
expected = rng.next_bytes(length) |
|
if payload != expected: |
|
off = 0 |
|
for i in range(length): |
|
if payload[i] != expected[i]: |
|
off = i |
|
break |
|
raise RuntimeError("byte mismatch at offset %d: got %02x want %02x" % |
|
(received + off, payload[off], expected[off])) |
|
received += length |
|
nframes += 1 |
|
if time.monotonic() - last_log >= 1.0: |
|
print("[%s] recv %d/%d bytes" % (name, received, total), flush=True) |
|
last_log = time.monotonic() |
|
|
|
|
|
def send_session(sock, send_seed, total, min_chunk, max_chunk, min_idle, max_idle): |
|
"""Send `total` deterministic bytes in random chunks with random idle, then END.""" |
|
rng = ByteStream(send_seed) |
|
rnd = random.Random(send_seed ^ 0xABCDEF) |
|
remaining = total |
|
nframes = 0 |
|
while remaining > 0: |
|
length = rand_chunk(rnd, min_chunk, max_chunk) |
|
if length > remaining: |
|
length = remaining |
|
if length < 1: |
|
length = 1 |
|
payload = rng.next_bytes(length) |
|
send_all(sock, struct.pack(">I", length) + payload) |
|
remaining -= length |
|
nframes += 1 |
|
if remaining > 0: |
|
time.sleep(rnd.uniform(min_idle, max_idle) / 1000.0) |
|
send_all(sock, struct.pack(">I", 0)) |
|
return nframes, total |
|
|
|
|
|
def run_peer(sock, args, role, name, idx): |
|
"""Run one full-duplex session. Returns (ok, recv_bytes, send_bytes, err).""" |
|
if role == "client": |
|
seed_out, seed_in = args.seed_c2s ^ idx, args.seed_s2c ^ idx |
|
else: |
|
seed_out, seed_in = args.seed_s2c ^ idx, args.seed_c2s ^ idx |
|
|
|
sender_res = {} |
|
|
|
def sender(): |
|
try: |
|
sender_res["frames"], sender_res["bytes"] = send_session( |
|
sock, seed_out, args.total, |
|
args.min_chunk, args.max_chunk, |
|
args.min_idle_ms, args.max_idle_ms) |
|
sender_res["ok"] = True |
|
except Exception as e: |
|
sender_res["ok"] = False |
|
sender_res["err"] = str(e) |
|
|
|
t = threading.Thread(target=sender, daemon=True) |
|
t.start() |
|
|
|
try: |
|
recv, _ = receive_session(sock, seed_in, args.total, args.stall_ms, name) |
|
except StuckError as e: |
|
return False, 0, 0, "STUCK: %s" % e |
|
except Exception as e: |
|
return False, 0, 0, str(e) |
|
|
|
t.join() |
|
if not sender_res.get("ok"): |
|
return False, recv, 0, "send error: %s" % sender_res.get("err") |
|
|
|
return True, recv, sender_res["bytes"], None |
|
|
|
|
|
def main(): |
|
p = argparse.ArgumentParser() |
|
p.add_argument("--server", action="store_true") |
|
p.add_argument("--host", required=True) |
|
p.add_argument("--port", type=int, required=True) |
|
p.add_argument("--name", default="burst") |
|
p.add_argument("--connections", type=int, default=1) |
|
p.add_argument("--total", type=int, default=4 * 1024 * 1024) |
|
p.add_argument("--min-chunk", type=int, default=10) |
|
p.add_argument("--max-chunk", type=int, default=1048576) |
|
p.add_argument("--min-idle-ms", type=int, default=100) |
|
p.add_argument("--max-idle-ms", type=int, default=500) |
|
p.add_argument("--stall-ms", type=int, default=2000) |
|
p.add_argument("--mark", type=int, default=1) |
|
p.add_argument("--seed-c2s", type=lambda x: int(x, 0), default=0x11443322778855A) |
|
p.add_argument("--seed-s2c", type=lambda x: int(x, 0), default=0x2233445511223344) |
|
args = p.parse_args() |
|
|
|
if args.server: |
|
run_server(args) |
|
else: |
|
run_client(args) |
|
|
|
|
|
def run_client(args): |
|
results = [] |
|
lock = threading.Lock() |
|
t0 = time.monotonic() |
|
|
|
def worker(i): |
|
name = "%s[%d]" % (args.name, i) |
|
s = None |
|
try: |
|
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) |
|
if args.mark: |
|
s.setsockopt(socket.SOL_SOCKET, SO_MARK, args.mark) |
|
s.connect((args.host, args.port)) |
|
s.sendall(struct.pack(">I", i)) |
|
ok, recv, send, err = run_peer(s, args, "client", name, i) |
|
if ok: |
|
print("[PASS] %s: recv %d sent %d" % (name, recv, send), flush=True) |
|
else: |
|
print("[FAIL] %s: %s" % (name, err), flush=True) |
|
with lock: |
|
results.append((ok, recv, send)) |
|
except OSError as e: |
|
print("[FAIL] %s: connect: %s" % (name, e), flush=True) |
|
with lock: |
|
results.append((False, 0, 0)) |
|
finally: |
|
if s: |
|
s.close() |
|
|
|
threads = [threading.Thread(target=worker, args=(i,)) for i in range(args.connections)] |
|
for t in threads: |
|
t.start() |
|
for t in threads: |
|
t.join() |
|
|
|
dt = time.monotonic() - t0 |
|
ok = len(results) == args.connections and all(r[0] for r in results) |
|
total_bytes = sum(r[1] + r[2] for r in results) |
|
if ok: |
|
print("[PASS] %s: %d/%d connections, %d bytes in %.2fs (%.2f MB/s)" % |
|
(args.name, len(results), args.connections, total_bytes, dt, |
|
total_bytes / 1e6 / dt if dt > 0 else 0), flush=True) |
|
sys.exit(0) |
|
failed = args.connections - sum(1 for r in results if r[0]) |
|
print("[FAIL] %s: %d/%d connections failed" % (args.name, failed, args.connections), flush=True) |
|
sys.exit(1) |
|
|
|
|
|
def run_server(args): |
|
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) |
|
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) |
|
s.bind((args.host, args.port)) |
|
s.listen(256) |
|
print("BURST: listening on %s:%d connections=%d total=%d" % |
|
(args.host, args.port, args.connections, args.total), flush=True) |
|
|
|
results = [] |
|
lock = threading.Lock() |
|
running = [True] |
|
|
|
def _stop(sig, frame): |
|
running[0] = False |
|
|
|
signal.signal(signal.SIGTERM, _stop) |
|
signal.signal(signal.SIGINT, _stop) |
|
|
|
def handle(conn): |
|
idx = 0 |
|
try: |
|
idxb = recv_exact(conn, 4) |
|
(idx,) = struct.unpack(">I", idxb) |
|
name = "%s[%d]" % (args.name, idx) |
|
ok, recv, send, err = run_peer(conn, args, "server", name, idx) |
|
if ok: |
|
print("[PASS] %s: recv %d sent %d" % (name, recv, send), flush=True) |
|
else: |
|
print("[FAIL] %s: %s" % (name, err), flush=True) |
|
with lock: |
|
results.append(ok) |
|
except OSError as e: |
|
print("[FAIL] %s: %s" % (args.name, e), flush=True) |
|
with lock: |
|
results.append(False) |
|
finally: |
|
conn.close() |
|
|
|
s.settimeout(1.0) |
|
threads = [] |
|
while running[0] and len(threads) < args.connections: |
|
try: |
|
conn, addr = s.accept() |
|
except socket.timeout: |
|
continue |
|
except OSError: |
|
break |
|
t = threading.Thread(target=handle, args=(conn,), daemon=True) |
|
t.start() |
|
threads.append(t) |
|
|
|
for t in threads: |
|
t.join() |
|
s.close() |
|
|
|
ok = len(results) == args.connections and all(results) |
|
if ok: |
|
print("[PASS] %s: %d/%d connections" % (args.name, len(results), args.connections), flush=True) |
|
sys.exit(0) |
|
print("[FAIL] %s: %d/%d connections ok" % |
|
(args.name, sum(1 for r in results if r), args.connections), flush=True) |
|
sys.exit(1) |
|
|
|
|
|
if __name__ == "__main__": |
|
main()
|
|
|