Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 8 additions & 6 deletions graphtage/dataclasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
35 changes: 35 additions & 0 deletions test/test_dataclasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading