mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
[MCP Gateway] Litellm mcp fixes team control (#15304)
* fix: _set_object_permission * fix: _set_object_permission on teams * fix: _set_object_permission * fixes for team/key permissions * statsh: object permission view * fix: MCPServerPermissions
This commit is contained in:
parent
49e04e0217
commit
7b56ba240e
9 changed files with 694 additions and 102 deletions
|
|
@ -50,6 +50,7 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
|
|||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
attach_object_permission_to_dict,
|
||||
handle_update_object_permission_common,
|
||||
_set_object_permission,
|
||||
)
|
||||
from litellm.proxy.management_helpers.team_member_permission_checks import (
|
||||
TeamMemberPermissionChecks,
|
||||
|
|
@ -1114,36 +1115,6 @@ def prepare_metadata_fields(
|
|||
return non_default_values
|
||||
|
||||
|
||||
async def _set_object_permission(
|
||||
data_json: dict,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
):
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the key.
|
||||
- Handles permissions for vector stores and mcp servers.
|
||||
"""
|
||||
if prisma_client is None:
|
||||
return data_json
|
||||
|
||||
if "object_permission" in data_json:
|
||||
# Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues
|
||||
# (e.g., server IDs starting with "3e64" being interpreted as floats)
|
||||
if "mcp_tool_permissions" in data_json["object_permission"]:
|
||||
data_json["object_permission"]["mcp_tool_permissions"] = safe_dumps(
|
||||
data_json["object_permission"]["mcp_tool_permissions"]
|
||||
)
|
||||
|
||||
created_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data_json["object_permission"],
|
||||
)
|
||||
)
|
||||
data_json["object_permission_id"] = (
|
||||
created_object_permission.object_permission_id
|
||||
)
|
||||
# delete the object_permission from the data_json
|
||||
data_json.pop("object_permission")
|
||||
return data_json
|
||||
|
||||
|
||||
async def prepare_key_update_data(
|
||||
|
|
|
|||
|
|
@ -56,6 +56,9 @@ from litellm.proxy._types import (
|
|||
UpdateTeamRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_set_object_permission,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
allowed_route_check_inside_route,
|
||||
can_org_access_model,
|
||||
|
|
@ -473,14 +476,15 @@ async def new_team( # noqa: PLR0915
|
|||
|
||||
_model_id = model_dict.id
|
||||
|
||||
## Create Team Member Budget Table
|
||||
data_json = data.json()
|
||||
|
||||
## Handle Object Permission - MCP, Vector Stores etc.
|
||||
object_permission_id = await _set_object_permission(
|
||||
data=data,
|
||||
data_json = await _set_object_permission(
|
||||
data_json=data_json,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
## Create Team Member Budget Table
|
||||
data_json = data.json()
|
||||
if TeamMemberBudgetHandler.should_create_budget(
|
||||
team_member_budget=data.team_member_budget,
|
||||
team_member_rpm_limit=data.team_member_rpm_limit,
|
||||
|
|
@ -499,7 +503,6 @@ async def new_team( # noqa: PLR0915
|
|||
complete_team_data = LiteLLM_TeamTable(
|
||||
**data_json,
|
||||
model_id=_model_id,
|
||||
object_permission_id=object_permission_id,
|
||||
)
|
||||
|
||||
# Set Management Endpoint Metadata Fields
|
||||
|
|
@ -616,29 +619,6 @@ async def _update_model_table(
|
|||
return _model_id
|
||||
|
||||
|
||||
async def _set_object_permission(
|
||||
data: NewTeamRequest,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the team.
|
||||
- Handles permissions for vector stores and mcp servers.
|
||||
|
||||
Returns the object_permission_id if created, otherwise None.
|
||||
"""
|
||||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
if data.object_permission is not None:
|
||||
created_object_permission = (
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=data.object_permission.model_dump(exclude_none=True),
|
||||
)
|
||||
)
|
||||
del data.object_permission
|
||||
return created_object_permission.object_permission_id
|
||||
return None
|
||||
|
||||
|
||||
def validate_team_org_change(
|
||||
team: LiteLLM_TeamTable, organization: LiteLLM_OrganizationTable, llm_router: Router
|
||||
|
|
|
|||
|
|
@ -143,3 +143,38 @@ async def handle_update_object_permission_common(
|
|||
)
|
||||
|
||||
return created_object_permission_row.object_permission_id
|
||||
|
||||
|
||||
async def _set_object_permission(
|
||||
data_json: dict,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
):
|
||||
"""
|
||||
Creates the LiteLLM_ObjectPermissionTable record for the key/team.
|
||||
Handles permissions for vector stores and mcp servers.
|
||||
"""
|
||||
if prisma_client is None or "object_permission" not in data_json:
|
||||
return data_json
|
||||
|
||||
permission_data = data_json["object_permission"]
|
||||
if not isinstance(permission_data, dict):
|
||||
data_json.pop("object_permission")
|
||||
return data_json
|
||||
|
||||
# Clean data: exclude None values and object_permission_id
|
||||
clean_data = {
|
||||
k: v for k, v in permission_data.items()
|
||||
if v is not None and k != "object_permission_id"
|
||||
}
|
||||
|
||||
# Serialize mcp_tool_permissions to JSON string for GraphQL compatibility
|
||||
if "mcp_tool_permissions" in clean_data:
|
||||
clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"])
|
||||
|
||||
created_permission = await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data=clean_data
|
||||
)
|
||||
|
||||
data_json["object_permission_id"] = created_permission.object_permission_id
|
||||
data_json.pop("object_permission")
|
||||
return data_json
|
||||
|
|
@ -1813,3 +1813,95 @@ def test_check_team_key_model_specific_limits_rpm_overallocation():
|
|||
"Allocated RPM limit=800 + Key RPM limit=300 is greater than team RPM limit=1000"
|
||||
in str(exc_info.value.detail)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_key_with_object_permission():
|
||||
"""
|
||||
Test that /key/generate correctly handles object_permission by:
|
||||
1. Creating a record in litellm_objectpermissiontable
|
||||
2. Passing the returned object_permission_id into the key insert payload
|
||||
3. NOT passing the object_permission dict to the key table
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._types import (
|
||||
GenerateKeyRequest,
|
||||
LiteLLM_ObjectPermissionBase,
|
||||
LitellmUserRoles,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_common_key_generation_helper,
|
||||
)
|
||||
|
||||
# Mock prisma client
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.jsonify_object = lambda x: x
|
||||
|
||||
# Mock object permission creation
|
||||
mock_object_perm_create = AsyncMock(
|
||||
return_value=MagicMock(object_permission_id="objperm_key_456")
|
||||
)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.create = mock_object_perm_create
|
||||
|
||||
# Mock key insertion
|
||||
mock_key_insert = AsyncMock(
|
||||
return_value=MagicMock(
|
||||
token="hashed_token_123",
|
||||
litellm_budget_table=None,
|
||||
created_at="2024-01-01T00:00:00Z",
|
||||
updated_at="2024-01-01T00:00:00Z",
|
||||
)
|
||||
)
|
||||
mock_prisma_client.insert_data = mock_key_insert
|
||||
|
||||
# Create request with object_permission
|
||||
key_request = GenerateKeyRequest(
|
||||
models=["gpt-4"],
|
||||
object_permission=LiteLLM_ObjectPermissionBase(
|
||||
vector_stores=["vector_store_1"],
|
||||
mcp_servers=["mcp_server_1"],
|
||||
),
|
||||
)
|
||||
|
||||
mock_admin_auth = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="admin_user",
|
||||
)
|
||||
|
||||
# Patch the prisma_client and other dependencies
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.prisma_client",
|
||||
mock_prisma_client,
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.llm_router", None
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.premium_user",
|
||||
False,
|
||||
), patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.litellm_proxy_admin_name",
|
||||
"admin",
|
||||
):
|
||||
# Execute
|
||||
result = await _common_key_generation_helper(
|
||||
data=key_request,
|
||||
user_api_key_dict=mock_admin_auth,
|
||||
litellm_changed_by=None,
|
||||
team_table=None,
|
||||
)
|
||||
|
||||
# Verify object permission creation was called
|
||||
mock_object_perm_create.assert_awaited_once()
|
||||
|
||||
# Verify key insertion was called
|
||||
assert mock_key_insert.call_count == 1
|
||||
key_insert_kwargs = mock_key_insert.call_args.kwargs
|
||||
key_data = key_insert_kwargs["data"]
|
||||
|
||||
# Verify object_permission_id is in the key data
|
||||
assert key_data.get("object_permission_id") == "objperm_key_456"
|
||||
|
||||
# Verify object_permission dict is NOT in the key data
|
||||
assert "object_permission" not in key_data
|
||||
|
|
|
|||
|
|
@ -254,33 +254,32 @@ async def test_update_team_permissions_success(mock_db_client, mock_admin_auth):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
|
||||
"""Ensure /team/new correctly handles `object_permission` by
|
||||
1. Creating a record in litellm_objectpermissiontable
|
||||
2. Passing the returned `object_permission_id` into the team insert payload
|
||||
"""
|
||||
# --- Configure mocked prisma client ---
|
||||
# Helper identity converters used by team logic
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data # type: ignore
|
||||
Test that /team/new correctly handles object_permission by:
|
||||
1. Creating a record in litellm_objectpermissiontable
|
||||
2. Passing the returned object_permission_id into the team insert payload
|
||||
3. NOT passing the object_permission dict to the team table
|
||||
"""
|
||||
# Configure mocked prisma client
|
||||
mock_db_client.jsonify_team_object = lambda db_data: db_data
|
||||
mock_db_client.get_data = AsyncMock(return_value=None)
|
||||
mock_db_client.update_data = AsyncMock(return_value=MagicMock())
|
||||
|
||||
# Mock DB structure under prisma_client.db
|
||||
mock_db_client.db = MagicMock()
|
||||
|
||||
# 1. Mock object permission table creation
|
||||
# Mock object permission table creation
|
||||
mock_object_perm_create = AsyncMock(
|
||||
return_value=MagicMock(object_permission_id="objperm123")
|
||||
)
|
||||
mock_db_client.db.litellm_objectpermissiontable = MagicMock()
|
||||
mock_db_client.db.litellm_objectpermissiontable.create = mock_object_perm_create
|
||||
|
||||
# 2. Mock model table creation (may be skipped but provided for safety)
|
||||
# Mock model table creation
|
||||
mock_db_client.db.litellm_modeltable = MagicMock()
|
||||
mock_db_client.db.litellm_modeltable.create = AsyncMock(
|
||||
return_value=MagicMock(id="model123")
|
||||
)
|
||||
|
||||
# 3. Capture team table creation and count
|
||||
# Capture team table creation
|
||||
team_create_result = MagicMock(
|
||||
team_id="team-456",
|
||||
object_permission_id="objperm123",
|
||||
|
|
@ -290,9 +289,7 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
|
|||
"object_permission_id": "objperm123",
|
||||
}
|
||||
mock_team_create = AsyncMock(return_value=team_create_result)
|
||||
mock_team_count = AsyncMock(
|
||||
return_value=0
|
||||
) # Mock count to return 0 (no existing teams)
|
||||
mock_team_count = AsyncMock(return_value=0)
|
||||
mock_db_client.db.litellm_teamtable = MagicMock()
|
||||
mock_db_client.db.litellm_teamtable.create = mock_team_create
|
||||
mock_db_client.db.litellm_teamtable.count = mock_team_count
|
||||
|
|
@ -300,23 +297,21 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
|
|||
return_value=team_create_result
|
||||
)
|
||||
|
||||
# 4. Mock user table update behaviour (called for each member)
|
||||
# Mock user table
|
||||
mock_db_client.db.litellm_usertable = MagicMock()
|
||||
mock_db_client.db.litellm_usertable.update = AsyncMock(return_value=MagicMock())
|
||||
|
||||
# --- Import after mocks applied ---
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewTeamRequest
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
||||
# Build request objects
|
||||
# Build request with object_permission
|
||||
team_request = NewTeamRequest(
|
||||
team_alias="my-team",
|
||||
object_permission=LiteLLM_ObjectPermissionBase(vector_stores=["my-vector"]),
|
||||
)
|
||||
|
||||
# Pass a dummy FastAPI Request object
|
||||
dummy_request = MagicMock(spec=Request)
|
||||
|
||||
# Execute the endpoint function
|
||||
|
|
@ -326,14 +321,19 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth):
|
|||
user_api_key_dict=mock_admin_auth,
|
||||
)
|
||||
|
||||
# --- Assertions ---
|
||||
# 1. Object permission creation should be called exactly once
|
||||
# Verify object permission creation was called
|
||||
mock_object_perm_create.assert_awaited_once()
|
||||
|
||||
# 2. Team creation payload should include the generated object_permission_id
|
||||
# Verify team creation was called
|
||||
assert mock_team_create.call_count == 1
|
||||
created_team_kwargs = mock_team_create.call_args.kwargs
|
||||
assert created_team_kwargs["data"].get("object_permission_id") == "objperm123"
|
||||
team_data = created_team_kwargs["data"]
|
||||
|
||||
# Verify object_permission_id is in the team data
|
||||
assert team_data.get("object_permission_id") == "objperm123"
|
||||
|
||||
# Verify object_permission dict is NOT in the team data
|
||||
assert "object_permission" not in team_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -0,0 +1,84 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
)
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_set_object_permission,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_object_permission():
|
||||
"""
|
||||
Test that _set_object_permission correctly:
|
||||
1. Creates an object permission record in the database
|
||||
2. Excludes None values from the data
|
||||
3. Excludes object_permission_id from the data sent to create
|
||||
4. Serializes mcp_tool_permissions to JSON string
|
||||
5. Returns data_json with object_permission_id set and object_permission removed
|
||||
"""
|
||||
# Mock prisma client
|
||||
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
|
||||
)
|
||||
|
||||
# Test data with object_permission
|
||||
data_json = {
|
||||
"user_id": "test_user",
|
||||
"models": ["gpt-4"],
|
||||
"object_permission": {
|
||||
"vector_stores": ["store_1", "store_2"],
|
||||
"mcp_servers": ["server_a"],
|
||||
"mcp_tool_permissions": {
|
||||
"server_a": ["tool1", "tool2"]
|
||||
},
|
||||
"object_permission_id": "should_be_excluded",
|
||||
"mcp_access_groups": None, # This should be excluded
|
||||
}
|
||||
}
|
||||
|
||||
# Call the function
|
||||
result = await _set_object_permission(
|
||||
data_json=data_json,
|
||||
prisma_client=mock_prisma_client
|
||||
)
|
||||
|
||||
# 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"]
|
||||
|
||||
|
|
@ -7,6 +7,7 @@ interface ObjectPermission {
|
|||
object_permission_id: string;
|
||||
mcp_servers: string[];
|
||||
mcp_access_groups?: string[];
|
||||
mcp_tool_permissions?: Record<string, string[]>;
|
||||
vector_stores: string[];
|
||||
}
|
||||
|
||||
|
|
@ -26,11 +27,17 @@ export function ObjectPermissionsView({
|
|||
const vectorStores = objectPermission?.vector_stores || [];
|
||||
const mcpServers = objectPermission?.mcp_servers || [];
|
||||
const mcpAccessGroups = objectPermission?.mcp_access_groups || [];
|
||||
const mcpToolPermissions = objectPermission?.mcp_tool_permissions || {};
|
||||
|
||||
const content = (
|
||||
<div className={variant === "card" ? "grid grid-cols-1 md:grid-cols-2 gap-6" : "space-y-4"}>
|
||||
<VectorStorePermissions vectorStores={vectorStores} accessToken={accessToken} />
|
||||
<MCPServerPermissions mcpServers={mcpServers} mcpAccessGroups={mcpAccessGroups} accessToken={accessToken} />
|
||||
<MCPServerPermissions
|
||||
mcpServers={mcpServers}
|
||||
mcpAccessGroups={mcpAccessGroups}
|
||||
mcpToolPermissions={mcpToolPermissions}
|
||||
accessToken={accessToken}
|
||||
/>
|
||||
</div>
|
||||
);
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,356 @@
|
|||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import MCPServerPermissions from "./MCPServerPermissions";
|
||||
import * as networking from "../networking";
|
||||
|
||||
vi.mock("../networking");
|
||||
|
||||
describe("MCPServerPermissions", () => {
|
||||
const mockAccessToken = "test-token";
|
||||
const mockServerId1 = "3e64bed6-57e1-4247-ad5a-4b1a47ae6583";
|
||||
const mockServerId2 = "server-456";
|
||||
const mockServerName1 = "DW_MCP";
|
||||
const mockServerName2 = "Test Server";
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should display MCP servers with their aliases and IDs", async () => {
|
||||
/**
|
||||
* Tests that MCP servers are displayed with their correct aliases and truncated IDs.
|
||||
* This verifies the basic rendering of server information.
|
||||
*/
|
||||
const mockServers = [
|
||||
{
|
||||
server_id: mockServerId1,
|
||||
server_name: mockServerName1,
|
||||
alias: mockServerName1,
|
||||
},
|
||||
{
|
||||
server_id: mockServerId2,
|
||||
server_name: mockServerName2,
|
||||
alias: mockServerName2,
|
||||
},
|
||||
];
|
||||
|
||||
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1, mockServerId2]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for servers to load and display
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(/DW_MCP/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(/Test Server/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify the count badge shows correct number
|
||||
expect(screen.getByText("2")).toBeInTheDocument();
|
||||
|
||||
// Verify API was called
|
||||
expect(networking.fetchMCPServers).toHaveBeenCalledWith(mockAccessToken);
|
||||
});
|
||||
|
||||
it("should display expandable tool permissions for servers when they exist", async () => {
|
||||
/**
|
||||
* Tests that tool permissions can be expanded/collapsed by clicking the server row
|
||||
* and that the tool count is displayed correctly.
|
||||
*/
|
||||
const mockServers = [
|
||||
{
|
||||
server_id: mockServerId1,
|
||||
server_name: mockServerName1,
|
||||
alias: mockServerName1,
|
||||
},
|
||||
];
|
||||
|
||||
const mockToolPermissions = {
|
||||
[mockServerId1]: ["read_wiki_structure", "read_wiki_contents", "ask_question"],
|
||||
};
|
||||
|
||||
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={mockToolPermissions}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for server to load
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(/DW_MCP/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify tool count is shown
|
||||
expect(screen.getByText("3 tools")).toBeInTheDocument();
|
||||
|
||||
// Tools should NOT be visible initially (collapsed state)
|
||||
expect(screen.queryByText("read_wiki_structure")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("read_wiki_contents")).not.toBeInTheDocument();
|
||||
expect(screen.queryByText("ask_question")).not.toBeInTheDocument();
|
||||
|
||||
// Click the server row to expand
|
||||
const serverRow = screen.getByText(/DW_MCP/).closest("div");
|
||||
await userEvent.click(serverRow!);
|
||||
|
||||
// Now tools should be visible
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("read_wiki_structure")).toBeInTheDocument();
|
||||
expect(screen.getByText("read_wiki_contents")).toBeInTheDocument();
|
||||
expect(screen.getByText("ask_question")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Click the server row again to collapse
|
||||
await userEvent.click(serverRow!);
|
||||
|
||||
// Tools should be hidden again
|
||||
await waitFor(() => {
|
||||
expect(screen.queryByText("read_wiki_structure")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should not display tool permissions section when no tools are configured", async () => {
|
||||
/**
|
||||
* Tests that the tool permissions section is not shown when
|
||||
* mcp_tool_permissions is empty or not provided for a server.
|
||||
*/
|
||||
const mockServers = [
|
||||
{
|
||||
server_id: mockServerId1,
|
||||
server_name: mockServerName1,
|
||||
alias: mockServerName1,
|
||||
},
|
||||
];
|
||||
|
||||
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for server to load
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(/DW_MCP/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify no tool count is shown (since there are no tools)
|
||||
expect(screen.queryByText(/tool/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display access groups correctly", async () => {
|
||||
/**
|
||||
* Tests that access groups are displayed with the correct styling
|
||||
* and indicator badges.
|
||||
*/
|
||||
const mockAccessGroups = ["production-group", "development-group"];
|
||||
|
||||
vi.mocked(networking.fetchMCPAccessGroups).mockResolvedValue(mockAccessGroups);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[]}
|
||||
mcpAccessGroups={mockAccessGroups}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for access groups to load
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("production-group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
expect(screen.getByText("development-group")).toBeInTheDocument();
|
||||
expect(screen.getAllByText("(Access Group)")).toHaveLength(2);
|
||||
|
||||
// Verify the count badge shows correct number
|
||||
expect(screen.getByText("2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display both servers and access groups together", async () => {
|
||||
/**
|
||||
* Tests that both MCP servers and access groups can be displayed
|
||||
* simultaneously in the same component.
|
||||
*/
|
||||
const mockServers = [
|
||||
{
|
||||
server_id: mockServerId1,
|
||||
server_name: mockServerName1,
|
||||
alias: mockServerName1,
|
||||
},
|
||||
];
|
||||
|
||||
const mockAccessGroups = ["production-group"];
|
||||
|
||||
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
|
||||
vi.mocked(networking.fetchMCPAccessGroups).mockResolvedValue(mockAccessGroups);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1]}
|
||||
mcpAccessGroups={mockAccessGroups}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for both to load
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(/DW_MCP/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("production-group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify total count is 2 (1 server + 1 access group)
|
||||
expect(screen.getByText("2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display empty state when no servers or access groups are configured", () => {
|
||||
/**
|
||||
* Tests that the empty state message is shown when there are no
|
||||
* MCP servers or access groups to display.
|
||||
*/
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Verify empty state message
|
||||
expect(screen.getByText("No MCP servers or access groups configured")).toBeInTheDocument();
|
||||
|
||||
// Verify count badge shows 0
|
||||
expect(screen.getByText("0")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should handle multiple servers with different tool permissions", async () => {
|
||||
/**
|
||||
* Tests that multiple servers can each have their own tool permissions
|
||||
* displayed correctly without mixing them up.
|
||||
*/
|
||||
const mockServers = [
|
||||
{
|
||||
server_id: mockServerId1,
|
||||
server_name: mockServerName1,
|
||||
alias: mockServerName1,
|
||||
},
|
||||
{
|
||||
server_id: mockServerId2,
|
||||
server_name: mockServerName2,
|
||||
alias: mockServerName2,
|
||||
},
|
||||
];
|
||||
|
||||
const mockToolPermissions = {
|
||||
[mockServerId1]: ["read_wiki_structure", "read_wiki_contents"],
|
||||
[mockServerId2]: ["ask_question"],
|
||||
};
|
||||
|
||||
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1, mockServerId2]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={mockToolPermissions}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Wait for servers to load
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(/DW_MCP/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Test Server/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify both servers show tool counts
|
||||
expect(screen.getByText("2 tools")).toBeInTheDocument(); // Server 1
|
||||
expect(screen.getByText("1 tool")).toBeInTheDocument(); // Server 2
|
||||
|
||||
// Expand both servers by clicking their rows
|
||||
const server1Row = screen.getByText(/DW_MCP/).closest("div");
|
||||
const server2Row = screen.getByText(/Test Server/).closest("div");
|
||||
|
||||
await userEvent.click(server1Row!); // Expand server 1
|
||||
await userEvent.click(server2Row!); // Expand server 2
|
||||
|
||||
// Verify server 1 tools are now visible
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("read_wiki_structure")).toBeInTheDocument();
|
||||
expect(screen.getByText("read_wiki_contents")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify server 2 tools are now visible
|
||||
expect(screen.getByText("ask_question")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should handle API errors gracefully", async () => {
|
||||
/**
|
||||
* Tests that the component doesn't crash when API calls fail
|
||||
* and falls back to showing server IDs instead of names.
|
||||
*/
|
||||
vi.mocked(networking.fetchMCPServers).mockRejectedValue(
|
||||
new Error("Failed to fetch servers")
|
||||
);
|
||||
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={mockAccessToken}
|
||||
/>
|
||||
);
|
||||
|
||||
// Should still render with server ID (fallback)
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText(mockServerId1)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Verify error was logged
|
||||
expect(networking.fetchMCPServers).toHaveBeenCalledWith(mockAccessToken);
|
||||
});
|
||||
|
||||
it("should not fetch server details when accessToken is not provided", () => {
|
||||
/**
|
||||
* Tests that the component doesn't attempt to fetch server details
|
||||
* when no access token is provided.
|
||||
*/
|
||||
render(
|
||||
<MCPServerPermissions
|
||||
mcpServers={[mockServerId1]}
|
||||
mcpAccessGroups={[]}
|
||||
mcpToolPermissions={{}}
|
||||
accessToken={null}
|
||||
/>
|
||||
);
|
||||
|
||||
// API should not be called without token
|
||||
expect(networking.fetchMCPServers).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
import React, { useState, useEffect } from "react";
|
||||
import { Text, Badge } from "@tremor/react";
|
||||
import { ServerIcon } from "@heroicons/react/outline";
|
||||
import { ServerIcon, ChevronDownIcon, ChevronRightIcon } from "@heroicons/react/outline";
|
||||
import { Tooltip } from "antd";
|
||||
import { fetchMCPServers } from "../networking";
|
||||
import { MCPServer } from "../mcp_tools/types";
|
||||
|
|
@ -8,12 +8,31 @@ import { MCPServer } from "../mcp_tools/types";
|
|||
interface MCPServerPermissionsProps {
|
||||
mcpServers: string[];
|
||||
mcpAccessGroups?: string[];
|
||||
mcpToolPermissions?: Record<string, string[]>;
|
||||
accessToken?: string | null;
|
||||
}
|
||||
|
||||
export function MCPServerPermissions({ mcpServers, mcpAccessGroups = [], accessToken }: MCPServerPermissionsProps) {
|
||||
export function MCPServerPermissions({
|
||||
mcpServers,
|
||||
mcpAccessGroups = [],
|
||||
mcpToolPermissions = {},
|
||||
accessToken
|
||||
}: MCPServerPermissionsProps) {
|
||||
const [mcpServerDetails, setMCPServerDetails] = useState<MCPServer[]>([]);
|
||||
const [accessGroupNames, setAccessGroupNames] = useState<string[]>([]);
|
||||
const [expandedServers, setExpandedServers] = useState<Set<string>>(new Set());
|
||||
|
||||
const toggleServerExpansion = (serverId: string) => {
|
||||
setExpandedServers((prev) => {
|
||||
const newSet = new Set(prev);
|
||||
if (newSet.has(serverId)) {
|
||||
newSet.delete(serverId);
|
||||
} else {
|
||||
newSet.add(serverId);
|
||||
}
|
||||
return newSet;
|
||||
});
|
||||
};
|
||||
|
||||
// Fetch MCP server details when component mounts
|
||||
useEffect(() => {
|
||||
|
|
@ -74,32 +93,80 @@ export function MCPServerPermissions({ mcpServers, mcpAccessGroups = [], accessT
|
|||
return (
|
||||
<div className="space-y-3">
|
||||
<div className="flex items-center gap-2">
|
||||
<ServerIcon className="h-4 w-4 text-gray-600" />
|
||||
<ServerIcon className="h-4 w-4 text-blue-600" />
|
||||
<Text className="font-semibold text-gray-900">MCP Servers</Text>
|
||||
<Badge color="gray" size="xs">
|
||||
<Badge color="blue" size="xs">
|
||||
{totalCount}
|
||||
</Badge>
|
||||
</div>
|
||||
|
||||
{totalCount > 0 ? (
|
||||
<div className="flex flex-wrap gap-2">
|
||||
{mergedItems.map((item, index) =>
|
||||
item.type === "server" ? (
|
||||
<Tooltip key={index} title={`Full ID: ${item.value}`} placement="top">
|
||||
<div className="inline-flex items-center px-3 py-1.5 rounded-lg bg-blue-50 border border-blue-200 text-blue-800 text-sm font-medium cursor-help">
|
||||
{getMCPServerDisplayName(item.value)}
|
||||
<div className="max-h-[400px] overflow-y-auto space-y-2 pr-1">
|
||||
{mergedItems.map((item, index) => {
|
||||
const toolsForServer = item.type === "server" ? mcpToolPermissions[item.value] : undefined;
|
||||
const hasToolRestrictions = toolsForServer && toolsForServer.length > 0;
|
||||
const isExpanded = expandedServers.has(item.value);
|
||||
|
||||
return (
|
||||
<div key={index} className="space-y-2">
|
||||
<div
|
||||
onClick={() => hasToolRestrictions && toggleServerExpansion(item.value)}
|
||||
className={`flex items-center gap-3 py-2 px-3 rounded-lg border border-gray-200 transition-all ${
|
||||
hasToolRestrictions
|
||||
? 'cursor-pointer hover:bg-gray-50 hover:border-gray-300'
|
||||
: 'bg-white'
|
||||
}`}
|
||||
>
|
||||
<div className="flex items-center gap-2 flex-1 min-w-0">
|
||||
{item.type === "server" ? (
|
||||
<Tooltip title={`Full ID: ${item.value}`} placement="top">
|
||||
<div className="inline-flex items-center gap-2 min-w-0">
|
||||
<span className="inline-block w-1.5 h-1.5 bg-blue-500 rounded-full flex-shrink-0"></span>
|
||||
<span className="text-sm font-medium text-gray-900 truncate">{getMCPServerDisplayName(item.value)}</span>
|
||||
</div>
|
||||
</Tooltip>
|
||||
) : (
|
||||
<div className="inline-flex items-center gap-2 min-w-0">
|
||||
<span className="inline-block w-1.5 h-1.5 bg-green-500 rounded-full flex-shrink-0"></span>
|
||||
<span className="text-sm font-medium text-gray-900 truncate">{getAccessGroupDisplayName(item.value)}</span>
|
||||
<span className="ml-1 px-1.5 py-0.5 text-[9px] font-semibold text-green-600 bg-green-50 border border-green-200 rounded uppercase tracking-wide flex-shrink-0">
|
||||
Group
|
||||
</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{hasToolRestrictions && (
|
||||
<div className="flex items-center gap-1 flex-shrink-0 whitespace-nowrap">
|
||||
<span className="text-xs font-medium text-gray-600">{toolsForServer.length}</span>
|
||||
<span className="text-xs text-gray-500">{toolsForServer.length === 1 ? "tool" : "tools"}</span>
|
||||
{isExpanded ? (
|
||||
<ChevronDownIcon className="h-3.5 w-3.5 text-gray-400 ml-0.5" />
|
||||
) : (
|
||||
<ChevronRightIcon className="h-3.5 w-3.5 text-gray-400 ml-0.5" />
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</Tooltip>
|
||||
) : (
|
||||
<div
|
||||
key={index}
|
||||
className="inline-flex items-center px-3 py-1.5 rounded-lg bg-green-50 border border-green-200 text-green-800 text-sm font-medium"
|
||||
>
|
||||
<span className="inline-block w-2 h-2 bg-green-500 rounded-full mr-2"></span>
|
||||
{getAccessGroupDisplayName(item.value)}{" "}
|
||||
<span className="ml-1 text-xs text-green-500">(Access Group)</span>
|
||||
|
||||
{/* Show tool permissions if expanded */}
|
||||
{hasToolRestrictions && isExpanded && (
|
||||
<div className="ml-4 pl-4 border-l-2 border-blue-200 pb-1">
|
||||
<div className="flex flex-wrap gap-1.5">
|
||||
{toolsForServer.map((tool, toolIndex) => (
|
||||
<span
|
||||
key={toolIndex}
|
||||
className="inline-flex items-center px-2.5 py-1 rounded-lg bg-blue-50 border border-blue-200 text-blue-800 text-xs font-medium"
|
||||
>
|
||||
{tool}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
),
|
||||
)}
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
) : (
|
||||
<div className="flex items-center gap-2 px-3 py-2 rounded-lg bg-gray-50 border border-gray-200">
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue