"""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 inspect
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:
store = self._store_factory()
if inspect.isawaitable(store): # the pool behind the store is async
store = await store
key = await store.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
async def _factory() -> McpKeyStore:
return McpKeyStore(await get_session_store()._get_pool())
_resolver = KeyResolver(_factory)
return _resolver