Remove reranker cache
This commit is contained in:
parent
e7ce9694b5
commit
b8cf8f5198
3 changed files with 17 additions and 50 deletions
|
|
@ -1,6 +1,8 @@
|
||||||
# Changelog
|
# Changelog
|
||||||
## [Unreleased]
|
## [Unreleased]
|
||||||
|
|
||||||
|
- **Reranker Factory**: Removed unreliable `id(config)`-based caching from `get_reranker()`; factory now always instantiates fresh
|
||||||
|
|
||||||
## [0.26.7] - 2026-01-20
|
## [0.26.7] - 2026-01-20
|
||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
@ -17,6 +19,8 @@
|
||||||
|
|
||||||
### Changed
|
### Changed
|
||||||
|
|
||||||
|
- **MCP Error Handling**: MCP tools now let exceptions propagate naturally; FastMCP converts them to proper MCP error responses
|
||||||
|
- **Chunk Contextualization**: Consolidated duplicate `contextualize` logic into `Chunk.contextualize_content()` method
|
||||||
- **Type Checker**: Replaced pyright with [ty](https://github.com/astral-sh/ty), Astral's extremely fast Python type checker
|
- **Type Checker**: Replaced pyright with [ty](https://github.com/astral-sh/ty), Astral's extremely fast Python type checker
|
||||||
- Added explicit `Agent[Deps, Output]` type annotations to all pydantic-ai agents for better type inference
|
- Added explicit `Agent[Deps, Output]` type annotations to all pydantic-ai agents for better type inference
|
||||||
- Removed ~24 unnecessary `# type: ignore` comments that ty correctly infers
|
- Removed ~24 unnecessary `# type: ignore` comments that ty correctly infers
|
||||||
|
|
|
||||||
|
|
@ -3,8 +3,6 @@ import os
|
||||||
from haiku.rag.config import AppConfig, Config
|
from haiku.rag.config import AppConfig, Config
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
|
|
||||||
_reranker_cache: dict[int, RerankerBase | None] = {}
|
|
||||||
|
|
||||||
|
|
||||||
def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
"""
|
"""
|
||||||
|
|
@ -17,50 +15,41 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
Returns:
|
Returns:
|
||||||
A reranker instance if configured, None otherwise.
|
A reranker instance if configured, None otherwise.
|
||||||
"""
|
"""
|
||||||
# Use config id as cache key to support multiple configs
|
|
||||||
config_id = id(config)
|
|
||||||
if config_id in _reranker_cache:
|
|
||||||
return _reranker_cache[config_id]
|
|
||||||
|
|
||||||
reranker: RerankerBase | None = None
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "mxbai":
|
if config.reranking.model and config.reranking.model.provider == "mxbai":
|
||||||
try:
|
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"
|
||||||
reranker = MxBAIReranker()
|
return MxBAIReranker()
|
||||||
except ImportError: # pragma: no cover
|
except ImportError: # pragma: no cover
|
||||||
reranker = None
|
return None
|
||||||
|
|
||||||
elif config.reranking.model and config.reranking.model.provider == "cohere":
|
if config.reranking.model and config.reranking.model.provider == "cohere":
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.cohere import CohereReranker
|
from haiku.rag.reranking.cohere import CohereReranker
|
||||||
|
|
||||||
reranker = CohereReranker()
|
return CohereReranker()
|
||||||
except ImportError: # pragma: no cover
|
except ImportError: # pragma: no cover
|
||||||
reranker = None
|
return None
|
||||||
|
|
||||||
elif config.reranking.model and config.reranking.model.provider == "vllm":
|
if config.reranking.model and config.reranking.model.provider == "vllm":
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
|
|
||||||
base_url = config.reranking.model.base_url
|
base_url = config.reranking.model.base_url
|
||||||
if not 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")
|
||||||
reranker = VLLMReranker(config.reranking.model.name, base_url)
|
return VLLMReranker(config.reranking.model.name, base_url)
|
||||||
except ImportError: # pragma: no cover
|
except ImportError: # pragma: no cover
|
||||||
reranker = None
|
return None
|
||||||
|
|
||||||
elif config.reranking.model and config.reranking.model.provider == "zeroentropy":
|
if config.reranking.model and config.reranking.model.provider == "zeroentropy":
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
||||||
|
|
||||||
# Use configured model or default to zerank-1
|
|
||||||
model = config.reranking.model.name or "zerank-1"
|
model = config.reranking.model.name or "zerank-1"
|
||||||
reranker = ZeroEntropyReranker(model)
|
return ZeroEntropyReranker(model)
|
||||||
except ImportError: # pragma: no cover
|
except ImportError: # pragma: no cover
|
||||||
reranker = None
|
return None
|
||||||
|
|
||||||
_reranker_cache[config_id] = reranker
|
return None
|
||||||
return reranker
|
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@ from pathlib import Path
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.config.models import AppConfig, ModelConfig, RerankingConfig
|
from haiku.rag.config.models import AppConfig, ModelConfig, RerankingConfig
|
||||||
from haiku.rag.reranking import _reranker_cache, get_reranker
|
from haiku.rag.reranking import get_reranker
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
|
@ -13,14 +13,6 @@ def vcr_cassette_dir():
|
||||||
return str(Path(__file__).parent / "cassettes" / "test_reranker")
|
return str(Path(__file__).parent / "cassettes" / "test_reranker")
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
|
||||||
def clear_reranker_cache():
|
|
||||||
"""Clear the reranker cache before each test."""
|
|
||||||
_reranker_cache.clear()
|
|
||||||
yield
|
|
||||||
_reranker_cache.clear()
|
|
||||||
|
|
||||||
|
|
||||||
chunks = [
|
chunks = [
|
||||||
Chunk(content=content, document_id=str(i))
|
Chunk(content=content, document_id=str(i))
|
||||||
for i, content in enumerate(
|
for i, content in enumerate(
|
||||||
|
|
@ -202,24 +194,6 @@ class TestGetReranker:
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pytest.skip("Zero Entropy package not installed")
|
pytest.skip("Zero Entropy package not installed")
|
||||||
|
|
||||||
def test_caching_returns_same_instance(self):
|
|
||||||
config = AppConfig(reranking=RerankingConfig(model=None))
|
|
||||||
result1 = get_reranker(config)
|
|
||||||
result2 = get_reranker(config)
|
|
||||||
assert result1 is result2
|
|
||||||
|
|
||||||
def test_different_configs_get_separate_cache_entries(self):
|
|
||||||
config1 = AppConfig(reranking=RerankingConfig(model=None))
|
|
||||||
config2 = AppConfig(reranking=RerankingConfig(model=None))
|
|
||||||
|
|
||||||
result1 = get_reranker(config1)
|
|
||||||
result2 = get_reranker(config2)
|
|
||||||
|
|
||||||
# Both return None, but they should be cached separately
|
|
||||||
assert result1 is None
|
|
||||||
assert result2 is None
|
|
||||||
assert len(_reranker_cache) == 2
|
|
||||||
|
|
||||||
def test_unknown_provider_returns_none(self):
|
def test_unknown_provider_returns_none(self):
|
||||||
config = AppConfig(
|
config = AppConfig(
|
||||||
reranking=RerankingConfig(
|
reranking=RerankingConfig(
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue