Files
EvoScientist/tests/test_tool_error_handler.py
dinos 690b903f85 test: standardize async tests on pytest-asyncio auto mode (#338)
* chore: add pytest-asyncio in auto mode

* test: migrate channel and stream tests to native async

Convert run_async() wrapper tests to plain 'async def test_*' under
pytest-asyncio auto mode. collect_events() in stream_v3_fakes becomes a
coroutine awaited at every call site.

* test: migrate command and model/middleware tests to native async

Convert run_async() wrappers (import, alias, and fixture forms) to plain
'async def test_*'. Multi-call tests merge onto one loop as sequential
awaits; none asserted on loop identity.

* test: migrate TUI, notifier, gateway, and session tests to native async

TUI/notifier/gateway files convert run_async wrappers to plain async
tests. test_sessions.py's unittest.TestCase classes move to
unittest.IsolatedAsyncioTestCase (pytest-asyncio does not await async
methods on plain TestCase; converting blindly would have made ~70 tests
silently vacuous). Its setUpClass keeps a one-shot asyncio.run() since
IsolatedAsyncioTestCase has no async class-level hook. TestLoadingWidget
in test_tui_widgets.py drops its TestCase base for the same reason.

* test: replace direct asyncio.run() calls with native async tests

Convert tests that called asyncio.run() (directly or via a local _run
helper) to plain 'async def test_*'; delete the local helpers.

* test: drop undeclared anyio markers and delete run_async helper

The @pytest.mark.anyio tests relied on anyio being a transitive dep of
httpx; auto-mode pytest-asyncio collects them natively. run_async() and
its fixture are unreferenced after the migration, so remove them —
pytest-asyncio's per-test loop teardown covers the pending-task
cancellation the helper existed for (verified: full suite runs with no
'Event loop is closed' errors or destroyed-task warnings).
2026-07-08 18:37:48 +00:00

216 lines
6.9 KiB
Python

"""Tests for ToolErrorHandlerMiddleware."""
from unittest.mock import MagicMock
import pytest
from langchain_core.messages import ToolMessage
from langgraph.types import Command
from EvoScientist.middleware.tool_error_handler import (
ToolErrorHandlerMiddleware,
_build_error_message,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_request(tool_name: str = "my_mcp_tool", call_id: str = "tc_001"):
"""Create a minimal ToolCallRequest-like object."""
req = MagicMock()
req.tool_call = {"id": call_id, "name": tool_name, "args": {"query": "test"}}
return req
# ---------------------------------------------------------------------------
# _build_error_message
# ---------------------------------------------------------------------------
class TestBuildErrorMessage:
def test_returns_tool_message(self):
req = _make_request("broken_tool", "tc_123")
try:
raise RuntimeError("connection refused")
except RuntimeError:
msg = _build_error_message(req)
assert isinstance(msg, ToolMessage)
assert msg.tool_call_id == "tc_123"
assert msg.name == "broken_tool"
assert msg.status == "error"
def test_content_includes_tool_name(self):
req = _make_request("broken_tool")
try:
raise ValueError("bad value")
except ValueError:
msg = _build_error_message(req)
assert "broken_tool" in msg.content
def test_content_includes_traceback(self):
req = _make_request()
try:
raise RuntimeError("something went wrong")
except RuntimeError:
msg = _build_error_message(req)
assert "RuntimeError" in msg.content
assert "something went wrong" in msg.content
assert "Traceback" in msg.content
def test_content_includes_retry_guidance(self):
req = _make_request()
try:
raise Exception("fail")
except Exception:
msg = _build_error_message(req)
assert "retry" in msg.content.lower()
# ---------------------------------------------------------------------------
# Sync: wrap_tool_call
# ---------------------------------------------------------------------------
class TestWrapToolCallSync:
def setup_method(self):
self.mw = ToolErrorHandlerMiddleware()
def test_success_passes_through(self):
expected = ToolMessage(content="ok", tool_call_id="tc_001", name="t")
handler = MagicMock(return_value=expected)
req = _make_request()
result = self.mw.wrap_tool_call(req, handler)
assert result is expected
handler.assert_called_once_with(req)
def test_command_passes_through(self):
cmd = Command(update={"messages": []})
handler = MagicMock(return_value=cmd)
req = _make_request()
result = self.mw.wrap_tool_call(req, handler)
assert result is cmd
def test_exception_returns_error_tool_message(self):
handler = MagicMock(side_effect=RuntimeError("MCP server crashed"))
req = _make_request("flaky_tool", "tc_999")
result = self.mw.wrap_tool_call(req, handler)
assert isinstance(result, ToolMessage)
assert result.status == "error"
assert result.tool_call_id == "tc_999"
assert result.name == "flaky_tool"
assert "MCP server crashed" in result.content
def test_exception_does_not_propagate(self):
handler = MagicMock(side_effect=ConnectionError("refused"))
req = _make_request()
# Should NOT raise
result = self.mw.wrap_tool_call(req, handler)
assert isinstance(result, ToolMessage)
def test_keyboard_interrupt_propagates(self):
"""KeyboardInterrupt should NOT be caught (it's BaseException, not Exception)."""
handler = MagicMock(side_effect=KeyboardInterrupt())
req = _make_request()
with pytest.raises(KeyboardInterrupt):
self.mw.wrap_tool_call(req, handler)
def test_various_exception_types(self):
"""Different exception types are all caught and reported."""
for exc_cls in (ValueError, TypeError, OSError, TimeoutError, ConnectionError):
handler = MagicMock(side_effect=exc_cls(f"{exc_cls.__name__} happened"))
req = _make_request()
result = self.mw.wrap_tool_call(req, handler)
assert isinstance(result, ToolMessage)
assert result.status == "error"
assert exc_cls.__name__ in result.content
# ---------------------------------------------------------------------------
# Async: awrap_tool_call
# ---------------------------------------------------------------------------
class TestWrapToolCallAsync:
def setup_method(self):
self.mw = ToolErrorHandlerMiddleware()
async def test_success_passes_through(self):
expected = ToolMessage(content="ok", tool_call_id="tc_001", name="t")
async def handler(req):
return expected
req = _make_request()
result = await self.mw.awrap_tool_call(req, handler)
assert result is expected
async def test_command_passes_through(self):
cmd = Command(update={"messages": []})
async def handler(req):
return cmd
req = _make_request()
result = await self.mw.awrap_tool_call(req, handler)
assert result is cmd
async def test_exception_returns_error_tool_message(self):
async def handler(req):
raise RuntimeError("MCP server timed out")
req = _make_request("slow_tool", "tc_async")
result = await self.mw.awrap_tool_call(req, handler)
assert isinstance(result, ToolMessage)
assert result.status == "error"
assert result.tool_call_id == "tc_async"
assert result.name == "slow_tool"
assert "MCP server timed out" in result.content
async def test_exception_does_not_propagate(self):
async def handler(req):
raise ConnectionError("connection lost")
req = _make_request()
result = await self.mw.awrap_tool_call(req, handler)
assert isinstance(result, ToolMessage)
async def test_keyboard_interrupt_propagates(self):
async def handler(req):
raise KeyboardInterrupt()
req = _make_request()
with pytest.raises(KeyboardInterrupt):
await self.mw.awrap_tool_call(req, handler)
# ---------------------------------------------------------------------------
# Middleware metadata
# ---------------------------------------------------------------------------
class TestMiddlewareMeta:
def test_name(self):
assert ToolErrorHandlerMiddleware.name == "tool_error_handler"
def test_instantiation(self):
mw = ToolErrorHandlerMiddleware()
assert mw.name == "tool_error_handler"