"""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
import navi.api.deps as deps
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"
class TestResolverWiring:
"""The lazily built process-wide resolver, on a real (fake) asyncpg pool."""
async def test_default_resolver_awaits_the_session_pool(self, encryptor, monkeypatch):
"""Regression: the factory used to wrap `_get_pool()`'s coroutine in
McpKeyStore, so `resolve` always raised AttributeError, swallowed it and
silently fell back to the default plaintext credential.
"""
import navi.mcp.keystore as ks
from tests.conftest_factory import FakeConnection, FakePool
conn = FakeConnection()
conn.enqueue(encryptor.encrypt("sk-user"))
pool = FakePool(conn)
class _SessionStore:
async def _get_pool(self):
return pool
monkeypatch.setattr(deps, "get_session_store", lambda: _SessionStore())
monkeypatch.setattr(ks, "_resolver", None)
resolver = ks.get_key_resolver()
assert await resolver.resolve("u1", "srv") == "sk-user"
assert [c[0] for c in conn.calls] == ["fetchval"]
async def test_missing_key_resolves_none(self, encryptor, monkeypatch):
"""The store really answers `None` (no row) — not an exception that the
resolver turned into a fallback: the query must have been issued.
"""
import navi.mcp.keystore as ks
from tests.conftest_factory import FakeConnection, FakePool
conn = FakeConnection() # nothing enqueued → no row
pool = FakePool(conn)
class _SessionStore:
async def _get_pool(self):
return pool
monkeypatch.setattr(deps, "get_session_store", lambda: _SessionStore())
monkeypatch.setattr(ks, "_resolver", None)
assert await ks.get_key_resolver().resolve("u1", "srv") is None
assert [c[0] for c in conn.calls] == ["fetchval"]