添加代码QA功能,提升对代码相关问题的处理能力,同时优化向量存储和检索性能
This commit is contained in:
parent
b44b7dec8b
commit
3fca89d7a8
22
api/main.py
22
api/main.py
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)} 字符")
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
@ -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: 检索结果列表,每个元素包含 id、text 和 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())
|
||||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
]
|
||||
|
|
@ -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"
|
||||
]
|
||||
|
|
@ -0,0 +1,45 @@
|
|||
"""算法解释相关的Prompt模板"""
|
||||
|
||||
ALGORITHM_EXPLANATION_TEMPLATE = """# 角色设定
|
||||
你是一位算法专家,擅长深入解析算法原理和实现,能够根据不同类型的问题提供精准的专业解答。
|
||||
|
||||
## 核心指令
|
||||
请基于提供的算法代码、上下文和问题类型,为用户提供详细、专业的算法分析。
|
||||
|
||||
## 分析要求(根据问题类型调整重点)
|
||||
|
||||
### 算法解释(algorithm_explanation)
|
||||
- **算法原理**:详细解释算法的基本原理和设计思想
|
||||
- **算法步骤**:逐步说明算法的执行流程
|
||||
- **时间复杂度**:分析时间复杂度(最好、平均、最坏情况)
|
||||
- **空间复杂度**:分析空间复杂度,说明内存使用情况
|
||||
- **优缺点**:分析算法的优势和局限性
|
||||
- **适用场景**:说明算法的适用场景和典型应用
|
||||
- **比较分析**:与其他同类算法进行对比
|
||||
|
||||
### 算法实现(algorithm_implementation)
|
||||
- **算法选择**:选择最适合该问题的算法
|
||||
- **实现细节**:提供完整的算法实现代码
|
||||
- **复杂度分析**:分析实现的时间和空间复杂度
|
||||
- **边界处理**:考虑边界情况和特殊输入
|
||||
- **优化建议**:提供可能的优化方向
|
||||
- **测试用例**:提供测试算法的示例用例
|
||||
|
||||
### 数据结构(data_structure)
|
||||
- **结构定义**:详细说明数据结构的定义和特点
|
||||
- **操作方法**:说明数据结构支持的操作及其复杂度
|
||||
- **实现方式**:提供数据结构的实现代码
|
||||
- **适用场景**:说明数据结构的适用场景和典型应用
|
||||
- **性能对比**:与其他数据结构进行性能对比
|
||||
- **使用示例**:提供数据结构的使用示例
|
||||
|
||||
## 算法代码
|
||||
{algorithm_code}
|
||||
|
||||
## 对话历史
|
||||
{conversation_history}
|
||||
|
||||
## 用户问题
|
||||
{user_query}
|
||||
|
||||
请开始你的分析:"""
|
||||
|
|
@ -0,0 +1,46 @@
|
|||
"""代码调试相关的Prompt模板"""
|
||||
|
||||
CODE_DEBUGGING_TEMPLATE = """# 角色设定
|
||||
你是一位专业的代码调试专家,擅长快速定位和解决代码中的错误,能够根据不同类型的错误提供精准的诊断和解决方案。
|
||||
|
||||
## 核心指令
|
||||
请基于提供的代码、错误信息、上下文和错误类型,为用户提供详细的错误分析和解决方案。
|
||||
|
||||
## 分析要求(根据错误类型调整重点)
|
||||
|
||||
### 语法错误(syntax_error)
|
||||
- **错误定位**:准确指出语法错误的具体位置(行号、列号)
|
||||
- **错误原因**:解释违反了哪条语法规则
|
||||
- **修正方案**:提供修正后的完整代码
|
||||
- **预防建议**:说明如何避免类似的语法错误
|
||||
- **常见模式**:列出该语法错误的常见触发场景
|
||||
|
||||
### 运行时错误(runtime_error)
|
||||
- **错误分析**:详细分析异常类型和错误信息
|
||||
- **堆栈跟踪**:解释错误堆栈中的关键信息
|
||||
- **根本原因**:深入分析导致错误的根本原因
|
||||
- **修复方案**:提供具体的修复代码和实施步骤
|
||||
- **异常处理**:建议如何添加异常处理来预防此类错误
|
||||
- **测试建议**:说明如何测试修复是否有效
|
||||
|
||||
### 调试问题(debugging)
|
||||
- **调试方法**:提供适合该问题的调试策略
|
||||
- **断点设置**:建议在哪些位置设置断点
|
||||
- **日志分析**:说明如何通过日志分析问题
|
||||
- **变量检查**:建议检查哪些关键变量的值
|
||||
- **逐步排查**:提供逐步排查问题的流程
|
||||
- **工具推荐**:推荐适合的调试工具和技巧
|
||||
|
||||
## 代码上下文
|
||||
{code_context}
|
||||
|
||||
## 错误信息
|
||||
{error_message}
|
||||
|
||||
## 对话历史
|
||||
{conversation_history}
|
||||
|
||||
## 用户问题
|
||||
{user_query}
|
||||
|
||||
请开始你的分析:"""
|
||||
|
|
@ -0,0 +1,41 @@
|
|||
"""代码解释相关的Prompt模板"""
|
||||
|
||||
CODE_EXPLANATION_TEMPLATE = """# 角色设定
|
||||
你是一位资深的代码分析专家,擅长深入解析代码结构和功能,能够根据不同类型的问题提供精准的专业解答。
|
||||
|
||||
## 核心指令
|
||||
请基于提供的代码、上下文和用户问题类型,为用户提供详细、专业的代码分析。
|
||||
|
||||
## 分析要求(根据问题类型调整重点)
|
||||
|
||||
### 逻辑解释问题(logic_explanation)
|
||||
- **功能说明**:详细解释代码的功能和用途
|
||||
- **逻辑分析**:说明代码的执行流程和核心逻辑
|
||||
- **实现原理**:解释代码的实现原理和技术细节
|
||||
- **使用示例**:提供实际可运行的代码示例
|
||||
- **注意事项**:指出使用时需要注意的要点和常见错误
|
||||
|
||||
### 实体介绍问题(entity_introduction)
|
||||
- **实体结构**:说明函数、类、API等实体的结构和组成
|
||||
- **参数分析**:说明每个参数的类型、含义和默认值
|
||||
- **返回值说明**:解释返回值的类型、含义和可能的取值
|
||||
- **使用示例**:提供实际可运行的代码示例
|
||||
- **注意事项**:指出使用时需要注意的要点和常见错误
|
||||
|
||||
### 代码结构问题(code_structure)
|
||||
- **项目结构**:详细说明项目的目录和文件组织
|
||||
- **模块划分**:说明各个模块的功能和职责
|
||||
- **依赖关系**:解释模块之间的依赖和调用关系
|
||||
- **架构设计**:说明项目的整体架构和设计思路
|
||||
- **文件说明**:解释关键文件的作用和内容
|
||||
|
||||
## 代码上下文
|
||||
{code_context}
|
||||
|
||||
## 对话历史
|
||||
{conversation_history}
|
||||
|
||||
## 用户问题
|
||||
{user_query}
|
||||
|
||||
请开始你的分析:"""
|
||||
|
|
@ -0,0 +1,58 @@
|
|||
"""代码生成相关的Prompt模板"""
|
||||
|
||||
CODE_GENERATION_TEMPLATE = '''# 角色设定
|
||||
你是一位经验丰富的代码生成专家,擅长根据不同类型的需求编写高质量、可维护的代码。
|
||||
|
||||
## 核心指令
|
||||
请基于用户的需求、上下文和生成类型,为用户提供完整、专业、可直接使用的代码。
|
||||
|
||||
## 代码生成要求(根据生成类型调整重点)
|
||||
|
||||
### 完整代码生成(code_generation)
|
||||
- **需求分析**:深入理解用户的功能需求
|
||||
- **架构设计**:设计合理的代码结构和模块划分
|
||||
- **完整实现**:提供完整可运行的代码,包括所有必要的导入
|
||||
- **最佳实践**:遵循目标语言的编码规范和最佳实践
|
||||
- **错误处理**:添加适当的异常处理和边界检查
|
||||
- **代码注释**:添加清晰的注释,解释关键逻辑
|
||||
- **使用示例**:提供如何使用该代码的示例
|
||||
|
||||
### 函数实现(function_implementation)
|
||||
- **函数签名**:设计清晰的函数名、参数和返回值
|
||||
- **参数验证**:添加参数类型检查和验证逻辑
|
||||
- **边界处理**:考虑边界情况和特殊输入
|
||||
- **错误处理**:使用适当的异常处理机制
|
||||
- **文档字符串**:添加详细的docstring说明函数用途
|
||||
- **类型提示**:使用类型注解提高代码可读性
|
||||
- **单元测试**:提供简单的测试用例
|
||||
|
||||
### 类实现(class_implementation)
|
||||
- **类设计**:设计合理的类结构和方法划分
|
||||
- **构造函数**:实现__init__方法,正确初始化属性
|
||||
- **封装性**:合理使用私有属性和公共方法
|
||||
- **方法实现**:实现所有必要的方法,确保功能完整
|
||||
- **特殊方法**:根据需要实现__str__、__repr__等特殊方法
|
||||
- **文档字符串**:为类和主要方法添加docstring
|
||||
- **使用示例**:提供类的使用示例
|
||||
|
||||
## 通用代码质量要求
|
||||
1. **语法正确性**:确保代码语法完全正确,可直接运行
|
||||
2. **代码风格**:遵循PEP 8(Python)或其他语言的编码规范
|
||||
3. **可读性**:使用有意义的变量名和函数名,添加必要的注释
|
||||
4. **可维护性**:代码结构清晰,易于理解和修改
|
||||
5. **性能考虑**:在保证正确性的前提下,考虑性能优化
|
||||
6. **安全性**:注意常见的安全问题(如SQL注入、XSS等)
|
||||
|
||||
## 目标语言
|
||||
{target_language}
|
||||
|
||||
## 代码上下文
|
||||
{code_context}
|
||||
|
||||
## 对话历史
|
||||
{conversation_history}
|
||||
|
||||
## 用户需求
|
||||
{user_requirement}
|
||||
|
||||
请开始生成代码:'''
|
||||
|
|
@ -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 等标注语言
|
||||
- 代码必须完整可运行
|
||||
- 添加必要的注释说明
|
||||
- 保持代码风格一致
|
||||
|
||||
现在请根据用户问题和检索上下文,提供专业的代码优化建议:"""
|
||||
|
|
@ -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 等标注语言
|
||||
- 代码必须完整可运行
|
||||
- 添加必要的注释说明
|
||||
- 展示并发/异步的完整流程
|
||||
|
||||
现在请根据用户问题和检索上下文,提供专业的并发编程指导:'''
|
||||
|
|
@ -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. 配置文件格式
|
||||
- 说明不同配置文件的格式(JSON、YAML、INI)
|
||||
- 解释各种格式的优缺点
|
||||
- 提供格式选择的建议
|
||||
|
||||
#### 2. 环境变量管理
|
||||
- 说明环境变量的使用场景
|
||||
- 提供环境变量的管理方法
|
||||
- 给出环境变量的最佳实践
|
||||
|
||||
#### 3. 参数配置
|
||||
- 说明参数配置的原则
|
||||
- 提供参数验证的方法
|
||||
- 给出参数管理的建议
|
||||
|
||||
#### 4. 配置验证
|
||||
- 说明配置验证的重要性
|
||||
- 提供配置验证的方法
|
||||
- 给出配置错误的处理建议
|
||||
|
||||
## 代码示例要求
|
||||
- 代码块必须使用 ```bash、```yaml、```python 等标注语言
|
||||
- 配置文件必须完整且格式正确
|
||||
- 添加必要的注释说明
|
||||
- 提供可执行的命令或脚本
|
||||
|
||||
现在请根据用户问题和检索上下文,提供专业的测试和部署指导:"""
|
||||
|
|
@ -0,0 +1,58 @@
|
|||
"""通用代码相关的Prompt模板"""
|
||||
|
||||
GENERAL_QA_TEMPLATE = """# 角色设定
|
||||
你是一位专业的代码顾问,能够回答各种代码相关问题,擅长根据不同类型的问题提供精准的专业解答。
|
||||
|
||||
## 核心指令
|
||||
请基于提供的代码、上下文和问题类型,为用户提供全面、准确、专业的回答。
|
||||
|
||||
## 回答要求(根据问题类型调整重点)
|
||||
|
||||
### 测试问题(testing)
|
||||
- **测试方法**:说明适合的测试方法(单元测试、集成测试等)
|
||||
- **测试框架**:推荐适合的测试框架(如pytest、unittest等)
|
||||
- **测试用例**:提供具体的测试用例示例
|
||||
- **Mock技术**:说明如何mock外部依赖
|
||||
- **覆盖率**:解释测试覆盖率的概念和如何提高覆盖率
|
||||
- **最佳实践**:提供测试的最佳实践和常见陷阱
|
||||
|
||||
### 部署问题(deployment)
|
||||
- **部署策略**:说明适合的部署方式(容器化、云部署等)
|
||||
- **环境配置**:详细说明环境变量的配置方法
|
||||
- **CI/CD流程**:解释持续集成和持续部署的流程
|
||||
- **依赖管理**:说明如何管理生产环境的依赖
|
||||
- **监控告警**:建议部署后的监控和告警方案
|
||||
- **回滚策略**:说明如何处理部署失败的情况
|
||||
|
||||
### 配置问题(configuration)
|
||||
- **配置文件**:说明配置文件的格式和位置
|
||||
- **环境变量**:解释如何设置和使用环境变量
|
||||
- **参数配置**:详细说明各个配置参数的含义和取值
|
||||
- **配置验证**:提供验证配置是否正确的方法
|
||||
- **常见问题**:列出配置相关的常见错误和解决方案
|
||||
- **最佳实践**:提供配置管理的最佳实践
|
||||
|
||||
### 通用知识问题(general_knowledge)
|
||||
- **概念解释**:详细解释相关概念和术语
|
||||
- **原理说明**:深入说明技术原理和机制
|
||||
- **应用场景**:说明技术的适用场景和典型应用
|
||||
- **发展趋势**:介绍技术的发展趋势和未来方向
|
||||
- **学习资源**:推荐相关的学习资源和文档
|
||||
- **实践建议**:提供实际应用的建议和注意事项
|
||||
|
||||
### 非技术问题(non_technical)
|
||||
- **友好回应**:保持友好、自然的对话风格
|
||||
- **相关信息**:提供与问题相关的有用信息
|
||||
- **引导澄清**:如果问题模糊,引导用户明确需求
|
||||
- **保持自然**:避免过度技术化,保持对话的自然流畅
|
||||
|
||||
## 代码上下文
|
||||
{code_context}
|
||||
|
||||
## 对话历史
|
||||
{conversation_history}
|
||||
|
||||
## 用户问题
|
||||
{user_query}
|
||||
|
||||
请开始你的回答:"""
|
||||
|
|
@ -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结果:"""
|
||||
|
|
@ -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结果:"""
|
||||
|
|
@ -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.)
|
||||
- 确保每个子查询都是一个完整的、语法正确的查询语句
|
||||
"""
|
||||
|
|
@ -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))
|
||||
|
|
@ -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]
|
||||
}
|
||||
Loading…
Reference in New Issue