mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): retry a scoped name as an access group when the server it names is ungranted
_scoped_server returned a denial whenever the registry placed a scoped name on a server the caller does not hold, so a key granted only access group docs lost the group's servers when an ungranted server was also named docs. It now returns None there, as it does for a name hidden from the caller's IP, and the caller retries the name as an access group it holds, which is what the merge base did. Regression tests pin the moved oauth_delegate and gateway oauth2 route shapes (alias case, server id, aggregate x-mcp-servers, authorization-server document, authorize relay) to the exact-name route, and the unselected aggregate connect to no challenge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
71b946cd99
commit
ae7b0fb8ea
4 changed files with 206 additions and 37 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue