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]