From c96fcb9ab5b1bc408faf49f3c6f268ac9ef8eff9 Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 3 Mar 2026 15:36:40 -0800 Subject: [PATCH] resolved greptile comemnt --- .../mcp_management_endpoints.py | 6 +- .../test_mcp_management_endpoints.py | 100 +++++++++++++++++- 2 files changed, 100 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 828d2a3fd62..428784a7ccc 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -673,6 +673,7 @@ if MCP_AVAILABLE: response_model=LiteLLM_MCPServerTable, ) async def fetch_mcp_server( + request: Request, server_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): @@ -695,11 +696,14 @@ if MCP_AVAILABLE: if mcp_server is None: # Fallback: check registry (config-based servers) - list endpoint uses get_registry() + from litellm.proxy.auth.ip_address_utils import IPAddressUtils + + client_ip = IPAddressUtils.get_mcp_client_ip(request) registry_server = global_mcp_server_manager.get_mcp_server_by_id(server_id) if registry_server is None: # Try lookup by server_name or alias (client may use display name in URL) registry_server = global_mcp_server_manager.get_mcp_server_by_name( - server_id, client_ip=None + server_id, client_ip=client_ip ) if registry_server is not None: mcp_server = global_mcp_server_manager._build_mcp_server_table( 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 29f98a581a4..653db775ef5 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 @@ -76,6 +76,15 @@ def generate_mock_mcp_server_config_record( ) +def _make_mock_request(ip: str = "127.0.0.1"): + """Create a mock Request for fetch_mcp_server tests (IP used for access control).""" + req = MagicMock() + req.client = MagicMock() + req.client.host = ip + req.headers = {} + return req + + def generate_mock_user_api_key_auth( user_role: LitellmUserRoles = LitellmUserRoles.PROXY_ADMIN, user_id: str = "test_user_id", @@ -735,7 +744,9 @@ class TestListMCPServers: ) result = await fetch_mcp_server( - server_id="server-1", user_api_key_dict=mock_user_auth + request=_make_mock_request(), + server_id="server-1", + user_api_key_dict=mock_user_auth, ) assert result.server_id == "server-1" @@ -788,7 +799,9 @@ class TestListMCPServers: ) result = await fetch_mcp_server( - server_id="server-2", user_api_key_dict=mock_user_auth + request=_make_mock_request(), + server_id="server-2", + user_api_key_dict=mock_user_auth, ) assert result.server_id == "server-2" @@ -861,7 +874,9 @@ class TestListMCPServers: ) result = await fetch_mcp_server( - server_id="serper_custom_dev", user_api_key_dict=mock_user_auth + request=_make_mock_request(), + server_id="serper_custom_dev", + user_api_key_dict=mock_user_auth, ) assert result.server_id == "serper_custom_dev" @@ -869,6 +884,77 @@ 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_by_name_passes_client_ip(self): + """ + When lookup by server_id fails, fallback to get_mcp_server_by_name. + Verify client_ip is passed for IP-based access control (security). + """ + config_server = generate_mock_mcp_server_config_record( + server_id="serper_custom_dev", + name="Serper MCP", + url="https://serper.example.com/mcp", + transport="http", + ) + + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id = MagicMock(return_value=None) + mock_manager.get_mcp_server_by_name = MagicMock(return_value=config_server) + mock_manager._build_mcp_server_table = MagicMock( + return_value=generate_mock_mcp_server_db_record( + server_id="serper_custom_dev", + alias="Serper MCP", + url="https://serper.example.com/mcp", + transport="http", + ) + ) + mock_manager.get_allowed_mcp_servers = AsyncMock( + return_value=["serper_custom_dev"] + ) + mock_manager.health_check_server = AsyncMock( + return_value=generate_mock_mcp_server_db_record( + server_id="serper_custom_dev", alias="Serper MCP" + ) + ) + + mock_user_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + + 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=True, + ), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_server, + ) + + result = await fetch_mcp_server( + request=_make_mock_request(ip="192.168.1.100"), + server_id="Serper MCP", + user_api_key_dict=mock_user_auth, + ) + + assert result.server_id == "serper_custom_dev" + mock_manager.get_mcp_server_by_id.assert_called_with("Serper MCP") + mock_manager.get_mcp_server_by_name.assert_called_once_with( + "Serper MCP", client_ip="192.168.1.100" + ) + @pytest.mark.asyncio async def test_fetch_single_mcp_server_from_registry_non_admin_denied(self): """ @@ -926,7 +1012,9 @@ class TestListMCPServers: with pytest.raises(HTTPException) as exc_info: await fetch_mcp_server( - server_id="restricted_server", user_api_key_dict=mock_user_auth + request=_make_mock_request(), + server_id="restricted_server", + user_api_key_dict=mock_user_auth, ) assert exc_info.value.status_code == 403 @@ -996,7 +1084,9 @@ class TestListMCPServers: ) result = await fetch_mcp_server( - server_id="allowed_config_server", user_api_key_dict=mock_user_auth + request=_make_mock_request(), + server_id="allowed_config_server", + user_api_key_dict=mock_user_auth, ) assert result.server_id == "allowed_config_server"