Newer
Older
navi-1 / navi / mcp / keystore.py
"""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