245 lines
7.3 KiB
Python
245 lines
7.3 KiB
Python
"""The loop must improve answers, and must be unable to run away.
|
|
|
|
Bounded is the load-bearing property: an unbounded self-improvement loop on a
|
|
paid provider is a bill and a latency incident, and on a clinical tool it is
|
|
also an answer nobody is waiting for any more.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
|
|
from rag.conversation import ConversationState, ResolvedQuestion
|
|
from rag.metrics import CLARIFY_ASKED, LOOP_REFINED, InMemoryMetrics
|
|
from rag.reasoning import (
|
|
ClarifyReason,
|
|
DeterministicAssessor,
|
|
LoopTrace,
|
|
Sufficiency,
|
|
TurnBudget,
|
|
run_turn,
|
|
)
|
|
|
|
ADULT = "Người lớn: uống 0,5 - 1 g/lần, cách 4 - 6 giờ; tối đa 4 g/ngày."
|
|
CHILD = "Trẻ em 6 - 12 tuổi: 240 - 250 mg mỗi lần."
|
|
|
|
|
|
def _q(text: str = "liều dùng paracetamol", population: str | None = None) -> ResolvedQuestion:
|
|
return ResolvedQuestion(
|
|
text=text,
|
|
drug_id="paracetamol",
|
|
section_key="lieu_luong_va_cach_dung",
|
|
population=population,
|
|
verbosity=None,
|
|
inherited_drug=False,
|
|
inherited_section=False,
|
|
)
|
|
|
|
|
|
def _state() -> ConversationState:
|
|
return ConversationState("c1", turn_count=1)
|
|
|
|
|
|
class _Retriever:
|
|
"""Returns a different evidence set on each round, recording calls."""
|
|
|
|
def __init__(self, *rounds: tuple[str, ...]) -> None:
|
|
self._rounds = list(rounds)
|
|
self.queries: list[str] = []
|
|
|
|
def __call__(self, resolved: ResolvedQuestion) -> tuple[str, ...]:
|
|
self.queries.append(resolved.text)
|
|
if self._rounds:
|
|
return self._rounds.pop(0)
|
|
return ()
|
|
|
|
|
|
def _generator(answer: str | None):
|
|
calls = {"n": 0}
|
|
|
|
def generate(resolved, evidence, state):
|
|
calls["n"] += 1
|
|
return answer
|
|
|
|
generate.calls = calls # type: ignore[attr-defined]
|
|
return generate
|
|
|
|
|
|
# --- the loop earns its rounds ------------------------------------------------
|
|
|
|
|
|
def test_a_named_gap_buys_exactly_one_more_round():
|
|
"""Asked for adults, first round returned only paediatric text."""
|
|
retriever = _Retriever((CHILD,), (ADULT, CHILD))
|
|
metrics = InMemoryMetrics()
|
|
|
|
outcome = run_turn(
|
|
_state(),
|
|
_q(population="nguoi_lon"),
|
|
retriever,
|
|
_generator("Người lớn: 0,5 - 1 g/lần [1]"),
|
|
metrics=metrics,
|
|
)
|
|
|
|
assert outcome.retrieval_rounds_used == 2
|
|
assert outcome.generated is True
|
|
assert metrics.total(LOOP_REFINED, missing="population:nguoi_lon") == 1
|
|
assert retriever.queries[1] != retriever.queries[0]
|
|
|
|
|
|
def test_a_satisfied_question_spends_one_round_only():
|
|
retriever = _Retriever((ADULT,))
|
|
|
|
outcome = run_turn(
|
|
_state(), _q(population="nguoi_lon"), retriever, _generator("ok [1]")
|
|
)
|
|
|
|
assert outcome.retrieval_rounds_used == 1
|
|
assert outcome.stopped_because == "sufficient"
|
|
|
|
|
|
def test_a_simple_question_does_not_loop():
|
|
"""No population asked for means nothing to be missing."""
|
|
retriever = _Retriever((ADULT, CHILD))
|
|
|
|
outcome = run_turn(_state(), _q(), retriever, _generator("ok [1]"))
|
|
|
|
assert outcome.retrieval_rounds_used == 1
|
|
|
|
|
|
# --- the loop cannot run away -------------------------------------------------
|
|
|
|
|
|
def test_retrieval_rounds_are_hard_capped():
|
|
"""Evidence never satisfies the assessor; the loop must still stop."""
|
|
retriever = _Retriever((CHILD,), (CHILD,), (CHILD,), (CHILD,), (CHILD,))
|
|
|
|
outcome = run_turn(
|
|
_state(),
|
|
_q(population="nguoi_lon"),
|
|
retriever,
|
|
_generator("ok [1]"),
|
|
budget=TurnBudget(retrieval_rounds=2),
|
|
)
|
|
|
|
assert outcome.retrieval_rounds_used == 2
|
|
assert outcome.stopped_because == "retrieval_budget"
|
|
assert len(retriever.queries) == 2
|
|
|
|
|
|
def test_repairs_are_hard_capped_and_degrade_to_no_answer():
|
|
"""`generate` returning None means verification refused it every time."""
|
|
generate = _generator(None)
|
|
|
|
outcome = run_turn(
|
|
_state(),
|
|
_q(),
|
|
_Retriever((ADULT,)),
|
|
generate,
|
|
budget=TurnBudget(repairs=1, llm_calls=4),
|
|
)
|
|
|
|
assert outcome.answer is None
|
|
assert outcome.repairs_used == 1
|
|
assert generate.calls["n"] == 2 # first attempt + one repair
|
|
assert outcome.stopped_because == "repair_budget"
|
|
|
|
|
|
def test_llm_call_budget_stops_generation_entirely():
|
|
generate = _generator(None)
|
|
|
|
outcome = run_turn(
|
|
_state(), _q(), _Retriever((ADULT,)), generate, budget=TurnBudget(llm_calls=0)
|
|
)
|
|
|
|
assert generate.calls["n"] == 0
|
|
assert outcome.stopped_because == "llm_budget"
|
|
|
|
|
|
def test_a_refinement_that_changes_nothing_stops_the_loop():
|
|
"""Guards against a loop that keeps re-issuing the same query."""
|
|
|
|
class _SameQuery:
|
|
def assess(self, resolved, evidence):
|
|
return Sufficiency(False, missing="x", refined_query=resolved.text)
|
|
|
|
retriever = _Retriever((CHILD,), (CHILD,))
|
|
|
|
outcome = run_turn(
|
|
_state(), _q(), retriever, _generator("ok [1]"), assessor=_SameQuery()
|
|
)
|
|
|
|
assert outcome.stopped_because == "query_unchanged"
|
|
assert len(retriever.queries) == 1
|
|
|
|
|
|
def test_an_unnamed_gap_does_not_buy_a_round():
|
|
""""Feels incomplete" is not a reason to spend the budget."""
|
|
|
|
class _Vague:
|
|
def assess(self, resolved, evidence):
|
|
return Sufficiency(False)
|
|
|
|
retriever = _Retriever((CHILD,), (CHILD,))
|
|
|
|
outcome = run_turn(_state(), _q(), retriever, _generator("ok [1]"), assessor=_Vague())
|
|
|
|
assert outcome.stopped_because == "no_actionable_gap"
|
|
assert len(retriever.queries) == 1
|
|
|
|
|
|
# --- clarify beats guessing ---------------------------------------------------
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"signal",
|
|
[ClarifyReason.NO_ATTRIBUTE, ClarifyReason.AMBIGUOUS_DRUG, ClarifyReason.MULTI_ATTRIBUTE],
|
|
)
|
|
def test_a_clarify_signal_short_circuits_before_any_spend(signal):
|
|
retriever = _Retriever((ADULT,))
|
|
generate = _generator("ok [1]")
|
|
metrics = InMemoryMetrics()
|
|
budget = TurnBudget()
|
|
|
|
outcome = run_turn(
|
|
_state(), _q(), retriever, generate, clarify_signals=(signal,), budget=budget, metrics=metrics
|
|
)
|
|
|
|
assert outcome.clarification is not None
|
|
assert outcome.clarification.reason == signal
|
|
assert outcome.answer is None
|
|
assert retriever.queries == []
|
|
assert generate.calls["n"] == 0
|
|
assert budget.llm_calls == 4 and budget.retrieval_rounds == 2
|
|
assert metrics.total(CLARIFY_ASKED, reason=signal) == 1
|
|
|
|
|
|
def test_no_evidence_at_all_asks_rather_than_abstaining_silently():
|
|
outcome = run_turn(_state(), _q(), _Retriever(()), _generator("ok [1]"))
|
|
|
|
assert outcome.clarification is not None
|
|
assert outcome.clarification.reason == ClarifyReason.STILL_INSUFFICIENT
|
|
assert outcome.stopped_because == "no_evidence"
|
|
|
|
|
|
# --- the deterministic assessor ----------------------------------------------
|
|
|
|
|
|
def test_assessor_only_reports_gaps_it_can_demonstrate():
|
|
assessor = DeterministicAssessor()
|
|
|
|
assert assessor.assess(_q(population="nguoi_lon"), (ADULT,)).sufficient is True
|
|
assert assessor.assess(_q(population="nguoi_lon"), (CHILD,)).sufficient is False
|
|
# No population asked for: nothing can be shown missing.
|
|
assert assessor.assess(_q(), (CHILD,)).sufficient is True
|
|
|
|
|
|
def test_trace_records_the_stages_walked():
|
|
trace = LoopTrace()
|
|
|
|
run_turn(_state(), _q(), _Retriever((ADULT,)), _generator("ok [1]"), trace=trace)
|
|
|
|
assert trace.stages[0] == "understand"
|
|
assert "retrieve" in trace.stages
|
|
assert "assess" in trace.stages
|
|
assert trace.stages[-1] == "generate"
|