"""写作模式:仅在开始写 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