多库检索与存储修复: - RRF 融合去重改用 (collection, chunk_id) 复合键,修复同名文件结果被吞 - DocStore 存储路径加 collection 前缀,修复跨库同名切片数据覆盖 - search_multiple 去重改用复合键 - chunk_id 解析改用 rsplit 兼容下划线文件名 上传与版本管理修复: - 同名文件上传改为覆盖模式,自动清理旧切片 - 修复首次上传不创建版本记录 - 修复覆盖上传版本号回退到 v1 - sync ADDED 分支改用动态版本号生成 - _generate_version_id 改为基于全部版本递增 - 废止/恢复操作同步 SQLite 版本记录 - mark_document_as_superseded 改为仅更新 SQLite 删除清理修复: - 删除文档时同步清理 SQLite 版本记录和变更日志 - 删除向量库时同步清理该库所有版本记录 - cleanup 改为清理 SQLite 记录而非 ChromaDB 测试: - test_version_management.py: 27 条版本管理单元测试 - test_edge_cases.py: 28 条边界用例测试 - test_upload_dedup.py: 5 条上传去重测试 - e2e_risk_test.py: 27 条端到端风险测试 文档: - 新增风险边界问题修复注意事项.md(面向后端的对接文档) - 新增向量库边界风险分析.md - 更新多篇现有文档
385 lines
13 KiB
Python
385 lines
13 KiB
Python
"""
|
||
知识库管理器 - 检索功能 Mixin
|
||
|
||
提供多源融合检索功能,支持:
|
||
- 单向量库检索:向量检索 + BM25 混合检索,RRF 融合排序
|
||
- 多向量库并行检索:跨多个向量库并行检索并合并结果
|
||
- 废止版本检测:查找与查询相关的已废止文档
|
||
|
||
检索流程:
|
||
1. 向量检索:使用 cosine 相似度在 ChromaDB 中检索
|
||
2. BM25 检索:使用关键词匹配在 BM25 索引中检索
|
||
3. RRF 融合:使用 Reciprocal Rank Fusion 合并两路结果
|
||
4. 过滤:排除已废止/已替代的文档
|
||
|
||
主要方法:
|
||
- search_single: 单向量库检索
|
||
- search_multiple: 多向量库并行检索
|
||
- find_deprecated_versions: 查找已废止版本
|
||
"""
|
||
|
||
import logging
|
||
from typing import List, Tuple, Dict, Optional
|
||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||
|
||
from .base import SearchResult
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class SearchMixin:
|
||
"""
|
||
检索功能 Mixin
|
||
|
||
提供混合检索(向量 + BM25)和多库并行检索能力。
|
||
|
||
依赖属性(需由主类提供):
|
||
- self.get_collection: 获取向量库集合的方法
|
||
- self.get_bm25_index: 获取 BM25 索引的方法
|
||
"""
|
||
|
||
def search_single(
|
||
self,
|
||
kb_name: str,
|
||
query_vector: List[float],
|
||
query_text: str,
|
||
top_k: int = 5,
|
||
use_bm25: bool = True,
|
||
include_deprecated: bool = False
|
||
) -> Optional[SearchResult]:
|
||
"""
|
||
单向量库检索
|
||
|
||
对单个向量库执行混合检索:向量检索 + BM25 检索,
|
||
使用 RRF (Reciprocal Rank Fusion) 融合排序。
|
||
|
||
Args:
|
||
kb_name: 向量库名称
|
||
query_vector: 查询向量(由 embedding 模型生成)
|
||
query_text: 查询文本(用于 BM25 检索)
|
||
top_k: 返回结果数量(默认 5)
|
||
use_bm25: 是否启用 BM25 混合检索(默认 True)
|
||
include_deprecated: 是否包含已废止/已替代的文档(默认 False)
|
||
|
||
Returns:
|
||
SearchResult 对象,包含:
|
||
- ids: 文档 ID 列表
|
||
- documents: 文档内容列表
|
||
- metadatas: 元数据列表
|
||
- distances: 距离/分数列表
|
||
- collection_name: 向量库名称
|
||
|
||
向量库为空时返回 None
|
||
"""
|
||
collection = self.get_collection(kb_name)
|
||
if not collection or collection.count() == 0:
|
||
return None
|
||
|
||
where_filter = None
|
||
if not include_deprecated:
|
||
where_filter = {"status": "active"}
|
||
|
||
vector_result = collection.query(
|
||
query_embeddings=[query_vector],
|
||
n_results=top_k,
|
||
where=where_filter
|
||
)
|
||
|
||
if not use_bm25:
|
||
return SearchResult(
|
||
ids=vector_result['ids'][0] if vector_result['ids'] else [],
|
||
documents=vector_result['documents'][0] if vector_result['documents'] else [],
|
||
metadatas=vector_result['metadatas'][0] if vector_result['metadatas'] else [],
|
||
distances=vector_result['distances'][0] if vector_result['distances'] else [],
|
||
collection_name=kb_name
|
||
)
|
||
|
||
bm25_index = self.get_bm25_index(kb_name)
|
||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = bm25_index.search(
|
||
query_text, top_k=min(top_k * 2, 20)
|
||
)
|
||
|
||
if not include_deprecated and bm25_metas:
|
||
filtered_bm25 = []
|
||
for i, meta in enumerate(bm25_metas):
|
||
if meta.get('status', 'active') == 'active':
|
||
filtered_bm25.append((bm25_ids[i], bm25_docs[i], bm25_metas[i], bm25_scores[i]))
|
||
|
||
if filtered_bm25:
|
||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = zip(*filtered_bm25)
|
||
else:
|
||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = [], [], [], []
|
||
|
||
return self._merge_results(
|
||
vector_result,
|
||
(list(bm25_ids), list(bm25_docs), list(bm25_metas), list(bm25_scores)),
|
||
top_k=top_k,
|
||
collection_name=kb_name
|
||
)
|
||
|
||
def search_multiple(
|
||
self,
|
||
kb_names: List[str],
|
||
query_vector: List[float],
|
||
query_text: str,
|
||
top_k: int = 5,
|
||
use_bm25: bool = True
|
||
) -> SearchResult:
|
||
"""
|
||
多向量库并行检索
|
||
|
||
同时在多个向量库中检索,使用线程池并行执行,
|
||
最终合并去重并按分数排序。
|
||
|
||
Args:
|
||
kb_names: 向量库名称列表
|
||
query_vector: 查询向量
|
||
query_text: 查询文本
|
||
top_k: 每个库返回的数量
|
||
use_bm25: 是否启用 BM25
|
||
|
||
Returns:
|
||
合并后的 SearchResult 对象
|
||
"""
|
||
if not kb_names:
|
||
return SearchResult(
|
||
ids=[], documents=[], metadatas=[], distances=[]
|
||
)
|
||
|
||
results = []
|
||
with ThreadPoolExecutor(max_workers=len(kb_names)) as executor:
|
||
futures = {
|
||
executor.submit(
|
||
self.search_single,
|
||
kb_name,
|
||
query_vector,
|
||
query_text,
|
||
top_k,
|
||
use_bm25
|
||
): kb_name for kb_name in kb_names
|
||
}
|
||
|
||
for future in as_completed(futures):
|
||
result = future.result()
|
||
if result:
|
||
results.append(result)
|
||
|
||
return self._merge_multiple_results(results, top_k)
|
||
|
||
def _merge_results(
|
||
self,
|
||
vector_result: dict,
|
||
bm25_result: Tuple,
|
||
top_k: int,
|
||
collection_name: str
|
||
) -> SearchResult:
|
||
"""
|
||
RRF 融合向量检索和 BM25 检索结果
|
||
|
||
使用 Reciprocal Rank Fusion 算法合并两路检索结果,
|
||
综合考虑向量相似度和 BM25 分数进行排序。
|
||
|
||
Args:
|
||
vector_result: ChromaDB 向量检索结果
|
||
bm25_result: BM25 检索结果元组 (ids, docs, metas, scores)
|
||
top_k: 返回数量
|
||
collection_name: 向量库名称
|
||
|
||
Returns:
|
||
融合后的 SearchResult 对象
|
||
|
||
Note:
|
||
RRF 参数 k=60,向量权重 0.5,BM25 权重 0.5。
|
||
"""
|
||
k = 60 # RRF 参数
|
||
|
||
doc_scores = {}
|
||
|
||
# 向量检索结果
|
||
if vector_result['ids'] and vector_result['ids'][0]:
|
||
for rank, (doc_id, doc, meta, dist) in enumerate(zip(
|
||
vector_result['ids'][0],
|
||
vector_result['documents'][0],
|
||
vector_result['metadatas'][0],
|
||
vector_result['distances'][0]
|
||
)):
|
||
rrf_score = 1 / (k + rank + 1)
|
||
sim_score = 1 - dist
|
||
combined = rrf_score * 0.5 + sim_score * 0.5
|
||
|
||
doc_scores[doc_id] = {
|
||
'score': combined,
|
||
'doc': doc,
|
||
'meta': meta
|
||
}
|
||
|
||
# BM25 结果
|
||
bm25_ids, bm25_docs, bm25_metas, bm25_scores = bm25_result
|
||
for rank, (doc_id, doc, meta, score) in enumerate(zip(
|
||
bm25_ids, bm25_docs, bm25_metas, bm25_scores
|
||
)):
|
||
rrf_score = 1 / (k + rank + 1)
|
||
norm_score = score / 10.0 if score > 0 else 0
|
||
combined = rrf_score * 0.5 + norm_score * 0.5
|
||
|
||
if doc_id in doc_scores:
|
||
doc_scores[doc_id]['score'] += combined
|
||
else:
|
||
doc_scores[doc_id] = {
|
||
'score': combined,
|
||
'doc': doc,
|
||
'meta': meta
|
||
}
|
||
|
||
# 排序
|
||
sorted_items = sorted(
|
||
doc_scores.items(),
|
||
key=lambda x: x[1]['score'],
|
||
reverse=True
|
||
)[:top_k]
|
||
|
||
return SearchResult(
|
||
ids=[item[0] for item in sorted_items],
|
||
documents=[item[1]['doc'] for item in sorted_items],
|
||
metadatas=[item[1]['meta'] for item in sorted_items],
|
||
distances=[item[1]['score'] for item in sorted_items],
|
||
collection_name=collection_name
|
||
)
|
||
|
||
def _merge_multiple_results(
|
||
self,
|
||
results: List[SearchResult],
|
||
top_k: int
|
||
) -> SearchResult:
|
||
"""
|
||
合并多个向量库的检索结果
|
||
|
||
将多个向量库的检索结果合并、去重、排序。
|
||
|
||
Args:
|
||
results: 各向量库的检索结果列表
|
||
top_k: 最终返回数量
|
||
|
||
Returns:
|
||
合并后的 SearchResult 对象
|
||
"""
|
||
if not results:
|
||
return SearchResult(
|
||
ids=[], documents=[], metadatas=[], distances=[]
|
||
)
|
||
|
||
if len(results) == 1:
|
||
return results[0]
|
||
|
||
all_items = []
|
||
for result in results:
|
||
for i, doc_id in enumerate(result.ids):
|
||
all_items.append({
|
||
'id': doc_id,
|
||
'doc': result.documents[i],
|
||
'meta': result.metadatas[i],
|
||
'score': result.distances[i],
|
||
'collection': result.collection_name
|
||
})
|
||
|
||
all_items.sort(key=lambda x: x['score'], reverse=True)
|
||
|
||
# 使用 (collection, id) 复合键去重,防止跨库同名文件的结果被吞
|
||
seen = set()
|
||
unique_items = []
|
||
for item in all_items:
|
||
composite_key = (item['collection'], item['id'])
|
||
if composite_key not in seen:
|
||
seen.add(composite_key)
|
||
unique_items.append(item)
|
||
|
||
unique_items = unique_items[:top_k]
|
||
|
||
return SearchResult(
|
||
ids=[item['id'] for item in unique_items],
|
||
documents=[item['doc'] for item in unique_items],
|
||
metadatas=[item['meta'] for item in unique_items],
|
||
distances=[item['score'] for item in unique_items],
|
||
collection_name="multiple"
|
||
)
|
||
|
||
def find_deprecated_versions(
|
||
self,
|
||
kb_names: List[str],
|
||
query_vector: List[float],
|
||
top_k: int = 3
|
||
) -> List[Dict]:
|
||
"""
|
||
查找与查询相关的已废止版本
|
||
|
||
在指定向量库中搜索已废止的文档,
|
||
当相似度 >= 0.7 时返回废止提示信息。
|
||
|
||
Args:
|
||
kb_names: 向量库名称列表
|
||
query_vector: 查询向量
|
||
top_k: 每个库返回的数量
|
||
|
||
Returns:
|
||
废止提示列表,每个元素包含:
|
||
- document: 文档来源
|
||
- collection: 向量库名称
|
||
- status: 状态("deprecated")
|
||
- deprecated_date: 废止日期
|
||
- deprecated_reason: 废止原因
|
||
- similarity: 相似度分数
|
||
- snippet: 内容摘要
|
||
- message: 废止提示消息
|
||
"""
|
||
hints = []
|
||
|
||
for kb_name in kb_names:
|
||
collection = self.get_collection(kb_name)
|
||
if not collection:
|
||
continue
|
||
|
||
result = collection.query(
|
||
query_embeddings=[query_vector],
|
||
n_results=top_k,
|
||
where={"status": "deprecated"}
|
||
)
|
||
|
||
if result['ids'] and result['ids'][0]:
|
||
for doc, meta, score in zip(
|
||
result['documents'][0],
|
||
result['metadatas'][0],
|
||
result['distances'][0]
|
||
):
|
||
sim_score = 1 - score
|
||
if sim_score >= 0.7:
|
||
hints.append({
|
||
"document": meta.get("source", ""),
|
||
"collection": kb_name,
|
||
"status": "deprecated",
|
||
"deprecated_date": meta.get("deprecated_date", ""),
|
||
"deprecated_reason": meta.get("deprecated_reason", ""),
|
||
"similarity": sim_score,
|
||
"snippet": doc[:100] + "..." if len(doc) > 100 else doc,
|
||
"message": self._build_deprecation_hint(meta)
|
||
})
|
||
|
||
return hints
|
||
|
||
def _build_deprecation_hint(self, metadata: Dict) -> str:
|
||
"""
|
||
构建废止提示消息
|
||
|
||
Args:
|
||
metadata: 切片元数据,包含 deprecated_date 和 deprecated_reason
|
||
|
||
Returns:
|
||
格式化的废止提示消息
|
||
"""
|
||
deprecated_date = metadata.get("deprecated_date", "")
|
||
deprecated_reason = metadata.get("deprecated_reason", "")
|
||
|
||
date_str = deprecated_date[:10] if deprecated_date else "未知日期"
|
||
reason_str = f",原因:{deprecated_reason}" if deprecated_reason else ""
|
||
|
||
return f"⚠️ 该文档已于 {date_str} 废止{reason_str},内容不再有效"
|