Skip to content

Commit aad228e

Browse files
MAINT: Minor Scenario class improvements (microsoft#1190)
Co-authored-by: Victor Valbuena <50061128+ValbuenaVC@users.noreply.github.com>
1 parent 577f507 commit aad228e

8 files changed

Lines changed: 105 additions & 21 deletions

File tree

pyrit/scenarios/scenario.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -79,7 +79,6 @@ def __init__(
7979
name: str,
8080
version: int,
8181
strategy_class: Type[ScenarioStrategy],
82-
default_aggregate: ScenarioStrategy,
8382
objective_scorer_identifier: Optional[Dict[str, str]] = None,
8483
include_default_baseline: bool = True,
8584
scenario_result_id: Optional[Union[uuid.UUID, str]] = None,
@@ -91,8 +90,6 @@ def __init__(
9190
name (str): Descriptive name for the scenario.
9291
version (int): Version number of the scenario.
9392
strategy_class (Type[ScenarioStrategy]): The strategy enum class for this scenario.
94-
default_aggregate (ScenarioStrategy): The default aggregate strategy to use when
95-
scenario_strategies is None in initialize_async.
9693
objective_scorer_identifier (Optional[Dict[str, str]]): Identifier for the objective scorer.
9794
include_default_baseline (bool): Whether to include a baseline atomic attack that sends all objectives
9895
from the first atomic attack without modifications. Most scenarios should have some kind of
@@ -119,7 +116,6 @@ def __init__(
119116

120117
# Store strategy configuration for use in initialize_async
121118
self._strategy_class = strategy_class
122-
self._default_aggregate = default_aggregate
123119

124120
# These will be set in initialize_async
125121
self._objective_target: Optional[PromptTarget] = None
@@ -284,7 +280,7 @@ async def initialize_async(
284280

285281
# Prepare scenario strategies using the stored configuration
286282
self._scenario_composites = self._strategy_class.prepare_scenario_strategies(
287-
scenario_strategies, default_aggregate=self._default_aggregate
283+
scenario_strategies, default_aggregate=self.get_default_strategy()
288284
)
289285

290286
self._atomic_attacks = await self._get_atomic_attacks_async()

pyrit/scenarios/scenario_strategy.py

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -393,6 +393,47 @@ def is_single_strategy(self) -> bool:
393393
"""Check if this composition contains only a single strategy."""
394394
return len(self._strategies) == 1
395395

396+
@staticmethod
397+
def extract_single_strategy_values(
398+
composites: Sequence["ScenarioCompositeStrategy"], *, strategy_type: type[T]
399+
) -> Set[str]:
400+
"""
401+
Extract strategy values from single-strategy composites.
402+
403+
This is a helper method for scenarios that don't support composition and need
404+
to filter or map strategies by their values. It flattens the composites into
405+
a simple set of strategy values.
406+
407+
This method enforces that all composites contain only a single strategy. If any
408+
composite contains multiple strategies, a ValueError is raised.
409+
410+
Args:
411+
composites (Sequence[ScenarioCompositeStrategy]): List of composite strategies.
412+
Each composite must contain only a single strategy.
413+
strategy_type (type[T]): The strategy enum type to filter by.
414+
415+
Returns:
416+
Set[str]: Set of strategy values (e.g., {"base64", "rot13", "morse_code"}).
417+
418+
Raises:
419+
ValueError: If any composite contains multiple strategies.
420+
"""
421+
# Check that all composites are single-strategy
422+
multi_strategy_composites = [comp for comp in composites if not comp.is_single_strategy]
423+
if multi_strategy_composites:
424+
composite_names = [comp.name for comp in multi_strategy_composites]
425+
raise ValueError(
426+
f"extract_single_strategy_values() requires all composites to contain a single strategy. "
427+
f"Found composites with multiple strategies: {composite_names}"
428+
)
429+
430+
return {
431+
strategy.value
432+
for composite in composites
433+
for strategy in composite.strategies
434+
if isinstance(strategy, strategy_type)
435+
}
436+
396437
@staticmethod
397438
def get_composite_name(strategies: Sequence[ScenarioStrategy]) -> str:
398439
"""

pyrit/scenarios/scenarios/encoding_scenario.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
from pyrit.scenarios.atomic_attack import AtomicAttack
3737
from pyrit.scenarios.scenario import Scenario
3838
from pyrit.scenarios.scenario_strategy import (
39+
ScenarioCompositeStrategy,
3940
ScenarioStrategy,
4041
)
4142
from pyrit.score import TrueFalseScorer
@@ -153,7 +154,6 @@ def __init__(
153154
name="Encoding Scenario",
154155
version=self.version,
155156
strategy_class=EncodingStrategy,
156-
default_aggregate=EncodingStrategy.ALL,
157157
objective_scorer_identifier=objective_scorer.get_identifier(),
158158
include_default_baseline=include_baseline,
159159
scenario_result_id=scenario_result_id,
@@ -224,8 +224,9 @@ def _get_converter_attacks(self) -> list[AtomicAttack]:
224224
]
225225

226226
# Filter to only include selected strategies
227-
# Extract strategy names from composites (each has exactly one strategy since composition not supported)
228-
selected_encoding_names = {comp.strategies[0].value for comp in self._scenario_composites if comp.strategies}
227+
selected_encoding_names = ScenarioCompositeStrategy.extract_single_strategy_values(
228+
self._scenario_composites, strategy_type=EncodingStrategy
229+
)
229230
converters_with_encodings = [
230231
(conv, name) for conv, name in all_converters_with_encodings if name in selected_encoding_names
231232
]

pyrit/scenarios/scenarios/foundry_scenario.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -275,7 +275,6 @@ def __init__(
275275
name="Foundry Scenario",
276276
version=self.version,
277277
strategy_class=FoundryStrategy,
278-
default_aggregate=FoundryStrategy.EASY,
279278
objective_scorer_identifier=self._objective_scorer.get_identifier(),
280279
include_default_baseline=include_baseline,
281280
scenario_result_id=scenario_result_id,

tests/unit/scenarios/test_scenario.py

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -94,18 +94,17 @@ def __init__(self, atomic_attacks_to_return=None, **kwargs):
9494
# Default include_default_baseline=False for tests unless explicitly specified
9595
kwargs.setdefault("include_default_baseline", False)
9696

97-
# Add required strategy_class and default_aggregate if not provided
97+
# Add required strategy_class if not provided
9898

9999
class TestStrategy(ScenarioStrategy):
100-
TEST = ("test", set())
100+
TEST = ("test", {"concrete"}) # Tagged as concrete, not aggregate
101101
ALL = ("all", {"all"})
102102

103103
@classmethod
104104
def get_aggregate_tags(cls) -> set[str]:
105105
return {"all"}
106106

107107
kwargs.setdefault("strategy_class", TestStrategy)
108-
kwargs.setdefault("default_aggregate", TestStrategy.ALL)
109108

110109
super().__init__(**kwargs)
111110
self._atomic_attacks_to_return = atomic_attacks_to_return or []
@@ -118,7 +117,7 @@ def get_strategy_class(cls):
118117

119118
# Return a simple mock strategy class for testing
120119
class TestStrategy(ScenarioStrategy):
121-
TEST = ("test", set())
120+
TEST = ("test", {"concrete"}) # Tagged as concrete, not aggregate
122121
ALL = ("all", {"all"})
123122

124123
@classmethod
@@ -130,7 +129,7 @@ def get_aggregate_tags(cls) -> set[str]:
130129
@classmethod
131130
def get_default_strategy(cls):
132131
"""Return the default strategy for testing."""
133-
return cls.get_strategy_class().TEST
132+
return cls.get_strategy_class().ALL
134133

135134
async def _get_atomic_attacks_async(self):
136135
return self._atomic_attacks_to_return

tests/unit/scenarios/test_scenario_partial_results.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -51,9 +51,8 @@ def __init__(self, *, atomic_attacks_to_return=None, **kwargs):
5151

5252
# Get strategy_class from kwargs or use default
5353
strategy_class = kwargs.pop("strategy_class", None) or self.get_strategy_class()
54-
default_aggregate = kwargs.pop("default_aggregate", None) or self.get_default_strategy()
5554

56-
super().__init__(strategy_class=strategy_class, default_aggregate=default_aggregate, **kwargs)
55+
super().__init__(strategy_class=strategy_class, **kwargs)
5756
self._test_atomic_attacks = atomic_attacks_to_return or []
5857

5958
async def _get_atomic_attacks_async(self):
@@ -63,7 +62,7 @@ async def _get_atomic_attacks_async(self):
6362
def get_strategy_class(cls):
6463

6564
class TestStrategy(ScenarioStrategy):
66-
CONCRETE = ("concrete", set())
65+
CONCRETE = ("concrete", {"concrete"})
6766
ALL = ("all", {"all"})
6867

6968
@classmethod

tests/unit/scenarios/test_scenario_retry.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -122,9 +122,8 @@ def __init__(self, atomic_attacks_to_return=None, **kwargs):
122122

123123
# Get strategy_class from kwargs or use default
124124
strategy_class = kwargs.pop("strategy_class", None) or self.get_strategy_class()
125-
default_aggregate = kwargs.pop("default_aggregate", None) or self.get_default_strategy()
126125

127-
super().__init__(strategy_class=strategy_class, default_aggregate=default_aggregate, **kwargs)
126+
super().__init__(strategy_class=strategy_class, **kwargs)
128127
self._atomic_attacks_to_return = atomic_attacks_to_return or []
129128

130129
@classmethod
@@ -133,7 +132,7 @@ def get_strategy_class(cls):
133132

134133
# Return a simple mock strategy class for testing
135134
class TestStrategy(ScenarioStrategy):
136-
CONCRETE = ("concrete", set())
135+
CONCRETE = ("concrete", {"concrete"})
137136
ALL = ("all", {"all"})
138137

139138
@classmethod

tests/unit/scenarios/test_strategy_validation.py

Lines changed: 51 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
import pytest
77

8-
from pyrit.scenarios import EncodingStrategy, FoundryStrategy
8+
from pyrit.scenarios import EncodingStrategy, FoundryStrategy, ScenarioCompositeStrategy
99

1010

1111
class TestStrategyValidation:
@@ -49,3 +49,53 @@ def test_foundry_validation_rejects_attacks_with_converters_and_another_attack(s
4949
FoundryStrategy.validate_composition(
5050
[FoundryStrategy.Base64, FoundryStrategy.Crescendo, FoundryStrategy.MultiTurn]
5151
)
52+
53+
54+
class TestScenarioCompositeStrategyExtraction:
55+
"""Test extraction of strategy values from composite strategies."""
56+
57+
def test_extract_single_strategy_values_with_single_strategies(self):
58+
"""Test extracting values from single-strategy composites."""
59+
composites = [
60+
ScenarioCompositeStrategy(strategies=[EncodingStrategy.Base64]),
61+
ScenarioCompositeStrategy(strategies=[EncodingStrategy.ROT13]),
62+
ScenarioCompositeStrategy(strategies=[EncodingStrategy.Atbash]),
63+
]
64+
65+
values = ScenarioCompositeStrategy.extract_single_strategy_values(composites, strategy_type=EncodingStrategy)
66+
67+
assert values == {"base64", "rot13", "atbash"}
68+
69+
def test_extract_single_strategy_values_filters_by_type(self):
70+
"""Test that extraction filters by strategy type."""
71+
composites = [
72+
ScenarioCompositeStrategy(strategies=[EncodingStrategy.Base64]),
73+
ScenarioCompositeStrategy(strategies=[FoundryStrategy.ROT13]),
74+
]
75+
76+
# Extract only EncodingStrategy values
77+
encoding_values = ScenarioCompositeStrategy.extract_single_strategy_values(
78+
composites, strategy_type=EncodingStrategy
79+
)
80+
assert encoding_values == {"base64"}
81+
82+
# Extract only FoundryStrategy values
83+
foundry_values = ScenarioCompositeStrategy.extract_single_strategy_values(
84+
composites, strategy_type=FoundryStrategy
85+
)
86+
assert foundry_values == {"rot13"}
87+
88+
def test_extract_single_strategy_values_rejects_multi_strategy_composites(self):
89+
"""Test that extraction raises error if any composite has multiple strategies."""
90+
composites = [
91+
ScenarioCompositeStrategy(strategies=[FoundryStrategy.Base64]),
92+
ScenarioCompositeStrategy(strategies=[FoundryStrategy.ROT13, FoundryStrategy.Atbash]), # Multi-strategy!
93+
]
94+
95+
with pytest.raises(ValueError, match="extract_single_strategy_values.*requires all composites"):
96+
ScenarioCompositeStrategy.extract_single_strategy_values(composites, strategy_type=FoundryStrategy)
97+
98+
def test_extract_single_strategy_values_with_empty_list(self):
99+
"""Test that extraction handles empty composite list."""
100+
values = ScenarioCompositeStrategy.extract_single_strategy_values([], strategy_type=EncodingStrategy)
101+
assert values == set()

0 commit comments

Comments
 (0)