kg_relations.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367
  1. """知识图谱关系管理 —— 节点/边数据模型、图谱指标统计与查询."""
  2. from __future__ import annotations
  3. import json
  4. import logging
  5. import re
  6. import sqlite3
  7. import threading
  8. import time
  9. from typing import Any
  10. from ...utils.common import suppress_exceptions
  11. from ..reader.paper_reader_context import extract_pdf_text_full_cached, extract_pdf_text_full, _cache_set
  12. from ...utils import normalize_arxiv_id as _norm_arxiv_id
  13. logger = logging.getLogger(__name__)
  14. _kg_infer_lock = threading.Lock()
  15. _kg_recent_fingerprints: dict[str, float] = {}
  16. _kg_fingerprints_lock = threading.Lock()
  17. _KG_DEDUP_WINDOW_SEC = 180.0
  18. _kg_metrics: dict[str, int] = {
  19. "build_ok": 0,
  20. "build_skip_no_candidates": 0,
  21. "build_skip_dedup": 0,
  22. "relations_upserted": 0,
  23. }
  24. _kg_metrics_lock = threading.Lock()
  25. def get_kg_metrics() -> dict[str, int]:
  26. with _kg_metrics_lock:
  27. return dict(_kg_metrics)
  28. def _prune_recent_fingerprints(now: float) -> None:
  29. cutoff = now - _KG_DEDUP_WINDOW_SEC * 2
  30. with _kg_fingerprints_lock:
  31. dead = [k for k, t in _kg_recent_fingerprints.items() if t < cutoff]
  32. for k in dead:
  33. _kg_recent_fingerprints.pop(k, None)
  34. from ...utils.common import exec_sql
  35. def ensure_tables(db_path: str) -> None:
  36. exec_sql(db_path,
  37. """CREATE TABLE IF NOT EXISTS paper_relations (
  38. source_paper_id INTEGER NOT NULL,
  39. target_paper_id INTEGER NOT NULL,
  40. relation TEXT NOT NULL,
  41. score REAL DEFAULT 0.0,
  42. evidence TEXT,
  43. created_at INTEGER NOT NULL,
  44. updated_at INTEGER NOT NULL,
  45. PRIMARY KEY (source_paper_id, target_paper_id, relation)
  46. )""",
  47. "CREATE INDEX IF NOT EXISTS idx_paper_relations_source ON paper_relations(source_paper_id)",
  48. "CREATE INDEX IF NOT EXISTS idx_paper_relations_target ON paper_relations(target_paper_id)",
  49. )
  50. @suppress_exceptions(default_return=[])
  51. def _loads_json_list(x: str | None) -> list[Any]:
  52. r = json.loads(x or "[]")
  53. return r if isinstance(r, list) else []
  54. def _row_to_paper_meta(row: sqlite3.Row) -> dict[str, Any]:
  55. keys = row.keys()
  56. def _col(name: str) -> Any:
  57. return row[name] if name in keys else None
  58. return {
  59. "id": row["id"],
  60. "title": row["title"],
  61. "abstract": row["abstract"],
  62. "keywords": _loads_json_list(_col("keywords")),
  63. "tags": _loads_json_list(_col("tags")),
  64. "journal": row["journal"],
  65. "venue_type": row["venue_type"] if "venue_type" in keys else None,
  66. "year": row["year"],
  67. "category": row["category"],
  68. "local_pdf_path": _col("local_pdf_path"),
  69. "doi": (_col("doi") or None),
  70. "arxiv_id": (_col("arxiv_id") or None),
  71. "references": _loads_json_list(_col("references")),
  72. }
  73. def _pdf_abspath_from_row(db_path: str, local_pdf_path: str | None) -> str | None:
  74. import os
  75. if not local_pdf_path or not str(local_pdf_path).strip():
  76. return None
  77. data_root = os.path.dirname(os.path.abspath(db_path))
  78. abspath = os.path.normpath(os.path.join(data_root, str(local_pdf_path).strip()))
  79. if os.path.isfile(abspath):
  80. return abspath
  81. return None
  82. def _extract_related_work_excerpt(text: str, max_chars: int = 1600) -> str:
  83. t = (text or "").strip()
  84. if not t:
  85. return ""
  86. patterns = [
  87. r"(?i)\brelated\s+work\b",
  88. r"(?i)\brelated\s+works\b",
  89. r"相关工作",
  90. ]
  91. for pat in patterns:
  92. m = re.search(pat, t)
  93. if m:
  94. return t[m.start() : m.start() + max_chars].strip()
  95. return t[: max_chars // 2].strip()
  96. _TOKEN_RE = re.compile(r"[^\w\u4e00-\u9fff]+", re.UNICODE)
  97. _LEX_STOP = frozenset({
  98. "the", "a", "an", "and", "or", "of", "to", "in", "for", "with", "on", "by",
  99. "we", "our", "is", "are", "be", "this", "that", "from", "as", "at", "it",
  100. })
  101. def _lexical_token_set(meta: dict[str, Any]) -> set[str]:
  102. parts = [
  103. str(meta.get("title") or ""),
  104. str(meta.get("abstract") or ""),
  105. *(str(k) for k in meta.get("keywords") or []),
  106. *(str(t) for t in meta.get("tags") or []),
  107. ]
  108. text = " ".join(parts).lower()
  109. text = _TOKEN_RE.sub(" ", text)
  110. return {s for s in text.split() if len(s) >= 2 and s not in _LEX_STOP}
  111. def _reference_signatures(refs: list[Any]) -> set[str]:
  112. sigs: set[str] = set()
  113. for r in refs or []:
  114. s = str(r).strip().lower()
  115. if not s:
  116. continue
  117. sigs.add(s)
  118. m = re.search(r"(10\.\d{4,9}/[^\s,;\"'<>]+)", s)
  119. if m:
  120. sigs.add(m.group(1).rstrip(").,]}"))
  121. m = re.search(r"arxiv[:/\s]*(\d{4}\.\d{4,5})(?:v\d+)?", s)
  122. if m:
  123. sigs.add(_norm_arxiv_id(m.group(1)) or m.group(1))
  124. return sigs
  125. def _ref_overlap_bonus(new_sigs: set[str], cand: dict[str, Any]) -> float:
  126. if not new_sigs:
  127. return 0.0
  128. doi = (cand.get("doi") or "").strip().lower()
  129. if doi and doi in new_sigs:
  130. return 1.0
  131. ax = _norm_arxiv_id(cand.get("arxiv_id"))
  132. if ax and ax in new_sigs:
  133. return 1.0
  134. if doi:
  135. for sig in new_sigs:
  136. if len(sig) > 8 and sig in doi:
  137. return 0.85
  138. return 0.0
  139. def fetch_new_and_candidates(
  140. db_path: str,
  141. new_paper_id: int,
  142. k: int = 32,
  143. *,
  144. sql_limit: int = 22,
  145. lexical_pool: int = 480,
  146. ) -> tuple[dict[str, Any | None], list[dict[str, Any]]]:
  147. conn = sqlite3.connect(db_path)
  148. conn.row_factory = sqlite3.Row
  149. cur = conn.cursor()
  150. cur.execute("SELECT * FROM papers WHERE id=?", (int(new_paper_id),))
  151. row = cur.fetchone()
  152. if not row:
  153. conn.close()
  154. return None, []
  155. new_meta = _row_to_paper_meta(row)
  156. cat = (new_meta.get("category") or "").strip()
  157. year = new_meta.get("year")
  158. params: list[Any] = []
  159. where = ["id != ?"]
  160. params.append(int(new_paper_id))
  161. if cat:
  162. where.append("(category = ? OR category LIKE ?)")
  163. params.extend([cat, f"{cat.split('/')[0]}%"])
  164. if year:
  165. try:
  166. y = int(year)
  167. where.append("(year IS NULL OR year >= ?)")
  168. params.append(max(1900, y - 8))
  169. except Exception:
  170. pass
  171. wsql = " AND ".join(where)
  172. cur.execute(
  173. f"SELECT * FROM papers WHERE {wsql} ORDER BY created_at DESC LIMIT ?",
  174. (*params, int(sql_limit)),
  175. )
  176. sql_metas = [_row_to_paper_meta(r) for r in cur.fetchall()]
  177. cur.execute(
  178. """
  179. SELECT * FROM papers
  180. WHERE id != ?
  181. ORDER BY created_at DESC
  182. LIMIT ?
  183. """,
  184. (int(new_paper_id), int(lexical_pool)),
  185. )
  186. pool_rows = cur.fetchall()
  187. conn.close()
  188. new_lex = _lexical_token_set(new_meta)
  189. new_ref_sigs = _reference_signatures(new_meta.get("references") or [])
  190. scored: list[tuple[float, dict[str, Any]]] = []
  191. for r in pool_rows:
  192. meta = _row_to_paper_meta(r)
  193. cand_tokens = _lexical_token_set(meta)
  194. j = len(new_lex & cand_tokens) / max(1, len(new_lex | cand_tokens))
  195. ro = _ref_overlap_bonus(new_ref_sigs, meta)
  196. comb = j + ro * 0.45
  197. scored.append((comb, meta))
  198. scored.sort(key=lambda x: -x[0])
  199. seen: set[int] = set()
  200. out: list[dict[str, Any]] = []
  201. for m in sql_metas:
  202. pid = int(m["id"])
  203. if pid not in seen:
  204. seen.add(pid)
  205. out.append(m)
  206. min_lex = 0.055
  207. for comb, m in scored:
  208. if len(out) >= int(k):
  209. break
  210. pid = int(m["id"])
  211. if pid in seen:
  212. continue
  213. ro = _ref_overlap_bonus(new_ref_sigs, m)
  214. if comb < min_lex and ro <= 0:
  215. continue
  216. seen.add(pid)
  217. out.append(m)
  218. return new_meta, out[: int(k)]
  219. def upsert_relations(db_path: str, source_paper_id: int, edges: list[dict[str, Any]]) -> int:
  220. ensure_tables(db_path)
  221. now = int(time.time())
  222. conn = sqlite3.connect(db_path)
  223. cur = conn.cursor()
  224. n = 0
  225. for e in edges:
  226. try:
  227. tid = int(e.get("target_paper_id"))
  228. except Exception:
  229. continue
  230. rel = str(e.get("relation") or "").strip()[:32]
  231. if len(rel) < 2 or len(rel) > 32:
  232. continue
  233. try:
  234. score = float(e.get("score") or 0.0)
  235. except Exception:
  236. score = 0.0
  237. ev = str(e.get("evidence") or "").strip()[:240]
  238. if tid <= 0 or tid == int(source_paper_id):
  239. continue
  240. cur.execute(
  241. """
  242. INSERT INTO paper_relations(source_paper_id, target_paper_id, relation, score, evidence, created_at, updated_at)
  243. VALUES(?, ?, ?, ?, ?, ?, ?)
  244. ON CONFLICT(source_paper_id, target_paper_id, relation) DO UPDATE SET
  245. score=excluded.score,
  246. evidence=excluded.evidence,
  247. updated_at=excluded.updated_at
  248. """,
  249. (int(source_paper_id), int(tid), rel, float(score), ev, now, now),
  250. )
  251. n += 1
  252. conn.commit()
  253. conn.close()
  254. return n
  255. def build_relations_for_new_paper(db_path: str, new_paper_id: int) -> int:
  256. from ...agents import get_knowledge_graph_agent
  257. ensure_tables(db_path)
  258. new_meta, cands = fetch_new_and_candidates(db_path, int(new_paper_id), k=32)
  259. if not new_meta:
  260. return 0
  261. if not cands:
  262. with _kg_metrics_lock:
  263. _kg_metrics["build_skip_no_candidates"] = _kg_metrics.get("build_skip_no_candidates", 0) + 1
  264. logger.info(
  265. "kg_build_skip_no_candidates",
  266. extra={"paper_id": int(new_paper_id)},
  267. )
  268. return 0
  269. now = time.time()
  270. _prune_recent_fingerprints(now)
  271. _ax = _norm_arxiv_id(new_meta.get("arxiv_id"))
  272. if _ax:
  273. fp = f"arxiv:{_ax}"
  274. else:
  275. _doi = (new_meta.get("doi") or "").strip().lower()
  276. if _doi:
  277. fp = f"doi:{_doi}"
  278. else:
  279. _t = (new_meta.get("title") or "").strip().lower()[:160]
  280. fp = f"title:{_t}"
  281. with _kg_fingerprints_lock:
  282. last = _kg_recent_fingerprints.get(fp)
  283. if last is not None and (now - last) < _KG_DEDUP_WINDOW_SEC:
  284. with _kg_metrics_lock:
  285. _kg_metrics["build_skip_dedup"] = _kg_metrics.get("build_skip_dedup", 0) + 1
  286. logger.info(
  287. "kg_build_skip_dedup",
  288. extra={"paper_id": int(new_paper_id), "fingerprint": fp},
  289. )
  290. return 0
  291. _kg_recent_fingerprints[fp] = now
  292. try:
  293. pdf_abspath = _pdf_abspath_from_row(db_path, new_meta.get("local_pdf_path"))
  294. if pdf_abspath:
  295. excerpt, _hit = extract_pdf_text_full_cached(
  296. db_path, int(new_paper_id), pdf_abspath, max_chars=9000
  297. )
  298. if not excerpt.strip():
  299. excerpt = extract_pdf_text_full(pdf_abspath, max_chars=9000)
  300. if excerpt.strip():
  301. _cache_set(db_path, int(new_paper_id), pdf_abspath, excerpt)
  302. if excerpt.strip():
  303. new_meta["pdf_excerpt"] = excerpt[:9000]
  304. new_meta["related_work_excerpt"] = _extract_related_work_excerpt(excerpt, max_chars=1600)
  305. except Exception as exc:
  306. logger.warning(
  307. "kg_pdf_excerpt_failed",
  308. extra={"paper_id": int(new_paper_id)},
  309. exc_info=exc,
  310. )
  311. edges: list[dict[str, Any]] = []
  312. with _kg_infer_lock:
  313. agent = get_knowledge_graph_agent()
  314. edges, _ = agent.infer_edges(new_paper=new_meta, candidates=cands)
  315. n = upsert_relations(db_path, int(new_paper_id), edges)
  316. if n:
  317. with _kg_metrics_lock:
  318. _kg_metrics["relations_upserted"] = _kg_metrics.get("relations_upserted", 0) + n
  319. with _kg_metrics_lock:
  320. _kg_metrics["build_ok"] = _kg_metrics.get("build_ok", 0) + 1
  321. return n