From 9c8998699cc69c60972ee78eb0430488fc911ab4 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 12 Apr 2026 23:38:13 +0800 Subject: [PATCH] =?UTF-8?q?fix(summary):=20=E4=BF=AE=E5=A4=8D=E5=AF=B9?= =?UTF-8?q?=E8=AF=9D=E6=91=98=E8=A6=81=E4=B8=AD=E9=97=B4=E4=BB=B6=E7=9A=84?= =?UTF-8?q?=E5=B7=A5=E5=85=B7=E7=BB=93=E6=9E=9C=E5=8D=B8=E8=BD=BD=E9=93=BE?= =?UTF-8?q?=E8=B7=AF=EF=BC=8C=E4=BF=AE=E6=AD=A3=E8=B7=AF=E5=BE=84=E6=8B=BC?= =?UTF-8?q?=E6=8E=A5=E9=94=99=E8=AF=AF=E5=B9=B6=E8=A1=A5=E5=85=85=E5=8D=95?= =?UTF-8?q?=E5=85=83=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../agents/middlewares/summary_middleware.py | 228 +++++++++++++----- .../middlewares/test_summary_middleware.py | 138 +++++++++++ docs/develop-guides/roadmap.md | 1 + 3 files changed, 303 insertions(+), 64 deletions(-) create mode 100644 backend/test/unit/middlewares/test_summary_middleware.py diff --git a/backend/package/yuxi/agents/middlewares/summary_middleware.py b/backend/package/yuxi/agents/middlewares/summary_middleware.py index 1fad4936..64b5ae1a 100644 --- a/backend/package/yuxi/agents/middlewares/summary_middleware.py +++ b/backend/package/yuxi/agents/middlewares/summary_middleware.py @@ -10,7 +10,6 @@ from __future__ import annotations import uuid from collections.abc import Callable, Iterable, Mapping from functools import partial -from pathlib import Path from typing import Any, Literal, cast, override 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.runtime import Runtime -from yuxi.utils.paths import OUTPUTS_DIR_NAME +from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS TokenCounter = Callable[[Iterable[MessageLikeRepresentation]], int] @@ -78,7 +77,7 @@ Messages to summarize: _DEFAULT_MESSAGES_TO_KEEP = 20 _DEFAULT_FALLBACK_MESSAGE_COUNT = 15 -_OFFLOAD_DIR = "/summary_offload" # 虚拟文件系统路径 +_OFFLOAD_DIR = "summary_offload" ContextFraction = tuple[Literal["fraction"], float] 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: - """Convert ToolMessage content to string for size checking.""" + """Convert plain-text ToolMessage content to string for size checking.""" if isinstance(content, str): return content if isinstance(content, list): if len(content) == 1 and isinstance(content[0], dict) and content[0].get("type") == "text": return str(content[0].get("text", "")) - return str(content) - return str(content) + return None + return None 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: @@ -141,10 +175,7 @@ def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenC tool_name = msg.name or "unknown" tool_call_id = msg.tool_call_id or "" - # 生成文件路径 (工具名称-xxx) - 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() + file_path = _build_offload_file_path(msg) # 构建文件头部信息 header_lines = [ @@ -156,17 +187,9 @@ def _offload_tool_result(msg: ToolMessage, threshold: int, token_counter: TokenC ] header = "\n".join(header_lines) - # 保存到 files 格式(包含头部信息) - from datetime import datetime - - timestamp = datetime.now().isoformat() - files_update = { - file_path: { - "content": [header + content_str], - "created_at": timestamp, - "modified_at": timestamp, - } - } + written, files_update = _write_offloaded_content(runtime, file_path, header + content_str) + if not written: + return None # 创建预览内容 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) - return files_update + return files_update or {} 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]]: """扫描消息列表,卸载所有超阈值的工具结果. @@ -198,14 +221,57 @@ def _offload_tool_results( if not isinstance(msg, ToolMessage): continue - result = _offload_tool_result(msg, threshold, token_counter) - if result: + result = _offload_tool_result(msg, threshold, token_counter, runtime) + if result is not None: files_update.update(result) modified_messages.append(msg) 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): """总结+工具结果卸载中间件. @@ -322,7 +388,9 @@ class SummaryOffloadMiddleware(AgentMiddleware): files_update: dict[str, Any] = {} 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 modified_messages = agg_msgs @@ -330,40 +398,55 @@ class SummaryOffloadMiddleware(AgentMiddleware): current_tokens = self.token_counter(messages) trigger_value = self._get_token_trigger_value() - retention_limit = float("inf") - if trigger_value: + if trigger_value is None: + 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 - if current_tokens <= retention_limit: - if files_update: - return {"files": files_update, "messages": modified_messages} - return None + if current_tokens <= retention_limit: + if files_update: + result: dict[str, Any] = {"messages": modified_messages} + if files_update: + result["files"] = files_update + return result + return None - # 4. 超过 limit,需要 Eviction (Summary) - system_msg_count = 0 - messages_to_process = messages + # 4. 超过 limit,需要 Eviction (Summary) + system_msg_count = 0 + messages_to_process = messages - if messages and messages[0].type == "system": - system_msg_count = 1 - messages_to_process = messages[1:] + if messages and messages[0].type == "system": + system_msg_count = 1 + messages_to_process = messages[1:] - cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit)) - cutoff_index = system_msg_count + cutoff_relative + cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit)) + cutoff_index = system_msg_count + cutoff_relative if cutoff_index <= system_msg_count: 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 - 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) new_messages = self._build_new_messages(summary) # 如果有 System Message,需要保留在最前面 final_messages = [] - if system_msg_count > 0: - final_messages.append(messages[0]) + if system_message is not None: + final_messages.append(system_message) final_messages.extend(new_messages) final_messages.extend(preserved_messages) @@ -393,7 +476,9 @@ class SummaryOffloadMiddleware(AgentMiddleware): files_update: dict[str, Any] = {} 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 modified_messages = agg_msgs @@ -401,40 +486,55 @@ class SummaryOffloadMiddleware(AgentMiddleware): current_tokens = self.token_counter(messages) trigger_value = self._get_token_trigger_value() - retention_limit = float("inf") - if trigger_value: + if trigger_value is None: + 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 - if current_tokens <= retention_limit: - if files_update: - return {"files": files_update, "messages": modified_messages} - return None + if current_tokens <= retention_limit: + if files_update: + result: dict[str, Any] = {"messages": modified_messages} + if files_update: + result["files"] = files_update + return result + return None - # 4. 超过 limit,需要 Eviction (Summary) - system_msg_count = 0 - messages_to_process = messages + # 4. 超过 limit,需要 Eviction (Summary) + system_msg_count = 0 + messages_to_process = messages - if messages and messages[0].type == "system": - system_msg_count = 1 - messages_to_process = messages[1:] + if messages and messages[0].type == "system": + system_msg_count = 1 + messages_to_process = messages[1:] - cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit)) - cutoff_index = system_msg_count + cutoff_relative + cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit)) + cutoff_index = system_msg_count + cutoff_relative if cutoff_index <= system_msg_count: 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 - 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) new_messages = self._build_new_messages(summary) final_messages = [] - if system_msg_count > 0: - final_messages.append(messages[0]) + if system_message is not None: + final_messages.append(system_message) final_messages.extend(new_messages) final_messages.extend(preserved_messages) diff --git a/backend/test/unit/middlewares/test_summary_middleware.py b/backend/test/unit/middlewares/test_summary_middleware.py new file mode 100644 index 00000000..935d393e --- /dev/null +++ b/backend/test/unit/middlewares/test_summary_middleware.py @@ -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" diff --git a/docs/develop-guides/roadmap.md b/docs/develop-guides/roadmap.md index ea8ef08f..a5d41adc 100644 --- a/docs/develop-guides/roadmap.md +++ b/docs/develop-guides/roadmap.md @@ -50,6 +50,7 @@ - 修复 DOCX 解析中的图片回插顺序:Docling 导出的多个 `` 占位符现在按文档图片顺序替换,避免多图文档中的图片链接前后颠倒。 - 修复前端依赖安全告警:通过 `pnpm.overrides` 将传递依赖 `flatted` 锁定到 `3.4.2`、`lodash-es` 锁定到 `4.18.1`,并同步更新 `pnpm-lock.yaml` 以消除 DriftGuard 报告的高危 CVE - 重写界面设计规范:参考 `DESIGN.md` 写法补充视觉气质、颜色 token、组件状态、布局层级、响应式与 Agent Prompt Guide,并基于该规范收敛首页视觉表现,移除装饰性渐变、重阴影、hover 位移和入场动画。 +- 修复对话摘要中间件的工具结果卸载链路:摘要触发时改为将大体积 `ToolMessage` 写入当前 agent 可见的 sandbox outputs 路径,修正 `summary_offload` 路径拼接错误、`messages` 触发条件下不会真正裁剪历史的问题,并避免将 system message 重复纳入摘要与最终消息列表;补充对应单元测试覆盖。 ---