nex_docus/backend/app/services/rag_service.py

135 lines
4.5 KiB
Python

"""
知识库 RAG 服务 - 向量检索 + 大模型生成
"""
from typing import List, Dict, Any, Optional
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
import json
from app.models.document_vector import DocumentVector
from app.models.llm_model_config import LLMModelConfig
from app.services.llm_provider_service import LLMProviderService
class RAGService:
"""RAG 知识库检索和生成服务"""
@staticmethod
async def retrieve_documents(
db: AsyncSession,
project_id: int,
query_vector: List[float],
top_k: int = 5,
similarity_threshold: float = 0.3,
) -> List[Dict[str, Any]]:
"""向量检索相关文档"""
if not query_vector or len(query_vector) == 0:
return []
stmt = select(DocumentVector).where(
DocumentVector.project_id == project_id,
DocumentVector.vector.isnot(None),
)
result = await db.execute(stmt)
all_vectors = result.scalars().all()
if not all_vectors:
return []
scored_docs = []
for doc in all_vectors:
try:
doc_vector = json.loads(doc.vector) if isinstance(doc.vector, str) else doc.vector
similarity = RAGService._cosine_similarity(query_vector, doc_vector)
if similarity >= similarity_threshold:
scored_docs.append({
"doc": doc,
"similarity": similarity,
"file_path": doc.file_path,
"content_preview": doc.content[:500] if doc.content else "",
})
except Exception:
continue
scored_docs.sort(key=lambda x: x["similarity"], reverse=True)
return scored_docs[:top_k]
@staticmethod
def _cosine_similarity(vec1: List[float], vec2: List[float]) -> float:
"""计算余弦相似度"""
if len(vec1) != len(vec2) or len(vec1) == 0:
return 0.0
dot_product = sum(a * b for a, b in zip(vec1, vec2))
magnitude1 = sum(a * a for a in vec1) ** 0.5
magnitude2 = sum(b * b for b in vec2) ** 0.5
if magnitude1 == 0 or magnitude2 == 0:
return 0.0
return dot_product / (magnitude1 * magnitude2)
@staticmethod
async def generate_response(
db: AsyncSession,
query: str,
query_vector: List[float],
project_id: int,
llm_config_id: int,
retrieved_docs: List[Dict[str, Any]],
conversation_history: List[Dict[str, str]],
) -> str:
"""基于检索文档生成对话回复"""
stmt = select(LLMModelConfig).where(LLMModelConfig.config_id == llm_config_id)
result = await db.execute(stmt)
llm_config = result.scalar_one_or_none()
if not llm_config:
raise ValueError(f"LLM配置不存在: {llm_config_id}")
context_text = RAGService._build_context(retrieved_docs)
system_prompt = f"""你是一个知识库助手。基于用户提供的知识库文档,回答用户的问题。
如果知识库中没有相关信息,请明确说明。
知识库文档内容:
{context_text}
请基于以上知识库内容,用中文回答用户的问题。"""
messages = list(conversation_history)
messages.append({"role": "user", "content": query})
response_text = await LLMProviderService.generate_text(
provider=llm_config.provider,
endpoint_url=llm_config.endpoint_url,
api_key=llm_config.api_key,
llm_model_name=llm_config.llm_model_name,
timeout=llm_config.llm_timeout,
temperature=float(llm_config.llm_temperature),
top_p=float(llm_config.llm_top_p),
max_tokens=llm_config.llm_max_tokens,
system_prompt=system_prompt,
messages=messages,
)
return response_text
@staticmethod
def _build_context(retrieved_docs: List[Dict[str, Any]]) -> str:
"""构建上下文文本"""
if not retrieved_docs:
return "(未找到相关知识库内容)"
context_parts = []
for i, doc_info in enumerate(retrieved_docs, 1):
context_parts.append(
f"文档 {i}: {doc_info['file_path']} (相似度: {doc_info['similarity']:.2%})\n"
f"{doc_info['content_preview']}..."
)
return "\n\n".join(context_parts)
rag_service = RAGService()