diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index fb6c11afba4..68af5b4c9f5 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index b8c0aefded2..d05a0259f46 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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 diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index fb2656b6018..9670cdf330a 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -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 \ No newline at end of file diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index b6c4c609bc9..3626e170782 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index f652819b317..929ec41e09a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py new file mode 100644 index 00000000000..07d89035dce --- /dev/null +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -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"] + diff --git a/ui/litellm-dashboard/src/components/object_permissions_view.tsx b/ui/litellm-dashboard/src/components/object_permissions_view.tsx index d689916e813..77727e06602 100644 --- a/ui/litellm-dashboard/src/components/object_permissions_view.tsx +++ b/ui/litellm-dashboard/src/components/object_permissions_view.tsx @@ -7,6 +7,7 @@ interface ObjectPermission { object_permission_id: string; mcp_servers: string[]; mcp_access_groups?: string[]; + mcp_tool_permissions?: Record; 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 = (
- +
); diff --git a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx new file mode 100644 index 00000000000..6d0018fa160 --- /dev/null +++ b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx @@ -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( + + ); + + // 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( + + ); + + // 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( + + ); + + // 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( + + ); + + // 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( + + ); + + // 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( + + ); + + // 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( + + ); + + // 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( + + ); + + // 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( + + ); + + // API should not be called without token + expect(networking.fetchMCPServers).not.toHaveBeenCalled(); + }); +}); + diff --git a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx index 29b02c12b93..962fc00670a 100644 --- a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx +++ b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx @@ -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; accessToken?: string | null; } -export function MCPServerPermissions({ mcpServers, mcpAccessGroups = [], accessToken }: MCPServerPermissionsProps) { +export function MCPServerPermissions({ + mcpServers, + mcpAccessGroups = [], + mcpToolPermissions = {}, + accessToken +}: MCPServerPermissionsProps) { const [mcpServerDetails, setMCPServerDetails] = useState([]); const [accessGroupNames, setAccessGroupNames] = useState([]); + const [expandedServers, setExpandedServers] = useState>(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 (
- + MCP Servers - + {totalCount}
+ {totalCount > 0 ? ( -
- {mergedItems.map((item, index) => - item.type === "server" ? ( - -
- {getMCPServerDisplayName(item.value)} +
+ {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 ( +
+
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' + }`} + > +
+ {item.type === "server" ? ( + +
+ + {getMCPServerDisplayName(item.value)} +
+
+ ) : ( +
+ + {getAccessGroupDisplayName(item.value)} + + Group + +
+ )} +
+ + {hasToolRestrictions && ( +
+ {toolsForServer.length} + {toolsForServer.length === 1 ? "tool" : "tools"} + {isExpanded ? ( + + ) : ( + + )} +
+ )}
- - ) : ( -
- - {getAccessGroupDisplayName(item.value)}{" "} - (Access Group) + + {/* Show tool permissions if expanded */} + {hasToolRestrictions && isExpanded && ( +
+
+ {toolsForServer.map((tool, toolIndex) => ( + + {tool} + + ))} +
+
+ )}
- ), - )} + ); + })}
) : (