diff --git a/tests/patch_attribute_testslide.py b/tests/patch_attribute_testslide.py index a15b13d..c3f9d07 100644 --- a/tests/patch_attribute_testslide.py +++ b/tests/patch_attribute_testslide.py @@ -217,3 +217,9 @@ def patch_attribute_passes_for_private_with_allow_private(self): self.patch_attribute( sample_module.SomeClass, "_private_attr", "notsoprivate", allow_private=True ) + + @context.example + def unpatching_restores_falsy_values(self): + self.patch_attribute(sample_module.SomeUnhashableClass, "class_attr", 1) + unpatch_all_mocked_attributes() + self.assertEqual(sample_module.SomeUnhashableClass.class_attr, 0) diff --git a/testslide/core/patch.py b/testslide/core/patch.py index 98d0566..8d3d0e9 100644 --- a/testslide/core/patch.py +++ b/testslide/core/patch.py @@ -96,7 +96,7 @@ def _patch( setattr(type(target), attribute, property(fget=lambda _: new_value)) def unpatcher() -> None: - if restore_value: + if restore or restore_value: setattr(type(target), attribute, original_property) else: delattr(target, attribute) @@ -105,7 +105,7 @@ def unpatcher() -> None: setattr(target, attribute, new_value) def unpatcher() -> None: - if restore_value: + if restore or restore_value: setattr(target, attribute, restore_value) else: