diff --git a/pyrit/scenario/core/_attack_constructor_compatibility.py b/pyrit/scenario/core/_attack_constructor_compatibility.py new file mode 100644 index 0000000000..eeac36f6c3 --- /dev/null +++ b/pyrit/scenario/core/_attack_constructor_compatibility.py @@ -0,0 +1,169 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +from __future__ import annotations + +import inspect +import logging +import sys +import typing +from enum import Enum +from typing import Any, Union + +from pyrit.executor.attack.core.attack_config import AttackScoringConfig + +logger = logging.getLogger(__name__) + + +class ScorerOverridePolicy(str, Enum): + """Policy for what to do when the scenario's scorer is incompatible with an attack's annotation.""" + + SKIP = "skip" + WARN = "warn" + RAISE = "raise" + + +class _ConstructorCompatibilityHelper: + """Evaluates constructor compatibility and extracts type annotations for an attack class.""" + + def __init__( + self, + *, + attack_class: type, + scorer_override_policy: ScorerOverridePolicy, + ) -> None: + self._attack_class = attack_class + self._scorer_override_policy = scorer_override_policy + + self.accepted_params = self._derive_accepted_params() + + @property + def scoring_config_type(self) -> type | None: + """The required ``attack_scoring_config`` subtype, or ``None`` if any config is accepted.""" + return self._derive_scoring_config_type() + + def _derive_accepted_params(self) -> set[str]: + """Return the set of keyword parameter names accepted by the attack class constructor.""" + sig = inspect.signature(self._attack_class.__init__) + return { + name + for name, param in sig.parameters.items() + if name != "self" + and param.kind + in ( + inspect.Parameter.KEYWORD_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD, + ) + } + + def should_apply_scoring_config(self, attack_scoring_config: AttackScoringConfig) -> bool: + """ + Determine whether the scoring config should be forwarded to the attack constructor. + + Checks two conditions: + 1. The attack class accepts an ``attack_scoring_config`` parameter. + 2. The provided config is type-compatible with the attack's annotation. + + When either condition fails, the ``scorer_override_policy`` determines + behavior: RAISE raises ValueError, WARN logs and returns False, SKIP + silently returns False. + + Args: + attack_scoring_config: The scoring config to evaluate. + + Returns: + True if the config should be applied, False otherwise. + + Raises: + ValueError: If the policy is RAISE and the config cannot be applied. + """ + if "attack_scoring_config" not in self.accepted_params: + self._apply_scorer_policy( + f"Scorer config provided but {self._attack_class.__name__} does not accept 'attack_scoring_config'." + ) + return False + + if self.scoring_config_type is None or isinstance(attack_scoring_config, self.scoring_config_type): + return True + + self._apply_scorer_policy( + f"Scorer config of type {type(attack_scoring_config).__name__} is incompatible " + f"with {self._attack_class.__name__} (requires {self.scoring_config_type.__name__})." + ) + return False + + def _apply_scorer_policy(self, message: str) -> None: + """ + Apply the scorer override policy for an incompatibility. + + Args: + message: Description of the incompatibility. + + Raises: + ValueError: If the policy is RAISE. + """ + if self._scorer_override_policy == ScorerOverridePolicy.RAISE: + raise ValueError(message) + if self._scorer_override_policy == ScorerOverridePolicy.WARN: + logger.warning(message) + + def _derive_scoring_config_type(self) -> type | None: + """ + Introspect the attack class to determine the required type for ``attack_scoring_config``. + + Resolves the type annotation (handling ``X | None`` / ``X | None``) and returns + the inner concrete type. Returns ``None`` if the annotation is the base + ``AttackScoringConfig`` or cannot be resolved — meaning any config is accepted. + + Returns: + The narrowed type if the annotation is narrower than the base, else None. + """ + try: + # get_type_hints resolves string annotations from __future__ annotations + hints = typing.get_type_hints( + self._attack_class.__init__, + globalns=getattr(sys.modules.get(self._attack_class.__module__, None), "__dict__", None), + ) + except Exception: + return None + + annotation = hints.get("attack_scoring_config") + if annotation is None: + return None + + inner = self._unwrap_optional(annotation) + if inner is None or inner is AttackScoringConfig: + # Base type or unresolvable — any config is accepted + return None + if not issubclass(inner, AttackScoringConfig): + return None + return inner + + @staticmethod + def _unwrap_optional(annotation: Any) -> type | None: + """ + Unwrap a union containing one concrete type and ``None`` to extract the concrete type. + + Returns: + The inner type X, or None if the annotation cannot be unwrapped to a single type. + """ + # Handle typing.Union and Optional annotations. + origin = typing.get_origin(annotation) + if origin is Union or (hasattr(annotation, "__args__") and origin is None and hasattr(annotation, "__or__")): + args = typing.get_args(annotation) + non_none = [a for a in args if a is not type(None)] + candidate = non_none[0] if len(non_none) == 1 else None + return candidate if isinstance(candidate, type) else None + + # Handle PEP 604 unions (X | None). + if hasattr(annotation, "__args__") and type(annotation).__name__ == "UnionType": + args = annotation.__args__ + non_none = [a for a in args if a is not type(None)] + candidate = non_none[0] if len(non_none) == 1 else None + return candidate if isinstance(candidate, type) else None + + # Plain type (not Optional) + if isinstance(annotation, type): + return annotation + + return None diff --git a/pyrit/scenario/core/attack_technique_factory.py b/pyrit/scenario/core/attack_technique_factory.py index 9ac1af5b0e..84d2f94c3d 100644 --- a/pyrit/scenario/core/attack_technique_factory.py +++ b/pyrit/scenario/core/attack_technique_factory.py @@ -20,11 +20,8 @@ import copy import inspect import logging -import sys -import typing -from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING, Any, Union +from typing import TYPE_CHECKING, Any from pyrit.common.path import EXECUTOR_SEED_PROMPT_PATH from pyrit.executor.attack import PromptSendingAttack @@ -47,6 +44,7 @@ resolve_prompt_source, ) from pyrit.models.seeds.seed_simulated_conversation import NextMessageSystemPromptPaths +from pyrit.scenario.core._attack_constructor_compatibility import ScorerOverridePolicy, _ConstructorCompatibilityHelper from pyrit.scenario.core.attack_technique import AttackTechnique from pyrit.scenario.core.scenario_target_defaults import get_default_adversarial_target @@ -59,14 +57,6 @@ logger = logging.getLogger(__name__) -class ScorerOverridePolicy(str, Enum): - """Policy for what to do when the scenario's scorer is incompatible with an attack's annotation.""" - - SKIP = "skip" - WARN = "warn" - RAISE = "raise" - - class AttackTechniqueFactory(Identifiable): """ A self-describing factory that produces AttackTechnique instances on demand. @@ -173,6 +163,11 @@ class constructor signature and seed-technique shape. self._supports_additional_request_converters = supports_additional_request_converters self._scorer_override_policy = scorer_override_policy + self._compatibility_helper = _ConstructorCompatibilityHelper( + attack_class=self._attack_class, + scorer_override_policy=self._scorer_override_policy, + ) + self._uses_adversarial = uses_adversarial if uses_adversarial is not None else self._derive_uses_adversarial() self._validate_kwargs() @@ -386,7 +381,7 @@ def _validate_converter_composition(self) -> None: """ if ( self._supports_additional_request_converters - and "attack_converter_config" not in self._get_accepted_params() + and "attack_converter_config" not in self._compatibility_helper.accepted_params ): raise ValueError( f"Factory '{self._name}' declares supports_additional_request_converters=True, " @@ -426,16 +421,7 @@ def _validate_kwargs(self) -> None: f"parameter validation. All attack constructor parameters must be explicitly named." ) - valid_params = { - name - for name, param in sig.parameters.items() - if name != "self" - and param.kind - in ( - inspect.Parameter.KEYWORD_ONLY, - inspect.Parameter.POSITIONAL_OR_KEYWORD, - ) - } + valid_params = self._compatibility_helper.accepted_params invalid = set(self._attack_kwargs) - valid_params if invalid: @@ -501,7 +487,7 @@ def can_append_request_converter(self, *, converter_type: type[Converter]) -> bo Returns: bool: ``True`` when the converter can be appended safely. """ - if "attack_converter_config" not in self._get_accepted_params(): + if "attack_converter_config" not in self._compatibility_helper.accepted_params: return False output_types: set[PromptDataType] = {"text"} @@ -579,7 +565,7 @@ def supports_additional_request_converters(self) -> bool: @property def scoring_config_type(self) -> type | None: """The required ``attack_scoring_config`` subtype, or ``None`` if any config is accepted.""" - return self._get_scoring_config_type() + return self._compatibility_helper.scoring_config_type def with_adversarial_system_prompt_prefix(self, prefix: str) -> AttackTechniqueFactory: """ @@ -608,7 +594,7 @@ def with_adversarial_system_prompt_prefix(self, prefix: str) -> AttackTechniqueF """ SeedPrompt.reject_jinja_syntax(prefix, component_name="adversarial_system_prompt_prefix") seed_technique, supports_simulated = self._copy_seed_technique_with_prefix(prefix=prefix) - accepts_adversarial_config = "attack_adversarial_config" in self._get_accepted_params() + accepts_adversarial_config = "attack_adversarial_config" in self._compatibility_helper.accepted_params if not accepts_adversarial_config and not supports_simulated: raise ValueError( f"Factory '{self._name}' cannot accept an adversarial system prompt prefix. " @@ -720,10 +706,9 @@ class constructor accepts ``attack_converter_config``. kwargs = dict(self._attack_kwargs) kwargs["objective_target"] = objective_target - accepted_params = self._get_accepted_params() - if self._should_apply_scoring_config( + accepted_params = self._compatibility_helper.accepted_params + if self._compatibility_helper.should_apply_scoring_config( attack_scoring_config=attack_scoring_config, - accepted_params=accepted_params, ): kwargs["attack_scoring_config"] = attack_scoring_config if "attack_adversarial_config" in accepted_params and ( @@ -870,139 +855,6 @@ def _copy_seed_technique_with_prefix( return self._seed_technique, False return self._seed_technique.model_copy(update={"seeds": seeds}, deep=True), True - def _get_accepted_params(self) -> set[str]: - """Return the set of keyword parameter names accepted by the attack class constructor.""" - sig = inspect.signature(self._attack_class.__init__) - return { - name - for name, param in sig.parameters.items() - if name != "self" - and param.kind - in ( - inspect.Parameter.KEYWORD_ONLY, - inspect.Parameter.POSITIONAL_OR_KEYWORD, - ) - } - - def _should_apply_scoring_config( - self, - *, - attack_scoring_config: AttackScoringConfig, - accepted_params: set[str], - ) -> bool: - """ - Determine whether the scoring config should be forwarded to the attack constructor. - - Checks two conditions: - 1. The attack class accepts an ``attack_scoring_config`` parameter. - 2. The provided config is type-compatible with the attack's annotation. - - When either condition fails, the ``scorer_override_policy`` determines - behavior: RAISE raises ValueError, WARN logs and returns False, SKIP - silently returns False. - - Args: - attack_scoring_config: The scoring config to evaluate. - accepted_params: The set of parameter names the attack class accepts. - - Returns: - True if the config should be applied, False otherwise. - - Raises: - ValueError: If the policy is RAISE and the config cannot be applied. - """ - if "attack_scoring_config" not in accepted_params: - self._apply_scorer_policy( - f"Scorer config provided but {self._attack_class.__name__} does not accept 'attack_scoring_config'." - ) - return False - - required_type = self._get_scoring_config_type() - if required_type is None or isinstance(attack_scoring_config, required_type): - return True - - self._apply_scorer_policy( - f"Scorer config of type {type(attack_scoring_config).__name__} is incompatible " - f"with {self._attack_class.__name__} (requires {required_type.__name__})." - ) - return False - - def _apply_scorer_policy(self, message: str) -> None: - """ - Apply the scorer override policy for an incompatibility. - - Args: - message: Description of the incompatibility. - - Raises: - ValueError: If the policy is RAISE. - """ - if self._scorer_override_policy == ScorerOverridePolicy.RAISE: - raise ValueError(message) - if self._scorer_override_policy == ScorerOverridePolicy.WARN: - logger.warning(message) - - def _get_scoring_config_type(self) -> type | None: - """ - Introspect the attack class to determine the required type for ``attack_scoring_config``. - - Resolves the type annotation (handling ``X | None`` / ``X | None``) and returns - the inner concrete type. Returns ``None`` if the annotation is the base - ``AttackScoringConfig`` or cannot be resolved — meaning any config is accepted. - - Returns: - The narrowed type if the annotation is narrower than the base, else None. - """ - try: - # get_type_hints resolves string annotations from __future__ annotations - hints = typing.get_type_hints( - self._attack_class.__init__, - globalns=getattr(sys.modules.get(self._attack_class.__module__, None), "__dict__", None), - ) - except Exception: - return None - - annotation = hints.get("attack_scoring_config") - if annotation is None: - return None - - inner = self._unwrap_optional(annotation) - if inner is None or inner is AttackScoringConfig: - # Base type or unresolvable — any config is accepted - return None - if not issubclass(inner, AttackScoringConfig): - return None - return inner - - @staticmethod - def _unwrap_optional(annotation: Any) -> type | None: - """ - Unwrap a union containing one concrete type and ``None`` to extract the concrete type. - - Returns: - The inner type X, or None if the annotation cannot be unwrapped to a single type. - """ - # Handle typing.Union and Optional annotations. - origin = typing.get_origin(annotation) - if origin is Union or (hasattr(annotation, "__args__") and origin is None and hasattr(annotation, "__or__")): - args = typing.get_args(annotation) - non_none = [a for a in args if a is not type(None)] - candidate = non_none[0] if len(non_none) == 1 else None - return candidate if isinstance(candidate, type) else None - - # Handle PEP 604 unions (X | None). - if hasattr(annotation, "__args__") and type(annotation).__name__ == "UnionType": - args = annotation.__args__ - non_none = [a for a in args if a is not type(None)] - candidate = non_none[0] if len(non_none) == 1 else None - return candidate if isinstance(candidate, type) else None - - # Plain type (not Optional) - if isinstance(annotation, type): - return annotation - - return None - @staticmethod def _serialize_value(value: Any) -> Any: """ diff --git a/tests/unit/scenario/core/test_attack_constructor_compatibility.py b/tests/unit/scenario/core/test_attack_constructor_compatibility.py new file mode 100644 index 0000000000..51a99a4f97 --- /dev/null +++ b/tests/unit/scenario/core/test_attack_constructor_compatibility.py @@ -0,0 +1,276 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +import typing +from typing import Any +from unittest.mock import MagicMock + +import pytest + +from pyrit.executor.attack.core.attack_config import AttackScoringConfig +from pyrit.scenario.core._attack_constructor_compatibility import ScorerOverridePolicy, _ConstructorCompatibilityHelper + +if typing.TYPE_CHECKING: + # The regression test binds this name only after constructing the helper. + DeferredConfig = AttackScoringConfig + + +class _StubAttack: + def __init__( + self, + *, + objective_target: Any, + attack_scoring_config: AttackScoringConfig | None = None, + ) -> None: + pass + + +class TestScorerPolicy: + """Tests for scorer override policy logic.""" + + def test_should_apply_returns_true_when_type_compatible(self): + """Config passes through when the attack accepts base AttackScoringConfig.""" + helper = _ConstructorCompatibilityHelper( + attack_class=_StubAttack, scorer_override_policy=ScorerOverridePolicy.WARN + ) + config = MagicMock(spec=AttackScoringConfig) + result = helper.should_apply_scoring_config(attack_scoring_config=config) + assert result is True + + def test_should_apply_returns_false_when_param_not_accepted(self): + """If the attack class doesn't accept attack_scoring_config, return False.""" + + class _NoScoringAttack: + def __init__(self, *, objective_target): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_NoScoringAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + config = MagicMock(spec=AttackScoringConfig) + result = helper.should_apply_scoring_config(attack_scoring_config=config) + assert result is False + + def test_should_apply_returns_false_when_type_incompatible_warn(self, caplog): + """When annotation is narrowed and config doesn't match, WARN returns False and logs.""" + + class _NarrowedScoringConfig(AttackScoringConfig): + pass + + class _NarrowedAttack: + def __init__(self, *, objective_target, attack_scoring_config: _NarrowedScoringConfig | None = None): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_NarrowedAttack, scorer_override_policy=ScorerOverridePolicy.WARN + ) + config = MagicMock(spec=AttackScoringConfig) + result = helper.should_apply_scoring_config(attack_scoring_config=config) + assert result is False + assert "incompatible" in caplog.text + + def test_should_apply_raises_when_type_incompatible_raise_policy(self): + """When annotation is narrowed and policy is RAISE, ValueError is raised.""" + + class _NarrowedScoringConfig(AttackScoringConfig): + pass + + class _NarrowedAttack: + def __init__(self, *, objective_target, attack_scoring_config: _NarrowedScoringConfig | None = None): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_NarrowedAttack, scorer_override_policy=ScorerOverridePolicy.RAISE + ) + config = MagicMock(spec=AttackScoringConfig) + with pytest.raises(ValueError, match="incompatible"): + helper.should_apply_scoring_config(attack_scoring_config=config) + + def test_should_apply_accepts_subclass_of_narrowed_type(self): + """A subclass of the narrowed annotation type should pass through.""" + + class _NarrowedScoringConfig(AttackScoringConfig): + pass + + class _NarrowedAttack: + def __init__(self, *, objective_target, attack_scoring_config: _NarrowedScoringConfig | None = None): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_NarrowedAttack, scorer_override_policy=ScorerOverridePolicy.RAISE + ) + config = MagicMock(spec=_NarrowedScoringConfig) + result = helper.should_apply_scoring_config(attack_scoring_config=config) + assert result is True + + def test_apply_scorer_policy_skip_is_silent(self, caplog): + helper = _ConstructorCompatibilityHelper( + attack_class=_StubAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + helper._apply_scorer_policy("some incompatibility message") + assert "some incompatibility message" not in caplog.text + + def test_apply_scorer_policy_warn_logs(self, caplog): + helper = _ConstructorCompatibilityHelper( + attack_class=_StubAttack, scorer_override_policy=ScorerOverridePolicy.WARN + ) + helper._apply_scorer_policy("scorer mismatch detail") + assert "scorer mismatch detail" in caplog.text + + def test_apply_scorer_policy_raise_raises(self): + helper = _ConstructorCompatibilityHelper( + attack_class=_StubAttack, scorer_override_policy=ScorerOverridePolicy.RAISE + ) + with pytest.raises(ValueError, match="error detail"): + helper._apply_scorer_policy("error detail") + + +class TestUnwrapOptional: + """Tests for _ConstructorCompatibilityHelper._unwrap_optional static method.""" + + def test_unwrap_union_with_none(self): + result = _ConstructorCompatibilityHelper._unwrap_optional(AttackScoringConfig | None) + assert result is AttackScoringConfig + + def test_unwrap_plain_type(self): + result = _ConstructorCompatibilityHelper._unwrap_optional(AttackScoringConfig) + assert result is AttackScoringConfig + + def test_unwrap_multi_union_returns_none(self): + result = _ConstructorCompatibilityHelper._unwrap_optional(int | str | None) + assert result is None + + def test_unwrap_none_type_alone(self): + result = _ConstructorCompatibilityHelper._unwrap_optional(type(None)) + assert result is type(None) + + def test_unwrap_non_type_annotation_returns_none(self): + result = _ConstructorCompatibilityHelper._unwrap_optional("SomeForwardRef") + assert result is None + + def test_unwrap_typing_optional_with_none(self): + result = _ConstructorCompatibilityHelper._unwrap_optional(typing.Optional[AttackScoringConfig]) # noqa: UP045 + assert result is AttackScoringConfig + + def test_unwrap_typing_union_multi_returns_none(self): + result = _ConstructorCompatibilityHelper._unwrap_optional(typing.Union[int, str, None]) # noqa: UP007 + assert result is None + + +class TestAcceptedParams: + """Tests for _ConstructorCompatibilityHelper.accepted_params.""" + + def test_accepted_params_discovers_positional_and_keyword_args(self): + class _ComplexAttack: + def __init__(self, objective_target, *, attack_scoring_config=None, custom_kwarg=42): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_ComplexAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.accepted_params == {"objective_target", "attack_scoring_config", "custom_kwarg"} + + def test_accepted_params_ignores_self(self): + class _SelfOnlyAttack: + def __init__(self): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_SelfOnlyAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.accepted_params == set() + + +class TestGetScoringConfigType: + """Tests for _ConstructorCompatibilityHelper.scoring_config_type.""" + + def test_returns_none_when_no_scoring_config_param(self): + class _NoScoringParamAttack: + def __init__(self, *, objective_target: Any): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_NoScoringParamAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.scoring_config_type is None + + def test_returns_none_when_annotation_is_base_attack_scoring_config(self): + class _BaseScoringAttack: + def __init__(self, *, attack_scoring_config: AttackScoringConfig): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_BaseScoringAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.scoring_config_type is None + + def test_returns_none_when_annotation_is_optional_base_attack_scoring_config(self): + class _OptionalBaseScoringAttack: + def __init__(self, *, attack_scoring_config: AttackScoringConfig | None = None): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_OptionalBaseScoringAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.scoring_config_type is None + + def test_returns_subclass_when_annotation_is_narrowed_attack_scoring_config(self): + class _CustomScoringConfig(AttackScoringConfig): + pass + + class _NarrowedScoringAttack: + def __init__(self, *, attack_scoring_config: _CustomScoringConfig): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_NarrowedScoringAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.scoring_config_type is _CustomScoringConfig + + def test_returns_none_when_annotation_is_not_attack_scoring_config_subclass(self): + class _WrongAnnotationAttack: + def __init__(self, *, objective_target, attack_scoring_config: int | None = None): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_WrongAnnotationAttack, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.scoring_config_type is None + + def test_returns_none_when_type_hints_fail_to_resolve(self, monkeypatch): + class _StubWithFailingHints: + def __init__(self, *, attack_scoring_config: AttackScoringConfig): + pass + + def mock_get_type_hints(*args, **kwargs): + raise NameError("Unresolvable type name") + + monkeypatch.setattr(typing, "get_type_hints", mock_get_type_hints) + + helper = _ConstructorCompatibilityHelper( + attack_class=_StubWithFailingHints, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + assert helper.scoring_config_type is None + + def test_resolves_deferred_forward_ref_after_init(self): + class _AttackWithDeferredRef: + def __init__(self, *, attack_scoring_config: "DeferredConfig | None" = None): + pass + + helper = _ConstructorCompatibilityHelper( + attack_class=_AttackWithDeferredRef, scorer_override_policy=ScorerOverridePolicy.SKIP + ) + # Initially unresolvable + assert helper.scoring_config_type is None + + class _DeferredConfig(AttackScoringConfig): + pass + + import sys + + mod_dict = sys.modules[_AttackWithDeferredRef.__module__].__dict__ + mod_dict["DeferredConfig"] = _DeferredConfig + try: + assert helper.scoring_config_type is _DeferredConfig + finally: + mod_dict.pop("DeferredConfig", None) diff --git a/tests/unit/scenario/core/test_attack_technique_factory.py b/tests/unit/scenario/core/test_attack_technique_factory.py index 0e5f6f8cbf..aef3080a34 100644 --- a/tests/unit/scenario/core/test_attack_technique_factory.py +++ b/tests/unit/scenario/core/test_attack_technique_factory.py @@ -3,8 +3,8 @@ """Tests for the AttackTechniqueFactory class.""" -import typing import warnings +from typing import TYPE_CHECKING, cast from unittest.mock import MagicMock, patch import pytest @@ -22,6 +22,12 @@ from pyrit.scenario.core.attack_technique import AttackTechnique from pyrit.scenario.core.attack_technique_factory import AttackTechniqueFactory, ScorerOverridePolicy +if TYPE_CHECKING: + # The regression tests bind these names only after constructing each factory. + DeferredNarrowConfig = AttackScoringConfig + DeferredWarnConfig = AttackScoringConfig + DeferredSkipConfig = AttackScoringConfig + def _make_seed_technique() -> AttackTechniqueSeedGroup: return AttackTechniqueSeedGroup( @@ -253,7 +259,7 @@ class TestFactoryCreate: """Tests for AttackTechniqueFactory.create().""" def _scoring(self) -> AttackScoringConfig: - return MagicMock(spec=AttackScoringConfig) + return cast("AttackScoringConfig", MagicMock(spec=AttackScoringConfig)) def test_create_produces_attack_technique(self): factory = AttackTechniqueFactory(name="test", attack_class=_StubAttack) @@ -477,6 +483,125 @@ def get_identifier(self): assert isinstance(technique, AttackTechnique) + def test_create_with_deferred_forward_ref_scoring_config_policy_raise(self): + """Forward-referenced scoring config defined after factory init resolves and raises on incompatible type.""" + import sys + + class _DeferredAttack: + def __init__( + self, + *, + objective_target: PromptTarget, + attack_scoring_config: "DeferredNarrowConfig | None" = None, + ): + self.objective_target = objective_target + self.attack_scoring_config = attack_scoring_config + + def get_identifier(self): + return ComponentIdentifier(class_name="_DeferredAttack", class_module=__name__) + + factory = AttackTechniqueFactory( + name="deferred_test", + attack_class=_DeferredAttack, + scorer_override_policy=ScorerOverridePolicy.RAISE, + uses_adversarial=False, + ) + + class _DeferredNarrowConfig(AttackScoringConfig): + pass + + mod_dict = sys.modules[_DeferredAttack.__module__].__dict__ + mod_dict["DeferredNarrowConfig"] = _DeferredNarrowConfig + + try: + target = MagicMock(spec=PromptTarget) + # Incompatible base config should raise ValueError + with pytest.raises(ValueError, match="incompatible"): + factory.create(objective_target=target, attack_scoring_config=AttackScoringConfig()) + + # Compatible config should succeed + narrow_config = _DeferredNarrowConfig() + technique = factory.create(objective_target=target, attack_scoring_config=narrow_config) + assert technique.attack.attack_scoring_config is narrow_config + finally: + mod_dict.pop("DeferredNarrowConfig", None) + + def test_create_with_deferred_forward_ref_scoring_config_policy_warn(self, caplog): + """Forward-referenced scoring config with WARN policy logs and omits incompatible config.""" + import sys + + class _DeferredWarnAttack: + def __init__( + self, + *, + objective_target: PromptTarget, + attack_scoring_config: "DeferredWarnConfig | None" = None, + ): + self.objective_target = objective_target + self.attack_scoring_config = attack_scoring_config + + def get_identifier(self): + return ComponentIdentifier(class_name="_DeferredWarnAttack", class_module=__name__) + + factory = AttackTechniqueFactory( + name="deferred_warn_test", + attack_class=_DeferredWarnAttack, + scorer_override_policy=ScorerOverridePolicy.WARN, + uses_adversarial=False, + ) + + class _DeferredWarnConfig(AttackScoringConfig): + pass + + mod_dict = sys.modules[_DeferredWarnAttack.__module__].__dict__ + mod_dict["DeferredWarnConfig"] = _DeferredWarnConfig + + try: + target = MagicMock(spec=PromptTarget) + technique = factory.create(objective_target=target, attack_scoring_config=AttackScoringConfig()) + assert technique.attack.attack_scoring_config is None + assert "incompatible" in caplog.text + finally: + mod_dict.pop("DeferredWarnConfig", None) + + def test_create_with_deferred_forward_ref_scoring_config_policy_skip(self, caplog): + """Forward-referenced scoring config with SKIP policy silently omits incompatible config.""" + import sys + + class _DeferredSkipAttack: + def __init__( + self, + *, + objective_target: PromptTarget, + attack_scoring_config: "DeferredSkipConfig | None" = None, + ): + self.objective_target = objective_target + self.attack_scoring_config = attack_scoring_config + + def get_identifier(self): + return ComponentIdentifier(class_name="_DeferredSkipAttack", class_module=__name__) + + factory = AttackTechniqueFactory( + name="deferred_skip_test", + attack_class=_DeferredSkipAttack, + scorer_override_policy=ScorerOverridePolicy.SKIP, + uses_adversarial=False, + ) + + class _DeferredSkipConfig(AttackScoringConfig): + pass + + mod_dict = sys.modules[_DeferredSkipAttack.__module__].__dict__ + mod_dict["DeferredSkipConfig"] = _DeferredSkipConfig + + try: + target = MagicMock(spec=PromptTarget) + technique = factory.create(objective_target=target, attack_scoring_config=AttackScoringConfig()) + assert technique.attack.attack_scoring_config is None + assert "incompatible" not in caplog.text + finally: + mod_dict.pop("DeferredSkipConfig", None) + class TestFactoryIdentifier: """Tests for AttackTechniqueFactory._build_identifier().""" @@ -606,162 +731,6 @@ def test_different_seed_techniques_produce_different_hashes(self): assert factory1.get_identifier().hash != factory2.get_identifier().hash -class TestScorerPolicy: - """Tests for scorer override policy logic (_should_apply_scoring_config, _apply_scorer_policy).""" - - def test_should_apply_returns_true_when_type_compatible(self): - """Config passes through when the attack accepts base AttackScoringConfig.""" - factory = AttackTechniqueFactory(name="test", attack_class=_StubAttack) - config = MagicMock(spec=AttackScoringConfig) - - result = factory._should_apply_scoring_config( - attack_scoring_config=config, - accepted_params=factory._get_accepted_params(), - ) - - assert result is True - - def test_should_apply_returns_false_when_param_not_accepted(self): - """If the attack class doesn't accept attack_scoring_config, return False.""" - - class _NoScoringAttack: - def __init__(self, *, objective_target): - pass - - def get_identifier(self): - return ComponentIdentifier(class_name="_NoScoringAttack", class_module="test") - - factory = AttackTechniqueFactory( - name="test", - attack_class=_NoScoringAttack, - scorer_override_policy=ScorerOverridePolicy.SKIP, - ) - config = MagicMock(spec=AttackScoringConfig) - - result = factory._should_apply_scoring_config( - attack_scoring_config=config, - accepted_params=factory._get_accepted_params(), - ) - - assert result is False - - def test_should_apply_returns_false_when_type_incompatible_warn(self, caplog): - """When annotation is narrowed and config doesn't match, WARN returns False and logs.""" - - class _NarrowedScoringConfig(AttackScoringConfig): - pass - - class _NarrowedAttack: - def __init__(self, *, objective_target, attack_scoring_config: _NarrowedScoringConfig | None = None): - pass - - def get_identifier(self): - return ComponentIdentifier(class_name="_NarrowedAttack", class_module="test") - - factory = AttackTechniqueFactory( - name="test", - attack_class=_NarrowedAttack, - scorer_override_policy=ScorerOverridePolicy.WARN, - ) - config = MagicMock(spec=AttackScoringConfig) - - result = factory._should_apply_scoring_config( - attack_scoring_config=config, - accepted_params=factory._get_accepted_params(), - ) - - assert result is False - assert "incompatible" in caplog.text - - def test_should_apply_raises_when_type_incompatible_raise_policy(self): - """When annotation is narrowed and policy is RAISE, ValueError is raised.""" - - class _NarrowedScoringConfig(AttackScoringConfig): - pass - - class _NarrowedAttack: - def __init__(self, *, objective_target, attack_scoring_config: _NarrowedScoringConfig | None = None): - pass - - def get_identifier(self): - return ComponentIdentifier(class_name="_NarrowedAttack", class_module="test") - - factory = AttackTechniqueFactory( - name="test", - attack_class=_NarrowedAttack, - scorer_override_policy=ScorerOverridePolicy.RAISE, - ) - config = MagicMock(spec=AttackScoringConfig) - - with pytest.raises(ValueError, match="incompatible"): - factory._should_apply_scoring_config( - attack_scoring_config=config, - accepted_params=factory._get_accepted_params(), - ) - - def test_should_apply_accepts_subclass_of_narrowed_type(self): - """A subclass of the narrowed annotation type should pass through.""" - - class _NarrowedScoringConfig(AttackScoringConfig): - pass - - class _NarrowedAttack: - def __init__(self, *, objective_target, attack_scoring_config: _NarrowedScoringConfig | None = None): - pass - - def get_identifier(self): - return ComponentIdentifier(class_name="_NarrowedAttack", class_module="test") - - factory = AttackTechniqueFactory( - name="test", - attack_class=_NarrowedAttack, - scorer_override_policy=ScorerOverridePolicy.RAISE, - ) - config = MagicMock(spec=_NarrowedScoringConfig) - - result = factory._should_apply_scoring_config( - attack_scoring_config=config, - accepted_params=factory._get_accepted_params(), - ) - - assert result is True - - def test_apply_scorer_policy_skip_is_silent(self, caplog): - """SKIP policy should not log or raise.""" - factory = AttackTechniqueFactory( - name="test", - attack_class=_StubAttack, - scorer_override_policy=ScorerOverridePolicy.SKIP, - ) - - factory._apply_scorer_policy("some incompatibility message") - - assert "some incompatibility message" not in caplog.text - - def test_apply_scorer_policy_warn_logs(self, caplog): - """WARN policy should log a warning.""" - factory = AttackTechniqueFactory( - name="test", - attack_class=_StubAttack, - scorer_override_policy=ScorerOverridePolicy.WARN, - ) - - factory._apply_scorer_policy("scorer mismatch detail") - - assert "scorer mismatch detail" in caplog.text - - def test_apply_scorer_policy_raise_raises(self): - """RAISE policy should raise ValueError with the message.""" - factory = AttackTechniqueFactory( - name="test", - attack_class=_StubAttack, - scorer_override_policy=ScorerOverridePolicy.RAISE, - ) - - with pytest.raises(ValueError, match="error detail"): - factory._apply_scorer_policy("error detail") - - class TestCustomAdversarialPrompt: """Tests for the adversarial_system_prompt / adversarial_seed_prompt params.""" @@ -1267,71 +1236,6 @@ def test_simulated_conversation_resolves_default_lazily(self): mock_default.assert_called_once() -class TestUnwrapOptional: - """Tests for AttackTechniqueFactory._unwrap_optional static method.""" - - def test_unwrap_union_with_none(self): - """X | None should unwrap to X.""" - result = AttackTechniqueFactory._unwrap_optional(AttackScoringConfig | None) - assert result is AttackScoringConfig - - def test_unwrap_plain_type(self): - """A bare type (no Optional wrapping) returns itself.""" - result = AttackTechniqueFactory._unwrap_optional(AttackScoringConfig) - assert result is AttackScoringConfig - - def test_unwrap_multi_union_returns_none(self): - """Union of more than one non-None type returns None (ambiguous).""" - result = AttackTechniqueFactory._unwrap_optional(int | str | None) - assert result is None - - def test_unwrap_none_type_alone(self): - """NoneType alone is a plain type — returns itself.""" - result = AttackTechniqueFactory._unwrap_optional(type(None)) - assert result is type(None) - - def test_unwrap_non_type_annotation_returns_none(self): - """A non-type annotation (e.g., string forward ref) returns None.""" - result = AttackTechniqueFactory._unwrap_optional("SomeForwardRef") - assert result is None - - def test_unwrap_typing_optional_with_none(self): - """typing.Optional[X] (legacy typing.Union syntax) should unwrap to X.""" - # Intentionally uses the legacy typing.Optional/Union construct (rather than `X | None`) - # to exercise _unwrap_optional's typing.Union-origin branch specifically. - result = AttackTechniqueFactory._unwrap_optional(typing.Optional[AttackScoringConfig]) # noqa: UP045 - assert result is AttackScoringConfig - - def test_unwrap_typing_union_multi_returns_none(self): - """typing.Union of more than one non-None type returns None (ambiguous).""" - # Intentionally uses typing.Union (rather than `X | Y`) to exercise _unwrap_optional's - # typing.Union-origin branch specifically. - result = AttackTechniqueFactory._unwrap_optional(typing.Union[int, str, None]) # noqa: UP007 - assert result is None - - -class TestGetScoringConfigType: - """Tests for AttackTechniqueFactory._get_scoring_config_type.""" - - def test_returns_none_when_annotation_is_not_attack_scoring_config_subclass(self): - """A resolved, narrowed annotation that isn't an AttackScoringConfig subclass yields None.""" - - class _WrongAnnotationAttack: - def __init__(self, *, objective_target, attack_scoring_config: int | None = None): - pass - - def get_identifier(self): - return ComponentIdentifier(class_name="_WrongAnnotationAttack", class_module="test") - - factory = AttackTechniqueFactory( - name="test", - attack_class=_WrongAnnotationAttack, - scorer_override_policy=ScorerOverridePolicy.SKIP, - ) - - assert factory._get_scoring_config_type() is None - - @pytest.mark.usefixtures("patch_central_database") class TestWithSimulatedConversationPromptSources: """Tests for the canonical prompt inputs on ``with_simulated_conversation``."""