"""Client MCP stdio tối giản cho `zalo-agent-cli mcp start` (thay `mcp_client.pool` của thansa-os).

Điểm then chốt: MCP của zalo-agent-cli giữ BỘ ĐỆM TIN trong chính tiến trình và đánh số thứ tự từ 1 mỗi lần
khởi động. Vì thế phải giữ tiến trình SỐNG lâu (một session cho mỗi tài khoản), và phải biết khi nào nó bị dựng
lại (`epoch` tăng) để bỏ con trỏ `since` cũ.
"""
from __future__ import annotations

import asyncio
import itertools
import json
import os
import time
from typing import Any, Dict, List, Optional

from .cli import kill_tree, no_window_kwargs

PROTOCOL = "2025-06-18"
_EPOCH = itertools.count(1)
_STREAM_LIMIT = 16 * 1024 * 1024


class McpError(RuntimeError):
    pass


class StdioMcp:
    def __init__(self, argv: List[str], env: Optional[Dict[str, str]] = None):
        self.argv = list(argv)
        self.env = dict(env or {})
        self.proc = None
        self.epoch = 0                      # 0 = chưa có phiên
        self._id = 0
        self._lock = asyncio.Lock()
        self._init = False
        self._stderr: List[str] = []

    def alive(self) -> bool:
        return self.proc is not None and self.proc.returncode is None

    async def _start(self):
        self.proc = await asyncio.create_subprocess_exec(
            *self.argv, stdin=asyncio.subprocess.PIPE, stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE,
            env=dict(os.environ, **self.env), limit=_STREAM_LIMIT, **no_window_kwargs())
        self.epoch = next(_EPOCH)
        asyncio.ensure_future(self._drain())

    async def _drain(self):
        try:
            while self.proc and self.proc.stderr:
                line = await self.proc.stderr.readline()
                if not line:
                    return
                self._stderr = (self._stderr + [line.decode("utf-8", "replace").rstrip()])[-40:]
        except Exception:
            pass

    async def _rpc(self, method: str, params=None, notify=False, timeout=120):
        if not self.alive():
            raise ConnectionError("tiến trình MCP đã chết: " + " | ".join(self._stderr[-5:]))
        self._id += 1
        msg: Dict[str, Any] = {"jsonrpc": "2.0", "method": method}
        if not notify:
            msg["id"] = self._id
        if params is not None:
            msg["params"] = params
        self.proc.stdin.write((json.dumps(msg, ensure_ascii=False) + "\n").encode("utf-8"))
        await self.proc.stdin.drain()
        if notify:
            return None
        deadline = time.time() + timeout
        while True:
            remain = deadline - time.time()
            if remain <= 0:
                raise TimeoutError(f"MCP không phản hồi sau {timeout}s")
            raw = await asyncio.wait_for(self.proc.stdout.readline(), timeout=remain)
            if not raw:
                raise ConnectionError("MCP đóng stdout: " + " | ".join(self._stderr[-5:]))
            try:
                obj = json.loads(raw.decode("utf-8", "replace").strip() or "null")
            except ValueError:
                continue                    # log lạc của npx
            if isinstance(obj, dict) and obj.get("id") == msg["id"]:
                return obj

    async def _ensure(self):
        if self.alive() and self._init:
            return
        await self.close()
        await self._start()
        # Lần đầu npx phải tải package -> chờ lâu.
        await self._rpc("initialize", {"protocolVersion": PROTOCOL, "capabilities": {},
                                       "clientInfo": {"name": "zalo-kit", "version": "1.0"}}, timeout=90)
        await self._rpc("notifications/initialized", notify=True)
        self._init = True

    async def call_tool(self, name: str, arguments: dict, timeout=120) -> Any:
        """Gọi tool, trả JSON đã bóc (dict/list) hoặc `{"text": ...}`. Lỗi tool -> McpError."""
        async with self._lock:
            try:
                await self._ensure()
                res = await self._rpc("tools/call", {"name": name, "arguments": arguments or {}}, timeout=timeout)
            except ConnectionError:
                # Chết TRƯỚC khi gửi: dựng lại một lần. (Tool ghi như send chỉ retry khi chắc chắn chưa gửi.)
                self._init = False
                await self._ensure()
                res = await self._rpc("tools/call", {"name": name, "arguments": arguments or {}}, timeout=timeout)
        if "error" in (res or {}):
            raise McpError(json.dumps(res["error"], ensure_ascii=False)[:500])
        result = (res or {}).get("result") or {}
        text = "\n".join(c.get("text", "") for c in (result.get("content") or [])
                         if isinstance(c, dict) and c.get("type") == "text")
        if result.get("isError"):
            raise McpError(text or "tool error")
        return parse_payload(text)

    async def list_tools(self) -> list:
        async with self._lock:
            await self._ensure()
            res = await self._rpc("tools/list", timeout=60)
        return ((res or {}).get("result") or {}).get("tools") or []

    async def close(self):
        self._init = False
        if self.proc is not None:
            try:
                kill_tree(self.proc)
                await asyncio.wait_for(self.proc.wait(), 5)
            except Exception:
                pass
        self.proc = None


def parse_payload(raw: Any) -> Any:
    if isinstance(raw, (dict, list)):
        return raw
    s = str(raw or "").strip()
    if not s:
        return {}
    try:
        return json.loads(s)
    except ValueError:
        pass
    for a_ch, b_ch in (("{", "}"), ("[", "]")):
        a, b = s.find(a_ch), s.rfind(b_ch)
        if 0 <= a < b:
            try:
                return json.loads(s[a:b + 1])
            except ValueError:
                continue
    return {"text": s}
