diff --git a/astrbot/core/knowledge_base/retrieval/rank_fusion.py b/astrbot/core/knowledge_base/retrieval/rank_fusion.py index 40afd97484..39f402d99c 100644 --- a/astrbot/core/knowledge_base/retrieval/rank_fusion.py +++ b/astrbot/core/knowledge_base/retrieval/rank_fusion.py @@ -66,7 +66,10 @@ async def fuse( dense_ranks = { r.data["doc_id"]: (idx + 1) for idx, r in enumerate(dense_results) } # 这里的 doc_id 实际上是 chunk_id - sparse_ranks = {r.chunk_id: (idx + 1) for idx, r in enumerate(sparse_results)} + sparse_ranks = { + r.chunk_id: r.rank if r.rank is not None else idx + 1 + for idx, r in enumerate(sparse_results) + } # 2. 收集所有唯一的 ID # 需要统一为 chunk_id diff --git a/astrbot/core/knowledge_base/retrieval/sparse_retriever.py b/astrbot/core/knowledge_base/retrieval/sparse_retriever.py index f06eb50909..316367991d 100644 --- a/astrbot/core/knowledge_base/retrieval/sparse_retriever.py +++ b/astrbot/core/knowledge_base/retrieval/sparse_retriever.py @@ -30,6 +30,7 @@ class SparseResult: kb_id: str content: str score: float + rank: int | None = None class SparseRetriever: @@ -87,7 +88,9 @@ async def retrieve( fallback_kb_ids.append(kb_id) continue - for doc in result: + # BM25 scores from independent FTS5 indexes are not comparable. + # Preserve each index's local rank for the later RRF stage. + for rank, doc in enumerate(result, start=1): chunk_md = json.loads(doc["metadata"]) fts_results.append( SparseResult( @@ -97,6 +100,7 @@ async def retrieve( kb_id=kb_id, content=doc["text"], score=-float(doc["score"]), + rank=rank, ), ) @@ -172,5 +176,7 @@ async def _retrieve_with_bm25( ) results.sort(key=lambda x: x.score, reverse=True) + for rank, result in enumerate(results, start=1): + result.rank = rank # return results[: len(results) // len(kb_ids)] return results[:top_k_sparse] diff --git a/tests/unit/test_rank_fusion.py b/tests/unit/test_rank_fusion.py new file mode 100644 index 0000000000..534944bad8 --- /dev/null +++ b/tests/unit/test_rank_fusion.py @@ -0,0 +1,59 @@ +import pytest + +from astrbot.core.db.vec_db.base import Result +from astrbot.core.knowledge_base.retrieval.rank_fusion import RankFusion +from astrbot.core.knowledge_base.retrieval.sparse_retriever import SparseResult + + +def make_dense_result(chunk_id: str, similarity: float) -> Result: + return Result( + similarity=similarity, + data={ + "doc_id": chunk_id, + "text": chunk_id, + "metadata": "{}", + }, + ) + + +def make_sparse_result( + chunk_id: str, + kb_id: str, + score: float, + rank: int, +) -> SparseResult: + return SparseResult( + chunk_index=0, + chunk_id=chunk_id, + doc_id=f"doc-{chunk_id}", + kb_id=kb_id, + content=chunk_id, + score=score, + rank=rank, + ) + + +@pytest.mark.asyncio +async def test_rank_fusion_uses_source_rank_for_independent_sparse_indexes(): + dense_results = [ + make_dense_result("small-exact", 0.99), + make_dense_result("large-1", 0.95), + make_dense_result("large-2", 0.90), + ] + sparse_results = [ + make_sparse_result("large-1", "kb-large", 12.0, 1), + make_sparse_result("large-2", "kb-large", 10.0, 2), + make_sparse_result("small-exact", "kb-small", 0.00001, 1), + ] + + results = await RankFusion(kb_db=None).fuse( + dense_results=dense_results, + sparse_results=sparse_results, + ) + + assert [result.chunk_id for result in results] == [ + "small-exact", + "large-1", + "large-2", + ] + assert results[0].score == pytest.approx(2 / 61) diff --git a/tests/unit/test_sparse_retriever.py b/tests/unit/test_sparse_retriever.py index 11c491b4d2..067fbb46b0 100644 --- a/tests/unit/test_sparse_retriever.py +++ b/tests/unit/test_sparse_retriever.py @@ -6,7 +6,12 @@ from astrbot.core.knowledge_base.retrieval.sparse_retriever import SparseRetriever -def make_doc(chunk_id: str, text: str, chunk_index: int = 0) -> dict: +def make_doc( + chunk_id: str, + text: str, + chunk_index: int = 0, + kb_id: str = "kb-1", +) -> dict: return { "doc_id": chunk_id, "text": text, @@ -14,7 +19,7 @@ def make_doc(chunk_id: str, text: str, chunk_index: int = 0) -> dict: { "chunk_index": chunk_index, "kb_doc_id": f"doc-{chunk_index}", - "kb_id": "kb-1", + "kb_id": kb_id, }, ), } @@ -59,6 +64,14 @@ async def get_documents(self, metadata_filters: dict, limit: int | None, offset) ] +class StaticFTSStorage: + def __init__(self, documents: list[dict]): + self.documents = documents + + async def search_sparse(self, query_tokens: list[str], limit: int): + return self.documents[:limit] + + @pytest.mark.asyncio async def test_sparse_retriever_uses_fts5_when_available(): storage = FTSStorage() @@ -91,3 +104,51 @@ async def test_sparse_retriever_falls_back_to_bm25_when_fts5_is_unavailable(): assert [result.chunk_id for result in results] == ["chunk-1"] assert storage.search_sparse_calls == 1 assert storage.get_documents_calls == 1 + + +@pytest.mark.asyncio +async def test_sparse_retriever_preserves_per_kb_fts_ranks(): + large_storage = StaticFTSStorage( + [ + { + **make_doc("large-1", "管理员账号安全说明", 0, "kb-large"), + "score": -12.0, + }, + { + **make_doc("large-2", "密码策略说明", 1, "kb-large"), + "score": -10.0, + }, + ], + ) + small_storage = StaticFTSStorage( + [ + { + **make_doc( + "small-exact", + "如何重置管理员密码?", + 0, + "kb-small", + ), + "score": -0.00001, + }, + ], + ) + retriever = SparseRetriever(kb_db=None) + + results = await retriever.retrieve( + query="如何重置管理员密码?", + kb_ids=["kb-large", "kb-small"], + kb_options={ + "kb-large": { + "vec_db": SimpleNamespace(document_storage=large_storage), + "top_k_sparse": 2, + }, + "kb-small": { + "vec_db": SimpleNamespace(document_storage=small_storage), + "top_k_sparse": 1, + }, + }, + ) + + ranks = {result.chunk_id: result.rank for result in results} + assert ranks == {"large-1": 1, "large-2": 2, "small-exact": 1}