#!/usr/bin/env python3
"""Loopback lab: "SSL: WRONG_VERSION_NUMBER" behind a proxy, at the proxy hop and at the target hop.

Everything listens on 127.0.0.1. No proxy account is needed and no packet leaves the machine.

  plain-400 listener   answers any input with a plaintext HTTP 400 (stands in for a plain HTTP proxy port)
  CONNECT proxy        minimal plain-HTTP forward proxy; tunnels to loopback destinations only
  plain HTTP target    python -m http.server
  TLS target           the same handler behind a throwaway self-signed certificate (the control)

Usage:  python3 lab.py [--out results.json] [--curl PATH ...] [--python PATH ...] [--skip-node] [--skip-playwright]
        python3 lab.py --serve     (start the listeners only)

--curl and --python add more builds (for example ones linked against another OpenSSL) to every case.
"""
import argparse
import http.server
import json
import os
import platform
import re
import select
import shutil
import socket
import socketserver
import ssl
import subprocess
import sys
import tempfile
import threading
import time
from pathlib import Path

HOST = "127.0.0.1"
PORT = {"plain400": 18400, "connect": 18480, "http_target": 18000, "tls_target": 18443}
HERE = Path(__file__).resolve().parent
BAD_REQUEST = (b"HTTP/1.1 400 Bad Request\r\nContent-Type: text/plain\r\n"
               b"Content-Length: 12\r\nConnection: close\r\n\r\nBad Request\n")
proxy_log = []
PYTHON_ABOUT = ("import json,platform,ssl,importlib.metadata as m;print(json.dumps({'python':platform.python_version(),"
                "'openssl':ssl.OPENSSL_VERSION,'packages':{n:m.version(n) for n in "
                "['requests','urllib3','httpx','httpcore','h11','certifi']}}))")


class Server(socketserver.ThreadingTCPServer):
    allow_reuse_address = True
    daemon_threads = True


def finish(sock, payload):
    """Send a plaintext reply, then drain so the kernel closes with FIN instead of RST."""
    try:
        sock.sendall(payload)
        sock.shutdown(socket.SHUT_WR)
        sock.settimeout(1)
        while sock.recv(4096):
            pass
    except OSError:
        pass


class Plain400(socketserver.BaseRequestHandler):
    def handle(self):
        self.request.settimeout(5)
        try:
            self.request.recv(4096)
        except OSError:
            return
        finish(self.request, BAD_REQUEST)


