Distance Between Two BST Keys
Problem
A binary search tree stores distinct keys, and two of its keys a and b are given, both guaranteed to be in the tree. Return the number of edges on the path that connects the node holding a to the node holding b. When a equals b the answer is 0.
Examples
Input: tree = 8, left 3 with children 1 and 6 (6 has children 4 and 7),
right 10 with a right child 14; a = 4, b = 14
Output: 5
Why: the path is 4, 6, 3, 8, 10, 14
Input: same tree; a = 6, b = 7
Output: 1
Why: 7 is a child of 6
Input: same tree; a = 10, b = 10
Output: 0
Why: edge case, a node is zero edges from itself
Hints
0 / 3
Every path between two nodes of a tree climbs up to one turning point and then goes down. Which node is that turning point?
It is the lowest common ancestor. In a BST you can find it from the root without any searching: keep stepping while both keys lie on the same side.
Walk down from the root until a and b stop being on the same side of the current node, or one of them equals it. From that node, count the steps a normal BST search takes to reach a, do the same for b, and add the two counts.
Solution
The path between two nodes passes through their lowest common ancestor, so its length is the depth of a below that node plus the depth of b below it. In a BST the ancestor is the first node on the way down where the two keys split to different sides or where one of them is found, exactly as in the lesson. From there, two ordinary BST searches count the steps to each key. Time is O(h) for a tree of height h, and space is O(1).
class T:
def __init__(self, val, left=None, right=None):
self.val, self.left, self.right = val, left, right
def bst_distance(root, a, b):
node = root
while (a < node.val and b < node.val) or (a > node.val and b > node.val):
node = node.left if a < node.val else node.right # both lie the same way
def steps(n, key): # edges from n down to key by BST search
d = 0
while n.val != key:
n = n.left if key < n.val else n.right
d += 1
return d
return steps(node, a) + steps(node, b)
tree = T(8, T(3, T(1), T(6, T(4), T(7))), T(10, None, T(14)))
print(bst_distance(tree, 4, 14)) # -> 5
print(bst_distance(tree, 6, 7)) # -> 1
print(bst_distance(tree, 10, 10)) # -> 0Stuck on the idea rather than the code? Lowest Common Ancestor covers it.