From 710e230ed730b0f96952510dd853537aae688468 Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 3 Mar 2026 15:14:51 -0800 Subject: [PATCH] added non-admin test --- .../test_mcp_management_endpoints.py | 134 ++++++++++++++++++ 1 file changed, 134 insertions(+) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 8a4e66377cf..29f98a581a4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -869,6 +869,140 @@ class TestListMCPServers: mock_manager.get_mcp_server_by_id.assert_called_with("serper_custom_dev") mock_manager._build_mcp_server_table.assert_called_once() + @pytest.mark.asyncio + async def test_fetch_single_mcp_server_from_registry_non_admin_denied(self): + """ + Non-admin user: config server NOT in allowed_server_ids -> 403. + """ + config_server = generate_mock_mcp_server_config_record( + server_id="restricted_server", + name="Restricted MCP", + url="https://restricted.example.com/mcp", + transport="http", + ) + + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id = MagicMock( + side_effect=lambda sid: config_server if sid == "restricted_server" else None + ) + mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) + mock_manager._build_mcp_server_table = MagicMock( + return_value=generate_mock_mcp_server_db_record( + server_id="restricted_server", + alias="Restricted MCP", + url="https://restricted.example.com/mcp", + transport="http", + ) + ) + mock_manager.get_allowed_mcp_servers = AsyncMock( + return_value=["other_server"] # restricted_server NOT in list + ) + + mock_user_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.INTERNAL_USER + ) + + 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.get_mcp_server", + AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + return_value=False, + ), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_server, + ) + + with pytest.raises(HTTPException) as exc_info: + await fetch_mcp_server( + server_id="restricted_server", user_api_key_dict=mock_user_auth + ) + + assert exc_info.value.status_code == 403 + mock_manager.get_allowed_mcp_servers.assert_called_once_with(mock_user_auth) + + @pytest.mark.asyncio + async def test_fetch_single_mcp_server_from_registry_non_admin_granted(self): + """ + Non-admin user: config server IS in allowed_server_ids -> 200. + """ + config_server = generate_mock_mcp_server_config_record( + server_id="allowed_config_server", + name="Allowed MCP", + url="https://allowed.example.com/mcp", + transport="http", + ) + + mock_health_result = generate_mock_mcp_server_db_record( + server_id="allowed_config_server", alias="Allowed MCP" + ) + mock_health_result.status = "healthy" + mock_health_result.last_health_check = datetime.now() + mock_health_result.health_check_error = None + + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id = MagicMock( + side_effect=lambda sid: config_server if sid == "allowed_config_server" else None + ) + mock_manager.get_mcp_server_by_name = MagicMock(return_value=None) + mock_manager._build_mcp_server_table = MagicMock( + return_value=generate_mock_mcp_server_db_record( + server_id="allowed_config_server", + alias="Allowed MCP", + url="https://allowed.example.com/mcp", + transport="http", + ) + ) + mock_manager.get_allowed_mcp_servers = AsyncMock( + return_value=["allowed_config_server"] + ) + mock_manager.health_check_server = AsyncMock(return_value=mock_health_result) + + mock_user_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.INTERNAL_USER + ) + + 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.get_mcp_server", + AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + return_value=False, + ), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_server, + ) + + result = await fetch_mcp_server( + server_id="allowed_config_server", user_api_key_dict=mock_user_auth + ) + + assert result.server_id == "allowed_config_server" + assert result.status == "healthy" + mock_manager.get_allowed_mcp_servers.assert_called_once_with(mock_user_auth) + class TestTemporaryMCPSessionEndpoints: def test_inherit_credentials_from_existing_server(self):