diff --git a/docs/configuration.md b/docs/configuration.md index d44b31b4..5e78ee8a 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -105,16 +105,20 @@ ANTHROPIC_API_KEY="your-api-key" ## Reranking -Reranking is **enabled by default** and improves search quality by re-ordering the initial search results using specialized models. When enabled, the system retrieves more candidates (3x the requested limit) and then reranks them to return the most relevant results. +Reranking improves search quality by re-ordering the initial search results using specialized models. When enabled, the system retrieves more candidates (3x the requested limit) and then reranks them to return the most relevant results. -If you use the default reranked (running locally), it can slow down searching significantly. To disable reranking for faster searches: - -```bash -RERANK=false -``` +Reranking is **automatically enabled** if you install the appropriate reranking provider package. ### MixedBread AI (Default) +For MxBAI reranking, install with mxbai extras: + +```bash +uv pip install haiku.rag[mxbai] +``` + +Then configure: + ```bash RERANK_PROVIDER="mxbai" RERANK_MODEL="mixedbread-ai/mxbai-rerank-base-v2" diff --git a/pyproject.toml b/pyproject.toml index 025cdbbb..c5fd668b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,7 +25,6 @@ dependencies = [ "docling>=2.15.0", "fastmcp>=2.8.1", "httpx>=0.28.1", - "mxbai-rerank>=0.1.6", "ollama>=0.5.1", "pydantic>=2.11.7", "python-dotenv>=1.1.0", @@ -41,6 +40,7 @@ voyageai = ["voyageai>=0.3.2"] openai = ["openai>=1.0.0"] anthropic = ["anthropic>=0.56.0"] cohere = ["cohere>=5.16.1"] +mxbai = ["mxbai-rerank>=0.1.6"] [project.scripts] haiku-rag = "haiku.rag.cli:cli" diff --git a/src/haiku/rag/client.py b/src/haiku/rag/client.py index 1e7ebaae..51694138 100644 --- a/src/haiku/rag/client.py +++ b/src/haiku/rag/client.py @@ -319,7 +319,7 @@ class HaikuRAG: return await self.document_repository.list_all(limit=limit, offset=offset) async def search( - self, query: str, limit: int = 5, k: int = 60, rerank=Config.RERANK + self, query: str, limit: int = 5, k: int = 60 ) -> list[tuple[Chunk, float]]: """Search for relevant chunks using hybrid search (vector similarity + full-text search) with reranking. @@ -331,8 +331,10 @@ class HaikuRAG: Returns: List of (chunk, score) tuples ordered by relevance. """ + # Get reranker if available + reranker = get_reranker() - if not rerank: + if reranker is None: return await self.chunk_repository.search_chunks_hybrid(query, limit, k) # Get more initial results (3X) for reranking @@ -340,7 +342,6 @@ class HaikuRAG: query, limit * 3, k ) # Apply reranking - reranker = get_reranker() chunks = [chunk for chunk, _ in search_results] reranked_results = await reranker.rerank(query, chunks, top_n=limit) diff --git a/src/haiku/rag/config.py b/src/haiku/rag/config.py index 3671ed19..251d4a46 100644 --- a/src/haiku/rag/config.py +++ b/src/haiku/rag/config.py @@ -19,7 +19,6 @@ class AppConfig(BaseModel): EMBEDDINGS_MODEL: str = "mxbai-embed-large" EMBEDDINGS_VECTOR_DIM: int = 1024 - RERANK: bool = True RERANK_PROVIDER: str = "mxbai" RERANK_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2" diff --git a/src/haiku/rag/reranking/__init__.py b/src/haiku/rag/reranking/__init__.py index 1a3920ca..80df6bcb 100644 --- a/src/haiku/rag/reranking/__init__.py +++ b/src/haiku/rag/reranking/__init__.py @@ -9,29 +9,30 @@ except ImportError: _reranker: RerankerBase | None = None -def get_reranker() -> RerankerBase: +def get_reranker() -> RerankerBase | None: """ Factory function to get the appropriate reranker based on the configuration. + Returns None if the required package is not available. """ global _reranker if _reranker is not None: return _reranker if Config.RERANK_PROVIDER == "mxbai": - from haiku.rag.reranking.mxbai import MxBAIReranker + try: + from haiku.rag.reranking.mxbai import MxBAIReranker - _reranker = MxBAIReranker() - return _reranker + _reranker = MxBAIReranker() + return _reranker + except ImportError: + return None if Config.RERANK_PROVIDER == "cohere": try: from haiku.rag.reranking.cohere import CohereReranker - except ImportError: - raise ImportError( - "Cohere reranker requires the 'cohere' package. " - "Please install haiku.rag with the 'cohere' extra:" - "uv pip install haiku.rag[cohere]" - ) - _reranker = CohereReranker() - return _reranker - raise ValueError(f"Unsupported reranker provider: {Config.RERANK_PROVIDER}") + _reranker = CohereReranker() + return _reranker + except ImportError: + return None + + return None diff --git a/tests/test_reranker.py b/tests/test_reranker.py index b3de96cf..00e28b3b 100644 --- a/tests/test_reranker.py +++ b/tests/test_reranker.py @@ -1,7 +1,6 @@ import pytest from haiku.rag.reranking.base import RerankerBase -from haiku.rag.reranking.mxbai import MxBAIReranker from haiku.rag.store.models.chunk import Chunk chunks = [ @@ -30,12 +29,17 @@ async def test_reranker_base(): @pytest.mark.asyncio async def test_mxbai_reranker(): - reranker = MxBAIReranker() - reranked = await reranker.rerank( - "Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2 - ) - assert [chunk.document_id for chunk, score in reranked] == [0, 2] - assert all(isinstance(score, float) for chunk, score in reranked) + try: + from haiku.rag.reranking.mxbai import MxBAIReranker + + reranker = MxBAIReranker() + reranked = await reranker.rerank( + "Who wrote 'To Kill a Mockingbird'?", chunks, top_n=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("MxBAI package not installed") @pytest.mark.asyncio diff --git a/uv.lock b/uv.lock index ffa6a2d6..64107a5d 100644 --- a/uv.lock +++ b/uv.lock @@ -886,7 +886,6 @@ dependencies = [ { name = "docling" }, { name = "fastmcp" }, { name = "httpx" }, - { name = "mxbai-rerank" }, { name = "ollama" }, { name = "pydantic" }, { name = "python-dotenv" }, @@ -904,6 +903,9 @@ anthropic = [ cohere = [ { name = "cohere" }, ] +mxbai = [ + { name = "mxbai-rerank" }, +] openai = [ { name = "openai" }, ] @@ -931,7 +933,7 @@ requires-dist = [ { name = "docling", specifier = ">=2.15.0" }, { name = "fastmcp", specifier = ">=2.8.1" }, { name = "httpx", specifier = ">=0.28.1" }, - { name = "mxbai-rerank", specifier = ">=0.1.6" }, + { name = "mxbai-rerank", marker = "extra == 'mxbai'", specifier = ">=0.1.6" }, { name = "ollama", specifier = ">=0.5.1" }, { name = "openai", marker = "extra == 'openai'", specifier = ">=1.0.0" }, { name = "pydantic", specifier = ">=2.11.7" }, @@ -943,7 +945,7 @@ requires-dist = [ { name = "voyageai", marker = "extra == 'voyageai'", specifier = ">=0.3.2" }, { name = "watchfiles", specifier = ">=1.1.0" }, ] -provides-extras = ["voyageai", "openai", "anthropic", "cohere"] +provides-extras = ["voyageai", "openai", "anthropic", "cohere", "mxbai"] [package.metadata.requires-dev] dev = [