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

372 lines
14 KiB
Python

"""Conversation state, and the rules for carrying context across turns.
Pure domain. Everything here works without an LLM, which is deliberate: the
part of "understanding a follow-up" that matters clinically — *which drug is
this still about* — must be deterministic and testable, not inferred.
Two structures with different jobs:
`Focus` is structured and drives routing. It is what makes "còn trẻ em thì
sao?" resolvable at all.
`summary` is prose for the generator. It records **what was discussed**, never
clinical content: a dose restated from a summary carries no citation and could
not be grounding-verified, because that check compares against retrieved
evidence and a summary is not evidence.
"""
from __future__ import annotations
from dataclasses import dataclass, field, replace
from typing import Literal, Protocol
# A drug named six turns ago is not context, it is a hazard: conversations
# drift, and inheriting a stale drug produces a confident answer about the
# wrong medicine.
FOCUS_TTL_TURNS = 6
# Three exchanges kept verbatim; older turns are folded into the summary.
RECENT_TURNS = 6
Role = Literal["user", "assistant"]
Verbosity = Literal["concise", "detailed"]
@dataclass(frozen=True)
class Turn:
role: Role
text: str
at: str
drug_id: str | None = None
section_key: str | None = None
# Storing what answered a turn is what lets the planner reuse evidence
# instead of retrieving the same section again.
evidence_ids: tuple[str, ...] = ()
@dataclass(frozen=True)
class Focus:
"""The entities a follow-up may inherit, each with the turn that set it."""
drug_id: str | None = None
drug_name: str | None = None
section_key: str | None = None
population: str | None = None
verbosity: Verbosity | None = None
set_at_turn: dict[str, int] = field(default_factory=dict)
def age_of(self, name: str, turn_count: int) -> int | None:
set_at = self.set_at_turn.get(name)
return None if set_at is None else turn_count - set_at
def is_fresh(self, name: str, turn_count: int, ttl: int = FOCUS_TTL_TURNS) -> bool:
age = self.age_of(name, turn_count)
return age is not None and age <= ttl
def with_field(self, name: str, value, turn: int) -> "Focus":
stamps = dict(self.set_at_turn)
stamps[name] = turn
return replace(self, **{name: value}, set_at_turn=stamps)
def expire(self, turn_count: int, ttl: int = FOCUS_TTL_TURNS) -> "Focus":
"""Drops every field older than the TTL, stamps included."""
kept = {
name: getattr(self, name)
for name in ("drug_id", "drug_name", "section_key", "population", "verbosity")
if self.is_fresh(name, turn_count, ttl)
}
stamps = {
name: at for name, at in self.set_at_turn.items() if name in kept
}
return Focus(**kept, set_at_turn=stamps)
@dataclass(frozen=True)
class ConversationState:
conversation_id: str
recent: tuple[Turn, ...] = ()
summary: str = ""
focus: Focus = field(default_factory=Focus)
turn_count: int = 0
# Turns evicted from `recent` since the last time `overflow()` was
# consumed and cleared (by the caller passing `pending_overflow=()` to
# `replace()` after folding them into the summary). NOT derivable from
# `recent` alone — `recent` is already capped at `window`, so comparing
# its length against `window` can never find anything (see the bug note
# on `overflow` below). Plumbing, not conversation content.
pending_overflow: tuple[Turn, ...] = ()
def append(self, turn: Turn, window: int = RECENT_TURNS) -> "ConversationState":
"""Adds a turn and evicts the oldest beyond the window.
Eviction accumulates the dropped turns into `pending_overflow` for
the caller's summariser to fold via `overflow()`, rather than
discarding them here — this type does not decide what a summary
says. Accumulates rather than overwrites because one turn commonly
triggers two `append()` calls in a row (user, then assistant); each
can evict at most one turn, and the second call must not lose the
first's.
"""
combined = (*self.recent, turn)
recent = combined[-window:]
dropped = combined[:-window] if len(combined) > window else ()
return replace(
self,
recent=recent,
turn_count=self.turn_count + 1,
pending_overflow=(*self.pending_overflow, *dropped),
)
def overflow(self) -> tuple[Turn, ...]:
"""Turns evicted from `recent` and not yet folded into the summary.
Bug fixed 2026-08-06 (Codex review, F-06): this used to check
`len(self.recent) > window`, but `recent` is already truncated to
`window` by every `append()` call, so that comparison could never be
true — dropped turns were silently discarded and the summariser
never received them, no matter how long a conversation ran. The
caller must clear `pending_overflow` (pass `pending_overflow=()` to
`replace()`) after folding, or the same turns fold again next time.
"""
return self.pending_overflow
def inherited(self, name: str):
"""A focus value only if it is still fresh; otherwise None."""
return getattr(self.focus, name) if self.focus.is_fresh(name, self.turn_count) else None
# --- follow-up resolution -----------------------------------------------------
# Phrases that mean "same question, different population". Longest-first for the
# same reason `sections.py` sorts that way: "phụ nữ cho con bú" must be tested
# before "phụ nữ", or the more specific reading is never reached.
POPULATION_PHRASES: dict[str, str] = {
"phụ nữ cho con bú": "phu_nu_cho_con_bu",
"người cao tuổi": "nguoi_cao_tuoi",
"phụ nữ có thai": "phu_nu_co_thai",
"người suy thận": "suy_than",
"người suy gan": "suy_gan",
"trẻ sơ sinh": "tre_so_sinh",
"người lớn": "nguoi_lon",
"bà bầu": "phu_nu_co_thai",
"trẻ nhỏ": "tre_em",
"trẻ em": "tre_em",
"người già": "nguoi_cao_tuoi",
}
VERBOSITY_PHRASES: dict[str, Verbosity] = {
"giải thích kỹ hơn": "detailed",
"nói rõ hơn": "detailed",
"chi tiết hơn": "detailed",
"ngắn gọn": "concise",
"tóm tắt": "concise",
}
# A turn that is only a qualifier — no drug, no attribute — is a follow-up by
# construction. These are the openers that mark one.
FOLLOWUP_MARKERS = ("còn", "thế còn", "vậy còn", "so với", "thuốc vừa", "cái đó", "")
# Greetings, thanks, farewells and bare acknowledgements. A turn made up only of
# these is social, not a failed drug lookup: answering "Chưa xác định được
# thuốc" to "chào bạn" reads as broken. Longest-first so "cảm ơn nhiều" is
# stripped before "cảm ơn".
SMALLTALK_PHRASES = (
"xin chào", "chào bạn", "chào ad", "cảm ơn nhiều", "cảm ơn bạn", "cám ơn",
"cảm ơn", "tạm biệt", "hay quá", "tuyệt vời", "hiểu rồi", "được rồi",
"chào", "hello", "hi", "alo", "thanks", "thank", "ok", "oke", "okie",
"", "uh", "haha", "hihi", "bye",
)
def is_smalltalk(text: str) -> bool:
"""True when a turn carries nothing but social phrases.
Deliberately conservative: it strips every known social phrase and returns
True only if what remains is empty. "chào bạn, liều paracetamol?" keeps
"liều paracetamol" after stripping, so it is treated as a real question —
a greeting must never swallow the medical part of a turn.
"""
remainder = _normalise(text).strip(" .,!?;:")
for phrase in sorted(SMALLTALK_PHRASES, key=len, reverse=True):
# Space-pad both sides so a short phrase ("hi", "ok") matches a whole
# word only, never a substring of "chi" or "block".
remainder = f" {remainder} ".replace(f" {phrase} ", " ").strip(" .,!?;:")
return not remainder
def _normalise(text: str) -> str:
return " ".join(text.casefold().split())
def _longest_first(phrases: dict[str, str]) -> list[tuple[str, str]]:
return sorted(phrases.items(), key=lambda item: -len(item[0]))
def detect_population(text: str) -> str | None:
normalised = _normalise(text)
for phrase, tag in _longest_first(POPULATION_PHRASES):
if phrase in normalised:
return tag
return None
def detect_verbosity(text: str) -> Verbosity | None:
normalised = _normalise(text)
for phrase, level in _longest_first(VERBOSITY_PHRASES):
if phrase in normalised:
return level
return None
def looks_like_followup(text: str) -> bool:
normalised = _normalise(text)
return any(normalised.startswith(marker) for marker in FOLLOWUP_MARKERS)
@dataclass(frozen=True)
class ResolvedQuestion:
"""What this turn is asking, after the conversation is taken into account."""
text: str
drug_id: str | None
section_key: str | None
population: str | None
verbosity: Verbosity | None
inherited_drug: bool
inherited_section: bool
@property
def needs_carry_over_notice(self) -> bool:
"""Whether the answer must name what it inherited.
An inherited drug that is wrong is a wrong-drug answer, so the answer
has to say which drug it decided this was about.
"""
return self.inherited_drug
def resolve_against(
state: ConversationState,
text: str,
drug_id: str | None,
section_key: str | None,
) -> ResolvedQuestion:
"""Fills gaps in this turn from conversation focus, freshness permitting.
`drug_id` and `section_key` are what this turn resolved on its own — the
existing resolvers decide those, unchanged. Only what the turn left blank
is inherited, so an explicit mention always wins over context.
"""
inherited_drug = False
inherited_section = False
if drug_id is None:
carried = state.inherited("drug_id")
if carried is not None:
drug_id, inherited_drug = carried, True
if section_key is None:
carried = state.inherited("section_key")
if carried is not None:
section_key, inherited_section = carried, True
population = detect_population(text) or state.inherited("population")
verbosity = detect_verbosity(text) or state.inherited("verbosity")
return ResolvedQuestion(
text=text,
drug_id=drug_id,
section_key=section_key,
population=population,
verbosity=verbosity,
inherited_drug=inherited_drug,
inherited_section=inherited_section,
)
def update_focus(
state: ConversationState,
resolved: ResolvedQuestion,
) -> Focus:
"""Focus after this turn, stamped with the current turn index."""
focus = state.focus.expire(state.turn_count)
turn = state.turn_count
for name, value in (
("drug_id", resolved.drug_id),
("section_key", resolved.section_key),
("population", resolved.population),
("verbosity", resolved.verbosity),
):
if value is not None:
focus = focus.with_field(name, value, turn)
return focus
# --- persistence and summary --------------------------------------------------
#
# Protocol + no-LLM default co-located, matching how `reasoning.py` ships
# `SufficiencyAssessor`/`DeterministicAssessor` and `metrics.py` ships
# `Metrics`/`NullMetrics`. The Postgres-backed store lives in `adapters/`.
class ConversationStore(Protocol):
"""Loads and persists one conversation's state.
`load` returns a fresh empty state for an unknown id rather than raising: a
first turn has no prior state, and that is not an error.
"""
def load(self, conversation_id: str) -> "ConversationState": ...
def save(self, state: "ConversationState") -> None: ...
class InMemoryConversationStore:
"""Reference implementation and the offline/test default."""
def __init__(self) -> None:
self._states: dict[str, ConversationState] = {}
def load(self, conversation_id: str) -> ConversationState:
return self._states.get(conversation_id, ConversationState(conversation_id))
def save(self, state: ConversationState) -> None:
self._states[state.conversation_id] = state
class Summariser(Protocol):
"""Folds turns evicted from the recent window into rolling prose.
Contract, load-bearing for safety: the summary records *what was discussed*,
never a clinical value. A dose copied into a summary carries no citation and
cannot be grounding-verified — the check compares against retrieved
evidence, and a summary is not evidence.
"""
def fold(self, prev_summary: str, dropped: tuple["Turn", ...]) -> str: ...
class DeterministicSummariser:
"""No-LLM default: one topic line per evicted user turn, capped.
Records only the drug and section a turn was *about* — labels, never cell
values — so the no-clinical-content rule holds by construction rather than
by trusting a generator not to leak a dose.
"""
MAX_CHARS = 1600 # ~400 tokens, per ADR 0007 §2
def fold(self, prev_summary: str, dropped: tuple[Turn, ...]) -> str:
lines = [prev_summary] if prev_summary else []
for turn in dropped:
if turn.role != "user":
continue
drug = turn.drug_id or "thuốc chưa xác định"
section = turn.section_key or "thông tin chung"
lines.append(f"- đã hỏi {section} của {drug}")
text = "\n".join(lines)
# Keep the most recent topics when over budget: drop oldest lines, not
# mid-line characters, so the summary never ends on a fragment.
while len(text) > self.MAX_CHARS and len(lines) > 1:
lines.pop(0)
text = "\n".join(lines)
return text