fix(middleware): mount tool-history repair on sync subagents and harden raw tool-call vetting (#393)

* feat(middleware): add ToolHistoryRepairMiddleware and enhance tool call validation

* fix(tests): add test for dropping non-list raw tool calls in repair_tool_history
This commit is contained in:
Xi Zhang
2026-07-28 18:14:37 +01:00
committed by GitHub
parent a6a8e19dcc
commit 562ce0eb83
3 changed files with 175 additions and 8 deletions
+3
View File
@@ -318,6 +318,7 @@ def _inject_subagent_middleware(
ContextOverflowMapperMiddleware,
ErrorNormalizationMiddleware,
ToolErrorHandlerMiddleware,
ToolHistoryRepairMiddleware,
create_context_editing_middleware,
create_memory_lifecycle_middleware,
create_memory_middleware,
@@ -350,6 +351,8 @@ def _inject_subagent_middleware(
# them into a non-dataclass envelope wrapper before
# anything downstream sees them.
ErrorNormalizationMiddleware(),
# Sync subagents replay their own history to strict providers too.
ToolHistoryRepairMiddleware(),
# Subagents share the main agent's model: use the threaded
# ``chat_model`` on the pure path, else defer to the factory's
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
+25 -7
View File
@@ -91,15 +91,34 @@ def repair_tool_history(
additional_kwargs = message.additional_kwargs
raw_calls = additional_kwargs.get("tool_calls")
valid_raw_calls = []
if isinstance(raw_calls, list):
for call in raw_calls:
function = call.get("function") if isinstance(call, dict) else None
if isinstance(function, dict) and function.get("name"):
valid_raw_calls.append(call)
if "tool_calls" in additional_kwargs:
if isinstance(raw_calls, list):
for call in raw_calls:
function = (
call.get("function") if isinstance(call, dict) else None
)
name = (
function.get("name") if isinstance(function, dict) else None
)
tool_call_id = (
call.get("id") if isinstance(call, dict) else None
)
# Entries without a usable str name AND id can never be
# closed by a result, so keeping them would leave
# provider-invalid history in the payload.
if (
isinstance(name, str)
and name
and isinstance(tool_call_id, str)
and tool_call_id
):
valid_raw_calls.append(call)
additional_kwargs = dict(additional_kwargs)
if valid_raw_calls:
additional_kwargs["tool_calls"] = valid_raw_calls
else:
# Also covers non-list junk (dict/str/int), which would
# otherwise crash langchain's serializer downstream.
additional_kwargs.pop("tool_calls", None)
message = message.model_copy(
@@ -114,8 +133,7 @@ def repair_tool_history(
if tool_call_id := call.get("id"):
pending[tool_call_id] = call.get("name")
for call in valid_raw_calls:
if tool_call_id := call.get("id"):
pending[tool_call_id] = call["function"]["name"]
pending[call["id"]] = call["function"]["name"]
repaired.append(message)
if pending:
+147 -1
View File
@@ -1,4 +1,4 @@
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch
from langchain.agents.middleware.types import ModelRequest
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
@@ -108,6 +108,152 @@ def test_removes_unnamed_calls_before_serialization():
assert [message.tool_call_id for message in repaired[1:]] == ["raw-good"]
def test_malformed_raw_entries_are_dropped_without_crashing():
message = AIMessage(content="").model_copy(
update={
"additional_kwargs": {
"tool_calls": [
{"id": ["a"], "function": {"name": "x", "arguments": "{}"}},
{"id": "c1", "function": {"name": ["evil"], "arguments": "{}"}},
{"id": "c2", "function": {"name": 7, "arguments": "{}"}},
]
},
}
)
repaired = repair_tool_history([message])
assert "tool_calls" not in repaired[0].additional_kwargs
assert not any(isinstance(m, ToolMessage) for m in repaired)
def test_non_str_raw_id_entry_is_dropped_with_its_result():
message = AIMessage(content="").model_copy(
update={
"additional_kwargs": {
"tool_calls": [
{"id": 123, "function": {"name": "f", "arguments": "{}"}}
]
},
}
)
repaired = repair_tool_history([message, ToolMessage("real", tool_call_id="123")])
assert "tool_calls" not in repaired[0].additional_kwargs
assert not any(isinstance(m, ToolMessage) for m in repaired)
def test_non_list_raw_tool_calls_value_is_dropped():
for junk in ({"id": "bad", "function": {"name": "x"}}, "bad", 1):
message = AIMessage(content="").model_copy(
update={"additional_kwargs": {"extra": "kept", "tool_calls": junk}}
)
repaired = repair_tool_history([message])
assert "tool_calls" not in repaired[0].additional_kwargs
assert repaired[0].additional_kwargs["extra"] == "kept"
def test_removes_raw_tool_calls_key_when_all_entries_invalid():
message = AIMessage(content="").model_copy(
update={
"additional_kwargs": {
"extra": "kept",
"tool_calls": [{"id": "x", "function": {"arguments": "{}"}}],
},
}
)
repaired = repair_tool_history([message, ToolMessage("x", tool_call_id="x")])
assert "tool_calls" not in repaired[0].additional_kwargs
assert repaired[0].additional_kwargs["extra"] == "kept"
assert len(repaired) == 1
def test_mixed_named_and_unnamed_parsed_calls():
message = AIMessage(content="").model_copy(
update={
"tool_calls": [
{"id": "good", "name": "execute", "args": {}},
{"id": "bad", "name": "", "args": {}},
],
}
)
messages = [
message,
ToolMessage("ok", tool_call_id="good"),
ToolMessage("junk", tool_call_id="bad"),
]
repaired = repair_tool_history(messages)
assert [call["id"] for call in repaired[0].tool_calls] == ["good"]
assert [m.tool_call_id for m in repaired[1:]] == ["good"]
def test_synthesizes_result_for_unanswered_raw_call():
message = AIMessage(content="").model_copy(
update={
"additional_kwargs": {
"tool_calls": [
{"id": "raw-1", "function": {"name": "grep", "arguments": "{}"}}
]
},
}
)
repaired = repair_tool_history([message])
assert repaired[-1].tool_call_id == "raw-1"
assert repaired[-1].name == "grep"
assert repaired[-1].status == "error"
def test_repair_is_idempotent():
messages = [
AIMessage(content="").model_copy(
update={
"tool_calls": [
_tool_call("kept"),
{"id": "bad", "name": "", "args": {}},
],
"additional_kwargs": {
"tool_calls": [
{
"id": "raw-1",
"function": {"name": "grep", "arguments": "{}"},
},
{"id": "raw-2", "function": {"arguments": "{}"}},
]
},
}
),
ToolMessage("done", tool_call_id="kept"),
ToolMessage("junk", tool_call_id="bad"),
]
once = repair_tool_history(messages)
assert repair_tool_history(once) == once
@patch("EvoScientist.EvoScientist._ensure_chat_model")
def test_inject_subagent_includes_tool_history_repair(mock_model):
mock_model.return_value = MagicMock(profile={"max_input_tokens": 200_000})
from EvoScientist.EvoScientist import _inject_subagent_middleware
subs = [{"name": "test-agent"}]
_inject_subagent_middleware(subs)
assert any(
isinstance(m, ToolHistoryRepairMiddleware) for m in subs[0]["middleware"]
)
def test_wrap_model_call_repairs_request():
request = _request(
[