fix(mcp): persist alias MCP grants verbatim instead of rewriting to local server ids

Since PR #29128, key create/update/regenerate resolved every
object_permission.mcp_servers entry against the saving instance's
DB + config registry and persisted the resolved server ids. For
config-loaded servers the id is derived from a hash of the regional
URL, so in a shared-database multi-region deployment the rewrite baked
one region's ids into the row and every other region denied the key.
Grants written before v1.88.0 kept the raw alias and kept working,
which is why only newly provisioned keys broke.

Keep the validation and the stale-entry drop (the LIT-3278 fix), but
persist the caller's original identifiers for everything that resolves.
Read-time expand_permission_list already maps a name to each region's
local server id.
This commit is contained in:
ryan-crabbe-berri 2026-09-01 08:10:36 -07:00
parent ec3f8183c3
commit 0d7035989c
2 changed files with 66 additions and 35 deletions

View file

@ -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,
)

View file

@ -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,44 @@ 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"]}
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value=set(),
)
@patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_alias_grant_expands_on_other_region_after_save(mock_access_groups, mock_allow_all):
"""Full cross-region flow: save a key on the west instance (alias resolves to
west's hash-derived id), then expand the persisted grant on the central
instance, whose registry maps the same alias to a different id."""
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")])
team_obj = _make_team_obj(mcp_servers=["west-id"])
object_permission = {"mcp_servers": ["github-mcp"]}
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
west_mgr,
):
await validate_key_mcp_servers_against_team(
object_permission=object_permission,
team_obj=team_obj,
)
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