class ConnectProxy(socketserver.BaseRequestHandler):
    def handle(self):
        client = self.request
        client.settimeout(5)
        head = b""
        try:
            while b"\r\n\r\n" not in head and len(head) < 8192:
                chunk = client.recv(4096)
                if not chunk:
                    break
                head += chunk
                if not re.match(rb"[A-Z]+ ", head):  # not HTTP at all, e.g. a TLS ClientHello
                    break
        except OSError:
            return
        match = re.match(rb"CONNECT ([^\s:]+):(\d+) HTTP/1\.[01]\r\n", head)
        if not match:
            proxy_log.append({"event": "not-a-connect-request", "first_bytes_hex": head[:5].hex()})
            return finish(client, BAD_REQUEST)
        host, port = match.group(1).decode(), int(match.group(2))
        if host not in ("127.0.0.1", "localhost"):
            proxy_log.append({"event": "refused-non-loopback", "authority": f"{host}:{port}"})
            return finish(client, b"HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
        try:
            upstream = socket.create_connection((HOST, port), timeout=5)
        except OSError:
            proxy_log.append({"event": "upstream-failed", "authority": f"{host}:{port}"})
            return finish(client, b"HTTP/1.1 502 Bad Gateway\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
        proxy_log.append({"event": "tunnel-open", "authority": f"{host}:{port}"})
        client.sendall(b"HTTP/1.1 200 Connection established\r\n\r\n")
        client.settimeout(None)
        with upstream:
            pair = {client: upstream, upstream: client}
            while True:
                readable, _, _ = select.select(list(pair), [], [], 15)
                if not readable:
                    return
                for sock in readable:
                    try:
                        data = sock.recv(65536)
                    except OSError:
                        return
                    if not data:
                        return
                    pair[sock].sendall(data)


class Quiet(http.server.SimpleHTTPRequestHandler):
    def log_message(self, *args):
        pass


def start(server):
    threading.Thread(target=server.serve_forever, daemon=True).start()
    return server


def wait_port(port):
    for _ in range(50):
        try:
            socket.create_connection((HOST, port), timeout=0.2).close()
            return
        except OSError:
            time.sleep(0.1)
    raise SystemExit(f"port {port} did not open")


def run(argv, env=None, timeout=40):
    base = {"PATH": os.environ.get("PATH", "/usr/bin:/bin"), "HOME": os.environ.get("HOME", "/tmp"), "LC_ALL": "C"}
    done = subprocess.run(argv, env={**base, **(env or {})}, capture_output=True, text=True, timeout=timeout, cwd=HERE)
    return {"exit": done.returncode, "stdout": done.stdout.strip(), "stderr": done.stderr.strip()}


def as_json(result):
    try:
        return json.loads(result["stdout"].splitlines()[-1])
    except (ValueError, IndexError):
        return result


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--out", default="results.json")
    parser.add_argument("--curl", action="append", default=[], help="extra curl binaries to run")
    parser.add_argument("--python", action="append", default=[], help="extra Python interpreters with requests and httpx installed")
    parser.add_argument("--serve", action="store_true", help="only start the listeners and wait, to try commands by hand")
    parser.add_argument("--skip-node", action="store_true")
    parser.add_argument("--skip-playwright", action="store_true")
    args = parser.parse_args()

    work = Path(tempfile.mkdtemp(prefix="wrong-version-lab-"))
    (work / "www").mkdir()
    (work / "www" / "index.html").write_text("lab target ok\n")
    cert, key = work / "cert.pem", work / "key.pem"
    subprocess.run(["openssl", "req", "-x509", "-newkey", "rsa:2048", "-nodes", "-days", "2", "-subj", "/CN=127.0.0.1",
                    "-addext", "subjectAltName=IP:127.0.0.1,DNS:localhost", "-addext", "basicConstraints=critical,CA:TRUE",
                    "-keyout", str(key), "-out", str(cert)], check=True, capture_output=True)

    start(Server((HOST, PORT["plain400"]), Plain400))
    start(Server((HOST, PORT["connect"]), ConnectProxy))
    plain_target = subprocess.Popen([sys.executable, "-m", "http.server", str(PORT["http_target"]), "--bind", HOST,
                                     "--directory", str(work / "www")], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
    tls = Server((HOST, PORT["tls_target"]), lambda *a: Quiet(*a, directory=str(work / "www")))
    context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
    context.load_cert_chain(cert, key)
    tls.socket = context.wrap_socket(tls.socket, server_side=True)
    start(tls)
    for port in PORT.values():
        wait_port(port)

    plain400 = f"{HOST}:{PORT['plain400']}"
    connect = f"{HOST}:{PORT['connect']}"
    tls_url = f"https://{HOST}:{PORT['tls_target']}/"
    wrong_url = f"https://{HOST}:{PORT['http_target']}/"  # https:// URL on a port that speaks plain HTTP
    cases = [
        {"id": "case1-proxy-hop", "what": "https:// proxy URL, proxy port answers in plaintext (plain-400 listener)",
         "proxy": f"https://{plain400}", "target": tls_url},
        {"id": "case1-proxy-hop-connect-proxy", "what": "https:// proxy URL against the working plain-HTTP CONNECT proxy",
         "proxy": f"https://{connect}", "target": tls_url},
        {"id": "case2-target-hop", "what": "http:// proxy URL, CONNECT succeeds, target port speaks plain HTTP",
         "proxy": f"http://{connect}", "target": wrong_url},
        {"id": "case2-direct", "what": "the case 2 target with no proxy at all", "proxy": None, "target": wrong_url},
        {"id": "control-fixed", "what": "http:// proxy URL and a real TLS target", "proxy": f"http://{connect}", "target": tls_url},
    ]

    if args.serve:
        print(f"listening on {HOST}; certificate for the TLS target: {cert}\n" + "\n".join(
            f"  {case['id']}: proxy {case['proxy']}  target {case['target']}" for case in cases) + "\nCtrl+C to stop")
        try:
            threading.Event().wait()
        except KeyboardInterrupt:
            plain_target.terminate()
            shutil.rmtree(work, ignore_errors=True)
            return

    curls = [shutil.which("curl")] + args.curl
    pythons = []
    for binary in [sys.executable] + args.python:
        about = as_json(run([binary, "-c", PYTHON_ABOUT]))
        pythons.append((binary, about))
    node = None if args.skip_node else shutil.which("node")
    results = []
    for case in cases:
        proxy, target = case["proxy"], case["target"]
        entry = {**case, "clients": {}}
        before = len(proxy_log)
        for binary in filter(None, curls):
            label = run([binary, "--version"])["stdout"].splitlines()[0]
            route = ["--noproxy", "*"] if proxy is None else ["--noproxy", "", "--proxy", proxy]
            common = [binary, "--disable", "--silent", "--show-error", "--output", "/dev/null", "--max-time", "10",
                      "--cacert", str(cert), *route]
            flag = run([*common, "--write-out", "http_code=%{http_code} http_connect=%{http_connect} exit=%{exitcode}", target])
            verbose = run([*common, "--verbose", target])
            keep = re.compile(r"^(\* (Establish HTTP proxy tunnel|CONNECT tunnel|Proxy replied|.*(TLS|SSL|error|rror:|Closing)).*|[<>] (CONNECT|HTTP/).*|curl: .*)$")
            entry["clients"].setdefault("curl", []).append({
                "version": label, "exit": flag["exit"], "write_out": flag["stdout"], "stderr": flag["stderr"],
                "verbose_excerpt": [line for line in verbose["stderr"].splitlines() if keep.match(line)]})
            if proxy is not None and binary == curls[0]:
                env = run([binary, "--disable", "--silent", "--show-error", "--output", "/dev/null", "--max-time", "10",
                           "--cacert", str(cert), target], env={"HTTPS_PROXY": proxy, "https_proxy": proxy})
                entry["clients"]["curl_env_HTTPS_PROXY"] = {"exit": env["exit"], "stderr": env["stderr"]}
        for name in ("requests", "httpx"):
            script = str(HERE / f"client_{name}.py")
            for python, about in pythons:
                entry["clients"].setdefault(name, []).append(
                    {"python": about["python"], "openssl": about["openssl"], **as_json(run([python, script, proxy or "-", target, str(cert)]))})
            if proxy is not None:
                entry["clients"][f"{name}_env_HTTPS_PROXY"] = as_json(run(
                    [pythons[0][0], script, "--env", target, str(cert)], env={"HTTPS_PROXY": proxy, "https_proxy": proxy}))
        if proxy is not None:
            debug = run([pythons[0][0], str(HERE / "client_httpx_debug.py"), proxy, target, str(cert)])
            wanted = re.compile(r"CONNECT|start_tls|receive_response_headers\.complete|^raised")
            entry["clients"]["httpx_debug_log"] = [re.sub(r" object at 0x[0-9a-f]+", "", line)
                                                   for line in debug["stdout"].splitlines() if wanted.search(line)]
        if node:
            env = {"NODE_EXTRA_CA_CERTS": str(cert)}
            if proxy is not None:
                env.update({"NODE_USE_ENV_PROXY": "1", "HTTPS_PROXY": proxy, "HTTP_PROXY": proxy})
            entry["clients"]["node_fetch"] = as_json(run([node, str(HERE / "client_fetch.mjs"), target], env=env))
        if node and not args.skip_playwright and proxy is not None:
            entry["clients"]["playwright"] = as_json(run([node, str(HERE / "client_playwright.cjs"), proxy, target],
                                                         env={"NODE_PATH": os.environ.get("NODE_PATH", "")}, timeout=90))
        entry["connect_proxy_log"] = proxy_log[before:]
        results.append(entry)
        print(f"{case['id']}: done", file=sys.stderr)

    versions = {
        "recorded_utc": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
        "platform": platform.platform(),
        "pythons": [about for _, about in pythons],
        "openssl_cli": run(["openssl", "version"])["stdout"],
    }
    if node:
        versions["node"] = as_json(run([node, "-p", "JSON.stringify({node:process.version,openssl:process.versions.openssl,undici:process.versions.undici})"]))
    Path(args.out).write_text(json.dumps({"versions": versions, "ports": PORT, "cases": results}, indent=1) + "\n")
    plain_target.terminate()
    shutil.rmtree(work, ignore_errors=True)
    print(f"wrote {args.out}", file=sys.stderr)


if __name__ == "__main__":
    main()
