Newer
Older
navi-1 / tests / unit / tools / test_mcp_meta_tools.py
""""test_mcp_tool / mcp_status after the BYOK change — user key semantics."""

from unittest.mock import MagicMock

from navi.mcp.config import McpServerConfig, McpUserKey
from navi.tools._internal.base import current_user_id, ToolContext
from navi.tools.mcp_status import McpStatusTool
from navi.tools.test_mcp_tool import TestMcpToolTool


def _configs():
    return {
        "http-server": McpServerConfig(
            transport="streamable_http",
            url="https://x",
            user_key=McpUserKey(header="Authorization", prefix="Bearer "),
        ),
        "plain": McpServerConfig(transport="sse", url="https://p"),
    }


class RecordingManager:
    def __init__(self):
        self.calls: list = []
        self._configs_dict = _configs()

    @property
    def clients(self):
        client = MagicMock()
        client.connected = True
        return {"http-server": client, "plain": client}

    def _get_configs(self):
        return self._configs_dict

    async def call_tool(self, server_name, tool_name, arguments=None, *, user_id=None):
        self.calls.append((server_name, tool_name, arguments, user_id))
        return ("ok", False)


class TestTestMcpToolForwarding:
    async def test_forwards_ctx_user_id(self):
        manager = RecordingManager()
        tool = TestMcpToolTool(mcp_manager=manager)
        result = await tool.execute(
            {"server_name": "http-server", "tool_name": "t"},
            ToolContext(session_id="s1", user_id="u9"),
        )
        assert result.success
        assert manager.calls == [("http-server", "t", {}, "u9")]

    async def test_forwards_contextvar_user_id_without_ctx(self):
        manager = RecordingManager()
        tool = TestMcpToolTool(mcp_manager=manager)
        token = current_user_id.set("u7")
        try:
            await tool.execute({"server_name": "http-server", "tool_name": "t"})
        finally:
            current_user_id.reset(token)
        assert manager.calls == [("http-server", "t", {}, "u7")]

    async def test_no_user_resolves_none(self):
        manager = RecordingManager()
        tool = TestMcpToolTool(mcp_manager=manager)
        await tool.execute({"server_name": "http-server", "tool_name": "t"})
        assert manager.calls == [("http-server", "t", {}, None)]


class TestMcpStatusByok:
    async def test_marks_byok_servers(self):
        manager = RecordingManager()
        tool = McpStatusTool(mcp_manager=manager)
        result = await tool.execute({})
        lines = result.output.splitlines()
        idx_http = next(i for i, l in enumerate(lines) if "http-server" in l)
        idx_plain = next(i for i, l in enumerate(lines) if "plain " in l)
        assert "accepts a per-user key via Authorization" in lines[idx_http + 1]
        assert "accepts a per-user key" not in lines[idx_plain]
        assert "accepts a per-user key" not in lines[idx_plain + 1]