from __future__ import annotations

import asyncio
import logging
import re
import time
from collections import OrderedDict
from pathlib import Path
from typing import Any, Awaitable, Callable

from .client import McpClient
from .config import McpServerConfig, load_mcp_servers

logger = logging.getLogger(__name__)

# Longest one-line server summary kept in the system prompt. Servers whose first
# sentence is a poor hook should set ``summary`` in their config instead.
_SUMMARY_LIMIT = 200


def summarize_instructions(text: str, limit: int = _SUMMARY_LIMIT) -> str:
    """The first sentence of a server's instructions, capped at *limit* chars.

    This is the line the system prompt keeps for a server; the rest of the text is
    one ``tool_manual("<server>")`` call away. Everything a server says about itself
    is aimed at a model deciding *whether* to reach for it, and that decision only
    needs the opening claim — the query mappings and workflows underneath it are
    read when the server is actually used.
    """
    flat = " ".join((text or "").split())
    if not flat:
        return ""
    head = re.split(r"(?<=[.!?])\s+", flat, maxsplit=1)[0]
    if len(head) > limit:
        head = head[:limit].rsplit(" ", 1)[0].rstrip(" ,;:.—-") + "…"
    return head

# Per-user client cache (BYOK): hard size cap to bound stdio subprocesses.
_USER_CLIENT_LIMIT = 32
# Per-user clients unused for this long are dropped (health-check tick).
_USER_CLIENT_IDLE_SEC = 30.0 * 60.0


class McpManager:
    """Holds a pool of :class:`McpClient` instances and manages their lifecycle.

    Typical usage at application startup::

        manager = McpManager()
        await manager.load_all()
        # … later, on reload_tools built-in invocation …
        await manager.reload_all()
    """

    def __init__(self, config_path: str | Path | None = None) -> None:
        self.config_path = config_path
        self._clients: dict[str, McpClient] = {}
        self._configs: dict[str, McpServerConfig] | None = None

        # Callback invoked after a server successfully connects (or reconnects).
        # Signature: async def callback(server_name: str) -> None
        self._on_server_connected: Callable[[str], Awaitable[None]] | None = None

        # Background health-check task
        self._health_check_task: asyncio.Task | None = None
        self._health_check_interval: float = 30.0

        # Last known connected status per server — used to suppress duplicate
        # "connected" toast notifications on every health-check poll.
        self._connected_status: dict[str, bool] = {}

        # BYOK: per-user client clones keyed (server_name, user_id) ->
        # (client, key it was built with, monotonic last-used). Ordered
        # oldest-first for LRU eviction.
        self._user_clients: OrderedDict[
            tuple[str, str], tuple[McpClient, str, float]
        ] = OrderedDict()
        self._key_resolver: Any = None

    @property
    def clients(self) -> dict[str, McpClient]:
        return self._clients

    def set_on_server_connected(self, callback: Callable[[str], Awaitable[None]] | None) -> None:
        """Set a callback that is called whenever an MCP server comes online."""
        self._on_server_connected = callback

    def start_health_check(self) -> None:
        """Start the background health-check loop if not already running."""
        if self._health_check_task is None:
            self._health_check_task = asyncio.create_task(self._health_check_loop())

    async def stop_health_check(self) -> None:
        """Stop the background health-check loop."""
        if self._health_check_task:
            self._health_check_task.cancel()
            try:
                await self._health_check_task
            except asyncio.CancelledError:
                pass
            self._health_check_task = None

    def _get_configs(self) -> dict[str, McpServerConfig]:
        """Return cached configs, falling back to disk if not yet loaded."""
        if self._configs is None:
            self._configs = load_mcp_servers(self.config_path)
        return self._configs

    async def load_all(self, configs: dict[str, McpServerConfig] | None = None) -> None:
        """Connect to every server in *configs* (or load from disk).

        Existing clients are disconnected first so that a reload is clean.
        Servers that fail to connect are still added to the pool so the
        health-check loop can retry them later.
        """
        if configs is None:
            configs = load_mcp_servers(self.config_path)
        self._configs = configs

        # disconnect old
        await self.disconnect_all()

        # connect new
        for name, cfg in configs.items():
            client = McpClient(name, cfg)
            try:
                await client.connect()
                self._clients[name] = client
                self._connected_status[name] = True
                if self._on_server_connected:
                    await self._on_server_connected(name)
            except Exception as exc:
                logger.warning("MCP server %r failed to connect: %s", name, exc)
                # Keep the client in the pool so health-check can retry later
                self._clients[name] = client
                self._connected_status[name] = False

    async def reload_all(self) -> None:
        """Re-read the config file and reconnect every server."""
        self._configs = None  # bust cache
        configs = self._get_configs()
        await self.load_all(configs)

    async def disconnect_all(self) -> None:
        """Close every open connection and clear the client pool."""
        if not self._clients and not self._user_clients:
            return
        for name, client in list(self._clients.items()):
            try:
                await client.disconnect()
            except asyncio.CancelledError:
                pass
            except Exception as exc:
                logger.warning("MCP server %r disconnect error: %s", name, exc)
        self._clients.clear()
        self._connected_status.clear()
        for (server, user), (client, _, _) in list(self._user_clients.items()):
            try:
                await client.disconnect()
            except asyncio.CancelledError:
                pass
            except Exception as exc:
                logger.warning(
                    "MCP user client (server=%s user=%s) disconnect error: %s",
                    server, user, exc,
                )
        self._user_clients.clear()

    async def get_all_tools(self) -> list[tuple[str, Any]]:
        """Return ``(server_name, mcp_tool)`` for every tool on every server.

        Servers that fail to list tools are skipped gracefully.
        """
        out: list[tuple[str, Any]] = []
        for name, client in self._clients.items():
            try:
                tools = await client.list_tools()
                out.extend((name, t) for t in tools)
            except Exception as exc:
                logger.warning("MCP server %r list_tools failed: %s", name, exc)
        return out

    def resolve_group(self, server_name: str, group_name: str) -> list[str]:
        """Return the list of tool names in a server group.

        Reads from the static config (``mcp_servers.d/*.json``), not from the live
        server, so it works even when the server is temporarily disconnected.
        """
        configs = self._get_configs()
        cfg = configs.get(server_name)
        if cfg is None:
            return []
        return list(cfg.groups.get(group_name, []))

    def get_instructions(self, server_names: list[str] | set[str] | None = None) -> dict[str, str]:
        """Return combined instructions for every connected server.

        Server-provided instructions (from MCP initialize handshake) are merged
        with the overlay ``instructions`` field from ``mcp_servers.d/*.json``.
        If a selected server is disconnected, only the config overlay is returned.
        """
        configs = self._get_configs()
        out: dict[str, str] = {}
        if server_names is None:
            names = set(self._clients.keys())
        else:
            names = set(server_names)
        for name in names:
            client = self._clients.get(name)
            parts: list[str] = []
            if client and client.instructions:
                parts.append(client.instructions)
            cfg = configs.get(name)
            if cfg and cfg.instructions:
                if parts:
                    parts.append("")
                parts.append(cfg.instructions)
            if parts:
                out[name] = "\n".join(parts)
        return out

    def configured_servers(self) -> list[str]:
        """Every server named in ``mcp_servers.d/*.json``, connected or not."""
        return sorted(self._get_configs())

    def get_summaries(self, server_names: list[str] | set[str] | None = None) -> dict[str, str]:
        """One line per server, for the system prompt.

        The config's explicit ``summary`` wins; otherwise it is the first sentence
        of the server's instructions. Servers with neither are omitted rather than
        listed as a bare name — a name with no claim attached tells the agent nothing
        it could not read off the tool prefix.

        Full instructions stay reachable: ``tool_manual("<server>")``.
        """
        configs = self._get_configs()
        instructions = self.get_instructions(server_names)
        if server_names is None:
            names = set(instructions)
        else:
            names = set(server_names)
        names |= {name for name, cfg in configs.items() if cfg.summary}

        out: dict[str, str] = {}
        for name in sorted(names):
            cfg = configs.get(name)
            line = (cfg.summary if cfg and cfg.summary else "") or summarize_instructions(
                instructions.get(name, "")
            )
            if line:
                out[name] = line
        return out

    def server_tool_names(self, server_name: str) -> list[str]:
        """Tool names this server exposes, from the static config groups.

        Read from the config rather than the live server so that a disconnected
        server still documents what it would offer once it reconnects, and so a
        profile that enables only part of it sees the whole catalogue.
        """
        cfg = self._get_configs().get(server_name)
        if cfg is None:
            return []
        seen: list[str] = []
        for group_tools in cfg.groups.values():
            for tool_name in group_tools:
                if tool_name not in seen:
                    seen.append(tool_name)
        return sorted(seen)

    async def call_tool(
        self,
        server_name: str,
        tool_name: str,
        arguments: dict[str, Any] | None = None,
        *,
        user_id: str | None = None,
    ) -> tuple[str, bool]:
        """Proxy a tool call to the named server.

        *user_id* selects that user's BYOK credential when the server config
        declares a ``user_key`` slot; None (or no saved key) falls back to the
        default shared client. Keyword-only, so existing positional callers
        are unaffected.

        Returns (output_text, is_error) so the caller knows whether the MCP
        tool itself reported a failure.
        """
        client = await self._client_for(server_name, user_id)
        return await client.call_tool(tool_name, arguments)

    # ── Per-user (BYOK) clients ──────────────────────────────────────────────

    def set_key_resolver(self, resolver: Any) -> None:
        """Wire the KeyResolver (navi.mcp.keystore) into the call path."""
        self._key_resolver = resolver
        resolver.set_on_change(self._on_key_change)

    def _on_key_change(self, user_id: str | None, server_name: str | None) -> None:
        """A key changed — per-user clients built on it become stale."""
        stale = [
            k
            for k in self._user_clients
            if (user_id is None or k[1] == user_id)
            and (server_name is None or k[0] == server_name)
        ]
        for key in stale:
            client = self._user_clients.pop(key)[0]
            logger.info("MCP user client dropped (key changed): %s %s", key[0], key[1])
            self._disconnect_quietly(client)

    def _disconnect_quietly(self, client: McpClient) -> None:
        try:
            asyncio.get_running_loop().create_task(self._quiet_disconnect(client))
        except RuntimeError:
            # No running loop (interpreter shutdown) — the transport dies with it.
            pass

    @staticmethod
    async def _quiet_disconnect(client: McpClient) -> None:
        try:
            await client.disconnect()
        except asyncio.CancelledError:
            pass
        except Exception:
            logger.debug("MCP user-client disconnect error", exc_info=True)

    async def _client_for(self, server_name: str, user_id: str | None) -> McpClient:
        """The client to call with: default, or a per-user BYOK clone.

        The default path is checked first and never touches the resolver —
        servers without a ``user_key`` slot and user-less calls behave
        byte-for-byte as before.
        """
        if user_id is not None:
            cfg = self._get_configs().get(server_name)
            if cfg is not None and cfg.user_key is not None and self._key_resolver is not None:
                try:
                    user_key = await self._key_resolver.resolve(user_id, server_name)
                except Exception:
                    # KeyResolver guarantees None on failure; belt-and-suspenders:
                    # the tool call must never break because of the BYOK layer.
                    logger.warning(
                        "MCP key resolver error for server=%s user=%s — using default",
                        server_name, user_id, exc_info=True,
                    )
                    user_key = None
                if user_key:
                    return await self._user_client(server_name, user_id, cfg, user_key)
        client = self._clients.get(server_name)
        if client is None:
            raise RuntimeError(f"MCP server {server_name!r} is not connected")
        return client

    async def _user_client(
        self, server_name: str, user_id: str, cfg: McpServerConfig, key: str
    ) -> McpClient:
        """A per-user McpClient clone; rebuilt transparently when the key changes."""
        entry = self._user_clients.get((server_name, user_id))
        if entry is not None and entry[1] == key:
            client, _, _ = entry
        else:
            if entry is not None:
                # Key rotated between resolutions — tear the old one down.
                self._user_clients.pop((server_name, user_id))
                self._disconnect_quietly(entry[0])
            client = McpClient(server_name, cfg.with_user_key(key))
            logger.info("MCP user client created: server=%s user=%s", server_name, user_id)
        self._user_clients[(server_name, user_id)] = (client, key, time.monotonic())
        self._user_clients.move_to_end((server_name, user_id))

        # Hard LRU cap: bound stdio subprocesses and idle HTTP clients.
        while len(self._user_clients) > _USER_CLIENT_LIMIT:
            evicted_key, (evicted_client, _, _) = self._user_clients.popitem(last=False)
            logger.info(
                "MCP user client evicted (LRU): server=%s user=%s", *evicted_key
            )
            self._disconnect_quietly(evicted_client)
        return client

    # ── Health check ─────────────────────────────────────────────────────────

    async def _health_check_loop(self) -> None:
        """Background task that periodically probes every configured server."""
        while True:
            try:
                await asyncio.sleep(self._health_check_interval)
                await self._run_health_check()
            except asyncio.CancelledError:
                raise
            except Exception:
                logger.exception("MCP health-check loop error")
                # Brief sleep to avoid tight error loops
                await asyncio.sleep(5.0)

    async def _run_health_check(self) -> None:
        """Probe all servers: reconnect dead ones, verify live ones are still alive."""
        from navi.core.event_bus import get_event_bus
        from navi.core.events import McpStatusUpdate

        for name, client in list(self._clients.items()):
            if not client.connected:
                # Dead server — try to bring it back
                try:
                    await client.connect()
                    logger.info("MCP server %r reconnected by health check", name)
                    self._connected_status[name] = True
                    if self._on_server_connected:
                        await self._on_server_connected(name)
                except Exception as exc:
                    logger.debug("MCP health-check reconnect failed for %r: %s", name, exc)
                    continue
                continue

            # Live server — make sure it still responds
            try:
                await client.list_tools()
            except Exception:
                logger.warning("MCP server %r dropped during health check", name)
                await client.mark_disconnected()
                self._connected_status[name] = False
                await get_event_bus().publish(
                    McpStatusUpdate(server_name=name, status="disconnected")
                )
                continue

            # Server is still connected — only notify if it was previously
            # known as disconnected (recovery), not on every routine poll.
            if self._connected_status.get(name):
                continue
            self._connected_status[name] = True
            try:
                tools = await client.list_tools()
                await get_event_bus().publish(
                    McpStatusUpdate(
                        server_name=name,
                        status="connected",
                        tool_count=len(tools),
                    )
                )
            except Exception:
                pass

        # Idle per-user clients: not part of the health-check itself (they
        # reconnect lazily with backoff when next called), only reaped here.
        self._evict_idle_user_clients()

    def _evict_idle_user_clients(self) -> None:
        now = time.monotonic()
        idle_keys = [
            k
            for k, (_, _, last_used) in self._user_clients.items()
            if now - last_used > _USER_CLIENT_IDLE_SEC
        ]
        for key in idle_keys:
            client = self._user_clients.pop(key)[0]
            logger.info(
                "MCP user client evicted (idle): server=%s user=%s", *key
            )
            self._disconnect_quietly(client)
