"""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")]