feat(mcp): support x-mcp-exclude-servers for subtractive MCP server scoping

Closes #34910
This commit is contained in:
Devin AI 2026-07-28 11:42:00 +00:00
parent daf22ec871
commit 8c47a5c1ce
6 changed files with 284 additions and 2 deletions

View file

@ -1,3 +1,4 @@
from collections.abc import Sequence
from typing import Dict, List, Optional
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
@ -16,6 +17,7 @@ class MCPAuthenticatedUser(AuthenticatedUser):
4. Server-specific authentication headers
5. OAuth2 headers
6. Raw headers - allows forwarding specific headers to the MCP server, specified by the admin.
7. MCP servers the caller asked to drop from the session (subtractive scope)
"""
def __init__(
@ -28,6 +30,7 @@ class MCPAuthenticatedUser(AuthenticatedUser):
mcp_protocol_version: Optional[str] = None,
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
mcp_excluded_servers: "Sequence[str] | None" = None,
):
self.user_api_key_auth = user_api_key_auth
self.mcp_auth_header = mcp_auth_header
@ -37,3 +40,4 @@ class MCPAuthenticatedUser(AuthenticatedUser):
self.oauth2_headers = oauth2_headers
self.raw_headers = raw_headers
self.client_ip = client_ip
self.mcp_excluded_servers = tuple(mcp_excluded_servers or ())

View file

@ -256,6 +256,19 @@ class MCPRequestHandler:
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
LITELLM_MCP_EXCLUDE_SERVERS_HEADER_NAME = SpecialHeaders.mcp_exclude_servers.value
@staticmethod
def get_mcp_excluded_servers_from_scope(scope: Scope) -> "tuple[str, ...]":
"""Server names, aliases, or access groups the caller asked to drop from the session
via ``x-mcp-exclude-servers``. Subtractive only: it can never widen scope beyond what
the caller is already allowed, so an unknown name is simply a no-op."""
headers = MCPRequestHandler._safe_get_headers_from_scope(scope)
raw_value = headers.get(MCPRequestHandler.LITELLM_MCP_EXCLUDE_SERVERS_HEADER_NAME)
if raw_value is None:
return ()
return tuple(name.strip() for name in raw_value.split(",") if name.strip())
@staticmethod
async def process_mcp_request(
scope: Scope,

View file

@ -13,7 +13,7 @@ import time
import traceback
import types
import uuid
from collections.abc import AsyncIterator, Callable, Mapping
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
from datetime import datetime
from typing import (
TYPE_CHECKING,
@ -1414,6 +1414,54 @@ if MCP_AVAILABLE:
return allowed_mcp_servers
async def _resolve_scope_name_to_server_ids(
server_or_group: str,
allowed_mcp_servers: Sequence[MCPServer],
) -> frozenset[str]:
"""Resolve one caller-supplied scope name (server name, alias, or access group)
to the ids it selects out of ``allowed_mcp_servers``."""
direct = frozenset(
server.server_id
for server in allowed_mcp_servers
if server
and server_or_group.lower() in {prefix.lower() for prefix in iter_known_server_prefixes(server) if prefix}
)
if direct:
return direct
try:
access_group_server_ids = frozenset(
await MCPRequestHandler._get_mcp_servers_from_access_groups([server_or_group])
)
except Exception as e: # noqa: BLE001 # access-group lookup is best-effort; an unresolvable name is a no-op
verbose_logger.debug(f"Could not resolve '{server_or_group}' as access group: {e}")
return frozenset()
return frozenset(
server.server_id for server in allowed_mcp_servers if server.server_id in access_group_server_ids
)
async def _exclude_mcp_servers_by_names(
excluded_servers: Sequence[str] | None,
allowed_mcp_servers: Sequence[MCPServer],
) -> tuple[MCPServer, ...]:
"""Drop the servers named by ``x-mcp-exclude-servers`` from an already authorized set.
Purely subtractive: names that do not resolve are ignored rather than failing closed,
and an empty result is a legitimate "everything was excluded" outcome.
"""
if not excluded_servers:
return tuple(allowed_mcp_servers)
excluded_ids = frozenset[str]().union(
*[
await _resolve_scope_name_to_server_ids(server_or_group, allowed_mcp_servers)
for server_or_group in excluded_servers
]
)
if not excluded_ids:
return tuple(allowed_mcp_servers)
verbose_logger.debug("MCP exclude filter: dropping server ids %s", sorted(excluded_ids))
return tuple(server for server in allowed_mcp_servers if server.server_id not in excluded_ids)
def _tool_name_matches(tool_name: str, filter_list: list[str]) -> bool:
"""
Check if a tool name matches any name in the filter list.
@ -1504,6 +1552,11 @@ if MCP_AVAILABLE:
tool.description = description_map[lookup_key]
return tools
def _get_excluded_servers_from_context() -> Sequence[str]:
"""Excluded server names from the request's auth context; empty when unset."""
auth_user = get_active_auth_context()
return auth_user.mcp_excluded_servers if auth_user else []
def _get_client_ip_from_context() -> str | None:
"""
Extract client_ip from auth context.
@ -1576,7 +1629,12 @@ if MCP_AVAILABLE:
allowed_mcp_servers=allowed_mcp_servers,
)
return allowed_mcp_servers
return list(
await _exclude_mcp_servers_by_names(
excluded_servers=_get_excluded_servers_from_context(),
allowed_mcp_servers=allowed_mcp_servers,
)
)
def _client_has_per_server_auth_header(
server: MCPServer,
@ -3020,6 +3078,12 @@ if MCP_AVAILABLE:
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
)
allowed_mcp_servers = list(
await _exclude_mcp_servers_by_names(
excluded_servers=_get_excluded_servers_from_context(),
allowed_mcp_servers=allowed_mcp_servers,
)
)
if not allowed_mcp_servers:
raise HTTPException(
status_code=403,
@ -4278,6 +4342,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
mcp_excluded_servers=MCPRequestHandler.get_mcp_excluded_servers_from_scope(scope),
session_id=session_id if use_stateful else None,
touch_last_seen=(scope.get("method") or "").upper() != "DELETE",
copy_existing_session_auth_context=is_initialize,
@ -4419,6 +4484,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_sse_client_ip,
mcp_excluded_servers=MCPRequestHandler.get_mcp_excluded_servers_from_scope(scope),
)
if not _SESSION_MANAGERS_INITIALIZED:
@ -4505,6 +4571,7 @@ if MCP_AVAILABLE:
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
mcp_excluded_servers: Sequence[str] | None = None,
) -> None:
auth_user.user_api_key_auth = user_api_key_auth
auth_user.mcp_auth_header = mcp_auth_header
@ -4513,6 +4580,7 @@ if MCP_AVAILABLE:
auth_user.oauth2_headers = oauth2_headers
auth_user.raw_headers = raw_headers
auth_user.client_ip = client_ip
auth_user.mcp_excluded_servers = mcp_excluded_servers or []
def set_auth_context(
user_api_key_auth: UserAPIKeyAuth | None,
@ -4522,6 +4590,7 @@ if MCP_AVAILABLE:
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
mcp_excluded_servers: Sequence[str] | None = None,
) -> MCPAuthenticatedUser:
"""
Set the UserAPIKeyAuth in the auth context variable.
@ -4541,6 +4610,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
mcp_excluded_servers=mcp_excluded_servers,
)
auth_context_var.set(auth_user)
return auth_user
@ -4553,6 +4623,7 @@ if MCP_AVAILABLE:
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
mcp_excluded_servers: Sequence[str] | None = None,
session_id: str | None = None,
touch_last_seen: bool = True,
copy_existing_session_auth_context: bool = False,
@ -4570,6 +4641,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
mcp_excluded_servers=mcp_excluded_servers,
)
_update_auth_context(
auth_user=auth_user,
@ -4580,6 +4652,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
mcp_excluded_servers=mcp_excluded_servers,
)
auth_context_var.set(auth_user)
return auth_user
@ -4591,6 +4664,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
mcp_excluded_servers=mcp_excluded_servers,
)
def _wrap_send_with_stateful_session_auth_context(

View file

@ -3944,6 +3944,7 @@ class SpecialHeaders(enum.Enum):
mcp_auth = "x-mcp-auth"
mcp_servers = "x-mcp-servers"
mcp_access_groups = "x-mcp-access-groups"
mcp_exclude_servers = "x-mcp-exclude-servers"
@classmethod
def litellm_credential_header_names(cls) -> "frozenset[str]":

View file

@ -7645,3 +7645,21 @@ class TestSessionBearerEgressScrub:
assert oauth2 is None
assert "authorization" not in {k.lower() for k in raw}
assert per_server == {"github": {"Authorization": "Bearer gh_injected_upstream"}}
class TestGetMcpExcludedServersFromScope:
"""``x-mcp-exclude-servers`` parsing: a comma-separated, subtractive scope list."""
@pytest.mark.parametrize(
"headers,expected",
[
([], ()),
([(b"x-mcp-exclude-servers", b"")], ()),
([(b"x-mcp-exclude-servers", b"ui_server")], ("ui_server",)),
([(b"x-mcp-exclude-servers", b" server_a , server_b ,")], ("server_a", "server_b")),
([(b"X-MCP-Exclude-Servers", b"server_a")], ("server_a",)),
],
)
def test_parses_excluded_servers(self, headers, expected):
scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": headers}
assert MCPRequestHandler.get_mcp_excluded_servers_from_scope(scope) == expected

View file

@ -6836,6 +6836,178 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_empty_list_fails_cl
assert result == []
class TestExcludeMcpServers:
"""``x-mcp-exclude-servers`` drops servers from the aggregated session while keeping
everything else the caller is otherwise allowed to see.
"""
@pytest.mark.asyncio
async def test_excludes_named_alias_and_keeps_the_rest(self):
from litellm.proxy._experimental.mcp_server.server import (
_exclude_mcp_servers_by_names,
)
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
_make_mcp_server_for_scope_filter("id-c", "gamma"),
]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
):
result = await _exclude_mcp_servers_by_names(
excluded_servers=["BETA"],
allowed_mcp_servers=allowed,
)
assert [s.server_id for s in result] == ["id-a", "id-c"]
@pytest.mark.asyncio
async def test_unknown_name_is_a_no_op(self):
from litellm.proxy._experimental.mcp_server.server import (
_exclude_mcp_servers_by_names,
)
allowed = [_make_mcp_server_for_scope_filter("id-a", "alpha")]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
):
result = await _exclude_mcp_servers_by_names(
excluded_servers=["does-not-exist"],
allowed_mcp_servers=allowed,
)
assert [s.server_id for s in result] == ["id-a"]
@pytest.mark.asyncio
async def test_excludes_every_server_in_an_access_group(self):
from litellm.proxy._experimental.mcp_server.server import (
_exclude_mcp_servers_by_names,
)
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["id-b"],
):
result = await _exclude_mcp_servers_by_names(
excluded_servers=["group-name"],
allowed_mcp_servers=allowed,
)
assert [s.server_id for s in result] == ["id-a"]
@pytest.mark.asyncio
async def test_excluding_everything_yields_an_empty_scope(self):
from litellm.proxy._experimental.mcp_server.server import (
_exclude_mcp_servers_by_names,
)
allowed = [_make_mcp_server_for_scope_filter("id-a", "alpha")]
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
):
result = await _exclude_mcp_servers_by_names(
excluded_servers=["alpha"],
allowed_mcp_servers=allowed,
)
assert result == ()
@pytest.mark.asyncio
async def test_no_exclusions_returns_full_scope(self):
from litellm.proxy._experimental.mcp_server.server import (
_exclude_mcp_servers_by_names,
)
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
assert await _exclude_mcp_servers_by_names(excluded_servers=[], allowed_mcp_servers=allowed) == tuple(allowed)
assert await _exclude_mcp_servers_by_names(excluded_servers=None, allowed_mcp_servers=allowed) == tuple(allowed)
@pytest.mark.asyncio
async def test_request_scope_drops_excluded_server_without_an_allowlist(self):
"""End-to-end through the scope resolver: the caller sends only
``x-mcp-exclude-servers``, so every other allowed server survives.
"""
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
MCPAuthenticatedUser,
)
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers,
auth_context_var,
)
allowed = [
_make_mcp_server_for_scope_filter("id-a", "alpha"),
_make_mcp_server_for_scope_filter("id-b", "beta"),
]
by_id = {server.server_id: server for server in allowed}
user_api_key_auth = UserAPIKeyAuth(api_key="sk-test", user_id="u1")
token = auth_context_var.set(
MCPAuthenticatedUser(
user_api_key_auth=user_api_key_auth,
mcp_excluded_servers=["beta"],
)
)
try:
with (
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager."
"get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=["id-a", "id-b"],
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager."
"filter_server_ids_by_ip_with_info",
side_effect=lambda server_ids, client_ip: (server_ids, 0),
),
patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager."
"get_mcp_server_by_id",
side_effect=by_id.get,
),
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
result = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=None,
client_ip=None,
)
finally:
auth_context_var.reset(token)
assert [s.server_id for s in result] == ["id-a"]
class TestProxyExceptionToHttpException:
"""Auth failures reach the MCP ASGI handlers as ProxyException, not
HTTPException. The handlers must map them back to their real status and