Skip to content
Merged
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
8 changes: 8 additions & 0 deletions pyrit/scenario/core/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
95 changes: 95 additions & 0 deletions pyrit/scenario/core/_technique_resolution.py
Original file line number Diff line number Diff line change
@@ -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(
Comment thread
romanlutz marked this conversation as resolved.
*,
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}
96 changes: 16 additions & 80 deletions pyrit/scenario/core/matrix_atomic_attack_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,31 +26,36 @@
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

from pyrit.converter import Converter
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:
Expand Down Expand Up @@ -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,
Expand Down
12 changes: 4 additions & 8 deletions pyrit/scenario/core/scenario.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
81 changes: 19 additions & 62 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,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")
Expand Down
Loading
Loading