| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129 |
- """菌群知识库加载器: 从 Java 后端同步文章到 ChromaDB"""
- import logging
- import httpx
- from langchain_chroma import Chroma
- from langchain_core.documents import Document
- from app.config import settings
- from app.rag.embeddings import get_embeddings
- from app.rag.splitter import get_knowledge_splitter
- logger = logging.getLogger(__name__)
- COLLECTION_NAME = "microbiome_kb"
- class MicrobiomeLoader:
- def __init__(self):
- self.embeddings = get_embeddings()
- self.vectorstore = Chroma(
- collection_name=COLLECTION_NAME,
- embedding_function=self.embeddings,
- persist_directory=settings.chroma_db_path,
- )
- self.splitter = get_knowledge_splitter()
- async def fetch_articles(self, keyword=None) -> list[dict]:
- client = httpx.AsyncClient(base_url=settings.java_base_url, timeout=120)
- try:
- if keyword:
- resp = await client.post("/api/microbiome/article/search", json={"keyword": keyword})
- data = resp.json()
- if data.get("code") == 200:
- return data.get("data", [])
- else:
- page = 1
- size = 50
- records = []
- while True:
- resp = await client.post("/api/microbiome/article/list",
- json={"page": page, "size": size, "keyword": "", "category": ""})
- data = resp.json()
- if data.get("code") != 200:
- break
- body = data.get("data", {}) or {}
- items = body.get("records", []) or []
- records.extend(items)
- total = body.get("total", 0)
- if page * size >= int(total or 0) or not items:
- break
- page += 1
- return records
- except Exception as e:
- logger.warning("获取菌群文章失败: %s", e)
- finally:
- await client.aclose()
- return []
- async def fetch_article(self, article_id: int) -> dict | None:
- client = httpx.AsyncClient(base_url=settings.java_base_url, timeout=120)
- try:
- resp = await client.post("/api/microbiome/article/detail", json={"id": article_id})
- data = resp.json()
- if data.get("code") == 200:
- return data.get("data")
- except Exception as e:
- logger.warning("获取菌群文章详情失败: %s", e)
- finally:
- await client.aclose()
- return None
- def _build_chunks(self, article: dict) -> list[Document]:
- content = f"{article.get('title', '')}\n\n{article.get('content', '')}"
- split_texts = self.splitter.split_text(content)
- category = article.get("category", "科普")
- tags = article.get("tags", "")
- docs = []
- for i, text in enumerate(split_texts):
- docs.append(Document(
- page_content=text,
- metadata={
- "id": f"article_{article['id']}_{i}",
- "postId": str(article["id"]),
- "title": article.get("title", ""),
- "category": category,
- "tags": tags,
- "source_url": article.get("sourceUrl", ""),
- "pub_date": article.get("pubDate", ""),
- "chunk_index": i,
- "chunk_total": len(split_texts),
- "source": "microbiome_kb",
- },
- ))
- return docs
- async def sync_all(self):
- articles = await self.fetch_articles()
- if not articles:
- logger.info("菌群知识库同步: 无文章")
- return
- try:
- existing_ids = self.vectorstore.get()["ids"]
- if existing_ids:
- self.vectorstore.delete(existing_ids)
- except Exception:
- pass
- all_docs = []
- for article in articles:
- all_docs.extend(self._build_chunks(article))
- for i in range(0, len(all_docs), 200):
- await self.vectorstore.aadd_documents(all_docs[i:i + 200])
- logger.info("菌群知识库全量同步完成: %d 篇文章, %d 个chunk", len(articles), len(all_docs))
- async def sync_one(self, article_id: int):
- article = await self.fetch_article(article_id)
- if not article:
- logger.warning("菌群文章不存在: %d", article_id)
- return
- prefix = f"article_{article_id}_"
- try:
- existing = self.vectorstore.get()
- ids_to_delete = [i for i in existing["ids"] if i.startswith(prefix)]
- if ids_to_delete:
- self.vectorstore.delete(ids_to_delete)
- except Exception:
- pass
- docs = self._build_chunks(article)
- if docs:
- await self.vectorstore.aadd_documents(docs)
- logger.info("菌群文章同步完成: %d (%s), %d 个chunk",
- article_id, article.get("title", ""), len(docs))
|