#!/usr/bin/env python3
"""Chatties inbox agent sidecar — outbound HTTPS/WSS only.

Stay-open WebSocket first. Reconnects with jittered backoff only when the
socket drops or errors. After any message, reconnect immediately if dropped.
Long-poll GET /pending is fallback only (HTTP 426 or WS handshake fail) —
a held long-poll bills Durable Object duration while waiting.

The iPhone never needs an inbound port on this box.

Required env:
  CHATTIES_INBOX_URL      e.g. https://inbox.chatties.app
  CHATTIES_MAILBOX_ID     mb_...
  CHATTIES_AGENT_TOKEN    ag_...  (shown once at pair time)

Optional:
  CHATTIES_INBOX_HANDLER      command; stdin = message text, stdout = reply
  CHATTIES_INBOX_AUTO_REPLY   template; {text} is replaced with the phone message
  CHATTIES_INBOX_WAIT         long-poll seconds if WS is unavailable (default 25)

Usage:
  python3 chatties-inbox-agent.py
  python3 chatties-inbox-agent.py --echo
"""

from __future__ import annotations

import base64
import hashlib
import json
import os
import random
import socket
import ssl
import struct
import subprocess
import sys
import time
import urllib.error
import urllib.parse
import urllib.request

DEFAULT_URL = "https://inbox.chatties.app"
DEFAULT_WAIT = 25
BACKOFF_SECONDS = (1, 2, 5, 15, 30)
# Browser-like UA: Cloudflare 1010s datacenter/script UAs on Bot Fight Mode.
BROWSER_UA = (
    "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) "
    "AppleWebKit/537.36 (KHTML, like Gecko) "
    "Chrome/128.0.0.0 Safari/537.36"
)
PING_INTERVAL = 20.0


def env(name: str, default: str = "") -> str:
    return os.environ.get(name, default).strip()


def die(msg: str, code: int = 2) -> None:
    print(f"chatties-inbox-agent: {msg}", file=sys.stderr)
    raise SystemExit(code)


def jittered_backoff(failures: int) -> float:
    """1s, 2s, 5s, 15s, cap 30s, with ±30% jitter."""
    idx = min(max(failures, 1), len(BACKOFF_SECONDS)) - 1
    delay = float(BACKOFF_SECONDS[idx])
    return delay * (0.7 + random.random() * 0.6)


class WebSocketError(RuntimeError):
    def __init__(self, message: str, status: int | None = None) -> None:
        super().__init__(message)
        self.status = status


