From 2343ab7751e9898d599c404fc8c56ff415c73e6c Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Sat, 19 Jul 2025 12:16:17 +0300 Subject: [PATCH] Return tuple (Chunk, score,) in rerank() --- src/haiku/rag/reranking/base.py | 2 +- src/haiku/rag/reranking/cohere.py | 4 ++-- src/haiku/rag/reranking/mxbai.py | 4 ++-- tests/test_reranker.py | 8 +++++--- 4 files changed, 10 insertions(+), 8 deletions(-) diff --git a/src/haiku/rag/reranking/base.py b/src/haiku/rag/reranking/base.py index 72b3a3df..0e95e26b 100644 --- a/src/haiku/rag/reranking/base.py +++ b/src/haiku/rag/reranking/base.py @@ -7,7 +7,7 @@ class RerankerBase: async def rerank( self, query: str, chunks: list[Chunk], top_n: int = 10 - ) -> list[Chunk]: + ) -> list[tuple[Chunk, float]]: raise NotImplementedError( "Reranker is an abstract class. Please implement the rerank method in a subclass." ) diff --git a/src/haiku/rag/reranking/cohere.py b/src/haiku/rag/reranking/cohere.py index 19e91577..6d30952d 100644 --- a/src/haiku/rag/reranking/cohere.py +++ b/src/haiku/rag/reranking/cohere.py @@ -16,7 +16,7 @@ class CohereReranker(RerankerBase): async def rerank( self, query: str, chunks: list[Chunk], top_n: int = 10 - ) -> list[Chunk]: + ) -> list[tuple[Chunk, float]]: if not chunks: return [] @@ -29,6 +29,6 @@ class CohereReranker(RerankerBase): reranked_chunks = [] for result in response.results: original_chunk = chunks[result.index] - reranked_chunks.append(original_chunk) + reranked_chunks.append((original_chunk, result.relevance_score)) return reranked_chunks diff --git a/src/haiku/rag/reranking/mxbai.py b/src/haiku/rag/reranking/mxbai.py index 135727ed..032edac5 100644 --- a/src/haiku/rag/reranking/mxbai.py +++ b/src/haiku/rag/reranking/mxbai.py @@ -13,7 +13,7 @@ class MxBAIReranker(RerankerBase): async def rerank( self, query: str, chunks: list[Chunk], top_n: int = 10 - ) -> list[Chunk]: + ) -> list[tuple[Chunk, float]]: if not chunks: return [] @@ -23,6 +23,6 @@ class MxBAIReranker(RerankerBase): reranked_chunks = [] for result in results: original_chunk = chunks[result.index] - reranked_chunks.append(original_chunk) + reranked_chunks.append((original_chunk, result.score)) return reranked_chunks diff --git a/tests/test_reranker.py b/tests/test_reranker.py index c83c2d3d..b3de96cf 100644 --- a/tests/test_reranker.py +++ b/tests/test_reranker.py @@ -34,7 +34,8 @@ async def test_mxbai_reranker(): reranked = await reranker.rerank( "Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2 ) - assert [r.document_id for r in reranked] == [0, 2] + assert [chunk.document_id for chunk, score in reranked] == [0, 2] + assert all(isinstance(score, float) for chunk, score in reranked) @pytest.mark.asyncio @@ -43,12 +44,13 @@ async def test_cohere_reranker(): from haiku.rag.reranking.cohere import CohereReranker reranker = CohereReranker() - assert reranker._model == "rerank-v3.5" + reranker._model = "rerank-v3.5" reranked = await reranker.rerank( "Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2 ) - assert [r.document_id for r in reranked] == [0, 2] + assert [chunk.document_id for chunk, score in reranked] == [0, 2] + assert all(isinstance(score, float) for chunk, score in reranked) except ImportError: pytest.skip("Cohere package not installed")