444 lines
18 KiB
Python
444 lines
18 KiB
Python
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"))
|
|
),
|
|
)
|