Given the root of a binary search tree and an integer k, return the kth smallest key. It is one of the most common tree interview questions, and it is also a real data-structure problem: a leaderboard asking for the player in 1,000th place, a database asking for the median of an indexed column, or a scheduler asking for the tenth-earliest deadline are all order-statistic queries.
There are three good answers, and which one is right depends on how often you ask and how often the tree changes. A single query on a tree you do not own is an early-exit inorder traversal. A query under a strict memory limit is a Morris traversal, provided you handle one subtle trap. Many queries on a tree that keeps changing call for a tree that stores subtree sizes, so each query is a short walk from the root. This article builds all three and runs them on one worked tree.
Why inorder order is sorted order
A binary search tree keeps one invariant: every key in a node's left subtree is smaller than the node's key, and every key in its right subtree is larger. An inorder traversal visits the left subtree, then the node, then the right subtree. Apply the invariant at the root and it follows that everything visited before the root is smaller and everything after is larger; apply it recursively and the whole visit sequence is sorted. So the kth node an inorder traversal visits is the kth smallest key.
That gives the brute-force answer immediately: collect the inorder sequence into a list and return element k - 1. It is correct and costs O(n) time and O(n) space whatever k is.
The worked tree for the rest of the article is built by inserting 50, 30, 70, 20, 40, 60, 80, 35, 45 into an empty BST. Its inorder sequence is 20, 30, 35, 40, 45, 50, 60, 70, 80, so the answer for k = 4 is 40. Its height is 3 edges: the path 50, 30, 40, 35.
Early-exit inorder
The first improvement is to count as you visit and stop when the count reaches k. Before the first visit, the traversal has to walk down the left spine from the root to the minimum, which costs up to h steps for a tree of height h; after that each of the k visits costs amortised constant work. Total time is O(h + k) and extra space is O(h) for the explicit stack.
def kth_iterative(root, k):
stack, node = [], root
while stack or node:
while node: # push the left spine
stack.append(node)
node = node.left
node = stack.pop() # the next key in sorted order
k -= 1
if k == 0:
return node.val
node = node.right # continue into the right subtree
raise ValueError("k is larger than the number of nodes")A recursive version is shorter, but its depth equals the tree's height, and an unbalanced BST built from sorted input has height n - 1. That raises RecursionError in Python and overflows the thread stack in Java or C++; the explicit stack lives on the heap and simply grows.
Trace k = 4 on the worked tree. The inner loop pushes 50, 30, 20. Pop 20 (k becomes 3), it has no right child. Pop 30 (k becomes 2), move to 40 and push 40 and 35. Pop 35 (k becomes 1). Pop 40 (k becomes 0) and return 40. Four pops; 50 was pushed but never popped, and nothing to its right was touched.
Morris traversal: O(1) space, and the trap
Both versions above use O(h) memory. Morris traversal removes it by temporarily rewiring the tree. Before descending into a node's left subtree, it finds the node's inorder predecessor, the rightmost node of the left subtree, and sets that predecessor's right pointer to the current node. That pointer is a thread: when the traversal later reaches the predecessor and follows its right pointer, it arrives back at the node without a stack. On the second arrival it sees the thread, removes it, visits the node and moves right.
def kth_morris(root, k):
"""O(1) extra space. Keeps walking after the answer so every thread is removed."""
node, answer = root, None
while node:
if node.left is None:
k -= 1
if k == 0:
answer = node.val
node = node.right
else:
pred = node.left
while pred.right is not None and pred.right is not node:
pred = pred.right
if pred.right is None:
pred.right = node # first arrival: make the thread
node = node.left
else:
pred.right = None # second arrival: remove the thread
k -= 1
if k == 0:
answer = node.val
node = node.right
if answer is None:
raise ValueError("k is larger than the number of nodes")
return answerLook at what the function does not do: it does not return as soon as k reaches zero. That is the trap. At the moment the answer is found, some threads may still be in place. In the worked tree, when 40 is visited for k = 4, the thread from 45 back to 50 still exists, because 50 has not been reached for the second time. Returning there leaves node 45 with a right child that points up the tree. The next plain traversal of that tree loops forever or visits 50 twice; a size computation counts wrong; a serializer that walks right pointers never terminates.
There are two correct fixes. The version above keeps traversing to the end, which restores every thread and costs O(n) time instead of O(h + k): you trade time for the constant space. The alternative is to record the threads you created on the current path and undo them before returning, but tracking them needs memory proportional to the path, which is exactly what Morris was supposed to avoid. So in practice: use Morris only when memory is truly the binding constraint, finish the walk, and test it by snapshotting the tree's shape before and after the call. The tests for this article do exactly that on 300 random trees.
Morris also writes to the tree, so it is unsafe while another thread reads it and impossible on an immutable tree.
When one query is not enough: the order-statistic tree
The follow-up question matters more than the original: what if the tree is modified often and you need the kth smallest often? Each traversal costs O(h + k), which is O(n) for k near n. The fix is to store, in every node, the size of the subtree rooted there. With sizes, the rank of a node inside its subtree is size(left) + 1, and a query can decide at each node which way to go without visiting anything else.
class ONode:
__slots__ = ("key", "left", "right", "size")
def __init__(self, key):
self.key, self.left, self.right, self.size = key, None, None, 1
def size(n):
return n.size if n else 0
def select(root, k):
"""k-th smallest, 1-based, in O(h)."""
if not 1 <= k <= size(root):
raise IndexError(k)
node = root
while True:
left = size(node.left)
if k <= left:
node = node.left
elif k == left + 1:
return node.key
else:
k -= left + 1
node = node.right
def rank(root, key):
"""Number of keys strictly smaller than key, in O(h)."""
r, node = 0, root
while node:
if key <= node.key:
node = node.left
else:
r += size(node.left) + 1
node = node.right
return rThe inverse operation, rank, comes free with the same field: it counts how many keys are smaller than a value. Together they answer median, percentile and "what place is this player in" queries in O(h). On the worked tree, select(root, 4) returns 40 after touching 50, 30 and 40, and rank(root, 45) returns 4, because 20, 30, 35 and 40 are smaller.
Keeping sizes correct under updates
The size field is only useful if every update maintains it. The rule is local: after any change below a node, recompute size = 1 + size(left) + size(right) on the way back up. Recursive insert and delete do this naturally, because each frame fixes its own node after the child call returns.
def insert(root, key):
if root is None:
return ONode(key)
if key < root.key:
root.left = insert(root.left, key)
elif key > root.key:
root.right = insert(root.right, key)
else:
return root # set semantics: a duplicate changes nothing
root.size = 1 + size(root.left) + size(root.right)
return root
def delete(root, key):
if root is None:
return None
if key < root.key:
root.left = delete(root.left, key)
elif key > root.key:
root.right = delete(root.right, key)
else:
if root.left is None:
return root.right
if root.right is None:
return root.left
succ = root.right # two children: copy the successor up
while succ.left:
succ = succ.left
root.key = succ.key
root.right = delete(root.right, succ.key)
root.size = 1 + size(root.left) + size(root.right)
return rootThree details are worth stating. First, the duplicate case returns before recomputing, which is correct because nothing changed, but if your tree is a multiset you need a count field per node and size must add counts, not 1. Second, deleting 30 from the worked tree copies its successor 35 into that node and deletes 35 from the right subtree; afterwards select(root, 4) returns 45, which the tests confirm. Third, the height h in every bound is the real height. On a plain BST fed sorted keys, it is n - 1 and the augmented tree is no faster than a list.
So in production the augmentation goes on a balanced tree. A red-black or AVL tree keeps height O(log n), and each rotation changes the subtrees of exactly two nodes, so it only has to recompute those two sizes, lower node first. A treap or a weight-balanced tree works the same way. The standard libraries mostly do not expose this: Java's TreeMap and C++'s std::set keep no subtree sizes, so kth element there is a linear walk. GCC ships a policy-based tree_order_statistics_node_update extension, and many languages have third-party sorted containers; check whether yours stores sizes before assuming O(log n) selection.
Choosing between the approaches
| Approach | Time per query | Extra space | Mutates tree | Use when |
|---|---|---|---|---|
| Full inorder into a list | O(n) | O(n) | no | you need many ranks from one frozen snapshot: sort once, index many times |
| Early-exit iterative inorder | O(h + k) | O(h) | no | a single query, small k, a tree you do not control |
| Morris, finishing the walk | O(n) | O(1) | temporarily | memory is the hard limit and no other reader shares the tree |
| Size-augmented balanced tree | O(log n) | one integer per node | no | many queries interleaved with inserts and deletes |
| Heap of size k | O(n log k) | O(k) | no | the input is not a BST at all, for example a stream |
If the keys are small integers in a known range, you may not need a tree at all: a Fenwick tree over key counts answers rank with a prefix sum and select with a binary descent, both in O(log U) for a universe of size U. A segment tree over counts does the same and extends to range-restricted versions of the question. For unordered input, a bounded max-heap of size k (see heap operations) or quickselect, the selection cousin of quicksort, is the right tool.
Failure modes
- k out of range.
k <= 0ork > nshould fail loudly. The iterative version above raises; a version that returnsNoneor-1silently turns into wrong data downstream. Decide whetherkis 0-based or 1-based at the API boundary and name it. - Threads left behind. An early return from Morris corrupts the tree's shape. Finish the walk and test with a before-and-after shape snapshot.
- Stale sizes. One code path that forgets to recompute a size, typically a rotation or a bulk-load helper, makes
selectreturn plausible wrong answers. Add a debug-mode checker that verifiessizeon every node after each mutation in tests. - Duplicates. The BST definition above has no equal keys. If yours allows them, decide which side they go to, and in an augmented tree store a count per key so that
selectstays exact. - Concurrent modification. A traversal that interleaves with inserts can skip or repeat keys. Take a read lock, use a persistent tree, or snapshot first.
What to do next
- Write the iterative early-exit version from memory, then test it against
sorted(values)[k - 1]on a few hundred random trees, including one built from sorted input. - Implement Morris, snapshot the tree's shape before and after each call, and confirm the snapshot is unchanged for every
k. - Add a
sizefield to your BST, implementselectandrank, and verify both after every random insert and delete. - Move the augmentation onto a balanced tree (AVL, red-black or treap) and recompute the two affected sizes in each rotation.
- In your own systems, find the queries that ask for a position, median or percentile, and check whether the underlying structure stores subtree sizes or is walking linearly.