"""Owned loopback TLS origins for the download clients; they never dial a destination."""
from contextlib import contextmanager
import hashlib
import json
import os
from pathlib import Path
import socket
import socketserver
import ssl
import struct
import subprocess
import threading
import urllib.request

from test_flutter_cache_setup import SETUP_SHA, SETUP_URL

HOST = "manifest.fixture.invalid"
MANIFEST = json.dumps({"releases": [{"channel": "stable", "version": "3.44.0", "dart_sdk_arch": "x64",
    "hash": "a" * 40, "sha256": "b" * 64, "archive": "owned-sdk.tar.xz"}]}).encode()


def official_setup() -> bytes:
    fixture = os.environ.get("GOWALK_CACHE_ACTION_FIXTURES")
    if fixture:
        source = (Path(fixture) / "flutter-setup.sh").read_bytes()
    else:
        with urllib.request.urlopen(SETUP_URL, timeout=30) as response:
            source = response.read(128 * 1024)
    assert hashlib.sha256(source).hexdigest() == SETUP_SHA
    return source


def response(server, mode, headers, path):
    requested = headers.get(b"range", b"").decode()
    server.ranges.append(requested)
    # A transient origin refusal the client is expected to retry (flutter_manifest.RETRY_HTTP).
    if mode.startswith("origin_") and mode.split("_", 1)[1].isdigit():
        return (b"HTTP/1.1 " + mode.split("_", 1)[1].encode()
                + b" Fixture refusal\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")
    complete = server.body[path] if isinstance(server.body, dict) else server.body
    extra = b""
    status = b"401 Unauthorized" if mode == "origin_auth" else b"200 OK"
    if mode.startswith("range"):
        if not requested.startswith("bytes=") or not requested.endswith("-"):
            return b"HTTP/1.1 416 Range Not Satisfiable\r\nContent-Length: 0\r\n\r\n"
        start = int(requested[6:-1])
        extra = f"Content-Range: bytes {start}-{len(complete) - 1}/{len(complete)}\r\n".encode()
        complete, status = complete[start:], b"206 Partial Content"
    body = complete[:31] if mode in {"partial", "range_partial", "large_partial", "too_large"} else complete
    if mode == "range_corrupt":
        body = bytes([body[0] ^ 1]) + body[1:]
    size = {"large_partial": 2 * 1024 ** 3 + 1, "too_large": 9 * 1024 ** 3}.get(mode, len(complete))
    return (b"HTTP/1.1 " + status + b"\r\n" + extra + b"Content-Length: "
            + str(size).encode() + b"\r\nConnection: close\r\n\r\n" + body)


class Proxy(socketserver.StreamRequestHandler):
    def handle(self):
        first = self.rfile.readline(4096).split()
        while self.rfile.readline(4096) not in (b"\r\n", b"\n", b""):
            pass
        if first != [b"CONNECT", (HOST + ":443").encode(), b"HTTP/1.1"]:
            return
        with self.server.lock:
            self.server.connects += 1
            mode = self.server.modes.pop(0)
        proxy_status = {"proxy_auth": 407, "proxy_401": 401, "proxy_403": 403,
                        "proxy_502": 502, "proxy_503": 503, "proxy_504": 504}.get(mode)
        if proxy_status:
            self.wfile.write(f"HTTP/1.1 {proxy_status} Fixture refusal\r\nContent-Length: 0\r\n\r\n".encode())
            return
        self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n")
        self.wfile.flush()
        if mode == "reset":
            self.connection.setsockopt(socket.SOL_SOCKET, socket.SO_LINGER, struct.pack("ii", 1, 0))
            return
        try:
            with self.server.context.wrap_socket(self.connection, server_side=True) as stream:
                stream.settimeout(3)
                with stream.makefile("rb") as reader:
                    request = reader.readline(4096).split()
                    headers = {}
                    while (line := reader.readline(4096)) not in (b"\r\n", b"\n", b""):
                        name, _, value = line.partition(b":")
                        headers[name.lower()] = value.strip()
                    self.server.requests.append(request[1].decode())
                    stream.sendall(response(self.server, mode, headers, request[1].decode()))
        except (OSError, ssl.SSLError):
            pass  # An untrusted fixture certificate must be rejected by the real client.


@contextmanager
def tls_proxy(root: Path, modes: list[str], body: bytes | dict[str, bytes] = MANIFEST):
    cert, key = root / "cert.pem", root / "key.pem"
    command = ["openssl", "req", "-x509", "-newkey", "rsa:2048", "-nodes", "-days", "1",
               "-keyout", str(key), "-out", str(cert), "-subj", "/CN=" + HOST,
               "-addext", "subjectAltName=DNS:" + HOST]
    subprocess.run(command, check=True, capture_output=True, timeout=10)
    context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
    context.load_cert_chain(cert, key)
    server = socketserver.ThreadingTCPServer(("127.0.0.1", 0), Proxy)
    server.daemon_threads = True
    server.context, server.modes, server.body = context, list(modes), body
    server.connects, server.requests, server.lock = 0, [], threading.Lock()
    server.ranges = []
    thread = threading.Thread(target=server.serve_forever, daemon=True)
    thread.start()
    try:
        yield server, cert
    finally:
        server.shutdown(); server.server_close(); thread.join(3)


ORIGIN_HOST = "127.0.0.1"


class Origin(socketserver.StreamRequestHandler):
    """The same fixture responses, served directly rather than behind a CONNECT hop."""

    def handle(self):
        with self.server.lock:
            self.server.connects += 1
            mode = self.server.modes.pop(0)
        if mode == "reset":
            self.connection.setsockopt(socket.SOL_SOCKET, socket.SO_LINGER, struct.pack("ii", 1, 0))
            return
        try:
            with self.server.context.wrap_socket(self.connection, server_side=True) as stream:
                stream.settimeout(3)
                with stream.makefile("rb") as reader:
                    request = reader.readline(4096).split()
                    headers = {}
                    while (line := reader.readline(4096)) not in (b"\r\n", b"\n", b""):
                        name, _, value = line.partition(b":")
                        headers[name.lower()] = value.strip()
                    self.server.requests.append(request[1].decode())
                    stream.sendall(response(self.server, mode, headers, request[1].decode()))
        except (OSError, ssl.SSLError):
            pass  # An untrusted fixture certificate must be rejected by the real client.


@contextmanager
def tls_origin(root: Path, modes: list[str], body: bytes | dict[str, bytes] = MANIFEST):
    """A loopback HTTPS origin. ``server.url(path)`` is what the client fetches."""
    cert, key = root / "origin-cert.pem", root / "origin-key.pem"
    command = ["openssl", "req", "-x509", "-newkey", "rsa:2048", "-nodes", "-days", "1",
               "-keyout", str(key), "-out", str(cert), "-subj", "/CN=" + ORIGIN_HOST,
               "-addext", "subjectAltName=IP:" + ORIGIN_HOST]
    subprocess.run(command, check=True, capture_output=True, timeout=10)
    context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
    context.load_cert_chain(cert, key)
    server = socketserver.ThreadingTCPServer((ORIGIN_HOST, 0), Origin)
    server.daemon_threads = True
    server.context, server.modes, server.body = context, list(modes), body
    server.connects, server.requests, server.lock = 0, [], threading.Lock()
    server.ranges = []
    server.url = lambda path="/owned-sdk.tar.xz", _s=server: (
        f"https://{ORIGIN_HOST}:{_s.server_address[1]}{path}")
    thread = threading.Thread(target=server.serve_forever, daemon=True)
    thread.start()
    try:
        yield server, cert
    finally:
        server.shutdown(); server.server_close(); thread.join(3)
