haiku.rag/src/haiku/rag/reranking/__init__.py
2025-10-23 13:30:21 +03:00

37 lines
969 B
Python

import os
from haiku.rag.config import Config
from haiku.rag.reranking.base import RerankerBase
_reranker: RerankerBase | None = None
def get_reranker() -> RerankerBase | None:
"""
Factory function to get the appropriate reranker based on the configuration.
Returns None if if reranking is disabled.
"""
global _reranker
if _reranker is not None:
return _reranker
if Config.reranking.provider == "mxbai":
try:
from haiku.rag.reranking.mxbai import MxBAIReranker
os.environ["TOKENIZERS_PARALLELISM"] = "true"
_reranker = MxBAIReranker()
return _reranker
except ImportError:
return None
if Config.reranking.provider == "cohere":
try:
from haiku.rag.reranking.cohere import CohereReranker
_reranker = CohereReranker()
return _reranker
except ImportError:
return None
return None