"""Serve the owner's ChatGPT web session as an OpenAI-compatible endpoint.

One browser session for the daemon's lifetime, so the ~27s cold start is paid once
and the profile lock is held by exactly one process. Requests are served from a
single FIFO queue because the browser is single-seated — this is a property of the
adapter, not a bug to engineer around.

Filesystem confinement is the calling agent's own permission system; nothing here
executes a tool. `sandbox.py`/`workspace.py` guard the MCP tunnel path, not this one.
"""

from __future__ import annotations

import argparse
import json
import os
import queue
import secrets
import statistics
import sys
import threading
import time
import uuid
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path

import chat
import protocol
import translate
from browser import MarionetteError, Session
from protocol import Final, Malformed, ToolCall

STATE = Path.home() / ".overdeck" / "gptbridge"
TOKEN_FILE = STATE / "solwebd.token"
DEFAULT_PORT = 8791
MODEL_PREFIX = "sol-web-"
MODELS = [MODEL_PREFIX + effort for effort in chat.EFFORTS]
KEEPALIVE_SECONDS = 5.0
# A client that walked away cannot cancel a Marionette call in flight, so the turn
# timeout is also how long one abandoned request can hold the single browser seat.
TURN_TIMEOUT = float(os.environ.get("SOLWEBD_TURN_TIMEOUT", "300"))

# ChatGPT renders fenced code as DOM nodes, so fence markers never reach innerText.
# Tool calls are read out of the reply turn's `pre code` elements instead.
REPLY_BLOCKS = """
let last = null;
for (const t of document.querySelectorAll(arguments[0])) {
  if (t.querySelector(arguments[1])) continue;
  last = t;
}
if (!last) return null;
return {blocks: Array.from(last.querySelectorAll('pre code')).map(c => c.textContent),
        text: (last.innerText || '').replace(/^ChatGPT said:\\s*/, '').trim(),
        url: location.href};
"""


class Capped(Exception):
    def __init__(self, state: dict):
        super().__init__(state.get("text", "usage cap"))
        self.state = state


class LoggedOut(Exception):
    pass


class ProtocolBroken(Exception):
    pass


def read_token() -> str:
    """One bearer token per machine; any local process could otherwise drive the
    owner's real ChatGPT account."""
    STATE.mkdir(parents=True, exist_ok=True)
    if not TOKEN_FILE.exists():
        TOKEN_FILE.touch(mode=0o600)
        TOKEN_FILE.write_text(secrets.token_urlsafe(32) + "\n")
    TOKEN_FILE.chmod(0o600)
    return TOKEN_FILE.read_text().strip()


def effort_of(model: str) -> str:
    effort = model.split(",")[-1].strip()
    if not effort.startswith(MODEL_PREFIX):
        raise ValueError(f"model must be one of {MODELS}")
    effort = effort[len(MODEL_PREFIX):]
    if effort not in chat.EFFORTS:
        raise ValueError(f"model must be one of {MODELS}")
    return effort


def split_messages(messages: list[dict]) -> tuple[str, list[dict]]:
    system = "\n\n".join(translate._text(m) for m in messages if m.get("role") == "system")
    return system, [m for m in messages if m.get("role") != "system"]


class BrowserSeat:
    """The single live conversation surface. Owns the browser; knows nothing about HTTP."""

    def __init__(self, mode: str = "virtual", session_factory=None):
        self.mode = mode
        self._factory = session_factory or (lambda: Session(mode=mode))
        self._session = None
        self._m = None

    def open(self):
        if self._m is None:
            self._session = self._factory()
            self._m = self._session.__enter__()
            self._m.navigate("https://chatgpt.com/")
            chat.wait_ready(self._m)
            chat.use_chat_surface(self._m)
        return self._m

    def close(self) -> None:
        if self._session is not None:
            self._session.__exit__(None, None, None)
        self._session, self._m = None, None

    def rebuild(self):
        """A dead browser is rebuilt once; a second failure is reported, not retried."""
        self.close()
        return self.open()

    def new_conversation(self) -> str:
        m = self.open()
        m.navigate("https://chatgpt.com/")
        chat.wait_ready(m)
        chat.use_chat_surface(m)
        return m.url()

    def turn(self, text: str, effort: str, timeout: float = 900.0) -> dict:
        """Type one delta, send it, and read the reply's code blocks back."""
        m = self.open()
        cap = chat.detect_cap(m)
        if cap:
            raise Capped(cap)
        # Re-asserted every send: whether the effort control is per-conversation or
        # composer-global is a UI detail this daemon must not bet on.
        chat.set_effort(m, effort)
        chat.type_prompt(m, text)
        before = int(m.sync_script(chat.TURN_COUNT, [chat.TURN, chat.USER]) or 0)
        try:
            chat.send(m, before, timeout=timeout)
        except chat.ChatError:
            cap = chat.detect_cap(m)
            if cap:
                raise Capped(cap) from None
            raise
        data = m.sync_script(REPLY_BLOCKS, [chat.TURN, chat.USER])
        if not data:
            raise chat.ChatError("no reply turn appeared in the conversation")
        return data


class Engine:
    """Serial request service: one queue, one browser, one conversation per session."""

    def __init__(self, seat: BrowserSeat):
        self.seat = seat
        self.registry = translate.Registry()
        self.lock = threading.Lock()
        self.waiting = 0
        self.durations: list[float] = []

    @property
    def queue_depth(self) -> int:
        return self.waiting

    def health(self) -> dict:
        return {
            "ok": True,
            "conversations": len(self.registry),
            "queue": self.queue_depth,
            "turn_p50_s": round(statistics.median(self.durations), 1) if self.durations else 0.0,
        }

    def complete(self, body: dict) -> dict:
        effort = effort_of(str(body.get("model", "")))
        tools = body.get("tools") or []
        messages = body.get("messages") or []
        system, history = split_messages(messages)
        key = translate.session_key(messages, tools, effort)

        self.waiting += 1
        started = time.monotonic()
        try:
            with self.lock:
                return self._serve(key, system, tools, history, effort)
        finally:
            self.waiting -= 1
            self.durations.append(time.monotonic() - started)
            del self.durations[:-50]

    def ask(self, body: dict) -> dict:
        """One-shot ask on a fresh conversation, under the SAME lock agentic turns
        take — the browser has one composer and two writers would interleave."""
        prompt = str(body.get("prompt") or "").strip()
        if not prompt:
            raise ValueError("prompt is required")
        effort = body.get("effort") or None
        if effort and effort not in chat.EFFORTS:
            raise ValueError(f"effort must be one of {list(chat.EFFORTS)}")
        attachments = [Path(p) for p in body.get("attach") or []]
        missing = [str(p) for p in attachments if not p.is_file()]
        if missing:
            raise ValueError(f"no such file: {', '.join(missing)}")
        out_dir = Path(body.get("out") or chat.OUT_DIR)
        timeout = float(body.get("timeout") or 900.0)

        self.waiting += 1
        started = time.monotonic()
        try:
            with self.lock:
                self.seat.new_conversation()
                m = self.seat.open()
                cap = chat.detect_cap(m)
                if cap:
                    raise Capped(cap)
                reply = chat.ask_on(m, prompt, effort=effort, attachments=attachments,
                                    out_dir=out_dir, stamp=str(body.get("stamp") or ""),
                                    timeout=timeout)
        finally:
            self.waiting -= 1
            self.durations.append(time.monotonic() - started)
            del self.durations[:-50]
        return {"text": reply.text, "images": [str(p) for p in reply.images],
                "files": [str(p) for p in reply.files], "conversation": reply.url}

    def _serve(self, key, system, tools, history, effort) -> dict:
        conv = self.registry.get(key)
        preamble = protocol.render_preamble(system, tools)
        if conv is None:
            conv = self._open(key, preamble, effort)
            text = translate.deliver_all(conv, history)
        else:
            try:
                text = translate.deliver(conv, history)
            except translate.ReplayNeeded:
                conv = self._open(key, preamble, effort)
                text = translate.deliver_all(conv, history)
        return self._exchange(conv, text, tools, effort)

    def _open(self, key, preamble, effort):
        url = self.seat.new_conversation()
        conv = self.registry.open(key, preamble, effort, opened_at=time.time(), url=url)
        conv.pending_preamble = True  # type: ignore[attr-defined]
        return conv

    def _exchange(self, conv, text: str, tools: list[dict], effort: str) -> dict:
        if getattr(conv, "pending_preamble", False):
            text = f"{conv.preamble}\n\n{text}"
            conv.pending_preamble = False  # type: ignore[attr-defined]
        reply = self._send(conv, text, effort)
        parsed = protocol.parse_reply(reply["blocks"], tools, reply["text"])
        if isinstance(parsed, Malformed):
            reply = self._send(conv, protocol.correction_prompt(parsed.reason), effort)
            parsed = protocol.parse_reply(reply["blocks"], tools, reply["text"])
            if isinstance(parsed, Malformed):
                raise ProtocolBroken(parsed.reason)
        conv.turns += 1
        return self._render(parsed, conv)

    def _send(self, conv, text: str, effort: str) -> dict:
        try:
            return self.seat.turn(text, effort, timeout=TURN_TIMEOUT)
        except MarionetteError:
            self.seat.rebuild()
            return self.seat.turn(text, effort, timeout=TURN_TIMEOUT)

    def _render(self, parsed, conv) -> dict:
        base = {"id": f"chatcmpl-{uuid.uuid4().hex[:16]}", "object": "chat.completion",
                "created": int(time.time()), "model": MODEL_PREFIX + conv.effort,
                "usage": {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}}
        if isinstance(parsed, ToolCall):
            call_id = f"call_{conv.turns}_{uuid.uuid4().hex[:8]}"
            message = {"role": "assistant", "content": None, "tool_calls": [
                {"id": call_id, "type": "function",
                 "function": {"name": parsed.name,
                              "arguments": json.dumps(parsed.arguments)}}]}
            finish = "tool_calls"
        else:
            message = {"role": "assistant", "content": parsed.text}
            finish = "stop"
        base["choices"] = [{"index": 0, "message": message, "finish_reason": finish}]
        return base


