haiku.rag.config exported two configuration instances: the lazy _config behind get_config/set_config, and Config, loaded at import time. Nothing linked them, and eleven signatures captured Config as a default argument, so set_config could not reach the factories, the client, the store or the MCP server. reranking/base.py went further and snapshotted the configured reranker name into a class attribute at import. Config is removed. Internal defaults are config: AppConfig | None = None, resolved through get_config() per call. RerankerBase._model is None and CohereReranker takes its model name as an argument, like every other reranker. The suite patched attributes on Config while production read the instance get_config() returns, a different object, so those patches were no-ops waiting to happen. They now go through get_config().
55 lines
2 KiB
Python
55 lines
2 KiB
Python
from haiku.rag.config import AppConfig, get_config
|
|
from haiku.rag.reranking.base import RerankerBase
|
|
|
|
|
|
def get_reranker(config: AppConfig | None = None) -> RerankerBase | None:
|
|
"""Build the configured reranker, or None if reranking is disabled or its
|
|
optional dependency is not installed."""
|
|
config = config if config is not None else get_config()
|
|
model = config.reranking.model
|
|
if model is None:
|
|
return None
|
|
|
|
if config.reranking.multimodal and model.provider != "vllm":
|
|
raise ValueError("reranking.multimodal is only supported on the vllm provider")
|
|
|
|
try:
|
|
if model.provider == "cohere":
|
|
from haiku.rag.reranking.cohere import CohereReranker
|
|
|
|
return CohereReranker(model.name)
|
|
|
|
if model.provider == "vllm":
|
|
if not model.base_url:
|
|
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
|
from haiku.rag.reranking.vllm import VLLMReranker
|
|
|
|
return VLLMReranker(model.name, model.base_url)
|
|
|
|
if model.provider == "zeroentropy":
|
|
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
|
|
|
return ZeroEntropyReranker(model.name or "zerank-1")
|
|
|
|
if model.provider == "jina":
|
|
from haiku.rag.reranking.jina import JinaReranker
|
|
|
|
return JinaReranker(model.name or "jina-reranker-v3")
|
|
|
|
if model.provider == "jina-local":
|
|
from haiku.rag.reranking.jina_local import JinaLocalReranker
|
|
|
|
return JinaLocalReranker(model.name or "jinaai/jina-reranker-v3")
|
|
|
|
if model.provider == "cross-encoder":
|
|
if not model.name:
|
|
raise ValueError(
|
|
"cross-encoder reranker requires name in reranking.model"
|
|
)
|
|
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
|
|
|
|
return CrossEncoderReranker(model.name)
|
|
except ImportError: # pragma: no cover
|
|
return None
|
|
|
|
raise ValueError(f"Unknown reranking provider: {model.provider}")
|