#!/usr/bin/env python3
"""
AgentRouter <-> OpenCode proxy (OpenAI-compatible endpoints).

Fixes three incompatibilities:
  1. Injects the User-Agent required by AgentRouter (avoids 401).
  2. Strips null fields from the request body (avoids 400).
  3. Filters billing.summary and `data: null` events from the response.

Usage:
    export AGENTROUTER_API_KEY="sk-..."
    python3 agentrouter-proxy.py [port]      # default port: 4182

Environment variables:
    AGENTROUTER_API_KEY    AgentRouter API key (required)
    AGENTROUTER_UPSTREAM   Upstream URL (default: https://agentrouter.org)
    AGENTROUTER_DEBUG=1    Dump requests and SSE events to stderr
"""

import http.client
import http.server
import json
import os
import ssl
import sys
import threading
import urllib.error
import urllib.request

UPSTREAM = os.environ.get("AGENTROUTER_UPSTREAM", "https://agentrouter.org")
PORT = int(sys.argv[1]) if len(sys.argv) > 1 else 4182
API_KEY = os.environ.get("AGENTROUTER_API_KEY", "")
DEBUG = os.environ.get("AGENTROUTER_DEBUG", "") == "1"

# Without this exact User-Agent, AgentRouter returns 401 even with a valid key.
USER_AGENT = "opencode/1.17.12"
CTX = ssl.create_default_context()


def strip_nulls(value):
    """Recursively remove dict entries whose value is None."""
    if isinstance(value, dict):
        return {k: strip_nulls(v) for k, v in value.items() if v is not None}
    if isinstance(value, list):
        return [strip_nulls(item) for item in value]
    return value


def clean_request_body(raw: bytes) -> bytes:
    """Drop null fields from the outgoing JSON body."""
    try:
        payload = json.loads(raw)
    except (json.JSONDecodeError, ValueError):
        return raw

    if not isinstance(payload, dict):
        return raw

    return json.dumps(strip_nulls(payload), separators=(",", ":")).encode()


def is_billing_summary(payload) -> bool:
    """AgentRouter injects billing.summary objects the AI SDK cannot parse."""
    return isinstance(payload, dict) and payload.get("object") == "billing.summary"


def clean_sse_event(event: bytes):
    """Filter a single SSE event. Return None to drop it."""
    stripped = event.strip()
    if not stripped:
        return event  # SSE separator, keep as-is

    if stripped.startswith(b"data:"):
        payload = stripped[len(b"data:"):].strip()
        if payload in (b"", b"null"):
            return None
        try:
            parsed = json.loads(payload)
        except (json.JSONDecodeError, ValueError):
            return event  # e.g. "[DONE]"
        return None if is_billing_summary(parsed) else event

    # Bare JSON with no `data:` prefix
    if stripped[:1] in (b"{", b"["):
        try:
            parsed = json.loads(stripped)
        except (json.JSONDecodeError, ValueError):
            return event
        if is_billing_summary(parsed):
            return None

    return event


def clean_nonstream_body(raw: bytes) -> bytes:
    """Replace a billing.summary response with a readable JSON error."""
    if not raw:
        return raw
    try:
        payload = json.loads(raw)
    except (json.JSONDecodeError, ValueError):
        return raw

    if is_billing_summary(payload):
        sys.stderr.write("[ar] filtered billing.summary from non-stream response\n")
        return json.dumps(
            {
                "error": {
                    "message": "agentrouter billing.summary is not a chat response",
                    "type": "proxy_filtered",
                }
            }
        ).encode()
    return raw


class ProxyHandler(http.server.BaseHTTPRequestHandler):
    protocol_version = "HTTP/1.1"

    def do_POST(self):
        self._proxy()

    def do_GET(self):
        self._proxy()

    def do_OPTIONS(self):
        self.send_response(200)
        self.send_header("Access-Control-Allow-Origin", "*")
        self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
        self.send_header("Access-Control-Allow-Headers", "*")
        self.send_header("Content-Length", "0")
        self.end_headers()

    def _proxy(self):
        content_len = int(self.headers.get("Content-Length", 0))
        raw_body = self.rfile.read(content_len) if content_len > 0 else None

        if DEBUG and raw_body is not None:
            try:
                payload = json.loads(raw_body)
                sys.stderr.write(
                    f"[ar-dbg] req model={payload.get('model')} "
                    f"stream={payload.get('stream')} "
                    f"messages={len(payload.get('messages', []))} "
                    f"bytes={len(raw_body)} keys={list(payload.keys())}\n"
                )
            except (json.JSONDecodeError, ValueError):
                sys.stderr.write(f"[ar-dbg] req {len(raw_body)}b (not JSON)\n")

        body = clean_request_body(raw_body) if raw_body else None

        headers = {
            k: v
            for k, v in self.headers.items()
            if k.lower()
            not in ("host", "connection", "transfer-encoding", "accept-encoding")
        }
        headers["Authorization"] = f"Bearer {API_KEY}"
        headers["User-Agent"] = USER_AGENT
        # Force an uncompressed response: gzip/br cannot be filtered line by line.
        headers["Accept-Encoding"] = "identity"
        if body is not None:
            headers["Content-Length"] = str(len(body))

        url = f"{UPSTREAM}{self.path}"
        request = urllib.request.Request(
            url, data=body, headers=headers, method=self.command
        )
        is_stream = bool(body) and b'"stream":true' in body.replace(b" ", b"")

        try:
            response = urllib.request.urlopen(request, context=CTX, timeout=300)

            if is_stream:
                self.send_response(response.status)
                for k, v in response.headers.items():
                    if k.lower() not in (
                        "transfer-encoding",
                        "connection",
                        "content-length",
                    ):
                        self.send_header(k, v)
                self.send_header("Transfer-Encoding", "chunked")
                self.end_headers()
                self._stream_sse(response)
            else:
                try:
                    raw_response = response.read()
                except http.client.IncompleteRead as exc:
                    raw_response = exc.partial
                payload = clean_nonstream_body(raw_response)
                self.send_response(response.status)
                for k, v in response.headers.items():
                    if k.lower() not in (
                        "transfer-encoding",
                        "connection",
                        "content-length",
                    ):
                        self.send_header(k, v)
                self.send_header("Content-Length", str(len(payload)))
                self.end_headers()
                self.wfile.write(payload)

        except urllib.error.HTTPError as exc:
            err_body = exc.read()
            sys.stderr.write(
                f"[ar] ERROR {exc.code}: {err_body.decode(errors='replace')[:500]}\n"
            )
            self.send_response(exc.code)
            self.send_header(
                "Content-Type", exc.headers.get("Content-Type", "application/json")
            )
            self.send_header("Content-Length", str(len(err_body)))
            self.end_headers()
            self.wfile.write(err_body)

        except Exception as exc:  # noqa: BLE001 - proxy must never crash
            sys.stderr.write(f"[ar] EXCEPTION: {exc}\n")
            payload = json.dumps({"error": {"message": str(exc)}}).encode()
            self.send_response(502)
            self.send_header("Content-Type", "application/json")
            self.send_header("Content-Length", str(len(payload)))
            self.end_headers()
            self.wfile.write(payload)

    def _stream_sse(self, response):
        """Forward the SSE stream, dropping `data: null` and billing.summary."""
        buffer = b""
        try:
            while True:
                try:
                    chunk = response.read(1)
                except http.client.IncompleteRead as exc:
                    chunk = exc.partial
                if not chunk:
                    break

                buffer += chunk
                if not (buffer.endswith(b"\n\n") or buffer.endswith(b"\r\n\r\n")):
                    continue

                if DEBUG:
                    sys.stderr.write(
                        f"[ar-dbg] event ({len(buffer)}b): {buffer[:400]!r}\n"
                    )

                cleaned = clean_sse_event(buffer)
                buffer = b""
                if cleaned is None:
                    if DEBUG:
                        sys.stderr.write("[ar-dbg] -> dropped by filter\n")
                    continue
                self._write_chunk(cleaned)

            if buffer:
                cleaned = clean_sse_event(buffer)
                if cleaned is not None:
                    self._write_chunk(cleaned)

            self.wfile.write(b"0\r\n\r\n")
            self.wfile.flush()
        except (ConnectionResetError, BrokenPipeError):
            return

    def _write_chunk(self, data: bytes):
        self.wfile.write(f"{len(data):X}\r\n".encode())
        self.wfile.write(data)
        self.wfile.write(b"\r\n")
        self.wfile.flush()

    def log_message(self, fmt, *args):
        sys.stderr.write(f"[ar] {args[0]}\n")


class ThreadedHTTPServer(http.server.HTTPServer):
    """Handle each request in its own thread so streaming never blocks."""

    daemon_threads = True

    def process_request(self, request, client_address):
        thread = threading.Thread(
            target=self.process_request_thread, args=(request, client_address)
        )
        thread.daemon = True
        thread.start()

    def process_request_thread(self, request, client_address):
        try:
            self.finish_request(request, client_address)
        except Exception:  # noqa: BLE001
            self.handle_error(request, client_address)
        finally:
            self.shutdown_request(request)


if __name__ == "__main__":
    if not API_KEY:
        print("ERROR: set AGENTROUTER_API_KEY first", file=sys.stderr)
        sys.exit(1)
    server = ThreadedHTTPServer(("127.0.0.1", PORT), ProxyHandler)
    print(f"agentrouter-proxy :{PORT} -> {UPSTREAM}")
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        print("\nProxy stopped.")
        server.server_close()
