Kth Smallest BST Key, Iteratively
Problem
A binary search tree holds distinct integer keys. Given its root and a number k, return the k-th smallest key, counting from 1, or None if the tree has fewer than k keys. Do it without recursion, since the tree may be a chain far deeper than the call stack allows, and stop as soon as the answer is known.
Examples
Input: tree = 8, left 3 with children 1 and 6 (6 has children 4 and 7),
right 10 with a right child 14; k = 4
Output: 6
Why: in order the keys read 1, 3, 4, 6, 7, 8, 10, 14
Input: same tree; k = 8
Output: 14
Why: the largest key is the last one in order
Input: same tree; k = 9
Output: None
Why: edge case, the tree holds only eight keys
Hints
0 / 3
An inorder walk of a search tree visits the keys from smallest to largest, so the answer is simply the k-th node that walk visits.
The recursive walk keeps pending nodes in call frames. A list used as a stack can hold exactly those nodes instead: the ones you passed on the way down and still owe a visit.
Push nodes while stepping left until you run out. Pop one, count it, and return its key if the count reaches k; otherwise move to its right child and repeat. If the stack and the current node both run out, return None.
Solution
This is the inorder walk with the call stack made explicit: stepping left pushes each node that still owes a visit, and popping one means everything smaller has already been reported. Counting down k as nodes are popped finds the answer at the k-th pop, and returning there skips the rest of the tree. The stack never holds more than one path from the root, and the depth of the tree no longer matters to the interpreter. Time is O(h + k) and space is O(h), for a tree of height h.
class T:
def __init__(self, val, left=None, right=None):
self.val, self.left, self.right = val, left, right
def kth_smallest_key(root, k):
stack, node = [], root
while stack or node:
while node: # step left, noting the way back
stack.append(node)
node = node.left
node = stack.pop() # everything smaller is already counted
k -= 1
if k == 0:
return node.val # stop early: the rest is never visited
node = node.right
return None # fewer than k keys
tree = T(8, T(3, T(1), T(6, T(4), T(7))), T(10, None, T(14)))
print(kth_smallest_key(tree, 4)) # -> 6
print(kth_smallest_key(tree, 8)) # -> 14
print(kth_smallest_key(tree, 9)) # -> NoneStuck on the idea rather than the code? Your Own Call Stack covers it.