| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687 |
- from langchain_chroma import Chroma
- from langchain_openai import ChatOpenAI, OpenAIEmbeddings
- 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"""
- def __init__(self, collection_name: str = "cfc_knowledge"):
- embeddings = OpenAIEmbeddings(
- model=settings.embedding_model,
- api_key=settings.effective_embedding_api_key,
- base_url=settings.effective_embedding_base_url,
- )
- self.vectorstore = Chroma(
- collection_name=collection_name,
- embedding_function=embeddings,
- persist_directory=settings.chroma_db_path,
- )
- self.java_client = JavaClient()
- async def initialize(self):
- """从 Java 侧拉取知识库,更新到向量库"""
- try:
- articles = await self.java_client.get_published_articles()
- if articles:
- # 构建文本和内容元数据
- texts = [f"{a['title']} {a.get('summary', '')} {' '.join(a.get('tags', []))}"
- for a in articles]
- metadatas = [
- {
- "id": a["id"],
- "title": a["title"],
- "summary": a.get("summary", ""),
- "tags": a.get("tags", []),
- "type": "article",
- }
- for a in articles
- ]
- # 添加或更新向量数据库
- await self.vectorstore.aadd_texts(texts=texts, metadatas=metadatas)
- logger.info("知识库向量化完成:%d 篇文章", len(articles))
- except Exception as e:
- logger.warning("知识库初始化失败:%s", e)
- async def retrieve(
- self,
- query: str,
- filters: Optional[dict] = None,
- k: int = 5,
- ) -> list[dict]:
- """向量相似度检索"""
- results = []
-
- # 向量检索
- filter_query = {"user_id": filters["user_id"]} if filters and "user_id" in filters else None
-
- docs = self.vectorstore.similarity_search(
- query,
- k=k * 2, # 多取一些
- filter=filter_query,
- )
-
- # 格式化结果并去重
- 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,
- })
-
- if len(results) >= k:
- break
-
- return results[:k]
|