amplegest / tests /test_reranker.py
Viney's picture
deploy: Primer app for HF Spaces β€” clean orphan history
35676b4
Raw
History Blame Contribute Delete
2.86 kB
"""tests/test_reranker.py β€” unit tests for storage/reranker.py.
We don't require the actual cross-encoder model to be downloaded; the
tests use a fake model object that scores by simple keyword overlap so
we can assert deterministic ordering changes.
"""
from __future__ import annotations
import pytest
from storage import reranker
class _FakeModel:
"""Minimal stand-in for sentence_transformers.CrossEncoder.
Returns higher scores for candidates whose text shares more words with the query.
"""
def predict(self, pairs):
scores = []
for query, text in pairs:
q_tokens = set(query.lower().split())
t_tokens = set(text.lower().split())
scores.append(float(len(q_tokens & t_tokens)))
return scores
@pytest.fixture(autouse=True)
def _reset(monkeypatch):
"""Reset the module-level model cache between tests."""
reranker.reset_for_test()
yield
reranker.reset_for_test()
def _patch_with_fake(monkeypatch, model=None):
monkeypatch.setattr(reranker, "_get_model", lambda: model or _FakeModel())
def test_passthrough_when_few_candidates(monkeypatch):
_patch_with_fake(monkeypatch)
candidates = [{"text": "a"}, {"text": "b"}]
out = reranker.rerank("anything", candidates, top_k=3)
assert out == candidates
def test_truncates_to_top_k(monkeypatch):
_patch_with_fake(monkeypatch)
candidates = [{"text": f"chunk {i}"} for i in range(10)]
out = reranker.rerank("query", candidates, top_k=3)
assert len(out) == 3
def test_relevant_chunk_promoted_to_top(monkeypatch):
"""Candidate with most query-token overlap should rank #1 after rerank."""
_patch_with_fake(monkeypatch)
candidates = [
{"text": "the quick brown fox"},
{"text": "completely unrelated content here"},
{"text": "gross margin pressure from input costs and tariff exposure"},
{"text": "share buybacks resumed this quarter"},
{"text": "another unrelated thing"},
]
query = "gross margin pressure tariff"
out = reranker.rerank(query, candidates, top_k=1)
assert out[0]["text"].startswith("gross margin pressure")
def test_falls_back_when_model_unavailable(monkeypatch):
"""If model loading fails, return original order truncated to top_k."""
monkeypatch.setattr(reranker, "_get_model", lambda: None)
candidates = [{"text": f"c{i}"} for i in range(5)]
out = reranker.rerank("q", candidates, top_k=2)
assert out == candidates[:2]
def test_falls_back_when_predict_raises(monkeypatch):
class _BadModel:
def predict(self, pairs):
raise RuntimeError("boom")
_patch_with_fake(monkeypatch, model=_BadModel())
candidates = [{"text": f"c{i}"} for i in range(5)]
out = reranker.rerank("q", candidates, top_k=2)
assert out == candidates[:2]