#!/usr/bin/env python3

"""Bridge adapter for the ExternalChannel JSON-RPC/stdio plugin protocol.

This keeps nullclaw core on a single external-channel contract while still
supporting the HTTP bridge shape introduced by the whatsmeow example:

- GET /health
- POST /poll
- POST /send
"""

from __future__ import annotations

import json
import os
import sys
import threading
import time
import urllib.error
import urllib.parse
import urllib.request
from collections import deque
from pathlib import Path
from typing import Any

MAX_SEEN_MESSAGE_IDS = 256
MAX_MESSAGE_LEN = 3500


class JsonRpcError(Exception):
    def __init__(self, code: int, message: str) -> None:
        super().__init__(message)
        self.code = code
        self.message = message


class WhatsAppWebPlugin:
    def __init__(self) -> None:
        self.write_lock = threading.Lock()
        self.state_lock = threading.Lock()
        self.poll_thread: threading.Thread | None = None
        self.poll_thread_stop_event: threading.Event | None = None

        self.runtime_name = os.environ.get("NULLCLAW_EXTERNAL_RUNTIME_NAME", "whatsapp_web")
        self.account_id = "default"
        self.bridge_url = "http://127.0.0.1:3301"
        self.api_key: str | None = None
        self.allow_from: list[str] = []
        self.group_allow_from: list[str] = []
        self.group_policy = "allowlist"
        self.poll_interval_ms = 1500
        self.bridge_timeout_ms = 10000
        self.state_dir: str | None = None

        self.cursor: str | None = None
        self.seen_message_ids: deque[str] = deque(maxlen=MAX_SEEN_MESSAGE_IDS)
        self.seen_message_id_index: set[str] = set()

    def run(self) -> None:
        try:
            for line in sys.stdin:
                line = line.strip()
                if not line:
                    continue
                self.handle_line(line)
        finally:
            self.stop()

    def handle_line(self, line: str) -> None:
        try:
            request = json.loads(line)
        except json.JSONDecodeError as exc:
            self.log(f"ignoring malformed JSON-RPC line: {exc}")
            return

        request_id = request.get("id")
        method = request.get("method")
        params = request.get("params")
        if not isinstance(method, str):
            if request_id is not None:
                self.respond_error(request_id, -32600, "invalid request")
            return
        if params is None:
            params = {}
        if not isinstance(params, dict):
            if request_id is not None:
                self.respond_error(request_id, -32602, "params must be an object")
            return

        try:
            result = self.dispatch(method, params)
        except JsonRpcError as exc:
            if request_id is not None:
                self.respond_error(request_id, exc.code, exc.message)
            return
        except Exception as exc:  # pragma: no cover - defensive runtime path
            self.log(f"request failed for method={method}: {exc}")
            if request_id is not None:
                self.respond_error(request_id, -32000, str(exc))
            return

        if request_id is not None:
            self.respond_result(request_id, result)

    def dispatch(self, method: str, params: dict[str, Any]) -> Any:
        if method == "get_manifest":
            return {
                "protocol_version": 2,
                "capabilities": {
                    "health": True,
                    "streaming": False,
                    "send_rich": False,
                    "typing": False,
                },
            }
        if method == "start":
            return self.start(params)
        if method == "stop":
            return {"stopped": self.stop()}
        if method == "health":
            return self.health()
        if method == "send":
            return self.send(params)
        raise JsonRpcError(-32601, f"unknown method: {method}")

    def start(self, params: dict[str, Any]) -> dict[str, Any]:
        runtime = params.get("runtime")
        if not isinstance(runtime, dict):
            raise JsonRpcError(-32602, "runtime must be an object")

        config = params.get("config")
        if config is None:
            config = {}
        if not isinstance(config, dict):
            raise JsonRpcError(-32602, "config must be an object")

        runtime_name = str(runtime.get("name") or self.runtime_name).strip()
        if not runtime_name:
            raise JsonRpcError(-32602, "runtime.name is required")

        self.runtime_name = runtime_name
        self.account_id = str(runtime.get("account_id") or "default").strip() or "default"
        self.bridge_url = str(config.get("bridge_url") or self.bridge_url).strip().rstrip("/")
        if not self.bridge_url:
            raise JsonRpcError(-32602, "config.bridge_url is required")
        if not self.is_valid_bridge_url(self.bridge_url):
            raise JsonRpcError(-32602, "config.bridge_url must be https:// or local loopback http://")

        api_key = config.get("api_key")
        self.api_key = str(api_key).strip() if isinstance(api_key, str) and api_key.strip() else None
        self.allow_from = self.parse_string_list(config.get("allow_from"))
        self.group_allow_from = self.parse_string_list(config.get("group_allow_from"))
        self.group_policy = str(config.get("group_policy") or "allowlist").strip() or "allowlist"
        self.poll_interval_ms = self.parse_positive_int(config.get("poll_interval_ms"), 1500)
        self.bridge_timeout_ms = self.parse_positive_int(config.get("timeout_ms"), 10000)
        state_dir = runtime.get("state_dir")
        self.state_dir = str(state_dir).strip() if isinstance(state_dir, str) and state_dir.strip() else None

        if not self.stop():
            raise JsonRpcError(-32030, "previous poll thread did not stop cleanly")
        self.cursor = None
        with self.state_lock:
            self.seen_message_ids.clear()
            self.seen_message_id_index.clear()
        self.load_state()

        self.poll_thread_stop_event = threading.Event()
        self.poll_thread = threading.Thread(
            target=self.poll_loop,
            args=(self.poll_thread_stop_event,),
            name="whatsapp-web-plugin-poll",
            daemon=True,
        )
        self.poll_thread.start()
        return {"started": True, "runtime": {"name": self.runtime_name, "account_id": self.account_id}}

    def stop(self) -> bool:
        if self.poll_thread_stop_event is not None:
            self.poll_thread_stop_event.set()
        if self.poll_thread is not None:
            self.poll_thread.join(timeout=max(self.bridge_timeout_ms / 1000.0, 1.0))
            if self.poll_thread.is_alive():
                self.log("poll thread did not stop within timeout")
                return False
            self.poll_thread = None
        self.poll_thread_stop_event = None
        return True

    def health(self) -> dict[str, Any]:
        payload = self.bridge_request("GET", "/health")
        ok = bool(payload.get("ok", True))
        connected = bool(payload.get("connected", True))
        logged_in = bool(payload.get("logged_in", True))
        return {
            "healthy": ok and connected and logged_in,
            "ok": ok,
            "connected": connected,
            "logged_in": logged_in,
        }

    def send(self, params: dict[str, Any]) -> dict[str, Any]:
        runtime = params.get("runtime")
        if not isinstance(runtime, dict):
            raise JsonRpcError(-32602, "runtime must be an object")
        if str(runtime.get("name") or "").strip() != self.runtime_name:
            raise JsonRpcError(-32602, "runtime.name mismatch")
        if str(runtime.get("account_id") or "").strip() != self.account_id:
            raise JsonRpcError(-32602, "runtime.account_id mismatch")

        message_payload = params.get("message")
        if not isinstance(message_payload, dict):
            raise JsonRpcError(-32602, "message must be an object")

        stage = str(message_payload.get("stage") or "final")
        if stage != "final":
            return {"accepted": False, "ignored_stage": stage}

        target = str(message_payload.get("target") or "").strip()
        if not target:
            raise JsonRpcError(-32602, "target is required")

        message = str(message_payload.get("text") or "")
        media = message_payload.get("media")
        if isinstance(media, list) and media:
            self.log("ignoring media attachments; bridge adapter is text-only")

        sent_chunks = 0
        for chunk in self.split_message(message):
            self.bridge_request(
                "POST",
                "/send",
                {
                    "account_id": self.account_id,
                    "to": target,
                    "text": chunk,
                },
            )
            sent_chunks += 1

        return {"accepted": True, "sent_chunks": sent_chunks}

    def poll_loop(self, stop_event: threading.Event) -> None:
        while not stop_event.is_set():
            dirty = False
            try:
                response = self.bridge_request(
                    "POST",
                    "/poll",
                    {
                        "account_id": self.account_id,
                        "cursor": self.cursor or "",
                    },
                )

                next_cursor = response.get("next_cursor")
                if next_cursor is not None:
                    next_cursor = str(next_cursor)
                    if next_cursor != (self.cursor or ""):
                        self.cursor = next_cursor
                        dirty = True

                messages = response.get("messages")
                if isinstance(messages, list):
                    for item in messages:
                        if not isinstance(item, dict):
                            continue
                        if self.publish_bridge_message(item):
                            dirty = True

                if dirty:
                    self.save_state()
            except Exception as exc:
                self.log(f"poll failed: {exc}")

            stop_event.wait(max(self.poll_interval_ms, 100) / 1000.0)

    def publish_bridge_message(self, item: dict[str, Any]) -> bool:
        sender = str(item.get("from") or item.get("sender_id") or "").strip()
        if not sender:
            return False

        text = str(item.get("text") or item.get("content") or "").strip()
        if not text:
            return False

        is_group = bool(item.get("is_group"))
        if not self.is_sender_allowed(sender, is_group):
            return False

        chat_id = str(item.get("chat_id") or item.get("group_id") or sender).strip()
        if not chat_id:
            return False

        message_id = str(item.get("id") or "").strip()
        if message_id and self.has_seen_message_id(message_id):
            return False

        peer_kind = "group" if is_group else "direct"
        if is_group:
            peer_id = str(item.get("group_id") or chat_id).strip() or chat_id
        else:
            peer_id = sender

        metadata = {
            "account_id": self.account_id,
            "is_group": is_group,
            "peer_kind": peer_kind,
            "peer_id": peer_id,
        }
        if message_id:
            metadata["message_id"] = message_id

        session_key = f"{self.runtime_name}:{self.account_id}:{peer_kind}:{peer_id}"
        self.notify(
            "inbound_message",
            {
                    "message": {
                        "sender_id": sender,
                        "chat_id": chat_id,
                        "text": text,
                        "session_key": session_key,
                        "metadata": metadata,
                    }
            },
        )

        if message_id:
            self.remember_message_id(message_id)
            return True
        return False

    def bridge_request(self, method: str, path: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
        url = f"{self.bridge_url}{path}"
        headers = {"Content-Type": "application/json"}
        if self.api_key:
            headers["Authorization"] = f"Bearer {self.api_key}"

        data = None
        if payload is not None:
            data = json.dumps(payload, separators=(",", ":"), ensure_ascii=True).encode("utf-8")

        request = urllib.request.Request(url, data=data, headers=headers, method=method)
        try:
            with urllib.request.urlopen(request, timeout=max(self.bridge_timeout_ms / 1000.0, 1.0)) as response:
                body = response.read()
        except urllib.error.HTTPError as exc:
            detail = exc.read().decode("utf-8", "replace")
            raise JsonRpcError(-32020, f"bridge HTTP {exc.code}: {detail or exc.reason}") from exc
        except urllib.error.URLError as exc:
            raise JsonRpcError(-32021, f"bridge request failed: {exc.reason}") from exc

        if not body:
            return {}
        try:
            parsed = json.loads(body.decode("utf-8"))
        except json.JSONDecodeError as exc:
            raise JsonRpcError(-32022, f"bridge returned invalid JSON: {exc}") from exc
        if not isinstance(parsed, dict):
            raise JsonRpcError(-32023, "bridge response must be a JSON object")
        return parsed

    def respond_result(self, request_id: Any, result: Any) -> None:
        self.write_line({"jsonrpc": "2.0", "id": request_id, "result": result})

    def respond_error(self, request_id: Any, code: int, message: str) -> None:
        self.write_line(
            {
                "jsonrpc": "2.0",
                "id": request_id,
                "error": {
                    "code": code,
                    "message": message,
                },
            }
        )

    def notify(self, method: str, params: dict[str, Any]) -> None:
        self.write_line({"jsonrpc": "2.0", "method": method, "params": params})

    def write_line(self, payload: dict[str, Any]) -> None:
        encoded = json.dumps(payload, separators=(",", ":"), ensure_ascii=True)
        with self.write_lock:
            sys.stdout.write(encoded)
            sys.stdout.write("\n")
            sys.stdout.flush()

    def log(self, message: str) -> None:
        print(f"[whatsapp-web-plugin] {message}", file=sys.stderr, flush=True)

    def parse_string_list(self, value: Any) -> list[str]:
        if not isinstance(value, list):
            return []
        result: list[str] = []
        for item in value:
            if not isinstance(item, str):
                continue
            normalized = item.strip()
            if normalized:
                result.append(normalized)
        return result

    def parse_positive_int(self, value: Any, default: int) -> int:
        try:
            parsed = int(value)
            if parsed > 0:
                return parsed
        except (TypeError, ValueError):
            pass
        return default

    def is_valid_bridge_url(self, raw_url: str) -> bool:
        try:
            parsed = urllib.parse.urlparse(raw_url)
        except ValueError:
            return False
        if parsed.scheme == "https":
            return bool(parsed.netloc)
        if parsed.scheme != "http":
            return False
        return parsed.hostname in {"127.0.0.1", "localhost", "::1"}

    def is_sender_allowed(self, sender: str, is_group: bool) -> bool:
        if not is_group:
            if not self.allow_from:
                return True
            return self.is_allowed(self.allow_from, sender)

        if self.group_policy == "disabled":
            return False
        if self.group_policy == "open":
            return True

        effective_allowlist = self.group_allow_from or self.allow_from
        if not effective_allowlist:
            return False
        return self.is_allowed(effective_allowlist, sender)

    def is_allowed(self, allowlist: list[str], sender: str) -> bool:
        for entry in allowlist:
            if entry == "*" or entry == sender:
                return True
        return False

    def split_message(self, text: str) -> list[str]:
        if not text:
            return []
        return [text[i : i + MAX_MESSAGE_LEN] for i in range(0, len(text), MAX_MESSAGE_LEN)]

    def has_seen_message_id(self, message_id: str) -> bool:
        with self.state_lock:
            return message_id in self.seen_message_id_index

    def remember_message_id(self, message_id: str) -> None:
        with self.state_lock:
            if message_id in self.seen_message_id_index:
                return
            if len(self.seen_message_ids) == self.seen_message_ids.maxlen:
                evicted = self.seen_message_ids.popleft()
                self.seen_message_id_index.discard(evicted)
            self.seen_message_ids.append(message_id)
            self.seen_message_id_index.add(message_id)

    def load_state(self) -> None:
        path = self.state_path()
        if not path.exists():
            return

        try:
            payload = json.loads(path.read_text(encoding="utf-8"))
        except Exception as exc:
            self.log(f"failed to load state from {path}: {exc}")
            return

        if not isinstance(payload, dict):
            return
        if payload.get("runtime_name") != self.runtime_name:
            return
        if payload.get("account_id") != self.account_id:
            return
        if payload.get("bridge_url") != self.bridge_url:
            return

        cursor = payload.get("cursor")
        if isinstance(cursor, str) and cursor:
            self.cursor = cursor

        seen = payload.get("seen_message_ids")
        if not isinstance(seen, list):
            return

        with self.state_lock:
            self.seen_message_ids.clear()
            self.seen_message_id_index.clear()
            for item in seen[-MAX_SEEN_MESSAGE_IDS:]:
                if not isinstance(item, str) or not item:
                    continue
                self.seen_message_ids.append(item)
                self.seen_message_id_index.add(item)

    def save_state(self) -> None:
        path = self.state_path()
        path.parent.mkdir(parents=True, exist_ok=True)
        with self.state_lock:
            payload = {
                "runtime_name": self.runtime_name,
                "account_id": self.account_id,
                "bridge_url": self.bridge_url,
                "cursor": self.cursor or "",
                "seen_message_ids": list(self.seen_message_ids),
            }
        path.write_text(json.dumps(payload, separators=(",", ":"), ensure_ascii=True) + "\n", encoding="utf-8")

    def state_path(self) -> Path:
        root = self.state_dir
        if root:
            base = Path(root)
        else:
            xdg_state_home = os.environ.get("XDG_STATE_HOME")
            if xdg_state_home:
                base = Path(xdg_state_home) / "nullclaw" / "external"
            else:
                base = Path.home() / ".local" / "state" / "nullclaw" / "external"
        return base / f"{self.runtime_name}-{self.normalized_account_id()}.json"

    def normalized_account_id(self) -> str:
        raw = self.account_id.strip() or "default"
        return "".join(ch if ch.isalnum() or ch in "._-" else "_" for ch in raw)


def main() -> int:
    plugin = WhatsAppWebPlugin()
    plugin.run()
    return 0


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