"""McpManager BYOK runtime — per-user clients, caching, fallback, LRU."""

import asyncio
from unittest.mock import MagicMock

import pytest

import navi.mcp.manager as manager_mod
from navi.mcp.config import McpServerConfig, McpUserKey


class FakeClient:
    """Stand-in for McpClient: records construction and disconnects."""

    instances: list["FakeClient"] = []

    def __init__(self, name, config):
        self.name = name
        self.config = config
        self.disconnected = False
        FakeClient.instances.append(self)

    async def call_tool(self, tool_name, arguments=None):
        return ("ok", False)

    async def disconnect(self):
        self.disconnected = True


class FakeResolver:
    def __init__(self, keys: dict | None = None):
        self.keys = keys or {}
        self.on_change = None
        self.resolves: list = []

    async def resolve(self, user_id, server_name):
        self.resolves.append((user_id, server_name))
        return self.keys.get((user_id, server_name))

    async def list_servers(self, user_id):
        return {server: None for (uid, server) in self.keys if uid == user_id}

    def set_on_change(self, cb):
        self.on_change = cb

    def invalidate(self, user_id, server_name=None):
        if self.on_change:
            self.on_change(user_id, server_name)


@pytest.fixture(autouse=True)
def fake_client(monkeypatch):
    FakeClient.instances = []
    monkeypatch.setattr(manager_mod, "McpClient", FakeClient)
    return FakeClient


@pytest.fixture
def manager():
    m = manager_mod.McpManager()
    m._configs = _configs()
    m._clients = {"http-server": FakeClient("http-server", _configs()["http-server"])}
    return m


def _configs():
    return {
        "http-server": McpServerConfig(
            transport="streamable_http",
            url="https://x",
            headers={"Authorization": "Bearer default"},
            user_key=McpUserKey(header="Authorization", prefix="Bearer "),
        ),
        "std-server": McpServerConfig(
            transport="stdio",
            command="uvx",
            env={"API_KEY": "default"},
            user_key=McpUserKey(env="API_KEY"),
        ),
        "plain-server": McpServerConfig(transport="sse", url="https://plain"),
    }


async def test_call_without_user_uses_default_client(manager):
    manager._clients["plain-server"] = FakeClient("plain-server", _configs()["plain-server"])
    out, is_err = await manager.call_tool("plain-server", "t")
    assert (out, is_err) == ("ok", False)


