fix: surface platform mcp in key grants

This commit is contained in:
Krrish Dholakia 2026-06-22 21:32:59 -07:00
parent 6ad764423f
commit 9a360efcca
6 changed files with 221 additions and 7 deletions

View file

@ -4532,13 +4532,20 @@ class MCPServerManager:
async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:
"""Return all MCP servers from registry without applying access controls."""
from litellm.proxy._experimental.mcp_server.platform_mcp import (
get_platform_mcp_enabled,
is_platform_mcp_server,
)
registry = self.get_registry()
if not registry:
return []
platform_mcp_enabled = await get_platform_mcp_enabled()
servers: List[LiteLLM_MCPServerTable] = []
for server in registry.values():
if is_platform_mcp_server(server) and not platform_mcp_enabled:
continue
servers.append(self._build_mcp_server_table(server))
return servers
@ -4546,11 +4553,22 @@ class MCPServerManager:
self, server_ids: Optional[List[str]] = None
) -> List[LiteLLM_MCPServerTable]:
"""Return health info for all servers in registry regardless of user access."""
from litellm.proxy._experimental.mcp_server.platform_mcp import (
get_platform_mcp_enabled,
is_platform_mcp_server,
)
registry = self.get_registry()
if not registry:
return []
if not await get_platform_mcp_enabled():
registry = {
server_id: server
for server_id, server in registry.items()
if not is_platform_mcp_server(server)
}
if server_ids:
target_server_ids = [sid for sid in server_ids if sid in registry]
else:

View file

@ -20,7 +20,7 @@ import secrets
import traceback
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, cast
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, cast
import fastapi
import yaml
@ -211,6 +211,29 @@ def _set_key_rotation_fields(
)
def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
return user_api_key_dict.user_role in {
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN.value,
}
async def _get_admin_assignable_mcp_server_ids(
user_api_key_dict: UserAPIKeyAuth,
) -> Set[str]:
if not _is_proxy_admin(user_api_key_dict):
return set()
from litellm.proxy._experimental.mcp_server.platform_mcp import (
PLATFORM_MCP_SERVER_ID,
get_platform_mcp_enabled,
)
if await get_platform_mcp_enabled():
return {PLATFORM_MCP_SERVER_ID}
return set()
def _is_allowed_to_make_key_request(
user_api_key_dict: UserAPIKeyAuth,
user_id: Optional[str],
@ -874,6 +897,9 @@ async def _common_key_generation_helper(
object_permission=data_json.get("object_permission"),
team_obj=team_table,
prisma_client=prisma_client,
additional_allowed_mcp_server_ids=await _get_admin_assignable_mcp_server_ids(
user_api_key_dict
),
)
if normalized_object_permission is not None:
data_json["object_permission"] = normalized_object_permission
@ -2164,6 +2190,7 @@ async def _validate_mcp_servers_for_key_update(
existing_key_row: Any,
prisma_client: Any,
user_api_key_cache: Any,
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[dict]:
"""Validate MCP servers in object_permission against the effective team."""
effective_team_obj = team_obj
@ -2186,6 +2213,9 @@ async def _validate_mcp_servers_for_key_update(
object_permission=object_permission_dict,
team_obj=effective_team_obj,
prisma_client=prisma_client,
additional_allowed_mcp_server_ids=await _get_admin_assignable_mcp_server_ids(
user_api_key_dict
),
)
await validate_key_search_tools_against_team(
object_permission=object_permission_dict,
@ -2422,6 +2452,7 @@ async def _validate_update_key_data(
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_api_key_dict=user_api_key_dict,
)
if normalized_object_permission is not None:
data.object_permission = LiteLLM_ObjectPermissionBase(

View file

@ -507,7 +507,10 @@ if MCP_AVAILABLE:
must see them, but a read-only admin gets the same redacted view as
any other non-managing caller.
"""
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
return user_api_key_dict.user_role in {
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN.value,
}
def _is_restricted_virtual_key_request(user_api_key_dict: UserAPIKeyAuth) -> bool:
"""Best-effort detection for route-restricted virtual keys.
@ -843,8 +846,47 @@ if MCP_AVAILABLE:
return "view_all"
return "restricted"
async def _get_admin_assignable_platform_mcp_server_ids(
user_api_key_dict: UserAPIKeyAuth,
) -> Set[str]:
if not _user_is_full_admin(user_api_key_dict):
return set()
from litellm.proxy._experimental.mcp_server.platform_mcp import (
PLATFORM_MCP_SERVER_ID,
get_platform_mcp_enabled,
)
if await get_platform_mcp_enabled():
return {PLATFORM_MCP_SERVER_ID}
return set()
async def _append_admin_assignable_platform_mcp_server(
servers: List[LiteLLM_MCPServerTable],
user_api_key_dict: UserAPIKeyAuth,
) -> List[LiteLLM_MCPServerTable]:
platform_server_ids = await _get_admin_assignable_platform_mcp_server_ids(
user_api_key_dict
)
if not platform_server_ids:
return servers
if any(server.server_id in platform_server_ids for server in servers):
return servers
platform_server = global_mcp_server_manager.get_mcp_server_by_id(
next(iter(platform_server_ids))
)
if platform_server is None:
return servers
return [
*servers,
global_mcp_server_manager._build_mcp_server_table(platform_server),
]
async def _get_team_scoped_mcp_server_list(
team_id: str,
additional_server_ids: Optional[Set[str]] = None,
) -> List[LiteLLM_MCPServerTable]:
"""
Return MCP servers scoped to a team: team's allowed servers + allow_all_keys servers.
@ -866,7 +908,20 @@ if MCP_AVAILABLE:
team_server_ids = await _get_team_allowed_mcp_servers(team_obj)
allow_all_server_ids = _get_allow_all_keys_server_ids()
all_allowed_ids = team_server_ids | allow_all_server_ids
all_allowed_ids = (
team_server_ids | allow_all_server_ids | (additional_server_ids or set())
)
from litellm.proxy._experimental.mcp_server.platform_mcp import (
get_platform_mcp_enabled,
is_platform_mcp_server_identifier,
)
if not await get_platform_mcp_enabled():
all_allowed_ids = {
server_id
for server_id in all_allowed_ids
if not is_platform_mcp_server_identifier(server_id)
}
if not all_allowed_ids:
return []
@ -977,11 +1032,18 @@ if MCP_AVAILABLE:
)
redacted_mcp_servers = await _get_team_scoped_mcp_server_list(
sanitized_team_id
sanitized_team_id,
additional_server_ids=await _get_admin_assignable_platform_mcp_server_ids(
user_api_key_dict
),
)
else:
servers = await _resolve_accessible_mcp_servers(user_api_key_dict)
redacted_mcp_servers = _redact_mcp_credentials_list(servers)
redacted_mcp_servers = _redact_mcp_credentials_list(
await _append_admin_assignable_platform_mcp_server(
servers, user_api_key_dict
)
)
# augment the mcp servers with public status
if litellm.public_mcp_servers is not None:

View file

@ -464,6 +464,7 @@ async def validate_key_mcp_servers_against_team(
object_permission: Optional[dict],
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
prisma_client: Optional[PrismaClient] = None,
additional_allowed_mcp_server_ids: Optional[Set[str]] = None,
) -> Optional[dict]:
"""
Validate that MCP servers requested on a key are within the allowed scope.
@ -492,8 +493,11 @@ async def validate_key_mcp_servers_against_team(
prisma_client=prisma_client,
)
# Combined allowed set = team servers + allow_all_keys servers
all_allowed_servers = team_allowed_servers | allow_all_keys_servers
all_allowed_servers = (
team_allowed_servers
| allow_all_keys_servers
| (additional_allowed_mcp_server_ids or set())
)
# Validate requested server IDs
if requested_servers:

View file

@ -301,6 +301,45 @@ class TestListMCPServers:
assert len(result) == 2
assert {server.server_id for server in result} == {"server-1", "server-2"}
@pytest.mark.asyncio
async def test_list_mcp_servers_admin_includes_enabled_platform_mcp(self):
admin = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
server = generate_mock_mcp_server_db_record(server_id="server-1", alias="One")
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_user_mcp_management_mode",
return_value="restricted",
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts",
AsyncMock(return_value=[admin]),
),
patch.object(
mgmt_endpoints.global_mcp_server_manager,
"get_all_allowed_mcp_servers",
AsyncMock(return_value=[server]),
),
patch(
"litellm.proxy._experimental.mcp_server.platform_mcp.get_platform_mcp_enabled",
AsyncMock(return_value=True),
),
):
result = await mgmt_endpoints.fetch_all_mcp_servers(
user_api_key_dict=admin
)
ids = {server.server_id for server in result}
assert ids == {"server-1", "platform_mcp"}
platform_server = next(
server for server in result if server.server_id == "platform_mcp"
)
assert platform_server.server_name == "platform_mcp"
assert platform_server.mcp_info is not None
assert platform_server.mcp_info["is_platform_mcp"] is True
@pytest.mark.asyncio
async def test_list_mcp_servers_view_all_mode_virtual_key_is_sanitized(self):
"""Issue #20325: virtual keys should get a safe discovery view."""
@ -1327,6 +1366,40 @@ class TestTeamScopedMCPServerAccess:
)
assert len(result) == 1
@pytest.mark.asyncio
async def test_admin_team_query_adds_enabled_platform_mcp_assignment_option(self):
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin_user",
)
scoped_list = AsyncMock(return_value=[])
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=True,
),
patch(
"litellm.proxy._experimental.mcp_server.platform_mcp.get_platform_mcp_enabled",
AsyncMock(return_value=True),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_scoped_mcp_server_list",
scoped_list,
),
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
await fetch_all_mcp_servers(
user_api_key_dict=mock_user_auth, team_id="any-team-id"
)
scoped_list.assert_awaited_once_with(
"any-team-id", additional_server_ids={"platform_mcp"}
)
@pytest.mark.asyncio
async def test_restricted_virtual_key_cannot_use_team_id_filter(self):
"""Restricted virtual keys must not bypass access limits via team_id."""

View file

@ -276,6 +276,32 @@ async def test_validate_allow_all_keys_servers_always_allowed(
)
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
new=_make_mock_mcp_manager("platform_mcp"),
)
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_additional_allowed_server_without_team(
mock_access_groups, mock_allow_all
):
object_permission = {"mcp_servers": ["platform_mcp"]}
await validate_key_mcp_servers_against_team(
object_permission=object_permission,
team_obj=None,
additional_allowed_mcp_server_ids={"platform_mcp"},
)
assert object_permission["mcp_servers"] == ["platform_mcp"]
@pytest.mark.asyncio
@patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",