DSA
Burning Tree
Binary Trees problem — solution with code and analysis.
Practice Link
Given a binary tree and a node data called target. Find the minimum time required to burn the complete binary tree if the target is set on fire. It is known that in 1 second all nodes connected to a given node get burned. That is its left child, right child, and parent. Note: The tree contains unique values.
Approach#
A tree only lets you traverse downward, but fire spreads in all three directions (left child, right child, and parent). The solution converts the tree into an undirected graph so BFS can explore in every direction from the target.
- Intuition: Model each tree edge as a bidirectional graph edge. Once the graph is built, standard multi-source BFS starting at the target will burn all reachable nodes level by level, where each BFS level corresponds to exactly one second.
- Mechanics: A DFS pass builds an adjacency list — for every parent-child link, both graph[parent] and graph[child] get each other's value. A BFS from the target then processes nodes in waves, using a visited set to avoid re-burning. The number of completed BFS iterations minus one gives the burn time.
- Trade-off: The two-phase approach (O(n) DFS + O(n) BFS) keeps the solution at O(n) overall and is much cleaner than trying to track parent pointers during the BFS itself.
cpp
class Solution {
public:
void generateGraph(unordered_map<int,vector<int>> &graph, Node* root)
{
if(!root)
return;
if(root->left){
graph[root->data].push_back(root->left->data);
graph[root->left->data].push_back(root->data);
generateGraph(graph, root->left);
}
if(root->right){
graph[root->data].push_back(root->right->data);
graph[root->right->data].push_back(root->data);
generateGraph(graph, root->right);
}
if(graph.find(root->data) == graph.end())
graph[root->data] = {};
}
int minTime(Node* root, int target) {
unordered_map<int,vector<int>> graph;
generateGraph(graph, root);
int time = -1;
queue<int> q;
unordered_set<int> visited;
q.push(target);
visited.insert(target);
while(!q.empty())
{
int size = q.size();
for(int i=0;i<size;i++)
{
int curr = q.front();
q.pop();
for(int next: graph[curr])
if(visited.find(next)==visited.end()){
q.push(next);
visited.insert(next);
}
}
time++;
}
return time;
}
};