| 188 | |
| 189 | |
| 190 | def topo_sort(downstream_nodes): |
| 191 | marked_nodes = [] |
| 192 | sorted_nodes = [] |
| 193 | outgoing_edge_maps = {} |
| 194 | |
| 195 | def visit( |
| 196 | upstream_node, |
| 197 | upstream_label, |
| 198 | downstream_node, |
| 199 | downstream_label, |
| 200 | downstream_selector=None, |
| 201 | ): |
| 202 | if upstream_node in marked_nodes: |
| 203 | raise RuntimeError('Graph is not a DAG') |
| 204 | |
| 205 | if downstream_node is not None: |
| 206 | outgoing_edge_map = outgoing_edge_maps.get(upstream_node, {}) |
| 207 | outgoing_edge_infos = outgoing_edge_map.get(upstream_label, []) |
| 208 | outgoing_edge_infos += [ |
| 209 | (downstream_node, downstream_label, downstream_selector) |
| 210 | ] |
| 211 | outgoing_edge_map[upstream_label] = outgoing_edge_infos |
| 212 | outgoing_edge_maps[upstream_node] = outgoing_edge_map |
| 213 | |
| 214 | if upstream_node not in sorted_nodes: |
| 215 | marked_nodes.append(upstream_node) |
| 216 | for edge in upstream_node.incoming_edges: |
| 217 | visit( |
| 218 | edge.upstream_node, |
| 219 | edge.upstream_label, |
| 220 | edge.downstream_node, |
| 221 | edge.downstream_label, |
| 222 | edge.upstream_selector, |
| 223 | ) |
| 224 | marked_nodes.remove(upstream_node) |
| 225 | sorted_nodes.append(upstream_node) |
| 226 | |
| 227 | unmarked_nodes = [(node, None) for node in downstream_nodes] |
| 228 | while unmarked_nodes: |
| 229 | upstream_node, upstream_label = unmarked_nodes.pop() |
| 230 | visit(upstream_node, upstream_label, None, None) |
| 231 | return sorted_nodes, outgoing_edge_maps |