mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: surface platform mcp in key grants
This commit is contained in:
parent
6ad764423f
commit
9a360efcca
6 changed files with 221 additions and 7 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue