Skip to content
Closed
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
81 changes: 1 addition & 80 deletions pyrit/scenario/core/matrix_atomic_attack_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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__)
Expand Down Expand Up @@ -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,
Expand Down
4 changes: 4 additions & 0 deletions pyrit/scenario/core/scenario.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)

Expand Down Expand Up @@ -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,
)

Expand Down
100 changes: 100 additions & 0 deletions pyrit/scenario/core/technique_resolution.py
Original file line number Diff line number Diff line change
@@ -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}
4 changes: 3 additions & 1 deletion pyrit/scenario/scenarios/benchmark/adversarial.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
65 changes: 0 additions & 65 deletions tests/unit/scenario/core/test_matrix_atomic_attack_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down
80 changes: 80 additions & 0 deletions tests/unit/scenario/core/test_technique_resolution.py
Original file line number Diff line number Diff line change
@@ -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
Loading