Files
duocthu/apps/ai-service/tests/test_reasoning_loop.py
T

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"