mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #39119 from BerriAI/litellm_fix_mcp_alias_grant_persistence
fix(mcp): persist alias MCP grants verbatim instead of rewriting to local server ids
This commit is contained in:
commit
4db165379a
4 changed files with 78 additions and 35 deletions
|
|
@ -286,7 +286,7 @@ async def _resolve_mcp_server_identifiers_to_ids(
|
|||
return resolved
|
||||
|
||||
|
||||
def _rewrite_object_permission_mcp_servers(
|
||||
def _drop_stale_object_permission_mcp_servers(
|
||||
object_permission: ObjectPermissionDict,
|
||||
identifier_to_server_ids: dict[str, set[str]],
|
||||
) -> None:
|
||||
|
|
@ -294,16 +294,18 @@ def _rewrite_object_permission_mcp_servers(
|
|||
if not isinstance(mcp_servers, list):
|
||||
return
|
||||
|
||||
normalized_servers: Final[list[str]] = []
|
||||
for identifier in mcp_servers:
|
||||
if identifier == SpecialMCPServerNames.no_mcp_servers.value:
|
||||
normalized_servers.append(SpecialMCPServerNames.no_mcp_servers.value)
|
||||
continue
|
||||
normalized_servers.extend(sorted(identifier_to_server_ids.get(identifier, [])))
|
||||
object_permission["mcp_servers"] = _dedupe_preserving_order(normalized_servers)
|
||||
# Persist original identifiers, never resolved ids: shared-DB multi-region
|
||||
# instances each expand a name/alias to their own local server id at read
|
||||
# time. Only entries resolving to nothing (deleted servers, typos) drop.
|
||||
kept_servers: Final = [
|
||||
identifier
|
||||
for identifier in mcp_servers
|
||||
if identifier == SpecialMCPServerNames.no_mcp_servers.value or identifier_to_server_ids.get(identifier)
|
||||
]
|
||||
object_permission["mcp_servers"] = _dedupe_preserving_order(kept_servers)
|
||||
|
||||
|
||||
def _rewrite_object_permission_mcp_tool_permissions(
|
||||
def _drop_stale_object_permission_mcp_tool_permissions(
|
||||
object_permission: ObjectPermissionDict,
|
||||
identifier_to_server_ids: dict[str, set[str]],
|
||||
) -> None:
|
||||
|
|
@ -311,31 +313,25 @@ def _rewrite_object_permission_mcp_tool_permissions(
|
|||
if not isinstance(mcp_tool_permissions, dict):
|
||||
return
|
||||
|
||||
normalized_tool_permissions: Final[dict[str, list[str]]] = {}
|
||||
for identifier, tools in mcp_tool_permissions.items():
|
||||
if not isinstance(tools, list):
|
||||
tools = []
|
||||
for server_id in sorted(identifier_to_server_ids.get(identifier, [])):
|
||||
normalized_tool_permissions.setdefault(server_id, [])
|
||||
normalized_tool_permissions[server_id].extend(tools)
|
||||
|
||||
object_permission["mcp_tool_permissions"] = {
|
||||
server_id: _dedupe_preserving_order(tools) for server_id, tools in normalized_tool_permissions.items()
|
||||
identifier: _dedupe_preserving_order(tools if isinstance(tools, list) else [])
|
||||
for identifier, tools in mcp_tool_permissions.items()
|
||||
if identifier_to_server_ids.get(identifier)
|
||||
}
|
||||
|
||||
|
||||
def _rewrite_object_permission_mcp_identifiers(
|
||||
def _drop_stale_object_permission_mcp_identifiers(
|
||||
object_permission: ObjectPermissionDict | None,
|
||||
identifier_to_server_ids: dict[str, set[str]],
|
||||
) -> None:
|
||||
if not object_permission or not isinstance(object_permission, dict):
|
||||
return
|
||||
|
||||
_rewrite_object_permission_mcp_servers(
|
||||
_drop_stale_object_permission_mcp_servers(
|
||||
object_permission=object_permission,
|
||||
identifier_to_server_ids=identifier_to_server_ids,
|
||||
)
|
||||
_rewrite_object_permission_mcp_tool_permissions(
|
||||
_drop_stale_object_permission_mcp_tool_permissions(
|
||||
object_permission=object_permission,
|
||||
identifier_to_server_ids=identifier_to_server_ids,
|
||||
)
|
||||
|
|
@ -615,7 +611,7 @@ async def validate_key_mcp_servers_against_team(
|
|||
"validate_key_mcp_servers_against_team: ignoring stale MCP server identifiers (no longer in registry or DB): %s",
|
||||
sorted(stale_identifiers),
|
||||
)
|
||||
_rewrite_object_permission_mcp_identifiers(
|
||||
_drop_stale_object_permission_mcp_identifiers(
|
||||
object_permission=object_permission,
|
||||
identifier_to_server_ids=identifier_to_server_ids,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -30,6 +30,34 @@ def _key(client: McpClient, resources: ResourceManager, *, mcp_servers: list[str
|
|||
return key
|
||||
|
||||
|
||||
class TestMcpKeyGrantByAlias:
|
||||
def test_alias_grant_persists_verbatim_and_lists_tools(
|
||||
self,
|
||||
client: McpClient,
|
||||
resources: ResourceManager,
|
||||
) -> None:
|
||||
"""A key granted an MCP server by its alias must store the alias, not the
|
||||
resolved server_id: in a shared-DB multi-region deployment each instance
|
||||
derives a different id for the same config server, so only the alias
|
||||
grants access on every region. The same key must still see the server's
|
||||
tools, proving the alias grant is honored at request time."""
|
||||
server_id = register_datadog_mcp(client, resources)
|
||||
client.await_registered(server_id)
|
||||
alias = next(row.alias for row in client.registered_servers() if row.server_id == server_id)
|
||||
assert alias, f"registered server {server_id} has no alias to grant by"
|
||||
|
||||
key = _key(client, resources, mcp_servers=[alias])
|
||||
|
||||
stored = client.proxy.key_info(key).object_permission
|
||||
assert stored is not None and stored.mcp_servers == [alias], (
|
||||
f"alias grant was rewritten before persisting (expected [{alias!r}]): "
|
||||
f"{stored.mcp_servers if stored else None}. A stored server_id is region-local "
|
||||
f"and breaks the grant on every other instance sharing this database"
|
||||
)
|
||||
|
||||
_ = client.await_tool(key, server_id, SEARCH_LOGS_TOOL)
|
||||
|
||||
|
||||
class TestMcpKeyWithoutAccessIsDenied:
|
||||
@pytest.mark.covers("mcp.list_tools.api_key.denied_without_permission")
|
||||
def test_list_tools_denied_without_permission(
|
||||
|
|
|
|||
|
|
@ -114,6 +114,7 @@ class KeyInfo(BaseModel):
|
|||
budget_id: str | None = None
|
||||
litellm_budget_table: LiteLLMBudgetTable | None = None
|
||||
budget_limits: list[BudgetWindowState] | None = None
|
||||
object_permission: ObjectPermission | None = None
|
||||
|
||||
|
||||
class KeyInfoResponse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -1,11 +1,9 @@
|
|||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionBase,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -13,10 +11,10 @@ from litellm.proxy._types import (
|
|||
SpecialMCPServerName,
|
||||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_drop_stale_object_permission_mcp_servers,
|
||||
_extract_requested_mcp_access_groups,
|
||||
_extract_requested_mcp_server_ids,
|
||||
_resolve_team_allowed_mcp_servers,
|
||||
_rewrite_object_permission_mcp_servers,
|
||||
_set_object_permission,
|
||||
enforce_all_proxy_mcp_servers_grant_is_admin_only,
|
||||
validate_key_mcp_servers_against_team,
|
||||
|
|
@ -153,10 +151,10 @@ def test_extract_requested_mcp_server_ids_excludes_no_mcp_servers_sentinel():
|
|||
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1"}
|
||||
|
||||
|
||||
def test_rewrite_object_permission_mcp_servers_preserves_sentinel():
|
||||
obj_perm = {"mcp_servers": ["no-mcp-servers", "alias-1"]}
|
||||
_rewrite_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}})
|
||||
assert obj_perm["mcp_servers"] == ["no-mcp-servers", "server-1"]
|
||||
def test_drop_stale_object_permission_mcp_servers_preserves_sentinel_and_alias():
|
||||
obj_perm = {"mcp_servers": ["no-mcp-servers", "alias-1", "gone-id"]}
|
||||
_drop_stale_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}, "gone-id": set()})
|
||||
assert obj_perm["mcp_servers"] == ["no-mcp-servers", "alias-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -692,9 +690,10 @@ async def test_validate_mcp_server_alias_outside_team_scope_raises(
|
|||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
)
|
||||
async def test_validate_mcp_server_alias_is_normalized_before_save(
|
||||
mock_access_groups, mock_allow_all
|
||||
):
|
||||
async def test_validate_mcp_server_alias_persists_verbatim(mock_access_groups, mock_allow_all):
|
||||
"""Regression for the multi-region shared-DB setup: an alias grant must be
|
||||
stored as the alias, so every instance can expand it to its own local id.
|
||||
Rewriting to this instance's server_id breaks access on the other region."""
|
||||
team_obj = _make_team_obj(mcp_servers=["allowed-server-id"])
|
||||
object_permission = {
|
||||
"mcp_servers": ["allowed-alias"],
|
||||
|
|
@ -706,8 +705,27 @@ async def test_validate_mcp_server_alias_is_normalized_before_save(
|
|||
team_obj=team_obj,
|
||||
)
|
||||
|
||||
assert object_permission["mcp_servers"] == ["allowed-server-id"]
|
||||
assert object_permission["mcp_tool_permissions"] == {"allowed-server-id": ["tool1"]}
|
||||
assert object_permission["mcp_servers"] == ["allowed-alias"]
|
||||
assert object_permission["mcp_tool_permissions"] == {"Allowed Server": ["tool1"]}
|
||||
|
||||
|
||||
def test_alias_grant_expands_on_other_region_after_save():
|
||||
"""Cross-region flow: the west instance saves an alias grant (its resolver maps
|
||||
the alias to west's hash-derived id), then the central instance, whose registry
|
||||
maps the same alias to a different id, expands the persisted grant. Rewriting
|
||||
to west's id at save time is exactly the regression this guards against."""
|
||||
west_mgr = _make_mock_mcp_manager(servers=[_make_mock_mcp_server("west-id", alias="github-mcp")])
|
||||
central_mgr = _make_mock_mcp_manager(servers=[_make_mock_mcp_server("central-id", alias="github-mcp")])
|
||||
|
||||
object_permission = {"mcp_servers": ["github-mcp"]}
|
||||
_drop_stale_object_permission_mcp_servers(object_permission, {"github-mcp": {"west-id"}})
|
||||
assert object_permission["mcp_servers"] == ["github-mcp"]
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
||||
expand = MCPServerManager.expand_permission_list
|
||||
assert expand(west_mgr, object_permission["mcp_servers"]) == ["west-id"]
|
||||
assert expand(central_mgr, object_permission["mcp_servers"]) == ["central-id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue