search.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284
  1. """智能搜索 API 路由 —— 自然语言论文搜索与 SSE 流式响应."""
  2. from __future__ import annotations
  3. import asyncio
  4. import logging
  5. import os
  6. import time
  7. from typing import Any, Dict, List, Optional
  8. import anyio
  9. from fastapi import APIRouter, Depends, HTTPException
  10. from fastapi.responses import StreamingResponse
  11. from pydantic import BaseModel, Field
  12. from ...agents.search_agent import SearchAgent, SearchIntent, get_search_agent
  13. from ...api.dependencies import get_searcher
  14. from ...api.search_route_support import (
  15. ToolCallInfo,
  16. last_pipeline_tool_error,
  17. normalize_tool_calls,
  18. track_tool_call,
  19. user_facing_error_message,
  20. )
  21. from ...models.schemas import Paper
  22. from ...services.papers.papers_converters import litpapers_to_api_papers
  23. from ...services.retrieval.search_plan import ResolvedSearchPlan
  24. from ...services.retrieval.search_pipeline import run_search_pipeline_async
  25. from ..tool_events import ToolCallTracker, sse_pack
  26. router = APIRouter(prefix="/papers", tags=["智能搜索"])
  27. logger = logging.getLogger(__name__)
  28. _SSE_QUEUE_SIZE = 128
  29. _SEARCH_AGENT_WALL_SEC = max(
  30. 120.0,
  31. min(900.0, float(os.getenv("PAPERGRAPH_SEARCH_AGENT_WALL_SEC") or 420.0)),
  32. )
  33. _SEARCH_AGENT_INIT_SEC = 25.0
  34. _PREFIX_CONFLICT_MARKER = "为您找到"
  35. class SearchAgentMessage(BaseModel):
  36. message: str = Field(..., min_length=1, max_length=2000, description="用户搜索需求")
  37. mode: str = Field(default="accuracy", description="accuracy=准确性优先, novelty=新颖性优先")
  38. use_tavily: bool = Field(default=False, description="是否使用 Tavily 预搜索")
  39. history: List[Dict[str, str]] = Field(default_factory=list, description="对话历史")
  40. class SearchAgentResponse(BaseModel):
  41. success: bool
  42. response: str
  43. search_params: Optional[Dict[str, Any]] = None
  44. tool_calls: List[ToolCallInfo] = Field(default_factory=list)
  45. papers: List[Paper] = Field(default_factory=list)
  46. total: int = 0
  47. message: Optional[str] = None
  48. def _search_params_from_intent(intent: SearchIntent, **extra: Any) -> Dict[str, Any]:
  49. yf, yt = intent.year_from, intent.year_to
  50. if isinstance(yf, int) and isinstance(yt, int) and yf > yt:
  51. yf, yt = yt, yf
  52. out: Dict[str, Any] = {
  53. "query": intent.query,
  54. "keywords": intent.keywords,
  55. "authors": getattr(intent, "authors", []) or [],
  56. "arxiv_id_list": getattr(intent, "arxiv_id_list", []) or [],
  57. "venues": intent.venues,
  58. "year_from": yf,
  59. "year_to": yt,
  60. "sort": intent.sort,
  61. "use_llm_rank": intent.use_llm_rank,
  62. "rerank_recall_max": intent.rerank_recall_max,
  63. "ranking_rationale": intent.ranking_rationale or None,
  64. }
  65. out.update(extra)
  66. return out
  67. def _generate_suggestions(intent: SearchIntent, papers: List[Paper]) -> List[str]:
  68. if len(papers) < 5:
  69. return [f"扩大搜索:尝试「{intent.query}」而不限定会议"]
  70. return []
  71. def _strip_conflicting_search_summary_prefix(text: str) -> str:
  72. t = (text or "").strip()
  73. if not t or _PREFIX_CONFLICT_MARKER not in t:
  74. return t
  75. return t[: t.find(_PREFIX_CONFLICT_MARKER)].rstrip()
  76. def _explanation_with_suggestions(
  77. agent: SearchAgent,
  78. intent: SearchIntent,
  79. papers: List[Paper],
  80. profile_mode: str,
  81. *,
  82. prefix_plain: str = "",
  83. ) -> str:
  84. base = _strip_conflicting_search_summary_prefix((prefix_plain or "").strip())
  85. expl = agent.explain_results(intent, papers, profile_mode)
  86. explanation = base + "\n\n---\n\n" + expl if (papers and base) else (base or expl)
  87. if papers:
  88. sug = _generate_suggestions(intent, papers)
  89. if sug:
  90. explanation += "\n\n🔍 **您可以这样优化**:\n" + "".join(
  91. f"{i}. {s}\n" for i, s in enumerate(sug, 1)
  92. )
  93. return explanation
  94. def _error_response(msg: str) -> SearchAgentResponse:
  95. return SearchAgentResponse(
  96. success=False,
  97. response=user_facing_error_message(msg),
  98. message=msg,
  99. )
  100. async def _prepare_agent_and_query(request: SearchAgentMessage) -> tuple[SearchAgent, str]:
  101. try:
  102. with anyio.fail_after(_SEARCH_AGENT_INIT_SEC):
  103. agent = await anyio.to_thread.run_sync(get_search_agent)
  104. except TimeoutError as exc:
  105. logger.warning("search-agent init timeout after %.0fs", _SEARCH_AGENT_INIT_SEC, exc_info=exc)
  106. raise HTTPException(status_code=504, detail="search_agent_init_timeout") from exc
  107. return agent, (request.message or "").strip()
  108. async def _run_search_agent_core(
  109. *,
  110. agent: SearchAgent,
  111. request: SearchAgentMessage,
  112. merged_query: str,
  113. searcher: Any,
  114. ) -> SearchAgentResponse:
  115. tool_calls: List[ToolCallInfo] = []
  116. intent = agent.understand_intent(merged_query, request.mode)
  117. with track_tool_call(tool_calls, "understand_intent", {"query": merged_query}) as tc:
  118. tc.result_summary = f"sort={intent.sort}, venues={intent.venues}, yf={intent.year_from}, kw={intent.keywords}"
  119. plan = ResolvedSearchPlan.from_search_intent(intent)
  120. with track_tool_call(tool_calls, "search_pipeline", {"query": intent.query or merged_query}) as tc:
  121. tc.result_summary = "intent→SearchPlan→pipeline"
  122. mr = int(getattr(plan, "max_results", None) or intent.max_results or 10)
  123. pip = await run_search_pipeline_async(searcher=searcher, plan=plan, max_results=mr)
  124. tc.result_summary = f"ranked={len(pip.ranked or [])}"
  125. papers = litpapers_to_api_papers(rp.paper for rp in (pip.ranked or []))
  126. prefix = f"为您找到 {len(papers)} 篇论文。" if papers else "未找到相关论文。"
  127. pipeline_err = last_pipeline_tool_error(tool_calls)
  128. if not papers and pipeline_err:
  129. body = (
  130. "主检索未成功返回论文(多源召回或精排阶段出错),与「数据库里确实没有匹配文献」不同。\n\n"
  131. f"**错误摘要**:{pipeline_err}\n\n"
  132. "建议稍后重试,或略微改写查询;若频繁出现请查看服务端日志。"
  133. )
  134. return SearchAgentResponse(
  135. success=False,
  136. response=body,
  137. search_params=_search_params_from_intent(intent, mode=request.mode),
  138. tool_calls=normalize_tool_calls(tool_calls),
  139. papers=[],
  140. total=0,
  141. message="search_pipeline_error",
  142. )
  143. body = _explanation_with_suggestions(agent, intent, papers, request.mode, prefix_plain=prefix)
  144. return SearchAgentResponse(
  145. success=True,
  146. response=body,
  147. search_params=_search_params_from_intent(intent, mode=request.mode),
  148. tool_calls=normalize_tool_calls(tool_calls),
  149. papers=papers,
  150. total=len(papers),
  151. )
  152. async def _search_agent_impl(request: SearchAgentMessage, searcher: Any):
  153. try:
  154. agent, merged_query = await _prepare_agent_and_query(request)
  155. resp = await asyncio.wait_for(
  156. _run_search_agent_core(
  157. agent=agent,
  158. request=request,
  159. merged_query=merged_query,
  160. searcher=searcher,
  161. ),
  162. timeout=_SEARCH_AGENT_WALL_SEC,
  163. )
  164. return resp, None
  165. except asyncio.TimeoutError as exc:
  166. logger.warning("search-agent timeout after %.0fs", _SEARCH_AGENT_WALL_SEC, exc_info=exc)
  167. return _error_response("search_agent_timeout"), HTTPException(status_code=504, detail="search_agent_timeout")
  168. except HTTPException as e:
  169. return _error_response(str(e.detail or "search_agent_http_error")), e
  170. except Exception:
  171. logger.exception("search-agent unexpected failure")
  172. return _error_response("search_agent_internal_error"), None
  173. @router.post("/search-agent/stream")
  174. async def search_agent_chat_stream(
  175. request: SearchAgentMessage,
  176. searcher=Depends(get_searcher),
  177. ):
  178. async def gen():
  179. # SSE 流式生成器:通过 anyio 内存通道实现事件驱动的流式推送
  180. send, recv = anyio.create_memory_object_stream(_SSE_QUEUE_SIZE)
  181. tracker = ToolCallTracker(sink=lambda ev: send.send_nowait(ev))
  182. tracker.emit("status", {"message": "search-agent 已接入,开始处理"})
  183. async def run_once() -> SearchAgentResponse:
  184. tracker.emit("status", {"message": f"初始化 SearchAgent(mode={request.mode})"})
  185. t0 = time.time()
  186. tracker.emit("status", {"message": "正在检索论文…"})
  187. resp, exc = await _search_agent_impl(request, searcher)
  188. if exc:
  189. code = (
  190. str(exc.detail or "search_agent_http_error")
  191. if isinstance(exc, HTTPException)
  192. else "search_agent_internal_error"
  193. )
  194. msg = user_facing_error_message(code)
  195. tracker.emit("error", {"message": msg})
  196. if not isinstance(exc, HTTPException):
  197. logger.exception("search-agent stream run loop failed")
  198. return _error_response(msg)
  199. tracker.emit(
  200. "final",
  201. {"elapsed_ms": int((time.time() - t0) * 1000), "success": bool(resp.success)},
  202. )
  203. return resp
  204. box: Dict[str, Any] = {"resp": None}
  205. cancelled_exc = anyio.get_cancelled_exc_class()
  206. try:
  207. async with anyio.create_task_group() as tg:
  208. async def _run() -> None:
  209. try:
  210. box["resp"] = await run_once()
  211. finally:
  212. try:
  213. await send.aclose()
  214. except Exception:
  215. pass
  216. tg.start_soon(_run)
  217. async for ev in recv:
  218. try:
  219. yield sse_pack(ev)
  220. except (cancelled_exc, asyncio.CancelledError):
  221. return
  222. except Exception:
  223. return
  224. except (cancelled_exc, asyncio.CancelledError):
  225. return
  226. finally:
  227. try:
  228. await recv.aclose()
  229. except Exception:
  230. pass
  231. resp: Optional[SearchAgentResponse] = box.get("resp")
  232. if resp is None:
  233. resp = _error_response("search_agent_stream_incomplete")
  234. tracker.emit("error", {"message": resp.message or "search_agent_stream_incomplete"})
  235. yield sse_pack({"type": "final_result", "result": resp.model_dump(mode="json")})
  236. return StreamingResponse(
  237. gen(),
  238. media_type="text/event-stream",
  239. headers={
  240. "Cache-Control": "no-cache, no-transform",
  241. "Connection": "keep-alive",
  242. "X-Accel-Buffering": "no",
  243. },
  244. )