arxiv_normalization.py 2.8 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889
  1. """arXiv 规范化 —— arXiv ID 格式识别、URL 提取与新旧格式转换."""
  2. from __future__ import annotations
  3. from ...utils.common import dedupe_strings_preserve_order
  4. _ALLOWED_ARXIV_CAT_REST = frozenset("abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789.-")
  5. def _valid_arxiv_category_token(s: str) -> bool:
  6. t = (s or "").strip()
  7. if len(t) > 56 or len(t) < 2:
  8. return False
  9. if ".." in t or t.startswith(".") or t.endswith("."):
  10. return False
  11. first = t[0]
  12. if not ("A" <= first <= "Z" or "a" <= first <= "z"):
  13. return False
  14. return all(c in _ALLOWED_ARXIV_CAT_REST for c in t[1:])
  15. def sanitize_arxiv_categories(raw: list[str | None]) -> list[str]:
  16. stripped = [str(x).strip() for x in raw or []]
  17. return dedupe_strings_preserve_order([t for t in stripped if _valid_arxiv_category_token(t)], max_n=12)
  18. def extract_arxiv_id_from_org_url(t: str) -> str | None:
  19. low = t.lower()
  20. key = "arxiv.org"
  21. idx = low.find(key)
  22. if idx < 0:
  23. return None
  24. after = t[idx + len(key) :]
  25. la = after.lower()
  26. for pref in ("/abs/", "/pdf/"):
  27. p = la.find(pref)
  28. if p < 0:
  29. continue
  30. seg = after[p + len(pref) :]
  31. seg = seg.split("?", 1)[0].split("#", 1)[0].strip().rstrip("/")
  32. return seg if seg else None
  33. return None
  34. def strip_trailing_arxiv_version(u: str) -> str:
  35. s = u.strip().lower()
  36. idx = s.rfind("v")
  37. if idx > 0 and idx < len(s) - 1 and s[idx + 1 :].isdigit():
  38. return s[:idx]
  39. return s
  40. def parse_new_style_arxiv_id(t: str) -> str | None:
  41. u = strip_trailing_arxiv_version(t)
  42. if len(u) < 10 or "." not in u:
  43. return None
  44. dot = u.find(".")
  45. if dot != 4:
  46. return None
  47. ym, tail = u[:4], u[5:]
  48. if not ym.isdigit() or not tail.isdigit() or not (4 <= len(tail) <= 5):
  49. return None
  50. return u
  51. def parse_legacy_arxiv_id(t: str) -> str | None:
  52. u = strip_trailing_arxiv_version(t)
  53. slash = u.find("/")
  54. if slash <= 0:
  55. return None
  56. prefix, digits = u[:slash], u[slash + 1 :]
  57. if len(digits) != 7 or not digits.isdigit():
  58. return None
  59. if not prefix or not all(ch.islower() or ch in ".-" for ch in prefix):
  60. return None
  61. return u
  62. def sanitize_arxiv_id_list(raw: list[str | None]) -> list[str]:
  63. out: list[str] = []
  64. for x in raw or []:
  65. t = str(x).strip().replace(" ", "")
  66. if not t:
  67. continue
  68. extracted = extract_arxiv_id_from_org_url(t)
  69. if extracted is not None:
  70. t = extracted
  71. t = t.replace("arXiv:", "").strip()
  72. nid = parse_new_style_arxiv_id(t)
  73. if nid is not None:
  74. out.append(nid)
  75. continue
  76. lid = parse_legacy_arxiv_id(t)
  77. if lid is not None:
  78. out.append(lid)
  79. return dedupe_strings_preserve_order(out, max_n=8)