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
8 changes: 3 additions & 5 deletions docs/builders.rst
Original file line number Diff line number Diff line change
Expand Up @@ -160,7 +160,8 @@ Initializing a Data Class Node
:meth:`DataClassNode.__init__ <graphtage.dataclasses.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

Expand All @@ -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.
22 changes: 15 additions & 7 deletions graphtage/dataclasses.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -88,25 +89,32 @@ 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

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
Expand Down
53 changes: 47 additions & 6 deletions test/test_dataclasses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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)
Expand All @@ -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):
Expand Down