fix(mcp): clear active MCP servers after cleanup (#4586)
This commit is contained in:
committed by
GitHub
parent
5b8f6c7174
commit
3e6715573d
@@ -367,6 +367,7 @@ class MCPServerManager(AbstractAsyncContextManager["MCPServerManager"]):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
async def _cleanup_all(self) -> None:
|
async def _cleanup_all(self) -> None:
|
||||||
|
try:
|
||||||
for server in reversed(self._all_servers):
|
for server in reversed(self._all_servers):
|
||||||
try:
|
try:
|
||||||
await self._cleanup_server(server)
|
await self._cleanup_server(server)
|
||||||
@@ -386,6 +387,8 @@ class MCPServerManager(AbstractAsyncContextManager["MCPServerManager"]):
|
|||||||
exc,
|
exc,
|
||||||
)
|
)
|
||||||
self._errors[server] = exc
|
self._errors[server] = exc
|
||||||
|
finally:
|
||||||
|
self._refresh_active_servers()
|
||||||
|
|
||||||
async def _run_with_timeout(
|
async def _run_with_timeout(
|
||||||
self, func: Callable[[], Awaitable[Any]], timeout_seconds: float | None
|
self, func: Callable[[], Awaitable[Any]], timeout_seconds: float | None
|
||||||
@@ -420,8 +423,9 @@ class MCPServerManager(AbstractAsyncContextManager["MCPServerManager"]):
|
|||||||
|
|
||||||
def _refresh_active_servers(self) -> None:
|
def _refresh_active_servers(self) -> None:
|
||||||
if self.drop_failed_servers:
|
if self.drop_failed_servers:
|
||||||
failed = set(self._failed_server_set)
|
self._active_servers = [
|
||||||
self._active_servers = [server for server in self._all_servers if server not in failed]
|
server for server in self._all_servers if server in self._connected_servers
|
||||||
|
]
|
||||||
else:
|
else:
|
||||||
self._active_servers = list(self._all_servers)
|
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