3proxy/tests/harness.py
Vladimir Dubrovin 7011e78ece Add TLS tests: a wrapped proxy, a TLS chain, and MITM
Three arrangements, with key material generated for the run rather than
kept in the tree: a proxy wrapped in TLS, a proxy that reaches a TLS parent
and verifies it against the CA, and MITM.

The MITM case checks what interception is for: the decrypted request line,
URI and all, reaches the log, where the same request through a plain
CONNECT tunnel leaves only the host and port.

The origin runs in its own process there so the proxy log holds only what
the proxy saw, and log assertions wait, since a record is written when the
connection finishes rather than when the reply arrives.

Verification of the spoofed certificate is deliberately not strict: 3proxy
issues those without an Authority Key Identifier, which Python rejects
under its 3.13 defaults.
2026-08-25 23:01:45 +03:00

631 lines
24 KiB
Python

"""Support code for the 3proxy regression tests.
Everything here is standard library, so the suite runs wherever 3proxy
builds: no shell, no curl, no netcat.
A test case is a module under tests/cases/ exporting run(t). It writes the
configurations it needs, starts them, and states what it expects:
def run(t):
srv = t.free_port()
t.start("echo", f'''
log
auth iponly
allow *
http * /echo echo
httpsrv -p{srv}
''', ports=[srv])
r = t.http(f"http://127.0.0.1:{srv}/echo")
t.eq(200, r.status, "the server answers")
"""
import base64
import http.client
import os
import shutil
import socket
import ssl
import struct
import subprocess
import sys
import textwrap
import time
class Response:
"""A reply, or the reason there wasn't one."""
def __init__(self, status=None, body=b"", headers=None, error=None):
self.status = status
self.body = body
self.headers = headers or {}
self.error = error
@property
def text(self):
return self.body.decode("utf-8", "replace")
@property
def length(self):
return len(self.body)
def header(self, name):
for k, v in self.headers.items():
if k.lower() == name.lower():
return v
return None
def __repr__(self):
if self.error:
return f"<no reply: {self.error}>"
return f"<{self.status}, {len(self.body)} bytes>"
class Server:
"""A running 3proxy, with the configuration it was given."""
def __init__(self, name, path, proc, logfile):
self.name = name
self.path = path
self.proc = proc
self.logfile = logfile
def output(self):
try:
with open(self.logfile, "rb") as fp:
return fp.read().decode("utf-8", "replace")
except OSError:
return ""
def stop(self):
if self.proc.poll() is None:
self.proc.terminate()
try:
self.proc.wait(timeout=5)
except subprocess.TimeoutExpired:
self.proc.kill()
self.proc.wait(timeout=5)
class Certs:
"""A test CA, a certificate it signed, and somewhere to cache spoofed ones.
Paths use forward slashes: they are written into configurations read by
3proxy, and ssl_certcache insists on a trailing separator.
"""
def __init__(self, directory):
self.dir = directory.replace("\\", "/")
self.ca = self.dir + "/ca.pem"
self.ca_key = self.dir + "/ca.key"
self.server = self.dir + "/server.pem"
self.server_key = self.dir + "/server.key"
# a second CA nothing is signed by, for the cases that must fail
self.other = self.dir + "/other.pem"
self.other_key = self.dir + "/other.key"
self.cache = self.dir + "/cache/"
class Failure(Exception):
"""Raised when a case cannot go on, e.g. a server refused to start."""
class Tester:
"""The API a case runs against: start servers, make requests, assert."""
def __init__(self, binary, tmpdir, case):
self.binary = binary
self.tmpdir = tmpdir
self.case = case
self.servers = []
self.checks = []
self.timeout = 10
self._skipped = 0
self._certs = None
# ---- servers -----------------------------------------------------
def free_port(self):
"""A port nothing is listening on. Closed again before it is used,
which is racy in principle and reliable enough in practice."""
s = socket.socket()
try:
s.bind(("127.0.0.1", 0))
return s.getsockname()[1]
finally:
s.close()
def write_config(self, name, config):
path = os.path.join(self.tmpdir, name + ".cfg")
text = textwrap.dedent(config).strip() + "\n"
# newline="" keeps the line endings as written, rather than letting
# Windows turn them into CRLF behind the parser's back
with open(path, "w", newline="") as fp:
fp.write(text)
return path
def start(self, name, config, ports=()):
"""Write a configuration, run it, and wait for its ports to open."""
path = self.write_config(name, config)
logfile = os.path.join(self.tmpdir, name + ".out")
with open(logfile, "wb") as out:
proc = subprocess.Popen([self.binary, path], stdout=out,
stderr=subprocess.STDOUT)
server = Server(name, path, proc, logfile)
self.servers.append(server)
for port in ports:
if not self.wait_port(port):
code = proc.poll()
if code is None:
died = "the process is still running"
else:
died = f"the process exited with code {code}"
if os.name == "nt" and code is not None and code & 0xFFFFFFFF == 0xC0000135:
died += " (a DLL it needs was not found)"
raise Failure(
f"{name} never listened on port {port}: {died}\n"
f"--- configuration ---\n{open(path).read()}"
f"--- output ---\n{server.output()}")
return server
def run_config(self, name, config):
"""Run a configuration expected to be rejected; return its output."""
path = self.write_config(name, config)
done = subprocess.run([self.binary, path], stdout=subprocess.PIPE,
stderr=subprocess.STDOUT, timeout=15)
return done.stdout.decode("utf-8", "replace")
def wait_port(self, port, timeout=5.0):
deadline = time.time() + timeout
while time.time() < deadline:
try:
with socket.create_connection(("127.0.0.1", port), 0.25):
return True
except OSError:
time.sleep(0.02)
return False
def wait_output(self, server, needle, timeout=5.0, since=0):
"""Wait for a server to log something.
A record is written when the connection it describes finishes, not
when the reply reaches the client, so reading straight after a
request usually finds nothing yet.
"""
deadline = time.time() + timeout
while True:
text = server.output()[since:]
if needle in text or time.time() > deadline:
return text
time.sleep(0.05)
def stop_all(self):
for server in self.servers:
server.stop()
self.servers = []
# ---- requests ----------------------------------------------------
def http(self, url, proxy=None, socks=None, socks4=False,
remote_dns=False, method="GET", body=None, headers=None,
auth=None, proxy_auth=None, tunnel=False, conn=None):
"""Make a request, directly or through a proxy, and read the reply.
proxy "host:port" of an HTTP proxy
socks "host:port" of a SOCKS proxy
tunnel reach the origin with CONNECT rather than an absolute URI
conn reuse a connection returned by connection()
"""
host, port, path = self._split(url)
headers = dict(headers or {})
if auth:
headers["Authorization"] = self._basic(auth)
if proxy_auth:
headers["Proxy-Authorization"] = self._basic(proxy_auth)
own = conn is None
try:
if own:
conn = self.connection(host, port, proxy=proxy, socks=socks,
socks4=socks4, remote_dns=remote_dns,
tunnel=tunnel)
target = path
if proxy and not tunnel:
target = f"http://{host}:{port}{path}"
if body is not None and not isinstance(body, bytes):
body = body.encode()
conn.request(method, target, body=body, headers=headers)
reply = conn.getresponse()
data = reply.read()
return Response(reply.status, data, dict(reply.getheaders()))
except (OSError, http.client.HTTPException) as exc:
return Response(error=f"{type(exc).__name__}: {exc}")
finally:
if own and conn is not None:
try:
conn.close()
except OSError:
pass
def connection(self, host, port, proxy=None, socks=None, socks4=False,
remote_dns=False, tunnel=False):
"""A connection to an origin, kept open for reuse."""
if socks:
shost, sport = self._hostport(socks)
sock = self._socks_connect(shost, sport, host, port,
socks4=socks4, remote_dns=remote_dns)
conn = http.client.HTTPConnection(host, port, timeout=self.timeout)
conn.sock = sock
return conn
if proxy:
phost, pport = self._hostport(proxy)
conn = http.client.HTTPConnection(phost, pport, timeout=self.timeout)
if tunnel:
conn.set_tunnel(host, port)
return conn
return http.client.HTTPConnection(host, port, timeout=self.timeout)
def raw(self, port, request, host="127.0.0.1"):
"""Send bytes as they are and return whatever comes back."""
if not isinstance(request, bytes):
request = request.encode("latin-1")
try:
with socket.create_connection((host, port), self.timeout) as sock:
sock.settimeout(self.timeout)
sock.sendall(request)
chunks = []
while True:
try:
piece = sock.recv(65536)
except OSError:
# a timeout, or a reset once the server is done:
# either way keep whatever already arrived
break
if not piece:
break
chunks.append(piece)
return b"".join(chunks).decode("utf-8", "replace")
except OSError as exc:
return f"<no reply: {exc}>"
# ---- SOCKS -------------------------------------------------------
def _socks_connect(self, shost, sport, host, port, socks4=False,
remote_dns=False, auth=None):
sock = socket.create_connection((shost, sport), self.timeout)
sock.settimeout(self.timeout)
try:
if socks4:
addr = socket.inet_aton(socket.gethostbyname(host))
sock.sendall(b"\x04\x01" + struct.pack("!H", port) + addr + b"\x00")
reply = self._recvall(sock, 8)
if len(reply) < 2 or reply[1] != 0x5a:
raise OSError("SOCKS4 request refused")
return sock
if auth:
sock.sendall(b"\x05\x02\x00\x02")
else:
sock.sendall(b"\x05\x01\x00")
reply = self._recvall(sock, 2)
if len(reply) < 2 or reply[0] != 5:
raise OSError("SOCKS5 handshake failed")
if reply[1] == 0x02:
if not auth:
raise OSError("SOCKS5 server demands credentials")
user, password = auth
sock.sendall(b"\x01" + bytes([len(user)]) + user.encode() +
bytes([len(password)]) + password.encode())
status = self._recvall(sock, 2)
if len(status) < 2 or status[1] != 0:
raise OSError("SOCKS5 credentials refused")
elif reply[1] != 0x00:
raise OSError("SOCKS5 offered no acceptable method")
if remote_dns:
target = b"\x03" + bytes([len(host)]) + host.encode()
else:
target = b"\x01" + socket.inet_aton(socket.gethostbyname(host))
sock.sendall(b"\x05\x01\x00" + target + struct.pack("!H", port))
reply = self._recvall(sock, 4)
if len(reply) < 4 or reply[1] != 0:
raise OSError("SOCKS5 request refused")
self._read_socks_addr(sock, reply[3])
return sock
except Exception:
sock.close()
raise
def socks_connect(self, socks, host, port, socks4=False, remote_dns=False,
auth=None):
"""Open a SOCKS connection, reporting failure rather than raising."""
shost, sport = self._hostport(socks)
try:
sock = self._socks_connect(shost, sport, host, port, socks4=socks4,
remote_dns=remote_dns, auth=auth)
sock.close()
return None
except OSError as exc:
return str(exc)
def socks_http(self, socks, url, auth=None, **kwargs):
"""A request through SOCKS, with optional SOCKS credentials."""
host, port, path = self._split(url)
shost, sport = self._hostport(socks)
try:
sock = self._socks_connect(shost, sport, host, port, auth=auth,
**kwargs)
except OSError as exc:
return Response(error=str(exc))
conn = http.client.HTTPConnection(host, port, timeout=self.timeout)
conn.sock = sock
try:
conn.request("GET", path)
reply = conn.getresponse()
return Response(reply.status, reply.read(), dict(reply.getheaders()))
except (OSError, http.client.HTTPException) as exc:
return Response(error=str(exc))
finally:
conn.close()
def socks_udp_associate(self, port, host="127.0.0.1"):
"""Ask for a UDP association and report the port handed back.
That socket is allocated per association, which is where an intport
range has to take effect.
"""
try:
with socket.create_connection((host, port), self.timeout) as sock:
sock.settimeout(self.timeout)
sock.sendall(b"\x05\x01\x00")
if self._recvall(sock, 2) != b"\x05\x00":
return None
sock.sendall(b"\x05\x03\x00\x01\x00\x00\x00\x00" +
struct.pack("!H", 0))
reply = self._recvall(sock, 4)
if len(reply) < 4 or reply[1] != 0:
return None
_, bound = self._read_socks_addr(sock, reply[3])
return bound
except OSError:
return None
def _read_socks_addr(self, sock, atyp):
if atyp == 1:
addr = socket.inet_ntoa(self._recvall(sock, 4))
elif atyp == 3:
length = self._recvall(sock, 1)[0]
addr = self._recvall(sock, length).decode()
elif atyp == 4:
addr = self._recvall(sock, 16).hex()
else:
raise OSError(f"unknown SOCKS address type {atyp}")
port = struct.unpack("!H", self._recvall(sock, 2))[0]
return addr, port
@staticmethod
def _recvall(sock, count):
data = b""
while len(data) < count:
piece = sock.recv(count - len(data))
if not piece:
break
data += piece
return data
# ---- TLS ---------------------------------------------------------
def certs(self):
"""A CA and a certificate for 127.0.0.1, generated once per run.
Returns None when openssl is unavailable, so a case can skip rather
than fail on a machine that cannot make key material.
"""
if self._certs is not None:
return self._certs or None
if not shutil.which("openssl"):
self._certs = False
return None
c = Certs(os.path.join(self.tmpdir, "certs"))
os.makedirs(c.cache, exist_ok=True)
csr = c.dir + "/server.csr"
ext = c.dir + "/server.ext"
with open(ext, "w") as fp:
fp.write("subjectAltName=IP:127.0.0.1,DNS:localhost\n")
# OpenSSL 3 refuses to trust a CA without these extensions
ca_ext = ["-addext", "basicConstraints=critical,CA:TRUE",
"-addext", "keyUsage=critical,keyCertSign,cRLSign"]
steps = [
["openssl", "genrsa", "-out", c.ca_key, "2048"],
["openssl", "req", "-x509", "-new", "-nodes", "-key", c.ca_key,
"-sha256", "-days", "3650", "-subj", "/CN=3proxy-test-ca",
"-out", c.ca] + ca_ext,
["openssl", "genrsa", "-out", c.other_key, "2048"],
["openssl", "req", "-x509", "-new", "-nodes", "-key", c.other_key,
"-sha256", "-days", "3650", "-subj", "/CN=3proxy-test-other-ca",
"-out", c.other] + ca_ext,
["openssl", "genrsa", "-out", c.server_key, "2048"],
["openssl", "req", "-new", "-key", c.server_key,
"-subj", "/CN=127.0.0.1", "-out", csr],
["openssl", "x509", "-req", "-in", csr, "-CA", c.ca,
"-CAkey", c.ca_key, "-CAcreateserial", "-out", c.server,
"-days", "3650", "-sha256", "-extfile", ext],
]
for step in steps:
done = subprocess.run(step, stdout=subprocess.PIPE,
stderr=subprocess.STDOUT, timeout=60)
if done.returncode:
self._certs = False
return None
self._certs = c
return c
def _context(self, ca=None, strict=True):
"""A client context. strict=False drops the RFC 5280 checks Python
turns on by default from 3.13, which reject a certificate with no
Authority Key Identifier."""
if ca:
context = ssl.create_default_context(cafile=ca)
if not strict:
context.verify_flags &= ~getattr(ssl, "VERIFY_X509_STRICT", 0)
return context
context = ssl.create_default_context()
context.check_hostname = False
context.verify_mode = ssl.CERT_NONE
return context
def tls_proxy_http(self, proxy, url, ca=None, strict=True, method="GET",
body=None, headers=None):
"""A request to a proxy that is itself wrapped in TLS (ssl_serv)."""
host, port, path = self._split(url)
phost, pport = self._hostport(proxy)
try:
raw = socket.create_connection((phost, pport), self.timeout)
sock = self._context(ca, strict).wrap_socket(raw, server_hostname=phost)
except (OSError, ssl.SSLError) as exc:
return Response(error=f"{type(exc).__name__}: {exc}")
conn = http.client.HTTPConnection(host, port, timeout=self.timeout)
conn.sock = sock
try:
if body is not None and not isinstance(body, bytes):
body = body.encode()
conn.request(method, f"http://{host}:{port}{path}", body=body,
headers=headers or {})
reply = conn.getresponse()
return Response(reply.status, reply.read(), dict(reply.getheaders()))
except (OSError, http.client.HTTPException) as exc:
return Response(error=f"{type(exc).__name__}: {exc}")
finally:
conn.close()
def https(self, url, proxy=None, ca=None, strict=True, method="GET",
headers=None):
"""An https:// request, optionally tunnelled through a proxy."""
host, port, path = self._split(url, default_port=443)
context = self._context(ca, strict)
try:
if proxy:
phost, pport = self._hostport(proxy)
conn = http.client.HTTPSConnection(phost, pport, context=context,
timeout=self.timeout)
conn.set_tunnel(host, port)
else:
conn = http.client.HTTPSConnection(host, port, context=context,
timeout=self.timeout)
conn.request(method, path, headers=headers or {})
reply = conn.getresponse()
return Response(reply.status, reply.read(), dict(reply.getheaders()))
except (OSError, ssl.SSLError, http.client.HTTPException) as exc:
return Response(error=f"{type(exc).__name__}: {exc}")
finally:
try:
conn.close()
except (OSError, NameError, UnboundLocalError):
pass
# ---- helpers -----------------------------------------------------
@staticmethod
def _basic(credentials):
user, password = credentials
token = base64.b64encode(f"{user}:{password}".encode()).decode()
return "Basic " + token
@staticmethod
def _hostport(value):
host, _, port = value.rpartition(":")
return host or "127.0.0.1", int(port)
@staticmethod
def _split(url, default_port=80):
for prefix in ("http://", "https://"):
if url.startswith(prefix):
url = url[len(prefix):]
break
authority, _, path = url.partition("/")
if ":" in authority:
host, _, port = authority.rpartition(":")
else:
host, port = authority, default_port
return host or "127.0.0.1", int(port), "/" + path
# ---- assertions --------------------------------------------------
def _record(self, passed, label, expected=None, actual=None):
self.checks.append((passed, label, expected, actual))
return passed
def ok(self, label):
return self._record(True, label)
def fail(self, label, expected=None, actual=None):
return self._record(False, label, expected, actual)
def eq(self, expected, actual, label):
return self._record(expected == actual, label, expected, actual)
def ne(self, unexpected, actual, label):
return self._record(unexpected != actual, label,
f"anything but {unexpected!r}", actual)
@staticmethod
def _as_text(value):
"""A reply that never arrived has no text, so report the reason."""
if isinstance(value, Response):
if value.error:
return f"<no reply: {value.error}>"
if not value.body and value.status is not None:
return f"<{value.status}, empty body>"
return value.text
return value
def contains(self, haystack, needle, label):
haystack = self._as_text(haystack)
return self._record(needle in haystack, label,
f"text containing {needle!r}", self._clip(haystack))
def not_contains(self, haystack, needle, label):
haystack = self._as_text(haystack)
return self._record(needle not in haystack, label,
f"text without {needle!r}", self._clip(haystack))
def in_range(self, value, low, high, label):
good = isinstance(value, int) and low <= value <= high
return self._record(good, label, f"between {low} and {high}", value)
def not_in_range(self, value, low, high, label):
good = isinstance(value, int) and not (low <= value <= high)
return self._record(good, label, f"outside {low}-{high}", value)
def skip(self, label):
self._skipped += 1
self.checks.append((None, label, None, None))
@staticmethod
def _clip(text, limit=200):
text = str(text).replace("\r\n", " ").replace("\n", " ")
return text[:limit] + ("..." if len(text) > limit else "")
def field(response, name):
"""Pull one 'key=value' line out of an echo reply."""
text = response.text if isinstance(response, Response) else response
for line in text.splitlines():
key, _, value = line.partition("=")
if key == name:
return value
return None
def int_field(response, name):
value = field(response, name)
try:
return int(value)
except (TypeError, ValueError):
return None