Merge pull request #23323 from BerriAI/litellm_mcp_permissions_yj

[Feature] MCP Server Team-Scoped Filtering for Key Creation
This commit is contained in:
yuneng-jiang
2026-03-11 00:20:19 -07:00
committed by GitHub
10 changed files with 804 additions and 30 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,
@@ -638,6 +639,12 @@ async def _common_key_generation_helper( # noqa: PLR0915
data_json.pop("tags")
# Validate MCP servers in object_permission are within team scope
await validate_key_mcp_servers_against_team(
object_permission=data_json.get("object_permission"),
team_obj=team_table,
)
data_json = await _set_object_permission(
data_json=data_json,
prisma_client=prisma_client,
@@ -1947,6 +1954,27 @@ async def update_key_fn(
# Set Management Endpoint Metadata Fields
# Validate MCP servers in object_permission against the effective team
if data.object_permission is not None:
effective_team_obj = team_obj
# If team_id isn't being changed, resolve the existing key's team
if effective_team_obj is None and existing_key_row.team_id:
effective_team_obj = await get_team_object(
team_id=existing_key_row.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
object_permission_dict = (
data.object_permission.model_dump()
if hasattr(data.object_permission, "model_dump")
else data.object_permission
)
await validate_key_mcp_servers_against_team(
object_permission=object_permission_dict,
team_obj=effective_team_obj,
)
non_default_values = await prepare_key_update_data(
data=data, existing_key_row=existing_key_row
)
@@ -615,6 +615,46 @@ if MCP_AVAILABLE:
return "view_all"
return "restricted"
async def _get_team_scoped_mcp_server_list(
team_id: str,
) -> List[LiteLLM_MCPServerTable]:
"""
Return MCP servers scoped to a team: team's allowed servers + allow_all_keys servers.
Used by the Create Key UI to populate the MCP server dropdown.
"""
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.management_helpers.object_permission_utils import (
_get_allow_all_keys_server_ids,
_get_team_allowed_mcp_servers,
)
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_server_ids = await _get_team_allowed_mcp_servers(team_obj)
allow_all_server_ids = _get_allow_all_keys_server_ids()
all_allowed_ids = team_server_ids | allow_all_server_ids
if not all_allowed_ids:
return []
# Collect servers from registry
servers: List[LiteLLM_MCPServerTable] = []
for server_id in all_allowed_ids:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if server is not None:
mcp_server_table = global_mcp_server_manager._build_mcp_server_table(
server
)
servers.append(mcp_server_table)
return _redact_mcp_credentials_list(servers)
@router.get(
"/server",
description="Returns the mcp server list with associated teams",
@@ -623,38 +663,88 @@ if MCP_AVAILABLE:
)
async def fetch_all_mcp_servers(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
team_id: Optional[str] = Query(
None,
description="Filter MCP servers by team scope. When provided, returns only "
"servers the team has access to plus globally available (allow_all_keys) servers. "
"Used by the Create Key UI to show team-scoped MCP servers.",
),
):
"""
Get all of the configured mcp servers for the user in the db with their associated teams
```
curl --location 'http://localhost:4000/v1/mcp/server' \
--header 'Authorization: Bearer your_api_key_here'
# Filter by team scope (for Create Key UI)
curl --location 'http://localhost:4000/v1/mcp/server?team_id=team-123' \
--header 'Authorization: Bearer your_api_key_here'
```
"""
user_mcp_management_mode = _get_user_mcp_management_mode()
# If team_id is provided, return team-scoped servers + allow_all_keys servers
is_restricted_virtual_key = _is_restricted_virtual_key_request(
user_api_key_dict
)
if user_mcp_management_mode == "view_all" and not is_restricted_virtual_key:
servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered()
redacted_mcp_servers = _redact_mcp_credentials_list(servers)
else:
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
for auth_context in auth_contexts:
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
user_api_key_auth=auth_context
if team_id is not None and isinstance(team_id, str) and team_id.strip():
# Restricted virtual keys must not use the team_id filter to
# bypass their own access limitations.
if is_restricted_virtual_key:
raise HTTPException(
status_code=403,
detail="Restricted virtual keys cannot query team-scoped MCP servers.",
)
for server in servers:
if server.server_id not in aggregated_servers:
aggregated_servers[server.server_id] = server
redacted_mcp_servers = _redact_mcp_credentials_list(
aggregated_servers.values()
)
# Only proxy admins may query another team's MCP servers.
# Non-admins must belong to the requested team.
sanitized_team_id = team_id.strip()
is_admin = _user_has_admin_view(user_api_key_dict)
if not is_admin:
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=sanitized_team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
user_in_team = any(
m.user_id is not None
and m.user_id == user_api_key_dict.user_id
for m in team_obj.members_with_roles
)
if not user_in_team:
raise HTTPException(
status_code=403,
detail="You do not have permission to view MCP servers for this team.",
)
redacted_mcp_servers = await _get_team_scoped_mcp_server_list(sanitized_team_id)
else:
user_mcp_management_mode = _get_user_mcp_management_mode()
if user_mcp_management_mode == "view_all" and not is_restricted_virtual_key:
servers = await global_mcp_server_manager.get_all_mcp_servers_unfiltered()
redacted_mcp_servers = _redact_mcp_credentials_list(servers)
else:
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
aggregated_servers: Dict[str, LiteLLM_MCPServerTable] = {}
for auth_context in auth_contexts:
servers = await global_mcp_server_manager.get_all_allowed_mcp_servers(
user_api_key_auth=auth_context
)
for server in servers:
if server.server_id not in aggregated_servers:
aggregated_servers[server.server_id] = server
redacted_mcp_servers = _redact_mcp_credentials_list(
aggregated_servers.values()
)
# augment the mcp servers with public status
if litellm.public_mcp_servers is not None:
@@ -4,12 +4,14 @@ 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, status
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.utils import PrismaClient
@@ -177,4 +179,177 @@ async def _set_object_permission(
data_json["object_permission_id"] = created_permission.object_permission_id
data_json.pop("object_permission")
return data_json
return data_json
async def _resolve_team_allowed_mcp_servers(
team_object_permission: "LiteLLM_ObjectPermissionTable",
) -> Set[str]:
"""
Resolve the full set of MCP server IDs a team has access to.
Combines:
- Direct mcp_servers list
- Servers from mcp_access_groups
- Server IDs referenced in mcp_tool_permissions keys
"""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
direct_servers: List[str] = team_object_permission.mcp_servers or []
access_group_servers: List[str] = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
team_object_permission.mcp_access_groups or []
)
)
tool_perm_servers: List[str] = list(
(team_object_permission.mcp_tool_permissions or {}).keys()
)
return set(direct_servers + access_group_servers + tool_perm_servers)
def _get_allow_all_keys_server_ids() -> Set[str]:
"""Return the set of MCP server IDs marked with allow_all_keys=True."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
return set(global_mcp_server_manager.get_allow_all_keys_server_ids())
async def _get_team_allowed_mcp_servers(
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
) -> Set[str]:
"""
Get the full set of MCP server IDs a team allows.
If team has no object_permission or no MCP config, returns empty set
(meaning only allow_all_keys servers are permitted).
"""
if team_obj is None:
return set()
team_object_permission = team_obj.object_permission
if team_object_permission is None:
return set()
return await _resolve_team_allowed_mcp_servers(team_object_permission)
def _extract_requested_mcp_server_ids(
object_permission: Optional[dict],
) -> Set[str]:
"""
Extract all MCP server IDs referenced in a key's object_permission dict.
Includes:
- mcp_servers list
- Keys from mcp_tool_permissions
"""
if not object_permission or not isinstance(object_permission, dict):
return set()
server_ids: Set[str] = set()
mcp_servers = object_permission.get("mcp_servers")
if isinstance(mcp_servers, list):
server_ids.update(mcp_servers)
mcp_tool_permissions = object_permission.get("mcp_tool_permissions")
if isinstance(mcp_tool_permissions, dict):
server_ids.update(mcp_tool_permissions.keys())
return server_ids
def _extract_requested_mcp_access_groups(
object_permission: Optional[dict],
) -> Set[str]:
"""Extract MCP access groups from a key's object_permission dict."""
if not object_permission or not isinstance(object_permission, dict):
return set()
groups = object_permission.get("mcp_access_groups")
if isinstance(groups, list):
return set(groups)
return set()
async def validate_key_mcp_servers_against_team(
object_permission: Optional[dict],
team_obj: Optional["LiteLLM_TeamTableCachedObj"],
):
"""
Validate that MCP servers requested on a key are within the allowed scope.
Rules:
- If key is in a team: key's mcp_servers must be a subset of
(team's allowed servers + allow_all_keys servers)
- If key is NOT in a team: key's mcp_servers must only contain
allow_all_keys servers
- If team has no MCP config: key can only use allow_all_keys servers
Raises HTTPException(403) if validation fails.
"""
requested_servers = _extract_requested_mcp_server_ids(object_permission)
requested_access_groups = _extract_requested_mcp_access_groups(object_permission)
# Nothing to validate
if not requested_servers and not requested_access_groups:
return
allow_all_keys_servers = _get_allow_all_keys_server_ids()
team_allowed_servers = await _get_team_allowed_mcp_servers(team_obj)
# Combined allowed set = team servers + allow_all_keys servers
all_allowed_servers = team_allowed_servers | allow_all_keys_servers
# Validate requested server IDs
if requested_servers:
disallowed_servers = requested_servers - all_allowed_servers
if disallowed_servers:
if team_obj is not None:
detail = (
f"Key requests MCP servers not allowed by team '{team_obj.team_id}': "
f"{sorted(disallowed_servers)}. "
f"Team allows: {sorted(team_allowed_servers)}. "
f"Global (allow_all_keys) servers: {sorted(allow_all_keys_servers)}."
)
else:
detail = (
f"Key is not in a team. Only globally available (allow_all_keys) MCP servers "
f"can be assigned: {sorted(allow_all_keys_servers)}. "
f"Disallowed servers: {sorted(disallowed_servers)}."
)
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": detail},
)
# Validate requested access groups (must be subset of team's access groups)
if requested_access_groups:
team_access_groups: Set[str] = set()
if (
team_obj is not None
and team_obj.object_permission is not None
and team_obj.object_permission.mcp_access_groups
):
team_access_groups = set(team_obj.object_permission.mcp_access_groups)
disallowed_groups = requested_access_groups - team_access_groups
if disallowed_groups:
if team_obj is not None:
detail = (
f"Key requests MCP access groups not allowed by team '{team_obj.team_id}': "
f"{sorted(disallowed_groups)}. "
f"Team allows: {sorted(team_access_groups)}."
)
else:
detail = (
f"Key is not in a team. MCP access groups cannot be assigned to "
f"keys outside of a team. Disallowed groups: {sorted(disallowed_groups)}."
)
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": detail},
)
@@ -605,6 +605,10 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch):
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team",
AsyncMock(),
)
from litellm.proxy._types import (
GenerateKeyRequest,
@@ -2346,6 +2350,9 @@ async def test_generate_key_with_object_permission():
), patch(
"litellm.proxy.proxy_server.litellm_proxy_admin_name",
"admin",
), patch(
"litellm.proxy.management_endpoints.key_management_endpoints.validate_key_mcp_servers_against_team",
new_callable=AsyncMock,
):
# Execute
result = await _common_key_generation_helper(
@@ -797,6 +797,157 @@ class TestListMCPServers:
assert result.status == "healthy"
class TestTeamScopedMCPServerAccess:
"""Tests for cross-team information disclosure and restricted key bypass fixes."""
@pytest.mark.asyncio
async def test_non_member_cannot_query_foreign_team(self):
"""Non-admin user who is NOT a member of the target team should get 403."""
from litellm.proxy._types import Member
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="attacker_user",
)
# Team with a different member
mock_team_obj = MagicMock()
mock_team_obj.members_with_roles = [
Member(user_id="legitimate_user", role="admin"),
]
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=False,
),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
AsyncMock(return_value=mock_team_obj),
),
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
with pytest.raises(HTTPException) as exc_info:
await fetch_all_mcp_servers(
user_api_key_dict=mock_user_auth, team_id="foreign-team-id"
)
assert exc_info.value.status_code == 403
assert "permission" in str(exc_info.value.detail).lower()
@pytest.mark.asyncio
async def test_team_member_can_query_own_team(self):
"""User who IS a member of the team should be able to query it."""
from litellm.proxy._types import Member
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="team_member",
)
mock_team_obj = MagicMock()
mock_team_obj.members_with_roles = [
Member(user_id="team_member", role="user"),
]
mock_team_obj.object_permission = MagicMock(mcp_servers=["server-1"])
mock_server = generate_mock_mcp_server_config_record(
server_id="server-1", name="Team Server"
)
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server)
mock_manager._build_mcp_server_table = MagicMock(
return_value=generate_mock_mcp_server_db_record(server_id="server-1")
)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=False,
),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
AsyncMock(return_value=mock_team_obj),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_scoped_mcp_server_list",
AsyncMock(
return_value=[
generate_mock_mcp_server_db_record(server_id="server-1")
]
),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
mock_manager,
),
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
result = await fetch_all_mcp_servers(
user_api_key_dict=mock_user_auth, team_id="my-team-id"
)
assert len(result) == 1
assert result[0].server_id == "server-1"
@pytest.mark.asyncio
async def test_admin_can_query_any_team(self):
"""Proxy admins should be able to query any team's MCP servers."""
mock_user_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin_user",
)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view",
return_value=True,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_scoped_mcp_server_list",
AsyncMock(
return_value=[
generate_mock_mcp_server_db_record(server_id="server-1")
]
),
),
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
# Admin should NOT need to be a team member
result = await fetch_all_mcp_servers(
user_api_key_dict=mock_user_auth, team_id="any-team-id"
)
assert len(result) == 1
@pytest.mark.asyncio
async def test_restricted_virtual_key_cannot_use_team_id_filter(self):
"""Restricted virtual keys must not bypass access limits via team_id."""
mock_user_auth = UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER,
user_id="vkey_user",
api_key="sk-restricted",
allowed_routes=["mcp_routes"],
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
fetch_all_mcp_servers,
)
with pytest.raises(HTTPException) as exc_info:
await fetch_all_mcp_servers(
user_api_key_dict=mock_user_auth, team_id="some-team"
)
assert exc_info.value.status_code == 403
assert "Restricted virtual key" in str(exc_info.value.detail)
class TestTemporaryMCPSessionEndpoints:
def test_inherit_credentials_from_existing_server(self):
payload = NewMCPServerRequest(
@@ -3,15 +3,20 @@ 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_ObjectPermissionTable
from litellm.proxy.management_helpers.object_permission_utils import (
_extract_requested_mcp_access_groups,
_extract_requested_mcp_server_ids,
_set_object_permission,
validate_key_mcp_servers_against_team,
)
@@ -82,3 +87,311 @@ async def test_set_object_permission():
assert result["user_id"] == "test_user"
assert result["models"] == ["gpt-4"]
# ---- Tests for _extract_requested_mcp_server_ids ----
def test_extract_requested_mcp_server_ids_from_mcp_servers():
obj_perm = {"mcp_servers": ["server-1", "server-2"]}
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1", "server-2"}
def test_extract_requested_mcp_server_ids_from_tool_permissions():
obj_perm = {"mcp_tool_permissions": {"server-a": ["tool1"], "server-b": ["tool2"]}}
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-a", "server-b"}
def test_extract_requested_mcp_server_ids_combined():
obj_perm = {
"mcp_servers": ["server-1"],
"mcp_tool_permissions": {"server-2": ["tool1"]},
}
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1", "server-2"}
def test_extract_requested_mcp_server_ids_none():
assert _extract_requested_mcp_server_ids(None) == set()
assert _extract_requested_mcp_server_ids({}) == set()
# ---- Tests for _extract_requested_mcp_access_groups ----
def test_extract_requested_mcp_access_groups():
obj_perm = {"mcp_access_groups": ["group-a", "group-b"]}
assert _extract_requested_mcp_access_groups(obj_perm) == {"group-a", "group-b"}
def test_extract_requested_mcp_access_groups_none():
assert _extract_requested_mcp_access_groups(None) == set()
assert _extract_requested_mcp_access_groups({}) == set()
# ---- Tests for validate_key_mcp_servers_against_team ----
def _make_team_obj(
team_id="team-1",
mcp_servers=None,
mcp_access_groups=None,
mcp_tool_permissions=None,
):
"""Create a mock team object with the given MCP permissions."""
mock_team = MagicMock()
mock_team.team_id = team_id
if mcp_servers is not None or mcp_access_groups is not None or mcp_tool_permissions is not None:
mock_team.object_permission = MagicMock(spec=LiteLLM_ObjectPermissionTable)
mock_team.object_permission.mcp_servers = mcp_servers or []
mock_team.object_permission.mcp_access_groups = mcp_access_groups or []
mock_team.object_permission.mcp_tool_permissions = mcp_tool_permissions or {}
else:
mock_team.object_permission = None
return mock_team
@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_no_object_permission(mock_access_groups, mock_allow_all):
"""No object_permission on key — should pass without error."""
await validate_key_mcp_servers_against_team(
object_permission=None,
team_obj=_make_team_obj(mcp_servers=["server-1"]),
)
@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_key_servers_within_team_scope(mock_access_groups, mock_allow_all):
"""Key requests servers that are in the team's scope — should pass."""
team_obj = _make_team_obj(mcp_servers=["server-1", "server-2", "server-3"])
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["server-1", "server-2"]},
team_obj=team_obj,
)
@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_key_servers_outside_team_scope_raises(mock_access_groups, mock_allow_all):
"""Key requests servers NOT in the team's scope — should raise 403."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["server-1", "server-outside"]},
team_obj=team_obj,
)
assert exc_info.value.status_code == 403
assert "server-outside" in str(exc_info.value.detail)
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value={"global-server"},
)
@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_allow_all_keys_servers_always_allowed(mock_access_groups, mock_allow_all):
"""allow_all_keys servers should be accessible even if not in team scope."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["server-1", "global-server"]},
team_obj=team_obj,
)
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value={"global-server"},
)
@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_no_team_only_allow_all_keys(mock_access_groups, mock_allow_all):
"""Key without a team can only use allow_all_keys servers."""
# This should pass — requesting a global server without a team
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["global-server"]},
team_obj=None,
)
@pytest.mark.asyncio
@patch(
"litellm.proxy.management_helpers.object_permission_utils._get_allow_all_keys_server_ids",
return_value={"global-server"},
)
@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_no_team_non_global_server_raises(mock_access_groups, mock_allow_all):
"""Key without a team requesting a non-global server — should raise 403."""
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["private-server"]},
team_obj=None,
)
assert exc_info.value.status_code == 403
assert "not in a team" in str(exc_info.value.detail)
@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_team_no_mcp_config_blocks_all(mock_access_groups, mock_allow_all):
"""Team with no object_permission — key can't use any non-global MCP servers."""
team_obj = _make_team_obj() # No object_permission
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["some-server"]},
team_obj=team_obj,
)
assert exc_info.value.status_code == 403
@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_tool_permissions_validated_against_team(mock_access_groups, mock_allow_all):
"""Server IDs in mcp_tool_permissions should also be validated."""
team_obj = _make_team_obj(mcp_servers=["server-1"])
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={
"mcp_tool_permissions": {"server-outside": ["tool1"]}
},
team_obj=team_obj,
)
assert exc_info.value.status_code == 403
assert "server-outside" in str(exc_info.value.detail)
@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_access_groups_within_team_scope(mock_access_groups, mock_allow_all):
"""Key requests access groups that are in the team's scope — should pass."""
team_obj = _make_team_obj(mcp_access_groups=["group-a", "group-b"])
await validate_key_mcp_servers_against_team(
object_permission={"mcp_access_groups": ["group-a"]},
team_obj=team_obj,
)
@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_access_groups_outside_team_scope_raises(mock_access_groups, mock_allow_all):
"""Key requests access groups NOT in the team's scope — should raise 403."""
team_obj = _make_team_obj(mcp_access_groups=["group-a"])
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_access_groups": ["group-outside"]},
team_obj=team_obj,
)
assert exc_info.value.status_code == 403
assert "group-outside" in str(exc_info.value.detail)
@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_access_groups_no_team_raises(mock_access_groups, mock_allow_all):
"""Key without a team requesting access groups — should raise 403."""
with pytest.raises(HTTPException) as exc_info:
await validate_key_mcp_servers_against_team(
object_permission={"mcp_access_groups": ["group-a"]},
team_obj=None,
)
assert exc_info.value.status_code == 403
assert "not in a team" in str(exc_info.value.detail)
@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=["server-from-group"],
)
async def test_validate_team_access_groups_resolve_to_servers(mock_access_groups, mock_allow_all):
"""Team access groups should resolve to server IDs and be included in allowed set."""
team_obj = _make_team_obj(mcp_access_groups=["group-a"])
# Key requests a server that comes from the team's access group
await validate_key_mcp_servers_against_team(
object_permission={"mcp_servers": ["server-from-group"]},
team_obj=team_obj,
)
@@ -6,11 +6,11 @@ import useAuthorized from "../useAuthorized";
const mcpServersKeys = createQueryKeys("mcpServers");
export const useMCPServers = () => {
export const useMCPServers = (teamId?: string | null) => {
const { accessToken } = useAuthorized();
return useQuery<MCPServer[]>({
queryKey: mcpServersKeys.list({}),
queryFn: async () => await fetchMCPServers(accessToken!),
queryKey: mcpServersKeys.list({ teamId: teamId ?? undefined }),
queryFn: async () => await fetchMCPServers(accessToken!, teamId),
enabled: !!accessToken,
});
};
@@ -13,6 +13,7 @@ interface MCPServerSelectorProps {
accessToken: string;
placeholder?: string;
disabled?: boolean;
teamId?: string | null;
}
const MCPServerSelector: React.FC<MCPServerSelectorProps> = ({
@@ -22,8 +23,9 @@ const MCPServerSelector: React.FC<MCPServerSelectorProps> = ({
accessToken,
placeholder = "Select MCP servers",
disabled = false,
teamId,
}) => {
const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers();
const { data: mcpServers = [], isLoading: serversLoading } = useMCPServers(teamId);
const { data: accessGroups = [], isLoading: groupsLoading } = useMCPAccessGroups();
const loading = serversLoading || groupsLoading;
@@ -6320,10 +6320,15 @@ export const fetchDiscoverableMCPServers = async (accessToken: string) => {
}
};
export const fetchMCPServers = async (accessToken: string) => {
export const fetchMCPServers = async (accessToken: string, teamId?: string | null) => {
try {
// Construct base URL
const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server` : `/v1/mcp/server`;
// Construct base URL with optional team_id filter
let url = proxyBaseUrl ? `${proxyBaseUrl}/v1/mcp/server` : `/v1/mcp/server`;
if (teamId) {
const params = new URLSearchParams();
params.append("team_id", teamId);
url = `${url}?${params.toString()}`;
}
console.log("Fetching MCP servers from:", url);
@@ -563,6 +563,8 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
if (!pendingPrefillModels) {
form.setFieldValue("models", []);
}
// Clear MCP server selection when team changes (available servers may differ)
form.setFieldValue("allowed_mcp_servers_and_groups", { servers: [], accessGroups: [] });
}, [selectedCreateKeyTeam, selectedProjectId, accessToken, userID, userRole, form]);
// Apply deferred model prefill once the available model list arrives.
@@ -1323,6 +1325,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
onChange={(val: any) => form.setFieldValue("allowed_mcp_servers_and_groups", val)}
value={form.getFieldValue("allowed_mcp_servers_and_groups")}
accessToken={accessToken}
teamId={selectedCreateKeyTeam?.team_id ?? null}
placeholder="Select MCP servers or access groups (optional)"
/>
</Form.Item>