fix(mcp): scope allow-all servers to virtual keys (#39531)

This commit is contained in:
tin-berri 2026-09-03 10:32:03 -07:00 committed by GitHub
parent 27274f65e4
commit 7256bd307a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 214 additions and 32 deletions

View file

@ -3,7 +3,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, cast
from typing import TYPE_CHECKING, Final, Literal, cast
from fastapi import HTTPException
from starlette.datastructures import Headers
@ -305,6 +305,12 @@ def _admission_failure_fallback(
raise exc
@dataclass(frozen=True, slots=True)
class MCPServerAccess:
server_ids: tuple[str, ...]
scope: Literal["unscoped", "scoped", "unresolved"] = "unscoped"
@dataclass(frozen=True, slots=True)
class DcrBridgeTarget:
"""The single DCR-bridge server a request targets, paired with the exact name the caller
@ -1456,6 +1462,18 @@ class MCPRequestHandler:
*,
keyless_source: bool = False,
) -> list[str]:
access: Final = await MCPRequestHandler.get_mcp_server_access(
user_api_key_auth,
keyless_source=keyless_source,
)
return list(access.server_ids)
@staticmethod
async def get_mcp_server_access(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
) -> MCPServerAccess:
"""
Get list of allowed MCP servers for the given user/key based on permissions.
@ -1478,13 +1496,17 @@ class MCPRequestHandler:
"""
from litellm.proxy.proxy_server import general_settings
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
try:
# A keyless admitted subject resolves per source BEFORE any single-source rule here. Ordering
# matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above
# this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is
# independent; an opt-out silences only its own source, inside the recursive call).
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
return await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)
return MCPServerAccess(
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
)
# Get allowed servers from key and team
allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
@ -1492,7 +1514,7 @@ class MCPRequestHandler:
# The key explicitly opted out of every MCP server. This overrides
# team inheritance and additive grants (mirrors no-default-models).
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key:
return []
return MCPServerAccess(server_ids=(), scope="scoped")
allowed_mcp_servers_for_team = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
@ -1572,7 +1594,7 @@ class MCPRequestHandler:
"require_end_user_mcp_access_defined=True and end_user %s has no MCP permissions - blocking MCP access",
user_api_key_auth.end_user_id,
)
return []
return MCPServerAccess(server_ids=(), scope="scoped")
#########################################################
# Check agent permissions if agent_id is set on the key
@ -1601,14 +1623,22 @@ class MCPRequestHandler:
#########################################################
# Apply org-level ceiling if org_id is set
#########################################################
allowed_mcp_servers = await MCPRequestHandler._apply_primary_org_ceiling(
allowed_mcp_servers, org_restricts = await MCPRequestHandler._apply_primary_org_ceiling(
allowed_mcp_servers,
user_api_key_auth,
has_lower_level_mcp_restrictions,
keyless_source=keyless_source,
)
return list(set(allowed_mcp_servers))
declares_key_mcp_scope: Final = getattr(key_object_permission, "mcp_servers", None) is not None
return MCPServerAccess(
server_ids=tuple(set(allowed_mcp_servers)),
scope=(
"scoped"
if has_lower_level_mcp_restrictions or org_restricts or declares_key_mcp_scope
else "unscoped"
),
)
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
# A ceiling we KNOW exists and cannot read. Denying is the only answer that does not
@ -1616,7 +1646,10 @@ class MCPRequestHandler:
verbose_logger.warning("Denying MCP access, entitlement unreadable: %s", e)
else:
verbose_logger.warning("Failed to get allowed MCP servers: %s", e)
return []
return MCPServerAccess(
server_ids=(),
scope="scoped" if getattr(key_object_permission, "mcp_servers", None) is not None else "unresolved",
)
@staticmethod
async def _apply_primary_org_ceiling(
@ -1624,7 +1657,7 @@ class MCPRequestHandler:
user_api_key_auth: UserAPIKeyAuth | None,
has_lower_level_mcp_restrictions: bool,
keyless_source: bool = False,
) -> list[str]:
) -> tuple[list[str], bool]:
"""Cap the resolved server list by this caller's org ceiling: an explicit org list intersects
lower-level restrictions (else becomes the ceiling); no org or an empty list leaves it unchanged.
@ -1638,7 +1671,7 @@ class MCPRequestHandler:
cannot be read raises out of ``_get_allowed_mcp_servers_for_org`` and never arrives here as
``None``, so key auth cannot silently shed a ceiling an operator did configure."""
if not (user_api_key_auth and user_api_key_auth.org_id):
return allowed_mcp_servers
return allowed_mcp_servers, False
allowed_mcp_servers_for_org: Final = await MCPRequestHandler._get_allowed_mcp_servers_for_org(user_api_key_auth)
if allowed_mcp_servers_for_org is None:
verbose_logger.warning(
@ -1646,9 +1679,9 @@ class MCPRequestHandler:
user_api_key_auth.org_id,
"denying (keyless admitted subject)" if keyless_source else "leaving uncapped (key auth)",
)
return [] if keyless_source else allowed_mcp_servers
return ([] if keyless_source else allowed_mcp_servers), False
if len(allowed_mcp_servers_for_org) == 0:
return allowed_mcp_servers
return allowed_mcp_servers, False
if has_lower_level_mcp_restrictions or keyless_source:
# Org can only cap lower-level restrictions. A keyless admitted source ALWAYS takes this
# arm: its model unions GRANTS, so an org list may only narrow a source, never become one.
@ -1657,7 +1690,7 @@ class MCPRequestHandler:
# No lower-level restrictions → org list becomes the ceiling.
capped = allowed_mcp_servers_for_org
verbose_logger.debug("Applied org ceiling filter. Final allowed servers: %s", capped)
return capped
return capped, True
@staticmethod
def _scoped_source_auth(

View file

@ -55,6 +55,7 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
@ -2958,7 +2959,13 @@ class MCPServerManager:
return None
return user_api_key_auth.mcp_session_resource_server_id
async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]:
async def get_allowed_mcp_servers(
self,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
access: MCPServerAccess | None = None,
general_settings: Mapping[str, object] | None = None,
) -> list[str]:
"""
Get the allowed MCP Servers for the user.
@ -2967,6 +2974,9 @@ class MCPServerManager:
2. If admin and no object_permission, return all servers
3. Otherwise, use standard permission checks
"""
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
# A keyless admitted subject is resolved per grant source, and channel decisions that are
@ -3007,11 +3017,16 @@ class MCPServerManager:
# whole registry, for keys AND admitted session subjects alike (one predicate owns the
# question). Seeded into the union rather than returned early so the session resource
# scope below still bounds a per-server envelope held by an admin.
combined_servers: Final = (
set(self.get_registry().keys())
if await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
else set(await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth))
admin_unscoped: Final = await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
resolved_access: Final = (
MCPServerAccess(server_ids=())
if admin_unscoped
else access or await MCPRequestHandler.get_mcp_server_access(user_api_key_auth)
)
resolved_server_ids: Final = (
set(self.get_registry().keys()) if admin_unscoped else set(resolved_access.server_ids)
)
combined_servers: Final = set(resolved_server_ids)
verbose_logger.debug("Allowed MCP Servers for user api key auth: %s", combined_servers)
combined_servers.update(
await self.operator_open_server_ids(
@ -3052,6 +3067,18 @@ class MCPServerManager:
]
combined_servers.update(delegate_server_ids)
restrict_allow_all: Final = (
resolved_general_settings.get("mcp_allow_all_keys_respects_mcp_scope", False)
and user_api_key_auth is not None
and user_api_key_auth.via_virtual_key
and resolved_access.scope != "unscoped"
)
if restrict_allow_all:
combined_servers.difference_update(
set(allow_all_server_ids)
- resolved_server_ids
- (set(submitted_server_ids) if resolved_access.scope != "unresolved" else set())
)
if len(combined_servers) == 0:
verbose_logger.debug("No allowed MCP Servers found for user api key auth.")
scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth)

View file

@ -2815,6 +2815,7 @@ async def test_mcp_server_manager_with_access_groups_integration():
"""Integration test for MCPServerManager with access group filtering"""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -2848,11 +2849,11 @@ async def test_mcp_server_manager_with_access_groups_integration():
)
# Mock the permission lookup to return staff access group
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers") as mock_get_allowed:
mock_get_allowed.return_value = [
"staff-server-id",
"ops-server-id",
] # User has access to staff and ops
with patch.object(MCPRequestHandler, "get_mcp_server_access") as mock_get_allowed: # test-quality-ok: manager resolver seam
mock_get_allowed.return_value = MCPServerAccess(
server_ids=("staff-server-id", "ops-server-id"),
scope="scoped",
)
allowed_servers = await test_manager.get_allowed_mcp_servers(user_auth)
@ -2901,6 +2902,7 @@ async def test_get_allowed_mcp_servers_returns_empty_for_non_admin_without_permi
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
)
test_manager = MCPServerManager()
@ -2923,9 +2925,9 @@ async def test_get_allowed_mcp_servers_returns_empty_for_non_admin_without_permi
)
with patch.object(
MCPRequestHandler, "get_allowed_mcp_servers", new_callable=AsyncMock
MCPRequestHandler, "get_mcp_server_access", new_callable=AsyncMock
) as mock_permission_lookup:
mock_permission_lookup.return_value = []
mock_permission_lookup.return_value = MCPServerAccess(server_ids=())
allowed_servers = await test_manager.get_allowed_mcp_servers(user_auth)
assert allowed_servers == []

View file

@ -667,6 +667,44 @@ class TestMCPRequestHandler:
assert result == []
async def test_db_default_empty_key_scope_keeps_org_substitution(self):
"""A key whose object_permission row carries only the DB-default empty mcp_servers
list (e.g. a vector-stores-only key) places no lower-level MCP restriction: the org
list still substitutes with the flag off, while the access result stays scoped so
the opt-in allow-all ceiling can still bind"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", org_id="org-1")
key_object_permission = self._toolset_only_object_permission([])
key_object_permission.mcp_toolsets = None
mock_manager = self._mock_manager_with_toolsets({})
with (
patch.object( # test-quality-ok: stub the level's perm loader; the resolver reads module globals with no injection seam
MCPRequestHandler, "_get_key_object_permission", return_value=key_object_permission
),
patch.object( # test-quality-ok: team resolution has its own tests; pin it empty here
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here
MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=[])
),
patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here
MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])
),
patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch.object( # test-quality-ok: stub the level's perm loader; the resolver reads module globals with no injection seam
MCPRequestHandler,
"_get_allowed_mcp_servers_for_org",
AsyncMock(return_value=["server-x", "server-y"]),
),
):
access = await MCPRequestHandler.get_mcp_server_access(user_api_key_auth)
assert sorted(access.server_ids) == ["server-x", "server-y"]
assert access.scope == "scoped"
async def test_team_dangling_toolset_denies_key_own_grants(self):
"""A team toolset that cannot be resolved must deny on the SERVER axis too,
not silently drop the team ceiling and pass the key's own grants through"""

