diff --git a/src/mcp/client/session_group.py b/src/mcp/client/session_group.py index a544cecbe8..f602edae4b 100644 --- a/src/mcp/client/session_group.py +++ b/src/mcp/client/session_group.py @@ -281,8 +281,8 @@ async def disconnect_from_server(self, session: mcp.ClientSession) -> None: # Clean up the session's resources via its dedicated exit stack if session_known_for_stack: - session_stack_to_close = self._session_exit_stacks.pop(session) # pragma: no cover - await session_stack_to_close.aclose() # pragma: no cover + session_stack_to_close = self._session_exit_stacks.pop(session) + await session_stack_to_close.aclose() async def connect_with_session( self, server_info: types.Implementation, session: mcp.ClientSession @@ -414,11 +414,6 @@ async def _aggregate_components(self, server_info: types.Implementation, session except MCPError as err: # pragma: no cover logging.warning(f"Could not fetch tools: {err}") - # Clean up exit stack for session if we couldn't retrieve anything - # from the server. - if not any((prompts_temp, resources_temp, tools_temp)): - del self._session_exit_stacks[session] # pragma: no cover - # Check for duplicates. matching_prompts = prompts_temp.keys() & self._prompts.keys() if matching_prompts: diff --git a/tests/client/test_session_group.py b/tests/client/test_session_group.py index b75d22b7a0..d4745d4c03 100644 --- a/tests/client/test_session_group.py +++ b/tests/client/test_session_group.py @@ -402,3 +402,67 @@ async def test_client_session_group_establish_session_parameterized( # 3. Assert returned values assert returned_server_info is mock_initialize_result.server_info assert returned_session is mock_entered_session + + +@pytest.mark.anyio +async def test_connect_with_session_empty_server_does_not_raise(): + """A server exposing no components connects cleanly via connect_with_session. + + Regression: the caller-supplied session is never registered in + _session_exit_stacks, so the old empty-server cleanup deleted a missing key + and raised KeyError. + """ + server_info = mock.Mock(spec=types.Implementation) + server_info.name = "EmptyServer" + session = mock.AsyncMock(spec=mcp.ClientSession) + session.list_tools.return_value = mock.AsyncMock(tools=[]) + session.list_resources.return_value = mock.AsyncMock(resources=[]) + session.list_prompts.return_value = mock.AsyncMock(prompts=[]) + + group = ClientSessionGroup() + await group.connect_with_session(server_info, session) + + assert session in group._sessions + assert not group.tools + assert not group.resources + assert not group.prompts + assert session not in group._session_exit_stacks + + +@pytest.mark.anyio +async def test_connect_to_server_empty_server_keeps_exit_stack( + mock_exit_stack: contextlib.AsyncExitStack, +): + """An empty server connected via connect_to_server retains its exit stack. + + Regression: the old cleanup dropped the freshly-registered stack, so a later + disconnect could not close the transport. + """ + server_info = mock.Mock(spec=types.Implementation) + server_info.name = "EmptyServer" + session = mock.AsyncMock(spec=mcp.ClientSession) + session.list_tools.return_value = mock.AsyncMock(tools=[]) + session.list_resources.return_value = mock.AsyncMock(resources=[]) + session.list_prompts.return_value = mock.AsyncMock(prompts=[]) + session_stack = mock.AsyncMock(spec=contextlib.AsyncExitStack) + + group = ClientSessionGroup(exit_stack=mock_exit_stack) + + async def fake_establish( + server_params: StdioServerParameters, + session_params: ClientSessionParameters, + ) -> tuple[types.Implementation, mcp.ClientSession]: + group._session_exit_stacks[session] = session_stack + return server_info, session + + with mock.patch.object(group, "_establish_session", side_effect=fake_establish): + await group.connect_to_server(StdioServerParameters(command="test")) + + assert session in group._sessions + assert group._session_exit_stacks[session] is session_stack + + await group.disconnect_from_server(session) + + assert session not in group._sessions + assert session not in group._session_exit_stacks + session_stack.aclose.assert_awaited_once()