diff --git a/graphtage/dataclasses.py b/graphtage/dataclasses.py index b8d94db..b5ed482 100644 --- a/graphtage/dataclasses.py +++ b/graphtage/dataclasses.py @@ -120,11 +120,13 @@ def __init_subclass__(cls, **kwargs): for a in ancestors for name in a._SLOTS } - if not hasattr(cls, "_SLOT_ANNOTATIONS") or cls._SLOT_ANNOTATIONS is None: - cls._SLOT_ANNOTATIONS = {} - cls._SLOTS = () - else: - cls._SLOT_ANNOTATIONS = dict(cls._SLOT_ANNOTATIONS) + # Collect the inherited slots from *all* data-class ancestors, in reverse-MRO order. + # Reading the inherited `_SLOT_ANNOTATIONS`/`_SLOTS` attributes instead would follow + # only the first inheritance chain, silently dropping the slots of any additional bases. + inherited_slot_annotations: dict[str, type[TreeNode]] = {} + for ancestor in reversed(ancestors): + inherited_slot_annotations.update(ancestor._SLOT_ANNOTATIONS) + cls._SLOT_ANNOTATIONS = inherited_slot_annotations new_slots = [] for name, slot_type in cls.__annotations__.items(): # get_origin() screens out subscripted generics before issubclass() sees them. On Python @@ -138,7 +140,7 @@ def __init_subclass__(cls, **kwargs): f"defined in its superclass {ancestor_slot_names[name].__name__}") new_slots.append(name) cls._SLOT_ANNOTATIONS[name] = slot_type - cls._SLOTS = cls._SLOTS + tuple(new_slots) + cls._SLOTS = tuple(cls._SLOT_ANNOTATIONS) def __hash__(self): return self.__hash diff --git a/test/test_dataclasses.py b/test/test_dataclasses.py index bd60df8..7ed712f 100644 --- a/test/test_dataclasses.py +++ b/test/test_dataclasses.py @@ -91,6 +91,41 @@ def post_init(self): self.name.quoted = False self.assertFalse(Unquoted(StringNode("name")).name.quoted) + + def test_multiple_inheritance(self): + """Slots from *all* bases must survive diamond inheritance, not just the first chain.""" + class Foo(DataClassNode): + foo: IntegerNode + + class Bar(Foo): + bar: StringNode + + class Baz(Foo): + baz: StringNode + + class Quux(Bar, Baz): + quux: IntegerNode + + self.assertEqual(("foo",), Foo._SLOTS) + self.assertEqual(("foo", "bar"), Bar._SLOTS) + self.assertEqual(("foo", "baz"), Baz._SLOTS) + self.assertEqual(("foo", "baz", "bar", "quux"), Quux._SLOTS) + self.assertEqual( + {"foo": IntegerNode, "baz": StringNode, "bar": StringNode, "quux": IntegerNode}, + Quux._SLOT_ANNOTATIONS + ) + + node = Quux(foo=IntegerNode(1), bar=StringNode("bar"), baz=StringNode("baz"), quux=IntegerNode(4)) + self.assertEqual(1, node.foo.object) + self.assertEqual("bar", node.bar.object) + self.assertEqual("baz", node.baz.object) + self.assertEqual(4, node.quux.object) + self.assertEqual({"foo", "bar", "baz", "quux"}, set(node.to_obj())) + + # diffing against an identical node yields a DataClassEdit, not a Replace + twin = Quux(foo=IntegerNode(1), bar=StringNode("bar"), baz=StringNode("baz"), quux=IntegerNode(4)) + self.assertIsInstance(node.edits(twin), DataClassEdit) + def test_print_renders_slots(self): """:meth:`DataClassNode.print` is the fallback when no formatter resolves the node type.""" class Foo(DataClassNode):