"""Adaptive local cross-encoder reranking for genuinely ambiguous retrievals.""" from __future__ import annotations import os import threading from dataclasses import dataclass from typing import Sequence from app.rag.models import EvidenceRequest, EvidenceType, RetrievedEvidence @dataclass(frozen=True) class RerankDecision: use: bool reason: str class EvidenceReranker: _model = None _lock = threading.Lock() def decide( self, request: EvidenceRequest, candidates: Sequence[RetrievedEvidence], ) -> RerankDecision: if request.anchor_evidence_ids or request.selection_anchors: return RerankDecision(False, "explicit_anchor") if len(candidates) < 9: return RerankDecision(False, "small_candidate_pool") documents = {item.evidence.document_id for item in candidates[:24]} lowered = request.query.casefold() visual = any( term in lowered for term in ("table", "figure", "plot", "chart", "diagram", "metric") ) or any( item.evidence.element_type in {EvidenceType.TABLE, EvidenceType.FIGURE, EvidenceType.PLOT, EvidenceType.DIAGRAM} for item in candidates[:16] ) comparative = len(documents) > 1 and any( term in lowered for term in ("compare", "contrast", "across", "difference", "versus", "synthesis") ) if visual: return RerankDecision(True, "visual_or_table_ambiguity") if comparative: return RerankDecision(True, "cross_document_comparison") if len(candidates) >= 18 and len(documents) > 1: return RerankDecision(True, "large_cross_document_pool") return RerankDecision(False, "fused_ranking_sufficient") def rerank( self, query: str, candidates: Sequence[RetrievedEvidence], *, limit: int = 24, ) -> list[RetrievedEvidence]: head = list(candidates[:limit]) if not head: return list(candidates) model = self._get_model() scores = list(model.rerank(query, [item.evidence.index_text for item in head], batch_size=16)) for item, score in zip(head, scores): item.rerank_score = float(score) head.sort( key=lambda item: ( -(item.rerank_score if item.rerank_score is not None else float("-inf")), -item.fused_score, ) ) return [*head, *candidates[limit:]] @classmethod def _get_model(cls): if cls._model is not None: return cls._model with cls._lock: if cls._model is None: from fastembed.rerank.cross_encoder import TextCrossEncoder cls._model = TextCrossEncoder( model_name=os.getenv( "RAG_RERANK_MODEL", "Xenova/ms-marco-MiniLM-L-6-v2", ), lazy_load=True, ) return cls._model