fix(mcp): evaluate tool permissions as convention over discovered inventory

Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
Devin AI 2026-09-25 00:27:39 +00:00
parent 4aa3ff47fe
commit 49b7bd8da3
23 changed files with 1645 additions and 465 deletions

View file

@ -0,0 +1,4 @@
-- AlterTable
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "mcp_tool_overrides" JSONB;
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "mcp_permission_version" INTEGER NOT NULL DEFAULT 0;
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "mcp_tool_permissions_archive" JSONB;

View file

@ -282,6 +282,9 @@ model LiteLLM_ObjectPermissionTable {
mcp_servers String[] @default([])
mcp_access_groups String[] @default([])
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
mcp_tool_overrides Json? // Per-server tool allow/deny overrides. Format: {"server_id": {"allow": ["t1"], "deny": ["t2"]}}
mcp_permission_version Int @default(0) // 0 = pre-overrides row (legacy evaluation), 1 = written after overrides shipped
mcp_tool_permissions_archive Json? // Backfill's snapshot of the pre-conversion mcp_tool_permissions, recovery only
vector_stores String[] @default([])
agents String[] @default([])
agent_access_groups String[] @default([])

View file

@ -6,6 +6,7 @@ Canonical definition for ``litellm_objectpermissiontable``. Re-exported from
"""
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.mcp import MCPToolOverrideEntry
class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
@ -15,6 +16,9 @@ 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_permission_version: int | None = None
vector_stores: list[str] | None = []
agents: list[str] | None = []
agent_access_groups: list[str] | None = []

View file

@ -3,7 +3,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, cast
from typing import TYPE_CHECKING, Final, Literal
from fastapi import HTTPException
from starlette.datastructures import Headers
@ -76,6 +76,60 @@ if TYPE_CHECKING:
_EMPTY_TOOLSET_GRANTS: Final[Mapping[str, Sequence[str]]] = MappingProxyType({})
def level_allowed_tools(
*,
row: LiteLLM_ObjectPermissionTable | None,
server_id: str,
grants_server: bool,
toolset_tools: Sequence[str] | None,
inventory: Mapping[str, str | None],
) -> frozenset[str] | None:
"""One permission level's effective tool allowlist on ``server_id``.
``None`` means this level places no restriction (it inherits whatever the
other levels decide); a frozenset — including the empty one — is a closed
answer from this level and intersects with the rest.
1. A legacy ``mcp_tool_permissions`` entry for the server stays a closed
allowlist (``[]`` denies all): allowed = legacy ∪ toolset tools.
2. An unconverted row (``mcp_permission_version`` falsy) keeps pre-overrides
behavior: unrestricted unless a toolset names the server.
3. A converted row that does not grant the server places no restriction.
4. A converted row granting the server applies the convention: every
inventory tool that is not delete-classified and not explicitly denied,
plus explicit allows and toolset tools. Deny beats allow on overlap.
"""
if row is None:
return None
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
toolset: Final[frozenset[str]] = frozenset(toolset_tools or ())
legacy: Final = global_mcp_server_manager.expand_tool_permissions(row.mcp_tool_permissions).get(server_id)
if legacy is not None:
return frozenset(legacy) | toolset
if not getattr(row, "mcp_permission_version", None):
return frozenset(toolset) if toolset_tools is not None else None
if not grants_server:
return None
overrides: Final = (
global_mcp_server_manager.expand_tool_overrides(getattr(row, "mcp_tool_overrides", None)).get(server_id) or {}
)
allow: Final[frozenset[str]] = frozenset(overrides.get("allow") or ())
deny: Final[frozenset[str]] = frozenset(overrides.get("deny") or ())
from litellm.proxy._experimental.mcp_server.tool_classification import (
classify_tool_op,
)
convention: Final[frozenset[str]] = frozenset(
tool_name
for tool_name, description in inventory.items()
if tool_name not in deny and classify_tool_op(tool_name, description) != "delete"
)
return ((convention | allow) - deny) | toolset
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"."""
@ -2121,26 +2175,6 @@ class MCPRequestHandler:
)
return resolved
@staticmethod
async def _toolset_tools_for_server(
object_permission: LiteLLM_ObjectPermissionTable | None,
server_id: str,
) -> Sequence[str] | None:
"""Tool names this row's toolsets grant on ``server_id``, ``None`` when its toolsets place
no restriction on that server (it declares no toolsets, or none of them name it)."""
return (await MCPRequestHandler._toolset_tool_permissions(object_permission)).get(server_id)
@staticmethod
def _union_tool_grants(
direct: Sequence[str] | None,
via_toolsets: Sequence[str] | None,
) -> Sequence[str] | None:
"""Union of one level's direct tool grants and its toolset-granted tools on one server,
``None`` when neither source restricts (allow-all from this level)."""
if direct is None and via_toolsets is None:
return None
return tuple({*(direct or ()), *(via_toolsets or ())})
@staticmethod
async def _key_object_permission_hydrated(
user_api_key_auth: UserAPIKeyAuth,
@ -2193,12 +2227,78 @@ class MCPRequestHandler:
verbose_logger.warning("Failed to check declared MCP toolsets, org ceiling unchanged: %s", e)
return False
@staticmethod
async def _row_level_tools(
row: LiteLLM_ObjectPermissionTable | None,
server_id: str,
inventory: Mapping[str, str | None],
) -> frozenset[str] | None:
"""Level evaluation for one stored permission row: resolves the row's
toolset grants and whether it grants ``server_id``, then defers to the
pure ``level_allowed_tools``. A declared toolset that resolves to no
grants raises ``UnloadableEntitlementError`` as before."""
if row is None:
return None
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
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 []
)
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(getattr(row, "mcp_tool_overrides", None)).keys(),
*toolset_perms.keys(),
}
return level_allowed_tools(
row=row,
server_id=server_id,
grants_server=grants_server,
toolset_tools=toolset_perms.get(server_id),
inventory=inventory,
)
@staticmethod
async def _any_ceiling_row_present(user_api_key_auth: UserAPIKeyAuth, keyless_source: bool) -> bool:
"""Whether any ceiling level (end user, user, agent, org) holds a
permission row at all. Only consulted when every level answered
``None``; lookup faults propagate so the caller's own fault arms
decide (fail-open for key auth, deny for a keyless source). The user
level is skipped for a keyless source exactly as its ceiling arm is:
each admitted source's own grants must not cap the others."""
from litellm.proxy.proxy_server import prisma_client
if (
user_api_key_auth.end_user_id
and prisma_client is not None
and await MCPRequestHandler._get_end_user_object_permission(user_api_key_auth, prisma_client) is not None
):
return True
if not keyless_source and await MCPRequestHandler._get_user_object_permission(user_api_key_auth) is not None:
return True
if (
user_api_key_auth.agent_id
and await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) is not None
):
return True
if (
user_api_key_auth.org_id
and await MCPRequestHandler._get_org_object_permission(user_api_key_auth) is not None
):
return True
return False
@staticmethod
async def get_allowed_tools_for_server(
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
inventory: Mapping[str, str | None] | None = None,
) -> list[str] | None:
"""
Get list of allowed tool names for a specific server based on key/team permissions.
@ -2207,6 +2307,10 @@ class MCPRequestHandler:
Args:
server_id: Server ID to check permissions for
user_api_key_auth: User auth
inventory: Bare tool name -> description catalog for the server.
Defaults to the manager's last discovered catalog; the listing
path passes the tools it just fetched. The convention only
ever classifies names that appear in it.
Returns:
List[str] if restrictions exist, None if no restrictions (allow all)
@ -2221,76 +2325,68 @@ class MCPRequestHandler:
if _is_mcp_admitted_user_subject(user_api_key_auth):
return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth)
# Get key and team object permissions (already loaded in main auth flow)
key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(user_api_key_auth)
# Extract tool permissions for this server. Dict keys may be
# server_ids OR names/aliases; normalize to server_id-keyed form
# before lookup so a name-based key does not silently drop its
# tool restrictions when server_id is the resolved uuid.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
key_direct_tools: Final = (
global_mcp_server_manager.expand_tool_permissions(key_obj_perm.mcp_tool_permissions).get(server_id)
if key_obj_perm
else None
resolved_inventory: Final[Mapping[str, str | None]] = (
global_mcp_server_manager.discovered_inventory(server_id) if inventory is None else inventory
)
# Tools granted through the key's toolsets restrict this server exactly
# as direct tool permissions do; union with any direct grants so the
# tool-level check sees the key's full effective tool scope
key_toolset_ids: Final = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
key_toolset_tools: Final = (
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
server_id
)
if key_toolset_ids
else None
)
# Get key and team object permissions (already loaded in main auth flow)
key_obj_perm: Final = await MCPRequestHandler._key_object_permission_hydrated(user_api_key_auth)
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(user_api_key_auth)
key_tools: Final = (
list(set(key_direct_tools or []) | set(key_toolset_tools or []))
if key_direct_tools is not None or key_toolset_tools is not None
else None
)
team_direct_tools: Final = (
global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id)
if team_obj_perm
else None
)
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)
# Tools granted through the team's toolsets restrict this server exactly
# as the team's direct tool permissions do, mirroring the key path above
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
# Apply same inheritance logic as get_allowed_mcp_servers
if team_tools:
if key_tools:
# Both have restrictions → intersection
allowed_tools = list(set(team_tools) & set(key_tools))
else:
# Only team has restrictions → inherit from team
allowed_tools = team_tools
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:
# No team restrictions → use key restrictions
allowed_tools = cast(list[str], key_tools)
allowed_tools = list(team_level) if team_level is not None else None
allowed_tools = _as_list(
await MCPRequestHandler._apply_end_user_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
await MCPRequestHandler._apply_end_user_tool_ceiling(
allowed_tools, server_id, user_api_key_auth, inventory=resolved_inventory
)
)
allowed_tools = _as_list(
await MCPRequestHandler._apply_user_tool_ceiling(
allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source
allowed_tools,
server_id,
user_api_key_auth,
keyless_source=keyless_source,
inventory=resolved_inventory,
)
)
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source
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
)
if allowed_tools is not None:
return allowed_tools
# Every level answered "no restriction". With a row present at any
# level that is the legacy unrestricted answer; with no row
# anywhere the convention defaults to non-delete inventory tools.
# A keyless source keeps the legacy answer outright: its own grants
# ARE the source, and the subject-level union narrows per source —
# applying the convention here would deny tools its own row grants.
if key_obj_perm is not None or team_obj_perm is not None or keyless_source:
return None
if await MCPRequestHandler._any_ceiling_row_present(user_api_key_auth, keyless_source):
return None
from litellm.proxy._experimental.mcp_server.tool_classification import (
classify_tool_op,
)
return sorted(
tool_name
for tool_name, description in resolved_inventory.items()
if classify_tool_op(tool_name, description) != "delete"
)
except Exception as e:
@ -2314,6 +2410,8 @@ class MCPRequestHandler:
server_id: str,
user_api_key_auth: UserAPIKeyAuth,
keyless_source: bool = False,
*,
inventory: Mapping[str, str | None] | None = None,
) -> list[str] | None:
"""Narrow a key/team tool allowlist by the agent's tool permissions and the caller's org tool
ceiling. Each level only intersects; None at a level means no restriction from it.
@ -2322,8 +2420,8 @@ class MCPRequestHandler:
fail-open (skip the org step, keep the key/team/agent restrictions; letting the raise escape
would collapse them to allow-all, WIDER than before the fault), while a keyless source re-raises
so the outer handler denies that one source (its only org bound is this ceiling)."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
resolved_inventory: Final = (
inventory if inventory is not None else (await MCPRequestHandler._manager_inventory(server_id))
)
if user_api_key_auth.agent_id:
@ -2333,6 +2431,7 @@ class MCPRequestHandler:
server_id=server_id,
user_api_key_auth=user_api_key_auth,
agent_object_permission=agent_obj_perm,
inventory=resolved_inventory,
)
if agent_tools is not None:
allowed_tools = (
@ -2355,13 +2454,7 @@ class MCPRequestHandler:
e,
)
return allowed_tools
org_direct_tools: Final = (
global_mcp_server_manager.expand_tool_permissions(org_obj_perm.mcp_tool_permissions).get(server_id)
if org_obj_perm and org_obj_perm.mcp_tool_permissions
else None
)
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(org_obj_perm, server_id)
org_tools: Final = MCPRequestHandler._union_tool_grants(org_direct_tools, org_toolset_tools)
org_tools: Final = await MCPRequestHandler._row_level_tools(org_obj_perm, server_id, resolved_inventory)
if org_tools is not None:
allowed_tools = (
list(set(allowed_tools) & set(org_tools)) if allowed_tools is not None else list(org_tools)
@ -2369,6 +2462,15 @@ class MCPRequestHandler:
return allowed_tools
@staticmethod
async def _manager_inventory(server_id: str) -> Mapping[str, str | None]:
"""The manager's last discovered tool catalog for ``server_id``."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
return global_mcp_server_manager.discovered_inventory(server_id)
@staticmethod
def tool_is_granted(bare_tool_name: str, allowed_tool_names: list[str] | None) -> bool:
"""Whether key/team tool permissions reach ``bare_tool_name`` on one server.
@ -2516,10 +2618,15 @@ class MCPRequestHandler:
key_object_permission.mcp_access_groups or []
)
# servers referenced in tool permissions should also be accessible
# servers referenced in tool permissions or tool overrides should also be accessible
tool_perm_servers: Final = list(
global_mcp_server_manager.expand_tool_permissions(key_object_permission.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(key_object_permission, "mcp_tool_overrides", None)
).keys()
)
# servers referenced by the key's toolset grants are part of the key's
# scope on every path (list, call, REST), subject to the same team/org
@ -2532,7 +2639,9 @@ class MCPRequestHandler:
)
# Combine all lists
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers + toolset_servers
all_servers: Final = (
direct_mcp_servers + access_group_servers + tool_perm_servers + override_servers + toolset_servers
)
return list(set(all_servers))
except Exception as e:
verbose_logger.warning("Failed to get allowed MCP servers for key: %s", e)
@ -2621,6 +2730,11 @@ class MCPRequestHandler:
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
| set(legacy_access_group_servers)
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
| set(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permissions, "mcp_tool_overrides", None)
).keys()
)
| (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()
| set(team_access_group_servers)
)
@ -2812,13 +2926,18 @@ class MCPRequestHandler:
tool_perm_servers: Final = list(
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permissions, "mcp_tool_overrides", None)
).keys()
)
# servers referenced by the org's toolset grants are part of the org ceiling,
# exactly as servers referenced by its inline tool permissions are
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
all_servers: Final = tuple(
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *override_servers, *toolset_grants}
)
return list(set(all_servers))
except Exception as e:
@ -2910,13 +3029,18 @@ class MCPRequestHandler:
object_permission.mcp_access_groups or []
)
# servers referenced in tool permissions should also be accessible
# servers referenced in tool permissions or overrides should also be accessible
tool_perm_servers: Final = list(
global_mcp_server_manager.expand_tool_permissions(object_permission.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permission, "mcp_tool_overrides", None)
).keys()
)
# Combine all lists
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers + override_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning("Failed to get allowed MCP servers for end_user: %s", e)
@ -3032,8 +3156,15 @@ class MCPRequestHandler:
tool_perm_servers: Final = list(
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
override_servers: Final = list(
global_mcp_server_manager.expand_tool_overrides(
getattr(object_permissions, "mcp_tool_overrides", None)
).keys()
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
return tuple(
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *override_servers, *toolset_grants}
)
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
return None
@ -3135,6 +3266,7 @@ class MCPRequestHandler:
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
inventory: Mapping[str, str | None] | None = None,
) -> Sequence[str] | None:
"""Narrow a key/team tool allowlist by the internal user's own tool entitlement.
@ -3143,10 +3275,6 @@ class MCPRequestHandler:
restriction. Returns ``[]`` (deny every tool on this server) when the entitlement cannot be
resolved, because the caller's own except-handler treats a raise as allow-all for key auth.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
if keyless_source:
return allowed_tools
@ -3159,11 +3287,10 @@ class MCPRequestHandler:
if object_permissions is None:
return allowed_tools
user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).get(server_id)
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
user_tools: Final = MCPRequestHandler._union_tool_grants(user_direct_tools, user_toolset_tools)
resolved_inventory: Final = (
inventory if inventory is not None else await MCPRequestHandler._manager_inventory(server_id)
)
user_tools: Final = await MCPRequestHandler._row_level_tools(object_permissions, server_id, resolved_inventory)
if user_tools is None:
return allowed_tools
if allowed_tools is None:
@ -3175,11 +3302,10 @@ class MCPRequestHandler:
allowed_tools: Sequence[str] | None,
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
inventory: Mapping[str, str | None] | None = None,
) -> Sequence[str] | None:
"""Narrow a key/team tool allowlist by the end user's (customer's) tool entitlement."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.proxy_server import prisma_client
if user_api_key_auth is None or not user_api_key_auth.end_user_id or prisma_client is None:
@ -3191,11 +3317,12 @@ class MCPRequestHandler:
if object_permissions is None:
return allowed_tools
end_user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).get(server_id)
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
end_user_tools: Final = MCPRequestHandler._union_tool_grants(end_user_direct_tools, end_user_toolset_tools)
resolved_inventory: Final = (
inventory if inventory is not None else await MCPRequestHandler._manager_inventory(server_id)
)
end_user_tools: Final = await MCPRequestHandler._row_level_tools(
object_permissions, server_id, resolved_inventory
)
if end_user_tools is None:
return allowed_tools
if allowed_tools is None:
@ -3347,6 +3474,8 @@ class MCPRequestHandler:
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
*,
inventory: Mapping[str, str | None] | None = None,
) -> list[str] | None:
"""
Get allowed tool names for a server from the agent's object_permission: the union of its
@ -3374,18 +3503,11 @@ class MCPRequestHandler:
return None
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
resolved_inventory: Final = (
inventory if inventory is not None else await MCPRequestHandler._manager_inventory(server_id)
)
direct_tools: Final = (
global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions).get(server_id)
if obj_perm.mcp_tool_permissions
else None
)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
return list(agent_tools) if agent_tools else None
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
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
raise

View file

@ -185,6 +185,7 @@ from litellm.types.mcp import (
MCPAuth,
MCPStdioConfig,
MCPTokenEndpointAuthMethod,
MCPToolOverrideEntry,
has_header,
without_header,
)
@ -1921,6 +1922,10 @@ class MCPServerManager:
"gmail_send_email": "zapier_mcp_server",
}
"""
# 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._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
@ -2821,6 +2826,8 @@ class MCPServerManager:
)
self._invalidate_discovery_lists(server.server_id)
if server.server_id:
self.discovered_tool_inventory.pop(server.server_id, None)
prefix_root: Final = normalize_server_name(get_server_prefix(server))
if server.spec_path and prefix_root:
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
@ -5378,8 +5385,15 @@ class MCPServerManager:
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
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}
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, {})
def _create_prefixed_prompts(
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
) -> list[Prompt]:
@ -6829,6 +6843,30 @@ class MCPServerManager:
result.setdefault(server_id, []).extend(tools or [])
return result
def expand_tool_overrides(
self,
tool_overrides: dict[str, MCPToolOverrideEntry] | None,
) -> dict[str, MCPToolOverrideEntry]:
"""
Rewrite an ``mcp_tool_overrides`` dict keyed by id/name/alias so every
key is a concrete server_id, same expansion as
``expand_tool_permissions``. Entries resolving to the same server
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 [])],
}
for server_id, entries in grouped.items()
}
def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None:
"""
Get the MCP Server from the server name.

