| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109 |
- 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 langchain_openai import ChatOpenAI
- 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:
- """升级版混合检索器: 向量 + 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: Optional[BM25Retriever] = None
- self._bm25_texts: list[str] = []
- 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 = True,
- ) -> list[dict]:
- """混合检索 + 可选 LLM 压缩重排序"""
- retrievers = []
- # 1. 向量检索 (多取一些方便后续 ensemble 排序)
- vector_retriever = self.vectorstore.as_retriever(
- search_kwargs={"k": k * 2, "filter": filters},
- )
- retrievers.append(vector_retriever)
- # 2. BM25 关键词检索
- if self._bm25_retriever:
- bm25_k = self._bm25_retriever.k
- self._bm25_retriever.k = k * 2
- retrievers.append(self._bm25_retriever)
- self._bm25_retriever.k = bm25_k
- if len(retrievers) == 1:
- docs = await retrievers[0].ainvoke(query)
- ensemble = retrievers[0]
- else:
- ensemble = EnsembleRetriever(
- retrievers=retrievers,
- weights=[0.6, 0.4],
- )
- docs = await ensemble.ainvoke(query)
- # 3. LLM 压缩 (剔除不相关内容)
- if use_compression and docs:
- llm = ChatOpenAI(
- model=settings.llm_model,
- api_key=settings.llm_api_key,
- base_url=settings.llm_base_url,
- temperature=0,
- )
- compressor = LLMChainExtractor.from_llm(llm)
- compression_retriever = ContextualCompressionRetriever(
- base_compressor=compressor,
- base_retriever=ensemble if len(retrievers) > 1 else retrievers[0],
- )
- docs = await compression_retriever.ainvoke(query)
- # 4. 格式化为统一输出 + 去重
- results = []
- seen = set()
- for doc in docs:
- content_hash = hash(doc.page_content[:100])
- if content_hash in seen:
- continue
- seen.add(content_hash)
- results.append({
- "content": doc.page_content,
- "metadata": doc.metadata,
- "score": doc.metadata.get("score", 0) if hasattr(doc, "metadata") else 0,
- })
- return results[:k]
|