diff --git a/navi/core/container.py b/navi/core/container.py index 4641dfd..55751a0 100644 --- a/navi/core/container.py +++ b/navi/core/container.py @@ -114,6 +114,11 @@ mcp_manager = McpManager() await mcp_manager.load_all() + # BYOK: per-user MCP credentials resolve through the process-wide resolver. + from navi.mcp.keystore import get_key_resolver + + mcp_manager.set_key_resolver(get_key_resolver()) + from navi.tools._internal.terminal_manager import TerminalManager terminal_manager = TerminalManager() terminal_manager.start() diff --git a/navi/mcp/manager.py b/navi/mcp/manager.py index 80641eb..87e8b51 100644 --- a/navi/mcp/manager.py +++ b/navi/mcp/manager.py @@ -2,6 +2,8 @@ import asyncio import logging +import time +from collections import OrderedDict from pathlib import Path from typing import Any, Awaitable, Callable @@ -10,6 +12,11 @@ logger = logging.getLogger(__name__) +# 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. @@ -39,6 +46,14 @@ # "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 @@ -105,7 +120,7 @@ async def disconnect_all(self) -> None: """Close every open connection and clear the client pool.""" - if not self._clients: + if not self._clients and not self._user_clients: return for name, client in list(self._clients.items()): try: @@ -116,6 +131,17 @@ 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. @@ -170,16 +196,115 @@ out[name] = "\n".join(parts) return out - async def call_tool(self, server_name: str, tool_name: str, arguments: dict[str, Any] | None = None) -> tuple[str, bool]: + 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 await client.call_tool(tool_name, arguments) + 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 ───────────────────────────────────────────────────────── @@ -243,3 +368,21 @@ ) 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) diff --git a/navi/mcp/tools.py b/navi/mcp/tools.py index e3d7218..ee32f97 100644 --- a/navi/mcp/tools.py +++ b/navi/mcp/tools.py @@ -6,7 +6,13 @@ from pathlib import Path from navi.config import settings -from navi.tools._internal.base import Tool, ToolContext, ToolResult, current_session_id +from navi.tools._internal.base import ( + Tool, + ToolContext, + ToolResult, + current_session_id, + current_user_id, +) from .manager import McpManager @@ -96,8 +102,11 @@ forwarded[key] = self._normalize_path_param(forwarded[key]) try: + uid = ctx.user_id if ctx else None + if uid is None: + uid = current_user_id.get() output, is_error = await self._manager.call_tool( - self.server_name, self.tool_name, forwarded + self.server_name, self.tool_name, forwarded, user_id=uid ) if is_error: return ToolResult( diff --git a/tests/unit/mcp/test_manager_byok.py b/tests/unit/mcp/test_manager_byok.py new file mode 100644 index 0000000..0d37b2e --- /dev/null +++ b/tests/unit/mcp/test_manager_byok.py @@ -0,0 +1,228 @@ +"""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")] \ No newline at end of file diff --git a/tests/unit/test_mcp.py b/tests/unit/test_mcp.py index 9efd011..2f799b6 100644 --- a/tests/unit/test_mcp.py +++ b/tests/unit/test_mcp.py @@ -120,7 +120,9 @@ result = await tool.execute({"query": "foo"}) assert result.success assert result.output == "found 3 results" - mock_manager.call_tool.assert_awaited_once_with("book", "search", {"query": "foo"}) + mock_manager.call_tool.assert_awaited_once_with( + "book", "search", {"query": "foo"}, user_id=None + ) async def test_execute_mcp_error(self): mock_manager = AsyncMock(spec=McpManager) @@ -168,7 +170,8 @@ result = await tool.execute({"query": "foo"}) assert result.success mock_manager.call_tool.assert_awaited_once_with( - "book", "search", {"query": "foo", "session_id": "real-session-id"} + "book", "search", {"query": "foo", "session_id": "real-session-id"}, + user_id=None, ) finally: current_session_id.reset(token)