Recursion & backtracking in Java
Subsets, permutations, pruning.
A smaller version of yourself
A recursive method solves a problem by calling itself on a smaller input until it hits a base case that answers directly. Each call waits on the call stack until the calls it made return.
int sum(int n) {
if (n == 0) return 0; // base case
return n + sum(n - 1); // smaller input
}
// sum(3) = 3 + 2 + 1 + 0 = 6Going down, coming back
What does this print?
void go(int n) {
if (n == 0) return;
IO.print(n);
go(n - 1);
IO.print(n);
}
void main() { go(3); }321123321123321
Show the answer
On the way down, each call prints before recursing: 3, 2, 1. At 0 the base case returns, and each waiting call resumes and prints again in reverse: 1, 2, 3.
A base case you can skip over
Every path must reach a base case. Here 7, 5, 3, 1, -1, -3... skips over 0, so the recursion never stops until the call stack runs out: StackOverflowError. Fix it by covering every path: if (n <= 0) return 0;
int down(int n) {
if (n == 0) return 0;
return down(n - 2);
}
// down(7) -> StackOverflowErrorBacktracking: choose, explore, undo
Backtracking builds a candidate step by step. At each item: skip it and recurse, then choose it, recurse, and undo the choice. Like exploring a maze with chalk: at a dead end, walk back to the last fork. n items have 2^n subsets.
void subsets(int i, List<Integer> cur) {
if (i == nums.length) { use(cur); return; }
subsets(i + 1, cur); // skip
cur.add(nums[i]); // choose
subsets(i + 1, cur); // explore
cur.remove(cur.size() - 1); // undo
}Which subset comes first?
This version takes each character before skipping it. What does this print?
void sub(String s, int i, String cur) {
if (i == s.length()) {
IO.print("[" + cur + "]");
return;
}
sub(s, i + 1, cur + s.charAt(i));
sub(s, i + 1, cur);
}
void main() { sub("xy", 0, ""); }[xy][x][y][][][y][x][xy][x][y][xy]
Show the answer
Take x, take y: [xy]. Take x, skip y: [x]. Skip x, take y: [y]. Skip both: []. The last decision (about y) changes fastest. Swap the two calls to skip first and you'd get [][y][x][xy]. Either way: all 2^2 = 4 subsets.
Permutations explode
n distinct items have n! permutations: 3 items give 6, but 10 items give 3,628,800. Generating all of them takes at least O(n!) time, because you can't output n! results in fewer steps. Exhaustive search only works for small n, or with heavy pruning.
Pruning
Pruning abandons a partial candidate as soon as it can't lead to a valid answer. Looking for subsets of positive numbers summing to 10, a partial sum of 12 can never come back down, so return immediately. That skips the whole subtree below it.
void search(int i, int sum) {
if (sum > target) return; // prune
if (sum == target) { found++; return; }
if (i == nums.length) return;
search(i + 1, sum + nums[i]); // take
search(i + 1, sum); // skip
}Where it shows up
Puzzle solvers, scheduling and constraint problems, and "generate all combinations" interview questions are backtracking. Java doesn't eliminate tail calls, so very deep recursion can overflow the stack; production code often converts deep recursion into a loop with an explicit stack.
Key takeaways
- Every recursion needs a base case that is always reached
- n items have 2ⁿ subsets and n! permutations
- Backtracking: choose → recurse → undo
- Prune early when a partial path is already invalid
💡 Backtracking is exploring a maze with chalk: at a dead end, walk back to the last fork and try another way.
The classic eight queens puzzle, placing 8 queens so none attack each other, has exactly 92 solutions. Backtracking with pruning finds them all almost instantly.
Practice questions
What does this print?
void sub(String s, int i, String cur) {
if (i == s.length()) {
System.out.print("[" + cur + "]");
return;
}
sub(s, i + 1, cur);
sub(s, i + 1, cur + s.charAt(i));
}
void main() { sub("ab", 0, ""); }- [][b][a][ab]
- [][a][b][ab]
- [a][b][ab]
- [ab][a][b][]
Check your answer
[][b][a][ab]. At each character the method first skips it, then takes it. The last decision (about b) changes fastest, giving [], [b], [a], [ab]: all 2² subsets.
What does this print?
int f(int n) {
if (n <= 1) return 1;
return f(n - 1) + f(n - 2);
}
void main() { System.out.println(f(5)); }- 5
- 8
- 13
- 120
Check your answer
8. f(0) = f(1) = 1, then each value is the sum of the previous two: 2, 3, 5, 8. The recursion recomputes the same values many times; dynamic programming fixes that.