File size: 2,411 Bytes
8eb009a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
from core.config import ProjectSettings
from serving.model_loader import _load_reranker_runtime


def test_load_reranker_runtime_none_backend() -> None:
    settings = ProjectSettings(reranker_backend="none", use_cross_encoder=False)

    runtime = _load_reranker_runtime(settings)

    assert runtime.reranker is None
    assert runtime.reranker_backend == "none"
    assert runtime.fallback_used is False


def test_load_reranker_runtime_heuristic_backend() -> None:
    settings = ProjectSettings(reranker_backend="heuristic", use_cross_encoder=False)

    runtime = _load_reranker_runtime(settings)

    assert runtime.reranker is not None
    assert runtime.reranker_backend == "heuristic"
    assert runtime.fallback_used is False


def test_load_reranker_runtime_cross_encoder_backend_uses_mock(monkeypatch) -> None:
    class FakeCrossEncoderReranker:
        def __init__(self, model_name: str) -> None:
            self.model_name = model_name
            self.backend_name = "cross-encoder"

        def rank(self, claim, evidence_list):  # noqa: ANN001
            return list(reversed(evidence_list))

    monkeypatch.setattr("serving.model_loader.CrossEncoderReranker", FakeCrossEncoderReranker)

    settings = ProjectSettings(
        reranker_backend="cross_encoder",
        use_cross_encoder=True,
        cross_encoder_model="fake-model",
    )

    runtime = _load_reranker_runtime(settings)

    assert runtime.reranker is not None
    assert getattr(runtime.reranker, "model_name", None) == "fake-model"
    assert runtime.reranker_backend == "cross_encoder"
    assert runtime.cross_encoder_model == "fake-model"
    assert runtime.fallback_used is False


def test_load_reranker_runtime_cross_encoder_falls_back(monkeypatch) -> None:
    class ExplodingCrossEncoderReranker:
        def __init__(self, model_name: str) -> None:  # noqa: ARG002
            raise RuntimeError("download failed")

    monkeypatch.setattr("serving.model_loader.CrossEncoderReranker", ExplodingCrossEncoderReranker)

    settings = ProjectSettings(
        reranker_backend="cross_encoder",
        use_cross_encoder=True,
        cross_encoder_model="fake-model",
    )

    runtime = _load_reranker_runtime(settings)

    assert runtime.reranker is not None
    assert runtime.reranker_backend == "heuristic"
    assert runtime.cross_encoder_model == "fake-model"
    assert runtime.fallback_used is True