{}DSA Atlas

Recursion & Divide and Conquer

Think in terms of smaller subproblems. Base cases, the call stack, recursion trees, and the divide-and-conquer strategy behind merge sort, quickselect and more.

Beginner9 practice problems3 easy4 medium2 hard

The idea

A recursive function solves a problem by calling itself on a smaller version of the same problem, then combining the answer. Every correct recursive function has two parts:

  1. Base case — an input small enough to answer directly (empty list, n = 0, null node).
  2. Recursive case — reduce the problem, recurse, and build the answer from the result.
Python
def total(nums, i=0):
    if i == len(nums):          # base case: nothing left
        return 0
    return nums[i] + total(nums, i + 1)   # trust total() for the rest

The call stack

Every call gets a stack frame holding its parameters and local variables. Frames pile up until a base case returns, then unwind. That's why:

  • Recursion uses O(depth) space even if you allocate nothing.
  • Too-deep recursion causes a stack overflow (Python's default limit is 1000 frames; C++ usually handles ~10⁵–10⁶ depending on frame size).
Python
import sys
sys.setrecursionlimit(10**6)   # Python: raise the limit for deep DFS

Thinking recursively: a recipe

  1. Define the function precisely in words. “height(node) returns the number of nodes on the longest path from node down to a leaf.”
  2. Find the base case(s). What input needs no recursion?
  3. Express the answer using smaller answers. height(node) = 1 + max(height(node.left), height(node.right)).
  4. Make sure every call moves toward a base case.
  5. Analyze: time = calls × work per call; space = max depth.

Recursion tree

Drawing the calls as a tree is the best way to understand both correctness and complexity.

Text
fib(5)
├── fib(4)
│   ├── fib(3)
│   │   ├── fib(2) ...
│   │   └── fib(1)
│   └── fib(2) ...
└── fib(3)          ← computed again! overlapping subproblems
    ├── fib(2) ...
    └── fib(1)

Repeated nodes in the tree mean overlapping subproblems — cache them (memoization) and you've invented dynamic programming. See DP Fundamentals.

Types of recursion

TypeShapeExample
Linearone recursive callsum of list, factorial, linked list traversal
Binary / treetwo or more callsfibonacci, tree traversal, merge sort
Tailrecursive call is the last operationgcd(a, b) = gcd(b, a % b)
Mutualf calls g calls fparsing expressions (expr → term → factor)
Backtrackingtry choice, recurse, undopermutations, N-Queens

Divide and conquer

A specific recursive strategy:

  1. Divide the input into (usually two) independent halves.
  2. Conquer each half recursively.
  3. Combine the results.

The combine step is where the cleverness lives.

Merge sort

def merge_sort(a):
    if len(a) <= 1:
        return a
    mid = len(a) // 2
    left = merge_sort(a[:mid])
    right = merge_sort(a[mid:])
    # combine: merge two sorted lists in O(n)
    out, i, j = [], 0, 0
    while i < len(left) and j < len(right):
        if left[i] <= right[j]:
            out.append(left[i]); i += 1
        else:
            out.append(right[j]); j += 1
    out.extend(left[i:])
    out.extend(right[j:])
    return out
void mergeSort(vector<int>& a, int lo, int hi, vector<int>& tmp) {  // sorts a[lo, hi)
    if (hi - lo <= 1) return;
    int mid = lo + (hi - lo) / 2;
    mergeSort(a, lo, mid, tmp);
    mergeSort(a, mid, hi, tmp);
    int i = lo, j = mid, k = lo;
    while (i < mid && j < hi) tmp[k++] = (a[i] <= a[j]) ? a[i++] : a[j++];
    while (i < mid) tmp[k++] = a[i++];
    while (j < hi) tmp[k++] = a[j++];
    for (int t = lo; t < hi; t++) a[t] = tmp[t];
}

T(n) = 2T(n/2) + O(n) → O(n log n) time, O(n) extra space.

Counting during the merge (inversion count)

The combine step of merge sort sees every “left element vs right element” relationship exactly once. That makes it perfect for counting pairs (i < j) with some property:

Python
def count_inversions(a):
    """Number of pairs i < j with a[i] > a[j]."""
    def solve(lo, hi):          # sorts a[lo:hi], returns inversions inside
        if hi - lo <= 1:
            return 0
        mid = (lo + hi) // 2
        inv = solve(lo, mid) + solve(mid, hi)
        merged, i, j = [], lo, mid
        while i < mid and j < hi:
            if a[i] <= a[j]:
                merged.append(a[i]); i += 1
            else:
                inv += mid - i      # a[j] is smaller than all remaining left items
                merged.append(a[j]); j += 1
        merged += a[i:mid] + a[j:hi]
        a[lo:hi] = merged
        return inv
    return solve(0, len(a))

This pattern solves Count of Smaller Numbers After Self, Reverse Pairs, and Count of Range Sum.

Quickselect: k-th smallest in O(n) average

Partition around a pivot like quicksort, but only recurse into the side containing k.

import random

def quickselect(a, k):          # k is 0-indexed: k=0 → minimum
    lo, hi = 0, len(a) - 1
    while True:
        p = random.randint(lo, hi)
        a[p], a[hi] = a[hi], a[p]
        pivot, store = a[hi], lo
        for i in range(lo, hi):
            if a[i] < pivot:
                a[i], a[store] = a[store], a[i]
                store += 1
        a[store], a[hi] = a[hi], a[store]
        if store == k:
            return a[store]
        if store < k:
            lo = store + 1
        else:
            hi = store - 1
int quickselect(vector<int>& a, int k) {   // k-th smallest, 0-indexed
    int lo = 0, hi = a.size() - 1;
    mt19937 rng(42);
    while (true) {
        int p = lo + rng() % (hi - lo + 1);
        swap(a[p], a[hi]);
        int store = lo;
        for (int i = lo; i < hi; i++)
            if (a[i] < a[hi]) swap(a[i], a[store++]);
        swap(a[store], a[hi]);
        if (store == k) return a[store];
        if (store < k) lo = store + 1; else hi = store - 1;
    }
}

Average T(n) = T(n/2) + O(n) → O(n). Worst case O(n²), made unlikely by the random pivot. In C++ you can just call nth_element.

Fast exponentiation

Python
def power(x, n):
    if n < 0:
        return 1 / power(x, -n)
    if n == 0:
        return 1
    half = power(x, n // 2)
    return half * half if n % 2 == 0 else half * half * x

T(n) = T(n/2) + O(1) → O(log n). Compute half once — calling power(x, n//2) twice would make it O(n).

Recursion → iteration

Any recursion can be converted to a loop with an explicit stack. This is useful to avoid stack overflows:

Python
# recursive pre-order
def preorder(node, out):
    if not node: return
    out.append(node.val)
    preorder(node.left, out)
    preorder(node.right, out)

# iterative version: push right first so left is processed first
def preorder_iter(root):
    out, stack = [], [root] if root else []
    while stack:
        node = stack.pop()
        out.append(node.val)
        if node.right: stack.append(node.right)
        if node.left: stack.append(node.left)
    return out

Common mistakes

  • Missing or wrong base case → infinite recursion. Check the smallest inputs: empty, size 1, zero, null.
  • Not progressing toward the base case (e.g. f(n) calls f(n)), or calling f(n//2) when n = 1 and your base case is only n = 0 — fine — but f((n+1)//2) with n = 1 loops forever.
  • Recomputing the same call twice in one frame (the power bug above).
  • Slicing lists in every call (Python a[1:]) → silently O(n²). Pass indices.
  • Mutating shared state without undoing it (see Backtracking).
  • Returning the wrong thing from some branches — make every path return a value of the same meaning.

Summary

  • Define the function in words, write base cases first, then trust the recursive call.
  • Complexity = number of calls × work per call; space = recursion depth.
  • Divide and conquer: split → solve halves → combine. The combine step can count cross pairs (merge sort trick).
  • Overlapping subproblems in the recursion tree → memoize → dynamic programming.

Practice problems

Ordered roughly from warm-up to hard. Try each one for 25–40 minutes before peeking at the key idea; if you needed the hint, mark it “review” and redo it in a few days.

0 / 9 solved
swap(s[l], s[r]) then recurse on (l+1, r-1). Base case l >= r.
Easy
The smaller head wins; its next = merge(rest of its list, other list).
Easy
depth(node) = 1 + max(depth(left), depth(right)); depth(None) = 0.
Easy
half = pow(x, n//2); return half*half (times x if n odd). Handle negative n with 1/x.
Medium
Row n's k-th symbol comes from its parent at row n-1, position (k+1)//2; flip if k is even.
Medium
Split at every operator, recursively compute all results on each side, combine. Memoize on substring.
Medium
Implement merge sort: split, sort halves recursively, merge in O(n).
Medium
Merge sort on (value, index) pairs: when taking from the left half, add how many right-half elements already went before it.
Hard
Merge sort; before merging, count pairs i in left, j in right with a[i] > 2*a[j] using a moving pointer.
Hard

Click the box to cycle: solved ✓ → needs review ↺ → not started. Links open LeetCode.