mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge pull request #24171 from BerriAI/litellm_/awesome-dhawan
[Feature] Team MCP Server Manager Role
This commit is contained in:
commit
b36269e2c1
5 changed files with 772 additions and 31 deletions
|
|
@ -40,8 +40,8 @@ def _prepare_mcp_server_data(
|
|||
"""
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
# Convert model to dict
|
||||
data_dict = data.model_dump(exclude_none=True)
|
||||
# Convert model to dict, excluding fields not in the DB schema
|
||||
data_dict = data.model_dump(exclude_none=True, exclude={"team_id"})
|
||||
# Ensure alias is always present in the dict (even if None)
|
||||
if "alias" not in data_dict:
|
||||
data_dict["alias"] = getattr(data, "alias", None)
|
||||
|
|
|
|||
|
|
@ -1144,6 +1144,10 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="Server-managed: set by the endpoint; caller values are overridden.",
|
||||
)
|
||||
team_id: Optional[str] = Field(
|
||||
None,
|
||||
description="Team ID to assign the MCP server to. Required for team MCP managers.",
|
||||
)
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
@ -1608,8 +1612,9 @@ class Member(MemberBase):
|
|||
role: Literal[
|
||||
"admin",
|
||||
"user",
|
||||
"mcp_server_manager",
|
||||
] = Field(
|
||||
description="The role of the user within the team. 'admin' users can manage team settings and members, 'user' is a regular team member"
|
||||
description="The role of the user within the team. 'admin' users can manage team settings and members, 'user' is a regular team member, 'mcp_server_manager' can manage MCP servers for the team"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -41,6 +41,18 @@ def _is_user_team_admin(
|
|||
return False
|
||||
|
||||
|
||||
def _is_user_team_mcp_manager(
|
||||
user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable
|
||||
) -> bool:
|
||||
for member in team_obj.members_with_roles:
|
||||
if (
|
||||
member.user_id is not None and member.user_id == user_api_key_dict.user_id
|
||||
) and member.role == "mcp_server_manager":
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
async def _is_user_org_admin_for_team(
|
||||
user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable
|
||||
) -> bool:
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ import json
|
|||
import os
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, Iterable, List, Literal, Optional
|
||||
from typing import Any, Dict, Iterable, List, Literal, Optional, Tuple
|
||||
|
||||
from fastapi import (
|
||||
APIRouter,
|
||||
|
|
@ -112,6 +112,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LitellmUserRoles,
|
||||
MakeMCPServersPublicRequest,
|
||||
MCPApprovalStatus,
|
||||
|
|
@ -130,11 +131,142 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_is_user_team_mcp_manager,
|
||||
_user_has_admin_view,
|
||||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_get_team_allowed_mcp_servers,
|
||||
handle_update_object_permission_common,
|
||||
)
|
||||
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
|
||||
from litellm.types.mcp import MCPCredentials
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
async def _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
team_id: Optional[str] = None,
|
||||
server_id: Optional[str] = None,
|
||||
) -> Tuple[str, LiteLLM_TeamTableCachedObj]:
|
||||
"""
|
||||
Verify that the caller is an MCP server manager for a team and (for edit/delete)
|
||||
that the target server belongs to that team.
|
||||
|
||||
Returns a tuple of (team_id, team_obj) for downstream use.
|
||||
Raises HTTPException(400) if no team_id can be determined.
|
||||
Raises HTTPException(403) if the caller is not an MCP manager or server not in team.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
resolved_team_id = team_id or user_api_key_dict.team_id
|
||||
if not resolved_team_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "team_id is required for MCP server manager operations."},
|
||||
)
|
||||
|
||||
# When the API key is team-scoped, ensure the request team_id matches
|
||||
if (
|
||||
team_id
|
||||
and user_api_key_dict.team_id
|
||||
and team_id != user_api_key_dict.team_id
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "team_id does not match the API key's team."},
|
||||
)
|
||||
|
||||
team_obj = await get_team_object(
|
||||
team_id=resolved_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_db_only=True, # bypass cache to get fresh object_permission
|
||||
)
|
||||
|
||||
if not _is_user_team_mcp_manager(user_api_key_dict, team_obj):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"User does not have mcp_server_manager role in team {resolved_team_id}."
|
||||
},
|
||||
)
|
||||
|
||||
if server_id is not None:
|
||||
team_server_ids = await _get_team_allowed_mcp_servers(team_obj)
|
||||
if server_id not in team_server_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"MCP server {server_id} is not assigned to team {resolved_team_id}."
|
||||
},
|
||||
)
|
||||
|
||||
return resolved_team_id, team_obj
|
||||
|
||||
async def _auto_assign_mcp_server_to_team(
|
||||
server_id: str,
|
||||
team_id: str,
|
||||
team_obj: LiteLLM_TeamTableCachedObj,
|
||||
prisma_client: Any,
|
||||
) -> None:
|
||||
"""
|
||||
Add an MCP server to a team's ObjectPermissionTable and link the
|
||||
permission back to the team if it didn't have one yet.
|
||||
|
||||
Uses handle_update_object_permission (the team-endpoint helper) so
|
||||
the object_permission_id linkage follows the same pattern as
|
||||
team_endpoints.update_team.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.team_endpoints import (
|
||||
handle_update_object_permission,
|
||||
)
|
||||
|
||||
existing_mcp_servers: list = []
|
||||
if team_obj.object_permission is not None:
|
||||
existing_mcp_servers = team_obj.object_permission.mcp_servers or []
|
||||
updated_mcp_servers = list(set(existing_mcp_servers + [server_id]))
|
||||
|
||||
# Build the data dict in the same shape team_endpoints uses
|
||||
data_json: Dict[str, Any] = {
|
||||
"object_permission": {"mcp_servers": updated_mcp_servers},
|
||||
}
|
||||
data_json = await handle_update_object_permission(
|
||||
data_json=data_json,
|
||||
existing_team_row=team_obj,
|
||||
)
|
||||
|
||||
# If handle_update_object_permission produced an object_permission_id,
|
||||
# persist it on the team row (it sets data_json["object_permission_id"]).
|
||||
if "object_permission_id" in data_json:
|
||||
await prisma_client.db.litellm_teamtable.update(
|
||||
where={"team_id": team_id},
|
||||
data={"object_permission_id": data_json["object_permission_id"]},
|
||||
)
|
||||
|
||||
async def _remove_mcp_server_from_team(
|
||||
server_id: str,
|
||||
team_obj: LiteLLM_TeamTableCachedObj,
|
||||
) -> None:
|
||||
"""Remove a server ID from a team's ObjectPermissionTable.mcp_servers list."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
existing_permission_id = getattr(team_obj, "object_permission_id", None)
|
||||
if (
|
||||
existing_permission_id is None
|
||||
or team_obj.object_permission is None
|
||||
):
|
||||
return
|
||||
|
||||
existing_mcp_servers = team_obj.object_permission.mcp_servers or []
|
||||
updated_mcp_servers = [s for s in existing_mcp_servers if s != server_id]
|
||||
|
||||
await handle_update_object_permission_common(
|
||||
data_json={"object_permission": {"mcp_servers": updated_mcp_servers}},
|
||||
existing_object_permission_id=existing_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
@dataclass
|
||||
class _TemporaryMCPServerEntry:
|
||||
server: MCPServer
|
||||
|
|
@ -1206,16 +1338,27 @@ if MCP_AVAILABLE:
|
|||
# Validate and normalize payload fields
|
||||
validate_and_normalize_mcp_server_payload(payload)
|
||||
|
||||
# AuthZ - restrict only proxy admins to create mcp servers
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "User does not have permission to create mcp servers. You can only create mcp servers if you are a PROXY_ADMIN."
|
||||
},
|
||||
# AuthZ - proxy admins or team MCP managers can create MCP servers
|
||||
is_proxy_admin = LitellmUserRoles.PROXY_ADMIN == user_api_key_dict.user_role
|
||||
manager_team_id: Optional[str] = None
|
||||
|
||||
manager_team_obj = None
|
||||
if not is_proxy_admin:
|
||||
# Check if the user is an MCP manager for the specified team
|
||||
if payload.team_id is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail={
|
||||
"error": "team_id is required when creating MCP servers as a team MCP manager."
|
||||
},
|
||||
)
|
||||
manager_team_id, manager_team_obj = await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_id=payload.team_id,
|
||||
)
|
||||
elif payload.server_id is not None:
|
||||
# fail if the mcp server with id already exists
|
||||
|
||||
# Fail if the MCP server with this id already exists
|
||||
if payload.server_id is not None:
|
||||
mcp_server = await get_mcp_server(prisma_client, payload.server_id)
|
||||
if mcp_server is not None:
|
||||
raise HTTPException(
|
||||
|
|
@ -1224,7 +1367,8 @@ if MCP_AVAILABLE:
|
|||
"error": f"MCP Server with id {payload.server_id} already exists. Cannot create another."
|
||||
},
|
||||
)
|
||||
elif (
|
||||
|
||||
if (
|
||||
SpecialMCPServerName.all_team_servers == payload.server_id
|
||||
or SpecialMCPServerName.all_proxy_servers == payload.server_id
|
||||
):
|
||||
|
|
@ -1255,12 +1399,28 @@ if MCP_AVAILABLE:
|
|||
|
||||
# Ensure registry is up to date by reloading from database
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating mcp server: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail={"error": f"Error creating mcp server: {str(e)}"},
|
||||
)
|
||||
|
||||
# Auto-assign the server to the manager's team (separate from create
|
||||
# so a failure here doesn't mask the successfully created server).
|
||||
if manager_team_id is not None and manager_team_obj is not None:
|
||||
try:
|
||||
await _auto_assign_mcp_server_to_team(
|
||||
server_id=new_mcp_server.server_id,
|
||||
team_id=manager_team_id,
|
||||
team_obj=manager_team_obj,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
f"MCP server created but failed to auto-assign to team {manager_team_id}: {str(e)}"
|
||||
)
|
||||
return _redact_mcp_credentials(new_mcp_server)
|
||||
|
||||
@router.post(
|
||||
|
|
@ -1473,15 +1633,12 @@ if MCP_AVAILABLE:
|
|||
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
|
||||
)
|
||||
|
||||
# Authz - restrict only admins to delete mcp servers
|
||||
# Authz - proxy admins or team MCP managers can delete MCP servers
|
||||
manager_team_obj = None
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "Call not allowed to delete MCP server. User is not a proxy admin. route={}".format(
|
||||
"DELETE /v1/mcp/server"
|
||||
)
|
||||
},
|
||||
_, manager_team_obj = await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
server_id=server_id,
|
||||
)
|
||||
|
||||
# try to delete the mcp server
|
||||
|
|
@ -1497,6 +1654,18 @@ if MCP_AVAILABLE:
|
|||
# Ensure registry is up to date by reloading from database
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
|
||||
# Remove server from the manager's team permission list
|
||||
if manager_team_obj is not None:
|
||||
try:
|
||||
await _remove_mcp_server_from_team(
|
||||
server_id=server_id,
|
||||
team_obj=manager_team_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
f"MCP server deleted but failed to remove from team permissions: {str(e)}"
|
||||
)
|
||||
|
||||
# TODO: Enterprise: Finish audit log trail
|
||||
if litellm.store_audit_logs:
|
||||
pass
|
||||
|
|
@ -1798,15 +1967,11 @@ if MCP_AVAILABLE:
|
|||
# Validate and normalize payload fields
|
||||
validate_and_normalize_mcp_server_payload(payload)
|
||||
|
||||
# Authz - restrict only admins to delete mcp servers
|
||||
# Authz - proxy admins or team MCP managers can update MCP servers
|
||||
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={
|
||||
"error": "Call not allowed to update MCP server. User is not a proxy admin. route={}".format(
|
||||
"PUT /v1/mcp/server"
|
||||
)
|
||||
},
|
||||
await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
server_id=payload.server_id,
|
||||
)
|
||||
|
||||
# try to update the mcp server
|
||||
|
|
|
|||
|
|
@ -0,0 +1,559 @@
|
|||
import pytest
|
||||
|
||||
pytest.importorskip("mcp", reason="mcp package not installed")
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTable,
|
||||
LitellmUserRoles,
|
||||
Member,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_is_user_team_mcp_manager,
|
||||
)
|
||||
from litellm.proxy.management_endpoints import (
|
||||
mcp_management_endpoints as mgmt_endpoints,
|
||||
)
|
||||
|
||||
|
||||
class TestIsUserTeamMcpManager:
|
||||
def test_mcp_server_manager_role_returns_true(self):
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
)
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager")
|
||||
],
|
||||
)
|
||||
assert _is_user_team_mcp_manager(user_auth, team) is True
|
||||
|
||||
def test_regular_user_role_returns_false(self):
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
)
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="user")],
|
||||
)
|
||||
assert _is_user_team_mcp_manager(user_auth, team) is False
|
||||
|
||||
def test_admin_role_returns_false(self):
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
)
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="admin")],
|
||||
)
|
||||
assert _is_user_team_mcp_manager(user_auth, team) is False
|
||||
|
||||
def test_user_not_in_team_returns_false(self):
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user2",
|
||||
api_key="sk-test",
|
||||
)
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager")
|
||||
],
|
||||
)
|
||||
assert _is_user_team_mcp_manager(user_auth, team) is False
|
||||
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from fastapi import HTTPException
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed"
|
||||
)
|
||||
class TestAssertCanManageTeamMcpServer:
|
||||
async def test_mcp_manager_with_team_id_succeeds(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_assert_can_manage_team_mcp_server,
|
||||
)
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test",
|
||||
)
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="mcp_server_manager")],
|
||||
)
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=mock_team),
|
||||
):
|
||||
team_id, team_obj = await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_auth, team_id="team1"
|
||||
)
|
||||
assert team_id == "team1"
|
||||
assert team_obj == mock_team
|
||||
|
||||
async def test_regular_user_gets_403(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_assert_can_manage_team_mcp_server,
|
||||
)
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test",
|
||||
)
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="user")],
|
||||
)
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=mock_team),
|
||||
):
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_auth, team_id="team1"
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
async def test_admin_gets_403(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_assert_can_manage_team_mcp_server,
|
||||
)
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test",
|
||||
)
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="admin")],
|
||||
)
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=mock_team),
|
||||
):
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_auth, team_id="team1"
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
async def test_mcp_manager_server_in_team_succeeds(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_assert_can_manage_team_mcp_server,
|
||||
)
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", team_id="team1",
|
||||
)
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="mcp_server_manager")],
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=mock_team),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_allowed_mcp_servers",
|
||||
AsyncMock(return_value={"server1", "server2"}),
|
||||
),
|
||||
):
|
||||
team_id, team_obj = await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_auth, server_id="server1"
|
||||
)
|
||||
assert team_id == "team1"
|
||||
assert team_obj == mock_team
|
||||
|
||||
async def test_mcp_manager_server_not_in_team_gets_403(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_assert_can_manage_team_mcp_server,
|
||||
)
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user1", api_key="sk-test", team_id="team1",
|
||||
)
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[Member(user_id="user1", role="mcp_server_manager")],
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_team_object",
|
||||
AsyncMock(return_value=mock_team),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_team_allowed_mcp_servers",
|
||||
AsyncMock(return_value={"server2", "server3"}),
|
||||
),
|
||||
):
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await _assert_can_manage_team_mcp_server(
|
||||
user_api_key_dict=user_auth, server_id="server1"
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed"
|
||||
)
|
||||
class TestCreateMcpServerAsManager:
|
||||
async def test_create_auto_assigns_to_team(self):
|
||||
"""MCP manager creating a server should auto-assign it to their team's ObjectPermissionTable."""
|
||||
from litellm.proxy._types import NewMCPServerRequest
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
add_mcp_server,
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
team_id="team1",
|
||||
)
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="test-server",
|
||||
url="https://example.com/mcp",
|
||||
team_id="team1",
|
||||
)
|
||||
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager"),
|
||||
],
|
||||
object_permission_id="perm1",
|
||||
)
|
||||
mock_team.object_permission = MagicMock(mcp_servers=["existing_server"])
|
||||
|
||||
created_server = MagicMock()
|
||||
created_server.server_id = "new_server_id"
|
||||
created_server.credentials = None
|
||||
|
||||
mock_auto_assign = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server",
|
||||
AsyncMock(return_value=("team1", mock_team)),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
|
||||
AsyncMock(return_value=created_server),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
MagicMock(
|
||||
add_server=AsyncMock(),
|
||||
reload_servers_from_database=AsyncMock(),
|
||||
),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._auto_assign_mcp_server_to_team",
|
||||
mock_auto_assign,
|
||||
),
|
||||
):
|
||||
await add_mcp_server(payload=payload, user_api_key_dict=user_auth)
|
||||
|
||||
# Verify _auto_assign_mcp_server_to_team was called with the right args
|
||||
mock_auto_assign.assert_called_once()
|
||||
call_kwargs = mock_auto_assign.call_args.kwargs
|
||||
assert call_kwargs["server_id"] == "new_server_id"
|
||||
assert call_kwargs["team_id"] == "team1"
|
||||
assert call_kwargs["team_obj"] == mock_team
|
||||
|
||||
async def test_auto_assign_links_new_permission_to_team(self):
|
||||
"""_auto_assign_mcp_server_to_team should create permission and link to team."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_auto_assign_mcp_server_to_team,
|
||||
)
|
||||
|
||||
# Team with NO object_permission_id
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager"),
|
||||
],
|
||||
object_permission_id=None,
|
||||
)
|
||||
mock_team.object_permission = None
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.update = AsyncMock()
|
||||
|
||||
# handle_update_object_permission sets object_permission_id in data_json
|
||||
async def fake_handle(data_json, existing_team_row):
|
||||
data_json["object_permission_id"] = "new_perm_id"
|
||||
return data_json
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.team_endpoints.handle_update_object_permission",
|
||||
side_effect=fake_handle,
|
||||
):
|
||||
await _auto_assign_mcp_server_to_team(
|
||||
server_id="new_server_id",
|
||||
team_id="team1",
|
||||
team_obj=mock_team,
|
||||
prisma_client=mock_prisma,
|
||||
)
|
||||
|
||||
# Verify the team was updated with the new object_permission_id
|
||||
mock_prisma.db.litellm_teamtable.update.assert_called_once_with(
|
||||
where={"team_id": "team1"},
|
||||
data={"object_permission_id": "new_perm_id"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed"
|
||||
)
|
||||
class TestEditMcpServerAsManager:
|
||||
async def test_edit_succeeds_for_mcp_manager(self):
|
||||
"""MCP manager should be able to edit a server assigned to their team."""
|
||||
from litellm.proxy._types import UpdateMCPServerRequest
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
edit_mcp_server,
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
team_id="team1",
|
||||
)
|
||||
|
||||
payload = UpdateMCPServerRequest(
|
||||
server_id="server1",
|
||||
description="Updated description",
|
||||
)
|
||||
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager"),
|
||||
],
|
||||
)
|
||||
|
||||
updated_server = MagicMock()
|
||||
updated_server.server_id = "server1"
|
||||
updated_server.credentials = None
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server",
|
||||
AsyncMock(return_value=("team1", mock_team)),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
|
||||
AsyncMock(return_value=updated_server),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
MagicMock(
|
||||
update_server=AsyncMock(),
|
||||
reload_servers_from_database=AsyncMock(),
|
||||
),
|
||||
),
|
||||
):
|
||||
result = await edit_mcp_server(payload=payload, user_api_key_dict=user_auth)
|
||||
# Should not raise — edit succeeded
|
||||
assert result is not None
|
||||
|
||||
async def test_edit_fails_for_server_not_in_team(self):
|
||||
"""MCP manager should get 403 when editing a server not in their team."""
|
||||
from litellm.proxy._types import UpdateMCPServerRequest
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
edit_mcp_server,
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
team_id="team1",
|
||||
)
|
||||
|
||||
payload = UpdateMCPServerRequest(
|
||||
server_id="server_not_in_team",
|
||||
description="Updated description",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server",
|
||||
AsyncMock(side_effect=HTTPException(status_code=403, detail="Not in team")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await edit_mcp_server(payload=payload, user_api_key_dict=user_auth)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not mgmt_endpoints.MCP_AVAILABLE, reason="MCP module not installed"
|
||||
)
|
||||
class TestDeleteMcpServerAsManager:
|
||||
async def test_delete_succeeds_and_cleans_up_team(self):
|
||||
"""MCP manager deleting a server should also remove it from team permissions."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
remove_mcp_server,
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
team_id="team1",
|
||||
)
|
||||
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager"),
|
||||
],
|
||||
)
|
||||
|
||||
deleted_server = MagicMock()
|
||||
deleted_server.server_id = "server1"
|
||||
|
||||
mock_remove_from_team = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server",
|
||||
AsyncMock(return_value=("team1", mock_team)),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.delete_mcp_server",
|
||||
AsyncMock(return_value=deleted_server),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
MagicMock(
|
||||
remove_server=MagicMock(),
|
||||
reload_servers_from_database=AsyncMock(),
|
||||
),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._remove_mcp_server_from_team",
|
||||
mock_remove_from_team,
|
||||
),
|
||||
):
|
||||
response = await remove_mcp_server(
|
||||
server_id="server1", user_api_key_dict=user_auth
|
||||
)
|
||||
assert response.status_code == 202
|
||||
|
||||
# Verify team cleanup was called
|
||||
mock_remove_from_team.assert_called_once()
|
||||
call_kwargs = mock_remove_from_team.call_args.kwargs
|
||||
assert call_kwargs["server_id"] == "server1"
|
||||
assert call_kwargs["team_obj"] == mock_team
|
||||
|
||||
async def test_delete_fails_for_server_not_in_team(self):
|
||||
"""MCP manager should get 403 when deleting a server not in their team."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
remove_mcp_server,
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="user1",
|
||||
api_key="sk-test",
|
||||
team_id="team1",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._assert_can_manage_team_mcp_server",
|
||||
AsyncMock(side_effect=HTTPException(status_code=403, detail="Not in team")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await remove_mcp_server(
|
||||
server_id="server_not_in_team", user_api_key_dict=user_auth
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
async def test_remove_mcp_server_from_team_helper(self):
|
||||
"""_remove_mcp_server_from_team should update the permission list without the deleted server."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_remove_mcp_server_from_team,
|
||||
)
|
||||
|
||||
mock_team = LiteLLM_TeamTableCachedObj(
|
||||
team_id="team1",
|
||||
members_with_roles=[
|
||||
Member(user_id="user1", role="mcp_server_manager"),
|
||||
],
|
||||
object_permission_id="perm1",
|
||||
)
|
||||
mock_team.object_permission = MagicMock(
|
||||
mcp_servers=["server1", "server2", "server3"]
|
||||
)
|
||||
|
||||
mock_handle_common = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.prisma_client",
|
||||
MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.handle_update_object_permission_common",
|
||||
mock_handle_common,
|
||||
),
|
||||
):
|
||||
await _remove_mcp_server_from_team(
|
||||
server_id="server2",
|
||||
team_obj=mock_team,
|
||||
)
|
||||
|
||||
mock_handle_common.assert_called_once()
|
||||
call_kwargs = mock_handle_common.call_args.kwargs
|
||||
updated_servers = call_kwargs["data_json"]["object_permission"]["mcp_servers"]
|
||||
assert "server2" not in updated_servers
|
||||
assert "server1" in updated_servers
|
||||
assert "server3" in updated_servers
|
||||
assert call_kwargs["existing_object_permission_id"] == "perm1"
|
||||
Loading…
Add table
Reference in a new issue