Trees and BST

Trees & BSTs

DFS/BFS on trees, path problems, and BST invariants — with recursion patterns interviewers expect — with plain-language explanation before templates.

Recursion contract

Every tree solution states what one call returns (‘height’, ‘is-balanced pair’, ‘best path through me’). BST validation must pass a legal (low, high) range — comparing only to the parent misses left-subtree vs aunt violations.

Analogy: org chart

Tree org chart
One manager, many reports — no cycles.

A tree is a company org chart: every person has one manager (parent), except the CEO (root). DFS is walking a reporting line to the leaves; BFS is “everyone at this level, then the next.” A BST is an org chart ordered so left reports are “smaller,” right are “larger.”

Traverse with intent

Tree DFS BFS
DFS recursion vs level-order queue.
def max_depth(root):
    if not root: return 0
    return 1 + max(max_depth(root.left), max_depth(root.right))
int maxDepth(TreeNode root) {
    if (root == null) return 0;
    return 1 + Math.max(maxDepth(root.left), maxDepth(root.right));
}
from collections import deque
def level_order(root):
    if not root: return []
    q, out = deque([root]), []
    while q:
        level = []
        for _ in range(len(q)):
            n = q.popleft(); level.append(n.val)
            if n.left: q.append(n.left)
            if n.right: q.append(n.right)
        out.append(level)
    return out
List<List<Integer>> levelOrder(TreeNode root) {
    List<List<Integer>> out = new ArrayList<>();
    if (root == null) return out;
    ArrayDeque<TreeNode> q = new ArrayDeque<>();
    q.add(root);
    while (!q.isEmpty()) {
        List<Integer> level = new ArrayList<>();
        int sz = q.size();
        for (int i = 0; i < sz; i++) {
            TreeNode n = q.poll();
            level.add(n.val);
            if (n.left != null) q.add(n.left);
            if (n.right != null) q.add(n.right);
        }
        out.add(level);
    }
    return out;
}

BST checks

def is_bst(root, lo=float('-inf'), hi=float('inf')):
    if not root: return True
    if not (lo < root.val < hi): return False
    return is_bst(root.left, lo, root.val) and is_bst(root.right, root.val, hi)
boolean isBst(TreeNode root) {
    return isBst(root, Long.MIN_VALUE, Long.MAX_VALUE);
}
boolean isBst(TreeNode root, long lo, long hi) {
    if (root == null) return true;
    if (!(lo < root.val && root.val < hi)) return false;
    return isBst(root.left, lo, root.val) && isBst(root.right, root.val, hi);
}

Pattern menu

  • Pure recurse aggregate (depth, tilt)
  • Pass context down (BST bounds, path sum remaining)
  • Return multiple values (balanced? height)
  • BFS for levels / rightmost / width

Worked trap: false BST

Tree: 5 → left 1, right 7 → left 4. Parent checks look fine (4<7), but 4 is not >5. Range validation catches it: when going right of 5, low=5, so 4 fails. Mention this trap — it’s a common interviewer follow-up.

When to reach for this pattern

Every recursive tree solution needs a return contract: what does one call return about its subtree? Interviews fail candidates who recurse without stating that.

Core template + worked trace

Height / balanced check with an explicit contract:

def is_balanced(root) -> bool:
    def height(node):
        # returns height, or -1 if subtree unbalanced
        if not node:
            return 0
        lh = height(node.left)
        if lh == -1:
            return -1
        rh = height(node.right)
        if rh == -1:
            return -1
        if abs(lh - rh) > 1:
            return -1
        return 1 + max(lh, rh)
    return height(root) != -1

# Tree: 1 / \ 2 3 / 4  — heights bubble up; imbalance returns -1 early
boolean isBalanced(TreeNode root) {
    return height(root) != -1;
}
int height(TreeNode node) {
    // returns height, or -1 if subtree unbalanced
    if (node == null) return 0;
    int lh = height(node.left);
    if (lh == -1) return -1;
    int rh = height(node.right);
    if (rh == -1) return -1;
    if (Math.abs(lh - rh) > 1) return -1;
    return 1 + Math.max(lh, rh);
}

// Tree: 1 / \ 2 3 / 4  — heights bubble up; imbalance returns -1 early

BST validate with bounds (not ‘left < me < right’ only — that misses ancestors):

def is_valid_bst(root) -> bool:
    def ok(node, lo, hi):
        if not node:
            return True
        if not (lo < node.val < hi):
            return False
        return ok(node.left, lo, node.val) and ok(node.right, node.val, hi)
    return ok(root, float('-inf'), float('inf'))

# Classic trap: 5 / \ 1 6 / \ 4 7 — 4 is in right subtree of 5 but 4 < 5 → invalid
boolean isValidBst(TreeNode root) {
    return ok(root, Long.MIN_VALUE, Long.MAX_VALUE);
}
boolean ok(TreeNode node, long lo, long hi) {
    if (node == null) return true;
    if (!(lo < node.val && node.val < hi)) return false;
    return ok(node.left, lo, node.val) && ok(node.right, node.val, hi);
}

// Classic trap: 5 / \ 1 6 / \ 4 7 — 4 is in right subtree of 5 but 4 < 5 → invalid

Edge cases & common bugs

Complexity — say it aloud

Interview talk track

You: ‘My helper returns the height, or −1 if the subtree is already unbalanced, so I prune early. For BST validity I’ll pass exclusive bounds from ancestors — a local left/right check isn’t enough.’

