diff --git a/src/taskgraph/graph.py b/src/taskgraph/graph.py index 573275087..576325514 100644 --- a/src/taskgraph/graph.py +++ b/src/taskgraph/graph.py @@ -61,20 +61,24 @@ def transitive_closure(self, nodes, reverse=False): f"Unknown nodes in transitive closure: {nodes - self.nodes}" ) - # generate a new graph by expanding along edges until reaching a fixed - # point - new_nodes, new_edges = nodes, set() - nodes, edges = set(), set() - while (new_nodes, new_edges) != (nodes, edges): - nodes, edges = new_nodes, new_edges - add_edges = { - (left, right, name) - for (left, right, name) in self.edges - if (right if reverse else left) in nodes - } - add_nodes = {(left if reverse else right) for (left, right, _) in add_edges} - new_nodes = nodes | add_nodes - new_edges = edges | add_edges + # Build an adjacency map keyed by the node to expand from. This reduces + # traversal below from O(V·E) -> O(V+E). + adjacency = collections.defaultdict(list) + for edge in self.edges: + left, right, _ = edge + adjacency[right if reverse else left].append(edge) + + new_nodes = set(nodes) + new_edges = set() + queue = collections.deque(nodes) + while queue: + node = queue.popleft() + for edge in adjacency.get(node, ()): + new_edges.add(edge) + neighbor = edge[0] if reverse else edge[1] + if neighbor not in new_nodes: + new_nodes.add(neighbor) + queue.append(neighbor) return Graph(new_nodes, new_edges) def _visit(self, reverse): diff --git a/test/test_graph.py b/test/test_graph.py index f1d683cc4..81d339baa 100644 --- a/test/test_graph.py +++ b/test/test_graph.py @@ -95,6 +95,21 @@ def test_transitive_closure_trees(self): ), ) + def test_transitive_closure_trees_reverse(self): + "reverse transitive closure of a tree, at two leaves, is their ancestors" + self.assertEqual( + self.tree.transitive_closure({"d", "f"}, reverse=True), + Graph( + {"a", "b", "c", "d", "f"}, + { + ("a", "b", "L"), + ("a", "c", "L"), + ("b", "d", "K"), + ("c", "f", "N"), + }, + ), + ) + def test_transitive_closure_multi_edges(self): "transitive closure of a tree with multiple edges between nodes keeps those edges" self.assertEqual( @@ -111,6 +126,23 @@ def test_transitive_closure_multi_edges(self): ), ) + def test_transitive_closure_multi_edges_reverse(self): + "reverse transitive closure of a tree with multiple edges between nodes keeps those edges" + self.assertEqual( + self.multi_edges.transitive_closure({"1"}, reverse=True), + Graph( + {"1", "2", "3", "4"}, + { + ("2", "1", "red"), + ("2", "1", "blue"), + ("3", "1", "red"), + ("3", "2", "blue"), + ("3", "2", "green"), + ("4", "3", "green"), + }, + ), + ) + def test_transitive_closure_disjoint_edges(self): "transitive closure of a disjoint graph keeps those edges" self.assertEqual( @@ -126,14 +158,42 @@ def test_transitive_closure_disjoint_edges(self): ), ) + def test_transitive_closure_disjoint_edges_reverse(self): + "reverse transitive closure of a disjoint graph keeps those edges" + self.assertEqual( + self.disjoint.transitive_closure({"1", "γ"}, reverse=True), + Graph( + {"1", "2", "3", "4", "α", "β", "γ"}, + { + ("2", "1", "red"), + ("3", "1", "red"), + ("4", "3", "green"), + ("3", "2", "green"), + ("α", "β", "πράσινο"), + ("β", "γ", "κόκκινο"), + ("α", "γ", "μπλε"), + }, + ), + ) + def test_transitive_closure_linear(self): "transitive closure of a linear graph includes all nodes in the line" self.assertEqual(self.linear.transitive_closure({"1"}), self.linear) + def test_transitive_closure_linear_reverse(self): + "reverse transitive closure of a linear graph includes all nodes in the line" + self.assertEqual( + self.linear.transitive_closure({"4"}, reverse=True), self.linear + ) + def test_transitive_closure_loopy(self): "transitive closure of a loop is the whole loop" self.assertEqual(self.loopy.transitive_closure({"A"}), self.loopy) + def test_transitive_closure_loopy_reverse(self): + "reverse transitive closure of a loop is the whole loop" + self.assertEqual(self.loopy.transitive_closure({"A"}, reverse=True), self.loopy) + def test_visit_postorder_empty(self): "postorder visit of an empty graph is empty" self.assertEqual(list(Graph(set(), set()).visit_postorder()), [])