|
|
@@ -0,0 +1,95 @@
|
|
|
+from langchain_chroma import Chroma
|
|
|
+from langchain.retrievers import EnsembleRetriever
|
|
|
+from langchain_community.retrievers import BM25Retriever
|
|
|
+from langchain.retrievers.document_compressors import LLMChainExtractor
|
|
|
+from langchain.retrievers import ContextualCompressionRetriever
|
|
|
+from .embeddings import get_embeddings
|
|
|
+from app.config import settings
|
|
|
+from app.tools.java_client import JavaClient
|
|
|
+from typing import Optional
|
|
|
+import logging
|
|
|
+
|
|
|
+logger = logging.getLogger(__name__)
|
|
|
+
|
|
|
+
|
|
|
+class RagRetriever:
|
|
|
+ """混合检索器: ChromaDB 向量 + BM25 关键词 + 可选 LLM 压缩"""
|
|
|
+
|
|
|
+ def __init__(self, collection_name: str = "cfc_knowledge"):
|
|
|
+ embeddings = get_embeddings()
|
|
|
+ self.vectorstore = Chroma(
|
|
|
+ collection_name=collection_name,
|
|
|
+ embedding_function=embeddings,
|
|
|
+ persist_directory=settings.chroma_db_path,
|
|
|
+ )
|
|
|
+ self.java_client = JavaClient()
|
|
|
+ self._bm25_retriever = None
|
|
|
+ self._bm25_texts = []
|
|
|
+
|
|
|
+ async def initialize(self):
|
|
|
+ """从 Java 侧拉取知识库数据, 构建 BM25 索引"""
|
|
|
+ try:
|
|
|
+ articles = await self.java_client.get_published_articles()
|
|
|
+ self._bm25_texts = [
|
|
|
+ f"{a['title']} {a['summary']} {a.get('tags', '')}"
|
|
|
+ for a in articles
|
|
|
+ ]
|
|
|
+ if self._bm25_texts:
|
|
|
+ self._bm25_retriever = BM25Retriever.from_texts(
|
|
|
+ self._bm25_texts, metadatas=articles
|
|
|
+ )
|
|
|
+ logger.info("BM25 索引就绪: %d 条", len(self._bm25_texts))
|
|
|
+ except Exception as e:
|
|
|
+ logger.warning("BM25 初始化失败(不影响向量检索): %s", e)
|
|
|
+
|
|
|
+ async def retrieve(
|
|
|
+ self,
|
|
|
+ query: str,
|
|
|
+ filters: Optional[dict] = None,
|
|
|
+ k: int = 5,
|
|
|
+ use_compression: bool = False,
|
|
|
+ ) -> list[dict]:
|
|
|
+ """混合检索, 返回 [{content, metadata, score}]"""
|
|
|
+ retrievers = []
|
|
|
+
|
|
|
+ # 向量检索
|
|
|
+ vector_retriever = self.vectorstore.as_retriever(
|
|
|
+ search_kwargs={"k": k, "filter": filters}
|
|
|
+ )
|
|
|
+ retrievers.append(vector_retriever)
|
|
|
+
|
|
|
+ # BM25 检索
|
|
|
+ if self._bm25_retriever:
|
|
|
+ retrievers.append(self._bm25_retriever)
|
|
|
+
|
|
|
+ if len(retrievers) == 1:
|
|
|
+ docs = await retrievers[0].ainvoke(query)
|
|
|
+ else:
|
|
|
+ ensemble = EnsembleRetriever(
|
|
|
+ retrievers=retrievers, weights=[0.6, 0.4]
|
|
|
+ )
|
|
|
+ docs = await ensemble.ainvoke(query)
|
|
|
+
|
|
|
+ # 可选: LLM 压缩去噪
|
|
|
+ if use_compression and docs:
|
|
|
+ from langchain_openai import ChatOpenAI
|
|
|
+ llm = ChatOpenAI(
|
|
|
+ model=settings.llm_model,
|
|
|
+ api_key=settings.llm_api_key,
|
|
|
+ base_url=settings.llm_base_url,
|
|
|
+ )
|
|
|
+ compressor = LLMChainExtractor.from_llm(llm)
|
|
|
+ compression_retriever = ContextualCompressionRetriever(
|
|
|
+ base_compressor=compressor,
|
|
|
+ base_retriever=self.vectorstore.as_retriever(),
|
|
|
+ )
|
|
|
+ docs = await compression_retriever.ainvoke(query)
|
|
|
+
|
|
|
+ results = []
|
|
|
+ for doc in docs:
|
|
|
+ results.append({
|
|
|
+ "content": doc.page_content,
|
|
|
+ "metadata": doc.metadata,
|
|
|
+ "score": getattr(doc, "metadata", {}).get("score", 0),
|
|
|
+ })
|
|
|
+ return results[:k]
|