"""Small dependency-free lexical retriever for the recall demo.""" from __future__ import annotations import math import re from collections import Counter from rag_cut.models import Chunk LATIN_TOKEN_RE = re.compile(r"[a-z0-9_]+", re.I) CJK_RE = re.compile(r"[\u3400-\u9fff]") def _tokens(text: str) -> list[str]: normalized = text.lower() tokens = LATIN_TOKEN_RE.findall(normalized) cjk = CJK_RE.findall(normalized) tokens.extend(cjk) tokens.extend("".join(cjk[i : i + 2]) for i in range(len(cjk) - 1)) return tokens def _search_text(chunk: Chunk) -> str: meta = chunk.meta fields = [ meta.get("heading") or "", meta.get("nearest_heading") or "", meta.get("embedding_text") or "", " ".join(meta.get("keywords") or []), chunk.content, ] return "\n".join(str(value) for value in fields if value) def recall_chunks(query: str, chunks: list[Chunk], top_k: int = 5) -> tuple[list[dict], int]: """Rank retrievable chunks with a compact BM25-style lexical score.""" candidates = [chunk for chunk in chunks if chunk.meta.get("retrieval", True)] query_tokens = _tokens(query) if not candidates or not query_tokens: return [], len(candidates) documents = [_tokens(_search_text(chunk)) for chunk in candidates] document_frequency = Counter(token for tokens in documents for token in set(tokens)) average_length = sum(len(tokens) for tokens in documents) / len(documents) query_frequency = Counter(query_tokens) scored: list[tuple[float, Chunk]] = [] for chunk, tokens in zip(candidates, documents): frequency = Counter(tokens) length_normalizer = 1.2 * (0.25 + 0.75 * len(tokens) / max(average_length, 1)) score = 0.0 for token, query_count in query_frequency.items(): term_count = frequency[token] if not term_count: continue inverse_frequency = math.log(1 + (len(documents) - document_frequency[token] + 0.5) / (document_frequency[token] + 0.5)) score += inverse_frequency * ((term_count * 2.2) / (term_count + length_normalizer)) * min(query_count, 2) if score > 0: scored.append((score, chunk)) scored.sort(key=lambda item: (-item[0], item[1].index)) results = [] for rank, (score, chunk) in enumerate(scored[:top_k], start=1): results.append( { "rank": rank, "chunk_index": chunk.index, "score": round(score, 6), "content": chunk.content, "heading": chunk.meta.get("heading") or chunk.meta.get("nearest_heading"), "pages": chunk.meta.get("pages") or ([chunk.meta["page"]] if chunk.meta.get("page") is not None else []), "block_types": chunk.block_types, "parent_chunk_id": chunk.meta.get("parent_chunk_id"), "is_sub_chunk": bool(chunk.meta.get("is_sub_chunk")), } ) return results, len(candidates)