Files
duocthu/apps/ai-service/rag/routing.py
T

312 lines
12 KiB
Python

from __future__ import annotations
import functools
import re
from dataclasses import dataclass, replace
from difflib import SequenceMatcher
from enum import StrEnum
from .models import EvidenceDecision, QueryIntent, RetrievalResult, SubjectScope
from .service import RetrievalService
from .text import WORD_RE, normalize_name
__all__ = [
"WORD_RE",
"CatalogDrugResolver",
"DrugResolution",
"DrugResolutionStatus",
"QueryRoutingService",
"normalize_name",
]
class DrugResolutionStatus(StrEnum):
RESOLVED = "resolved"
NOT_FOUND = "not_found"
AMBIGUOUS = "ambiguous"
@dataclass(frozen=True)
class DrugResolution:
status: DrugResolutionStatus
drug_id: str | None = None
score: float | None = None
candidate_drug_ids: tuple[str, ...] = ()
class CatalogDrugResolver:
def __init__(
self,
catalog: dict[str, set[str]],
fuzzy_threshold: float = 0.84,
ambiguity_margin: float = 0.04,
) -> None:
self._catalog = {
drug_id: {
normalized for alias in aliases if (normalized := normalize_name(alias))
}
for drug_id, aliases in catalog.items()
}
self._aliases = [
(drug_id, normalized)
for drug_id, aliases in catalog.items()
for alias in aliases
if (normalized := normalize_name(alias))
]
self._fuzzy_threshold = fuzzy_threshold
self._ambiguity_margin = ambiguity_margin
# Measured live 2026-08-07: a single `resolve()` call over the real
# ~10,164-alias catalog costs ~0.65-0.7s, `suggest()` ~0.94-0.97s — both
# O(aliases) regex/SequenceMatcher work, pure functions of their
# arguments (only `self._aliases` et al, fixed at construction, feed
# them). `understanding.py`'s `_candidate_ids` calls both PER HISTORY
# LINE on every single turn — so the SAME already-seen history lines
# were being re-resolved from scratch every turn a conversation grew,
# ~1.6-1.7s of pure CPU per repeated line. A real user's ordinary
# multi-turn conversation was enough to exceed the 20s F-08 budget
# before the first Bedrock call ever ran, surfacing as a false
# "Dịch vụ đang gặp sự cố" — not a provider outage at all. Caching by
# exact input turns all but the newest turn's own text into a dict
# lookup on every subsequent call.
@functools.lru_cache(maxsize=4096)
def resolve(self, query: str) -> DrugResolution:
normalized_query = normalize_name(query)
query_tokens = normalized_query.split()
exact = [
(drug_id, alias, match.start(1), match.end(1))
for drug_id, alias in self._aliases
for match in [
re.search(rf"(?:^| )({re.escape(alias)})(?:$| )", normalized_query)
]
if match
]
if exact:
maximal = [
row for row in exact
if not any(
other[2] <= row[2] and row[3] <= other[3]
and (other[2], other[3]) != (row[2], row[3])
for other in exact
)
]
drug_ids = {drug_id for drug_id, _, _, _ in maximal}
if len(drug_ids) == 1:
return DrugResolution(
DrugResolutionStatus.RESOLVED, next(iter(drug_ids)), 1.0,
)
return DrugResolution(
DrugResolutionStatus.AMBIGUOUS,
candidate_drug_ids=tuple(sorted(drug_ids)),
)
scores: dict[str, float] = {}
for drug_id, alias in self._aliases:
width = len(alias.split())
if width > len(query_tokens):
continue
spans = (
" ".join(query_tokens[start:start + width])
for start in range(len(query_tokens) - width + 1)
)
score = max(
(SequenceMatcher(None, alias, span).ratio() for span in spans),
default=0.0,
)
scores[drug_id] = max(scores.get(drug_id, 0.0), score)
ranked = sorted(scores.items(), key=lambda item: (-item[1], item[0]))
if not ranked or ranked[0][1] < self._fuzzy_threshold:
return DrugResolution(DrugResolutionStatus.NOT_FOUND)
if len(ranked) > 1 and ranked[0][1] - ranked[1][1] < self._ambiguity_margin:
return DrugResolution(
DrugResolutionStatus.AMBIGUOUS,
candidate_drug_ids=(ranked[0][0], ranked[1][0]),
)
return DrugResolution(DrugResolutionStatus.RESOLVED, *ranked[0])
def aliases_for(self, drug_id: str) -> set[str]:
return self._catalog.get(drug_id, set())
def complete(self, prefix: str, k: int = 8) -> list[str]:
"""Drug ids whose alias contains `prefix`, for as-you-type autocomplete.
Substring match on the normalized alias, ranked prefix-first then by
alias length, so "para" surfaces "paracetamol" ahead of a drug that only
contains "para" mid-word. Distinct drug ids, best first.
"""
needle = normalize_name(prefix)
if not needle:
return []
matches: list[tuple[tuple[int, int], str]] = []
for drug_id, alias in self._aliases:
position = alias.find(needle)
if position < 0:
continue
matches.append(((0 if position == 0 else 1, len(alias)), drug_id))
matches.sort()
ordered: list[str] = []
seen: set[str] = set()
for _, drug_id in matches:
if drug_id in seen:
continue
seen.add(drug_id)
ordered.append(drug_id)
if len(ordered) >= k:
break
return ordered
# See the comment on `resolve` above — same cost, same fix, same
# single-caller read-only usage (safe to hand back a cached list).
@functools.lru_cache(maxsize=4096)
def suggest(
self, query: str, k: int = 3, min_score: float = 0.5
) -> list[tuple[str, float]]:
"""Closest drug ids by fuzzy score, for a 'did you mean' on a miss.
Uses the same windowed SequenceMatcher scoring as `resolve`, but returns
the top-k *below* the resolution threshold too, so a typo that does not
confidently resolve ("metfomin") can still be offered as a suggestion.
`min_score` keeps a genuinely non-drug query ("cái này thế nào") from
surfacing spurious suggestions.
"""
query_tokens = normalize_name(query).split()
if not query_tokens:
return []
scores: dict[str, float] = {}
for drug_id, alias in self._aliases:
width = len(alias.split())
if width > len(query_tokens):
continue
spans = (
" ".join(query_tokens[start:start + width])
for start in range(len(query_tokens) - width + 1)
)
score = max(
(SequenceMatcher(None, alias, span).ratio() for span in spans),
default=0.0,
)
scores[drug_id] = max(scores.get(drug_id, 0.0), score)
ranked = sorted(scores.items(), key=lambda item: (-item[1], item[0]))
return [(drug_id, score) for drug_id, score in ranked[:k] if score >= min_score]
class QueryRoutingService:
def __init__(
self,
retrieval: RetrievalService,
resolver: CatalogDrugResolver,
) -> None:
self._retrieval = retrieval
self._resolver = resolver
@staticmethod
def _scope_gate(
subject_scope: SubjectScope, intent: QueryIntent
) -> RetrievalResult | None:
# Scope comes from the API/policy layer. Unknown is deliberately
# fail-closed; retrieval must not infer clinical scope from keywords.
if subject_scope == SubjectScope.NON_HUMAN:
return RetrievalResult(EvidenceDecision.ABSTAIN, "out_of_scope_non_human")
if subject_scope == SubjectScope.UNKNOWN:
return RetrievalResult(EvidenceDecision.ABSTAIN, "subject_scope_unknown")
if intent == QueryIntent.RECOMMENDATION:
return RetrievalResult(EvidenceDecision.ABSTAIN, "recommendation_out_of_scope")
if intent == QueryIntent.UNKNOWN:
return RetrievalResult(EvidenceDecision.ABSTAIN, "query_intent_unknown")
return None
def retrieve_for_drug(
self,
query: str,
drug_id: str,
subject_scope: SubjectScope = SubjectScope.UNKNOWN,
intent: QueryIntent = QueryIntent.UNKNOWN,
) -> RetrievalResult:
"""Retrieve for an ALREADY-resolved drug, skipping name resolution.
The conversational layer has already resolved (and possibly inherited)
the drug; re-resolving from the rewritten turn text is what produced the
`drug_resolution_ambiguous` empty answers on follow-ups. The query text
still drives section routing and the intro/overview decision.
"""
gate = self._scope_gate(subject_scope, intent)
if gate is not None:
return gate
result = self._retrieval.retrieve(query, drug_id)
return replace(
result,
resolved_drug_id=drug_id,
drug_resolution_status=DrugResolutionStatus.RESOLVED,
)
def retrieve(
self,
query: str,
subject_scope: SubjectScope = SubjectScope.UNKNOWN,
intent: QueryIntent = QueryIntent.UNKNOWN,
) -> RetrievalResult:
gate = self._scope_gate(subject_scope, intent)
if gate is not None:
return gate
resolution = self._resolver.resolve(query)
if resolution.status == DrugResolutionStatus.NOT_FOUND:
return RetrievalResult(
EvidenceDecision.ABSTAIN,
"drug_not_resolved",
drug_resolution_status=resolution.status,
)
if resolution.status == DrugResolutionStatus.AMBIGUOUS:
disambiguated = self._disambiguate_with_evidence(
query, resolution.candidate_drug_ids,
)
if disambiguated is not None:
drug_id, result = disambiguated
return replace(
result,
resolved_drug_id=drug_id,
drug_resolution_status=DrugResolutionStatus.RESOLVED,
)
return RetrievalResult(
EvidenceDecision.ABSTAIN, "drug_resolution_ambiguous",
drug_resolution_status=resolution.status,
)
if resolution.drug_id is None:
return RetrievalResult(
EvidenceDecision.ABSTAIN,
"drug_resolution_invalid_state",
drug_resolution_status="invalid_state",
)
result = self._retrieval.retrieve(query, resolution.drug_id)
return replace(
result,
resolved_drug_id=resolution.drug_id,
drug_resolution_status=resolution.status,
)
def _disambiguate_with_evidence(
self,
query: str,
candidate_ids: tuple[str, ...],
) -> tuple[str, RetrievalResult] | None:
"""Resolve subject-vs-component ambiguity through asymmetric evidence.
A candidate wins only when its retrieved evidence explicitly contains
every other mentioned entity, while the reverse direction does not.
This keeps genuine multi-drug questions ambiguous.
"""
if len(candidate_ids) < 2:
return None
winners = []
for candidate_id in candidate_ids:
result = self._retrieval.retrieve(query, candidate_id)
if result.decision == EvidenceDecision.ABSTAIN:
continue
evidence_text = normalize_name(" ".join(item.text for item in result.evidence))
others = [item for item in candidate_ids if item != candidate_id]
if all(any(
re.search(rf"(?:^| ){re.escape(alias)}(?:$| )", evidence_text)
for alias in self._resolver.aliases_for(other)
) for other in others):
winners.append((candidate_id, result))
return winners[0] if len(winners) == 1 else None