Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
10 changes: 6 additions & 4 deletions simpleaudit/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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]:
Expand All @@ -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
Expand Down
43 changes: 38 additions & 5 deletions simpleaudit/model_auditor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
8 changes: 4 additions & 4 deletions simpleaudit/single_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
56 changes: 55 additions & 1 deletion tests/test_single_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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 -------------------------------------------------------


Expand Down
Loading