"""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"