650 lines
27 KiB
Python
650 lines
27 KiB
Python
"""写作模式内联深度研究报告工具的单元测试。
|
||
|
||
覆盖:默认配置组装(通用分析结构 / 继承聊天模型)、事件转发(custom 通道帧
|
||
格式/过滤/容错)、工具主体(收割优先——资料足够直接写 / 不足也直接开写 /
|
||
retrieve 参数被忽略、按当前主题的相关度强过滤、线程解析、失败路径、
|
||
artifacts 更新),以及 lead / agent-only 两条 prompt 模板的写作模式首轮合约。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import json
|
||
from types import SimpleNamespace
|
||
from typing import Any
|
||
|
||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||
|
||
from deerflow.agents.deep_research.collection.thread_harvester import StaticHarvestMaterialProvider
|
||
from deerflow.agents.deep_research.types import DeepResearchRequest, DeepResearchResult
|
||
from deerflow.persistence.report_structures.defaults import GENERAL_OUTLINE
|
||
from deerflow.tools.builtins import deep_research_report_tool as tool_module
|
||
from deerflow.tools.builtins.deep_research_report_tool import (
|
||
_KNOWLEDGE_WRITE_INSTRUCTION,
|
||
StreamWriterEventSink,
|
||
build_writing_research_config,
|
||
deep_research_report_tool,
|
||
)
|
||
|
||
_DEFAULT_TOPIC = "新能源汽车市场分析"
|
||
_OFF_TOPIC = "养老产业发展研究"
|
||
|
||
# ── 配置组装 ─────────────────────────────────────────────────────────────────
|
||
|
||
|
||
def test_writing_research_config_defaults_to_general_outline_basic() -> None:
|
||
cfg = build_writing_research_config()
|
||
|
||
assert cfg.mode == "basic"
|
||
assert cfg.custom_outline == GENERAL_OUTLINE
|
||
assert cfg.structure_mode == "adaptive"
|
||
assert cfg.generate_images is False
|
||
assert cfg.report_instruction == _KNOWLEDGE_WRITE_INSTRUCTION
|
||
assert "必须基于模型自身的领域知识" in cfg.report_instruction
|
||
|
||
|
||
def test_writing_research_config_maps_focus_to_report_instruction() -> None:
|
||
cfg = build_writing_research_config(focus=" 侧重对比分析,3000 字左右 ")
|
||
|
||
assert cfg.report_instruction is not None
|
||
assert cfg.report_instruction.startswith("侧重对比分析,3000 字左右")
|
||
assert _KNOWLEDGE_WRITE_INSTRUCTION in cfg.report_instruction
|
||
# 空白 focus 仍保留知识补全指令,不产生空行前缀
|
||
empty_focus = build_writing_research_config(focus=" ")
|
||
assert empty_focus.report_instruction == _KNOWLEDGE_WRITE_INSTRUCTION
|
||
|
||
|
||
# ── 事件转发 sink ────────────────────────────────────────────────────────────
|
||
|
||
|
||
def _capture_writer() -> tuple[list[dict], object]:
|
||
frames: list[dict] = []
|
||
|
||
def writer(frame: dict) -> None:
|
||
frames.append(frame)
|
||
|
||
return frames, writer
|
||
|
||
|
||
def test_event_sink_wraps_forwardable_events() -> None:
|
||
frames, writer = _capture_writer()
|
||
sink = StreamWriterEventSink(writer, session_id="chat-t1", job_id="call-1")
|
||
|
||
asyncio.run(sink.emit("phase_changed", phase="planning", payload={"to": "planning"}))
|
||
|
||
assert len(frames) == 1
|
||
frame = frames[0]
|
||
assert frame["type"] == "deep_research_report"
|
||
assert frame["event"]["type"] == "phase_changed"
|
||
assert frame["event"]["phase"] == "planning"
|
||
assert frame["event"]["payload"] == {"to": "planning"}
|
||
assert frame["event"]["jobId"] == "call-1"
|
||
|
||
|
||
def test_event_sink_drops_noise_events() -> None:
|
||
frames, writer = _capture_writer()
|
||
sink = StreamWriterEventSink(writer, session_id="s", job_id="j")
|
||
|
||
asyncio.run(sink.emit("model_usage", payload={"tokens": 1}))
|
||
asyncio.run(sink.emit("heartbeat"))
|
||
|
||
assert frames == []
|
||
|
||
|
||
def test_event_sink_forwards_live_report_deltas_only() -> None:
|
||
frames, writer = _capture_writer()
|
||
sink = StreamWriterEventSink(writer, session_id="s", job_id="j")
|
||
|
||
asyncio.run(sink.publish_live("report_delta", phase="writing", payload={"delta": "第一"}))
|
||
asyncio.run(sink.publish_live("summary_delta", phase="summarizing", payload={"delta": "x"}))
|
||
|
||
assert len(frames) == 1
|
||
assert frames[0]["event"]["type"] == "report_delta"
|
||
assert frames[0]["event"]["live"] is True
|
||
|
||
|
||
def test_event_sink_survives_writer_failure() -> None:
|
||
def broken_writer(frame: dict) -> None:
|
||
raise RuntimeError("stream closed")
|
||
|
||
sink = StreamWriterEventSink(broken_writer, session_id="s", job_id="j")
|
||
|
||
event = asyncio.run(sink.emit("report_completed", payload={"sourceCount": 3}))
|
||
|
||
assert event.type == "report_completed" # 事件本体仍然返回
|
||
asyncio.run(sink.publish_live("report_delta", payload={"delta": "x"})) # 不抛异常
|
||
|
||
|
||
# ── 工具主体 ─────────────────────────────────────────────────────────────────
|
||
|
||
|
||
def _material_message(index: int, topic: str = _DEFAULT_TOPIC) -> ToolMessage:
|
||
"""一条可被收割器识别、且与 ``topic`` 词面相关的工具结果消息。"""
|
||
return ToolMessage(
|
||
content=json.dumps(
|
||
{
|
||
"results": [
|
||
{
|
||
"title": f"{topic}:资料{index}",
|
||
"content": f"第 {index} 条关于{topic}的对话检索资料,包含要点分析。",
|
||
"url": f"https://example.com/{index}",
|
||
}
|
||
]
|
||
},
|
||
ensure_ascii=False,
|
||
),
|
||
name="weknora_search",
|
||
tool_call_id=f"call-m{index}",
|
||
)
|
||
|
||
|
||
def _fake_runtime(
|
||
thread_id: str = "thread-1",
|
||
messages: list | None = None,
|
||
) -> SimpleNamespace:
|
||
return SimpleNamespace(
|
||
context={"thread_id": thread_id},
|
||
config={"configurable": {"thread_id": thread_id}},
|
||
state={"messages": messages if messages is not None else [_material_message(i) for i in range(1, 4)]},
|
||
)
|
||
|
||
|
||
class _FakeEngine:
|
||
"""替代 DeepResearchEngine:记录请求与 adapters 并返回可控结果。"""
|
||
|
||
instances: list[_FakeEngine] = []
|
||
|
||
def __init__(self) -> None:
|
||
self.requests: list[DeepResearchRequest] = []
|
||
self.adapters: Any = None # noqa: RUF012 - 测试桩
|
||
self.result = DeepResearchResult(
|
||
report_markdown="# 报告\n正文",
|
||
source_ids=["s1", "s2", "s3"],
|
||
)
|
||
self.error: Exception | None = None
|
||
_FakeEngine.instances.append(self)
|
||
|
||
async def run(self, request, *, adapters, cancellation=None) -> DeepResearchResult:
|
||
self.requests.append(request)
|
||
self.adapters = adapters
|
||
if self.error is not None:
|
||
raise self.error
|
||
return self.result
|
||
|
||
|
||
def _patch_tool_dependencies(monkeypatch, *, user_id: str = "u1", engine=None) -> list[dict]:
|
||
frames, writer = _capture_writer()
|
||
monkeypatch.setattr(tool_module, "DeepResearchEngine", lambda: engine or _FakeEngine())
|
||
monkeypatch.setattr(tool_module, "resolve_path_user_id", lambda thread_id: user_id)
|
||
monkeypatch.setattr(tool_module, "_safe_stream_writer", lambda: writer)
|
||
monkeypatch.setattr(
|
||
tool_module,
|
||
"DeerFlowCompletionBackend",
|
||
lambda *args, **kwargs: SimpleNamespace(usage_totals={}),
|
||
)
|
||
monkeypatch.setattr(
|
||
tool_module,
|
||
"DeerFlowContextCompressor",
|
||
lambda **kwargs: SimpleNamespace(captured=kwargs),
|
||
)
|
||
monkeypatch.setattr(tool_module, "ThreadOutputsArtifactWriter", lambda *args, **kwargs: SimpleNamespace(captured=(args, kwargs)))
|
||
monkeypatch.setattr(tool_module, "build_image_provider", lambda: None)
|
||
return frames
|
||
|
||
|
||
def test_tool_writes_with_harvested_materials_when_sufficient(monkeypatch) -> None:
|
||
frames = _patch_tool_dependencies(monkeypatch)
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-9", messages=[_material_message(i) for i in range(1, 5)]),
|
||
"call-1",
|
||
_DEFAULT_TOPIC,
|
||
focus="侧重销量对比",
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert len(engine.requests) == 1
|
||
request = engine.requests[0]
|
||
assert request.query == "新能源汽车市场分析"
|
||
assert request.config.custom_outline == GENERAL_OUTLINE
|
||
assert request.config.report_instruction is not None
|
||
assert request.config.report_instruction.startswith("侧重销量对比")
|
||
assert _KNOWLEDGE_WRITE_INSTRUCTION in request.config.report_instruction
|
||
assert request.session_id == "chat-thread-9"
|
||
|
||
# 资料足够 → 只用收割池,不触发自检索
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
assert engine.adapters.materials.material_count == 4
|
||
|
||
update = command.update
|
||
assert update["artifacts"] == ["/mnt/user-data/outputs/report.md"]
|
||
tool_message = update["messages"][0]
|
||
assert tool_message.tool_call_id == "call-1"
|
||
assert "对话资料" in tool_message.content
|
||
assert "3" in tool_message.content # 来源条数写入总结
|
||
|
||
# 终态帧已发出(success)
|
||
assert any(f.get("event", {}).get("type") == "tool_finished" and f["event"]["payload"]["status"] == "success" for f in frames)
|
||
|
||
|
||
def test_tool_writes_immediately_when_harvested_materials_insufficient(monkeypatch) -> None:
|
||
"""资料不足也不再询问用户,直接用现有对话资料开写。"""
|
||
frames = _patch_tool_dependencies(monkeypatch)
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-a", messages=[_material_message(1, topic="课题")]),
|
||
"call-2",
|
||
"课题",
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
assert engine.adapters.materials.material_count == 1
|
||
|
||
update = command.update
|
||
assert update["artifacts"] == ["/mnt/user-data/outputs/report.md"]
|
||
tool_message = update["messages"][0]
|
||
assert tool_message.tool_call_id == "call-2"
|
||
assert "对话资料" in tool_message.content
|
||
assert "ask_clarification" not in tool_message.content
|
||
assert any(f.get("event", {}).get("type") == "tool_finished" and f["event"]["payload"]["status"] == "success" for f in frames)
|
||
|
||
|
||
def test_tool_retrieve_true_still_writes_from_conversation_only(monkeypatch) -> None:
|
||
"""写作管线忽略 retrieve=true:检索留给对话工具,这里只写文件。"""
|
||
_patch_tool_dependencies(monkeypatch)
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-b", messages=[_material_message(1, topic="课题")]),
|
||
"call-3",
|
||
"课题",
|
||
retrieve=True,
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
assert engine.adapters.materials.material_count == 1
|
||
|
||
update = command.update
|
||
assert update["artifacts"] == ["/mnt/user-data/outputs/report.md"]
|
||
assert "对话资料" in update["messages"][0].content
|
||
assert "补充检索" not in update["messages"][0].content
|
||
|
||
|
||
def test_tool_drops_off_topic_materials_from_pool_and_threshold(monkeypatch) -> None:
|
||
"""课题切换:旧课题检索结果既不计入门槛,也不进报告资料池。"""
|
||
_patch_tool_dependencies(monkeypatch)
|
||
messages = [
|
||
*[_material_message(i) for i in range(1, 6)], # 旧课题(新能源)5 条
|
||
*[_material_message(i, topic=_OFF_TOPIC) for i in range(11, 14)], # 当前课题 3 条
|
||
]
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-t", messages=messages),
|
||
"call-t",
|
||
_OFF_TOPIC,
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
# 只有当前课题的 3 条入池;旧课题 5 条被强过滤剔除
|
||
assert engine.adapters.materials.material_count == 3
|
||
assert command.update["artifacts"] == ["/mnt/user-data/outputs/report.md"]
|
||
|
||
|
||
def test_tool_writes_immediately_when_all_materials_off_topic(monkeypatch) -> None:
|
||
"""收割池非空但全部偏题 → 相关数 0,仍直接开写,不再问用户。"""
|
||
_patch_tool_dependencies(monkeypatch)
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-o", messages=[_material_message(i) for i in range(1, 6)]),
|
||
"call-o",
|
||
_OFF_TOPIC,
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
assert engine.adapters.materials.material_count == 0
|
||
assert command.update["artifacts"] == ["/mnt/user-data/outputs/report.md"]
|
||
assert "ask_clarification" not in command.update["messages"][0].content
|
||
|
||
|
||
def test_relevance_filter_rules() -> None:
|
||
"""词面相关度判定的最小规则集(强词 / 标题 / 正文 / 短主题兜底)。"""
|
||
from deerflow.agents.deep_research.types import ResearchMaterial
|
||
|
||
strong, weak = tool_module._topic_terms(_DEFAULT_TOPIC)
|
||
assert _DEFAULT_TOPIC in strong # 整段短语是强词
|
||
assert "新能" in weak and "分析" in weak
|
||
|
||
def _mat(title: str, content: str) -> ResearchMaterial:
|
||
return ResearchMaterial(
|
||
id=title,
|
||
query="q",
|
||
title=title,
|
||
url=None,
|
||
raw_content=content,
|
||
snippet=None,
|
||
source="t",
|
||
source_type="other",
|
||
content_hash=title,
|
||
)
|
||
|
||
# 强词命中标题
|
||
assert tool_module._is_topic_relevant(_mat(f"{_DEFAULT_TOPIC}:综述", "无关键词"), strong, weak) is True
|
||
# 标题两个不同 bigram
|
||
assert tool_module._is_topic_relevant(_mat("新能源汽车销量榜", "……"), strong, weak) is True
|
||
# 完全偏题
|
||
assert tool_module._is_topic_relevant(_mat("养老产业白皮书", "养老服务体系建设"), strong, weak) is False
|
||
# 正文密集命中(≥4 个不同 bigram)
|
||
assert tool_module._is_topic_relevant(_mat("行业观察", "新能源汽车市场分析"), strong, weak) is True
|
||
# 短主题兜底:唯一词命中即相关
|
||
s2, w2 = tool_module._topic_terms("课题")
|
||
assert tool_module._is_topic_relevant(_mat("课题:资料1", "关于课题的要点"), s2, w2) is True
|
||
|
||
|
||
# ── 课题锚点:请求之后的对话检索整批保留(语言/词面无关) ────────────────────
|
||
|
||
|
||
def _search_envelope_message(index: int, *, query: str, titles: list[str]) -> ToolMessage:
|
||
"""一条 web_search 工具结果信封(收割器走最高保真映射分支)。"""
|
||
return ToolMessage(
|
||
content=json.dumps(
|
||
{
|
||
"query": query,
|
||
"results": [
|
||
{
|
||
"title": title,
|
||
"content": f"{title} — key findings and technical details.",
|
||
"url": f"https://example.com/{index}-{i}",
|
||
}
|
||
for i, title in enumerate(titles)
|
||
],
|
||
},
|
||
ensure_ascii=False,
|
||
),
|
||
name="web_search",
|
||
tool_call_id=f"call-s{index}",
|
||
)
|
||
|
||
|
||
def test_topic_anchor_index_rules() -> None:
|
||
"""锚点 = 最新的非澄清回答人类消息;澄清回答不充当新课题。"""
|
||
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="方向:AIGC;篇幅:3000字"),
|
||
HumanMessage(content="写人工智能领域的论文md"),
|
||
]
|
||
# 文本前缀命中课题原文
|
||
assert tool_module.topic_anchor_index(messages, topic="写人工智能领域的论文md") == 4
|
||
# 无 topic 匹配时回落「最新的非回答人类消息」——回答被跳过
|
||
messages[-1] = HumanMessage(content="补充:要英文文献")
|
||
assert tool_module.topic_anchor_index(messages) == 4
|
||
assert tool_module.topic_anchor_index([]) == -1
|
||
|
||
|
||
def test_tool_keeps_conversation_searches_after_topic_request(monkeypatch) -> None:
|
||
"""课题请求之后的对话检索结果整批入池——中文课题 + 英文检索结果不再被
|
||
纯词面匹配误杀(线上 35 条只留 4 条的根因)。"""
|
||
_patch_tool_dependencies(monkeypatch)
|
||
messages = [
|
||
HumanMessage(content="写一篇人工智能领域的论文md"),
|
||
_search_envelope_message(
|
||
1,
|
||
query="AIGC generative AI survey 2024",
|
||
titles=[
|
||
"Sora: A Review of Text-to-Video Generation Models",
|
||
"Diffusion Transformers for Image Synthesis",
|
||
"Challenges and Ethics of AI-Generated Content",
|
||
],
|
||
),
|
||
_search_envelope_message(
|
||
2,
|
||
query="LLM alignment RLHF DPO",
|
||
titles=["Direct Preference Optimization: A Survey"],
|
||
),
|
||
]
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-keep", messages=messages),
|
||
"call-keep",
|
||
"写一篇人工智能领域的论文md",
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
# 4 条英文资料全部进入写作池(锚点=课题请求,之后的检索一律保留)
|
||
assert engine.adapters.materials.material_count == 4
|
||
assert command.update["artifacts"] == ["/mnt/user-data/outputs/report.md"]
|
||
|
||
|
||
def test_tool_still_drops_old_topic_searches_before_new_request(monkeypatch) -> None:
|
||
"""课题切换:新请求之前的旧课题检索仍被词面强过滤,不凑门槛不入报告。"""
|
||
_patch_tool_dependencies(monkeypatch)
|
||
messages = [
|
||
HumanMessage(content="写养老产业发展研究报告"),
|
||
*[_material_message(i, topic=_OFF_TOPIC) for i in range(1, 4)], # 旧课题 3 条
|
||
HumanMessage(content="改成写新能源汽车市场分析"),
|
||
*[_material_message(i, topic=_DEFAULT_TOPIC) for i in range(11, 14)], # 当前课题 3 条
|
||
]
|
||
|
||
asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-switch", messages=messages),
|
||
"call-switch",
|
||
_DEFAULT_TOPIC,
|
||
)
|
||
)
|
||
|
||
engine = _FakeEngine.instances[-1]
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
# 旧课题 3 条(锚点之前、词面不命中)被剔除;当前课题 3 条保留
|
||
assert engine.adapters.materials.material_count == 3
|
||
|
||
|
||
def test_tool_retrieve_false_writes_with_existing_materials(monkeypatch) -> None:
|
||
frames = _patch_tool_dependencies(monkeypatch)
|
||
engines_before = len(_FakeEngine.instances)
|
||
|
||
command = asyncio.run(
|
||
deep_research_report_tool.coroutine(
|
||
_fake_runtime("thread-c", messages=[]), # 极端:一条资料也没有
|
||
"call-4",
|
||
"课题",
|
||
retrieve=False,
|
||
)
|
||
)
|
||
|
||
# 用户明确拒绝检索 → 即使收割为空也直接进入管线(runner 有空资料兜底)
|
||
engine = _FakeEngine.instances[-1]
|
||
assert len(_FakeEngine.instances) == engines_before + 1
|
||
assert isinstance(engine.adapters.materials, StaticHarvestMaterialProvider)
|
||
assert engine.adapters.materials.material_count == 0
|
||
assert "对话资料" in command.update["messages"][0].content
|
||
assert any(f.get("event", {}).get("type") == "tool_finished" and f["event"]["payload"]["status"] == "success" for f in frames)
|
||
|
||
|
||
def test_tool_reports_engine_failure_without_artifacts(monkeypatch) -> None:
|
||
failing = _FakeEngine()
|
||
failing.error = RuntimeError("检索服务不可用")
|
||
frames = _patch_tool_dependencies(monkeypatch, engine=failing)
|
||
|
||
command = asyncio.run(deep_research_report_tool.coroutine(_fake_runtime("thread-x"), "call-5", _DEFAULT_TOPIC))
|
||
|
||
update = command.update
|
||
assert "artifacts" not in update
|
||
assert "检索服务不可用" in update["messages"][0].content
|
||
assert any(f.get("event", {}).get("type") == "tool_finished" and f["event"]["payload"]["status"] == "error" for f in frames)
|
||
|
||
|
||
def test_tool_errors_cleanly_without_thread_id(monkeypatch) -> None:
|
||
_patch_tool_dependencies(monkeypatch)
|
||
runtime = SimpleNamespace(context={}, config={}, state={"messages": []})
|
||
engines_before = len(_FakeEngine.instances)
|
||
|
||
command = asyncio.run(deep_research_report_tool.coroutine(runtime, "call-6", "课题"))
|
||
|
||
update = command.update
|
||
assert "artifacts" not in update
|
||
assert "thread id" in update["messages"][0].content
|
||
# 未触达引擎
|
||
assert len(_FakeEngine.instances) == engines_before
|
||
|
||
|
||
# ── 模型继承(P2:管线用聊天选的模型与思考开关) ─────────────────────────────
|
||
|
||
|
||
def test_writing_research_config_inherits_chat_model() -> None:
|
||
cfg = build_writing_research_config(thinking_enabled=True, model_name="qwen3.6")
|
||
|
||
assert cfg.fast_model == "qwen3.6"
|
||
assert cfg.smart_model == "qwen3.6"
|
||
assert cfg.strategic_model == "qwen3.6"
|
||
assert cfg.thinking_enabled is True
|
||
# 缺省回落系统默认模型 + 关思考
|
||
defaults = build_writing_research_config()
|
||
assert defaults.fast_model is None and defaults.smart_model is None and defaults.strategic_model is None
|
||
assert defaults.thinking_enabled is False
|
||
|
||
|
||
def _fake_app_config(monkeypatch, *, known: str = "qwen3.6", supports_thinking: bool = True) -> None:
|
||
import deerflow.config.app_config as app_config_module
|
||
|
||
def get_model_config(name: str):
|
||
if name != known:
|
||
return None
|
||
return SimpleNamespace(supports_thinking=supports_thinking)
|
||
|
||
monkeypatch.setattr(
|
||
app_config_module,
|
||
"get_app_config",
|
||
lambda: SimpleNamespace(get_model_config=get_model_config),
|
||
)
|
||
|
||
|
||
def test_resolve_chat_model_passes_through_known_model(monkeypatch) -> None:
|
||
_fake_app_config(monkeypatch, known="qwen3.6", supports_thinking=True)
|
||
|
||
name, thinking, force = tool_module._resolve_chat_model({"model_name": "qwen3.6", "thinking_enabled": True})
|
||
|
||
assert (name, thinking, force) == ("qwen3.6", True, False)
|
||
|
||
|
||
def test_resolve_chat_model_falls_back_on_unknown_or_unsupported(monkeypatch) -> None:
|
||
_fake_app_config(monkeypatch, known="qwen3.6", supports_thinking=False)
|
||
|
||
# 不在 config.yaml 的模型 → 回落系统默认
|
||
assert tool_module._resolve_chat_model({"model_name": "no-such-model"})[0] is None
|
||
# 模型不支持思考 → 关思考
|
||
name, thinking, _ = tool_module._resolve_chat_model({"model_name": "qwen3.6", "thinking_enabled": True})
|
||
assert (name, thinking) == ("qwen3.6", False)
|
||
# 未选模型 → None + 透传 force 开关
|
||
assert tool_module._resolve_chat_model({"thinking_force_disabled": True}) == (None, True, True)
|
||
|
||
|
||
def test_tool_threads_chat_model_into_research_config(monkeypatch) -> None:
|
||
_patch_tool_dependencies(monkeypatch)
|
||
_fake_app_config(monkeypatch, known="qwen3.6", supports_thinking=True)
|
||
runtime = SimpleNamespace(
|
||
context={
|
||
"thread_id": "thread-m",
|
||
"model_name": "qwen3.6",
|
||
"thinking_enabled": True,
|
||
"thinking_force_disabled": False,
|
||
},
|
||
config={},
|
||
state={"messages": [_material_message(i) for i in range(1, 4)]},
|
||
)
|
||
|
||
asyncio.run(deep_research_report_tool.coroutine(runtime, "call-m", _DEFAULT_TOPIC))
|
||
|
||
config = _FakeEngine.instances[-1].requests[0].config
|
||
assert config.smart_model == "qwen3.6"
|
||
assert config.strategic_model == "qwen3.6"
|
||
assert config.thinking_enabled is True
|
||
|
||
|
||
# ── prompt 合约 ──────────────────────────────────────────────────────────────
|
||
|
||
|
||
def _patch_prompt_sections(monkeypatch, module) -> None:
|
||
monkeypatch.setattr(module, "_get_memory_context", lambda *a, **k: "")
|
||
monkeypatch.setattr(module, "get_skills_prompt_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(module, "get_deferred_tools_prompt_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(module, "_build_acp_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(module, "_build_custom_mounts_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(module, "_current_time_context", lambda: "")
|
||
|
||
|
||
def test_lead_prompt_first_turn_uses_research_tool_contract(monkeypatch) -> None:
|
||
from deerflow.agents.lead_agent import prompt as prompt_module
|
||
|
||
_patch_prompt_sections(monkeypatch, prompt_module)
|
||
|
||
prompt = prompt_module.apply_prompt_template(writing_mode=True)
|
||
|
||
assert "deep_research_report" in prompt
|
||
assert "Work normally until delivery time" in prompt
|
||
assert "Do NOT ask the user whether to search for more" in prompt
|
||
assert "first report turn" in prompt
|
||
assert "do NOT call `read_file`" in prompt
|
||
assert "Current Markdown Artifact" not in prompt
|
||
|
||
|
||
def test_lead_prompt_follow_up_keeps_manual_edit_contract(monkeypatch) -> None:
|
||
from deerflow.agents.lead_agent import prompt as prompt_module
|
||
|
||
_patch_prompt_sections(monkeypatch, prompt_module)
|
||
|
||
prompt = prompt_module.apply_prompt_template(
|
||
writing_mode=True,
|
||
writing_artifact_path="/mnt/user-data/outputs/report.md",
|
||
)
|
||
|
||
# 续写轮保持原合约:编辑现有文件 + present_files,不再要求调用管线工具
|
||
assert "Current Markdown Artifact" in prompt
|
||
assert "present_files" in prompt
|
||
assert "deep_research_report" not in prompt
|
||
|
||
|
||
def test_agent_only_prompt_mirrors_both_contracts(monkeypatch) -> None:
|
||
from deerflow.agents.lead_agent import prompt as prompt_module
|
||
|
||
monkeypatch.setattr(prompt_module, "get_agent_soul", lambda *a, **k: "SOUL")
|
||
monkeypatch.setattr(prompt_module, "get_skills_prompt_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(prompt_module, "get_deferred_tools_prompt_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(prompt_module, "_build_acp_section", lambda *a, **k: "")
|
||
monkeypatch.setattr(prompt_module, "_build_custom_mounts_section", lambda *a, **k: "")
|
||
|
||
first_turn = prompt_module.apply_agent_only_prompt_template(agent_id="a1", agent_name="A1", writing_mode=True)
|
||
follow_up = prompt_module.apply_agent_only_prompt_template(
|
||
agent_id="a1",
|
||
agent_name="A1",
|
||
writing_mode=True,
|
||
writing_artifact_path="/mnt/user-data/outputs/report.md",
|
||
)
|
||
|
||
assert "deep_research_report" in first_turn
|
||
assert "Work normally (chat, search, clarify) until the report is due" in first_turn
|
||
assert "Do NOT ask the user whether to search for more" in first_turn
|
||
assert "do NOT call `read_file`" in first_turn
|
||
assert "Current Markdown Artifact" in follow_up
|
||
assert "deep_research_report" not in follow_up
|