File size: 14,459 Bytes
3fc8e60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
"""Cross-encoder reranking, and the refusal gate that rides on its scores.

Bi-encoder vs cross-encoder β€” the distinction the whole stage rests on
---------------------------------------------------------------------
The retriever is a **bi-encoder**: query and passage are embedded *separately* and
compared by cosine. That is what makes it fast β€” every passage vector is precomputed β€”
and it is also its ceiling, because the passage was encoded without ever having seen the
query. "Can my employer keep my passport?" and a passage about custody of documents get
one shared summary vector each, and nuance is averaged away.

A **cross-encoder** concatenates the pair and runs them through the transformer together,
so attention operates across both. It cannot precompute anything β€” cost is O(candidates)
per query, not O(1) β€” which is precisely why it is used on 20 candidates and not on 181
chunks. Retrieve wide with the cheap model, judge narrow with the expensive one.

The output is an unbounded relevance logit, not a probability. That is useful twice: it
orders the candidates, and its absolute value is a calibrated signal of whether the
corpus contains an answer at all. :func:`apply_refusal_gate` uses the second property β€”
knowing when *not* to answer is the feature this project exists to demonstrate.
"""

from __future__ import annotations

import logging
import threading
from dataclasses import dataclass

from tokenizers import Tokenizer

from app.core.models import ScoredChunk
from app.core.settings import Settings, get_settings

logger = logging.getLogger(__name__)

_LOCK = threading.Lock()
_ENCODER: object | None = None


class RerankerUnavailableError(RuntimeError):
    """The cross-encoder could not be loaded."""


def get_cross_encoder(settings: Settings | None = None) -> object:
    """Process-wide cross-encoder.

    Two interchangeable backends. ``fastembed`` runs the cross-encoder as ONNX on CPU
    with no PyTorch in the image; ``sentence-transformers`` offers the more familiar API
    at the cost of a multi-gigabyte torch dependency. ONNX is the default because image
    size is what decides whether this fits a free Cloud Run or Hugging Face Space at all.
    Either backend serves whichever checkpoint ``reranker_model`` names.
    """
    global _ENCODER  # noqa: PLW0603 - deliberate process-wide singleton
    if _ENCODER is not None:
        return _ENCODER
    cfg = settings or get_settings()
    with _LOCK:
        if _ENCODER is None:
            cfg.models_cache_dir.mkdir(parents=True, exist_ok=True)
            logger.info("loading reranker %s via %s", cfg.reranker_model, cfg.reranker_backend)
            if cfg.reranker_backend == "fastembed":
                from fastembed.rerank.cross_encoder import TextCrossEncoder

                _ENCODER = TextCrossEncoder(
                    model_name=cfg.reranker_model,
                    cache_dir=str(cfg.models_cache_dir),
                )
            else:  # pragma: no cover - optional heavyweight backend
                try:
                    from sentence_transformers import CrossEncoder
                except ImportError as exc:
                    raise RerankerUnavailableError(
                        "reranker_backend='sentence-transformers' requires the optional "
                        "`sentence-transformers` extra; the default 'fastembed' backend "
                        "runs the same weights as ONNX with no torch dependency."
                    ) from exc
                _ENCODER = CrossEncoder(cfg.reranker_model)
    return _ENCODER


def _prepare(documents: list[str], settings: Settings) -> list[tuple[str, int]]:
    """Cap each passage at ``rerank_max_tokens`` and return it with its token length.

    The length is returned rather than recomputed later so bucketing sorts on the real
    padded cost instead of a character-count proxy.
    """
    tokenizer = get_reranker_tokenizer(settings)
    if tokenizer is None:
        return [(document, len(document)) for document in documents]
    limit = settings.rerank_max_tokens
    prepared: list[tuple[str, int]] = []
    for document in documents:
        ids = tokenizer.encode(document, add_special_tokens=False).ids
        if 0 < limit < len(ids):
            prepared.append((tokenizer.decode(ids[:limit]), limit))
        else:
            prepared.append((document, len(ids)))
    return prepared


def _score_batch(query: str, documents: list[str], settings: Settings) -> list[float]:
    encoder = get_cross_encoder(settings)
    if settings.reranker_backend == "fastembed":
        scores = encoder.rerank(query, documents, batch_size=len(documents))  # type: ignore[attr-defined]
        return [float(value) for value in scores]
    pairs = [(query, document) for document in documents]  # pragma: no cover
    return [float(value) for value in encoder.predict(pairs)]  # type: ignore[attr-defined]


