Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 18 additions & 11 deletions src/taskgraph/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
27 changes: 27 additions & 0 deletions test/test_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()), [])
Expand Down