papers_converters.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134
  1. """论文数据转换器 —— LiteraturePaper ↔ API Paper 格式互转."""
  2. from __future__ import annotations
  3. import logging
  4. from typing import Any, Iterable
  5. from pydantic import TypeAdapter
  6. from app.core.paper import Paper as LitPaper
  7. from app.core.search.paper_searcher import abbreviate_journal as _abbrev
  8. from ...models.schemas import Author, Paper, PaperSource, ReadStatus
  9. logger = logging.getLogger(__name__)
  10. _paper_list_adapter = TypeAdapter(list[Paper])
  11. def _coerce_paper_source(val: Any) -> PaperSource:
  12. if isinstance(val, PaperSource):
  13. return val
  14. try:
  15. return PaperSource(str(val or "unknown").lower().strip())
  16. except ValueError:
  17. return PaperSource.UNKNOWN
  18. def _normalize_author_entries(authors_in: list[Any]) -> list[dict[str, Any]]:
  19. norm: list[dict[str, Any]] = []
  20. for a in authors_in:
  21. if isinstance(a, str):
  22. norm.append({"name": a})
  23. elif isinstance(a, dict):
  24. norm.append(a)
  25. else:
  26. norm.append({"name": getattr(a, "name", "") or ""})
  27. return norm
  28. def litpaper_to_api_paper(p: LitPaper) -> Paper:
  29. d = p.to_dict()
  30. d["journal"] = _abbrev(d.get("journal"))
  31. try:
  32. return Paper.model_validate(d)
  33. except Exception:
  34. d = p.to_dict()
  35. src = d.get("source") or "unknown"
  36. try:
  37. ps = PaperSource(src)
  38. except ValueError:
  39. ps = PaperSource.UNKNOWN
  40. rs = d.get("read_status") or "unread"
  41. try:
  42. rse = ReadStatus(rs)
  43. except ValueError:
  44. rse = ReadStatus.UNREAD
  45. return Paper(
  46. id=d.get("id"),
  47. title=d.get("title", ""),
  48. authors=[
  49. Author(**a) if isinstance(a, dict) else Author(name=str(a)) for a in d.get("authors", [])
  50. ],
  51. abstract=d.get("abstract"),
  52. doi=d.get("doi"),
  53. pmid=d.get("pmid"),
  54. arxiv_id=d.get("arxiv_id"),
  55. pmc_id=d.get("pmc_id"),
  56. journal=_abbrev(d.get("journal")),
  57. year=d.get("year"),
  58. volume=d.get("volume"),
  59. issue=d.get("issue"),
  60. pages=d.get("pages"),
  61. publisher=d.get("publisher"),
  62. pdf_url=d.get("pdf_url"),
  63. source_url=d.get("source_url"),
  64. local_pdf_path=d.get("local_pdf_path"),
  65. keywords=d.get("keywords") or [],
  66. mesh_terms=d.get("mesh_terms") or [],
  67. references=d.get("references") or [],
  68. citations=d.get("citations") or 0,
  69. source=ps,
  70. relevance_score=d.get("relevance_score") or 0,
  71. notes=d.get("notes"),
  72. tags=d.get("tags") or [],
  73. category=d.get("category"),
  74. rating=d.get("rating"),
  75. read_status=rse,
  76. importance=d.get("importance") or "normal",
  77. )
  78. def api_paper_to_litpaper(p: Paper) -> LitPaper:
  79. d = p.model_dump(mode="json", exclude_none=False, exclude_unset=False)
  80. d.pop("id", None)
  81. d.pop("local_pdf_path", None)
  82. d.pop("category", None)
  83. d.pop("created_at", None)
  84. d.pop("updated_at", None)
  85. d["source"] = p.source.value
  86. d["read_status"] = p.read_status.value
  87. return LitPaper.from_dict(d)
  88. def normalize_papers_for_api(papers: Iterable[Any] | None) -> list[Paper]:
  89. """统一 API 层 Paper 列表:接受 Paper / LitPaper / dict,返回校验后的 list[Paper]。"""
  90. if not papers:
  91. return []
  92. blobs: list[Any] = []
  93. for raw in papers:
  94. if isinstance(raw, Paper):
  95. blobs.append(raw)
  96. continue
  97. if isinstance(raw, LitPaper):
  98. blobs.append(litpaper_to_api_paper(raw))
  99. continue
  100. d = raw.model_dump() if hasattr(raw, "model_dump") else (dict(raw) if isinstance(raw, dict) else None)
  101. if not d or not str(d.get("title") or "").strip():
  102. continue
  103. d["authors"] = _normalize_author_entries(d.get("authors") or [])
  104. d["source"] = _coerce_paper_source(d.get("source"))
  105. if "journal" not in d and d.get("venue") is not None:
  106. d["journal"] = d.get("venue")
  107. if "source_url" not in d and d.get("url") is not None:
  108. d["source_url"] = d.get("url")
  109. blobs.append(d)
  110. try:
  111. return _paper_list_adapter.validate_python(blobs, strict=False)
  112. except Exception as ex:
  113. logger.exception("paper normalization failed")
  114. raise ValueError("paper_normalization_failed") from ex
  115. def litpapers_to_api_papers(papers: Iterable[LitPaper]) -> list[Paper]:
  116. return normalize_papers_for_api([litpaper_to_api_paper(p) for p in papers])