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.
 
 
 
 
 
 

133 lines
3.9 KiB

#!/usr/bin/env python3
"""TCP test client for tcp_proxy recv→send integration test.
Connects to host:port with SO_MARK=1, sends 1MB random data,
receives 4-byte sentinel + 1MB response, verifies.
Supports --count N for multiple sequential requests.
"""
import os
import socket
import struct
import sys
import time
SO_MARK = 36
ONE_MB = 1048576
SENTINEL = 0xBEEF0102
def recv_exact(sock: socket.socket, n: int) -> bytes:
data = bytearray()
while len(data) < n:
chunk = sock.recv(min(65536, n - len(data)))
if not chunk:
break
data += chunk
return bytes(data)
def send_exact(sock: socket.socket, data: bytes) -> None:
sent = 0
while sent < len(data):
n = sock.send(data[sent:])
if n <= 0:
raise OSError("send failed")
sent += n
def run_one(args, payload: bytes, idx: int) -> int:
t0 = time.monotonic()
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
if args.mark != 0:
s.setsockopt(socket.SOL_SOCKET, SO_MARK, args.mark)
s.settimeout(args.timeout)
try:
s.connect((args.host, args.port))
except OSError as e:
print(f"[FAIL] {args.name}[{idx}]: connect: {e}", flush=True)
s.close()
return 1
try:
send_exact(s, payload)
except OSError as e:
print(f"[FAIL] {args.name}[{idx}]: send: {e}", flush=True)
s.close()
return 1
try:
got = recv_exact(s, 4 + args.size)
except OSError as e:
print(f"[FAIL] {args.name}[{idx}]: recv: {e}", flush=True)
s.close()
return 1
s.close()
elapsed = time.monotonic() - t0
if len(got) != 4 + args.size:
print(f"[FAIL] {args.name}[{idx}]: len={len(got)}/{4 + args.size} in {elapsed:.3f}s", flush=True)
return 1
sentinel, = struct.unpack(">I", got[:4])
if sentinel != SENTINEL:
print(f"[FAIL] {args.name}[{idx}]: sentinel={sentinel:#010x} expected={SENTINEL:#010x}", flush=True)
return 1
received = got[4:]
if not args.verify:
throughput = (args.size * 2) / 1e6 / elapsed
print(f"[PASS] {args.name}[{idx}]: {args.size} bytes in {elapsed:.3f}s ({throughput:.2f} MB/s)", flush=True)
return 0
if payload != received:
for i in range(min(len(payload), len(received))):
if payload[i] != received[i]:
print(f"[FAIL] {args.name}[{idx}]: mismatch at {i}: "
f"sent={payload[i]:02x} recv={received[i]:02x}", flush=True)
return 1
print(f"[FAIL] {args.name}[{idx}]: size mismatch", flush=True)
return 1
throughput = (args.size * 2) / 1e6 / elapsed
print(f"[PASS] {args.name}[{idx}]: {args.size} bytes in {elapsed:.3f}s ({throughput:.2f} MB/s)", flush=True)
return 0
def main():
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--host", required=True)
parser.add_argument("--port", type=int, required=True)
parser.add_argument("--size", type=int, default=ONE_MB)
parser.add_argument("--verify", action="store_true")
parser.add_argument("--timeout", type=float, default=60.0)
parser.add_argument("--mark", type=int, default=1, help="SO_MARK (0=off)")
parser.add_argument("--count", type=int, default=1, help="number of sequential requests")
parser.add_argument("--name", default="test")
parser.add_argument("--seed", type=int, default=0, help="random seed")
args = parser.parse_args()
import random
if args.seed:
random.seed(args.seed)
payloads = [os.urandom(args.size) for _ in range(args.count)]
fails = 0
for i in range(args.count):
if run_one(args, payloads[i], i + 1) != 0:
fails += 1
time.sleep(0.1)
if fails > 0:
print(f"[FAIL] {args.name}: {fails}/{args.count} failed", flush=True)
sys.exit(1)
print(f"[PASS] {args.name}: all {args.count} ok", flush=True)
if __name__ == "__main__":
main()