graph_service.py 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184
  1. """知识图谱构建服务 —— 论文关系抽取、图谱数据管理与查询."""
  2. from __future__ import annotations
  3. import logging
  4. from typing import Any
  5. from fastapi import HTTPException
  6. from ...api.repo import RelationRepository
  7. from ...models.schemas import GraphEdge, GraphNode, LibraryGraphResponse
  8. from ...services.papers.papers_helpers import graph_author_label, graph_author_node_id
  9. from .kg_relations import ensure_tables
  10. logger = logging.getLogger(__name__)
  11. def build_library_graph(
  12. *,
  13. db: Any,
  14. limit: int,
  15. category: str | None,
  16. include_authors: bool,
  17. include_keywords: bool,
  18. relation_edge_limit: int,
  19. focus_paper_id: int | None,
  20. ) -> LibraryGraphResponse:
  21. try:
  22. ensure_tables(db.db_path)
  23. repo = RelationRepository(db.db_path)
  24. focus_id = int(focus_paper_id) if focus_paper_id is not None else None
  25. if focus_id is not None:
  26. fp = db.get_paper_by_id(int(focus_id))
  27. papers = [fp] if fp else []
  28. if not papers:
  29. return LibraryGraphResponse(success=True, nodes=[], edges=[])
  30. paper_ids_in_view: set[int] = {int(focus_id)}
  31. else:
  32. papers = db.get_all_papers(limit=int(limit), order_by="created_at DESC")
  33. cat = (category or "").strip() or None
  34. if cat:
  35. papers = [p for p in papers if (getattr(p, "category", None) or "").strip() == cat]
  36. paper_ids_in_view = set()
  37. for p in papers:
  38. pid = int(getattr(p, "id") or 0)
  39. if pid > 0:
  40. paper_ids_in_view.add(pid)
  41. nodes: dict[str, GraphNode] = {}
  42. edges: dict[tuple[str, str, str], GraphEdge] = {}
  43. def up_node(n: GraphNode):
  44. if n.id in nodes:
  45. nodes[n.id].weight = float(nodes[n.id].weight) + float(n.weight or 1.0)
  46. return
  47. nodes[n.id] = n
  48. def up_edge(e: GraphEdge):
  49. k = (e.source, e.target, e.type)
  50. if k in edges:
  51. edges[k].weight = float(edges[k].weight) + float(e.weight or 1.0)
  52. return
  53. edges[k] = e
  54. for p in papers:
  55. pid = int(getattr(p, "id") or 0)
  56. if pid <= 0:
  57. continue
  58. paper_node_id = f"paper:{pid}"
  59. up_node(GraphNode(
  60. id=paper_node_id, type="paper",
  61. label=str(getattr(p, "title", "") or f"Paper {pid}")[:140],
  62. paper_id=pid, year=getattr(p, "year", None),
  63. category=getattr(p, "category", None),
  64. journal=(getattr(p, "journal", None) or "").strip() or None,
  65. venue_type=(getattr(p, "venue_type", None) or "").strip() or None,
  66. weight=3.0,
  67. ))
  68. if include_authors:
  69. for idx, a in enumerate(getattr(p, "authors", []) or []):
  70. name = (getattr(a, "name", None) or "").strip()
  71. if not name:
  72. continue
  73. aid = graph_author_node_id(pid, idx, a)
  74. alabel = graph_author_label(a, idx, aid)
  75. up_node(GraphNode(id=aid, type="author", label=alabel, weight=1.0))
  76. up_edge(GraphEdge(source=aid, target=paper_node_id, type="authored_by", weight=1.0))
  77. kws: list[str] = []
  78. if include_keywords:
  79. for kw in (getattr(p, "keywords", None) or [])[:24]:
  80. k = str(kw or "").strip()
  81. if not k:
  82. continue
  83. kws.append(k)
  84. kid = f"kw:{k.lower()}"
  85. up_node(GraphNode(id=kid, type="keyword", label=k, weight=1.0))
  86. up_edge(GraphEdge(source=kid, target=paper_node_id, type="has_keyword", weight=1.0))
  87. if include_keywords and len(kws) > 1:
  88. base = [f"kw:{k.lower()}" for k in kws[:12]]
  89. for i in range(len(base)):
  90. for j in range(i + 1, len(base)):
  91. s, t = base[i], base[j]
  92. if s == t:
  93. continue
  94. if s > t:
  95. s, t = t, s
  96. up_edge(GraphEdge(source=s, target=t, type="co_keyword", weight=0.5))
  97. try:
  98. rel_rows: list[tuple[int, int, str, float, str]] = []
  99. if focus_id is not None:
  100. rel_rows = repo.fetch_relation_rows(focus_id=int(focus_id), paper_ids=None, limit=int(relation_edge_limit))
  101. rel_paper_ids: set[int] = {int(focus_id)}
  102. for sid, tid, _, _, _ in rel_rows:
  103. rel_paper_ids.add(int(sid))
  104. rel_paper_ids.add(int(tid))
  105. meta = repo.papers_minimal_by_ids(rel_paper_ids)
  106. for pid, (title, year, cat) in meta.items():
  107. nid = f"paper:{pid}"
  108. if nid in nodes:
  109. continue
  110. up_node(GraphNode(
  111. id=nid, type="paper",
  112. label=(title or f"Paper {pid}")[:140],
  113. paper_id=int(pid), year=year, category=cat, weight=2.0,
  114. ))
  115. else:
  116. rel_rows = repo.fetch_relation_rows(
  117. focus_id=None, paper_ids=paper_ids_in_view, limit=int(relation_edge_limit),
  118. )
  119. for sid, tid, rel, score, evidence in rel_rows:
  120. s = f"paper:{int(sid)}"
  121. t = f"paper:{int(tid)}"
  122. if s not in nodes or t not in nodes:
  123. continue
  124. up_edge(GraphEdge(
  125. source=s, target=t,
  126. type=f"paper_{str(rel or 'related')}",
  127. weight=float(score or 0.6),
  128. evidence=(str(evidence or "").strip()[:240] or None),
  129. ))
  130. rev = f"rev_{rel}" if rel else "related_to"
  131. up_edge(GraphEdge(source=t, target=s, type=f"paper_{rev}", weight=float(score or 0.6) * 0.8))
  132. except Exception:
  133. logger.warning("graph_service: paper-paper relation fetch failed", exc_info=True)
  134. if len(papers) > 1:
  135. paper_kw: dict[int, set[str]] = {}
  136. paper_au: dict[int, set[str]] = {}
  137. for p in papers:
  138. pid = int(getattr(p, "id") or 0)
  139. if pid <= 0:
  140. continue
  141. if include_authors:
  142. paper_au[pid] = {graph_author_label(a, i, "") for i, a in enumerate(getattr(p, "authors", []) or []) if (getattr(a, "name", None) or "").strip()}
  143. if include_keywords:
  144. paper_kw[pid] = {str(k).strip().lower() for k in (getattr(p, "keywords", None) or [])[:16] if str(k).strip()}
  145. if include_authors:
  146. pids = list(paper_au.keys())
  147. for i in range(len(pids)):
  148. for j in range(i + 1, len(pids)):
  149. shared = paper_au[pids[i]] & paper_au[pids[j]]
  150. if shared:
  151. up_edge(GraphEdge(source=f"paper:{pids[i]}", target=f"paper:{pids[j]}",
  152. type="shared_author", weight=min(2.0, len(shared) * 0.6)))
  153. if include_keywords:
  154. pids = list(paper_kw.keys())
  155. for i in range(len(pids)):
  156. for j in range(i + 1, len(pids)):
  157. shared = paper_kw[pids[i]] & paper_kw[pids[j]]
  158. if shared:
  159. up_edge(GraphEdge(source=f"paper:{pids[i]}", target=f"paper:{pids[j]}",
  160. type="shared_keyword", weight=min(2.0, len(shared) * 0.35)))
  161. return LibraryGraphResponse(success=True, nodes=list(nodes.values()), edges=list(edges.values()))
  162. except Exception as e:
  163. logger.exception("graph_service.build_library_graph_failed")
  164. raise HTTPException(status_code=500, detail=str(e))