Skip to content
Open
Show file tree
Hide file tree
Changes from 2 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
9 changes: 2 additions & 7 deletions src/strands_evals/experimental/redteam/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,9 +72,6 @@ def __init__(
self._attack_strategies = attack_strategies or []
self._by_label = self._build_by_label(self._attack_strategies)
self._model = model
# case name -> strategy run metadata; the base Experiment drops task-returned
# metadata, so we join this onto the report ourselves.
self._run_meta: dict[str, dict[str, Any]] = {}

@property
def agent(self) -> Agent | MultiAgentBase | TargetSession | None:
Expand Down Expand Up @@ -177,12 +174,11 @@ async def run_evaluations_async( # type: ignore[override]
"""
if max_workers < 1:
raise ValueError(f"max_workers must be >= 1, got {max_workers}")
self._run_meta.clear()
if task is None:
task = self._default_task(parallel=max_workers > 1)
# Swap _cases for the expanded cross-product. Each case has a unique name (case x strategy
# label), so parallel workers never collide on the same key in the base runner's results
# buffer or in `self._run_meta`.
# buffer.
original_cases = self._cases
self._cases = self._expand_cross_product()
try:
Expand All @@ -191,7 +187,7 @@ async def run_evaluations_async( # type: ignore[override]
)
finally:
self._cases = original_cases
return RedTeamReport.from_evaluation_report(report, run_meta=self._run_meta)
return RedTeamReport.from_evaluation_report(report)

def _default_task(self, *, parallel: bool = False) -> Callable[[Case[InputT, OutputT]], Any]:
if self._agent is None and self._agent_factory is None:
Expand All @@ -206,7 +202,6 @@ def _default_task(self, *, parallel: bool = False) -> Callable[[Case[InputT, Out
self._by_label,
agent_factory=self._agent_factory,
model=self._model,
run_meta=self._run_meta,
parallel=parallel,
),
)
Expand Down
28 changes: 23 additions & 5 deletions src/strands_evals/experimental/redteam/report.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from __future__ import annotations

import warnings
from collections.abc import Callable
from dataclasses import dataclass, field

Expand Down Expand Up @@ -91,9 +92,17 @@ def from_evaluation_report(cls, report: EvaluationReport, run_meta: dict[str, di

Args:
report: Flattened report from the base experiment, one row per (case, evaluator).
run_meta: Per-case strategy run metadata keyed by case name; merged into each case's metadata so
the report sees it.
run_meta: Deprecated. Per-case strategy run metadata keyed by case name; merged into each case's
metadata so the report sees it. Return the stats as a `run_results` entry in the task's
`environment_state` instead.
"""
if run_meta is not None:
warnings.warn(
"`run_meta` is deprecated and will be removed in a future release. Return run stats as "
'`EnvironmentState(name="run_results", state=...)` in the task\'s `environment_state` instead.',
DeprecationWarning,
stacklevel=2,
)
run_meta = run_meta or {}
n = len(report.cases)
if not (len(report.scores) == n and len(report.test_passes) == n and len(report.reasons) == n):
Expand Down Expand Up @@ -125,6 +134,7 @@ def attack_results(self) -> list[AttackResult]:
name = case_data.get("name", f"case_{i}")
evaluator = case_data.get("evaluator", "evaluator")
metadata = case_data.get("metadata") or {}
run_results = _run_results(case_data)
result = by_case.setdefault(
name,
AttackResult(
Expand All @@ -133,10 +143,10 @@ def attack_results(self) -> list[AttackResult]:
strategy=metadata.get("strategy", "unknown"),
severity=metadata.get("severity", "unknown"),
objective=metadata.get("actor_goal", ""),
turns_used=metadata.get("turns_used"),
backtracks=metadata.get("backtracks"),
turns_used=run_results.get("turns_used"),
backtracks=run_results.get("backtracks"),
conversation=case_data.get("actual_output") or [],
pruned_branches=metadata.get("pruned_branches") or [],
pruned_branches=run_results.get("pruned_branches") or [],
),
)
result.reasons[evaluator] = self.reasons[i]
Expand Down Expand Up @@ -350,6 +360,14 @@ def _base_case_is_unique(results: list[AttackResult]) -> bool:
return len(set(keys)) == len(keys)


def _run_results(case_data: dict) -> dict:
"""Return the row's `run_results` environment state, falling back to `metadata` for older reports."""
for state in case_data.get("actual_environment_state") or []:
if state.get("name") == "run_results":
return state.get("state") or {}
Comment thread
nhungbi marked this conversation as resolved.
Outdated
return case_data.get("metadata") or {}


def _format_run_stats(result: AttackResult) -> str:
"""Render the strategy's per-run stats when present."""
parts = []
Expand Down
8 changes: 8 additions & 0 deletions src/strands_evals/experimental/redteam/strategies/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,14 @@ class AttackRunResult:
metadata: dict[str, Any] = field(default_factory=dict)
pruned_branches: list[dict[str, Any]] = field(default_factory=list)

def to_metadata(self) -> dict[str, Any]:
Comment thread
nhungbi marked this conversation as resolved.
Outdated
"""Return the run details `RedTeamReport` reads: `turns_used`, `backtracks` and `pruned_branches`."""
return {
"turns_used": self.metadata.get("turns_used"),
"backtracks": self.metadata.get("backtracks"),
"pruned_branches": self.pruned_branches,
}


class AttackStrategy(ABC):
"""Base class for red team attack strategies.
Expand Down
30 changes: 13 additions & 17 deletions src/strands_evals/experimental/redteam/task.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from strands.models.model import Model
from strands.multiagent.base import MultiAgentBase

from ...types import EnvironmentState
from .case import RedTeamCase
from .strategies import AttackStrategy
from .strategies.target_session import StrandsAgentSession, StrandsMultiAgentSession, TargetSession
Expand All @@ -25,22 +26,25 @@ def _build_attacker_task(
*,
agent_factory: Callable[[], Agent | MultiAgentBase | TargetSession] | None = None,
model: Model | str | None = None,
run_meta: dict[str, dict[str, Any]] | None = None,
parallel: bool = False,
) -> Callable[[RedTeamCase], dict]:
"""Build a `task(case) -> {"output": conversation, "trajectory": tool_uses}` callable.
"""Build a `task(case) -> {"output", "trajectory", "environment_state"}` callable.

Looks up each case's strategy by `metadata["strategy"]` and delegates the multi-turn loop to
`strategy.run_attack`, injecting a `TargetSession`. `MAX_ALLOWED_TURNS` is the hard ceiling. Run metadata
is recorded into `run_meta` keyed by case name.
`strategy.run_attack`, injecting a `TargetSession`. `MAX_ALLOWED_TURNS` is the hard ceiling.

The returned dict holds:
output: The attacker/target conversation.
trajectory: The target's tool uses.
environment_state: `[EnvironmentState(name="run_results", state=result.to_metadata())]`, the run
stats `RedTeamReport` reads.

Args:
agent: The shared target for sequential runs. Required when `agent_factory` is None.
by_label: Strategy registry keyed by `metadata["strategy"]` label.
agent_factory: Zero-arg callable returning a fresh target for each case. Required for parallel
runs (`parallel=True`); takes precedence over `agent` when both are set.
model: Model passed through to `strategy.run_attack` for strategy-internal LLM calls.
run_meta: Per-case strategy metadata sink, keyed by `case.name`.
parallel: When True, every case is built from `agent_factory` so concurrent cases never share
mutable state. When False, all cases share one target and rewind to a once-captured baseline
between cases.
Expand All @@ -57,7 +61,6 @@ def _build_attacker_task(
agent_factory=agent_factory,
by_label=by_label,
model=model,
run_meta=run_meta,
)

# Sequential path: `agent` is guaranteed non-None here -- the upfront check rejects (None,
Expand All @@ -66,7 +69,6 @@ def _build_attacker_task(
agent=agent, # type: ignore[arg-type]
by_label=by_label,
model=model,
run_meta=run_meta,
)


Expand All @@ -75,7 +77,6 @@ def _build_shared_target_task_fn(
agent: Agent | MultiAgentBase | TargetSession,
by_label: dict[str, AttackStrategy],
model: Model | str | None,
run_meta: dict[str, dict[str, Any]] | None,
) -> Callable[[RedTeamCase], dict]:
"""Build a task fn that drives one shared target across cases, rewinding to a fixed baseline.

Expand All @@ -97,7 +98,7 @@ def task_fn(case: RedTeamCase) -> dict:
session = _build_session(agent, baseline=initial_snapshot)
session.reset()

return _run_attack(strategy, case, session, model=model, run_meta=run_meta)
return _run_attack(strategy, case, session, model=model)

return task_fn

Expand All @@ -108,7 +109,6 @@ def _build_per_case_task_fn(
agent_factory: Callable[[], Agent | MultiAgentBase | TargetSession] | None,
by_label: dict[str, AttackStrategy],
model: Model | str | None,
run_meta: dict[str, dict[str, Any]] | None,
) -> Callable[[RedTeamCase], dict]:
"""Build a per-case task fn that constructs its own target every case.

Expand All @@ -127,9 +127,7 @@ def task_fn(case: RedTeamCase) -> dict:
session = _build_session(make_target(), baseline=None)
session.reset()

# CPython dict assignment for a single distinct key is atomic, and case names are unique
# per cross-product expansion, so concurrent writers never target the same key.
return _run_attack(strategy, case, session, model=model, run_meta=run_meta)
return _run_attack(strategy, case, session, model=model)

return task_fn

Expand All @@ -140,20 +138,18 @@ def _run_attack(
session: TargetSession,
*,
model: Model | str | None,
run_meta: dict[str, dict[str, Any]] | None,
) -> dict:
"""Run one attack and record its strategy metadata into `run_meta`.
"""Run one attack and return its conversation, tool uses and run stats.

Errors propagate: the base `Experiment` retries throttling, records any other failure as an
error reason (which the report classifies as errored), and skips caching the failed case.
"""
result = strategy.run_attack(case, session, max_turns=MAX_ALLOWED_TURNS, model=model)
if run_meta is not None and case.name is not None:
run_meta[case.name] = {**result.metadata, "pruned_branches": result.pruned_branches}
return {
"output": result.conversation,
# Snapshot of the trace; the next case's session.reset() clears this list in place.
"trajectory": list(session.trace),
"environment_state": [EnvironmentState(name="run_results", state=result.to_metadata())],
Comment thread
nhungbi marked this conversation as resolved.
Outdated
}


Expand Down
70 changes: 69 additions & 1 deletion tests/strands_evals/experimental/redteam/test_experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
import pytest
from strands.models.model import Model

from strands_evals import LocalFileTaskResultStore
from strands_evals.evaluators import Evaluator
from strands_evals.experimental.redteam.case import RedTeamCase
from strands_evals.experimental.redteam.evaluators import AttackSuccessEvaluator
from strands_evals.experimental.redteam.experiment import RedTeamExperiment
Expand All @@ -17,6 +19,7 @@
)
from strands_evals.experimental.redteam.strategies.base import AttackRunResult, AttackStrategy
from strands_evals.experimental.redteam.types import AttackGoal, RedTeamConfig
from strands_evals.types import EnvironmentState, EvaluationOutput


class _StubModel(Model):
Expand Down Expand Up @@ -193,7 +196,7 @@ async def test_parallel_report_per_case_isolation():
"""End-to-end: under `max_workers > 1`, each case's output lands on its own report row.

Pins the property the parallel path is *for* -- that concurrent cases don't bleed each
other's outputs through the shared strategy / shared `_run_meta` dict.
other's outputs through the shared strategy.
"""

class _EchoStrategy(AttackStrategy):
Expand Down Expand Up @@ -494,3 +497,68 @@ def _task(case, _cap=captured):
assert runs[0] == runs[1] # second run identical, not squared
# held cases were never mutated
assert [c.name for c in exp.cases] == ["c0"]


class _PassEvaluator(Evaluator):
"""Deterministic evaluator so report tests don't call a judge model."""

def evaluate(self, evaluation_case):
return [EvaluationOutput(score=0.0, test_pass=True, reason="defended")]


_PRUNED = [{"role": "attacker", "content": "direct ask"}, {"role": "target", "content": "no"}]


class _RunStatsStrategy(AttackStrategy):
"""Returns fixed run stats so tests can check they reach the report."""

@property
def name(self) -> str:
return "stats"

def run_attack(self, case, target_session, *, max_turns, model=None, **kwargs) -> AttackRunResult:
return AttackRunResult(conversation=[], metadata={"turns_used": 3, "backtracks": 1}, pruned_branches=_PRUNED)


def test_custom_task_run_results_reach_report():
"""A user task that returns `run_results` fills the report's run stats; no side channel needed."""

def task(case):
return {
"output": [],
"environment_state": [
EnvironmentState(name="run_results", state={"turns_used": 4, "backtracks": 0, "pruned_branches": []})
],
}

exp = RedTeamExperiment(cases=[_case("c0")], attack_strategies=[_StubStrategy()], evaluators=[_PassEvaluator()])
(result,) = exp.run_evaluations(task=task).attack_results()

assert result.turns_used == 4
assert result.backtracks == 0


def test_cached_rerun_keeps_run_stats(tmp_path):
"""A rerun served from the evaluation data store still shows turns and blocked attempts."""
runs = 0

def factory():
nonlocal runs
runs += 1
return _FakeSession()

store = LocalFileTaskResultStore(tmp_path)
exp = RedTeamExperiment(
cases=[_case("c0")],
agent_factory=factory,
attack_strategies=[_RunStatsStrategy()],
evaluators=[_PassEvaluator()],
)
first = exp.run_evaluations(evaluation_data_store=store).attack_results()
second = exp.run_evaluations(evaluation_data_store=store).attack_results()

assert runs == 1 # the second run came from the cache
for (result,) in (first, second):
assert result.turns_used == 3
assert result.backtracks == 1
assert result.pruned_branches == _PRUNED
50 changes: 50 additions & 0 deletions tests/strands_evals/experimental/redteam/test_report.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
"""Tests for RedTeamReport."""

import warnings

import pytest

from strands_evals.experimental.redteam.report import AttackResult, RedTeamReport
from strands_evals.types.evaluation import NOT_APPLICABLE, EvaluationOutput
from strands_evals.types.evaluation_report import EvaluationReport
Expand Down Expand Up @@ -87,6 +91,52 @@ def test_single_evaluator_single_case(self):
assert r.passes == {"judge": False}
assert r.reasons == {"judge": "bypassed"}

def test_run_stats_read_from_run_results_environment_state(self):
pruned = [{"role": "attacker", "content": "a"}, {"role": "target", "content": "no"}]
case = _case("c0", "guideline_bypass", "crescendo", "high", turns_used=99)
case["actual_environment_state"] = [
{"name": "other", "state": {"turns_used": 7}},
{"name": "run_results", "state": {"turns_used": 3, "backtracks": 1, "pruned_branches": pruned}},
]
report = RedTeamReport.from_evaluation_report(
_flatten(_eval_report("judge", [case], scores=[0.0], passes=[True], reasons=[""]))
)

(r,) = report.attack_results()
# run_results wins over metadata and over other environment states
assert r.turns_used == 3
assert r.backtracks == 1
assert r.pruned_branches == pruned

def test_run_stats_fall_back_to_metadata_without_run_results(self):
case = _case("c0", "guideline_bypass", "crescendo", "high", turns_used=5, backtracks=2)
case["actual_environment_state"] = [{"name": "other", "state": {"turns_used": 7}}]
report = RedTeamReport.from_evaluation_report(
_flatten(_eval_report("judge", [case], scores=[0.0], passes=[True], reasons=[""]))
)

(r,) = report.attack_results()
assert r.turns_used == 5
assert r.backtracks == 2
assert r.pruned_branches == []

def test_run_meta_is_deprecated_but_still_merged(self):
case = _case("c0", "guideline_bypass", "crescendo", "high")
eval_report = _eval_report("judge", [case], scores=[0.0], passes=[True], reasons=[""])

with pytest.warns(DeprecationWarning, match="run_meta"):
report = RedTeamReport.from_evaluation_report(eval_report, run_meta={"c0": {"turns_used": 4}})

assert report.attack_results()[0].turns_used == 4

def test_no_deprecation_warning_without_run_meta(self):
case = _case("c0", "guideline_bypass", "crescendo", "high")
eval_report = _eval_report("judge", [case], scores=[0.0], passes=[True], reasons=[""])

with warnings.catch_warnings():
warnings.simplefilter("error", DeprecationWarning)
RedTeamReport.from_evaluation_report(eval_report)

def test_multiple_evaluators_merge_on_case_name(self):
cases = [_case("c0", "guideline_bypass", "gradual_escalation", "high")]
r1 = _eval_report("judge", cases, scores=[0.0], passes=[False], reasons=["bypassed"])
Expand Down
Loading
Loading