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