fix(summary): 修复对话摘要中间件的工具结果卸载链路,修正路径拼接错误并补充单元测试
This commit is contained in:
parent
63786b7c66
commit
9c8998699c
@ -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)
|
||||||
|
|||||||
138
backend/test/unit/middlewares/test_summary_middleware.py
Normal file
138
backend/test/unit/middlewares/test_summary_middleware.py
Normal 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"
|
||||||
@ -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 重复纳入摘要与最终消息列表;补充对应单元测试覆盖。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user