View file

@ -1418,16 +1418,19 @@ async def filter_tools_by_key_team_permissions(
but tool names from MCP servers are prefixed. We need to strip
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}
# Filter by key/team tool-level permissions
allowed_tool_names: Final = await MCPRequestHandler.get_allowed_tools_for_server(
server_id=server_id,
user_api_key_auth=user_api_key_auth,
inventory=inventory,
)
# Tools arrive prefixed with the server's own prefix; strip exactly that
# prefix (resolved from the server) rather than the first separator, so a
# prefix containing the separator still reduces to the stored bare name.
server: Final = global_mcp_server_manager.get_mcp_server_by_id(server_id)
return [
t
for t in tools

View file

@ -0,0 +1,140 @@
"""Classify an MCP tool name (+ optional description) into a coarse operation kind.
The MCP tool-permission convention treats ``delete`` differently from
everything else (create/read/update/unknown all count as non-delete), so the
single exported function is all callers need.
"""
import re
from typing import Final, Literal
ToolOperation = Literal["read", "create", "update", "delete", "unknown"]
_READ_TOKENS: Final = frozenset(
{
"get",
"read",
"list",
"fetch",
"search",
"find",
"query",
"retrieve",
"show",
"view",
"check",
"describe",
"info",
"lookup",
"count",
"export",
"download",
}
)
_DELETE_TOKENS: Final = frozenset(
{
"delete",
"remove",
"destroy",
"purge",
"drop",
"erase",
"unlink",
"wipe",
"clear",
"revoke",
"uninstall",
"trash",
"truncate",
"rm",
"del",
}
)
_UPDATE_TOKENS: Final = frozenset(
{
"update",
"edit",
"modify",
"change",
"patch",
"put",
"set",
"rename",
"move",
"transform",
"toggle",
"enable",
"disable",
"archive",
"restore",
}
)
_CREATE_TOKENS: Final = frozenset(
{
"create",
"add",
"insert",
"new",
"post",
"submit",
"register",
"make",
"generate",
"write",
"upload",
"send",
"publish",
}
)
_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 _description_tokens(description: str) -> list[str]:
return [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"),
}
return frozenset(variant for variant in variants if variant)
def _classify_tokens(tokens: list[str]) -> ToolOperation:
token_set: Final = frozenset(variant for token in tokens for variant in _token_variants(token))
if token_set & _READ_TOKENS:
return "read"
if token_set & _DELETE_TOKENS:
return "delete"
if token_set & _UPDATE_TOKENS:
return "update"
if token_set & _CREATE_TOKENS:
return "create"
return "unknown"
def classify_tool_op(name: str, description: str | None = None) -> ToolOperation:
"""Classify a tool by exact token match on its name, falling back to the
description's words only when the name yields no recognized token.
Precedence is read > delete > update > create, so ``get_removed_entries``
is read and a misleading description cannot override a recognized name."""
by_name: Final = _classify_tokens(_name_tokens(name))
if by_name != "unknown":
return by_name
if description:
return _classify_tokens(_description_tokens(description))
return "unknown"

View file

@ -41,6 +41,7 @@ from litellm.types.mcp import (
MCPAuth,
MCPAuthType,
MCPCredentials,
MCPToolOverrideEntry,
MCPTransport,
MCPTransportType,
)
@ -1194,6 +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_toolsets: list[str] | None = None
blocked_tools: list[str] | None = None
vector_stores: list[str] | None = None

View file

@ -54,6 +54,7 @@ from litellm.proxy.management_endpoints.common_utils import (
from litellm.proxy.management_helpers.object_permission_utils import (
handle_update_object_permission_common,
prepare_object_permission_upsert,
reject_ambiguous_mcp_tool_override_keys,
reject_ambiguous_mcp_tool_permission_keys,
)
from litellm.proxy.management_helpers.utils import (
@ -649,8 +650,16 @@ async def _set_object_permission(
existing_mcp_tool_permissions=None,
prisma_client=prisma_client,
)
await reject_ambiguous_mcp_tool_override_keys(
new_mcp_tool_overrides=getattr(data.object_permission, "mcp_tool_overrides", None),
existing_mcp_tool_overrides=None,
prisma_client=prisma_client,
)
created_object_permission: Final = await _table(ObjectPermissionRepository(prisma_client)).create(
data=data.object_permission.model_dump(exclude_none=True),
data={
**data.object_permission.model_dump(exclude_none=True),
"mcp_permission_version": 1,
},
)
del data.object_permission
return created_object_permission.object_permission_id

View file

@ -22,6 +22,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, objec
from litellm.proxy.utils import PrismaClient
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.table_repositories import MCPServerRepository
from litellm.types.mcp import MCPToolOverrideEntry
if TYPE_CHECKING:
from prisma import models as prisma_models
@ -113,10 +114,16 @@ async def prepare_object_permission_upsert(
existing_mcp_tool_permissions=existing_fields.get("mcp_tool_permissions"),
prisma_client=prisma_client,
)
await reject_ambiguous_mcp_tool_override_keys(
new_mcp_tool_overrides=new_object_permission.get("mcp_tool_overrides"),
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,
@ -125,6 +132,7 @@ async def prepare_object_permission_upsert(
if "mcp_tool_permissions" in merged
else {}
),
**({"mcp_tool_overrides": safe_dumps(merged["mcp_tool_overrides"])} if "mcp_tool_overrides" in merged else {}),
}
return ObjectPermissionUpsert(object_permission_id=object_permission_id, record=record)
@ -225,10 +233,18 @@ async def _set_object_permission(
existing_mcp_tool_permissions=None,
prisma_client=prisma_client,
)
await reject_ambiguous_mcp_tool_override_keys(
new_mcp_tool_overrides=clean_data.get("mcp_tool_overrides"),
existing_mcp_tool_overrides=None,
prisma_client=prisma_client,
)
# Serialize mcp_tool_permissions to JSON string for GraphQL compatibility
# Serialize mcp_tool_permissions / mcp_tool_overrides to JSON strings for GraphQL compatibility
if "mcp_tool_permissions" in clean_data:
clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"])
if "mcp_tool_overrides" in clean_data:
clean_data["mcp_tool_overrides"] = safe_dumps(clean_data["mcp_tool_overrides"])
clean_data["mcp_permission_version"] = 1
created_permission: Final = await ObjectPermissionRepository(prisma_client).table.create(data=clean_data)
@ -332,6 +348,54 @@ def _mcp_tool_permission_entries(raw: object) -> Mapping[str, frozenset[str]]:
return MappingProxyType({identifier: frozenset(tools or ()) for identifier, tools in parsed.items()})
_MCP_TOOL_OVERRIDES_ADAPTER: Final = TypeAdapter(dict[str, MCPToolOverrideEntry | None])
def _mcp_tool_override_entries(raw: object) -> Mapping[str, object]:
parsed: Final[Mapping[str, MCPToolOverrideEntry | None]] = (
_MCP_TOOL_OVERRIDES_ADAPTER.validate_json(raw)
if isinstance(raw, str)
else _MCP_TOOL_OVERRIDES_ADAPTER.validate_python(raw)
if isinstance(raw, Mapping)
else MappingProxyType({})
)
return parsed
async def reject_ambiguous_mcp_tool_override_keys(
new_mcp_tool_overrides: object,
existing_mcp_tool_overrides: object,
prisma_client: PrismaClient | None,
) -> None:
"""
Same ambiguity rule as ``reject_ambiguous_mcp_tool_permission_keys`` for
``mcp_tool_overrides`` keys: a name or alias matching several servers
cannot key an override entry. Raises HTTPException(400) on collision.
"""
requested: Final = _mcp_tool_override_entries(new_mcp_tool_overrides)
stored: Final = _mcp_tool_override_entries(existing_mcp_tool_overrides)
resolved: Final = await _resolve_mcp_server_identifiers_to_ids(
identifiers=frozenset(identifier for identifier, entry in requested.items() if stored.get(identifier) != entry),
prisma_client=prisma_client,
)
collisions: Final = "; ".join(
f"'{identifier}' matches MCP servers {sorted(server_ids)}"
for identifier, server_ids in sorted(resolved.items())
if identifier not in server_ids and len(server_ids) > 1
)
if not collisions:
return
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={ # mutable-ok: HTTPException.detail has no immutable form; same shape as the sibling errors here
"error": (
f"Ambiguous mcp_tool_overrides key: {collisions}. "
"Key tool overrides by server_id when servers share a name or alias."
)
},
)
async def reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions: object,
existing_mcp_tool_permissions: object,
@ -405,6 +469,21 @@ 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]],
) -> None:
mcp_tool_overrides: Final = object_permission.get("mcp_tool_overrides")
if not isinstance(mcp_tool_overrides, dict):
return
object_permission["mcp_tool_overrides"] = {
identifier: entry
for identifier, entry in mcp_tool_overrides.items()
if identifier_to_server_ids.get(identifier)
}
def _drop_stale_object_permission_mcp_identifiers(
object_permission: ObjectPermissionDict | None,
identifier_to_server_ids: dict[str, set[str]],
@ -420,6 +499,10 @@ def _drop_stale_object_permission_mcp_identifiers(
object_permission=object_permission,
identifier_to_server_ids=identifier_to_server_ids,
)
_drop_stale_object_permission_mcp_tool_overrides(
object_permission=object_permission,
identifier_to_server_ids=identifier_to_server_ids,
)
def _flatten_resolved_mcp_server_ids(
@ -453,7 +536,8 @@ async def _resolve_team_allowed_mcp_servers(
raw_tool_perms = team_object_permission.mcp_tool_permissions or {}
if isinstance(raw_tool_perms, str):
raw_tool_perms = json.loads(raw_tool_perms)
tool_perm_servers: Final[list[str]] = list(raw_tool_perms.keys())
raw_tool_overrides: Final = _mcp_tool_override_entries(getattr(team_object_permission, "mcp_tool_overrides", None))
tool_perm_servers: Final[list[str]] = list(raw_tool_perms.keys()) + list(raw_tool_overrides.keys())
raw_servers: Final = set(direct_servers + access_group_servers + tool_perm_servers)
resolved_servers: Final = await _resolve_mcp_server_identifiers_to_ids(
identifiers=raw_servers,
@ -541,9 +625,12 @@ async def _get_grandfathered_key_mcp_server_ids(
if existing_object_permission is None or prisma_client is None:
return frozenset()
raw_tool_perms: Final = existing_object_permission.mcp_tool_permissions or {}
raw_tool_overrides: Final = _mcp_tool_override_entries(
getattr(existing_object_permission, "mcp_tool_overrides", None)
)
tool_perm_keys: Final[frozenset[str]] = frozenset(
json.loads(raw_tool_perms).keys() if isinstance(raw_tool_perms, str) else raw_tool_perms.keys()
)
) | frozenset(raw_tool_overrides.keys())
identifiers: Final = (frozenset(existing_object_permission.mcp_servers or []) | tool_perm_keys) - {
SpecialMCPServerNames.no_mcp_servers.value,
SpecialMCPServerName.all_proxy_servers.value,
@ -604,6 +691,10 @@ def _extract_requested_mcp_server_ids(
if isinstance(mcp_tool_permissions, dict):
server_ids.update(mcp_tool_permissions.keys())
mcp_tool_overrides: Final = object_permission.get("mcp_tool_overrides")
if isinstance(mcp_tool_overrides, dict):
server_ids.update(mcp_tool_overrides.keys())
return server_ids

View file

@ -282,6 +282,9 @@ model LiteLLM_ObjectPermissionTable {
mcp_servers String[] @default([])
mcp_access_groups String[] @default([])
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
mcp_tool_overrides Json? // Per-server tool allow/deny overrides. Format: {"server_id": {"allow": ["t1"], "deny": ["t2"]}}
mcp_permission_version Int @default(0) // 0 = pre-overrides row (legacy evaluation), 1 = written after overrides shipped
mcp_tool_permissions_archive Json? // Backfill's snapshot of the pre-conversion mcp_tool_permissions, recovery only
vector_stores String[] @default([])
agents String[] @default([])
agent_access_groups String[] @default([])

View file

@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any, Final
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.repositories.base_repository import BaseRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.types.mcp import MCPToolOverrideEntry
if TYPE_CHECKING:
from prisma import models as prisma_models
@ -33,6 +34,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,
vector_stores: list[str] | None = None,
agents: list[str] | None = None,
agent_access_groups: list[str] | None = None,
@ -43,13 +45,15 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
skills: list[str] | None = None,
) -> LiteLLM_ObjectPermissionTable:
"""Create a new object permission record."""
data: Final[dict[str, Any]] = {}
data: Final[dict[str, Any]] = {"mcp_permission_version": 1}
if mcp_servers is not None:
data["mcp_servers"] = mcp_servers
if mcp_access_groups is not None:
data["mcp_access_groups"] = mcp_access_groups
if mcp_tool_permissions is not None:
data["mcp_tool_permissions"] = mcp_tool_permissions
if mcp_tool_overrides is not None:
data["mcp_tool_overrides"] = mcp_tool_overrides
if vector_stores is not None:
data["vector_stores"] = vector_stores
if agents is not None:
@ -75,6 +79,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,
vector_stores: list[str] | None = None,
agents: list[str] | None = None,
agent_access_groups: list[str] | None = None,
@ -85,13 +90,15 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
skills: list[str] | None = None,
) -> LiteLLM_ObjectPermissionTable | None:
"""Update an object permission record."""
data: Final[dict[str, Any]] = {}
data: Final[dict[str, Any]] = {"mcp_permission_version": 1}
if mcp_servers is not None:
data["mcp_servers"] = mcp_servers
if mcp_access_groups is not None:
data["mcp_access_groups"] = mcp_access_groups
if mcp_tool_permissions is not None:
data["mcp_tool_permissions"] = mcp_tool_permissions
if mcp_tool_overrides is not None:
data["mcp_tool_overrides"] = mcp_tool_overrides
if vector_stores is not None:
data["vector_stores"] = vector_stores
if agents is not None:

View file

@ -9,7 +9,7 @@ from urllib.parse import urlsplit
import httpx
from pydantic import BaseModel, ConfigDict, Field
from typing_extensions import TypedDict
from typing_extensions import ReadOnly, TypedDict
from litellm.types.llms.base import HiddenParams
@ -491,3 +491,12 @@ class MCPGatewaySessionsTerminateResponse(BaseModel):
worker_pid: int
terminated_sessions: int
sessions: list[MCPGatewaySession] = Field(default_factory=list)
class MCPToolOverrideEntry(TypedDict, total=False):
"""Per-server tool overrides stored on an object permission row's
``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]]

View file

@ -10,11 +10,14 @@ layering rule.
from typing_extensions import ReadOnly, TypedDict
from litellm.types.mcp import MCPToolOverrideEntry
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_toolsets: list[str] | None
blocked_tools: list[str] | None
vector_stores: list[str] | None

View file

@ -282,6 +282,9 @@ model LiteLLM_ObjectPermissionTable {
mcp_servers String[] @default([])
mcp_access_groups String[] @default([])
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
mcp_tool_overrides Json? // Per-server tool allow/deny overrides. Format: {"server_id": {"allow": ["t1"], "deny": ["t2"]}}
mcp_permission_version Int @default(0) // 0 = pre-overrides row (legacy evaluation), 1 = written after overrides shipped
mcp_tool_permissions_archive Json? // Backfill's snapshot of the pre-conversion mcp_tool_permissions, recovery only
vector_stores String[] @default([])
agents String[] @default([])
agent_access_groups String[] @default([])

View file

@ -331,14 +331,18 @@ class TestMCPRequestHandler:
key_object_permission.mcp_servers = []
key_object_permission.mcp_access_groups = []
key_object_permission.mcp_tool_permissions = None
key_object_permission.mcp_tool_overrides = None
key_object_permission.mcp_permission_version = None
key_object_permission.mcp_toolsets = toolset_ids
return key_object_permission
def _mock_manager_with_toolsets(self, toolset_perms):
def _mock_manager_with_toolsets(self, toolset_perms, inventory=None):
mock_manager = MagicMock()
mock_manager.expand_permission_list = MagicMock(side_effect=lambda servers: servers)
mock_manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {})
mock_manager.expand_tool_overrides = MagicMock(side_effect=lambda overrides: overrides or {})
mock_manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms)
mock_manager.discovered_inventory = MagicMock(return_value=inventory or {})
return mock_manager
async def test_get_allowed_mcp_servers_for_key_includes_toolset_servers(self):
@ -434,6 +438,9 @@ class TestMCPRequestHandler:
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
key_object_permission = self._toolset_only_object_permission(["toolset-1"])
key_object_permission.mcp_tool_permissions = {"server-a": ["direct_tool"]}
key_object_permission.mcp_tool_overrides = None
key_object_permission.mcp_permission_version = None
key_object_permission.mcp_access_groups = []
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
with (
@ -518,6 +525,9 @@ class TestMCPRequestHandler:
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", team_id="team-1")
team_object_permission = self._toolset_only_object_permission(["toolset-1"])
team_object_permission.mcp_tool_permissions = {"server-a": ["direct_tool"]}
team_object_permission.mcp_tool_overrides = None
team_object_permission.mcp_permission_version = None
team_object_permission.mcp_access_groups = []
mock_manager = self._mock_manager_with_toolsets({"server-a": ["search_channels", "read_thread"]})
with (
@ -3917,6 +3927,10 @@ async def test_get_allowed_tools_for_server_ui_session_team_keeps_key_restrictio
)
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a"]}
key_perm.mcp_tool_overrides = None
key_perm.mcp_permission_version = None
key_perm.mcp_toolsets = None
key_perm.mcp_access_groups = []
mock_prisma = MagicMock()
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
@ -4354,8 +4368,7 @@ class TestAgentMCPPermissions:
assert result == frozenset({"ag-server-id"})
assert asked == ["agent-ag"]
assert (
await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k"), resolve)
is None
await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k"), resolve) is None
)
assert asked == ["agent-ag"]
@ -4384,6 +4397,10 @@ class TestAgentMCPPermissions:
)
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
key_perm.mcp_tool_overrides = None
key_perm.mcp_permission_version = None
key_perm.mcp_toolsets = None
key_perm.mcp_access_groups = []
team_perm = None
with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm):
with patch.object(
@ -4417,6 +4434,10 @@ class TestAgentMCPPermissions:
)
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
key_perm.mcp_tool_overrides = None
key_perm.mcp_permission_version = None
key_perm.mcp_toolsets = None
key_perm.mcp_access_groups = []
with patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm):
with patch.object(
MCPRequestHandler,
@ -4492,7 +4513,9 @@ class TestAgentMCPPermissions:
stack.enter_context(patcher)
stack.enter_context(
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"])
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
AsyncMock(return_value=["server-a", "server-b"]),
)
)
stack.enter_context(
@ -4518,7 +4541,9 @@ class TestAgentMCPPermissions:
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
stack.enter_context(
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"])
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
AsyncMock(return_value=["server-a", "server-b"]),
)
)
stack.enter_context(
@ -4542,9 +4567,15 @@ class TestAgentMCPPermissions:
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-a", user_api_key_auth)
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-b", user_api_key_auth)
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-c", user_api_key_auth)
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
"server-a", user_api_key_auth
)
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
"server-b", user_api_key_auth
)
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
"server-c", user_api_key_auth
)
assert sorted(server_a_tools) == ["tool_direct", "tool_via_toolset"]
assert server_b_tools == ["tool_b"]
@ -4677,6 +4708,10 @@ async def test_tool_permission_servers_included_in_allowed_servers():
perm.mcp_servers = []
perm.mcp_access_groups = []
perm.mcp_tool_permissions = {"server_id_123": ["tool_a", "tool_b"]}
perm.mcp_tool_overrides = None
perm.mcp_permission_version = None
perm.mcp_toolsets = None
perm.mcp_access_groups = []
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
@ -4828,6 +4863,10 @@ class TestOrgMCPPermissions:
mock_perm.mcp_servers = ["org_server_1", "org_server_2"]
mock_perm.mcp_access_groups = []
mock_perm.mcp_tool_permissions = {}
mock_perm.mcp_tool_overrides = None
mock_perm.mcp_permission_version = None
mock_perm.mcp_toolsets = None
mock_perm.mcp_access_groups = []
with (
patch.object(
@ -4854,6 +4893,10 @@ class TestOrgMCPPermissions:
mock_perm.mcp_servers = []
mock_perm.mcp_access_groups = ["group-a"]
mock_perm.mcp_tool_permissions = {}
mock_perm.mcp_tool_overrides = None
mock_perm.mcp_permission_version = None
mock_perm.mcp_toolsets = None
mock_perm.mcp_access_groups = []
with (
patch.object(
@ -4880,6 +4923,10 @@ class TestOrgMCPPermissions:
mock_perm.mcp_servers = []
mock_perm.mcp_access_groups = []
mock_perm.mcp_tool_permissions = {"tool_only_server": ["tool_x"]}
mock_perm.mcp_tool_overrides = None
mock_perm.mcp_permission_version = None
mock_perm.mcp_toolsets = None
mock_perm.mcp_access_groups = []
with (
patch.object(
@ -4915,10 +4962,18 @@ class TestOrgMCPPermissions:
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b", "tool_c"]}
key_perm.mcp_tool_overrides = None
key_perm.mcp_permission_version = None
key_perm.mcp_toolsets = None
key_perm.mcp_access_groups = []
org_perm = MagicMock()
org_perm.mcp_toolsets = None # a bare MagicMock attr reads as a DECLARED toolset and now denies
org_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
org_perm.mcp_tool_overrides = None
org_perm.mcp_permission_version = None
org_perm.mcp_toolsets = None
org_perm.mcp_access_groups = []
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm),
@ -4946,10 +5001,18 @@ class TestOrgMCPPermissions:
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
key_perm.mcp_tool_overrides = None
key_perm.mcp_permission_version = None
key_perm.mcp_toolsets = None
key_perm.mcp_access_groups = []
org_perm = MagicMock()
org_perm.mcp_toolsets = None # a bare MagicMock attr reads as a DECLARED toolset and now denies
org_perm.mcp_tool_permissions = {}
org_perm.mcp_tool_overrides = None
org_perm.mcp_permission_version = None
org_perm.mcp_toolsets = None
org_perm.mcp_access_groups = []
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_perm),
@ -9832,3 +9895,150 @@ class TestScopedSessionAdmission:
def test_scope_field_cannot_be_forged_through_construction(self):
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
assert forged.mcp_session_resource_server_id is None
def _converted_row(
mcp_servers,
mcp_tool_permissions=None,
mcp_tool_overrides=None,
mcp_toolsets=None,
mcp_access_groups=None,
):
row = MagicMock()
row.mcp_servers = mcp_servers
row.mcp_access_groups = mcp_access_groups or []
row.mcp_tool_permissions = mcp_tool_permissions
row.mcp_tool_overrides = mcp_tool_overrides
row.mcp_permission_version = 1
row.mcp_toolsets = mcp_toolsets
return row
def _manager(toolset_perms=None, inventory=None):
manager = MagicMock()
manager.expand_permission_list = MagicMock(side_effect=lambda servers: servers)
manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {})
manager.expand_tool_overrides = MagicMock(side_effect=lambda overrides: overrides or {})
manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms or {})
manager.discovered_inventory = MagicMock(return_value=inventory or {})
return manager
@pytest.mark.asyncio
class TestLevelAllowedToolsConvention:
"""Convention semantics for converted rows: non-delete inventory tools are
allowed by default, deletes are denied, overrides adjust on top"""
_INVENTORY = {"list_items": "list things", "delete_item": "remove one", "search_notes": None}
async def _allowed(self, key_row, team_row=None, inventory=_INVENTORY, toolset_perms=None):
manager = _manager(toolset_perms=toolset_perms, inventory=inventory)
with (
patch.object(MCPRequestHandler, "_get_key_object_permission", return_value=key_row),
patch.object(MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_row)),
patch.object(MCPRequestHandler, "_toolset_tool_permissions", AsyncMock(return_value=toolset_perms or {})),
patch.object(MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])),
patch.object(
MCPRequestHandler,
"_apply_end_user_tool_ceiling",
AsyncMock(side_effect=lambda allowed, *_a, **_k: allowed),
),
patch.object(
MCPRequestHandler,
"_apply_user_tool_ceiling",
AsyncMock(side_effect=lambda allowed, *_a, **_k: allowed),
),
patch.object(
MCPRequestHandler,
"_apply_agent_and_org_tool_ceilings",
AsyncMock(side_effect=lambda allowed, *_a, **_k: allowed),
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
manager,
),
):
auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
return await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a",
user_api_key_auth=auth,
)
async def test_converted_empty_overrides_allows_new_nondelete_denies_delete(self):
result = await self._allowed(_converted_row(["server-a"]))
assert result is not None
assert set(result) == {"list_items", "search_notes"}
async def test_explicit_allow_of_delete_tool(self):
row = _converted_row(["server-a"], mcp_tool_overrides={"server-a": {"allow": ["delete_item"]}})
result = await self._allowed(row)
assert set(result) == {"list_items", "search_notes", "delete_item"}
async def test_explicit_deny_of_nondelete_tool(self):
row = _converted_row(["server-a"], mcp_tool_overrides={"server-a": {"deny": ["search_notes"]}})
result = await self._allowed(row)
assert set(result) == {"list_items"}
async def test_tool_in_allow_and_deny_is_denied(self):
row = _converted_row(
["server-a"],
mcp_tool_overrides={"server-a": {"allow": ["list_items"], "deny": ["list_items"]}},
)
result = await self._allowed(row)
assert "list_items" not in (result or [])
async def test_legacy_entry_stays_closed_allowlist(self):
row = _converted_row(["server-a"], mcp_tool_permissions={"server-a": ["list_items"]})
result = await self._allowed(row)
assert set(result) == {"list_items"}
async def test_legacy_empty_list_denies_all(self):
row = _converted_row(["server-a"], mcp_tool_permissions={"server-a": []})
result = await self._allowed(row)
assert result == []
async def test_unconverted_row_places_no_restriction(self):
row = _converted_row(["server-a"])
row.mcp_permission_version = 0
result = await self._allowed(row)
assert result is None
async def test_converted_row_not_granting_server_places_no_restriction(self):
row = _converted_row(["other-server"])
result = await self._allowed(row)
assert result is None
async def test_no_rows_anywhere_uses_convention_over_inventory(self):
result = await self._allowed(None, team_row=None)
assert result is not None
assert set(result) == {"list_items", "search_notes"}
async def test_empty_inventory_only_explicit_names(self):
row = _converted_row(["server-a"], mcp_tool_overrides={"server-a": {"allow": ["pinned_tool"]}})
result = await self._allowed(row, inventory={})
assert set(result) == {"pinned_tool"}
async def test_empty_inventory_converted_no_overrides_allows_nothing(self):
result = await self._allowed(_converted_row(["server-a"]), inventory={})
assert result == []
async def test_team_convention_deny_intersects_with_key_allow(self):
key_row = _converted_row(["server-a"], mcp_tool_overrides={"server-a": {"allow": ["delete_item"]}})
team_row = _converted_row(["server-a"])
result = await self._allowed(key_row, team_row=team_row)
assert "delete_item" not in (result or [])
assert set(result) == {"list_items", "search_notes"}
async def test_toolset_tools_allowed_even_if_delete_classified(self):
row = _converted_row([], mcp_toolsets=["toolset-1"])
result = await self._allowed(row, toolset_perms={"server-a": ["delete_item"]})
assert "delete_item" in (result or [])
async def test_legacy_entry_ignores_overrides(self):
row = _converted_row(
["server-a"],
mcp_tool_permissions={"server-a": ["list_items"]},
mcp_tool_overrides={"server-a": {"allow": ["delete_item"]}},
)
result = await self._allowed(row)
assert set(result) == {"list_items"}

View file

@ -0,0 +1,25 @@
[
{"name": "delete_link", "description": null, "expected": "delete"},
{"name": "delete-link", "description": null, "expected": "delete"},
{"name": "deleteLink", "description": null, "expected": "delete"},
{"name": "DeleteLink", "description": null, "expected": "delete"},
{"name": "DELETE_LINK", "description": null, "expected": "delete"},
{"name": "remove_user", "description": null, "expected": "delete"},
{"name": "purge_cache", "description": null, "expected": "delete"},
{"name": "get_removed_entries", "description": null, "expected": "read"},
{"name": "list_deleted_items", "description": null, "expected": "read"},
{"name": "find_deleted", "description": null, "expected": "read"},
{"name": "describe_purge_job", "description": null, "expected": "read"},
{"name": "updateItem", "description": null, "expected": "update"},
{"name": "create-record", "description": null, "expected": "create"},
{"name": "foo", "description": "Deletes the file", "expected": "delete"},
{"name": "foo", "description": "Lists files", "expected": "read"},
{"name": "delete_link", "description": "Reads a link", "expected": "delete"},
{"name": "get_removed_entries", "description": "Permanently deletes", "expected": "read"},
{"name": "unlinkNode", "description": null, "expected": "delete"},
{"name": "checkDeletion", "description": null, "expected": "read"},
{"name": "rm_file", "description": null, "expected": "delete"},
{"name": "x", "description": null, "expected": "unknown"},
{"name": "info", "description": null, "expected": "read"},
{"name": "settings", "description": null, "expected": "unknown"}
]

View file

@ -19,7 +19,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -1073,7 +1073,16 @@ class TestOpenApiByokCallTool:
spec_path="https://example.com/openapi.json",
is_byok=True,
)
user_auth = UserAPIKeyAuth(user_id="default_user_id", api_key="sk-dashboard")
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
user_auth = UserAPIKeyAuth(
user_id="default_user_id",
api_key="sk-dashboard",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-unconverted",
mcp_permission_version=0,
),
)
captured_auth: dict[str, Optional[str]] = {}
async def fake_openapi_handler(_server, _name, _arguments):
@ -1313,7 +1322,14 @@ class TestOpenApiResolvedUpstreamAuth:
manager = MCPServerManager()
server = self._oauth_server()
user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user")
user_auth = UserAPIKeyAuth(
user_id="alice",
api_key="sk-user",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-unconverted",
mcp_permission_version=0,
),
)
captured: Dict[str, Any] = {}
async def fake_openapi_handler(_server, _name, _arguments):
@ -1361,7 +1377,14 @@ class TestOpenApiResolvedUpstreamAuth:
server_name=server.server_name,
name="get_values",
arguments={},
user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="sk-user"),
user_api_key_auth=UserAPIKeyAuth(
user_id="alice",
api_key="sk-user",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-unconverted",
mcp_permission_version=0,
),
),
)
called.assert_not_awaited()

View file

@ -83,9 +83,6 @@ def cleanup_mcp_global_state():
yield
def _call_tool_params(name, arguments=None):
from mcp.types import CallToolRequestParams
@ -97,6 +94,7 @@ def _paged_params():
return PaginatedRequestParams()
@pytest.mark.asyncio
async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx):
"""Test that proxy_server_request body contains name and arguments"""
@ -293,7 +291,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r
):
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
result = await mcp_server_tool_call(
_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})
)
assert result.is_error is True
# The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this
@ -1165,20 +1165,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind,
else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata)
)
with (
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))),
patch.object(
server,
"get_or_extract_auth_context",
AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)),
),
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])),
patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))),
patch.object(
operations.global_mcp_server_manager,
"read_resource_from_server",
AsyncMock(return_value=ReadResourceResult(contents=[content])),
),
):
result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri))
assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == {
"cacheScope": "private", "resultType": "complete", "ttlMs": 0,
"contents": [{
"uri": uri,
"mimeType": "text/plain" if kind == "text" else "image/png",
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
**({"_meta": metadata} if metadata is not None else {}),
}],
"cacheScope": "private",
"resultType": "complete",
"ttlMs": 0,
"contents": [
{
"uri": uri,
"mimeType": "text/plain" if kind == "text" else "image/png",
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
**({"_meta": metadata} if metadata is not None else {}),
}
],
}
@ -1672,7 +1684,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
with (
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None),
new=AsyncMock(
return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None
),
),
patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam
"litellm.proxy._experimental.mcp_server.operations._list_mcp_tools",
@ -1911,8 +1925,8 @@ async def test_streamable_http_session_manager_is_stateless():
("DELETE", b"", False),
),
)
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx,
debug: bool, method: str, request_body: bytes, stateful: bool
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
_mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool
) -> None:
from starlette.requests import Request
from starlette.types import Message, Receive, Scope, Send
@ -4054,7 +4068,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock(
# parsed, with a nested "method" key in the first bytes to trip a flat
# substring heuristic.
response_prefix: Final = (
'{"jsonrpc":"2.0","id":99,"' + response_field
'{"jsonrpc":"2.0","id":99,"'
+ response_field
+ '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"'
).encode()
response_body: Final = (
@ -7411,7 +7426,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
return_value=oauth_server,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -7657,7 +7673,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
return_value=alias_less_server,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -7932,7 +7949,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req
return_value=None,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -8097,7 +8115,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste
return_value=None,
),
patch.object(
mcp_operations, "_handle_managed_mcp_tool",
mcp_operations,
"_handle_managed_mcp_tool",
new=fake_handle_managed_mcp_tool,
),
patch.object(
@ -8596,7 +8615,9 @@ class TestMCPMetaTraceCarrier:
assert _mcp_meta_trace_carrier(None) is None
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta
only_progress = CallToolRequestParams.model_validate(
{"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False
).meta
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None
@ -10117,6 +10138,49 @@ class TestListFiltersHonorThePrefixBoundary:
assert listed == callable_, f"grants={grants!r} listed={listed} callable={callable_}"
assert listed is expected, f"grants={grants!r} expected={expected} got={listed}"
@pytest.mark.asyncio
async def test_key_team_listing_passes_fetched_tools_as_inventory(self):
"""The listing path must hand the just-fetched tools (bare names plus
descriptions) to the evaluator as its inventory, so convention
semantics classify exactly the tools the server advertised."""
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.server import (
filter_tools_by_key_team_permissions,
)
from litellm.proxy._types import UserAPIKeyAuth
server = MCPServer(
server_id=self.SERVER_ID,
name=self.SERVER_ID,
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
)
published = MCPTool(
name=f"{self.SERVER_ID}-read_wiki_contents",
description="reads a wiki page",
inputSchema={"type": "object"},
)
auth = UserAPIKeyAuth(api_key="sk-test")
captured: dict = {}
async def _capture(server_id, user_api_key_auth, *, keyless_source=False, inventory=None):
captured["inventory"] = inventory
return ["read_wiki_contents"]
with (
patch.object(MCPRequestHandler, "get_allowed_tools_for_server", AsyncMock(side_effect=_capture)),
patch("litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager") as mock_manager,
):
mock_manager.get_mcp_server_by_id.return_value = server
listed = await filter_tools_by_key_team_permissions([published], self.SERVER_ID, auth)
assert listed == [published]
assert captured["inventory"] == {"read_wiki_contents": "reads a wiki page"}
@pytest.mark.asyncio
@pytest.mark.parametrize(
@ -10247,7 +10311,9 @@ async def test_mcp_origin_admission_precedes_authentication(
patch("litellm.proxy.proxy_server.origins", allowed_origins),
patch.object(server, "extract_mcp_auth_context", authenticate),
):
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=server.app), base_url="http://gateway"
) as client:
response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers))
assert response.status_code == expected_status
@ -10329,12 +10395,15 @@ async def test_streamable_http_rejects_modern_protocol_version(
@pytest.mark.asyncio
@pytest.mark.parametrize("handler_name,field", [
("handle_list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
])
@pytest.mark.parametrize(
"handler_name,field",
[
("handle_list_tools", "tools"),
("list_prompts", "prompts"),
("list_resources", "resources"),
("list_resource_templates", "resource_templates"),
],
)
async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field):
from litellm.proxy._experimental.mcp_server import server
@ -10352,7 +10421,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
auth = UserAPIKeyAuth(user_id="denied-caller")
denial = HTTPException(status_code=403, detail="scope denied")
logger = MagicMock()
logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None)
logger.post_call_failure_hook = AsyncMock(
side_effect=RuntimeError("log unavailable") if failure_hook_raises else None
)
upstream = AsyncMock()
with (
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)),
@ -10361,7 +10432,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream),
):
with pytest.raises(HTTPException) as rejected:
await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True)
await operations._get_tools_from_mcp_servers(
user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True
)
assert rejected.value is denial
upstream.assert_not_awaited()
logger.post_call_failure_hook.assert_awaited_once()
@ -10460,7 +10533,16 @@ async def test_legacy_sse_mount_emits_message_endpoint(prefix: str, suffix: str)
patch.object(
mcp_server,
"extract_mcp_auth_context",
AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})),
AsyncMock(
return_value=(
post_auth,
None,
[marker],
{marker: {"Authorization": marker}},
{"Authorization": marker},
{"x-request-marker": marker},
)
),
),
patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing),
):

View file

@ -0,0 +1,17 @@
import json
from pathlib import Path
import pytest
from litellm.proxy._experimental.mcp_server.tool_classification import classify_tool_op
_FIXTURE_PATH: Path = Path(__file__).parent / "fixtures" / "mcp_tool_classification_cases.json"
_CASES: list[dict] = json.loads(_FIXTURE_PATH.read_text())
@pytest.mark.parametrize(
("name", "description", "expected"),
[pytest.param(c["name"], c.get("description"), c["expected"], id=c["name"]) for c in _CASES],
)
def test_classify_tool_op_fixture(name: str, description: str | None, expected: str) -> None:
assert classify_tool_op(name, description) == expected

View file

@ -39,9 +39,7 @@ async def test_set_object_permission():
mock_created_permission = MagicMock()
mock_created_permission.object_permission_id = "test_perm_id_123"
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=mock_created_permission
)
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(return_value=mock_created_permission)
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
# Test data with object_permission
@ -58,9 +56,7 @@ async def test_set_object_permission():
}
# Call the function
result = await _set_object_permission(
data_json=data_json, prisma_client=mock_prisma_client
)
result = await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client)
# Verify object_permission_id was added to result
assert result["object_permission_id"] == "test_perm_id_123"
@ -101,9 +97,7 @@ async def test_set_object_permission_persists_mcp_tool_search_enabled():
mock_prisma_client = MagicMock()
mock_created_permission = MagicMock()
mock_created_permission.object_permission_id = "perm_id"
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=mock_created_permission
)
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(return_value=mock_created_permission)
data_json = {
"object_permission": {
@ -114,11 +108,7 @@ async def test_set_object_permission_persists_mcp_tool_search_enabled():
await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client)
created_data = (
mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[
"data"
]
)
created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"]
assert created_data["mcp_tool_search_enabled"] is True
@ -127,9 +117,7 @@ async def test_set_object_permission_persists_skills():
mock_prisma_client = MagicMock()
mock_created_permission = MagicMock()
mock_created_permission.object_permission_id = "perm_id"
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=mock_created_permission
)
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(return_value=mock_created_permission)
data_json = {
"object_permission": LiteLLM_ObjectPermissionBase(skills=["private-skill"]).model_dump(),
@ -137,11 +125,7 @@ async def test_set_object_permission_persists_skills():
await _set_object_permission(data_json=data_json, prisma_client=mock_prisma_client)
created_data = (
mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs[
"data"
]
)
created_data = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args.kwargs["data"]
assert created_data["skills"] == ["private-skill"]
@ -222,11 +206,7 @@ def _make_team_obj(
mock_team = MagicMock()
mock_team.team_id = team_id
if (
mcp_servers is not None
or mcp_access_groups is not None
or mcp_tool_permissions is not None
):
if mcp_servers is not None or mcp_access_groups is not None or mcp_tool_permissions is not None:
mock_team.object_permission = MagicMock(spec=LiteLLM_ObjectPermissionTable)
mock_team.object_permission.mcp_servers = mcp_servers or []
mock_team.object_permission.mcp_access_groups = mcp_access_groups or []
@ -297,9 +277,7 @@ async def test_validate_no_object_permission(mock_access_groups, mock_allow_all)
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_key_servers_within_team_scope(
mock_access_groups, mock_allow_all
):
async def test_validate_key_servers_within_team_scope(mock_access_groups, mock_allow_all):
"""Key requests servers that are in the team's scope — should pass."""
team_obj = _make_team_obj(mcp_servers=["server-1", "server-2", "server-3"])
await validate_key_mcp_servers_against_team(
@ -322,9 +300,7 @@ async def test_validate_key_servers_within_team_scope(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_key_servers_outside_team_scope_raises(
mock_access_groups, mock_allow_all
):
async def test_validate_key_servers_outside_team_scope_raises(mock_access_groups, mock_allow_all):
"""Key requests a server that exists but is NOT in the team's scope — should raise 403."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
with pytest.raises(HTTPException) as exc_info:
@ -350,9 +326,7 @@ async def test_validate_key_servers_outside_team_scope_raises(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_allow_all_keys_servers_always_allowed(
mock_access_groups, mock_allow_all
):
async def test_validate_allow_all_keys_servers_always_allowed(mock_access_groups, mock_allow_all):
"""allow_all_keys servers should be accessible even if not in team scope."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
await validate_key_mcp_servers_against_team(
@ -397,9 +371,7 @@ async def test_validate_no_team_only_allow_all_keys(mock_access_groups, mock_all
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_non_global_server_raises(
mock_access_groups, mock_allow_all
):
async def test_validate_no_team_non_global_server_raises(mock_access_groups, mock_allow_all):
"""Key without a team requesting an existing non-global server — should raise 403."""
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
@ -424,9 +396,7 @@ async def test_validate_no_team_non_global_server_raises(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_proxy_admin_can_assign_private_server(
mock_access_groups, mock_allow_all
):
async def test_validate_no_team_proxy_admin_can_assign_private_server(mock_access_groups, mock_allow_all):
"""Proxy admin assigning a non-global server to a teamless key — should pass (LIT-3815)."""
result = await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["private-server"]},
@ -450,9 +420,7 @@ async def test_validate_no_team_proxy_admin_can_assign_private_server(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_non_admin_private_server_still_raises(
mock_access_groups, mock_allow_all
):
async def test_validate_no_team_non_admin_private_server_still_raises(mock_access_groups, mock_allow_all):
"""The teamless override is gated on proxy admin — a non-admin still gets 403."""
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
@ -473,9 +441,7 @@ async def test_validate_no_team_non_admin_private_server_still_raises(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_no_team_proxy_admin_can_assign_access_group(
mock_access_groups, mock_allow_all
):
async def test_validate_no_team_proxy_admin_can_assign_access_group(mock_access_groups, mock_allow_all):
"""Proxy admin assigning an access group to a teamless key — should pass (LIT-3815)."""
result = await validate_key_mcp_servers_against_team(
object_permission={"mcp_access_groups": ["group-1"]},
@ -499,9 +465,7 @@ async def test_validate_no_team_proxy_admin_can_assign_access_group(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_proxy_admin_still_bounded_by_team_scope(
mock_access_groups, mock_allow_all
):
async def test_validate_proxy_admin_still_bounded_by_team_scope(mock_access_groups, mock_allow_all):
"""The override is scoped to teamless keys — an admin assigning beyond a team's scope still raises."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
with pytest.raises(HTTPException) as exc_info:
@ -528,9 +492,7 @@ async def test_validate_proxy_admin_still_bounded_by_team_scope(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_team_no_mcp_config_blocks_all(
mock_access_groups, mock_allow_all
):
async def test_validate_team_no_mcp_config_blocks_all(mock_access_groups, mock_allow_all):
"""Team with no object_permission — key can't use any non-global MCP servers."""
team_obj = _make_team_obj() # No object_permission
with pytest.raises(HTTPException) as exc_info:
@ -555,9 +517,7 @@ async def test_validate_team_no_mcp_config_blocks_all(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_tool_permissions_validated_against_team(
mock_access_groups, mock_allow_all
):
async def test_validate_tool_permissions_validated_against_team(mock_access_groups, mock_allow_all):
"""Server IDs in mcp_tool_permissions should also be validated when they exist."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
with pytest.raises(HTTPException) as exc_info:
@ -583,9 +543,7 @@ async def test_validate_tool_permissions_validated_against_team(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_stale_mcp_server_ids_are_silently_dropped(
mock_access_groups, mock_allow_all
):
async def test_validate_stale_mcp_server_ids_are_silently_dropped(mock_access_groups, mock_allow_all):
"""
Stale MCP server IDs (servers deleted and no longer in the registry) must not
block a key save with a 403. They are silently stripped instead.
@ -615,9 +573,7 @@ async def test_validate_stale_mcp_server_ids_are_silently_dropped(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_stale_ids_in_mcp_tool_permissions_silently_dropped(
mock_access_groups, mock_allow_all
):
async def test_validate_stale_ids_in_mcp_tool_permissions_silently_dropped(mock_access_groups, mock_allow_all):
"""
Stale server IDs referenced only as keys in mcp_tool_permissions (not in
mcp_servers) must also be silently stripped rather than raising a 403.
@ -645,9 +601,7 @@ async def test_validate_stale_ids_in_mcp_tool_permissions_silently_dropped(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_stale_mcp_server_ids_are_removed_from_object_permission(
mock_access_groups, mock_allow_all
):
async def test_validate_stale_mcp_server_ids_are_removed_from_object_permission(mock_access_groups, mock_allow_all):
team_obj = _make_team_obj(mcp_servers=["s3", "s4"])
object_permission = {"mcp_servers": ["s1-stale", "s2-stale"]}
await validate_key_mcp_servers_against_team(
@ -680,9 +634,7 @@ async def test_validate_stale_mcp_server_ids_are_removed_from_object_permission(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_mcp_server_alias_outside_team_scope_raises(
mock_access_groups, mock_allow_all
):
async def test_validate_mcp_server_alias_outside_team_scope_raises(mock_access_groups, mock_allow_all):
team_obj = _make_team_obj(mcp_servers=["team-server"])
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
@ -775,9 +727,7 @@ async def test_validate_db_mcp_server_alias_outside_team_scope_raises_when_regis
mock_db_server.server_id = "private-server-id"
mock_db_server.alias = "private-alias"
mock_db_server.server_name = "Private Server"
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[mock_db_server]
)
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[mock_db_server])
team_obj = _make_team_obj(mcp_servers=[])
with pytest.raises(HTTPException) as exc_info:
@ -801,9 +751,7 @@ async def test_validate_db_mcp_server_alias_outside_team_scope_raises_when_regis
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_access_groups_within_team_scope(
mock_access_groups, mock_allow_all
):
async def test_validate_access_groups_within_team_scope(mock_access_groups, mock_allow_all):
"""Key requests access groups that are in the team's scope — should pass."""
team_obj = _make_team_obj(mcp_access_groups=["group-a", "group-b"])
await validate_key_mcp_servers_against_team(
@ -822,9 +770,7 @@ async def test_validate_access_groups_within_team_scope(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_access_groups_outside_team_scope_raises(
mock_access_groups, mock_allow_all
):
async def test_validate_access_groups_outside_team_scope_raises(mock_access_groups, mock_allow_all):
"""Key requests access groups NOT in the team's scope — should raise 403."""
team_obj = _make_team_obj(mcp_access_groups=["group-a"])
with pytest.raises(HTTPException) as exc_info:
@ -846,9 +792,7 @@ async def test_validate_access_groups_outside_team_scope_raises(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_access_groups_no_team_raises(
mock_access_groups, mock_allow_all
):
async def test_validate_access_groups_no_team_raises(mock_access_groups, mock_allow_all):
"""Key without a team requesting access groups — should raise 403."""
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
@ -869,9 +813,7 @@ async def test_validate_access_groups_no_team_raises(
new_callable=AsyncMock,
return_value=["server-from-group"],
)
async def test_validate_team_access_groups_resolve_to_servers(
mock_access_groups, mock_allow_all
):
async def test_validate_team_access_groups_resolve_to_servers(mock_access_groups, mock_allow_all):
"""Team access groups should resolve to server IDs and be included in allowed set."""
team_obj = _make_team_obj(mcp_access_groups=["group-a"])
# Key requests a server that comes from the team's access group
@ -976,9 +918,7 @@ async def test_resolve_team_all_proxy_sentinel_resolves_dynamically(mock_access_
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_key_scoped_to_server_added_after_team_all_proxy(
mock_access_groups, mock_allow_all
):
async def test_validate_key_scoped_to_server_added_after_team_all_proxy(mock_access_groups, mock_allow_all):
"""The exact user scenario: a team scoped to the all-proxy sentinel, a server
(srv-z) registered afterwards, and a key scoped to just srv-z. Because the
team ceiling resolves to every registered server, the key passes validation
@ -1007,9 +947,7 @@ async def test_validate_key_scoped_to_server_added_after_team_all_proxy(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_key_scoped_to_server_rejected_when_team_not_all_proxy(
mock_access_groups, mock_allow_all
):
async def test_validate_key_scoped_to_server_rejected_when_team_not_all_proxy(mock_access_groups, mock_allow_all):
"""Contrast with the sentinel case: a team scoped to a concrete server list
(srv-x, not the sentinel) does NOT unlock srv-z for a key. It is the sentinel
specifically, not a blanket allow, that widens the team ceiling."""
@ -1150,9 +1088,7 @@ async def test_validate_search_tools_raises_when_not_subset():
new_callable=AsyncMock,
return_value=[],
)
async def test_personal_non_admin_cannot_assign_mcp_toolsets(
mock_access_groups, mock_allow_all
):
async def test_personal_non_admin_cannot_assign_mcp_toolsets(mock_access_groups, mock_allow_all):
with pytest.raises(HTTPException) as exc:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_toolsets": ["ts-private"]},
@ -1173,9 +1109,7 @@ async def test_personal_non_admin_cannot_assign_mcp_toolsets(
new_callable=AsyncMock,
return_value=[],
)
async def test_personal_admin_can_assign_mcp_toolsets(
mock_access_groups, mock_allow_all
):
async def test_personal_admin_can_assign_mcp_toolsets(mock_access_groups, mock_allow_all):
await validate_key_mcp_servers_against_team(
object_permission={"mcp_toolsets": ["ts-private"]},
team_obj=None,
@ -1344,9 +1278,7 @@ async def test_validate_key_update_grandfathers_tool_permission_keys(monkeypatch
(stored as a JSON string) are grandfathered too."""
_patch_grandfather_env(monkeypatch, _make_mock_mcp_manager("server-a"))
team_obj = _make_team_obj(mcp_servers=[])
mock_prisma, existing_row = _make_grandfather_fixtures(
mcp_tool_permissions=json.dumps({"server-a": ["tool1"]})
)
mock_prisma, existing_row = _make_grandfather_fixtures(mcp_tool_permissions=json.dumps({"server-a": ["tool1"]}))
result = await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["server-a"]},
team_obj=team_obj,
@ -1504,3 +1436,78 @@ def test_object_permission_dict_mirrors_pydantic_model():
f"Only in Pydantic model: {sorted(pydantic_fields - typeddict_fields)}\n"
f"Only in TypedDict: {sorted(typeddict_fields - pydantic_fields)}"
)
@pytest.mark.asyncio
async def test_set_object_permission_accepts_mcp_tool_overrides_and_marks_converted():
"""mcp_tool_overrides writes through _set_object_permission, JSON-serialized
like mcp_tool_permissions, and every write stamps mcp_permission_version=1."""
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=MagicMock(object_permission_id="perm-id")
)
overrides = {"server_a": {"allow": ["search_notes"], "deny": ["delete_item"]}}
data_json = {
"object_permission": {
"mcp_servers": ["server_a"],
"mcp_tool_overrides": overrides,
},
}
result = await _set_object_permission(data_json=data_json, prisma_client=mock_prisma)
created_data = mock_prisma.db.litellm_objectpermissiontable.create.call_args.kwargs["data"]
assert json.loads(created_data["mcp_tool_overrides"]) == overrides
assert created_data["mcp_permission_version"] == 1
assert result["object_permission_id"] == "perm-id"
@pytest.mark.asyncio
@pytest.mark.parametrize("identifier", ["wiki", "github"])
async def test_set_object_permission_rejects_ambiguous_mcp_tool_override_key(identifier):
"""An alias or server_name shared by several servers cannot key mcp_tool_overrides
on create; the write is rejected with 400 and nothing is persisted."""
from litellm.proxy._types import SpecialMCPServerName # noqa: F401 # keep import parity cheap
mock_prisma = _make_ambiguity_prisma()
data_json = {"object_permission": {"mcp_tool_overrides": {identifier: {"allow": ["t"]}}}}
with pytest.raises(HTTPException) as exc_info:
await _set_object_permission(data_json=data_json, prisma_client=mock_prisma)
assert exc_info.value.status_code == 400
assert "mcp_tool_overrides" in str(exc_info.value.detail)
mock_prisma.db.litellm_objectpermissiontable.create.assert_not_called()
@pytest.mark.asyncio
async def test_prepare_object_permission_upsert_marks_converted_and_merges_overrides():
"""The update seam stamps mcp_permission_version=1 and JSON-serializes the
merged mcp_tool_overrides."""
mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"solo-id": ["tool1"]})
upsert = await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_overrides": {"solo-id": {"deny": ["delete_item"]}}},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert upsert.record["mcp_permission_version"] == 1
assert json.loads(upsert.record["mcp_tool_overrides"]) == {"solo-id": {"deny": ["delete_item"]}}
@pytest.mark.asyncio
async def test_prepare_object_permission_upsert_rejects_ambiguous_mcp_tool_override_key():
mock_prisma = _make_ambiguity_prisma()
with pytest.raises(HTTPException) as exc_info:
await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_overrides": {"wiki": {"deny": ["t"]}}},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert exc_info.value.status_code == 400
assert "'wiki'" in str(exc_info.value.detail)

View file

@ -1275,6 +1275,16 @@ class TestObjectPermissionRepository:
models=["gpt-4"],
)
assert perm.mcp_servers == ["server1"]
assert perm.mcp_permission_version == 1
@pytest.mark.asyncio
async def test_create_permission_with_tool_overrides(self, repo):
perm = await repo.create_permission(
mcp_servers=["server1"],
mcp_tool_overrides={"server1": {"allow": ["search_notes"], "deny": ["delete_item"]}},
)
assert perm.mcp_tool_overrides == {"server1": {"allow": ["search_notes"], "deny": ["delete_item"]}}
assert perm.mcp_permission_version == 1
@pytest.mark.asyncio
async def test_create_permission_all_fields(self, repo):
@ -1304,6 +1314,7 @@ class TestObjectPermissionRepository:
models=["gpt-4"],
)
assert updated.models == ["gpt-4"]
assert updated.mcp_permission_version == 1
@pytest.mark.asyncio
async def test_update_permission_all_fields(self, repo):