From 5b1076d156b7df62593c1885994847d95213fdb7 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 02:16:43 +0000 Subject: [PATCH] refactor(mcp): adopt immutable collection types for permission defaults Co-Authored-By: bot_apk --- litellm/models/object_permission.py | 6 +- .../mcp_server/auth/user_api_key_auth_mcp.py | 88 ++++++----- .../mcp_server/mcp_server_manager.py | 68 +++++--- .../_experimental/mcp_server/operations.py | 2 +- .../mcp_server/tool_classification.py | 34 ++-- .../mcp_server/tool_permission_backfill.py | 145 ++++++++++-------- litellm/proxy/_types.py | 2 +- .../organization_endpoints.py | 2 +- .../object_permission_utils.py | 83 +++++----- .../object_permission_repository.py | 13 +- litellm/types/agents.py | 2 +- litellm/types/mcp.py | 6 +- litellm/types/object_permission.py | 4 +- .../test_tool_permission_backfill.py | 4 +- 14 files changed, 261 insertions(+), 198 deletions(-) diff --git a/litellm/models/object_permission.py b/litellm/models/object_permission.py index 32f1f7be2cb..96bc4aedd39 100644 --- a/litellm/models/object_permission.py +++ b/litellm/models/object_permission.py @@ -5,6 +5,8 @@ Canonical definition for ``litellm_objectpermissiontable``. Re-exported from ``litellm.proxy._types`` for backwards compatibility. """ +from collections.abc import Mapping, Sequence + from litellm.types.llms.base import LiteLLMPydanticObjectBase from litellm.types.mcp import MCPToolOverrideEntry @@ -16,8 +18,8 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase): mcp_servers: list[str] | None = [] mcp_access_groups: list[str] | None = [] mcp_tool_permissions: dict[str, list[str]] | None = None - mcp_tool_overrides: dict[str, MCPToolOverrideEntry] | None = None - mcp_tool_permissions_archive: dict[str, list[str]] | None = None + mcp_tool_overrides: Mapping[str, MCPToolOverrideEntry] | None = None + mcp_tool_permissions_archive: Mapping[str, Sequence[str]] | None = None mcp_permission_version: int | None = None vector_stores: list[str] | None = [] agents: list[str] | None = [] 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 1610f9cf02f..55a00403ca5 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 @@ -71,6 +71,7 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient + from litellm.types.mcp import MCPToolOverrideEntry _EMPTY_TOOLSET_GRANTS: Final[Mapping[str, Sequence[str]]] = MappingProxyType({}) @@ -113,7 +114,9 @@ def level_allowed_tools( 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(row.mcp_tool_overrides).get(server_id) or {} + overrides: Final[MCPToolOverrideEntry | Mapping[str, Sequence[str]]] = ( + global_mcp_server_manager.expand_tool_overrides(row.mcp_tool_overrides).get(server_id) or _EMPTY_OVERRIDE + ) allow: Final[frozenset[str]] = frozenset(overrides.get("allow") or ()) deny: Final[frozenset[str]] = frozenset(overrides.get("deny") or ()) if toolset_tools is not None: @@ -130,6 +133,9 @@ def level_allowed_tools( return (convention | allow) - deny +_EMPTY_OVERRIDE: Final[Mapping[str, Sequence[str]]] = MappingProxyType({}) + + def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: resolver returns a list """Widen a read-only allowlist back to the mutable list the resolver's own contract returns, preserving the ``None`` that means "no restriction".""" @@ -2245,15 +2251,19 @@ class MCPRequestHandler: toolset_perms: Final = await MCPRequestHandler._toolset_tool_permissions(row) access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( - row.mcp_access_groups or [] + sorted(row.mcp_access_groups or ()) + ) + grants_server: Final = SpecialMCPServerName.all_proxy_servers.value in ( + row.mcp_servers or () + ) or server_id in frozenset( + { + *global_mcp_server_manager.expand_permission_list(sorted(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(row.mcp_tool_overrides).keys(), + *toolset_perms.keys(), + } ) - grants_server: Final = SpecialMCPServerName.all_proxy_servers.value in (row.mcp_servers or []) or server_id in { - *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(row.mcp_tool_overrides).keys(), - *toolset_perms.keys(), - } return level_allowed_tools( row=row, server_id=server_id, @@ -2340,31 +2350,29 @@ class MCPRequestHandler: key_level: Final = await MCPRequestHandler._row_level_tools(key_obj_perm, server_id, resolved_inventory) team_level: Final = await MCPRequestHandler._row_level_tools(team_obj_perm, server_id, resolved_inventory) - if key_level is not None and team_level is not None: - allowed_tools: list[str] | None = list(key_level & team_level) - elif key_level is not None: - allowed_tools = list(key_level) - else: - allowed_tools = list(team_level) if team_level is not None else None - - allowed_tools = _as_list( + level_tools: Final = ( + sorted(key_level & team_level) + if key_level is not None and team_level is not None + else sorted(key_level) + if key_level is not None + else (sorted(team_level) if team_level is not None else None) + ) + after_end_user: Final = _as_list( await MCPRequestHandler._apply_end_user_tool_ceiling( - allowed_tools, server_id, user_api_key_auth, inventory=resolved_inventory + level_tools, server_id, user_api_key_auth, inventory=resolved_inventory ) ) - - allowed_tools = _as_list( + after_user: Final = _as_list( await MCPRequestHandler._apply_user_tool_ceiling( - allowed_tools, + after_end_user, server_id, user_api_key_auth, keyless_source=keyless_source, inventory=resolved_inventory, ) ) - - allowed_tools = await MCPRequestHandler._apply_agent_and_org_tool_ceilings( - allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source, inventory=resolved_inventory + allowed_tools: Final = await MCPRequestHandler._apply_agent_and_org_tool_ceilings( + after_user, server_id, user_api_key_auth, keyless_source=keyless_source, inventory=resolved_inventory ) if allowed_tools is not None: return allowed_tools @@ -2619,10 +2627,10 @@ class MCPRequestHandler: ) # servers referenced in tool permissions or tool overrides should also be accessible - tool_perm_servers: Final = list( + tool_perm_servers: Final = sorted( global_mcp_server_manager.expand_tool_permissions(key_object_permission.mcp_tool_permissions).keys() ) - override_servers: Final = list( + override_servers: Final = sorted( global_mcp_server_manager.expand_tool_overrides(key_object_permission.mcp_tool_overrides).keys() ) @@ -2725,12 +2733,14 @@ class MCPRequestHandler: object_permissions.mcp_access_groups or [] ) return ( - 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(object_permissions.mcp_tool_overrides).keys()) - | (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys() - | set(team_access_group_servers) + frozenset(global_mcp_server_manager.expand_permission_list(sorted(object_permissions.mcp_servers or ()))) + | frozenset(legacy_access_group_servers) + | frozenset( + global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() + ) + | frozenset(global_mcp_server_manager.expand_tool_overrides(object_permissions.mcp_tool_overrides).keys()) + | frozenset((await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()) + | frozenset(team_access_group_servers) ) @staticmethod @@ -2917,10 +2927,10 @@ class MCPRequestHandler: object_permissions.mcp_access_groups or [] ) - tool_perm_servers: Final = list( + tool_perm_servers: Final = sorted( global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() ) - override_servers: Final = list( + override_servers: Final = sorted( global_mcp_server_manager.expand_tool_overrides(object_permissions.mcp_tool_overrides).keys() ) @@ -3022,10 +3032,10 @@ class MCPRequestHandler: ) # servers referenced in tool permissions or overrides should also be accessible - tool_perm_servers: Final = list( + tool_perm_servers: Final = sorted( global_mcp_server_manager.expand_tool_permissions(object_permission.mcp_tool_permissions).keys() ) - override_servers: Final = list( + override_servers: Final = sorted( global_mcp_server_manager.expand_tool_overrides(object_permission.mcp_tool_overrides).keys() ) @@ -3143,10 +3153,10 @@ class MCPRequestHandler: access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups( object_permissions.mcp_access_groups or [] ) - tool_perm_servers: Final = list( + tool_perm_servers: Final = sorted( global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys() ) - override_servers: Final = list( + override_servers: Final = sorted( global_mcp_server_manager.expand_tool_overrides(object_permissions.mcp_tool_overrides).keys() ) toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions) @@ -3495,7 +3505,7 @@ class MCPRequestHandler: inventory if inventory is not None else await MCPRequestHandler._manager_inventory(server_id) ) agent_tools: Final = await MCPRequestHandler._row_level_tools(obj_perm, server_id, resolved_inventory) - return list(agent_tools) if agent_tools is not None else None + return sorted(agent_tools) if agent_tools is not None else None except Exception as e: if isinstance(e, UnloadableEntitlementError): raise diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c3de9564bb5..22a532fe78f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1925,7 +1925,9 @@ class MCPServerManager: # Bare tool name -> description per server_id, replaced wholesale each # time _create_prefixed_tools runs for that server. Feeds the # tool-permission convention (allow non-deletes, deny deletes). - self.discovered_tool_inventory: dict[str, dict[str, str | None]] = {} + self.discovered_tool_inventory: dict[ # mutable-ok: inventory entries are populated per server as tools are discovered + str, Mapping[str, str | None] + ] = {} # mutable-ok: inventory entries are populated per server as tools are discovered self._upstream_initialize_instructions_by_server_id: dict[str, str] = {} # Per-server monotonic timestamp of last upstream prefetch attempt (success, # empty result, or failure). Used to throttle re-probes for servers that do @@ -5386,13 +5388,15 @@ class MCPServerManager: verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) if server.server_id: - self.discovered_tool_inventory[server.server_id] = {tool.name: tool.description for tool in tools} + self.discovered_tool_inventory[server.server_id] = MappingProxyType( + {tool.name: tool.description for tool in tools} + ) return prefixed_tools def discovered_inventory(self, server_id: str) -> Mapping[str, str | None]: """The last tool catalog discovered for ``server_id``, bare names to descriptions; empty when the server is unknown or never listed.""" - return self.discovered_tool_inventory.get(server_id, {}) + return self.discovered_tool_inventory.get(server_id) or MappingProxyType({}) 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. @@ -5403,12 +5407,14 @@ class MCPServerManager: 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, - }: + if server.auth_type in frozenset( + { + MCPAuth.oauth2_token_exchange, + MCPAuth.oauth_delegate, + MCPAuth.oauth2_id_jag, + MCPAuth.true_passthrough, + } + ): return None try: await self._get_tools_from_server(server) @@ -6802,7 +6808,9 @@ class MCPServerManager: if server.available_on_public_internet or server.server_id in public_ids ] - def expand_permission_list(self, identifiers: list[str]) -> list[str]: + def expand_permission_list( + self, identifiers: Sequence[str] + ) -> list[str]: # mutable-ok: callers concatenate the returned list """ Expand a permission list of server_ids/names/aliases into concrete server_ids against the current region's config + DB registry union. @@ -6862,14 +6870,14 @@ class MCPServerManager: return {} result: Final[dict[str, list[str]]] = {} for key, tools in tool_permissions.items(): - for server_id in self.expand_permission_list([key]): + for server_id in self.expand_permission_list((key,)): result.setdefault(server_id, []).extend(tools or []) return result def expand_tool_overrides( self, - tool_overrides: dict[str, MCPToolOverrideEntry] | None, - ) -> dict[str, MCPToolOverrideEntry]: + tool_overrides: Mapping[str, MCPToolOverrideEntry] | None, + ) -> Mapping[str, MCPToolOverrideEntry]: """ Rewrite an ``mcp_tool_overrides`` dict keyed by id/name/alias so every key is a concrete server_id, same expansion as @@ -6877,18 +6885,30 @@ class MCPServerManager: merge their allow/deny lists (deny still wins at evaluation time). """ if not tool_overrides: - return {} - grouped: Final[dict[str, list[MCPToolOverrideEntry | None]]] = {} - for key, entry in tool_overrides.items(): - for server_id in self.expand_permission_list([key]): - grouped.setdefault(server_id, []).append(entry) - return { - server_id: { - "allow": [tool for entry in entries if entry for tool in (entry.get("allow") or [])], - "deny": [tool for entry in entries if entry for tool in (entry.get("deny") or [])], + return MappingProxyType({}) + expanded: Final[tuple[tuple[str, MCPToolOverrideEntry], ...]] = tuple( + (server_id, entry) + for key, entry in tool_overrides.items() + for server_id in self.expand_permission_list((key,)) + if entry is not None + ) + return MappingProxyType( + { + server_id: MCPToolOverrideEntry( + allow=sorted( + frozenset( + tool for sid, entry in expanded if sid == server_id for tool in (entry.get("allow") or ()) + ) + ), + deny=sorted( + frozenset( + tool for sid, entry in expanded if sid == server_id for tool in (entry.get("deny") or ()) + ) + ), + ) + for server_id in frozenset(sid for sid, _ in expanded) } - for server_id, entries in grouped.items() - } + ) def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: """ diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 5fa2871864b..d0a4faa4e0e 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1419,7 +1419,7 @@ async def filter_tools_by_key_team_permissions( the prefix before comparing. """ server: Final = global_mcp_server_manager.get_mcp_server_by_id(server_id) - inventory: Final = {strip_known_server_prefix(t.name, server): t.description for t in tools} + inventory: Final = types.MappingProxyType({strip_known_server_prefix(t.name, server): t.description for t in tools}) # Filter by key/team tool-level permissions allowed_tool_names: Final = await MCPRequestHandler.get_allowed_tools_for_server( diff --git a/litellm/proxy/_experimental/mcp_server/tool_classification.py b/litellm/proxy/_experimental/mcp_server/tool_classification.py index 0db917ebd64..cbbc76cdc95 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_classification.py +++ b/litellm/proxy/_experimental/mcp_server/tool_classification.py @@ -6,9 +6,10 @@ single exported function is all callers need. """ import re -from typing import Final, Literal +from collections.abc import Iterable +from typing import Final, Literal, TypeAlias -ToolOperation = Literal["read", "create", "update", "delete", "unknown"] +ToolOperation: TypeAlias = Literal["read", "create", "update", "delete", "unknown"] _READ_TOKENS: Final = frozenset( { @@ -91,29 +92,28 @@ _SPLIT_RE: Final = re.compile(r"[_\-./\s]+") _CAMEL_BOUNDARY_RE: Final = re.compile(r"(?<=[a-z0-9])(?=[A-Z])|(?<=[A-Z])(?=[A-Z][a-z])") -def _name_tokens(name: str) -> list[str]: - pieces: Final[list[str]] = [] - for chunk in _SPLIT_RE.split(name): - pieces.extend(_CAMEL_BOUNDARY_RE.split(chunk)) - return [token.lower() for token in pieces if token] +def _name_tokens(name: str) -> tuple[str, ...]: + return tuple(token.lower() for chunk in _SPLIT_RE.split(name) for token in _CAMEL_BOUNDARY_RE.split(chunk) if token) -def _description_tokens(description: str) -> list[str]: - return [token.lower() for token in re.split(r"[^\w]+", description) if token] +def _description_tokens(description: str) -> tuple[str, ...]: + return tuple(token.lower() for token in re.split(r"[^\w]+", description) if token) def _token_variants(token: str) -> frozenset[str]: - variants: Final = { - token, - token.removesuffix("es"), - token.removesuffix("s"), - token.removesuffix("ed"), - token.removesuffix("ing"), - } + variants: Final = frozenset( + ( + token, + token.removesuffix("es"), + token.removesuffix("s"), + token.removesuffix("ed"), + token.removesuffix("ing"), + ) + ) return frozenset(variant for variant in variants if variant) -def _classify_tokens(tokens: list[str]) -> ToolOperation: +def _classify_tokens(tokens: Iterable[str]) -> ToolOperation: token_set: Final = frozenset(variant for token in tokens for variant in _token_variants(token)) if token_set & _READ_TOKENS: return "read" diff --git a/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py b/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py index a7f5a425a9f..01e00db21de 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/tool_permission_backfill.py @@ -14,9 +14,10 @@ import json from collections.abc import Mapping, Sequence from dataclasses import dataclass from types import MappingProxyType -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, TypeAlias from pydantic import TypeAdapter +from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_logger from litellm.models.object_permission import LiteLLM_ObjectPermissionTable @@ -27,14 +28,15 @@ from litellm.types.mcp import MCPToolOverrideEntry if TYPE_CHECKING: from prisma import models as prisma_models + from prisma import types as prisma_types from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.proxy.utils import PrismaClient - RawRow = LiteLLM_ObjectPermissionTable | prisma_models.LiteLLM_ObjectPermissionTable + RawRow: TypeAlias = LiteLLM_ObjectPermissionTable | prisma_models.LiteLLM_ObjectPermissionTable -ToolInventory = Mapping[str, str | None] -Inventories = Mapping[str, ToolInventory | None] +ToolInventory: TypeAlias = Mapping[str, str | None] +Inventories: TypeAlias = Mapping[str, ToolInventory | None] _ROW_DUMP_ADAPTER: Final = TypeAdapter(dict[str, object]) @@ -52,7 +54,14 @@ class Unavailable: server_ids: frozenset[str] -ConversionResult = ConvertedRow | Unavailable +class _ConvertedRowRecord(TypedDict, total=False): + mcp_tool_overrides: ReadOnly[Mapping[str, MCPToolOverrideEntry]] + mcp_tool_permissions: ReadOnly[Mapping[str, Sequence[str]]] + mcp_tool_permissions_archive: ReadOnly[Mapping[str, Sequence[str]]] + mcp_permission_version: ReadOnly[int] + + +ConversionResult: TypeAlias = ConvertedRow | Unavailable def _server_override_entry( @@ -79,7 +88,8 @@ def _server_override_entry( ) if not allow and not deny: return None - return {"allow": sorted(allow), "deny": sorted(deny)} + entry: Final[MCPToolOverrideEntry] = {"allow": sorted(allow), "deny": sorted(deny)} + return entry @dataclass(frozen=True) @@ -108,7 +118,7 @@ def convert_row( if missing: return Unavailable(server_ids=missing) - legacy: Final = row.mcp_tool_permissions or {} + 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} ) @@ -116,7 +126,7 @@ def convert_row( { server_id: entry for server_id, inventory in inventories.items() - if legacy.get(server_id) != [] + if server_id not in legacy or legacy[server_id] for entry in (_server_override_entry(legacy.get(server_id), inventory or MappingProxyType({})),) if entry is not None } @@ -137,47 +147,52 @@ async def resolve_granted_server_ids( 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) + direct: Final = frozenset( + server_id + for identifier in manager.expand_permission_list(sorted(row.mcp_servers or ())) + for server_id in ( + manager.get_registry().keys() + if identifier == SpecialMCPServerName.all_proxy_servers.value + else (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 [] + sorted(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()) + | frozenset(access_group_servers) + | frozenset(manager.expand_tool_permissions(row.mcp_tool_permissions).keys()) + | frozenset(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]: + cache: dict[str, ToolInventory | None] | None = None, # mutable-ok: caller-owned fetch cache updated in place +) -> Mapping[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 + return MappingProxyType( + {server_id: await manager.fetch_unfiltered_inventory(server_id) for server_id in server_ids} ) - for server_id in server_ids - } + return MappingProxyType( + { + server_id: ( + cache[server_id] + if server_id in cache + else cache.setdefault(server_id, await manager.fetch_unfiltered_inventory(server_id)) + ) + 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), +def converted_row_record(conversion: ConvertedRow) -> Mapping[str, object]: + record: Final[_ConvertedRowRecord] = { + "mcp_tool_overrides": {**conversion.mcp_tool_overrides}, + "mcp_tool_permissions": {**conversion.mcp_tool_permissions}, + "mcp_tool_permissions_archive": {**conversion.mcp_tool_permissions_archive}, "mcp_permission_version": 1, } return record @@ -190,10 +205,12 @@ def _to_model(row: "RawRow") -> LiteLLM_ObjectPermissionTable: "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() - } + normalized: Final = MappingProxyType( + { + 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) @@ -201,13 +218,15 @@ async def _converted_page( rows: Sequence["RawRow"], prisma_client: "PrismaClient", manager: "MCPServerManager", - inventory_cache: dict[str, ToolInventory | None], + inventory_cache: dict[str, ToolInventory | None], # mutable-ok: caller-owned fetch cache updated in place ) -> 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) - ] + models: Final[Sequence[LiteLLM_ObjectPermissionTable]] = tuple(_to_model(row) for row in rows) + outcomes: Final = tuple( + [ + 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 @@ -237,7 +256,7 @@ async def _convert_one_row( row: LiteLLM_ObjectPermissionTable, prisma_client: "PrismaClient", manager: "MCPServerManager", - inventory_cache: dict[str, ToolInventory | None], + inventory_cache: dict[str, ToolInventory | None], # mutable-ok: caller-owned fetch cache updated in place ) -> 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.""" @@ -253,19 +272,21 @@ async def _convert_one_row( ("mcp_toolsets", raw_row.mcp_toolsets), ("mcp_tool_permissions", raw_row.mcp_tool_permissions), ) - updated: Final = await ObjectPermissionRepository(prisma_client).table.update_many( - where={ - "object_permission_id": row.object_permission_id, - "mcp_permission_version": {"in": [0, None]}, - **{field: {"equals": value} for field, value in stored_fields if value is not None}, - }, - 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, - }, + equals_filters: Final = MappingProxyType( + {field: MappingProxyType({"equals": value}) for field, value in stored_fields if value is not None} ) + where: Final[prisma_types.LiteLLM_ObjectPermissionTableWhereInput] = { + "object_permission_id": row.object_permission_id, + "mcp_permission_version": 0, + **equals_filters, + } + data: Final[prisma_types.LiteLLM_ObjectPermissionTableUpdateManyMutationInput] = { + "mcp_tool_overrides": json.dumps({**conversion.mcp_tool_overrides}), + "mcp_tool_permissions": json.dumps({**conversion.mcp_tool_permissions}), + "mcp_tool_permissions_archive": json.dumps({**conversion.mcp_tool_permissions_archive}), + "mcp_permission_version": 1, + } + updated: Final = await ObjectPermissionRepository(prisma_client).table.update_many(where=where, data=data) return "cas_missed" if updated == 0 else "converted" @@ -294,10 +315,10 @@ async def run_mcp_tool_permission_backfill( 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"}, + where={"mcp_permission_version": 0}, # mutable-ok: prisma where kwarg + order={"object_permission_id": "asc"}, # mutable-ok: prisma order kwarg take=batch_size, - cursor={"object_permission_id": cursor} if cursor is not None else None, + cursor={"object_permission_id": cursor} if cursor is not None else None, # mutable-ok: prisma cursor kwarg skip=1 if cursor is not None else None, ) if not rows: @@ -313,6 +334,6 @@ async def run_mcp_tool_permission_backfill( 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()}, + MappingProxyType({row_id: tuple(sorted(server_ids)) for row_id, server_ids in report.unavailable.items()}), ) return report diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 906eeb734be..afc07d3a3c0 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1195,7 +1195,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase): mcp_servers: list[str] | None = None mcp_access_groups: list[str] | None = None mcp_tool_permissions: dict[str, list[str]] | None = None - mcp_tool_overrides: dict[str, MCPToolOverrideEntry] | None = None + mcp_tool_overrides: Mapping[str, MCPToolOverrideEntry] | None = None mcp_toolsets: list[str] | None = None blocked_tools: list[str] | None = None vector_stores: list[str] | None = None diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 73cfb344976..d1fbd765035 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -656,7 +656,7 @@ async def _set_object_permission( prisma_client=prisma_client, ) created_object_permission: Final = await _table(ObjectPermissionRepository(prisma_client)).create( - data={ + data={ # mutable-ok: prisma create payload **data.object_permission.model_dump(exclude_none=True), "mcp_permission_version": 1, }, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index bfebbdb322f..4f8f403b6ed 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -82,13 +82,13 @@ class ObjectPermissionUpsert: async def _convert_unversioned_object_permission( existing_object_permission: "LiteLLM_ObjectPermissionTable | None", -) -> dict[str, object]: +) -> Mapping[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 {} + return MappingProxyType({}) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) @@ -102,14 +102,14 @@ async def _convert_unversioned_object_permission( granted: Final = await resolve_granted_server_ids(existing_object_permission, global_mcp_server_manager) if not granted: - return {} + return MappingProxyType({}) 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={ + detail={ # mutable-ok: HTTPException detail must be a JSON-serializable dict "error": "Could not discover tools for servers granted by this permission; retry once the MCP servers are reachable.", "server_ids": sorted(conversion.server_ids), }, @@ -141,15 +141,17 @@ 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_raw: Final[dict[str, object]] = ( + existing_fields_raw: Final[Mapping[str, object]] = ( existing_object_permission.model_dump(exclude_unset=True, exclude_none=True) if existing_object_permission is not None - else {} + else MappingProxyType({}) + ) + existing_fields: Final = MappingProxyType( + { + **existing_fields_raw, + **await _convert_unversioned_object_permission(existing_object_permission), + } ) - 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"), @@ -160,26 +162,22 @@ async def prepare_object_permission_upsert( existing_mcp_tool_overrides=existing_fields.get("mcp_tool_overrides"), prisma_client=prisma_client, ) - merged: Final[dict[str, object]] = { - **existing_fields, - **new_object_permission, - "object_permission_id": object_permission_id, - "mcp_permission_version": 1, - } - record: Final[dict[str, object]] = { - **merged, - **( - {"mcp_tool_permissions": safe_dumps(merged["mcp_tool_permissions"])} - if "mcp_tool_permissions" in merged - 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 {} - ), - } + merged: Final = MappingProxyType( + { + **existing_fields, + **new_object_permission, + "object_permission_id": object_permission_id, + "mcp_permission_version": 1, + } + ) + json_fields: Final = MappingProxyType( + { + field: safe_dumps(merged[field]) + for field in ("mcp_tool_permissions", "mcp_tool_overrides", "mcp_tool_permissions_archive") + if field in merged + } + ) + record: Final[dict[str, object]] = {**merged, **json_fields} # mutable-ok: prisma upsert payload return ObjectPermissionUpsert(object_permission_id=object_permission_id, record=record) @@ -483,7 +481,7 @@ async def reject_ambiguous_mcp_tool_permission_keys( def _drop_stale_object_permission_mcp_servers( object_permission: ObjectPermissionDict, - identifier_to_server_ids: dict[str, set[str]], + identifier_to_server_ids: Mapping[str, AbstractSet[str]], ) -> None: mcp_servers: Final = object_permission.get("mcp_servers") if not isinstance(mcp_servers, list): @@ -502,7 +500,7 @@ def _drop_stale_object_permission_mcp_servers( def _drop_stale_object_permission_mcp_tool_permissions( object_permission: ObjectPermissionDict, - identifier_to_server_ids: dict[str, set[str]], + identifier_to_server_ids: Mapping[str, AbstractSet[str]], ) -> None: mcp_tool_permissions: Final = object_permission.get("mcp_tool_permissions") if not isinstance(mcp_tool_permissions, dict): @@ -517,13 +515,15 @@ def _drop_stale_object_permission_mcp_tool_permissions( def _drop_stale_object_permission_mcp_tool_overrides( object_permission: ObjectPermissionDict, - identifier_to_server_ids: dict[str, set[str]], + identifier_to_server_ids: Mapping[str, AbstractSet[str]], ) -> None: mcp_tool_overrides: Final = object_permission.get("mcp_tool_overrides") if not isinstance(mcp_tool_overrides, dict): return - object_permission["mcp_tool_overrides"] = { + object_permission[ # rebind-ok: in-place prune of stale keys is this helper's contract + "mcp_tool_overrides" + ] = { # mutable-ok: dict stored into the permission field identifier: entry for identifier, entry in mcp_tool_overrides.items() if identifier_to_server_ids.get(identifier) @@ -579,17 +579,20 @@ async def _resolve_team_allowed_mcp_servers( access_group_servers: Final[list[str]] = await MCPRequestHandler._get_mcp_servers_from_access_groups( team_object_permission.mcp_access_groups or [] ) - raw_tool_perms = team_object_permission.mcp_tool_permissions or {} - if isinstance(raw_tool_perms, str): - raw_tool_perms = json.loads(raw_tool_perms) + stored_tool_perms: Final[Mapping[str, Sequence[str]]] = ( + team_object_permission.mcp_tool_permissions or MappingProxyType({}) + ) + raw_tool_perms: Final[Mapping[str, Sequence[str]]] = ( + json.loads(stored_tool_perms) if isinstance(stored_tool_perms, str) else stored_tool_perms + ) 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) + tool_perm_servers: Final[tuple[str, ...]] = tuple(raw_tool_perms.keys()) + tuple(raw_tool_overrides.keys()) + raw_servers: Final = frozenset((*direct_servers, *access_group_servers, *tool_perm_servers)) resolved_servers: Final = await _resolve_mcp_server_identifiers_to_ids( identifiers=raw_servers, prisma_client=prisma_client, ) - unresolved_servers: Final = {server_id for server_id in raw_servers if not resolved_servers.get(server_id)} + unresolved_servers: Final = frozenset(server_id for server_id in raw_servers if not resolved_servers.get(server_id)) return _flatten_resolved_mcp_server_ids(resolved_servers) | unresolved_servers diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 173945b88d4..c0f43bfe27b 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -2,6 +2,7 @@ ObjectPermission repository for database operations on LiteLLM_ObjectPermissionTable. """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final from litellm.models.object_permission import LiteLLM_ObjectPermissionTable @@ -34,7 +35,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): mcp_servers: list[str] | None = None, mcp_access_groups: list[str] | None = None, mcp_tool_permissions: dict[str, list[str]] | None = None, - mcp_tool_overrides: dict[str, MCPToolOverrideEntry] | None = None, + mcp_tool_overrides: Mapping[str, MCPToolOverrideEntry] | None = None, vector_stores: list[str] | None = None, agents: list[str] | None = None, agent_access_groups: list[str] | None = None, @@ -45,7 +46,9 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable: """Create a new object permission record.""" - data: Final[dict[str, Any]] = {"mcp_permission_version": 1} + data: Final[dict[str, Any]] = { # mutable-ok: prisma payload fields assigned conditionally + "mcp_permission_version": 1 + } if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: @@ -79,7 +82,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): mcp_servers: list[str] | None = None, mcp_access_groups: list[str] | None = None, mcp_tool_permissions: dict[str, list[str]] | None = None, - mcp_tool_overrides: dict[str, MCPToolOverrideEntry] | None = None, + mcp_tool_overrides: Mapping[str, MCPToolOverrideEntry] | None = None, vector_stores: list[str] | None = None, agents: list[str] | None = None, agent_access_groups: list[str] | None = None, @@ -90,7 +93,9 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): skills: list[str] | None = None, ) -> LiteLLM_ObjectPermissionTable | None: """Update an object permission record.""" - data: Final[dict[str, Any]] = {"mcp_permission_version": 1} + data: Final[dict[str, Any]] = { # mutable-ok: prisma payload fields assigned conditionally + "mcp_permission_version": 1 + } if mcp_servers is not None: data["mcp_servers"] = mcp_servers if mcp_access_groups is not None: diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 004cd03ecf2..7206d10ca6e 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -175,7 +175,7 @@ class AgentObjectPermission(TypedDict, total=False): mcp_access_groups: list[str] | None mcp_toolsets: ReadOnly[Sequence[str] | None] mcp_tool_permissions: dict[str, list[str]] | None - mcp_tool_overrides: dict[str, MCPToolOverrideEntry] | None + mcp_tool_overrides: ReadOnly[Mapping[str, MCPToolOverrideEntry] | None] models: list[str] | None agents: list[str] | None diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 2a2212274b8..27f79d7f12c 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -2,7 +2,7 @@ from __future__ import annotations import enum import re -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable, Mapping, Sequence from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal from urllib.parse import urlsplit @@ -498,5 +498,5 @@ class MCPToolOverrideEntry(TypedDict, total=False): ``mcp_tool_overrides``: ``allow`` re-arms names the convention denies, ``deny`` disables names the convention or an allowlist would permit.""" - allow: ReadOnly[list[str]] - deny: ReadOnly[list[str]] + allow: ReadOnly[Sequence[str]] + deny: ReadOnly[Sequence[str]] diff --git a/litellm/types/object_permission.py b/litellm/types/object_permission.py index ff246843c43..96487546fce 100644 --- a/litellm/types/object_permission.py +++ b/litellm/types/object_permission.py @@ -8,6 +8,8 @@ can adopt the type without violating the SDK-must-not-import-from-proxy layering rule. """ +from collections.abc import Mapping + from typing_extensions import ReadOnly, TypedDict from litellm.types.mcp import MCPToolOverrideEntry @@ -17,7 +19,7 @@ class ObjectPermissionDict(TypedDict, total=False): mcp_servers: list[str] | None mcp_access_groups: list[str] | None mcp_tool_permissions: dict[str, list[str]] | None - mcp_tool_overrides: dict[str, MCPToolOverrideEntry] | None + mcp_tool_overrides: Mapping[str, MCPToolOverrideEntry] | None # writable-ok: stale keys pruned in place mcp_toolsets: list[str] | None blocked_tools: list[str] | None vector_stores: list[str] | None 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 index 776c5ed3929..0013d073de0 100644 --- 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 @@ -140,7 +140,7 @@ async def test_runner_converts_row_with_cas_update(): 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]} + assert where["mcp_permission_version"] == 0 @pytest.mark.asyncio @@ -156,7 +156,7 @@ async def test_runner_omits_cas_equals_filter_for_null_fields(): assert report.converted == {"perm-1"} where = prisma.db.litellm_objectpermissiontable.update_many.await_args.kwargs["where"] assert "mcp_tool_permissions" not in where - assert where["mcp_permission_version"] == {"in": [0, None]} + assert where["mcp_permission_version"] == 0 @pytest.mark.asyncio