microbiome_loader.py 5.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. """菌群知识库加载器: 从 Java 后端同步文章到 ChromaDB"""
  2. import logging
  3. import httpx
  4. from langchain_chroma import Chroma
  5. from langchain_core.documents import Document
  6. from app.config import settings
  7. from app.rag.embeddings import get_embeddings
  8. from app.rag.splitter import get_knowledge_splitter
  9. logger = logging.getLogger(__name__)
  10. COLLECTION_NAME = "microbiome_kb"
  11. class MicrobiomeLoader:
  12. def __init__(self):
  13. self.embeddings = get_embeddings()
  14. self.vectorstore = Chroma(
  15. collection_name=COLLECTION_NAME,
  16. embedding_function=self.embeddings,
  17. persist_directory=settings.chroma_db_path,
  18. )
  19. self.splitter = get_knowledge_splitter()
  20. async def fetch_articles(self, keyword=None) -> list[dict]:
  21. client = httpx.AsyncClient(base_url=settings.java_base_url, timeout=120)
  22. try:
  23. if keyword:
  24. resp = await client.post("/api/microbiome/article/search", json={"keyword": keyword})
  25. data = resp.json()
  26. if data.get("code") == 200:
  27. return data.get("data", [])
  28. else:
  29. page = 1
  30. size = 50
  31. records = []
  32. while True:
  33. resp = await client.post("/api/microbiome/article/list",
  34. json={"page": page, "size": size, "keyword": "", "category": ""})
  35. data = resp.json()
  36. if data.get("code") != 200:
  37. break
  38. body = data.get("data", {}) or {}
  39. items = body.get("records", []) or []
  40. records.extend(items)
  41. total = body.get("total", 0)
  42. if page * size >= int(total or 0) or not items:
  43. break
  44. page += 1
  45. return records
  46. except Exception as e:
  47. logger.warning("获取菌群文章失败: %s", e)
  48. finally:
  49. await client.aclose()
  50. return []
  51. async def fetch_article(self, article_id: int) -> dict | None:
  52. client = httpx.AsyncClient(base_url=settings.java_base_url, timeout=120)
  53. try:
  54. resp = await client.post("/api/microbiome/article/detail", json={"id": article_id})
  55. data = resp.json()
  56. if data.get("code") == 200:
  57. return data.get("data")
  58. except Exception as e:
  59. logger.warning("获取菌群文章详情失败: %s", e)
  60. finally:
  61. await client.aclose()
  62. return None
  63. def _build_chunks(self, article: dict) -> list[Document]:
  64. content = f"{article.get('title', '')}\n\n{article.get('content', '')}"
  65. split_texts = self.splitter.split_text(content)
  66. category = article.get("category", "科普")
  67. tags = article.get("tags", "")
  68. docs = []
  69. for i, text in enumerate(split_texts):
  70. docs.append(Document(
  71. page_content=text,
  72. metadata={
  73. "id": f"article_{article['id']}_{i}",
  74. "postId": str(article["id"]),
  75. "title": article.get("title", ""),
  76. "category": category,
  77. "tags": tags,
  78. "source_url": article.get("sourceUrl", ""),
  79. "pub_date": article.get("pubDate", ""),
  80. "chunk_index": i,
  81. "chunk_total": len(split_texts),
  82. "source": "microbiome_kb",
  83. },
  84. ))
  85. return docs
  86. async def sync_all(self):
  87. articles = await self.fetch_articles()
  88. if not articles:
  89. logger.info("菌群知识库同步: 无文章")
  90. return
  91. try:
  92. existing_ids = self.vectorstore.get()["ids"]
  93. if existing_ids:
  94. self.vectorstore.delete(existing_ids)
  95. except Exception:
  96. pass
  97. all_docs = []
  98. for article in articles:
  99. all_docs.extend(self._build_chunks(article))
  100. for i in range(0, len(all_docs), 200):
  101. await self.vectorstore.aadd_documents(all_docs[i:i + 200])
  102. logger.info("菌群知识库全量同步完成: %d 篇文章, %d 个chunk", len(articles), len(all_docs))
  103. async def sync_one(self, article_id: int):
  104. article = await self.fetch_article(article_id)
  105. if not article:
  106. logger.warning("菌群文章不存在: %d", article_id)
  107. return
  108. prefix = f"article_{article_id}_"
  109. try:
  110. existing = self.vectorstore.get()
  111. ids_to_delete = [i for i in existing["ids"] if i.startswith(prefix)]
  112. if ids_to_delete:
  113. self.vectorstore.delete(ids_to_delete)
  114. except Exception:
  115. pass
  116. docs = self._build_chunks(article)
  117. if docs:
  118. await self.vectorstore.aadd_documents(docs)
  119. logger.info("菌群文章同步完成: %d (%s), %d 个chunk",
  120. article_id, article.get("title", ""), len(docs))