RAG/utils/query_processor.py

419 lines
16 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""
查询处理器
整合意图识别filter生成和查询转换功能使用一次LLM调用生成所有结果
"""
import json
from typing import Dict, Any, Optional, List
from loguru import logger
from config import settings
from llama_index.llms.ollama import Ollama
class CodeIntentCategory:
"""代码意图分类枚举 - 细粒度分类体系"""
# 代码解释与逻辑类 (Existing Code Focus)
LOGIC_EXPLANATION = "logic_explanation" # 解释既有代码的底层逻辑
ENTITY_INTRODUCTION = "entity_introduction" # 介绍具体的代码实体函数定义、类属性、API参数
CODE_STRUCTURE = "code_structure" # 询问项目组织
# 代码生成与实现类 (New Code Focus)
CODE_GENERATION = "code_generation" # 请求从零编写完整代码或功能块
BOILERPLATE_IMPLEMENTATION = "boilerplate_implementation" # 请求提供标准算法/模板
# 调试、优化与理论类
ERROR_DEBUGGING = "error_debugging" # 排查 Bug 或异常
CODE_OPTIMIZATION = "code_optimization" # 改进既有代码的性能或质量
ALGORITHM_THEORY = "algorithm_theory" # 算法原理或复杂度分析
# 非代码类
GENERAL_TECHNICAL = "general_technical" # 通用技术咨询
NON_TECHNICAL = "non_technical" # 非技术问题
UNKNOWN = "unknown" # 未知类型
class PromptTemplateType:
"""Prompt模板类型枚举"""
CODE_EXPLANATION = "code_explanation" # 代码解释模板(逻辑解释、实体介绍、代码结构)
CODE_GENERATION = "code_generation" # 代码生成模板(代码生成、模板实现)
CODE_DEBUGGING = "code_debugging" # 代码调试模板(错误调试)
CODE_OPTIMIZATION = "code_optimization" # 代码优化模板(代码优化)
ALGORITHM_EXPLANATION = "algorithm_explanation" # 算法解释模板(算法理论)
GENERAL_QA = "general_qa" # 通用问答模板(通用技术咨询、非技术问题)
class CodeIntentResult:
"""代码意图识别结果"""
def __init__(
self,
is_code_related: bool,
category: str,
confidence: float,
prompt_template_type: str,
keywords: List[str],
reasoning: str,
requires_code_context: bool,
suggested_search_terms: List[str],
):
self.is_code_related = is_code_related # 是否与代码相关True/False
self.category = category # 代码意图分类
self.confidence = confidence # 置信度分数范围0-1之间
self.prompt_template_type = prompt_template_type # Prompt模板类型
self.keywords = keywords # 相关关键词列表
self.reasoning = reasoning # 解释或理由
self.requires_code_context = requires_code_context # 是否需要代码上下文True/False
self.suggested_search_terms = suggested_search_terms # 建议搜索条款列表
def to_dict(self) -> Dict[str, Any]:
"""转换为字典格式"""
return {
"is_code_related": self.is_code_related,
"category": self.category,
"confidence": self.confidence,
"prompt_template_type": self.prompt_template_type,
"keywords": self.keywords,
"reasoning": self.reasoning,
"requires_code_context": self.requires_code_context,
"suggested_search_terms": self.suggested_search_terms,
}
def to_json(self) -> str:
"""转换为JSON格式"""
return json.dumps(self.to_dict(), ensure_ascii=False, indent=2)
class QueryProcessor:
"""
查询处理器
整合意图识别filter生成和查询转换功能
"""
def __init__(self, llm: Optional[Ollama] = None):
"""
初始化查询处理器
Args:
llm: LLM实例如果为None则使用默认配置
"""
if llm is None:
self.llm = Ollama(
model=settings.OLLAMA_MODEL,
base_url=settings.OLLAMA_BASE_URL,
temperature=0.1, # 低温度确保输出稳定
request_timeout=1200.0
)
else:
self.llm = llm
logger.info("查询处理器初始化完成")
def _build_integrated_prompt(self, query: str, history: Optional[str] = None) -> str:
"""
构建集成Prompt一次调用生成所有结果
Args:
query: 用户问题
history: 对话历史
Returns:
str: 集成Prompt
"""
from utils.prompt.integrated_query_processing import INTEGRATED_QUERY_PROCESSING_TEMPLATE
history_str = "" if not history else history
prompt = INTEGRATED_QUERY_PROCESSING_TEMPLATE.format(
history_str=history_str,
query=query
)
return prompt
def _map_to_prompt_template_type(self, category: str) -> str:
"""
根据分类确定Prompt模板类型
Args:
category: 意图分类
Returns:
str: Prompt模板类型
"""
# 代码解释与逻辑类 -> CODE_EXPLANATION
if category in [CodeIntentCategory.LOGIC_EXPLANATION, CodeIntentCategory.ENTITY_INTRODUCTION, CodeIntentCategory.CODE_STRUCTURE]:
return PromptTemplateType.CODE_EXPLANATION
# 代码生成与实现类 -> CODE_GENERATION
elif category in [CodeIntentCategory.CODE_GENERATION, CodeIntentCategory.BOILERPLATE_IMPLEMENTATION]:
return PromptTemplateType.CODE_GENERATION
# 调试、优化与理论类
elif category == CodeIntentCategory.ERROR_DEBUGGING:
return PromptTemplateType.CODE_DEBUGGING
elif category == CodeIntentCategory.CODE_OPTIMIZATION:
return PromptTemplateType.CODE_OPTIMIZATION
elif category == CodeIntentCategory.ALGORITHM_THEORY:
return PromptTemplateType.ALGORITHM_EXPLANATION
# 非代码问题 -> GENERAL_QA
elif category in [CodeIntentCategory.GENERAL_TECHNICAL, CodeIntentCategory.NON_TECHNICAL, CodeIntentCategory.UNKNOWN]:
return PromptTemplateType.GENERAL_QA
else:
return PromptTemplateType.GENERAL_QA
def _parse_llm_response(self, response_text: str) -> Optional[Dict[str, Any]]:
"""
解析LLM响应
Args:
response_text: LLM响应文本
Returns:
解析后的字典如果解析失败返回None
"""
try:
response_text = response_text.strip()
# 尝试提取JSON部分
json_start = response_text.find('{')
json_end = response_text.rfind('}')
if json_start == -1 or json_end == -1:
logger.warning(f"未找到JSON格式响应: {response_text}")
return None
json_str = response_text[json_start:json_end + 1]
logger.debug(f"提取的JSON字符串: {json_str}")
result = json.loads(json_str)
return result
except json.JSONDecodeError as e:
logger.error(f"JSON解析失败: {e}, 响应: {response_text}")
return None
except Exception as e:
logger.error(f"解析响应失败: {e}")
return None
def process_query(self, query: str, history: Optional[str] = None) -> Dict[str, Any]:
"""
处理查询一次调用生成所有结果
Args:
query: 用户问题
history: 对话历史
Returns:
包含意图识别filter生成和查询转换结果的字典
"""
try:
# 构建集成Prompt
prompt = self._build_integrated_prompt(query, history)
# 调用LLM
response = self.llm.complete(prompt)
response_text = response.text
logger.debug(f"LLM响应: {response_text}")
# 解析响应
parsed_result = self._parse_llm_response(response_text)
if parsed_result is None:
logger.error("解析LLM响应失败使用默认结果")
return self._get_default_result(query)
# 验证并处理结果
result = {
"intent": None,
"filters": {},
"transformed": {
"rewritten": query,
"backward": query,
"sub_queries": [query]
}
}
# 处理意图识别结果
try:
if "intent" in parsed_result:
intent_data = parsed_result["intent"]
# 确保所有必需字段都存在
intent_data.setdefault("is_code_related", False)
intent_data.setdefault("category", CodeIntentCategory.UNKNOWN)
intent_data.setdefault("confidence", 0.5)
intent_data.setdefault("keywords", [])
intent_data.setdefault("reasoning", "")
intent_data.setdefault("requires_code_context", False)
intent_data.setdefault("suggested_search_terms", [])
# 确定Prompt模板类型
prompt_template_type = self._map_to_prompt_template_type(intent_data["category"])
# 构建CodeIntentResult对象
intent_result = CodeIntentResult(
is_code_related=intent_data["is_code_related"],
category=intent_data["category"],
confidence=intent_data["confidence"],
prompt_template_type=prompt_template_type,
keywords=intent_data["keywords"],
reasoning=intent_data["reasoning"],
requires_code_context=intent_data["requires_code_context"],
suggested_search_terms=intent_data["suggested_search_terms"]
)
result["intent"] = intent_result
except Exception as e:
logger.error(f"处理意图识别结果失败: {e}")
# 处理过滤条件
try:
if "filters" in parsed_result:
filters = parsed_result["filters"]
if isinstance(filters, dict):
# 验证并过滤结果确保只包含有效的metadata key
valid_keys = ['func_id', 'func_name', 'class_name', 'file_path', 'lang', 'params', 'return_type', 'docstring', 'start_line', 'end_line', 'repo_id', 'branch', 'func_body']
filtered_filters = {}
for key, value in filters.items():
if key in valid_keys and value:
# 确保value是字符串类型
if isinstance(value, str):
filtered_filters[key] = value
result["filters"] = filtered_filters
except Exception as e:
logger.error(f"处理过滤条件失败: {e}")
# 处理查询转换结果
try:
if "transformed" in parsed_result:
transformed = parsed_result["transformed"]
if isinstance(transformed, dict):
result["transformed"].update({
"rewritten": transformed.get("rewritten", query),
"backward": transformed.get("backward", query),
"sub_queries": transformed.get("sub_queries", [query])
})
except Exception as e:
logger.error(f"处理查询转换结果失败: {e}")
logger.info("查询处理完成")
return result
except Exception as e:
logger.error(f"查询处理失败: {e}")
return self._get_default_result(query)
def _get_default_result(self, query: str) -> Dict[str, Any]:
"""
获取默认结果当处理失败时使用
Args:
query: 用户问题
Returns:
默认结果
"""
logger.warning(f"使用默认查询处理结果: {query}")
# 构建默认的意图识别结果
default_intent = CodeIntentResult(
is_code_related=False,
category=CodeIntentCategory.UNKNOWN,
confidence=0.0,
prompt_template_type=PromptTemplateType.GENERAL_QA,
keywords=[],
reasoning="查询处理失败,使用默认结果",
requires_code_context=False,
suggested_search_terms=[query]
)
return {
"intent": default_intent,
"filters": {},
"transformed": {
"rewritten": query,
"backward": query,
"sub_queries": [query]
}
}
# 辅助函数
def create_query_processor(llm: Optional[Ollama] = None) -> QueryProcessor:
"""
创建查询处理器实例
Args:
llm: LLM实例如果为None则使用默认配置
Returns:
QueryProcessor: 处理器实例
"""
return QueryProcessor(llm)
class MetadataFilter:
"""
Metadata过滤工具类
"""
@staticmethod
def apply_filter(metadata: Dict[str, Any], filters: Dict[str, Any]) -> bool:
"""
应用过滤条件到metadata
Args:
metadata: 文档的metadata
filters: 过滤条件
Returns:
bool: 如果metadata符合过滤条件返回True否则返回False
"""
if not filters:
return True
for key, value in filters.items():
if key not in metadata:
return False
metadata_value = metadata[key]
if isinstance(metadata_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in metadata_value.lower():
return False
else:
if metadata_value != value:
return False
return True
if __name__ == "__main__":
"""测试查询处理器"""
processor = QueryProcessor()
test_queries = [
"在 data_structures 目录下二叉搜索树Binary Search Tree的删除操作依赖于哪些辅助方法来寻找后继节点",
"如何实现快速排序算法?",
"kth_number的时间复杂度是多少",
"今天天气怎么样?"
]
for query in test_queries:
print(f"\n{'='*60}")
print(f"问题: {query}")
print('='*60)
result = processor.process_query(query)
print("意图识别结果:")
if result["intent"]:
print(result["intent"].to_json())
print("\n过滤条件:")
import json
print(json.dumps(result["filters"], ensure_ascii=False, indent=2))
print("\n查询转换结果:")
print(json.dumps(result["transformed"], ensure_ascii=False, indent=2))