diff --git a/docs/builders.rst b/docs/builders.rst index 68699ad..26d655a 100644 --- a/docs/builders.rst +++ b/docs/builders.rst @@ -160,7 +160,8 @@ Initializing a Data Class Node :meth:`DataClassNode.__init__ ` assigns the slots from its positional and keyword arguments, so overriding it means reimplementing that assignment. Override :meth:`graphtage.dataclasses.DataClassNode.post_init` instead. It is called once the slots have been assigned, and it -should not call ``super().post_init()``: each ancestor's implementation is called in turn, in order of the ``__mro__``. +should not call ``super().post_init()``: every implementation in the class hierarchy is called automatically, +starting with the least derived data class and ending with the class being instantiated. .. code-block:: python @@ -171,7 +172,4 @@ should not call ``super().post_init()``: each ancestor's implementation is calle self.name.quoted = False .. note:: - As of Graphtage 0.3.1, ``post_init()`` is only called for the ancestors of the class being instantiated, never for - the class itself. ``UnquotedName(StringNode("x"))`` leaves ``quoted`` set to :const:`True`; the callback runs only - when a subclass of ``UnquotedName`` is instantiated. Code that must run for the class itself still has to go in - ``__init__``. + An implementation that a subclass inherits without overriding runs once, not once per class that inherits it. diff --git a/graphtage/dataclasses.py b/graphtage/dataclasses.py index 3458eb0..b8d94db 100644 --- a/graphtage/dataclasses.py +++ b/graphtage/dataclasses.py @@ -1,4 +1,4 @@ -from collections.abc import Iterator +from collections.abc import Callable, Iterator from typing import get_origin from . import AbstractCompoundEdit, Edit, Range, Replace @@ -51,6 +51,7 @@ class DataClassNode(ContainerNode): _SLOTS: tuple[str, ...] _SLOT_ANNOTATIONS: dict[str, type[TreeNode]] _DATA_CLASS_ANCESTORS: list[type["DataClassNode"]] + _POST_INITS: tuple[Callable[["DataClassNode"], None], ...] def __init__(self, *args, **kwargs): """Be careful extending __init__; consider using :func:`DataClassNode.post_init` instead.""" @@ -88,14 +89,15 @@ def __init__(self, *args, **kwargs): setattr(self, s, value) # self.__hash__ gets called so often, we cache the result: self.__hash = hash(tuple(self)) - for ancestor in self._DATA_CLASS_ANCESTORS: - ancestor.post_init(self) + for post_init in self._POST_INITS: + post_init(self) def post_init(self): - """Callback called after this class's members have been initialized. + """Callback called after this node's slots have been initialized. - This callback should not call `super().post_init()`. Each superclass's `post_init()` will be automatically - called in order of the `__mro__`. + This callback should not call `super().post_init()`. Every implementation in the class hierarchy is called + automatically, starting with the least derived data class and ending with the class being instantiated. An + implementation that a subclass inherits without overriding is called only once. """ pass @@ -103,10 +105,16 @@ def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) ancestors = [ c - for c in cls.__mro__ + for c in reversed(cls.__mro__) if c is not cls and issubclass(c, DataClassNode) and c is not DataClassNode ] cls._DATA_CLASS_ANCESTORS = ancestors + # Selecting on __dict__ keeps an inherited implementation from being called once per class that inherits it. + cls._POST_INITS = tuple( + c.__dict__["post_init"] + for c in (*ancestors, cls) + if "post_init" in c.__dict__ + ) ancestor_slot_names = { name: a for a in ancestors diff --git a/test/test_dataclasses.py b/test/test_dataclasses.py index 217ae27..bd60df8 100644 --- a/test/test_dataclasses.py +++ b/test/test_dataclasses.py @@ -10,17 +10,17 @@ class TestDataclasses(TestCase): def test_inheritance(self): class Foo(DataClassNode): foo: IntegerNode - initialized = False + foo_initialized = False def post_init(self): - self.initialized = True + self.foo_initialized = True class Bar(Foo): bar: StringNode - initialized = False + bar_initialized = False def post_init(self): - self.initialized = True + self.bar_initialized = True self.assertEqual(("foo",), Foo._SLOTS) self.assertEqual(0, len(Foo._DATA_CLASS_ANCESTORS)) @@ -30,13 +30,15 @@ def post_init(self): b = Bar(foo=IntegerNode(10), bar=StringNode("bar")) self.assertEqual(10, b.foo.object) self.assertEqual("bar", b.bar.object) - self.assertTrue(b.initialized) + self.assertTrue(b.foo_initialized) + self.assertTrue(b.bar_initialized) # now test a mixture of positional and keyword arguments b = Bar(StringNode("bar"), foo=IntegerNode(10)) self.assertEqual(10, b.foo.object) self.assertEqual("bar", b.bar.object) - self.assertTrue(b.initialized) + self.assertTrue(b.foo_initialized) + self.assertTrue(b.bar_initialized) # test equality self.assertEqual(Bar(IntegerNode(10), StringNode("bar")), b) @@ -50,6 +52,45 @@ def post_init(self): edit = f.edits(c) self.assertIsInstance(edit, DataClassEdit) + def test_post_init_runs_once_per_implementation(self): + calls: list[tuple[str, str]] = [] + + class Base(DataClassNode): + base: IntegerNode + + def post_init(self): + calls.append(("Base", type(self).__name__)) + + class Middle(Base): + middle: StringNode + + class Derived(Middle): + derived: IntegerNode + + def post_init(self): + calls.append(("Derived", type(self).__name__)) + + Base(IntegerNode(1)) + self.assertEqual([("Base", "Base")], calls) + + # Middle inherits Base.post_init without overriding it, so it must run exactly once + calls.clear() + Middle(IntegerNode(1), StringNode("middle")) + self.assertEqual([("Base", "Middle")], calls) + + # each implementation runs once, least derived first + calls.clear() + Derived(IntegerNode(1), StringNode("middle"), IntegerNode(2)) + self.assertEqual([("Base", "Derived"), ("Derived", "Derived")], calls) + + def test_post_init_runs_for_direct_subclass(self): + class Unquoted(DataClassNode): + name: StringNode + + def post_init(self): + self.name.quoted = False + + self.assertFalse(Unquoted(StringNode("name")).name.quoted) def test_print_renders_slots(self): """:meth:`DataClassNode.print` is the fallback when no formatter resolves the node type.""" class Foo(DataClassNode):