mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
65bfd449e9
commit
b8c9bf7d25
2 changed files with 96 additions and 112 deletions
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue