Nan Gao cf556fa9d4
feat(agents): elide superseded write_file payloads from model-bound requests (#5374)
* feat(agents): elide superseded write_file payloads from model-bound requests

Step 2 of #5328. After a successful write_file the file on disk is the source
of truth, and the read-before-write gate forces a read_file before the next
modification of that path, so once a later successful read or write of the
same path exists the historical `content` argument is redundant with it. Long
report-writing runs (append-in-chunks) therefore carried every section twice,
once as the write argument and once as the following read output, until
summarization compacted the whole turn.

- ToolOutputBudgetMiddleware's model-call hooks now replace such superseded
  content with a short deterministic placeholder pointing at read_file, in the
  model-bound request only: state["messages"], checkpoints, receipts, loop
  detection, and the run journal keep the original arguments, and nothing is
  externalized to disk. The newest `keep_recent_writes` successful writes
  (default 1) always stay visible; str_replace payloads are never touched; a
  same-turn read never supersedes (parallel calls run in no fixed order); only
  results stamped deerflow_tool_meta.status == "success" count, so failed,
  gate-blocked, partial, or unstamped writes are never candidates.
- New `tool_output.elide_superseded_writes` (default on),
  `tool_output.superseded_write_min_chars` (default 2000), and
  `tool_output.keep_recent_writes` (default 1); config_version 41 -> 42 in
  config.example.yaml and the Helm chart.
- The per-occurrence call/result pairing the gate introduced in #5329 moves
  into the shared `tool_call_args.pair_tool_call_results` helper so both
  policies pair the same way; the gate now uses it.

* fix(agents): scope tool-call result pairing to the issuing turn

Review finding on #5374 (P2): pair_tool_call_results consumed results from a
history-wide per-id queue, so an interrupted write_file with no result whose
tool-call id a later turn reused inherited that later call's success. With
the default elision the unconfirmed draft was then replaced by a placeholder
claiming the write succeeded, and the gate's blocked-call pairing had the
mirror-image hole.

Pair results the way DanglingToolCallMiddleware does: walk in document order,
open each AIMessage's calls, and let a ToolMessage answer only a still-open
call of the most recent preceding AIMessage. A result never answers a call
from an earlier turn, so the interrupted call stays unanswered (never a
candidate, never labeled blocked) and stray or duplicate results are ignored.
Regressions cover the helper, the superseded-write policy, and the gate.

* fix(agents): never rewrite tool-call ids duplicated within one AIMessage

Review finding on #5374 (P2): the policies select calls per occurrence, but
every provider surface is addressed by tool-call id, so when a malformed
provider payload repeats an id inside one assistant turn the rewriter could
only replace all of its occurrences at once. A failed write_file sibling then
took on the superseded successful call's path and elided content and was
presented as a success; the gate's blocked-call elision had the mirror-image
hole (a successful sibling rewritten into the blocked call).

rewrite_messages_tool_call_args now never offers an id that repeats within
its message to the selector and leaves those calls untouched on every
surface. Both policies are covered by the shared helper; regressions cover
the helper, the superseded-write policy, and the gate.

* fix(agents): skip unhashable tool-call ids in the duplicate-id guard

Review finding on #5374 (round 3): _duplicated_call_ids fed every id into a
Counter before the string guard, so a list or dict id from a malformed
provider payload raised TypeError out of wrap_model_call and failed the whole
model call whenever the history also held a rewrite candidate. The pre-PR
loop and pair_tool_call_results skip such ids; only this helper regressed.

Count non-empty string ids only, and pin it with regressions for the helper,
the superseded-write policy, and the gate.

* docs(agents): keep middleware guidance within size limit

---------

Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
2026-09-13 18:04:51 +08:00

284 lines
14 KiB
Python

"""Rewrite AIMessage tool-call arguments on every provider surface at once.
Middlewares that shrink or replace a historical tool call's arguments in the
*model-bound request* (never in graph state) share one hazard: a LangChain
``AIMessage`` carries the same arguments on up to four surfaces, and provider
adapters do not all read the same one —
- ``tool_calls``: the structured list most adapters prefer;
- ``additional_kwargs["tool_calls"]``: the raw provider payload (OpenAI
``function.arguments`` JSON string) some adapters fall back to;
- ``content`` blocks that carry their own copy of the arguments: Anthropic
``tool_use`` (``input`` + ``partial_json``), OpenAI Responses
``function_call`` (``arguments`` string, matched by ``call_id``; the
``fc_…`` item id is preserved), and LangChain standard-content
``tool_call`` / ``tool_call_chunk`` (``args`` plus ``extras.arguments``);
- ``tool_call_chunks`` on an ``AIMessageChunk``.
Rewriting only one surface leaves the original payload reachable through the
others and can hand a strict provider a request whose surfaces disagree. The
content surfaces matter most: ``langchain_openai``'s Responses input builder
emits a content ``function_call`` block *instead of* the structured call
whose ``call_id`` it already carries, and prefers ``extras.arguments`` over
the structured args when translating a v1 ``tool_call`` block, so a rewrite
that touched ``tool_calls`` alone would still send the original payload.
:func:`rewrite_tool_call_args` rewrites them together and returns a
``model_copy`` (or the same object when nothing matched), so callers never
mutate state and the result is identical across model calls. Policy — which
calls, and what replaces their arguments — stays with the caller; see
``read_before_write_middleware.elide_blocked_write_payloads`` and
``tool_output_budget_middleware.elide_superseded_write_payloads``. Both decide
per call *occurrence*, pairing each AIMessage call with the ToolMessage that
answered it through :func:`pair_tool_call_results`, because tool-call ids may
repeat across assistant turns.
A rewrite also invalidates server-side continuation. With
``use_previous_response_id`` the OpenAI adapter sends only the messages after
the last AIMessage carrying a ``resp_…`` ``response_metadata["id"]`` and lets
the server rebuild the rest from *its* stored copy of the conversation, which
still holds the original arguments; stored responses cannot be edited, and
every response produced after the rewritten call chains back to that history.
So whenever anything was rewritten, :func:`rewrite_messages_tool_call_args`
drops every ``resp_`` id from the model-bound copy and the adapter falls back
to replaying the full rewritten history (the same request shape as
``use_previous_response_id=False``; per OpenAI's docs chained input tokens are
billed either way, so replay costs no more).
"""
from __future__ import annotations
import json
from collections import Counter, defaultdict, deque
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass, replace
from typing import Any
from langchain_core.messages import AIMessage, ToolMessage
#: Replacement args keyed by tool-call id.
ArgsReplacements = Mapping[str, dict[str, Any]]
#: ``(message, tool_call) -> new_args`` or ``None`` to leave the call alone.
ReplacementSelector = Callable[[AIMessage, dict[str, Any]], dict[str, Any] | None]
def rewrite_messages_tool_call_args(messages: list[Any], replacement_for: ReplacementSelector) -> list[Any] | None:
"""Apply ``replacement_for(message, tool_call)`` to every AIMessage tool call in ``messages``.
Returns a new list with the rewritten AIMessages, or ``None`` when no call
was replaced. Untouched messages pass through by identity, except that once
anything was rewritten every AIMessage loses its ``resp_`` response id (see
the module docstring: the server-side history behind that id still holds
the original arguments). Only calls with a non-empty string id that is
unique within its message are offered to the selector: every surface is
addressed by id, so nothing else can be matched across surfaces, and an id
a malformed provider payload repeats inside one AIMessage could only be
rewritten for *all* of its occurrences at once — a failed sibling would
take on the successful call's arguments (review on #5374). Such calls are
conservatively left alone.
"""
updated: list[Any] = []
changed = False
for message in messages:
patched = message
if isinstance(message, AIMessage) and message.tool_calls:
replacements: dict[str, dict[str, Any]] = {}
duplicated = _duplicated_call_ids(message.tool_calls)
for tool_call in message.tool_calls:
if not isinstance(tool_call, dict):
continue
call_id = tool_call.get("id")
if not isinstance(call_id, str) or not call_id or call_id in duplicated:
continue
new_args = replacement_for(message, tool_call)
if new_args is not None:
replacements[call_id] = new_args
if replacements:
patched = rewrite_tool_call_args(message, replacements)
if patched is not message:
changed = True
updated.append(patched)
if not changed:
return None
return [_without_response_chain_id(message) for message in updated]
def _duplicated_call_ids(tool_calls: Sequence[Any]) -> set[str]:
"""Ids that occur more than once in one message's structured tool-call list (the list every surface mirrors).
Only non-empty string ids are counted: a list or dict id from a malformed
provider payload is unhashable and must be skipped, never hashed, or the
whole model call would fail (review on #5374).
"""
counts = Counter(call_id for tool_call in tool_calls if isinstance(tool_call, dict) and isinstance(call_id := tool_call.get("id"), str) and call_id)
return {call_id for call_id, count in counts.items() if count > 1}
@dataclass(frozen=True, slots=True)
class ToolCallOccurrence:
"""One tool call on one AIMessage, paired with the ToolMessage that answered it (``None`` if unanswered)."""
#: Position of ``message`` in the history it was paired from.
index: int
message: AIMessage
tool_call: dict[str, Any]
result: ToolMessage | None
@property
def call_id(self) -> str:
return self.tool_call["id"]
@property
def name(self) -> str:
name = self.tool_call.get("name")
return name if isinstance(name, str) else ""
@property
def args(self) -> dict[str, Any]:
args = self.tool_call.get("args")
return args if isinstance(args, dict) else {}
def pair_tool_call_results(messages: Sequence[Any]) -> list[ToolCallOccurrence]:
"""Pair every AIMessage tool call carrying a non-empty string id with the ToolMessage that answered it.
Walks ``messages`` in document order. Each AIMessage opens its own calls,
and a ToolMessage answers the still-open call with its id from the *most
recent preceding* AIMessage only — the rule ``DanglingToolCallMiddleware``
applies: a result never answers a call from an earlier turn. So ids that
repeat across turns pair per occurrence, an interrupted call whose id a
later turn reused stays unanswered instead of inheriting that turn's result
(review on #5374), and stray or duplicate results are ignored. ``index`` is
the AIMessage's position in ``messages``, so callers can order events
across turns; the calls of one AIMessage share an index because they ran
concurrently, in no fixed order.
"""
occurrences: list[ToolCallOccurrence] = []
# Unanswered calls of the most recent AIMessage: id -> positions in ``occurrences``.
open_calls: dict[str, deque[int]] = defaultdict(deque)
for index, message in enumerate(messages):
if isinstance(message, AIMessage):
open_calls = defaultdict(deque)
for tool_call in message.tool_calls or ():
if not isinstance(tool_call, dict):
continue
call_id = tool_call.get("id")
if not isinstance(call_id, str) or not call_id:
continue
open_calls[call_id].append(len(occurrences))
occurrences.append(ToolCallOccurrence(index, message, tool_call, None))
elif isinstance(message, ToolMessage):
queue = open_calls.get(message.tool_call_id) if isinstance(message.tool_call_id, str) else None
if queue:
position = queue.popleft()
occurrences[position] = replace(occurrences[position], result=message)
return occurrences
def _without_response_chain_id(message: Any) -> Any:
"""Drop an OpenAI ``resp_`` response id so the adapter replays history instead of chaining to it."""
if not isinstance(message, AIMessage):
return message
response_metadata = message.response_metadata or {}
response_id = response_metadata.get("id")
if not (isinstance(response_id, str) and response_id.startswith("resp_")):
return message
return message.model_copy(update={"response_metadata": {key: value for key, value in response_metadata.items() if key != "id"}})
def rewrite_tool_call_args(message: AIMessage, replacements: ArgsReplacements) -> AIMessage:
"""Return ``message`` with the args of every tool call in ``replacements`` (by id) rewritten on all surfaces.
``message`` is never mutated; the same object comes back when no id matches.
"""
if not replacements:
return message
update: dict[str, Any] = {}
tool_calls = message.tool_calls or []
rewritten_calls = [dict(tool_call, args=new_args) if isinstance(tool_call, dict) and (new_args := _replacement_for_id(tool_call.get("id"), replacements)) is not None else tool_call for tool_call in tool_calls]
if _any_replaced(rewritten_calls, tool_calls):
update["tool_calls"] = rewritten_calls
tool_call_chunks = getattr(message, "tool_call_chunks", None)
if isinstance(tool_call_chunks, list):
rewritten_chunks = [dict(chunk, args=_serialize(new_args)) if isinstance(chunk, dict) and (new_args := _replacement_for_id(chunk.get("id"), replacements)) is not None else chunk for chunk in tool_call_chunks]
if _any_replaced(rewritten_chunks, tool_call_chunks):
update["tool_call_chunks"] = rewritten_chunks
additional_kwargs = message.additional_kwargs or {}
raw_tool_calls = additional_kwargs.get("tool_calls")
if isinstance(raw_tool_calls, list):
rewritten_raw = [_rewrite_raw_tool_call(entry, replacements) for entry in raw_tool_calls]
if _any_replaced(rewritten_raw, raw_tool_calls):
update["additional_kwargs"] = {**additional_kwargs, "tool_calls": rewritten_raw}
if isinstance(message.content, list):
rewritten_content = [_rewrite_content_block(block, replacements) for block in message.content]
if _any_replaced(rewritten_content, message.content):
update["content"] = rewritten_content
return message.model_copy(update=update) if update else message
def _replacement_for_id(identifier: Any, replacements: ArgsReplacements) -> dict[str, Any] | None:
"""Non-string ids (malformed provider payloads) never match, and never raise from a membership probe."""
return replacements.get(identifier) if isinstance(identifier, str) else None
def _any_replaced(rewritten: Sequence[Any], original: Sequence[Any]) -> bool:
return any(new is not old for new, old in zip(rewritten, original, strict=True))
def _serialize(args: dict[str, Any]) -> str:
return json.dumps(args, ensure_ascii=False)
def _rewrite_raw_tool_call(entry: Any, replacements: ArgsReplacements) -> Any:
"""Rewrite one raw provider tool-call payload (OpenAI ``function.arguments`` JSON string, or flattened variants)."""
if not isinstance(entry, dict):
return entry
new_args = _replacement_for_id(entry.get("id"), replacements)
if new_args is None:
return entry
function = entry.get("function")
if isinstance(function, dict):
return {**entry, "function": {**function, "arguments": _serialize(new_args)}}
if isinstance(entry.get("arguments"), str):
return {**entry, "arguments": _serialize(new_args)}
if isinstance(entry.get("args"), dict):
return {**entry, "args": new_args}
return entry
def _rewrite_content_block(block: Any, replacements: ArgsReplacements) -> Any:
"""Rewrite one content block that carries tool-call arguments; anything else passes through by identity."""
if not isinstance(block, dict):
return block
block_type = block.get("type")
if block_type == "tool_use":
# Anthropic: ``partial_json`` is dropped so it cannot leak the old payload.
new_args = _replacement_for_id(block.get("id"), replacements)
if new_args is None:
return block
rewritten = {key: value for key, value in block.items() if key != "partial_json"}
rewritten["input"] = new_args
return rewritten
if block_type == "function_call":
# OpenAI Responses (``responses/v1``): matched by ``call_id``; the ``fc_…`` item id and status are kept.
new_args = _replacement_for_id(block.get("call_id"), replacements)
if new_args is None:
return block
return {**block, "arguments": _serialize(new_args)}
if block_type in ("tool_call", "tool_call_chunk"):
# LangChain standard content (``v1``): ``args`` is a dict on tool_call and a JSON string on
# tool_call_chunk; ``extras.arguments`` (raw provider string) wins in the Responses translator.
new_args = _replacement_for_id(block.get("id"), replacements)
if new_args is None:
return block
rewritten = {**block, "args": new_args if block_type == "tool_call" else _serialize(new_args)}
extras = block.get("extras")
if isinstance(extras, dict) and "arguments" in extras:
rewritten["extras"] = {**extras, "arguments": _serialize(new_args)}
return rewritten
return block