knowledge_base.py 8.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263
  1. """向量知识库管理接口 - 直接操作 ChromaDB"""
  2. from fastapi import APIRouter, HTTPException
  3. from pydantic import BaseModel
  4. from typing import Optional, List
  5. import logging
  6. import asyncio
  7. import os
  8. logger = logging.getLogger(__name__)
  9. router = APIRouter(prefix="/api/v1/knowledge-base", tags=["knowledge-base"])
  10. # 后台同步任务引用
  11. _sync_task = None
  12. class DocumentInfo(BaseModel):
  13. id: str
  14. source: str
  15. title: str
  16. content: str = ""
  17. content_preview: str = ""
  18. metadata: dict = {}
  19. class ListRequest(BaseModel):
  20. source: Optional[str] = None
  21. keyword: Optional[str] = None
  22. page: int = 1
  23. size: int = 20
  24. class UpdateDocumentRequest(BaseModel):
  25. content: str
  26. title: Optional[str] = None
  27. metadata_override: Optional[dict] = None
  28. @router.post("/stats")
  29. async def get_stats():
  30. """获取向量库统计信息"""
  31. try:
  32. from app.rag.retriever import RagRetriever
  33. from app.rag.embeddings import get_embeddings
  34. from app.config import settings
  35. from langchain_chroma import Chroma
  36. retriever = RagRetriever()
  37. collection = retriever.vectorstore._collection
  38. count = collection.count()
  39. # 按 source 统计(兼容新旧两种 metadata 字段:source / type)
  40. all_data = collection.get(include=["metadatas"])
  41. metas = all_data.get("metadatas") or []
  42. source_counts = {}
  43. type_counts = {}
  44. for meta in metas:
  45. if not meta:
  46. continue
  47. src = meta.get("source") or meta.get("type") or "unknown"
  48. source_counts[src] = source_counts.get(src, 0) + 1
  49. doc_type = meta.get("type") or meta.get("source") or "unknown"
  50. type_counts[doc_type] = type_counts.get(doc_type, 0) + 1
  51. return {
  52. "total": count,
  53. "by_source": source_counts,
  54. "by_type": type_counts,
  55. "last_sync": _load_sync_state(),
  56. }
  57. except Exception as e:
  58. logger.error("获取向量库统计失败: %s", e)
  59. raise HTTPException(status_code=500, detail=str(e))
  60. @router.post("/documents")
  61. async def list_documents(req: ListRequest):
  62. """分页列出向量库文档"""
  63. try:
  64. from app.rag.retriever import RagRetriever
  65. from app.rag.embeddings import get_embeddings
  66. from app.config import settings
  67. from langchain_chroma import Chroma
  68. retriever = RagRetriever()
  69. collection = retriever.vectorstore._collection
  70. # 构建过滤条件(兼容新旧两种 metadata 字段:source / type)
  71. where_filter = None
  72. if req.source:
  73. where_filter = {"$or": [{"source": req.source}, {"type": req.source}]}
  74. # 使用 get() 分页获取,避免加载全部 embedding
  75. limit = req.size
  76. offset = (req.page - 1) * req.size
  77. result = collection.get(
  78. where=where_filter,
  79. include=["documents", "metadatas"],
  80. limit=limit,
  81. offset=offset,
  82. )
  83. ids = result.get("ids") or []
  84. docs = result.get("documents") or []
  85. metas = result.get("metadatas") or []
  86. documents = []
  87. for i, doc_id in enumerate(ids):
  88. meta = metas[i] if i < len(metas) else {}
  89. content = docs[i] if i < len(docs) else ""
  90. # 截取内容预览
  91. preview = content[:200] + "..." if len(content) > 200 else content
  92. documents.append(DocumentInfo(
  93. id=doc_id,
  94. source=meta.get("source") or meta.get("type") or "unknown",
  95. title=meta.get("title", ""),
  96. content=content,
  97. content_preview=preview,
  98. metadata=meta,
  99. ))
  100. # 获取总数
  101. count_result = collection.get(
  102. where=where_filter,
  103. include=[],
  104. )
  105. total = len(count_result.get("ids", []))
  106. return {
  107. "total": total,
  108. "page": req.page,
  109. "size": req.size,
  110. "documents": documents,
  111. }
  112. except Exception as e:
  113. logger.error("列出文档失败: %s", e)
  114. raise HTTPException(status_code=500, detail=str(e))
  115. @router.delete("/source/{source_prefix}")
  116. async def delete_by_source(source_prefix: str):
  117. """删除指定 source 前缀的所有文档"""
  118. try:
  119. from app.rag.retriever import RagRetriever
  120. from app.rag.embeddings import get_embeddings
  121. from app.config import settings
  122. from langchain_chroma import Chroma
  123. retriever = RagRetriever()
  124. collection = retriever.vectorstore._collection
  125. # 兼容新旧两种 metadata 字段:source / type
  126. existing = collection.get(
  127. where={"$or": [{"source": source_prefix}, {"type": source_prefix}]},
  128. include=[],
  129. )
  130. ids = existing.get("ids", []) or []
  131. if not ids:
  132. return {"deleted": 0, "message": f"无 {source_prefix} 文档"}
  133. collection.delete(ids=ids)
  134. logger.info("已删除 source=%s 的 %d 个文档", source_prefix, len(ids))
  135. return {"deleted": len(ids), "source": source_prefix}
  136. except Exception as e:
  137. logger.error("删除文档失败: %s", e)
  138. raise HTTPException(status_code=500, detail=str(e))
  139. @router.post("/sync")
  140. async def trigger_sync():
  141. """触发一次完整知识库同步"""
  142. global _sync_task
  143. try:
  144. # 如果已有同步任务在运行,返回提示
  145. if _sync_task and not _sync_task.done():
  146. return {"status": "running", "message": "同步任务已在运行中"}
  147. _sync_task = asyncio.create_task(_run_sync())
  148. return {"status": "started", "message": "已启动知识库同步任务"}
  149. except Exception as e:
  150. logger.error("触发同步失败: %s", e)
  151. raise HTTPException(status_code=500, detail=str(e))
  152. async def _run_sync():
  153. """执行完整同步"""
  154. try:
  155. from app.tasks.knowledge_sync import sync_knowledge_base
  156. await sync_knowledge_base()
  157. logger.info("知识库同步完成")
  158. except Exception as e:
  159. logger.error("知识库同步失败: %s", e)
  160. @router.put("/document/{doc_id}")
  161. async def update_document(doc_id: str, req: UpdateDocumentRequest):
  162. """更新单个文档内容(先获取原 metadata,再删除重建,保持 ID 不变)"""
  163. try:
  164. from app.rag.retriever import RagRetriever
  165. from app.rag.embeddings import get_embeddings
  166. from app.config import settings
  167. from langchain_chroma import Chroma
  168. retriever = RagRetriever()
  169. collection = retriever.vectorstore._collection
  170. # 先获取原文档 metadata,避免删除后丢失
  171. old_result = collection.get(ids=[doc_id], include=["metadatas", "documents"])
  172. old_metas = old_result.get("metadatas") or []
  173. old_docs = old_result.get("documents") or []
  174. if not old_metas:
  175. raise HTTPException(status_code=404, detail="文档不存在")
  176. original_meta = old_metas[0]
  177. original_content = old_docs[0] if old_docs else ""
  178. # 删除旧文档
  179. collection.delete(ids=[doc_id])
  180. logger.info("已删除旧文档: %s", doc_id)
  181. # 构建新文档:保留原有 metadata,仅覆盖 content/title
  182. meta = req.metadata_override or {}
  183. # 保留原有 source/type/title 等字段,除非显式覆盖
  184. for k, v in original_meta.items():
  185. if k not in meta:
  186. meta[k] = v
  187. if req.title:
  188. meta["title"] = req.title
  189. # 使用原 ID 重新写入(显式传 ids,避免生成新 ID)
  190. from langchain_core.documents import Document
  191. doc = Document(page_content=req.content, metadata=meta)
  192. await retriever.vectorstore.aadd_documents([doc], ids=[doc_id])
  193. logger.info("已更新文档: %s", doc_id)
  194. return {"success": True, "message": "文档已更新", "doc_id": doc_id}
  195. except HTTPException:
  196. raise
  197. except Exception as e:
  198. logger.error("更新文档失败: %s", e)
  199. raise HTTPException(status_code=500, detail=str(e))
  200. def _load_sync_state() -> str:
  201. """加载上次同步时间"""
  202. try:
  203. import json
  204. from app.config import settings
  205. import os
  206. sync_file = os.path.join(settings.chroma_db_path, ".sync_state")
  207. if os.path.exists(sync_file):
  208. with open(sync_file) as f:
  209. state = json.load(f)
  210. return state.get("last_sync", "未知")
  211. except Exception:
  212. pass
  213. return "未知"