mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(mcp): keep wildcard tool grants as legacy entries during conversion
Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
parent
6d42ebd991
commit
27207bae53
4 changed files with 168 additions and 9 deletions
|
|
@ -22,6 +22,7 @@ from pydantic import TypeAdapter
|
|||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_ALL_TOOLS_WILDCARD
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy._experimental.mcp_server.tool_classification import classify_tool_op
|
||||
from litellm.proxy._types import SpecialMCPServerName
|
||||
|
|
@ -101,6 +102,17 @@ class BackfillReport:
|
|||
skipped_no_grants: frozenset[str]
|
||||
|
||||
|
||||
def wildcard_permission_keys(
|
||||
tool_permissions: Mapping[str, Sequence[str]] | None,
|
||||
) -> frozenset[str]:
|
||||
"""Keys whose stored tool list contains ``MCP_ALL_TOOLS_WILDCARD``."""
|
||||
return frozenset(
|
||||
server_id
|
||||
for server_id, stored in (tool_permissions or {}).items()
|
||||
if stored and MCP_ALL_TOOLS_WILDCARD in stored
|
||||
)
|
||||
|
||||
|
||||
def convert_row(
|
||||
row: LiteLLM_ObjectPermissionTable,
|
||||
inventories: Inventories,
|
||||
|
|
@ -112,22 +124,33 @@ def convert_row(
|
|||
all-proxy sentinel already expanded); ``None`` marks a server whose
|
||||
catalog could not be discovered and makes the row Unavailable. Servers
|
||||
granted only through toolsets are excluded by the caller and stay closed.
|
||||
Wildcard entries are explicit grants of every current and future tool, so
|
||||
they stay in the retained legacy map verbatim and need no inventory.
|
||||
"""
|
||||
legacy: Final = row.mcp_tool_permissions or MappingProxyType({})
|
||||
wildcard_servers: Final[frozenset[str]] = wildcard_permission_keys(legacy)
|
||||
missing: Final[frozenset[str]] = frozenset(
|
||||
server_id for server_id, inventory in inventories.items() if inventory is None
|
||||
server_id
|
||||
for server_id, inventory in inventories.items()
|
||||
if inventory is None and server_id not in wildcard_servers
|
||||
)
|
||||
if missing:
|
||||
return Unavailable(server_ids=missing)
|
||||
|
||||
legacy: Final = row.mcp_tool_permissions or MappingProxyType({})
|
||||
remaining_permissions: Final[Mapping[str, Sequence[str]]] = MappingProxyType(
|
||||
{server_id: stored for server_id, stored in legacy.items() if server_id not in inventories or not stored}
|
||||
{
|
||||
server_id: stored
|
||||
for server_id, stored in legacy.items()
|
||||
if server_id not in inventories or not stored or server_id in wildcard_servers
|
||||
}
|
||||
)
|
||||
override_entries: Final = MappingProxyType(
|
||||
{
|
||||
server_id: _server_override_entry(legacy.get(server_id), inventory or MappingProxyType({}))
|
||||
for server_id, inventory in inventories.items()
|
||||
if server_id not in legacy or legacy[server_id]
|
||||
if (server_id not in legacy or legacy[server_id])
|
||||
and server_id not in wildcard_servers
|
||||
and MCP_ALL_TOOLS_WILDCARD not in (legacy.get(server_id) or ())
|
||||
}
|
||||
)
|
||||
overrides: Final[Mapping[str, MCPToolOverrideEntry]] = MappingProxyType(
|
||||
|
|
@ -265,7 +288,12 @@ async def _convert_one_row(
|
|||
granted: Final = await resolve_granted_server_ids(row, manager)
|
||||
if not granted:
|
||||
return "skipped"
|
||||
conversion: Final = convert_row(row, await gather_inventories(granted, manager, inventory_cache))
|
||||
wildcard_servers: Final[frozenset[str]] = frozenset(
|
||||
server_id
|
||||
for server_id, tools in manager.expand_tool_permissions(row.mcp_tool_permissions).items()
|
||||
if tools and MCP_ALL_TOOLS_WILDCARD in tools
|
||||
)
|
||||
conversion: Final = convert_row(row, await gather_inventories(granted - wildcard_servers, manager, inventory_cache))
|
||||
if isinstance(conversion, Unavailable):
|
||||
return conversion.server_ids
|
||||
stored_fields: Final = (
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from pydantic import TypeAdapter
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import MCP_ALL_TOOLS_WILDCARD
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import ObjectPermissionDict, SpecialMCPServerName, SpecialMCPServerNames
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
|
||||
|
|
@ -117,8 +118,16 @@ async def _convert_unversioned_object_permission(
|
|||
remaining_grants: Final = await resolve_granted_server_ids(post_update_row, global_mcp_server_manager)
|
||||
if not remaining_grants:
|
||||
return MappingProxyType({})
|
||||
wildcard_servers: Final[frozenset[str]] = frozenset(
|
||||
server_id
|
||||
for server_id, tools in global_mcp_server_manager.expand_tool_permissions(
|
||||
post_update_row.mcp_tool_permissions
|
||||
).items()
|
||||
if tools and MCP_ALL_TOOLS_WILDCARD in tools
|
||||
)
|
||||
conversion: Final = convert_row(
|
||||
existing_object_permission, await gather_inventories(remaining_grants, global_mcp_server_manager)
|
||||
existing_object_permission,
|
||||
await gather_inventories(remaining_grants - wildcard_servers, global_mcp_server_manager),
|
||||
)
|
||||
if isinstance(conversion, Unavailable):
|
||||
raise HTTPException(
|
||||
|
|
@ -441,6 +450,21 @@ async def reject_ambiguous_mcp_tool_override_keys(
|
|||
"""
|
||||
requested: Final = _mcp_tool_override_entries(new_mcp_tool_overrides)
|
||||
stored: Final = _mcp_tool_override_entries(existing_mcp_tool_overrides)
|
||||
wildcard_servers: Final = sorted(
|
||||
identifier
|
||||
for identifier, entry in requested.items()
|
||||
if MCP_ALL_TOOLS_WILDCARD in (entry.get("allow") or ()) or MCP_ALL_TOOLS_WILDCARD in (entry.get("deny") or ())
|
||||
)
|
||||
if wildcard_servers:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail={ # mutable-ok: HTTPException.detail has no immutable form; same shape as the sibling errors here
|
||||
"error": (
|
||||
f"mcp_tool_overrides entries for {wildcard_servers} may not contain '{MCP_ALL_TOOLS_WILDCARD}': "
|
||||
"the wildcard belongs in mcp_tool_permissions, where it grants every current and future tool."
|
||||
)
|
||||
},
|
||||
)
|
||||
resolved: Final = await _resolve_mcp_server_identifiers_to_ids(
|
||||
identifiers=frozenset(identifier for identifier, entry in requested.items() if stored.get(identifier) != entry),
|
||||
prisma_client=prisma_client,
|
||||
|
|
|
|||
|
|
@ -251,3 +251,50 @@ async def test_resolve_granted_server_ids_ignores_tool_overrides():
|
|||
):
|
||||
granted = await resolve_granted_server_ids(row, manager)
|
||||
assert granted == frozenset()
|
||||
|
||||
|
||||
def test_convert_wildcard_entry_stays_legacy_and_needs_no_inventory():
|
||||
"""A ``["*"]`` entry is an explicit grant of every current and future tool,
|
||||
so conversion keeps it in the legacy map verbatim even without inventory."""
|
||||
row = _row(mcp_tool_permissions={"s1": ["*"], "s2": ["a"]})
|
||||
result = convert_row(row, {"s1": None, "s2": {"a": "read a", "b": "write b"}})
|
||||
assert isinstance(result, ConvertedRow)
|
||||
assert result.mcp_tool_permissions == {"s1": ["*"]}
|
||||
assert "s1" not in result.mcp_tool_overrides
|
||||
assert result.mcp_tool_overrides["s2"] == {"allow": [], "deny": ["b"]}
|
||||
assert result.mcp_tool_permissions_archive == {"s1": ["*"], "s2": ["a"]}
|
||||
assert result.mcp_permission_version == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_keeps_wildcard_entries_and_evaluator_reads_them_unrestricted():
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
level_allowed_tools,
|
||||
)
|
||||
|
||||
row = _row(mcp_servers=[], mcp_tool_permissions={"s1": ["*"], "s2": ["a"]})
|
||||
prisma = _prisma([row])
|
||||
manager = _manager(inventories={"s2": {"a": "read a", "b": "write b"}})
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
AsyncMock(return_value=[]),
|
||||
):
|
||||
report = await run_mcp_tool_permission_backfill(prisma, manager)
|
||||
assert report.converted == {"perm-1"}
|
||||
manager.fetch_unfiltered_inventory.assert_awaited_once_with("s2")
|
||||
|
||||
import json
|
||||
|
||||
data = prisma.db.litellm_objectpermissiontable.update_many.await_args.kwargs["data"]
|
||||
assert json.loads(data["mcp_tool_permissions"]) == {"s1": ["*"]}
|
||||
assert json.loads(data["mcp_tool_overrides"]) == {"s2": {"allow": [], "deny": ["b"]}}
|
||||
|
||||
converted = _row(mcp_servers=[], mcp_tool_permissions={"s1": ["*"]}, mcp_permission_version=1)
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
manager,
|
||||
):
|
||||
allowed = level_allowed_tools(
|
||||
row=converted, server_id="s1", grants_server=True, toolset_tools=None, inventory={}
|
||||
)
|
||||
assert allowed is None
|
||||
|
|
|
|||
|
|
@ -1685,7 +1685,67 @@ async def test_upsert_revoked_server_needs_no_inventory():
|
|||
|
||||
manager.fetch_unfiltered_inventory.assert_awaited_once_with("server-a")
|
||||
assert upsert.record["mcp_permission_version"] == 1
|
||||
assert json.loads(upsert.record["mcp_tool_overrides"]) == {
|
||||
"server-a": {"allow": ["delete_item"], "deny": []}
|
||||
}
|
||||
assert json.loads(upsert.record["mcp_tool_overrides"]) == {"server-a": {"allow": ["delete_item"], "deny": []}}
|
||||
assert upsert.record["mcp_servers"] == ["server-a"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upsert_wildcard_server_needs_no_inventory():
|
||||
"""A v0 row whose residual grant is a ``["*"]`` wildcard converts without
|
||||
discovering that server: the entry is an explicit all-tools grant that
|
||||
stays in the retained legacy map."""
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
|
||||
existing_row = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-id",
|
||||
mcp_servers=[],
|
||||
mcp_tool_permissions={"s1": ["*"], "s2": ["a"]},
|
||||
mcp_permission_version=0,
|
||||
)
|
||||
mock_prisma = _make_ambiguity_prisma()
|
||||
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_row)
|
||||
|
||||
manager = MagicMock()
|
||||
manager.expand_permission_list = MagicMock(side_effect=lambda servers: servers)
|
||||
manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {})
|
||||
manager.expand_tool_overrides = MagicMock(side_effect=lambda overrides: overrides or {})
|
||||
manager.get_registry = MagicMock(return_value={})
|
||||
manager.fetch_unfiltered_inventory = AsyncMock(return_value={"a": "read a", "b": "write b"})
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
manager,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
AsyncMock(return_value=[]),
|
||||
),
|
||||
):
|
||||
upsert = await prepare_object_permission_upsert(
|
||||
new_object_permission={"mcp_tool_permissions": {"s1": ["*"], "s2": ["a"]}},
|
||||
existing_object_permission_id="perm-id",
|
||||
prisma_client=mock_prisma,
|
||||
)
|
||||
|
||||
manager.fetch_unfiltered_inventory.assert_awaited_once_with("s2")
|
||||
assert upsert.record["mcp_permission_version"] == 1
|
||||
assert json.loads(upsert.record["mcp_tool_permissions"]) == {"s1": ["*"], "s2": ["a"]}
|
||||
assert "*" not in str(upsert.record["mcp_tool_overrides"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_overrides_rejects_wildcard_entries():
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
reject_ambiguous_mcp_tool_override_keys,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await reject_ambiguous_mcp_tool_override_keys(
|
||||
new_mcp_tool_overrides={"server-a": {"allow": ["*"], "deny": []}},
|
||||
existing_mcp_tool_overrides=None,
|
||||
prisma_client=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "server-a" in str(exc_info.value.detail)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue