fix(mcp): resolve /mcp/{name} routes through one exact-first lookup for connect, discovery and the scoped router

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-30 10:47:23 +00:00
parent 23eb6bd8a5
commit f89b92763b
5 changed files with 80 additions and 10 deletions

View file

@ -7200,8 +7200,12 @@ class MCPServerManager:
return None
def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None:
"""The server a scoped ``/mcp/{name}`` connect resolves to: alias, then server_name, then name, each
case-insensitive so ``/mcp/GH`` and ``/mcp/gh`` agree, then the router's prefix match."""
"""The one server a ``/mcp/{name}`` segment denotes, shared by the connect preflight, the scoped
router, and RFC 9728 discovery so all three name the same server: the exact ``get_mcp_server_by_name``
priority first, then the same priority case-insensitively, then any prefix form routing accepts."""
exact: Final = self.get_mcp_server_by_name(name, client_ip=client_ip)
if exact is not None:
return exact
requested: Final = name.lower()
servers: Final = tuple(self.get_registry().values())
identifiers: Final[tuple[Callable[[MCPServer], str | None], ...]] = (

View file

@ -455,15 +455,10 @@ async def _get_allowed_mcp_servers_from_mcp_server_names(
# Filter servers based on mcp_servers parameter if provided
if mcp_servers is not None:
for server_or_group in mcp_servers:
server_name_matched = False
if (scoped := _scoped_server(server_or_group, allowed_mcp_servers)) is not None:
filtered_server[scoped.server_id] = scoped
for server in allowed_mcp_servers:
if server and _server_answers_to(server, server_or_group):
filtered_server[server.server_id] = server
server_name_matched = True
break
if not server_name_matched:
if scoped is None:
try:
access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups(
[server_or_group]
@ -500,6 +495,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]) -> MCPServer | None:
"""The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick when
the caller holds it, so the router agrees with the connect preflight and discovery, and none when the
registry names a server the caller does not hold. Names the registry cannot place fall back to the first
granted server answering to them."""
registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name)
if registry_pick is not None:
return next((s for s in allowed_mcp_servers if s.server_id == registry_pick.server_id), None)
return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None)
async def raise_denied_scoped_mcp_access(
requested_names: Sequence[str],
user_api_key_auth: UserAPIKeyAuth | None,

View file

@ -1279,6 +1279,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
mock_manager.get_mcp_server_by_id = lambda server_id: (
working_server if server_id == "working_server" else failing_server
)
mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
server_ids,
@ -6764,6 +6765,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
):
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"])
mock_manager.get_mcp_server_by_id = MagicMock(return_value=legacy_server)
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["legacy-m2m-id"], 0))
mock_manager._get_tools_from_server = AsyncMock(side_effect=capture_extra_headers)
@ -8588,6 +8590,39 @@ async def test_get_allowed_mcp_servers_from_mcp_server_names_known_alias_returns
assert [s.server_id for s in result] == ["id-a"]
@pytest.mark.asyncio
@pytest.mark.parametrize("alias_server_first", [True, False], ids=["alias-granted-first", "server-name-granted-first"])
async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(alias_server_first):
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
by_alias = MCPServer(server_id="a-id", name="a", server_name="a", alias="gh", transport=MCPTransport.http)
by_server_name = MCPServer(server_id="b-id", name="b", server_name="Gh", transport=MCPTransport.http)
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry.update({"a-id": by_alias, "b-id": by_server_name})
granted_both = [by_alias, by_server_name] if alias_server_first else [by_server_name, by_alias]
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=[],
):
for name in ("Gh", "gh", "GH"):
expected = global_mcp_server_manager.get_mcp_server_answering_to(name)
selected = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=[name], allowed_mcp_servers=granted_both
)
assert [s.server_id for s in selected] == [expected.server_id], name
only_b = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=["gh"], allowed_mcp_servers=[by_server_name]
)
finally:
global_mcp_server_manager.registry.clear()
assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another"
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown():
"""
@ -9729,6 +9764,7 @@ async def test_aggregate_listing_reports_per_server_outcomes():
mock_manager.get_mcp_server_by_id = lambda server_id: (
working_server if server_id == "working_server" else broken_server
)
mock_manager.get_mcp_server_answering_to = lambda name, client_ip=None: None
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
async def mock_get_tools_from_server(server, **kwargs):

View file

@ -6231,6 +6231,28 @@ class TestMCPServerManager:
assert manager.get_mcp_server_answering_to("GH_PUBLIC") is gh_public
assert manager.get_mcp_server_answering_to("gh-id") is gh
@pytest.mark.parametrize("gh_first", [True, False], ids=["alias-listed-first", "server-name-listed-first"])
def test_answering_to_agrees_with_exact_name_before_case_folding(self, gh_first):
manager = MCPServerManager()
by_alias = MCPServer(
server_id="a-id",
name="a",
server_name="a",
alias="gh",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
by_server_name = MCPServer(server_id="b-id", name="b", server_name="Gh", transport=MCPTransport.http)
manager.registry = (
{"a-id": by_alias, "b-id": by_server_name} if gh_first else {"b-id": by_server_name, "a-id": by_alias}
)
for name in ("Gh", "gh"):
assert manager.get_mcp_server_answering_to(name) is manager.get_mcp_server_by_name(name), name
assert manager.get_mcp_server_answering_to("Gh") is by_server_name
assert manager.get_mcp_server_answering_to("gh") is by_alias
assert manager.get_mcp_server_answering_to("GH") is by_alias
def test_remove_server_drops_only_its_own_tool_mapping_rows(self):
manager = self._manager_with_deepwiki_and_huggingface()

View file

@ -924,6 +924,7 @@ async def test_get_tools_from_mcp_servers():
mock_manager.get_mcp_server_by_id = lambda server_id: (
mock_server_1 if server_id == "server1_id" else mock_server_2
)
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(
@ -1002,6 +1003,7 @@ async def test_get_tools_from_mcp_servers():
if server_id == "server1_id"
else (mock_server_2 if server_id == "server2_id" else mock_server_3)
)
mock_manager.get_mcp_server_answering_to = MagicMock(return_value=None)
mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1])
# Mock filter_server_ids_by_ip_with_info to return input unchanged (no IP filtering in test)
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(