papers_library_service.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361
  1. """文献库管理服务 —— 论文存储、分类、标签、阅读记录与 PDF 下载管理."""
  2. import logging
  3. import os
  4. import time
  5. from fastapi import BackgroundTasks, HTTPException, Request
  6. from ...settings import get_settings
  7. from ...models.schemas import (
  8. DeletePaperResponse,
  9. LibraryCategoriesResponse,
  10. LibraryCategoryFolder,
  11. Paper,
  12. PapersResponse,
  13. SavePapersRequest,
  14. SavePapersResponse,
  15. UpdatePaperRequest,
  16. UpdatePaperResponse,
  17. )
  18. logger = logging.getLogger(__name__)
  19. def _merge_tag_lists(base: list[str], extra: list[str], max_n: int = 24) -> list[str]:
  20. out: list[str] = []
  21. seen: set[str] = set()
  22. for t in (base or []) + (extra or []):
  23. k = (t or "").strip()
  24. if not k:
  25. continue
  26. low = k.lower()
  27. if low in seen:
  28. continue
  29. seen.add(low)
  30. out.append(k)
  31. if len(out) >= max_n:
  32. break
  33. return out
  34. def list_library_categories(*, db) -> LibraryCategoriesResponse:
  35. try:
  36. from app.core.paper_paths import LIBRARY_PDF_ROOT_DIR
  37. items = db.list_library_category_folders()
  38. return LibraryCategoriesResponse(
  39. success=True,
  40. store_root=LIBRARY_PDF_ROOT_DIR,
  41. folders=[LibraryCategoryFolder(**x) for x in items],
  42. )
  43. except Exception as e:
  44. raise HTTPException(status_code=500, detail=str(e))
  45. def get_library(
  46. *,
  47. db,
  48. litpaper_to_api_paper_fn,
  49. limit: int = 50,
  50. offset: int = 0,
  51. q: str | None = None,
  52. year_from: int | None = None,
  53. year_to: int | None = None,
  54. read_status=None,
  55. tags: str | None = None,
  56. category: str | None = None,
  57. ) -> PapersResponse:
  58. try:
  59. tag_list = [t.strip() for t in tags.split(",")] if tags else None
  60. cat = (category or "").strip() or None
  61. if q or year_from or year_to or read_status or tag_list or cat:
  62. papers_data = db.search_library(
  63. query=q,
  64. tags=tag_list,
  65. year_from=year_from,
  66. year_to=year_to,
  67. read_status=read_status.value if read_status else None,
  68. category=cat,
  69. limit=limit,
  70. offset=offset,
  71. )
  72. else:
  73. papers_data = db.get_all_papers(limit=limit, offset=offset, order_by="created_at DESC")
  74. ids_missing_pdf = [
  75. int(p.id)
  76. for p in papers_data
  77. if p.id is not None and not (getattr(p, "local_pdf_path", None) or "").strip()
  78. ]
  79. if ids_missing_pdf:
  80. repaired = db.repair_library_local_pdf_paths_batch(ids_missing_pdf)
  81. for p in papers_data:
  82. if p.id is not None and int(p.id) in repaired:
  83. p.local_pdf_path = repaired[int(p.id)]
  84. total = db.count_papers() if not (q or year_from or year_to or read_status or tag_list or cat) else len(papers_data) + (1 if len(papers_data) >= limit else 0)
  85. papers = [litpaper_to_api_paper_fn(p) for p in papers_data]
  86. return PapersResponse(success=True, total=total or len(papers), papers=papers)
  87. except Exception as e:
  88. raise HTTPException(status_code=500, detail=str(e))
  89. async def save_papers(
  90. *,
  91. db,
  92. request: SavePapersRequest,
  93. background_tasks: BackgroundTasks,
  94. api_to_lit_fn,
  95. litpaper_to_api_paper_fn,
  96. ) -> SavePapersResponse:
  97. try:
  98. t0 = time.perf_counter()
  99. from ..graph.kg_relations import build_relations_for_new_paper
  100. from app.core.paper_paths import (
  101. LIBRARY_PDF_ROOT_DIR,
  102. library_pdf_relative_path,
  103. normalize_library_category_display,
  104. )
  105. from app.core.pdf_download import download_paper_pdf_to_path, resolve_paper_pdf_url
  106. lit_list = [api_to_lit_fn(p) for p in request.papers]
  107. for api_p, lit_p in zip(request.papers, lit_list):
  108. if (api_p.source_url or "").strip():
  109. lit_p.source_url = (api_p.source_url or "").strip()
  110. if (api_p.pdf_url or "").strip():
  111. lit_p.pdf_url = (api_p.pdf_url or "").strip()
  112. if (api_p.doi or "").strip():
  113. lit_p.doi = (api_p.doi or "").strip()
  114. if (api_p.arxiv_id or "").strip():
  115. lit_p.arxiv_id = (api_p.arxiv_id or "").strip()
  116. # Fill missing arXiv ID from DOI/source URL.
  117. if not (getattr(lit_p, "arxiv_id", None) or "").strip():
  118. import re
  119. for field_val in ((api_p.doi or "").strip(), (api_p.source_url or "").strip()):
  120. m = re.search(r"arxiv/([\d.]+)", field_val, re.I)
  121. if m:
  122. lit_p.arxiv_id = m.group(1)
  123. break
  124. # Infer venue type from journal name when absent.
  125. if not (getattr(lit_p, "venue_type", None) or "").strip():
  126. j = (getattr(lit_p, "journal", None) or "").strip()
  127. if j:
  128. import re
  129. if re.search(r"(?i)\b(journal|transactions|letters|magazine|review|annals|acta|bulletin)\b", j):
  130. lit_p.venue_type = "journal"
  131. else:
  132. lit_p.venue_type = "conference"
  133. # Fill missing DOI from source URL.
  134. if not (getattr(lit_p, "doi", None) or "").strip():
  135. su = (api_p.source_url or "").strip()
  136. import re
  137. m = re.search(r"doi\.org/(10\.\S+)", su, re.I)
  138. if m:
  139. lit_p.doi = m.group(1)
  140. if (api_p.category or "").strip():
  141. lit_p.category = normalize_library_category_display(api_p.category)
  142. llm_classified = 0
  143. if lit_list and getattr(request, "llm_classify", True):
  144. try:
  145. existing_categories = []
  146. try:
  147. existing_categories = db.list_library_categories_by_count(limit=80)
  148. except Exception:
  149. existing_categories = []
  150. from ...agents import get_paper_analysis_agent
  151. agent = get_paper_analysis_agent()
  152. for lit_p in lit_list:
  153. cat, extra = agent.classify_for_library(
  154. lit_p.title,
  155. lit_p.abstract,
  156. lit_p.journal,
  157. getattr(lit_p, "keywords", None) or [],
  158. existing_categories=existing_categories,
  159. )
  160. lit_p.category = cat
  161. lit_p.tags = _merge_tag_lists(lit_p.tags or [], extra)
  162. if not getattr(lit_p, "venue_type", None):
  163. lit_p.venue_type = agent.classify_venue_type(lit_p.journal)
  164. llm_classified += 1
  165. except Exception as e:
  166. logger.warning("大模型归类未执行:%s", e)
  167. for lit_p in lit_list:
  168. lit_p.category = normalize_library_category_display(getattr(lit_p, "category", None))
  169. elif lit_list:
  170. for lit_p in lit_list:
  171. lit_p.category = normalize_library_category_display(getattr(lit_p, "category", None))
  172. for lit_p in lit_list:
  173. lit_p.category = normalize_library_category_display(getattr(lit_p, "category", None))
  174. t_after_classify = time.perf_counter()
  175. # Backfill missing abstracts via Tavily.
  176. for lit_p in lit_list:
  177. if not (lit_p.abstract or "").strip() and lit_p.title:
  178. try:
  179. from ...settings import get_settings as _gs
  180. _ak = getattr(_gs(), "tavily_api_key", "").strip()
  181. if _ak:
  182. import httpx
  183. _resp = httpx.post("https://api.tavily.com/search", json={
  184. "api_key": _ak, "query": f"{lit_p.title} paper abstract", "max_results": 3}, timeout=15.0)
  185. for _it in (_resp.json().get("results") or []):
  186. if len(_it.get("content","")) > 100:
  187. lit_p.abstract = _it["content"][:2000]
  188. break
  189. except Exception: pass
  190. ids, added_new, updated_existing = db.add_papers(lit_list)
  191. t_after_db = time.perf_counter()
  192. t_after_memory = time.perf_counter()
  193. for pid in ids or []:
  194. try:
  195. if pid is None or int(pid) <= 0:
  196. continue
  197. build_relations_for_new_paper(db.db_path, int(pid))
  198. except Exception:
  199. continue
  200. pdf_downloaded = 0
  201. if request.download_pdfs and lit_list and ids:
  202. s = get_settings()
  203. mail = (s.ncbi_email or "").strip()
  204. data_root = os.path.dirname(os.path.abspath(db.db_path))
  205. os.makedirs(os.path.join(data_root, LIBRARY_PDF_ROOT_DIR), exist_ok=True)
  206. for lit_p, pid in zip(lit_list, ids):
  207. if pid is None or pid < 0:
  208. continue
  209. relpath = library_pdf_relative_path(
  210. getattr(lit_p, "category", None), int(pid), getattr(lit_p, "title", None)
  211. )
  212. dest = os.path.join(data_root, relpath)
  213. os.makedirs(os.path.dirname(dest), exist_ok=True)
  214. try:
  215. resolved = resolve_paper_pdf_url(lit_p, email=mail)
  216. except Exception as ex:
  217. logger.warning("解析 PDF 链接异常(已跳过该条 PDF): %s", ex, exc_info=True)
  218. resolved = None
  219. if not resolved:
  220. logger.warning(
  221. "保存跳过 PDF:无可用链接 title=%r doi=%r",
  222. lit_p.title,
  223. lit_p.doi,
  224. )
  225. if os.path.isfile(dest) and os.path.getsize(dest) >= 256:
  226. db.set_local_pdf_path(int(pid), relpath)
  227. pdf_downloaded += 1
  228. continue
  229. if resolved and download_paper_pdf_to_path(lit_p, dest, email=mail):
  230. db.set_local_pdf_path(int(pid), relpath)
  231. pdf_downloaded += 1
  232. elif resolved:
  233. logger.warning("保存 PDF 下载失败 title=%r", lit_p.title)
  234. t_after_pdf = time.perf_counter()
  235. ids_ok = [int(x) for x in ids if x is not None and int(x) >= 0]
  236. if ids_ok:
  237. need_repair = []
  238. for pid in ids_ok:
  239. row = db.get_paper_by_id(pid)
  240. if row and not (getattr(row, "local_pdf_path", None) or "").strip():
  241. need_repair.append(pid)
  242. if need_repair:
  243. db.repair_library_local_pdf_paths_batch(need_repair)
  244. msg = None
  245. if (
  246. request.download_pdfs
  247. and lit_list
  248. and pdf_downloaded == 0
  249. and any(pid is not None and pid >= 0 for pid in ids)
  250. ):
  251. msg = "未能写入本地 PDF:请确认含 arXiv / pdf_url 等可下载链接。"
  252. logger.info(
  253. "POST /api/papers/save timing total=%.3fs classify=%.3fs db=%.3fs memory=%.3fs pdf=%.3fs"
  254. " papers=%d llm_classify=%s download_pdfs=%s",
  255. (t_after_pdf - t0),
  256. (t_after_classify - t0),
  257. (t_after_db - t_after_classify),
  258. (t_after_memory - t_after_db),
  259. (t_after_pdf - t_after_memory),
  260. len(lit_list),
  261. getattr(request, "llm_classify", True),
  262. getattr(request, "download_pdfs", False),
  263. )
  264. return SavePapersResponse(
  265. success=True,
  266. added=int(added_new),
  267. updated=int(updated_existing),
  268. ids=ids,
  269. pdf_downloaded=pdf_downloaded,
  270. llm_classified=llm_classified,
  271. message=msg,
  272. )
  273. except Exception as e:
  274. logger.exception("POST /api/papers/save 失败")
  275. raise HTTPException(status_code=500, detail=str(e))
  276. def get_paper_by_id(*, db, paper_id: int, litpaper_to_api_paper_fn) -> Paper:
  277. p = db.get_paper_by_id(paper_id)
  278. if not p:
  279. raise HTTPException(status_code=404, detail="文献不存在")
  280. if p.id is not None and not (getattr(p, "local_pdf_path", None) or "").strip():
  281. repaired = db.repair_library_local_pdf_paths_batch([int(p.id)])
  282. if repaired:
  283. p2 = db.get_paper_by_id(paper_id)
  284. if p2:
  285. p = p2
  286. return litpaper_to_api_paper_fn(p)
  287. def update_paper_by_id(*, db, paper_id: int, body: UpdatePaperRequest) -> UpdatePaperResponse:
  288. try:
  289. from app.core.paper_paths import normalize_library_category_display
  290. fields = {}
  291. if body.notes is not None:
  292. fields["notes"] = body.notes
  293. if body.tags is not None:
  294. fields["tags"] = body.tags
  295. if body.category is not None:
  296. fields["category"] = normalize_library_category_display(body.category)
  297. if body.rating is not None:
  298. fields["rating"] = body.rating
  299. if body.read_status is not None:
  300. fields["read_status"] = body.read_status.value
  301. if body.importance is not None:
  302. fields["importance"] = body.importance
  303. ok = db.update_paper(paper_id, **fields)
  304. if not ok:
  305. raise HTTPException(status_code=404, detail="未更新或文献不存在")
  306. return UpdatePaperResponse(success=True, updated_fields=list(fields.keys()))
  307. except HTTPException:
  308. raise
  309. except Exception as e:
  310. raise HTTPException(status_code=500, detail=str(e))
  311. def delete_paper_by_id(*, db, paper_id: int) -> DeletePaperResponse:
  312. ok = db.delete_paper(paper_id)
  313. if not ok:
  314. raise HTTPException(status_code=404, detail="文献不存在")
  315. return DeletePaperResponse(success=True, message="已删除")
  316. def build_library_pdf_response_service(*, paper_id: int, request: Request, db_path: str, logger_obj):
  317. from ..pdf.pdf_service import build_library_pdf_response
  318. return build_library_pdf_response(
  319. paper_id=int(paper_id),
  320. request=request,
  321. db_path=db_path,
  322. logger=logger_obj,
  323. )