mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-21 03:56:20 +00:00
* 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>
284 lines
14 KiB
Python
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
|