From 6eb0c782f158ac7745128752b217e4e67a9147a9 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 10 Mar 2026 17:20:05 -0700 Subject: [PATCH] feat: validate MCP server permissions on key create/update and scope UI dropdowns by team Adds backend validation to ensure keys can only be assigned MCP servers that their team has access to, and scopes the UI MCP server selector dropdown to the selected team's allowed servers. Co-Authored-By: Claude Opus 4.6 --- .../key_management_endpoints.py | 25 ++ .../mcp_management_endpoints.py | 118 ++++++++++ .../object_permission_utils.py | 197 +++++++++++++++- .../test_mcp_management_endpoints.py | 124 ++++++++++ .../test_object_permission_utils.py | 215 +++++++++++++++++- .../hooks/mcpServers/useMCPAccessGroups.ts | 6 +- .../hooks/mcpServers/useMCPServers.ts | 6 +- .../MCPServerSelector.tsx | 6 +- .../src/components/networking.tsx | 14 +- .../organisms/create_key_button.tsx | 1 + .../components/templates/key_edit_view.tsx | 2 + 11 files changed, 690 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 68f997e29cb..d385abd0e87 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, @@ -638,6 +639,13 @@ async def _common_key_generation_helper( # noqa: PLR0915 data_json.pop("tags") + # Validate MCP server permissions against the key's team + await validate_key_mcp_servers_against_team( + object_permission=data.object_permission, + team_obj=team_table, + user_api_key_dict=user_api_key_dict, + ) + data_json = await _set_object_permission( data_json=data_json, prisma_client=prisma_client, @@ -1947,6 +1955,23 @@ async def update_key_fn( # Set Management Endpoint Metadata Fields + # Validate MCP server permissions against the key's team on update + if data.object_permission is not None: + 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( + object_permission=data.object_permission, + team_obj=effective_team_obj, + user_api_key_dict=user_api_key_dict, + ) + non_default_values = await prepare_key_update_data( data=data, existing_key_row=existing_key_row ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f7a4cec301b..2b4cb6a4434 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -280,6 +280,94 @@ if MCP_AVAILABLE: allowed_routes = getattr(user_api_key_dict, "allowed_routes", None) return isinstance(allowed_routes, list) and len(allowed_routes) > 0 + async def _verify_team_membership( + user_api_key_dict: UserAPIKeyAuth, team_id: str + ) -> None: + """Verify the caller is a member of the specified team or is an admin.""" + if _user_has_admin_view(user_api_key_dict): + return + + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + check_db_only=True, + ) + if team_obj is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail={"error": f"Team {team_id} not found"}, + ) + + # Check if user is a member of the team + user_id = user_api_key_dict.user_id + is_member = False + for member in team_obj.members_with_roles or []: + member_user_id = getattr(member, "user_id", None) + if member_user_id == user_id: + is_member = True + break + + if not is_member: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": f"User is not a member of team {team_id}" + }, + ) + + async def _get_team_scoped_mcp_servers( + team_id: str, + user_api_key_dict: UserAPIKeyAuth, + ) -> List[LiteLLM_MCPServerTable]: + """Return MCP servers allowed by the team + allow_all_keys servers.""" + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.management_helpers.object_permission_utils import ( + get_team_mcp_permissions, + ) + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + check_db_only=True, + ) + + team_mcp = await get_team_mcp_permissions(team_obj) if team_obj else None + allow_all_ids = set(global_mcp_server_manager.get_allow_all_keys_server_ids()) + + if team_mcp is not None: + allowed_ids = set(team_mcp["mcp_servers"]) | allow_all_ids + else: + # Team has no MCP config - only allow_all_keys servers + allowed_ids = allow_all_ids + + all_servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered() + filtered = [s for s in all_servers if s.server_id in allowed_ids] + return _redact_mcp_credentials_list(filtered) + + async def _get_team_scoped_access_groups(team_id: str) -> dict: + """Return access groups available to the specified team.""" + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + team_obj = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + check_db_only=True, + ) + + if team_obj is None or getattr(team_obj, "object_permission", None) is None: + return {"access_groups": []} + + team_access_groups = team_obj.object_permission.mcp_access_groups or [] + return {"access_groups": sorted(team_access_groups)} + def _sanitize_mcp_server_for_virtual_key( mcp_server: LiteLLM_MCPServerTable, ) -> LiteLLM_MCPServerTable: @@ -456,6 +544,10 @@ if MCP_AVAILABLE: dependencies=[Depends(user_api_key_auth)], ) async def get_mcp_access_groups( + team_id: Optional[str] = Query( + None, + description="Filter access groups by team. Caller must be a member of the team or an admin.", + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -466,6 +558,11 @@ if MCP_AVAILABLE: ) from litellm.proxy.proxy_server import prisma_client + # If team_id is provided, return team-scoped access groups + if isinstance(team_id, str): + await _verify_team_membership(user_api_key_dict, team_id) + return await _get_team_scoped_access_groups(team_id) + access_groups = set() # Get from config-loaded servers @@ -486,6 +583,18 @@ if MCP_AVAILABLE: except Exception as e: verbose_proxy_logger.debug(f"Error getting MCP access groups: {e}") + # Filter for non-admins + from litellm.proxy.management_helpers.object_permission_utils import ( + get_allowed_mcp_access_groups_for_user, + ) + + if prisma_client is not None: + allowed = await get_allowed_mcp_access_groups_for_user( + user_api_key_dict, prisma_client + ) + if allowed is not None: + access_groups = access_groups & allowed + # Convert to sorted list access_groups_list = sorted(list(access_groups)) return {"access_groups": access_groups_list} @@ -561,6 +670,10 @@ if MCP_AVAILABLE: response_model=List[LiteLLM_MCPServerTable], ) async def fetch_all_mcp_servers( + team_id: Optional[str] = Query( + None, + description="Filter MCP servers by team. Caller must be a member of the team or an admin.", + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -571,6 +684,11 @@ if MCP_AVAILABLE: ``` """ + # If team_id is provided, verify membership and return team-scoped results + if isinstance(team_id, str): + await _verify_team_membership(user_api_key_dict, team_id) + return await _get_team_scoped_mcp_servers(team_id, user_api_key_dict) + user_mcp_management_mode = _get_user_mcp_management_mode() is_restricted_virtual_key = _is_restricted_virtual_key_request( user_api_key_dict diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 9670cdf330a..2af28790d2f 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -4,12 +4,20 @@ organizations, teams, and keys. """ import json -from litellm._uuid import uuid -from typing import Dict, Optional, Union +from typing import Dict, List, Optional, Set, Union + +from fastapi import HTTPException from litellm._logging import verbose_proxy_logger -from litellm.proxy.utils import PrismaClient +from litellm._uuid import uuid from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_TeamTableCachedObj, + UserAPIKeyAuth, +) +from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.utils import PrismaClient @@ -177,4 +185,185 @@ async def _set_object_permission( 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 get_team_mcp_permissions( + team_obj: LiteLLM_TeamTableCachedObj, +) -> Optional[Dict]: + """ + Returns the team's MCP permissions: {"mcp_servers": [...], "mcp_access_groups": [...]}. + Resolves access groups to server IDs. + Returns None if team has no object_permission. + """ + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + object_permission = getattr(team_obj, "object_permission", None) + if object_permission is None: + return None + + direct_servers: List[str] = object_permission.mcp_servers or [] + access_groups: List[str] = object_permission.mcp_access_groups or [] + tool_permission_keys: List[str] = list( + (object_permission.mcp_tool_permissions or {}).keys() + ) + + # Resolve access groups to server IDs + resolved_from_groups: List[str] = [] + if access_groups: + resolved_from_groups = ( + await MCPRequestHandler._get_mcp_servers_from_access_groups(access_groups) + ) + + all_servers: Set[str] = set(direct_servers + resolved_from_groups + tool_permission_keys) + + return { + "mcp_servers": list(all_servers), + "mcp_access_groups": access_groups, + } + + +async def validate_key_mcp_servers_against_team( + object_permission: Optional[LiteLLM_ObjectPermissionBase], + team_obj: Optional[LiteLLM_TeamTableCachedObj], + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """ + Validates that the requested MCP servers/access groups/tool permissions + on a key are allowed by the key's team. + + - Admin bypass: skips validation for proxy admins. + - Deny-by-default: if team has no MCP config, only allow_all_keys servers pass. + - Server validation: requested mcp_servers must be subset of team's expanded set + allow_all_keys. + - Access group validation: resolve requested mcp_access_groups to server IDs, + check those are subset of team's expanded server set. + - Tool permission validation: server IDs in mcp_tool_permissions keys must be in team's allowed set. + + Raises HTTPException(403) on failure. + """ + if object_permission is None: + return + + # Admin bypass + if _user_has_admin_view(user_api_key_dict): + return + + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + allow_all_keys_ids: Set[str] = set( + global_mcp_server_manager.get_allow_all_keys_server_ids() + ) + + requested_servers: List[str] = object_permission.mcp_servers or [] + requested_access_groups: List[str] = object_permission.mcp_access_groups or [] + requested_tool_perm_keys: List[str] = list( + (object_permission.mcp_tool_permissions or {}).keys() + ) + + # Nothing requested - nothing to validate + if not requested_servers and not requested_access_groups and not requested_tool_perm_keys: + return + + # Build the team's allowed server set + team_mcp = await get_team_mcp_permissions(team_obj) if team_obj else None + + if team_mcp is None: + # Team has no MCP config - deny-by-default: only allow_all_keys servers pass + disallowed = ( + set(requested_servers) | set(requested_tool_perm_keys) + ) - allow_all_keys_ids + if disallowed: + raise HTTPException( + status_code=403, + detail={ + "error": f"Team has no MCP configuration. The following MCP servers are not allowed: {sorted(disallowed)}" + }, + ) + # Also resolve requested access groups to server IDs and check those + if requested_access_groups: + resolved = await MCPRequestHandler._get_mcp_servers_from_access_groups( + requested_access_groups + ) + disallowed_from_groups = set(resolved) - allow_all_keys_ids + if disallowed_from_groups: + raise HTTPException( + status_code=403, + detail={ + "error": f"Team has no MCP configuration. Access groups resolve to unauthorized servers: {sorted(disallowed_from_groups)}" + }, + ) + return + + team_allowed_servers: Set[str] = set(team_mcp["mcp_servers"]) | allow_all_keys_ids + + # Validate direct server IDs + disallowed_servers = set(requested_servers) - team_allowed_servers + if disallowed_servers: + raise HTTPException( + status_code=403, + detail={ + "error": f"Key requests MCP servers not allowed by team: {sorted(disallowed_servers)}" + }, + ) + + # Validate access groups: resolve to server IDs, then check subset + if requested_access_groups: + resolved_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups( + requested_access_groups + ) + disallowed_from_groups = set(resolved_servers) - team_allowed_servers + if disallowed_from_groups: + raise HTTPException( + status_code=403, + detail={ + "error": f"Key's MCP access groups resolve to servers not allowed by team: {sorted(disallowed_from_groups)}" + }, + ) + + # Validate tool permission keys + disallowed_tool_keys = set(requested_tool_perm_keys) - team_allowed_servers + if disallowed_tool_keys: + raise HTTPException( + status_code=403, + detail={ + "error": f"Key's mcp_tool_permissions reference servers not allowed by team: {sorted(disallowed_tool_keys)}" + }, + ) + + +async def get_allowed_mcp_access_groups_for_user( + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, +) -> Optional[Set[str]]: + """ + Returns the set of MCP access groups the user can see (via their teams + key). + Returns None for admins (all groups allowed). + """ + if _user_has_admin_view(user_api_key_dict): + return None + + from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + build_effective_auth_contexts, + ) + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + allowed_groups: Set[str] = set() + auth_contexts = await build_effective_auth_contexts(user_api_key_dict) + + for auth_context in auth_contexts: + groups = await MCPRequestHandler.get_mcp_access_groups(auth_context) + allowed_groups.update(groups) + + return allowed_groups \ No newline at end of file diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index e81c6264f7b..10e2fdb05d0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1512,3 +1512,127 @@ class TestManagementPayloadValidation: assert len(result) == 1 assert result[0]["server_id"] == "server-1" assert result[0]["status"] == "healthy" + + +class TestTeamIdParam: + """Tests for team_id query param on MCP endpoints.""" + + @pytest.mark.asyncio + async def test_fetch_mcp_servers_team_id_returns_team_scoped_results(self): + """team_id param returns only the team's allowed servers + allow_all_keys.""" + mock_team = MagicMock() + mock_team.object_permission = MagicMock() + mock_team.object_permission.mcp_servers = ["server-1"] + mock_team.object_permission.mcp_access_groups = [] + mock_team.object_permission.mcp_tool_permissions = {} + mock_team.members_with_roles = [MagicMock(user_id="user-1")] + + admin_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-1", + ) + + server_1 = generate_mock_mcp_server_db_record(server_id="server-1") + server_2 = generate_mock_mcp_server_db_record(server_id="server-2") + server_public = generate_mock_mcp_server_db_record(server_id="server-public") + + mock_manager = MagicMock() + mock_manager.get_allow_all_keys_server_ids.return_value = ["server-public"] + mock_manager.get_all_mcp_servers_unfiltered = AsyncMock( + return_value=[server_1, server_2, server_public] + ) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._verify_team_membership", + AsyncMock(), + ), + patch( + "litellm.proxy.auth.auth_checks.get_team_object", + AsyncMock(return_value=mock_team), + ), + patch( + "litellm.proxy.management_helpers.object_permission_utils.get_team_mcp_permissions", + AsyncMock(return_value={"mcp_servers": ["server-1"], "mcp_access_groups": []}), + ), + patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + ), + patch( + "litellm.proxy.proxy_server.user_api_key_cache", + MagicMock(), + ), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_all_mcp_servers, + ) + + result = await fetch_all_mcp_servers( + team_id="team-1", + user_api_key_dict=admin_auth, + ) + + server_ids = {s.server_id for s in result} + assert "server-1" in server_ids + assert "server-public" in server_ids + assert "server-2" not in server_ids + + @pytest.mark.asyncio + async def test_fetch_mcp_servers_team_id_non_member_rejected(self): + """team_id param with non-member user -> rejected.""" + mock_team = MagicMock() + mock_team.members_with_roles = [MagicMock(user_id="other-user")] + + non_member_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-test", + user_id="user-not-member", + ) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + MagicMock(), + ), + patch( + "litellm.proxy.auth.auth_checks.get_team_object", + AsyncMock(return_value=mock_team), + ), + patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + ), + patch( + "litellm.proxy.proxy_server.user_api_key_cache", + MagicMock(), + ), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _verify_team_membership, + ) + + with pytest.raises(HTTPException) as exc: + await _verify_team_membership(non_member_auth, "team-1") + assert exc.value.status_code == 403 + + @pytest.mark.asyncio + async def test_fetch_mcp_servers_team_id_admin_bypasses_membership(self): + """Admin can use team_id param without being a team member.""" + admin_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-1", + ) + + # Should not raise for admin + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _verify_team_membership, + ) + + await _verify_team_membership(admin_auth, "team-1") diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 07d89035dce..9fa3492edcd 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -3,15 +3,24 @@ import os import sys import pytest +from fastapi import HTTPException sys.path.insert( 0, os.path.abspath("../../../..") ) -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionBase, + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + UserAPIKeyAuth, +) from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, + validate_key_mcp_servers_against_team, ) @@ -29,7 +38,7 @@ async def test_set_object_permission(): mock_prisma_client = MagicMock() mock_created_permission = MagicMock() mock_created_permission.object_permission_id = "test_perm_id_123" - + mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=mock_created_permission ) @@ -57,28 +66,218 @@ async def test_set_object_permission(): # Verify object_permission_id was added to result assert result["object_permission_id"] == "test_perm_id_123" - + # Verify object_permission was removed from result assert "object_permission" not in result - + # Verify create was called mock_prisma_client.db.litellm_objectpermissiontable.create.assert_called_once() - + # Verify the data passed to create excludes None values and object_permission_id call_args = mock_prisma_client.db.litellm_objectpermissiontable.create.call_args created_data = call_args.kwargs["data"] - + assert "object_permission_id" not in created_data assert "mcp_access_groups" not in created_data # None value should be excluded assert created_data["vector_stores"] == ["store_1", "store_2"] assert created_data["mcp_servers"] == ["server_a"] - + # Verify mcp_tool_permissions was serialized to JSON string assert isinstance(created_data["mcp_tool_permissions"], str) mcp_tools_parsed = json.loads(created_data["mcp_tool_permissions"]) assert mcp_tools_parsed == {"server_a": ["tool1", "tool2"]} - + # Verify other fields remain in result assert result["user_id"] == "test_user" assert result["models"] == ["gpt-4"] + +def _make_team_obj( + team_id: str = "team-1", + mcp_servers: list = None, + mcp_access_groups: list = None, + mcp_tool_permissions: dict = None, +) -> LiteLLM_TeamTableCachedObj: + """Helper to create a team object with object_permission.""" + obj_perm = None + if mcp_servers is not None or mcp_access_groups is not None: + obj_perm = LiteLLM_ObjectPermissionTable( + object_permission_id="team-perm-1", + mcp_servers=mcp_servers or [], + mcp_access_groups=mcp_access_groups or [], + mcp_tool_permissions=mcp_tool_permissions or {}, + vector_stores=[], + agents=[], + agent_access_groups=[], + ) + return LiteLLM_TeamTableCachedObj( + team_id=team_id, + object_permission=obj_perm, + ) + + +def _make_user_auth(role=LitellmUserRoles.INTERNAL_USER) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_role=role, + api_key="sk-test", + user_id="user-1", + ) + + +def _mcp_patches( + allow_all_keys_ids=None, + access_group_resolver=None, +): + """Return patch objects for the lazy-imported MCP dependencies.""" + mock_handler = MagicMock() + mock_handler._get_mcp_servers_from_access_groups = AsyncMock( + side_effect=access_group_resolver or (lambda groups: []) + ) + + mock_manager = MagicMock() + mock_manager.get_allow_all_keys_server_ids.return_value = allow_all_keys_ids or [] + + return ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler", + mock_handler, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + mock_manager, + ), + ) + + +@pytest.mark.asyncio +async def test_validate_mcp_servers_allowed_by_team(): + """Key creation allowed when MCP server is in team's list.""" + team_obj = _make_team_obj(mcp_servers=["server-1", "server-2"]) + obj_perm = LiteLLM_ObjectPermissionBase(mcp_servers=["server-1"]) + + p1, p2 = _mcp_patches() + with p1, p2: + # Should not raise + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(), + ) + + +@pytest.mark.asyncio +async def test_validate_mcp_servers_rejected_when_not_in_team(): + """Key creation rejected when MCP server not in team's list.""" + team_obj = _make_team_obj(mcp_servers=["server-1"]) + obj_perm = LiteLLM_ObjectPermissionBase(mcp_servers=["server-1", "server-999"]) + + p1, p2 = _mcp_patches() + with p1, p2: + with pytest.raises(HTTPException) as exc: + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(), + ) + assert exc.value.status_code == 403 + assert "server-999" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_validate_mcp_servers_admin_bypasses(): + """Admin bypasses validation.""" + team_obj = _make_team_obj(mcp_servers=["server-1"]) + obj_perm = LiteLLM_ObjectPermissionBase(mcp_servers=["server-999"]) + + # Admin should not raise even with disallowed server + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(role=LitellmUserRoles.PROXY_ADMIN), + ) + + +@pytest.mark.asyncio +async def test_validate_mcp_servers_deny_by_default_no_team_config(): + """Team with no MCP config -> deny-by-default (non-allow_all_keys blocked).""" + team_obj = _make_team_obj() # No object_permission + obj_perm = LiteLLM_ObjectPermissionBase(mcp_servers=["server-1"]) + + p1, p2 = _mcp_patches() + with p1, p2: + with pytest.raises(HTTPException) as exc: + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(), + ) + assert exc.value.status_code == 403 + assert "no MCP configuration" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_validate_mcp_servers_allow_all_keys_always_passes(): + """allow_all_keys server always passes even without team config.""" + team_obj = _make_team_obj() # No object_permission + obj_perm = LiteLLM_ObjectPermissionBase(mcp_servers=["public-server"]) + + p1, p2 = _mcp_patches(allow_all_keys_ids=["public-server"]) + with p1, p2: + # Should not raise + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(), + ) + + +@pytest.mark.asyncio +async def test_validate_mcp_access_groups_resolve_to_unauthorized_servers(): + """Access groups resolved to unauthorized servers -> rejected.""" + team_obj = _make_team_obj(mcp_servers=["server-1"]) + obj_perm = LiteLLM_ObjectPermissionBase(mcp_access_groups=["group-evil"]) + + p1, p2 = _mcp_patches( + access_group_resolver=lambda groups: ( + ["server-unauthorized"] if "group-evil" in groups else [] + ), + ) + with p1, p2: + with pytest.raises(HTTPException) as exc: + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(), + ) + assert exc.value.status_code == 403 + assert "server-unauthorized" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_validate_mcp_tool_permissions_unauthorized_server(): + """mcp_tool_permissions with unauthorized server key -> rejected.""" + team_obj = _make_team_obj(mcp_servers=["server-1"]) + obj_perm = LiteLLM_ObjectPermissionBase( + mcp_tool_permissions={"server-1": ["tool1"], "server-999": ["tool2"]} + ) + + p1, p2 = _mcp_patches() + with p1, p2: + with pytest.raises(HTTPException) as exc: + await validate_key_mcp_servers_against_team( + object_permission=obj_perm, + team_obj=team_obj, + user_api_key_dict=_make_user_auth(), + ) + assert exc.value.status_code == 403 + assert "server-999" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_validate_mcp_none_object_permission_passes(): + """No object_permission -> no validation needed.""" + await validate_key_mcp_servers_against_team( + object_permission=None, + team_obj=_make_team_obj(), + user_api_key_dict=_make_user_auth(), + ) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts index 0e88b62b0f3..603dedbdf6e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPAccessGroups.ts @@ -4,11 +4,11 @@ import { fetchMCPAccessGroups } from "@/components/networking"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; const mcpAccessGroupsKeys = createQueryKeys("mcpAccessGroups"); -export const useMCPAccessGroups = () => { +export const useMCPAccessGroups = (teamId?: string) => { const { accessToken } = useAuthorized(); return useQuery({ - queryKey: mcpAccessGroupsKeys.list({}), - queryFn: async () => await fetchMCPAccessGroups(accessToken!), + queryKey: mcpAccessGroupsKeys.list({ teamId }), + queryFn: async () => await fetchMCPAccessGroups(accessToken!, teamId), enabled: Boolean(accessToken), }); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts index 8746baae148..9aff24f13b4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpServers/useMCPServers.ts @@ -6,11 +6,11 @@ import useAuthorized from "../useAuthorized"; const mcpServersKeys = createQueryKeys("mcpServers"); -export const useMCPServers = () => { +export const useMCPServers = (teamId?: string) => { const { accessToken } = useAuthorized(); return useQuery({ - queryKey: mcpServersKeys.list({}), - queryFn: async () => await fetchMCPServers(accessToken!), + queryKey: mcpServersKeys.list({ teamId }), + queryFn: async () => await fetchMCPServers(accessToken!, teamId), enabled: !!accessToken, }); }; diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx index d94a80e502c..63666847b02 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPServerSelector.tsx @@ -13,6 +13,7 @@ interface MCPServerSelectorProps { accessToken: string; placeholder?: string; disabled?: boolean; + teamId?: string; } const MCPServerSelector: React.FC = ({ @@ -22,9 +23,10 @@ const MCPServerSelector: React.FC = ({ accessToken, placeholder = "Select MCP servers", disabled = false, + teamId, }) => { - const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(); - const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups(); + const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(teamId); + const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups(teamId); const loading = serversLoading || groupsLoading; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 3ec877f03f4..5d6218e368e 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -6286,10 +6286,13 @@ export const fetchDiscoverableMCPServers = async (accessToken: string) => { } }; -export const fetchMCPServers = async (accessToken: string) => { +export const fetchMCPServers = async (accessToken: string, teamId?: string) => { try { // Construct base URL - const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server` : `/v1/mcp/server`; + let url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server` : `/v1/mcp/server`; + if (teamId) { + url += `?team_id=${encodeURIComponent(teamId)}`; + } console.log("Fetching MCP servers from:", url); @@ -6355,10 +6358,13 @@ export const fetchMCPServerHealth = async (accessToken: string, serverIds?: stri } }; -export const fetchMCPAccessGroups = async (accessToken: string) => { +export const fetchMCPAccessGroups = async (accessToken: string, teamId?: string) => { try { // Construct base URL - const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/access_groups` : `/v1/mcp/access_groups`; + let url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/access_groups` : `/v1/mcp/access_groups`; + if (teamId) { + url += `?team_id=${encodeURIComponent(teamId)}`; + } console.log("Fetching MCP access groups from:", url); diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 78ff01d6419..de1978cf2e5 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -1324,6 +1324,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp value={form.getFieldValue("allowed_mcp_servers_and_groups")} accessToken={accessToken} placeholder="Select MCP servers or access groups (optional)" + teamId={selectedCreateKeyTeam?.team_id} /> diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index b6c00577c9b..28860be5f88 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -87,6 +87,7 @@ export function KeyEditView({ }: KeyEditViewProps) { const canEditGuardrails = premiumUser || (userRole != null && rolesWithWriteAccess.includes(userRole)); const [form] = Form.useForm(); + const watchedTeamId = Form.useWatch("team_id", form); const [promptsList, setPromptsList] = useState([]); const [tagsList, setTagsList] = useState>({}); const team = teams?.find((team) => team.team_id === keyData.team_id); @@ -574,6 +575,7 @@ export function KeyEditView({ value={form.getFieldValue("mcp_servers_and_groups")} accessToken={accessToken || ""} placeholder="Select MCP servers or access groups (optional)" + teamId={watchedTeamId} />