diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 65f772d1664..84099995f49 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2d49297c8e9..8a1dd5348a8 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index e86982307e7..611c88a6dd9 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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: diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index f2ddae40d8c..dd7ac41a6e5 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 0b5b5fb6ceb..a4beb841d78 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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.""" diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 965580e8758..0d6f0eeed00 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -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",