ForcePilot/backend/test/unit/middlewares/test_summary_middleware.py

139 lines
4.8 KiB
Python

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"