135 lines
4.5 KiB
Python
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()
|