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

139 lines
5.0 KiB
Python

"""Built-in forced-research responder seed and runtime-policy regressions."""
from __future__ import annotations
from types import SimpleNamespace
import pytest
import yaml
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from app.gateway.routers import _forced_research_seed as seed
from deerflow.agents.middlewares import forced_research_middleware as policy
from deerflow.config.paths import Paths
AGENT_ID = "forced-research-responder"
def _tool_turn(name: str, args: dict, result: str = "OK") -> list:
return [
HumanMessage(content="请查询并按我的格式输出"),
AIMessage(
content="",
tool_calls=[{"id": "call-1", "name": name, "args": args, "type": "tool_call"}],
),
ToolMessage(content=result, tool_call_id="call-1", name=name),
]
def test_seed_assets_define_forced_research_agent():
assert policy.FORCED_RESEARCH_AGENT_ID == AGENT_ID
src = seed._ASSETS_DIR / AGENT_ID
config = yaml.safe_load((src / "config.yaml").read_text(encoding="utf-8"))
soul = (src / "SOUL.md").read_text(encoding="utf-8")
assert config["id"] == AGENT_ID
assert "knowledge-base-search" in config["skills"]
assert "每轮" in soul and "实际" in soul and "检索" in soul
assert "/mnt/user-data/workspace/" in soul
assert "禁止创建或修改 `.md`" in soul
@pytest.fixture()
def tmp_paths(tmp_path, monkeypatch):
paths = Paths(tmp_path)
monkeypatch.setattr(seed, "get_paths", lambda: paths)
return paths
def test_seed_creates_once_and_never_overwrites(tmp_paths):
assert seed.ensure_forced_research_agent() is True
agent_dir = tmp_paths.agent_dir(AGENT_ID)
assert (agent_dir / "config.yaml").is_file()
assert (agent_dir / "SOUL.md").is_file()
(agent_dir / "SOUL.md").write_text("管理员现场修改", encoding="utf-8")
assert seed.ensure_forced_research_agent() is False
assert (agent_dir / "SOUL.md").read_text(encoding="utf-8") == "管理员现场修改"
def test_skill_discovery_and_python_creation_do_not_satisfy_collection_gate():
messages = _tool_turn("skill_view", {"name": "knowledge-base-search"}, "技能说明")
assert policy._turn_collection_state(messages) == (False, 1)
messages = _tool_turn(
"write_file",
{"path": "/mnt/user-data/workspace/retrieve.py", "content": "print('x')"},
"OK",
)
assert policy._turn_collection_state(messages) == (False, 1)
def test_real_search_or_python_skill_execution_satisfies_collection_gate():
messages = _tool_turn("web_search", {"query": "测试检索"}, "真实搜索结果")
assert policy._turn_collection_state(messages) == (True, 1)
messages = _tool_turn(
"bash",
{"command": "cd /mnt/skills/demo && python scripts/search.py 测试"},
'{"results": [{"title": "命中"}]}',
)
assert policy._turn_collection_state(messages) == (True, 1)
@pytest.mark.parametrize(
("name", "args"),
[
("write_file", {"path": "/mnt/user-data/workspace/query.py"}),
("str_replace", {"path": "/mnt/user-data/workspace/query.py"}),
("bash", {"command": "cd /mnt/skills/demo && python3 scripts/query.py 关键词"}),
("web_search", {"query": "关键词"}),
],
)
def test_policy_allows_retrieval_and_temporary_python(name, args):
assert policy.validate_forced_research_tool_call({"name": name, "args": args}) is None
@pytest.mark.parametrize(
("name", "args"),
[
("write_file", {"path": "/mnt/user-data/outputs/report.md"}),
("write_file", {"path": "/mnt/user-data/workspace/report.md"}),
("str_replace", {"path": "/mnt/user-data/workspace/result.json"}),
("present_files", {"paths": ["/mnt/user-data/outputs/report.md"]}),
("bash", {"command": "python scripts/query.py > result.md"}),
("bash", {"command": "echo hello"}),
],
)
def test_policy_rejects_document_delivery_and_non_python_shell(name, args):
assert policy.validate_forced_research_tool_call({"name": name, "args": args})
def test_forced_research_tool_filter_removes_unrelated_write_surfaces():
tools = [
SimpleNamespace(name="web_search"),
SimpleNamespace(name="write_file"),
SimpleNamespace(name="bash"),
SimpleNamespace(name="present_files"),
SimpleNamespace(name="skill_manage"),
SimpleNamespace(name="skill_view"),
]
names = [tool.name for tool in policy.filter_forced_research_tools(tools)]
assert names == ["web_search", "write_file", "bash", "skill_view"]
def test_provider_plain_answer_is_sent_back_to_model_before_collection():
middleware = policy.ForcedResearchMiddleware()
state = {
"messages": [
HumanMessage(content="问题"),
AIMessage(content="没有检索就直接给出的答案"),
]
}
result = middleware._prevent_premature_exit(state)
assert result is not None
assert result["jump_to"] == "model"
suppressed = result["messages"][0]
assert suppressed.content == ""
assert suppressed.additional_kwargs["hide_from_ui"] is True