File size: 3,424 Bytes
dc1b199
 
 
 
 
 
 
 
 
 
 
3f6fdc5
dc1b199
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Reranker that combines vector similarity, lexical overlap, and numeric overlap.

The combined rerank score is:
    score = α·vector_score + β·lexical_score + γ·numeric_score

Default weights: α=0.6, β=0.3, γ=0.1.
"""

import logging
import re

from app.config import settings
from app.models.schemas import RerankedResult, SearchResult

logger = logging.getLogger(__name__)

# Scoring weights
_ALPHA = 0.6  # vector similarity
_BETA = 0.30  # lexical term overlap
_GAMMA = 0.10  # numeric overlap

# Regex for extracting number-like tokens (integers, decimals, percentages, areas)
_NUM_PATTERN = re.compile(r"\b\d[\d,]*\.?\d*\b")


def rerank(
    query: str,
    results: list[SearchResult],
    top_n: int | None = None,
) -> list[RerankedResult]:
    """Rerank retrieval results using a weighted combination of signals.

    Signals used:

    * **Vector score** — cosine similarity from the vector store (primary signal).
    * **Lexical overlap** — proportion of query tokens present in the chunk text.
    * **Numeric overlap** — proportion of numbers in the query also found in chunk.

    Args:
        query: The original query string (usually concatenated bullets).
        results: Candidates from the retriever, already ordered by vector score.
        top_n: Number of results to return.  Defaults to
            ``settings.rerank_top_n``.

    Returns:
        Top-``top_n`` :class:`~app.models.schemas.RerankedResult` objects,
        ordered by ``rerank_score`` descending.

    Example::

        reranked = rerank(query="95 sqm semi-detached", results=candidates, top_n=3)
    """
    if top_n is None:
        top_n = settings.rerank_top_n

    query_tokens = _tokenise(query)
    query_numbers = set(_NUM_PATTERN.findall(query))

    reranked: list[RerankedResult] = []
    for r in results:
        chunk_tokens = _tokenise(r.text)
        chunk_numbers = set(_NUM_PATTERN.findall(r.text))

        lex = _jaccard(query_tokens, chunk_tokens)
        num = _overlap(query_numbers, chunk_numbers) if query_numbers else 0.0

        combined = _ALPHA * r.score + _BETA * lex + _GAMMA * num
        reranked.append(
            RerankedResult(
                **r.model_dump(),
                rerank_score=round(combined, 6),
            )
        )

    reranked.sort(key=lambda x: x.rerank_score, reverse=True)
    top = reranked[:top_n]
    logger.debug("Reranked %d → %d results", len(results), len(top))
    return top


def _tokenise(text: str) -> set[str]:
    """Lower-case and split ``text`` into a set of alpha/digit tokens.

    Args:
        text: Any plain-text string.

    Returns:
        Set of normalised tokens.
    """
    return set(re.findall(r"[a-z0-9]+", text.lower()))


def _jaccard(a: set[str], b: set[str]) -> float:
    """Compute Jaccard similarity between two token sets.

    Args:
        a: First token set.
        b: Second token set.

    Returns:
        Float in [0, 1].
    """
    if not a and not b:
        return 0.0
    return len(a & b) / len(a | b)


def _overlap(query_set: set[str], chunk_set: set[str]) -> float:
    """Fraction of ``query_set`` items found in ``chunk_set``.

    Args:
        query_set: Numbers extracted from the query.
        chunk_set: Numbers extracted from the chunk.

    Returns:
        Float in [0, 1].
    """
    if not query_set:
        return 0.0
    return len(query_set & chunk_set) / len(query_set)