diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 354f677469c..1610f9cf02f 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0ddf78c2e53..c3de9564bb5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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]: diff --git a/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py b/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py new file mode 100644 index 00000000000..fbdfd62562e --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py @@ -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 diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 13925d37bc8..73cfb344976 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -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, ) diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 99827dc56aa..bfebbdb322f 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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()) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c0071aa7c81..ce061819ffd 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 5727a77d6f8..420b09d7e6e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_tool_permission_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_tool_permission_backfill.py new file mode 100644 index 00000000000..c8355a0e925 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_tool_permission_backfill.py @@ -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() diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index c8d842bcea6..f3cfb670fda 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -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)