Practice set

  • Maximum Depth of Binary Tree
  • Balanced Binary Tree
  • Validate Binary Search Tree
  • Lowest Common Ancestor of a BST / Binary Tree
  • Binary Tree Level Order Traversal
  • Diameter of Binary Tree
  • Kth Smallest Element in a BST
  • Serialize and Deserialize Binary Tree

Harder follow-up

Harder variant: Binary Tree Maximum Path Sum — path can bend through a node. Helper returns best downward gain; global tracks best bend. Speak the contract split clearly.

Pattern bank: more questions + efficient solutions

Five tree/BST staples with efficient interview solutions. State the recursion contract before coding.

Q: Given a binary tree, compute its maximum depth, and check whether it is height-balanced (for every node, |leftHeight − rightHeight| ≤ 1).

# Max depth + balanced in one pass style helpers — O(n) time, O(h) space
def max_depth(root):
    if not root:
        return 0
    return 1 + max(max_depth(root.left), max_depth(root.right))

def is_balanced(root) -> bool:
    def height(node):
        if not node:
            return 0
        lh = height(node.left)
        if lh == -1:
            return -1
        rh = height(node.right)
        if rh == -1:
            return -1
        if abs(lh - rh) > 1:
            return -1
        return 1 + max(lh, rh)
    return height(root) != -1
// Max depth + balanced — O(n) time, O(h) space
int maxDepth(TreeNode root) {
    if (root == null) return 0;
    return 1 + Math.max(maxDepth(root.left), maxDepth(root.right));
}
boolean isBalanced(TreeNode root) {
    return height(root) != -1;
}
int height(TreeNode node) {
    if (node == null) return 0;
    int lh = height(node.left);
    if (lh == -1) return -1;
    int rh = height(node.right);
    if (rh == -1) return -1;
    if (Math.abs(lh - rh) > 1) return -1;
    return 1 + Math.max(lh, rh);
}

Q: Find the lowest common ancestor of two nodes p and q in a BST (both exist in the tree).

# LCA in BST — O(h) time, O(1) space iterative
def lowest_common_ancestor(root, p, q):
    while root:
        if p.val < root.val and q.val < root.val:
            root = root.left
        elif p.val > root.val and q.val > root.val:
            root = root.right
        else:
            return root
    return None
// LCA in BST — O(h) time, O(1) space
TreeNode lowestCommonAncestor(TreeNode root, TreeNode p, TreeNode q) {
    while (root != null) {
        if (p.val < root.val && q.val < root.val) root = root.left;
        else if (p.val > root.val && q.val > root.val) root = root.right;
        else return root;
    }
    return null;
}

Q: Validate whether a binary tree is a valid BST: every node’s value is strictly greater than all values in its left subtree and strictly less than all in its right.

# Validate BST with bounds — O(n) time, O(h) space
def is_valid_bst(root) -> bool:
    def ok(node, lo, hi):
        if not node:
            return True
        if not (lo < node.val < hi):
            return False
        return ok(node.left, lo, node.val) and ok(node.right, node.val, hi)
    return ok(root, float("-inf"), float("inf"))
// Validate BST with bounds — O(n) time, O(h) space
boolean isValidBST(TreeNode root) {
    return ok(root, Long.MIN_VALUE, Long.MAX_VALUE);
}
boolean ok(TreeNode node, long lo, long hi) {
    if (node == null) return true;
    if (!(lo < node.val && node.val < hi)) return false;
    return ok(node.left, lo, node.val) && ok(node.right, node.val, hi);
}

Q: Serialize a binary tree to a string and deserialize it back (or: return level-order values grouped by level). Prefer a preorder with null markers for round-trip fidelity.

# Serialize / deserialize (preorder) — O(n) time & space
def serialize(root) -> str:
    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: str):
    it = iter(data.split(","))
    def dfs():
        tok = next(it)
        if tok == "#":
            return None
        node = TreeNode(int(tok))
        node.left = dfs()
        node.right = dfs()
        return node
    return dfs()
// Serialize / deserialize (preorder) — O(n) time & space
String serialize(TreeNode root) {
    StringBuilder sb = new StringBuilder();
    ser(root, sb);
    return sb.toString();
}
void ser(TreeNode node, StringBuilder sb) {
    if (node == null) { sb.append("#,"); return; }
    sb.append(node.val).append(',');
    ser(node.left, sb);
    ser(node.right, sb);
}
TreeNode deserialize(String data) {
    Queue<String> q = new ArrayDeque<>(Arrays.asList(data.split(",")));
    return des(q);
}
TreeNode des(Queue<String> q) {
    String tok = q.poll();
    if (tok == null || tok.equals("#") || tok.isEmpty()) return null;
    TreeNode node = new TreeNode(Integer.parseInt(tok));
    node.left = des(q);
    node.right = des(q);
    return node;
}

Q: Return the diameter of a binary tree: the length (number of edges) of the longest path between any two nodes (path may or may not pass through the root).

# Diameter (edge count) — O(n) time, O(h) space
def diameter_of_binary_tree(root) -> int:
    best = 0
    def height(node):
        nonlocal best
        if not node:
            return 0
        lh, rh = height(node.left), height(node.right)
        best = max(best, lh + rh)
        return 1 + max(lh, rh)
    height(root)
    return best
// Diameter (edge count) — O(n) time, O(h) space
int diameter = 0;
int diameterOfBinaryTree(TreeNode root) {
    diameter = 0;
    height(root);
    return diameter;
}
int height(TreeNode node) {
    if (node == null) return 0;
    int lh = height(node.left), rh = height(node.right);
    diameter = Math.max(diameter, lh + rh);
    return 1 + Math.max(lh, rh);
}

45-minute pattern drill

← Lattice