async def test_no_user_key_slot_resolves_never(manager):
    """A key-less server must not touch the resolver at all (zero regression)."""
    manager._clients["plain-server"] = FakeClient("plain-server", _configs()["plain-server"])
    resolver = FakeResolver(keys={("u1", "plain-server"): "sk"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("plain-server", "t", user_id="u1")
    assert resolver.resolves == []


async def test_no_user_id_resolves_never(manager):
    """A user-less, role-less call cannot be attributed to anyone, so it is
    refused rather than served from the owner's credential."""
    resolver = FakeResolver(keys={("u1", "http-server"): "sk"})
    manager.set_key_resolver(resolver)
    with pytest.raises(RuntimeError, match="not available to this account"):
        await manager.call_tool("http-server", "t")
    assert resolver.resolves == []


async def test_user_key_builds_per_user_client(manager, fake_client):
    resolver = FakeResolver(keys={("u1", "http-server"): "sk-user"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("http-server", "t", user_id="u1")

    user_clients = [c for c in fake_client.instances if c is not manager._clients["http-server"]]
    assert len(user_clients) == 1
    assert user_clients[0].config.headers["Authorization"] == "Bearer sk-user"
    # The default config itself is untouched.
    assert manager._clients["http-server"].config.headers["Authorization"] == "Bearer default"


async def test_user_key_env_substitution(manager, fake_client):
    manager._clients["std-server"] = FakeClient("std-server", _configs()["std-server"])
    resolver = FakeResolver(keys={("u1", "std-server"): "sk-env"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("std-server", "t", user_id="u1")

    user_clients = [c for c in fake_client.instances if c.config is not None and c.config.env]
    assert user_clients[-1].config.env["API_KEY"] == "sk-env"
    assert user_clients[-1].config.env["API_KEY"] != "default"


async def test_a_user_without_a_key_is_refused_not_served_the_owner_credential(manager, fake_client):
    """This is the rule, not an edge case: the silent fallback to the shared
    credential is what made six of the owner's servers reachable by any account."""
    resolver = FakeResolver(keys={})  # user never saved a key
    manager.set_key_resolver(resolver)

    with pytest.raises(RuntimeError, match="not available to this account"):
        await manager.call_tool("http-server", "t", user_id="u1", role="user")

    assert len(fake_client.instances) == 1  # only the default client exists
    assert manager._clients["http-server"].config.headers["Authorization"] == "Bearer default"


async def test_an_admin_without_a_key_still_gets_the_shared_client(manager):
    """The shared credential in the config is the admin's own, so nothing about
    an admin's calls may change."""
    resolver = FakeResolver(keys={})
    manager.set_key_resolver(resolver)

    out, is_err = await manager.call_tool("http-server", "t", user_id="u1", role="admin")

    assert (out, is_err) == ("ok", False)


async def test_user_client_cached_until_key_changes(manager, fake_client):
    resolver = FakeResolver(keys={("u1", "http-server"): "sk-a"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("http-server", "t", user_id="u1")
    await manager.call_tool("http-server", "t", user_id="u1")
    assert len(fake_client.instances) == 2  # default + one user client

    resolver.keys[("u1", "http-server")] = "sk-b"
    await manager.call_tool("http-server", "t", user_id="u1")
    assert len(fake_client.instances) == 3  # rebuilt on key change


async def test_resolver_error_refuses_an_ordinary_user(manager):
    """A key we cannot read is a key we do not have: the BYOK layer failing must
    never degrade into the owner's credential."""
    class BrokenResolver:
        async def resolve(self, user_id, server_name):
            raise RuntimeError("db down")

        def set_on_change(self, cb):
            pass

    manager.set_key_resolver(BrokenResolver())
    with pytest.raises(RuntimeError, match="not available to this account"):
        await manager.call_tool("http-server", "t", user_id="u1", role="user")


async def test_resolver_error_still_falls_back_for_an_admin(manager, fake_client):
    class BrokenResolver:
        async def resolve(self, user_id, server_name):
            raise RuntimeError("db down")

        def set_on_change(self, cb):
            pass

    manager.set_key_resolver(BrokenResolver())
    out, is_err = await manager.call_tool("http-server", "t", user_id="u1", role="admin")
    assert (out, is_err) == ("ok", False)


async def test_on_change_drops_user_clients(manager, fake_client):
    resolver = FakeResolver(keys={("u1", "http-server"): "sk-a"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("http-server", "t", user_id="u1")
    uc = [c for c in fake_client.instances if c.__class__ is FakeClient and "http-server" in c.name]
    user_client = uc[-1]

    resolver.invalidate("u1", "http-server")  # fires on_change
    assert (None, None) not in manager._user_clients
    assert ("http-server", "u1") not in manager._user_clients
    await asyncio.sleep(0)
    assert user_client.disconnected


async def test_disconnect_all_tears_down_user_clients(manager, fake_client):
    resolver = FakeResolver(keys={("u1", "http-server"): "sk-a"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("http-server", "t", user_id="u1")

    await manager.disconnect_all()
    assert not manager._user_clients
    assert all(c.disconnected for c in fake_client.instances)


async def test_lru_cap_evicts_oldest(manager, fake_client):
    resolver = FakeResolver(keys={(f"u{i}", "http-server"): f"sk-{i}" for i in range(40)})
    manager.set_key_resolver(resolver)
    for i in range(40):
        await manager.call_tool("http-server", "t", user_id=f"u{i}")
    assert len(manager._user_clients) == manager_mod._USER_CLIENT_LIMIT


async def test_idle_eviction(manager, fake_client):
    resolver = FakeResolver(keys={("u1", "http-server"): "sk"})
    manager.set_key_resolver(resolver)
    await manager.call_tool("http-server", "t", user_id="u1")

    manager._evict_idle_user_clients()
    assert ("http-server", "u1") in manager._user_clients

    # Backdate last-used past the idle threshold, then reap again.
    client, key, _ = manager._user_clients[("http-server", "u1")]
    manager._user_clients[("http-server", "u1")] = (client, key, 0.0)
    manager._evict_idle_user_clients()
    assert ("http-server", "u1") not in manager._user_clients


async def test_mcp_tool_forwards_user_id_and_role():
    """McpTool.execute прокидывает ctx.user_id и ctx.user_role в call_tool."""
    from navi.mcp.tools import McpTool
    from navi.tools._internal.base import ToolContext

    calls: list = []

    class RecordingManager:
        async def call_tool(self, server_name, tool_name, arguments=None, *, user_id=None, role=None):
            calls.append((tool_name, arguments, user_id, role))
            return ("ok", False)

    tool = McpTool(
        server_name="http-server",
        tool_name="t1",
        description="d",
        parameters={},
        manager=RecordingManager(),
    )
    ctx = ToolContext(session_id="s1", user_id="u9", user_role="admin")
    await tool.execute({"x": 1}, ctx)
    assert calls == [("t1", {"x": 1, "session_id": "s1"}, "u9", "admin")]


async def test_mcp_tool_defaults_the_role_to_user():
    """A ctx carrying no role must not read as admin — the default denies."""
    from navi.mcp.tools import McpTool
    from navi.tools._internal.base import ToolContext

    seen: list = []

    class RecordingManager:
        async def call_tool(self, server_name, tool_name, arguments=None, *, user_id=None, role=None):
            seen.append(role)
            return ("ok", False)

    tool = McpTool(
        server_name="http-server",
        tool_name="t1",
        description="d",
        parameters={},
        manager=RecordingManager(),
    )
    await tool.execute({}, ToolContext(session_id="s1", user_id="u9"))
    assert seen == ["user"]


class TestVisibleServers:
    """What build_tool_list is allowed to offer. A gated server must be absent
    from the list until the user has a key, or the model is shown a tool the
    manager refuses."""

    async def test_an_admin_is_unrestricted(self, manager):
        manager.set_key_resolver(FakeResolver(keys={}))
        assert await manager.visible_servers("admin", "u1") is None

    async def test_the_no_auth_local_user_is_unrestricted(self, manager):
        """NAVI_AUTH_ENABLED=false resolves everyone to role="admin"."""
        manager.set_key_resolver(FakeResolver(keys={}))
        assert await manager.visible_servers("admin", None) is None

    async def test_a_user_without_keys_keeps_only_the_ungated_servers(self, manager):
        manager.set_key_resolver(FakeResolver(keys={}))
        assert await manager.visible_servers("user", "u1") == frozenset({"plain-server"})

    async def test_a_user_with_a_key_gains_that_server(self, manager):
        manager.set_key_resolver(FakeResolver(keys={("u1", "http-server"): "sk"}))
        assert await manager.visible_servers("user", "u1") == frozenset(
            {"plain-server", "http-server"}
        )

    async def test_without_a_resolver_only_ungated_servers_are_offered(self, manager):
        """No BYOK layer means no way to prove a key exists — deny, never assume."""
        assert await manager.visible_servers("user", "u1") == frozenset({"plain-server"})

    async def test_saving_a_key_in_the_panel_returns_the_server(self, manager):
        """PUT /mcp-keys calls resolver.invalidate, which is the only hook the
        manager has — without it the panel would need a restart."""
        resolver = FakeResolver(keys={})
        manager.set_key_resolver(resolver)
        assert "http-server" not in await manager.visible_servers("user", "u1")

        resolver.keys[("u1", "http-server")] = "sk"
        resolver.invalidate("u1", "http-server")

        assert "http-server" in await manager.visible_servers("user", "u1")

    async def test_a_failing_key_list_offers_only_ungated_servers(self, manager):
        class BrokenResolver:
            async def list_servers(self, user_id):
                raise RuntimeError("db down")

            def set_on_change(self, cb):
                pass

        manager.set_key_resolver(BrokenResolver())
        assert await manager.visible_servers("user", "u1") == frozenset({"plain-server"})