search_pipeline.py 9.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262
  1. """检索流水线 —— 多源召回 → 去重过滤 → LLM 精排 → 结果输出."""
  2. from __future__ import annotations
  3. import asyncio
  4. from collections import Counter
  5. from dataclasses import dataclass
  6. from typing import Any
  7. import anyio
  8. from ...core.paper import Paper as LitPaper
  9. from ...settings import get_settings
  10. from .paper_filters import should_exclude_main_conference_paper
  11. from .paper_ranker import LlmPaperRanker, RankedPaper
  12. from .pipeline_runtime import SearchRuntimeConfig
  13. from .plan_helpers import is_venue_browse_plan, method_acronym_for, primary_venue
  14. from .recall_context import RecallContext, build_recall_context, enrich_recall_context_from_tavily
  15. from .recall_jobs import build_recall_jobs, dedupe_papers, execute_recall_jobs, merge_candidates
  16. from .relevance_guard import apply_relevance_guard
  17. from .search_plan import ResolvedSearchPlan
  18. @dataclass
  19. class SearchPipelineResult:
  20. effective_query: str
  21. total_candidates: int
  22. ranking_method: str
  23. ranked: list[RankedPaper]
  24. metadata: dict[str, Any]
  25. plan: dict[str, Any]
  26. plan_explanation: str
  27. def _merge_pinned_papers(candidates: list[LitPaper], pinned_ids: list[str], searcher: Any) -> list[LitPaper]:
  28. if not pinned_ids:
  29. return candidates
  30. try:
  31. if hasattr(searcher, "search_by_arxiv_ids"):
  32. pinned = searcher.search_by_arxiv_ids(pinned_ids)
  33. elif hasattr(searcher, "search_async"):
  34. loop = asyncio.new_event_loop()
  35. try:
  36. pinned = loop.run_until_complete(
  37. searcher.search_async(
  38. "",
  39. sources=["arxiv"],
  40. arxiv_id_list=pinned_ids,
  41. max_results=len(pinned_ids) * 2,
  42. )
  43. )
  44. finally:
  45. loop.close()
  46. else:
  47. pinned = searcher.search(
  48. "",
  49. sources=["arxiv"],
  50. arxiv_id_list=pinned_ids,
  51. max_results=len(pinned_ids) * 2,
  52. )
  53. pinned = pinned or []
  54. except Exception:
  55. pinned = []
  56. return merge_candidates(candidates, list(pinned), "prepend")
  57. def _merge_target_titles(plan: ResolvedSearchPlan, ctx: RecallContext) -> list[str]:
  58. seen: set[str] = set()
  59. out: list[str] = []
  60. for t in list(plan.target_titles or []) + list(ctx.canonical_titles or []):
  61. tl = (t or "").strip()
  62. if tl and tl.lower() not in seen:
  63. seen.add(tl.lower())
  64. out.append(tl)
  65. return out[:6]
  66. def normalize_and_filter_candidates(
  67. candidates: list[LitPaper],
  68. *,
  69. plan: ResolvedSearchPlan,
  70. ctx: RecallContext,
  71. recall_cap: int,
  72. meta: dict[str, Any],
  73. ) -> list[LitPaper]:
  74. venue = primary_venue(plan)
  75. ma = method_acronym_for(plan, ctx) or None
  76. candidates = dedupe_papers(candidates)
  77. if ma:
  78. from .method_acronym import paper_matches_method_query
  79. narrowed = [
  80. p
  81. for p in candidates
  82. if paper_matches_method_query(
  83. p,
  84. ma,
  85. canonical_titles=ctx.canonical_titles,
  86. pinned_arxiv_ids=ctx.pinned_arxiv_ids,
  87. venue=venue,
  88. )
  89. ]
  90. if narrowed:
  91. candidates = narrowed
  92. guard_threshold = max(36, recall_cap + 8)
  93. if ma:
  94. guard_threshold = max(10, min(guard_threshold, len(candidates) + 2))
  95. if not is_venue_browse_plan(plan):
  96. candidates, guard_applied = apply_relevance_guard(candidates, plan=plan, guard_threshold=guard_threshold)
  97. if guard_applied:
  98. meta["relevance_guard"] = True
  99. if plan.main_conference_proceedings_only and venue:
  100. pin_y = plan.year_from if plan.year_from == plan.year_to else None
  101. # Only require strong venue signal if we actually found venue-verified papers
  102. from .paper_filters import has_strong_main_conference_venue_signal
  103. venue_verified_count = sum(1 for p in candidates if has_strong_main_conference_venue_signal(p, venue))
  104. require_venue_signal = venue_verified_count >= 3
  105. candidates = [
  106. p
  107. for p in candidates
  108. if not should_exclude_main_conference_paper(
  109. p,
  110. venue,
  111. pinned_year=pin_y,
  112. require_venue_signal=require_venue_signal,
  113. )
  114. ]
  115. return candidates
  116. async def rank_candidates(
  117. candidates: list[LitPaper],
  118. *,
  119. plan: ResolvedSearchPlan,
  120. ctx: RecallContext,
  121. runtime: SearchRuntimeConfig,
  122. meta: dict[str, Any],
  123. ) -> tuple[list[RankedPaper], str, dict[str, Any]]:
  124. if not candidates:
  125. return [], "recall_only", {}
  126. if not plan.use_llm_rank:
  127. return [RankedPaper(paper=p) for p in candidates[: runtime.max_results]], "recall_only", {}
  128. venue = primary_venue(plan)
  129. ranker = LlmPaperRanker(recall_max=runtime.recall_max, fine_top_k=runtime.max_results)
  130. prefer_rec = (plan.sort or "").strip().lower() == "date" or bool(plan.year_from) or bool(venue)
  131. try:
  132. with anyio.fail_after(runtime.rank_wall):
  133. ranked, ranking_metadata = await anyio.to_thread.run_sync(
  134. lambda: ranker.rank(
  135. candidates,
  136. ctx.rank_query,
  137. runtime.max_results,
  138. ranking_profile=ctx.ranking_profile,
  139. target_venue=venue,
  140. target_titles=_merge_target_titles(plan, ctx),
  141. authors=list(plan.authors or []),
  142. venues=list(plan.venues or []),
  143. year_from=plan.year_from,
  144. year_to=plan.year_to,
  145. sort=plan.sort,
  146. prefer_recency=prefer_rec,
  147. main_conference_proceedings_only=bool(plan.main_conference_proceedings_only),
  148. intent_source_message=ctx.intent_source_message,
  149. method_acronym=ctx.search_kwargs.get("method_acronym"),
  150. )
  151. )
  152. return ranked, ranking_metadata.get("ranking_method", "llm_rank"), ranking_metadata
  153. except TimeoutError:
  154. meta["ranking_timeout"] = True
  155. from .paper_ranker import _papers_to_ranked_pool
  156. pool = _papers_to_ranked_pool(candidates, cap=runtime.recall_max, prefer_recency=prefer_rec)
  157. return pool[: runtime.max_results], "recall_fallback_timeout", {}
  158. async def run_search_pipeline_async(
  159. *,
  160. searcher: Any,
  161. plan: ResolvedSearchPlan,
  162. max_results: int | None = None,
  163. ) -> SearchPipelineResult:
  164. runtime = SearchRuntimeConfig.from_settings(get_settings(), plan, max_results)
  165. ctx = await enrich_recall_context_from_tavily(build_recall_context(plan), plan)
  166. meta: dict[str, Any] = {
  167. "ranking_profile": ctx.ranking_profile,
  168. "source_plan": ctx.source_plan,
  169. "recall_context": {
  170. "effective_query": ctx.effective_query[:200],
  171. "rank_query": ctx.rank_query[:200],
  172. "merged_keywords": ctx.merged_keywords[:12],
  173. },
  174. "search_recipe": plan.recipe.value,
  175. }
  176. fallbacks: list[dict[str, Any]] = []
  177. # 阶段 1: 多源并行召回
  178. constraint_kwargs = {**ctx.search_kwargs, "sort": plan.sort or ctx.search_kwargs.get("sort") or "relevance"}
  179. jobs = build_recall_jobs(plan, ctx, runtime=runtime, constraint_kwargs=constraint_kwargs)
  180. candidates = await execute_recall_jobs(
  181. searcher, jobs, plan=plan, ctx=ctx, runtime=runtime, meta=meta, fallbacks=fallbacks
  182. )
  183. # 阶段 2: 补回用户指定的 arXiv ID(pinned papers)
  184. pinned_ids = list(ctx.pinned_arxiv_ids or [])
  185. if pinned_ids and searcher is not None:
  186. try:
  187. with anyio.fail_after(3.0):
  188. candidates = await anyio.to_thread.run_sync(
  189. _merge_pinned_papers, candidates, pinned_ids, searcher
  190. )
  191. except TimeoutError:
  192. pass
  193. # 阶段 3: 去重、过滤非主会论文、相关性守卫
  194. candidates = normalize_and_filter_candidates(
  195. candidates, plan=plan, ctx=ctx, recall_cap=runtime.recall_cap, meta=meta
  196. )
  197. # 阶段 4: LLM 精排(或召回直接截断)
  198. ranked, ranking_method, ranking_metadata = await rank_candidates(
  199. candidates, plan=plan, ctx=ctx, runtime=runtime, meta=meta
  200. )
  201. if not candidates and plan.fallback.allow_arxiv_only:
  202. fallbacks.append({"type": "arxiv_only", "reason": "no_candidates_after_recall"})
  203. sc = Counter(getattr(p, "source", "unknown") or "unknown" for p in candidates)
  204. rsc = Counter(getattr(rp.paper, "source", "unknown") or "unknown" for rp in ranked)
  205. metadata = {
  206. "tavily_enabled": plan.use_tavily,
  207. "tavily_keywords_count": len(ctx.tavily_keywords),
  208. "anchor_title": ctx.canonical_titles[0] if ctx.canonical_titles else None,
  209. "anchor_arxiv_ids": pinned_ids,
  210. "pinned_arxiv_ids": pinned_ids,
  211. "fallbacks": fallbacks,
  212. "candidates_by_source": dict(sc),
  213. "ranked_by_source": dict(rsc),
  214. "deduped_total": len(candidates),
  215. "final_ranked": len(ranked),
  216. **meta,
  217. }
  218. if ranking_metadata:
  219. metadata["ranking"] = ranking_metadata
  220. return SearchPipelineResult(
  221. effective_query=ctx.effective_query,
  222. total_candidates=len(candidates),
  223. ranking_method=ranking_method,
  224. ranked=ranked,
  225. metadata=metadata,
  226. plan={
  227. "llm_keywords": ctx.merged_keywords,
  228. "tavily_keywords": ctx.tavily_keywords,
  229. "canonical_titles": ctx.canonical_titles,
  230. "recipe": plan.recipe.value,
  231. },
  232. plan_explanation="",
  233. )