From 3fca89d7a8a765bebc359c1992effffaebae9749 Mon Sep 17 00:00:00 2001 From: gu0weix1n <2844905945@qq.com> Date: Fri, 27 Feb 2026 10:44:35 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E4=BB=A3=E7=A0=81QA=E5=8A=9F?= =?UTF-8?q?=E8=83=BD=EF=BC=8C=E6=8F=90=E5=8D=87=E5=AF=B9=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=E7=9B=B8=E5=85=B3=E9=97=AE=E9=A2=98=E7=9A=84=E5=A4=84=E7=90=86?= =?UTF-8?q?=E8=83=BD=E5=8A=9B=EF=BC=8C=E5=90=8C=E6=97=B6=E4=BC=98=E5=8C=96?= =?UTF-8?q?=E5=90=91=E9=87=8F=E5=AD=98=E5=82=A8=E5=92=8C=E6=A3=80=E7=B4=A2?= =?UTF-8?q?=E6=80=A7=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/main.py | 22 +- .../git_git@gitee_com_maxine2_testrepo_git | 1 + ...git_git@gitee_com_thealgorithms_python_git | 1 + ...git@github_com_secondmind-labs_trieste_git | 1 + rag/rag_engine.py | 258 ++++++- rag/vector_store.py | 321 +++++++- sync/base_sync.py | 2 +- utils/code_intent.py | 403 ++++++++++ utils/code_prompt_manager.py | 725 ++++++++++++++++++ utils/git_tool.py | 2 + utils/metadata_filter.py | 160 ++++ utils/prompt/__init__.py | 23 + utils/prompt/answer_generator/__init__.py | 17 + .../answer_generator/algorithm_explanation.py | 45 ++ .../prompt/answer_generator/code_debugging.py | 46 ++ .../answer_generator/code_explanation.py | 41 + .../answer_generator/code_generation.py | 58 ++ .../answer_generator/code_optimization.py | 76 ++ .../concurrency_explanation.py | 99 +++ .../deployment_explanation.py | 99 +++ utils/prompt/answer_generator/general_qa.py | 58 ++ utils/prompt/integrated_query_processing.py | 148 ++++ utils/prompt/intent_detector.py | 82 ++ utils/prompt/query_trans.py | 64 ++ utils/query_processor.py | 418 ++++++++++ utils/query_transformer.py | 150 ++++ 26 files changed, 3279 insertions(+), 41 deletions(-) create mode 160000 git_repos/default/git_git@gitee_com_maxine2_testrepo_git create mode 160000 git_repos/default/git_git@gitee_com_thealgorithms_python_git create mode 160000 git_repos/default/git_git@github_com_secondmind-labs_trieste_git create mode 100644 utils/code_intent.py create mode 100644 utils/code_prompt_manager.py create mode 100644 utils/metadata_filter.py create mode 100644 utils/prompt/__init__.py create mode 100644 utils/prompt/answer_generator/__init__.py create mode 100644 utils/prompt/answer_generator/algorithm_explanation.py create mode 100644 utils/prompt/answer_generator/code_debugging.py create mode 100644 utils/prompt/answer_generator/code_explanation.py create mode 100644 utils/prompt/answer_generator/code_generation.py create mode 100644 utils/prompt/answer_generator/code_optimization.py create mode 100644 utils/prompt/answer_generator/concurrency_explanation.py create mode 100644 utils/prompt/answer_generator/deployment_explanation.py create mode 100644 utils/prompt/answer_generator/general_qa.py create mode 100644 utils/prompt/integrated_query_processing.py create mode 100644 utils/prompt/intent_detector.py create mode 100644 utils/prompt/query_trans.py create mode 100644 utils/query_processor.py create mode 100644 utils/query_transformer.py diff --git a/api/main.py b/api/main.py index 60396d2..ad000e1 100644 --- a/api/main.py +++ b/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}") diff --git a/git_repos/default/git_git@gitee_com_maxine2_testrepo_git b/git_repos/default/git_git@gitee_com_maxine2_testrepo_git new file mode 160000 index 0000000..c22e97f --- /dev/null +++ b/git_repos/default/git_git@gitee_com_maxine2_testrepo_git @@ -0,0 +1 @@ +Subproject commit c22e97f255c92b61ad6014ef74ec662d9f5f6325 diff --git a/git_repos/default/git_git@gitee_com_thealgorithms_python_git b/git_repos/default/git_git@gitee_com_thealgorithms_python_git new file mode 160000 index 0000000..678dedb --- /dev/null +++ b/git_repos/default/git_git@gitee_com_thealgorithms_python_git @@ -0,0 +1 @@ +Subproject commit 678dedbbf94be54b3c9c258368e28bb8e7736d62 diff --git a/git_repos/default/git_git@github_com_secondmind-labs_trieste_git b/git_repos/default/git_git@github_com_secondmind-labs_trieste_git new file mode 160000 index 0000000..1e967fc --- /dev/null +++ b/git_repos/default/git_git@github_com_secondmind-labs_trieste_git @@ -0,0 +1 @@ +Subproject commit 1e967fc87c2761167e0a0a0e84dd2d213c0e1186 diff --git a/rag/rag_engine.py b/rag/rag_engine.py index 151c522..2dd6d24 100644 --- a/rag/rag_engine.py +++ b/rag/rag_engine.py @@ -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 diff --git a/rag/vector_store.py b/rag/vector_store.py index ea11182..a1cc8b7 100644 --- a/rag/vector_store.py +++ b/rag/vector_store.py @@ -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) diff --git a/sync/base_sync.py b/sync/base_sync.py index 8b5f386..5ab1815 100644 --- a/sync/base_sync.py +++ b/sync/base_sync.py @@ -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)} 字符") diff --git a/utils/code_intent.py b/utils/code_intent.py new file mode 100644 index 0000000..ee6f63b --- /dev/null +++ b/utils/code_intent.py @@ -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()) diff --git a/utils/code_prompt_manager.py b/utils/code_prompt_manager.py new file mode 100644 index 0000000..bdae2ba --- /dev/null +++ b/utils/code_prompt_manager.py @@ -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()) diff --git a/utils/git_tool.py b/utils/git_tool.py index 87bf866..7ac7bc1 100644 --- a/utils/git_tool.py +++ b/utils/git_tool.py @@ -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" diff --git a/utils/metadata_filter.py b/utils/metadata_filter.py new file mode 100644 index 0000000..9a04323 --- /dev/null +++ b/utils/metadata_filter.py @@ -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) diff --git a/utils/prompt/__init__.py b/utils/prompt/__init__.py new file mode 100644 index 0000000..9f5753d --- /dev/null +++ b/utils/prompt/__init__.py @@ -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" +] diff --git a/utils/prompt/answer_generator/__init__.py b/utils/prompt/answer_generator/__init__.py new file mode 100644 index 0000000..b89ca1f --- /dev/null +++ b/utils/prompt/answer_generator/__init__.py @@ -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" +] diff --git a/utils/prompt/answer_generator/algorithm_explanation.py b/utils/prompt/answer_generator/algorithm_explanation.py new file mode 100644 index 0000000..685a1b3 --- /dev/null +++ b/utils/prompt/answer_generator/algorithm_explanation.py @@ -0,0 +1,45 @@ +"""算法解释相关的Prompt模板""" + +ALGORITHM_EXPLANATION_TEMPLATE = """# 角色设定 +你是一位算法专家,擅长深入解析算法原理和实现,能够根据不同类型的问题提供精准的专业解答。 + +## 核心指令 +请基于提供的算法代码、上下文和问题类型,为用户提供详细、专业的算法分析。 + +## 分析要求(根据问题类型调整重点) + +### 算法解释(algorithm_explanation) +- **算法原理**:详细解释算法的基本原理和设计思想 +- **算法步骤**:逐步说明算法的执行流程 +- **时间复杂度**:分析时间复杂度(最好、平均、最坏情况) +- **空间复杂度**:分析空间复杂度,说明内存使用情况 +- **优缺点**:分析算法的优势和局限性 +- **适用场景**:说明算法的适用场景和典型应用 +- **比较分析**:与其他同类算法进行对比 + +### 算法实现(algorithm_implementation) +- **算法选择**:选择最适合该问题的算法 +- **实现细节**:提供完整的算法实现代码 +- **复杂度分析**:分析实现的时间和空间复杂度 +- **边界处理**:考虑边界情况和特殊输入 +- **优化建议**:提供可能的优化方向 +- **测试用例**:提供测试算法的示例用例 + +### 数据结构(data_structure) +- **结构定义**:详细说明数据结构的定义和特点 +- **操作方法**:说明数据结构支持的操作及其复杂度 +- **实现方式**:提供数据结构的实现代码 +- **适用场景**:说明数据结构的适用场景和典型应用 +- **性能对比**:与其他数据结构进行性能对比 +- **使用示例**:提供数据结构的使用示例 + +## 算法代码 +{algorithm_code} + +## 对话历史 +{conversation_history} + +## 用户问题 +{user_query} + +请开始你的分析:""" diff --git a/utils/prompt/answer_generator/code_debugging.py b/utils/prompt/answer_generator/code_debugging.py new file mode 100644 index 0000000..adaa49f --- /dev/null +++ b/utils/prompt/answer_generator/code_debugging.py @@ -0,0 +1,46 @@ +"""代码调试相关的Prompt模板""" + +CODE_DEBUGGING_TEMPLATE = """# 角色设定 +你是一位专业的代码调试专家,擅长快速定位和解决代码中的错误,能够根据不同类型的错误提供精准的诊断和解决方案。 + +## 核心指令 +请基于提供的代码、错误信息、上下文和错误类型,为用户提供详细的错误分析和解决方案。 + +## 分析要求(根据错误类型调整重点) + +### 语法错误(syntax_error) +- **错误定位**:准确指出语法错误的具体位置(行号、列号) +- **错误原因**:解释违反了哪条语法规则 +- **修正方案**:提供修正后的完整代码 +- **预防建议**:说明如何避免类似的语法错误 +- **常见模式**:列出该语法错误的常见触发场景 + +### 运行时错误(runtime_error) +- **错误分析**:详细分析异常类型和错误信息 +- **堆栈跟踪**:解释错误堆栈中的关键信息 +- **根本原因**:深入分析导致错误的根本原因 +- **修复方案**:提供具体的修复代码和实施步骤 +- **异常处理**:建议如何添加异常处理来预防此类错误 +- **测试建议**:说明如何测试修复是否有效 + +### 调试问题(debugging) +- **调试方法**:提供适合该问题的调试策略 +- **断点设置**:建议在哪些位置设置断点 +- **日志分析**:说明如何通过日志分析问题 +- **变量检查**:建议检查哪些关键变量的值 +- **逐步排查**:提供逐步排查问题的流程 +- **工具推荐**:推荐适合的调试工具和技巧 + +## 代码上下文 +{code_context} + +## 错误信息 +{error_message} + +## 对话历史 +{conversation_history} + +## 用户问题 +{user_query} + +请开始你的分析:""" diff --git a/utils/prompt/answer_generator/code_explanation.py b/utils/prompt/answer_generator/code_explanation.py new file mode 100644 index 0000000..8c8be92 --- /dev/null +++ b/utils/prompt/answer_generator/code_explanation.py @@ -0,0 +1,41 @@ +"""代码解释相关的Prompt模板""" + +CODE_EXPLANATION_TEMPLATE = """# 角色设定 +你是一位资深的代码分析专家,擅长深入解析代码结构和功能,能够根据不同类型的问题提供精准的专业解答。 + +## 核心指令 +请基于提供的代码、上下文和用户问题类型,为用户提供详细、专业的代码分析。 + +## 分析要求(根据问题类型调整重点) + +### 逻辑解释问题(logic_explanation) +- **功能说明**:详细解释代码的功能和用途 +- **逻辑分析**:说明代码的执行流程和核心逻辑 +- **实现原理**:解释代码的实现原理和技术细节 +- **使用示例**:提供实际可运行的代码示例 +- **注意事项**:指出使用时需要注意的要点和常见错误 + +### 实体介绍问题(entity_introduction) +- **实体结构**:说明函数、类、API等实体的结构和组成 +- **参数分析**:说明每个参数的类型、含义和默认值 +- **返回值说明**:解释返回值的类型、含义和可能的取值 +- **使用示例**:提供实际可运行的代码示例 +- **注意事项**:指出使用时需要注意的要点和常见错误 + +### 代码结构问题(code_structure) +- **项目结构**:详细说明项目的目录和文件组织 +- **模块划分**:说明各个模块的功能和职责 +- **依赖关系**:解释模块之间的依赖和调用关系 +- **架构设计**:说明项目的整体架构和设计思路 +- **文件说明**:解释关键文件的作用和内容 + +## 代码上下文 +{code_context} + +## 对话历史 +{conversation_history} + +## 用户问题 +{user_query} + +请开始你的分析:""" diff --git a/utils/prompt/answer_generator/code_generation.py b/utils/prompt/answer_generator/code_generation.py new file mode 100644 index 0000000..ed60259 --- /dev/null +++ b/utils/prompt/answer_generator/code_generation.py @@ -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} + +请开始生成代码:''' diff --git a/utils/prompt/answer_generator/code_optimization.py b/utils/prompt/answer_generator/code_optimization.py new file mode 100644 index 0000000..1e03852 --- /dev/null +++ b/utils/prompt/answer_generator/code_optimization.py @@ -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 等标注语言 +- 代码必须完整可运行 +- 添加必要的注释说明 +- 保持代码风格一致 + +现在请根据用户问题和检索上下文,提供专业的代码优化建议:""" diff --git a/utils/prompt/answer_generator/concurrency_explanation.py b/utils/prompt/answer_generator/concurrency_explanation.py new file mode 100644 index 0000000..1132390 --- /dev/null +++ b/utils/prompt/answer_generator/concurrency_explanation.py @@ -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 等标注语言 +- 代码必须完整可运行 +- 添加必要的注释说明 +- 展示并发/异步的完整流程 + +现在请根据用户问题和检索上下文,提供专业的并发编程指导:''' diff --git a/utils/prompt/answer_generator/deployment_explanation.py b/utils/prompt/answer_generator/deployment_explanation.py new file mode 100644 index 0000000..9513bca --- /dev/null +++ b/utils/prompt/answer_generator/deployment_explanation.py @@ -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 等标注语言 +- 配置文件必须完整且格式正确 +- 添加必要的注释说明 +- 提供可执行的命令或脚本 + +现在请根据用户问题和检索上下文,提供专业的测试和部署指导:""" diff --git a/utils/prompt/answer_generator/general_qa.py b/utils/prompt/answer_generator/general_qa.py new file mode 100644 index 0000000..8c7537e --- /dev/null +++ b/utils/prompt/answer_generator/general_qa.py @@ -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} + +请开始你的回答:""" diff --git a/utils/prompt/integrated_query_processing.py b/utils/prompt/integrated_query_processing.py new file mode 100644 index 0000000..b976e6f --- /dev/null +++ b/utils/prompt/integrated_query_processing.py @@ -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结果:""" diff --git a/utils/prompt/intent_detector.py b/utils/prompt/intent_detector.py new file mode 100644 index 0000000..0c74f9e --- /dev/null +++ b/utils/prompt/intent_detector.py @@ -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结果:""" diff --git a/utils/prompt/query_trans.py b/utils/prompt/query_trans.py new file mode 100644 index 0000000..f6c050a --- /dev/null +++ b/utils/prompt/query_trans.py @@ -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.) +- 确保每个子查询都是一个完整的、语法正确的查询语句 +""" diff --git a/utils/query_processor.py b/utils/query_processor.py new file mode 100644 index 0000000..1a61abe --- /dev/null +++ b/utils/query_processor.py @@ -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)) diff --git a/utils/query_transformer.py b/utils/query_transformer.py new file mode 100644 index 0000000..398153c --- /dev/null +++ b/utils/query_transformer.py @@ -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] + }