From 17dcdf41dc4f8666383ceb98c6e821666ee206c0 Mon Sep 17 00:00:00 2001 From: Eirik Botten Nicolaysen Date: Wed, 7 Oct 2026 14:42:49 +0200 Subject: [PATCH] Document the on_turn callback contract on_turn had a type in each signature and a three-line note in one docstring (AuditExperiment.run_scenario_reps), and nothing in the README. Which roles fire when, what turn_index and max_turns mean, and what happens on a failure were only recoverable by reading run_scenario and SingleTurnAuditor side by side. OnTurn in model_auditor.py now names the callback type and documents the contract as the code behaves today, measured on both runners: "auditor" only when the auditor wrote the probe, so not on a test_prompt turn; "target" per turn; "judge" once, at max_turns - 1, including on a parse-failure judgment; nothing from a call that raised, and nothing after it. It also records what the arguments do not carry: no scenario name, so calls interleave under max_workers > 1, and a retried rep in run_scenario_reps fires again. The signatures use the alias, the README parameter table gets an on_turn row, and the run_scenario_reps docstring points to the alias. No behaviour change: with docstrings stripped, the AST of the three changed modules differs only in the nine annotations, the alias and the imports that bring it in. One test runs ModelAuditor and SingleTurnAuditor and asserts the full sequence for each, including the target-down and judge-down paths. It passes on the unchanged code too, since it pins existing behaviour. --- README.md | 1 + simpleaudit/experiment.py | 10 ++++--- simpleaudit/model_auditor.py | 43 +++++++++++++++++++++++---- simpleaudit/single_turn.py | 8 +++--- tests/test_single_turn.py | 56 +++++++++++++++++++++++++++++++++++- 5 files changed, 104 insertions(+), 14 deletions(-) diff --git a/README.md b/README.md index a9b7e06..9257ce5 100644 --- a/README.md +++ b/README.md @@ -491,6 +491,7 @@ auditor = ModelAuditor( | `show_progress` | Show tqdm progress bars | No (default: `True`) | | `max_retries` | Retries per API call for transient failures | No (default: 2) | | `retry_backoff` | Initial retry delay in seconds, doubled per attempt (exponential backoff) | No (default: 0.5) | +| `on_turn` | Progress callback `on_turn(turn_index, max_turns, role)`, fired after each phase with `role` `"auditor"`, `"target"` or `"judge"`. `turn_index` is 0-based; the judge is reported once, at `max_turns - 1`, and not at all if a target or judge call raised. Also accepted per call by `run` / `run_async`, where it overrides this one. Full contract: `simpleaudit.model_auditor.OnTurn` | No | ## Scenario Packs diff --git a/simpleaudit/experiment.py b/simpleaudit/experiment.py index 6ee87dd..aafe680 100644 --- a/simpleaudit/experiment.py +++ b/simpleaudit/experiment.py @@ -10,7 +10,7 @@ from tqdm.auto import tqdm from simpleaudit._event_loop import run_sync from simpleaudit.results import AuditResult, AuditResults -from simpleaudit.model_auditor import ModelAuditor +from simpleaudit.model_auditor import ModelAuditor, OnTurn from simpleaudit.repeated_results import RepeatedExperimentResults @@ -328,7 +328,7 @@ async def _run_single_rep( max_turns: Optional[int], language: str, max_workers: int, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, ) -> AuditResults: @@ -367,7 +367,7 @@ async def run_scenario_reps( scenario: Dict[str, Any], max_turns: Optional[int] = None, language: str = "English", - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, ) -> List[AuditResult]: @@ -386,7 +386,9 @@ async def run_scenario_reps( on_turn: Optional callback fired at each phase boundary with ``(turn_index, max_turns, role)`` where role is "auditor", "target", or "judge". Called synchronously from within the - asyncio event loop. + asyncio event loop. The full contract, including what fires on + a failure and on a retried rep, is documented on + :data:`simpleaudit.model_auditor.OnTurn`. audit_run_id: Optional run id propagated to each rep's ``run_async`` for trace correlation. trace_correlation: Optional :class:`TraceCorrelation` shared diff --git a/simpleaudit/model_auditor.py b/simpleaudit/model_auditor.py index 041eb81..de1abe4 100644 --- a/simpleaudit/model_auditor.py +++ b/simpleaudit/model_auditor.py @@ -43,6 +43,38 @@ ) from .utils import parse_json_response as _parse_json_response +#: Progress callback: ``on_turn(turn_index, max_turns, role)``. +#: +#: Called synchronously, from inside the event loop, after each phase of a +#: scenario completes: +#: +#: - ``"auditor"``: the auditor model wrote a probe. Not fired for a turn whose +#: probe is the scenario's ``test_prompt`` (turn 0, when one is set), so a +#: turn reports either ``"auditor"`` then ``"target"``, or ``"target"`` alone. +#: - ``"target"``: the target answered. +#: - ``"judge"``: the judge call returned. Once per scenario, after the last +#: turn, reported at ``max_turns - 1``. It also fires when the judgment is a +#: parse-failure ERROR: it marks that the call returned, not that the verdict +#: is usable. +#: +#: ``turn_index`` is 0-based. ``max_turns`` is the turn count the run was given +#: (the per-call ``max_turns`` if passed), not the number of turns that ran, +#: and is assumed to be at least 1. ``SingleTurnAuditor`` reports its one +#: exchange as turn 0 of 1 whatever ``max_turns`` says. +#: +#: A phase that fails fires nothing, and nothing fires after it: a target error +#: ends the scenario's events with no ``"judge"``, and so does a judge call that +#: raises. The scenario's result is then ERROR. +#: +#: The arguments carry no scenario name. With ``max_workers > 1``, calls from +#: different scenarios interleave, and ``AuditExperiment.run_scenario_reps`` +#: fires a scenario's events again when it retries a rep that came back ERROR. +#: +#: The return value is ignored. An exception raised by the callback is logged +#: and swallowed. The callback must be a plain function: it is not awaited, +#: so an ``async def`` callback's body never runs. +OnTurn = Callable[[int, int, str], None] + def _user_agent() -> str: try: @@ -317,7 +349,7 @@ def __init__( target_params: Optional[Dict[str, Any]] = None, judge_params: Optional[Dict[str, Any]] = None, auditor_params: Optional[Dict[str, Any]] = None, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, ): if max_retries < 0: raise ValueError(f"max_retries must be >= 0, got {max_retries}") @@ -547,13 +579,14 @@ def _fire_on_turn( turn_index: int, max_turns: int, role: str, - callback: Optional[Callable[[int, int, str], None]] = None, + callback: Optional[OnTurn] = None, ) -> None: """Invoke an on_turn callback if set. Never raises. Uses ``callback`` when provided, otherwise falls back to the construction-time ``self.on_turn``. A raising callback is logged, not propagated — a progress observer must never fail an audit. + The contract callers can rely on is documented on ``OnTurn``. """ cb = callback if callback is not None else self.on_turn if cb is not None: @@ -872,7 +905,7 @@ async def run_scenario( target_params: Optional[Dict[str, Any]] = None, judge_params: Optional[Dict[str, Any]] = None, auditor_params: Optional[Dict[str, Any]] = None, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, evidence_spans: Optional[List[Dict[str, Any]]] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, @@ -1081,7 +1114,7 @@ async def run_async( target_params: Optional[Dict[str, Any]] = None, judge_params: Optional[Dict[str, Any]] = None, auditor_params: Optional[Dict[str, Any]] = None, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, evidence_spans: Optional[List[Dict[str, Any]]] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, @@ -1215,7 +1248,7 @@ def run( target_params: Optional[Dict[str, Any]] = None, judge_params: Optional[Dict[str, Any]] = None, auditor_params: Optional[Dict[str, Any]] = None, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, evidence_spans: Optional[List[Dict[str, Any]]] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, diff --git a/simpleaudit/single_turn.py b/simpleaudit/single_turn.py index ce377fc..7020139 100644 --- a/simpleaudit/single_turn.py +++ b/simpleaudit/single_turn.py @@ -38,7 +38,7 @@ import asyncio from datetime import date -from typing import Any, Callable, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union from tqdm.auto import tqdm @@ -47,7 +47,7 @@ from .context_findings import FINDING_SEVERITY, derive_findings from .context_marks import DocumentMark, parse_as_of, parse_documents, render_documents from .judges import get_judge -from .model_auditor import ModelAuditor +from .model_auditor import ModelAuditor, OnTurn from .results import AuditResult, AuditResults from .targets.base import TargetContext from .tracing.context import make_traceparent, new_trace_id @@ -264,7 +264,7 @@ async def _run_one_scenario( params: Optional[Dict[str, Any]] = None, target_params: Optional[Dict[str, Any]] = None, judge_params: Optional[Dict[str, Any]] = None, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, evidence_spans: Optional[List[Dict[str, Any]]] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, @@ -495,7 +495,7 @@ async def run_async( target_params: Optional[Dict[str, Any]] = None, judge_params: Optional[Dict[str, Any]] = None, auditor_params: Optional[Dict[str, Any]] = None, - on_turn: Optional[Callable[[int, int, str], None]] = None, + on_turn: Optional[OnTurn] = None, evidence_spans: Optional[List[Dict[str, Any]]] = None, audit_run_id: Optional[str] = None, trace_correlation: Optional[Any] = None, diff --git a/tests/test_single_turn.py b/tests/test_single_turn.py index c36b2c9..3e88c7e 100644 --- a/tests/test_single_turn.py +++ b/tests/test_single_turn.py @@ -25,7 +25,7 @@ from simpleaudit.results import AuditResults from simpleaudit.single_turn import SingleTurnAuditor, _build_judge_context -from .fakes import FakeClient +from .fakes import FakeClient, make_auditor JUDGE_JSON = json.dumps({ "severity": "pass", @@ -434,6 +434,60 @@ def test_the_exchange_carries_and_records_its_trace_context(self, wired): assert context.trace_headers["traceparent"].startswith(f"00-{trace_id}-") +# --- the on_turn contract --------------------------------------------------- + + +class TestOnTurnContract: + """The sequence documented on ``model_auditor.OnTurn``, pinned on both + runners at once so the two implementations cannot drift apart.""" + + def test_both_runners_report_the_documented_sequence(self): + def answer(**_): + return TARGET_ANSWER + + def down(**_): + raise RuntimeError("down") + + def multi(target, judge, turns): + return make_auditor( + target=FakeClient(target), judge=FakeClient(judge), + auditor=FakeClient(lambda **_: "Tell me more."), max_turns=turns, + ) + + def single(target, judge): + return make_single_turn_auditor( + target=FakeClient(target), judge=FakeClient(judge), max_turns=5, + ) + + def events(auditor, scenario): + seen = [] + asyncio.run(auditor.run_async([scenario], on_turn=lambda *e: seen.append(e))) + return seen + + def judged(**_): + return JUDGE_JSON + + probed = {"name": "probed", "description": "d"} + prompted = {"name": "prompted", "description": "d", "test_prompt": "Hei"} + + # Multi-turn without a test_prompt: the auditor writes every probe. + assert events(multi(answer, judged, 2), probed) == [ + (0, 2, "auditor"), (0, 2, "target"), + (1, 2, "auditor"), (1, 2, "target"), + (1, 2, "judge"), + ] + # A test_prompt replaces the turn-0 probe, so turn 0 has no "auditor". + # One such turn is exactly what single-turn reports, max_turns aside. + one_turn = [(0, 1, "target"), (0, 1, "judge")] + assert events(multi(answer, judged, 1), prompted) == one_turn + assert events(single(answer, judged), prompted) == one_turn + # A failing call fires nothing, and nothing fires after it. + assert events(multi(down, judged, 2), prompted) == [] + assert events(single(down, judged), prompted) == [] + assert events(multi(answer, down, 1), prompted) == [(0, 1, "target")] + assert events(single(answer, down), prompted) == [(0, 1, "target")] + + # --- token accounting -------------------------------------------------------