From f2a876f23f6c3a875c7c3b4444c2213fd9ccfbe4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 12:24:16 +0000 Subject: [PATCH] fix(mcp): keep a route hidden from a client ip from rerouting to a case variant of its name get_mcp_server_answering_to stops at the pass that finds an exact name or id and hides it from client_ip instead of falling through to the case-insensitive and prefix passes, and _scoped_server treats a name the registry knows for some caller but not this one as denied, so connect, discovery and scoped routing all refuse the hidden route instead of serving a public server whose alias only differs by case Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 11 +++-- .../_experimental/mcp_server/operations.py | 7 ++- .../mcp/test_mcp_caller_sign_in.py | 48 +++++++++++++++++++ .../mcp_server/test_discoverable_endpoints.py | 20 ++++---- .../mcp_server/test_mcp_server.py | 6 ++- .../mcp_server/test_mcp_server_manager.py | 18 +++++++ 6 files changed, 92 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8fbb075c9e2..5282b228f92 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -7203,13 +7203,14 @@ class MCPServerManager: """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 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) + form routing accepts. A name that denotes a server hidden from ``client_ip`` resolves to ``None`` at the + pass that found it: it never falls through to a looser pass that could name another server.""" + exact: Final = self.get_mcp_server_by_name(name) if exact is not None: - return exact - by_id: Final = self.get_mcp_server_by_id(name, client_ip=client_ip) + return exact if self._is_server_accessible_from_ip(exact, client_ip) else None + by_id: Final = self.get_mcp_server_by_id(name) if by_id is not None: - return by_id + return by_id if self._is_server_accessible_from_ip(by_id, client_ip) else None 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 c04c32a6f29..15a9183ef7a 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -503,11 +503,14 @@ def _scoped_server( ) -> 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.""" + when the registry names a server the caller does not hold, or one hidden from ``client_ip``, so the name + is neither rerouted to another granted server nor retried as an access group. ``None`` when the registry + cannot place the name for any caller, 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), "denied") + if global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: + return "denied" return next((s for s in allowed_mcp_servers if s and _server_answers_to(s, name)), None) diff --git a/tests/integration/mcp/test_mcp_caller_sign_in.py b/tests/integration/mcp/test_mcp_caller_sign_in.py index bf5d7de3e32..9755031fe3c 100644 --- a/tests/integration/mcp/test_mcp_caller_sign_in.py +++ b/tests/integration/mcp/test_mcp_caller_sign_in.py @@ -182,6 +182,54 @@ def test_exact_name_wins_over_a_case_folded_config_alias_for_connect_discovery_a assert _advertised(candidate, stem) != _advertised(candidate, cased) +def test_name_of_a_server_hidden_from_an_external_ip_does_not_reroute_to_a_case_variant( + gateway: Gateway, tmp_path: Path +) -> None: + stem: Final = "gh" + uuid.uuid4().hex[:6] + cased: Final = stem.capitalize() + external: Final = {"X-Forwarded-For": "203.0.113.7"} + with mcp_peer() as hidden_peer, mcp_peer() as public_peer: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config.get("general_settings", {}), + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["127.0.0.0/8"], + } + config["mcp_servers"] = { + stem: {"transport": "http", "url": hidden_peer.url, "available_on_public_internet": False} + } + path: Final = tmp_path / "external-ip.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + public_id: Final = register_mcp(scenario, public_peer, cased) + key: Final = scenario.key(object_permission={"mcp_servers": [public_id]}) + + hidden_name: Final = _rpc(candidate, f"/mcp/{stem}", key, external) + assert hidden_name.status_code == 403, hidden_name.text + assert "www-authenticate" not in hidden_name.headers, hidden_name.headers + assert ( + candidate.client.get(f"/.well-known/oauth-protected-resource/mcp/{stem}", headers=external).status_code + == 404 + ) + assert _rpc(candidate, f"/mcp/{stem}", key, {}).status_code == 403 + + own_name: Final = _rpc(candidate, f"/mcp/{cased}", key, external) + assert own_name.status_code == 200, own_name.text + listed: Final = _rpc(candidate, f"/mcp/{cased}", key, external, method="tools/list") + add: Final = next( + tool["name"] + for tool in json.loads(_sse_data(listed))["result"]["tools"] + if tool["name"].endswith("add") + ) + called: Final = _rpc( + candidate, f"/mcp/{cased}", key, external, 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(public_peer.drain())) == 1 + assert tool_calls(hidden_peer.drain()) == () + + 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_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 361d78059ad..ccf968e8179 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -3642,8 +3642,8 @@ async def test_authorize_resolves_server_by_id_when_name_lookup_fails(): assert response.status_code == 307 assert "https://provider.com/oauth/authorize" in response.headers["location"] - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3679,8 +3679,8 @@ async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails(): ) assert json.loads(result.body)["access_token"] == "token" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3711,8 +3711,8 @@ async def test_register_client_resolves_server_by_id_when_name_lookup_fails(): result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id) assert json.loads(result.body)["client_id"] == "registered-client" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3739,8 +3739,8 @@ async def test_protected_resource_metadata_resolves_server_by_id_when_name_looku assert result["authorization_servers"] == ["https://llm.example.com/mcp"] assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio @@ -3815,8 +3815,8 @@ def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fa assert result["scopes_supported"] == server.scopes assert result["issuer"] == f"https://llm.example.com/{server.server_id}" - by_name.assert_called_once_with(server.server_id, client_ip=None) - by_id.assert_called_once_with(server.server_id, client_ip=None) + by_name.assert_called_once_with(server.server_id) + by_id.assert_called_once_with(server.server_id) @pytest.mark.asyncio 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 7da59cae4d0..5cf35fb597c 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 @@ -8682,13 +8682,17 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the internal = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=["gh"], allowed_mcp_servers=[public], client_ip=None ) + by_own_name = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["Gh"], allowed_mcp_servers=[public], client_ip="203.0.113.7" + ) 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 external == [], "a name the preflight hides from this IP must not reroute to a case variant" assert internal == [] + assert [s.server_id for s in by_own_name] == ["u-id"] @pytest.mark.asyncio 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 87d87a3e74c..d7a1b3090a0 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,24 @@ 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("hidden_first", [True, False], ids=["hidden-listed-first", "public-listed-first"]) + def test_answering_to_never_reroutes_a_name_hidden_from_an_ip_to_a_case_variant(self, hidden_first): + manager = MCPServerManager() + hidden = MCPServer( + server_id="p-id", + name="gh", + server_name="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) + manager.registry = {"p-id": hidden, "u-id": public} if hidden_first else {"u-id": public, "p-id": hidden} + + assert manager.get_mcp_server_answering_to("gh", client_ip="203.0.113.7") is None + assert manager.get_mcp_server_answering_to("p-id", client_ip="203.0.113.7") is None + assert manager.get_mcp_server_answering_to("gh", client_ip="10.0.0.7") is hidden + assert manager.get_mcp_server_answering_to("Gh", client_ip="203.0.113.7") is public + @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