def _score(query: str, documents: list[str], settings: Settings) -> list[float]:
    """Score every (query, passage) pair, length-bucketed for throughput.

    A transformer batch is padded to its longest member, so scoring one 506-token
    passage alongside nineteen 90-token ones costs as much as twenty long ones. Sorting
    by length and scoring in small batches removes that waste. The transformation is
    purely a reordering β€” scores are bit-identical to a single large batch, which
    ``tests/test_rerank.py`` asserts β€” so it buys latency and changes nothing else.
    """
    if not documents:
        return []
    prepared = _prepare(documents, settings)
    batch_size = max(1, settings.rerank_batch_size)
    order = sorted(range(len(prepared)), key=lambda i: prepared[i][1])
    scores = [0.0] * len(prepared)
    for start in range(0, len(order), batch_size):
        indices = order[start : start + batch_size]
        batch = [prepared[i][0] for i in indices]
        for index, value in zip(indices, _score_batch(query, batch, settings), strict=True):
            scores[index] = value
    return scores


def rerank(
    query: str,
    candidates: list[ScoredChunk],
    settings: Settings | None = None,
    top_k: int | None = None,
) -> list[ScoredChunk]:
    """Rescore candidates with the cross-encoder and keep the best ``top_k``."""
    cfg = settings or get_settings()
    limit = top_k if top_k is not None else cfg.rerank_top_k
    if not candidates:
        return []

    scores = _score(query, [candidate.chunk.text for candidate in candidates], cfg)
    if len(scores) != len(candidates):  # pragma: no cover - defensive
        raise RerankerUnavailableError(
            f"reranker returned {len(scores)} scores for {len(candidates)} candidates"
        )

    scored = [
        candidate.model_copy(update={"rerank_score": score})
        for candidate, score in zip(candidates, scores, strict=True)
    ]
    # Ties broken by chunk_id so the same input always yields the same output order.
    scored.sort(key=lambda item: (-(item.rerank_score or 0.0), item.chunk.chunk_id))
    return [
        item.model_copy(update={"final_rank": position + 1})
        for position, item in enumerate(scored[:limit])
    ]


def passthrough(
    candidates: list[ScoredChunk],
    settings: Settings | None = None,
    top_k: int | None = None,
) -> list[ScoredChunk]:
    """Take the top fused candidates without reranking.

    This is the ``--no-rerank`` arm of the evaluation. It exists so the reranker's
    contribution is a measured delta rather than an assertion.
    """
    cfg = settings or get_settings()
    limit = top_k if top_k is not None else cfg.rerank_top_k
    return [
        candidate.model_copy(update={"final_rank": position + 1})
        for position, candidate in enumerate(candidates[:limit])
    ]


@dataclass(frozen=True, slots=True)
class GateOutcome:
    """Whether the corpus covers the question, and the evidence either way."""

    covered: bool
    evidence: tuple[ScoredChunk, ...]
    near_misses: tuple[ScoredChunk, ...]
    best_score: float | None
    floor: float
    best_dense: float | None = None
    dense_floor: float | None = None
    reason: str = ""
    signal: str = ""


def apply_refusal_gate(
    reranked: list[ScoredChunk],
    settings: Settings | None = None,
    *,
    best_dense: float | None = None,
    scope: object | None = None,
) -> GateOutcome:
    """Decide whether the retrieved evidence is good enough to answer from.

    Retrieval always returns *something*: nearest-neighbour search over a non-empty index
    cannot return nothing, and a question the corpus has never heard of still comes back
    with five confidently-ranked passages. Generating from them is exactly how a RAG
    system produces a fluent, well-cited, wrong answer.

    Three independent signals must all pass, because each catches what the others miss β€”
    all three thresholds were fitted on the labelled eval set, not chosen by intuition:

    1. **Scope.** Does the question name a legal system the corpus does not contain?
       No similarity score can answer this; see ``app.rag.scope``.
    2. **Domain floor** on the best dense similarity. Answers "is this question even
       about the corpus's subject matter?"
    3. **Relevance floor** on the best cross-encoder score. Answers "is the single best
       passage actually responsive?"

    A refusal hands back the near misses, so it is auditable: the user sees what was
    considered and can judge the call themselves.

    With reranking disabled the cross-encoder signal does not exist, so only the scope
    and domain checks apply. AUDIT.md reports refusal accuracy per configuration rather
    than implying the un-reranked path is equally protected β€” measured, it is not.
    """
    cfg = settings or get_settings()
    if scope is not None:
        reason = getattr(scope, "reason", "")
        signal = getattr(scope, "signal", "out-of-scope")
        return GateOutcome(
            covered=False,
            evidence=(),
            near_misses=tuple(reranked[: cfg.refusal_near_miss_count]),
            best_score=reranked[0].rerank_score if reranked else None,
            floor=cfg.refusal_score_floor,
            best_dense=best_dense,
            dense_floor=cfg.refusal_dense_floor,
            reason=reason,
            signal=signal,
        )

    if not reranked:
        return GateOutcome(
            covered=False,
            evidence=(),
            near_misses=(),
            best_score=None,
            floor=cfg.refusal_score_floor,
            best_dense=best_dense,
            dense_floor=cfg.refusal_dense_floor,
            reason="Retrieval returned no candidates.",
            signal="empty-retrieval",
        )

    if best_dense is not None and best_dense < cfg.refusal_dense_floor:
        return GateOutcome(
            covered=False,
            evidence=(),
            near_misses=tuple(reranked[: cfg.refusal_near_miss_count]),
            best_score=reranked[0].rerank_score,
            floor=cfg.refusal_score_floor,
            best_dense=best_dense,
            dense_floor=cfg.refusal_dense_floor,
            reason="No indexed passage is close enough to this question's subject matter.",
            signal="below-domain-floor",
        )

    best = reranked[0].rerank_score
    if best is None or best >= cfg.refusal_score_floor:
        return GateOutcome(
            covered=True,
            evidence=tuple(reranked),
            near_misses=(),
            best_score=best,
            floor=cfg.refusal_score_floor,
            best_dense=best_dense,
            dense_floor=cfg.refusal_dense_floor,
        )
    return GateOutcome(
        covered=False,
        evidence=(),
        near_misses=tuple(reranked[: cfg.refusal_near_miss_count]),
        best_score=best,
        floor=cfg.refusal_score_floor,
        best_dense=best_dense,
        dense_floor=cfg.refusal_dense_floor,
        reason="The closest passage is not responsive enough to answer from.",
        signal="below-relevance-floor",
    )


_TOKENIZERS: dict[str, Tokenizer] = {}


def get_reranker_tokenizer(settings: Settings | None = None) -> Tokenizer | None:
    """The reranker's own tokenizer, used to cap passage length before scoring.

    Two details here are load-bearing, and both were found by a CI failure that could
    not be reproduced locally:

    1. **The cross-encoder is loaded first.** The tokenizer is located by globbing the
       model cache, so on a cold cache β€” a fresh container, a CI runner with no restored
       cache β€” the file does not exist yet because nothing has downloaded it. Forcing the
       encoder to load first guarantees the files are on disk before the search.

    2. **A miss is never cached.** With ``lru_cache`` a single early miss was memoised for
       the lifetime of the process, so truncation stayed silently disabled long after the
       model had arrived. Only successful lookups are cached.

    Why it matters that truncation actually happens: without it, passages exceed the
    model's window and the runtime truncates per batch, which makes a score depend on
    which other documents happen to share its batch. Length bucketing then stops being a
    pure reordering. With truncation on, scores are identical across batch sizes β€”
    measured at max |delta| 0.0000 over the corpus.
    """
    cfg = settings or get_settings()
    cached = _TOKENIZERS.get(cfg.reranker_model)
    if cached is not None:
        return cached

    # Ensure the model β€” and therefore its tokenizer.json β€” is on disk before looking.
    get_cross_encoder(cfg)

    stem = cfg.reranker_model.split("/")[-1]
    candidates = [
        path for path in cfg.models_cache_dir.rglob("tokenizer.json") if stem in str(path)
    ]
    if not candidates:
        logger.warning(
            "no cached tokenizer for %s under %s; reranking without truncation, which "
            "makes scores batch-dependent",
            cfg.reranker_model,
            cfg.models_cache_dir,
        )
        return None

    tokenizer = Tokenizer.from_file(str(max(candidates, key=lambda path: path.stat().st_mtime)))
    _TOKENIZERS[cfg.reranker_model] = tokenizer
    return tokenizer


def reset_cross_encoder() -> None:
    """Drop the cached encoder. Used by tests."""
    global _ENCODER  # noqa: PLW0603 - mirrors get_cross_encoder
    with _LOCK:
        _ENCODER = None
    _TOKENIZERS.clear()