Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions src/agents/mcp/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,7 +234,8 @@ async def reconnect(self, *, failed_only: bool = True) -> list[MCPServer]:
If False, cleanup and retry all servers.
"""
if failed_only:
servers_to_retry = self._unique_servers(self.failed_servers)
failed_servers = self._unique_servers(self.failed_servers)
servers_to_retry = await self._cleanup_servers(failed_servers)
else:
await self.cleanup_all()
servers_to_retry = list(self._all_servers)
Expand Down Expand Up @@ -349,8 +350,10 @@ async def _cleanup_server(self, server: MCPServer) -> None:
finally:
self._connected_servers.discard(server)

async def _cleanup_servers(self, servers: Iterable[MCPServer]) -> None:
for server in reversed(list(servers)):
async def _cleanup_servers(self, servers: Iterable[MCPServer]) -> list[MCPServer]:
servers_list = list(servers)
cleaned_servers: set[MCPServer] = set()
for server in reversed(servers_list):
try:
await self._cleanup_server(server)
except asyncio.CancelledError as exc:
Expand All @@ -369,6 +372,9 @@ async def _cleanup_servers(self, servers: Iterable[MCPServer]) -> None:
exc,
)
self.errors[server] = exc
else:
cleaned_servers.add(server)
return [server for server in servers_list if server in cleaned_servers]

async def _connect_all_parallel(self, servers: list[MCPServer]) -> None:
tasks = [
Expand Down
80 changes: 80 additions & 0 deletions tests/mcp/test_mcp_server_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,36 @@ async def read_resource(self, uri: str) -> ReadResourceResult:
return ReadResourceResult(contents=[])


class PartialFailureServer(FlakyServer):
def __init__(self, *, fail_cleanup: bool = False) -> None:
super().__init__(failures=0)
self.fail_cleanup = fail_cleanup
self.cleanup_calls = 0
self.resource_open = False
self._connect_task: asyncio.Task[object] | None = None

@property
def name(self) -> str:
return "partial-failure"

async def connect(self) -> None:
self.connect_calls += 1
self._connect_task = asyncio.current_task()
if self.resource_open:
raise RuntimeError("connect called without cleanup")
self.resource_open = True
if self.connect_calls == 1:
raise RuntimeError("connect failed after opening resource")

async def cleanup(self) -> None:
self.cleanup_calls += 1
if asyncio.current_task() is not self._connect_task:
raise RuntimeError("Attempted to exit cancel scope in a different task")
if self.fail_cleanup:
raise RuntimeError("cleanup failed")
self.resource_open = False


class SensitiveNamedServer(FlakyServer):
def __init__(self, name: str) -> None:
super().__init__(failures=1)
Expand Down Expand Up @@ -400,6 +430,56 @@ async def test_manager_reconnect_failed_only() -> None:
assert manager.failed_servers == []


@pytest.mark.asyncio
@pytest.mark.parametrize("connect_in_parallel", [False, True])
async def test_manager_reconnect_cleans_partial_failure_before_retry(
connect_in_parallel: bool,
) -> None:
healthy_server = CleanupAwareServer()
failed_server = PartialFailureServer()
manager = MCPServerManager(
[healthy_server, failed_server], connect_in_parallel=connect_in_parallel
)
try:
await manager.connect_all()

assert manager.active_servers == [healthy_server]
assert manager.failed_servers == [failed_server]

await manager.reconnect()

assert manager.active_servers == [healthy_server, failed_server]
assert manager.failed_servers == []
assert failed_server not in manager.errors
assert failed_server.connect_calls == 2
assert failed_server.cleanup_calls == 1
assert failed_server.resource_open is True
assert healthy_server.connect_calls == 1
assert healthy_server.cleanup_calls == 0
finally:
await manager.cleanup_all()


@pytest.mark.asyncio
@pytest.mark.parametrize("connect_in_parallel", [False, True])
async def test_manager_reconnect_does_not_retry_after_cleanup_failure(
connect_in_parallel: bool,
) -> None:
server = PartialFailureServer(fail_cleanup=True)
manager = MCPServerManager([server], connect_in_parallel=connect_in_parallel)

await manager.connect_all()
await manager.reconnect()

assert manager.active_servers == []
assert manager.failed_servers == [server]
assert server.connect_calls == 1
assert server.cleanup_calls == 1
assert server.resource_open is True
assert str(manager.errors[server]) == "cleanup failed"
assert manager._workers == {}


@pytest.mark.asyncio
async def test_manager_reconnect_deduplicates_failures() -> None:
server = FlakyServer(failures=2)
Expand Down