View file

@ -3367,6 +3367,7 @@ async def test_mcp_manager_returns_public_when_permission_lookup_fails():
@pytest.mark.asyncio
async def test_mcp_manager_merges_public_and_restricted_servers():
try:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
@ -3398,8 +3399,8 @@ async def test_mcp_manager_merges_public_and_restricted_servers():
return_value=False,
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
AsyncMock(return_value=["restricted"]),
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_mcp_server_access",
AsyncMock(return_value=MCPServerAccess(server_ids=("restricted",), scope="scoped")),
),
):
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())

View file

@ -5884,6 +5884,7 @@ class TestMCPServerManager:
"""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
@ -5903,19 +5904,22 @@ class TestMCPServerManager:
object_permission_id="perm_123",
)
# Mock MCPRequestHandler.get_allowed_mcp_servers to verify it receives user_api_key_auth
# Mock MCPRequestHandler.get_mcp_server_access to verify it receives user_api_key_auth
with patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
"get_mcp_server_access",
new_callable=AsyncMock,
) as mock_get_allowed:
# Configure mock to return servers from object_permission
mock_get_allowed.return_value = ["test_server_1", "test_server_2"]
mock_get_allowed.return_value = MCPServerAccess(
server_ids=("test_server_1", "test_server_2"),
scope="scoped",
)
# Call get_allowed_mcp_servers with user_api_key_auth
result = await manager.get_allowed_mcp_servers(user_api_key_auth)
# Verify MCPRequestHandler.get_allowed_mcp_servers was called with user_api_key_auth
# Verify MCPRequestHandler.get_mcp_server_access was called with user_api_key_auth
mock_get_allowed.assert_called_once()
call_args = mock_get_allowed.call_args
assert call_args[0][0] is user_api_key_auth # First positional arg should be user_api_key_auth
@ -6072,6 +6076,7 @@ class TestMCPServerManager:
from litellm.proxy import proxy_server as proxy_server_module
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
)
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
@ -6104,9 +6109,12 @@ class TestMCPServerManager:
patch.object(manager, "get_allow_all_keys_server_ids", return_value=["global-server"]),
patch.object(
MCPRequestHandler,
"get_allowed_mcp_servers",
"get_mcp_server_access",
new_callable=AsyncMock,
return_value=["toolset-server"],
return_value=MCPServerAccess(
server_ids=("toolset-server",),
scope="scoped",
),
),
):
result = await manager.get_allowed_mcp_servers(user_api_key_auth)
@ -6205,6 +6213,79 @@ class TestMCPServerManager:
assert set(result) == {"global-server", "submitted-server"}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"flag_enabled, via_virtual_key, resolved_server_ids, scope, submitted_server_ids, expected_server_ids",
[
(False, True, ("granted",), "scoped", (), {"granted", "public"}),
(True, True, ("granted",), "scoped", (), {"granted"}),
(True, True, ("granted", "public"), "scoped", (), {"granted", "public"}),
(True, True, (), "scoped", (), set()),
(True, True, (), "unscoped", (), {"public"}),
(
True,
True,
("team-granted",),
"scoped",
("submitted",),
{"team-granted", "submitted"},
),
(
True,
True,
("team-granted",),
"scoped",
("public",),
{"team-granted", "public"},
),
(True, False, ("granted",), "scoped", (), {"granted", "public"}),
],
ids=(
"flag_off_preserves_allow_all",
"flag_on_scoped_key_excludes_allow_all",
"flag_on_keeps_allow_all_when_granted",
"flag_on_restricted_empty_excludes_allow_all",
"flag_on_unscoped_key_preserves_allow_all",
"flag_on_preserves_submitted_byom",
"flag_on_preserves_submitted_byom_when_it_is_allow_all",
"flag_on_non_virtual_key_preserves_allow_all",
),
)
async def test_allow_all_keys_scope_flag(
self,
flag_enabled,
via_virtual_key,
resolved_server_ids,
scope,
submitted_server_ids,
expected_server_ids,
): # test-quality-ok: parameterized matrix covers the scope state machine
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
manager = MCPServerManager()
auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-123")
auth.via_virtual_key = via_virtual_key
access = MCPServerAccess(server_ids=resolved_server_ids, scope=scope)
with (
patch.object(manager, "get_allow_all_keys_server_ids", return_value=["public"]),
patch.object(
manager,
"_get_active_submitted_mcp_server_ids_for_user",
new=AsyncMock(return_value=list(submitted_server_ids)),
),
):
assert (
set(
await manager.get_allowed_mcp_servers(
auth,
access=access,
general_settings={"mcp_allow_all_keys_respects_mcp_scope": flag_enabled},
)
)
== expected_server_ids
)
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_anonymous_delegate_requires_oauth2(self):
"""Anonymous delegated auth listing should only include oauth2 servers."""