refactor: 优化RAG检索策略,提升查询性能
- 小模型返回置信度和提取的filter - 根据置信度分流到不同collection - 用纯向量检索替代融合检索 - 删除了意图识别步骤,减少LLM调用 - 新增配置项支持阈值配置
This commit is contained in:
parent
ced7021e16
commit
99115692e5
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Reference in New Issue