diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 5ca381c..f278712 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -687,6 +687,7 @@ def _get_default_middleware( ErrorNormalizationMiddleware, ModelFallbackMiddleware, ToolErrorHandlerMiddleware, + ToolHistoryRepairMiddleware, create_code_interpreter_middleware, create_context_editing_middleware, create_memory_lifecycle_middleware, @@ -759,6 +760,7 @@ def _get_default_middleware( # middlewares) and normalizes them into a non-dataclass # envelope wrapper before anything downstream sees them. ErrorNormalizationMiddleware(), + ToolHistoryRepairMiddleware(), ConfigurableModelMiddleware(), create_context_editing_middleware(model), ModelFallbackMiddleware(events=events), diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index bb8afcb..e26f0fc 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -35,6 +35,7 @@ from .scheduler import ( create_scheduler_middleware, ) from .tool_error_handler import ToolErrorHandlerMiddleware +from .tool_history_repair import ToolHistoryRepairMiddleware from .tool_selector import create_tool_selector_middleware from .utils import disable_thinking @@ -53,6 +54,7 @@ __all__ = [ "RuntimeContextMiddleware", "SchedulerMiddleware", "ToolErrorHandlerMiddleware", + "ToolHistoryRepairMiddleware", "compute_context_editing_trigger", "create_code_interpreter_middleware", "create_context_editing_middleware", diff --git a/EvoScientist/middleware/tool_history_repair.py b/EvoScientist/middleware/tool_history_repair.py new file mode 100644 index 0000000..c60f87e --- /dev/null +++ b/EvoScientist/middleware/tool_history_repair.py @@ -0,0 +1,135 @@ +"""Repair interrupted tool-call history before provider requests. + +Strict providers (OpenAI, etc.) reject a message thread in which an assistant +tool call has no matching tool result. That happens whenever a run is +interrupted (cancelled, crashed, timed out) after the model emitted tool calls +but before those tools produced results. This middleware rewrites the outgoing +request so every dangling tool call is closed with a synthetic error result and +every orphan ``ToolMessage`` (a result whose originating call is gone) is +dropped. + +It covers cases that deepagents' ``PatchToolCallsMiddleware`` does not: + +1. Orphan ``ToolMessage`` dropping -- a tool result whose originating tool call + is no longer present in history is removed, rather than left to trip strict + providers. +2. Mid-run coverage -- repair runs at the model boundary on every request + (including malformed / ``invalid_tool_calls``), not only at agent start, so + interruptions that happen partway through a run are healed too. + +Because the middleware only rewrites the request and cannot mutate thread +state, the repaired synthetic results are recomputed on every model call. To +avoid re-logging the same repair forever, warnings are deduplicated per unique +tool-call id via a ``warned`` set owned by the middleware instance. +""" + +from __future__ import annotations + +import logging +from collections.abc import Awaitable, Callable, Sequence + +from langchain.agents.middleware.types import ( + AgentMiddleware, + ModelRequest, + ModelResponse, +) +from langchain_core.messages import AIMessage, AnyMessage, ToolMessage + +logger = logging.getLogger(__name__) +_INTERRUPTED_RESULT = "Tool execution was interrupted before completion." + + +def repair_tool_history( + messages: Sequence[AnyMessage], + warned: set[str] | None = None, +) -> list[AnyMessage]: + """Return provider-valid history, preserving every complete tool exchange. + + When ``warned`` is provided, repair warnings are emitted only for tool-call + ids not already present in it; newly-warned ids are added. This keeps the + warning to once per unique interrupted/malformed call even though the + middleware re-runs on every model call. + """ + repaired: list[AnyMessage] = [] + pending: dict[str, str | None] = {} + synthesized: list[str] = [] + dropped: list[str] = [] + + def close_pending() -> None: + for tool_call_id, tool_name in pending.items(): + repaired.append( + ToolMessage( + content=_INTERRUPTED_RESULT, + tool_call_id=tool_call_id, + name=tool_name, + status="error", + ) + ) + synthesized.append(tool_call_id) + pending.clear() + + for message in messages: + if isinstance(message, ToolMessage): + tool_call_id = message.tool_call_id + if tool_call_id in pending: + repaired.append(message) + pending.pop(tool_call_id) + else: + dropped.append(tool_call_id) + continue + + if pending: + close_pending() + repaired.append(message) + if isinstance(message, AIMessage): + all_calls = list(message.tool_calls) + list( + getattr(message, "invalid_tool_calls", []) or [] + ) + for call in all_calls: + if tool_call_id := call.get("id"): + pending[tool_call_id] = call.get("name") + + if pending: + close_pending() + + if warned is not None: + synthesized = [tid for tid in synthesized if tid not in warned] + dropped = [tid for tid in dropped if tid not in warned] + warned.update(synthesized) + warned.update(dropped) + + if synthesized or dropped: + logger.warning( + "Repaired interrupted tool history: synthesized=%s dropped=%s", + synthesized, + dropped, + ) + return repaired + + +class ToolHistoryRepairMiddleware(AgentMiddleware): + """Repair dangling calls and orphan results at the model boundary.""" + + name = "tool_history_repair" + + def __init__(self) -> None: + super().__init__() + self._warned: set[str] = set() + + def modify_request(self, request: ModelRequest) -> ModelRequest: + messages = repair_tool_history(request.messages, warned=self._warned) + return request.override(messages=messages) + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + return handler(self.modify_request(request)) + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + return await handler(self.modify_request(request)) diff --git a/tests/test_tool_history_repair_middleware.py b/tests/test_tool_history_repair_middleware.py new file mode 100644 index 0000000..b4d976b --- /dev/null +++ b/tests/test_tool_history_repair_middleware.py @@ -0,0 +1,173 @@ +from unittest.mock import AsyncMock, MagicMock + +from langchain.agents.middleware.types import ModelRequest +from langchain_core.messages import AIMessage, HumanMessage, ToolMessage + +from EvoScientist.middleware.tool_history_repair import ( + ToolHistoryRepairMiddleware, + repair_tool_history, +) + + +def _request(messages): + return ModelRequest( + messages=messages, + model=MagicMock(), + state={}, + runtime=MagicMock(), + system_message=MagicMock(), + ) + + +def _tool_call(tool_call_id): + return {"id": tool_call_id, "name": "execute", "args": {}} + + +def _invalid_tool_call(tool_call_id): + return { + "id": tool_call_id, + "name": "execute", + "args": "{not valid json", + "error": "could not parse args", + } + + +def test_synthesizes_results_for_interrupted_tool_calls(): + messages = [ + HumanMessage("run tools"), + AIMessage(content="", tool_calls=[_tool_call("one"), _tool_call("two")]), + HumanMessage("continue"), + ] + + repaired = repair_tool_history(messages) + + assert [type(message) for message in repaired] == [ + HumanMessage, + AIMessage, + ToolMessage, + ToolMessage, + HumanMessage, + ] + assert [message.tool_call_id for message in repaired[2:4]] == ["one", "two"] + assert all(message.status == "error" for message in repaired[2:4]) + + +def test_drops_orphan_tool_results(): + messages = [ + HumanMessage("old request"), + ToolMessage("late result", tool_call_id="orphan"), + HumanMessage("continue"), + ] + + assert repair_tool_history(messages) == [messages[0], messages[2]] + + +def test_preserves_complete_tool_exchanges(): + messages = [ + HumanMessage("run tool"), + AIMessage(content="", tool_calls=[_tool_call("complete")]), + ToolMessage("done", tool_call_id="complete"), + HumanMessage("continue"), + ] + + assert repair_tool_history(messages) == messages + + +def test_wrap_model_call_repairs_request(): + request = _request( + [ + ToolMessage("late result", tool_call_id="orphan"), + HumanMessage("continue"), + ] + ) + handler = MagicMock(return_value="ok") + + assert ToolHistoryRepairMiddleware().wrap_model_call(request, handler) == "ok" + assert handler.call_args.args[0].messages == [request.messages[1]] + + +async def test_awrap_model_call_repairs_request(): + request = _request( + [ + AIMessage(content="", tool_calls=[_tool_call("interrupted")]), + HumanMessage("continue"), + ] + ) + handler = AsyncMock(return_value="ok") + + assert ( + await ToolHistoryRepairMiddleware().awrap_model_call(request, handler) == "ok" + ) + repaired = handler.call_args.args[0].messages + assert isinstance(repaired[1], ToolMessage) + assert repaired[1].tool_call_id == "interrupted" + + +def test_synthesizes_results_for_invalid_tool_calls(): + messages = [ + HumanMessage("run tools"), + AIMessage( + content="", + tool_calls=[_tool_call("good")], + invalid_tool_calls=[_invalid_tool_call("bad")], + ), + HumanMessage("continue"), + ] + + repaired = repair_tool_history(messages) + + assert [type(message) for message in repaired] == [ + HumanMessage, + AIMessage, + ToolMessage, + ToolMessage, + HumanMessage, + ] + assert [message.tool_call_id for message in repaired[2:4]] == ["good", "bad"] + assert all(message.status == "error" for message in repaired[2:4]) + + +def test_preserves_tool_call_name_in_synthesized_result(): + messages = [ + HumanMessage("run tool"), + AIMessage(content="", tool_calls=[_tool_call("one")]), + ] + + repaired = repair_tool_history(messages) + + assert repaired[-1].name == "execute" + + +def test_warning_deduplicates_across_calls(caplog): + messages = [ + HumanMessage("run tools"), + AIMessage(content="", tool_calls=[_tool_call("one")]), + ] + warned: set[str] = set() + + with caplog.at_level("WARNING"): + repair_tool_history(messages, warned=warned) + first_warnings = len(caplog.records) + repair_tool_history(messages, warned=warned) + second_warnings = len(caplog.records) + + assert first_warnings == 1 + assert second_warnings == 1 + assert warned == {"one"} + + +def test_middleware_warns_once_per_thread(caplog): + middleware = ToolHistoryRepairMiddleware() + request = _request( + [ + AIMessage(content="", tool_calls=[_tool_call("interrupted")]), + HumanMessage("continue"), + ] + ) + handler = MagicMock(return_value="ok") + + with caplog.at_level("WARNING"): + middleware.wrap_model_call(request, handler) + middleware.wrap_model_call(request, handler) + + assert len(caplog.records) == 1