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