From 75eacdbb1556337f1ab1c121efb754f68b83f422 Mon Sep 17 00:00:00 2001 From: Henry Su Date: Tue, 28 Jul 2026 10:20:27 -0500 Subject: [PATCH] fix(core): isolate shared initialization cancellation --- .../mcp_transport/transport_base.py | 2 +- .../tests/mcp_transport/test_base.py | 34 +++++++++++++++++++ 2 files changed, 35 insertions(+), 1 deletion(-) diff --git a/packages/toolbox-core/src/toolbox_core/mcp_transport/transport_base.py b/packages/toolbox-core/src/toolbox_core/mcp_transport/transport_base.py index e8ac7784c..f979deac5 100644 --- a/packages/toolbox-core/src/toolbox_core/mcp_transport/transport_base.py +++ b/packages/toolbox-core/src/toolbox_core/mcp_transport/transport_base.py @@ -92,7 +92,7 @@ async def _ensure_initialized( self._init_task = asyncio.create_task( self._initialize_session(headers=headers) ) - await self._init_task + await asyncio.shield(self._init_task) @property def base_url(self) -> str: diff --git a/packages/toolbox-core/tests/mcp_transport/test_base.py b/packages/toolbox-core/tests/mcp_transport/test_base.py index 581998455..f56aee504 100644 --- a/packages/toolbox-core/tests/mcp_transport/test_base.py +++ b/packages/toolbox-core/tests/mcp_transport/test_base.py @@ -152,6 +152,40 @@ async def slow_init(*args, **kwargs): transport._initialize_session.assert_called_once() + @pytest.mark.asyncio + async def test_cancelled_waiter_does_not_cancel_shared_initialization(self, mocker): + """Cancelling one waiter leaves the shared initialization running.""" + session = AsyncMock(spec=ClientSession) + transport = ConcreteTransport("http://fake-server.com", session=session) + init_started = asyncio.Event() + allow_init = asyncio.Event() + + async def slow_init(*args, **kwargs): + init_started.set() + await allow_init.wait() + + mocker.patch.object( + transport, + "_initialize_session", + new_callable=AsyncMock, + side_effect=slow_init, + ) + + cancelled_waiter = asyncio.create_task(transport._ensure_initialized()) + await init_started.wait() + surviving_waiter = asyncio.create_task(transport._ensure_initialized()) + await asyncio.sleep(0) + + cancelled_waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled_waiter + + allow_init.set() + assert await asyncio.gather(surviving_waiter, return_exceptions=True) == [None] + + await transport._ensure_initialized() + transport._initialize_session.assert_awaited_once() + def test_convert_tool_schema_valid(self, transport): """Test converting a valid MCP tool schema.""" raw_tool = {