diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 283b139320c..9f31263e38e 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, Literal, NoReturn, TypeAlias, overload +from typing import Any, Final, NoReturn, TypeAlias, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -453,9 +453,7 @@ async def _get_allowed_mcp_servers_from_mcp_server_names( if mcp_servers is not None: for server_or_group in mcp_servers: 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: + if scoped is not None: filtered_server[scoped.server_id] = scoped else: try: @@ -494,25 +492,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], client_ip: str | None -) -> MCPServer | Literal["denied"] | None: - """The server a scoped ``name`` selects for the caller, in this order. ``"denied"`` when the registry's - ``get_mcp_server_answering_to`` pick, made with the same ``client_ip`` the connect preflight and discovery - use, is a server hidden from that IP. Otherwise the granted server answering to ``name``: the registry's - own pass order run over ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or - case variant the registry would pick. With no granted server answering: ``"denied"`` when the registry - places the name on a server the caller does not hold, so it is not retried as an access group; ``None`` - when the registry cannot place the name for any caller.""" - registry_pick: Final = global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) - if registry_pick is None and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None: - return "denied" - granted: Final = global_mcp_server_manager.get_mcp_server_answering_to( - name, client_ip=client_ip, among=allowed_mcp_servers - ) - if granted is not None: - return granted - return "denied" if registry_pick is not None else None +def _scoped_server(name: str, allowed_mcp_servers: Sequence[MCPServer], client_ip: str | None) -> MCPServer | None: + """The granted server a scoped ``name`` selects for the caller: the registry's own pass order run over + ``allowed_mcp_servers`` alone, so a granted server wins over an ungranted alias or case variant the registry + would pick. ``None`` when no granted server answers, or when the registry's own pick for ``name`` is a server + hidden from ``client_ip``; the caller then retries the name as an access group it holds.""" + if ( + global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip) is None + and global_mcp_server_manager.get_mcp_server_answering_to(name) is not None + ): + return None + return global_mcp_server_manager.get_mcp_server_answering_to(name, client_ip=client_ip, among=allowed_mcp_servers) async def raise_denied_scoped_mcp_access( diff --git a/tests/integration/mcp/test_mcp_access_matrix.py b/tests/integration/mcp/test_mcp_access_matrix.py index 13759ce245c..d199a3c34ff 100644 --- a/tests/integration/mcp/test_mcp_access_matrix.py +++ b/tests/integration/mcp/test_mcp_access_matrix.py @@ -1,3 +1,4 @@ +import json import uuid from typing import Final @@ -79,6 +80,36 @@ def test_subject_grant_lists_only_reachable_tools_and_denies_the_rest( assert not any(name.startswith(denied_alias) for name in denied_listed.tools), denied_listed.tools +@pytest.mark.parametrize("entry", ("mcp", "server_mcp")) +def test_access_group_named_like_an_ungranted_server_still_routes_the_groups_servers( + gateway: Gateway, entry: EntryPoint +) -> None: + with peer_of("http") as shadow, peer_of("http") as member, gateway.scenario() as scenario: + group: Final = "docs" + uuid.uuid4().hex[:8] + member_alias: Final = "mem" + uuid.uuid4().hex[:8] + register_mcp(scenario, shadow, group) + register_mcp(scenario, member, member_alias, mcp_access_groups=[group]) + key: Final = scenario.key(object_permission={"mcp_access_groups": [group]}) + unmatched: Final = "none" + uuid.uuid4().hex[:8] + denied: Final = McpCaller( + gateway, key, entry, unmatched, {"x-mcp-servers": unmatched} if entry == "mcp" else {} + ).list_tools() + assert (denied.status, json.loads(denied.error or "null"), denied.tools) == ( + ( + 200, + {"code": -32600, "message": f"The key is not allowed to access the requested MCP servers: {unmatched}"}, + (), + ) + if entry == "mcp" + else (404, {"detail": f"MCP server, toolset, or access group '{unmatched}' not found"}, ()) + ), denied.raw + selection: Final = {"x-mcp-servers": group} if entry == "mcp" else {} + listed: Final = McpCaller(gateway, key, entry, group, selection).list_tools() + assert listed.ok, listed.raw + assert set(listed.tools) == {f"{member_alias}-{tool}" for tool in ("add", "multiply", "fail")}, listed.tools + assert tool_calls(shadow.drain()) == () + + @pytest.mark.parametrize("entry", ("mcp", "server_mcp", "rest")) def test_key_without_any_grant_sees_no_scoped_server(gateway: Gateway, entry: EntryPoint) -> None: with peer_of("http") as peer, gateway.scenario() as scenario: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index ccf968e8179..b2753eebc14 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7,6 +7,7 @@ from base64 import urlsafe_b64encode from datetime import datetime, timedelta, timezone from typing import TYPE_CHECKING, Final from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import parse_qs, urlparse import pytest from fastapi import HTTPException @@ -392,6 +393,84 @@ def trust_xff(): yield +def _registered_gateway_oauth2_server(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server: Final = MCPServer( + server_id="oid-7f3a", + name="gwx", + server_name="gwx", + alias="gwx", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + client_id="gw-client", + client_secret="gw-secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read"], + ) + global_mcp_server_manager.registry.clear() + global_mcp_server_manager.registry[server.server_id] = server + return server + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"]) +async def test_authorization_server_doc_for_a_moved_lookup_matches_the_exact_name_doc(lookup): + """A case variant or server id now resolves like the exact name, so its AS metadata is the + exact-name doc with the requested spelling in the issuer and endpoint paths.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_authorization_server_mcp_standard, + ) + + _registered_gateway_oauth2_server() + request: Final = _mock_callback_request("http://litellm.example.com/") + + exact: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name="gwx") + moved: Final = await oauth_authorization_server_mcp_standard(request=request, mcp_server_name=lookup) + + assert exact["issuer"] == "http://litellm.example.com/mcp/gwx" + assert moved == { + key: value.replace("/gwx", f"/{lookup}") if isinstance(value, str) else value for key, value in exact.items() + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lookup", ["GWX", "oid-7f3a"], ids=["alias_case", "server_id"]) +async def test_authorize_relay_for_a_moved_lookup_redirects_like_the_exact_name(lookup): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize + + _registered_gateway_oauth2_server() + request: Final = _mock_callback_request("http://litellm.example.com/") + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper", return_value="sealed" + ): + exact: Final = await authorize( + request=request, mcp_server_name="gwx", redirect_uri="http://127.0.0.1:60108/callback", state="s1" + ) + moved: Final = await authorize( + request=request, mcp_server_name=lookup, redirect_uri="http://127.0.0.1:60108/callback", state="s1" + ) + + exact_target: Final = urlparse(exact.headers["location"]) + moved_target: Final = urlparse(moved.headers["location"]) + exact_query: Final = parse_qs(exact_target.query) + moved_query: Final = parse_qs(moved_target.query) + assert exact.status_code == 307 + assert exact_target._replace(query="") == urlparse("https://provider.com/oauth/authorize") + assert exact_query["client_id"] == ["gw-client"] + assert len(exact_query.pop("state")) == 1 and len(moved_query.pop("state")) == 1 + assert (moved.status_code, moved_target._replace(query=""), moved_query) == ( + exact.status_code, + exact_target._replace(query=""), + exact_query, + ) + + @pytest.mark.asyncio async def test_authorize_endpoint_includes_response_type(): """Test that authorize endpoint includes response_type=code parameter (fixes #15684)""" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 258d24ff619..4f8578e917b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -8723,33 +8723,36 @@ async def test_scoped_router_selects_the_server_the_connect_preflight_resolves(a @pytest.mark.asyncio -async def test_scoped_name_of_an_ungranted_server_is_not_retried_as_an_access_group(): +async def test_scoped_name_of_an_ungranted_server_is_retried_as_an_access_group_the_key_holds(): 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) + shadow = MCPServer(server_id="s-id", name="docs", server_name="docs", 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}) + global_mcp_server_manager.registry.update({"s-id": shadow, "m-id": member}) + group_members = {"docs": ["m-id"]} 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"], + side_effect=lambda names: [sid for name in names for sid in group_members.get(name, [])], ) as groups: - denied = await _get_allowed_mcp_servers_from_mcp_server_names( - mcp_servers=["shared"], allowed_mcp_servers=[member] + collided = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["docs"], allowed_mcp_servers=[member] ) - unknown = await _get_allowed_mcp_servers_from_mcp_server_names( + unmatched = 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"])] + assert [s.server_id for s in collided] == ["m-id"], ( + "a name owned by an ungranted server must still resolve to the access group of that name the key holds" + ) + assert unmatched == [], "a name matching neither a granted server nor a granted access group stays denied" + assert groups.await_args_list == [call(["docs"]), call(["team"])] @pytest.mark.asyncio @@ -8801,7 +8804,6 @@ async def test_scoped_router_hides_a_private_server_from_an_external_ip_like_the pytest.param(("a-id", "d-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-alias-holder-listed-first"), pytest.param(("d-id", "a-id"), "docs", ("d-id",), ["d-id"], id="alias-collision-exact-name-listed-first"), pytest.param(("g1", "g2"), "GITHUB", ("g2",), ["g2"], id="case-variant-collision"), - pytest.param(("p-id", "m-id"), "shared", ("m-id",), [], id="ungranted-only-name-stays-denied"), ], ) async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(registry, scope, granted, expected): @@ -8815,8 +8817,6 @@ async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(re "d-id": MCPServer(server_id="d-id", name="docs", server_name="docs", transport=MCPTransport.http), "g1": MCPServer(server_id="g1", name="GitHub", server_name="GitHub", transport=MCPTransport.http), "g2": MCPServer(server_id="g2", name="github", server_name="github", transport=MCPTransport.http), - "p-id": MCPServer(server_id="p-id", name="p", server_name="p", alias="shared", transport=MCPTransport.http), - "m-id": 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({server_id: servers[server_id] for server_id in registry}) @@ -8825,7 +8825,7 @@ async def test_scoped_router_prefers_the_granted_server_answering_to_the_name(re "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"], + return_value=[], ) as groups: selected: Final = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=[scope], allowed_mcp_servers=[servers[server_id] for server_id in granted] @@ -10401,6 +10401,75 @@ class TestPreemptive401ModeAware: client_ip=None, ) + async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException: + from litellm.proxy._experimental.mcp_server import server as server_module + + manager = mcp_operations.global_mcp_server_manager + manager.registry.clear() + manager.registry[server.server_id] = server + with ( + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + manager, "get_allowed_mcp_servers", AsyncMock(return_value=[server.server_id]) + ), + patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False), + pytest.raises(HTTPException) as exc, + ): + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": path, "headers": [(b"host", b"testserver")]}, + mcp_servers=[requested], + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), + client_ip=None, + ) + return exc.value + + @pytest.mark.asyncio + @pytest.mark.parametrize("delegate", [True, False], ids=["oauth_delegate", "gateway_interactive"]) + @pytest.mark.parametrize("shape", ["alias_case", "server_id", "x_mcp_servers"]) + async def test_moved_connect_shapes_get_the_exact_name_routes_challenge(self, delegate, shape, monkeypatch): + """Alias-case, server-id and x-mcp-servers connects resolve the same server the router serves, so + they answer the exact-name route's 401 with the requested spelling in the route segment.""" + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + server = _make_oauth2_server("gwx", oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate) + requested, path, exact_path = { + "alias_case": ("GWX", "/mcp/GWX", "/mcp/gwx"), + "server_id": (server.server_id, f"/mcp/{server.server_id}", "/mcp/gwx"), + "x_mcp_servers": ("GWX", "/mcp", "/mcp"), + }[shape] + + exact = await self._connect_with_a_grant(server, "gwx", exact_path) + moved = await self._connect_with_a_grant(server, requested, path) + + assert exact.status_code == 401 + assert (moved.status_code, moved.detail) == (exact.status_code, exact.detail) + exact_header = {k.lower(): v for k, v in (exact.headers or {}).items()}["www-authenticate"] + moved_header = {k.lower(): v for k, v in (moved.headers or {}).items()}["www-authenticate"] + assert "/gwx" in exact_header + assert moved_header == exact_header.replace("/gwx", f"/{requested}") + + @pytest.mark.asyncio + async def test_aggregate_connect_without_a_server_selection_is_not_challenged(self): + from litellm.proxy._experimental.mcp_server import server as server_module + + manager = mcp_operations.global_mcp_server_manager + manager.registry.clear() + for alias, delegate in (("gwx", False), ("relay", True)): + server = _make_oauth2_server(alias, oauth2_flow="authorization_code", delegate_auth_to_upstream=delegate) + manager.registry[server.server_id] = server + + with patch.object(manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False) as tokens: + await server_module._raise_preemptive_401_for_unauthenticated_servers( + scope={"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"testserver")]}, + mcp_servers=None, + oauth2_headers=None, + mcp_server_auth_headers=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), + client_ip=None, + ) + + assert tokens.await_count == 0, "an unselected aggregate connect must not probe any server for a token" + @pytest.mark.asyncio async def test_deferred_discovery_runs_before_delegate_challenge(self): from litellm.proxy._experimental.mcp_server import server as server_module