Merge pull request #41609 from BerriAI/litellm_fix_mcp_health_permissions_4504

fix(mcp): restrict health discovery to virtual key grants
This commit is contained in:
joshua-berri 2026-09-17 19:03:50 +00:00 committed by GitHub
commit f075417643
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 161 additions and 24 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,26 +201,32 @@ class McpClient:
)
).root
def await_registered(self, server_id: str) -> None:
"""Poll /v1/mcp/server until `server_id` is listed. Fails at poll_timeout.
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,
)
The DB row exists the moment registration returns, but a data-plane pod
answers the listing from a registry it refreshes on a periodic DB sync, so a
pod that joined the load balancer after the write reports the server as
absent until its first sync.
"""
deadline = time.monotonic() + self.proxy.poll_timeout
while True:
registered = frozenset(row.server_id for row in self.registered_servers())
if server_id in registered:
return
if time.monotonic() >= deadline:
raise AssertionError(
f"registered server {server_id} still absent from /v1/mcp/server "
f"{self.proxy.poll_timeout}s after registration (the data plane never synced "
f"the row): {registered}"
)
time.sleep(self.proxy.poll_interval)
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) -> McpServerRow:
"""Wait for every configured replica to list the server and return its row."""
registered = self.proxy.read_body_back_everywhere(
"/v1/mcp/server",
McpServerListResponse,
settled=lambda response: any(row.server_id == server_id for row in response.root),
)
return next(
row for response in registered.values() for row in response.root if row.server_id == server_id
)
def generate_key(
self,

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
@ -42,8 +44,8 @@ class TestMcpKeyGrantByAlias:
grants access on every region. The same key must still see the server's
tools, proving the alias grant is honored at request time."""
server_id = register_datadog_mcp(client, resources)
client.await_registered(server_id)
alias = next(row.alias for row in client.registered_servers() if row.server_id == server_id)
registered = client.await_registered(server_id)
alias = registered.alias
assert alias, f"registered server {server_id} has no alias to grant by"
key = _key(client, resources, mcp_servers=[alias])
@ -108,3 +110,42 @@ 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)
owned: Final = {server_x, 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}.intersection(owned) == 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}.intersection(owned) == expected, (
f"health disclosed servers outside grants {grants}, requested {requested}: {health}"
)
assert all(row.status == "healthy" for row in health if row.server_id in owned), (
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,81 @@ 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( # test-quality-ok: TQ008 inject real registry into legacy route binding
mgmt_endpoints, "global_mcp_server_manager", manager,
),
patch.object( # test-quality-ok: TQ008 inject shared registry without mocking permission policy
mcp_server_manager, "global_mcp_server_manager", manager,
),
patch( # test-quality-ok: TQ008 configure mode without mocking authorization
"litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": mode},
),
):
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()