fix(summary): 修复对话摘要中间件的工具结果卸载链路,修正路径拼接错误并补充单元测试

This commit is contained in:
Wenjie Zhang 2026-04-12 23:38:13 +08:00
parent 63786b7c66
commit 9c8998699c
3 changed files with 303 additions and 64 deletions

View File

@ -10,7 +10,6 @@ from __future__ import annotations
import uuid import uuid
from collections.abc import Callable, Iterable, Mapping from collections.abc import Callable, Iterable, Mapping
from functools import partial from functools import partial
from pathlib import Path
from typing import Any, Literal, cast, override from typing import Any, Literal, cast, override
from langchain.agents import AgentState from langchain.agents import AgentState
@ -32,7 +31,7 @@ from langchain_core.messages.utils import (
from langgraph.graph.message import REMOVE_ALL_MESSAGES from langgraph.graph.message import REMOVE_ALL_MESSAGES
from langgraph.runtime import Runtime from langgraph.runtime import Runtime
from yuxi.utils.paths import OUTPUTS_DIR_NAME from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS
TokenCounter = Callable[[Iterable[MessageLikeRepresentation]], int] TokenCounter = Callable[[Iterable[MessageLikeRepresentation]], int]
@ -78,7 +77,7 @@ Messages to summarize:
_DEFAULT_MESSAGES_TO_KEEP = 20 _DEFAULT_MESSAGES_TO_KEEP = 20
_DEFAULT_FALLBACK_MESSAGE_COUNT = 15 _DEFAULT_FALLBACK_MESSAGE_COUNT = 15
_OFFLOAD_DIR = "/summary_offload" # 虚拟文件系统路径 _OFFLOAD_DIR = "summary_offload"
ContextFraction = tuple[Literal["fraction"], float] ContextFraction = tuple[Literal["fraction"], float]
ContextTokens = tuple[Literal["tokens"], int] ContextTokens = tuple[Literal["tokens"], int]
@ -95,14 +94,14 @@ def _get_approximate_token_counter(model: BaseChatModel) -> TokenCounter:
def _get_content_str(content: Any) -> str | None: def _get_content_str(content: Any) -> str | None:
"""Convert ToolMessage content to string for size checking.""" """Convert plain-text ToolMessage content to string for size checking."""
if isinstance(content, str): if isinstance(content, str):
return content return content
if isinstance(content, list): if isinstance(content, list):
if len(content) == 1 and isinstance(content[0], dict) and content[0].get("type") == "text": if len(content) == 1 and isinstance(content[0], dict) and content[0].get("type") == "text":
return str(content[0].get("text", "")) return str(content[0].get("text", ""))
return str(content) return None
return str(content) return None
def _format_offload_placeholder(file_path: str, content_sample: str) -> str: def _format_offload_placeholder(file_path: str, content_sample: str) -> str:
@ -115,7 +114,42 @@ def _format_offload_placeholder(file_path: str, content_sample: str) -> str:
) )
def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenCounter) -> dict[str, Any] | None: def _build_offload_file_path(msg: ToolMessage) -> str:
"""Build a read_file-compatible virtual path for offloaded tool output."""
tool_name = msg.name or "unknown"
message_id = msg.id or str(uuid.uuid4())[:8]
safe_name = "".join(c if c.isalnum() or c in "-_" else "_" for c in tool_name)
return f"{VIRTUAL_PATH_OUTPUTS}/{_OFFLOAD_DIR}/{safe_name}-{message_id}.txt"
def _write_offloaded_content(runtime: Runtime, file_path: str, content: str) -> tuple[bool, dict[str, Any]]:
"""Persist offloaded tool output into the active filesystem backend."""
from yuxi.agents.backends.composite import create_agent_composite_backend
backend = create_agent_composite_backend(runtime)
result = backend.write(file_path, content)
if result.error:
return False, {}
return True, result.files_update or {}
async def _awrite_offloaded_content(runtime: Runtime, file_path: str, content: str) -> tuple[bool, dict[str, Any]]:
"""Async variant of _write_offloaded_content."""
from yuxi.agents.backends.composite import create_agent_composite_backend
backend = create_agent_composite_backend(runtime)
result = await backend.awrite(file_path, content)
if result.error:
return False, {}
return True, result.files_update or {}
def _offload_tool_result(
msg: ToolMessage,
threshold: int,
token_counter: TokenCounter,
runtime: Runtime,
) -> dict[str, Any] | None:
"""卸载单个超阈值的工具结果. """卸载单个超阈值的工具结果.
Args: Args:
@ -141,10 +175,7 @@ def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenC
tool_name = msg.name or "unknown" tool_name = msg.name or "unknown"
tool_call_id = msg.tool_call_id or "" tool_call_id = msg.tool_call_id or ""
# 生成文件路径 (工具名称-xxx) file_path = _build_offload_file_path(msg)
message_id = msg.id or str(uuid.uuid4())[:8]
safe_name = "".join(c if c.isalnum() or c in "-_" else "_" for c in tool_name)
file_path = (Path(OUTPUTS_DIR_NAME) / f"{_OFFLOAD_DIR}/{safe_name}-{message_id}").as_posix()
# 构建文件头部信息 # 构建文件头部信息
header_lines = [ header_lines = [
@ -156,17 +187,9 @@ def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenC
] ]
header = "\n".join(header_lines) header = "\n".join(header_lines)
# 保存到 files 格式(包含头部信息) written, files_update = _write_offloaded_content(runtime, file_path, header + content_str)
from datetime import datetime if not written:
return None
timestamp = datetime.now().isoformat()
files_update = {
file_path: {
"content": [header + content_str],
"created_at": timestamp,
"modified_at": timestamp,
}
}
# 创建预览内容 # 创建预览内容
preview_lines = content_str.splitlines()[:10] preview_lines = content_str.splitlines()[:10]
@ -175,11 +198,11 @@ def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenC
# 替换消息内容为占位符 # 替换消息内容为占位符
msg.content = _format_offload_placeholder(file_path, content_sample) msg.content = _format_offload_placeholder(file_path, content_sample)
return files_update return files_update or {}
def _offload_tool_results( def _offload_tool_results(
messages: list[AnyMessage], threshold: int, token_counter: TokenCounter messages: list[AnyMessage], threshold: int, token_counter: TokenCounter, runtime: Runtime
) -> tuple[dict[str, Any], list[AnyMessage]]: ) -> tuple[dict[str, Any], list[AnyMessage]]:
"""扫描消息列表,卸载所有超阈值的工具结果. """扫描消息列表,卸载所有超阈值的工具结果.
@ -198,14 +221,57 @@ def _offload_tool_results(
if not isinstance(msg, ToolMessage): if not isinstance(msg, ToolMessage):
continue continue
result = _offload_tool_result(msg, threshold, token_counter) result = _offload_tool_result(msg, threshold, token_counter, runtime)
if result: if result is not None:
files_update.update(result) files_update.update(result)
modified_messages.append(msg) modified_messages.append(msg)
return files_update, modified_messages return files_update, modified_messages
async def _aoffload_tool_results(
messages: list[AnyMessage], threshold: int, token_counter: TokenCounter, runtime: Runtime
) -> tuple[dict[str, Any], list[AnyMessage]]:
"""Async variant of _offload_tool_results."""
files_update: dict[str, Any] = {}
modified_messages: list[AnyMessage] = []
for msg in messages:
if not isinstance(msg, ToolMessage):
continue
content_str = _get_content_str(msg.content)
if content_str is None:
continue
msg_tokens = token_counter([msg])
if msg_tokens <= threshold:
continue
tool_name = msg.name or "unknown"
tool_call_id = msg.tool_call_id or ""
file_path = _build_offload_file_path(msg)
header_lines = [
"=== Tool Invocation ===",
f"Tool: {tool_name}",
f"Tool Call ID: {tool_call_id}",
"=" * 40,
"",
]
header = "\n".join(header_lines)
written, result = await _awrite_offloaded_content(runtime, file_path, header + content_str)
if not written:
continue
preview_lines = content_str.splitlines()[:10]
content_sample = "\n".join(line[:500] for line in preview_lines)
msg.content = _format_offload_placeholder(file_path, content_sample)
files_update.update(result)
modified_messages.append(msg)
return files_update, modified_messages
class SummaryOffloadMiddleware(AgentMiddleware): class SummaryOffloadMiddleware(AgentMiddleware):
"""总结+工具结果卸载中间件. """总结+工具结果卸载中间件.
@ -322,7 +388,9 @@ class SummaryOffloadMiddleware(AgentMiddleware):
files_update: dict[str, Any] = {} files_update: dict[str, Any] = {}
modified_messages: list[AnyMessage] = [] modified_messages: list[AnyMessage] = []
agg_files, agg_msgs = _offload_tool_results(messages, self.summary_offload_threshold, self.token_counter) agg_files, agg_msgs = _offload_tool_results(
messages, self.summary_offload_threshold, self.token_counter, runtime
)
files_update = agg_files files_update = agg_files
modified_messages = agg_msgs modified_messages = agg_msgs
@ -330,40 +398,55 @@ class SummaryOffloadMiddleware(AgentMiddleware):
current_tokens = self.token_counter(messages) current_tokens = self.token_counter(messages)
trigger_value = self._get_token_trigger_value() trigger_value = self._get_token_trigger_value()
retention_limit = float("inf") if trigger_value is None:
if trigger_value: system_msg_count = 1 if messages and messages[0].type == "system" else 0
messages_to_process = messages[1:] if system_msg_count else messages
cutoff_relative = self._determine_cutoff_index(messages_to_process)
cutoff_index = system_msg_count + cutoff_relative
else:
retention_limit = trigger_value * self.max_retention_ratio retention_limit = trigger_value * self.max_retention_ratio
if current_tokens <= retention_limit: if current_tokens <= retention_limit:
if files_update: if files_update:
return {"files": files_update, "messages": modified_messages} result: dict[str, Any] = {"messages": modified_messages}
return None if files_update:
result["files"] = files_update
return result
return None
# 4. 超过 limit需要 Eviction (Summary) # 4. 超过 limit需要 Eviction (Summary)
system_msg_count = 0 system_msg_count = 0
messages_to_process = messages messages_to_process = messages
if messages and messages[0].type == "system": if messages and messages[0].type == "system":
system_msg_count = 1 system_msg_count = 1
messages_to_process = messages[1:] messages_to_process = messages[1:]
cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit)) cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit))
cutoff_index = system_msg_count + cutoff_relative cutoff_index = system_msg_count + cutoff_relative
if cutoff_index <= system_msg_count: if cutoff_index <= system_msg_count:
if files_update: if files_update:
return {"files": files_update, "messages": modified_messages} result = {"messages": modified_messages}
if files_update:
result["files"] = files_update
return result
return None return None
messages_to_summarize, preserved_messages = self._partition_messages(messages, cutoff_index) system_message = messages[0] if messages and messages[0].type == "system" else None
conversation_messages = messages[1:] if system_message is not None else messages
messages_to_summarize, preserved_messages = self._partition_messages(
conversation_messages, cutoff_index - system_msg_count
)
summary = self._create_summary(messages_to_summarize) summary = self._create_summary(messages_to_summarize)
new_messages = self._build_new_messages(summary) new_messages = self._build_new_messages(summary)
# 如果有 System Message需要保留在最前面 # 如果有 System Message需要保留在最前面
final_messages = [] final_messages = []
if system_msg_count > 0: if system_message is not None:
final_messages.append(messages[0]) final_messages.append(system_message)
final_messages.extend(new_messages) final_messages.extend(new_messages)
final_messages.extend(preserved_messages) final_messages.extend(preserved_messages)
@ -393,7 +476,9 @@ class SummaryOffloadMiddleware(AgentMiddleware):
files_update: dict[str, Any] = {} files_update: dict[str, Any] = {}
modified_messages: list[AnyMessage] = [] modified_messages: list[AnyMessage] = []
agg_files, agg_msgs = _offload_tool_results(messages, self.summary_offload_threshold, self.token_counter) agg_files, agg_msgs = await _aoffload_tool_results(
messages, self.summary_offload_threshold, self.token_counter, runtime
)
files_update = agg_files files_update = agg_files
modified_messages = agg_msgs modified_messages = agg_msgs
@ -401,40 +486,55 @@ class SummaryOffloadMiddleware(AgentMiddleware):
current_tokens = self.token_counter(messages) current_tokens = self.token_counter(messages)
trigger_value = self._get_token_trigger_value() trigger_value = self._get_token_trigger_value()
retention_limit = float("inf") if trigger_value is None:
if trigger_value: system_msg_count = 1 if messages and messages[0].type == "system" else 0
messages_to_process = messages[1:] if system_msg_count else messages
cutoff_relative = self._determine_cutoff_index(messages_to_process)
cutoff_index = system_msg_count + cutoff_relative
else:
retention_limit = trigger_value * self.max_retention_ratio retention_limit = trigger_value * self.max_retention_ratio
if current_tokens <= retention_limit: if current_tokens <= retention_limit:
if files_update: if files_update:
return {"files": files_update, "messages": modified_messages} result: dict[str, Any] = {"messages": modified_messages}
return None if files_update:
result["files"] = files_update
return result
return None
# 4. 超过 limit需要 Eviction (Summary) # 4. 超过 limit需要 Eviction (Summary)
system_msg_count = 0 system_msg_count = 0
messages_to_process = messages messages_to_process = messages
if messages and messages[0].type == "system": if messages and messages[0].type == "system":
system_msg_count = 1 system_msg_count = 1
messages_to_process = messages[1:] messages_to_process = messages[1:]
cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit)) cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit))
cutoff_index = system_msg_count + cutoff_relative cutoff_index = system_msg_count + cutoff_relative
if cutoff_index <= system_msg_count: if cutoff_index <= system_msg_count:
if files_update: if files_update:
return {"files": files_update, "messages": modified_messages} result = {"messages": modified_messages}
if files_update:
result["files"] = files_update
return result
return None return None
messages_to_summarize, preserved_messages = self._partition_messages(messages, cutoff_index) system_message = messages[0] if messages and messages[0].type == "system" else None
conversation_messages = messages[1:] if system_message is not None else messages
messages_to_summarize, preserved_messages = self._partition_messages(
conversation_messages, cutoff_index - system_msg_count
)
summary = await self._acreate_summary(messages_to_summarize) summary = await self._acreate_summary(messages_to_summarize)
new_messages = self._build_new_messages(summary) new_messages = self._build_new_messages(summary)
final_messages = [] final_messages = []
if system_msg_count > 0: if system_message is not None:
final_messages.append(messages[0]) final_messages.append(system_message)
final_messages.extend(new_messages) final_messages.extend(new_messages)
final_messages.extend(preserved_messages) final_messages.extend(preserved_messages)

