From 8cf62e2b616c84df35be278505c7b964c39ddd3c Mon Sep 17 00:00:00 2001 From: Manohar Paturi <186662190+ManoharPaturi@users.noreply.github.com> Date: Sat, 26 Sep 2026 12:37:34 +0530 Subject: [PATCH] MAINT: extract technique factory resolution into its own module resolve_technique_factories, resolve_technique_factories_for_techniques and TechniqueResolutionError move from matrix_atomic_attack_builder to a new pyrit/scenario/core/technique_resolution module, so non-matrix scenario shapes can resolve techniques without importing the matrix builder (issue #2869). consumers and patch targets updated; resolution tests split into their own module with full parity. Signed-off-by: Manohar Paturi <186662190+ManoharPaturi@users.noreply.github.com> --- .../core/matrix_atomic_attack_builder.py | 81 +------------- pyrit/scenario/core/scenario.py | 4 + pyrit/scenario/core/technique_resolution.py | 100 ++++++++++++++++++ .../scenarios/benchmark/adversarial.py | 4 +- .../core/test_matrix_atomic_attack_builder.py | 65 ------------ .../core/test_technique_resolution.py | 80 ++++++++++++++ .../test_default_run_size_estimates.py | 12 +-- 7 files changed, 194 insertions(+), 152 deletions(-) create mode 100644 pyrit/scenario/core/technique_resolution.py create mode 100644 tests/unit/scenario/core/test_technique_resolution.py diff --git a/pyrit/scenario/core/matrix_atomic_attack_builder.py b/pyrit/scenario/core/matrix_atomic_attack_builder.py index fdef49e858..08862be4ce 100644 --- a/pyrit/scenario/core/matrix_atomic_attack_builder.py +++ b/pyrit/scenario/core/matrix_atomic_attack_builder.py @@ -28,16 +28,7 @@ from pyrit.prompt_normalizer import ConverterConfiguration 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``. - """ - +from pyrit.scenario.core.technique_resolution import resolve_technique_factories if TYPE_CHECKING: from collections.abc import Callable, Mapping, Sequence @@ -46,7 +37,6 @@ 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__) @@ -142,75 +132,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..a347a15d24 100644 --- a/pyrit/scenario/core/scenario.py +++ b/pyrit/scenario/core/scenario.py @@ -710,6 +710,8 @@ def _build_technique_size_components( from pyrit.scenario.core.matrix_atomic_attack_builder import ( filter_compatible_seed_groups, + ) + from pyrit.scenario.core.technique_resolution import ( resolve_technique_factories_for_techniques, ) @@ -755,6 +757,8 @@ def _get_technique_compatibility_bounds( """ from pyrit.scenario.core.matrix_atomic_attack_builder import ( filter_compatible_seed_groups, + ) + from pyrit.scenario.core.technique_resolution import ( resolve_technique_factories_for_techniques, ) diff --git a/pyrit/scenario/core/technique_resolution.py b/pyrit/scenario/core/technique_resolution.py new file mode 100644 index 0000000000..4a7f067707 --- /dev/null +++ b/pyrit/scenario/core/technique_resolution.py @@ -0,0 +1,100 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +""" +Technique factory resolution for scenarios. + +Scenarios of different shapes (matrix atomic builders, adaptive scenarios, +composite builders, per-objective builders) all need to resolve selected +:class:`~pyrit.scenario.core.scenario_technique.ScenarioTechnique` entries to +their registered ``AttackTechniqueFactory`` instances. Keeping that logic here +lets non-matrix scenarios use it without importing the matrix builder. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from collections.abc import Sequence + + 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. + + 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} diff --git a/pyrit/scenario/scenarios/benchmark/adversarial.py b/pyrit/scenario/scenarios/benchmark/adversarial.py index b0a7d05dd7..9cd674c261 100644 --- a/pyrit/scenario/scenarios/benchmark/adversarial.py +++ b/pyrit/scenario/scenarios/benchmark/adversarial.py @@ -31,10 +31,12 @@ from pyrit.scenario.core.matrix_atomic_attack_builder import ( MatrixAtomicAttackBuilder, filter_compatible_seed_groups, +) +from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario +from pyrit.scenario.core.technique_resolution import ( resolve_technique_factories, resolve_technique_factories_for_techniques, ) -from pyrit.scenario.core.scenario import BaselineAttackPolicy, Scenario if TYPE_CHECKING: from pyrit.prompt_target import PromptTarget 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..70815a85f7 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,69 +386,6 @@ 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"), - ] - ) - 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") class TestBuildMatrixAtomicAttacks: """``build_matrix_atomic_attacks`` wires the context into the builder in one call.""" 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..76bd19da69 --- /dev/null +++ b/tests/unit/scenario/core/test_technique_resolution.py @@ -0,0 +1,80 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Tests for scenario technique factory resolution (``technique_resolution``).""" + +import pytest + +from pyrit.scenario.core.technique_resolution import ( + TechniqueResolutionError, + resolve_technique_factories, +) +from tests.unit.scenario.core.test_matrix_atomic_attack_builder import ( + _context, + _mock_factory, + _patch_registry, + _technique, +) + + +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"), + ] + ) + 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 diff --git a/tests/unit/scenario/test_default_run_size_estimates.py b/tests/unit/scenario/test_default_run_size_estimates.py index d34c56ff9d..e0c46950c8 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