diff --git a/pyrit/scenario/core/__init__.py b/pyrit/scenario/core/__init__.py index 7e1852274e..992393a906 100644 --- a/pyrit/scenario/core/__init__.py +++ b/pyrit/scenario/core/__init__.py @@ -10,6 +10,11 @@ if TYPE_CHECKING: from pyrit.models.parameter import Parameter + from pyrit.scenario.core._technique_resolution import ( + TechniqueResolutionError, + resolve_technique_factories, + resolve_technique_factories_for_techniques, + ) from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory, ScorerOverridePolicy @@ -48,9 +53,12 @@ "Scenario": "pyrit.scenario.core.scenario", "ScenarioTechnique": "pyrit.scenario.core.scenario_technique", "ScorerOverridePolicy": "pyrit.scenario.core.attack_technique_factory", + "TechniqueResolutionError": "pyrit.scenario.core._technique_resolution", "get_default_scorer_target": "pyrit.scenario.core.scenario_target_defaults", "get_default_adversarial_target": "pyrit.scenario.core.scenario_target_defaults", "override_default_adversarial_target": "pyrit.scenario.core.scenario_target_defaults", + "resolve_technique_factories": "pyrit.scenario.core._technique_resolution", + "resolve_technique_factories_for_techniques": "pyrit.scenario.core._technique_resolution", } __all__ = list(_LAZY_EXPORTS) diff --git a/pyrit/scenario/core/_technique_resolution.py b/pyrit/scenario/core/_technique_resolution.py new file mode 100644 index 0000000000..426daf7920 --- /dev/null +++ b/pyrit/scenario/core/_technique_resolution.py @@ -0,0 +1,95 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from collections.abc import Sequence + + from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory + from pyrit.scenario.core.scenario_context import ScenarioContext + from pyrit.scenario.core.scenario_technique import ScenarioTechnique + + +class TechniqueResolutionError(ValueError): + """ + Raised when a selected scenario technique has no registered factory. + + Subclasses ``ValueError`` so existing ``except ValueError`` handlers keep working, + mirroring ``DatasetConstraintError``. + """ + + +def resolve_technique_factories( + *, + context: ScenarioContext, + extra_factories: dict[str, AttackTechniqueFactory] | None = None, +) -> dict[str, AttackTechniqueFactory]: + """ + Resolve a run's selected techniques to their registered ``AttackTechniqueFactory`` instances. + + Reads the ``AttackTechniqueRegistry`` singleton and keeps only the factories whose name + matches a selected technique, preserving selection order. Raises if any selected + technique has no registered factory so the run cannot silently omit requested work. + + Args: + context (ScenarioContext): The resolved runtime inputs for this run. + extra_factories (dict[str, AttackTechniqueFactory] | None): Scenario-local factories + merged on top of the registry before filtering, so a scenario can offer techniques + without registering them globally. Entries override registry factories of the same + name. + + Returns: + dict[str, AttackTechniqueFactory]: Mapping of technique name to factory, ordered by + the selected techniques. + + Raises: + TechniqueResolutionError: If any selected technique has no registered factory. + """ + return resolve_technique_factories_for_techniques( + scenario_techniques=context.scenario_techniques, + extra_factories=extra_factories, + ) + + +def resolve_technique_factories_for_techniques( + *, + scenario_techniques: Sequence[ScenarioTechnique], + extra_factories: dict[str, AttackTechniqueFactory] | None = None, +) -> dict[str, AttackTechniqueFactory]: + """ + Resolve selected concrete techniques to their canonical factories. + + Accepts any sequence of scenario techniques so callers without a full + `ScenarioContext` (e.g. run-size estimators, dry runs, scenario builders) can + inspect factory metadata or validate coverage. + + Args: + scenario_techniques (Sequence[ScenarioTechnique]): Concrete techniques to resolve. + extra_factories (dict[str, AttackTechniqueFactory] | None): Scenario-local factories + merged on top of the registry. Entries override registry factories of the same name. + + Returns: + dict[str, AttackTechniqueFactory]: Selected factories in technique order. + + Raises: + TechniqueResolutionError: If any selected technique has no registered factory. + """ + from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry + + all_factories = dict(AttackTechniqueRegistry.get_registry_singleton().get_factories_or_raise()) + if extra_factories: + all_factories.update(extra_factories) + + missing = list(dict.fromkeys(t.value for t in scenario_techniques if t.value not in all_factories)) + + if missing: + raise TechniqueResolutionError( + "The following selected attack techniques have no registered factory: " + f"{', '.join(missing)}. Register the techniques (or pass them via " + "extra_factories) before starting the run." + ) + + return {technique.value: all_factories[technique.value] for technique in scenario_techniques} diff --git a/pyrit/scenario/core/matrix_atomic_attack_builder.py b/pyrit/scenario/core/matrix_atomic_attack_builder.py index fdef49e858..163fe7c299 100644 --- a/pyrit/scenario/core/matrix_atomic_attack_builder.py +++ b/pyrit/scenario/core/matrix_atomic_attack_builder.py @@ -26,19 +26,14 @@ from pyrit.executor.attack.single_turn.prompt_sending import PromptSendingAttack from pyrit.models import AttackSeedGroup from pyrit.prompt_normalizer import ConverterConfiguration +from pyrit.scenario.core._technique_resolution import ( + TechniqueResolutionError, + resolve_technique_factories, + resolve_technique_factories_for_techniques, +) from pyrit.scenario.core.atomic_attack import AtomicAttack from pyrit.scenario.core.attack_technique import AttackTechnique - -class TechniqueResolutionError(ValueError): - """ - Raised when a selected scenario technique has no registered factory. - - Subclasses ``ValueError`` so existing ``except ValueError`` handlers keep working, - mirroring ``DatasetConstraintError``. - """ - - if TYPE_CHECKING: from collections.abc import Callable, Mapping, Sequence @@ -46,11 +41,21 @@ class TechniqueResolutionError(ValueError): from pyrit.prompt_target import PromptTarget from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory from pyrit.scenario.core.scenario_context import ScenarioContext - from pyrit.scenario.core.scenario_technique import ScenarioTechnique from pyrit.score import Scorer, TrueFalseScorer logger = logging.getLogger(__name__) +__all__ = [ + "MatrixAtomicAttackBuilder", + "MatrixCombo", + "TechniqueResolutionError", + "build_baseline_atomic_attack", + "build_matrix_atomic_attacks", + "filter_compatible_seed_groups", + "resolve_technique_factories", + "resolve_technique_factories_for_techniques", +] + @dataclass(frozen=True) class MatrixCombo: @@ -142,75 +147,6 @@ def build_baseline_atomic_attack( ) -def resolve_technique_factories( - *, - context: ScenarioContext, - extra_factories: dict[str, AttackTechniqueFactory] | None = None, -) -> dict[str, AttackTechniqueFactory]: - """ - Resolve a run's selected techniques to their registered ``AttackTechniqueFactory`` instances. - - Reads the ``AttackTechniqueRegistry`` singleton and keeps only the factories whose name - matches a selected technique, preserving selection order. Raises if any selected - technique has no registered factory so the run cannot silently omit requested work. - - Args: - context (ScenarioContext): The resolved runtime inputs for this run. - extra_factories (dict[str, AttackTechniqueFactory] | None): Scenario-local factories - merged on top of the registry before filtering, so a scenario can offer techniques - without registering them globally. Entries override registry factories of the same - name. - - Returns: - dict[str, AttackTechniqueFactory]: Mapping of technique name to factory, ordered by - the selected techniques. - - Raises: - TechniqueResolutionError: If any selected technique has no registered factory. - """ - return resolve_technique_factories_for_techniques( - scenario_techniques=context.scenario_techniques, - extra_factories=extra_factories, - ) - - -def resolve_technique_factories_for_techniques( - *, - scenario_techniques: Sequence[ScenarioTechnique], - extra_factories: dict[str, AttackTechniqueFactory] | None = None, -) -> dict[str, AttackTechniqueFactory]: - """ - Resolve selected concrete techniques to their canonical factories. - - Args: - scenario_techniques (Sequence[ScenarioTechnique]): Concrete techniques to resolve. - extra_factories (dict[str, AttackTechniqueFactory] | None): Scenario-local factories - merged on top of the registry. - - Returns: - dict[str, AttackTechniqueFactory]: Selected factories in technique order. - - Raises: - TechniqueResolutionError: If any selected technique has no registered factory. - """ - from pyrit.registry.components.attack_technique_registry import AttackTechniqueRegistry - - all_factories = dict(AttackTechniqueRegistry.get_registry_singleton().get_factories_or_raise()) - if extra_factories: - all_factories.update(extra_factories) - - missing = list(dict.fromkeys(t.value for t in scenario_techniques if t.value not in all_factories)) - - if missing: - raise TechniqueResolutionError( - "The following selected attack techniques have no registered factory: " - f"{', '.join(missing)}. Register the techniques (or pass them via " - "extra_factories) before starting the run." - ) - - return {technique.value: all_factories[technique.value] for technique in scenario_techniques} - - def filter_compatible_seed_groups( *, factory: AttackTechniqueFactory, diff --git a/pyrit/scenario/core/scenario.py b/pyrit/scenario/core/scenario.py index 3984f62c0a..348e3a1010 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -708,10 +708,8 @@ def _build_technique_size_components( ) ] - from pyrit.scenario.core.matrix_atomic_attack_builder import ( - filter_compatible_seed_groups, - resolve_technique_factories_for_techniques, - ) + from pyrit.scenario.core._technique_resolution import resolve_technique_factories_for_techniques + from pyrit.scenario.core.matrix_atomic_attack_builder import filter_compatible_seed_groups factories = resolve_technique_factories_for_techniques( scenario_techniques=self._scenario_techniques, @@ -753,10 +751,8 @@ def _get_technique_compatibility_bounds( dict[str, tuple[int, int]] | None: Technique names mapped to minimum and maximum compatible counts, or ``None`` when the configured sampling shape is unsupported. """ - from pyrit.scenario.core.matrix_atomic_attack_builder import ( - filter_compatible_seed_groups, - resolve_technique_factories_for_techniques, - ) + from pyrit.scenario.core._technique_resolution import resolve_technique_factories_for_techniques + from pyrit.scenario.core.matrix_atomic_attack_builder import filter_compatible_seed_groups summaries = {dataset.name: dataset for dataset in datasets} factories = resolve_technique_factories_for_techniques( diff --git a/tests/unit/scenario/core/test_matrix_atomic_attack_builder.py b/tests/unit/scenario/core/test_matrix_atomic_attack_builder.py index 56413f62ac..79c93e658d 100644 --- a/tests/unit/scenario/core/test_matrix_atomic_attack_builder.py +++ b/tests/unit/scenario/core/test_matrix_atomic_attack_builder.py @@ -26,10 +26,8 @@ from pyrit.scenario.core.matrix_atomic_attack_builder import ( MatrixAtomicAttackBuilder, MatrixCombo, - TechniqueResolutionError, build_baseline_atomic_attack, build_matrix_atomic_attacks, - resolve_technique_factories, ) from pyrit.scenario.core.scenario_context import ScenarioContext from pyrit.score import TrueFalseScorer @@ -388,67 +386,26 @@ def _patch_registry(factories: dict): @pytest.mark.usefixtures("patch_central_database") -class TestResolveTechniqueFactories: - """``resolve_technique_factories`` filters the registry to the selected techniques.""" - - def test_keeps_only_selected_in_order(self): - factories = { - "alpha": _mock_factory(name="alpha"), - "beta": _mock_factory(name="beta"), - "gamma": _mock_factory(name="gamma"), - } - context = _context(techniques=[_technique("beta"), _technique("alpha")]) - with _patch_registry(factories): - resolved = resolve_technique_factories(context=context) - assert list(resolved.keys()) == ["beta", "alpha"] - - def test_raises_when_any_selected_technique_is_missing(self): - factories = {"alpha": _mock_factory(name="alpha")} - context = _context(techniques=[_technique("alpha"), _technique("missing")]) - with _patch_registry(factories), pytest.raises(TechniqueResolutionError, match="missing"): - resolve_technique_factories(context=context) - - def test_raises_when_all_selected_techniques_missing(self): - """A nonempty selection resolving to nothing must fail loudly, not run baseline-only.""" - factories = {"alpha": _mock_factory(name="alpha")} - context = _context(techniques=[_technique("missing_a"), _technique("missing_b")]) - with _patch_registry(factories), pytest.raises(TechniqueResolutionError, match="missing_a"): - resolve_technique_factories(context=context) - - def test_empty_selection_resolves_without_error(self): - context = _context(techniques=[]) - with _patch_registry({}): - assert resolve_technique_factories(context=context) == {} - - def test_error_lists_each_missing_technique_once_in_selection_order(self): - factories = {"alpha": _mock_factory(name="alpha")} - context = _context( - techniques=[ - _technique("missing_a"), - _technique("alpha"), - _technique("missing_b"), - _technique("missing_a"), - ] +class TestLegacyReExports: + """Ensure backward compatibility of re-exports in matrix_atomic_attack_builder and pyrit.scenario.core.""" + + def test_legacy_re_exports_are_identical_objects(self): + import pyrit.scenario.core as scenario_core + import pyrit.scenario.core._technique_resolution as new_module + import pyrit.scenario.core.matrix_atomic_attack_builder as legacy_module + + assert legacy_module.TechniqueResolutionError is new_module.TechniqueResolutionError + assert legacy_module.resolve_technique_factories is new_module.resolve_technique_factories + assert ( + legacy_module.resolve_technique_factories_for_techniques + is new_module.resolve_technique_factories_for_techniques + ) + assert scenario_core.TechniqueResolutionError is new_module.TechniqueResolutionError + assert scenario_core.resolve_technique_factories is new_module.resolve_technique_factories + assert ( + scenario_core.resolve_technique_factories_for_techniques + is new_module.resolve_technique_factories_for_techniques ) - with _patch_registry(factories), pytest.raises(TechniqueResolutionError) as exc_info: - resolve_technique_factories(context=context) - message = str(exc_info.value) - assert message.index("missing_a") < message.index("missing_b") - assert message.count("missing_a") == 1 - - def test_extra_factories_merged_and_override_registry(self): - registry_factories = {"alpha": _mock_factory(name="alpha")} - local_alpha = _mock_factory(name="alpha") - local_only = _mock_factory(name="local") - context = _context(techniques=[_technique("alpha"), _technique("local")]) - with _patch_registry(registry_factories): - resolved = resolve_technique_factories( - context=context, - extra_factories={"alpha": local_alpha, "local": local_only}, - ) - assert list(resolved.keys()) == ["alpha", "local"] - assert resolved["alpha"] is local_alpha # extra overrides the registry factory of the same name - assert resolved["local"] is local_only # local-only factory is selectable without global registration @pytest.mark.usefixtures("patch_central_database") diff --git a/tests/unit/scenario/core/test_technique_resolution.py b/tests/unit/scenario/core/test_technique_resolution.py new file mode 100644 index 0000000000..717e4a6be8 --- /dev/null +++ b/tests/unit/scenario/core/test_technique_resolution.py @@ -0,0 +1,146 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Unit tests for pyrit.scenario.core._technique_resolution.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from pyrit.prompt_target import PromptTarget +from pyrit.scenario.core._technique_resolution import ( + TechniqueResolutionError, + resolve_technique_factories, + resolve_technique_factories_for_techniques, +) +from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory +from pyrit.scenario.core.scenario_context import ScenarioContext + + +def _mock_factory(*, name: str) -> MagicMock: + factory = MagicMock(spec=AttackTechniqueFactory) + factory.name = name + return factory + + +def _technique(value: str) -> SimpleNamespace: + return SimpleNamespace(value=value) + + +def _context(*, techniques) -> ScenarioContext: + return ScenarioContext( + objective_target=MagicMock(spec=PromptTarget), + scenario_techniques=techniques, + dataset_config=MagicMock(), + memory_labels={"op": "unit"}, + include_baseline=False, + seed_groups_by_dataset={}, + ) + + +def _patch_registry(factories: dict): + registry = MagicMock() + registry.get_factories_or_raise.return_value = factories + return patch( + "pyrit.registry.components.attack_technique_registry.AttackTechniqueRegistry.get_registry_singleton", + return_value=registry, + ) + + +@pytest.mark.usefixtures("patch_central_database") +class TestTechniqueResolutionError: + """TechniqueResolutionError is a ValueError subclass.""" + + def test_subclasses_value_error(self): + assert issubclass(TechniqueResolutionError, ValueError) + err = TechniqueResolutionError("missing factory") + assert isinstance(err, ValueError) + + +@pytest.mark.usefixtures("patch_central_database") +class TestResolveTechniqueFactories: + """resolve_technique_factories resolves context scenario techniques.""" + + def test_keeps_only_selected_in_order(self): + factories = { + "alpha": _mock_factory(name="alpha"), + "beta": _mock_factory(name="beta"), + "gamma": _mock_factory(name="gamma"), + } + context = _context(techniques=[_technique("beta"), _technique("alpha")]) + with _patch_registry(factories): + resolved = resolve_technique_factories(context=context) + assert list(resolved.keys()) == ["beta", "alpha"] + + def test_raises_when_any_selected_technique_is_missing(self): + factories = {"alpha": _mock_factory(name="alpha")} + context = _context(techniques=[_technique("alpha"), _technique("missing")]) + with _patch_registry(factories), pytest.raises(TechniqueResolutionError, match="missing"): + resolve_technique_factories(context=context) + + def test_raises_when_all_selected_techniques_missing(self): + factories = {"alpha": _mock_factory(name="alpha")} + context = _context(techniques=[_technique("missing_a"), _technique("missing_b")]) + with _patch_registry(factories), pytest.raises(TechniqueResolutionError, match="missing_a"): + resolve_technique_factories(context=context) + + def test_empty_selection_resolves_without_error(self): + context = _context(techniques=[]) + with _patch_registry({}): + assert resolve_technique_factories(context=context) == {} + + def test_error_lists_each_missing_technique_once_in_selection_order(self): + factories = {"alpha": _mock_factory(name="alpha")} + context = _context( + techniques=[ + _technique("missing_a"), + _technique("alpha"), + _technique("missing_b"), + _technique("missing_a"), + ] + ) + with _patch_registry(factories), pytest.raises(TechniqueResolutionError) as exc_info: + resolve_technique_factories(context=context) + message = str(exc_info.value) + assert message.index("missing_a") < message.index("missing_b") + assert message.count("missing_a") == 1 + assert "Register the techniques (or pass them via extra_factories)" in message + + def test_extra_factories_merged_and_override_registry(self): + registry_factories = {"alpha": _mock_factory(name="alpha")} + local_alpha = _mock_factory(name="alpha") + local_only = _mock_factory(name="local") + context = _context(techniques=[_technique("alpha"), _technique("local")]) + with _patch_registry(registry_factories): + resolved = resolve_technique_factories( + context=context, + extra_factories={"alpha": local_alpha, "local": local_only}, + ) + assert list(resolved.keys()) == ["alpha", "local"] + assert resolved["alpha"] is local_alpha + assert resolved["local"] is local_only + + +@pytest.mark.usefixtures("patch_central_database") +class TestResolveTechniqueFactoriesForTechniques: + """resolve_technique_factories_for_techniques accepts raw sequences of ScenarioTechnique.""" + + def test_resolves_raw_sequence_preserving_order(self): + factories = { + "tech1": _mock_factory(name="tech1"), + "tech2": _mock_factory(name="tech2"), + } + techniques = [_technique("tech2"), _technique("tech1")] + with _patch_registry(factories): + resolved = resolve_technique_factories_for_techniques(scenario_techniques=techniques) + assert list(resolved.keys()) == ["tech2", "tech1"] + + def test_supports_extra_factories_without_context(self): + local_factory = _mock_factory(name="custom") + with _patch_registry({}): + resolved = resolve_technique_factories_for_techniques( + scenario_techniques=[_technique("custom")], + extra_factories={"custom": local_factory}, + ) + assert resolved["custom"] is local_factory diff --git a/tests/unit/scenario/test_default_run_size_estimates.py b/tests/unit/scenario/test_default_run_size_estimates.py index d34c56ff9d..d592c699e5 100644 --- a/tests/unit/scenario/test_default_run_size_estimates.py +++ b/tests/unit/scenario/test_default_run_size_estimates.py @@ -320,7 +320,7 @@ async def test_matrix_estimate_filters_each_technique_seed_population_like_execu ) with patch( - "pyrit.scenario.core.matrix_atomic_attack_builder.resolve_technique_factories_for_techniques", + "pyrit.scenario.core._technique_resolution.resolve_technique_factories_for_techniques", return_value={"one": plain_factory, "two": conversation_factory}, ): estimate = await scenario.get_run_size_estimate_async() @@ -363,7 +363,7 @@ async def resolve_groups() -> tuple[dict[str, list[AttackSeedGroup]], list[Scena factory.seed_technique = None with patch( - "pyrit.scenario.core.matrix_atomic_attack_builder.resolve_technique_factories_for_techniques", + "pyrit.scenario.core._technique_resolution.resolve_technique_factories_for_techniques", return_value={"one": factory, "two": factory}, ): estimate = await scenario.get_run_size_estimate_async() @@ -423,7 +423,7 @@ async def resolve_groups() -> tuple[dict[str, list[AttackSeedGroup]], list[Scena ) with patch( - "pyrit.scenario.core.matrix_atomic_attack_builder.resolve_technique_factories_for_techniques", + "pyrit.scenario.core._technique_resolution.resolve_technique_factories_for_techniques", return_value={"one": plain_factory, "two": conversation_factory}, ): estimate = await scenario.get_run_size_estimate_async() @@ -469,7 +469,7 @@ async def resolve_groups() -> tuple[dict[str, list[AttackSeedGroup]], list[Scena factory.seed_technique = None with patch( - "pyrit.scenario.core.matrix_atomic_attack_builder.resolve_technique_factories_for_techniques", + "pyrit.scenario.core._technique_resolution.resolve_technique_factories_for_techniques", return_value={"one": factory, "two": factory}, ): estimate = await scenario.get_run_size_estimate_async() @@ -488,7 +488,7 @@ def test_compatibility_bounds_skip_missing_factories_and_require_dataset_summari scenario._estimate_full_groups_by_dataset = {"sample": [_seed_group("one")]} with patch( - "pyrit.scenario.core.matrix_atomic_attack_builder.resolve_technique_factories_for_techniques", + "pyrit.scenario.core._technique_resolution.resolve_technique_factories_for_techniques", return_value={}, ): assert scenario._get_technique_compatibility_bounds(datasets=[]) == {} @@ -496,7 +496,7 @@ def test_compatibility_bounds_skip_missing_factories_and_require_dataset_summari factory = MagicMock() factory.seed_technique = None with patch( - "pyrit.scenario.core.matrix_atomic_attack_builder.resolve_technique_factories_for_techniques", + "pyrit.scenario.core._technique_resolution.resolve_technique_factories_for_techniques", return_value={"one": factory}, ): assert scenario._get_technique_compatibility_bounds(datasets=[]) is None