diff --git a/src/agents/mcp/manager.py b/src/agents/mcp/manager.py index de19fd74..94df096d 100644 --- a/src/agents/mcp/manager.py +++ b/src/agents/mcp/manager.py @@ -367,25 +367,28 @@ class MCPServerManager(AbstractAsyncContextManager["MCPServerManager"]): return True async def _cleanup_all(self) -> None: - for server in reversed(self._all_servers): - try: - await self._cleanup_server(server) - except asyncio.CancelledError as exc: - if not self.suppress_cancelled_error: - raise - log_tool_action_debug( - logger, - get_mcp_server_log_message("Cleanup cancelled for MCP server", server), - exc, - ) - self._errors[server] = exc - except Exception as exc: - log_tool_action_error( - logger, - get_mcp_server_log_message("Failed to cleanup MCP server", server), - exc, - ) - self._errors[server] = exc + try: + for server in reversed(self._all_servers): + try: + await self._cleanup_server(server) + except asyncio.CancelledError as exc: + if not self.suppress_cancelled_error: + raise + log_tool_action_debug( + logger, + get_mcp_server_log_message("Cleanup cancelled for MCP server", server), + exc, + ) + self._errors[server] = exc + except Exception as exc: + log_tool_action_error( + logger, + get_mcp_server_log_message("Failed to cleanup MCP server", server), + exc, + ) + self._errors[server] = exc + finally: + self._refresh_active_servers() async def _run_with_timeout( self, func: Callable[[], Awaitable[Any]], timeout_seconds: float | None @@ -420,8 +423,9 @@ class MCPServerManager(AbstractAsyncContextManager["MCPServerManager"]): def _refresh_active_servers(self) -> None: if self.drop_failed_servers: - failed = set(self._failed_server_set) - self._active_servers = [server for server in self._all_servers if server not in failed] + self._active_servers = [ + server for server in self._all_servers if server in self._connected_servers + ] else: self._active_servers = list(self._all_servers) diff --git a/tests/mcp/test_mcp_server_manager_cleanup_state.py b/tests/mcp/test_mcp_server_manager_cleanup_state.py new file mode 100644 index 00000000..bf6c82ef --- /dev/null +++ b/tests/mcp/test_mcp_server_manager_cleanup_state.py @@ -0,0 +1,45 @@ +import asyncio +from typing import cast +from unittest.mock import AsyncMock, Mock + +import pytest + +from agents.mcp import MCPServer, MCPServerManager + + +@pytest.mark.asyncio +async def test_cleanup_all_removes_cleaned_servers_from_active_servers() -> None: + server = cast(MCPServer, Mock(spec=MCPServer)) + server.connect = AsyncMock() + server.cleanup = AsyncMock() + + manager = MCPServerManager([server]) + assert await manager.connect_all() == [server] + + await manager.cleanup_all() + + assert manager.active_servers == [] + assert manager._connected_servers == set() + + assert await manager.reconnect() == [] + assert manager.active_servers == [] + assert server.connect.await_count == 1 + + assert await manager.connect_all() == [server] + assert server.connect.await_count == 2 + + +@pytest.mark.asyncio +async def test_cleanup_all_refreshes_active_servers_when_cancellation_propagates() -> None: + server = cast(MCPServer, Mock(spec=MCPServer)) + server.connect = AsyncMock() + server.cleanup = AsyncMock(side_effect=asyncio.CancelledError) + + manager = MCPServerManager([server], suppress_cancelled_error=False) + assert await manager.connect_all() == [server] + + with pytest.raises(asyncio.CancelledError): + await manager.cleanup_all() + + assert manager.active_servers == [] + assert manager._connected_servers == set()