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

#!/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()