fix(proxy): reject ambiguous name or alias keys in mcp_tool_permissions on write (#39947)

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-05 14:22:26 -07:00 committed by GitHub
parent 9832d6e4a6
commit a46a076b2a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 229 additions and 2 deletions

View file

@ -46,6 +46,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_permission_keys,
)
from litellm.proxy.management_helpers.utils import (
get_new_internal_user_defaults,
@ -606,6 +607,11 @@ async def _set_object_permission(
return None
if data.object_permission is not None:
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=data.object_permission.mcp_tool_permissions,
existing_mcp_tool_permissions=None,
prisma_client=prisma_client,
)
created_object_permission: Final = await _table(ObjectPermissionRepository(prisma_client)).create(
data=data.object_permission.model_dump(exclude_none=True),
)

View file

@ -5,10 +5,13 @@ organizations, teams, and keys.
import json
from collections.abc import Mapping, Sequence
from collections.abc import Set as AbstractSet
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Optional
from fastapi import HTTPException, status
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -103,6 +106,11 @@ async def prepare_object_permission_upsert(
if existing_object_permission is not None
else {}
)
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=new_object_permission.get("mcp_tool_permissions"),
existing_mcp_tool_permissions=existing_fields.get("mcp_tool_permissions"),
prisma_client=prisma_client,
)
merged: Final[dict[str, object]] = {
**existing_fields,
**new_object_permission,
@ -194,6 +202,12 @@ async def _set_object_permission(
k: v for k, v in permission_data.items() if v is not None and k != "object_permission_id"
}
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=clean_data.get("mcp_tool_permissions"),
existing_mcp_tool_permissions=None,
prisma_client=prisma_client,
)
# Serialize mcp_tool_permissions to JSON string for GraphQL compatibility
if "mcp_tool_permissions" in clean_data:
clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"])
@ -226,7 +240,7 @@ def _mcp_server_identifier_matches(server: Any, identifier: str) -> bool:
async def _get_db_mcp_servers_by_identifiers(
identifiers: set[str],
identifiers: AbstractSet[str],
prisma_client: PrismaClient | None,
) -> "Sequence[prisma_models.LiteLLM_MCPServerTable]":
if prisma_client is None or not identifiers:
@ -245,7 +259,7 @@ async def _get_db_mcp_servers_by_identifiers(
async def _resolve_mcp_server_identifiers_to_ids(
identifiers: set[str],
identifiers: AbstractSet[str],
prisma_client: PrismaClient | None,
) -> dict[str, set[str]]:
"""
@ -286,6 +300,59 @@ async def _resolve_mcp_server_identifiers_to_ids(
return resolved
_MCP_TOOL_PERMISSIONS_ADAPTER: Final = TypeAdapter(dict[str, list[str] | None])
def _mcp_tool_permission_entries(raw: object) -> Mapping[str, frozenset[str]]:
parsed: Final[Mapping[str, Sequence[str] | None]] = (
_MCP_TOOL_PERMISSIONS_ADAPTER.validate_json(raw)
if isinstance(raw, str)
else _MCP_TOOL_PERMISSIONS_ADAPTER.validate_python(raw)
if isinstance(raw, Mapping)
else MappingProxyType({})
)
return MappingProxyType({identifier: frozenset(tools or ()) for identifier, tools in parsed.items()})
async def reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions: object,
existing_mcp_tool_permissions: object,
prisma_client: PrismaClient | None,
) -> None:
"""
A name or alias shared by several MCP servers cannot key ``mcp_tool_permissions``:
the read path unions the entry into every match, so no edit can narrow one of
those servers without also changing the other. An exact server_id is never
ambiguous, even when another server uses that string as its alias. Entries the
row already stores with the same tool list are left alone, so unrelated edits
to such an entity still succeed.
Raises HTTPException(400) naming the colliding servers.
"""
requested: Final = _mcp_tool_permission_entries(new_mcp_tool_permissions)
stored: Final = _mcp_tool_permission_entries(existing_mcp_tool_permissions)
resolved: Final = await _resolve_mcp_server_identifiers_to_ids(
identifiers=frozenset(identifier for identifier, tools in requested.items() if stored.get(identifier) != tools),
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_permissions key: {collisions}. "
"Key tool permissions by server_id when servers share a name or alias."
)
},
)
def _drop_stale_object_permission_mcp_servers(
object_permission: ObjectPermissionDict,
identifier_to_server_ids: dict[str, set[str]],

View file

@ -3917,6 +3917,7 @@ def _object_permission_mocks(mocker, existing_object_permission_id=None):
mock_prisma_client.db.litellm_objectpermissiontable.upsert = mocker.AsyncMock(
return_value=SimpleNamespace(object_permission_id="perm-new")
)
mock_prisma_client.db.litellm_mcpservertable.find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.update_data = mocker.AsyncMock(
return_value={"user_id": "target-user"}
)
@ -4146,6 +4147,7 @@ async def test_new_user_persists_the_requested_mcp_entitlement(mocker):
mock_prisma_client.db.litellm_objectpermissiontable.create = mocker.AsyncMock(
return_value=SimpleNamespace(object_permission_id="perm-created")
)
mock_prisma_client.db.litellm_mcpservertable.find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(
return_value=None
)

View file

@ -963,6 +963,7 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch):
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable.create = mock_create
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
async def _insert_data_side_effect(*args, **kwargs):
table_name = kwargs.get("table_name")

View file

@ -1183,6 +1183,37 @@ async def test_find_member_if_email_missing_row_raises_documented_400():
}
@pytest.mark.asyncio
async def test_new_organization_rejects_shared_alias_tool_permission_key():
"""/organization/new creates its permission row through its own helper, so the
ambiguous mcp_tool_permissions key check (LIT-4982) has to run there too."""
from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewOrganizationRequest
from litellm.proxy.management_endpoints.organization_endpoints import (
_set_object_permission,
)
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[
MagicMock(server_id="wiki-a-id", alias="wiki", server_name="wiki_a"),
MagicMock(server_id="wiki-b-id", alias="wiki", server_name="wiki_b"),
]
)
prisma_client.db.litellm_objectpermissiontable.create = AsyncMock()
data = NewOrganizationRequest(
organization_alias="org",
object_permission=LiteLLM_ObjectPermissionBase(mcp_tool_permissions={"wiki": ["ask_question"]}),
)
with pytest.raises(HTTPException) as exc_info:
await _set_object_permission(data=data, prisma_client=prisma_client)
assert exc_info.value.status_code == 400
assert "wiki-a-id" in str(exc_info.value.detail)
assert "wiki-b-id" in str(exc_info.value.detail)
prisma_client.db.litellm_objectpermissiontable.create.assert_not_called()
def test_v2_update_organization_is_in_openapi_schema():
"""PATCH /v2/organization/{organization_id} is documented in the generated OpenAPI spec."""
from fastapi import FastAPI

View file

@ -651,6 +651,7 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut
mock_db_client.db.litellm_objectpermissiontable = MagicMock()
mock_db_client.db.litellm_objectpermissiontable.create = mock_obj_perm_create
mock_db_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
# Mock model table
mock_db_client.db.litellm_modeltable = MagicMock()

View file

@ -17,6 +17,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
_resolve_team_allowed_mcp_servers,
_set_object_permission,
enforce_all_proxy_mcp_servers_grant_is_admin_only,
prepare_object_permission_upsert,
validate_key_mcp_servers_against_team,
validate_key_search_tools_against_team,
validate_key_vector_stores_against_team,
@ -41,6 +42,7 @@ async def test_set_object_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
data_json = {
@ -1349,6 +1351,123 @@ async def test_validate_key_update_sentinels_do_not_grandfather(monkeypatch):
assert exc_info.value.status_code == 403
# ---- Tests for rejecting ambiguous mcp_tool_permissions keys on write (LIT-4982) ----
_SHARED_ALIAS_DB_SERVERS = (
_make_mock_mcp_server("wiki-a-id", alias="wiki", server_name="wiki_a"),
_make_mock_mcp_server("wiki-b-id", alias="wiki", server_name="wiki_b"),
_make_mock_mcp_server("gh-a-id", alias="gh_a", server_name="github"),
_make_mock_mcp_server("gh-b-id", alias="gh_b", server_name="github"),
_make_mock_mcp_server("solo-id", alias="solo", server_name="Solo Server"),
_make_mock_mcp_server("shadow-id", alias="solo-id", server_name="shadow"),
)
def _make_ambiguity_prisma(existing_tool_permissions=None):
"""Mock prisma client whose MCP server table holds _SHARED_ALIAS_DB_SERVERS and whose
object permission row (if any) stores the given mcp_tool_permissions JSON string."""
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(_SHARED_ALIAS_DB_SERVERS))
mock_prisma.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=MagicMock(object_permission_id="perm-id")
)
existing_row = None
if existing_tool_permissions is not None:
existing_row = MagicMock()
existing_row.model_dump.return_value = {
"object_permission_id": "perm-id",
"mcp_tool_permissions": json.dumps(existing_tool_permissions),
}
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_row)
return mock_prisma
@pytest.mark.asyncio
@pytest.mark.parametrize(
"identifier, colliding_ids",
[("wiki", ("wiki-a-id", "wiki-b-id")), ("github", ("gh-a-id", "gh-b-id"))],
)
async def test_set_object_permission_rejects_shared_alias_or_name_tool_permission_key(identifier, colliding_ids):
"""An alias or server_name two servers share cannot key mcp_tool_permissions on
create: the write is rejected with 400 naming both servers and nothing is persisted."""
mock_prisma = _make_ambiguity_prisma()
data_json = {"object_permission": {"mcp_tool_permissions": {identifier: ["read_wiki_structure"]}}}
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 all(server_id in str(exc_info.value.detail) for server_id in colliding_ids)
mock_prisma.db.litellm_objectpermissiontable.create.assert_not_called()
@pytest.mark.asyncio
async def test_prepare_object_permission_upsert_rejects_shared_alias_tool_permission_key():
"""The update seam shared by key/team/org/user/customer/agent rejects a new
shared-alias key when the existing row does not already hold it."""
mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"solo-id": ["tool1"]})
with pytest.raises(HTTPException) as exc_info:
await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_permissions": {"wiki": ["ask_question"]}},
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)
@pytest.mark.asyncio
async def test_unambiguous_tool_permission_keys_persist_verbatim():
"""Exact ids (even when another server uses that id string as its alias),
unique aliases, and an id plus alias pointing at one server all still write."""
mock_prisma = _make_ambiguity_prisma()
tool_permissions = {
"wiki-a-id": ["ask_question"],
"wiki-b-id": ["read_wiki_structure"],
"solo-id": ["tool1"],
"solo": ["tool2"],
"Solo Server": ["tool3"],
}
upsert = await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_permissions": dict(tool_permissions)},
existing_object_permission_id=None,
prisma_client=mock_prisma,
)
assert json.loads(upsert.record["mcp_tool_permissions"]) == tool_permissions
@pytest.mark.asyncio
async def test_stored_ambiguous_tool_permission_key_is_grandfathered_until_changed():
"""A shared-alias entry already on the row may be re-sent unchanged so unrelated
edits succeed, but changing its tool list is rejected."""
mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"wiki": ["read_wiki_structure"]})
upsert = await prepare_object_permission_upsert(
new_object_permission={
"mcp_tool_permissions": {"wiki": ["read_wiki_structure"], "solo-id": ["tool1"]},
},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert json.loads(upsert.record["mcp_tool_permissions"]) == {
"wiki": ["read_wiki_structure"],
"solo-id": ["tool1"],
}
with pytest.raises(HTTPException) as exc_info:
await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_permissions": {"wiki": ["ask_question"]}},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert exc_info.value.status_code == 400
def test_object_permission_dict_mirrors_pydantic_model():
"""ObjectPermissionDict must stay field-for-field aligned with
LiteLLM_ObjectPermissionBase. If a new field is added to the Pydantic