154 lines
7.3 KiB
Python
154 lines
7.3 KiB
Python
import logging
|
|
from types import SimpleNamespace
|
|
|
|
from deerflow.agents.lead_agent import agent as agent_module
|
|
from deerflow.agents.lead_agent import prompt as prompt_module
|
|
from deerflow.config.skills_config import EsQueryRoutingConfig, SkillsConfig
|
|
|
|
|
|
def _app_config(*, enabled: bool, prompt: str = "自定义 ES 路由规则") -> SimpleNamespace:
|
|
return SimpleNamespace(
|
|
skills=SkillsConfig(
|
|
es_query_routing=EsQueryRoutingConfig(enabled=enabled, prompt=prompt),
|
|
)
|
|
)
|
|
|
|
|
|
def test_es_query_routing_config_is_disabled_by_default() -> None:
|
|
config = SkillsConfig()
|
|
|
|
assert config.es_query_routing.enabled is False
|
|
assert "es_query" in config.es_query_routing.prompt
|
|
|
|
|
|
def test_es_query_routing_section_honors_switch_prompt_and_allowlist() -> None:
|
|
enabled = _app_config(enabled=True, prompt=" 自定义 ES 路由规则 ")
|
|
|
|
assert prompt_module._build_es_query_routing_section(app_config=enabled) == (
|
|
'<skill_routing skill="es_query">\n自定义 ES 路由规则\n</skill_routing>'
|
|
)
|
|
assert (
|
|
prompt_module._build_es_query_routing_section(
|
|
app_config=enabled,
|
|
available_skills={"other_skill"},
|
|
)
|
|
== ""
|
|
)
|
|
assert "自定义 ES 路由规则" in prompt_module._build_es_query_routing_section(
|
|
app_config=enabled,
|
|
available_skills={"es_query"},
|
|
)
|
|
assert prompt_module._build_es_query_routing_section(app_config=_app_config(enabled=False)) == ""
|
|
assert prompt_module._build_es_query_routing_section(app_config=_app_config(enabled=True, prompt=" ")) == ""
|
|
|
|
|
|
def test_es_query_routing_is_injected_into_lead_and_allowed_custom_agent_prompts(monkeypatch) -> None:
|
|
config = _app_config(enabled=True)
|
|
monkeypatch.setattr(prompt_module, "_get_memory_context", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_agent_soul", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_skills_prompt_section", lambda *args, **kwargs: "<skill_system />")
|
|
monkeypatch.setattr(prompt_module, "get_deferred_tools_prompt_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_build_acp_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_build_custom_mounts_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_current_time_context", lambda: "")
|
|
|
|
lead_prompt = prompt_module.apply_prompt_template(app_config=config)
|
|
allowed_custom_prompt = prompt_module.apply_agent_only_prompt_template(
|
|
agent_id="agent-with-es",
|
|
agent_name="Agent with ES",
|
|
available_skills={"es_query"},
|
|
app_config=config,
|
|
)
|
|
denied_custom_prompt = prompt_module.apply_agent_only_prompt_template(
|
|
agent_id="agent-without-es",
|
|
agent_name="Agent without ES",
|
|
available_skills={"other_skill"},
|
|
app_config=config,
|
|
)
|
|
|
|
assert '<skill_routing skill="es_query">' in lead_prompt
|
|
assert "自定义 ES 路由规则" in lead_prompt
|
|
assert '<skill_routing skill="es_query">' in allowed_custom_prompt
|
|
assert '<skill_routing skill="es_query">' not in denied_custom_prompt
|
|
|
|
|
|
def test_es_query_routing_registration_log_confirms_exact_final_prompt(caplog) -> None:
|
|
config = _app_config(enabled=True, prompt="自定义 ES 路由规则")
|
|
system_prompt = '<skill_routing skill="es_query">\n自定义 ES 路由规则\n</skill_routing>'
|
|
|
|
with caplog.at_level(logging.INFO, logger=agent_module.__name__):
|
|
agent_module._log_es_query_routing_registration(
|
|
system_prompt=system_prompt,
|
|
app_config=config,
|
|
agent_id=None,
|
|
available_skills=None,
|
|
)
|
|
|
|
assert "ES query routing prompt registration" in caplog.text
|
|
assert "agent_id=default" in caplog.text
|
|
assert "enabled=True" in caplog.text
|
|
assert "eligible=True" in caplog.text
|
|
assert "registered=True" in caplog.text
|
|
assert "prompt_chars=11" in caplog.text
|
|
assert "prompt_sha256=" in caplog.text
|
|
|
|
|
|
def test_es_query_routing_registration_log_explains_custom_agent_exclusion(caplog) -> None:
|
|
config = _app_config(enabled=True, prompt="自定义 ES 路由规则")
|
|
|
|
with caplog.at_level(logging.INFO, logger=agent_module.__name__):
|
|
agent_module._log_es_query_routing_registration(
|
|
system_prompt="custom agent prompt",
|
|
app_config=config,
|
|
agent_id="agent-without-es",
|
|
available_skills={"other_skill"},
|
|
)
|
|
|
|
assert "agent_id=agent-without-es" in caplog.text
|
|
assert "enabled=True" in caplog.text
|
|
assert "eligible=False" in caplog.text
|
|
assert "registered=False" in caplog.text
|
|
|
|
|
|
def test_formal_markdown_rules_are_scoped_to_ordinary_qa(monkeypatch) -> None:
|
|
config = _app_config(enabled=False)
|
|
monkeypatch.setattr(prompt_module, "_ordinary_qa_markdown_format_enabled", lambda: True)
|
|
monkeypatch.setattr(prompt_module, "_get_memory_context", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_agent_soul", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_skills_prompt_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_deferred_tools_prompt_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_build_acp_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_build_custom_mounts_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_current_time_context", lambda: "")
|
|
|
|
ordinary_prompt = prompt_module.apply_prompt_template(app_config=config)
|
|
writing_prompt = prompt_module.apply_prompt_template(app_config=config, writing_mode=True)
|
|
notebook_prompt = prompt_module.apply_prompt_template(app_config=config, is_notebook_mode=True)
|
|
scheduled_prompt = prompt_module.apply_prompt_template(app_config=config, is_scheduled_run=True)
|
|
|
|
assert '<ordinary_qa_markdown_format>' in ordinary_prompt
|
|
assert "## 一、标题" in ordinary_prompt
|
|
assert "### (一) 标题" in ordinary_prompt
|
|
assert "#### 1. 标题" in ordinary_prompt
|
|
assert "2.1.1" in ordinary_prompt
|
|
assert "无序圆点列表" in ordinary_prompt
|
|
assert '<ordinary_qa_markdown_format>' not in writing_prompt
|
|
assert '<ordinary_qa_markdown_format>' not in notebook_prompt
|
|
assert '<ordinary_qa_markdown_format>' not in scheduled_prompt
|
|
|
|
|
|
def test_formal_markdown_rules_are_disabled_when_switch_is_off(monkeypatch) -> None:
|
|
config = _app_config(enabled=False)
|
|
monkeypatch.setattr(prompt_module, "_ordinary_qa_markdown_format_enabled", lambda: False)
|
|
monkeypatch.setattr(prompt_module, "_get_memory_context", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_agent_soul", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_skills_prompt_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "get_deferred_tools_prompt_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_build_acp_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_build_custom_mounts_section", lambda *args, **kwargs: "")
|
|
monkeypatch.setattr(prompt_module, "_current_time_context", lambda: "")
|
|
|
|
ordinary_prompt = prompt_module.apply_prompt_template(app_config=config)
|
|
|
|
assert '<ordinary_qa_markdown_format>' not in ordinary_prompt
|