class Node: def __init__(self, value, left=None, right=None, parent=None): self.value = value self.left: Node | None = left self.right: Node | None = right # required to re-set the child later self.parent: Node | None = parent class RedBlackTree: def __init__(self): self.root: Node | None = None def insert(self, value): # assumption: inserting the same value twice will # insert nothing the second time node = Node(value) if self.root is None: self.root = node return # insert like a normal binary search tree prev: Node | None = None current = self.root while current: if current.value == value: return prev = current if current.value > value: current = current.left elif current.value < value: current = current.right if prev and prev.value > value: prev.left = node elif prev and prev.value < value: prev.right = node def delete(self, value, root=None, parent=None): # does nothing if the value doesn't exist if parent is None: root = self.root if root is None: return if root.value == value: new_root: Node | None if root.left is None: new_root = root.right elif root.right is None: new_root = root.left else: # both left and right subtrees exist # make left subtree root the tree root, and re-attach # the right subtree right_subtree_root = root.right new_root = root.left prev = None current = new_root while current: prev = current current = current.right prev.right = right_subtree_root if parent is None: self.root = new_root else: if parent.value > root.value: parent.left = new_root else: parent.right = new_root elif root.value > value: self.delete(value, root.left, root) elif root.value < value: self.delete(value, root.right, root) def rebalance(self): # this is what we'll be implementing pass def visualize(self) -> str: # renders the tree top-down with the root on top and branches # ( / and \ ) drawn down to each child. spacing is computed so that # subtrees never overlap, no matter the shape of the tree. if self.root is None: return "" # render() returns, for the subtree rooted at `node`: # lines - the block of text drawing the subtree # width - how many characters wide that block is # height - how many lines tall that block is # middle - the column where this node's value is centred (so the # caller knows where to attach its branch) def render(node: Node) -> tuple[list[str], int, int, int]: label = str(node.value) label_width = len(label) # leaf: just the value on a single line if node.left is None and node.right is None: return [label], label_width, 1, label_width // 2 # only a left child if node.right is None: lines, width, height, mid = render(node.left) first = (mid + 1) * " " + (width - mid - 1) * "_" + label second = mid * " " + "/" + (width - mid - 1 + label_width) * " " shifted = [line + label_width * " " for line in lines] return [first, second] + shifted, width + label_width, height + 2, width + label_width // 2 # only a right child if node.left is None: lines, width, height, mid = render(node.right) first = label + mid * "_" + (width - mid) * " " second = (label_width + mid) * " " + "\\" + (width - mid - 1) * " " shifted = [label_width * " " + line for line in lines] return [first, second] + shifted, width + label_width, height + 2, label_width // 2 # two children: render each side, then place this node between them left_lines, left_w, left_h, left_mid = render(node.left) right_lines, right_w, right_h, right_mid = render(node.right) first = ( (left_mid + 1) * " " + (left_w - left_mid - 1) * "_" + label + right_mid * "_" + (right_w - right_mid) * " " ) second = ( left_mid * " " + "/" + (left_w - left_mid - 1 + label_width + right_mid) * " " + "\\" + (right_w - right_mid - 1) * " " ) # pad the shorter side so the two blocks line up row-for-row if left_h < right_h: left_lines += [left_w * " "] * (right_h - left_h) elif right_h < left_h: right_lines += [right_w * " "] * (left_h - right_h) merged = [l + label_width * " " + r for l, r in zip(left_lines, right_lines)] return ( [first, second] + merged, left_w + right_w + label_width, max(left_h, right_h) + 2, left_w + label_width // 2, ) lines, _, _, _ = render(self.root) return "\n".join(lines) # Sample visualizations tree = RedBlackTree() tree.insert(20) tree.insert(10) tree.insert(5) tree.insert(30) tree.insert(40) print(tree.visualize()) tree.delete(20) print(tree.visualize()) tree.delete(40) tree.delete(5) print(tree.visualize()) tree.delete(10) print(tree.visualize()) tree.delete(30) print(tree.visualize()) # attempt duplicate delete tree.delete(30) print(tree.visualize())