knowledge_sync.py 2.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677
  1. """知识库同步定时任务: 定期从 Java 拉取文章, 更新 ChromaDB"""
  2. from app.rag.loader import KnowledgeLoader
  3. from app.rag.splitter import get_knowledge_splitter
  4. from app.rag.embeddings import get_embeddings
  5. from app.config import settings
  6. from langchain_chroma import Chroma
  7. from langchain_core.documents import Document
  8. import logging
  9. import os
  10. import json
  11. logger = logging.getLogger(__name__)
  12. SYNC_STATE_FILE = os.path.join(settings.chroma_db_path, ".sync_state")
  13. def _load_sync_state() -> dict:
  14. try:
  15. if os.path.exists(SYNC_STATE_FILE):
  16. with open(SYNC_STATE_FILE) as f:
  17. return json.load(f)
  18. except Exception:
  19. pass
  20. return {"last_sync": "2000-01-01T00:00:00"}
  21. def _save_sync_state(state: dict):
  22. os.makedirs(os.path.dirname(SYNC_STATE_FILE), exist_ok=True)
  23. with open(SYNC_STATE_FILE, "w") as f:
  24. json.dump(state, f)
  25. async def sync_knowledge_base():
  26. """执行知识库同步 (全量+增量)"""
  27. loader = KnowledgeLoader()
  28. splitter = get_knowledge_splitter()
  29. embeddings = get_embeddings()
  30. vectorstore = Chroma(
  31. collection_name="cfc_knowledge",
  32. embedding_function=embeddings,
  33. persist_directory=settings.chroma_db_path,
  34. )
  35. state = _load_sync_state()
  36. articles = await loader.load_updated_since(state["last_sync"])
  37. if not articles:
  38. logger.info("知识库同步: 无更新内容")
  39. return
  40. docs = loader.format_for_indexing(articles)
  41. chunks = []
  42. for doc in docs:
  43. split_texts = splitter.split_text(doc["content"])
  44. for i, text in enumerate(split_texts):
  45. metadata = dict(doc["metadata"])
  46. metadata["chunk_index"] = i
  47. chunks.append(Document(page_content=text, metadata=metadata))
  48. if not chunks:
  49. logger.info("知识库同步: 无新增块")
  50. return
  51. BATCH_SIZE = 100
  52. for i in range(0, len(chunks), BATCH_SIZE):
  53. batch = chunks[i:i + BATCH_SIZE]
  54. await vectorstore.aadd_documents(batch)
  55. vectorstore.persist()
  56. logger.debug("知识库同步进度: %d/%d", i + len(batch), len(chunks))
  57. import datetime
  58. state["last_sync"] = datetime.datetime.now().isoformat()
  59. _save_sync_state(state)
  60. logger.info("知识库同步完成: 新增 %d 篇文章, %d 个块", len(articles), len(chunks))