#!/usr/bin/env python3
from __future__ import annotations

import http.client
import json
import socketserver
import threading
import urllib.parse
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import Final

ANTHROPIC_HOST: Final[str] = "api.anthropic.com"
CODEX_MODEL_PREFIX: Final[str] = "gpt-"
PICKER_MODEL_PREFIX: Final[str] = "anthropic."
UPSTREAM_TIMEOUT_SECS: Final[float] = 900.0
HOP_BY_HOP: Final[frozenset[str]] = frozenset({
    "connection",
    "keep-alive",
    "proxy-authenticate",
    "proxy-authorization",
    "te",
    "trailer",
    "transfer-encoding",
    "upgrade",
})


@dataclass(frozen=True, slots=True)
class RouteTarget:
    codex: bool
    model: str | None


def route_target(model: object) -> RouteTarget:
    if not isinstance(model, str):
        return RouteTarget(False, None)
    bare = model[len(PICKER_MODEL_PREFIX) :] if model.startswith(PICKER_MODEL_PREFIX) else model
    if bare.startswith(CODEX_MODEL_PREFIX):
        return RouteTarget(True, bare)
    return RouteTarget(False, None)


class _Handler(BaseHTTPRequestHandler):
    protocol_version = "HTTP/1.1"
    server: HybridRouter

    def do_GET(self) -> None:
        self._forward(b"")

    def do_DELETE(self) -> None:
        self._forward(b"")

    def do_POST(self) -> None:
        self._forward(self._read_body())

    def do_PUT(self) -> None:
        self._forward(self._read_body())

    def do_PATCH(self) -> None:
        self._forward(self._read_body())

    def log_message(self, format: str, *args: object) -> None:
        return

    def _read_body(self) -> bytes:
        length = self.headers.get("content-length")
        if length is None:
            return b""
        try:
            return self.rfile.read(int(length))
        except ValueError:
            return b""

    def _forward(self, body: bytes) -> None:
        body, codex = self._rewrite_for_codex(body)
        try:
            connection = self._connect(codex)
        except OSError as exc:
            self._send_error_json(502, f"claudex router: upstream connect failed: {exc}")
            return
        try:
            connection.request(
                self.command,
                self.server.upstream_path(self.path, codex=codex),
                body=body,
                headers=self._upstream_headers(body, codex),
            )
            response = connection.getresponse()
            self._relay(response)
        except (OSError, http.client.HTTPException) as exc:
            self._send_error_json(502, f"claudex router: upstream request failed: {exc}")
        finally:
            connection.close()

    def _rewrite_for_codex(self, body: bytes) -> tuple[bytes, bool]:
        if not body:
            return body, False
        try:
            payload = json.loads(body)
        except (UnicodeDecodeError, json.JSONDecodeError):
            return body, False
        if not isinstance(payload, dict):
            return body, False
        target = route_target(payload.get("model"))
        if not target.codex or target.model is None:
            return body, False
        payload["model"] = target.model
        return json.dumps(payload).encode("utf-8"), True

    def _connect(self, codex: bool) -> http.client.HTTPConnection:
        if codex:
            return self.server.connect_codex()
        return self.server.connect_anthropic()

    def _upstream_headers(self, body: bytes, codex: bool) -> dict[str, str]:
        headers = {
            key: value
            for key, value in self.headers.items()
            if key.lower() not in HOP_BY_HOP and key.lower() not in ("host", "content-length")
        }
        headers["Host"] = (
            f"127.0.0.1:{self.server.codex_port}"
            if codex
            else self.server.anthropic_host_header
        )
        headers["Accept-Encoding"] = "identity"
        if body:
            headers["Content-Length"] = str(len(body))
        return headers

    def _relay(self, response: http.client.HTTPResponse) -> None:
        content_length = response.getheader("content-length")
        self.send_response(response.status)
        for key, value in response.getheaders():
            if key.lower() in HOP_BY_HOP or key.lower() == "content-length":
                continue
            self.send_header(key, value)
        if content_length is not None:
            self.send_header("Content-Length", content_length)
        else:
            self.send_header("Transfer-Encoding", "chunked")
        self.end_headers()
        if self.command == "HEAD" or response.status in (204, 304):
            return
        try:
            while True:
                chunk = response.read(8192)
                if not chunk:
                    break
                if content_length is None:
                    self.wfile.write(f"{len(chunk):x}\r\n".encode() + chunk + b"\r\n")
                else:
                    self.wfile.write(chunk)
                self.wfile.flush()
            if content_length is None:
                self.wfile.write(b"0\r\n\r\n")
                self.wfile.flush()
        except (BrokenPipeError, ConnectionResetError):
            self.close_connection = True

    def _send_error_json(self, status: int, message: str) -> None:
        payload = json.dumps(
            {"type": "error", "error": {"type": "api_error", "message": message}}
        ).encode("utf-8")
        try:
            self.send_response(status)
            self.send_header("Content-Type", "application/json")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)
        except (BrokenPipeError, ConnectionResetError):
            self.close_connection = True


