From 174e10f24777d1a41371145507943e91f47d5dc9 Mon Sep 17 00:00:00 2001 From: Priya Sundaram Date: Thu, 10 Sep 2026 08:12:06 +0000 Subject: [PATCH 1/2] Add splay tree (self-adjusting BST) Implements a splay tree with insert, search and delete, each of which splays the accessed node to the root for amortized O(log n) access and locality of reference. Includes full type hints and doctests. Closes #13760 --- data_structures/binary_tree/splay_tree.py | 231 ++++++++++++++++++++++ 1 file changed, 231 insertions(+) create mode 100644 data_structures/binary_tree/splay_tree.py diff --git a/data_structures/binary_tree/splay_tree.py b/data_structures/binary_tree/splay_tree.py new file mode 100644 index 000000000000..e8ca0ff02cb1 --- /dev/null +++ b/data_structures/binary_tree/splay_tree.py @@ -0,0 +1,231 @@ +""" +Splay Tree - a self-adjusting binary search tree. + +A splay tree is a binary search tree with the additional property that +recently accessed elements are quick to access again. Every access (search, +insert or delete) moves the target node to the root through a sequence of +rotations called "splaying". This gives an amortized time complexity of +O(log n) per operation and makes the tree very efficient when the access +pattern has locality of reference (a small subset of keys is touched often). + +Reference: https://en.wikipedia.org/wiki/Splay_tree +""" + +from __future__ import annotations + +from collections.abc import Iterator + + +class Node: + """A single node of a splay tree.""" + + def __init__(self, key: int) -> None: + self.key = key + self.left: Node | None = None + self.right: Node | None = None + + def __repr__(self) -> str: + return f"Node({self.key})" + + +class SplayTree: + """ + A self-adjusting binary search tree. + + >>> tree = SplayTree() + >>> tree.insert(10) + >>> tree.insert(20) + >>> tree.insert(30) + >>> tree.root.key # last inserted key is splayed to the root + 30 + >>> tree.search(10) + True + >>> tree.root.key # the searched key is now the root + 10 + >>> tree.search(99) + False + >>> list(tree) + [10, 20, 30] + """ + + def __init__(self) -> None: + self.root: Node | None = None + + def _rotate_right(self, node: Node) -> Node: + """ + Perform a right rotation around ``node`` and return the new subtree root. + + node left + / \\ / \\ + left c --> a node + / \\ / \\ + a b b c + """ + left = node.left + assert left is not None + node.left = left.right + left.right = node + return left + + def _rotate_left(self, node: Node) -> Node: + """ + Perform a left rotation around ``node`` and return the new subtree root. + + node right + / \\ / \\ + a right --> node c + / \\ / \\ + b c a b + """ + right = node.right + assert right is not None + node.right = right.left + right.left = node + return right + + def _splay(self, root: Node | None, key: int) -> Node | None: + """ + Splay the node with ``key`` (or the last node on the search path if + ``key`` is absent) to the root of the subtree and return the new root. + This uses the classic bottom-up recursive formulation. + """ + if root is None or root.key == key: + return root + + if key < root.key: + if root.left is None: + return root + if key < root.left.key: + # Zig-Zig (left left) + root.left.left = self._splay(root.left.left, key) + root = self._rotate_right(root) + elif key > root.left.key: + # Zig-Zag (left right) + root.left.right = self._splay(root.left.right, key) + if root.left.right is not None: + root.left = self._rotate_left(root.left) + return root if root.left is None else self._rotate_right(root) + else: + if root.right is None: + return root + if key > root.right.key: + # Zig-Zig (right right) + root.right.right = self._splay(root.right.right, key) + root = self._rotate_left(root) + elif key < root.right.key: + # Zig-Zag (right left) + root.right.left = self._splay(root.right.left, key) + if root.right.left is not None: + root.right = self._rotate_right(root.right) + return root if root.right is None else self._rotate_left(root) + + def insert(self, key: int) -> None: + """ + Insert ``key`` into the tree and splay it to the root. + + >>> tree = SplayTree() + >>> for key in (5, 3, 8, 3): # duplicate keys are ignored + ... tree.insert(key) + >>> list(tree) + [3, 5, 8] + >>> tree.root.key # the duplicate access splays 3 back to the root + 3 + """ + if self.root is None: + self.root = Node(key) + return + + self.root = self._splay(self.root, key) + assert self.root is not None + if self.root.key == key: + return # key already present, it is now at the root + + node = Node(key) + if key < self.root.key: + node.right = self.root + node.left = self.root.left + self.root.left = None + else: + node.left = self.root + node.right = self.root.right + self.root.right = None + self.root = node + + def search(self, key: int) -> bool: + """ + Return whether ``key`` is present and splay the last accessed node. + + >>> tree = SplayTree() + >>> tree.search(1) + False + >>> for key in (40, 20, 60): + ... tree.insert(key) + >>> tree.search(20) + True + >>> tree.root.key + 20 + """ + self.root = self._splay(self.root, key) + return self.root is not None and self.root.key == key + + def delete(self, key: int) -> None: + """ + Remove ``key`` from the tree if it is present. + + >>> tree = SplayTree() + >>> for key in (10, 20, 30, 40): + ... tree.insert(key) + >>> tree.delete(20) + >>> list(tree) + [10, 30, 40] + >>> tree.delete(99) # deleting an absent key is a no-op + >>> list(tree) + [10, 30, 40] + >>> for key in (10, 30, 40): + ... tree.delete(key) + >>> list(tree) + [] + """ + if self.root is None: + return + + self.root = self._splay(self.root, key) + assert self.root is not None + if self.root.key != key: + return # key not found + + left, right = self.root.left, self.root.right + if left is None: + self.root = right + else: + # Splay the maximum of the left subtree to its root; it has no + # right child, so the right subtree can be attached there. + left = self._splay(left, key) + assert left is not None + left.right = right + self.root = left + + def __iter__(self) -> Iterator[int]: + """ + Yield the keys of the tree in ascending (in-order) order. + + >>> tree = SplayTree() + >>> for key in (7, 2, 9, 4, 1): + ... tree.insert(key) + >>> list(tree) + [1, 2, 4, 7, 9] + """ + + def in_order(node: Node | None) -> Iterator[int]: + if node is not None: + yield from in_order(node.left) + yield node.key + yield from in_order(node.right) + + yield from in_order(self.root) + + +if __name__ == "__main__": + import doctest + + doctest.testmod() From 4d86ab2b6479d9dd9c8a1db9298c180515a53523 Mon Sep 17 00:00:00 2001 From: Priya Sundaram Date: Thu, 10 Sep 2026 08:42:36 +0000 Subject: [PATCH 2/2] Make Node a dataclass and add repr doctest (review feedback) --- data_structures/binary_tree/splay_tree.py | 20 +++++++++++++------- 1 file changed, 13 insertions(+), 7 deletions(-) diff --git a/data_structures/binary_tree/splay_tree.py b/data_structures/binary_tree/splay_tree.py index e8ca0ff02cb1..c200bad761a5 100644 --- a/data_structures/binary_tree/splay_tree.py +++ b/data_structures/binary_tree/splay_tree.py @@ -14,18 +14,24 @@ from __future__ import annotations from collections.abc import Iterator +from dataclasses import dataclass, field +@dataclass class Node: - """A single node of a splay tree.""" + """ + A single node of a splay tree. + + The ``left`` and ``right`` children are excluded from ``repr`` so that a + node prints compactly instead of recursively dumping the whole subtree. - def __init__(self, key: int) -> None: - self.key = key - self.left: Node | None = None - self.right: Node | None = None + >>> Node(10) + Node(key=10) + """ - def __repr__(self) -> str: - return f"Node({self.key})" + key: int + left: Node | None = field(default=None, repr=False) + right: Node | None = field(default=None, repr=False) class SplayTree: