139 lines
5.0 KiB
Python
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
|