hataa 266883b3dd
fix(subagents): inherit summarization middleware and harden step capture (#3875 Phase 3) (#4009)
Phase 3 of #3875 — subagents previously inherited none of the lead's
context-compaction, so a deep-research subagent (max_turns up to 150)
could accumulate >1M cumulative input before max_turns/timeout/token_budget
engaged, even after Phase 2's budget capped the pathological tail.

- Gate the subagent runtime chain on the SAME ``app_config.summarization.enabled``
  switch the lead reads (per maintainer guidance in #3875), via the shared
  ``create_summarization_middleware`` factory. One config covers both chains;
  no separate ``subagents.summarization`` field. No-op when summarization is
  off (factory returns None).
- ``skip_memory_flush=True`` on the subagent path: the factory otherwise
  attaches ``memory_flush_hook`` (when memory.enabled), which flushes
  pre-compaction messages into durable memory keyed by thread_id. Subagents
  share the parent's thread_id, so without skipping the hook a subagent's
  internal turns would pollute the PARENT thread's durable memory
  (#3875 Phase 3 review point).
- Harden ``capture_new_step_messages`` to tolerate history contraction:
  summarization rewrites the messages channel via
  ``RemoveMessage(id=REMOVE_ALL_MESSAGES)``, shrinking len(messages) below
  the step-capture cursor. Without a reset, every step appended after the
  compaction point was dropped until length overtook the stale cursor (#3845
  interaction, maintainer validation point (a)). Cursor now resets to the
  new tail; id/content dedup prevents re-emitting pre-compaction steps.
- Couple the DEFAULT token-budget ceiling to ``summarization.enabled``
  (#3875 Phase 3 review point): 1M when compaction is on, 2M when off
  (preserves Phase 2's deliberate headroom for summarization-off
  deep-research runs that can exceed 1M). A user-set budget (global or
  per-agent) always wins regardless of the switch. Flagged tunable.

The summarization middleware does not implement ``consume_stop_reason``, so
the Phase 2 guard-cap stop-reason channel is unaffected.

Refs: https://github.com/bytedance/deer-flow/issues/3875
2026-07-10 11:17:35 +08:00

504 lines
21 KiB
Python

"""Summarization middleware extensions for DeerFlow."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Any, Protocol, override, runtime_checkable
from langchain.agents import AgentState
from langchain.agents.middleware import SummarizationMiddleware
from langchain_core.messages import AnyMessage, HumanMessage, RemoveMessage, get_buffer_string, trim_messages
from langgraph.config import get_config
from langgraph.constants import TAG_NOSTREAM
from langgraph.graph.message import REMOVE_ALL_MESSAGES
from langgraph.runtime import Runtime
from deerflow.agents.middlewares.dynamic_context_middleware import is_dynamic_context_reminder
from deerflow.config.app_config import get_app_config
from deerflow.models import create_chat_model
logger = logging.getLogger(__name__)
_SUMMARY_TRIGGER_MESSAGE_NAME = "summary"
@dataclass(frozen=True)
class SummarizationEvent:
"""Context emitted before conversation history is summarized away."""
messages_to_summarize: tuple[AnyMessage, ...]
preserved_messages: tuple[AnyMessage, ...]
thread_id: str | None
agent_name: str | None
runtime: Runtime
@dataclass(frozen=True)
class ContextCompactionResult:
"""Result of summarizing old context and retaining the active tail."""
summary_text: str
messages_to_summarize: tuple[AnyMessage, ...]
preserved_messages: tuple[AnyMessage, ...]
total_tokens: int
@runtime_checkable
class BeforeSummarizationHook(Protocol):
"""Hook invoked before summarization removes messages from state."""
def __call__(self, event: SummarizationEvent) -> None: ...
def _resolve_thread_id(runtime: Runtime) -> str | None:
"""Resolve the current thread ID from runtime context or LangGraph config."""
thread_id = runtime.context.get("thread_id") if runtime.context else None
if thread_id is None:
try:
config_data = get_config()
except RuntimeError:
return None
thread_id = config_data.get("configurable", {}).get("thread_id")
return thread_id
def _resolve_agent_name(runtime: Runtime) -> str | None:
"""Resolve the current agent name from runtime context or LangGraph config."""
agent_name = runtime.context.get("agent_name") if runtime.context else None
if agent_name is None:
try:
config_data = get_config()
except RuntimeError:
return None
agent_name = config_data.get("configurable", {}).get("agent_name")
return agent_name
class DeerFlowSummarizationMiddleware(SummarizationMiddleware):
"""Summarization middleware with pre-compression hook dispatch."""
def __init__(
self,
*args,
before_summarization: list[BeforeSummarizationHook] | None = None,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
self._before_summarization_hooks = before_summarization or []
# The summary LLM call runs inside a LangGraph middleware hook, so its token
# stream would otherwise be captured by the messages-tuple stream callback and
# broadcast to the frontend as a phantom AI message. Tag a dedicated model copy
# with TAG_NOSTREAM so the streaming handler skips it.
# Keep self.model untagged so the parent's profile / ls_params inspection still works.
#
# Preserve any tags already bound on the model (e.g. "middleware:summarize" set in
# lead_agent/agent.py for RunJournal attribution): RunnableBinding.with_config does a
# shallow merge that would otherwise overwrite the existing tags list entirely.
existing_tags = list((getattr(self.model, "config", None) or {}).get("tags") or [])
merged_tags = [*existing_tags, TAG_NOSTREAM] if TAG_NOSTREAM not in existing_tags else existing_tags
self._summary_model = self.model.with_config(tags=merged_tags)
@override
def _create_summary(self, messages_to_summarize: list[AnyMessage]) -> str | None:
return self._summarize_with(messages_to_summarize)
@override
async def _acreate_summary(self, messages_to_summarize: list[AnyMessage]) -> str | None:
return await self._asummarize_with(messages_to_summarize)
def _summarize_with(self, messages_to_summarize: list[AnyMessage], previous_summary: str | None = None) -> str | None:
"""Mirror the parent ``_create_summary`` but invoke the nostream-tagged model.
We do not swap ``self.model`` at the instance level: the agent/middleware is
cached and reused across concurrent runs, so a temporary swap would leak the
``RunnableBinding`` to other coroutines during ``await`` and break parent logic
that inspects the raw model (``profile`` / ``_get_ls_params``).
"""
if not messages_to_summarize:
return "No previous conversation history."
prompt = self._build_summary_prompt(messages_to_summarize, previous_summary=previous_summary)
if prompt is None:
return "Previous conversation was too long to summarize."
try:
response = self._summary_model.invoke(
prompt,
config={"metadata": {"lc_source": "summarization"}},
)
return response.text.strip()
except Exception:
logger.exception("Summary generation failed; skipping compaction this turn")
return None
async def _asummarize_with(self, messages_to_summarize: list[AnyMessage], previous_summary: str | None = None) -> str | None:
"""Async counterpart of :meth:`_summarize_with` using the nostream model."""
if not messages_to_summarize:
return "No previous conversation history."
prompt = self._build_summary_prompt(messages_to_summarize, previous_summary=previous_summary)
if prompt is None:
return "Previous conversation was too long to summarize."
try:
response = await self._summary_model.ainvoke(
prompt,
config={"metadata": {"lc_source": "summarization"}},
)
return response.text.strip()
except Exception:
logger.exception("Summary generation failed; skipping compaction this turn")
return None
@staticmethod
def _summary_count_message(summary_text: str) -> HumanMessage:
return HumanMessage(content=summary_text, name=_SUMMARY_TRIGGER_MESSAGE_NAME)
def _messages_for_trigger_count(self, messages: list[AnyMessage], summary_text: str | None) -> list[AnyMessage]:
if not summary_text:
return messages
return [*messages, self._summary_count_message(summary_text)]
@staticmethod
def _bound_text(text: str, cap: int) -> str:
if len(text) <= cap:
return text
if cap <= 0:
return ""
head = cap * 2 // 3
omitted_marker = "\n...\n"
if cap <= len(omitted_marker):
return text[:cap]
tail = max(0, cap - head - len(omitted_marker))
if tail == 0:
return text[:cap]
return f"{text[:head]}{omitted_marker}{text[-tail:]}"
def _trim_summary_section_text(self, text: str, max_tokens: int, *, strategy: str) -> str:
if not text.strip():
return ""
max_tokens = max(1, max_tokens)
try:
trimmed = trim_messages(
[HumanMessage(content=text)],
max_tokens=max_tokens,
token_counter=self.token_counter,
strategy=strategy,
allow_partial=True,
text_splitter=list,
)
if trimmed:
content = trimmed[-1].content
if isinstance(content, str) and content.strip():
return content
except Exception:
logger.debug("Failed to trim summary prompt section with token counter; falling back to deterministic text cap", exc_info=True)
return self._bound_text(text, max_tokens)
def _build_summary_input_text(self, formatted_messages: str, previous_summary: str | None = None) -> str | None:
if self.trim_tokens_to_summarize is None:
trimmed_new_messages = formatted_messages
trimmed_previous_summary = previous_summary.strip() if previous_summary else ""
else:
max_tokens = max(1, self.trim_tokens_to_summarize)
if previous_summary:
new_message_tokens = max(1, max_tokens // 2)
previous_summary_tokens = max(1, max_tokens - new_message_tokens)
trimmed_previous_summary = self._trim_summary_section_text(
previous_summary.strip(),
previous_summary_tokens,
strategy="last",
)
trimmed_new_messages = self._trim_summary_section_text(
formatted_messages,
new_message_tokens,
strategy="first",
)
else:
trimmed_previous_summary = ""
trimmed_new_messages = self._trim_summary_section_text(
formatted_messages,
max_tokens,
strategy="first",
)
parts: list[str] = []
if trimmed_previous_summary:
parts.extend(
[
"<existing_summary>",
trimmed_previous_summary,
"</existing_summary>",
"",
]
)
if trimmed_new_messages:
parts.extend(
[
"<new_messages>",
trimmed_new_messages,
"</new_messages>",
]
)
if not parts:
return None
return "\n".join(parts)
def _build_summary_prompt(self, messages_to_summarize: list[AnyMessage], previous_summary: str | None = None) -> str | None:
"""Build the summary prompt, returning ``None`` when trimming leaves nothing."""
trimmed_messages = self._trim_messages_for_summary(messages_to_summarize)
if not trimmed_messages:
trimmed_messages = messages_to_summarize[-1:]
if not trimmed_messages:
return None
# Format messages to avoid token inflation from metadata when str() is called on
# message objects.
formatted_messages = get_buffer_string(trimmed_messages)
formatted_messages = self._build_summary_input_text(formatted_messages, previous_summary=previous_summary)
if not formatted_messages:
return None
return self.summary_prompt.format(messages=formatted_messages).rstrip()
def before_model(self, state: AgentState, runtime: Runtime) -> dict | None:
return self._maybe_summarize(state, runtime)
async def abefore_model(self, state: AgentState, runtime: Runtime) -> dict | None:
return await self._amaybe_summarize(state, runtime)
def _prepare_compaction(
self,
state: AgentState,
*,
force: bool = False,
) -> tuple[list[AnyMessage], list[AnyMessage], str | None, int] | None:
messages = state["messages"]
self._ensure_message_ids(messages)
previous_summary = state.get("summary_text") if isinstance(state.get("summary_text"), str) else None
trigger_messages = self._messages_for_trigger_count(messages, previous_summary)
total_tokens = self.token_counter(trigger_messages)
if not force and not self._should_summarize(trigger_messages, total_tokens):
return None
cutoff_index = self._determine_cutoff_index(messages)
if cutoff_index <= 0:
return None
messages_to_summarize, preserved_messages = self._partition_messages(messages, cutoff_index)
messages_to_summarize, preserved_messages = self._preserve_dynamic_context_reminders(messages_to_summarize, preserved_messages)
if not messages_to_summarize:
return None
return messages_to_summarize, preserved_messages, previous_summary, total_tokens
def compact_state(
self,
state: AgentState,
runtime: Runtime,
*,
force: bool = False,
) -> ContextCompactionResult | None:
prepared = self._prepare_compaction(state, force=force)
if prepared is None:
return None
messages_to_summarize, preserved_messages, previous_summary, total_tokens = prepared
self._fire_hooks(messages_to_summarize, preserved_messages, runtime)
summary = self._summarize_with(messages_to_summarize, previous_summary=previous_summary)
if summary is None:
return None
return ContextCompactionResult(
summary_text=summary,
messages_to_summarize=tuple(messages_to_summarize),
preserved_messages=tuple(preserved_messages),
total_tokens=total_tokens,
)
async def acompact_state(
self,
state: AgentState,
runtime: Runtime,
*,
force: bool = False,
) -> ContextCompactionResult | None:
prepared = self._prepare_compaction(state, force=force)
if prepared is None:
return None
messages_to_summarize, preserved_messages, previous_summary, total_tokens = prepared
self._fire_hooks(messages_to_summarize, preserved_messages, runtime)
summary = await self._asummarize_with(messages_to_summarize, previous_summary=previous_summary)
if summary is None:
return None
return ContextCompactionResult(
summary_text=summary,
messages_to_summarize=tuple(messages_to_summarize),
preserved_messages=tuple(preserved_messages),
total_tokens=total_tokens,
)
def _maybe_summarize(self, state: AgentState, runtime: Runtime) -> dict | None:
result = self.compact_state(state, runtime, force=False)
if result is None:
return None
return {
"messages": [
RemoveMessage(id=REMOVE_ALL_MESSAGES),
*result.preserved_messages,
],
"summary_text": result.summary_text,
}
async def _amaybe_summarize(self, state: AgentState, runtime: Runtime) -> dict | None:
result = await self.acompact_state(state, runtime, force=False)
if result is None:
return None
return {
"messages": [
RemoveMessage(id=REMOVE_ALL_MESSAGES),
*result.preserved_messages,
],
"summary_text": result.summary_text,
}
def _preserve_dynamic_context_reminders(
self,
messages_to_summarize: list[AnyMessage],
preserved_messages: list[AnyMessage],
) -> tuple[list[AnyMessage], list[AnyMessage]]:
"""Keep hidden dynamic-context reminders and their ID-swap peers out of summary compression.
These reminders carry the current date and optional memory. If summarization
removes them, DynamicContextMiddleware can lose the already-injected reminder
and inject a replacement into the wrong point of the conversation.
The ID-swap triplet produced by ``_make_reminder_and_user_messages`` contains
three messages: ``SystemMessage(id=X)`` and ``HumanMessage(id=X__memory)`` are
both tagged with ``dynamic_context_reminder=True``, but ``HumanMessage(id=X__user)``
carries the original user content and is **not** tagged. Without peer rescue,
``__user`` would stay in ``to_summarize`` and be compressed into prose — orphaning
the tagged messages and losing the user question from the model's direct context.
This method rescues tagged reminders and also rescues any untagged messages whose
``id`` shares the same ``stable_id`` prefix (i.e. ``X__user``, ``X__memory``).
"""
reminders = [msg for msg in messages_to_summarize if is_dynamic_context_reminder(msg)]
if not reminders:
return messages_to_summarize, preserved_messages
# Collect the base IDs (the stable_id prefix) from tagged reminders.
# For a reminder with id="ctx-001__memory", the base is "ctx-001".
# For a reminder with id="ctx-001" (SystemMessage), the base is "ctx-001".
# removesuffix is suffix-only — it won't strip a "__" that sits in the
# middle of a stable_id (e.g. "ctx__001" stays intact, unlike rsplit
# which would mis-derive "ctx"). Only known ID-swap suffixes (__memory,
# __user) are stripped; __user is not tagged so won't appear in reminders,
# but is included defensively.
reminder_base_ids: set[str] = set()
for msg in reminders:
if msg.id:
base = msg.id.removesuffix("__memory").removesuffix("__user")
reminder_base_ids.add(base)
# Single-pass partition: walk messages_to_summarize in chronological order
# and rescue both tagged reminders and untagged ID-swap peers (whose id
# starts with a known base + "__"). This preserves the original message
# order within rescued — critical when multiple triplets land in one
# summarization window — and eliminates the need for id(m)-based dedup
# that the previous reminders+peers concatenation required.
rescued: list[AnyMessage] = []
remaining: list[AnyMessage] = []
for msg in messages_to_summarize:
if is_dynamic_context_reminder(msg) or (msg.id and any(msg.id.startswith(b + "__") for b in reminder_base_ids)):
rescued.append(msg)
else:
remaining.append(msg)
return remaining, rescued + preserved_messages
def _fire_hooks(
self,
messages_to_summarize: list[AnyMessage],
preserved_messages: list[AnyMessage],
runtime: Runtime,
) -> None:
if not self._before_summarization_hooks:
return
event = SummarizationEvent(
messages_to_summarize=tuple(messages_to_summarize),
preserved_messages=tuple(preserved_messages),
thread_id=_resolve_thread_id(runtime),
agent_name=_resolve_agent_name(runtime),
runtime=runtime,
)
for hook in self._before_summarization_hooks:
try:
hook(event)
except Exception:
hook_name = getattr(hook, "__name__", None) or type(hook).__name__
logger.exception("before_summarization hook %s failed", hook_name)
def create_summarization_middleware(
*,
app_config: Any | None = None,
keep: tuple[str, int | float] | None = None,
skip_memory_flush: bool = False,
) -> DeerFlowSummarizationMiddleware | None:
"""Create the configured summarization middleware.
Both the lead-agent automatic path and the manual context-compaction path
use this factory so model resolution, hooks, prompt config, and retention
defaults cannot drift.
``skip_memory_flush`` omits the ``memory_flush_hook`` that otherwise
flushes pre-compaction messages into the durable memory queue. The lead
chain keeps it (research should persist); the subagent chain sets it so a
subagent's INTERNAL turns (the "Task" human message + intermediate AI/tool
turns) are not written into the PARENT thread's durable memory — the hook
is keyed by ``thread_id`` and subagents share the parent's ``thread_id``
(#3875 Phase 3 review).
"""
resolved_app_config = app_config or get_app_config()
config = resolved_app_config.summarization
if not config.enabled:
return None
trigger = None
if config.trigger is not None:
if isinstance(config.trigger, list):
trigger = [item.to_tuple() for item in config.trigger]
else:
trigger = config.trigger.to_tuple()
if config.model_name:
model = create_chat_model(
name=config.model_name,
thinking_enabled=False,
app_config=resolved_app_config,
attach_tracing=False,
)
else:
model = create_chat_model(
thinking_enabled=False,
app_config=resolved_app_config,
attach_tracing=False,
)
model = model.with_config(tags=["middleware:summarize"])
kwargs: dict[str, Any] = {
"model": model,
"trigger": trigger,
"keep": keep or config.keep.to_tuple(),
}
if config.trim_tokens_to_summarize is not None:
kwargs["trim_tokens_to_summarize"] = config.trim_tokens_to_summarize
if config.summary_prompt is not None:
kwargs["summary_prompt"] = config.summary_prompt
hooks: list[BeforeSummarizationHook] = []
if resolved_app_config.memory.enabled and not skip_memory_flush:
from deerflow.agents.memory.summarization_hook import memory_flush_hook
hooks.append(memory_flush_hook)
return DeerFlowSummarizationMiddleware(
**kwargs,
before_summarization=hooks,
)