Wire the guarded conversational RAG answer layer end-to-end
This commit is contained in:
@@ -0,0 +1,81 @@
|
||||
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"
|
||||
Reference in New Issue
Block a user