diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 4d44006369b..4edec6fe9bc 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 ) diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 9670cdf330a..069509a8fd3 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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 \ No newline at end of file + 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)}" + ) + }, + ) \ No newline at end of file diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 11c80351839..13cc3983c20 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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")