def keepalive_chunk(model: str, first: bool = False) -> dict:
    """An empty-but-well-formed delta.

    SSE comments (`: ping`) are legal but a translating proxy in the middle parses
    every data event and cancels the stream when nothing arrives for a browser turn
    (10-120s). A real chunk survives the translation.
    """
    delta = {"role": "assistant", "content": ""} if first else {"content": ""}
    return {"id": "chatcmpl-keepalive", "object": "chat.completion.chunk",
            "created": int(time.time()), "model": model,
            "choices": [{"index": 0, "delta": delta, "finish_reason": None}]}


def as_chunk(completion: dict) -> dict:
    """The whole answer as one streamed delta — the browser has no token stream."""
    choice = completion["choices"][0]
    return {"id": completion["id"], "object": "chat.completion.chunk",
            "created": completion["created"], "model": completion["model"],
            "choices": [{"index": 0, "finish_reason": choice["finish_reason"],
                         "delta": choice["message"]}]}


class Handler(BaseHTTPRequestHandler):
    engine: Engine
    token: str
    protocol_version = "HTTP/1.1"

    def log_message(self, fmt, *args):
        sys.stderr.write("solwebd %s\n" % (fmt % args))

    def _send_json(self, status: int, payload: dict, headers: dict | None = None) -> None:
        body = json.dumps(payload).encode()
        self.send_response(status)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(body)))
        for name, value in (headers or {}).items():
            self.send_header(name, value)
        self.end_headers()
        self.wfile.write(body)

    def _authed(self) -> bool:
        if self.headers.get("Authorization") == f"Bearer {self.token}":
            return True
        self._send_json(401, {"error": {"message": "missing or wrong bearer token"}})
        return False

    def do_GET(self):
        if not self._authed():
            return
        path = self.path.split("?")[0].rstrip("/")
        if path.endswith("/v1/models"):
            self._send_json(200, {"object": "list",
                                  "data": [{"id": m, "object": "model"} for m in MODELS]})
        elif path.endswith("/healthz"):
            self._send_json(200, self.engine.health())
        else:
            self._send_json(404, {"error": {"message": "not found"}})

    def do_POST(self):
        if not self._authed():
            return
        path = self.path.split("?")[0].rstrip("/")
        if path.endswith("/shutdown"):
            self._send_json(200, {"ok": True})
            threading.Thread(target=self.server.shutdown, daemon=True).start()
            return
        if not (path.endswith("/chat/completions") or path.endswith("/v1/ask")):
            self._send_json(404, {"error": {"message": "not found"}})
            return
        try:
            body = json.loads(self.rfile.read(int(self.headers.get("Content-Length", 0))) or b"{}")
        except ValueError as exc:
            self._send_json(400, {"error": {"message": f"invalid json: {exc}"}})
            return
        if path.endswith("/v1/ask"):
            status, payload, headers = self._run(body, self.engine.ask)
            self._send_json(status, payload, headers)
            return
        if body.get("stream"):
            self._stream(body)
        else:
            self._blocking(body)

    def _run(self, body: dict, call=None) -> tuple[int, dict, dict]:
        try:
            return 200, (call or self.engine.complete)(body), {}
        except ValueError as exc:
            return 400, {"error": {"message": str(exc)}}, {}
        except Capped as exc:
            payload = {"error": {"message": str(exc), "type": "rate_limit"},
                       "resume_at": exc.state.get("resume_at")}
            return 429, payload, {"Retry-After": "900"}
        except (LoggedOut, MarionetteError, chat.ChatError) as exc:
            return 503, {"error": {"message": f"chatgpt session unavailable: {exc}"}}, {}
        except ProtocolBroken as exc:
            return 502, {"error": {"message": f"model would not hold the protocol: {exc}"}}, {}

    def _blocking(self, body: dict) -> None:
        status, payload, headers = self._run(body)
        self._send_json(status, payload, headers)

    def _stream(self, body: dict) -> None:
        result: dict = {}
        worker = threading.Thread(target=lambda: result.update(
            zip(("status", "payload", "headers"), self._run(body))), daemon=True)
        worker.start()
        # A turn takes 10-20s; a silent socket for that long reads as a hung provider.
        worker.join(timeout=KEEPALIVE_SECONDS)
        if worker.is_alive():
            model = str(body.get("model", "")).split(",")[-1].strip() or MODEL_PREFIX + "medium"
            self.send_response(200)
            self.send_header("Content-Type", "text/event-stream")
            self.send_header("Cache-Control", "no-cache")
            self.send_header("Connection", "close")
            self.end_headers()
            first = True
            while worker.is_alive():
                self._write_event(keepalive_chunk(model, first))
                first = False
                worker.join(timeout=KEEPALIVE_SECONDS)
            self._write_stream_tail(result)
            return
        if result.get("status") != 200:
            self._send_json(result["status"], result["payload"], result.get("headers"))
            return
        self.send_response(200)
        self.send_header("Content-Type", "text/event-stream")
        self.send_header("Cache-Control", "no-cache")
        self.send_header("Connection", "close")
        self.end_headers()
        self._write_stream_tail(result)

    def _write_event(self, payload: dict) -> None:
        self.wfile.write(f"data: {json.dumps(payload)}\n\n".encode())
        self.wfile.flush()

    def _write_stream_tail(self, result: dict) -> None:
        if result.get("status") != 200:
            self._write_event({"error": result.get("payload", {}).get("error",
                                                                      {"message": "failed"})})
        else:
            self._write_event(as_chunk(result["payload"]))
        self.wfile.write(b"data: [DONE]\n\n")
        self.wfile.flush()


def serve(port: int = DEFAULT_PORT, mode: str = "virtual", seat: BrowserSeat | None = None):
    engine = Engine(seat or BrowserSeat(mode=mode))
    handler = type("BoundHandler", (Handler,), {"engine": engine, "token": read_token()})
    # 127.0.0.1 only: this endpoint drives the owner's real account.
    server = ThreadingHTTPServer(("127.0.0.1", port), handler)
    return server, engine


def main() -> int:
    parser = argparse.ArgumentParser(prog="solwebd")
    parser.add_argument("--port", type=int, default=DEFAULT_PORT)
    parser.add_argument("--mode", default="virtual", choices=Session.MODES)
    args = parser.parse_args()
    server, engine = serve(port=args.port, mode=args.mode)
    print(json.dumps({"listening": f"127.0.0.1:{args.port}", "models": MODELS,
                      "token_file": str(TOKEN_FILE)}), flush=True)
    try:
        server.serve_forever()
    except KeyboardInterrupt:
        pass
    finally:
        engine.seat.close()
    return 0


if __name__ == "__main__":
    sys.exit(main())
