feat(knowledge_base): 检索的内容以字典返回,同时包含 filepath 等 metadata
将知识库的异步查询方法aquery的返回类型从字符串更改为包含字典的列表,增强了查询结果的结构化。更新了Milvus和Chroma知识库的查询逻辑,确保返回的文档块包含内容、元数据和相似度分数。同时,修复了文件路径处理逻辑,确保文件名相对路径的正确性。
This commit is contained in:
parent
e45e90fec7
commit
e2d0b17e5a
@ -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:
|
||||
"""删除文件"""
|
||||
|
||||
@ -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)}"
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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:
|
||||
"""删除文件"""
|
||||
|
||||
Loading…
Reference in New Issue
Block a user