mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
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:
parent
9832d6e4a6
commit
a46a076b2a
7 changed files with 229 additions and 2 deletions
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue