Skip to content
Draft
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
44 changes: 31 additions & 13 deletions build_scripts/export_adversarial_benchmark_result.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,15 +5,12 @@

import argparse
import asyncio
import contextlib
import csv
import json
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any

from pyrit.cli._output import print_attacks_table
from pyrit.cli._results import build_attacks_table_payload
from pyrit.memory import CentralMemory
from pyrit.models import ScenarioResult
from pyrit.output.scenario_result.pretty import PrettyScenarioResultMemoryPrinter
Expand Down Expand Up @@ -46,16 +43,37 @@ async def _write_overview_async(*, result: ScenarioResult, output_dir: Path) ->
await printer.write_async(result)


def _write_attacks(*, result: ScenarioResult, output_dir: Path) -> None:
def _attack_rows(*, result: ScenarioResult) -> list[dict[str, Any]]:
"""Build machine-readable per-attack rows from the embedded attack results."""
rows: list[dict[str, Any]] = []
for atomic_attack_name, attacks in result.attack_results.items():
for attack in attacks:
score = attack.last_score
score_value = None
if score is not None:
score_value = score.score_value if score.score_value is not None else score.status.value
rows.append(
{
"attack_result_id": attack.attack_result_id,
"atomic_attack_name": atomic_attack_name,
"objective": attack.objective,
"outcome": attack.outcome.value,
"executed_turns": attack.executed_turns,
"score_value": score_value,
}
)
return rows


async def _write_attacks_async(*, result: ScenarioResult, output_dir: Path) -> None:
"""Write machine-readable and console-style partial attack tables."""
payload = build_attacks_table_payload(
result=result,
scenario_result_id=str(result.id),
)
(output_dir / "attacks.json").write_text(payload.model_dump_json(indent=2), encoding="utf-8")
with open(output_dir / "attacks.txt", "w", encoding="utf-8") as output:
with contextlib.redirect_stdout(output):
print_attacks_table(payload=payload)
rows = _attack_rows(result=result)
document = {"scenario_result_id": str(result.id), "rows": rows, "total": len(rows)}
(output_dir / "attacks.json").write_text(json.dumps(document, indent=2), encoding="utf-8")

sink = FileSink(path=output_dir / "attacks.txt")
printer = PrettyScenarioResultMemoryPrinter(sink=sink, enable_colors=False)
await printer.write_async(result, view="attacks")


def _build_technique_metrics(*, result: ScenarioResult) -> list[dict[str, Any]]:
Expand Down Expand Up @@ -143,7 +161,7 @@ async def _export_async(*, scenario_result_id: str, output_dir: Path) -> None:
result = await _load_result_async(scenario_result_id=scenario_result_id)
output_dir.mkdir(parents=True, exist_ok=True)
await _write_overview_async(result=result, output_dir=output_dir)
await asyncio.to_thread(_write_attacks, result=result, output_dir=output_dir)
await _write_attacks_async(result=result, output_dir=output_dir)
await asyncio.to_thread(_write_technique_metrics, result=result, output_dir=output_dir)


Expand Down
127 changes: 50 additions & 77 deletions pyrit/cli/_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from typing import TYPE_CHECKING, Any

if TYPE_CHECKING:
from pyrit.cli._results import AttacksTablePayload, ConversationsPayload, TranscriptMessage
from pyrit.cli.api_client import PyRITApiClient
from pyrit.models import ScenarioResult
from pyrit.models.catalog import (
RegisteredInitializer,
Expand Down Expand Up @@ -425,94 +425,67 @@ async def print_scenario_result_async(*, result: ScenarioResult) -> None:
"undetermined": None,
}

# Per-role transcript colors, mirroring PrettyConversationPrinter's palette so the
# thin-client transcript reads like the framework's own conversation output.
_ROLE_COLORS = {
"user": "blue",
"assistant": "yellow",
"system": "magenta",
}


def print_attacks_table(*, payload: AttacksTablePayload) -> None:
"""
Print the per-attack table for a scenario run.

Args:
payload (AttacksTablePayload): The rows to render plus the pre-limit total.
async def print_conversations_async(
*,
result: ScenarioResult,
client: PyRITApiClient,
scenario_result_id: str,
attack_result_ids: list[str] | None = None,
limit: int | None = None,
) -> None:
"""
if not payload.rows:
print(f"\nNo attack results found for scenario {payload.scenario_result_id}.")
return

_header(f"Attack Results — scenario {payload.scenario_result_id}")
for index, row in enumerate(payload.rows, start=1):
outcome = row.outcome.upper()
score = row.score_value if row.score_value is not None else "—"
_cprint(
f" {index}. [{outcome}] turns={row.executed_turns} score={score}",
color=_OUTCOME_COLORS.get(row.outcome),
bold=True,
)
print(f" id: {row.attack_result_id}")
print(f" technique: {row.atomic_attack_name}")
print(f" objective: {row.objective}")

shown = len(payload.rows)
if shown < payload.total:
print(f"\nShowing {shown} of {payload.total} attacks (use --limit to change).")
else:
print(f"\nTotal attacks: {payload.total}")

Print each attack's main-conversation transcript, rendered by the framework.

def print_conversations(*, payload: ConversationsPayload) -> None:
"""
Print the per-attack main-conversation transcripts for a scenario run.
Reuses ``pyrit.output``'s conversation printer via a REST-backed source, so the
CLI transcript matches the framework's own conversation output. The per-attack
fetch loop is gated by *limit* (network calls, not just rendered rows).

Args:
payload (ConversationsPayload): The transcripts to render plus the
pre-limit total.
"""
if not payload.conversations:
print(f"\nNo conversations found for scenario {payload.scenario_result_id}.")
result (ScenarioResult): The scenario result whose attacks to inspect.
client (PyRITApiClient): Client used to fetch each conversation's messages.
scenario_result_id (str): The run id, echoed in the header.
attack_result_ids (list[str] | None): Restrict to these attack ids. Defaults to None.
limit (int | None): Maximum number of attacks to fetch and render. Defaults to None.
"""
from pyrit.cli._results import _objective_scorer_key, _select_attacks
from pyrit.cli._sources import RestApiConversationSource
from pyrit.output.conversation.pretty import PrettyConversationPrinter

selected = _select_attacks(result=result, attack_result_ids=attack_result_ids)
total = len(selected)
if limit is not None:
selected = selected[:limit]

if not selected:
print(f"\nNo conversations found for scenario {scenario_result_id}.")
return

_header(f"Conversations — scenario {payload.scenario_result_id}")
for index, convo in enumerate(payload.conversations, start=1):
objective_hash, objective_class = _objective_scorer_key(result=result)
_header(f"Conversations — scenario {scenario_result_id}")
for index, (atomic_attack_name, attack_result) in enumerate(selected, start=1):
_cprint(
f" {index}. [{convo.outcome.upper()}] {convo.atomic_attack_name}",
color=_OUTCOME_COLORS.get(convo.outcome),
f" {index}. [{attack_result.outcome.value.upper()}] {atomic_attack_name}",
color=_OUTCOME_COLORS.get(attack_result.outcome.value),
bold=True,
)
print(f" id: {convo.attack_result_id}")
print(f" objective: {convo.objective}")
_print_transcript(messages=convo.messages)
print(f" id: {attack_result.attack_result_id}")
print(f" objective: {attack_result.objective}")
source = RestApiConversationSource(
client=client,
attack_result_id=attack_result.attack_result_id,
objective_hash=objective_hash,
objective_class=objective_class,
)
messages = await source.get_messages_async(conversation_id=attack_result.conversation_id)
printer = PrettyConversationPrinter(source=source)
print(await printer.render_async(messages, include_scores=True))

shown = len(payload.conversations)
if shown < payload.total:
print(f"\nShowing {shown} of {payload.total} attacks (use --limit or --attack-result-ids to change).")
shown = len(selected)
if shown < total:
print(f"\nShowing {shown} of {total} attacks (use --limit or --attack-result-ids to change).")
else:
print(f"\nTotal attacks: {payload.total}")


def _print_transcript(*, messages: list[TranscriptMessage]) -> None:
"""Print one attack's ordered messages with their optional scores."""
if not messages:
print(" (no messages)")
return
for message in messages:
_cprint(
f" [{message.role.upper()}] (turn {message.turn})",
color=_ROLE_COLORS.get(message.role.lower()),
bold=True,
)
print(_wrap(text=message.text, indent=" "))
if message.score is not None:
value = message.score.value if message.score.value is not None else "—"
label = f"SCORE [{message.score.scorer}]" if message.score.scorer else "SCORE"
_cprint(f" {label}: {value}", color="magenta", bold=True)
if message.score.rationale:
print(_wrap(text=f"rationale: {message.score.rationale}", indent=" "))
print(f"\nTotal attacks: {total}")


# ---------------------------------------------------------------------------
Expand Down
Loading
Loading