feat(mcp): backfill v0 tool permission rows and close toolset widening

Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
Devin AI 2026-09-25 00:39:40 +00:00
parent 49b7bd8da3
commit 766438bcc1
9 changed files with 751 additions and 29 deletions

View file

@ -109,15 +109,15 @@ def level_allowed_tools(
legacy: Final = global_mcp_server_manager.expand_tool_permissions(row.mcp_tool_permissions).get(server_id)
if legacy is not None:
return frozenset(legacy) | toolset
if not getattr(row, "mcp_permission_version", None):
if not row.mcp_permission_version:
return frozenset(toolset) if toolset_tools is not None else None
if not grants_server:
return None
overrides: Final = (
global_mcp_server_manager.expand_tool_overrides(getattr(row, "mcp_tool_overrides", None)).get(server_id) or {}
)
overrides: Final = global_mcp_server_manager.expand_tool_overrides(row.mcp_tool_overrides).get(server_id) or {}
allow: Final[frozenset[str]] = frozenset(overrides.get("allow") or ())
deny: Final[frozenset[str]] = frozenset(overrides.get("deny") or ())
if toolset_tools is not None:
return toolset - deny
from litellm.proxy._experimental.mcp_server.tool_classification import (
classify_tool_op,
)
@ -127,7 +127,7 @@ def level_allowed_tools(
for tool_name, description in inventory.items()
if tool_name not in deny and classify_tool_op(tool_name, description) != "delete"
)
return ((convention | allow) - deny) | toolset
return (convention | allow) - deny
def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: resolver returns a list
@ -2251,7 +2251,7 @@ class MCPRequestHandler:
*global_mcp_server_manager.expand_permission_list(row.mcp_servers or []),
*access_group_servers,
*global_mcp_server_manager.expand_tool_permissions(row.mcp_tool_permissions).keys(),
*global_mcp_server_manager.expand_tool_overrides(getattr(row, "mcp_tool_overrides", None)).keys(),
*global_mcp_server_manager.expand_tool_overrides(row.mcp_tool_overrides).keys(),
*toolset_perms.keys(),
}
return level_allowed_tools(
@ -2623,9 +2623,7 @@ class MCPRequestHandler:
global_mcp_server_manager.expand_tool_permissions(key_object_permission.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(key_object_permission, "mcp_tool_overrides", None)
).keys()
global_mcp_server_manager.expand_tool_overrides(key_object_permission.mcp_tool_overrides).keys()
)
# servers referenced by the key's toolset grants are part of the key's
@ -2730,11 +2728,7 @@ class MCPRequestHandler:
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
| set(legacy_access_group_servers)
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
| set(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permissions, "mcp_tool_overrides", None)
).keys()
)
| set(global_mcp_server_manager.expand_tool_overrides(object_permissions.mcp_tool_overrides).keys())
| (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()
| set(team_access_group_servers)
)
@ -2927,9 +2921,7 @@ class MCPRequestHandler:
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permissions, "mcp_tool_overrides", None)
).keys()
global_mcp_server_manager.expand_tool_overrides(object_permissions.mcp_tool_overrides).keys()
)
# servers referenced by the org's toolset grants are part of the org ceiling,
@ -3034,9 +3026,7 @@ class MCPRequestHandler:
global_mcp_server_manager.expand_tool_permissions(object_permission.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permission, "mcp_tool_overrides", None)
).keys()
global_mcp_server_manager.expand_tool_overrides(object_permission.mcp_tool_overrides).keys()
)
# Combine all lists
@ -3157,9 +3147,7 @@ class MCPRequestHandler:
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permissions, "mcp_tool_overrides", None)
).keys()
global_mcp_server_manager.expand_tool_overrides(object_permissions.mcp_tool_overrides).keys()
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
return tuple(

View file

@ -5394,6 +5394,29 @@ class MCPServerManager:
descriptions; empty when the server is unknown or never listed."""
return self.discovered_tool_inventory.get(server_id, {})
async def fetch_unfiltered_inventory(self, server_id: str) -> Mapping[str, str | None] | None:
"""List ``server_id``'s tools unfiltered by any caller's permissions.
Returns ``None`` when the server is unknown, when its auth mode needs a
per-user credential the proxy does not hold, or when discovery fails;
an empty mapping means the server answered tools/list with no tools."""
server: Final = self.get_mcp_server_by_id(server_id)
if server is None:
return None
if server.auth_type in {
MCPAuth.oauth2_token_exchange,
MCPAuth.oauth_delegate,
MCPAuth.oauth2_id_jag,
MCPAuth.true_passthrough,
}:
return None
try:
await self._get_tools_from_server(server)
except Exception as e: # noqa: BLE001 # any discovery failure means "inventory unavailable", never a partial empty
verbose_logger.warning("Backfill inventory fetch failed for server %s: %s", server_id, e)
return None
return self.discovered_inventory(server_id)
def _create_prefixed_prompts(
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
) -> list[Prompt]:

View file

@ -0,0 +1,315 @@
"""
Convert pre-overrides object_permission rows (``mcp_permission_version`` falsy)
to the mcp_tool_overrides storage so the convention evaluator can take over.
``convert_row`` is pure: given the row and the discovered tool inventory per
granted server it produces the fields to persist, or reports which servers
could not supply an inventory. ``run_mcp_tool_permission_backfill`` pages the
table at proxy boot, converts each residual row, and writes with an
update_many compare-and-set so a row edited mid-backfill is skipped rather
than clobbered.
"""
import json
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from pydantic import TypeAdapter
from litellm._logging import verbose_logger
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
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.types.mcp import MCPToolOverrideEntry
if TYPE_CHECKING:
from prisma import models as prisma_models
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy.utils import PrismaClient
RawRow = LiteLLM_ObjectPermissionTable | prisma_models.LiteLLM_ObjectPermissionTable
ToolInventory = Mapping[str, str | None]
Inventories = Mapping[str, ToolInventory | None]
_ROW_DUMP_ADAPTER: Final = TypeAdapter(dict[str, object])
@dataclass(frozen=True)
class ConvertedRow:
mcp_tool_overrides: Mapping[str, MCPToolOverrideEntry]
mcp_tool_permissions: Mapping[str, Sequence[str]]
mcp_tool_permissions_archive: Mapping[str, Sequence[str]]
mcp_permission_version: int = 1
@dataclass(frozen=True)
class Unavailable:
server_ids: frozenset[str]
ConversionResult = ConvertedRow | Unavailable
def _server_override_entry(
stored: Sequence[str] | None,
discovered: ToolInventory,
) -> MCPToolOverrideEntry | None:
allow: Final[frozenset[str]] = (
frozenset(
name for name in stored if name not in discovered or classify_tool_op(name, discovered[name]) == "delete"
)
if stored
else frozenset(
name for name, description in discovered.items() if classify_tool_op(name, description) == "delete"
)
)
deny: Final[frozenset[str]] = (
frozenset(
name
for name, description in discovered.items()
if name not in stored and classify_tool_op(name, description) != "delete"
)
if stored
else frozenset()
)
if not allow and not deny:
return None
return {"allow": sorted(allow), "deny": sorted(deny)}
@dataclass(frozen=True)
class BackfillReport:
converted: frozenset[str]
cas_missed: frozenset[str]
unavailable: Mapping[str, frozenset[str]]
skipped_no_grants: frozenset[str]
def convert_row(
row: LiteLLM_ObjectPermissionTable,
inventories: Inventories,
) -> ConversionResult:
"""Compute the converted fields for one row.
``inventories`` must carry an entry for every server the row grants
(expanded mcp_servers, access-group servers, tool-permission keys, the
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.
"""
missing: Final[frozenset[str]] = frozenset(
server_id for server_id, inventory in inventories.items() if inventory is None
)
if missing:
return Unavailable(server_ids=missing)
legacy: Final = row.mcp_tool_permissions or {}
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}
)
overrides: Final[Mapping[str, MCPToolOverrideEntry]] = MappingProxyType(
{
server_id: entry
for server_id, inventory in inventories.items()
if legacy.get(server_id) != []
for entry in (_server_override_entry(legacy.get(server_id), inventory or MappingProxyType({})),)
if entry is not None
}
)
return ConvertedRow(
mcp_tool_overrides=overrides,
mcp_tool_permissions=remaining_permissions,
mcp_tool_permissions_archive=MappingProxyType(dict(legacy)),
)
async def resolve_granted_server_ids(
row: LiteLLM_ObjectPermissionTable,
manager: "MCPServerManager",
) -> frozenset[str]:
"""Every concrete server_id the row grants, toolset-only servers excluded."""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
direct: Final[set[str]] = set() # mutable-ok: accumulation
for identifier in manager.expand_permission_list(list(row.mcp_servers or ())):
if identifier == SpecialMCPServerName.all_proxy_servers.value:
direct.update(manager.get_registry().keys())
else:
direct.add(identifier)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( # pyright: ignore[reportPrivateUsage] # shared access-group resolution owned by MCPRequestHandler
row.mcp_access_groups or []
)
return frozenset(
direct
| set(access_group_servers)
| set(manager.expand_tool_permissions(row.mcp_tool_permissions).keys())
| set(manager.expand_tool_overrides(row.mcp_tool_overrides).keys())
)
async def gather_inventories(
server_ids: frozenset[str],
manager: "MCPServerManager",
cache: dict[str, ToolInventory | None] | None = None,
) -> dict[str, ToolInventory | None]:
if cache is None:
return {server_id: await manager.fetch_unfiltered_inventory(server_id) for server_id in server_ids}
return {
server_id: (
cache[server_id]
if server_id in cache
else cache.setdefault(
server_id, await manager.fetch_unfiltered_inventory(server_id)
) # mutable-ok: run-scoped fetch cache
)
for server_id in server_ids
}
def converted_row_record(conversion: ConvertedRow) -> dict[str, object]:
record: Final[dict[str, object]] = {
"mcp_tool_overrides": dict(conversion.mcp_tool_overrides),
"mcp_tool_permissions": dict(conversion.mcp_tool_permissions),
"mcp_tool_permissions_archive": dict(conversion.mcp_tool_permissions_archive),
"mcp_permission_version": 1,
}
return record
def _to_model(row: "RawRow") -> LiteLLM_ObjectPermissionTable:
data: Final[Mapping[str, object]] = _ROW_DUMP_ADAPTER.validate_python(row.model_dump())
json_fields: Final = (
"mcp_tool_permissions",
"mcp_tool_overrides",
"mcp_tool_permissions_archive",
)
normalized: Final = {
key: (json.loads(value) if key in json_fields and isinstance(value, str) else value)
for key, value in data.items()
}
return LiteLLM_ObjectPermissionTable.model_validate(normalized)
async def _converted_page(
rows: Sequence["RawRow"],
prisma_client: "PrismaClient",
manager: "MCPServerManager",
inventory_cache: dict[str, ToolInventory | None],
) -> BackfillReport:
models: Final[Sequence[LiteLLM_ObjectPermissionTable]] = [_to_model(row) for row in rows]
outcomes: Final = [
await _convert_one_row(raw, model, prisma_client, manager, inventory_cache)
for raw, model in zip(rows, models, strict=True)
]
return BackfillReport(
converted=frozenset(
model.object_permission_id
for model, outcome in zip(models, outcomes, strict=True)
if outcome == "converted"
),
cas_missed=frozenset(
model.object_permission_id
for model, outcome in zip(models, outcomes, strict=True)
if outcome == "cas_missed"
),
unavailable=MappingProxyType(
{
model.object_permission_id: outcome
for model, outcome in zip(models, outcomes, strict=True)
if isinstance(outcome, frozenset)
}
),
skipped_no_grants=frozenset(
model.object_permission_id for model, outcome in zip(models, outcomes, strict=True) if outcome == "skipped"
),
)
async def _convert_one_row(
raw_row: "RawRow",
row: LiteLLM_ObjectPermissionTable,
prisma_client: "PrismaClient",
manager: "MCPServerManager",
inventory_cache: dict[str, ToolInventory | None],
) -> str | frozenset[str]:
"""Convert one row and CAS-write it. Returns the outcome tag, or the
unavailable server ids as a frozenset when the row cannot be converted."""
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))
if isinstance(conversion, Unavailable):
return conversion.server_ids
updated: Final = await ObjectPermissionRepository(prisma_client).table.update_many(
where={
"object_permission_id": row.object_permission_id,
"mcp_permission_version": {"in": [0, None]},
"mcp_servers": {"equals": raw_row.mcp_servers},
"mcp_access_groups": {"equals": raw_row.mcp_access_groups},
"mcp_toolsets": {"equals": raw_row.mcp_toolsets},
"mcp_tool_permissions": {"equals": raw_row.mcp_tool_permissions},
},
data={
"mcp_tool_overrides": json.dumps(dict(conversion.mcp_tool_overrides)),
"mcp_tool_permissions": json.dumps(dict(conversion.mcp_tool_permissions)),
"mcp_tool_permissions_archive": json.dumps(dict(conversion.mcp_tool_permissions_archive)),
"mcp_permission_version": 1,
},
)
return "cas_missed" if updated == 0 else "converted"
def _merge_reports(reports: Sequence[BackfillReport]) -> BackfillReport:
return BackfillReport(
converted=frozenset(row_id for report in reports for row_id in report.converted),
cas_missed=frozenset(row_id for report in reports for row_id in report.cas_missed),
unavailable=MappingProxyType(
{row_id: server_ids for report in reports for row_id, server_ids in report.unavailable.items()}
),
skipped_no_grants=frozenset(row_id for report in reports for row_id in report.skipped_no_grants),
)
async def run_mcp_tool_permission_backfill(
prisma_client: "PrismaClient",
manager: "MCPServerManager",
batch_size: int = 200,
) -> BackfillReport:
"""Convert every residual v0 permission row. Idempotent: converted rows
stop matching the version filter, and CAS-missed or Unavailable rows are
left for the next boot."""
table: Final = ObjectPermissionRepository(prisma_client).table
reports: Final[list[BackfillReport]] = [] # mutable-ok: accumulated per page
inventory_cache: Final[dict[str, ToolInventory | None]] = {} # mutable-ok: one fetch per server per run
cursor: str | None = None # rebind-ok: cursor pagination
while True:
rows = await table.find_many(
where={"OR": [{"mcp_permission_version": 0}, {"mcp_permission_version": None}]},
order={"object_permission_id": "asc"},
take=batch_size,
cursor={"object_permission_id": cursor} if cursor is not None else None,
skip=1 if cursor is not None else None,
)
if not rows:
break
reports.append(await _converted_page(rows, prisma_client, manager, inventory_cache))
if len(rows) < batch_size:
break
cursor = rows[-1].object_permission_id
report: Final = _merge_reports(reports)
verbose_logger.info(
"MCP tool permission backfill: converted=%s cas_missed=%s skipped_no_grants=%s unavailable=%s",
sorted(report.converted),
sorted(report.cas_missed),
sorted(report.skipped_no_grants),
{row_id: sorted(server_ids) for row_id, server_ids in report.unavailable.items()},
)
return report

View file

@ -651,7 +651,7 @@ async def _set_object_permission(
prisma_client=prisma_client,
)
await reject_ambiguous_mcp_tool_override_keys(
new_mcp_tool_overrides=getattr(data.object_permission, "mcp_tool_overrides", None),
new_mcp_tool_overrides=data.object_permission.mcp_tool_overrides,
existing_mcp_tool_overrides=None,
prisma_client=prisma_client,
)

View file

@ -80,6 +80,43 @@ class ObjectPermissionUpsert:
record: dict[str, object]
async def _convert_unversioned_object_permission(
existing_object_permission: "LiteLLM_ObjectPermissionTable | None",
) -> dict[str, object]:
"""Manual conversion path: a save over a stored v0 row first converts its
residual grants against discovered inventory, then applies the update on
top. Servers whose catalog cannot be discovered reject the save with 503
rather than silently dropping the tools admins had granted."""
if existing_object_permission is None or existing_object_permission.mcp_permission_version:
return {}
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.tool_permission_backfill import (
Unavailable,
convert_row,
converted_row_record,
gather_inventories,
resolve_granted_server_ids,
)
granted: Final = await resolve_granted_server_ids(existing_object_permission, global_mcp_server_manager)
if not granted:
return {}
conversion: Final = convert_row(
existing_object_permission, await gather_inventories(granted, global_mcp_server_manager)
)
if isinstance(conversion, Unavailable):
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail={
"error": "Could not discover tools for servers granted by this permission; retry once the MCP servers are reachable.",
"server_ids": sorted(conversion.server_ids),
},
)
return converted_row_record(conversion)
async def prepare_object_permission_upsert(
new_object_permission: Mapping[str, object],
existing_object_permission_id: str | None,
@ -104,11 +141,15 @@ async def prepare_object_permission_upsert(
existing_object_permission: Final = await ObjectPermissionRepository(prisma_client).table.find_unique(
where={"object_permission_id": object_permission_id},
)
existing_fields: Final[dict[str, object]] = (
existing_fields_raw: Final[dict[str, object]] = (
existing_object_permission.model_dump(exclude_unset=True, exclude_none=True)
if existing_object_permission is not None
else {}
)
existing_fields: Final[dict[str, object]] = {
**existing_fields_raw,
**await _convert_unversioned_object_permission(existing_object_permission),
}
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=new_object_permission.get("mcp_tool_permissions"),
existing_mcp_tool_permissions=existing_fields.get("mcp_tool_permissions"),
@ -133,6 +174,11 @@ async def prepare_object_permission_upsert(
else {}
),
**({"mcp_tool_overrides": safe_dumps(merged["mcp_tool_overrides"])} if "mcp_tool_overrides" in merged else {}),
**(
{"mcp_tool_permissions_archive": safe_dumps(merged["mcp_tool_permissions_archive"])}
if "mcp_tool_permissions_archive" in merged
else {}
),
}
return ObjectPermissionUpsert(object_permission_id=object_permission_id, record=record)
@ -536,7 +582,7 @@ async def _resolve_team_allowed_mcp_servers(
raw_tool_perms = team_object_permission.mcp_tool_permissions or {}
if isinstance(raw_tool_perms, str):
raw_tool_perms = json.loads(raw_tool_perms)
raw_tool_overrides: Final = _mcp_tool_override_entries(getattr(team_object_permission, "mcp_tool_overrides", None))
raw_tool_overrides: Final = _mcp_tool_override_entries(team_object_permission.mcp_tool_overrides)
tool_perm_servers: Final[list[str]] = list(raw_tool_perms.keys()) + list(raw_tool_overrides.keys())
raw_servers: Final = set(direct_servers + access_group_servers + tool_perm_servers)
resolved_servers: Final = await _resolve_mcp_server_identifiers_to_ids(
@ -625,9 +671,7 @@ async def _get_grandfathered_key_mcp_server_ids(
if existing_object_permission is None or prisma_client is None:
return frozenset()
raw_tool_perms: Final = existing_object_permission.mcp_tool_permissions or {}
raw_tool_overrides: Final = _mcp_tool_override_entries(
getattr(existing_object_permission, "mcp_tool_overrides", None)
)
raw_tool_overrides: Final = _mcp_tool_override_entries(existing_object_permission.mcp_tool_overrides)
tool_perm_keys: Final[frozenset[str]] = frozenset(
json.loads(raw_tool_perms).keys() if isinstance(raw_tool_perms, str) else raw_tool_perms.keys()
) | frozenset(raw_tool_overrides.keys())

View file

@ -8592,6 +8592,24 @@ class ProxyConfig:
except Exception as e:
verbose_proxy_logger.exception("litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db - %s", e)
async def _run_mcp_tool_permission_backfill() -> None:
from litellm.proxy._experimental.mcp_server.tool_permission_backfill import (
run_mcp_tool_permission_backfill,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_prisma_client_or_throw,
)
try:
await run_mcp_tool_permission_backfill(
prisma_client=get_prisma_client_or_throw("Database not connected"),
manager=global_mcp_server_manager,
)
except Exception as e: # noqa: BLE001 # startup backfill must never block boot; the next boot retries the residual
verbose_proxy_logger.warning("MCP tool permission backfill failed: %s", e)
asyncio.create_task(_run_mcp_tool_permission_backfill())
async def init_mcp_servers_from_db(self) -> None:
if self._should_load_db_object(object_type="mcp"):
await self._init_mcp_servers_in_db()

View file

@ -10042,3 +10042,21 @@ class TestLevelAllowedToolsConvention:
)
result = await self._allowed(row)
assert set(result) == {"list_items"}
async def test_toolset_grant_stays_closed_does_not_widen_to_convention(self):
row = _converted_row(
["server-a"],
mcp_toolsets=["toolset-1"],
mcp_tool_overrides={"server-a": {"allow": ["delete_item"]}},
)
result = await self._allowed(row, toolset_perms={"server-a": ["list_items"]})
assert set(result) == {"list_items"}
async def test_toolset_grant_deny_narrows_to_empty(self):
row = _converted_row(
["server-a"],
mcp_toolsets=["toolset-1"],
mcp_tool_overrides={"server-a": {"deny": ["list_items"]}},
)
result = await self._allowed(row, toolset_perms={"server-a": ["list_items"]})
assert set(result) == set()

View file

@ -0,0 +1,222 @@
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy._experimental.mcp_server.tool_permission_backfill import (
ConvertedRow,
Unavailable,
convert_row,
resolve_granted_server_ids,
run_mcp_tool_permission_backfill,
)
INVENTORY = {"list_items": "list things", "search_notes": "find notes", "delete_item": "remove one"}
def _row(**overrides):
fields = {
"object_permission_id": "perm-1",
"mcp_servers": ["server-a"],
"mcp_access_groups": [],
"mcp_tool_permissions": None,
"mcp_toolsets": None,
"mcp_permission_version": 0,
}
fields.update(overrides)
return LiteLLM_ObjectPermissionTable(**fields)
def test_convert_selected_delete_tool_goes_to_allow_and_unselected_nondelete_to_deny():
row = _row(mcp_tool_permissions={"server-a": ["list_items", "delete_item"]})
result = convert_row(row, {"server-a": INVENTORY})
assert isinstance(result, ConvertedRow)
assert result.mcp_tool_overrides["server-a"] == {
"allow": ["delete_item"],
"deny": ["search_notes"],
}
assert result.mcp_tool_permissions == {}
assert result.mcp_tool_permissions_archive == {"server-a": ["list_items", "delete_item"]}
assert result.mcp_permission_version == 1
def test_convert_preserves_undiscovered_stored_allows():
row = _row(mcp_tool_permissions={"server-a": ["ghost_tool"]})
result = convert_row(row, {"server-a": INVENTORY})
assert isinstance(result, ConvertedRow)
assert result.mcp_tool_overrides["server-a"]["allow"] == ["ghost_tool"]
def test_convert_empty_legacy_list_stays_deny_all_untouched():
row = _row(mcp_tool_permissions={"server-a": []})
result = convert_row(row, {"server-a": INVENTORY})
assert isinstance(result, ConvertedRow)
assert "server-a" not in result.mcp_tool_overrides
assert result.mcp_tool_permissions == {"server-a": []}
def test_convert_unrestricted_grant_keeps_only_current_deletes():
row = _row()
result = convert_row(row, {"server-a": INVENTORY})
assert isinstance(result, ConvertedRow)
assert result.mcp_tool_overrides["server-a"] == {"allow": ["delete_item"], "deny": []}
assert result.mcp_tool_permissions == {}
@pytest.mark.asyncio
async def test_converted_row_denies_new_delete_allows_new_nondelete():
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
level_allowed_tools,
)
row = _row()
conversion = convert_row(row, {"server-a": INVENTORY})
assert isinstance(conversion, ConvertedRow)
converted = _row(
mcp_tool_overrides=dict(conversion.mcp_tool_overrides),
mcp_permission_version=1,
)
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 {})
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
manager,
):
allowed = level_allowed_tools(
row=converted,
server_id="server-a",
grants_server=True,
toolset_tools=None,
inventory={**INVENTORY, "delete_link": "unlink", "create_item": "make one"},
)
assert allowed is not None
assert "create_item" in allowed
assert "delete_item" in allowed
assert "delete_link" not in allowed
def test_convert_inventory_unavailable_marks_row_without_writes():
row = _row()
result = convert_row(row, {"server-a": None})
assert isinstance(result, Unavailable)
assert result.server_ids == {"server-a"}
def _manager(inventories=None, registry=None):
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=registry or {"server-a": MagicMock(server_id="server-a")})
manager.fetch_unfiltered_inventory = AsyncMock(side_effect=lambda server_id: (inventories or {}).get(server_id))
return manager
def _prisma(rows, update_count=1):
prisma = MagicMock()
table = prisma.db.litellm_objectpermissiontable
table.find_many = AsyncMock(return_value=rows)
table.update_many = AsyncMock(return_value=update_count)
return prisma
@pytest.mark.asyncio
async def test_runner_converts_row_with_cas_update():
row = _row()
prisma = _prisma([row])
manager = _manager(inventories={"server-a": INVENTORY})
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"}
import json
data = prisma.db.litellm_objectpermissiontable.update_many.await_args.kwargs["data"]
assert json.loads(data["mcp_tool_overrides"])["server-a"]["allow"] == ["delete_item"]
assert data["mcp_permission_version"] == 1
where = prisma.db.litellm_objectpermissiontable.update_many.await_args.kwargs["where"]
assert where["object_permission_id"] == "perm-1"
assert where["mcp_permission_version"] == {"in": [0, None]}
@pytest.mark.asyncio
async def test_runner_unavailable_server_skips_row_no_write():
prisma = _prisma([_row()])
manager = _manager(inventories={"server-a": None})
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.unavailable == {"perm-1": frozenset({"server-a"})}
prisma.db.litellm_objectpermissiontable.update_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_runner_cas_miss_is_counted_without_retry():
prisma = _prisma([_row()], update_count=0)
manager = _manager(inventories={"server-a": INVENTORY})
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.cas_missed == {"perm-1"}
assert prisma.db.litellm_objectpermissiontable.update_many.await_count == 1
@pytest.mark.asyncio
async def test_runner_second_run_is_noop():
prisma = _prisma([])
manager = _manager()
report = await run_mcp_tool_permission_backfill(prisma, _manager())
assert report.converted == frozenset()
assert report.cas_missed == frozenset()
prisma.db.litellm_objectpermissiontable.update_many.assert_not_awaited()
manager.fetch_unfiltered_inventory.assert_not_awaited()
@pytest.mark.asyncio
async def test_runner_skips_row_with_no_grants():
row = _row(mcp_servers=[], mcp_tool_permissions=None)
prisma = _prisma([row])
manager = _manager()
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.skipped_no_grants == {"perm-1"}
prisma.db.litellm_objectpermissiontable.update_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_resolve_granted_server_ids_expands_all_proxy_sentinel():
row = _row(mcp_servers=["all-proxy-mcpservers"])
registry = {
"server-a": MagicMock(server_id="server-a"),
"server-b": MagicMock(server_id="server-b"),
}
manager = _manager(registry=registry)
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
AsyncMock(return_value=[]),
):
granted = await resolve_granted_server_ids(row, manager)
assert granted == frozenset({"server-a", "server-b"})
@pytest.mark.asyncio
async def test_resolve_granted_server_ids_excludes_toolset_only_servers():
row = _row(mcp_servers=[], mcp_toolsets=["toolset-1"])
manager = _manager()
with patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
AsyncMock(return_value=[]),
):
granted = await resolve_granted_server_ids(row, manager)
assert granted == frozenset()

View file

@ -211,6 +211,7 @@ def _make_team_obj(
mock_team.object_permission.mcp_servers = mcp_servers or []
mock_team.object_permission.mcp_access_groups = mcp_access_groups or []
mock_team.object_permission.mcp_tool_permissions = mcp_tool_permissions or {}
mock_team.object_permission.mcp_tool_overrides = None
else:
mock_team.object_permission = None
@ -840,6 +841,7 @@ async def test_resolve_team_allowed_mcp_servers_string_tool_permissions(
mock_perm.mcp_servers = ["server-1"]
mock_perm.mcp_access_groups = []
mock_perm.mcp_tool_permissions = json.dumps({"server-2": ["tool1"]})
mock_perm.mcp_tool_overrides = None
result = await _resolve_team_allowed_mcp_servers(mock_perm)
assert result == {"server-1", "server-2"}
@ -859,6 +861,7 @@ async def test_resolve_team_allowed_mcp_servers_dict_tool_permissions(
mock_perm.mcp_servers = []
mock_perm.mcp_access_groups = []
mock_perm.mcp_tool_permissions = {"server-a": ["tool1"]}
mock_perm.mcp_tool_overrides = None
result = await _resolve_team_allowed_mcp_servers(mock_perm)
assert result == {"server-a"}
@ -889,6 +892,7 @@ async def test_resolve_team_all_proxy_sentinel_resolves_dynamically(mock_access_
team_perm.mcp_servers = [SpecialMCPServerName.all_proxy_servers.value]
team_perm.mcp_access_groups = []
team_perm.mcp_tool_permissions = {}
team_perm.mcp_tool_overrides = None
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
@ -1035,6 +1039,7 @@ def _make_team_obj_search(team_id="team-1", search_tools=None):
if search_tools is not None:
mock_team.object_permission = MagicMock(spec=LiteLLM_ObjectPermissionTable)
mock_team.object_permission.search_tools = search_tools
mock_team.object_permission.mcp_tool_overrides = None
else:
mock_team.object_permission = None
return mock_team
@ -1511,3 +1516,92 @@ async def test_prepare_object_permission_upsert_rejects_ambiguous_mcp_tool_overr
assert exc_info.value.status_code == 400
assert "'wiki'" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_upsert_converts_stored_v0_row_before_applying_update():
"""Saving over a residual v0 row converts its legacy snapshot first, so the
previously granted server keeps its current deletes while the new update
lands on the converted row."""
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
existing_row = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-id",
mcp_servers=["server-a"],
mcp_tool_permissions={"server-a": ["list_items"]},
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={"list_items": "list", "delete_item": "remove one", "search_notes": "find"}
)
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={},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert upsert.record["mcp_permission_version"] == 1
assert json.loads(upsert.record["mcp_tool_overrides"]) == {"server-a": {"allow": [], "deny": ["search_notes"]}}
assert json.loads(upsert.record["mcp_tool_permissions"]) == {}
assert json.loads(upsert.record["mcp_tool_permissions_archive"]) == {"server-a": ["list_items"]}
@pytest.mark.asyncio
async def test_upsert_rejects_503_when_inventory_unavailable():
"""A v0 row whose granted server cannot supply a tool catalog fails the save
with 503 naming the server instead of converting against nothing."""
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
existing_row = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-id",
mcp_servers=["server-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=None)
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=[]),
),
):
with pytest.raises(HTTPException) as exc_info:
await prepare_object_permission_upsert(
new_object_permission={},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert exc_info.value.status_code == 503
assert "server-a" in str(exc_info.value.detail)