diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0611b153de5..707fa7625d7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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], ...]] = ( diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 50101975dca..9473f139def 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index c7dff4fa82e..03d4414ec7a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 1a4dd2cb93f..ccf9d7ce516 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index f8bf72428aa..cb4c690e18c 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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(