class HybridRouter(socketserver.ThreadingMixIn, HTTPServer):
    daemon_threads = True
    allow_reuse_address = False

    def __init__(self, codex_port: int, *, anthropic_base_url: str | None = None) -> None:
        super().__init__(("127.0.0.1", 0), _Handler)
        self.codex_port = codex_port
        parsed = urllib.parse.urlsplit(anthropic_base_url or f"https://{ANTHROPIC_HOST}")
        if (
            parsed.scheme not in {"http", "https"}
            or not parsed.hostname
            or parsed.username is not None
            or parsed.password is not None
            or parsed.query
            or parsed.fragment
        ):
            self.server_close()
            raise ValueError("invalid Anthropic upstream base URL")
        self.anthropic_scheme = parsed.scheme
        self.anthropic_host = parsed.hostname
        self.anthropic_port = parsed.port or (443 if parsed.scheme == "https" else 80)
        self.anthropic_path_prefix = parsed.path.rstrip("/")
        default_port = 443 if parsed.scheme == "https" else 80
        self.anthropic_host_header = (
            self.anthropic_host
            if self.anthropic_port == default_port
            else f"{self.anthropic_host}:{self.anthropic_port}"
        )
        self._thread: threading.Thread | None = None

    @property
    def port(self) -> int:
        return int(self.server_address[1])

    @property
    def base_url(self) -> str:
        return f"http://127.0.0.1:{self.port}"

    def connect_codex(self) -> http.client.HTTPConnection:
        return http.client.HTTPConnection(
            "127.0.0.1", self.codex_port, timeout=UPSTREAM_TIMEOUT_SECS
        )

    def connect_anthropic(self) -> http.client.HTTPConnection:
        connection_cls = (
            http.client.HTTPSConnection
            if self.anthropic_scheme == "https"
            else http.client.HTTPConnection
        )
        return connection_cls(
            self.anthropic_host, self.anthropic_port, timeout=UPSTREAM_TIMEOUT_SECS
        )

    def upstream_path(self, request_path: str, *, codex: bool) -> str:
        if codex or not self.anthropic_path_prefix:
            return request_path
        suffix = request_path if request_path.startswith("/") else f"/{request_path}"
        return f"{self.anthropic_path_prefix}{suffix}"

    def start(self) -> None:
        if self._thread is not None:
            raise RuntimeError("claudex router already started")
        self._thread = threading.Thread(target=self.serve_forever, daemon=True)
        self._thread.start()

    def stop(self) -> None:
        self.shutdown()
        if self._thread is not None:
            self._thread.join(timeout=5)
            self._thread = None
        self.server_close()

    def handle_error(self, request: object, client_address: object) -> None:
        return


__all__ = [
    "ANTHROPIC_HOST",
    "CODEX_MODEL_PREFIX",
    "HybridRouter",
    "PICKER_MODEL_PREFIX",
    "RouteTarget",
    "route_target",
]
