mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-03 13:53:27 +00:00
fix: repair orphaned tool calls before model calls (#1604)
When a run is cancelled or the sandbox dies mid-tool-call, LangGraph persists the AIMessage tool_call but never the matching ToolMessage. The next run sends the provider an orphaned tool_use (Anthropic 400: "tool_use ids were found without tool_result blocks"), permanently wedging the thread on every retry. Add RepairOrphanedToolCallsMiddleware, which inserts a synthetic error ToolMessage immediately after any tool_call lacking a result so the agent can retry instead of dying. Wired into the agent and reviewer graphs. Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
parent
534f402f48
commit
39681102d6
5 changed files with 234 additions and 0 deletions
|
|
@ -5,6 +5,7 @@ from .notify_step_limit import notify_step_limit_reached
|
||||||
from .plan_mode import PlanModeMiddleware
|
from .plan_mode import PlanModeMiddleware
|
||||||
from .refresh_github_proxy import refresh_github_proxy_before_model
|
from .refresh_github_proxy import refresh_github_proxy_before_model
|
||||||
from .refresh_slack_status import SlackAssistantStatusMiddleware
|
from .refresh_slack_status import SlackAssistantStatusMiddleware
|
||||||
|
from .repair_orphaned_tool_calls import RepairOrphanedToolCallsMiddleware
|
||||||
from .sandbox_circuit_breaker import SandboxCircuitBreakerMiddleware
|
from .sandbox_circuit_breaker import SandboxCircuitBreakerMiddleware
|
||||||
from .sanitize_thinking_blocks import SanitizeThinkingBlocksMiddleware
|
from .sanitize_thinking_blocks import SanitizeThinkingBlocksMiddleware
|
||||||
from .sanitize_tool_inputs import SanitizeToolInputsMiddleware
|
from .sanitize_tool_inputs import SanitizeToolInputsMiddleware
|
||||||
|
|
@ -16,6 +17,7 @@ __all__ = [
|
||||||
"ExcludeToolsMiddleware",
|
"ExcludeToolsMiddleware",
|
||||||
"ModelFallbackMiddleware",
|
"ModelFallbackMiddleware",
|
||||||
"PlanModeMiddleware",
|
"PlanModeMiddleware",
|
||||||
|
"RepairOrphanedToolCallsMiddleware",
|
||||||
"SanitizeThinkingBlocksMiddleware",
|
"SanitizeThinkingBlocksMiddleware",
|
||||||
"SanitizeToolInputsMiddleware",
|
"SanitizeToolInputsMiddleware",
|
||||||
"ToolArtifactMiddleware",
|
"ToolArtifactMiddleware",
|
||||||
|
|
|
||||||
119
agent/middleware/repair_orphaned_tool_calls.py
Normal file
119
agent/middleware/repair_orphaned_tool_calls.py
Normal file
|
|
@ -0,0 +1,119 @@
|
||||||
|
"""Middleware that repairs orphaned tool calls before model calls.
|
||||||
|
|
||||||
|
When a run is cancelled or the sandbox dies mid-tool-call, LangGraph persists the
|
||||||
|
``AIMessage`` with the ``tool_call`` but never the matching ``ToolMessage``. On the
|
||||||
|
next run a fresh human message lands where the tool result should be, so the
|
||||||
|
provider rejects the request (Anthropic: ``messages.N: `tool_use` ids were found
|
||||||
|
without `tool_result` blocks``; OpenAI raises the equivalent). That permanently
|
||||||
|
wedges the thread — every retry hits the same error.
|
||||||
|
|
||||||
|
This middleware scans the outgoing message list and inserts a synthetic error
|
||||||
|
``ToolMessage`` immediately after any ``tool_call`` whose id has no corresponding
|
||||||
|
``ToolMessage``. The agent then sees the interrupted tool as a normal tool error
|
||||||
|
and can retry instead of dying.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain.agents.middleware import AgentMiddleware
|
||||||
|
from langchain.agents.middleware.types import ModelCallResult, ModelRequest, ModelResponse
|
||||||
|
from langchain_core.messages import AIMessage, ToolMessage
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
INTERRUPTED_TOOL_RECOVERY = "tool_call_interrupted"
|
||||||
|
|
||||||
|
_INTERRUPTED_TOOL_ERROR = (
|
||||||
|
"The previous tool call did not complete — the run was interrupted (cancelled "
|
||||||
|
"or the sandbox became unavailable) before a result was returned. No output was "
|
||||||
|
"captured. Retry the tool call if you still need it; if repository files are "
|
||||||
|
"missing, re-clone or reinitialize the workspace first."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _iter_tool_calls(message: AIMessage) -> list[tuple[str, str | None]]:
|
||||||
|
"""Return ``(id, name)`` for each well-formed tool call on the message."""
|
||||||
|
calls: list[tuple[str, str | None]] = []
|
||||||
|
for tool_call in message.tool_calls or []:
|
||||||
|
if isinstance(tool_call, dict):
|
||||||
|
call_id = tool_call.get("id")
|
||||||
|
name = tool_call.get("name")
|
||||||
|
else:
|
||||||
|
call_id = getattr(tool_call, "id", None)
|
||||||
|
name = getattr(tool_call, "name", None)
|
||||||
|
if isinstance(call_id, str) and call_id:
|
||||||
|
calls.append((call_id, name if isinstance(name, str) and name else None))
|
||||||
|
return calls
|
||||||
|
|
||||||
|
|
||||||
|
def _synthetic_tool_message(call_id: str, name: str | None) -> ToolMessage:
|
||||||
|
payload: dict[str, str] = {
|
||||||
|
"status": "error",
|
||||||
|
"error_type": "InterruptedToolCall",
|
||||||
|
"recovery": INTERRUPTED_TOOL_RECOVERY,
|
||||||
|
"error": _INTERRUPTED_TOOL_ERROR,
|
||||||
|
}
|
||||||
|
if name:
|
||||||
|
payload["name"] = name
|
||||||
|
return ToolMessage(
|
||||||
|
content=json.dumps(payload),
|
||||||
|
tool_call_id=call_id,
|
||||||
|
name=name,
|
||||||
|
status="error",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _repair_messages(messages: list[Any]) -> list[Any] | None:
|
||||||
|
"""Insert synthetic results for orphaned tool calls; return new list or None."""
|
||||||
|
satisfied = {
|
||||||
|
message.tool_call_id
|
||||||
|
for message in messages
|
||||||
|
if isinstance(message, ToolMessage) and isinstance(message.tool_call_id, str)
|
||||||
|
}
|
||||||
|
|
||||||
|
repaired: list[Any] = []
|
||||||
|
inserted = 0
|
||||||
|
for message in messages:
|
||||||
|
repaired.append(message)
|
||||||
|
if not isinstance(message, AIMessage):
|
||||||
|
continue
|
||||||
|
for call_id, name in _iter_tool_calls(message):
|
||||||
|
if call_id in satisfied:
|
||||||
|
continue
|
||||||
|
repaired.append(_synthetic_tool_message(call_id, name))
|
||||||
|
satisfied.add(call_id)
|
||||||
|
inserted += 1
|
||||||
|
|
||||||
|
if not inserted:
|
||||||
|
return None
|
||||||
|
logger.warning("Repaired %d orphaned tool call(s) before model call", inserted)
|
||||||
|
return repaired
|
||||||
|
|
||||||
|
|
||||||
|
class RepairOrphanedToolCallsMiddleware(AgentMiddleware):
|
||||||
|
"""Insert synthetic tool results for interrupted tool calls before model calls."""
|
||||||
|
|
||||||
|
def wrap_model_call(
|
||||||
|
self,
|
||||||
|
request: ModelRequest,
|
||||||
|
handler: Callable[[ModelRequest], ModelResponse],
|
||||||
|
) -> ModelCallResult:
|
||||||
|
repaired = _repair_messages(request.messages)
|
||||||
|
if repaired is not None:
|
||||||
|
request.messages[:] = repaired
|
||||||
|
return handler(request)
|
||||||
|
|
||||||
|
async def awrap_model_call(
|
||||||
|
self,
|
||||||
|
request: ModelRequest,
|
||||||
|
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||||
|
) -> Any:
|
||||||
|
repaired = _repair_messages(request.messages)
|
||||||
|
if repaired is not None:
|
||||||
|
request.messages[:] = repaired
|
||||||
|
return await handler(request)
|
||||||
|
|
@ -39,6 +39,7 @@ from .dashboard.team_settings import (
|
||||||
get_team_default_model_pair,
|
get_team_default_model_pair,
|
||||||
)
|
)
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
|
RepairOrphanedToolCallsMiddleware,
|
||||||
SanitizeThinkingBlocksMiddleware,
|
SanitizeThinkingBlocksMiddleware,
|
||||||
SanitizeToolInputsMiddleware,
|
SanitizeToolInputsMiddleware,
|
||||||
SlackAssistantStatusMiddleware,
|
SlackAssistantStatusMiddleware,
|
||||||
|
|
@ -1166,6 +1167,7 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel:
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
SlackAssistantStatusMiddleware(),
|
SlackAssistantStatusMiddleware(),
|
||||||
SanitizeThinkingBlocksMiddleware(),
|
SanitizeThinkingBlocksMiddleware(),
|
||||||
|
RepairOrphanedToolCallsMiddleware(),
|
||||||
settle_review_check_on_exit,
|
settle_review_check_on_exit,
|
||||||
],
|
],
|
||||||
).with_config(config)
|
).with_config(config)
|
||||||
|
|
|
||||||
|
|
@ -56,6 +56,7 @@ from .integrations.notion_mcp import load_notion_tools
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
ModelFallbackMiddleware,
|
ModelFallbackMiddleware,
|
||||||
PlanModeMiddleware,
|
PlanModeMiddleware,
|
||||||
|
RepairOrphanedToolCallsMiddleware,
|
||||||
SandboxCircuitBreakerMiddleware,
|
SandboxCircuitBreakerMiddleware,
|
||||||
SanitizeThinkingBlocksMiddleware,
|
SanitizeThinkingBlocksMiddleware,
|
||||||
SanitizeToolInputsMiddleware,
|
SanitizeToolInputsMiddleware,
|
||||||
|
|
@ -840,6 +841,7 @@ async def get_agent(config: RunnableConfig) -> Pregel:
|
||||||
*fallback_middleware,
|
*fallback_middleware,
|
||||||
*plan_mode_middleware,
|
*plan_mode_middleware,
|
||||||
SanitizeThinkingBlocksMiddleware(),
|
SanitizeThinkingBlocksMiddleware(),
|
||||||
|
RepairOrphanedToolCallsMiddleware(),
|
||||||
],
|
],
|
||||||
).with_config(config)
|
).with_config(config)
|
||||||
|
|
||||||
|
|
|
||||||
109
tests/test_repair_orphaned_tool_calls.py
Normal file
109
tests/test_repair_orphaned_tool_calls.py
Normal file
|
|
@ -0,0 +1,109 @@
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||||
|
|
||||||
|
from agent.middleware.repair_orphaned_tool_calls import (
|
||||||
|
INTERRUPTED_TOOL_RECOVERY,
|
||||||
|
RepairOrphanedToolCallsMiddleware,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_request(messages: list[object]) -> MagicMock:
|
||||||
|
request = MagicMock()
|
||||||
|
request.model = MagicMock()
|
||||||
|
request.messages = messages
|
||||||
|
return request
|
||||||
|
|
||||||
|
|
||||||
|
def _ai_with_tool_call(call_id: str, name: str = "execute") -> AIMessage:
|
||||||
|
return AIMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[{"name": name, "args": {"command": "ls"}, "id": call_id, "type": "tool_call"}],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRepairOrphanedToolCallsMiddleware:
|
||||||
|
def test_inserts_synthetic_result_for_orphaned_tool_call(self) -> None:
|
||||||
|
ai = _ai_with_tool_call("call_1")
|
||||||
|
follow_up = HumanMessage(content="continue")
|
||||||
|
request = _make_request([HumanMessage(content="hi"), ai, follow_up])
|
||||||
|
response = MagicMock()
|
||||||
|
|
||||||
|
result = RepairOrphanedToolCallsMiddleware().wrap_model_call(request, lambda req: response)
|
||||||
|
|
||||||
|
assert result is response
|
||||||
|
messages = request.messages
|
||||||
|
assert len(messages) == 4
|
||||||
|
synthetic = messages[2]
|
||||||
|
assert isinstance(synthetic, ToolMessage)
|
||||||
|
assert synthetic.tool_call_id == "call_1"
|
||||||
|
assert synthetic.status == "error"
|
||||||
|
assert messages[3] is follow_up
|
||||||
|
payload = json.loads(synthetic.content)
|
||||||
|
assert payload["recovery"] == INTERRUPTED_TOOL_RECOVERY
|
||||||
|
assert payload["name"] == "execute"
|
||||||
|
|
||||||
|
def test_leaves_satisfied_tool_calls_untouched(self) -> None:
|
||||||
|
ai = _ai_with_tool_call("call_1")
|
||||||
|
tool = ToolMessage(content="done", tool_call_id="call_1")
|
||||||
|
original = [HumanMessage(content="hi"), ai, tool]
|
||||||
|
request = _make_request(list(original))
|
||||||
|
|
||||||
|
RepairOrphanedToolCallsMiddleware().wrap_model_call(request, lambda req: MagicMock())
|
||||||
|
|
||||||
|
assert request.messages == original
|
||||||
|
|
||||||
|
def test_repairs_multiple_orphans_on_one_message(self) -> None:
|
||||||
|
ai = AIMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
{"name": "execute", "args": {}, "id": "call_1", "type": "tool_call"},
|
||||||
|
{"name": "grep", "args": {}, "id": "call_2", "type": "tool_call"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
request = _make_request([ai])
|
||||||
|
|
||||||
|
RepairOrphanedToolCallsMiddleware().wrap_model_call(request, lambda req: MagicMock())
|
||||||
|
|
||||||
|
messages = request.messages
|
||||||
|
assert [getattr(m, "tool_call_id", None) for m in messages[1:]] == ["call_1", "call_2"]
|
||||||
|
assert all(isinstance(m, ToolMessage) for m in messages[1:])
|
||||||
|
|
||||||
|
def test_partial_repair_keeps_existing_result(self) -> None:
|
||||||
|
ai = AIMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
{"name": "execute", "args": {}, "id": "call_1", "type": "tool_call"},
|
||||||
|
{"name": "grep", "args": {}, "id": "call_2", "type": "tool_call"},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
tool = ToolMessage(content="done", tool_call_id="call_1")
|
||||||
|
request = _make_request([ai, tool])
|
||||||
|
|
||||||
|
RepairOrphanedToolCallsMiddleware().wrap_model_call(request, lambda req: MagicMock())
|
||||||
|
|
||||||
|
synthetic = [
|
||||||
|
m for m in request.messages if isinstance(m, ToolMessage) and m.status == "error"
|
||||||
|
]
|
||||||
|
assert len(synthetic) == 1
|
||||||
|
assert synthetic[0].tool_call_id == "call_2"
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_async_inserts_synthetic_result(self) -> None:
|
||||||
|
ai = _ai_with_tool_call("call_1")
|
||||||
|
request = _make_request([ai, HumanMessage(content="hi")])
|
||||||
|
response = MagicMock()
|
||||||
|
|
||||||
|
async def handler(req: object) -> object:
|
||||||
|
assert req is request
|
||||||
|
return response
|
||||||
|
|
||||||
|
result = await RepairOrphanedToolCallsMiddleware().awrap_model_call(request, handler)
|
||||||
|
|
||||||
|
assert result is response
|
||||||
|
assert isinstance(request.messages[1], ToolMessage)
|
||||||
|
assert request.messages[1].tool_call_id == "call_1"
|
||||||
Loading…
Add table
Reference in a new issue