fix(mcp): clear active MCP servers after cleanup (#4586)
This commit is contained in:
committed by
GitHub
parent
5b8f6c7174
commit
3e6715573d
+25
-21
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user