class WebSocketClient:
    """Minimal RFC 6455 client (stdlib only). Client frames are masked."""

    def __init__(self, url: str, token: str, origin: str, timeout: float = 30.0) -> None:
        self.url = url
        self.token = token
        self.origin = origin
        self.timeout = timeout
        self.sock: ssl.SSLSocket | socket.socket | None = None
        self._buf = b""

    def connect(self) -> None:
        parsed = urllib.parse.urlparse(self.url)
        scheme = parsed.scheme.lower()
        host = parsed.hostname or ""
        if not host:
            raise WebSocketError("missing host")
        use_tls = scheme in ("https", "wss")
        port = parsed.port or (443 if use_tls else 80)
        path = parsed.path or "/"
        if parsed.query:
            path = f"{path}?{parsed.query}"
        raw = socket.create_connection((host, port), timeout=self.timeout)
        sock: ssl.SSLSocket | socket.socket = raw
        if use_tls:
            ctx = ssl.create_default_context()
            sock = ctx.wrap_socket(raw, server_hostname=host)
        key = base64.b64encode(os.urandom(16)).decode("ascii")
        headers = (
            f"GET {path} HTTP/1.1\r\n"
            f"Host: {host}\r\n"
            f"Upgrade: websocket\r\n"
            f"Connection: Upgrade\r\n"
            f"Sec-WebSocket-Key: {key}\r\n"
            f"Sec-WebSocket-Version: 13\r\n"
            f"Authorization: Bearer {self.token}\r\n"
            f"User-Agent: {BROWSER_UA}\r\n"
            f"Origin: {self.origin}\r\n"
            f"Accept: */*\r\n"
            f"Accept-Language: en-US,en;q=0.9\r\n"
            f"Cache-Control: no-cache\r\n"
            f"Pragma: no-cache\r\n"
            f"\r\n"
        )
        sock.sendall(headers.encode("ascii"))
        buf = b""
        while b"\r\n\r\n" not in buf:
            chunk = sock.recv(4096)
            if not chunk:
                sock.close()
                raise WebSocketError("connection closed during handshake")
            buf += chunk
            if len(buf) > 64 * 1024:
                sock.close()
                raise WebSocketError("handshake too large")
        header, rest = buf.split(b"\r\n\r\n", 1)
        status_line = header.split(b"\r\n", 1)[0].decode("latin1", "replace")
        parts = status_line.split()
        status = int(parts[1]) if len(parts) > 1 and parts[1].isdigit() else 0
        if status != 101:
            sock.close()
            raise WebSocketError(f"HTTP {status or status_line}", status=status or None)
        expected = base64.b64encode(hashlib.sha1((key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11").encode()).digest()).decode()
        got_accept = ""
        for line in header.split(b"\r\n")[1:]:
            if b":" not in line:
                continue
            name, value = line.split(b":", 1)
            if name.decode("latin1").strip().lower() == "sec-websocket-accept":
                got_accept = value.decode("latin1").strip()
        if got_accept and got_accept != expected:
            sock.close()
            raise WebSocketError("bad Sec-WebSocket-Accept")
        self.sock = sock
        self._buf = rest
        sock.settimeout(PING_INTERVAL)

    def send_text(self, text: str) -> None:
        self._send_frame(0x1, text.encode("utf-8"))

    def ping(self) -> None:
        self._send_frame(0x9, b"")

    def close(self) -> None:
        sock = self.sock
        self.sock = None
        if sock is None:
            return
        try:
            self._send_frame(0x8, struct.pack("!H", 1000), sock=sock)
        except Exception:
            pass
        try:
            sock.close()
        except Exception:
            pass

    def recv(self) -> str:
        while True:
            opcode, payload = self._read_frame()
            if opcode == 0x8:
                raise WebSocketError("closed", status=None)
            if opcode == 0x9:
                self._send_frame(0xA, payload)
                continue
            if opcode == 0xA:
                continue
            if opcode == 0x1:
                return payload.decode("utf-8", "replace")
            if opcode == 0x2:
                return payload.decode("utf-8", "replace")

    def _send_frame(self, opcode: int, payload: bytes, sock: ssl.SSLSocket | socket.socket | None = None) -> None:
        target = sock or self.sock
        if target is None:
            raise WebSocketError("not connected")
        mask = os.urandom(4)
        masked = bytes(b ^ mask[i % 4] for i, b in enumerate(payload))
        header = bytes([0x80 | opcode])
        n = len(payload)
        if n < 126:
            header += bytes([0x80 | n])
        elif n < 65536:
            header += bytes([0x80 | 126]) + struct.pack("!H", n)
        else:
            header += bytes([0x80 | 127]) + struct.pack("!Q", n)
        target.sendall(header + mask + masked)

    def _recv_exact(self, n: int) -> bytes:
        if self.sock is None:
            raise WebSocketError("not connected")
        while len(self._buf) < n:
            try:
                chunk = self.sock.recv(max(4096, n - len(self._buf)))
            except socket.timeout as exc:
                raise TimeoutError("recv timeout") from exc
            if not chunk:
                raise WebSocketError("closed")
            self._buf += chunk
        data, self._buf = self._buf[:n], self._buf[n:]
        return data

    def _read_frame(self) -> tuple[int, bytes]:
        b1, b2 = self._recv_exact(2)
        opcode = b1[0] & 0x0F
        masked = b2[0] & 0x80
        length = b2[0] & 0x7F
        if length == 126:
            length = struct.unpack("!H", self._recv_exact(2))[0]
        elif length == 127:
            length = struct.unpack("!Q", self._recv_exact(8))[0]
        mask = self._recv_exact(4) if masked else None
        payload = self._recv_exact(length)
        if mask:
            payload = bytes(b ^ mask[i % 4] for i, b in enumerate(payload))
        return opcode, payload


class InboxClient:
    def __init__(self, base: str, mailbox_id: str, token: str, wait: int) -> None:
        self.base = base.rstrip("/")
        self.mailbox_id = mailbox_id
        self.token = token
        self.wait = wait
        self.ctx = ssl.create_default_context()
        parsed = urllib.parse.urlparse(self.base)
        self.origin = f"{parsed.scheme}://{parsed.netloc}"

    def _headers(self) -> dict[str, str]:
        return {
            "Authorization": f"Bearer {self.token}",
            "Accept": "application/json",
            "Content-Type": "application/json",
            "User-Agent": BROWSER_UA,
            "Origin": self.origin,
            "Accept-Language": "en-US,en;q=0.9",
        }

    def _request(self, method: str, path: str, body: dict | None = None, timeout: float = 35.0) -> dict:
        url = f"{self.base}{path}"
        data = None if body is None else json.dumps(body).encode("utf-8")
        req = urllib.request.Request(url, data=data, method=method, headers=self._headers())
        try:
            with urllib.request.urlopen(req, timeout=timeout, context=self.ctx) as resp:
                raw = resp.read().decode("utf-8", "replace")
                return json.loads(raw) if raw else {}
        except urllib.error.HTTPError as exc:
            raw = exc.read().decode("utf-8", "replace")
            try:
                payload = json.loads(raw)
            except json.JSONDecodeError:
                payload = {"error": raw[:200] or exc.reason}
            raise RuntimeError(f"HTTP {exc.code}: {payload.get('error', raw[:200])}") from exc

    def agent_ws_url(self) -> str:
        parsed = urllib.parse.urlparse(self.base)
        scheme = "wss" if parsed.scheme == "https" else "ws"
        netloc = parsed.netloc
        return f"{scheme}://{netloc}/v1/mailboxes/{self.mailbox_id}/agent"

    def pending(self) -> list[dict]:
        path = f"/v1/mailboxes/{self.mailbox_id}/pending?wait={self.wait}"
        payload = self._request("GET", path, timeout=self.wait + 10)
        return list(payload.get("messages") or [])

    def reply(self, text: str, in_reply_to: str | None = None) -> dict:
        body: dict = {"text": text}
        if in_reply_to:
            body["in_reply_to"] = in_reply_to
        return self._request("POST", f"/v1/mailboxes/{self.mailbox_id}/replies", body)

    def ack(self, ids: list[str]) -> dict:
        if not ids:
            return {"acked": 0}
        return self._request("POST", f"/v1/mailboxes/{self.mailbox_id}/ack", {"ids": ids})

    def open_agent_socket(self) -> WebSocketClient:
        ws = WebSocketClient(self.agent_ws_url(), self.token, origin=self.origin)
        ws.connect()
        return ws


def build_reply(text: str, args: list[str]) -> str:
    handler = env("CHATTIES_INBOX_HANDLER")
    auto = env("CHATTIES_INBOX_AUTO_REPLY")
    if "--echo" in args:
        return f"[inbox] {text}"
    if handler:
        proc = subprocess.run(
            handler,
            input=text,
            text=True,
            capture_output=True,
            shell=True,
            timeout=120,
            check=False,
        )
        out = (proc.stdout or "").strip()
        if proc.returncode != 0 and not out:
            err = (proc.stderr or "").strip() or f"handler exit {proc.returncode}"
            return f"(handler failed: {err})"
        return out or "(empty handler reply)"
    if auto:
        return auto.replace("{text}", text)
    return (
        "Chatties inbox received your message. "
        "Set CHATTIES_INBOX_HANDLER on this box to reply with the local agent."
    )


def handle_messages(client: InboxClient, messages: list[dict], args: list[str], via_ws: WebSocketClient | None) -> None:
    acked: list[str] = []
    for msg in messages:
        mid = str(msg.get("id") or "")
        text = str(msg.get("text") or "")
        print(f"<< {mid} {text}", flush=True)
        reply = build_reply(text, args)
        if not reply.strip():
            continue
        try:
            if via_ws is not None:
                payload: dict = {"type": "reply", "text": reply}
                if mid:
                    payload["in_reply_to"] = mid
                via_ws.send_text(json.dumps(payload))
                print(f">> {reply}", flush=True)
            else:
                posted = client.reply(reply, in_reply_to=mid or None)
                rid = (posted.get("message") or {}).get("id", "")
                print(f">> {rid} {reply}", flush=True)
            if mid:
                acked.append(mid)
        except Exception as exc:  # noqa: BLE001
            print(f"reply error: {exc}", file=sys.stderr, flush=True)
    if acked:
        try:
            if via_ws is not None:
                via_ws.send_text(json.dumps({"type": "ack", "ids": acked}))
            else:
                client.ack(acked)
        except Exception as exc:  # noqa: BLE001
            print(f"ack error: {exc}", file=sys.stderr, flush=True)


def run_websocket(client: InboxClient, args: list[str]) -> bool:
    """Stay open until the socket dies. Returns True if any message was handled."""
    ws = client.open_agent_socket()
    got = False
    print(f"chatties-inbox-agent websocket connected {client.agent_ws_url()}", flush=True)
    try:
        while True:
            try:
                raw = ws.recv()
            except TimeoutError:
                try:
                    ws.ping()
                except Exception as exc:  # noqa: BLE001
                    raise WebSocketError(f"ping failed: {exc}") from exc
                continue
            except WebSocketError:
                return got
            if raw in ("ping", "pong"):
                continue
            try:
                body = json.loads(raw) if raw else {}
            except json.JSONDecodeError:
                print(f"bad ws frame: {raw[:120]}", file=sys.stderr, flush=True)
                continue
            kind = str(body.get("type") or "")
            if kind in ("hello", "ack", "error", "pong", "ping"):
                if kind == "hello":
                    got = True
                    print(f"hello mailbox {body.get('mailbox_id') or client.mailbox_id}", flush=True)
                elif kind == "error":
                    print(f"ws error: {body.get('error')}", file=sys.stderr, flush=True)
                continue
            messages = list(body.get("messages") or [])
            if kind in ("pending", "message") and messages:
                got = True
                handle_messages(client, messages, args, via_ws=ws)
    finally:
        ws.close()
    return got


def run_long_poll_once(client: InboxClient, args: list[str]) -> bool:
    """One pending wait. Held HTTP bills DO duration — fallback only."""
    messages = client.pending()
    if not messages:
        return False
    handle_messages(client, messages, args, via_ws=None)
    return True


def main(argv: list[str] | None = None) -> int:
    args = list(argv if argv is not None else sys.argv[1:])
    if "--help" in args or "-h" in args:
        print(__doc__)
        return 0
    base = env("CHATTIES_INBOX_URL", DEFAULT_URL)
    mailbox_id = env("CHATTIES_MAILBOX_ID")
    token = env("CHATTIES_AGENT_TOKEN")
    if not mailbox_id or not token:
        die("set CHATTIES_MAILBOX_ID and CHATTIES_AGENT_TOKEN")
    if not base.startswith("https://") and "--allow-http" not in args:
        die("CHATTIES_INBOX_URL must be https:// (TLS only)")
    try:
        wait = int(env("CHATTIES_INBOX_WAIT", str(DEFAULT_WAIT)))
    except ValueError:
        wait = DEFAULT_WAIT
    wait = max(0, min(wait, 30))
    client = InboxClient(base, mailbox_id, token, wait)
    print(f"chatties-inbox-agent websocket-first {base} mailbox {mailbox_id}", flush=True)
    failures = 0
    while True:
        got = False
        try:
            got = run_websocket(client, args)
            failures = 0
            continue
        except WebSocketError as exc:
            status = exc.status
            print(f"websocket: {exc}", file=sys.stderr, flush=True)
            if status in (426, 404, 405) or (status is not None and status >= 400):
                try:
                    print(
                        "websocket unavailable; long-poll fallback (bills duration while waiting)",
                        file=sys.stderr,
                        flush=True,
                    )
                    got = run_long_poll_once(client, args) or got
                except Exception as poll_exc:  # noqa: BLE001
                    print(f"poll error: {poll_exc}", file=sys.stderr, flush=True)
        except Exception as exc:  # noqa: BLE001
            print(f"websocket error: {exc}", file=sys.stderr, flush=True)
            try:
                got = run_long_poll_once(client, args) or got
            except Exception as poll_exc:  # noqa: BLE001
                print(f"poll error: {poll_exc}", file=sys.stderr, flush=True)
        if got:
            failures = 0
            continue
        failures += 1
        delay = jittered_backoff(failures)
        print(f"reconnect in {delay:.1f}s", file=sys.stderr, flush=True)
        time.sleep(delay)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
