Collapse get_reranker into a single guard and import-guarded dispatch

Replace the seven repeated `config.reranking.model and ... == provider`
checks and six per-branch ImportError handlers with one None guard and one
try/except around the provider dispatch.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Yiorgis Gozadinos 2026-06-29 14:24:43 +03:00
parent 38079ff89a
commit 654cb2b94c
No known key found for this signature in database

View file

@ -5,79 +5,55 @@ from haiku.rag.reranking.base import RerankerBase
def get_reranker(config: AppConfig = Config) -> RerankerBase | None: def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
""" """Build the configured reranker, or None if reranking is disabled or its
Factory function to get the appropriate reranker based on the configuration. optional dependency is not installed."""
Returns None if reranking is disabled. model = config.reranking.model
if model is None:
return None
Args: try:
config: Configuration to use. Defaults to global Config. if model.provider == "mxbai":
Returns:
A reranker instance if configured, None otherwise.
"""
if config.reranking.model and config.reranking.model.provider == "mxbai":
try:
from haiku.rag.reranking.mxbai import MxBAIReranker from haiku.rag.reranking.mxbai import MxBAIReranker
os.environ["TOKENIZERS_PARALLELISM"] = "true" os.environ["TOKENIZERS_PARALLELISM"] = "true"
return MxBAIReranker() return MxBAIReranker()
except ImportError: # pragma: no cover
return None
if config.reranking.model and config.reranking.model.provider == "cohere": if model.provider == "cohere":
try:
from haiku.rag.reranking.cohere import CohereReranker from haiku.rag.reranking.cohere import CohereReranker
return CohereReranker() return CohereReranker()
except ImportError: # pragma: no cover
return None
if config.reranking.model and config.reranking.model.provider == "vllm": if model.provider == "vllm":
try:
from haiku.rag.reranking.vllm import VLLMReranker from haiku.rag.reranking.vllm import VLLMReranker
base_url = config.reranking.model.base_url if not model.base_url:
if not base_url:
raise ValueError("vLLM reranker requires base_url in reranking.model") raise ValueError("vLLM reranker requires base_url in reranking.model")
return VLLMReranker(config.reranking.model.name, base_url) return VLLMReranker(model.name, model.base_url)
except ImportError: # pragma: no cover
return None
if config.reranking.model and config.reranking.model.provider == "zeroentropy": if model.provider == "zeroentropy":
try:
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
model = config.reranking.model.name or "zerank-1" return ZeroEntropyReranker(model.name or "zerank-1")
return ZeroEntropyReranker(model)
except ImportError: # pragma: no cover
return None
if config.reranking.model and config.reranking.model.provider == "jina": if model.provider == "jina":
from haiku.rag.reranking.jina import JinaReranker from haiku.rag.reranking.jina import JinaReranker
model = config.reranking.model.name or "jina-reranker-v3" return JinaReranker(model.name or "jina-reranker-v3")
return JinaReranker(model)
if config.reranking.model and config.reranking.model.provider == "jina-local": if model.provider == "jina-local":
try:
from haiku.rag.reranking.jina_local import JinaLocalReranker from haiku.rag.reranking.jina_local import JinaLocalReranker
model = config.reranking.model.name or "jinaai/jina-reranker-v3" return JinaLocalReranker(model.name or "jinaai/jina-reranker-v3")
return JinaLocalReranker(model)
except ImportError: # pragma: no cover
return None
if config.reranking.model and config.reranking.model.provider == "cross-encoder": if model.provider == "cross-encoder":
try:
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
name = config.reranking.model.name if not model.name:
if not name:
raise ValueError( raise ValueError(
"cross-encoder reranker requires name in reranking.model" "cross-encoder reranker requires name in reranking.model"
) )
return CrossEncoderReranker(name) return CrossEncoderReranker(model.name)
except ImportError: # pragma: no cover except ImportError: # pragma: no cover
return None return None
return None return None