refactor: 优化RAG检索策略,提升查询性能

- 小模型返回置信度和提取的filter
- 根据置信度分流到不同collection
- 用纯向量检索替代融合检索
- 删除了意图识别步骤,减少LLM调用
- 新增配置项支持阈值配置
This commit is contained in:
linlin 2026-03-13 13:41:16 +08:00
parent ced7021e16
commit 99115692e5
2 changed files with 202 additions and 193 deletions

View File

@ -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

View File

@ -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: