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>
This commit is contained in:
yucheng 2026-09-30 11:26:45 +00:00
parent f89b92763b
commit 2172b8b60e
6 changed files with 190 additions and 22 deletions

View file

@ -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)

View file

@ -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], ...]] = (

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, 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(

View file

@ -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()))

View file

@ -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"])

View file

@ -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()