View File

@ -0,0 +1,138 @@
from __future__ import annotations
from types import SimpleNamespace
import pytest
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
import yuxi.agents.middlewares.summary_middleware as summary_middleware
from yuxi.agents.middlewares.summary_middleware import SummaryOffloadMiddleware
from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS
class _DummyModel:
_llm_type = "test-chat"
profile = {"max_input_tokens": 128000}
def invoke(self, _prompt: str) -> SimpleNamespace:
return SimpleNamespace(text="summary")
@pytest.mark.unit
def test_offload_tool_result_writes_readable_outputs_path(monkeypatch: pytest.MonkeyPatch) -> None:
captured: dict[str, str] = {}
def _fake_write(_runtime, file_path: str, content: str) -> tuple[bool, dict]:
captured["file_path"] = file_path
captured["content"] = content
return True, {}
monkeypatch.setattr(summary_middleware, "_write_offloaded_content", _fake_write)
message = ToolMessage(content="line1\nline2", tool_call_id="tool-1", name="search", id="msg-1")
result = summary_middleware._offload_tool_result(
message,
threshold=1,
token_counter=lambda _messages: 10,
runtime=SimpleNamespace(),
)
assert result == {}
assert captured["file_path"] == f"{VIRTUAL_PATH_OUTPUTS}/summary_offload/search-msg-1.txt"
assert captured["content"].startswith("=== Tool Invocation ===\nTool: search\nTool Call ID: tool-1\n")
assert "文件路径: " in str(message.content)
assert captured["file_path"] in str(message.content)
@pytest.mark.unit
def test_offload_tool_result_skips_non_text_content(monkeypatch: pytest.MonkeyPatch) -> None:
called = False
def _fake_write(_runtime, _file_path: str, _content: str) -> tuple[bool, dict]:
nonlocal called
called = True
return True, {}
monkeypatch.setattr(summary_middleware, "_write_offloaded_content", _fake_write)
message = ToolMessage(
content=[{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "abc"}}],
tool_call_id="tool-1",
name="read_file",
id="msg-2",
)
result = summary_middleware._offload_tool_result(
message,
threshold=1,
token_counter=lambda _messages: 10,
runtime=SimpleNamespace(),
)
assert result is None
assert called is False
assert message.content[0]["type"] == "image"
@pytest.mark.unit
def test_before_model_excludes_system_message_from_summary(monkeypatch: pytest.MonkeyPatch) -> None:
middleware = SummaryOffloadMiddleware(
model=_DummyModel(),
trigger=("tokens", 10),
keep=("messages", 1),
token_counter=lambda _messages: 100,
summary_offload_threshold=10_000,
max_retention_ratio=0.5,
)
captured_ids: list[str] = []
monkeypatch.setattr(summary_middleware, "_offload_tool_results", lambda *args, **kwargs: ({}, []))
monkeypatch.setattr(middleware, "_find_cutoff_by_token_limit", lambda _messages, _limit: 2)
monkeypatch.setattr(
middleware,
"_create_summary",
lambda messages: captured_ids.extend([str(message.id) for message in messages]) or "summary",
)
messages = [
SystemMessage(content="sys", id="sys-1"),
HumanMessage(content="human-1", id="human-1"),
AIMessage(content="ai-1", id="ai-1"),
HumanMessage(content="human-2", id="human-2"),
]
result = middleware.before_model({"messages": messages}, SimpleNamespace())
assert captured_ids == ["human-1", "ai-1"]
assert result is not None
new_messages = result["messages"]
assert new_messages[1].id == "sys-1"
assert new_messages[2].content == "Here is a summary of the conversation to date:\n\nsummary"
assert new_messages[3].id == "human-2"
@pytest.mark.unit
def test_before_model_uses_keep_cutoff_for_message_trigger(monkeypatch: pytest.MonkeyPatch) -> None:
middleware = SummaryOffloadMiddleware(
model=_DummyModel(),
trigger=("messages", 3),
keep=("messages", 1),
token_counter=lambda _messages: 1,
summary_offload_threshold=10_000,
)
monkeypatch.setattr(summary_middleware, "_offload_tool_results", lambda *args, **kwargs: ({}, []))
monkeypatch.setattr(middleware, "_create_summary", lambda _messages: "summary")
messages = [
HumanMessage(content="human-1", id="human-1"),
AIMessage(content="ai-1", id="ai-1"),
HumanMessage(content="human-2", id="human-2"),
]
result = middleware.before_model({"messages": messages}, SimpleNamespace())
assert result is not None
new_messages = result["messages"]
assert new_messages[1].content == "Here is a summary of the conversation to date:\n\nsummary"
assert new_messages[2].id == "human-2"

