mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
23eb6bd8a5
commit
f89b92763b
5 changed files with 80 additions and 10 deletions
|
|
@ -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], ...]] = (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue