from __future__ import annotations from dataclasses import dataclass from .models import Evidence, EvidenceDecision, RetrievalResult, SearchHit from .ports import ( ParentStore, QueryEmbeddingUnavailable, Reranker, RerankUnavailable, Retriever, ) from .sections import SectionResolver # The sections that introduce a drug: what it is, its class, its main use, its # mechanism — in book order. A bare drug name is answered from these, not from # the dosage-forms table that happens to sit near the top of the monograph. INTRO_SECTIONS = ( "ten_chung_quoc_te", "loai_thuoc", "chi_dinh", "duoc_ly_va_co_che_tac_dung", ) @dataclass(frozen=True) class EvidencePolicy: minimum_score: float = 0.12 candidate_limit: int = 5 evidence_limit: int = 3 # A free-form question about a resolved drug otherwise hands the LLM the # entire monograph; rerank trims it to the sections that actually answer. rerank_top_k: int = 6 class RetrievalService: """Section-filtered retrieval when the question names its attribute. Similarity is the fallback, not the default. Measured 2026-08-04, letting similarity choose the section answers "chống chỉ định" correctly 1 time in 20, because the largest section (`duoc_ly_va_co_che_tac_dung`) sits close to any question about the drug. When the question says which section it wants, filtering answers it exactly. """ def __init__( self, retriever: Retriever, parent_store: ParentStore, policy: EvidencePolicy | None = None, section_resolver: SectionResolver | None = None, reranker: Reranker | None = None, ) -> None: self._retriever = retriever self._parent_store = parent_store self._policy = policy or EvidencePolicy() self._section_resolver = section_resolver self._reranker = reranker def retrieve(self, query: str, drug_id: str) -> RetrievalResult: if not query.strip() or not drug_id.strip(): return RetrievalResult(EvidenceDecision.ABSTAIN, "missing_query_or_drug") section_hits = self._section_hits(query, drug_id) if section_hits is not None: # No `evidence_limit` here: the whole section is the answer, and a # truncated list of contraindications reads as a complete one. return self._decide(self._hydrate(section_hits, limit=None)) # Drug resolved but no attribute named. A *bare* drug name ("PARACETAMOL") # shows the whole monograph in book order. A *free-form question* about # the drug ("sốt cao uống được không?") would otherwise dump all ~29 # sections at the model; rerank keeps only the sections that answer it. overview_hits = self._drug_overview(drug_id) if overview_hits is not None: if self._is_question(query): overview_hits = self._rerank(query, overview_hits) return self._decide(self._hydrate(overview_hits, limit=None)) # A bare drug name is not a question — introduce the drug from its # identity sections (what it is, its class, its main indication), # not the whole monograph starting with the dosage-forms table. intro = [h for h in overview_hits if h.document.section_key in INTRO_SECTIONS] return self._decide( self._hydrate(intro or overview_hits, limit=None), is_drug_overview=True, ) try: hits = self._retriever.search( query=query, drug_id=drug_id, limit=self._policy.candidate_limit, ) except QueryEmbeddingUnavailable: # Fail closed. The section route needs no embedder, so this only # ever narrows the fallback: the caller is told nothing was found # rather than being shown an error page or, worse, an answer built # from a search that never ran. return RetrievalResult( EvidenceDecision.ABSTAIN, "query_embedding_unavailable" ) if not hits or hits[0].score < self._policy.minimum_score: return RetrievalResult(EvidenceDecision.ABSTAIN, "insufficient_retrieval_score") return self._decide(self._hydrate(self._rerank(query, hits))) @staticmethod def _is_question(query: str) -> bool: """A bare drug name (one or two tokens) wants the whole monograph; more than that is a question whose overview should be reranked to the point.""" return len(query.split()) > 2 def _rerank(self, query: str, hits: list[SearchHit]) -> list[SearchHit]: """Reorder hits by cross-encoder relevance, keep the top-k. Fail-open: no reranker configured, or the provider is unreachable, and the original order is returned unchanged — an ordering aid must never be able to lose an answer. The section route never reaches this. """ if self._reranker is None or len(hits) <= 1: return hits try: order = self._reranker.rerank( query, [hit.document.text for hit in hits], top_n=self._policy.rerank_top_k, ) except RerankUnavailable: return hits reranked = [hits[index] for index in order if 0 <= index < len(hits)] return reranked[: self._policy.rerank_top_k] or hits def _drug_overview(self, drug_id: str) -> list[SearchHit] | None: """Every prose section of the drug, or None if the store cannot scroll.""" find_by_drug = getattr(self._retriever, "find_by_drug", None) if find_by_drug is None: return None hits = find_by_drug(drug_id) return hits or None def _section_hits(self, query: str, drug_id: str) -> list[SearchHit] | None: """Hits for an explicitly named section, or None to fall back. Returns None — not an empty list — when this route does not apply, so "no section named" stays distinguishable from "section named but empty". """ if self._section_resolver is None: return None find_by_section = getattr(self._retriever, "find_by_section", None) if find_by_section is None: return None match = self._section_resolver.resolve(query) if match is None: return None hits = find_by_section(drug_id, match.section_key) return hits or None def _decide( self, evidence: tuple[Evidence, ...], is_drug_overview: bool = False ) -> RetrievalResult: if not evidence: return RetrievalResult(EvidenceDecision.ABSTAIN, "parent_hydration_failed") if any(not item.source_refs for item in evidence): return RetrievalResult(EvidenceDecision.ABSTAIN, "missing_provenance") if any(item.requires_visual_check for item in evidence): return RetrievalResult(EvidenceDecision.VERIFY_PDF, "visual_verification_required", evidence) return RetrievalResult( EvidenceDecision.ANSWERABLE, "grounded_evidence_available", evidence, is_drug_overview=is_drug_overview, ) def _hydrate( self, hits: list[SearchHit], limit: int | None = -1 ) -> tuple[Evidence, ...]: output: list[Evidence] = [] seen: set[str] = set() for hit in hits: document = hit.document evidence_id = document.parent_id or document.doc_id if evidence_id in seen: continue seen.add(evidence_id) if document.parent_id: parent = self._parent_store.get(document.parent_id) if parent is None: continue output.append(Evidence( evidence_id=parent.parent_id, matched_doc_id=document.doc_id, kind=parent.kind, text=parent.text, score=hit.score, source_refs=parent.source_refs, hydrated_from_parent=True, requires_visual_check=( document.requires_visual_check or parent.requires_visual_check ), )) else: output.append(Evidence( evidence_id=document.doc_id, matched_doc_id=document.doc_id, kind=document.kind, text=document.text, score=hit.score, source_refs=document.source_refs, hydrated_from_parent=False, requires_visual_check=document.requires_visual_check, )) cap = self._policy.evidence_limit if limit == -1 else limit if cap is not None and len(output) >= cap: break return tuple(output)