refactor(mcp): adopt immutable collection types for permission defaults

Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
Devin AI 2026-09-25 02:16:43 +00:00
parent dda8295af7
commit 5b1076d156
14 changed files with 261 additions and 198 deletions

View file

@ -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 = []

View file

@ -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

View file

@ -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:
"""

View file

@ -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(

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -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,
},

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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]]

View file

@ -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

View file

@ -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