Newer
Older
navi-1 / tests / unit / mcp / test_manager_byok.py
"""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))

    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):
    resolver = FakeResolver(keys={("u1", "http-server"): "sk"})
    manager.set_key_resolver(resolver)
    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_no_saved_key_falls_back_to_default(manager, fake_client):
    resolver = FakeResolver(keys={})  # user never saved a key
    manager.set_key_resolver(resolver)
    await manager.call_tool("http-server", "t", user_id="u1")
    default = manager._clients["http-server"]
    assert default.config.headers["Authorization"] == "Bearer default"
    assert len(fake_client.instances) == 1


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_falls_back_to_default(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")
    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():
    """McpTool.execute прокидывает ctx.user_id в manager.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):
            calls.append((tool_name, arguments, user_id))
            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")
    await tool.execute({"x": 1}, ctx)
    assert calls == [("t1", {"x": 1, "session_id": "s1"}, "u9")]