添加代码QA功能,提升对代码相关问题的处理能力,同时优化向量存储和检索性能

This commit is contained in:
gu0weix1n 2026-02-27 10:44:35 +08:00
parent b44b7dec8b
commit 3fca89d7a8
26 changed files with 3279 additions and 41 deletions

View File

@ -258,6 +258,10 @@ class QueryRequest(BaseModel):
query: str = Field(..., description="User query string", min_length=1)
top_k: Optional[int] = Field(None, description="Number of documents to retrieve", ge=1, le=20)
stream: bool = Field(True, description="Whether to stream the response")
repo: Optional[str] = Field(None, description="Git repository name (optional)")
branch: Optional[str] = Field(None, description="Git branch name (optional)")
is_code_related: Optional[bool] = Field(None, description="Whether the query is code related (optional)")
history: Optional[List[Dict[str, str]]] = Field(None, description="Conversation history (optional)")
class RetrieveRequest(BaseModel):
@ -606,11 +610,23 @@ async def query(request: QueryRequest):
raise HTTPException(status_code=503, detail="RAG engine not initialized")
try:
# Build conversation history
history_str = ""
if request.history:
history_parts = []
for msg in request.history:
role = msg.get('role', 'user')
content = msg.get('content', '')
if role == 'user':
history_parts.append(f"用户: {content}")
else:
history_parts.append(f"助手: {content}")
history_str = "\n".join(history_parts)
if request.stream:
# Stream response
return StreamingResponse(
rag_engine.query_stream(request.query, None, request.top_k),
rag_engine.query_stream(request.query, history_str, request.top_k),
media_type="text/event-stream",
headers={
"X-Accel-Buffering": "no",
@ -620,7 +636,7 @@ async def query(request: QueryRequest):
)
else:
# Return complete response (run in thread pool for better concurrency)
response = await rag_engine.query(request.query, None, request.top_k)
response = await rag_engine.query(request.query, history_str, request.top_k)
return response
except Exception as e:
logger.error(f"Error processing query: {e}")

@ -0,0 +1 @@
Subproject commit c22e97f255c92b61ad6014ef74ec662d9f5f6325

@ -0,0 +1 @@
Subproject commit 678dedbbf94be54b3c9c258368e28bb8e7736d62

@ -0,0 +1 @@
Subproject commit 1e967fc87c2761167e0a0a0e84dd2d213c0e1186

View File

@ -13,6 +13,8 @@ from loguru import logger
from config import settings
from .vector_store import VectorStoreManager
from .chunk_handler import OptimizedDeltaThinkFilter
from utils.query_processor import QueryProcessor
from utils.code_prompt_manager import generate_dynamic_code_prompt
# Single module-level prompt string for easy editing in one place
@ -126,7 +128,7 @@ class RAGEngine:
prompt_template: Optional[PromptTemplate] = None,
system_prompt: Optional[str] = None,
temperature: float = 0.7,
request_timeout: float = 120.0,
request_timeout: float = 1200.0,
):
self.vector_store_manager = vector_store_manager
# Configurable LLM / prompt parameters
@ -157,6 +159,9 @@ class RAGEngine:
temperature=self._temperature,
request_timeout=self._request_timeout,
)
# 初始化查询处理器
self.query_processor = QueryProcessor(llm=self.llm)
def extract_text_from_chunk(self, chunk) -> Optional[str]:
if hasattr(chunk, 'delta'):
@ -187,29 +192,127 @@ class RAGEngine:
Response text chunks
"""
try:
# Create query engine with streaming mode
retriever = self.vector_store_manager.get_retriever(top_k=top_k or settings.TOP_K)
# query index
retrieved_nodes = await retriever.aretrieve(query)
# 1. 处理查询整合意图识别、filter生成和查询转换
logger.info(f"开始查询处理: {query}")
process_result = self.query_processor.process_query(query, history)
# 2. 构建上下文
# 提取结果
intent_result = process_result['intent']
filters = process_result['filters']
transformed_queries = [process_result['transformed']['rewritten']] + process_result['transformed']['sub_queries']
logger.info(f"代码意图识别结果: {intent_result.category}")
logger.info(f"过滤条件: {filters}")
logger.info(f"查询转换结果: {transformed_queries}")
# 4. 使用融合检索,对每个转换后的查询进行检索
logger.info("使用融合检索策略")
all_hybrid_results = []
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
)
all_hybrid_results.extend(hybrid_results)
# 去重并按得分排序
seen_doc_ids = set()
unique_hybrid_results = []
for doc_id, score, metadata in all_hybrid_results:
if doc_id not in seen_doc_ids:
seen_doc_ids.add(doc_id)
unique_hybrid_results.append((doc_id, score, metadata))
# 按得分排序
unique_hybrid_results.sort(key=lambda x: x[1], reverse=True)
# 限制结果数量
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)
# 3. 构建上下文
context_parts = []
for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1): # 限制数量
max_nodes = top_k or settings.TOP_K
# 构建原始的检索结果列表,包含 text 和 metadata
retrieved_results = []
for i, node in enumerate(retrieved_nodes[:max_nodes], 1):
text = node.text if hasattr(node, 'text') else str(node)
# 清理和截断
text = text.strip()
if len(text) > 400:
text = text[:400] + "..."
context_parts.append(f"【参考信息{i}{text}")
# 获取原始 metadata
metadata = getattr(node, 'metadata', {})
# 保存原始的 text 和 metadata
retrieved_results.append({
'id': i,
'text': text,
'metadata': metadata
})
# 构建简单的上下文字符串,只包含基本信息
context_parts = []
for result in retrieved_results:
metadata = result['metadata']
file_path = metadata.get('file_path', '')
func_name = metadata.get('func_name', '')
context_part = f"【参考信息{result['id']}"
if file_path:
context_part += f"(来源:{file_path}"
if func_name:
context_part += f"\n函数:{func_name}"
context_part += f"\n{result['text']}"
context_parts.append(context_part)
context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息"
if history is not None:
qa_prompt = QA_PROMPT_HISTORY
filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query)
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}")
else:
qa_prompt = QA_PROMPT_NO_HISTORY
filled_prompt = qa_prompt.format(context_str=context_str, query_str=query)
# 对于非代码问题使用通用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
@ -220,7 +323,6 @@ class RAGEngine:
async for chunk in stream_response:
# 提取文本内容
# text_chunk = self.extract_text_from_chunk(chunk)
delta, full_text, has_output = think_filter.process_delta_robust(chunk)
if delta is not None:
@ -229,10 +331,12 @@ class RAGEngine:
await asyncio.sleep(0.001) # slight delay to yield control
logger.info(f"响应完成,长度: {len(full_response)}字符")
print(full_response)
except Exception as e:
logger.error(f"Error in RAG query: {e}")
async def query(self, query: str, history: str, top_k: Optional[int] = None) -> str:
"""
Query the RAG system and return complete response
@ -246,29 +350,113 @@ class RAGEngine:
Complete response string
"""
try:
# Create query engine with streaming mode
retriever = self.vector_store_manager.get_retriever(top_k=top_k or settings.TOP_K)
# query index
retrieved_nodes = await retriever.aretrieve(query)
# 1. 处理查询整合意图识别、filter生成和查询转换
logger.info(f"开始查询处理: {query[:50]}...")
process_result = self.query_processor.process_query(query, history)
# 2. 构建上下文
context_parts = []
for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1): # 限制数量
# 提取结果
intent_result = process_result['intent']
filters = process_result['filters']
transformed_queries = [process_result['transformed']['rewritten']] + process_result['transformed']['sub_queries']
logger.info(f"代码意图识别结果: {intent_result.category}")
logger.info(f"过滤条件: {filters}")
logger.info(f"查询转换完成,生成了 {len(transformed_queries)} 个转换后的查询")
# 2. 使用融合检索,对每个转换后的查询进行检索
logger.info("使用融合检索策略")
all_hybrid_results = []
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
)
all_hybrid_results.extend(hybrid_results)
# 去重并按得分排序
seen_doc_ids = set()
unique_hybrid_results = []
for doc_id, score, metadata in all_hybrid_results:
if doc_id not in seen_doc_ids:
seen_doc_ids.add(doc_id)
unique_hybrid_results.append((doc_id, score, metadata))
# 按得分排序
unique_hybrid_results.sort(key=lambda x: x[1], reverse=True)
# 限制结果数量
hybrid_results = unique_hybrid_results[:top_k or settings.TOP_K]
# 根据文档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
retrieved_results = []
for i, node in enumerate(retrieved_nodes[:top_k or settings.TOP_K], 1):
text = node.text if hasattr(node, 'text') else str(node)
# 清理和截断
text = text.strip()
if len(text) > 400:
text = text[:400] + "..."
context_parts.append(f"【参考信息{i}{text}")
# 获取原始 metadata
metadata = getattr(node, 'metadata', {})
# 保存原始的 text 和 metadata
retrieved_results.append({
'id': i,
'text': text,
'metadata': metadata
})
# 构建简单的上下文字符串,只包含基本信息
context_parts = []
for result in retrieved_results:
metadata = result['metadata']
file_path = metadata.get('file_path', '')
func_name = metadata.get('func_name', '')
context_part = f"【参考信息{result['id']}"
if file_path:
context_part += f"(来源:{file_path}"
if func_name:
context_part += f"\n函数:{func_name}"
context_part += f"\n{result['text']}"
context_parts.append(context_part)
context_str = "\n\n".join(context_parts) if context_parts else "未找到相关参考信息"
if history is not None:
qa_prompt = QA_PROMPT_HISTORY
filled_prompt = qa_prompt.format(history=history, context_str=context_str, query_str=query)
# 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}")
else:
qa_prompt = QA_PROMPT_NO_HISTORY
filled_prompt = qa_prompt.format(context_str=context_str, query_str=query)
# 对于非代码问题使用通用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

View File

@ -5,7 +5,7 @@ import os
import time
from datetime import datetime, date
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import List, Dict, Any
from typing import List, Dict, Any, Tuple
import chromadb
from chromadb.config import Settings as ChromaSettings
from llama_index.vector_stores.chroma import ChromaVectorStore
@ -13,6 +13,11 @@ from llama_index.core import VectorStoreIndex, StorageContext
from llama_index.embeddings.ollama import OllamaEmbedding
from loguru import logger
from config import settings
import numpy as np
import re
import string
from rank_bm25 import BM25Okapi
from utils.query_processor import MetadataFilter
class VectorStoreManager:
@ -822,12 +827,13 @@ class VectorStoreManager:
logger.warning(f"Error checking document count: {e}")
return False
def get_retriever(self, top_k: int = None):
def get_retriever(self, top_k: int = None, filters: dict = None):
"""
Get a retriever for querying the vector store
Args:
top_k: Number of documents to retrieve (defaults to settings.TOP_K)
filters: Metadata filters to apply before vector search
Returns:
VectorStoreRetriever instance
@ -884,3 +890,314 @@ class VectorStoreManager:
except Exception as e:
logger.error(f"获取db_source({target_db_source}) metadata中content_column失败: {e}")
return ''
def _preprocess_text(self, text: str) -> List[str]:
"""
预处理文本用于BM25搜索
Args:
text: 原始文本
Returns:
分词后的文本列表
"""
text = text.lower() # 转换为小写
text = text.translate(str.maketrans('', '', string.punctuation)) # 移除标点符号
tokens = re.findall(r'\b\w+\b', text) # 分词
return tokens
def keyword_search(self, query: str, top_k: int = 5, filters: dict = None) -> List[Tuple[str, float, Dict[str, Any]]]:
"""
使用BM25算法进行关键词搜索
Args:
query: 搜索查询
top_k: 返回结果数量
filters: metadata过滤条件
Returns:
排序后的结果列表每个元素包含(doc_id, score, metadata)
"""
try:
# 获取文档
results = self.collection.get(
include=['documents', 'metadatas']
)
documents = results.get('documents', [])
metadatas = results.get('metadatas', [])
ids = results.get('ids', [])
if not documents:
return []
# 应用过滤条件
filtered_documents = []
filtered_metadatas = []
filtered_ids = []
for doc, meta, doc_id in zip(documents, metadatas, ids):
if not filters:
# 没有过滤条件,直接添加
filtered_documents.append(doc)
filtered_metadatas.append(meta)
filtered_ids.append(doc_id)
else:
# 应用过滤条件
match = True
for key, value in filters.items():
if key not in meta:
match = False
break
meta_value = meta[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if match:
filtered_documents.append(doc)
filtered_metadatas.append(meta)
filtered_ids.append(doc_id)
if not filtered_documents:
return []
tokenized_docs = [self._preprocess_text(doc) for doc in filtered_documents] # 预处理文档
bm25 = BM25Okapi(tokenized_docs) # 初始化BM25
tokenized_query = self._preprocess_text(query) # 预处理查询
scores = bm25.get_scores(tokenized_query) # 计算BM25得分
sorted_indices = np.argsort(scores)[::-1][:top_k] # 排序并获取top_k结果
# 构建结果列表
search_results = []
for idx in sorted_indices:
if scores[idx] > 0: # 只返回得分大于0的结果
doc_id = filtered_ids[idx]
score = float(scores[idx])
metadata = filtered_metadatas[idx]
# 由于已经在获取文档时应用了过滤条件,这里不需要再次应用
search_results.append((doc_id, score, metadata))
return search_results
except Exception as e:
logger.error(f"关键词搜索失败: {e}")
return []
def hybrid_search(self, query: str, top_k: int = 5, vector_weight: float = 0.6, keyword_weight: float = 0.4, filters: dict = None) -> List[Tuple[str, float, Dict[str, Any]]]:
"""
融合检索结合向量搜索和关键词搜索
Args:
query: 搜索查询
top_k: 返回结果数量
vector_weight: 向量搜索权重
keyword_weight: 关键词搜索权重
filters: metadata过滤条件
Returns:
排序后的结果列表每个元素包含(doc_id, score, metadata)
"""
try:
# 1. 执行向量搜索
# 获取更多结果以确保有足够的候选
vector_retriever = self.get_retriever(top_k=top_k * 4) # 获取更多结果
vector_nodes = vector_retriever.retrieve(query)
vector_results = {}
for node in vector_nodes:
if hasattr(node, 'id_'):
doc_id = node.id_
elif hasattr(node, 'node_id'):
doc_id = node.node_id
else:
continue
vector_results[doc_id] = {
'score': node.score if hasattr(node, 'score') else 0.5,
'metadata': node.metadata if hasattr(node, 'metadata') else {},
'text': node.text if hasattr(node, 'text') else ''
}
# 2. 执行关键词搜索
keyword_results = self.keyword_search(query, top_k=top_k * 4, filters=filters)
keyword_scores = {}
keyword_metadata = {}
for doc_id, score, metadata in keyword_results:
keyword_scores[doc_id] = score
keyword_metadata[doc_id] = metadata
# 3. 归一化得分
# 归一化向量得分
if vector_results:
vector_scores = list(vector_results.values())
vector_min = min(item['score'] for item in vector_scores)
vector_max = max(item['score'] for item in vector_scores)
vector_range = vector_max - vector_min if vector_max > vector_min else 1
for doc_id in vector_results:
vector_results[doc_id]['normalized_score'] = (vector_results[doc_id]['score'] - vector_min) / vector_range
# 归一化关键词得分
if keyword_scores:
keyword_min = min(keyword_scores.values())
keyword_max = max(keyword_scores.values())
keyword_range = keyword_max - keyword_min if keyword_max > keyword_min else 1
for doc_id in keyword_scores:
keyword_scores[doc_id] = (keyword_scores[doc_id] - keyword_min) / keyword_range
# 4. 融合得分
hybrid_results = {}
# 合并向量搜索结果
for doc_id, info in vector_results.items():
# 应用过滤条件
if filters:
metadata = info['metadata']
match = True
for key, value in filters.items():
if key not in metadata:
match = False
break
meta_value = metadata[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if not match:
continue
vector_score = info.get('normalized_score', 0)
keyword_score = keyword_scores.get(doc_id, 0)
# 计算融合得分
hybrid_score = vector_weight * vector_score + keyword_weight * keyword_score
hybrid_results[doc_id] = {
'score': hybrid_score,
'metadata': info['metadata'],
'text': info['text']
}
# 合并关键词搜索结果(不在向量搜索结果中的)
for doc_id, score in keyword_scores.items():
if doc_id not in hybrid_results:
# 应用过滤条件
if filters:
metadata = keyword_metadata.get(doc_id, {})
match = True
for key, value in filters.items():
if key not in metadata:
match = False
break
meta_value = metadata[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if not match:
continue
hybrid_score = keyword_weight * score
hybrid_results[doc_id] = {
'score': hybrid_score,
'metadata': keyword_metadata.get(doc_id, {}),
'text': ''
}
# 5. 排序并获取top_k结果
sorted_results = sorted(
hybrid_results.items(),
key=lambda x: x[1]['score'],
reverse=True
)[:top_k]
# 6. 构建最终结果
final_results = []
for doc_id, info in sorted_results:
metadata = info['metadata']
final_results.append((doc_id, info['score'], metadata))
return final_results
except Exception as e:
logger.error(f"融合检索失败: {e}")
# 失败时回退到向量搜索
vector_retriever = self.get_retriever(top_k=top_k)
vector_nodes = vector_retriever.retrieve(query)
fallback_results = []
for node in vector_nodes:
if hasattr(node, 'id_'):
doc_id = node.id_
elif hasattr(node, 'node_id'):
doc_id = node.node_id
else:
continue
# 应用过滤条件
if filters:
metadata = node.metadata if hasattr(node, 'metadata') else {}
match = True
for key, value in filters.items():
if key not in metadata:
match = False
break
meta_value = metadata[key]
if isinstance(meta_value, str) and isinstance(value, str):
# 对于字符串类型,使用大小写不敏感的模糊匹配
if value.lower() not in meta_value.lower():
match = False
break
else:
# 对于其他类型,使用精确匹配
if meta_value != value:
match = False
break
if not match:
continue
score = node.score if hasattr(node, 'score') else 0.5
metadata = node.metadata if hasattr(node, 'metadata') else {}
fallback_results.append((doc_id, score, metadata))
return fallback_results
async def ahybrid_search(self, query: str, top_k: int = 5, vector_weight: float = 0.6, keyword_weight: float = 0.4, filters: dict = None) -> List[Tuple[str, float, Dict[str, Any]]]:
"""
异步融合检索结合向量搜索和关键词搜索
Args:
query: 搜索查询
top_k: 返回结果数量
vector_weight: 向量搜索权重
keyword_weight: 关键词搜索权重
filters: metadata过滤条件
Returns:
排序后的结果列表每个元素包含(doc_id, score, metadata)
"""
# 由于BM25搜索是CPU密集型的这里使用同步方法
# 在实际生产环境中,可以使用线程池来异步执行
return self.hybrid_search(query, top_k, vector_weight, keyword_weight, filters)

View File

@ -134,7 +134,7 @@ class BaseSync(ABC):
for doc in docs:
try:
llamaindex_doc = self.doc_to_llamaindex_doc(doc)
if len(llamaindex_doc.text.strip()) >= 100:
if len(llamaindex_doc.text.strip()) >= 10:
documents.append(llamaindex_doc)
else:
logger.warning(f"跳过过短文档 (id: {llamaindex_doc.id_}),内容长度: {len(llamaindex_doc.text)} 字符")

403
utils/code_intent.py Normal file
View File

@ -0,0 +1,403 @@
"""
代码意图识别模块
用于识别用户问题是否与代码相关并进行详细的意图分类
支持多种代码问题类型的识别为后续检索和Prompt生成提供指导
"""
import sys
import os
# 添加项目根目录到 Python 模块搜索路径
current_dir = os.path.dirname(os.path.abspath(__file__))
project_root = os.path.dirname(current_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
import json
import asyncio
from typing import Dict, Any, Optional, List
from enum import Enum
from functools import lru_cache
from loguru import logger
from config import settings
from llama_index.llms.ollama import Ollama
from utils.prompt import INTENT_DETECTION_TEMPLATE
class CodeIntentCategory(Enum):
"""代码意图分类枚举 - 细粒度分类体系"""
# 代码解释与逻辑类 (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(Enum):
"""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: CodeIntentCategory,
confidence: float,
prompt_template_type: PromptTemplateType,
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.value,
"confidence": self.confidence,
"prompt_template_type": self.prompt_template_type.value,
"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 CodeIntentDetector:
"""代码意图检测器"""
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_classification_prompt(self, query: str, history: Optional[str] = None) -> str:
"""
构建分类Prompt
Args:
query: 用户问题
history: 对话历史
Returns:
str: 分类Prompt
"""
history_str = "" if not history else history
prompt = INTENT_DETECTION_TEMPLATE.format(
history_str=history_str,
query=query
)
return prompt
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]
result = json.loads(json_str)
# 验证必需字段并使用默认值
required_fields = {
'is_code_related': False,
'category': 'unknown',
'confidence': 0.5,
'keywords': [],
'reasoning': '',
'requires_code_context': False,
'suggested_search_terms': []
}
for field, default_value in required_fields.items():
if field not in result:
logger.warning(f"缺少必需字段: {field}, 使用默认值: {default_value}")
result[field] = default_value
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 _map_to_enums(
self,
parsed_result: Dict[str, Any]
) -> tuple[CodeIntentCategory, PromptTemplateType]:
"""
将解析结果映射到枚举类型
Args:
parsed_result: 解析后的结果字典
Returns:
(CodeIntentCategory, PromptTemplateType)
"""
category_str = parsed_result.get('category', 'unknown')
try:
category = CodeIntentCategory(category_str)
except ValueError:
logger.warning(f"未知的分类: {category_str}, 使用默认值")
category = CodeIntentCategory.UNKNOWN
# 根据分类确定Prompt模板类型
if not parsed_result.get('is_code_related', False):
prompt_template_type = PromptTemplateType.GENERAL_QA
# 代码解释与逻辑类 -> CODE_EXPLANATION
elif category in [CodeIntentCategory.LOGIC_EXPLANATION, CodeIntentCategory.ENTITY_INTRODUCTION,
CodeIntentCategory.CODE_STRUCTURE]:
prompt_template_type = PromptTemplateType.CODE_EXPLANATION
# 代码生成与实现类 -> CODE_GENERATION
elif category in [CodeIntentCategory.CODE_GENERATION, CodeIntentCategory.BOILERPLATE_IMPLEMENTATION]:
prompt_template_type = PromptTemplateType.CODE_GENERATION
# 调试、优化与理论类
elif category == CodeIntentCategory.ERROR_DEBUGGING:
prompt_template_type = PromptTemplateType.CODE_DEBUGGING
elif category == CodeIntentCategory.CODE_OPTIMIZATION:
prompt_template_type = PromptTemplateType.CODE_OPTIMIZATION
elif category == CodeIntentCategory.ALGORITHM_THEORY:
prompt_template_type = PromptTemplateType.ALGORITHM_EXPLANATION
# 非代码问题 -> GENERAL_QA
elif category in [CodeIntentCategory.GENERAL_TECHNICAL, CodeIntentCategory.NON_TECHNICAL,
CodeIntentCategory.UNKNOWN]:
prompt_template_type = PromptTemplateType.GENERAL_QA
else:
prompt_template_type = PromptTemplateType.GENERAL_QA
return category, prompt_template_type
async def detect_intent_async(self, query: str, history: Optional[str] = None) -> CodeIntentResult:
"""
异步检测代码意图
Args:
query: 用户问题
history: 对话历史
Returns:
CodeIntentResult: 意图识别结果
"""
try:
# 构建Prompt
prompt = self._build_classification_prompt(query, history)
# 调用LLM
response = await self.llm.acomplete(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)
# 映射到枚举类型
category, prompt_template_type = self._map_to_enums(parsed_result)
# 构建结果对象
result = CodeIntentResult(
is_code_related=parsed_result.get('is_code_related', False),
category=category,
confidence=parsed_result.get('confidence', 0.5),
prompt_template_type=prompt_template_type,
keywords=parsed_result.get('keywords', []),
reasoning=parsed_result.get('reasoning', ''),
requires_code_context=parsed_result.get('requires_code_context', False),
suggested_search_terms=parsed_result.get('suggested_search_terms', [])
)
logger.info(f"代码意图检测完成: {result.to_json()}")
return result
except Exception as e:
logger.error(f"代码意图检测失败: {e}")
return self._get_default_result(query)
def detect_intent(self, query: str) -> CodeIntentResult:
"""
同步检测代码意图包装异步方法
Args:
query: 用户问题
Returns:
CodeIntentResult: 意图识别结果
"""
try:
loop = asyncio.get_event_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
return loop.run_until_complete(self.detect_intent_async(query))
def _get_default_result(self, query: str) -> CodeIntentResult:
"""
获取默认结果当检测失败时使用
Args:
query: 用户问题
Returns:
CodeIntentResult: 默认结果
"""
logger.warning(f"使用默认意图识别结果: {query}")
return 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]
)
@lru_cache(maxsize=100)
def detect_code_intent_cached(query: str) -> CodeIntentResult:
"""
带缓存的代码意图检测同步版本
Args:
query: 用户问题
Returns:
CodeIntentResult: 意图识别结果
"""
detector = CodeIntentDetector()
return detector.detect_intent(query)
async def detect_code_intent_async_cached(query: str) -> CodeIntentResult:
"""
带缓存的代码意图检测异步版本
Args:
query: 用户问题
Returns:
CodeIntentResult: 意图识别结果
"""
detector = CodeIntentDetector()
return await detector.detect_intent_async(query)
def create_intent_detector(llm: Optional[Ollama] = None) -> CodeIntentDetector:
"""
创建代码意图检测器实例
Args:
llm: LLM实例如果为None则使用默认配置
Returns:
CodeIntentDetector: 检测器实例
"""
return CodeIntentDetector(llm)
if __name__ == "__main__":
import asyncio
async def test_intent_detection():
"""测试代码意图检测"""
detector = CodeIntentDetector()
test_queries = [
"kth_number怎么实现的",
"kth_number的时间复杂度是多少",
"kth_number的测试用例有哪些",
"kth_number的代码质量如何",
"trieste实现的采集函数有哪些",
"这个函数是做什么的?",
"如何优化这段代码?",
"Python中如何实现异步",
"今天天气怎么样?",
"为什么会报这个错误?",
"项目结构是怎样的?",
"如何调用这个API"
]
for query in test_queries:
print(f"\n{'='*60}")
print(f"问题: {query}")
print('='*60)
result = await detector.detect_intent_async(query)
print(result.to_json())
asyncio.run(test_intent_detection())

View File

@ -0,0 +1,725 @@
"""
代码Prompt管理模块
用于生成和管理代码相关问答的专属Prompt模板
支持基于代码意图分类的动态Prompt选择和生成
"""
import json
import re
import sys
import os
from typing import Dict, Any, Optional, List, Tuple
from enum import Enum
from loguru import logger
# 添加项目根目录到 Python 模块搜索路径
current_dir = os.path.dirname(os.path.abspath(__file__))
project_root = os.path.dirname(current_dir)
if project_root not in sys.path:
sys.path.insert(0, project_root)
from utils.query_processor import CodeIntentCategory, PromptTemplateType
from utils.prompt import (
INTENT_DETECTION_TEMPLATE,
CODE_EXPLANATION_TEMPLATE,
CODE_DEBUGGING_TEMPLATE,
CODE_GENERATION_TEMPLATE,
ALGORITHM_EXPLANATION_TEMPLATE,
CODE_OPTIMIZATION_TEMPLATE,
GENERAL_QA_TEMPLATE,
)
class CodePromptManager:
"""代码Prompt管理类"""
def __init__(self):
"""
初始化代码Prompt管理器
"""
self._templates = self._load_templates()
self._context_cache: Dict[str, List[Dict[str, str]]] = {}
logger.info("代码Prompt管理器初始化完成")
def _load_templates(self) -> Dict[PromptTemplateType, str]:
"""
加载Prompt模板
Returns:
Dict[PromptTemplateType, str]: Prompt模板字典
"""
return {
PromptTemplateType.CODE_EXPLANATION: CODE_EXPLANATION_TEMPLATE,
PromptTemplateType.CODE_DEBUGGING: CODE_DEBUGGING_TEMPLATE,
PromptTemplateType.CODE_GENERATION: CODE_GENERATION_TEMPLATE,
PromptTemplateType.ALGORITHM_EXPLANATION: ALGORITHM_EXPLANATION_TEMPLATE,
PromptTemplateType.CODE_OPTIMIZATION: CODE_OPTIMIZATION_TEMPLATE,
PromptTemplateType.GENERAL_QA: GENERAL_QA_TEMPLATE
}
def _map_intent_to_prompt_type(self, intent_category: str) -> PromptTemplateType:
"""
将代码意图分类映射到Prompt类型
Args:
intent_category: 代码意图分类字符串
Returns:
PromptTemplateType: 对应的Prompt类型
"""
mapping = {
# 代码解释与逻辑类 -> CODE_EXPLANATION
"logic_explanation": PromptTemplateType.CODE_EXPLANATION,
"entity_introduction": PromptTemplateType.CODE_EXPLANATION,
"code_structure": PromptTemplateType.CODE_EXPLANATION,
# 代码生成与实现类 -> CODE_GENERATION
"code_generation": PromptTemplateType.CODE_GENERATION,
"boilerplate_implementation": PromptTemplateType.CODE_GENERATION,
# 调试、优化与理论类
"error_debugging": PromptTemplateType.CODE_DEBUGGING,
"code_optimization": PromptTemplateType.CODE_OPTIMIZATION,
"algorithm_theory": PromptTemplateType.ALGORITHM_EXPLANATION,
# 非代码问题 -> GENERAL_QA
"general_technical": PromptTemplateType.GENERAL_QA,
"non_technical": PromptTemplateType.GENERAL_QA,
"unknown": PromptTemplateType.GENERAL_QA
}
return mapping.get(intent_category, PromptTemplateType.GENERAL_QA)
def _build_conversation_history(self, history: Optional[Any]) -> str:
"""
构建对话历史字符串
Args:
history: 对话历史可以是字符串或字典列表
Returns:
str: 格式化的对话历史
"""
if not history:
return ""
# 如果是字符串,直接返回
if isinstance(history, str):
return history
# 如果是字典列表,格式化为字符串
if isinstance(history, list):
history_str = []
for item in history:
if isinstance(item, dict):
role = item.get('role', 'user')
content = item.get('content', '')
if role == 'user':
history_str.append(f"用户: {content}")
else:
history_str.append(f"助手: {content}")
return "\n".join(history_str)
# 其他类型,转换为字符串
return str(history)
def _extract_code_from_context(self, code_context: str) -> str:
"""
从上下文中提取代码
Args:
code_context: 代码上下文
Returns:
str: 提取的代码
"""
if not code_context:
return ""
# 尝试提取代码块
code_blocks = re.findall(r'```[\w]*\n[\s\S]*?```', code_context)
if code_blocks:
# 提取所有代码块并合并
extracted_code = []
for block in code_blocks:
# 提取语言标记
lang_match = re.match(r'```([\w]*)\n', block)
language = lang_match.group(1) if lang_match else ""
# 去除代码块标记
code = re.sub(r'```[\w]*\n|```', '', block)
code = code.strip()
if code:
if language:
extracted_code.append(f"语言: {language}\n{code}")
else:
extracted_code.append(code)
return "\n\n".join(extracted_code)
# 如果没有代码块标记,尝试提取看起来像代码的部分
# 查找连续的多行代码(以缩进或常见代码关键字开头)
lines = code_context.split('\n')
code_lines = []
in_code = False
for line in lines:
# 检查是否是代码行
line_stripped = line.strip()
if (line_stripped and
(line.startswith(' ') or line.startswith('\t') or # 缩进
line_stripped.startswith('def ') or line_stripped.startswith('class ') or # Python关键字
line_stripped.startswith('import ') or line_stripped.startswith('from ') or # 导入
line_stripped.startswith('if ') or line_stripped.startswith('for ') or # 控制流
line_stripped.startswith('while ') or line_stripped.startswith('try ') or
line_stripped.startswith('except ') or line_stripped.startswith('finally ') or
line_stripped.startswith('return ') or line_stripped.startswith('print(') or
line_stripped.startswith('// ') or line_stripped.startswith('# ') or # 注释
line_stripped.endswith(';') or # 分号结尾如Java、C++等)
line_stripped.startswith('{') or line_stripped.startswith('}') or # 大括号
re.match(r'^[\w_]+\s*=\s*', line_stripped) or # 变量赋值
re.match(r'^[\w_]+\s*\(.*\)\s*\{{?', line_stripped))): # 函数定义
code_lines.append(line)
in_code = True
elif in_code and line.strip() == '':
# 保留代码中的空行
code_lines.append(line)
elif in_code and len(code_lines) > 3:
# 如果已经收集了多行代码,并且遇到非代码行,停止收集
break
else:
# 非代码行,重置
code_lines = []
in_code = False
if len(code_lines) > 3:
return "\n".join(code_lines)
return code_context
def _format_code_block(self, code: str, language: str = "") -> str:
"""
格式化代码块提高显示质量
Args:
code: 代码内容
language: 代码语言
Returns:
str: 格式化的代码块
"""
if not code:
return ""
# 添加语言标记
lang_tag = language if language else ""
# 确保代码块格式正确
formatted_code = f"```{lang_tag}\n{code}\n```"
return formatted_code
def _build_enhanced_context(self, retrieved_results: List[Dict[str, Any]], intent_result: Optional[Dict[str, Any]]) -> str:
"""
根据意图和检索结果构建增强的上下文
Args:
retrieved_results: 检索结果列表每个元素包含 idtext metadata
intent_result: 意图识别结果
Returns:
str: 增强的上下文
"""
if not retrieved_results:
return "未找到相关参考信息"
context_parts = []
intent = intent_result.get('intent', '') if intent_result else ''
for result in retrieved_results:
metadata = result.get('metadata', {})
text = result.get('text', '')
result_id = result.get('id', 1)
# 提取所有 metadata 字段
func_id = metadata.get('func_id', '')
func_name = metadata.get('func_name', '')
class_name = metadata.get('class_name', 'None')
file_path = metadata.get('file_path', '')
lang = metadata.get('lang', '')
params = metadata.get('params', 0)
return_type = metadata.get('return_type', 'None')
docstring = metadata.get('docstring', '')
start_line = metadata.get('start_line', '')
end_line = metadata.get('end_line', '')
repo_id = metadata.get('repo_id', '')
branch = metadata.get('branch', '')
func_body = metadata.get('func_body', '')
# 根据意图构建不同的上下文
if intent == "code_understanding":
# 代码理解意图,强调语言、函数名、类名、参数、返回类型和函数体
context_part = f"【参考信息{result_id}】这是由{lang}实现的函数{func_name}"
if class_name and class_name != "None":
context_part += f",属于{class_name}"
context_part += f",它接收{params}个参数,返回类型为{return_type}"
if docstring:
context_part += f"。函数说明:{docstring}"
context_part += f"\n文件路径:{file_path},位置:{start_line}-{end_line}\n"
context_part += f"具体实现:\n{func_body}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"原始文本:\n{text}"
elif intent == "code_modification":
# 代码修改意图,强调文件路径、位置和函数体
context_part = f"【参考信息{result_id}】需要修改的代码位于文件:{file_path},位置:{start_line}-{end_line}"
context_part += f"\n函数名:{func_name}"
if class_name and class_name != "None":
context_part += f"{class_name}类)"
context_part += f",由{lang}实现\n"
context_part += f"函数签名:接收{params}个参数,返回类型为{return_type}\n"
if docstring:
context_part += f"函数说明:{docstring}\n"
context_part += f"具体实现:\n{func_body}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"原始文本:\n{text}"
elif intent == "functionality_question":
# 功能询问意图,强调函数名、文档、参数和返回类型
context_part = f"【参考信息{result_id}】函数{func_name}"
if class_name and class_name != "None":
context_part += f"{class_name}类)"
context_part += f"的功能说明:\n{docstring}\n"
context_part += f"{lang}实现,接收{params}个参数,返回类型为{return_type}\n"
context_part += f"文件路径:{file_path},位置:{start_line}-{end_line}\n"
context_part += f"具体实现:\n{func_body}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"原始文本:\n{text}"
else:
# 其他意图,综合所有信息
context_part = f"【参考信息{result_id}】(来源:{file_path}"
context_part += f"\n函数:{func_name}"
if class_name and class_name != "None":
context_part += f"{class_name}类)"
context_part += f",语言:{lang}\n"
context_part += f"参数:{params}个,返回类型:{return_type}\n"
if docstring:
context_part += f"说明:{docstring}\n"
context_part += f"位置:{start_line}-{end_line}\n"
context_part += f"仓库:{repo_id},分支:{branch}\n"
context_part += f"实现:\n{func_body}\n"
context_part += f"原始文本:\n{text}"
context_parts.append(context_part)
return "\n\n".join(context_parts)
def generate_prompt(
self,
user_query: str,
intent_category: CodeIntentCategory,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
error_message: Optional[str] = None,
target_language: Optional[str] = None,
user_requirement: Optional[str] = None
) -> str:
"""
生成代码专用Prompt
Args:
user_query: 用户问题
intent_category: 代码意图分类
code_context: 代码上下文
conversation_history: 对话历史
error_message: 错误信息仅Bug修复场景
target_language: 目标编程语言仅代码生成场景
user_requirement: 用户需求仅代码生成场景
Returns:
str: 生成的Prompt
"""
try:
logger.info(f"生成代码Prompt意图分类: {intent_category}")
# 映射意图到Prompt类型
prompt_type = self._map_intent_to_prompt_type(intent_category)
logger.info(f"选择Prompt类型: {prompt_type}")
# 获取对应模板
template = self._templates.get(prompt_type)
if not template:
logger.warning(f"未找到对应Prompt模板: {prompt_type}")
template = self._templates[PromptTemplateType.GENERAL_QA]
# 准备参数
params = {
"user_query": user_query,
"code_context": code_context or "",
"conversation_history": self._build_conversation_history(conversation_history),
"error_message": error_message or "",
"target_language": target_language or "根据上下文判断",
"user_requirement": user_requirement or user_query,
"algorithm_code": self._extract_code_from_context(code_context) if code_context else ""
}
# 填充模板
prompt = template
for key, value in params.items():
placeholder = f"{{{key}}}"
prompt = prompt.replace(placeholder, value)
logger.info(f"Prompt生成完成长度: {len(prompt)}字符")
return prompt
except Exception as e:
logger.error(f"生成Prompt失败: {e}")
# 返回通用模板
return self._templates[PromptTemplateType.GENERAL_QA].format(
user_query=user_query,
code_context=code_context or "",
conversation_history=self._build_conversation_history(conversation_history),
error_message="",
target_language="根据上下文判断",
user_requirement=user_query,
algorithm_code=""
)
def generate_dynamic_prompt(
self,
user_query: str,
intent_result: Optional[Dict[str, Any]] = None,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
**kwargs
) -> str:
"""
生成动态Prompt基于意图识别结果
Args:
user_query: 用户问题
intent_result: 意图识别结果
code_context: 代码上下文
conversation_history: 对话历史
**kwargs: 其他参数
Returns:
str: 生成的动态Prompt
"""
try:
if intent_result:
# 从意图结果中提取分类
category_str = intent_result.get('category', 'unknown')
# 直接使用category_str因为CodeIntentCategory是一个普通的类不是枚举类型
intent_category = category_str
else:
# 默认使用通用分类
intent_category = CodeIntentCategory.UNKNOWN
# 提取其他参数
error_message = kwargs.get('error_message')
target_language = kwargs.get('target_language')
user_requirement = kwargs.get('user_requirement')
retrieved_results = kwargs.get('retrieved_results', [])
# 根据意图和 retrieved_results 构建增强的上下文
enhanced_context = code_context
if retrieved_results:
enhanced_context = self._build_enhanced_context(retrieved_results, intent_result)
# 生成Prompt
return self.generate_prompt(
user_query=user_query,
intent_category=intent_category,
code_context=enhanced_context,
conversation_history=conversation_history,
error_message=error_message,
target_language=target_language,
user_requirement=user_requirement
)
except Exception as e:
logger.error(f"生成动态Prompt失败: {e}")
# 返回通用Prompt
return self._templates[PromptTemplateType.GENERAL_QA].format(
user_query=user_query,
code_context=code_context or "",
conversation_history=self._build_conversation_history(conversation_history),
error_message="",
target_language="根据上下文判断",
user_requirement=user_query,
algorithm_code=""
)
def optimize_prompt(
self,
prompt: str,
max_length: int = 4000,
preserve_structure: bool = True
) -> str:
"""
优化Prompt长度
Args:
prompt: 原始Prompt
max_length: 最大长度
preserve_structure: 是否保留结构
Returns:
str: 优化后的Prompt
"""
if len(prompt) <= max_length:
return prompt
logger.warning(f"Prompt过长 ({len(prompt)} > {max_length}),需要优化")
if preserve_structure:
# 保留结构,只优化内容部分
# 1. 保留角色设定和核心指令
# 2. 精简分析要求
# 3. 缩短代码上下文
# 提取角色设定和核心指令
role_match = re.search(r'# 角色设定[\s\S]*?# 核心指令[\s\S]*?\n', prompt)
if role_match:
role_section = role_match.group(0)
else:
role_section = ""
# 提取分析要求
req_match = re.search(r'# 分析要求[\s\S]*?(?=# |$)', prompt)
if req_match:
req_section = req_match.group(0)
# 精简分析要求
req_lines = req_section.split('\n')
# 只保留前3条要求
req_section = '\n'.join(req_lines[:4]) # 保留标题和前3条
else:
req_section = ""
# 提取其他部分
rest_match = re.search(r'# (代码上下文|错误信息|对话历史|用户问题|输出格式)[\s\S]*$', prompt)
if rest_match:
rest_section = rest_match.group(0)
# 缩短代码上下文
if '# 代码上下文' in rest_section:
code_match = re.search(r'# 代码上下文[\s\S]*?(?=# |$)', rest_section)
if code_match:
code_section = code_match.group(0)
# 只保留前500个字符
if len(code_section) > 600:
code_lines = code_section.split('\n')
if len(code_lines) > 3:
# 保留标题和前几行
code_section = '\n'.join(code_lines[:2]) + '\n...\n(代码已截断)'
rest_section = rest_section.replace(code_match.group(0), code_section)
else:
rest_section = ""
optimized = role_section + '\n' + req_section + '\n' + rest_section
if len(optimized) > max_length:
# 进一步缩短
optimized = optimized[:max_length - 3] + '...'
else:
# 直接截断
optimized = prompt[:max_length - 3] + '...'
logger.info(f"Prompt优化完成长度: {len(optimized)}字符")
return optimized
def save_prompt_template(
self,
template_type: PromptTemplateType,
template_content: str,
description: Optional[str] = None
) -> bool:
"""
保存自定义Prompt模板
Args:
template_type: Prompt类型
template_content: 模板内容
description: 模板描述
Returns:
bool: 保存是否成功
"""
try:
# 这里可以扩展为持久化存储
# 目前只是在内存中更新
self._templates[template_type] = template_content
logger.info(f"保存Prompt模板成功: {template_type.value}")
return True
except Exception as e:
logger.error(f"保存Prompt模板失败: {e}")
return False
def get_prompt_template(self, template_type: PromptTemplateType) -> Optional[str]:
"""
获取Prompt模板
Args:
template_type: Prompt类型
Returns:
Optional[str]: 模板内容
"""
return self._templates.get(template_type)
def list_available_templates(self) -> List[Dict[str, Any]]:
"""
列出可用的Prompt模板
Returns:
List[Dict[str, Any]]: 模板列表
"""
templates = []
for template_type, content in self._templates.items():
templates.append({
"type": template_type.value,
"name": template_type.name,
"length": len(content),
"sample": content[:100] + "..." if len(content) > 100 else content
})
return templates
# 全局Prompt管理器实例
_prompt_manager = None
def get_prompt_manager() -> CodePromptManager:
"""
获取全局Prompt管理器实例
Returns:
CodePromptManager: Prompt管理器实例
"""
global _prompt_manager
if _prompt_manager is None:
_prompt_manager = CodePromptManager()
return _prompt_manager
def generate_code_prompt(
user_query: str,
intent_category: CodeIntentCategory,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
**kwargs
) -> str:
"""
生成代码专用Prompt
Args:
user_query: 用户问题
intent_category: 代码意图分类
code_context: 代码上下文
conversation_history: 对话历史
**kwargs: 其他参数
Returns:
str: 生成的Prompt
"""
manager = get_prompt_manager()
return manager.generate_prompt(
user_query=user_query,
intent_category=intent_category,
code_context=code_context,
conversation_history=conversation_history,
**kwargs
)
def generate_dynamic_code_prompt(
user_query: str,
intent_result: Optional[Dict[str, Any]] = None,
code_context: Optional[str] = None,
conversation_history: Optional[List[Dict[str, str]]] = None,
**kwargs
) -> str:
"""
生成动态代码Prompt
Args:
user_query: 用户问题
intent_result: 意图识别结果
code_context: 代码上下文
conversation_history: 对话历史
**kwargs: 其他参数
Returns:
str: 生成的动态Prompt
"""
manager = get_prompt_manager()
return manager.generate_dynamic_prompt(
user_query=user_query,
intent_result=intent_result,
code_context=code_context,
conversation_history=conversation_history,
**kwargs
)
if __name__ == "__main__":
"""测试代码"""
import asyncio
from utils.code_intent import CodeIntentDetector
async def test_prompt_generation():
"""测试Prompt生成"""
print("=" * 80)
print("测试代码Prompt生成")
print("=" * 80)
# 初始化管理器
manager = CodePromptManager()
detector = CodeIntentDetector()
# 测试用例
test_cases = [
{
"query": "这个函数是做什么的?如何使用它?",
"code": "def calculate_factorial(n):\n if n <= 1:\n return 1\n return n * calculate_factorial(n-1)",
"category": CodeIntentCategory.ENTITY_INTRODUCTION
},
{
"query": "为什么会报语法错误?",
"code": "for i in range(10)\n print(i)",
"error": "SyntaxError: invalid syntax",
"category": CodeIntentCategory.ERROR_DEBUGGING
},
{
"query": "如何实现快速排序算法?",
"category": CodeIntentCategory.CODE_GENERATION
},
{
"query": "如何优化这段代码的性能?",
"code": "def slow_function():\n result = []\n for i in range(100000):\n result.append(i * 2)\n return result",
"category": CodeIntentCategory.CODE_OPTIMIZATION
}
]
for i, test_case in enumerate(test_cases):
print(f"\n测试用例 {i+1}: {test_case['query']}")
print("-" * 60)
# 生成Prompt
prompt = manager.generate_prompt(
user_query=test_case['query'],
intent_category=test_case['category'],
code_context=test_case.get('code'),
error_message=test_case.get('error')
)
# 打印结果
print(f"Prompt类型: {manager._map_intent_to_prompt_type(test_case['category']).value}")
print(f"Prompt长度: {len(prompt)}字符")
print("\nPrompt内容:")
print(prompt[:300] + "..." if len(prompt) > 300 else prompt)
print("-" * 60)
print("\n" + "=" * 80)
print("测试完成")
print("=" * 80)
asyncio.run(test_prompt_generation())

View File

@ -61,6 +61,8 @@ class GitTool:
ssh_key_path = f"/tmp/ssh_key_{self.user_id}_{self.repo_id}"
with open(ssh_key_path, "w") as f:
f.write(self.ssh_key)
if not self.ssh_key.endswith('\n'):
f.write('\n')
os.chmod(ssh_key_path, 0o600)
os.environ["GIT_SSH_COMMAND"] = f"ssh -i {ssh_key_path} -o StrictHostKeyChecking=no"

160
utils/metadata_filter.py Normal file
View File

@ -0,0 +1,160 @@
#!/usr/bin/env python3
"""
Metadata过滤工具
"""
import re
import json
from typing import Dict, Any, Optional
from llama_index.llms.ollama import Ollama
from config import settings
class MetadataFilter:
"""
Metadata过滤工具类
"""
@staticmethod
def extract_filters(query: str) -> Dict[str, Any]:
"""
从查询字符串中提取metadata过滤条件
Args:
query: 用户查询字符串
Returns:
包含过滤条件的字典格式为 {metadata_key: filter_value}
"""
filters = {}
# 提取路径信息,例如:"在 data_structures 目录下"
path_patterns = [
r'\s*([^\s]+)\s*目录下',
r'\s*([^\s]+)\s*文件夹下',
r'路径\s*([^\s]+)'
]
for pattern in path_patterns:
match = re.search(pattern, query)
if match:
path = match.group(1)
filters['file_path'] = path
break
# 提取其他metadata字段的过滤条件
# 例如:"类型为python的文件"、"语言为java的代码"
type_patterns = [
r'类型为\s*([^\s]+)',
r'语言为\s*([^\s]+)'
]
for pattern in type_patterns:
match = re.search(pattern, query)
if match:
lang = match.group(1)
filters['lang'] = lang
break
# 可以根据需要添加更多的过滤条件提取规则
return filters
@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):
if value not in metadata_value:
return False
else:
if metadata_value != value:
return False
return True
@staticmethod
def extract_filters_with_llm(query: str, llm: Optional[Ollama] = None) -> Dict[str, Any]:
"""
使用LLM从查询字符串中提取metadata过滤条件
Args:
query: 用户查询字符串
llm: LLM实例如果为None则创建新实例
Returns:
包含过滤条件的字典格式为 {metadata_key: filter_value}
"""
if not llm:
llm = Ollama(
model=settings.OLLAMA_MODEL,
base_url=settings.OLLAMA_BASE_URL,
temperature=0.1,
request_timeout=30.0
)
# 构建提示词
prompt = f"""
请从以下用户查询中提取metadata过滤条件
{query}
请严格按照以下规则提取
1. 只提取明确在查询中提到的过滤条件不要进行任何推测
2. 只提取与以下metadata key一致的过滤条件
- func_id: 函数ID
- func_name: 函数名
- class_name: 类名
- file_path: 文件路径
- lang: 编程语言
- params: 参数数量
- return_type: 返回类型
- docstring: 文档字符串
- start_line: 开始行号
- end_line: 结束行号
- repo_id: 仓库ID
- branch: 分支名
- func_body: 函数体
3. 只有当查询中明确提到某个key的值时才将其包含在结果中
4. 例如对于查询"在 algorithms 目录下用java实现的排序算法"
只提取 {{"file_path": "algorithms", "lang": "java"}}不要提取其他任何key
请以JSON格式返回提取的过滤条件格式为
{{"file_path": "value", "lang": "value", ...}}
如果没有找到任何过滤条件请返回空对象{{}}
"""
# 调用LLM
response = llm.complete(prompt)
# 解析响应
try:
filters = json.loads(response.text)
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:
filtered_filters[key] = value
return filtered_filters
else:
return {}
except json.JSONDecodeError:
# 如果解析失败,回退到正则表达式提取
return MetadataFilter.extract_filters(query)

23
utils/prompt/__init__.py Normal file
View File

@ -0,0 +1,23 @@
"""Prompt模板包"""
from .intent_detector import INTENT_DETECTION_TEMPLATE
from .answer_generator import (
CODE_EXPLANATION_TEMPLATE,
CODE_DEBUGGING_TEMPLATE,
CODE_GENERATION_TEMPLATE,
ALGORITHM_EXPLANATION_TEMPLATE,
CODE_OPTIMIZATION_TEMPLATE,
GENERAL_QA_TEMPLATE
)
from .integrated_query_processing import INTEGRATED_QUERY_PROCESSING_TEMPLATE
__all__ = [
"INTENT_DETECTION_TEMPLATE",
"CODE_EXPLANATION_TEMPLATE",
"CODE_DEBUGGING_TEMPLATE",
"CODE_GENERATION_TEMPLATE",
"ALGORITHM_EXPLANATION_TEMPLATE",
"CODE_OPTIMIZATION_TEMPLATE",
"GENERAL_QA_TEMPLATE",
"INTEGRATED_QUERY_PROCESSING_TEMPLATE"
]

View File

@ -0,0 +1,17 @@
"""答案生成相关的Prompt模板"""
from .code_explanation import CODE_EXPLANATION_TEMPLATE
from .code_generation import CODE_GENERATION_TEMPLATE
from .algorithm_explanation import ALGORITHM_EXPLANATION_TEMPLATE
from .code_debugging import CODE_DEBUGGING_TEMPLATE
from .code_optimization import CODE_OPTIMIZATION_TEMPLATE
from .general_qa import GENERAL_QA_TEMPLATE
__all__ = [
"CODE_EXPLANATION_TEMPLATE",
"CODE_DEBUGGING_TEMPLATE",
"CODE_GENERATION_TEMPLATE",
"ALGORITHM_EXPLANATION_TEMPLATE",
"CODE_OPTIMIZATION_TEMPLATE",
"GENERAL_QA_TEMPLATE"
]

View File

@ -0,0 +1,45 @@
"""算法解释相关的Prompt模板"""
ALGORITHM_EXPLANATION_TEMPLATE = """# 角色设定
你是一位算法专家擅长深入解析算法原理和实现能够根据不同类型的问题提供精准的专业解答
## 核心指令
请基于提供的算法代码上下文和问题类型为用户提供详细专业的算法分析
## 分析要求(根据问题类型调整重点)
### 算法解释algorithm_explanation
- **算法原理**详细解释算法的基本原理和设计思想
- **算法步骤**逐步说明算法的执行流程
- **时间复杂度**分析时间复杂度最好平均最坏情况
- **空间复杂度**分析空间复杂度说明内存使用情况
- **优缺点**分析算法的优势和局限性
- **适用场景**说明算法的适用场景和典型应用
- **比较分析**与其他同类算法进行对比
### 算法实现algorithm_implementation
- **算法选择**选择最适合该问题的算法
- **实现细节**提供完整的算法实现代码
- **复杂度分析**分析实现的时间和空间复杂度
- **边界处理**考虑边界情况和特殊输入
- **优化建议**提供可能的优化方向
- **测试用例**提供测试算法的示例用例
### 数据结构data_structure
- **结构定义**详细说明数据结构的定义和特点
- **操作方法**说明数据结构支持的操作及其复杂度
- **实现方式**提供数据结构的实现代码
- **适用场景**说明数据结构的适用场景和典型应用
- **性能对比**与其他数据结构进行性能对比
- **使用示例**提供数据结构的使用示例
## 算法代码
{algorithm_code}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的分析"""

View File

@ -0,0 +1,46 @@
"""代码调试相关的Prompt模板"""
CODE_DEBUGGING_TEMPLATE = """# 角色设定
你是一位专业的代码调试专家擅长快速定位和解决代码中的错误能够根据不同类型的错误提供精准的诊断和解决方案
## 核心指令
请基于提供的代码错误信息上下文和错误类型为用户提供详细的错误分析和解决方案
## 分析要求(根据错误类型调整重点)
### 语法错误syntax_error
- **错误定位**准确指出语法错误的具体位置行号列号
- **错误原因**解释违反了哪条语法规则
- **修正方案**提供修正后的完整代码
- **预防建议**说明如何避免类似的语法错误
- **常见模式**列出该语法错误的常见触发场景
### 运行时错误runtime_error
- **错误分析**详细分析异常类型和错误信息
- **堆栈跟踪**解释错误堆栈中的关键信息
- **根本原因**深入分析导致错误的根本原因
- **修复方案**提供具体的修复代码和实施步骤
- **异常处理**建议如何添加异常处理来预防此类错误
- **测试建议**说明如何测试修复是否有效
### 调试问题debugging
- **调试方法**提供适合该问题的调试策略
- **断点设置**建议在哪些位置设置断点
- **日志分析**说明如何通过日志分析问题
- **变量检查**建议检查哪些关键变量的值
- **逐步排查**提供逐步排查问题的流程
- **工具推荐**推荐适合的调试工具和技巧
## 代码上下文
{code_context}
## 错误信息
{error_message}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的分析"""

View File

@ -0,0 +1,41 @@
"""代码解释相关的Prompt模板"""
CODE_EXPLANATION_TEMPLATE = """# 角色设定
你是一位资深的代码分析专家擅长深入解析代码结构和功能能够根据不同类型的问题提供精准的专业解答
## 核心指令
请基于提供的代码上下文和用户问题类型为用户提供详细专业的代码分析
## 分析要求(根据问题类型调整重点)
### 逻辑解释问题logic_explanation
- **功能说明**详细解释代码的功能和用途
- **逻辑分析**说明代码的执行流程和核心逻辑
- **实现原理**解释代码的实现原理和技术细节
- **使用示例**提供实际可运行的代码示例
- **注意事项**指出使用时需要注意的要点和常见错误
### 实体介绍问题entity_introduction
- **实体结构**说明函数API等实体的结构和组成
- **参数分析**说明每个参数的类型含义和默认值
- **返回值说明**解释返回值的类型含义和可能的取值
- **使用示例**提供实际可运行的代码示例
- **注意事项**指出使用时需要注意的要点和常见错误
### 代码结构问题code_structure
- **项目结构**详细说明项目的目录和文件组织
- **模块划分**说明各个模块的功能和职责
- **依赖关系**解释模块之间的依赖和调用关系
- **架构设计**说明项目的整体架构和设计思路
- **文件说明**解释关键文件的作用和内容
## 代码上下文
{code_context}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的分析"""

View File

@ -0,0 +1,58 @@
"""代码生成相关的Prompt模板"""
CODE_GENERATION_TEMPLATE = '''# 角色设定
你是一位经验丰富的代码生成专家擅长根据不同类型的需求编写高质量可维护的代码
## 核心指令
请基于用户的需求上下文和生成类型为用户提供完整专业可直接使用的代码
## 代码生成要求(根据生成类型调整重点)
### 完整代码生成code_generation
- **需求分析**深入理解用户的功能需求
- **架构设计**设计合理的代码结构和模块划分
- **完整实现**提供完整可运行的代码包括所有必要的导入
- **最佳实践**遵循目标语言的编码规范和最佳实践
- **错误处理**添加适当的异常处理和边界检查
- **代码注释**添加清晰的注释解释关键逻辑
- **使用示例**提供如何使用该代码的示例
### 函数实现function_implementation
- **函数签名**设计清晰的函数名参数和返回值
- **参数验证**添加参数类型检查和验证逻辑
- **边界处理**考虑边界情况和特殊输入
- **错误处理**使用适当的异常处理机制
- **文档字符串**添加详细的docstring说明函数用途
- **类型提示**使用类型注解提高代码可读性
- **单元测试**提供简单的测试用例
### 类实现class_implementation
- **类设计**设计合理的类结构和方法划分
- **构造函数**实现__init__方法正确初始化属性
- **封装性**合理使用私有属性和公共方法
- **方法实现**实现所有必要的方法确保功能完整
- **特殊方法**根据需要实现__str____repr__等特殊方法
- **文档字符串**为类和主要方法添加docstring
- **使用示例**提供类的使用示例
## 通用代码质量要求
1. **语法正确性**确保代码语法完全正确可直接运行
2. **代码风格**遵循PEP 8Python或其他语言的编码规范
3. **可读性**使用有意义的变量名和函数名添加必要的注释
4. **可维护性**代码结构清晰易于理解和修改
5. **性能考虑**在保证正确性的前提下考虑性能优化
6. **安全性**注意常见的安全问题如SQL注入XSS等
## 目标语言
{target_language}
## 代码上下文
{code_context}
## 对话历史
{conversation_history}
## 用户需求
{user_requirement}
请开始生成代码'''

View File

@ -0,0 +1,76 @@
"""代码优化相关的Prompt模板"""
CODE_OPTIMIZATION_TEMPLATE = """你是一个专业的代码优化专家。请根据用户的问题和检索到的代码上下文,提供专业的优化建议。
## 意图上下文
当前意图{{intent_type}}
检索策略{{retrieval_strategy}}
## 对话历史
{{conversation_history}}
## 检索上下文
{{code_context}}
## 响应指南
1. **必须使用标准 Markdown 格式**
2. **确保流式输出体验**首句直接入题段落间使用 \n\n 分隔
3. **中英文之间自动添加空格**
4. **代码块必须闭合且标注语言**
5. **针对 {{intent_type}} 采用对应的回答模板**
## 回答模板
### 代码优化问题
当用户询问如何优化代码时请按以下结构回答
#### 1. 代码分析
- 指出当前代码存在的问题
- 分析性能瓶颈
- 说明可优化的地方
#### 2. 优化建议
- 提供具体的优化方案
- 说明优化原理
- 给出优化后的代码示例
#### 3. 性能对比
- 对比优化前后的性能
- 说明优化的效果
- 给出具体的性能指标
#### 4. 最佳实践
- 提供相关的编程最佳实践
- 说明代码规范
- 给出可维护性建议
### 性能调优问题
当用户询问性能调优时请按以下结构回答
#### 1. 性能分析
- 分析当前性能问题
- 定位性能瓶颈
- 说明影响性能的因素
#### 2. 调优策略
- 提供具体的调优方案
- 说明调优的原理
- 给出调优的步骤
#### 3. 优化效果
- 说明调优后的效果
- 给出性能提升的数据
- 对比调优前后的差异
#### 4. 注意事项
- 说明调优时的注意事项
- 提供避免问题的建议
- 给出监控和评估方法
## 代码示例要求
- 代码块必须使用 ```python ```javascript 等标注语言
- 代码必须完整可运行
- 添加必要的注释说明
- 保持代码风格一致
现在请根据用户问题和检索上下文提供专业的代码优化建议"""

View File

@ -0,0 +1,99 @@
"""并发编程相关的Prompt模板"""
CONCURRENCY_EXPLANATION_TEMPLATE = '''你是一个专业的并发编程专家。请根据用户的问题和检索到的代码上下文,提供专业的并发编程指导。
## 意图上下文
当前意图{{intent_type}}
检索策略{{retrieval_strategy}}
## 对话历史
{{conversation_history}}
## 检索上下文
{{code_context}}
## 响应指南
1. **必须使用标准 Markdown 格式**
2. **确保流式输出体验**首句直接入题段落间使用 \n\n 分隔
3. **中英文之间自动添加空格**
4. **代码块必须闭合且标注语言**
5. **针对 {{intent_type}} 采用对应的回答模板**
## 回答模板
### 并发问题
当用户询问并发问题时请按以下结构回答
#### 1. 并发模型
- 说明并发的基本概念
- 解释并发模型
- 对比不同并发模型
#### 2. 实现方式
- 提供具体的实现代码
- 说明实现原理
- 给出使用示例
#### 3. 同步机制
- 说明同步的必要性
- 提供同步方法
- 解释同步原理
#### 4. 竞争条件处理
- 说明竞争条件的概念
- 提供避免竞争条件的方法
- 给出并发设计模式
### 线程问题
当用户询问线程问题时请按以下结构回答
#### 1. 线程基础
- 说明线程的概念
- 解释线程的生命周期
- 说明线程的创建和管理
#### 2. 线程同步
- 说明线程同步的必要性
- 提供同步方法信号量等
- 给出同步示例代码
#### 3. 死锁预防
- 说明死锁的概念
- 提供死锁预防方法
- 给出避免死锁的最佳实践
#### 4. 线程安全
- 说明线程安全的概念
- 提供线程安全的实现方法
- 给出线程安全的编程建议
### 异步编程问题
当用户询问异步编程时请按以下结构回答
#### 1. 异步编程模型
- 说明异步编程的概念
- 解释异步与同步的区别
- 说明异步的优势
#### 2. 异步实现
- 提供异步编程的代码示例
- 说明 async/await 的使用
- 解释事件循环的原理
#### 3. 异步 I/O
- 说明异步 I/O 的概念
- 提供异步 I/O 的实现方法
- 给出异步 I/O 的使用示例
#### 4. 错误处理
- 说明异步编程中的错误处理
- 提供异常处理的方法
- 给出调试和测试建议
## 代码示例要求
- 代码块必须使用 ```python ```javascript 等标注语言
- 代码必须完整可运行
- 添加必要的注释说明
- 展示并发/异步的完整流程
现在请根据用户问题和检索上下文提供专业的并发编程指导'''

View File

@ -0,0 +1,99 @@
"""测试部署相关的Prompt模板"""
DEPLOYMENT_EXPLANATION_TEMPLATE = """你是一个专业的测试和部署专家。请根据用户的问题和检索到的代码上下文,提供专业的测试和部署指导。
## 意图上下文
当前意图{{intent_type}}
检索策略{{retrieval_strategy}}
## 对话历史
{{conversation_history}}
## 检索上下文
{{code_context}}
## 响应指南
1. **必须使用标准 Markdown 格式**
2. **确保流式输出体验**首句直接入题段落间使用 \n\n 分隔
3. **中英文之间自动添加空格**
4. **代码块必须闭合且标注语言**
5. **针对 {{intent_type}} 采用对应的回答模板**
## 回答模板
### 测试问题
当用户询问测试问题时请按以下结构回答
#### 1. 测试类型
- 说明不同类型的测试单元测试集成测试端到端测试
- 解释各种测试的适用场景
- 提供测试策略建议
#### 2. 测试框架
- 推荐适合的测试框架
- 说明框架的特点和优势
- 提供框架的使用示例
#### 3. 测试实现
- 提供具体的测试代码
- 说明测试的编写方法
- 给出测试的最佳实践
#### 4. Mock 和测试数据
- 说明 Mock 的使用场景
- 提供 Mock 的实现方法
- 给出测试数据的准备策略
### 部署问题
当用户询问部署问题时请按以下结构回答
#### 1. 部署策略
- 说明不同的部署方式手动部署自动化部署
- 解释部署的流程
- 提供部署策略建议
#### 2. 环境配置
- 说明开发测试生产环境的配置
- 提供环境变量的管理方法
- 给出配置文件的组织方式
#### 3. CI/CD 流程
- 说明 CI/CD 的概念
- 提供主流 CI/CD 工具的使用方法
- 给出 CI/CD 流程的配置示例
#### 4. 监控和告警
- 说明监控的重要性
- 提供监控工具的推荐
- 给出告警策略的配置方法
### 配置问题
当用户询问配置问题时请按以下结构回答
#### 1. 配置文件格式
- 说明不同配置文件的格式JSONYAMLINI
- 解释各种格式的优缺点
- 提供格式选择的建议
#### 2. 环境变量管理
- 说明环境变量的使用场景
- 提供环境变量的管理方法
- 给出环境变量的最佳实践
#### 3. 参数配置
- 说明参数配置的原则
- 提供参数验证的方法
- 给出参数管理的建议
#### 4. 配置验证
- 说明配置验证的重要性
- 提供配置验证的方法
- 给出配置错误的处理建议
## 代码示例要求
- 代码块必须使用 ```bash```yaml```python 等标注语言
- 配置文件必须完整且格式正确
- 添加必要的注释说明
- 提供可执行的命令或脚本
现在请根据用户问题和检索上下文提供专业的测试和部署指导"""

View File

@ -0,0 +1,58 @@
"""通用代码相关的Prompt模板"""
GENERAL_QA_TEMPLATE = """# 角色设定
你是一位专业的代码顾问能够回答各种代码相关问题擅长根据不同类型的问题提供精准的专业解答
## 核心指令
请基于提供的代码上下文和问题类型为用户提供全面准确专业的回答
## 回答要求(根据问题类型调整重点)
### 测试问题testing
- **测试方法**说明适合的测试方法单元测试集成测试等
- **测试框架**推荐适合的测试框架如pytestunittest等
- **测试用例**提供具体的测试用例示例
- **Mock技术**说明如何mock外部依赖
- **覆盖率**解释测试覆盖率的概念和如何提高覆盖率
- **最佳实践**提供测试的最佳实践和常见陷阱
### 部署问题deployment
- **部署策略**说明适合的部署方式容器化云部署等
- **环境配置**详细说明环境变量的配置方法
- **CI/CD流程**解释持续集成和持续部署的流程
- **依赖管理**说明如何管理生产环境的依赖
- **监控告警**建议部署后的监控和告警方案
- **回滚策略**说明如何处理部署失败的情况
### 配置问题configuration
- **配置文件**说明配置文件的格式和位置
- **环境变量**解释如何设置和使用环境变量
- **参数配置**详细说明各个配置参数的含义和取值
- **配置验证**提供验证配置是否正确的方法
- **常见问题**列出配置相关的常见错误和解决方案
- **最佳实践**提供配置管理的最佳实践
### 通用知识问题general_knowledge
- **概念解释**详细解释相关概念和术语
- **原理说明**深入说明技术原理和机制
- **应用场景**说明技术的适用场景和典型应用
- **发展趋势**介绍技术的发展趋势和未来方向
- **学习资源**推荐相关的学习资源和文档
- **实践建议**提供实际应用的建议和注意事项
### 非技术问题non_technical
- **友好回应**保持友好自然的对话风格
- **相关信息**提供与问题相关的有用信息
- **引导澄清**如果问题模糊引导用户明确需求
- **保持自然**避免过度技术化保持对话的自然流畅
## 代码上下文
{code_context}
## 对话历史
{conversation_history}
## 用户问题
{user_query}
请开始你的回答"""

View File

@ -0,0 +1,148 @@
"""集成查询处理Prompt模板
整合意图识别Metadata过滤条件提取和查询转换功能
"""
INTEGRATED_QUERY_PROCESSING_TEMPLATE = """### 角色定义
你是一个全面的查询处理助手需要完成以下三个任务
1. 代码意图识别分析用户问题的意图类型
2. 元数据过滤条件提取从问题中提取显示限制的过滤条件
3. 查询转换重写查询生成更广泛的查询分解复杂查询
### 对话历史
{history_str}
### 当前用户问题
{query}
---
### 任务1代码意图识别
请分析用户问题的意图判断其属于以下分类之一
- logic_explanation解释既有代码的底层逻辑
- entity_introduction介绍具体的代码实体
- code_structure询问项目组织
- code_generation请求从零编写完整代码或功能块
- boilerplate_implementation请求提供标准算法/模板
- error_debugging排查 Bug 或异常
- code_optimization改进既有代码的性能或质量
- algorithm_theory算法原理或复杂度分析
- general_technical通用技术咨询
- non_technical非技术问题
- unknown未知类型
#### 分类决策树 (判定逻辑)
在判定分类前请严格执行以下优先级逻辑
1. **上下文回溯**如果 query 中提到的实体函数变量类名在对话历史或上下文代码中出现过优先判定为代码解释/架构类
2. **句式辨析**
- **[实体/功能] 是怎么实现的/怎么做的** -> 倾向于代码解释语态为"对既有状态的追溯"
- **怎么实现 [功能]/ 帮我写一个...** -> 倾向于代码生成语态为"对未知实现的请求"
3. **理论深度**若问题涉及性能瓶颈数学原理或复杂度优先归类为算法与优化类
#### 语义微调示例 (Few-Shot)
- **输入**: "find_median 是如何实现的?"
**判定**: logic_explanation | **原因**: 指向特定函数名且询问其现状
- **输入**: "如何实现查找中位数的算法?"
**判定**: code_generation | **原因**: 泛指功能实现表现为编程请求
- **输入**: "这段代码能跑快一点吗?"
**判定**: code_optimization | **原因**: 基于现有代码的性能改进请求
- **输入**: "什么是深度优先搜索?"
**判定**: algorithm_theory | **原因**: 概念性理论询问
---
### 任务2严格元数据过滤条件提取
请从用户问题中提取显示限制的metadata过滤条件只提取与以下key一致的条件
- func_id: 函数ID
- func_name: 函数名
- class_name: 类名
- file_path: 文件路径
- lang: 编程语言
- params: 参数数量
- return_type: 返回类型
- docstring: 文档字符串
- start_line: 开始行号
- end_line: 结束行号
- repo_id: 仓库ID
- branch: 分支名
- func_body: 函数体
**重要规则**
- 只提取查询中**明确提到**的条件不要进行任何推测
- 只有当查询中明确使用了与某个key相关的词汇时才提取该key的value
- **value必须为小写**
- **一个key只对应一个value**
- **value的字符串长度尽可能短**
- 例如对于查询"在 algorithms 目录下用java实现的排序算法"
只提取 {{"lang": "java"}}不要提取其他任何key
**强制性约束**
1. **零推测原则**仅提取用户明确指定的属性限定若用户说计算斐波那契的函数由于未指定函数名文件名或语言提取结果应为空 `{{}}`
2. **关键词触发**
- 提取 `file_path`原文必须包含路径特征 .py, /path, 文件夹等
- 提取 `func_name` / `class_name`原文必须包含名为叫作或明显的标识符引用
- 提取 `return_type` / `params`原文必须明确提到返回类型为...参数个数为...
3. **格式规范**value 一律小写保持极简严禁包含任何描述性文字
---
### 任务3查询转换
请完成以下三个转换
#### 3.1 重写查询
将查询重写为更具体详细且对RAG系统中的信息检索更有效的形式
- 更具体和详细
- 如果适用包含来自对话历史的相关上下文
- 保持原始意图
- 适合向量搜索
#### 3.2 生成更广泛的查询
生成给定用户查询的更广泛版本以帮助在RAG系统中检索更全面的上下文信息
- 涵盖与原始查询相关的更一般方面
- 能够帮助检索相关的背景信息
- 保持原始查询的核心主题
- 适合向量搜索
#### 3.3 分解查询
将复杂用户查询分解为更简单更集中的子查询这些子查询可用于RAG系统中的全面信息检索
- 2-5个更简单的子查询
- 每个子查询应关注原始查询的特定方面
- 所有子查询一起应涵盖整个原始查询
- 每个子查询应适合向量搜索
---
### 输出格式要求
请以JSON格式返回所有结果包含以下字段
{{
"intent": {{
"is_code_related": true/false,
"category": "分类名称",
"confidence": 0.0-1.0,
"keywords": ["关键词列表"],
"reasoning": "分类理由",
"requires_code_context": true/false,
"suggested_search_terms": ["搜索词列表"]
}},
"filters": {{
"file_path": "value",
"lang": "value",
...
}},
"transformed": {{
"rewritten": "重写后的查询",
"backward": "更广泛的查询",
"sub_queries": ["子查询1", "子查询2", ...]
}}
}}
### 输出规则
1. 必须输出有效的JSON格式不要包含其他内容
2. confidence表示分类的置信度范围0.0-1.0
3. keywords从问题和对话历史中提取的关键词最多8个
4. reasoning简要说明为什么这样分类要考虑对话历史的内容
5. requires_code_context表示是否需要代码上下文来回答
6. suggested_search_terms建议的检索词最多5个要考虑对话历史中提到的技术或库
7. 只提取明确提到的信息不要进行推测
8. 确保所有字段都有合理的值
现在请分析用户问题和对话历史并输出JSON结果"""

View File

@ -0,0 +1,82 @@
"""意图识别相关的Prompt模板"""
INTENT_DETECTION_TEMPLATE = """### 角色定义
你是一个高精度的代码意图识别助手你的核心任务是分析用户问题与对话历史精准区分用户是想了解现有的代码逻辑还是请求编写新的代码
### 对话历史
{history_str}
### 当前用户问题
{query}
---
### 分类决策树 (判定逻辑)
在判定分类前请严格执行以下优先级逻辑
1. **上下文回溯**如果 query 中提到的实体函数变量类名在对话历史或上下文代码中出现过优先判定为代码解释/架构类
2. **句式辨析**
- **[实体/功能] 是怎么实现的/怎么做的** -> 倾向于代码解释语态为对既有状态的追溯
- **怎么实现 [功能]/ 帮我写一个...** -> 倾向于代码生成语态为对未知实现的请求
3. **理论深度**若问题涉及性能瓶颈数学原理或复杂度优先归类为算法与优化类
---
### 详细分类标准
#### 1. 代码解释与逻辑类 (Existing Code Focus)
- **logic_explanation**解释**既有代码**的底层逻辑例如这段循环是怎么工作的kth_number 是怎么实现的
- **entity_introduction**介绍具体的**代码实体**函数定义类属性API参数例如这个类的构造函数接收什么参数
- **code_structure**询问**项目组织**例如数据处理模块在哪个文件里这个项目的架构是怎么设计的
#### 2. 代码生成与实现类 (New Code Focus)
- **code_generation**请求**从零编写**完整代码或功能块例如帮我写一个解析 XML 的脚本怎么实现一个 LRU 缓存
- **boilerplate_implementation**请求提供**标准算法/模板**例如请给出快速排序的 Python 实现
#### 3. 调试、优化与理论类
- **error_debugging**排查 **Bug 或异常**例如这段代码为什么报空指针
- **code_optimization****改进**既有代码的性能或质量例如如何重构这段代码以减少内存占用
- **algorithm_theory**算法**原理或复杂度**分析例如这个排序的时间复杂度是多少
#### 4. 非代码类
- **general_technical**通用技术咨询例如环境配置Git 命令部署流程
- **non_technical**非技术问题例如天气闲聊常识
---
### 语义微调示例 (Few-Shot)
- **输入**: "find_median 是如何实现的?"
**判定**: logic_explanation | **原因**: 指向特定函数名且询问其现状
- **输入**: "如何实现查找中位数的算法?"
**判定**: code_generation | **原因**: 泛指功能实现表现为编程请求
- **输入**: "这段代码能跑快一点吗?"
**判定**: code_optimization | **原因**: 基于现有代码的性能改进请求
- **输入**: "什么是深度优先搜索?"
**判定**: algorithm_theory | **原因**: 概念性理论询问
---
### 输出要求
请直接输出 JSON 格式禁止包含任何说明文本
{{
"is_code_related": true/false,
"category": "上述分类名称",
"confidence": 0.0-1.0,
"analysis": {{
"referring_to_context": true/false, // 是否引用了上下文中已有的代码实体
"intent_type": "query_existing_logic | request_new_code | other"
}},
"reasoning": "结合语境说明分类依据(需体现对‘解释既有’与‘请求生成’的辨析)",
"keywords": ["关键词"],
"requires_code_context": true/false, // 是否需要代码上下文来回答
"suggested_search_terms": ["搜索词"]
}}
## 输出规则
1. 必须输出有效的JSON格式不要包含其他内容
2. confidence表示分类的置信度范围0.0-1.0
3. keywords从问题和对话历史中提取的关键词最多8个
4. reasoning简要说明为什么这样分类要考虑对话历史的内容
5. requires_code_context表示是否需要代码上下文来回答
6. suggested_search_terms建议的检索词最多5个要考虑对话历史中提到的技术或库
现在请分析用户问题和对话历史并输出JSON结果"""

View File

@ -0,0 +1,64 @@
"""
Chinese prompts for query transformation
"""
# Query rewriting prompt in Chinese
REWRITE_PROMPT = """
你是一名查询重写专家你的任务是将给定的用户查询重写为更具体详细且对RAG系统中的信息检索更有效的形式
原始查询{query}
对话历史如果有{history}
请将查询重写为
1. 更具体和详细
2. 如果适用包含来自对话历史的相关上下文
3. 保持原始意图
4. 适合向量搜索
**输出格式要求**
- 只输出重写后的查询内容不要包含任何其他文字或解释
- 不要包含"重写后的查询:"这样的前缀
- 确保输出是一个完整的语法正确的查询语句
"""
# Backward prompt generation prompt in Chinese
BACKWARD_PROMPT = """
你是一名上下文检索专家你的任务是生成给定用户查询的更广泛版本以帮助在RAG系统中检索更全面的上下文信息
原始查询{query}
对话历史如果有{history}
请生成一个更广泛的查询该查询
1. 涵盖与原始查询相关的更一般方面
2. 能够帮助检索相关的背景信息
3. 保持原始查询的核心主题
4. 适合向量搜索
**输出格式要求**
- 只输出更广泛的查询内容不要包含任何其他文字或解释
- 不要包含"更广泛的查询:"这样的前缀
- 确保输出是一个完整的语法正确的查询语句
"""
# Query decomposition prompt in Chinese
DECOMPOSE_PROMPT = """
你是一名查询分解专家你的任务是将给定的复杂用户查询分解为更简单更集中的子查询这些子查询可用于RAG系统中的全面信息检索
原始查询{query}
对话历史如果有{history}
请将查询分解为
1. 2-5个更简单的子查询
2. 每个子查询应关注原始查询的特定方面
3. 所有子查询一起应涵盖整个原始查询
4. 每个子查询应适合向量搜索
**输出格式要求**
- 只输出编号列表的子查询每行一个
- 不要包含"子查询:"这样的前缀
- 每个子查询应以数字加英文点号开头1. 2. 3.
- 确保每个子查询都是一个完整的语法正确的查询语句
"""

418
utils/query_processor.py Normal file
View File

@ -0,0 +1,418 @@
#!/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))

150
utils/query_transformer.py Normal file
View File

@ -0,0 +1,150 @@
"""
Query transformer for RAG system
Implements three query transformation techniques:
1. Query rewriting: Make queries more specific and detailed
2. Backward prompt generation: Generate broader queries for context retrieval
3. Sub-query decomposition: Break down complex queries into simpler components
"""
from typing import List, Dict, Optional, Tuple
from loguru import logger
from config import settings
from llama_index.llms.ollama import Ollama
from utils.prompt.query_trans import REWRITE_PROMPT, BACKWARD_PROMPT, DECOMPOSE_PROMPT
class QueryTransformer:
"""
Query transformer for RAG system
"""
def __init__(self, llm: Ollama):
"""
Initialize query transformer
Args:
llm: Ollama LLM instance for generating transformations
"""
self.llm = llm
async def rewrite_query(self, query: str, history: str = "") -> str:
"""
Rewrite query to be more specific and detailed
Args:
query: Original user query
history: Chat history (optional)
Returns:
Rewritten query
"""
try:
prompt = REWRITE_PROMPT.format(query=query, history=history)
response = await self.llm.acomplete(prompt=prompt)
rewritten_query = response.text.strip()
logger.info(f"Rewritten query: {rewritten_query}")
return rewritten_query
except Exception as e:
logger.error(f"Error in query rewriting: {e}")
return query
async def backward_query(self, query: str, history: str = "") -> str:
"""
Generate broader query for context retrieval
Args:
query: Original user query
history: Chat history (optional)
Returns:
Broader query for context retrieval
"""
try:
prompt = BACKWARD_PROMPT.format(query=query, history=history)
response = await self.llm.acomplete(prompt=prompt)
backward_query = response.text.strip()
logger.info(f"Backward query: {backward_query}")
return backward_query
except Exception as e:
logger.error(f"Error in backward query generation: {e}")
return query
async def decompose_query(self, query: str, history: str = "") -> List[str]:
"""
Decompose complex query into simpler sub-queries
Args:
query: Original user query
history: Chat history (optional)
Returns:
List of sub-queries
"""
try:
prompt = DECOMPOSE_PROMPT.format(query=query, history=history)
response = await self.llm.acomplete(prompt=prompt)
sub_queries_text = response.text.strip()
# Parse sub-queries from response
sub_queries = []
for line in sub_queries_text.split('\n'):
line = line.strip()
if line and (line.startswith('1.') or line.startswith('2.') or line.startswith('3.') or line.startswith('4.') or line.startswith('5.')):
sub_query = line.split('.', 1)[1].strip()
if sub_query:
sub_queries.append(sub_query)
# If no valid sub-queries found, return original query as single sub-query
if not sub_queries:
sub_queries = [query]
logger.info(f"Original query: {query}")
logger.info(f"Decomposed sub-queries: {sub_queries}")
return sub_queries
except Exception as e:
logger.error(f"Error in query decomposition: {e}")
return [query]
async def transform_query(self, query: str, history: str = "") -> Dict[str, any]:
"""
Transform query using all three techniques
Args:
query: Original user query
history: Chat history (optional)
Returns:
Dict with transformed queries
"""
try:
logger.info(f"Transforming query: {query}")
# Run all transformations in parallel
from asyncio import gather
rewritten, backward, sub_queries = await gather(
self.rewrite_query(query, history),
self.backward_query(query, history),
self.decompose_query(query, history)
)
# Combine all transformed queries for comprehensive retrieval
all_transformed_queries = [rewritten, backward] + sub_queries
return {
"original": query,
"rewritten": rewritten,
"backward": backward,
"sub_queries": sub_queries,
"all_transformed": all_transformed_queries
}
except Exception as e:
logger.error(f"Error in query transformation: {e}")
return {
"original": query,
"rewritten": query,
"backward": query,
"sub_queries": [query],
"all_transformed": [query]
}