fix(mcp): reconnect terminated sessions
This commit is contained in:
@@ -13,6 +13,7 @@ from nanobot.agent.tools.mcp import (
|
||||
MCPPromptWrapper,
|
||||
MCPResourceWrapper,
|
||||
MCPToolWrapper,
|
||||
_is_session_terminated,
|
||||
_is_transient,
|
||||
)
|
||||
|
||||
@@ -35,6 +36,14 @@ class _FakeEndOfStreamError(Exception):
|
||||
_FakeEndOfStreamError.__name__ = "EndOfStream"
|
||||
|
||||
|
||||
def _session_terminated_error() -> McpError:
|
||||
return McpError(ErrorData(code=-32000, message="Session terminated"))
|
||||
|
||||
|
||||
def _connection_closed_error() -> McpError:
|
||||
return McpError(ErrorData(code=-32000, message="Connection closed"))
|
||||
|
||||
|
||||
def test_is_transient_recognizes_closed_resource():
|
||||
assert _is_transient(_FakeClosedResourceError("gone"))
|
||||
|
||||
@@ -67,6 +76,14 @@ def test_is_transient_rejects_timeout():
|
||||
assert not _is_transient(TimeoutError("timeout"))
|
||||
|
||||
|
||||
def test_is_session_terminated_recognizes_mcp_error():
|
||||
assert _is_session_terminated(_session_terminated_error())
|
||||
|
||||
|
||||
def test_is_session_terminated_recognizes_connection_closed_mcp_error():
|
||||
assert _is_session_terminated(_connection_closed_error())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MCPToolWrapper retry behaviour
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -219,6 +236,35 @@ async def test_tool_retry_on_end_of_stream():
|
||||
assert session.call_tool.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_reconnects_when_transient_retry_reveals_terminated_session():
|
||||
"""Tool should reconnect if a stale session reports termination after transient retry."""
|
||||
old_session = AsyncMock()
|
||||
old_session.call_tool = AsyncMock(
|
||||
side_effect=[_FakeClosedResourceError("closed"), _session_terminated_error()]
|
||||
)
|
||||
new_session = AsyncMock()
|
||||
new_session.call_tool = AsyncMock(return_value=_make_tool_result("fresh"))
|
||||
|
||||
wrapper = MCPToolWrapper(old_session, "test_server", _make_tool_def(), tool_timeout=5)
|
||||
replacement = MCPToolWrapper(new_session, "test_server", _make_tool_def(), tool_timeout=5)
|
||||
|
||||
async def reconnect(server_name: str, tool_name: str, stale_tool):
|
||||
assert server_name == "test_server"
|
||||
assert tool_name == "mcp_test_server_test_tool"
|
||||
assert stale_tool is wrapper
|
||||
return replacement
|
||||
|
||||
wrapper.set_reconnect_handler(reconnect)
|
||||
|
||||
with patch("nanobot.agent.tools.mcp.asyncio.sleep", new_callable=AsyncMock):
|
||||
output = await wrapper.execute(foo="bar")
|
||||
|
||||
assert output == "fresh"
|
||||
assert old_session.call_tool.call_count == 2
|
||||
assert new_session.call_tool.call_count == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MCPResourceWrapper retry behaviour
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -284,6 +330,32 @@ async def test_resource_no_retry_on_non_transient():
|
||||
assert session.read_resource.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resource_reconnects_on_session_terminated():
|
||||
"""Resource should reconnect once when the MCP SDK reports a dead session."""
|
||||
old_session = AsyncMock()
|
||||
old_session.read_resource = AsyncMock(side_effect=_session_terminated_error())
|
||||
new_session = AsyncMock()
|
||||
new_session.read_resource = AsyncMock(return_value=_make_resource_result("fresh"))
|
||||
|
||||
wrapper = MCPResourceWrapper(old_session, "test_server", _make_resource_def())
|
||||
replacement = MCPResourceWrapper(new_session, "test_server", _make_resource_def())
|
||||
|
||||
async def reconnect(server_name: str, tool_name: str, stale_tool):
|
||||
assert server_name == "test_server"
|
||||
assert tool_name == "mcp_test_server_resource_test_resource"
|
||||
assert stale_tool is wrapper
|
||||
return replacement
|
||||
|
||||
wrapper.set_reconnect_handler(reconnect)
|
||||
|
||||
output = await wrapper.execute()
|
||||
|
||||
assert output == "fresh"
|
||||
assert old_session.read_resource.call_count == 1
|
||||
assert new_session.read_resource.call_count == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MCPPromptWrapper retry behaviour
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -366,3 +438,29 @@ async def test_prompt_no_retry_on_non_transient():
|
||||
|
||||
assert "RuntimeError" in output
|
||||
assert session.get_prompt.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_reconnects_on_session_terminated():
|
||||
"""Prompt should reconnect once before falling back to McpError handling."""
|
||||
old_session = AsyncMock()
|
||||
old_session.get_prompt = AsyncMock(side_effect=_session_terminated_error())
|
||||
new_session = AsyncMock()
|
||||
new_session.get_prompt = AsyncMock(return_value=_make_prompt_result("fresh prompt"))
|
||||
|
||||
wrapper = MCPPromptWrapper(old_session, "test_server", _make_prompt_def())
|
||||
replacement = MCPPromptWrapper(new_session, "test_server", _make_prompt_def())
|
||||
|
||||
async def reconnect(server_name: str, tool_name: str, stale_tool):
|
||||
assert server_name == "test_server"
|
||||
assert tool_name == "mcp_test_server_prompt_test_prompt"
|
||||
assert stale_tool is wrapper
|
||||
return replacement
|
||||
|
||||
wrapper.set_reconnect_handler(reconnect)
|
||||
|
||||
output = await wrapper.execute()
|
||||
|
||||
assert output == "fresh prompt"
|
||||
assert old_session.get_prompt.call_count == 1
|
||||
assert new_session.get_prompt.call_count == 1
|
||||
|
||||
Reference in New Issue
Block a user