diff --git a/pyrit/common/apply_defaults.py b/pyrit/common/apply_defaults.py index 7fdb7b8b22..39baa7eae8 100644 --- a/pyrit/common/apply_defaults.py +++ b/pyrit/common/apply_defaults.py @@ -125,14 +125,18 @@ def get_default_value( Returns: Tuple of (found, value) where found indicates if a default was found. """ - # First, try exact match - scope = DefaultValueScope( - class_type=class_type, - parameter_name=parameter_name, - include_subclasses=True, - ) - if scope in self._default_values: - return True, self._default_values[scope] + # First, try exact match for both registration flags. A default + # registered with include_subclasses=False must still apply to the + # registered class itself - the flag only controls whether subclasses + # inherit the default. + for include_subclasses in (True, False): + scope = DefaultValueScope( + class_type=class_type, + parameter_name=parameter_name, + include_subclasses=include_subclasses, + ) + if scope in self._default_values: + return True, self._default_values[scope] # Then, check parent classes if include_subclasses is True for existing_scope, value in self._default_values.items(): diff --git a/tests/unit/common/test_apply_defaults.py b/tests/unit/common/test_apply_defaults.py index 77f472ee07..198a472cc4 100644 --- a/tests/unit/common/test_apply_defaults.py +++ b/tests/unit/common/test_apply_defaults.py @@ -98,6 +98,26 @@ def test_global_default_values_no_subclass_when_disabled(): assert found is False +def test_global_default_values_no_subclass_still_applies_to_base(): + registry = GlobalDefaultValues() + registry.set_default_value(class_type=_Base, parameter_name="name", value="no-inherit", include_subclasses=False) + found, val = registry.get_default_value(class_type=_Base, parameter_name="name") + assert found is True + assert val == "no-inherit" + child = _Child() + assert child.name is None + + +def test_global_default_values_mixed_flags_same_param(): + registry = GlobalDefaultValues() + registry.set_default_value(class_type=_Base, parameter_name="name", value="base-only", include_subclasses=False) + registry.set_default_value(class_type=_Child, parameter_name="name", value="child-default") + found, val = registry.get_default_value(class_type=_Base, parameter_name="name") + assert found is True and val == "base-only" + found, val = registry.get_default_value(class_type=_Child, parameter_name="name") + assert found is True and val == "child-default" + + def test_global_default_values_reset(): registry = GlobalDefaultValues() registry.set_default_value(class_type=_Base, parameter_name="name", value="x")