139 lines
5.1 KiB
Python
139 lines
5.1 KiB
Python
"""Optional retrieval directions and skills attached to a report structure.
|
||
|
||
检索方向 (``retrieval_directions``): when listed, material collection **splits
|
||
search terms along those facets** instead of treating the labels as queries.
|
||
检索来源 (``retrieval_skills``): when listed, collection calls those skills
|
||
instead of the collector's default web_search / Q&A research-skill stack.
|
||
Empty lists are no-ops: runners and the chat collector keep today's behaviour.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import Any
|
||
|
||
from deerflow.persistence.report_structures.directions import (
|
||
MAX_DIRECTION_CHARS,
|
||
MAX_RETRIEVAL_DIRECTIONS,
|
||
normalize_retrieval_directions,
|
||
)
|
||
from deerflow.persistence.report_structures.skills import (
|
||
MAX_RETRIEVAL_SKILLS,
|
||
MAX_SKILL_NAME_CHARS,
|
||
normalize_retrieval_skills,
|
||
)
|
||
|
||
# Must match the built-in collector agent id (seed + frontend). Extra skills
|
||
# from a report structure are only merged onto this agent so other custom
|
||
# agents cannot widen their whitelist via run context.
|
||
COLLECTOR_AGENT_ID = "deep-research-collector"
|
||
|
||
# Keep in sync with frontend ``COLLECTOR_DIRECTION_INTRO``.
|
||
_COLLECTOR_DIRECTION_INTRO = (
|
||
"请按以下检索方向拆分检索词(这是拆词约束,不是检索词本身):"
|
||
"每个方向拆成若干具体、可检索的检索词(含课题中的专名、时间、地点、同义/中英对照等),"
|
||
"不要直接把方向名称拿去搜索——方向标签往往搜不到结果。"
|
||
"覆盖列出的每一个方向;未列出的方面不必主动扩展。"
|
||
"某个检索词没有结果时,换更具体或更短的词再试,仍不要改用未列出的方向。"
|
||
)
|
||
|
||
_COLLECTOR_SKILL_INTRO = (
|
||
"请严格按照以下技能收集素材:先用 read_file 读取每个技能的 SKILL.md,"
|
||
"再按其说明调用工具检索。不要改用未列出的技能,也不要自行改成纯 web_search"
|
||
"(除非技能说明要求联网)。每个技能至少使用一轮。"
|
||
)
|
||
|
||
__all__ = [
|
||
"COLLECTOR_AGENT_ID",
|
||
"MAX_DIRECTION_CHARS",
|
||
"MAX_RETRIEVAL_DIRECTIONS",
|
||
"MAX_RETRIEVAL_SKILLS",
|
||
"MAX_SKILL_NAME_CHARS",
|
||
"format_collector_direction_block",
|
||
"format_collector_skill_block",
|
||
"format_direction_lines",
|
||
"merge_collector_skill_allowlist",
|
||
"normalize_retrieval_directions",
|
||
"normalize_retrieval_skills",
|
||
"queries_from_retrieval_directions",
|
||
"runtime_extra_skills",
|
||
]
|
||
|
||
|
||
def format_direction_lines(directions: Any) -> str:
|
||
"""Numbered direction list, or ``""`` when empty."""
|
||
dirs = normalize_retrieval_directions(directions)
|
||
return "\n".join(f"{i}. {item}" for i, item in enumerate(dirs, 1))
|
||
|
||
|
||
def queries_from_retrieval_directions(topic: str, directions: Any) -> list[str]:
|
||
"""Last-resort search strings when directed query-planning fails.
|
||
|
||
Prefer LLM/collector splitting: direction labels concatenated with the
|
||
topic are often not retrievable. Empty directions → ``[]``.
|
||
"""
|
||
topic_text = " ".join((topic or "").split()).strip()
|
||
dirs = normalize_retrieval_directions(directions)
|
||
if not dirs:
|
||
return []
|
||
|
||
seen: set[str] = set()
|
||
queries: list[str] = []
|
||
for direction in dirs:
|
||
if topic_text and topic_text not in direction:
|
||
query = f"{topic_text} {direction}".strip()
|
||
else:
|
||
query = direction
|
||
key = query.casefold()
|
||
if not query or key in seen:
|
||
continue
|
||
seen.add(key)
|
||
queries.append(query)
|
||
return queries
|
||
|
||
|
||
def format_collector_direction_block(directions: Any) -> str:
|
||
"""XML block appended to the collector's first user message, or ``""``."""
|
||
lines = format_direction_lines(directions)
|
||
if not lines:
|
||
return ""
|
||
return f"<retrieval-directions>\n{_COLLECTOR_DIRECTION_INTRO}\n{lines}\n</retrieval-directions>"
|
||
|
||
|
||
def format_collector_skill_block(skills: Any) -> str:
|
||
"""XML block appended to the collector's first user message, or ``""``."""
|
||
names = normalize_retrieval_skills(skills)
|
||
if not names:
|
||
return ""
|
||
lines = "\n".join(f"{i}. {item}" for i, item in enumerate(names, 1))
|
||
return f"<retrieval-skills>\n{_COLLECTOR_SKILL_INTRO}\n{lines}\n</retrieval-skills>"
|
||
|
||
|
||
def runtime_extra_skills(cfg: Any) -> list[str]:
|
||
"""Read extra skill names from a LangGraph run config / context dict."""
|
||
if not isinstance(cfg, dict):
|
||
return []
|
||
return normalize_retrieval_skills(cfg.get("extra_skills") or cfg.get("retrieval_skills"))
|
||
|
||
|
||
def merge_collector_skill_allowlist(
|
||
agent_id: str | None,
|
||
base: list[str] | None,
|
||
extra: Any,
|
||
) -> list[str] | None:
|
||
"""Widen the collector agent's skill allowlist for one run.
|
||
|
||
Other agents are unchanged. Empty extras are a no-op. ``None`` base means
|
||
"all skills" (default lead agent) and is left alone.
|
||
"""
|
||
extras = normalize_retrieval_skills(extra)
|
||
if not extras or agent_id != COLLECTOR_AGENT_ID:
|
||
return list(base) if base is not None else None
|
||
seen: set[str] = set()
|
||
out: list[str] = []
|
||
for name in [*(base or []), *extras]:
|
||
if not name or name in seen:
|
||
continue
|
||
seen.add(name)
|
||
out.append(name)
|
||
return out
|