RAG/utils/query_processor.py

419 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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))