From fe887a23c2a7b96a4c34d206c91f83aca427ac63 Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 3 Mar 2026 14:01:42 -0800 Subject: [PATCH] fixed mcp api --- .../mcp_management_endpoints.py | 35 +++++++-- .../test_mcp_management_endpoints.py | 73 +++++++++++++++++++ 2 files changed, 102 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 8c4d4e7937e..828d2a3fd62 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -689,8 +689,23 @@ if MCP_AVAILABLE: "Database not connected. Connect a database to your proxy" ) - # check to see if server exists for all users + # check to see if server exists (DB first, then registry for config-based servers) mcp_server = await get_mcp_server(prisma_client, server_id) + from_db = mcp_server is not None + + if mcp_server is None: + # Fallback: check registry (config-based servers) - list endpoint uses get_registry() + 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 + ) + if registry_server is not None: + mcp_server = global_mcp_server_manager._build_mcp_server_table( + registry_server + ) + if mcp_server is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -706,10 +721,17 @@ if MCP_AVAILABLE: if not is_admin_view: # Perform authz check BEFORE any health check (avoid side-effects for # unauthorized callers). - mcp_server_records = await get_all_mcp_servers_for_user( - prisma_client, user_api_key_dict - ) - exists = does_mcp_server_exist(mcp_server_records, server_id) + if from_db: + mcp_server_records = await get_all_mcp_servers_for_user( + prisma_client, user_api_key_dict + ) + exists = does_mcp_server_exist(mcp_server_records, server_id) + else: + # Registry/config server: use same access logic as list endpoint + allowed_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers( + user_api_key_dict + ) + exists = mcp_server.server_id in allowed_server_ids if not exists: raise HTTPException( @@ -723,7 +745,8 @@ if MCP_AVAILABLE: ) # At this point caller is authorized to view the server. - await global_mcp_server_manager.add_server(mcp_server) + if from_db: + await global_mcp_server_manager.add_server(mcp_server) # Perform health check on the server using server manager try: 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 e81c6264f7b..8a4e66377cf 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 @@ -796,6 +796,79 @@ class TestListMCPServers: assert not hasattr(result, "credentials") assert result.status == "healthy" + @pytest.mark.asyncio + async def test_fetch_single_mcp_server_from_registry_config_based(self): + """ + Test that fetch_mcp_server finds config-based servers when not in DB. + Config servers appear in list via get_registry() but were 404 on fetch. + """ + 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_health_result = generate_mock_mcp_server_db_record( + server_id="serper_custom_dev", alias="Serper 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 == "serper_custom_dev" 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="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=mock_health_result) + + 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( + server_id="serper_custom_dev", user_api_key_dict=mock_user_auth + ) + + assert result.server_id == "serper_custom_dev" + assert result.status == "healthy" + mock_manager.get_mcp_server_by_id.assert_called_with("serper_custom_dev") + mock_manager._build_mcp_server_table.assert_called_once() + class TestTemporaryMCPSessionEndpoints: def test_inherit_credentials_from_existing_server(self):