DSA sheet · Trees · Binary Search Tree pattern
Two Sum IV (BST)
This is the classic Two Sum question ("do two numbers add up to the target?"), but the numbers live in a binary search tree instead of an array. The teacher walks through four ideas: a first guess that fails, then inorder + two pointers, then a hash set, and finally the optimised one: two BST iterators, one walking up from the smallest value and one walking down from the largest. That last one needs the previous problem, BST Iterator. Without it, Part D won't make sense, so do that page first.
Every part below follows the same order:
① the question in simple words → ② what the constraints tell us → ③ intuition → ④ building the logic from examples → ⑤ approach steps → ⑥ code → ⑦ code line by line → ⑧ dry run → ⑨ complexity & remember
- Part 0 · What you must know before starting
- Part A · The question, and a first idea that fails
- Part B · Inorder into a sorted list + two pointers
- Part C · Hash set while traversing
- Part D · Two BST iterators (optimal)
- Part E · Revision page
Part 0 · Before starting
Tree node
class TreeNode:
def __init__(self, val=0, left=None, right=None):
self.val = val # the number in this node
self.left = left # left child, or None
self.right = right # right child, or NoneThe BST property
In a Binary Search Tree, at every node: all values in its left subtree (left child and everything under it) are smaller, and all values in its right subtree are bigger. It's about every node below, not just the direct children. In the tree below, 4 is the right child of 3 (bigger than 3 ✓) and it's also inside 5's left subtree, so it must be smaller than 5 too ✓.
5
/ \
3 6
/ \ \
2 4 7
This is LeetCode's example tree, and the one the teacher uses throughout.
Why inorder of a BST is sorted
Inorder = left subtree, then the node, then right subtree (and the same rule inside each subtree). At any node, the left side holds only smaller values and the right side only bigger ones, so inorder lists the smaller ones, then the node, then the bigger ones. Repeated everywhere, the whole output is in increasing order. For our tree: 2 3 4 5 6 7.
There's also a mirror version: right subtree, then node, then left subtree. That gives the values in decreasing order: 7 6 5 4 3 2. Part D uses both.
Quick recall: the BST Iterator
From the previous problem: keep a stack (a pile where the last thing pushed is the first popped). Push the root and keep going left, pushing each node. next() pops the top node, pushes its right child and that child's whole left chain, and returns the popped value. This gives inorder values one by one, while the stack never holds more than the tree's height (the number of levels on the longest root-to-leaf path).
Part A · The question, and a first idea that fails
LeetCode 653 · Two Sum IV - Input is a BST
1The question in simple words
You get the root of a BST and a number k. Return True if there are two different nodes whose values add up to k, otherwise False. We don't need to say which two; just yes or no.
- Tree above, k = 9 → True (for example 2 + 7, or 3 + 6, or 4 + 5).
- Tree above, k = 28 → False (the largest possible sum is 6 + 7 = 13).
The only difference from the array version: the numbers are stored in a tree, so we have to think about how to reach them.
2What the constraints tell us
- Number of nodes: 1 to 10⁴ → the tree is never empty. And 10⁴ is small: even O(n) + O(n) is only about 2·10⁴ steps.
- Values: −10⁴ to 10⁴, and k is within −10⁵ to 10⁵ → sums are tiny, a normal int (up to about 2·10⁹) is completely safe.
- The tree is guaranteed to be a valid BST, so all values are distinct. That matters for "two different nodes" later.
3Intuition for the first idea
Stand at the root, 5. We want 9, so we need a partner of 9 − 5 = 4. Since 4 < 5, the BST property says 4 can only be on the left. Go left to 3; 4 > 3, go right; there's 4 ✓. Found 5 + 4 = 9. That feels nice: use the BST to search for the partner.
4Why it fails
The trouble shows up when the partner isn't where we look. Imagine 4 was not in the tree:
5
/ \
3 6
/ \
2 7 ← no 4 this time
- At 5 we want 4. Go left to 3, then right → None. So 5 has no partner.
- Now we move on and try 3 as the first number: it needs 9 − 3 = 6. But if we keep searching only below 3 (its left and right), we can never see 6. 6 is up and over on the root's right side, and from 3 there's no way back up.
- So we would answer False, but 3 + 6 = 9 is a real pair. Wrong answer.
The teacher drops this idea here and moves to approaches that look at the whole tree.
→ Yes, if for every node you search for its partner starting again from the root, not from the current node. One BST search costs O(h), where h is the height, so for all n nodes it's O(n·h). That is correct but slower than the next parts. One extra care: the partner you find must be a different node (for k = 10, node 5 would "find" itself).
class Solution:
def findTarget(self, root, k):
def search(x): # normal BST search from the root
cur = root
while cur is not None and cur.val != x:
cur = cur.left if x < cur.val else cur.right
return cur
def dfs(node):
if node is None:
return False
partner = search(k - node.val)
if partner is not None and partner is not node:
return True
return dfs(node.left) or dfs(node.right)
return dfs(root)Part B · Inorder into a sorted list + two pointers
1The question (same as Part A)
Do two different nodes add up to k?
2Constraints
n ≤ 10⁴, so spending O(n) time and O(n) extra space is totally fine here.
3Intuition
Do the inorder traversal and save values in a list. Because it's a BST, the list comes out sorted: [2, 3, 4, 5, 6, 7]. And "find two numbers in a sorted array with a given sum" is the well-known two-pointer pattern from the array playlist: one pointer at the smallest value (left end), one at the largest (right end).
4Building the pointer moves from examples
k = 9
Left pointer on 2, right pointer on 7. 2 + 7 = 9 = k → True right away.
k = 10
2 + 7 = 9, which is less than 10. We need a bigger sum. Which pointer should move?
- Moving the right pointer left gives 6, 5, 4…: smaller numbers, so the sum would only get smaller. No use.
- Moving the left pointer right gives 3: a bigger number. 3 + 7 = 10 ✓ True.
→ Yes. The right pointer is on the biggest value still available. If even the smallest value + that biggest value is too small, then the smallest value plus anything left is too small. It can never be part of an answer. The same logic in reverse throws away the right value when the sum is too big.
while l < r and not l <= r?→ We need two different nodes. When l == r both pointers are on the same node, and using it twice (like 5 + 5 for k = 10) is not allowed.
5Approach steps
- Inorder the tree into a list
vals(sorted). l = 0,r = len(vals) - 1.- While
l < r: compute the sum; equal → True; too small → l += 1; too big → r -= 1. - Loop ended without a match → False.
6Code (Python)
class Solution:
def findTarget(self, root, k):
vals = []
def inorder(node):
if node is None:
return
inorder(node.left)
vals.append(node.val)
inorder(node.right)
inorder(root) # vals is sorted, because it's a BST
l, r = 0, len(vals) - 1
while l < r:
s = vals[l] + vals[r]
if s == k:
return True
elif s < k:
l += 1 # need a bigger sum
else:
r -= 1 # need a smaller sum
return False7Code line by line
| line | what it means |
|---|---|
| inorder(root) | Left, node, right. Fills vals in increasing order. |
| l, r = 0, len(vals) - 1 | Smallest value on the left, largest on the right. |
| while l < r: | Two different positions only. |
| if s == k: return True | Found a pair. |
| elif s < k: l += 1 | Too small. Only a bigger left value can help. |
| else: r -= 1 | Too big. Only a smaller right value can help. |
| return False | The pointers met, and no pair was found. |
8Dry run
vals = [2, 3, 4, 5, 6, 7].
| k | l, r | vals[l] + vals[r] | action |
|---|---|---|---|
| 9 | 0, 5 | 2 + 7 = 9 | equal → True |
| 10 | 0, 5 | 2 + 7 = 9 | < 10 → l = 1 |
| 10 | 1, 5 | 3 + 7 = 10 | equal → True |
| 28 | 0, 5 | 9 | < 28 → l = 1 |
| 28 | 1..4, 5 | 10, 11, 12, 13 | always < 28 → l keeps moving |
| 28 | 5, 5 | l == r → stop → False |
9Complexity & remember
- Time O(n): O(n) to fill the list, plus O(n) for the pointers (each step moves one pointer, so at most n steps; the teacher describes it as about n/2, which is still linear). O(n) + O(n) = O(n). With n = 10⁴ that easily runs.
- Space O(n) for the list (plus the recursion stack).
Part C · Hash set while traversing
1The question (same)
Do two different nodes add up to k?
2Constraints
Same as before. The teacher's motivation here: what if we don't want to build that extra sorted list?
3Intuition
Walk through the tree in any order. Keep a set of the values seen so far (a set answers "is x in here?" in about O(1)). At each node, ask: "is the partner k − value already in my set?" If yes, we found a pair. If not, add this value and continue. This is exactly the hashing version of the array Two Sum.
4Building it from the example (k = 9)
The teacher starts at the root and goes left first:
- At 5: the set is empty, so there's no partner yet. Add 5. Set = {5}.
- At 3: need 9 − 3 = 6. Is 6 in {5}? No. Add 3. Set = {5, 3}.
- At 2: need 7. Not in the set. Add 2. Set = {5, 3, 2}.
- At 4: need 5. 5 is in the set → True. 4 + 5 = 9.
The question only asks True/False, so we stop the moment we find one pair.
→ So a node can't pair with itself. Take k = 10 at node 5: if we added 5 first and then looked for 10 − 5 = 5, we'd "find" it, but that's the same node used twice. Checking first means the set only holds other nodes.
or here, when Same Tree used and?→ Here we only need one pair anywhere. If the left side finds it, we're done, so
or. (Python also skips the right call when the left already returned True.)5Approach steps
- Make an empty set
seen. - Visit nodes with DFS (node, then left, then right).
- At each node: if
k − node.valis inseen→ return True. Else addnode.val. - If the whole tree is visited with no match → False.
6Code (Python)
class Solution:
def findTarget(self, root, k):
seen = set()
def dfs(node):
if node is None:
return False
if k - node.val in seen: # partner already met?
return True
seen.add(node.val)
return dfs(node.left) or dfs(node.right)
return dfs(root)7Code line by line
| line | what it means |
|---|---|
| seen = set() | Values of nodes already visited. |
| if node is None: return False | An empty spot can't give a pair. It also ends the recursion. |
| if k - node.val in seen: | The partner we need was seen earlier → a pair of two different nodes exists. |
| seen.add(node.val) | Remember this value for nodes we visit later. |
| dfs(node.left) or dfs(node.right) | A pair found on either side is enough. |
8Dry run
| k | node | need | seen before | result |
|---|---|---|---|---|
| 9 | 5 | 4 | {} | no → add 5 |
| 9 | 3 | 6 | {5} | no → add 3 |
| 9 | 2 | 7 | {5, 3} | no → add 2 |
| 9 | 4 | 5 | {5, 3, 2} | yes → True |
| 28 | 5, 3, 2, 4, 6, 7 | 23, 25, 26, 24, 22, 21 | never present | False |
9Complexity & remember
- Time O(n): each node once, each set check about O(1).
- Space O(n): the set can hold every value. So we skipped the sorted list, but we still pay O(n) extra space. The teacher's question: can we do better?
k − val in the set first, then add. Works on any binary tree, doesn't even use the BST property.Part D · Two BST iterators (optimal)
1The question (same)
Same question. Goal: keep the two-pointer idea from Part B, but without storing the whole sorted list.
2Constraints
Same as before. We aim for O(n) time and only O(h) extra space.
3Intuition
In Part B, the two pointers only ever needed two things: "give me the next bigger value" (left pointer moving right) and "give me the next smaller value" (right pointer moving left). We never jumped around in the list.
- "Next bigger value, one at a time" is exactly what the BST Iterator gives: inorder, ascending.
- "Next smaller value, one at a time" is the same iterator mirrored: push the root and keep going right; after popping, go to the left child and push its right chain. That gives descending order.
So we run two iterators at once, one from each end, and they act as the two pointers i and j.
4Building the logic
Starting the two stacks
- Ascending iterator (left side): push 5, go left → 3, go left → 2. Stack = [5, 3, 2]. The top (2) is the smallest value.
- Descending iterator (right side): push 5, go right → 6, go right → 7. Stack = [5, 6, 7]. The top (7) is the largest value.
The teacher points out that each stack holds only one node per level, not the whole tree.
Getting i and j
The pointers are values we get by calling next() on each iterator: i = left.next() pops 2, j = right.next() pops 7. Now it's Part B again: compare i + j with k.
- Sum too small → we want a bigger
i→i = left.next(). - Sum too big → we want a smaller
j→j = right.next().
How the descending iterator refills itself
In ascending order, after popping a node we push its right child's left chain. In descending order it's mirrored: after popping a node we push its left child and that child's right chain. The teacher's example: in our tree, 6 has no left child, so popping 6 adds nothing. But if 6 had a left child (imagine some value between 5 and 6 hanging there), popping 6 would push that child, because it must come out next on the way down.
One class, two directions, using a flag
We could write two iterator classes. The teacher instead writes one class with a reverse flag:
reverse = False→ ascending: push_all moves left, and after a pop we start fromnode.right.reverse = True→ descending: push_all moves right, and after a pop we start fromnode.left.
When do we stop? while i < j
As the pointers move, i rises (2, 3, 4…) and j falls (7, 6, 5…). Once they meet on the same value, there is no pair of two different nodes left to try: using one node twice (5 + 5 = 10) is not allowed. Once they cross, every pair has already been considered. So we keep going only while i < j. If the loop ends with no match → False.
left.next() be called on an empty stack?→ No. We only call it inside the loop, where
i < j. The node holding j is bigger than i, so the ascending iterator hasn't reached it yet: there's at least one more value for it to give. The same reasoning protects right.next().i < j) instead of nodes?→ Yes, because a valid BST has distinct values. Equal values would mean the very same node.
5Approach steps
- Build
left = BSTIterator(root, reverse=False)andright = BSTIterator(root, reverse=True). i = left.next()(smallest),j = right.next()(largest).- While
i < j: sum equal → True; sum < k →i = left.next(); sum > k →j = right.next(). - Return False.
6Code (Python)
class BSTIterator:
def __init__(self, root, reverse):
self.stack = []
self.reverse = reverse # False: ascending, True: descending
self.push_all(root)
def push_all(self, node):
while node is not None:
self.stack.append(node)
if self.reverse:
node = node.right # descending: keep going right
else:
node = node.left # ascending: keep going left
def next(self):
node = self.stack.pop()
if self.reverse:
self.push_all(node.left)
else:
self.push_all(node.right)
return node.val
class Solution:
def findTarget(self, root, k):
left = BSTIterator(root, False) # gives 2, 3, 4, ...
right = BSTIterator(root, True) # gives 7, 6, 5, ...
i = left.next()
j = right.next()
while i < j:
s = i + j
if s == k:
return True
elif s < k:
i = left.next() # need bigger
else:
j = right.next() # need smaller
return False7Code line by line
| line | what it means |
|---|---|
| self.reverse = reverse | Which direction this iterator walks. |
| push_all: append, then go right or left | Push the node, then follow the chain toward the end we want first (left end for small values, right end for big values). |
| node = self.stack.pop() | The next value in this iterator's order. |
| push_all(node.left) / push_all(node.right) | Unlock the subtree that comes next: the right side when going up, the left side when going down. |
| i = left.next(); j = right.next() | The two pointers start at the smallest and the largest value. |
| while i < j: | Two different nodes are still available. |
| i = left.next() | Like l += 1 in Part B, but the value is fetched lazily. |
| j = right.next() | Like r -= 1 in Part B. |
| return False | The pointers met or crossed with no match. |
8Dry run: both stacks at every step
Tree 5 / 3, 6 / 2, 4, -, 7 (the example tree). Left of each picture = bottom of the stack, red = top.
Run 1: k = 9
Build both stacks, then i = left.next() pops 2 (no right child), and j = right.next() pops 7 (no left child). 2 + 7 = 9 → True at once.
Run 2: k = 11 (the teacher's walk-through)
- Build the iterators.
i = left.next(): pop 2, no right child → i = 2.j = right.next(): pop 7, no left child → j = 7.
- 2 < 7 ✓. 2 + 7 = 9 < 11 → need bigger →
i = left.next(): pop 3. 3 has right child 4 → push 4. i = 3.
- 3 < 7 ✓. 3 + 7 = 10 < 11 →
i = left.next(): pop 4, no right child. i = 4.
- 4 < 7 ✓. 4 + 7 = 11 = k → return True.
Run 3: k = 14 (no pair, so we see the stopping rule)
- Same start as Run 2: i = 2, j = 7. Stacks: left [5, 3], right [5, 6].
- 2 + 7 = 9 < 14 → i = 3 (pop 3, push 4). Left [5, 4].
- 3 + 7 = 10 < 14 → i = 4 (pop 4). Left [5].
- 4 + 7 = 11 < 14 → i = 5: pop 5, it has right child 6 → push 6 (6 has no left). Left [6].
- 5 + 7 = 12 < 14 → i = 6: pop 6, right child 7 → push 7. Left [7].
- 6 + 7 = 13 < 14 → i = 7: pop 7. Left stack is now empty.
- Check the loop: is 7 < 7? No, the pointers met on the same node → stop → return False. (Indeed the biggest sum of two different nodes is 6 + 7 = 13.)
At no moment did either stack hold more than 3 nodes (the height of the tree), even though both runs saw every value they needed.
9Complexity & remember
- Time O(n): each node is pushed and popped at most once by each iterator, the same as Part B.
- Space O(h) + O(h) = O(h): each stack holds at most one node per level. The teacher calls this "log n + log n". That's right for a balanced tree; for a tree shaped like a straight line, h = n. Either way it's never more than Part B's full list, and usually much less. That's why this is the optimised approach.
i < j.Part E · Revision page
| approach | idea | time | extra space | verdict |
|---|---|---|---|---|
| A. search under each node | look for k − val only below the node | - | - | wrong (misses partners elsewhere) |
| A fixed | search k − val from the root | O(n·h) | O(h) | correct, slower |
| B. inorder + two pointers | sorted list, pointers from both ends | O(n) | O(n) | good |
| C. hash set | is k − val already seen? | O(n) | O(n) | good, works for any tree |
| D. two BST iterators | two pointers fed by two stacks | O(n) | O(h) | best |
| ascending iterator | descending iterator | |
|---|---|---|
| push_all moves | left | right |
| after a pop, start from | node.right | node.left |
| first value | smallest | largest |
| plays the role of | left pointer (i) | right pointer (j) |
2. Sum too small → move the small pointer up. Too big → move the big pointer down.
3. A hash set also works: check
k − val first, then add.4. The optimal way replaces the sorted list with two BST iterators (left-going and right-going).
5. Stop when
i < j fails: one node can't be used twice.✗
l <= r / i <= j (pairs a node with itself)✗ adding to the set before checking (same node used twice)
✗ in the descending iterator, pushing
node.right after a pop (it must be node.left)✗ moving the wrong pointer when the sum is too small
root = TreeNode(5, TreeNode(3, TreeNode(2), TreeNode(4)), TreeNode(6, None, TreeNode(7))) s = Solution() print(s.findTarget(root, 9)) # True print(s.findTarget(root, 11)) # True (4 + 7) print(s.findTarget(root, 28)) # False print(s.findTarget(root, 10)) # True (3 + 7, not 5 + 5) print(s.findTarget(TreeNode(1), 2)) # False (only one node)
Based on this video: Two Sum IV - Input is a BST