deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/agents/deep_research/retrieval.py
2026-09-07 18:24:55 +08:00

139 lines
5.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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