mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
refactor(mcp): adopt immutable collection types for permission defaults
Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
parent
dda8295af7
commit
5b1076d156
14 changed files with 261 additions and 198 deletions
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue