from __future__ import annotations from typing import Any, Protocol, Sequence from rag.models import ParentDocument, RetrievalDocument, SearchHit, SourceRef # Vietnamese function words dropped before lexical matching — high-frequency, # low-signal; matching on these alone would make `search_lexical` return # near-arbitrary same-drug chunks instead of ones sharing real query terms. _LEXICAL_STOPWORDS = frozenset({ "cua", "va", "la", "cho", "khi", "co", "gi", "duoc", "voi", "the", "nao", "nhu", "o", "trong", "de", "hay", "mot", "nay", "day", "thi", "bi", "khong", "da", "se", "neu", "nen", "phai", "sao", }) class QueryEmbedder(Protocol): @property def dimensions(self) -> int: ... def embed_query(self, text: str) -> Sequence[float]: ... def _source_refs(payload: dict[str, Any]) -> tuple[SourceRef, ...]: explicit = payload.get("source_refs") or [] if explicit: return tuple( SourceRef( physical_page=int(item["physical_page"]), precision=item.get("precision", "region"), block_id=item.get("block_id"), bbox=tuple(item["bbox"]) if item.get("bbox") else None, source_crop=item.get("source_crop"), page_range=( tuple(item["page_range"]) if item.get("page_range") else None ), printed_page=item.get("printed_page"), printed_page_range=( tuple(item["printed_page_range"]) if item.get("printed_page_range") else None ), ) for item in explicit ) physical_range = payload.get("source_page_range") printed_range = payload.get("printed_page_range") attachments = payload.get("attachments") or [] attachment_refs = tuple( SourceRef( physical_page=int(item["physical_page"]), precision="region", block_id=item.get("block_id"), bbox=tuple(item["bbox"]) if item.get("bbox") else None, source_crop=item.get("source_crop"), page_range=(int(item["physical_page"]), int(item["physical_page"])), printed_page=( int(item["printed_page"]) if item.get("printed_page") is not None else None ), printed_page_range=( (int(item["printed_page"]), int(item["printed_page"])) if item.get("printed_page") is not None else None ), ) for item in attachments if item.get("physical_page") is not None ) if payload.get("chunk_kind") == "block_descriptor" and attachment_refs: return attachment_refs physical_page = ( physical_range[0] if physical_range else payload.get("heading_physical_page") ) base_refs: tuple[SourceRef, ...] = () if physical_page is not None: base_refs = ( SourceRef( physical_page=int(physical_page), precision="chunk_page_range", page_range=tuple(physical_range) if physical_range else None, printed_page=(int(printed_range[0]) if printed_range else None), printed_page_range=tuple(printed_range) if printed_range else None, ), ) return base_refs + attachment_refs def _document(payload: dict[str, Any]) -> RetrievalDocument: return RetrievalDocument( doc_id=payload["chunk_id"], drug_id=payload["drug_id"], drug_name=payload.get("drug_name"), kind=payload.get("chunk_kind", "prose"), text=payload["text"], section_key=payload["section_key"], source_refs=_source_refs(payload), parent_id=payload.get("parent_id"), requires_visual_check=( bool(payload.get("requires_visual_check")) or bool(payload.get("has_quarantined_content")) ), part_index=payload.get("part_index"), part_count=payload.get("part_count"), context_labels=tuple(payload.get("context_labels") or ()), ) class QdrantRetriever: def __init__( self, client: Any, collection_name: str, embedder: QueryEmbedder, ) -> None: self._client = client self._collection_name = collection_name self._embedder = embedder def search(self, query: str, drug_id: str, limit: int) -> list[SearchHit]: from qdrant_client.models import FieldCondition, Filter, MatchValue vector = list(self._embedder.embed_query(query)) if len(vector) != self._embedder.dimensions: raise ValueError( f"query vector has {len(vector)} dimensions; " f"expected {self._embedder.dimensions}" ) points = self._client.search( collection_name=self._collection_name, query_vector=vector, query_filter=Filter( must=[FieldCondition(key="drug_id", match=MatchValue(value=drug_id))] ), limit=limit, with_payload=True, ) return [ SearchHit(_document(dict(point.payload or {})), float(point.score)) for point in points ] def find_by_section(self, drug_id: str, section_key: str) -> list[SearchHit]: """Every chunk of one section, by payload filter — no vector involved. A `scroll`, not a `search`: this must not be a top-k. Paging continues until the offset is exhausted, because Qdrant's default page is 256 and a long section silently truncated would read as a complete answer. Score is 1.0 because the match is exact by construction. It is not a similarity and must not be compared against one. Results are re-sorted by `part_index` before returning. Qdrant scrolls in point-id order, and point ids are `uuid5(chunk_id)`, so the natural order is effectively random: PARACETAMOL's dosing section came back 3, 4, 1, 2, 0 — the answer opened mid-sentence on paediatric doses and buried "Liều lượng: Người lớn:" last. A section served out of order is a clinical hazard, not a formatting one: a reader who stops early stops in the middle of a different population's dose. """ from qdrant_client.models import FieldCondition, Filter, MatchValue scroll_filter = Filter( must=[ FieldCondition(key="drug_id", match=MatchValue(value=drug_id)), FieldCondition(key="section_key", match=MatchValue(value=section_key)), ] ) hits: list[tuple[dict, SearchHit]] = [] offset = None while True: points, offset = self._client.scroll( collection_name=self._collection_name, scroll_filter=scroll_filter, limit=256, offset=offset, with_payload=True, ) hits.extend( (dict(point.payload or {}), SearchHit(_document(dict(point.payload or {})), 1.0)) for point in points ) if offset is None: break # `part_index` is the chunker's own position within the section. A # payload missing it sorts last rather than raising: an unordered # section is worse than a scrambled one only if it also disappears. hits.sort(key=lambda item: item[0].get("part_index", 1 << 30)) return [hit for _, hit in hits] def search_lexical(self, query: str, drug_id: str, limit: int) -> list[SearchHit]: """Keyword/BM25-style candidates across ALL of one drug's sections, ranked by term overlap with `query`. Two live callers, both in `rag.service.RetrievalService`: (1) inside the deterministic `_section_hits` route, to find a NEIGHBOUR section whose text lexically matches strongly enough to pool in alongside the one the keyword route resolved (a "thận trọng" question can have its real answer only in "chống chỉ định" — see `_pooled_neighbour_ hits`'s docstring); (2) available for hybrid fusion with `search`'s dense results (`rag.fusion.reciprocal_rank_fusion`) in the similarity-fallback path, for a free-form question naming no section a paraphrase makes the exact-phrase `find_by_indication`- style match miss. Qdrant's `text` index tokenizes and matches individual query tokens (OR semantics across a `should` filter — no `min_should_match` needed since the caller fuses ranks, not raw hits). Common short function words are dropped before matching so they don't dilute every result with the same handful of low-signal hits; score is the count of distinct matched tokens, a transparent stand-in for a real BM25 score given no term-frequency/IDF statistics are computed here. """ from qdrant_client.models import FieldCondition, Filter, MatchText, MatchValue from rag.text import normalize_name tokens = sorted(set(normalize_name(query).split()) - _LEXICAL_STOPWORDS) tokens = [token for token in tokens if len(token) >= 2] if not tokens: return [] points, _ = self._client.scroll( collection_name=self._collection_name, scroll_filter=Filter( must=[FieldCondition(key="drug_id", match=MatchValue(value=drug_id))], should=[FieldCondition(key="text", match=MatchText(text=t)) for t in tokens], ), limit=max(limit * 4, 20), offset=None, with_payload=True, ) scored: list[tuple[int, SearchHit]] = [] for point in points: payload = dict(point.payload or {}) text_normalized = normalize_name(payload.get("text", "")) matched = sum(1 for t in tokens if t in text_normalized.split()) if matched == 0: continue scored.append((matched, SearchHit(_document(payload), float(matched)))) scored.sort(key=lambda item: -item[0]) return [hit for _, hit in scored[:limit]] def find_by_drug(self, drug_id: str) -> list[SearchHit]: """Every prose section of one drug, in book order — the monograph view. For a query that names the drug but no attribute ("PARACETAMOL"), a drug reference shows the whole monograph, not a "specify an attribute" prompt. A `scroll` filtered on `drug_id`, prose only (block descriptors stay out of a text answer), ordered by the book's section sequence then `part_index`. Each section's first chunk gets a `【heading】` so the result reads as a monograph, not a wall of text. """ from qdrant_client.models import FieldCondition, Filter, MatchValue from rag.sections import SECTION_ORDER scroll_filter = Filter( must=[ FieldCondition(key="drug_id", match=MatchValue(value=drug_id)), FieldCondition(key="chunk_kind", match=MatchValue(value="prose")), ] ) payloads: list[dict] = [] offset = None while True: points, offset = self._client.scroll( collection_name=self._collection_name, scroll_filter=scroll_filter, limit=256, offset=offset, with_payload=True, ) payloads.extend(dict(point.payload or {}) for point in points) if offset is None: break order = {key: index for index, key in enumerate(SECTION_ORDER)} payloads.sort( key=lambda p: ( order.get(p.get("section_key"), len(order)), p.get("part_index", 1 << 30), ) ) hits: list[SearchHit] = [] seen_sections: set[str] = set() for payload in payloads: section_key = payload.get("section_key") if section_key not in seen_sections: seen_sections.add(section_key) name = payload.get("section_display_name") or section_key or "" payload = {**payload, "text": f"【{name}】\n{payload.get('text', '')}"} hits.append(SearchHit(_document(payload), 1.0)) return hits def find_by_indication(self, indication_text: str, limit: int) -> list[SearchHit]: """Reverse lookup: every drug whose `chi_dinh` text mentions the given symptom/indication, keyword-matched. Deterministic, no fabrication risk — the same "exact match wins, no-match-means-None" philosophy `find_by_section` already uses, applied across drugs instead of within one. `limit` caps how many DRUGS are returned (one hit per drug, first match wins), not how many chunks are scanned — a common symptom can match far more drugs than is useful to show. Prose only: a `block_descriptor` chunk carries no real `chi_dinh` text (its text is built only from metadata per the quarantine contract), so keyword-matching it would be meaningless. `indication_text`, normalized, must appear as a CONTIGUOUS, word-boundary-anchored phrase in the chunk's text — not a plain substring (risks a false positive inside an unrelated longer word after diacritic-stripping) and not a scattered bag-of-words match either. Found live 2026-08-07: a token-SUBSET match (every word present *somewhere*, any order) let a long nonsense phrase built from common filler words ("bệnh chưa từng ghi nhận trong sách…") false-positive against real chi_dinh text, since words that common appear scattered through nearly everything — it reached generation before being caught, instead of failing here where it's cheap. A genuine paraphrase that doesn't share the book's exact wording is `search_indication`'s job (semantic), not this one's (lexical). """ import re from qdrant_client.models import FieldCondition, Filter, MatchValue from rag.text import normalize_name needle = normalize_name(indication_text) if not needle: return [] needle_pattern = re.compile(rf"(?:^| ){re.escape(needle)}(?:$| )") scroll_filter = Filter( must=[ FieldCondition(key="section_key", match=MatchValue(value="chi_dinh")), FieldCondition(key="chunk_kind", match=MatchValue(value="prose")), ] ) hits: list[SearchHit] = [] seen_drugs: set[str] = set() offset = None while True: points, offset = self._client.scroll( collection_name=self._collection_name, scroll_filter=scroll_filter, limit=256, offset=offset, with_payload=True, ) for point in points: payload = dict(point.payload or {}) drug_id = payload.get("drug_id") if drug_id in seen_drugs: continue text = normalize_name(payload.get("text", "")) if not needle_pattern.search(f" {text} "): continue seen_drugs.add(drug_id) hits.append(SearchHit(_document(payload), 1.0)) if len(hits) >= limit: return hits if offset is None: break return hits def search_indication(self, query: str, limit: int) -> list[SearchHit]: """Dense-vector fallback for `find_by_indication` when no exact keyword phrase match exists — catches paraphrases ("sốt cao" vs "thân nhiệt tăng") a literal phrase match cannot. Deliberately narrow (`section_key=chi_dinh` only, never the whole corpus) so this stays a targeted fallback for one specific gap, not a return to unranked similarity search — see ADR 0008 on why the live path otherwise avoids `search()`. One hit per drug, highest-scoring chunk kept (Qdrant returns points pre-sorted by score).""" from qdrant_client.models import FieldCondition, Filter, MatchValue vector = list(self._embedder.embed_query(query)) if len(vector) != self._embedder.dimensions: raise ValueError( f"query vector has {len(vector)} dimensions; " f"expected {self._embedder.dimensions}" ) points = self._client.search( collection_name=self._collection_name, query_vector=vector, query_filter=Filter( must=[ FieldCondition(key="section_key", match=MatchValue(value="chi_dinh")), FieldCondition(key="chunk_kind", match=MatchValue(value="prose")), ] ), limit=limit * 4, with_payload=True, ) hits: list[SearchHit] = [] seen_drugs: set[str] = set() for point in points: payload = dict(point.payload or {}) drug_id = payload.get("drug_id") if drug_id in seen_drugs: continue seen_drugs.add(drug_id) hits.append(SearchHit(_document(payload), float(point.score))) if len(hits) >= limit: break return hits class QdrantParentStore: def __init__(self, client: Any, collection_name: str) -> None: self._client = client self._collection_name = collection_name def get(self, parent_id: str) -> ParentDocument | None: from qdrant_client.models import FieldCondition, Filter, MatchValue points, _ = self._client.scroll( collection_name=self._collection_name, scroll_filter=Filter( must=[FieldCondition(key="chunk_id", match=MatchValue(value=parent_id))] ), limit=1, with_payload=True, ) if not points: return None payload = dict(points[0].payload or {}) return ParentDocument( parent_id=parent_id, kind=payload.get("chunk_kind", "parent"), text=payload["text"], source_refs=_source_refs(payload), requires_visual_check=( bool(payload.get("requires_visual_check")) or bool(payload.get("has_quarantined_content")) ), )