mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(mcp): evaluate tool permissions as convention over discovered inventory
Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
parent
4aa3ff47fe
commit
49b7bd8da3
23 changed files with 1645 additions and 465 deletions
|
|
@ -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;
|
||||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
140
litellm/proxy/_experimental/mcp_server/tool_classification.py
Normal file
140
litellm/proxy/_experimental/mcp_server/tool_classification.py
Normal 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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
):
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue