deerflow-code/offline-backend-20260512/backend/tests/test_writing_mode_report_middleware.py
2026-09-07 18:24:55 +08:00

708 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""写作模式:仅在开始写 md 时改走 deep_research_report,且不强制 tool_choice。"""
from __future__ import annotations
from types import SimpleNamespace
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from deerflow.agents.middlewares.writing_mode_report_middleware import (
TOOL_NAME,
WritingModeReportMiddleware,
extract_report_focus,
extract_report_topic,
is_handwritten_report_call,
is_report_reread_call,
pipeline_already_succeeded,
rewrite_handwritten_report_calls,
)
from deerflow.tools.builtins.deep_research_report_tool import is_writing_mode_report_path
def test_only_report_md_is_writing_mode_follow_up() -> None:
assert is_writing_mode_report_path("/mnt/user-data/outputs/report.md") is True
assert is_writing_mode_report_path("report.md") is True
assert is_writing_mode_report_path("/mnt/user-data/outputs/notes.md") is False
assert is_writing_mode_report_path("/mnt/user-data/outputs/analysis.markdown") is False
assert is_writing_mode_report_path("") is False
assert is_writing_mode_report_path(None) is False
def test_handwritten_report_detects_outputs_markdown() -> None:
assert is_handwritten_report_call({"name": "write_file", "args": {"path": "/mnt/user-data/outputs/report.md"}, "id": "c1"})
assert is_handwritten_report_call({"name": "write_file", "args": {"path": "/mnt/user-data/outputs/analysis.md"}, "id": "c1"})
assert is_handwritten_report_call({"name": "str_replace", "args": {"path": "report.md"}, "id": "c1"})
assert not is_handwritten_report_call({"name": "write_file", "args": {"path": "/mnt/user-data/workspace/notes.md"}, "id": "c1"})
assert not is_handwritten_report_call({"name": "web_search", "args": {"query": "课题"}, "id": "c1"})
def test_extract_topic_uses_latest_visible_user_text() -> None:
topic = extract_report_topic(
[
HumanMessage(content="旧课题"),
HumanMessage(content="hidden", additional_kwargs={"hide_from_ui": True}),
HumanMessage(content="写一份新能源汽车产业报告,重点看产能"),
]
)
assert topic.startswith("写一份新能源汽车产业报告")
def test_rewrite_write_file_to_pipeline() -> None:
last = AIMessage(
content="",
tool_calls=[
{
"id": "c1",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/report.md", "content": "# 草稿"},
"type": "tool_call",
},
{
"id": "c2",
"name": "present_files",
"args": {"path": "/mnt/user-data/outputs/report.md"},
"type": "tool_call",
},
],
)
patched = rewrite_handwritten_report_calls(last, topic="新能源汽车")
assert patched is not None
assert len(patched.tool_calls) == 1
assert patched.tool_calls[0]["name"] == TOOL_NAME
assert patched.tool_calls[0]["args"] == {"topic": "新能源汽车"}
assert patched.tool_calls[0]["id"] == "c1"
def test_rewrite_keeps_existing_pipeline_and_drops_handwritten() -> None:
last = AIMessage(
content="",
tool_calls=[
{
"id": "c1",
"name": TOOL_NAME,
"args": {"topic": "课题"},
"type": "tool_call",
},
{
"id": "c2",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/report.md", "content": "x"},
"type": "tool_call",
},
],
)
patched = rewrite_handwritten_report_calls(last, topic="课题")
assert patched is not None
assert [call["name"] for call in patched.tool_calls] == [TOOL_NAME]
def test_rewrite_leaves_chat_and_search_alone() -> None:
chat = AIMessage(content="先确认一下研究范围。")
assert rewrite_handwritten_report_calls(chat, topic="课题") is None
search = AIMessage(
content="",
tool_calls=[
{"id": "s1", "name": "web_search", "args": {"query": "课题"}, "type": "tool_call"},
],
)
assert rewrite_handwritten_report_calls(search, topic="课题") is None
pipeline = AIMessage(
content="",
tool_calls=[
{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"},
],
)
assert rewrite_handwritten_report_calls(pipeline, topic="课题") is None
def test_middleware_allows_chat_before_writing() -> None:
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="先聊聊这个课题"),
AIMessage(content="好的,我们先明确范围。"),
]
}
assert middleware.after_model(state, runtime=SimpleNamespace()) is None
def test_middleware_rewrites_write_file_on_after_model() -> None:
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写一份新能源汽车报告"),
AIMessage(
content="",
tool_calls=[
{
"id": "c1",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/report.md", "content": "# 报告"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
patched = update["messages"][0]
assert patched.tool_calls[0]["name"] == TOOL_NAME
assert patched.tool_calls[0]["args"]["topic"].startswith("写一份新能源汽车报告")
def test_wrap_tool_call_rejects_handwritten_report() -> None:
middleware = WritingModeReportMiddleware()
request = SimpleNamespace(
tool_call={
"id": "c1",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/report.md", "content": "x"},
}
)
called = {"handler": False}
def handler(_request):
called["handler"] = True
return ToolMessage(content="wrote", tool_call_id="c1", name="write_file")
result = middleware.wrap_tool_call(request, handler)
assert called["handler"] is False
assert isinstance(result, ToolMessage)
assert result.content.startswith("Error:")
assert TOOL_NAME in result.content
def test_pipeline_success_is_detected() -> None:
messages = [
HumanMessage(content="写一份报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:已写入 report.md。", tool_call_id="c1", name=TOOL_NAME),
]
assert pipeline_already_succeeded(messages) is True
assert pipeline_already_succeeded(messages[:-1]) is False
failed = [
*messages[:-1],
ToolMessage(content="Error: 深度研究报告生成失败。", tool_call_id="c1", name=TOOL_NAME),
]
assert pipeline_already_succeeded(failed) is False
def test_reread_detects_outputs_markdown() -> None:
assert is_report_reread_call({"name": "read_file", "args": {"path": "/mnt/user-data/outputs/report.md"}, "id": "c1"})
assert is_report_reread_call({"name": "read_file", "args": {"path": "/mnt/user-data/outputs/ai-agents-paper.md"}, "id": "c1"})
assert not is_report_reread_call({"name": "read_file", "args": {"path": "/mnt/user-data/workspace/notes.md"}, "id": "c1"})
assert not is_report_reread_call({"name": "write_file", "args": {"path": "/mnt/user-data/outputs/report.md"}, "id": "c1"})
def test_rewrite_after_success_drops_deliverable_tools() -> None:
last = AIMessage(
content="我再重写一版",
tool_calls=[
{
"id": "c2",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/ai-agents-paper.md", "content": "# 新稿"},
"type": "tool_call",
},
{
"id": "c3",
"name": TOOL_NAME,
"args": {"topic": "课题"},
"type": "tool_call",
},
],
)
patched = rewrite_handwritten_report_calls(last, topic="课题", pipeline_already_done=True)
assert patched is not None
assert patched.tool_calls == []
def test_rewrite_after_success_drops_reread() -> None:
last = AIMessage(
content="先看看报告",
tool_calls=[
{
"id": "c2",
"name": "read_file",
"args": {"path": "/mnt/user-data/outputs/report.md"},
"type": "tool_call",
}
],
)
patched = rewrite_handwritten_report_calls(last, topic="课题", pipeline_already_done=True)
assert patched is not None
assert patched.tool_calls == []
assert patched.content == ""
def test_middleware_nudges_summary_after_successful_pipeline() -> None:
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写一份新能源汽车报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:已写入 report.md。", tool_call_id="c1", name=TOOL_NAME),
AIMessage(
content="报告太泛化,我要重写",
tool_calls=[
{
"id": "c2",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/ai-agents-paper.md", "content": "# 重写"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
assert update["jump_to"] == "model"
assert update["messages"][0].tool_calls == []
assert update["messages"][0].content == ""
nudge = update["messages"][1]
assert nudge.additional_kwargs["hide_from_ui"] is True
assert nudge.additional_kwargs["writing_mode_summary_nudge"] is True
assert "2-3" in nudge.content
def test_middleware_ends_turn_after_ignored_summary_nudge() -> None:
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写一份新能源汽车报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:已写入 report.md。", tool_call_id="c1", name=TOOL_NAME),
HumanMessage(
content="请总结",
additional_kwargs={"hide_from_ui": True, "writing_mode_summary_nudge": True},
),
AIMessage(
content="",
tool_calls=[
{
"id": "c2",
"name": "read_file",
"args": {"path": "/mnt/user-data/outputs/report.md"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
assert update["jump_to"] == "end"
assert update["messages"][0].tool_calls == []
def test_wrap_tool_call_rejects_rerun_after_success() -> None:
middleware = WritingModeReportMiddleware()
prior = [
HumanMessage(content="写一份报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:已写入 report.md。", tool_call_id="c1", name=TOOL_NAME),
]
request = SimpleNamespace(
tool_call={"id": "c2", "name": TOOL_NAME, "args": {"topic": "课题"}},
state={"messages": prior},
)
called = {"handler": False}
def handler(_request):
called["handler"] = True
return ToolMessage(content="ran", tool_call_id="c2", name=TOOL_NAME)
result = middleware.wrap_tool_call(request, handler)
assert called["handler"] is False
assert isinstance(result, ToolMessage)
assert result.content.startswith("Error:")
assert "已经完成" in result.content
assert "请立即调用" not in result.content
def test_wrap_tool_call_rejects_reread_after_success() -> None:
middleware = WritingModeReportMiddleware()
prior = [
HumanMessage(content="写一份报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:已写入 report.md。", tool_call_id="c1", name=TOOL_NAME),
]
request = SimpleNamespace(
tool_call={
"id": "c2",
"name": "read_file",
"args": {"path": "/mnt/user-data/outputs/report.md"},
},
state={"messages": prior},
)
called = {"handler": False}
def handler(_request):
called["handler"] = True
return ToolMessage(content="read", tool_call_id="c2", name="read_file")
result = middleware.wrap_tool_call(request, handler)
assert called["handler"] is False
assert isinstance(result, ToolMessage)
assert result.content.startswith("Error:")
assert "已经完成" in result.content
def test_wrap_tool_call_allows_unrelated_writes() -> None:
middleware = WritingModeReportMiddleware()
request = SimpleNamespace(
tool_call={
"id": "c1",
"name": "write_file",
"args": {"path": "/mnt/user-data/workspace/notes.md", "content": "x"},
}
)
expected = ToolMessage(content="wrote", tool_call_id="c1", name="write_file")
result = middleware.wrap_tool_call(request, lambda _request: expected)
assert result is expected
def test_extract_topic_skips_clarification_answers() -> None:
"""澄清回答不是课题:向上取真正的写作请求原文。"""
topic = extract_report_topic(
[
HumanMessage(content="写人工智能领域的论文md"),
AIMessage(
content="",
tool_calls=[
{
"id": "q1",
"name": "ask_clarification",
"args": {"question": "方向?"},
"type": "tool_call",
}
],
),
ToolMessage(content="需要澄清:方向?", tool_call_id="q1", name="ask_clarification"),
HumanMessage(content="方向:AI Agent;用途:学术期刊论文;篇幅:短篇2000-4000字"),
AIMessage(
content="",
tool_calls=[
{
"id": "c1",
"name": "write_file",
"args": {"path": "/mnt/user-data/outputs/report.md"},
"type": "tool_call",
}
],
),
]
)
assert topic.startswith("写人工智能领域的论文md")
def test_bash_heredoc_write_is_redirected_to_pipeline() -> None:
"""模型被剥离不了工具时会绕道 bash heredoc 手写报告——同样改写为管线。"""
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写人工智能领域的论文md"),
AIMessage(
content="",
tool_calls=[
{
"id": "b1",
"name": "bash",
"args": {"command": "cat > /mnt/user-data/outputs/report.md << 'ENDOFFILE'\n# 报告全文……\nENDOFFILE"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
patched = update["messages"][0]
assert patched.tool_calls[0]["name"] == TOOL_NAME
assert patched.tool_calls[0]["args"]["topic"].startswith("写人工智能领域的论文md")
def test_bash_python_open_write_is_redirected_to_pipeline() -> None:
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写一份新能源汽车报告"),
AIMessage(
content="",
tool_calls=[
{
"id": "b2",
"name": "bash",
"args": {"command": ("python3 << 'PYEOF'\ncontent = r'''# 报告'''\nwith open('/mnt/user-data/outputs/report.md', 'w', encoding='utf-8') as f:\n f.write(content)\nPYEOF")},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
patched = update["messages"][0]
assert patched.tool_calls[0]["name"] == TOOL_NAME
def test_bash_readonly_and_non_outputs_commands_are_left_alone() -> None:
"""mkdir / wc / cat 读取、非 outputs 的 .md 写入都不是手写报告。"""
middleware = WritingModeReportMiddleware()
for command in (
"mkdir -p /mnt/user-data/outputs",
"wc -m /mnt/user-data/outputs/report.md",
"cat /mnt/user-data/outputs/report.md",
"echo hi > /tmp/scratch.md",
):
state = {
"messages": [
HumanMessage(content="写一份报告"),
AIMessage(
content="",
tool_calls=[{"id": "b3", "name": "bash", "args": {"command": command}, "type": "tool_call"}],
),
]
}
assert middleware.after_model(state, runtime=SimpleNamespace()) is None, command
def test_bash_write_after_success_is_dropped() -> None:
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写一份报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:基于 3 条资料…", tool_call_id="c1", name=TOOL_NAME),
AIMessage(
content="",
tool_calls=[
{
"id": "b4",
"name": "bash",
"args": {"command": "echo '# 新版' > /mnt/user-data/outputs/report.md"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
patched = update["messages"][0]
assert patched.tool_calls == []
nudged = update["messages"][1]
assert getattr(nudged, "additional_kwargs", {}).get("writing_mode_summary_nudge") is True
def test_wrap_tool_call_rejects_bash_report_write() -> None:
"""漏网到执行层的 bash 手写(如裸 report.md 重定向)也要拒绝并引导走管线。"""
middleware = WritingModeReportMiddleware()
request = SimpleNamespace(
tool_call={
"id": "b5",
"name": "bash",
"args": {"command": "cat > report.md << 'EOF'\n# x\nEOF"},
}
)
called = {"handler": False}
def handler(_request):
called["handler"] = True
return ToolMessage(content="ran", tool_call_id="b5", name="bash")
result = middleware.wrap_tool_call(request, handler)
assert called["handler"] is False
assert isinstance(result, ToolMessage)
assert result.content.startswith("Error:")
assert TOOL_NAME in result.content
def test_wrap_model_call_strips_handwrite_tools_from_binding() -> None:
"""首报轮从模型绑定中剥离 write_file/str_replace/bash——bash 是唯一能写
文件的执行类工具(实测模型会把全文起草进 heredoc 参数流式几十秒后才被
after_model 拦截),剥离后模型没有任何可发起文件写入的工具。"""
middleware = WritingModeReportMiddleware()
def make_request(tools):
return SimpleNamespace(
tools=tools,
override=lambda **kwargs: SimpleNamespace(tools=kwargs.get("tools", [])),
)
seen: dict[str, list] = {}
def handler(request):
seen["tools"] = list(request.tools)
return SimpleNamespace()
middleware.wrap_model_call(
make_request(
[
SimpleNamespace(name="write_file"),
SimpleNamespace(name="web_search"),
SimpleNamespace(name="str_replace"),
SimpleNamespace(name="bash"),
SimpleNamespace(name=TOOL_NAME),
]
),
handler,
)
assert [getattr(t, "name", "") for t in seen["tools"]] == ["web_search", TOOL_NAME]
def test_wrap_model_call_passes_through_when_no_handwrite_tools() -> None:
middleware = WritingModeReportMiddleware()
request = SimpleNamespace(
tools=[SimpleNamespace(name="web_search"), SimpleNamespace(name=TOOL_NAME)],
override=lambda **kwargs: SimpleNamespace(tools=kwargs.get("tools", [])),
)
def handler(passed):
assert passed is request
return "unchanged"
assert middleware.wrap_model_call(request, handler) == "unchanged"
def test_extract_focus_joins_clarification_answers() -> None:
focus = extract_report_focus(
[
HumanMessage(content="写人工智能领域的论文md"),
AIMessage(
content="",
tool_calls=[{"id": "q1", "name": "ask_clarification", "args": {"question": "方向?"}, "type": "tool_call"}],
),
ToolMessage(content="需要澄清", tool_call_id="q1", name="ask_clarification"),
HumanMessage(content="主题:LLM的能力边界;类型:观点评论型"),
],
topic="写人工智能领域的论文md",
)
assert "LLM的能力边界" in focus
assert "观点评论型" in focus
assert "写人工智能领域的论文md" not in focus
def test_rewrite_carries_focus_from_clarification_answers() -> None:
"""改写进管线的调用要带上澄清回答(方向/类型/篇幅),否则管线只拿到
原始请求一句泛写。"""
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写人工智能领域的论文md,有不懂的可以问我"),
AIMessage(
content="",
tool_calls=[{"id": "q1", "name": "ask_clarification", "args": {"question": "主题?"}, "type": "tool_call"}],
),
ToolMessage(content="需要澄清", tool_call_id="q1", name="ask_clarification"),
HumanMessage(content="主题:LLM的能力边界;篇幅:3000-5000字"),
AIMessage(
content="",
tool_calls=[
{
"id": "c1",
"name": "bash",
"args": {"command": "cat > /mnt/user-data/outputs/report.md << 'EOF'\n# x\nEOF"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
patched = update["messages"][0]
assert patched.tool_calls[0]["name"] == TOOL_NAME
assert patched.tool_calls[0]["args"]["topic"].startswith("写人工智能领域的论文md")
assert "LLM的能力边界" in patched.tool_calls[0]["args"]["focus"]
def test_second_offense_after_nudge_gets_fallback_summary() -> None:
"""模型无视总结提示仍要继续:丢弃调用后必须留一段兜底总结,不能静默结束。"""
middleware = WritingModeReportMiddleware()
state = {
"messages": [
HumanMessage(content="写一份新能源汽车报告"),
AIMessage(
content="",
tool_calls=[{"id": "c1", "name": TOOL_NAME, "args": {"topic": "课题"}, "type": "tool_call"}],
),
ToolMessage(content="研究报告已完成:已写入 report.md。", tool_call_id="c1", name=TOOL_NAME),
HumanMessage(
content="请总结",
additional_kwargs={"hide_from_ui": True, "writing_mode_summary_nudge": True},
),
AIMessage(
content="",
tool_calls=[
{
"id": "c2",
"name": "read_file",
"args": {"path": "/mnt/user-data/outputs/report.md"},
"type": "tool_call",
}
],
),
]
}
update = middleware.after_model(state, runtime=SimpleNamespace())
assert update is not None
assert update["jump_to"] == "end"
patched = update["messages"][0]
assert patched.tool_calls == []
assert "report.md" in patched.content
assert "已完成" in patched.content