mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
fix(mcp): scope health discovery for route-restricted keys
This commit is contained in:
parent
4b368bf066
commit
e21db01d67
4 changed files with 139 additions and 3 deletions
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue