fix(mcp): scope health discovery for route-restricted keys

This commit is contained in:
Joshua Valluru 2026-09-17 09:17:55 -07:00
parent 4b368bf066
commit e21db01d67
4 changed files with 139 additions and 3 deletions

View file

@ -1254,7 +1254,7 @@ if MCP_AVAILABLE:
"""
user_mcp_management_mode: Final = _get_user_mcp_management_mode()
if user_mcp_management_mode == "view_all":
if user_mcp_management_mode == "view_all" and not _is_restricted_virtual_key_request(user_api_key_dict):
servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(server_ids=server_ids)
return [{"server_id": server.server_id, "status": server.status} for server in servers]

View file

@ -15,8 +15,9 @@ import re
import time
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
from pydantic import BaseModel, ConfigDict, Field, RootModel
from e2e_config import settle_propagation
from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap
@ -46,6 +47,19 @@ class McpServerNewResponse(BaseModel):
server_id: str
class McpHealthParams(BaseModel):
server_ids: list[str] | None = None
class McpHealthRow(BaseModel):
server_id: str
status: Literal["healthy", "unhealthy", "unknown"] | None
class McpHealthResponse(RootModel[list[McpHealthRow]]):
pass
class McpToolMcpInfo(BaseModel):
server_id: str | None = None
alias: str | None = None
@ -187,6 +201,22 @@ class McpClient:
)
).root
def list_servers(self, key: str) -> Result[McpServerListResponse]:
return self.proxy.transport.get(
"/v1/mcp/server",
headers=ApiKeyHeaders(x_litellm_api_key=key),
params=NoBody(),
response_type=McpServerListResponse,
)
def server_health(self, key: str, server_ids: list[str] | None = None) -> Result[McpHealthResponse]:
return self.proxy.transport.get(
"/v1/mcp/server/health",
headers=ApiKeyHeaders(x_litellm_api_key=key),
params=McpHealthParams(server_ids=server_ids),
response_type=McpHealthResponse,
)
def await_registered(self, server_id: str) -> None:
"""Poll /v1/mcp/server until `server_id` is listed. Fails at poll_timeout.

View file

@ -13,12 +13,14 @@ and must be refused with a 403 on `tools/call`.
from __future__ import annotations
import pytest
from typing import Final
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
from e2e_config import DD_SEARCH_FROM, unique_marker
from e2e_http import unwrap
from lifecycle import ResourceManager
from mcp_client import McpClient
from models import KeyGenerateBody, ObjectPermission
pytestmark = pytest.mark.e2e
@ -108,3 +110,39 @@ class TestMcpKeyWithoutAccessIsDenied:
denied_key, server_id=server_id, name=tool_name, arguments=search_args
)
assert "access_denied" in denied.body, f"403 was not an MCP access denial: {denied.body}"
class TestMcpHealthVisibility:
def test_route_restricted_health_matches_server_grants(
self,
client: McpClient,
resources: ResourceManager,
) -> None:
server_x: Final = register_datadog_mcp(client, resources)
server_y: Final = register_datadog_mcp(client, resources)
client.await_registered(server_x)
client.await_registered(server_y)
permitted: Final = _key(client, resources, mcp_servers=[server_x])
tool: Final = client.await_tool(permitted, server_x, SEARCH_LOGS_TOOL)
result: Final = client.await_call_tool(
permitted, server_id=server_x, name=tool,
arguments={"query": "service:litellm", "from": DD_SEARCH_FROM, "to": "now", "max_tokens": 1000},
)
assert result.is_error is not True, f"permitted control failed: {result}"
for grants in ([server_x], [server_y], []):
key = client.proxy.generate_key(KeyGenerateBody(
user_id=f"e2e-mcp-health-{unique_marker()}",
allowed_routes=["/v1/mcp/server", "/v1/mcp/server/health"],
object_permission=ObjectPermission(mcp_servers=grants),
))
resources.defer(lambda key=key: client.proxy.delete_key(key))
listed = unwrap(client.list_servers(key)).root
assert {row.server_id for row in listed} == set(grants)
for requested in (None, [server_y], [server_x, server_y]):
health = unwrap(client.server_health(key, requested)).root
expected = set(grants) if requested is None else set(grants).intersection(requested)
assert {row.server_id for row in health} == expected, (
f"health disclosed servers outside grants {grants}, requested {requested}: {health}"
)
assert all(row.status == "healthy" for row in health), f"upstream control unhealthy: {health}"

View file

@ -10,6 +10,7 @@ from typing import List, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from respx import MockRouter
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
@ -4040,7 +4041,7 @@ class TestHealthCheckServers:
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[mock_user_auth]),
AsyncMock(return_value=[mock_user_auth, mock_user_auth]),
),
):
result = await health_check_servers(
@ -4056,6 +4057,73 @@ class TestHealthCheckServers:
assert result[1]["status"] == "unhealthy"
@pytest.mark.asyncio
@pytest.mark.respx(assert_all_called=False)
@pytest.mark.parametrize(
("mode", "restricted", "grants", "requested", "expected", "upstream_status"),
[
("view_all", True, ("server-x",), None, ("server-x",), 200),
("view_all", True, ("server-x",), ("server-y",), (), 200),
("view_all", True, ("server-x",), ("server-x", "server-y"), ("server-x",), 200),
("view_all", True, (), None, (), 200),
("view_all", True, ("server-y",), None, ("server-y",), 200),
("view_all", True, ("server-x",), (), ("server-x",), 200),
("view_all", False, ("server-x",), None, ("server-x", "server-y"), 200),
("restricted", False, ("server-x",), None, ("server-x",), 200),
("restricted", True, ("server-x",), None, ("server-x",), 200),
("view_all", True, ("server-x",), None, ("server-x",), 503),
],
)
async def test_health_discovery_respects_route_restricted_key_grants(
respx_mock: MockRouter,
monkeypatch: pytest.MonkeyPatch,
mode: str,
restricted: bool,
grants: tuple[str, ...],
requested: tuple[str, ...] | None,
expected: tuple[str, ...],
upstream_status: int,
) -> None:
from typing import Final
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
manager: Final = mcp_server_manager.MCPServerManager()
manager.registry = {
server_id: MCPServer(
server_id=server_id, name=server_id, transport=MCPTransport.http,
spec_path=f"https://93.184.216.34/{server_id}.json", auth_type=MCPAuth.none,
)
for server_id in ("server-x", "server-y")
}
routes: Final = {
server_id: respx_mock.get(server.spec_path).respond(upstream_status, json={"paths": {}})
for server_id, server in manager.registry.items()
}
caller: Final = UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
api_key="test-health-key",
allowed_routes=["/v1/mcp/server", "/v1/mcp/server/health"] if restricted else [],
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="health-permissions", mcp_servers=list(grants)),
)
with (
patch.object(mgmt_endpoints, "global_mcp_server_manager", manager), # test-quality-ok: TQ008 inject real registry into legacy route binding
patch.object(mcp_server_manager, "global_mcp_server_manager", manager), # test-quality-ok: TQ008 share real registry with unchanged permission resolver
patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": mode}), # test-quality-ok: TQ008 configure mode without mocking authorization
):
result: Final = await mgmt_endpoints.health_check_servers(
server_ids=list(requested) if requested is not None else None,
user_api_key_dict=caller,
)
assert {row["server_id"] for row in result} == set(expected)
assert {server_id for server_id, route in routes.items() if route.called} == set(expected)
expected_status: Final = {200: "healthy", 503: "unhealthy"}[upstream_status]
assert all(row["status"] == expected_status for row in result)
class TestMCPRegistryEndpoint:
def test_registry_returns_404_when_flag_missing(self):
client = create_mcp_router_test_client()