708 lines
25 KiB
Python
708 lines
25 KiB
Python
"""写作模式:仅在开始写 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
|