fix(mcp): retry a scoped name as an access group when the server it names is ungranted

_scoped_server returned a denial whenever the registry placed a scoped name on a server the caller
does not hold, so a key granted only access group docs lost the group's servers when an ungranted
server was also named docs. It now returns None there, as it does for a name hidden from the caller's
IP, and the caller retries the name as an access group it holds, which is what the merge base did.

Regression tests pin the moved oauth_delegate and gateway oauth2 route shapes (alias case, server id,
aggregate x-mcp-servers, authorization-server document, authorize relay) to the exact-name route, and
the unselected aggregate connect to no challenge

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-02 21:52:54 +00:00
parent 71b946cd99
commit ae7b0fb8ea
4 changed files with 206 additions and 37 deletions

View file

@ -6,7 +6,7 @@ import types
import uuid
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, Final, Literal, NoReturn, TypeAlias, overload
from typing import Any, Final, NoReturn, TypeAlias, overload
from fastapi import HTTPException
from mcp import ReadResourceResult, Resource
@ -453,9 +453,7 @@ async def _get_allowed_mcp_servers_from_mcp_server_names(
if mcp_servers is not None:
for server_or_group in mcp_servers:
scoped = _scoped_server(server_or_group, allowed_mcp_servers, client_ip)
if isinstance(scoped, str):
verbose_logger.debug("MCP scope name %s names a server the caller does not hold", server_or_group)
elif scoped is not None:
if scoped is not None:
filtered_server[scoped.server_id] = scoped
else:
try:
@ -494,25 +492,17 @@ def _server_answers_to(server: MCPServer, name: str) -> bool:
return server_answers_to_name(server, name)
def _scoped_server(
name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None
) -> MCPServer | Literal["denied"] | None:
"""The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's
``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery
use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: the registry's
own pass order run over ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or
case variant the registry would pick. With no granted server answering: ``"denied"`` when the registry
places the name on a server the caller does not hold, so it is not retried as an access group; ``None``
when the registry cannot place the name for any caller."""
registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip)
if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None:
return "denied"
granted: Final = global_mcp_server_manager.get_mcp_server_answering_to(
name, client_ip=client_ip, among=allowed_mcp_servers
)
if granted is not None:
return granted
return "denied" if registry_pick is not None else None
def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None) -> MCPServer | None:
"""The granted server a scoped ``name`` selects for the caller: the registry's own pass order run over
``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or case variant the registry
would pick. ``None`` when no granted server answers, or when the registry's own pick for ``name`` is a server
hidden from ``client_ip``; the caller then retries the name as an access group it holds."""
if (
global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) is None
and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None
):
return None
return global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip, among=allowed_mcp_servers)
async def raise_denied_scoped_mcp_access(

View file

@ -1,3 +1,4 @@
import json
import uuid
from typing import Final
@ -79,6 +80,36 @@ def test_subject_grant_lists_only_reachable_tools_and_denies_the_rest(
assert not any(name.startswith(denied_alias) for name in denied_listed.tools), denied_listed.tools
@pytest.mark.parametrize("entry", ("mcp", "server_mcp"))
def test_access_group_named_like_an_ungranted_server_still_routes_the_groups_servers(
gateway: Gateway, entry: EntryPoint
) -> None:
with peer_of("http") as shadow, peer_of("http") as member, gateway.scenario() as scenario:
group: Final = "docs" + uuid.uuid4().hex[:8]
member_alias: Final = "mem" + uuid.uuid4().hex[:8]
register_mcp(scenario, shadow, group)
register_mcp(scenario, member, member_alias, mcp_access_groups=[group])
key: Final = scenario.key(object_permission={"mcp_access_groups": [group]})
unmatched: Final = "none" + uuid.uuid4().hex[:8]
denied: Final = McpCaller(
gateway, key, entry, unmatched, {"x-mcp-servers": unmatched} if entry == "mcp" else {}
).list_tools()
assert (denied.status, json.loads(denied.error or "null"), denied.tools) == (
(
200,
{"code": -32600, "message": f"The key is not allowed to access the requested MCP servers: {unmatched}"},
(),
)
if entry == "mcp"
else (404, {"detail": f"MCP server, toolset, or access group '{unmatched}' not found"}, ())
), denied.raw
selection: Final = {"x-mcp-servers": group} if entry == "mcp" else {}
listed: Final = McpCaller(gateway, key, entry, group, selection).list_tools()
assert listed.ok, listed.raw
assert set(listed.tools) == {f"{member_alias}-{tool}" for tool in ("add", "multiply", "fail")}, listed.tools
assert tool_calls(shadow.drain()) == ()
@pytest.mark.parametrize("entry", ("mcp", "server_mcp", "rest"))
def test_key_without_any_grant_sees_no_scoped_server(gateway: Gateway, entry: EntryPoint) -> None:
with peer_of("http") as peer, gateway.scenario() as scenario:

View file

@ -7,6 +7,7 @@ from base64 import urlsafe_b64encode
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
from urllib.parse import parse_qs, urlparse
import pytest
from fastapi import HTTPException
@ -392,6 +393,84 @@ def trust_xff():
yield
def _registered_gateway_oauth2_server():
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server: Final = MCPServer(
server_id="oid-7f3a",
name="gwx",
server_name="gwx",
alias="gwx",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
client_id="gw-client",
client_secret="gw-secret",
authorization_url="https://provider.com/oauth/authorize",
token_url="https://provider.com/oauth/token",
scopes=["read"],
)
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry[server.server_id] = server
return server
@pytest.mark.asyncio
@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"])
async def test_authorization_server_doc_for_a_moved_lookup_matches_the_exact_name_doc(lookup):
"""A case variant or server id now resolves like the exact name, so its AS metadata is the
exact-name doc with the requested spelling in the issuer and endpoint paths."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
oauth_authorization_server_mcp_standard,
)
_registered_gateway_oauth2_server()
request: Final = _mock_callback_request("http://litellm.example.com/")
exact: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name="gwx")
moved: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name=lookup)
assert exact["issuer"] == "http://litellm.example.com/mcp/gwx"
assert moved == {
key: value.replace("/gwx", f"/{lookup}") if isinstance(value, str) else value for key, value in exact.items()
}
@pytest.mark.asyncio
@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"])
async def test_authorize_relay_for_a_moved_lookup_redirects_like_the_exact_name(lookup):
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize
_registered_gateway_oauth2_server()
request: Final = _mock_callback_request("http://litellm.example.com/")
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", return_value="sealed"
):
exact: Final = await authorize(
request=request, mcp_server_name="gwx", redirect_uri="http://127.0.0.1:60108/callback", state="s1"
)
moved: Final = await authorize(
request=request, mcp_server_name=lookup, redirect_uri="http://127.0.0.1:60108/callback", state="s1"
)
exact_target: Final = urlparse(exact.headers["location"])
moved_target: Final = urlparse(moved.headers["location"])
exact_query: Final = parse_qs(exact_target.query)
moved_query: Final = parse_qs(moved_target.query)
assert exact.status_code == 307
assert exact_target._replace(query="") == urlparse("https://provider.com/oauth/authorize")
assert exact_query["client_id"] == ["gw-client"]
assert len(exact_query.pop("state")) == 1 and len(moved_query.pop("state")) == 1
assert (moved.status_code, moved_target._replace(query=""), moved_query) == (
exact.status_code,
exact_target._replace(query=""),
exact_query,
)
@pytest.mark.asyncio
async def test_authorize_endpoint_includes_response_type():
"""Test that authorize endpoint includes response_type=code parameter (fixes #15684)"""

View file

@ -8723,33 +8723,36 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a
@pytest.mark.asyncio
async def test_scoped_name_of_an_ungranted_server_is_not_retried_as_an_access_group():
async def test_scoped_name_of_an_ungranted_server_is_retried_as_an_access_group_the_key_holds():
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.server import _get_allowed_mcp_servers_from_mcp_server_names
private = MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http)
shadow = MCPServer(server_id="s-id", name="docs", server_name="docs", transport=MCPTransport.http)
member = MCPServer(server_id="m-id", name="m", server_name="m", transport=MCPTransport.http)
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry.update({"p-id": private, "m-id": member})
global_mcp_server_manager.registry.update({"s-id": shadow, "m-id": member})
group_members = {"docs": ["m-id"]}
try:
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=["m-id"],
side_effect=lambda names: [sid for name in names for sid in group_members.get(name, [])],
) as groups:
denied = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["shared"], allowed_mcp_servers=[member]
collided = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["docs"], allowed_mcp_servers=[member]
)
unknown = await _get_allowed_mcp_servers_from_mcp_server_names(
unmatched = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["team"], allowed_mcp_servers=[member]
)
finally:
global_mcp_server_manager.registry.clear()
assert denied == [], "a denied server name must not widen to an access group of the same name"
assert [s.server_id for s in unknown] == ["m-id"]
assert groups.await_args_list == [call(["team"])]
assert [s.server_id for s in collided] == ["m-id"], (
"a name owned by an ungranted server must still resolve to the access group of that name the key holds"
)
assert unmatched == [], "a name matching neither a granted server nor a granted access group stays denied"
assert groups.await_args_list == [call(["docs"]), call(["team"])]
@pytest.mark.asyncio
@ -8801,7 +8804,6 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the
pytest.param(("a-id", "d-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-alias-holder-listed-first"),
pytest.param(("d-id", "a-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-exact-name-listed-first"),
pytest.param(("g1", "g2"), "GITHUB", ("g2",), ["g2"], id="case-variant-collision"),
pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"),
],
)
async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected):
@ -8815,8 +8817,6 @@ async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(re
"d-id": MCPServer(server_id="d-id", name="docs", server_name="docs", transport=MCPTransport.http),
"g1": MCPServer(server_id="g1", name="GitHub", server_name="GitHub", transport=MCPTransport.http),
"g2": MCPServer(server_id="g2", name="github", server_name="github", transport=MCPTransport.http),
"p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http),
"m-id": MCPServer(server_id="m-id", name="m", server_name="m", transport=MCPTransport.http),
}
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry.update({server_id: servers[server_id] for server_id in registry})
@ -8825,7 +8825,7 @@ async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(re
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp."
"MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["m-id"],
return_value=[],
) as groups:
selected: Final = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=[scope], allowed_mcp_servers=[servers[server_id] for server_id in granted]
@ -10401,6 +10401,75 @@ class TestPreemptive401ModeAware:
client_ip=None,
)
async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException:
from litellm.proxy._experimental.mcp_server import server as server_module
manager = mcp_operations.global_mcp_server_manager
manager.registry.clear()
manager.registry[server.server_id] = server
with (
patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it
manager, "get_allowed_mcp_servers", AsyncMock(return_value=[server.server_id])
),
patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False),
pytest.raises(HTTPException) as exc,
):
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": path, "headers": [(b"host", b"testserver")]},
mcp_servers=[requested],
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None,
)
return exc.value
@pytest.mark.asyncio
@pytest.mark.parametrize("delegate", [True, False], ids=["oauth_delegate", "gateway_interactive"])
@pytest.mark.parametrize("shape", ["alias_case", "server_id", "x_mcp_servers"])
async def test_moved_connect_shapes_get_the_exact_name_routes_challenge(self, delegate, shape, monkeypatch):
"""Alias-case, server-id and x-mcp-servers connects resolve the same server the router serves, so
they answer the exact-name route's 401 with the requested spelling in the route segment."""
monkeypatch.delenv("SERVER_ROOT_PATH", raising=False)
server = _make_oauth2_server("gwx", oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate)
requested, path, exact_path = {
"alias_case": ("GWX", "/mcp/GWX", "/mcp/gwx"),
"server_id": (server.server_id, f"/mcp/{server.server_id}", "/mcp/gwx"),
"x_mcp_servers": ("GWX", "/mcp", "/mcp"),
}[shape]
exact = await self._connect_with_a_grant(server, "gwx", exact_path)
moved = await self._connect_with_a_grant(server, requested, path)
assert exact.status_code == 401
assert (moved.status_code, moved.detail) == (exact.status_code, exact.detail)
exact_header = {k.lower(): v for k, v in (exact.headers or {}).items()}["www-authenticate"]
moved_header = {k.lower(): v for k, v in (moved.headers or {}).items()}["www-authenticate"]
assert "/gwx" in exact_header
assert moved_header == exact_header.replace("/gwx", f"/{requested}")
@pytest.mark.asyncio
async def test_aggregate_connect_without_a_server_selection_is_not_challenged(self):
from litellm.proxy._experimental.mcp_server import server as server_module
manager = mcp_operations.global_mcp_server_manager
manager.registry.clear()
for alias, delegate in (("gwx", False), ("relay", True)):
server = _make_oauth2_server(alias, oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate)
manager.registry[server.server_id] = server
with patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False) as tokens:
await server_module._raise_preemptive_401_for_unauthenticated_servers(
scope={"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"testserver")]},
mcp_servers=None,
oauth2_headers=None,
mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None,
)
assert tokens.await_count == 0, "an unselected aggregate connect must not probe any server for a token"
@pytest.mark.asyncio
async def test_deferred_discovery_runs_before_delegate_challenge(self):
from litellm.proxy._experimental.mcp_server import server as server_module