[Fix] Validate key MCP server assignments against team permissions

When a key has a team, validate that any MCP servers/access groups in
the key's object_permission are a subset of what the team allows. This
prevents users from assigning MCP servers to keys for teams that don't
have access to those servers.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-03-03 22:48:04 -08:00
parent 007bea10b8
commit 2868796980
3 changed files with 300 additions and 4 deletions

View file

@ -64,6 +64,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
_set_object_permission,
attach_object_permission_to_dict,
handle_update_object_permission_common,
validate_key_mcp_servers_against_team,
)
from litellm.proxy.management_helpers.team_member_permission_checks import (
TeamMemberPermissionChecks,
@ -634,6 +635,11 @@ async def _common_key_generation_helper( # noqa: PLR0915
data_json.pop("tags")
await validate_key_mcp_servers_against_team(
data_json.get("object_permission"),
team_table,
)
data_json = await _set_object_permission(
data_json=data_json,
prisma_client=prisma_client,
@ -1749,7 +1755,7 @@ async def _process_single_key_update(
"/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)]
)
@management_endpoint_wrapper
async def update_key_fn(
async def update_key_fn( # noqa: PLR0915
request: Request,
data: UpdateKeyRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -1943,6 +1949,22 @@ async def update_key_fn(
# Set Management Endpoint Metadata Fields
# Validate MCP servers in object_permission against the key's effective team
if "object_permission" in data_json:
effective_team_id = data.team_id or existing_key_row.team_id
effective_team_obj = team_obj
if effective_team_obj is None and effective_team_id is not None:
effective_team_obj = await get_team_object(
team_id=effective_team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
await validate_key_mcp_servers_against_team(
data_json.get("object_permission"),
effective_team_obj,
)
non_default_values = await prepare_key_update_data(
data=data, existing_key_row=existing_key_row
)

View file

@ -5,9 +5,12 @@ organizations, teams, and keys.
import json
from litellm._uuid import uuid
from typing import Dict, Optional, Union
from typing import Any, Dict, List, Optional, Union
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
from litellm.proxy.utils import PrismaClient
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -174,7 +177,112 @@ async def _set_object_permission(
created_permission = await prisma_client.db.litellm_objectpermissiontable.create(
data=clean_data
)
data_json["object_permission_id"] = created_permission.object_permission_id
data_json.pop("object_permission")
return data_json
return data_json
async def validate_key_mcp_servers_against_team(
object_permission: Optional[Union[Dict, Any]],
team_obj: Optional[LiteLLM_TeamTableCachedObj],
) -> None:
"""
Validate that a key's requested MCP servers/access groups are allowed by its team.
Mirrors the runtime intersection logic: only restricts when the team has restrictions
configured (non-empty). allow_all_keys servers always pass validation.
Raises HTTPException(403) if the key requests MCP servers or access groups that
the team does not allow.
"""
if object_permission is None or team_obj is None:
return
# Extract key's requested MCP servers and access groups
if isinstance(object_permission, dict):
key_mcp_servers: List[str] = object_permission.get("mcp_servers") or []
key_mcp_access_groups: List[str] = object_permission.get("mcp_access_groups") or []
else:
key_mcp_servers = getattr(object_permission, "mcp_servers", None) or []
key_mcp_access_groups = getattr(object_permission, "mcp_access_groups", None) or []
if not key_mcp_servers and not key_mcp_access_groups:
return
team_object_permission = team_obj.object_permission
if team_object_permission is None:
# Team has no MCP config - no restriction
return
# Build the team's allowed server set from direct servers + access group resolution + tool permissions
team_allowed_servers: List[str] = []
if team_object_permission.mcp_servers:
team_allowed_servers.extend(team_object_permission.mcp_servers)
if team_object_permission.mcp_access_groups:
try:
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
resolved = await MCPRequestHandler._get_mcp_servers_from_access_groups(
team_object_permission.mcp_access_groups
)
team_allowed_servers.extend(resolved)
except Exception as e:
verbose_proxy_logger.warning(
f"validate_key_mcp_servers_against_team: failed to resolve team MCP access groups: {e}"
)
if team_object_permission.mcp_tool_permissions:
team_allowed_servers.extend(team_object_permission.mcp_tool_permissions.keys())
# Get allow_all_keys server IDs - these bypass per-key restrictions
allow_all_server_ids: set = set()
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
allow_all_server_ids = set(global_mcp_server_manager.get_allow_all_keys_server_ids())
except Exception as e:
verbose_proxy_logger.warning(
f"validate_key_mcp_servers_against_team: failed to get allow_all_keys servers: {e}"
)
# Validate key's mcp_servers only when the team has server restrictions configured
if key_mcp_servers and team_allowed_servers:
team_allowed_set = set(team_allowed_servers)
disallowed = [
s for s in key_mcp_servers
if s not in team_allowed_set and s not in allow_all_server_ids
]
if disallowed:
raise HTTPException(
status_code=403,
detail={
"error": (
f"MCP servers not allowed by team: {disallowed}. "
f"Team allows: {sorted(team_allowed_set)}"
)
},
)
# Validate key's mcp_access_groups only when the team has access group restrictions configured
if key_mcp_access_groups and team_object_permission.mcp_access_groups:
team_access_group_set = set(team_object_permission.mcp_access_groups)
disallowed_groups = [
g for g in key_mcp_access_groups if g not in team_access_group_set
]
if disallowed_groups:
raise HTTPException(
status_code=403,
detail={
"error": (
f"MCP access groups not allowed by team: {disallowed_groups}. "
f"Team allows: {sorted(team_access_group_set)}"
)
},
)

View file

@ -6455,3 +6455,169 @@ class TestValidateKeyAliasFormat:
_validate_key_alias_format(alias)
assert str(exc.value.code) == "400"
assert "Invalid key_alias format" in str(exc.value.message)
# ============================================================
# Tests for validate_key_mcp_servers_against_team
# ============================================================
def _make_team_obj_with_mcp_servers(mcp_servers=None, mcp_access_groups=None):
"""Build a LiteLLM_TeamTableCachedObj with an object_permission."""
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTableCachedObj
if mcp_servers is None and mcp_access_groups is None:
obj_perm = None
else:
obj_perm = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-team-1",
mcp_servers=mcp_servers or [],
mcp_access_groups=mcp_access_groups or [],
)
return LiteLLM_TeamTableCachedObj(
team_id="team-mcp-1",
object_permission=obj_perm,
)
@pytest.mark.asyncio
async def test_mcp_validation_key_creation_rejects_disallowed_server():
"""Key creation with an MCP server not in the team's allowed list raises 403."""
from fastapi import HTTPException
from litellm.proxy.management_helpers.object_permission_utils import (
validate_key_mcp_servers_against_team,
)
team_obj = _make_team_obj_with_mcp_servers(mcp_servers=["server-allowed"])
object_permission = {"mcp_servers": ["server-not-allowed"]}
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(object_permission, team_obj)
assert exc_info.value.status_code == 403
assert "server-not-allowed" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_mcp_validation_key_creation_allows_permitted_server():
"""Key creation with an MCP server in the team's allowed list succeeds."""
from litellm.proxy.management_helpers.object_permission_utils import (
validate_key_mcp_servers_against_team,
)
team_obj = _make_team_obj_with_mcp_servers(mcp_servers=["server-allowed"])
object_permission = {"mcp_servers": ["server-allowed"]}
# Should not raise
await validate_key_mcp_servers_against_team(object_permission, team_obj)
@pytest.mark.asyncio
async def test_mcp_validation_no_restriction_when_team_has_no_mcp_config():
"""When the team has no object_permission, any MCP server is allowed."""
from litellm.proxy.management_helpers.object_permission_utils import (
validate_key_mcp_servers_against_team,
)
team_obj = _make_team_obj_with_mcp_servers() # no object_permission
object_permission = {"mcp_servers": ["any-server"]}
# Should not raise
await validate_key_mcp_servers_against_team(object_permission, team_obj)
@pytest.mark.asyncio
async def test_mcp_validation_allow_all_keys_server_always_passes():
"""A server with allow_all_keys=True passes even when team restricts servers."""
from unittest.mock import MagicMock, patch
from litellm.proxy.management_helpers.object_permission_utils import (
validate_key_mcp_servers_against_team,
)
team_obj = _make_team_obj_with_mcp_servers(mcp_servers=["server-allowed"])
object_permission = {"mcp_servers": ["server-allow-all"]}
mock_manager = MagicMock()
mock_manager.get_allow_all_keys_server_ids.return_value = ["server-allow-all"]
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
):
# Should not raise because "server-allow-all" is in allow_all_keys
await validate_key_mcp_servers_against_team(object_permission, team_obj)
@pytest.mark.asyncio
async def test_mcp_validation_key_update_rejects_disallowed_server(monkeypatch):
"""Updating a key's object_permission with a disallowed MCP server raises 403."""
from unittest.mock import AsyncMock, MagicMock
from fastapi import HTTPException
from litellm.proxy._types import (
LiteLLM_ObjectPermissionBase,
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamTableCachedObj,
LiteLLM_VerificationToken,
)
from litellm.proxy.management_endpoints.key_management_endpoints import update_key_fn
# Set up prisma mock
mock_prisma_client = AsyncMock()
existing_key = LiteLLM_VerificationToken(
token="hashed-key",
team_id="team-mcp-1",
user_id="user-1",
)
mock_prisma_client.get_data = AsyncMock(return_value=existing_key)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
# Team allows only "server-allowed"
team_obj = LiteLLM_TeamTableCachedObj(
team_id="team-mcp-1",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-team-1",
mcp_servers=["server-allowed"],
),
)
async def mock_get_team_object(**kwargs):
return team_obj
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object",
mock_get_team_object,
)
request = MagicMock()
request.body = AsyncMock(return_value=b"{}")
update_data = UpdateKeyRequest(
key="sk-test",
object_permission=LiteLLM_ObjectPermissionBase(
mcp_servers=["server-not-allowed"],
),
)
with pytest.raises(Exception) as exc_info:
await update_key_fn(
request=request,
data=update_data,
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="user-1",
),
)
exc = exc_info.value
# The 403 HTTPException is re-raised as a ProxyException
assert "403" in str(exc) or (hasattr(exc, "code") and str(exc.code) == "403")