From 013f97fd8a6679588c996b5cfb6ecdeb3ac79739 Mon Sep 17 00:00:00 2001 From: sobolevn Date: Fri, 15 Sep 2023 10:43:11 +0300 Subject: [PATCH 1/3] gh-109409: Fix inheritance of frozen dataclass from non-frozen dataclass mixins --- Lib/dataclasses.py | 11 +++-- Lib/test/test_dataclasses/__init__.py | 47 +++++++++++++++++++ ...-09-15-10-42-30.gh-issue-109409.RlffA3.rst | 2 + 3 files changed, 57 insertions(+), 3 deletions(-) create mode 100644 Misc/NEWS.d/next/Library/2023-09-15-10-42-30.gh-issue-109409.RlffA3.rst diff --git a/Lib/dataclasses.py b/Lib/dataclasses.py index 84f8d68ce092a45..c64ab9581f85da4 100644 --- a/Lib/dataclasses.py +++ b/Lib/dataclasses.py @@ -944,8 +944,9 @@ def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, # Find our base classes in reverse MRO order, and exclude # ourselves. In reversed order so that more derived classes # override earlier field definitions in base classes. As long as - # we're iterating over them, see if any are frozen. + # we're iterating over them, see if all or any of them are frozen. any_frozen_base = False + all_frozen_bases = None has_dataclass_bases = False for b in cls.__mro__[-1:0:-1]: # Only process classes that have been processed by our @@ -955,7 +956,11 @@ def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, has_dataclass_bases = True for f in base_fields.values(): fields[f.name] = f - if getattr(b, _PARAMS).frozen: + if all_frozen_bases is None: + all_frozen_bases = True + current_frozen = getattr(b, _PARAMS).frozen + all_frozen_bases = all_frozen_bases and current_frozen + if current_frozen: any_frozen_base = True # Annotations defined specifically in this class (not in base classes). @@ -1025,7 +1030,7 @@ def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, 'frozen one') # Raise an exception if we're frozen, but none of our bases are. - if not any_frozen_base and frozen: + if all_frozen_bases is False and frozen: raise TypeError('cannot inherit frozen dataclass from a ' 'non-frozen one') diff --git a/Lib/test/test_dataclasses/__init__.py b/Lib/test/test_dataclasses/__init__.py index 7c07dfc77de208d..d79b197002f1928 100644 --- a/Lib/test/test_dataclasses/__init__.py +++ b/Lib/test/test_dataclasses/__init__.py @@ -2863,6 +2863,53 @@ class C: class D(C): j: int + def test_inherit_frozen_mutliple_inheritance(self): + @dataclass + class A: + pass + + @dataclass(frozen=True) + class B: + pass + + with self.assertRaisesRegex(TypeError, + 'cannot inherit non-frozen dataclass from a frozen one'): + @dataclass + class C(A, B): + pass + with self.assertRaisesRegex(TypeError, + 'cannot inherit non-frozen dataclass from a frozen one'): + @dataclass + class C(B, A): + pass + + with self.assertRaisesRegex(TypeError, + 'cannot inherit frozen dataclass from a non-frozen one'): + @dataclass(frozen=True) + class C(A, B): + pass + with self.assertRaisesRegex(TypeError, + 'cannot inherit frozen dataclass from a non-frozen one'): + @dataclass(frozen=True) + class C(B, A): + pass + + def test_inherit_frozen_mutliple_inheritance_regular_mixins(self): + @dataclass(frozen=True) + class D: + pass + + class M: + pass + + class C1(D, M): + pass + self.assertEqual(C1.__mro__, (C1, D, M, object)) + + class C2(M, D): + pass + self.assertEqual(C2.__mro__, (C2, M, D, object)) + def test_inherit_nonfrozen_from_empty(self): @dataclass class C: diff --git a/Misc/NEWS.d/next/Library/2023-09-15-10-42-30.gh-issue-109409.RlffA3.rst b/Misc/NEWS.d/next/Library/2023-09-15-10-42-30.gh-issue-109409.RlffA3.rst new file mode 100644 index 000000000000000..eddad643e434d7b --- /dev/null +++ b/Misc/NEWS.d/next/Library/2023-09-15-10-42-30.gh-issue-109409.RlffA3.rst @@ -0,0 +1,2 @@ +Fix error when it was possible to inherit a frozen dataclass from multiple +parents some of which were possibly not frozen. From 0c38af922096adb78f629903d0046c45a2b4f7ec Mon Sep 17 00:00:00 2001 From: sobolevn Date: Fri, 15 Sep 2023 16:36:46 +0300 Subject: [PATCH 2/3] More tests --- Lib/test/test_dataclasses/__init__.py | 88 +++++++++++++++++++++------ 1 file changed, 68 insertions(+), 20 deletions(-) diff --git a/Lib/test/test_dataclasses/__init__.py b/Lib/test/test_dataclasses/__init__.py index d79b197002f1928..eacf175dd8ce565 100644 --- a/Lib/test/test_dataclasses/__init__.py +++ b/Lib/test/test_dataclasses/__init__.py @@ -2872,27 +2872,38 @@ class A: class B: pass - with self.assertRaisesRegex(TypeError, - 'cannot inherit non-frozen dataclass from a frozen one'): - @dataclass - class C(A, B): - pass - with self.assertRaisesRegex(TypeError, - 'cannot inherit non-frozen dataclass from a frozen one'): - @dataclass - class C(B, A): - pass + class M: + pass - with self.assertRaisesRegex(TypeError, - 'cannot inherit frozen dataclass from a non-frozen one'): - @dataclass(frozen=True) - class C(A, B): - pass - with self.assertRaisesRegex(TypeError, - 'cannot inherit frozen dataclass from a non-frozen one'): - @dataclass(frozen=True) - class C(B, A): - pass + for bases in ( + (A, B), + (B, A), + (B, M), + (M, B), + ): + with self.subTest(bases=bases): + with self.assertRaisesRegex( + TypeError, + 'cannot inherit non-frozen dataclass from a frozen one', + ): + @dataclass + class C(*bases): + pass + + for bases in ( + (A, B), + (B, A), + (A, M), + (M, A), + ): + with self.subTest(bases=bases): + with self.assertRaisesRegex( + TypeError, + 'cannot inherit frozen dataclass from a non-frozen one', + ): + @dataclass(frozen=True) + class C(*bases): + pass def test_inherit_frozen_mutliple_inheritance_regular_mixins(self): @dataclass(frozen=True) @@ -2910,6 +2921,43 @@ class C2(M, D): pass self.assertEqual(C2.__mro__, (C2, M, D, object)) + @dataclass(frozen=True) + class C3(D, M): + pass + self.assertEqual(C3.__mro__, (C3, D, M, object)) + + @dataclass(frozen=True) + class C4(M, D): + pass + self.assertEqual(C4.__mro__, (C4, M, D, object)) + + def test_multiple_frozen_dataclasses_inheritance(self): + @dataclass(frozen=True) + class A: + pass + + @dataclass(frozen=True) + class B: + pass + + class C1(A, B): + pass + self.assertEqual(C1.__mro__, (C1, A, B, object)) + + class C2(B, A): + pass + self.assertEqual(C2.__mro__, (C2, B, A, object)) + + @dataclass(frozen=True) + class C3(A, B): + pass + self.assertEqual(C3.__mro__, (C3, A, B, object)) + + @dataclass(frozen=True) + class C4(B, A): + pass + self.assertEqual(C4.__mro__, (C4, B, A, object)) + def test_inherit_nonfrozen_from_empty(self): @dataclass class C: From b09b25a974145f8b0bb14aaa92ffb995a3e872e8 Mon Sep 17 00:00:00 2001 From: sobolevn Date: Wed, 11 Oct 2023 19:50:43 +0300 Subject: [PATCH 3/3] Address review --- Lib/dataclasses.py | 5 +- Lib/test/test_dataclasses/__init__.py | 66 +++++++++++++-------------- 2 files changed, 36 insertions(+), 35 deletions(-) diff --git a/Lib/dataclasses.py b/Lib/dataclasses.py index c64ab9581f85da4..845f6e87d625671 100644 --- a/Lib/dataclasses.py +++ b/Lib/dataclasses.py @@ -946,6 +946,8 @@ def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, # override earlier field definitions in base classes. As long as # we're iterating over them, see if all or any of them are frozen. any_frozen_base = False + # By default `all_frozen_bases` is `None` to represent a case, + # where some dataclasses does not have any bases with `_FIELDS` all_frozen_bases = None has_dataclass_bases = False for b in cls.__mro__[-1:0:-1]: @@ -960,8 +962,7 @@ def _process_class(cls, init, repr, eq, order, unsafe_hash, frozen, all_frozen_bases = True current_frozen = getattr(b, _PARAMS).frozen all_frozen_bases = all_frozen_bases and current_frozen - if current_frozen: - any_frozen_base = True + any_frozen_base = any_frozen_base or current_frozen # Annotations defined specifically in this class (not in base classes). # diff --git a/Lib/test/test_dataclasses/__init__.py b/Lib/test/test_dataclasses/__init__.py index eacf175dd8ce565..e9d76aaf3db0b38 100644 --- a/Lib/test/test_dataclasses/__init__.py +++ b/Lib/test/test_dataclasses/__init__.py @@ -2865,21 +2865,21 @@ class D(C): def test_inherit_frozen_mutliple_inheritance(self): @dataclass - class A: + class NotFrozen: pass @dataclass(frozen=True) - class B: + class Frozen: pass - class M: + class NotDataclass: pass for bases in ( - (A, B), - (B, A), - (B, M), - (M, B), + (NotFrozen, Frozen), + (Frozen, NotFrozen), + (Frozen, NotDataclass), + (NotDataclass, Frozen), ): with self.subTest(bases=bases): with self.assertRaisesRegex( @@ -2887,14 +2887,14 @@ class M: 'cannot inherit non-frozen dataclass from a frozen one', ): @dataclass - class C(*bases): + class NotFrozenChild(*bases): pass for bases in ( - (A, B), - (B, A), - (A, M), - (M, A), + (NotFrozen, Frozen), + (Frozen, NotFrozen), + (NotFrozen, NotDataclass), + (NotDataclass, NotFrozen), ): with self.subTest(bases=bases): with self.assertRaisesRegex( @@ -2902,61 +2902,61 @@ class C(*bases): 'cannot inherit frozen dataclass from a non-frozen one', ): @dataclass(frozen=True) - class C(*bases): + class FrozenChild(*bases): pass def test_inherit_frozen_mutliple_inheritance_regular_mixins(self): @dataclass(frozen=True) - class D: + class Frozen: pass - class M: + class NotDataclass: pass - class C1(D, M): + class C1(Frozen, NotDataclass): pass - self.assertEqual(C1.__mro__, (C1, D, M, object)) + self.assertEqual(C1.__mro__, (C1, Frozen, NotDataclass, object)) - class C2(M, D): + class C2(NotDataclass, Frozen): pass - self.assertEqual(C2.__mro__, (C2, M, D, object)) + self.assertEqual(C2.__mro__, (C2, NotDataclass, Frozen, object)) @dataclass(frozen=True) - class C3(D, M): + class C3(Frozen, NotDataclass): pass - self.assertEqual(C3.__mro__, (C3, D, M, object)) + self.assertEqual(C3.__mro__, (C3, Frozen, NotDataclass, object)) @dataclass(frozen=True) - class C4(M, D): + class C4(NotDataclass, Frozen): pass - self.assertEqual(C4.__mro__, (C4, M, D, object)) + self.assertEqual(C4.__mro__, (C4, NotDataclass, Frozen, object)) def test_multiple_frozen_dataclasses_inheritance(self): @dataclass(frozen=True) - class A: + class FrozenA: pass @dataclass(frozen=True) - class B: + class FrozenB: pass - class C1(A, B): + class C1(FrozenA, FrozenB): pass - self.assertEqual(C1.__mro__, (C1, A, B, object)) + self.assertEqual(C1.__mro__, (C1, FrozenA, FrozenB, object)) - class C2(B, A): + class C2(FrozenB, FrozenA): pass - self.assertEqual(C2.__mro__, (C2, B, A, object)) + self.assertEqual(C2.__mro__, (C2, FrozenB, FrozenA, object)) @dataclass(frozen=True) - class C3(A, B): + class C3(FrozenA, FrozenB): pass - self.assertEqual(C3.__mro__, (C3, A, B, object)) + self.assertEqual(C3.__mro__, (C3, FrozenA, FrozenB, object)) @dataclass(frozen=True) - class C4(B, A): + class C4(FrozenB, FrozenA): pass - self.assertEqual(C4.__mro__, (C4, B, A, object)) + self.assertEqual(C4.__mro__, (C4, FrozenB, FrozenA, object)) def test_inherit_nonfrozen_from_empty(self): @dataclass