deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/agents/middlewares/summarization_middleware.py
2026-09-07 18:24:55 +08:00

822 lines
35 KiB
Python

"""Summarization middleware extensions for DeerFlow.
Unlike LangChain's stock ``SummarizationMiddleware`` (which runs in
``before_model`` and **persists** a destructive ``RemoveMessage(REMOVE_ALL_MESSAGES)``
to the checkpoint — wiping the original Q&A from both the model context *and*
the conversation the frontend renders), this middleware compresses the history
**transiently in ``wrap_model_call``**. The full original messages stay in state
(so the UI keeps showing the user's real conversation); only the message list
handed to the model for a single call is replaced with ``[summary, *recent]``.
The generated summary is cached per-thread so we don't pay an extra LLM call on
every model invocation: once a thread is summarized, the cached summary is reused
until enough *new* messages accumulate to push the compressed context back over
the trigger threshold (sticky boundary, mirroring the stock middleware's
post-compaction behavior).
"""
from __future__ import annotations
import logging
import threading
from collections import OrderedDict
from collections.abc import Awaitable, Callable, Collection
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.agents.middleware.types import ModelCallResult, ModelRequest, ModelResponse
from langchain_core.messages import (
AIMessage,
AnyMessage,
HumanMessage,
SystemMessage,
ToolMessage,
message_to_dict,
messages_from_dict,
)
from langgraph.config import get_config
from langgraph.runtime import Runtime
logger = logging.getLogger(__name__)
# Upper bound on the per-thread summary cache so a long-lived process serving
# many threads can't grow it without limit. Oldest entries are evicted first.
_SUMMARY_CACHE_MAX_THREADS = 512
# After compaction the preserved window must land *comfortably* below every
# trigger threshold, otherwise the "compressed" context is still over the
# trigger and re-summarizes on the very next turn (the sticky boundary never
# sticks). We force the kept suffix down to this fraction of each trigger so a
# re-compaction only happens once meaningfully *new* content accumulates — not
# on every follow-up question. ``keep`` only sets a *target*; this is the hard
# headroom guarantee on top of it.
_KEEP_HEADROOM_FRACTION = 0.75
@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
@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
def _tool_call_path(tool_call: dict[str, Any]) -> str | None:
"""Best-effort extraction of a file path argument from a read_file-like tool call."""
args = tool_call.get("args") or {}
if not isinstance(args, dict):
return None
for key in ("path", "file_path", "filepath"):
value = args.get(key)
if isinstance(value, str) and value:
return value
return None
def _clone_ai_message(
message: AIMessage,
tool_calls: list[dict[str, Any]],
*,
content: Any | None = None,
) -> AIMessage:
"""Clone an AIMessage while replacing its tool_calls list and optional content."""
update: dict[str, Any] = {"tool_calls": tool_calls}
if content is not None:
update["content"] = content
return message.model_copy(update=update)
@dataclass
class _SkillBundle:
"""Skill-related tool calls and tool results associated with one AIMessage."""
ai_index: int
skill_tool_indices: tuple[int, ...]
skill_tool_call_ids: frozenset[str]
skill_tool_tokens: int
skill_key: str
@dataclass(frozen=True)
class _SummaryCacheEntry:
"""A cached, reusable summary for one thread.
``summarized_ids`` are the ids of the original messages folded into
``summary_text``. The cache is only reusable while every one of those ids is
still present in the live message list and no AIMessage/ToolMessage pair was
split when the boundary was drawn (``reusable``) — otherwise filtering by id
could orphan a tool result, so we fall back to re-summarizing.
"""
summarized_ids: frozenset[str]
summary_text: str
preserved_ids: frozenset[str] = frozenset()
preserved_messages: tuple[AnyMessage, ...] = ()
# Process-wide summary cache shared across every middleware instance.
#
# ``make_lead_agent`` rebuilds the agent — and therefore a *fresh*
# ``DeerFlowSummarizationMiddleware`` — on every run. If the cache lived on the
# instance (``self._summary_cache``) it would be empty at the start of each
# user turn, forcing a brand-new (blocking) summarization LLM call on *every*
# turn of a long conversation. Hoisting it to module scope, keyed by
# ``thread_id``, lets a summary computed on one turn be reused on the next,
# so only the first threshold crossing (or a genuine re-compaction) pays for an
# LLM call. A lock guards the non-atomic LRU eviction against concurrent
# threads (subagents / parallel conversations share this process).
_GLOBAL_SUMMARY_CACHE: OrderedDict[str, _SummaryCacheEntry] = OrderedDict()
_GLOBAL_SUMMARY_CACHE_LOCK = threading.Lock()
_STATE_CACHE_KEY = "context_summary_cache"
def _clear_global_summary_cache() -> None:
"""Drop every cached summary. Intended for tests/diagnostics only."""
with _GLOBAL_SUMMARY_CACHE_LOCK:
_GLOBAL_SUMMARY_CACHE.clear()
def _serialize_cache_entry(entry: _SummaryCacheEntry) -> dict[str, Any]:
"""Serialize a summary cache entry so it can live in LangGraph state."""
return {
"version": 1,
"summarized_ids": sorted(entry.summarized_ids),
"summary_text": entry.summary_text,
"preserved_ids": sorted(entry.preserved_ids),
"preserved_messages": [message_to_dict(message) for message in entry.preserved_messages],
}
def _deserialize_cache_entry(raw: Any) -> _SummaryCacheEntry | None:
"""Best-effort restore of a state-backed summary cache entry."""
if not isinstance(raw, dict):
return None
try:
summary_text = raw.get("summary_text")
if not isinstance(summary_text, str) or not summary_text:
return None
summarized_ids = raw.get("summarized_ids") or []
preserved_ids = raw.get("preserved_ids") or []
preserved_messages_raw = raw.get("preserved_messages") or []
if not isinstance(summarized_ids, list) or not isinstance(preserved_ids, list) or not isinstance(preserved_messages_raw, list):
return None
preserved_messages = tuple(messages_from_dict(preserved_messages_raw))
return _SummaryCacheEntry(
summarized_ids=frozenset(str(item) for item in summarized_ids if item is not None),
summary_text=summary_text,
preserved_ids=frozenset(str(item) for item in preserved_ids if item is not None),
preserved_messages=preserved_messages,
)
except Exception:
logger.debug("summarization: failed to restore state-backed summary cache", exc_info=True)
return None
def _has_nonempty_human_message(message: AnyMessage) -> bool:
"""Whether a message is a real, non-empty user turn.
Some OpenAI-compatible gateways reject a completion request that contains
only ``system`` / ``assistant`` / ``tool`` roles. That can legitimately
happen after compaction: the active model call is continuing a tool turn,
while the user request that started the turn has already moved into the
internal summary.
"""
if not isinstance(message, HumanMessage):
return False
content = message.content
if isinstance(content, str):
return bool(content.strip())
if isinstance(content, list):
return any(bool(block.strip()) if isinstance(block, str) else bool((block.get("text") or block.get("content") or "").strip()) if isinstance(block, dict) else False for block in content)
return bool(str(content).strip())
def _ensure_latest_human_message(
source_messages: list[AnyMessage],
compressed_messages: list[AnyMessage],
) -> list[AnyMessage]:
"""Keep the latest real user request visible to strict chat gateways.
The message is inserted after leading system messages, so it neither
exposes the internal summary to the UI nor splits an AI tool-call from its
following ToolMessages. When there was no user turn in the source at all,
do not fabricate one: that is an upstream caller bug and should remain
diagnosable rather than silently changing the task.
"""
if any(_has_nonempty_human_message(message) for message in compressed_messages):
return compressed_messages
latest_human = next(
(message for message in reversed(source_messages) if _has_nonempty_human_message(message)),
None,
)
if latest_human is None:
logger.warning("summarization: compressed model request has no user message and source history contains none; passing it through unchanged")
return compressed_messages
insert_at = 0
while insert_at < len(compressed_messages) and isinstance(compressed_messages[insert_at], SystemMessage):
insert_at += 1
return [*compressed_messages[:insert_at], latest_human, *compressed_messages[insert_at:]]
class DeerFlowSummarizationMiddleware(SummarizationMiddleware):
"""Summarization middleware with pre-compression hook dispatch and skill rescue."""
def __init__(
self,
*args,
skills_container_path: str | None = None,
skill_file_read_tool_names: Collection[str] | None = None,
before_summarization: list[BeforeSummarizationHook] | None = None,
preserve_recent_skill_count: int = 5,
preserve_recent_skill_tokens: int = 25_000,
preserve_recent_skill_tokens_per_skill: int = 5_000,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
self._skills_container_path = skills_container_path or "/mnt/skills"
self._skill_file_read_tool_names = frozenset(skill_file_read_tool_names or {"read_file", "read", "view", "cat"})
self._before_summarization_hooks = before_summarization or []
self._preserve_recent_skill_count = max(0, preserve_recent_skill_count)
self._preserve_recent_skill_tokens = max(0, preserve_recent_skill_tokens)
self._preserve_recent_skill_tokens_per_skill = max(0, preserve_recent_skill_tokens_per_skill)
# Per-thread cache of the most recent summary: thread_id -> (summarized message ids, summary text).
# Points at the *process-wide* cache (not a fresh per-instance dict) so a
# summary survives across runs/turns — see ``_GLOBAL_SUMMARY_CACHE``.
self._summary_cache: OrderedDict[str, _SummaryCacheEntry] = _GLOBAL_SUMMARY_CACHE
self._pending_state_cache_update: dict[str, Any] | None = None
@override
def before_model(self, state: AgentState, runtime: Runtime) -> dict | None:
"""Disable the stock middleware's destructive ``before_model`` compaction.
Compression now happens transiently in :meth:`wrap_model_call`, so the
persisted conversation (and the frontend's view of it) is never mutated.
"""
return None
@override
async def abefore_model(self, state: AgentState, runtime: Runtime) -> dict | None:
"""Async counterpart of :meth:`before_model` — also a no-op."""
return None
@override
def after_model(self, state: AgentState, runtime: Runtime) -> dict | None:
"""Persist the latest reusable summary cache into thread state."""
if self._pending_state_cache_update is None:
return None
update = self._pending_state_cache_update
self._pending_state_cache_update = None
return {_STATE_CACHE_KEY: update}
@override
async def aafter_model(self, state: AgentState, runtime: Runtime) -> dict | None:
"""Async counterpart of :meth:`after_model`."""
return self.after_model(state, runtime)
# ---- wrap_model_call: transient, display-preserving compression ----------
@override
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelCallResult:
compressed = self._compressed_messages(request)
if compressed is not None:
request = request.override(messages=compressed)
return handler(request)
@override
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelCallResult:
compressed = await self._acompressed_messages(request)
if compressed is not None:
request = request.override(messages=compressed)
return await handler(request)
def _compressed_messages(self, request: ModelRequest) -> list[AnyMessage] | None:
"""Build the compressed message list for this call (sync summary path).
Returns ``None`` when no compression is needed (the request is passed
through untouched).
"""
plan = self._plan_compression(request)
if plan is None:
return None
action, payload = plan
if action == "reuse":
return payload # type: ignore[return-value]
# action == "summarize"
messages_to_summarize, preserved, thread_id = payload # type: ignore[misc]
self._emit_compacting_notice()
self._fire_hooks(messages_to_summarize, preserved, getattr(request, "runtime", None))
summary = self._create_summary(messages_to_summarize)
return self._finish_summarize(messages_to_summarize, preserved, summary, thread_id)
async def _acompressed_messages(self, request: ModelRequest) -> list[AnyMessage] | None:
"""Async counterpart of :meth:`_compressed_messages`."""
plan = self._plan_compression(request)
if plan is None:
return None
action, payload = plan
if action == "reuse":
return payload # type: ignore[return-value]
messages_to_summarize, preserved, thread_id = payload # type: ignore[misc]
self._emit_compacting_notice()
self._fire_hooks(messages_to_summarize, preserved, getattr(request, "runtime", None))
summary = await self._acreate_summary(messages_to_summarize)
return self._finish_summarize(messages_to_summarize, preserved, summary, thread_id)
@staticmethod
def _emit_compacting_notice() -> None:
"""Best-effort UI hint that a (blocking) summarization LLM call is starting.
Only the genuine summarize path calls this — the cheap cache-reuse path
does no LLM work and stays silent. Emitted on LangGraph's ``custom``
stream channel; the chat frontend renders it as a transient toast. A
no-op when there is no active stream writer (e.g. unit tests, non-stream
runs), so callers never branch.
"""
# Always leave a backend trace of the genuine (blocking) compaction, even
# when there is no stream writer to render the frontend toast (scheduled
# runs, channels, tests). Grep ``上下文压缩`` / ``context compaction`` in
# the Gateway log to confirm a real summarization LLM call fired.
logger.info("summarization: 上下文压缩(真正摘要,触发一次 LLM 调用) context compaction firing")
try:
from langgraph.config import get_stream_writer
writer = get_stream_writer()
except Exception:
return
if writer is None:
return
try:
writer({"type": "context_compacting", "message": "正在压缩上下文以容纳更长对话…"})
except Exception:
logger.debug("summarization: failed to emit context_compacting notice", exc_info=True)
def _plan_compression(
self,
request: ModelRequest,
) -> tuple[str, Any] | None:
"""Decide what to do for this model call, doing **no** LLM work.
Returns one of:
- ``None`` — pass the request through unchanged (no compression needed).
- ``("reuse", compressed_messages)`` — reuse the cached summary.
- ``("summarize", (messages_to_summarize, preserved, thread_id))`` — the
caller must run the (a)sync summarizer, then call
:meth:`_finish_summarize`.
"""
messages = list(getattr(request, "messages", None) or [])
if not messages:
return None
self._ensure_message_ids(messages)
# Threshold is the first gate: while the *full* current conversation is
# still under the configured trigger, there is nothing to compress —
# pass it through untouched. No cache lookup, no summary substitution.
# Compression only ever engages once the real context crosses the
# trigger, which is exactly when it is needed.
if not self._should_summarize(messages, self.token_counter(messages)):
return None
thread_id = self._resolve_request_thread_id(request)
if thread_id is None:
logger.debug("summarization: thread_id unresolved; cross-turn summary cache disabled for this call")
cached = self._get_cache_entry(thread_id) if thread_id else None
if cached is None:
cached = _deserialize_cache_entry(getattr(request, "state", {}).get(_STATE_CACHE_KEY))
if cached is not None and thread_id:
self._store_cache_entry(thread_id, cached)
present_ids = {m.id for m in messages}
cache_usable = cached is not None and cached.summarized_ids <= present_ids and cached.preserved_ids <= present_ids
if cache_usable:
preserved = list(cached.preserved_messages)
excluded_ids = cached.summarized_ids | cached.preserved_ids
preserved.extend(m for m in messages if m.id not in excluded_ids)
effective = [self._summary_message(cached.summary_text), *preserved]
effective = _ensure_latest_human_message(messages, effective)
# Decide re-summarization on the *counted* size of the compressed
# context only — NOT ``_should_summarize`` — because that folds in the
# last AIMessage's reported ``usage_metadata`` (system prompt + tool
# schemas, fixed overhead compression can't shrink). Using it here
# would re-summarize on every single turn once that overhead alone
# approaches the trigger.
if not self._exceeds_trigger_by_count(effective):
# Refresh LRU recency and reuse without an LLM call.
self._touch_cache_entry(thread_id)
logger.info(
"summarization: 上下文压缩(复用缓存摘要,无 LLM 调用) context compaction reusing cached summary thread=%s full_tokens=%d compressed_tokens=%d preserved_msgs=%d",
thread_id,
self.token_counter(messages),
self.token_counter(effective),
len(preserved),
)
return ("reuse", effective)
# New messages have pushed the compressed context back over the
# threshold — fall through and re-summarize a larger prefix.
# Genuine (blocking) LLM summarize path. Log *why* we are paying for it so
# an "it compacts every turn" report is diagnosable straight from the
# backend log: cache=miss → thread_id/key problem; unusable → message-id
# instability; stale → preserved genuinely grew back over the trigger.
cache_state = "miss" if cached is None else "stale"
logger.info(
"summarization: compacting thread=%s cache=%s counted_tokens=%d msgs=%d",
thread_id,
cache_state,
self.token_counter(messages),
len(messages),
)
cutoff_index = self._determine_cutoff_index(messages)
if cutoff_index <= 0:
return None
# Force the preserved window below every trigger (with headroom) so the
# post-compaction context doesn't immediately re-trigger. Without this,
# a ``keep`` expressed in *messages* (e.g. 20) can itself blow past a
# *token* trigger (e.g. 60000) when recent messages are large (tool
# outputs, pasted docs) — and then every follow-up question re-summarizes.
cutoff_index = self._apply_stickiness_headroom(messages, cutoff_index)
if cutoff_index <= 0:
return None
messages_to_summarize, preserved = self._partition_with_skill_rescue(messages, cutoff_index)
if not messages_to_summarize:
return None
return ("summarize", (messages_to_summarize, preserved, thread_id))
def _finish_summarize(
self,
messages_to_summarize: list[AnyMessage],
preserved: list[AnyMessage],
summary: str,
thread_id: str | None,
) -> list[AnyMessage]:
"""Assemble the compressed message list and update the per-thread cache."""
compressed = [self._summary_message(summary), *preserved]
compressed = _ensure_latest_human_message([*messages_to_summarize, *preserved], compressed)
if thread_id:
summarized_ids = frozenset(m.id for m in messages_to_summarize if m.id is not None)
preserved_ids = {m.id for m in preserved if m.id is not None}
self._store_cache_entry(
thread_id,
entry := _SummaryCacheEntry(
summarized_ids=summarized_ids,
summary_text=summary,
preserved_ids=frozenset(preserved_ids),
preserved_messages=tuple(preserved),
),
)
self._pending_state_cache_update = _serialize_cache_entry(entry)
return compressed
def _get_cache_entry(self, thread_id: str) -> _SummaryCacheEntry | None:
"""Return the cached entry for ``thread_id`` (thread-safe read)."""
with _GLOBAL_SUMMARY_CACHE_LOCK:
return self._summary_cache.get(thread_id)
def _touch_cache_entry(self, thread_id: str) -> None:
"""Refresh LRU recency for ``thread_id`` if still present (thread-safe)."""
with _GLOBAL_SUMMARY_CACHE_LOCK:
if thread_id in self._summary_cache:
self._summary_cache.move_to_end(thread_id)
def _store_cache_entry(self, thread_id: str, entry: _SummaryCacheEntry) -> None:
"""Insert/refresh a cache entry, evicting the oldest thread if over capacity."""
with _GLOBAL_SUMMARY_CACHE_LOCK:
self._summary_cache[thread_id] = entry
self._summary_cache.move_to_end(thread_id)
while len(self._summary_cache) > _SUMMARY_CACHE_MAX_THREADS:
self._summary_cache.popitem(last=False)
def _exceeds_trigger_by_count(self, messages: list[AnyMessage]) -> bool:
"""Trigger check using only the *counted* size of ``messages``.
Deliberately ignores the stock middleware's reported-``usage_metadata``
heuristic: that number includes fixed system-prompt + tool-schema
overhead which compression can never shrink, so relying on it to decide
*re*-summarization makes a long, tool-heavy thread re-compact on every
turn. The first-ever decision to compress still uses the full
``_should_summarize`` (reported usage is a fair signal there); only the
cache-reuse-vs-recompact decision uses this stricter, content-only check.
"""
total = self.token_counter(messages)
for kind, value in getattr(self, "_trigger_conditions", []):
if kind == "messages" and len(messages) >= value:
return True
if kind == "tokens" and total >= value:
return True
return False
def _apply_stickiness_headroom(self, messages: list[AnyMessage], cutoff_index: int) -> int:
"""Advance ``cutoff_index`` (keep fewer recent messages) so the preserved
suffix lands below every trigger threshold with headroom.
Only ever *raises* the cutoff (summarizes more / keeps less); never keeps
more than the caller's ``keep`` policy already chose. This is what makes
the sticky boundary actually stick: after compaction the kept context is
guaranteed under ~``_KEEP_HEADROOM_FRACTION`` of each trigger, so a
re-compaction waits for genuinely new content instead of firing every turn.
"""
n = len(messages)
if n <= 1:
return cutoff_index
for kind, value in getattr(self, "_trigger_conditions", []):
if kind == "messages":
keep_cap = max(1, int(int(value) * _KEEP_HEADROOM_FRACTION))
cutoff_index = max(cutoff_index, n - keep_cap)
elif kind == "tokens":
target = max(1, int(int(value) * _KEEP_HEADROOM_FRACTION))
cutoff_index = max(cutoff_index, self._suffix_token_cutoff(messages, target))
# "fraction" triggers depend on model profile limits; the message/token
# caps above (when present) already bound the window, so skip them here.
# Always keep at least the final message, and re-snap to a safe AI/Tool
# boundary so we never orphan a tool result.
cutoff_index = min(cutoff_index, n - 1)
return self._find_safe_cutoff_point(messages, cutoff_index)
def _suffix_token_cutoff(self, messages: list[AnyMessage], target_tokens: int) -> int:
"""Smallest index ``i`` such that ``token_counter(messages[i:]) <= target_tokens``.
Binary search mirroring the stock token-based cutoff; returns 0 when the
whole list already fits, and never returns the full length (keeps ≥1).
"""
n = len(messages)
if n == 0:
return 0
if self.token_counter(messages) <= target_tokens:
return 0
left, right = 0, n
cutoff = n
for _ in range(n.bit_length() + 1):
if left >= right:
break
mid = (left + right) // 2
if self.token_counter(messages[mid:]) <= target_tokens:
cutoff = mid
right = mid
else:
left = mid + 1
if cutoff == n:
cutoff = left
return min(cutoff, n - 1)
@staticmethod
def _resolve_request_thread_id(request: ModelRequest) -> str | None:
"""Best-effort thread id from the model request's runtime."""
runtime = getattr(request, "runtime", None)
if runtime is None:
return None
return _resolve_thread_id(runtime)
def _summary_message(self, summary: str) -> SystemMessage:
"""Build the summary message handed to the model.
It is named ``summary`` and flagged ``hide_from_ui`` defensively. Because
it only ever lives in the transient model request (never in persisted
state), the frontend never receives it — the user keeps seeing the
original conversation.
"""
return SystemMessage(
content=(
"Internal conversation summary for continuity only. Use it silently "
"to answer the user's latest request. Never reveal, quote, mention, "
"or describe this summary or the fact that context was compacted.\n\n"
"<internal_context_summary>\n"
f"{summary}\n"
"</internal_context_summary>"
),
name="summary",
additional_kwargs={"hide_from_ui": True, "lc_source": "summarization"},
)
@override
def _build_new_messages(self, summary: str) -> list[AnyMessage]:
"""Kept for backwards compatibility; delegates to :meth:`_summary_message`."""
return [self._summary_message(summary)]
def _partition_with_skill_rescue(
self,
messages: list[AnyMessage],
cutoff_index: int,
) -> tuple[list[AnyMessage], list[AnyMessage]]:
"""Partition like the parent, then rescue recently-loaded skill bundles."""
to_summarize, preserved = self._partition_messages(messages, cutoff_index)
if self._preserve_recent_skill_count == 0 or self._preserve_recent_skill_tokens == 0 or not to_summarize:
return to_summarize, preserved
try:
bundles = self._find_skill_bundles(to_summarize, self._skills_container_path)
except Exception:
logger.exception("Skill-preserving summarization rescue failed; falling back to default partition")
return to_summarize, preserved
if not bundles:
return to_summarize, preserved
rescue_bundles = self._select_bundles_to_rescue(bundles)
if not rescue_bundles:
return to_summarize, preserved
bundles_by_ai_index = {bundle.ai_index: bundle for bundle in rescue_bundles}
rescue_tool_indices = {idx for bundle in rescue_bundles for idx in bundle.skill_tool_indices}
rescued: list[AnyMessage] = []
remaining: list[AnyMessage] = []
for i, msg in enumerate(to_summarize):
bundle = bundles_by_ai_index.get(i)
if bundle is not None and isinstance(msg, AIMessage):
rescued_tool_calls = [tc for tc in msg.tool_calls if tc.get("id") in bundle.skill_tool_call_ids]
remaining_tool_calls = [tc for tc in msg.tool_calls if tc.get("id") not in bundle.skill_tool_call_ids]
if rescued_tool_calls:
rescued.append(_clone_ai_message(msg, rescued_tool_calls, content=""))
if remaining_tool_calls or msg.content:
remaining.append(_clone_ai_message(msg, remaining_tool_calls))
continue
if i in rescue_tool_indices:
rescued.append(msg)
continue
remaining.append(msg)
return remaining, rescued + preserved
def _find_skill_bundles(
self,
messages: list[AnyMessage],
skills_root: str,
) -> list[_SkillBundle]:
"""Locate AIMessage + paired ToolMessage groups that load skill files."""
bundles: list[_SkillBundle] = []
n = len(messages)
i = 0
while i < n:
msg = messages[i]
if not (isinstance(msg, AIMessage) and msg.tool_calls):
i += 1
continue
tool_calls = list(msg.tool_calls)
skill_paths_by_id: dict[str, str] = {}
for tc in tool_calls:
if self._is_skill_tool_call(tc, skills_root):
tc_id = tc.get("id")
path = _tool_call_path(tc)
if tc_id and path:
skill_paths_by_id[tc_id] = path
if not skill_paths_by_id:
i += 1
continue
skill_tool_tokens = 0
skill_key_parts: list[str] = []
skill_tool_indices: list[int] = []
matched_skill_call_ids: set[str] = set()
j = i + 1
while j < n and isinstance(messages[j], ToolMessage):
j += 1
for k in range(i + 1, j):
tool_msg = messages[k]
if isinstance(tool_msg, ToolMessage) and tool_msg.tool_call_id in skill_paths_by_id:
skill_tool_tokens += self.token_counter([tool_msg])
skill_key_parts.append(skill_paths_by_id[tool_msg.tool_call_id])
skill_tool_indices.append(k)
matched_skill_call_ids.add(tool_msg.tool_call_id)
if not skill_tool_indices:
i = j
continue
bundles.append(
_SkillBundle(
ai_index=i,
skill_tool_indices=tuple(skill_tool_indices),
skill_tool_call_ids=frozenset(matched_skill_call_ids),
skill_tool_tokens=skill_tool_tokens,
skill_key="|".join(sorted(skill_key_parts)),
)
)
i = j
return bundles
def _select_bundles_to_rescue(self, bundles: list[_SkillBundle]) -> list[_SkillBundle]:
"""Pick bundles to keep, walking newest-first under count/token budgets."""
selected: list[_SkillBundle] = []
if not bundles:
return selected
seen_skill_keys: set[str] = set()
total_tokens = 0
kept = 0
for bundle in reversed(bundles):
if kept >= self._preserve_recent_skill_count:
break
if bundle.skill_key in seen_skill_keys:
continue
if bundle.skill_tool_tokens > self._preserve_recent_skill_tokens_per_skill:
continue
if total_tokens + bundle.skill_tool_tokens > self._preserve_recent_skill_tokens:
continue
selected.append(bundle)
total_tokens += bundle.skill_tool_tokens
kept += 1
seen_skill_keys.add(bundle.skill_key)
selected.reverse()
return selected
def _is_skill_tool_call(self, tool_call: dict[str, Any], skills_root: str) -> bool:
"""Return True when ``tool_call`` reads a file under the configured skills root."""
name = tool_call.get("name") or ""
if name not in self._skill_file_read_tool_names:
return False
path = _tool_call_path(tool_call)
if not path:
return False
normalized_root = skills_root.rstrip("/")
return path == normalized_root or path.startswith(normalized_root + "/")
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)