From 52ff2b45df8293ce19cfdcc869d46d145f3268da Mon Sep 17 00:00:00 2001 From: Henry Su Date: Tue, 28 Jul 2026 10:25:24 -0500 Subject: [PATCH] fix(core): close session after cancelled initialization --- .../mcp_transport/transport_base.py | 2 +- .../tests/mcp_transport/test_base.py | 31 +++++++++++++++++++ 2 files changed, 32 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..1bd73d722 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 @@ -246,7 +246,7 @@ async def close(self): if self._init_task: try: await self._init_task - except Exception: + except (asyncio.CancelledError, Exception): # If initialization failed, we can still try to close. pass if self._manage_session and self._session and not self._session.closed: diff --git a/packages/toolbox-core/tests/mcp_transport/test_base.py b/packages/toolbox-core/tests/mcp_transport/test_base.py index 581998455..1bca6780b 100644 --- a/packages/toolbox-core/tests/mcp_transport/test_base.py +++ b/packages/toolbox-core/tests/mcp_transport/test_base.py @@ -315,6 +315,37 @@ async def test_close_managed_session(self, mocker): await transport.close() mock_close.assert_called_once() + @pytest.mark.asyncio + async def test_close_managed_session_after_cancelled_initialization(self): + transport = ConcreteTransport("http://fake-server.com") + transport._init_task = asyncio.create_task(asyncio.sleep(0)) + transport._init_task.cancel() + + with pytest.raises(asyncio.CancelledError): + await transport._init_task + + try: + await transport.close() + assert transport._session.closed + finally: + if not transport._session.closed: + await transport._session.close() + + @pytest.mark.asyncio + async def test_close_managed_session_after_initialization_error(self): + async def fail_initialization(): + raise RuntimeError("initialization failed") + + transport = ConcreteTransport("http://fake-server.com") + transport._init_task = asyncio.create_task(fail_initialization()) + + try: + await transport.close() + assert transport._session.closed + finally: + if not transport._session.closed: + await transport._session.close() + @pytest.mark.asyncio async def test_close_unmanaged_session(self): mock_session = AsyncMock(spec=ClientSession)