From 2172b8b60e7d9a7dd7baac99ac78818b82810ee5 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 11:26:45 +0000 Subject: [PATCH] fix(mcp): resolve scoped, connect and discovery routes through one exact-first, ip-aware lookup and stop denied names widening to access groups Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 6 -- .../mcp_server/mcp_server_manager.py | 6 +- .../_experimental/mcp_server/operations.py | 29 +++++--- .../mcp/test_mcp_caller_sign_in.py | 71 +++++++++++++++++- .../mcp_server/test_mcp_server.py | 72 ++++++++++++++++++- .../mcp_server/test_mcp_server_manager.py | 28 ++++++++ 6 files changed, 190 insertions(+), 22 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 5c4894c3bd6..9e262432ade 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -528,12 +528,6 @@ def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCP global_mcp_server_manager, ) - by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip) - if by_name is not None: - return by_name - by_id: Final = global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip) - if by_id is not None: - return by_id return global_mcp_server_manager.get_mcp_server_answering_to(lookup, client_ip=client_ip) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 707fa7625d7..8fbb075c9e2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7202,10 +7202,14 @@ class MCPServerManager: def get_mcp_server_answering_to(self, name: str, client_ip: str | None = None) -> MCPServer | None: """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.""" + priority first, then the exact ``server_id``, then the name 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 + by_id: Final = self.get_mcp_server_by_id(name, client_ip=client_ip) + if by_id is not None: + return by_id 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 9473f139def..c04c32a6f29 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -6,7 +6,7 @@ import types import uuid from collections.abc import Mapping, Sequence from datetime import datetime -from typing import Any, Final, NoReturn, TypeAlias, overload +from typing import Any, Final, Literal, NoReturn, TypeAlias, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -439,6 +439,7 @@ async def _dispatch_virtual_mcp_tool( async def _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers: Sequence[str] | None, allowed_mcp_servers: list[MCPServer], + client_ip: str | None = None, ) -> list[MCPServer]: """ Get the filtered MCP servers from the MCP server names. @@ -455,10 +456,12 @@ 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: - if (scoped := _scoped_server(server_or_group, allowed_mcp_servers)) is not None: + 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: filtered_server[scoped.server_id] = scoped - - if scoped is None: + else: try: access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( [server_or_group] @@ -495,14 +498,16 @@ 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) +def _scoped_server( + name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None +) -> MCPServer | Literal["denied"] | None: + """The granted server a scoped ``name`` selects: the registry's ``get_mcp_server_answering_to`` pick, made + with the same ``client_ip`` the connect preflight and discovery use, when the caller holds it. ``"denied"`` + when the registry names a server the caller does not hold, so the name is not retried as an access group. + ``None`` when the registry cannot place the name, after trying the granted servers answering to it.""" + registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) 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.server_id == registry_pick.server_id), "denied") return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) @@ -684,6 +689,7 @@ async def _get_allowed_mcp_servers( allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=mcp_servers, allowed_mcp_servers=allowed_mcp_servers, + client_ip=client_ip, ) return allowed_mcp_servers @@ -2431,6 +2437,7 @@ async def call_mcp_tool( allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=mcp_servers, allowed_mcp_servers=allowed_mcp_servers, + client_ip=client_ip, ) if mcp_servers and not allowed_mcp_servers: await raise_denied_scoped_mcp_access( diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index 9e940470031..bf5d7de3e32 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -24,14 +24,34 @@ ADD: Final = {"a": 2, "b": 3} ACCEPT: Final = {"Accept": "application/json, text/event-stream"} -def _rpc(gateway: Gateway, path: str, key: str, headers: dict[str, str]) -> httpx.Response: +def _rpc( + gateway: Gateway, + path: str, + key: str, + headers: dict[str, str], + method: str = "initialize", + params: Mapping[str, object] = INITIALIZE, +) -> httpx.Response: return gateway.client.post( path, - json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": INITIALIZE}, + json={"jsonrpc": "2.0", "id": 1, "method": method, "params": dict(params)}, headers={"x-litellm-api-key": key, **ACCEPT, **headers}, ) +def _sse_data(response: httpx.Response) -> str: + return next(line[5:].strip() for line in response.text.splitlines() if line.startswith("data:")) + + +def _advertised(gateway: Gateway, segment: str) -> tuple[int, tuple[str, ...], object]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/mcp/{segment}") + document: Final = response.json() + issuers: Final = tuple( + str(issuer).removesuffix(f"/{segment}") for issuer in document.get("authorization_servers", ()) + ) + return response.status_code, issuers, document.get("scopes_supported") + + def _sign_in_config(guardrail_params: dict[str, object], path: Path) -> Path: config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) config["guardrails"] = [{"guardrail_name": "signin" + uuid.uuid4().hex, "litellm_params": guardrail_params}] @@ -115,6 +135,53 @@ def test_alias_first_lookup_wins_over_a_server_whose_name_matches_the_alias(gate assert any(name.endswith("add") for name in names), names +def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_and_calls( + gateway: Gateway, tmp_path: Path +) -> None: + with mcp_peer() as obo_peer, mcp_peer() as math_peer: + stem: Final = "gh" + uuid.uuid4().hex[:6] + cased: Final = stem.capitalize() + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["mcp_servers"] = { + stem + "_obo": { + "alias": stem, + "transport": "http", + "url": obo_peer.url, + "auth_type": "oauth2_token_exchange", + "token_exchange_endpoint": "http://127.0.0.1:9/token", + "credentials": {"client_id": "obo-client", "client_secret": "obo-secret"}, + } + } + path: Final = tmp_path / "case-collision.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + by_name: Final = register_mcp(scenario, math_peer, cased) + key: Final = scenario.key(object_permission={"mcp_servers": [by_name]}) + + init: Final = _rpc(candidate, f"/mcp/{cased}", key, {}) + assert init.status_code == 200, init.text + listed: Final = _rpc(candidate, f"/mcp/{cased}", key, {}, method="tools/list") + assert listed.status_code == 200, listed.text + tools: Final = json.loads(_sse_data(listed))["result"]["tools"] + add: Final = next(tool["name"] for tool in tools if tool["name"].endswith("add")) + called: Final = _rpc( + candidate, f"/mcp/{cased}", key, {}, method="tools/call", params={"name": add, "arguments": ADD} + ) + assert called.status_code == 200, called.text + assert json.loads(_sse_data(called))["result"]["content"][0]["text"] == "5", called.text + assert len(tool_calls(math_peer.drain())) == 1 + assert tool_calls(obo_peer.drain()) == () + assert _advertised(candidate, cased) == _advertised(candidate, by_name) + + challenged: Final = _rpc(candidate, f"/mcp/{stem}", key, {}) + assert challenged.status_code == 401, challenged.text + assert f'resource_metadata="/.well-known/oauth-protected-resource/mcp/{stem}"' in challenged.headers.get( + "www-authenticate", "" + ) + assert _advertised(candidate, stem) == _advertised(candidate, stem + "_obo") + assert _advertised(candidate, stem) != _advertised(candidate, cased) + + def test_jwt_signer_verifies_the_bearer_that_admitted_the_call(gateway: Gateway, tmp_path: Path) -> None: signer_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) jwk: Final = json.loads(jwt_algorithms.RSAAlgorithm.to_jwk(signer_key.public_key())) 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 03d4414ec7a..7da59cae4d0 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 @@ -6,7 +6,7 @@ import os from datetime import datetime, timedelta from types import SimpleNamespace from typing import Final -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, call, patch import httpx import pytest @@ -8623,6 +8623,74 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a assert only_b == [], "a name the registry gives to an ungranted server must not fall through to another" +@pytest.mark.asyncio +async def test_scoped_name_of_an_ungranted_server_is_not_retried_as_an_access_group(): + 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) + 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}) + 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"], + ) as groups: + denied = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["shared"], allowed_mcp_servers=[member] + ) + unknown = 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"])] + + +@pytest.mark.asyncio +async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the_connect_preflight(): + 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="gh", + transport=MCPTransport.http, + available_on_public_internet=False, + ) + public = MCPServer(server_id="u-id", name="u", server_name="u", alias="Gh", transport=MCPTransport.http) + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update({"p-id": private, "u-id": public}) + 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=[], + ): + external = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["gh"], allowed_mcp_servers=[public], client_ip="203.0.113.7" + ) + internal = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["gh"], allowed_mcp_servers=[public], client_ip=None + ) + assert global_mcp_server_manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None + assert global_mcp_server_manager.get_mcp_server_answering_to("gh", client_ip=None) is private + finally: + global_mcp_server_manager.registry.clear() + + assert [s.server_id for s in external] == ["u-id"], "the router must apply the connect preflight's IP filter" + assert internal == [] + + @pytest.mark.asyncio async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): """ @@ -9470,7 +9538,7 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): ), patch( "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", - new=AsyncMock(side_effect=lambda mcp_servers, allowed_mcp_servers: allowed_mcp_servers), + new=AsyncMock(side_effect=lambda mcp_servers, allowed_mcp_servers, client_ip=None: allowed_mcp_servers), ), ): mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["legacy-m2m-id"]) 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 ccf9d7ce516..87d87a3e74c 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 @@ -6253,6 +6253,34 @@ class TestMCPServerManager: assert manager.get_mcp_server_answering_to("gh") is by_alias assert manager.get_mcp_server_answering_to("GH") is by_alias + @pytest.mark.parametrize("pinned_first", [True, False], ids=["pinned-id-listed-first", "alias-listed-first"]) + def test_answering_to_and_discovery_agree_on_a_pinned_id_that_another_alias_case_folds_to(self, pinned_first): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + pinned = MCPServer(server_id="foo", name="pinned", server_name="pinned", transport=MCPTransport.http) + by_alias = MCPServer( + server_id="b-id", + name="b", + server_name="b", + alias="Foo", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + registry = {"foo": pinned, "b-id": by_alias} if pinned_first else {"b-id": by_alias, "foo": pinned} + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry.update(registry) + try: + for name in ("foo", "Foo", "FOO"): + connected = global_mcp_server_manager.get_mcp_server_answering_to(name) + discovered = discoverable_endpoints._resolve_mcp_server_by_name_or_id(name, client_ip=None) + assert discovered is connected, name + assert global_mcp_server_manager.get_mcp_server_answering_to("foo") is pinned + assert global_mcp_server_manager.get_mcp_server_answering_to("Foo") is by_alias + assert global_mcp_server_manager.get_mcp_server_answering_to("FOO") is by_alias + finally: + global_mcp_server_manager.registry.clear() + def test_remove_server_drops_only_its_own_tool_mapping_rows(self): manager = self._manager_with_deepwiki_and_huggingface()