feat(knowledge_base): 检索的内容以字典返回,同时包含 filepath 等 metadata

将知识库的异步查询方法aquery的返回类型从字符串更改为包含字典的列表,增强了查询结果的结构化。更新了Milvus和Chroma知识库的查询逻辑,确保返回的文档块包含内容、元数据和相似度分数。同时,修复了文件路径处理逻辑,确保文件名相对路径的正确性。
This commit is contained in:
Wenjie Zhang 2025-08-01 17:04:26 +08:00
parent e45e90fec7
commit e2d0b17e5a
4 changed files with 56 additions and 67 deletions

View File

@ -227,59 +227,53 @@ class ChromaKB(KnowledgeBase):
return processed_items_info 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) collection = await self._get_chroma_collection(db_id)
if not collection: if not collection:
raise ValueError(f"Database {db_id} not found") raise ValueError(f"Database {db_id} not found")
try: try:
# 设置查询参数 - ChromaDB 知识库特有的参数
top_k = kwargs.get("top_k", 10) top_k = kwargs.get("top_k", 10)
similarity_threshold = kwargs.get("similarity_threshold", 0.0) # 相似度阈值 similarity_threshold = kwargs.get("similarity_threshold", 0.0)
include_distances = kwargs.get("include_distances", True) # 是否包含距离信息
# 执行相似性搜索
results = collection.query( results = collection.query(
query_texts=[query_text], query_texts=[query_text],
n_results=top_k, n_results=top_k,
include=["documents", "metadatas", "distances"] if include_distances else ["documents", "metadatas"] include=["documents", "metadatas", "distances"]
) )
# 处理结果 if not results or not results.get("documents") or not results["documents"][0]:
if results and results.get("documents") and results["documents"][0]: return []
documents = results["documents"][0]
metadatas = results["metadatas"][0] if results.get("metadatas") else []
distances = results["distances"][0] if results.get("distances") else []
# 构建上下文,应用相似度阈值过滤 documents = results["documents"][0]
contexts = [] metadatas = results["metadatas"][0] if results.get("metadatas") else []
for i, doc in enumerate(documents): distances = results["distances"][0] if results.get("distances") else []
# 计算相似度(距离越小相似度越高)
similarity = 1 - distances[i] if i < len(distances) else 1.0
# 应用相似度阈值过滤 retrieved_chunks = []
if similarity < similarity_threshold: for i, doc in enumerate(documents):
continue similarity = 1 - distances[i] if i < len(distances) else 1.0
context = f"[文档片段 {i+1}]:\n{doc}\n" if similarity < similarity_threshold:
if i < len(metadatas) and metadatas[i]: continue
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)
response = "\n".join(contexts) metadata = metadatas[i] if i < len(metadatas) else {}
logger.debug(f"ChromaDB query response: {len(contexts)} chunks found (after similarity filtering)") # 确保 file_id 在元数据中,并使用统一的键名
return response 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: except Exception as e:
logger.error(f"ChromaDB query error: {e}, {traceback.format_exc()}") logger.error(f"ChromaDB query error: {e}, {traceback.format_exc()}")
return "" return []
async def delete_file(self, db_id: str, file_id: str) -> None: async def delete_file(self, db_id: str, file_id: str) -> None:
"""删除文件""" """删除文件"""

View File

@ -49,7 +49,7 @@ def prepare_item_metadata(item: str, content_type: str, db_id: str) -> dict:
file_path = Path(item) file_path = Path(item)
file_id = f"file_{hashstr(str(file_path) + str(time.time()), 6)}" file_id = f"file_{hashstr(str(file_path) + str(time.time()), 6)}"
file_type = file_path.suffix.lower().replace(".", "") file_type = file_path.suffix.lower().replace(".", "")
filename = file_path.name filename = os.path.relpath(file_path, Path.cwd())
item_path = str(file_path) item_path = str(file_path)
else: # URL else: # URL
file_id = f"url_{hashstr(item + str(time.time()), 6)}" file_id = f"url_{hashstr(item + str(time.time()), 6)}"

