refactor: extract _auto_assign and _remove helpers, use team_endpoints helper

- Replace raw prisma_client.db.litellm_teamtable.update with
  handle_update_object_permission from team_endpoints (follows
  established helper-function pattern)
- Extract _auto_assign_mcp_server_to_team and
  _remove_mcp_server_from_team helpers for reuse and testability
- Update tests to mock at the correct boundaries

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-03-20 09:05:18 -07:00
parent 65bfd449e9
commit b8c9bf7d25
2 changed files with 96 additions and 112 deletions

View file

@ -203,6 +203,69 @@ if MCP_AVAILABLE:
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
@ -1347,36 +1410,12 @@ if MCP_AVAILABLE:
# 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:
existing_permission_id = getattr(
manager_team_obj, "object_permission_id", None
)
# Read existing mcp_servers list and append the new server
existing_mcp_servers: list = []
if manager_team_obj.object_permission is not None:
existing_mcp_servers = (
manager_team_obj.object_permission.mcp_servers or []
)
updated_mcp_servers = list(
set(existing_mcp_servers + [new_mcp_server.server_id])
)
new_permission_id = await handle_update_object_permission_common(
data_json={
"object_permission": {
"mcp_servers": updated_mcp_servers,
}
},
existing_object_permission_id=existing_permission_id,
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,
)
# If the team had no object_permission_id, link the new one
if existing_permission_id is None and new_permission_id is not None:
await prisma_client.db.litellm_teamtable.update(
where={"team_id": manager_team_id},
data={"object_permission_id": new_permission_id},
)
except Exception as e:
verbose_proxy_logger.exception(
f"MCP server created but failed to auto-assign to team {manager_team_id}: {str(e)}"
@ -1617,28 +1656,10 @@ if MCP_AVAILABLE:
# Remove server from the manager's team permission list
if manager_team_obj is not None:
try:
existing_permission_id = getattr(
manager_team_obj, "object_permission_id", None
await _remove_mcp_server_from_team(
server_id=server_id,
team_obj=manager_team_obj,
)
if (
existing_permission_id is not None
and manager_team_obj.object_permission is not None
):
existing_mcp_servers = (
manager_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,
)
except Exception as e:
verbose_proxy_logger.exception(
f"MCP server deleted but failed to remove from team permissions: {str(e)}"

View file

@ -233,7 +233,7 @@ class TestCreateMcpServerAsManager:
created_server.server_id = "new_server_id"
created_server.credentials = None
mock_handle_update = AsyncMock()
mock_auto_assign = AsyncMock()
with (
patch(
@ -263,38 +263,23 @@ class TestCreateMcpServerAsManager:
),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.handle_update_object_permission_common",
mock_handle_update,
"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 handle_update_object_permission_common was called with merged server list
mock_handle_update.assert_called_once()
call_kwargs = mock_handle_update.call_args.kwargs
mcp_servers = call_kwargs["data_json"]["object_permission"]["mcp_servers"]
assert "existing_server" in mcp_servers
assert "new_server_id" in mcp_servers
assert call_kwargs["existing_object_permission_id"] == "perm1"
# 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_create_links_new_permission_to_team_when_none_exists(self):
"""When team has no object_permission_id, create should link the new one."""
from litellm.proxy._types import NewMCPServerRequest
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 (
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",
_auto_assign_mcp_server_to_team,
)
# Team with NO object_permission_id
@ -307,46 +292,24 @@ class TestCreateMcpServerAsManager:
)
mock_team.object_permission = None
created_server = MagicMock()
created_server.server_id = "new_server_id"
created_server.credentials = None
mock_prisma = MagicMock()
mock_prisma.db.litellm_teamtable.update = AsyncMock()
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
),
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.handle_update_object_permission_common",
AsyncMock(return_value="new_perm_id"),
),
# 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 add_mcp_server(payload=payload, user_api_key_dict=user_auth)
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(