From f802a495351182cbcbd68e87134ce916be866efa Mon Sep 17 00:00:00 2001 From: Sanjay Santhanam <51058514+Sanjays2402@users.noreply.github.com> Date: Wed, 22 Jul 2026 07:43:58 -0700 Subject: [PATCH] fix: repair interrupted tool call history (#366) * fix: repair interrupted tool call history Normalize incomplete tool exchanges before model calls so strict providers do not reject resumed sessions. Preserve completed exchanges and cover sync and async model paths. * fix: repair malformed tool calls and dedupe repair warnings Track AIMessage.invalid_tool_calls alongside tool_calls so interrupted threads with syntactically invalid tool calls get synthesized error results and are accepted by strict providers. Preserve the originating tool call's name in the synthesized ToolMessage, and deduplicate repair warnings per unique tool-call id via a warned set owned by the middleware instance, since the middleware rewrites the request but not thread state. Document the middleware's scope versus deepagents' PatchToolCallsMiddleware (orphan ToolMessage dropping and mid-run coverage). --------- Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com> --- EvoScientist/EvoScientist.py | 2 + EvoScientist/middleware/__init__.py | 2 + .../middleware/tool_history_repair.py | 135 ++++++++++++++ tests/test_tool_history_repair_middleware.py | 173 ++++++++++++++++++ 4 files changed, 312 insertions(+) create mode 100644 EvoScientist/middleware/tool_history_repair.py create mode 100644 tests/test_tool_history_repair_middleware.py 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