Largest BST Inside a Tree
Problem
A binary tree holds whole-number keys, and keys may repeat. A subtree means a node together with all of its descendants. Return the number of nodes in the largest subtree that is a valid binary search tree, where every key in a node's left subtree is strictly smaller than the node and every key in its right subtree is strictly larger. An empty tree has answer 0.
Examples
Input: tree = 6, left 4 (children 2 and 5, and 2 has a left child 1),
right 9 (children 7 and 3)
Output: 4
Why: the subtree under 4 holds 1, 2, 4, 5 in order; 3 under 9 spoils the rest
Input: tree = 5, left 3, right 8 (children 7 and 9)
Output: 5
Why: the whole tree is already a valid search tree
Input: tree = empty
Output: 0
Why: edge case, no nodes at all
Hints
0 / 3
Checking every subtree from scratch repeats the same work at every level, which is quadratic on a tall tree.
Whether a node's subtree is valid depends only on facts about its two child subtrees. Which facts does a parent need from a child?
Walk the tree bottom up. Each call returns whether its subtree is valid, its size, and its smallest and largest keys. A node is the root of a valid subtree when both children are valid and the node's key is above the left side's largest and below the right side's smallest.
Solution
A post-order walk answers every subtree using only what its children already reported, so no subtree is examined twice. An empty child reports itself as valid with a smallest key of plus infinity and a largest key of minus infinity, which makes the range test pass for missing children. An invalid child makes the parent invalid too, since the parent's subtree contains it. Time is O(n), and space is O(h) for the recursion on 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 largest_bst(root):
best = 0
def walk(node): # -> (valid, size, smallest, largest)
nonlocal best
if node is None:
return True, 0, float("inf"), float("-inf")
l_ok, l_size, l_lo, l_hi = walk(node.left)
r_ok, r_size, r_lo, r_hi = walk(node.right)
if l_ok and r_ok and l_hi < node.val < r_lo:
size = l_size + r_size + 1
best = max(best, size)
return True, size, min(l_lo, node.val), max(r_hi, node.val)
return False, 0, 0, 0 # the numbers no longer matter
walk(root)
return best
print(largest_bst(T(6, T(4, T(2, T(1)), T(5)), T(9, T(7), T(3))))) # -> 4
print(largest_bst(T(5, T(3), T(8, T(7), T(9))))) # -> 5
print(largest_bst(None)) # -> 0Stuck on the idea rather than the code? Validate a BST covers it.