diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 6273678ac6e..5838541634d 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1471,23 +1471,34 @@ if MCP_AVAILABLE: payload.submitted_by = None payload.submitted_at = None - # Attempt to create the mcp server + # The database write is the commit point: if it fails nothing was + # persisted and the request is a genuine failure. try: new_mcp_server = await create_mcp_server( prisma_client, payload, touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, ) - await global_mcp_server_manager.add_server(new_mcp_server) - - # Ensure registry is up to date by reloading from database - await global_mcp_server_manager.reload_servers_from_database() except Exception as e: verbose_proxy_logger.exception(f"Error creating mcp server: {str(e)}") raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Error creating mcp server: {str(e)}"}, ) + + # Registry refresh is best-effort: the row is already committed, so a + # failure here (e.g. an unrelated malformed row in the table) must not + # surface as a 500 and orphan the created server, which would push the + # caller to retry and create duplicates. + try: + await global_mcp_server_manager.add_server(new_mcp_server) + await global_mcp_server_manager.reload_servers_from_database() + except Exception as e: + verbose_proxy_logger.exception( + f"MCP server {new_mcp_server.server_id} created but in-memory " + f"registry refresh failed: {str(e)}" + ) + return _redact_mcp_credentials(new_mcp_server) @router.post( 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 873831341b6..b21be0d1746 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 @@ -2407,6 +2407,109 @@ class TestUpdateMCPServer: assert result.alias == "Updated Test Server" +class TestAddMCPServerAtomicity: + """A committed MCP server must survive a post-write registry refresh failure. + + Regression: add_mcp_server inserted the row and then reloaded the whole + registry from the database inside the same try block. One unrelated malformed + row made the reload raise, so the endpoint returned 500 even though the new + row was already persisted. Callers assumed failure and retried, creating + duplicate servers. + """ + + @pytest.mark.asyncio + async def test_create_succeeds_when_registry_refresh_fails(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + add_mcp_server, + ) + + payload = NewMCPServerRequest( + alias="echo", + url="https://echo.example.com/mcp", + transport=MCPTransport.http, + ) + admin = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user" + ) + created_server = generate_mock_mcp_server_db_record( + server_id="created-1", alias="echo" + ) + + mock_manager = MagicMock() + mock_manager.add_server = AsyncMock() + mock_manager.reload_servers_from_database = AsyncMock( + side_effect=Exception("malformed pre-existing row") + ) + + 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.validate_and_normalize_mcp_server_payload", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + AsyncMock(return_value=created_server), + ) as create_mock, + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + ): + result = await add_mcp_server(payload=payload, user_api_key_dict=admin) + + create_mock.assert_awaited_once() + mock_manager.reload_servers_from_database.assert_awaited_once() + assert result.server_id == "created-1" + + @pytest.mark.asyncio + async def test_create_500s_and_skips_registry_when_db_write_fails(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + add_mcp_server, + ) + + payload = NewMCPServerRequest( + alias="echo", + url="https://echo.example.com/mcp", + transport=MCPTransport.http, + ) + admin = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user" + ) + + mock_manager = MagicMock() + mock_manager.add_server = AsyncMock() + mock_manager.reload_servers_from_database = AsyncMock() + + 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.validate_and_normalize_mcp_server_payload", + MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server", + AsyncMock(side_effect=Exception("db down")), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await add_mcp_server(payload=payload, user_api_key_dict=admin) + + assert exc_info.value.status_code == 500 + mock_manager.add_server.assert_not_awaited() + mock_manager.reload_servers_from_database.assert_not_awaited() + + class TestHealthCheckServers: """Test suite for health check servers endpoint"""