mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
[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:
parent
007bea10b8
commit
2868796980
3 changed files with 300 additions and 4 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
)
|
||||
},
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue