diff --git a/src/taskgraph/graph.py b/src/taskgraph/graph.py index 576325514..a59875295 100644 --- a/src/taskgraph/graph.py +++ b/src/taskgraph/graph.py @@ -50,32 +50,39 @@ def transitive_closure(self, nodes, reverse=False): | `-------> d - transitive_closure([b]).nodes == set([a, b]) - transitive_closure([c]).nodes == set([c, b, a]) - transitive_closure([c], reverse=True).nodes == set([c]) - transitive_closure([b], reverse=True).nodes == set([b, c, d]) + transitive_closure([b]).nodes == set([b, c, d]) + transitive_closure([c]).nodes == set([c]) + transitive_closure([c], reverse=True).nodes == set([c, b, a]) + transitive_closure([b], reverse=True).nodes == set([b, a]) """ - assert isinstance(nodes, set) + assert isinstance(nodes, (set, frozenset)) if not (nodes <= self.nodes): raise Exception( f"Unknown nodes in transitive closure: {nodes - self.nodes}" ) - # 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) + # An empty closure reaches nothing, and the adjacency map below costs + # O(E) to build. + if not nodes: + return Graph(set(), set()) + + # Index the edges by the endpoint the traversal expands from, keeping the + # traversal O(V+E). `src` is the endpoint an edge is looked up by, `dst` + # the one it expands to. + src, dst = (1, 0) if reverse else (0, 1) + adjacency = {} for edge in self.edges: - left, right, _ = edge - adjacency[right if reverse else left].append(edge) + adjacency.setdefault(edge[src], []).append(edge) new_nodes = set(nodes) new_edges = set() queue = collections.deque(nodes) while queue: node = queue.popleft() + # Nodes with no edges to expand from have no entry in the map. for edge in adjacency.get(node, ()): new_edges.add(edge) - neighbor = edge[0] if reverse else edge[1] + neighbor = edge[dst] if neighbor not in new_nodes: new_nodes.add(neighbor) queue.append(neighbor) diff --git a/test/test_graph.py b/test/test_graph.py index 81d339baa..9820b0a5d 100644 --- a/test/test_graph.py +++ b/test/test_graph.py @@ -194,6 +194,33 @@ 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_transitive_closure_diamond(self): + "transitive closure of a diamond reaches shared descendants by both paths" + self.assertEqual( + self.diamonds.transitive_closure({"A"}), + Graph( + {"A", "D", "F", "G", "I", "J"}, + { + ("A", "D", "L"), + ("A", "F", "L"), + ("D", "F", "L"), + ("D", "G", "L"), + ("F", "I", "L"), + ("G", "I", "L"), + ("G", "J", "L"), + }, + ), + ) + + def test_transitive_closure_frozenset(self): + "transitive closure accepts a frozenset, as returned by Graph.nodes" + self.assertEqual(self.tree.transitive_closure(self.tree.nodes), self.tree) + + def test_transitive_closure_unknown_nodes(self): + "transitive closure raises when given nodes not in the graph" + with pytest.raises(Exception, match="Unknown nodes"): + self.tree.transitive_closure({"z"}) + def test_visit_postorder_empty(self): "postorder visit of an empty graph is empty" self.assertEqual(list(Graph(set(), set()).visit_postorder()), [])