From e2d0b17e5a151b964c54bad1c5566dea3acac55d Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 1 Aug 2025 17:04:26 +0800 Subject: [PATCH] =?UTF-8?q?feat(knowledge=5Fbase):=20=E6=A3=80=E7=B4=A2?= =?UTF-8?q?=E7=9A=84=E5=86=85=E5=AE=B9=E4=BB=A5=E5=AD=97=E5=85=B8=E8=BF=94?= =?UTF-8?q?=E5=9B=9E=EF=BC=8C=E5=90=8C=E6=97=B6=E5=8C=85=E5=90=AB=20filepa?= =?UTF-8?q?th=20=E7=AD=89=20metadata?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 将知识库的异步查询方法aquery的返回类型从字符串更改为包含字典的列表,增强了查询结果的结构化。更新了Milvus和Chroma知识库的查询逻辑,确保返回的文档块包含内容、元数据和相似度分数。同时,修复了文件路径处理逻辑,确保文件名相对路径的正确性。 --- src/knowledge/chroma_kb.py | 58 +++++++++++++++------------------ src/knowledge/kb_utils.py | 2 +- src/knowledge/knowledge_base.py | 8 ++--- src/knowledge/milvus_kb.py | 55 ++++++++++++++----------------- 4 files changed, 56 insertions(+), 67 deletions(-) diff --git a/src/knowledge/chroma_kb.py b/src/knowledge/chroma_kb.py index 96d3175a..bbecf6c0 100644 --- a/src/knowledge/chroma_kb.py +++ b/src/knowledge/chroma_kb.py @@ -227,59 +227,53 @@ class ChromaKB(KnowledgeBase): return processed_items_info - async def aquery(self, query_text: str, db_id: str, **kwargs) -> str: + async def aquery(self, query_text: str, db_id: str, **kwargs) -> list[dict]: """异步查询知识库""" collection = await self._get_chroma_collection(db_id) if not collection: raise ValueError(f"Database {db_id} not found") try: - # 设置查询参数 - ChromaDB 知识库特有的参数 top_k = kwargs.get("top_k", 10) - similarity_threshold = kwargs.get("similarity_threshold", 0.0) # 相似度阈值 - include_distances = kwargs.get("include_distances", True) # 是否包含距离信息 + similarity_threshold = kwargs.get("similarity_threshold", 0.0) - # 执行相似性搜索 results = collection.query( query_texts=[query_text], n_results=top_k, - include=["documents", "metadatas", "distances"] if include_distances else ["documents", "metadatas"] + include=["documents", "metadatas", "distances"] ) - # 处理结果 - if results and results.get("documents") and results["documents"][0]: - documents = results["documents"][0] - metadatas = results["metadatas"][0] if results.get("metadatas") else [] - distances = results["distances"][0] if results.get("distances") else [] + if not results or not results.get("documents") or not results["documents"][0]: + return [] - # 构建上下文,应用相似度阈值过滤 - contexts = [] - for i, doc in enumerate(documents): - # 计算相似度(距离越小相似度越高) - similarity = 1 - distances[i] if i < len(distances) else 1.0 + documents = results["documents"][0] + metadatas = results["metadatas"][0] if results.get("metadatas") else [] + distances = results["distances"][0] if results.get("distances") else [] - # 应用相似度阈值过滤 - if similarity < similarity_threshold: - continue + retrieved_chunks = [] + for i, doc in enumerate(documents): + similarity = 1 - distances[i] if i < len(distances) else 1.0 - context = f"[文档片段 {i+1}]:\n{doc}\n" - if i < len(metadatas) and metadatas[i]: - source = metadatas[i].get("source", "未知来源") - chunk_id = metadatas[i].get("chunk_id", f"chunk_{i}") - context += f"来源: {source} ({chunk_id})\n" - if include_distances and i < len(distances): - context += f"相似度: {similarity:.3f}\n" - contexts.append(context) + if similarity < similarity_threshold: + continue - response = "\n".join(contexts) - logger.debug(f"ChromaDB query response: {len(contexts)} chunks found (after similarity filtering)") - return response + metadata = metadatas[i] if i < len(metadatas) else {} + # 确保 file_id 在元数据中,并使用统一的键名 + if 'full_doc_id' in metadata: + metadata['file_id'] = metadata.pop('full_doc_id') - return "" + retrieved_chunks.append({ + "content": doc, + "metadata": metadata, + "score": similarity + }) + + logger.debug(f"ChromaDB query response: {len(retrieved_chunks)} chunks found (after similarity filtering)") + return retrieved_chunks except Exception as e: logger.error(f"ChromaDB query error: {e}, {traceback.format_exc()}") - return "" + return [] async def delete_file(self, db_id: str, file_id: str) -> None: """删除文件""" diff --git a/src/knowledge/kb_utils.py b/src/knowledge/kb_utils.py index 4726d4ee..cebfd7f8 100644 --- a/src/knowledge/kb_utils.py +++ b/src/knowledge/kb_utils.py @@ -49,7 +49,7 @@ def prepare_item_metadata(item: str, content_type: str, db_id: str) -> dict: file_path = Path(item) file_id = f"file_{hashstr(str(file_path) + str(time.time()), 6)}" file_type = file_path.suffix.lower().replace(".", "") - filename = file_path.name + filename = os.path.relpath(file_path, Path.cwd()) item_path = str(file_path) else: # URL file_id = f"url_{hashstr(item + str(time.time()), 6)}" diff --git a/src/knowledge/knowledge_base.py b/src/knowledge/knowledge_base.py index af685f9d..d0580b3c 100644 --- a/src/knowledge/knowledge_base.py +++ b/src/knowledge/knowledge_base.py @@ -162,7 +162,7 @@ class KnowledgeBase(ABC): pass @abstractmethod - async def aquery(self, query_text: str, db_id: str, **kwargs) -> str: + async def aquery(self, query_text: str, db_id: str, **kwargs) -> list[dict]: """ 异步查询知识库 @@ -172,11 +172,11 @@ class KnowledgeBase(ABC): **kwargs: 查询参数 Returns: - 查询结果 + 一个包含字典的列表,每个字典代表一个检索到的文档块。 """ pass - def query(self, query_text: str, db_id: str, **kwargs) -> str: + def query(self, query_text: str, db_id: str, **kwargs) -> list[dict]: """ 同步查询知识库(兼容性方法) @@ -186,7 +186,7 @@ class KnowledgeBase(ABC): **kwargs: 查询参数 Returns: - 查询结果 + 一个包含字典的列表,每个字典代表一个检索到的文档块。 """ import asyncio logger.warning("query is deprecated, use aquery instead") diff --git a/src/knowledge/milvus_kb.py b/src/knowledge/milvus_kb.py index 049bc0bc..d1c5c1ea 100644 --- a/src/knowledge/milvus_kb.py +++ b/src/knowledge/milvus_kb.py @@ -304,25 +304,21 @@ class MilvusKB(KnowledgeBase): return processed_items_info - async def aquery(self, query_text: str, db_id: str, **kwargs) -> str: + async def aquery(self, query_text: str, db_id: str, **kwargs) -> list[dict]: """异步查询知识库""" collection = await self._get_milvus_collection(db_id) if not collection: raise ValueError(f"Database {db_id} not found") try: - # 设置查询参数 - Milvus 知识库特有的参数 top_k = kwargs.get("top_k", 30) - similarity_threshold = kwargs.get("similarity_threshold", 0.2) # 相似度阈值 - include_distances = kwargs.get("include_distances", True) # 是否包含距离信息 - metric_type = kwargs.get("metric_type", "COSINE") # 距离度量类型 + similarity_threshold = kwargs.get("similarity_threshold", 0.2) + metric_type = kwargs.get("metric_type", "COSINE") - # 生成查询向量 embed_info = self.databases_meta[db_id].get("embed_info", {}) embedding_function = self._get_embedding_function(embed_info) query_embedding = embedding_function([query_text]) - # 执行相似性搜索 search_params = {"metric_type": metric_type, "params": {"nprobe": 10}} results = collection.search( data=query_embedding, @@ -332,37 +328,36 @@ class MilvusKB(KnowledgeBase): output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"] ) - # 处理结果 - if results and len(results) > 0 and len(results[0]) > 0: - contexts = [] - for i, hit in enumerate(results[0]): - # 计算相似度 - similarity = 1 - hit.distance if metric_type == "COSINE" else 1 / (1 + hit.distance) + if not results or len(results) == 0 or len(results[0]) == 0: + return [] - # 应用相似度阈值过滤 - if similarity < similarity_threshold: - continue + retrieved_chunks = [] + for hit in results[0]: + similarity = 1 - hit.distance if metric_type == "COSINE" else 1 / (1 + hit.distance) - entity = hit.entity - content = entity.get("content", "") - source = entity.get("source", "未知来源") - chunk_id = entity.get("chunk_id", f"chunk_{i}") + if similarity < similarity_threshold: + continue - context = f"[文档片段 {i+1}]:\n{content}\n" - context += f"来源: {source} ({chunk_id})\n" - if include_distances: - context += f"相似度: {similarity:.3f}\n" - contexts.append(context) + entity = hit.entity + metadata = { + "source": entity.get("source", "未知来源"), + "chunk_id": entity.get("chunk_id"), + "file_id": entity.get("file_id"), + "chunk_index": entity.get("chunk_index") + } - response = "\n".join(contexts) - logger.debug(f"Milvus query response: {len(contexts)} chunks found (after similarity filtering)") - return response + retrieved_chunks.append({ + "content": entity.get("content", ""), + "metadata": metadata, + "score": similarity + }) - return "" + logger.debug(f"Milvus query response: {len(retrieved_chunks)} chunks found (after similarity filtering)") + return retrieved_chunks except Exception as e: logger.error(f"Milvus query error: {e}, {traceback.format_exc()}") - return "" + return [] async def delete_file(self, db_id: str, file_id: str) -> None: """删除文件"""