Flatten a Tree Into a Chain
Problem
Given the root of a binary tree, rewire it in place into a chain that follows preorder: the root first, then its left subtree, then its right subtree. In the chain every node's left link is None and its right link points to the next node in preorder. Only links may change; no new nodes may be created.
Examples
Input: tree = 7, left 3 (children 1 and 5), right 9 (left child 8)
Output: chain 7 -> 3 -> 1 -> 5 -> 9 -> 8
Why: that is the preorder of the tree
Input: tree = 4, left 2, whose left is 1
Output: chain 4 -> 2 -> 1
Why: the left spine becomes a right spine
Input: tree = empty
Output: empty chain
Why: edge case, there is nothing to rewire
Hints
0 / 3
Preorder visits a node before either of its subtrees, so by the time you reach a node you already know which node comes before it in the chain.
If you keep a pointer to the previously visited node, linking it to the current one is a single assignment. The danger is overwriting a child link you have not followed yet.
Walk the tree in preorder with an explicit stack, pushing the right child before the left one. A node's children are on the stack before any link changes, so it is then safe to set the previous node's left to None and its right to the current node.
Solution
An explicit stack produces preorder: pop a node, push its right child, then its left child, so the left side comes off first. By the time a node is linked into the chain, its children are already saved on the stack, so rewriting the previous node's links can never lose part of the tree. The last node in preorder is always a leaf, so its links are already None. Time is O(n) and space is O(h) for the stack, 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 flatten(root):
stack, prev = ([root] if root else []), None
while stack:
node = stack.pop()
if node.right: stack.append(node.right)
if node.left: stack.append(node.left) # the left side comes off first
if prev:
prev.left, prev.right = None, node # its children are already saved
prev = node
return root
def chain(node):
out = []
while node:
out.append(node.val)
node = node.right
return out
print(chain(flatten(T(7, T(3, T(1), T(5)), T(9, T(8)))))) # -> [7, 3, 1, 5, 9, 8]
print(chain(flatten(T(4, T(2, T(1)))))) # -> [4, 2, 1]
print(chain(flatten(None))) # -> []Stuck on the idea rather than the code? Tree Traversals covers it.