You're given a tree with n nodes labelled 0 … n-1, as two lists tree_from and tree_to of length n - 1: edge i joins tree_from[i] and tree_to[i]. The distance between two nodes is the number of edges on the path between them.
You're also given three nodes x, y and z (not necessarily different). For a node u, take its three distances dist(u, x), dist(u, y), dist(u, z) and sort them into a <= b <= c. Call u Pythagorean if a² + b² == c².
Write count_pythagorean_nodes(n, tree_from, tree_to, x, y, z) -> int: how many nodes are Pythagorean. u may be one of x, y, z (then one distance is 0, and it counts if the other two are equal).
# path 0 - 1 - 2 - 3 - 4 - 5, with x = 0, y = 3, z = 5
count_pythagorean_nodes(6, [0, 1, 2, 3, 4], [1, 2, 3, 4, 5], 0, 3, 5)
# node 3: distances 3, 0, 2 -> 0, 2, 3 0 + 4 != 9 no
# node 5: distances 5, 2, 0 -> 0, 2, 5 no
# node 4: distances 4, 1, 1 -> 1, 1, 4 no ... the answer is 0
# star: centre 0 joined to 1, 2, 3; x = 1, y = 2, z = 3
count_pythagorean_nodes(4, [0, 0, 0], [1, 2, 3], 1, 2, 3)
# node 0: 1, 1, 1 -> 1 + 1 != 1 no
# node 1: 0, 2, 2 -> 0 + 4 == 4 yes (and likewise nodes 2 and 3) -> 3
Constraints: 1 <= n <= 2 * 10^5. The tree may be a single long path, so don't recurse. Distances reach about 2 * 10^5, so their squares reach 4 * 10^10: in C++ or Java, square them as 64-bit integers (long), not int.
Show hint
Run one BFS from each of x, y, z to get three distance arrays (O(n) each), then check every node. Remember that b² == c² alone is not enough: a must be 0 then.