| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263 |
- """向量知识库管理接口 - 直接操作 ChromaDB"""
- from fastapi import APIRouter, HTTPException
- from pydantic import BaseModel
- from typing import Optional, List
- import logging
- import asyncio
- import os
- logger = logging.getLogger(__name__)
- router = APIRouter(prefix="/api/v1/knowledge-base", tags=["knowledge-base"])
- # 后台同步任务引用
- _sync_task = None
- class DocumentInfo(BaseModel):
- id: str
- source: str
- title: str
- content: str = ""
- content_preview: str = ""
- metadata: dict = {}
- class ListRequest(BaseModel):
- source: Optional[str] = None
- keyword: Optional[str] = None
- page: int = 1
- size: int = 20
- class UpdateDocumentRequest(BaseModel):
- content: str
- title: Optional[str] = None
- metadata_override: Optional[dict] = None
- @router.post("/stats")
- async def get_stats():
- """获取向量库统计信息"""
- try:
- from app.rag.retriever import RagRetriever
- from app.rag.embeddings import get_embeddings
- from app.config import settings
- from langchain_chroma import Chroma
- retriever = RagRetriever()
- collection = retriever.vectorstore._collection
- count = collection.count()
- # 按 source 统计(兼容新旧两种 metadata 字段:source / type)
- all_data = collection.get(include=["metadatas"])
- metas = all_data.get("metadatas") or []
- source_counts = {}
- type_counts = {}
- for meta in metas:
- if not meta:
- continue
- src = meta.get("source") or meta.get("type") or "unknown"
- source_counts[src] = source_counts.get(src, 0) + 1
- doc_type = meta.get("type") or meta.get("source") or "unknown"
- type_counts[doc_type] = type_counts.get(doc_type, 0) + 1
- return {
- "total": count,
- "by_source": source_counts,
- "by_type": type_counts,
- "last_sync": _load_sync_state(),
- }
- except Exception as e:
- logger.error("获取向量库统计失败: %s", e)
- raise HTTPException(status_code=500, detail=str(e))
- @router.post("/documents")
- async def list_documents(req: ListRequest):
- """分页列出向量库文档"""
- try:
- from app.rag.retriever import RagRetriever
- from app.rag.embeddings import get_embeddings
- from app.config import settings
- from langchain_chroma import Chroma
- retriever = RagRetriever()
- collection = retriever.vectorstore._collection
- # 构建过滤条件(兼容新旧两种 metadata 字段:source / type)
- where_filter = None
- if req.source:
- where_filter = {"$or": [{"source": req.source}, {"type": req.source}]}
- # 使用 get() 分页获取,避免加载全部 embedding
- limit = req.size
- offset = (req.page - 1) * req.size
- result = collection.get(
- where=where_filter,
- include=["documents", "metadatas"],
- limit=limit,
- offset=offset,
- )
- ids = result.get("ids") or []
- docs = result.get("documents") or []
- metas = result.get("metadatas") or []
- documents = []
- for i, doc_id in enumerate(ids):
- meta = metas[i] if i < len(metas) else {}
- content = docs[i] if i < len(docs) else ""
- # 截取内容预览
- preview = content[:200] + "..." if len(content) > 200 else content
- documents.append(DocumentInfo(
- id=doc_id,
- source=meta.get("source") or meta.get("type") or "unknown",
- title=meta.get("title", ""),
- content=content,
- content_preview=preview,
- metadata=meta,
- ))
- # 获取总数
- count_result = collection.get(
- where=where_filter,
- include=[],
- )
- total = len(count_result.get("ids", []))
- return {
- "total": total,
- "page": req.page,
- "size": req.size,
- "documents": documents,
- }
- except Exception as e:
- logger.error("列出文档失败: %s", e)
- raise HTTPException(status_code=500, detail=str(e))
- @router.delete("/source/{source_prefix}")
- async def delete_by_source(source_prefix: str):
- """删除指定 source 前缀的所有文档"""
- try:
- from app.rag.retriever import RagRetriever
- from app.rag.embeddings import get_embeddings
- from app.config import settings
- from langchain_chroma import Chroma
- retriever = RagRetriever()
- collection = retriever.vectorstore._collection
- # 兼容新旧两种 metadata 字段:source / type
- existing = collection.get(
- where={"$or": [{"source": source_prefix}, {"type": source_prefix}]},
- include=[],
- )
- ids = existing.get("ids", []) or []
- if not ids:
- return {"deleted": 0, "message": f"无 {source_prefix} 文档"}
- collection.delete(ids=ids)
- logger.info("已删除 source=%s 的 %d 个文档", source_prefix, len(ids))
- return {"deleted": len(ids), "source": source_prefix}
- except Exception as e:
- logger.error("删除文档失败: %s", e)
- raise HTTPException(status_code=500, detail=str(e))
- @router.post("/sync")
- async def trigger_sync():
- """触发一次完整知识库同步"""
- global _sync_task
- try:
- # 如果已有同步任务在运行,返回提示
- if _sync_task and not _sync_task.done():
- return {"status": "running", "message": "同步任务已在运行中"}
- _sync_task = asyncio.create_task(_run_sync())
- return {"status": "started", "message": "已启动知识库同步任务"}
- except Exception as e:
- logger.error("触发同步失败: %s", e)
- raise HTTPException(status_code=500, detail=str(e))
- async def _run_sync():
- """执行完整同步"""
- try:
- from app.tasks.knowledge_sync import sync_knowledge_base
- await sync_knowledge_base()
- logger.info("知识库同步完成")
- except Exception as e:
- logger.error("知识库同步失败: %s", e)
- @router.put("/document/{doc_id}")
- async def update_document(doc_id: str, req: UpdateDocumentRequest):
- """更新单个文档内容(先获取原 metadata,再删除重建,保持 ID 不变)"""
- try:
- from app.rag.retriever import RagRetriever
- from app.rag.embeddings import get_embeddings
- from app.config import settings
- from langchain_chroma import Chroma
- retriever = RagRetriever()
- collection = retriever.vectorstore._collection
- # 先获取原文档 metadata,避免删除后丢失
- old_result = collection.get(ids=[doc_id], include=["metadatas", "documents"])
- old_metas = old_result.get("metadatas") or []
- old_docs = old_result.get("documents") or []
- if not old_metas:
- raise HTTPException(status_code=404, detail="文档不存在")
- original_meta = old_metas[0]
- original_content = old_docs[0] if old_docs else ""
- # 删除旧文档
- collection.delete(ids=[doc_id])
- logger.info("已删除旧文档: %s", doc_id)
- # 构建新文档:保留原有 metadata,仅覆盖 content/title
- meta = req.metadata_override or {}
- # 保留原有 source/type/title 等字段,除非显式覆盖
- for k, v in original_meta.items():
- if k not in meta:
- meta[k] = v
- if req.title:
- meta["title"] = req.title
- # 使用原 ID 重新写入(显式传 ids,避免生成新 ID)
- from langchain_core.documents import Document
- doc = Document(page_content=req.content, metadata=meta)
- await retriever.vectorstore.aadd_documents([doc], ids=[doc_id])
- logger.info("已更新文档: %s", doc_id)
- return {"success": True, "message": "文档已更新", "doc_id": doc_id}
- except HTTPException:
- raise
- except Exception as e:
- logger.error("更新文档失败: %s", e)
- raise HTTPException(status_code=500, detail=str(e))
- def _load_sync_state() -> str:
- """加载上次同步时间"""
- try:
- import json
- from app.config import settings
- import os
- sync_file = os.path.join(settings.chroma_db_path, ".sync_state")
- if os.path.exists(sync_file):
- with open(sync_file) as f:
- state = json.load(f)
- return state.get("last_sync", "未知")
- except Exception:
- pass
- return "未知"
|