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
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:
"""删除文件"""

View File

@ -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)}"

View File

@ -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")

View File

@ -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:
"""删除文件"""