View File

@ -50,6 +50,7 @@
- 修复 DOCX 解析中的图片回插顺序Docling 导出的多个 `<!-- image -->` 占位符现在按文档图片顺序替换,避免多图文档中的图片链接前后颠倒。 - 修复 DOCX 解析中的图片回插顺序Docling 导出的多个 `<!-- image -->` 占位符现在按文档图片顺序替换,避免多图文档中的图片链接前后颠倒。
- 修复前端依赖安全告警:通过 `pnpm.overrides` 将传递依赖 `flatted` 锁定到 `3.4.2`、`lodash-es` 锁定到 `4.18.1`,并同步更新 `pnpm-lock.yaml` 以消除 DriftGuard 报告的高危 CVE - 修复前端依赖安全告警:通过 `pnpm.overrides` 将传递依赖 `flatted` 锁定到 `3.4.2`、`lodash-es` 锁定到 `4.18.1`,并同步更新 `pnpm-lock.yaml` 以消除 DriftGuard 报告的高危 CVE
- 重写界面设计规范:参考 `DESIGN.md` 写法补充视觉气质、颜色 token、组件状态、布局层级、响应式与 Agent Prompt Guide并基于该规范收敛首页视觉表现移除装饰性渐变、重阴影、hover 位移和入场动画。 - 重写界面设计规范:参考 `DESIGN.md` 写法补充视觉气质、颜色 token、组件状态、布局层级、响应式与 Agent Prompt Guide并基于该规范收敛首页视觉表现移除装饰性渐变、重阴影、hover 位移和入场动画。
- 修复对话摘要中间件的工具结果卸载链路:摘要触发时改为将大体积 `ToolMessage` 写入当前 agent 可见的 sandbox outputs 路径,修正 `summary_offload` 路径拼接错误、`messages` 触发条件下不会真正裁剪历史的问题,并避免将 system message 重复纳入摘要与最终消息列表;补充对应单元测试覆盖。
--- ---