diff --git a/src/haiku/rag/config.py b/src/haiku/rag/config.py index 8fb2bef1..d4fd7421 100644 --- a/src/haiku/rag/config.py +++ b/src/haiku/rag/config.py @@ -19,8 +19,8 @@ class AppConfig(BaseModel): EMBEDDINGS_MODEL: str = "mxbai-embed-large" EMBEDDINGS_VECTOR_DIM: int = 1024 - RERANK_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2" RERANK_PROVIDER: str = "mxbai" + RERANK_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2" QA_PROVIDER: str = "ollama" QA_MODEL: str = "qwen3" diff --git a/src/haiku/rag/reranking/__init__.py b/src/haiku/rag/reranking/__init__.py index 8fbdae31..ccef8b9f 100644 --- a/src/haiku/rag/reranking/__init__.py +++ b/src/haiku/rag/reranking/__init__.py @@ -6,11 +6,21 @@ try: except ImportError: pass +_reranker: RerankerBase | None = None + def get_reranker() -> RerankerBase: """ Factory function to get the appropriate reranker based on the configuration. """ + global _reranker + if _reranker is not None: + return _reranker + if Config.RERANK_PROVIDER == "mxbai": + from haiku.rag.reranking.mxbai import MxBAIReranker + + _reranker = MxBAIReranker() + return _reranker if Config.RERANK_PROVIDER == "cohere": try: @@ -21,6 +31,7 @@ def get_reranker() -> RerankerBase: "Please install haiku.rag with the 'cohere' extra:" "uv pip install haiku.rag --extra cohere" ) - return CohereReranker() + _reranker = CohereReranker() + return _reranker raise ValueError(f"Unsupported reranker provider: {Config.RERANK_PROVIDER}")