Construct a sample tree with the rebalance implementation missing

This commit is contained in:
2026-06-11 12:36:53 -04:00
parent 0fd98a46ca
commit 14c14d1797
2 changed files with 183 additions and 0 deletions

View File

@@ -0,0 +1,5 @@
## Outline
- Write a red-black tree with rebalancing
- Assignment: finish the red-black tree
- Next week: introduce graphs and graph algorithms

View File

@@ -0,0 +1,178 @@
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 "<empty tree>"
# 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())