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.
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:
- Base case — an input small enough to answer directly (empty list, n = 0, null node).
- Recursive case — reduce the problem, recurse, and build the answer from the result.
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).
import sys
sys.setrecursionlimit(10**6) # Python: raise the limit for deep DFS
Thinking recursively: a recipe
- Define the function precisely in words. “
height(node)returns the number of nodes on the longest path fromnodedown to a leaf.” - Find the base case(s). What input needs no recursion?
- Express the answer using smaller answers.
height(node) = 1 + max(height(node.left), height(node.right)). - Make sure every call moves toward a base case.
- 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.
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
| Type | Shape | Example |
|---|---|---|
| Linear | one recursive call | sum of list, factorial, linked list traversal |
| Binary / tree | two or more calls | fibonacci, tree traversal, merge sort |
| Tail | recursive call is the last operation | gcd(a, b) = gcd(b, a % b) |
| Mutual | f calls g calls f | parsing expressions (expr → term → factor) |
| Backtracking | try choice, recurse, undo | permutations, N-Queens |
Divide and conquer
A specific recursive strategy:
- Divide the input into (usually two) independent halves.
- Conquer each half recursively.
- 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:
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
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:
# 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)callsf(n)), or callingf(n//2)when n = 1 and your base case is only n = 0 — fine — butf((n+1)//2)with n = 1 loops forever. - Recomputing the same call twice in one frame (the
powerbug 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.
Click the box to cycle: solved ✓ → needs review ↺ → not started. Links open LeetCode.