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.

Advertisement

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.

Advertisement

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 answer

Look 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.

select(root, k = 4) on a size-augmented BST: three comparisons, no traversal50size 930size 570size 320size 140size 360size 180size 135size 145size 1at 50: left size 5, k = 4 <= 5, go leftat 30: left size 1, k = 4 > 2,k = 4 - 2 = 2, go rightat 40: left size 1, k = 2 = 1 + 1: answer 40Inorder: 20 30 35 40 45 50 60 70 80. Plain inorder visits 4 nodes after descending 3 levels; select touches 3 nodes.Each node stores size = 1 + size(left) + size(right). Insert and delete fix sizes on the way back up.
A size-augmented BST. The query compares k with the left subtree's size at each node and walks one path from the root: go left, subtract and go right, or stop.
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 r

The 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 root

Three 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

ApproachTime per queryExtra spaceMutates treeUse when
Full inorder into a listO(n)O(n)noyou need many ranks from one frozen snapshot: sort once, index many times
Early-exit iterative inorderO(h + k)O(h)noa single query, small k, a tree you do not control
Morris, finishing the walkO(n)O(1)temporarilymemory is the hard limit and no other reader shares the tree
Size-augmented balanced treeO(log n)one integer per nodenomany queries interleaved with inserts and deletes
Heap of size kO(n log k)O(k)nothe 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 <= 0 or k > n should fail loudly. The iterative version above raises; a version that returns None or -1 silently turns into wrong data downstream. Decide whether k is 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 select return plausible wrong answers. Add a debug-mode checker that verifies size on 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 select stays 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

  1. 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.
  2. Implement Morris, snapshot the tree's shape before and after each call, and confirm the snapshot is unchanged for every k.
  3. Add a size field to your BST, implement select and rank, and verify both after every random insert and delete.
  4. Move the augmentation onto a balanced tree (AVL, red-black or treap) and recompute the two affected sizes in each rotation.
  5. 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.
Key takeaway: Inorder traversal of a BST visits keys in sorted order, so the kth visit is the answer. For one query, stop the iterative traversal after k visits: O(h + k) time and O(h) space. Morris traversal removes the stack but rewires the tree, so it must finish the walk or it leaves threads behind. For repeated queries on a changing tree, store subtree sizes and walk one root-to-node path, on a balanced tree so the path is O(log n); the same field gives rank for free.