From 5279cfa367579756daea66f8adbcc0ce23f68324 Mon Sep 17 00:00:00 2001 From: fei <204683769+feiiiiii5@users.noreply.github.com> Date: Sun, 23 Aug 2026 18:59:25 +0800 Subject: [PATCH] fix(common): honor include_subclasses=False for the registered class itself GlobalDefaultValues.get_default_value() previously only looked up the exact scope with include_subclasses=True. A default registered with include_subclasses=False was therefore unreachable for every class - including the very class it was registered for - because the exact-match probe used the wrong flag and the parent-class fallback loop skips registrations whose flag is False. The flag's documented meaning is 'apply to subclasses as well', so a False registration must still apply to the registered class itself while stopping short of inheritance. The lookup now probes both flags for the exact (class_type, parameter_name) scope before falling back to the inheritance loop. Added regression tests: no-subclass default resolves for the base class but not the child, and mixed-flag registrations on base/child resolve independently. --- pyrit/common/apply_defaults.py | 20 ++++++++++++-------- tests/unit/common/test_apply_defaults.py | 20 ++++++++++++++++++++ 2 files changed, 32 insertions(+), 8 deletions(-) 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")