""" 知识库 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()