822 lines
35 KiB
Python
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)
|