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
|
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:
|
||||||
"""删除文件"""
|
"""删除文件"""
|
||||||
|
|||||||
@ -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)}"
|
||||||
|
|||||||
@ -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")
|
||||||
|
|||||||
@ -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:
|
||||||
"""删除文件"""
|
"""删除文件"""
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user