diff --git a/navi/api/routes/mcp_keys.py b/navi/api/routes/mcp_keys.py new file mode 100644 index 0000000..b611ff0 --- /dev/null +++ b/navi/api/routes/mcp_keys.py @@ -0,0 +1,138 @@ +"""Per-user MCP credentials (BYOK) — REST for the settings UI. + +A server config (``mcp_servers.d/*.json``) may declare a ``user_key`` slot; +only such servers appear here. Users save their own key, which overrides the +default plaintext credential for their tool calls; users without a key fall +back to the default. Keys are never returned — only presence and update time. +""" + +from typing import Annotated + +import structlog +from fastapi import APIRouter, Depends, HTTPException +from pydantic import BaseModel, Field + +from navi.api.deps import require_user +from navi.auth import User + +log = structlog.get_logger() + +router = APIRouter(prefix="/mcp-keys", tags=["mcp-keys"]) + + +def _pool(): + from navi.api.deps import get_session_store + + return get_session_store()._get_pool() + + +def _eligible_configs(): + """Configs that accept a user key and are referenced by at least one profile. + + A server nobody connects to is noise in the settings UI even if its + config declares a user_key slot. + """ + from navi.api.deps import get_profile_registry + from navi.mcp.config import load_mcp_servers + + referenced: set[str] = set() + for profile in get_profile_registry().all(): + referenced.update(profile.tools.agent.mcp) + referenced.update(profile.tools.subagent.mcp) + + return { + name: cfg + for name, cfg in load_mcp_servers().items() + if name in referenced and cfg.user_key is not None + } + + +class McpKeyItem(BaseModel): + server_name: str + transport: str + key_type: str # header | env + key_location: str # header name or env var name + prefix: str | None = None + instructions: str | None = None + has_key: bool + updated_at: str | None = None + + +@router.get("") +async def list_mcp_keys( + user: Annotated[User, Depends(require_user)], +) -> dict: + keyed = _eligible_configs() + if not keyed: + return {"items": []} + + from navi.mcp.keystore import McpKeyStore + + keys = await McpKeyStore(await _pool()).list_servers(user.id) + + def _item(name: str, cfg) -> dict: + return McpKeyItem( + server_name=name, + transport=cfg.transport, + key_type="header" if cfg.user_key.header else "env", + key_location=cfg.user_key.header or cfg.user_key.env, + prefix=cfg.user_key.prefix, + instructions=cfg.instructions, + has_key=name in keys, + updated_at=keys[name].isoformat() if name in keys else None, + ).model_dump() + + return {"items": [_item(name, cfg) for name, cfg in sorted(keyed.items())]} + + +class SaveMcpKeyRequest(BaseModel): + key: str = Field(min_length=1, max_length=4096) + + +def _get_config(server_name: str): + """The config for *server_name*; 404 when missing or key-less.""" + from navi.mcp.config import load_mcp_servers + + cfg = load_mcp_servers().get(server_name) + if cfg is None or cfg.user_key is None: + raise HTTPException(status_code=404, detail="Unknown server or it accepts no user key") + return cfg + + +@router.put("/{server_name}") +async def save_mcp_key( + server_name: str, + payload: SaveMcpKeyRequest, + user: Annotated[User, Depends(require_user)], +) -> dict: + _get_config(server_name) + + from navi.mcp.keystore import McpKeyStore, get_key_resolver + + updated_at = await McpKeyStore(await _pool()).set( + user.id, server_name, payload.key.strip(), + ) + get_key_resolver().invalidate(user.id, server_name) + log.info("mcp_key.saved", user_id=user.id, server=server_name) + return { + "server_name": server_name, + "has_key": True, + "updated_at": updated_at.isoformat(), + } + + +@router.delete("/{server_name}", status_code=204) +async def delete_mcp_key( + server_name: str, + user: Annotated[User, Depends(require_user)], +) -> None: + _get_config(server_name) + + from navi.mcp.keystore import McpKeyStore, get_key_resolver + + deleted = await McpKeyStore(await _pool()).delete(user.id, server_name) + get_key_resolver().invalidate(user.id, server_name) + if not deleted: + # Key never existed — still 204, the desired end state is reached. + return + log.info("mcp_key.deleted", user_id=user.id, server=server_name) \ No newline at end of file diff --git a/navi/main.py b/navi/main.py index e1856cf..e38b0ee 100644 --- a/navi/main.py +++ b/navi/main.py @@ -13,7 +13,7 @@ from fastapi.staticfiles import StaticFiles from navi.api.deps import require_admin -from navi.api.routes import agents, api_tokens, auth, health, messages, peer, push, sessions, webhooks, synapse +from navi.api.routes import agents, api_tokens, auth, health, messages, mcp_keys, peer, push, sessions, webhooks, synapse from navi.api.routes.admin import router as admin_router from navi.api.websocket import router as ws_router from navi.config import settings @@ -119,6 +119,7 @@ from navi.core.pg_session_store import pending_sweep_loop from navi.push._ddl import ensure_tables as ensure_push_tables from navi.synapse._ddl import ensure_tables as ensure_synapse_tables + from navi.mcp._ddl import ensure_tables as ensure_mcp_tables # Ensure auth tables first (navi_users is referenced by other DDL). for attempt in range(1, 6): @@ -128,6 +129,7 @@ pool = await container.database.pool() await ensure_push_tables(pool) await ensure_synapse_tables(pool) + await ensure_mcp_tables(pool) break except Exception as e: if attempt < 5: @@ -261,6 +263,7 @@ app.include_router(synapse.webhook_router) app.include_router(synapse.targets_router) app.include_router(synapse.settings_router) +app.include_router(mcp_keys.router) app.include_router(admin_router) # Eval endpoints spend LLM tokens and read session data — admin only. # (With auth disabled require_admin resolves to the anonymous admin user.) diff --git a/navi/mcp/__init__.py b/navi/mcp/__init__.py index 6444326..aa9b0c1 100644 --- a/navi/mcp/__init__.py +++ b/navi/mcp/__init__.py @@ -1,5 +1,11 @@ from .client import McpClient -from .config import McpServerConfig, load_mcp_servers +from .config import McpServerConfig, McpUserKey, load_mcp_servers from .manager import McpManager -__all__ = ["McpClient", "McpServerConfig", "load_mcp_servers", "McpManager"] +__all__ = [ + "McpClient", + "McpServerConfig", + "McpUserKey", + "load_mcp_servers", + "McpManager", +] \ No newline at end of file diff --git a/navi/mcp/_ddl.py b/navi/mcp/_ddl.py new file mode 100644 index 0000000..9bc5bfe --- /dev/null +++ b/navi/mcp/_ddl.py @@ -0,0 +1,21 @@ +"""DDL for mcp_user_keys — per-user credentials for MCP servers (BYOK). + +A server config declares a ``user_key`` slot (header or env destination). +Users may save their own key, which overrides the default plaintext +credential in ``mcp_servers.d/*.json`` for their tool calls; users without a +key fall back to the default. Keys are stored Fernet-encrypted. +""" + +_DDL = """ +CREATE TABLE IF NOT EXISTS mcp_user_keys ( + user_id TEXT NOT NULL REFERENCES navi_users(id) ON DELETE CASCADE, + server_name TEXT NOT NULL, + key_enc TEXT NOT NULL, + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (user_id, server_name) +); +""" + + +async def ensure_tables(pool) -> None: + await pool.execute(_DDL) \ No newline at end of file diff --git a/navi/mcp/config.py b/navi/mcp/config.py index e3a8eef..931a3ab 100644 --- a/navi/mcp/config.py +++ b/navi/mcp/config.py @@ -5,11 +5,37 @@ from pathlib import Path from typing import Literal -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator logger = logging.getLogger(__name__) +class McpUserKey(BaseModel): + """Declares where a per-user credential is injected into a server config. + + Exactly one destination. When a user has saved their own key (see + ``navi/mcp/keystore.py``), the runtime clones the server config with + ``headers[header] = prefix + key`` (sse/streamable_http) or + ``env[env] = key`` (stdio); users without a key fall back to the default + plaintext credential in the config. + """ + + header: str | None = None + prefix: str | None = None + env: str | None = None + + @model_validator(mode="after") + def _exactly_one_destination(self) -> "McpUserKey": + if bool(self.header) == bool(self.env): + raise ValueError("user_key must declare exactly one of 'header' or 'env'") + if self.prefix and not self.header: + raise ValueError("'prefix' is only valid together with 'header'") + return self + + +_HTTP_TRANSPORTS = ("sse", "streamable_http") + + class McpServerConfig(BaseModel): """Configuration for a single MCP server.""" @@ -33,6 +59,41 @@ # instructions provided by the MCP server itself during the initialize handshake. instructions: str | None = None + # Per-user credential slot (Bring Your Own Key). None = the server needs no + # user key; all users share the default plaintext credential above. + user_key: McpUserKey | None = None + + @property + def accepts_user_key(self) -> bool: + return self.user_key is not None + + def with_user_key(self, key: str) -> "McpServerConfig": + """Return a deep copy of this config with *key* injected per user_key.""" + if self.user_key is None: + raise ValueError("Server accepts no user key") + cfg = self.model_copy(deep=True) + assert cfg.user_key is not None + if cfg.user_key.header: + headers = dict(cfg.headers or {}) + headers[cfg.user_key.header] = (cfg.user_key.prefix or "") + key + cfg.headers = headers + else: + env = dict(cfg.env or {}) + env[cfg.user_key.env or ""] = key + cfg.env = env + return cfg + + @model_validator(mode="after") + def _user_key_transport_compatible(self) -> "McpServerConfig": + if self.user_key is not None: + if self.user_key.header and self.transport not in _HTTP_TRANSPORTS: + raise ValueError( + "user_key.header requires transport 'sse' or 'streamable_http'" + ) + if self.user_key.env and self.transport != "stdio": + raise ValueError("user_key.env requires transport 'stdio'") + return self + @property def is_stdio(self) -> bool: return self.transport == "stdio" diff --git a/navi/mcp/keystore.py b/navi/mcp/keystore.py new file mode 100644 index 0000000..addf076 --- /dev/null +++ b/navi/mcp/keystore.py @@ -0,0 +1,144 @@ +"""Storage for per-user MCP credentials (BYOK), postgres + Fernet. + +The key replaces the default plaintext credential from ``mcp_servers.d/*.json`` +for a given user's tool calls. It is never returned by any API — only its +presence and update time. +""" + +from __future__ import annotations + +import logging +import time +from datetime import datetime, timezone +from typing import Callable + +from navi.auth.encrypt import get_encryptor + +log = logging.getLogger(__name__) + +_UPSERT = """ +INSERT INTO mcp_user_keys (user_id, server_name, key_enc, updated_at) +VALUES ($1, $2, $3, $4) +ON CONFLICT (user_id, server_name) +DO UPDATE SET key_enc = EXCLUDED.key_enc, updated_at = EXCLUDED.updated_at +""" + +_DELETE = """ +DELETE FROM mcp_user_keys WHERE user_id = $1 AND server_name = $2 +RETURNING server_name +""" + +_GET = """ +SELECT key_enc FROM mcp_user_keys WHERE user_id = $1 AND server_name = $2 +""" + +_LIST_SERVERS = """ +SELECT server_name, updated_at FROM mcp_user_keys WHERE user_id = $1 +""" + + +class McpKeyStore: + def __init__(self, pool): + self._pool = pool + + async def set(self, user_id: str, server_name: str, key: str) -> datetime: + updated_at = datetime.now(timezone.utc) + await self._pool.execute( + _UPSERT, user_id, server_name, get_encryptor().encrypt(key), updated_at, + ) + return updated_at + + async def delete(self, user_id: str, server_name: str) -> bool: + """Remove the user's key. False when there was none.""" + deleted = await self._pool.fetchval(_DELETE, user_id, server_name) + return deleted is not None + + async def get(self, user_id: str, server_name: str) -> str | None: + """The decrypted key, or None when the user has none.""" + enc = await self._pool.fetchval(_GET, user_id, server_name) + if not enc: + return None + return get_encryptor().decrypt(enc) + + async def list_servers(self, user_id: str) -> dict[str, datetime]: + """{server_name: updated_at} — servers where the user has a key.""" + rows = await self._pool.fetch(_LIST_SERVERS, user_id) + return {r["server_name"]: r["updated_at"] for r in rows} + + +class KeyResolver: + """Cached (user_id, server_name) -> decrypted key, TTL-bounded. + + Every failure — DB error, missing encryption key — resolves to ``None`` + with a warning: the BYOK layer must never break a tool call, the caller + falls back to the default config credential. + """ + + def __init__( + self, + store_factory: Callable[[], McpKeyStore], + *, + ttl: float = 30.0, + now: Callable[[], float] = time.monotonic, + ) -> None: + self._store_factory = store_factory + self._ttl = ttl + self._now = now + self._cache: dict[tuple[str, str], tuple[float, str | None]] = {} + self._on_change: Callable[[str, str | None], None] | None = None + + async def resolve(self, user_id: str, server_name: str) -> str | None: + """The user's key for *server_name*, or None (no key / any failure).""" + entry = self._cache.get((user_id, server_name)) + if entry is not None: + expires_at, key = entry + if self._now() < expires_at: + return key + del self._cache[(user_id, server_name)] + try: + key = await self._store_factory().get(user_id, server_name) + except Exception: + log.warning( + "mcp user key resolve failed: user=%s server=%s — falling back to default credential", + user_id, + server_name, + exc_info=True, + ) + return None + self._cache[(user_id, server_name)] = (self._now() + self._ttl, key) + return key + + def invalidate(self, user_id: str, server_name: str | None = None) -> None: + """Drop cached key(s); *server_name* None clears every server of the user.""" + if server_name is None: + for k in [k for k in self._cache if k[0] == user_id]: + del self._cache[k] + else: + self._cache.pop((user_id, server_name), None) + if self._on_change is not None: + self._on_change(user_id, server_name) + + def invalidate_all(self) -> None: + self._cache.clear() + if self._on_change is not None: + self._on_change(None, None) + + def set_on_change(self, cb: Callable[[str | None, str | None], None]) -> None: + """Callback fired on invalidation — used by McpManager to drop + per-user clients built on the old key.""" + self._on_change = cb + + +_resolver: KeyResolver | None = None + + +def get_key_resolver() -> KeyResolver: + """Process-wide resolver, lazily wired to the session-store pool.""" + global _resolver + if _resolver is None: + from navi.api.deps import get_session_store + + _resolver = KeyResolver( + lambda: McpKeyStore(get_session_store()._get_pool()), + ) + return _resolver \ No newline at end of file diff --git a/tests/unit/api/test_mcp_keys.py b/tests/unit/api/test_mcp_keys.py new file mode 100644 index 0000000..4563847 --- /dev/null +++ b/tests/unit/api/test_mcp_keys.py @@ -0,0 +1,189 @@ +"""Per-user MCP keys REST — GET filter, PUT/DELETE round-trip.""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi.testclient import TestClient + +from navi.config import Settings +from navi.mcp.config import McpServerConfig, McpUserKey + + +@pytest.fixture +def mcp_keys_client(monkeypatch): + """Anonymous-admin test client with faked configs and key store pool.""" + import types + + import navi.api.routes.mcp_keys as keys_mod + import navi.auth.deps as auth_deps + import navi.mcp.config as cfg_mod + from navi.main import app + + # Anonymous admin (auth disabled) — routes resolve without a cookie. + monkeypatch.setattr( + auth_deps, "settings", + Settings(_env_file=None, navi_persona_file="", navi_auth_enabled=False), + ) + monkeypatch.setattr(keys_mod, "_pool", AsyncMock(return_value=MagicMock())) + + # Four servers: two eligible (http-header, stdio-env), one eligible but + # not referenced by any profile, one without a user_key slot. + cfgs = { + "http-server": McpServerConfig( + transport="streamable_http", + url="https://http", + headers={"Authorization": "Bearer default"}, + user_key=McpUserKey(header="Authorization", prefix="Bearer "), + ), + "stdio-server": McpServerConfig( + transport="stdio", + command="uvx", + env={"API_KEY": "default"}, + user_key=McpUserKey(env="API_KEY"), + ), + "unreferenced": McpServerConfig( + transport="sse", + url="https://unref", + user_key=McpUserKey(header="Authorization"), + ), + "no-slider": McpServerConfig(transport="sse", url="https://plain"), + } + monkeypatch.setattr(cfg_mod, "load_mcp_servers", lambda path=None: cfgs) + + profile = types.SimpleNamespace( + tools=types.SimpleNamespace( + agent=types.SimpleNamespace(mcp={"http-server": [], "stdio-server": []}), + subagent=types.SimpleNamespace(mcp={}), + ) + ) + + class _Registry: + def all(self): + return [profile] + + import navi.api.deps as deps_mod + + monkeypatch.setattr(deps_mod, "get_profile_registry", lambda: _Registry()) + + return TestClient(app) + + +def test_get_returns_only_keyed_referenced_servers(mcp_keys_client, monkeypatch): + import navi.mcp.keystore as keystore_mod + + store = MagicMock() + store.list_servers = AsyncMock(return_value={}) + monkeypatch.setattr(keystore_mod, "McpKeyStore", MagicMock(return_value=store)) + + data = mcp_keys_client.get("/mcp-keys").json() + names = [i["server_name"] for i in data["items"]] + assert names == ["http-server", "stdio-server"] + by_name = {i["server_name"]: i for i in data["items"]} + assert by_name["http-server"]["key_type"] == "header" + assert by_name["http-server"]["key_location"] == "Authorization" + assert by_name["http-server"]["prefix"] == "Bearer " + assert by_name["http-server"]["has_key"] is False + assert by_name["http-server"]["updated_at"] is None + assert by_name["stdio-server"]["key_type"] == "env" + # The key value itself is never present anywhere in the response. + assert "default" not in mcp_keys_client.get("/mcp-keys").text + + +def test_get_empty_when_no_keyed_servers(mcp_keys_client, monkeypatch): + import navi.mcp.config as cfg_mod + + monkeypatch.setattr( + cfg_mod, "load_mcp_servers", + lambda path=None: {"no-slider": McpServerConfig(transport="sse", url="https://x")}, + ) + assert mcp_keys_client.get("/mcp-keys").json() == {"items": []} + + +def test_get_lists_saved_key_presence(mcp_keys_client, monkeypatch): + import navi.mcp.keystore as keystore_mod + + ts = datetime(2026, 10, 7, tzinfo=timezone.utc) + store = MagicMock() + store.list_servers = AsyncMock(return_value={"http-server": ts}) + monkeypatch.setattr(keystore_mod, "McpKeyStore", MagicMock(return_value=store)) + + data = mcp_keys_client.get("/mcp-keys").json() + by_name = {i["server_name"]: i for i in data["items"]} + assert by_name["http-server"]["has_key"] is True + assert by_name["http-server"]["updated_at"] == ts.isoformat() + assert by_name["stdio-server"]["has_key"] is False + + +class TestPut: + def test_save_and_invalidate(self, mcp_keys_client, monkeypatch): + import navi.mcp.keystore as keystore_mod + + store = MagicMock() + store = MagicMock() + store.set = AsyncMock(return_value=datetime(2026, 10, 7, tzinfo=timezone.utc)) + invalidated: list = [] + + class FakeResolver: + def invalidate(self, user_id, server_name): + invalidated.append((user_id, server_name)) + + monkeypatch.setattr(keystore_mod, "McpKeyStore", MagicMock(return_value=store)) + monkeypatch.setattr(keystore_mod, "get_key_resolver", lambda: FakeResolver()) + + resp = mcp_keys_client.put("/mcp-keys/http-server", json={"key": "sk-user"}) + assert resp.status_code == 200 + body = resp.json() + assert body["server_name"] == "http-server" + assert body["has_key"] is True + store.set.assert_awaited_once() + # key passed stripped, value never logged in the response + assert store.set.await_args.args[2] == "sk-user" + assert "sk-user" not in resp.text + assert len(invalidated) == 1 + + def test_save_404_on_unknown_server(self, mcp_keys_client): + resp = mcp_keys_client.put("/mcp-keys/nope", json={"key": "sk"}) + assert resp.status_code == 404 + + def test_save_404_on_keyless_server(self, mcp_keys_client): + resp = mcp_keys_client.put("/mcp-keys/no-slider", json={"key": "sk"}) + assert resp.status_code == 404 + + def test_save_422_on_empty_key(self, mcp_keys_client): + resp = mcp_keys_client.put("/mcp-keys/http-server", json={"key": ""}) + assert resp.status_code == 422 + + +class TestDelete: + def test_delete_204(self, mcp_keys_client, monkeypatch): + import navi.mcp.keystore as keystore_mod + + store = MagicMock() + store.delete = AsyncMock(return_value=True) + invalidated: list = [] + + class FakeResolver: + def invalidate(self, user_id, server_name): + invalidated.append((user_id, server_name)) + + monkeypatch.setattr(keystore_mod, "McpKeyStore", MagicMock(return_value=store)) + monkeypatch.setattr(keystore_mod, "get_key_resolver", lambda: FakeResolver()) + + resp = mcp_keys_client.delete("/mcp-keys/http-server") + assert resp.status_code == 204 + assert len(invalidated) == 1 + + def test_missing_key_is_still_204(self, mcp_keys_client, monkeypatch): + import navi.mcp.keystore as keystore_mod + + store = MagicMock() + store.delete = AsyncMock(return_value=False) + monkeypatch.setattr(keystore_mod, "McpKeyStore", MagicMock(return_value=store)) + + resp = mcp_keys_client.delete("/mcp-keys/http-server") + assert resp.status_code == 204 + + def test_delete_404_on_unknown_server(self, mcp_keys_client): + resp = mcp_keys_client.delete("/mcp-keys/nope") + assert resp.status_code == 404 \ No newline at end of file diff --git a/tests/unit/mcp/test_keystore.py b/tests/unit/mcp/test_keystore.py new file mode 100644 index 0000000..3058f2c --- /dev/null +++ b/tests/unit/mcp/test_keystore.py @@ -0,0 +1,240 @@ +"""McpUserKey config validation + McpKeyStore / KeyResolver unit tests.""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock + +import pytest +from cryptography.fernet import Fernet + +from navi.auth.encrypt import TokenEncryptor +from navi.mcp.config import McpServerConfig, McpUserKey +from navi.mcp.keystore import KeyResolver, McpKeyStore + + +from tests.conftest_factory import FakeRecord + + +class FakePool: + """Pool-level asyncpg fake: captures queries, returns canned results.""" + + def __init__(self): + self.queries: list[tuple[str, tuple]] = [] + self.executes: list = [] + self.fetch_returns: list = [] + self.fetchval_returns: list = [] + + @staticmethod + def _take(queue: list): + if not queue: + return None + result = queue.pop(0) + if isinstance(result, Exception): + raise result + return result + + async def execute(self, query, *args): + self.queries.append((query, args)) + return self._take(self.executes) or "OK" + + async def fetchval(self, query, *args): + self.queries.append((query, args)) + return self._take(self.fetchval_returns) + + async def fetch(self, query, *args): + self.queries.append((query, args)) + return self._take(self.fetch_returns) or [] + + +class TestUserKeyValidation: + def test_header_slot(self): + cfg = McpServerConfig( + transport="streamable_http", + url="https://x", + headers={"Authorization": "Bearer default"}, + user_key=McpUserKey(header="Authorization", prefix="Bearer "), + ) + assert cfg.accepts_user_key + + def test_env_slot(self): + cfg = McpServerConfig( + transport="stdio", + command="uvx", + env={"API_KEY": "default"}, + user_key=McpUserKey(env="API_KEY"), + ) + assert cfg.accepts_user_key + + def test_both_destinations_rejected(self): + with pytest.raises(Exception): + McpServerConfig(user_key=McpUserKey(header="H", env="E")) + + def test_no_destination_rejected(self): + with pytest.raises(Exception): + McpServerConfig(user_key=McpUserKey()) + + def test_header_requires_http_transport(self): + with pytest.raises(Exception): + McpServerConfig(transport="stdio", user_key=McpUserKey(header="Authorization")) + + def test_env_requires_stdio_transport(self): + with pytest.raises(Exception): + McpServerConfig(transport="sse", user_key=McpUserKey(env="API_KEY")) + + def test_no_user_key_is_default(self): + assert McpServerConfig().accepts_user_key is False + + +class TestWithUserKey: + def _http_cfg(self): + return McpServerConfig( + transport="streamable_http", + url="https://x", + headers={"Authorization": "Bearer default", "X-Other": "1"}, + user_key=McpUserKey(header="Authorization", prefix="Bearer "), + ) + + def test_injects_header_and_prefix(self): + patched = self._http_cfg().with_user_key("sk-123") + assert patched.headers["Authorization"] == "Bearer sk-123" + assert patched.headers["X-Other"] == "1" + + def test_no_prefix_appends_raw(self): + cfg = self._http_cfg() + cfg.user_key.prefix = None + assert cfg.with_user_key("sk").headers["Authorization"] == "sk" + + def test_env_injection(self): + cfg = McpServerConfig( + transport="stdio", + command="uvx", + env={"API_KEY": "default"}, + user_key=McpUserKey(env="API_KEY"), + ) + assert cfg.with_user_key("abc").env["API_KEY"] == "abc" + + def test_original_config_untouched(self): + cfg = self._http_cfg() + cfg.with_user_key("sk-123") + assert cfg.headers["Authorization"] == "Bearer default" + + def test_raises_when_no_slot(self): + with pytest.raises(ValueError): + McpServerConfig().with_user_key("sk") + + +@pytest.fixture +def encryptor(monkeypatch): + enc = TokenEncryptor(Fernet.generate_key().decode()) + import navi.mcp.keystore as ks + + monkeypatch.setattr(ks, "get_encryptor", lambda: enc) + return enc + + +class TestMcpKeyStore: + async def test_set_persists_encrypted(self, encryptor): + pool = FakePool() + store = McpKeyStore(pool) + ts = await store.set("u1", "srv", "sk-plain") + assert isinstance(ts, datetime) + query, args = pool.queries[0] + assert "ON CONFLICT (user_id, server_name)" in query + assert args[0] == "u1" and args[1] == "srv" + assert encryptor.decrypt(args[2]) == "sk-plain" + + async def test_get_missing_returns_none(self, encryptor): + store = McpKeyStore(FakePool()) + assert await store.get("u1", "srv") is None + + async def test_get_decrypts(self, encryptor): + pool = FakePool() + store = McpKeyStore(pool) + pool.fetchval_returns.append(encryptor.encrypt("sk-plain")) + assert await store.get("u1", "srv") == "sk-plain" + + async def test_delete_true_when_row_returned(self, encryptor): + pool = FakePool() + store = McpKeyStore(pool) + pool.fetchval_returns.append("srv") + assert await store.delete("u1", "srv") is True + + async def test_delete_false_when_no_row(self, encryptor): + store = McpKeyStore(FakePool()) + assert await store.delete("u1", "srv") is False + + async def test_list_servers(self, encryptor): + pool = FakePool() + store = McpKeyStore(pool) + ts = datetime.now(timezone.utc) + pool.fetch_returns.append([FakeRecord(server_name="a", updated_at=ts)]) + assert await store.list_servers("u1") == {"a": ts} + + +class TestKeyResolver: + @staticmethod + def _store(key="sk-1"): + store = AsyncMock(spec=McpKeyStore) + store.get = AsyncMock(return_value=key) + return store + + async def test_resolve_delegates_to_store(self, encryptor): + store = self._store() + resolver = KeyResolver(lambda: store) + assert await resolver.resolve("u1", "srv") == "sk-1" + store.get.assert_awaited_once_with("u1", "srv") + + async def test_resolve_caches_within_ttl(self, encryptor): + store = self._store() + clock = [1_000.0] + resolver = KeyResolver(lambda: store, ttl=30.0, now=lambda: clock[0]) + await resolver.resolve("u1", "srv") + await resolver.resolve("u1", "srv") + store.get.assert_awaited_once() + + clock[0] += 31.0 # past TTL -> re-resolve + await resolver.resolve("u1", "srv") + assert store.get.await_count == 2 + + async def test_missing_key_resolves_none(self, encryptor): + store = self._store(key=None) + resolver = KeyResolver(lambda: store) + assert await resolver.resolve("u1", "srv") is None + + async def test_db_error_resolves_none(self, encryptor): + store = AsyncMock(spec=McpKeyStore) + store.get = AsyncMock(side_effect=RuntimeError("db down")) + resolver = KeyResolver(lambda: store) + assert await resolver.resolve("u1", "srv") is None + + async def test_invalidate_forces_re_resolve(self, encryptor): + store = self._store() + resolver = KeyResolver(lambda: store) + await resolver.resolve("u1", "srv") + resolver.invalidate("u1", "srv") + await resolver.resolve("u1", "srv") + assert store.get.await_count == 2 + + async def test_invalidate_all(self, encryptor): + store = self._store() + resolver = KeyResolver(lambda: store) + await resolver.resolve("u1", "srv") + await resolver.resolve("u2", "srv") + resolver.invalidate_all() + await resolver.resolve("u1", "srv") + assert store.get.await_count == 3 + + async def test_on_change_fired_with_scope(self, encryptor): + events: list = [] + resolver = KeyResolver(lambda: self._store()) + resolver.set_on_change(lambda user, srv: events.append((user, srv))) + resolver.invalidate("u1", "srv") + resolver.invalidate("u2", None) + resolver.invalidate_all() + assert events == [("u1", "srv"), ("u2", None), (None, None)] + + async def test_error_not_cached_then_recovery(self, encryptor): + store = AsyncMock(spec=McpKeyStore) + store.get = AsyncMock(side_effect=[RuntimeError("down"), "sk-recovered"]) + resolver = KeyResolver(lambda: store) + assert await resolver.resolve("u1", "srv") is None + assert await resolver.resolve("u1", "srv") == "sk-recovered" \ No newline at end of file