#!/usr/bin/env python3
"""Loopback HTTP proxy for the aiohttp lab: Basic authentication, CONNECT and absolute-form requests.

Without the expected Proxy-Authorization header it answers 407 with Proxy-Authenticate. Bytes that
are not an HTTP request (for example a TLS ClientHello sent to this plain-HTTP port) get a plaintext
400. Every request is appended to a JSON-lines log, so a run shows which requests reached the proxy.
Bind it to 127.0.0.1 only; it is a test fixture, not a production proxy.

    python3 auth_proxy.py 8888 --user user --password pass --log proxy.jsonl
"""
import argparse
import asyncio
import base64
import json
from urllib.parse import urlsplit

METHODS = (b"GET ", b"HEAD ", b"POST ", b"PUT ", b"DELETE ", b"OPTIONS ", b"PATCH ", b"CONNECT ")


def response(status: str, extra: str = "") -> bytes:
    return f"HTTP/1.1 {status}\r\n{extra}Content-Length: 0\r\nConnection: close\r\n\r\n".encode("ascii")


async def pipe(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
    try:
        while data := await reader.read(65536):
            writer.write(data)
            await writer.drain()
    except (ConnectionError, OSError):
        pass
    finally:
        writer.close()


class Proxy:
    def __init__(self, user: str, password: str, log_path: str) -> None:
        self.expected = "Basic " + base64.b64encode(f"{user}:{password}".encode()).decode("ascii")
        self.log_path = log_path

    def log(self, **entry: object) -> None:
        line = json.dumps(entry)
        print(line, flush=True)
        if self.log_path:
            with open(self.log_path, "a", encoding="utf-8") as handle:
                handle.write(line + "\n")

    async def handle(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
        try:
            await asyncio.wait_for(self.serve(reader, writer), timeout=30)
        except (asyncio.TimeoutError, ConnectionError, OSError, asyncio.IncompleteReadError, asyncio.LimitOverrunError):
            pass
        finally:
            writer.close()

    async def serve(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
        first = await reader.read(8)
        if not first.startswith(METHODS) and not any(method.startswith(first) for method in METHODS):
            self.log(event="not_http", first_bytes=first.hex())
            writer.write(response("400 Bad Request"))
            await writer.drain()
            return
        head = first + await reader.readuntil(b"\r\n\r\n")
        request_line, *header_lines = head.decode("latin-1").split("\r\n")
        method, target, _version = request_line.split(" ", 2)
        headers = dict(line.split(": ", 1) for line in header_lines if ": " in line)
        sent = next((value for name, value in headers.items() if name.lower() == "proxy-authorization"), None)
        auth = "ok" if sent == self.expected else "missing" if sent is None else "wrong"
        self.log(event="request", method=method, target=target, auth=auth, header_names=sorted(headers))
        if auth != "ok":
            writer.write(response("407 Proxy Authentication Required", 'Proxy-Authenticate: Basic realm="lab"\r\n'))
            await writer.drain()
            return
        if method == "CONNECT":
            host, _, port = target.rpartition(":")
            upstream_reader, upstream_writer = await asyncio.open_connection(host, int(port))
            writer.write(b"HTTP/1.1 200 Connection established\r\n\r\n")
            await writer.drain()
            await asyncio.gather(pipe(reader, upstream_writer), pipe(upstream_reader, writer))
            return
        url = urlsplit(target)
        if url.scheme != "http" or not url.hostname:
            writer.write(response("400 Bad Request"))
            await writer.drain()
            return
        upstream_reader, upstream_writer = await asyncio.open_connection(url.hostname, url.port or 80)
        path = url.path or "/"
        if url.query:
            path += "?" + url.query
        kept = [line for line in header_lines if line and not line.lower().startswith(("proxy-", "connection:"))]
        upstream_writer.write("\r\n".join([f"{method} {path} HTTP/1.1", *kept, "Connection: close", "", ""]).encode("latin-1"))
        await upstream_writer.drain()
        await pipe(upstream_reader, writer)
        upstream_writer.close()


async def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("port", type=int, nargs="?", default=8888)
    parser.add_argument("--user", default="user")
    parser.add_argument("--password", default="pass")
    parser.add_argument("--log", default="")
    args = parser.parse_args()
    proxy = Proxy(args.user, args.password, args.log)
    server = await asyncio.start_server(proxy.handle, "127.0.0.1", args.port)
    print(f"auth proxy listening on 127.0.0.1:{args.port}", flush=True)
    async with server:
        await server.serve_forever()


if __name__ == "__main__":
    asyncio.run(main())
