"""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