Newer
Older
navi-1 / tests / unit / api / test_mcp_keys.py
"""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