import_all.py 3.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116
  1. """
  2. 批量导入脚本 — 扫描 data/ 目录下所有兼容格式的 JSON 文件,逐个导入数据库。
  3. 兼容的文件格式:列表中每个元素包含 drug_id, name, sections, source 字段。
  4. 自动跳过 drug_index.json(只有索引无内容)等非兼容文件。
  5. """
  6. import json
  7. import asyncio
  8. import os
  9. import sys
  10. from pathlib import Path
  11. # 确保能导入 ingest 模块
  12. sys.path.insert(0, str(Path(__file__).resolve().parent))
  13. from ingest import ingest_drugs, load_env
  14. # ============================================
  15. # 配置:哪些文件跳过(非兼容格式)
  16. # ============================================
  17. SKIP_FILES = {
  18. "drug_index.json", # 只有药品名和页码,无 sections 内容
  19. "catalog_volume1.json", # 目录结构数据,非药品内容
  20. "catalog_volume2.json",
  21. "catalog_volume3.json",
  22. "catalog_volume4.json",
  23. }
  24. def is_compatible_drug_data(filepath: str) -> bool:
  25. """检查 JSON 文件是否为兼容的药品数据格式"""
  26. filename = os.path.basename(filepath)
  27. if filename in SKIP_FILES:
  28. return False
  29. try:
  30. with open(filepath, "r", encoding="utf-8") as f:
  31. data = json.load(f)
  32. if not isinstance(data, list) or len(data) == 0:
  33. return False
  34. # 检查第一个元素是否有必要字段
  35. first = data[0]
  36. required = ["drug_id", "name", "sections"]
  37. return all(k in first for k in required)
  38. except (json.JSONDecodeError, Exception):
  39. return False
  40. async def main():
  41. load_env()
  42. data_dir = Path(__file__).resolve().parent / "data"
  43. json_files = sorted(data_dir.glob("*.json"))
  44. print("=" * 60)
  45. print("📂 扫描数据目录:", data_dir)
  46. print(f" 发现 {len(json_files)} 个 JSON 文件")
  47. print("=" * 60)
  48. compatible = []
  49. skipped = []
  50. for f in json_files:
  51. if is_compatible_drug_data(str(f)):
  52. compatible.append(f)
  53. print(f" ✅ {f.name} — 兼容,将导入")
  54. else:
  55. skipped.append(f)
  56. print(f" ⏭️ {f.name} — 跳过(非药品数据格式)")
  57. if not compatible:
  58. print("\n❌ 没有找到可导入的数据文件!")
  59. print(" 需要包含 drug_id, name, sections, source 字段的 JSON 数组")
  60. return
  61. print(f"\n📦 共 {len(compatible)} 个文件待导入")
  62. print("-" * 60)
  63. total_drugs = 0
  64. total_chunks = 0
  65. for f in compatible:
  66. try:
  67. print(f"\n🚀 正在导入: {f.name} ...")
  68. await ingest_drugs(str(f))
  69. # ingest_drugs 已经打印了统计信息
  70. # 数一下这个文件有多少条
  71. with open(f, "r", encoding="utf-8") as fh:
  72. data = json.load(fh)
  73. drug_count = len(data) if isinstance(data, list) else 0
  74. total_drugs += drug_count
  75. except Exception as e:
  76. print(f" ❌ 导入失败: {e}")
  77. continue
  78. print("\n" + "=" * 60)
  79. print(f"🎉 全部导入完成!共处理 {len(compatible)} 个文件")
  80. print("=" * 60)
  81. # 提示:如果有旧数据没有 vec,执行修复 SQL
  82. print("""
  83. 💡 提示:如果之前导入的数据缺少 vec 列(pgvector),
  84. 请在 psql 中执行以下 SQL 修复:
  85. UPDATE drug_chunks
  86. SET vec = (embedding::text)::vector
  87. WHERE vec IS NULL AND embedding IS NOT NULL;
  88. """)
  89. if __name__ == "__main__":
  90. asyncio.run(main())