82 lines
2.9 KiB
Python
82 lines
2.9 KiB
Python
from rag.conversation import DeterministicSummariser, InMemoryConversationStore
|
|
from rag.conversational import ConversationalRagService, TurnResolution
|
|
from rag.reasoning import ClarifyReason, MAX_LLM_CALLS, MAX_RETRIEVAL_ROUNDS, TurnBudget
|
|
|
|
|
|
class FakeResolver:
|
|
"""Maps a turn's text to what it resolves on its own (no context)."""
|
|
|
|
def __init__(self, table):
|
|
self._table = table
|
|
|
|
def resolve_turn(self, text):
|
|
for needle, resolution in self._table:
|
|
if needle in text:
|
|
return resolution
|
|
return TurnResolution(drug_id=None, section_key=None, drug_status="not_found")
|
|
|
|
|
|
def _service(resolver, retrieve, generate):
|
|
return ConversationalRagService(
|
|
store=InMemoryConversationStore(),
|
|
summariser=DeterministicSummariser(),
|
|
resolver=resolver,
|
|
retrieve=retrieve,
|
|
generate=generate,
|
|
)
|
|
|
|
|
|
def test_followup_inherits_drug_and_answer_names_it():
|
|
resolver = FakeResolver([
|
|
("metformin", TurnResolution("metformin", "chong_chi_dinh", "resolved")),
|
|
# "còn trẻ em" names no drug on its own — must inherit.
|
|
("trẻ em", TurnResolution(None, None, "not_found")),
|
|
])
|
|
# Evidence mentions "trẻ em" so the population assessor is satisfied.
|
|
retrieve = lambda q: ("Ở trẻ em, liều metformin điều chỉnh theo cân nặng.",)
|
|
generate = lambda q, ev, st: "liều theo cân nặng"
|
|
svc = _service(resolver, retrieve, generate)
|
|
|
|
first = svc.answer("c1", "Chống chỉ định của metformin?")
|
|
assert first.inherited_drug is None
|
|
|
|
second = svc.answer("c1", "còn trẻ em thì sao?")
|
|
assert second.inherited_drug == "metformin"
|
|
assert second.answer.startswith("Về metformin:")
|
|
|
|
|
|
def test_no_drug_and_no_context_asks_without_spending_budget():
|
|
resolver = FakeResolver([]) # nothing resolves
|
|
calls = {"retrieve": 0, "generate": 0}
|
|
|
|
def retrieve(q):
|
|
calls["retrieve"] += 1
|
|
return ("x",)
|
|
|
|
def generate(q, ev, st):
|
|
calls["generate"] += 1
|
|
return "x"
|
|
|
|
svc = _service(resolver, retrieve, generate)
|
|
budget = TurnBudget()
|
|
out = svc.answer("c2", "cái này thế nào?", budget=budget)
|
|
|
|
assert out.answer is None
|
|
assert out.clarification is not None
|
|
assert out.clarification.reason == ClarifyReason.AMBIGUOUS_DRUG
|
|
# Asking short-circuits before any spend.
|
|
assert calls == {"retrieve": 0, "generate": 0}
|
|
assert budget.retrieval_rounds == MAX_RETRIEVAL_ROUNDS
|
|
assert budget.llm_calls == MAX_LLM_CALLS
|
|
|
|
|
|
def test_state_persists_across_turns():
|
|
resolver = FakeResolver([
|
|
("metformin", TurnResolution("metformin", "chi_dinh", "resolved")),
|
|
])
|
|
svc = _service(resolver, lambda q: ("Chỉ định của metformin.",), lambda q, ev, st: "ok")
|
|
svc.answer("c3", "chỉ định metformin?")
|
|
state = svc._store.load("c3")
|
|
assert state.turn_count == 2 # user + assistant
|
|
assert state.focus.drug_id == "metformin"
|