diff --git a/docs/source/index.rst b/docs/source/index.rst index c85e367..14c3a1c 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -703,8 +703,8 @@ Extend the Node class for specialized functionality: from treelib import Node, Tree class FileNode(Node): - def __init__(self, tag, identifier=None, size=0, file_type="unknown"): - super().__init__(tag, identifier) + def __init__(self, tag, identifier=None, size=0, file_type="unknown", **kwargs): + super().__init__(tag, identifier, **kwargs) self.size = size self.file_type = file_type diff --git a/tests/test_tree.py b/tests/test_tree.py index 79eeeb3..aacea09 100644 --- a/tests/test_tree.py +++ b/tests/test_tree.py @@ -711,6 +711,18 @@ class SubTree(Tree): node = tree.create_node() self.assertTrue(isinstance(node, SubNode)) + def test_create_node_forwards_kwargs_to_node_class(self): + class FileNode(Node): + def __init__(self, tag=None, identifier=None, size=0, **kwargs): + super().__init__(tag=tag, identifier=identifier, **kwargs) + self.size = size + + tree = Tree(node_class=FileNode) + node = tree.create_node("Large File", "big_file", size=1024) + self.assertTrue(isinstance(node, FileNode)) + self.assertEqual(node.size, 1024) + self.assertEqual(tree["big_file"].size, 1024) + def test_shallow_copy_hermetic_pointers(self): # tree 1 # Hárry diff --git a/treelib/tree.py b/treelib/tree.py index f9baa50..ef6935e 100644 --- a/treelib/tree.py +++ b/treelib/tree.py @@ -853,6 +853,7 @@ def create_node( identifier: Optional[str] = None, parent: Optional[Union[Node, str]] = None, data: Any = None, + **kwargs: Any, ) -> Node: """ Create and add a new node to the tree. @@ -867,6 +868,8 @@ def create_node( parent: Parent node identifier, Node object, or None for root. Must be None if tree is empty, must exist if tree has nodes. data: Optional user data to associate with this node. + **kwargs: Extra keyword arguments forwarded to the tree's + ``node_class`` constructor, for custom Node subclasses. Returns: Node: The newly created Node object. @@ -896,7 +899,7 @@ def create_node( tree.create_node("Baby", "baby", parent="me", data={"age": 1, "cute": True}) """ - node = self.node_class(tag=tag, identifier=identifier, data=data) + node = self.node_class(tag=tag, identifier=identifier, data=data, **kwargs) self.add_node(node, parent) return node