View File

@ -162,7 +162,7 @@ class KnowledgeBase(ABC):
pass pass
@abstractmethod @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: 查询参数 **kwargs: 查询参数
Returns: Returns:
查询结果 一个包含字典的列表每个字典代表一个检索到的文档块
""" """
pass 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: 查询参数 **kwargs: 查询参数
Returns: Returns:
查询结果 一个包含字典的列表每个字典代表一个检索到的文档块
""" """
import asyncio import asyncio
logger.warning("query is deprecated, use aquery instead") logger.warning("query is deprecated, use aquery instead")

View File

@ -304,25 +304,21 @@ class MilvusKB(KnowledgeBase):
return processed_items_info 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) collection = await self._get_milvus_collection(db_id)
if not collection: if not collection:
raise ValueError(f"Database {db_id} not found") raise ValueError(f"Database {db_id} not found")
try: try:
# 设置查询参数 - Milvus 知识库特有的参数
top_k = kwargs.get("top_k", 30) top_k = kwargs.get("top_k", 30)
similarity_threshold = kwargs.get("similarity_threshold", 0.2) # 相似度阈值 similarity_threshold = kwargs.get("similarity_threshold", 0.2)
include_distances = kwargs.get("include_distances", True) # 是否包含距离信息 metric_type = kwargs.get("metric_type", "COSINE")
metric_type = kwargs.get("metric_type", "COSINE") # 距离度量类型
# 生成查询向量
embed_info = self.databases_meta[db_id].get("embed_info", {}) embed_info = self.databases_meta[db_id].get("embed_info", {})
embedding_function = self._get_embedding_function(embed_info) embedding_function = self._get_embedding_function(embed_info)
query_embedding = embedding_function([query_text]) query_embedding = embedding_function([query_text])
# 执行相似性搜索
search_params = {"metric_type": metric_type, "params": {"nprobe": 10}} search_params = {"metric_type": metric_type, "params": {"nprobe": 10}}
results = collection.search( results = collection.search(
data=query_embedding, data=query_embedding,
@ -332,37 +328,36 @@ class MilvusKB(KnowledgeBase):
output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"] output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"]
) )
# 处理结果 if not results or len(results) == 0 or len(results[0]) == 0:
if results and len(results) > 0 and len(results[0]) > 0: return []
contexts = []
for i, hit in enumerate(results[0]):
# 计算相似度
similarity = 1 - hit.distance if metric_type == "COSINE" else 1 / (1 + hit.distance)
# 应用相似度阈值过滤 retrieved_chunks = []
if similarity < similarity_threshold: for hit in results[0]:
continue similarity = 1 - hit.distance if metric_type == "COSINE" else 1 / (1 + hit.distance)
entity = hit.entity if similarity < similarity_threshold:
content = entity.get("content", "") continue
source = entity.get("source", "未知来源")
chunk_id = entity.get("chunk_id", f"chunk_{i}")
context = f"[文档片段 {i+1}]:\n{content}\n" entity = hit.entity
context += f"来源: {source} ({chunk_id})\n" metadata = {
if include_distances: "source": entity.get("source", "未知来源"),
context += f"相似度: {similarity:.3f}\n" "chunk_id": entity.get("chunk_id"),
contexts.append(context) "file_id": entity.get("file_id"),
"chunk_index": entity.get("chunk_index")
}
response = "\n".join(contexts) retrieved_chunks.append({
logger.debug(f"Milvus query response: {len(contexts)} chunks found (after similarity filtering)") "content": entity.get("content", ""),
return response "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: except Exception as e:
logger.error(f"Milvus query error: {e}, {traceback.format_exc()}") logger.error(f"Milvus query error: {e}, {traceback.format_exc()}")
return "" return []
async def delete_file(self, db_id: str, file_id: str) -> None: async def delete_file(self, db_id: str, file_id: str) -> None:
"""删除文件""" """删除文件"""