From 99115692e5d1348c0d4e7edc0dca16be3dee9b49 Mon Sep 17 00:00:00 2001 From: linlin Date: Fri, 13 Mar 2026 13:41:16 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E4=BC=98=E5=8C=96RAG=E6=A3=80?= =?UTF-8?q?=E7=B4=A2=E7=AD=96=E7=95=A5=EF=BC=8C=E6=8F=90=E5=8D=87=E6=9F=A5?= =?UTF-8?q?=E8=AF=A2=E6=80=A7=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 小模型返回置信度和提取的filter - 根据置信度分流到不同collection - 用纯向量检索替代融合检索 - 删除了意图识别步骤,减少LLM调用 - 新增配置项支持阈值配置 --- config.py | 5 + rag/rag_engine.py | 390 +++++++++++++++++++++++----------------------- 2 files changed, 202 insertions(+), 193 deletions(-) diff --git a/config.py b/config.py index df8b372..4a51d49 100644 --- a/config.py +++ b/config.py @@ -185,6 +185,11 @@ class Settings(BaseSettings): CHUNK_OVERLAP: int = 200 TOP_K: int = 5 # Number of documents to retrieve + # RAG Query Classification Settings + CODE_RELATED_THRESHOLD_LOW: float = 0.3 # 低于此值认为是非代码问题 + CODE_RELATED_THRESHOLD_HIGH: float = 0.7 # 高于此值认为是代码相关问题 + FILTER_METADATA_FIELDS: str = "class_name,func_name,file_path" # 从query中提取的metadata字段 + # Sync Settings SYNC_INTERVAL: int = 300 # Sync interval in seconds AUTO_SYNC: bool = True diff --git a/rag/rag_engine.py b/rag/rag_engine.py index 6d2fe48..f7be8bd 100644 --- a/rag/rag_engine.py +++ b/rag/rag_engine.py @@ -6,7 +6,7 @@ from llama_index.core.response_synthesizers import ResponseMode from llama_index.core.base.response.schema import StreamingResponse from llama_index.llms.ollama import Ollama from llama_index.core import PromptTemplate -from typing import AsyncIterator, Optional, Tuple +from typing import AsyncIterator, Optional, Tuple, Dict, Any import asyncio import httpx from loguru import logger @@ -187,38 +187,70 @@ class RAGEngine: except: pass - async def is_code_related(self, query: str, history: str) -> bool: + async def is_code_related(self, query: str, history: str) -> Dict[str, Any]: """ - 使用小型模型判断是否是代码相关问题 + 使用小型模型判断是否是代码相关问题,并提取可能的filter Args: query: 用户查询字符串 history: 对话历史字符串 Returns: - bool: 是否是代码相关问题 + dict: 包含 confidence (float) 和 filters (dict) """ try: - prompt = f"""你是一个分类器,需要判断用户的问题是否与代码相关。 + # 获取需要提取的metadata字段 + filter_fields = settings.FILTER_METADATA_FIELDS.split(',') + filter_fields_str = ', '.join(filter_fields) + + prompt = f"""你是一个分类器,需要分析用户的问题。 用户问题:{query} 对话历史:{history} -请仅回答 '是' 或 '否',不要添加任何其他内容。""" +请分析这个问题并返回JSON格式的分析结果: +{{ + "confidence": 0.0-1.0之间的置信度,1表示完全确定是代码相关问题,0表示完全确定不是代码相关问题, + "filters": {{}} 或 {{"字段名": "从问题中提取的值"}},只有当问题中明确提到"xxx文件/xxx函数/xxx类"时才提取 +}} + +提取规则: +- 只有当用户明确说明了"xxx文件"、"xxx函数"、"xxx类"时才提取对应的metadata +- class_name: 用户提到具体类名时提取,如"User类"、"ArrayList" +- func_name: 用户提到具体函数名时提取,如"main函数"、"delete方法" +- file_path: 用户提到具体文件时提取,如"utils.py"、"config.json" + +请仅返回JSON,不要添加任何其他内容。""" response = await self.small_llm.acomplete(prompt=prompt) - answer = response.text.strip().lower() + response_text = response.text.strip() - logger.info(f"小型模型代码问题判断结果: {answer}") + logger.info(f"小型模型分析结果: {response_text}") - return answer == '是' + # 尝试解析JSON + import json + json_start = response_text.find('{') + json_end = response_text.rfind('}') + if json_start != -1 and json_end != -1: + json_str = response_text[json_start:json_end + 1] + result = json.loads(json_str) + confidence = float(result.get('confidence', 0.5)) + filters = result.get('filters', {}) + # 过滤掉空的filter + filters = {k: v for k, v in filters.items() if v} + logger.info(f"解析成功: confidence={confidence}, filters={filters}") + return {'confidence': confidence, 'filters': filters} + else: + logger.warning(f"无法解析JSON,使用默认结果") + return {'confidence': 0.5, 'filters': {}} + except Exception as e: logger.error(f"判断代码问题时出错: {e}") - # 出错时默认返回 False - return False + # 出错时返回中间值,不过滤 + return {'confidence': 0.5, 'filters': {}} - async def query_stream(self, query: str, history: str, top_k: Optional[int] = None, filters: Optional[dict] = None) -> AsyncIterator[str]: + async def query_stream(self, query: str, history: str, top_k: Optional[int] = None) -> AsyncIterator[str]: """ Query the RAG system and stream the response @@ -226,113 +258,89 @@ class RAGEngine: query: User query string history: Chat history string top_k: Number of documents to retrieve (optional) - filters: Metadata filters for pre-filtering documents (optional) Yields: Response text chunks """ try: - # 1. 使用小型模型判断是否是代码相关问题 - logger.info("使用小型模型判断是否是代码相关问题") - is_code = await self.is_code_related(query, history) - logger.info(f"是否是代码相关问题: {is_code}") - #TODO: 分支分流提前,如果是代码相关,走下面的代码相关检索和问答处理,否则做简单检索和问答 + # 1. 使用小型模型分析问题,返回置信度和可能的filter + logger.info("使用小型模型分析问题") + analysis_result = await self.is_code_related(query, history) + confidence = analysis_result['confidence'] + filters = analysis_result['filters'] + logger.info(f"分析结果: confidence={confidence}, filters={filters}") - # 2. 处理查询(整合意图识别、filter生成和查询转换) - logger.info(f"开始查询处理: {query}") - process_result = self.query_processor.process_query(query, history) + # 2. 根据置信度选择collection和检索策略 + threshold_low = settings.CODE_RELATED_THRESHOLD_LOW + threshold_high = settings.CODE_RELATED_THRESHOLD_HIGH - # 提取结果 - intent_result = process_result['intent'] - # 使用意图识别结果更新 is_code 变量,因为意图识别结果更准确 - is_code = intent_result.is_code_related - logger.info(f"根据意图识别结果更新 is_code: {is_code}") - filters = process_result['filters'] - transformed_queries = [process_result['transformed']['rewritten']] # + process_result['transformed']['sub_queries'] #NOTE: 太多queries,检索费时 + if confidence < threshold_low: + collection_key = 'non_code' + use_advanced_prompt = False + logger.info(f"置信度 {confidence} < {threshold_low},判定为非代码问题,使用non_code collection") + elif confidence > threshold_high: + collection_key = 'code' + use_advanced_prompt = True + logger.info(f"置信度 {confidence} > {threshold_high},判定为代码相关问题,使用code collection") + else: + collection_key = None + use_advanced_prompt = False + logger.info(f"置信度 {threshold_low} <= {confidence} <= {threshold_high},通用问题,搜索所有collection") - logger.info(f"代码意图识别结果: {intent_result.category}") - logger.info(f"过滤条件: {filters}") - logger.info(f"查询转换结果: {transformed_queries}") + # 3. 使用纯向量检索(替代融合检索) + logger.info("使用纯向量检索") + k = top_k or settings.TOP_K - # 4. 使用融合检索,对每个转换后的查询进行检索 - logger.info("使用融合检索策略") - all_hybrid_results = [] - # 根据是否是代码问题选择合适的 collection - collection_key = 'code' if is_code else 'non_code' - logger.info(f"使用 collection: {collection_key}") - for transformed_query in transformed_queries: - hybrid_results = await self.vector_store_manager.ahybrid_search( - query=transformed_query, - top_k=settings.TOP_K, - filters=filters, - collection_key=collection_key - ) - all_hybrid_results.extend(hybrid_results) + retriever = self.vector_store_manager.get_retriever( + top_k=k * 2, + filters=filters if filters else None, + collection_key=collection_key + ) + + if isinstance(retriever, list): + all_nodes = [] + for key, r in retriever: + nodes = r.retrieve(query) + for node in nodes: + actual_node = node.node if hasattr(node, 'node') else node + if hasattr(actual_node, 'metadata'): + actual_node.metadata['_collection_key'] = key + all_nodes.extend(nodes) + vector_nodes = all_nodes + else: + vector_nodes = retriever.retrieve(query) # 去重并按得分排序 seen_doc_ids = set() - unique_hybrid_results = [] - for doc_id, score, metadata in all_hybrid_results: - if doc_id not in seen_doc_ids: + unique_results = [] + for node in vector_nodes: + doc_id = getattr(node, 'id_', None) or getattr(node, 'node_id', None) + if doc_id and doc_id not in seen_doc_ids: seen_doc_ids.add(doc_id) - unique_hybrid_results.append((doc_id, score, metadata)) + unique_results.append(node) - # 按得分排序 - unique_hybrid_results.sort(key=lambda x: x[1], reverse=True) + unique_results.sort(key=lambda x: getattr(x, 'score', 0), reverse=True) + vector_nodes = unique_results[:k] - # 限制结果数量 - hybrid_results = unique_hybrid_results[:top_k or settings.TOP_K] - # 美化输出融合检索结果 - logger.info("融合检索结果:") - for i, (doc_id, score, metadata) in enumerate(hybrid_results): - func_name = metadata.get('func_name', 'N/A') - file_path = metadata.get('file_path', 'N/A') - lang = metadata.get('lang', 'N/A') - logger.info(f" [{i+1}] 相似度: {score:.4f}") - logger.info(f" 函数: {func_name}") - logger.info(f" 文件: {file_path}") - logger.info(f" 语言: {lang}") - logger.info(" " + "-" * 50) - # 根据文档ID获取完整的节点信息 - retrieved_nodes = [] - for doc_id, score, metadata in hybrid_results: - # 从向量存储中获取文档内容 - doc_chunks = self.vector_store_manager.get_document_by_id(doc_id) - for chunk in doc_chunks: - # 创建节点对象 - from llama_index.core.schema import TextNode - node = TextNode( - text=chunk['text'], - node_id=chunk['id'], - metadata=chunk['metadata'] - ) - retrieved_nodes.append(node) + logger.info(f"检索到 {len(vector_nodes)} 个结果") - # 3. 构建上下文 - context_parts = [] - max_nodes = top_k or settings.TOP_K - - # 构建原始的检索结果列表,包含 text 和 metadata + # 4. 构建检索结果 retrieved_results = [] - for i, node in enumerate(retrieved_nodes[:max_nodes], 1): - # 优先使用metadata中的func_body字段 - if hasattr(node, 'metadata') and 'func_body' in node.metadata: - text = node.metadata['func_body'] + for i, node in enumerate(vector_nodes, 1): + metadata = getattr(node, 'metadata', {}) + if 'func_body' in metadata: + text = metadata['func_body'] else: text = node.text if hasattr(node, 'text') else str(node) text = text.strip() - # 获取原始 metadata - metadata = getattr(node, 'metadata', {}) - - # 保存原始的 text 和 metadata retrieved_results.append({ 'id': i, 'text': text, 'metadata': metadata }) - # 构建简单的上下文字符串,只包含基本信息 + # 5. 构建上下文 context_parts = [] for result in retrieved_results: metadata = result['metadata'] @@ -348,44 +356,41 @@ class RAGEngine: context_parts.append(context_part) context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息" - logger.info(f"上下文: {context_str}") - # 4. 生成优化的Prompt - logger.info("生成优化的Prompt") - if intent_result.is_code_related: - # 对于代码相关问题,使用代码专用Prompt - filled_prompt = generate_dynamic_code_prompt( - user_query=query, - intent_result=intent_result.to_dict(), - code_context=context_str, - retrieved_results=retrieved_results, - conversation_history=history - ) - logger.info(f"使用代码专用Prompt,类型: {intent_result.category}") + + # 6. 生成Prompt + logger.info("生成Prompt") + if use_advanced_prompt and confidence > threshold_high: + filled_prompt = f"""基于以下参考信息回答用户问题。如果参考信息不足,请基于你的知识回答。 + +参考信息: +{context_str} + +用户问题:{query} + +请直接回答:""" else: - # 对于非代码问题,使用通用Prompt if history: qa_prompt = QA_PROMPT_HISTORY filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query) else: qa_prompt = QA_PROMPT_NO_HISTORY filled_prompt = qa_prompt.format(context_str=context_str, query_str=query) - logger.info("使用通用Prompt") - stream_response = await self.llm.astream_complete( - prompt=filled_prompt - ) + logger.info(f"使用Prompt类型: {'通用' if not use_advanced_prompt else '代码自由发挥'}") + + # 7. 流式生成回答 + stream_response = await self.llm.astream_complete(prompt=filled_prompt) full_response = "" think_filter = OptimizedDeltaThinkFilter() async for chunk in stream_response: - # 提取文本内容 delta, full_text, has_output = think_filter.process_delta_robust(chunk) if delta is not None: full_response += delta yield delta.encode('utf-8') - await asyncio.sleep(0.001) # slight delay to yield control + await asyncio.sleep(0.001) logger.info(f"响应完成,长度: {len(full_response)}字符") print(full_response) @@ -400,98 +405,97 @@ class RAGEngine: query: User query string history: Chat history string top_k: Number of documents to retrieve (optional) - filters: Metadata filters for pre-filtering documents (optional) Returns: Complete response string """ try: - # 1. 使用小型模型判断是否是代码相关问题 - logger.info("使用小型模型判断是否是代码相关问题") - is_code = await self.is_code_related(query, history) - logger.info(f"是否是代码相关问题: {is_code}") + # 1. 使用小型模型分析问题,返回置信度和可能的filter + logger.info("使用小型模型分析问题") + analysis_result = await self.is_code_related(query, history) + confidence = analysis_result['confidence'] + filters = analysis_result['filters'] + logger.info(f"分析结果: confidence={confidence}, filters={filters}") - # 2. 处理查询(整合意图识别、filter生成和查询转换) - logger.info(f"开始查询处理: {query[:50]}...") - process_result = self.query_processor.process_query(query, history) + # 2. 根据置信度选择collection和检索策略 + threshold_low = settings.CODE_RELATED_THRESHOLD_LOW + threshold_high = settings.CODE_RELATED_THRESHOLD_HIGH - # 提取结果 - intent_result = process_result['intent'] - # 使用意图识别结果更新 is_code 变量,因为意图识别结果更准确 - is_code = intent_result.is_code_related - logger.info(f"根据意图识别结果更新 is_code: {is_code}") - filters = process_result['filters'] - transformed_queries = [process_result['transformed']['rewritten']] + process_result['transformed']['sub_queries'] + if confidence < threshold_low: + # 置信度低:非代码问题,使用non_code collection + collection_key = 'non_code' + use_advanced_prompt = False + logger.info(f"置信度 {confidence} < {threshold_low},判定为非代码问题,使用non_code collection") + elif confidence > threshold_high: + # 置信度高:代码相关问题,使用code collection + collection_key = 'code' + use_advanced_prompt = True + logger.info(f"置信度 {confidence} > {threshold_high},判定为代码相关问题,使用code collection") + else: + # 中间区间:不区分,使用所有collection + collection_key = None + use_advanced_prompt = False + logger.info(f"置信度 {threshold_low} <= {confidence} <= {threshold_high},通用问题,搜索所有collection") - logger.info(f"代码意图识别结果: {intent_result.category}") - logger.info(f"过滤条件: {filters}") - logger.info(f"查询转换完成,生成了 {len(transformed_queries)} 个转换后的查询") + # 3. 使用纯向量检索(替代融合检索) + logger.info("使用纯向量检索") + k = top_k or settings.TOP_K - # 2. 使用融合检索,对每个转换后的查询进行检索 - logger.info("使用融合检索策略") - all_hybrid_results = [] - # 根据是否是代码问题选择合适的 collection - collection_key = 'code' if is_code else 'non_code' - logger.info(f"使用 collection: {collection_key}") - for transformed_query in transformed_queries: - hybrid_results = await self.vector_store_manager.ahybrid_search( - query=transformed_query, - top_k=top_k or settings.TOP_K, - filters=filters, - collection_key=collection_key - ) - all_hybrid_results.extend(hybrid_results) + # 获取retriever + retriever = self.vector_store_manager.get_retriever( + top_k=k * 2, # 获取更多结果用于去重 + filters=filters if filters else None, + collection_key=collection_key + ) + + # 执行检索 + if isinstance(retriever, list): + # 多个collection + all_nodes = [] + for key, r in retriever: + nodes = r.retrieve(query) + for node in nodes: + actual_node = node.node if hasattr(node, 'node') else node + if hasattr(actual_node, 'metadata'): + actual_node.metadata['_collection_key'] = key + all_nodes.extend(nodes) + vector_nodes = all_nodes + else: + vector_nodes = retriever.retrieve(query) # 去重并按得分排序 seen_doc_ids = set() - unique_hybrid_results = [] - for doc_id, score, metadata in all_hybrid_results: - if doc_id not in seen_doc_ids: + unique_results = [] + for node in vector_nodes: + doc_id = getattr(node, 'id_', None) or getattr(node, 'node_id', None) + if doc_id and doc_id not in seen_doc_ids: seen_doc_ids.add(doc_id) - unique_hybrid_results.append((doc_id, score, metadata)) + unique_results.append(node) - # 按得分排序 - unique_hybrid_results.sort(key=lambda x: x[1], reverse=True) + # 按得分排序,取top_k + unique_results.sort(key=lambda x: getattr(x, 'score', 0), reverse=True) + vector_nodes = unique_results[:k] - # 限制结果数量 - hybrid_results = unique_hybrid_results[:top_k or settings.TOP_K] + logger.info(f"检索到 {len(vector_nodes)} 个结果") - # 根据文档ID获取完整的节点信息 - retrieved_nodes = [] - for doc_id, score, metadata in hybrid_results: - # 从向量存储中获取文档内容 - doc_chunks = self.vector_store_manager.get_document_by_id(doc_id) - for chunk in doc_chunks: - # 创建节点对象 - from llama_index.core.schema import TextNode - node = TextNode( - text=chunk['text'], - node_id=chunk['id'], - metadata=chunk['metadata'] - ) - retrieved_nodes.append(node) - - # 3. 构建原始的检索结果列表,包含 text 和 metadata + # 4. 获取文档内容 retrieved_results = [] - for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1): + for i, node in enumerate(vector_nodes, 1): # 优先使用metadata中的func_body字段 - if hasattr(node, 'metadata') and 'func_body' in node.metadata: - text = node.metadata['func_body'] + metadata = getattr(node, 'metadata', {}) + if 'func_body' in metadata: + text = metadata['func_body'] else: text = node.text if hasattr(node, 'text') else str(node) text = text.strip() - # 获取原始 metadata - metadata = getattr(node, 'metadata', {}) - - # 保存原始的 text 和 metadata retrieved_results.append({ 'id': i, 'text': text, 'metadata': metadata }) - # 构建简单的上下文字符串,只包含基本信息 + # 5. 构建上下文 context_parts = [] for result in retrieved_results: metadata = result['metadata'] @@ -508,31 +512,31 @@ class RAGEngine: context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息" - # 4. 生成优化的Prompt - logger.info("生成优化的Prompt") - if intent_result.is_code_related: - # 对于代码相关问题,使用代码专用Prompt - filled_prompt = generate_dynamic_code_prompt( - user_query=query, - intent_result=intent_result.to_dict(), - code_context=context_str, - retrieved_results=retrieved_results, - conversation_history=history - ) - logger.info(f"使用代码专用Prompt,类型: {intent_result.category}") + # 6. 生成Prompt + logger.info("生成Prompt") + if use_advanced_prompt and confidence > threshold_high: + # 高置信度代码问题:让LLM自由发挥,不预设模板 + filled_prompt = f"""基于以下参考信息回答用户问题。如果参考信息不足,请基于你的知识回答。 + +参考信息: +{context_str} + +用户问题:{query} + +请直接回答:""" else: - # 对于非代码问题,使用通用Prompt + # 非代码问题或中间区间:使用通用Prompt if history: qa_prompt = QA_PROMPT_HISTORY filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query) else: qa_prompt = QA_PROMPT_NO_HISTORY filled_prompt = qa_prompt.format(context_str=context_str, query_str=query) - logger.info("使用通用Prompt") - response = await self.llm.acomplete( - prompt=filled_prompt - ) + logger.info(f"使用Prompt类型: {'通用' if not use_advanced_prompt else '代码自由发挥'}") + + # 7. 调用LLM生成回答 + response = await self.llm.acomplete(prompt=filled_prompt) return response.text except Exception as e: