| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596 |
- from langchain_chroma import Chroma
- from langchain_openai import ChatOpenAI
- from app.config import settings
- from app.rag.embeddings import get_embeddings
- 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 = get_embeddings()
- 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 not articles:
- logger.info("知识库初始化:无已发布文章")
- return
- # 防御式处理:容忍缺字段 / tags 为字符串或列表
- texts = []
- metadatas = []
- for a in articles:
- if not isinstance(a, dict):
- continue
- title = a.get("title", "")
- if not title:
- continue
- tags = a.get("tags", []) or []
- tags_str = tags if isinstance(tags, str) else " ".join(str(t) for t in tags)
- texts.append(f"{title} {a.get('summary', '')} {tags_str}")
- metadatas.append({
- "id": str(a["id"]),
- "title": title,
- "summary": a.get("summary", ""),
- "tags": tags,
- "type": "article",
- })
- if not texts:
- logger.info("知识库初始化:无有效文章内容")
- return
- # 添加或更新向量数据库
- await self.vectorstore.aadd_texts(texts=texts, metadatas=metadatas)
- logger.info("知识库向量化完成:%d 篇文章", len(metadatas))
- 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]
|