mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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 <noreply@anthropic.com>
This commit is contained in:
parent
6fe82d3886
commit
6eb0c782f1
11 changed files with 690 additions and 24 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,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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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<string[]>({
|
||||
queryKey: mcpAccessGroupsKeys.list({}),
|
||||
queryFn: async () => await fetchMCPAccessGroups(accessToken!),
|
||||
queryKey: mcpAccessGroupsKeys.list({ teamId }),
|
||||
queryFn: async () => await fetchMCPAccessGroups(accessToken!, teamId),
|
||||
enabled: Boolean(accessToken),
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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<MCPServer[]>({
|
||||
queryKey: mcpServersKeys.list({}),
|
||||
queryFn: async () => await fetchMCPServers(accessToken!),
|
||||
queryKey: mcpServersKeys.list({ teamId }),
|
||||
queryFn: async () => await fetchMCPServers(accessToken!, teamId),
|
||||
enabled: !!accessToken,
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ interface MCPServerSelectorProps {
|
|||
accessToken: string;
|
||||
placeholder?: string;
|
||||
disabled?: boolean;
|
||||
teamId?: string;
|
||||
}
|
||||
|
||||
const MCPServerSelector: React.FC<MCPServerSelectorProps> = ({
|
||||
|
|
@ -22,9 +23,10 @@ const MCPServerSelector: React.FC<MCPServerSelectorProps> = ({
|
|||
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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -1324,6 +1324,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ 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}
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
|
|
|
|||
|
|
@ -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<string[]>([]);
|
||||
const [tagsList, setTagsList] = useState<Record<string, Tag>>({});
|
||||
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}
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue