Trees and Binary Search Trees
Trees are the most common data structure in interviews after arrays, and nearly every tree problem is solved by recursion: solve the problem for the left subtree, solve it for the right subtree, and combine the two answers at the current node. Once you see that shape, a large family of problems becomes the same problem. This chapter teaches that shape, the traversals, and the patterns that build on them, including binary search trees.
1. The structure
class TreeNode:
def __init__(self, val=0, left=None, right=None):
self.val = val
self.left = left
self.right = right
def build_tree(values):
"""Build from a level-order list with None for missing children, e.g. [3, 9, 20, None, None, 15, 7]."""
if not values or values[0] is None:
return None
root = TreeNode(values[0])
queue = [root]
i = 1
for node in queue:
if i < len(values) and values[i] is not None:
node.left = TreeNode(values[i]); queue.append(node.left)
i += 1
if i < len(values) and values[i] is not None:
node.right = TreeNode(values[i]); queue.append(node.right)
i += 1
return root
root = build_tree([3, 9, 20, None, None, 15, 7])
assert root.val == 3 and root.left.val == 9 and root.right.right.val == 7
Vocabulary to use precisely: root, leaf (no children), depth of a node (edges from the root), height of a node (edges to its deepest leaf), subtree, balanced (heights of subtrees differ by at most one), complete and full. In a binary search tree (BST), everything in the left subtree is smaller than the node, and everything in the right subtree is larger.
2. The recursive template
For most problems, write a function that handles one node by trusting the recursion for its children.
def solve(node):
if node is None: # base case
return <answer for an empty tree>
left = solve(node.left)
right = solve(node.right)
return <combine node.val, left, right>
Decide three things: the base case, what each call returns, and how to combine. Recursion depth equals the tree height: for a balanced tree and for a skewed one, so the space is .
Maximum depth.
def max_depth(node):
if not node:
return 0
return 1 + max(max_depth(node.left), max_depth(node.right))
assert max_depth(root) == 3
assert max_depth(None) == 0
Same tree, invert tree, symmetric tree. All are the template with a different combine step.
def is_same_tree(a, b):
if not a and not b:
return True
if not a or not b or a.val != b.val:
return False
return is_same_tree(a.left, b.left) and is_same_tree(a.right, b.right)
def invert_tree(node):
if node:
node.left, node.right = invert_tree(node.right), invert_tree(node.left)
return node
def is_symmetric(node):
def mirror(a, b):
if not a and not b:
return True
if not a or not b or a.val != b.val:
return False
return mirror(a.left, b.right) and mirror(a.right, b.left)
return mirror(node.left, node.right) if node else True
assert is_same_tree(build_tree([1, 2, 3]), build_tree([1, 2, 3])) is True
assert is_same_tree(build_tree([1, 2]), build_tree([1, None, 2])) is False
assert is_symmetric(build_tree([1, 2, 2, 3, 4, 4, 3])) is True
assert is_symmetric(build_tree([1, 2, 2, None, 3, None, 3])) is False
3. Traversals
There are four standard orders. Know them recursively and iteratively.
<!--fig:traversals-->- Preorder: node, left, right. Copies a tree; serialises it.
- Inorder: left, node, right. On a BST, visits values in sorted order.
- Postorder: left, right, node. Children before parent; used to compute heights, delete, evaluate.
- Level order: breadth-first, row by row, using a queue.
def preorder(node, out=None):
out = [] if out is None else out
if node:
out.append(node.val); preorder(node.left, out); preorder(node.right, out)
return out
def inorder(node, out=None):
out = [] if out is None else out
if node:
inorder(node.left, out); out.append(node.val); inorder(node.right, out)
return out
def postorder(node, out=None):
out = [] if out is None else out
if node:
postorder(node.left, out); postorder(node.right, out); out.append(node.val)
return out
t = build_tree([4, 2, 6, 1, 3, 5, 7])
assert preorder(t) == [4, 2, 1, 3, 6, 5, 7]
assert inorder(t) == [1, 2, 3, 4, 5, 6, 7]
assert postorder(t) == [1, 3, 2, 5, 7, 6, 4]
Iterative inorder with an explicit stack avoids recursion limits:
def inorder_iterative(node):
out, stack = [], []
while node or stack:
while node:
stack.append(node)
node = node.left
node = stack.pop()
out.append(node.val)
node = node.right
return out
assert inorder_iterative(t) == [1, 2, 3, 4, 5, 6, 7]
Level order (BFS) uses a queue and processes one level at a time by taking the current queue length as the level size.
from collections import deque
def level_order(root):
if not root:
return []
result, queue = [], deque([root])
while queue:
level = []
for _ in range(len(queue)): # exactly the nodes of this level
node = queue.popleft()
level.append(node.val)
if node.left: queue.append(node.left)
if node.right: queue.append(node.right)
result.append(level)
return result
assert level_order(build_tree([3, 9, 20, None, None, 15, 7])) == [[3], [9, 20], [15, 7]]
assert level_order(None) == []
The same skeleton gives right side view (last node of each level), zigzag level order (reverse alternate levels), average of each level and minimum depth.
def right_side_view(root):
return [level[-1] for level in level_order(root)]
assert right_side_view(build_tree([1, 2, 3, None, 5, None, 4])) == [1, 3, 4]
4. Computing values bottom-up
Some answers need information from both subtrees, and the best answer may pass through a node without being returned to the parent. Handle that by keeping a separate variable for the best answer, and returning to the parent only what the parent can extend.
Diameter of a binary tree (the longest path between any two nodes, in edges). At each node, the longest path through it is left_height + right_height. Return the node's height to its parent.
def diameter(root):
best = 0
def height(node):
nonlocal best
if not node:
return 0
left, right = height(node.left), height(node.right)
best = max(best, left + right) # path through this node
return 1 + max(left, right) # height offered to the parent
height(root)
return best
assert diameter(build_tree([1, 2, 3, 4, 5])) == 3
assert diameter(build_tree([1, 2])) == 1
Balanced tree check. Return the height, and use a sentinel (here -1) to signal "unbalanced" so the check short-circuits in .
def is_balanced(root):
def height(node):
if not node:
return 0
left = height(node.left)
if left == -1:
return -1
right = height(node.right)
if right == -1 or abs(left - right) > 1:
return -1
return 1 + max(left, right)
return height(root) != -1
assert is_balanced(build_tree([3, 9, 20, None, None, 15, 7])) is True
assert is_balanced(build_tree([1, 2, 2, 3, 3, None, None, 4, 4])) is False
Maximum path sum (any path, node values may be negative) uses the same shape: at each node, the best path through it is node.val + max(left, 0) + max(right, 0), and the parent gets node.val + max(left, right, 0).
def max_path_sum(root):
best = float("-inf")
def gain(node):
nonlocal best
if not node:
return 0
left, right = max(gain(node.left), 0), max(gain(node.right), 0)
best = max(best, node.val + left + right)
return node.val + max(left, right)
gain(root)
return best
assert max_path_sum(build_tree([1, 2, 3])) == 6
assert max_path_sum(build_tree([-10, 9, 20, None, None, 15, 7])) == 42
assert max_path_sum(build_tree([-3])) == -3
5. Passing information down
When a node needs context from above, pass it as a parameter.
Path sum (does a root-to-leaf path add up to a target?): pass the remaining target down.
def has_path_sum(node, target):
if not node:
return False
if not node.left and not node.right:
return node.val == target
return has_path_sum(node.left, target - node.val) or has_path_sum(node.right, target - node.val)
assert has_path_sum(build_tree([5, 4, 8, 11, None, 13, 4, 7, 2, None, None, None, 1]), 22) is True
assert has_path_sum(build_tree([1, 2, 3]), 5) is False
assert has_path_sum(None, 0) is False
Good nodes (a node is good if no ancestor is larger): pass down the maximum so far.
def good_nodes(root):
def dfs(node, highest):
if not node:
return 0
count = 1 if node.val >= highest else 0
highest = max(highest, node.val)
return count + dfs(node.left, highest) + dfs(node.right, highest)
return dfs(root, float("-inf"))
assert good_nodes(build_tree([3, 1, 4, 3, None, 1, 5])) == 4
6. Lowest common ancestor
The LCA of two nodes is the deepest node that has both as descendants. In a general binary tree, recurse: if the current node is one of the targets, it is the answer for its subtree. If the two targets are found in different subtrees, the current node is the LCA.
def lowest_common_ancestor(root, p, q):
if not root or root is p or root is q:
return root
left = lowest_common_ancestor(root.left, p, q)
right = lowest_common_ancestor(root.right, p, q)
if left and right:
return root # p and q are on different sides
return left or right
r = build_tree([3, 5, 1, 6, 2, 0, 8, None, None, 7, 4])
p, q = r.left, r.right
assert lowest_common_ancestor(r, p, q) is r
assert lowest_common_ancestor(r, p, p.right.right) is p # a node can be its own ancestor
In a BST you can do better, using the ordering: if both values are smaller than the node, go left, if both are larger go right, otherwise the node is the LCA. That is and needs no recursion.
7. Binary search trees
The BST property gives search, insert and delete, and sorted order through inorder traversal.
def search_bst(node, target):
while node and node.val != target:
node = node.left if target < node.val else node.right
return node
def insert_bst(node, val):
if not node:
return TreeNode(val)
if val < node.val:
node.left = insert_bst(node.left, val)
else:
node.right = insert_bst(node.right, val)
return node
bst = None
for v in [5, 3, 8, 1, 4, 7, 9]:
bst = insert_bst(bst, v)
assert inorder(bst) == [1, 3, 4, 5, 7, 8, 9]
assert search_bst(bst, 4).val == 4 and search_bst(bst, 6) is None
Validate a BST. A common mistake is to check only that each node is larger than its left child and smaller than its right child. That misses cases where a deeper node violates an ancestor's bound. Pass down the allowed range.
def is_valid_bst(node, low=float("-inf"), high=float("inf")):
if not node:
return True
if not (low < node.val < high):
return False
return is_valid_bst(node.left, low, node.val) and is_valid_bst(node.right, node.val, high)
assert is_valid_bst(build_tree([2, 1, 3])) is True
assert is_valid_bst(build_tree([5, 1, 4, None, None, 3, 6])) is False
assert is_valid_bst(build_tree([5, 4, 6, None, None, 3, 7])) is False # 3 is in the right subtree of 5
Equivalent test: the inorder traversal is strictly increasing.
Kth smallest in a BST. Inorder visits in sorted order, so stop at the th visit. The iterative version stops early.
def kth_smallest(root, k):
stack, node = [], root
while node or stack:
while node:
stack.append(node)
node = node.left
node = stack.pop()
k -= 1
if k == 0:
return node.val
node = node.right
assert kth_smallest(build_tree([3, 1, 4, None, 2]), 1) == 1
assert kth_smallest(build_tree([5, 3, 6, 2, 4, None, None, 1]), 3) == 3
Delete a node from a BST has three cases: a leaf (remove it), one child (replace it with the child), two children (replace its value with the inorder successor, the smallest value in the right subtree, then delete that successor).
def delete_node(root, key):
if not root:
return None
if key < root.val:
root.left = delete_node(root.left, key)
elif key > root.val:
root.right = delete_node(root.right, key)
else:
if not root.left:
return root.right
if not root.right:
return root.left
successor = root.right
while successor.left:
successor = successor.left
root.val = successor.val
root.right = delete_node(root.right, successor.val)
return root
b = None
for v in [5, 3, 6, 2, 4, 7]:
b = insert_bst(b, v)
b = delete_node(b, 3)
assert inorder(b) == [2, 4, 5, 6, 7]
Convert a sorted array into a height-balanced BST by choosing the middle element as the root and recursing on each half.
def sorted_array_to_bst(nums):
if not nums:
return None
mid = len(nums) // 2
return TreeNode(nums[mid], sorted_array_to_bst(nums[:mid]), sorted_array_to_bst(nums[mid + 1:]))
assert inorder(sorted_array_to_bst([-10, -3, 0, 5, 9])) == [-10, -3, 0, 5, 9]
assert max_depth(sorted_array_to_bst(list(range(15)))) == 4
Plain BSTs can become skewed (inserting sorted data gives a linked list with operations). Self-balancing trees (AVL, red-black) guarantee . In interviews, know that they exist, how they rebalance in outline (rotations), and that standard library ordered maps use them, without implementing one unless asked.
8. Constructing and serialising trees
Build a tree from preorder and inorder traversals. The first preorder element is the root. Its position in the inorder list splits the left and right subtrees.
def build_from_preorder_inorder(preorder_list, inorder_list):
index = {v: i for i, v in enumerate(inorder_list)} # value -> inorder position
pre = iter(preorder_list)
def build(lo, hi):
if lo > hi:
return None
root = TreeNode(next(pre))
mid = index[root.val]
root.left = build(lo, mid - 1) # build left first: preorder order
root.right = build(mid + 1, hi)
return root
return build(0, len(inorder_list) - 1)
tree = build_from_preorder_inorder([3, 9, 20, 15, 7], [9, 3, 15, 20, 7])
assert level_order(tree) == [[3], [9, 20], [15, 7]]
The hash map makes it instead of .
Serialise and deserialise a binary tree with preorder and a marker for empty children.
def serialize(root):
out = []
def dfs(node):
if not node:
out.append("#")
return
out.append(str(node.val))
dfs(node.left)
dfs(node.right)
dfs(root)
return ",".join(out)
def deserialize(data):
tokens = iter(data.split(","))
def build():
tok = next(tokens)
if tok == "#":
return None
node = TreeNode(int(tok))
node.left = build()
node.right = build()
return node
return build()
tr = build_tree([1, 2, 3, None, None, 4, 5])
s = serialize(tr)
assert level_order(deserialize(s)) == level_order(tr)
assert deserialize("#") is None
9. Choosing among the techniques
| Signal in the problem | Technique |
|---|---|
| Property of the whole tree built from children | Postorder style recursion returning a value |
| Best path may pass through a node | Track a global best, return an extendable value |
| Needs information from ancestors | Pass a parameter down |
| Level by level, shortest depth | BFS with a queue |
| BST: order, th, range | Inorder traversal, or use the range bounds |
| Two nodes, ancestor | LCA recursion, or the BST bounds shortcut |
| Reconstruct from traversals | Hash map from value to inorder index |
| Avoid recursion depth issues | Explicit stack, or Morris traversal ( space) |
10. Common mistakes
- Forgetting the base case (
if not node). Most crashes are an attribute access onNone. - Validating a BST locally (only comparing with the children). Pass bounds.
- Mixing up depth, height and the number of nodes versus edges. State your convention.
- Not distinguishing a leaf from a node with one child in path problems. A root-to-leaf path ends at a node with no children.
- Returning the wrong thing in bottom-up problems: the answer through the node versus the value the parent can extend.
- Space analysis that ignores the recursion stack. It is .
- Assuming a balanced tree. State worst-case height when it is not guaranteed.
- Modifying a tree while traversing it without care for what the recursion still points to.
11. Practice set
- Maximum depth, minimum depth, invert tree, same tree, subtree of another tree.
- Level order, zigzag level order, right side view, binary tree from level order.
- Diameter, balanced tree, binary tree maximum path sum.
- Path sum I, II and III, sum root to leaf numbers.
- Lowest common ancestor of a binary tree and of a BST.
- Validate BST, kth smallest in BST, insert and delete in BST, convert sorted array to BST.
- Construct from preorder and inorder, and from inorder and postorder.
- Serialise and deserialise a binary tree, and a BST.
- Flatten a binary tree to a linked list, populating next right pointers.
- Binary tree cameras and house robber III (tree dynamic programming).