1#![warn(missing_docs)]
2
3extern crate proc_macro;
4
5use std::collections::{BTreeMap, BTreeSet};
6use std::fmt::Debug;
7use std::iter::FusedIterator;
8
9use itertools::Itertools;
10use proc_macro2::{Ident, Literal, Span, TokenStream};
11use quote::{ToTokens, format_ident, quote, quote_spanned};
12use serde::{Deserialize, Serialize};
13use slotmap::{Key, SecondaryMap, SlotMap, SparseSecondaryMap};
14use syn::spanned::Spanned;
15
16use super::graph_write::{Dot, GraphWrite, Mermaid};
17use super::ops::{
18 DelayType, FloType, OPERATORS, OperatorWriteOutput, WriteContextArgs, find_op_op_constraints,
19 null_write_iterator_fn,
20};
21use super::{
22 CONTEXT, Color, DiMulGraph, GRAPH, GraphEdgeId, GraphLoopId, GraphNode, GraphNodeId,
23 GraphSubgraphId, HANDOFF_NODE_STR, HandoffKind, MODULE_BOUNDARY_NODE_STR, OperatorInstance,
24 PortIndexValue, SINGLETON_SLOT_NODE_STR, Varname, change_spans, get_operator_generics,
25};
26use crate::diagnostic::{Diagnostic, Diagnostics, Level};
27use crate::pretty_span::{PrettyRowCol, PrettySpan};
28use crate::process_singletons;
29
30#[derive(Clone, Debug, Serialize, Deserialize)]
32pub struct ResolvedHandoffRef {
33 pub node_id: Option<GraphNodeId>,
35 pub is_mut: bool,
37 pub access_group: Option<u32>,
39}
40
41#[derive(Default, Debug, Serialize, Deserialize)]
51pub struct DfirGraph {
52 nodes: SlotMap<GraphNodeId, GraphNode>,
54
55 #[serde(skip)]
58 operator_instances: SecondaryMap<GraphNodeId, OperatorInstance>,
59 operator_tag: SecondaryMap<GraphNodeId, String>,
61 graph: DiMulGraph<GraphNodeId, GraphEdgeId>,
63 ports: SecondaryMap<GraphEdgeId, (PortIndexValue, PortIndexValue)>,
65
66 node_loops: SecondaryMap<GraphNodeId, GraphLoopId>,
68 loop_nodes: SlotMap<GraphLoopId, Vec<GraphNodeId>>,
70 loop_parent: SparseSecondaryMap<GraphLoopId, GraphLoopId>,
72 root_loops: Vec<GraphLoopId>,
74 loop_children: SecondaryMap<GraphLoopId, Vec<GraphLoopId>>,
76
77 node_subgraph: SecondaryMap<GraphNodeId, GraphSubgraphId>,
79
80 subgraph_nodes: SlotMap<GraphSubgraphId, Vec<GraphNodeId>>,
82 subgraph_toposort: Vec<GraphSubgraphId>,
84
85 node_handoff_references: SparseSecondaryMap<GraphNodeId, Vec<ResolvedHandoffRef>>,
87 node_varnames: SparseSecondaryMap<GraphNodeId, Varname>,
89
90 handoff_delay_type: SparseSecondaryMap<GraphNodeId, DelayType>,
94}
95
96impl DfirGraph {
98 pub fn new() -> Self {
100 Default::default()
101 }
102}
103
104impl DfirGraph {
106 pub fn node(&self, node_id: GraphNodeId) -> &GraphNode {
108 self.nodes.get(node_id).expect("Node not found.")
109 }
110
111 pub fn node_op_inst(&self, node_id: GraphNodeId) -> Option<&OperatorInstance> {
116 self.operator_instances.get(node_id)
117 }
118
119 pub fn node_varname(&self, node_id: GraphNodeId) -> Option<&Varname> {
121 self.node_varnames.get(node_id)
122 }
123
124 pub fn node_subgraph(&self, node_id: GraphNodeId) -> Option<GraphSubgraphId> {
126 self.node_subgraph.get(node_id).copied()
127 }
128
129 pub fn node_degree_in(&self, node_id: GraphNodeId) -> usize {
131 self.graph.degree_in(node_id)
132 }
133
134 pub fn node_degree_out(&self, node_id: GraphNodeId) -> usize {
136 self.graph.degree_out(node_id)
137 }
138
139 pub fn node_successors(
141 &self,
142 src: GraphNodeId,
143 ) -> impl '_
144 + DoubleEndedIterator<Item = (GraphEdgeId, GraphNodeId)>
145 + ExactSizeIterator
146 + FusedIterator
147 + Clone
148 + Debug {
149 self.graph.successors(src)
150 }
151
152 pub fn node_predecessors(
154 &self,
155 dst: GraphNodeId,
156 ) -> impl '_
157 + DoubleEndedIterator<Item = (GraphEdgeId, GraphNodeId)>
158 + ExactSizeIterator
159 + FusedIterator
160 + Clone
161 + Debug {
162 self.graph.predecessors(dst)
163 }
164
165 pub fn node_successor_edges(
167 &self,
168 src: GraphNodeId,
169 ) -> impl '_
170 + DoubleEndedIterator<Item = GraphEdgeId>
171 + ExactSizeIterator
172 + FusedIterator
173 + Clone
174 + Debug {
175 self.graph.successor_edges(src)
176 }
177
178 pub fn node_predecessor_edges(
180 &self,
181 dst: GraphNodeId,
182 ) -> impl '_
183 + DoubleEndedIterator<Item = GraphEdgeId>
184 + ExactSizeIterator
185 + FusedIterator
186 + Clone
187 + Debug {
188 self.graph.predecessor_edges(dst)
189 }
190
191 pub fn node_successor_nodes(
193 &self,
194 src: GraphNodeId,
195 ) -> impl '_
196 + DoubleEndedIterator<Item = GraphNodeId>
197 + ExactSizeIterator
198 + FusedIterator
199 + Clone
200 + Debug {
201 self.graph.successor_vertices(src)
202 }
203
204 pub fn node_predecessor_nodes(
206 &self,
207 dst: GraphNodeId,
208 ) -> impl '_
209 + DoubleEndedIterator<Item = GraphNodeId>
210 + ExactSizeIterator
211 + FusedIterator
212 + Clone
213 + Debug {
214 self.graph.predecessor_vertices(dst)
215 }
216
217 pub fn node_ids(&self) -> slotmap::basic::Keys<'_, GraphNodeId, GraphNode> {
219 self.nodes.keys()
220 }
221
222 pub fn nodes(&self) -> slotmap::basic::Iter<'_, GraphNodeId, GraphNode> {
224 self.nodes.iter()
225 }
226
227 pub fn insert_node(
229 &mut self,
230 node: GraphNode,
231 varname_opt: Option<Ident>,
232 loop_opt: Option<GraphLoopId>,
233 ) -> GraphNodeId {
234 let node_id = self.nodes.insert(node);
235 if let Some(varname) = varname_opt {
236 self.node_varnames.insert(node_id, Varname(varname));
237 }
238 if let Some(loop_id) = loop_opt {
239 self.node_loops.insert(node_id, loop_id);
240 self.loop_nodes[loop_id].push(node_id);
241 }
242 node_id
243 }
244
245 pub fn insert_node_op_inst(&mut self, node_id: GraphNodeId, op_inst: OperatorInstance) {
247 assert!(matches!(
248 self.nodes.get(node_id),
249 Some(GraphNode::Operator(_))
250 ));
251 let old_inst = self.operator_instances.insert(node_id, op_inst);
252 assert!(old_inst.is_none());
253 }
254
255 pub fn insert_node_op_insts_all(&mut self, diagnostics: &mut Diagnostics) {
257 let mut op_insts = Vec::new();
262 let mut handoff_nodes: Vec<(GraphNodeId, HandoffKind, Span)> = Vec::new();
264
265 for (node_id, node) in self.nodes() {
266 let GraphNode::Operator(operator) = node else {
267 continue;
268 };
269 if self.node_op_inst(node_id).is_some() {
270 continue;
271 };
272
273 let handoff_kind = match &*operator.name_string() {
275 "handoff" => Some(HandoffKind::Vec),
276 "singleton" => Some(HandoffKind::Singleton),
277 "optional" => Some(HandoffKind::Optional),
278 _ => None,
279 };
280 if let Some(kind) = handoff_kind {
281 if !operator.args.is_empty() {
282 diagnostics.push(Diagnostic::spanned(
283 operator.path.span(),
284 Level::Error,
285 format!("`{}` takes no arguments.", operator.name_string()),
286 ));
287 }
288 if operator.type_arguments().is_some() {
289 diagnostics.push(Diagnostic::spanned(
290 operator.path.span(),
291 Level::Error,
292 format!("`{}` takes no generic arguments.", operator.name_string()),
293 ));
294 }
295 handoff_nodes.push((node_id, kind, operator.path.span()));
296 continue;
297 }
298
299 let Some(op_constraints) = find_op_op_constraints(operator) else {
301 diagnostics.push(Diagnostic::spanned(
302 operator.path.span(),
303 Level::Error,
304 format!("Unknown operator `{}`", operator.name_string()),
305 ));
306 continue;
307 };
308
309 let (input_ports, output_ports) = {
311 let mut input_edges: Vec<(&PortIndexValue, GraphNodeId)> = self
312 .node_predecessors(node_id)
313 .map(|(edge_id, pred_id)| (self.edge_ports(edge_id).1, pred_id))
314 .collect();
315 input_edges.sort();
317 let input_ports: Vec<PortIndexValue> = input_edges
318 .into_iter()
319 .map(|(port, _pred)| port)
320 .cloned()
321 .collect();
322
323 let mut output_edges: Vec<(&PortIndexValue, GraphNodeId)> = self
325 .node_successors(node_id)
326 .map(|(edge_id, succ)| (self.edge_ports(edge_id).0, succ))
327 .collect();
328 output_edges.sort();
330 let output_ports: Vec<PortIndexValue> = output_edges
331 .into_iter()
332 .map(|(port, _succ)| port)
333 .cloned()
334 .collect();
335
336 (input_ports, output_ports)
337 };
338
339 let generics = get_operator_generics(diagnostics, operator);
341 {
343 let generics_span = generics
345 .generic_args
346 .as_ref()
347 .map(Spanned::span)
348 .unwrap_or_else(|| operator.path.span());
349
350 if !op_constraints
351 .persistence_args
352 .contains(&generics.persistence_args.len())
353 {
354 diagnostics.push(Diagnostic::spanned(
355 generics.persistence_args_span().unwrap_or(generics_span),
356 Level::Error,
357 format!(
358 "`{}` should have {} persistence lifetime arguments, actually has {}.",
359 op_constraints.name,
360 op_constraints.persistence_args.human_string(),
361 generics.persistence_args.len()
362 ),
363 ));
364 }
365 if !op_constraints.type_args.contains(&generics.type_args.len()) {
366 diagnostics.push(Diagnostic::spanned(
367 generics.type_args_span().unwrap_or(generics_span),
368 Level::Error,
369 format!(
370 "`{}` should have {} generic type arguments, actually has {}.",
371 op_constraints.name,
372 op_constraints.type_args.human_string(),
373 generics.type_args.len()
374 ),
375 ));
376 }
377 }
378
379 op_insts.push((
380 node_id,
381 OperatorInstance {
382 op_constraints,
383 input_ports,
384 output_ports,
385 singletons_referenced: operator.singletons_referenced.clone(),
386 generics,
387 arguments_pre: operator.args.clone(),
388 arguments_raw: operator.args_raw.clone(),
389 },
390 ));
391 }
392
393 for (node_id, op_inst) in op_insts {
394 self.insert_node_op_inst(node_id, op_inst);
395 }
396
397 for (node_id, kind, span) in handoff_nodes {
399 self.nodes[node_id] = GraphNode::Handoff {
400 kind,
401 src_span: span,
402 dst_span: span,
403 };
404 }
405 }
406
407 pub fn insert_intermediate_node(
419 &mut self,
420 edge_id: GraphEdgeId,
421 new_node: GraphNode,
422 ) -> (GraphNodeId, GraphEdgeId) {
423 let span = Some(new_node.span());
424
425 let op_inst_opt = 'oc: {
427 let GraphNode::Operator(operator) = &new_node else {
428 break 'oc None;
429 };
430 let Some(op_constraints) = find_op_op_constraints(operator) else {
431 break 'oc None;
432 };
433 let (input_port, output_port) = self.ports.get(edge_id).cloned().unwrap();
434
435 let mut dummy_diagnostics = Diagnostics::new();
436 let generics = get_operator_generics(&mut dummy_diagnostics, operator);
437 assert!(dummy_diagnostics.is_empty());
438
439 Some(OperatorInstance {
440 op_constraints,
441 input_ports: vec![input_port],
442 output_ports: vec![output_port],
443 singletons_referenced: operator.singletons_referenced.clone(),
444 generics,
445 arguments_pre: operator.args.clone(),
446 arguments_raw: operator.args_raw.clone(),
447 })
448 };
449
450 let node_id = self.nodes.insert(new_node);
452 if let Some(op_inst) = op_inst_opt {
454 self.operator_instances.insert(node_id, op_inst);
455 }
456 let (e0, e1) = self
458 .graph
459 .insert_intermediate_vertex(node_id, edge_id)
460 .unwrap();
461
462 let (src_idx, dst_idx) = self.ports.remove(edge_id).unwrap();
464 self.ports
465 .insert(e0, (src_idx, PortIndexValue::Elided(span)));
466 self.ports
467 .insert(e1, (PortIndexValue::Elided(span), dst_idx));
468
469 (node_id, e1)
470 }
471
472 pub fn remove_intermediate_node(&mut self, node_id: GraphNodeId) {
475 assert_eq!(
476 1,
477 self.node_degree_in(node_id),
478 "Removed intermediate node must have one predecessor"
479 );
480 assert_eq!(
481 1,
482 self.node_degree_out(node_id),
483 "Removed intermediate node must have one successor"
484 );
485 assert!(
486 self.node_subgraph.is_empty() && self.subgraph_nodes.is_empty(),
487 "Should not remove intermediate node after subgraph partitioning"
488 );
489
490 assert!(self.nodes.remove(node_id).is_some());
491 let (new_edge_id, (pred_edge_id, succ_edge_id)) =
492 self.graph.remove_intermediate_vertex(node_id).unwrap();
493 self.operator_instances.remove(node_id);
494 self.node_varnames.remove(node_id);
495
496 let (src_port, _) = self.ports.remove(pred_edge_id).unwrap();
497 let (_, dst_port) = self.ports.remove(succ_edge_id).unwrap();
498 self.ports.insert(new_edge_id, (src_port, dst_port));
499 }
500
501 pub(crate) fn node_color(&self, node_id: GraphNodeId) -> Option<Color> {
507 if matches!(self.node(node_id), GraphNode::Handoff { .. }) {
508 return Some(Color::Hoff);
509 }
510
511 if let GraphNode::Operator(op) = self.node(node_id)
513 && (op.name_string() == "resolve_futures_blocking"
514 || op.name_string() == "resolve_futures_blocking_ordered")
515 {
516 return Some(Color::Push);
517 }
518
519 let inn_degree = self.node_predecessor_nodes(node_id).len();
521 let out_degree = self.node_successor_nodes(node_id).len();
523
524 match (inn_degree, out_degree) {
525 (0, 0) => None, (0, 1) => Some(Color::Pull),
527 (1, 0) => Some(Color::Push),
528 (1, 1) => None, (_many, 0 | 1) => Some(Color::Pull),
530 (0 | 1, _many) => Some(Color::Push),
531 (_many, _to_many) => Some(Color::Comp),
532 }
533 }
534
535 pub fn set_operator_tag(&mut self, node_id: GraphNodeId, tag: String) {
537 self.operator_tag.insert(node_id, tag);
538 }
539}
540
541impl DfirGraph {
543 pub fn set_node_handoff_references(
546 &mut self,
547 node_id: GraphNodeId,
548 singletons_referenced: Vec<ResolvedHandoffRef>,
549 ) -> Option<Vec<ResolvedHandoffRef>> {
550 self.node_handoff_references
551 .insert(node_id, singletons_referenced)
552 }
553
554 pub fn node_handoff_references(&self, node_id: GraphNodeId) -> &[ResolvedHandoffRef] {
557 self.node_handoff_references
558 .get(node_id)
559 .map(std::ops::Deref::deref)
560 .unwrap_or_default()
561 }
562
563 pub fn node_handoff_reference_groups(&self) -> NodeHandoffReferenceGroups<'_> {
565 let mut handoff_references = NodeHandoffReferenceGroups::new();
566 for node_id in self.node_ids() {
567 if let GraphNode::Operator(operator) = self.node(node_id) {
568 let resolved = self.node_handoff_references(node_id);
569 for (resolved_ref, ref_token) in
570 resolved.iter().zip(operator.singletons_referenced.iter())
571 {
572 if let Some(target_nid) = resolved_ref.node_id {
573 handoff_references
574 .entry(target_nid)
575 .or_default()
576 .entry(resolved_ref.access_group)
577 .or_default()
578 .push((node_id, resolved_ref, ref_token.span()));
579 }
580 }
581 }
582 }
583 handoff_references
584 }
585}
586
587pub type NodeHandoffReferenceGroups<'a> =
590 BTreeMap<GraphNodeId, BTreeMap<Option<u32>, Vec<(GraphNodeId, &'a ResolvedHandoffRef, Span)>>>;
591
592impl DfirGraph {
594 pub fn merge_modules(&mut self) -> Result<(), Diagnostic> {
602 let mod_bound_nodes = self
603 .nodes()
604 .filter(|(_nid, node)| matches!(node, GraphNode::ModuleBoundary { .. }))
605 .map(|(nid, _node)| nid)
606 .collect::<Vec<_>>();
607
608 for mod_bound_node in mod_bound_nodes {
609 self.remove_module_boundary(mod_bound_node)?;
610 }
611
612 Ok(())
613 }
614
615 fn remove_module_boundary(&mut self, mod_bound_node: GraphNodeId) -> Result<(), Diagnostic> {
619 assert!(
620 self.node_subgraph.is_empty() && self.subgraph_nodes.is_empty(),
621 "Should not remove intermediate node after subgraph partitioning"
622 );
623
624 let mut mod_pred_ports = BTreeMap::new();
625 let mut mod_succ_ports = BTreeMap::new();
626
627 for mod_out_edge in self.node_predecessor_edges(mod_bound_node) {
628 let (pred_port, succ_port) = self.edge_ports(mod_out_edge);
629 mod_pred_ports.insert(succ_port.clone(), (mod_out_edge, pred_port.clone()));
630 }
631
632 for mod_inn_edge in self.node_successor_edges(mod_bound_node) {
633 let (pred_port, succ_port) = self.edge_ports(mod_inn_edge);
634 mod_succ_ports.insert(pred_port.clone(), (mod_inn_edge, succ_port.clone()));
635 }
636
637 if mod_pred_ports.keys().collect::<BTreeSet<_>>()
638 != mod_succ_ports.keys().collect::<BTreeSet<_>>()
639 {
640 let GraphNode::ModuleBoundary { input, import_expr } = self.node(mod_bound_node) else {
642 panic!();
643 };
644
645 if *input {
646 return Err(Diagnostic {
647 span: *import_expr,
648 level: Level::Error,
649 message: format!(
650 "The ports into the module did not match. input: {:?}, expected: {:?}",
651 mod_pred_ports.keys().map(|x| x.to_string()).join(", "),
652 mod_succ_ports.keys().map(|x| x.to_string()).join(", ")
653 ),
654 });
655 } else {
656 return Err(Diagnostic {
657 span: *import_expr,
658 level: Level::Error,
659 message: format!(
660 "The ports out of the module did not match. output: {:?}, expected: {:?}",
661 mod_succ_ports.keys().map(|x| x.to_string()).join(", "),
662 mod_pred_ports.keys().map(|x| x.to_string()).join(", "),
663 ),
664 });
665 }
666 }
667
668 for (port, (pred_edge, pred_port)) in mod_pred_ports {
669 let (succ_edge, succ_port) = mod_succ_ports.remove(&port).unwrap();
670
671 let (src, _) = self.edge(pred_edge);
672 let (_, dst) = self.edge(succ_edge);
673 self.remove_edge(pred_edge);
674 self.remove_edge(succ_edge);
675
676 let new_edge_id = self.graph.insert_edge(src, dst);
677 self.ports.insert(new_edge_id, (pred_port, succ_port));
678 }
679
680 self.graph.remove_vertex(mod_bound_node);
681 self.nodes.remove(mod_bound_node);
682
683 Ok(())
684 }
685}
686
687impl DfirGraph {
689 pub fn edge(&self, edge_id: GraphEdgeId) -> (GraphNodeId, GraphNodeId) {
691 let (src, dst) = self.graph.edge(edge_id).expect("Edge not found.");
692 (src, dst)
693 }
694
695 pub fn edge_ports(&self, edge_id: GraphEdgeId) -> (&PortIndexValue, &PortIndexValue) {
697 let (src_port, dst_port) = self.ports.get(edge_id).expect("Edge not found.");
698 (src_port, dst_port)
699 }
700
701 pub fn edge_ids(&self) -> slotmap::basic::Keys<'_, GraphEdgeId, (GraphNodeId, GraphNodeId)> {
703 self.graph.edge_ids()
704 }
705
706 pub fn edges(
708 &self,
709 ) -> impl '_
710 + ExactSizeIterator<Item = (GraphEdgeId, (GraphNodeId, GraphNodeId))>
711 + FusedIterator
712 + Clone
713 + Debug {
714 self.graph.edges()
715 }
716
717 pub fn insert_edge(
719 &mut self,
720 src: GraphNodeId,
721 src_port: PortIndexValue,
722 dst: GraphNodeId,
723 dst_port: PortIndexValue,
724 ) -> GraphEdgeId {
725 let edge_id = self.graph.insert_edge(src, dst);
726 self.ports.insert(edge_id, (src_port, dst_port));
727 edge_id
728 }
729
730 pub fn remove_edge(&mut self, edge: GraphEdgeId) {
732 let (_src, _dst) = self.graph.remove_edge(edge).unwrap();
733 let (_src_port, _dst_port) = self.ports.remove(edge).unwrap();
734 }
735}
736
737impl DfirGraph {
739 pub fn subgraph(&self, subgraph_id: GraphSubgraphId) -> &Vec<GraphNodeId> {
741 self.subgraph_nodes
742 .get(subgraph_id)
743 .expect("Subgraph not found.")
744 }
745
746 pub fn subgraph_ids(&self) -> slotmap::basic::Keys<'_, GraphSubgraphId, Vec<GraphNodeId>> {
748 self.subgraph_nodes.keys()
749 }
750
751 pub fn subgraph_toposort(&self) -> &[GraphSubgraphId] {
753 &self.subgraph_toposort
754 }
755
756 pub fn set_subgraph_toposort(&mut self, order: Vec<GraphSubgraphId>) {
758 self.subgraph_toposort = order;
759 }
760
761 pub fn subgraphs(&self) -> slotmap::basic::Iter<'_, GraphSubgraphId, Vec<GraphNodeId>> {
763 self.subgraph_nodes.iter()
764 }
765
766 pub fn insert_subgraph(
768 &mut self,
769 node_ids: Vec<GraphNodeId>,
770 ) -> Result<GraphSubgraphId, (GraphNodeId, GraphSubgraphId)> {
771 for &node_id in node_ids.iter() {
773 if let Some(&old_sg_id) = self.node_subgraph.get(node_id) {
774 return Err((node_id, old_sg_id));
775 }
776 }
777 let subgraph_id = self.subgraph_nodes.insert_with_key(|sg_id| {
778 for &node_id in node_ids.iter() {
779 self.node_subgraph.insert(node_id, sg_id);
780 }
781 node_ids
782 });
783
784 Ok(subgraph_id)
785 }
786
787 pub fn remove_from_subgraph(&mut self, node_id: GraphNodeId) -> bool {
789 if let Some(old_sg_id) = self.node_subgraph.remove(node_id) {
790 self.subgraph_nodes[old_sg_id].retain(|&other_node_id| other_node_id != node_id);
791 true
792 } else {
793 false
794 }
795 }
796
797 pub fn handoff_delay_type(&self, node_id: GraphNodeId) -> Option<DelayType> {
799 self.handoff_delay_type.get(node_id).copied()
800 }
801
802 pub fn set_handoff_delay_type(&mut self, node_id: GraphNodeId, delay_type: DelayType) {
804 self.handoff_delay_type.insert(node_id, delay_type);
805 }
806
807 fn find_pull_to_push_idx(&self, subgraph_nodes: &[GraphNodeId]) -> usize {
809 subgraph_nodes
810 .iter()
811 .position(|&node_id| {
812 self.node_color(node_id)
813 .is_some_and(|color| Color::Pull != color)
814 })
815 .unwrap_or(subgraph_nodes.len())
816 }
817}
818
819impl DfirGraph {
821 fn node_as_ident(&self, node_id: GraphNodeId, is_pred: bool) -> Ident {
823 let name = match &self.nodes[node_id] {
824 GraphNode::Operator(_) => format!("op_{:?}", node_id.data()),
825 GraphNode::Handoff {
826 kind: HandoffKind::Vec,
827 ..
828 } => format!(
829 "hoff_{:?}_{}",
830 node_id.data(),
831 if is_pred { "recv" } else { "send" }
832 ),
833 GraphNode::Handoff {
834 kind: HandoffKind::Singleton | HandoffKind::Optional,
835 ..
836 } => format!(
837 "singleton_{:?}_{}",
838 node_id.data(),
839 if is_pred { "recv" } else { "send" }
840 ),
841 GraphNode::ModuleBoundary { .. } => panic!(),
842 };
843 let span = match (is_pred, &self.nodes[node_id]) {
844 (_, GraphNode::Operator(operator)) => operator.span(),
845 (true, &GraphNode::Handoff { src_span, .. }) => src_span,
846 (false, &GraphNode::Handoff { dst_span, .. }) => dst_span,
847 (_, GraphNode::ModuleBoundary { .. }) => panic!(),
848 };
849 Ident::new(&name, span)
850 }
851
852 fn hoff_buf_ident(&self, hoff_id: GraphNodeId, span: Span) -> Ident {
854 Ident::new(&format!("hoff_{:?}_buf", hoff_id.data()), span)
855 }
856
857 fn hoff_back_ident(&self, hoff_id: GraphNodeId, span: Span) -> Ident {
859 Ident::new(&format!("hoff_{:?}_back", hoff_id.data()), span)
860 }
861
862 fn helper_resolve_singletons(&self, node_id: GraphNodeId, span: Span) -> Vec<TokenStream> {
871 self.node_handoff_references(node_id)
872 .iter()
873 .map(|resolved_ref| {
874 let ref_node_id = resolved_ref
876 .node_id
877 .expect("Expected singleton to be resolved but was not, this is a bug.");
878 let is_mut = resolved_ref.is_mut;
879 match self.node(ref_node_id) {
880 GraphNode::Handoff {
881 kind: HandoffKind::Singleton,
882 ..
883 } => {
884 let buf_ident = self.hoff_buf_ident(ref_node_id, span);
885 if is_mut {
886 quote_spanned! {span=> #buf_ident.as_mut().unwrap() }
887 } else {
888 quote_spanned! {span=> #buf_ident.as_ref().unwrap() }
889 }
890 }
891 GraphNode::Handoff {
892 kind: HandoffKind::Optional | HandoffKind::Vec,
893 ..
894 } => {
895 let buf_ident = self.hoff_buf_ident(ref_node_id, span);
896 if is_mut {
897 quote_spanned! {span=> &mut #buf_ident }
898 } else {
899 quote_spanned! {span=> &#buf_ident }
900 }
901 }
902 _ => {
903 unreachable!("Only handoff nodes should be reachable as handoff references")
904 }
905 }
906 })
907 .collect::<Vec<_>>()
908 }
909
910 fn helper_collect_subgraph_handoffs(
913 &self,
914 ) -> SecondaryMap<GraphSubgraphId, (Vec<GraphNodeId>, Vec<GraphNodeId>)> {
915 let mut subgraph_handoffs: SecondaryMap<
917 GraphSubgraphId,
918 (Vec<GraphNodeId>, Vec<GraphNodeId>),
919 > = self
920 .subgraph_nodes
921 .keys()
922 .map(|k| (k, Default::default()))
923 .collect();
924
925 for (hoff_id, hoff) in self.nodes() {
927 if !matches!(hoff, GraphNode::Handoff { .. }) {
928 continue;
929 }
930 for (_edge, succ_id) in self.node_successors(hoff_id) {
932 let succ_sg = self
933 .node_subgraph(succ_id)
934 .expect("bug: successor not in subgraph, may be a doubled/adjacent handoff");
935 subgraph_handoffs[succ_sg].0.push(hoff_id);
936 }
937 for (_edge, pred_id) in self.node_predecessors(hoff_id) {
939 let pred_sg = self
940 .node_subgraph(pred_id)
941 .expect("bug: predecessor not in subgraph, may be a doubled/adjacent handoff");
942 subgraph_handoffs[pred_sg].1.push(hoff_id);
943 }
944 }
945
946 subgraph_handoffs
947 }
948
949 fn helper_loop_output_handoffs(&self) -> SecondaryMap<GraphLoopId, Vec<GraphNodeId>> {
952 let mut loop_hoffs_out = SecondaryMap::<GraphLoopId, Vec<GraphNodeId>>::new();
953
954 for (hoff_id, hoff) in self.nodes() {
955 if !matches!(hoff, GraphNode::Handoff { .. }) {
956 continue;
957 }
958
959 let loop_pred = self
960 .node_predecessors(hoff_id)
961 .next()
962 .and_then(|(_, pred)| self.node_loop(pred));
963 let loop_succ = self
964 .node_successors(hoff_id)
965 .next()
966 .and_then(|(_, succ)| self.node_loop(succ));
967
968 if let Some(loop_pred) = loop_pred
969 && loop_succ == self.loop_parent(loop_pred)
970 {
971 loop_hoffs_out
973 .entry(loop_pred)
974 .expect("loop removed")
975 .or_default()
976 .push(hoff_id);
977 }
978 }
979
980 loop_hoffs_out
981 }
982
983 fn is_inside_loop(&self, node_loop: Option<GraphLoopId>, loop_id: GraphLoopId) -> bool {
985 let mut current = node_loop;
986 while let Some(l) = current {
987 if l == loop_id {
988 return true;
989 }
990 current = self.loop_parent(l);
991 }
992 false
993 }
994
995 fn emit_loop_gate(
1004 &self,
1005 loop_id: GraphLoopId,
1006 child_body: TokenStream,
1007 loop_input_handoffs: &SecondaryMap<GraphLoopId, Vec<GraphNodeId>>,
1008 back_edge_hoffs_and_lazyness: &SparseSecondaryMap<GraphNodeId, bool>,
1009 loop_swap_code: &std::collections::HashMap<GraphLoopId, Vec<TokenStream>>,
1010 output: &mut TokenStream,
1011 ) {
1012 let swap_code = loop_swap_code
1014 .get(&loop_id)
1015 .map(|v| v.as_slice())
1016 .unwrap_or(&[]);
1017
1018 let is_root_loop = self.loop_parent(loop_id).is_none();
1020
1021 let entry_handoffs = loop_input_handoffs.get(loop_id).expect("loop missing");
1023 let mut gate_checks: Vec<TokenStream> = entry_handoffs
1024 .iter()
1025 .filter(|&&hoff_id| {
1026 let is_lazy = self
1029 .node_successors(hoff_id)
1030 .next()
1031 .and_then(|(_, succ)| self.node_op_inst(succ))
1032 .is_some_and(|op_inst| {
1033 op_inst.op_constraints.flo_type == Some(FloType::WindowingLazy)
1034 });
1035 !is_lazy
1036 })
1037 .map(|&hoff_id| {
1038 let span = self.node(hoff_id).span();
1039 let buf_ident = self.hoff_buf_ident(hoff_id, span);
1040 if back_edge_hoffs_and_lazyness.contains_key(hoff_id) {
1041 let back_ident = self.hoff_back_ident(hoff_id, span);
1042 quote_spanned! {span=> !#back_ident.is_empty() }
1043 } else {
1044 quote_spanned! {span=> !#buf_ident.is_empty() }
1045 }
1046 })
1047 .collect();
1048
1049 if !is_root_loop {
1051 for (hoff_id, hoff) in self.nodes() {
1052 if !matches!(hoff, GraphNode::Handoff { .. }) {
1053 continue;
1054 }
1055 let Some(delay_type) = self.handoff_delay_type(hoff_id) else {
1056 continue;
1057 };
1058 if delay_type != DelayType::Loop {
1059 continue;
1060 }
1061 let hoff_loop = self
1063 .node_successors(hoff_id)
1064 .next()
1065 .and_then(|(_, succ)| self.node_subgraph(succ))
1066 .and_then(|sg| self.subgraph_loop(sg));
1067 if hoff_loop != Some(loop_id) {
1068 continue;
1069 }
1070 let span = self.node(hoff_id).span();
1071 let back_ident = self.hoff_back_ident(hoff_id, span);
1072 gate_checks.push(quote_spanned! {span=> !#back_ident.is_empty() });
1073 }
1074 }
1075
1076 if is_root_loop {
1079 for (hoff_id, hoff) in self.nodes() {
1080 if !matches!(hoff, GraphNode::Handoff { .. }) {
1081 continue;
1082 }
1083 let Some(delay_type) = self.handoff_delay_type(hoff_id) else {
1084 continue;
1085 };
1086 if delay_type != DelayType::Tick {
1087 continue;
1088 }
1089 let hoff_loop = self
1091 .node_successors(hoff_id)
1092 .next()
1093 .and_then(|(_, succ)| self.node_subgraph(succ))
1094 .and_then(|sg| self.subgraph_loop(sg));
1095 if hoff_loop != Some(loop_id) {
1096 continue;
1097 }
1098 let span = self.node(hoff_id).span();
1099 let back_ident = self.hoff_back_ident(hoff_id, span);
1100 gate_checks.push(quote_spanned! {span=> !#back_ident.is_empty() });
1101 }
1102 }
1103
1104 if gate_checks.is_empty() {
1105 output.extend(child_body);
1107 output.extend(quote! { #( #swap_code )* });
1108 } else if is_root_loop {
1109 output.extend(quote! {
1111 #[allow(clippy::nonminimal_bool, reason = "codegen")]
1112 if false #( || #gate_checks )* {
1113 #child_body
1114 #( #swap_code )*
1115 }
1116 });
1117 } else {
1118 output.extend(quote! {
1120 #[allow(clippy::nonminimal_bool, reason = "codegen")]
1121 while false #( || #gate_checks )* {
1122 #child_body
1123 #( #swap_code )*
1124 }
1125 });
1126 }
1127 }
1128
1129 fn helper_loop_input_handoffs(&self) -> SecondaryMap<GraphLoopId, Vec<GraphNodeId>> {
1131 let mut loop_hoffs_inn = SecondaryMap::<GraphLoopId, Vec<GraphNodeId>>::new();
1132
1133 for (hoff_id, hoff) in self.nodes() {
1135 if !matches!(hoff, GraphNode::Handoff { .. }) {
1136 continue;
1137 }
1138
1139 let loop_pred = self
1141 .node_predecessors(hoff_id)
1142 .next()
1143 .and_then(|(_, pred)| self.node_loop(pred));
1144 let loop_succ = self
1145 .node_successors(hoff_id)
1146 .next()
1147 .and_then(|(_, succ)| self.node_loop(succ));
1148
1149 if let Some(loop_succ) = loop_succ
1150 && loop_pred == self.loop_parent(loop_succ)
1151 {
1152 loop_hoffs_inn
1154 .entry(loop_succ)
1155 .expect("loop removed")
1156 .or_default()
1157 .push(hoff_id);
1158 }
1159 }
1160
1161 loop_hoffs_inn
1162 }
1163
1164 pub fn as_code(
1179 &self,
1180 root: &TokenStream,
1181 include_type_guards: bool,
1182 prefix: TokenStream,
1183 diagnostics: &mut Diagnostics,
1184 ) -> Result<TokenStream, Diagnostics> {
1185 self.as_code_with_options(root, include_type_guards, true, prefix, diagnostics)
1186 }
1187
1188 pub fn as_code_with_options(
1197 &self,
1198 root: &TokenStream,
1199 include_type_guards: bool,
1200 include_meta: bool,
1201 prefix: TokenStream,
1202 diagnostics: &mut Diagnostics,
1203 ) -> Result<TokenStream, Diagnostics> {
1204 let df = Ident::new(GRAPH, Span::call_site());
1205 let context = Ident::new(CONTEXT, Span::call_site());
1206 let bump_ident = Ident::new("__dfir_bump", Span::call_site());
1208
1209 let handoff_nodes = self
1211 .nodes
1212 .iter()
1213 .filter_map(|(node_id, node)| match node {
1214 &GraphNode::Handoff {
1215 kind,
1216 src_span,
1217 dst_span,
1218 } => Some((node_id, kind, (src_span, dst_span))),
1219 GraphNode::Operator(_) => None,
1220 GraphNode::ModuleBoundary { .. } => panic!(),
1221 })
1222 .collect::<Vec<_>>();
1223
1224 let back_edge_hoffs_and_lazyness = handoff_nodes
1228 .iter()
1229 .map(|&(node_id, _, _)| node_id)
1230 .filter_map(|node_id| {
1231 let delay_type = self.handoff_delay_type(node_id)?;
1232 Some((
1233 node_id,
1234 matches!(delay_type, DelayType::TickLazy | DelayType::LoopLazy),
1235 ))
1236 })
1237 .collect::<SparseSecondaryMap<_, _>>();
1238
1239 let back_buffer_idents_laziness = handoff_nodes
1241 .iter()
1242 .filter_map(|&(hoff_id, _kind, (src_span, dst_span))| {
1243 back_edge_hoffs_and_lazyness.get(hoff_id).map(|&is_lazy| {
1244 let span = src_span.join(dst_span).unwrap_or(src_span);
1245 let back_ident = self.hoff_back_ident(hoff_id, span);
1246 let buf_ident = self.hoff_buf_ident(hoff_id, span);
1247 (back_ident, buf_ident, is_lazy)
1248 })
1249 })
1250 .collect::<Vec<_>>();
1251
1252 let back_edge_swap_code = handoff_nodes
1259 .iter()
1260 .filter(|&&(node_id, _kind, _)| {
1261 self.handoff_delay_type(node_id)
1262 .is_some_and(|dt| matches!(dt, DelayType::Tick | DelayType::TickLazy))
1263 })
1264 .filter(|&&(hoff_id, _kind, _)| {
1265 let consumer_loop = self
1268 .node_successors(hoff_id)
1269 .next()
1270 .and_then(|(_, succ)| self.node_subgraph(succ))
1271 .and_then(|sg| self.subgraph_loop(sg));
1272 if let Some(loop_id) = consumer_loop {
1273 self.loop_parent(loop_id).is_some()
1275 } else {
1276 true
1278 }
1279 })
1280 .map(|&(hoff_id, _kind, _)| {
1281 let span = self.nodes[hoff_id].span();
1282 let buf_ident = self.hoff_buf_ident(hoff_id, span);
1283 let back_ident = self.hoff_back_ident(hoff_id, span);
1284 quote_spanned! {span=>
1285 ::std::mem::swap(&mut #buf_ident, &mut #back_ident);
1286 }
1287 })
1288 .collect::<Vec<_>>();
1289
1290 let mut loop_swap_code: std::collections::HashMap<GraphLoopId, Vec<TokenStream>> =
1294 std::collections::HashMap::new();
1295 for &(hoff_id, _kind, _) in handoff_nodes.iter() {
1296 let Some(delay_type) = self.handoff_delay_type(hoff_id) else {
1297 continue;
1298 };
1299 let loop_id = self
1301 .node_successors(hoff_id)
1302 .next()
1303 .and_then(|(_, succ)| self.node_subgraph(succ))
1304 .and_then(|sg| self.subgraph_loop(sg));
1305 let Some(loop_id) = loop_id else {
1306 continue;
1307 };
1308 let include = match delay_type {
1309 DelayType::Loop | DelayType::LoopLazy => true,
1310 DelayType::Tick | DelayType::TickLazy => {
1311 self.loop_parent(loop_id).is_none()
1313 }
1314 };
1315 if !include {
1316 continue;
1317 }
1318 let span = self.nodes[hoff_id].span();
1319 let buf_ident = self.hoff_buf_ident(hoff_id, span);
1320 let back_ident = self.hoff_back_ident(hoff_id, span);
1321 loop_swap_code
1322 .entry(loop_id)
1323 .or_default()
1324 .push(quote_spanned! {span=>
1325 ::std::mem::swap(&mut #buf_ident, &mut #back_ident);
1326 });
1327 }
1328
1329 let subgraph_handoffs = self.helper_collect_subgraph_handoffs();
1331
1332 let all_subgraphs: Vec<_> = self
1334 .subgraph_toposort()
1335 .iter()
1336 .map(|&sg_id| (sg_id, self.subgraph(sg_id)))
1337 .collect();
1338
1339 let mut op_prologue_code = Vec::new();
1343 let mut op_tick_end_code = Vec::new();
1344
1345 let mut loop_stack: Vec<(GraphLoopId, TokenStream)> = Vec::new();
1349 let mut current_output = TokenStream::new();
1350
1351 let loop_input_handoffs = self.helper_loop_input_handoffs();
1353 let loop_output_handoffs = self.helper_loop_output_handoffs();
1354
1355 {
1356 for &(subgraph_id, subgraph_nodes) in all_subgraphs.iter() {
1357 let sg_loop = self.subgraph_loop(subgraph_id);
1358
1359 while let Some(&(top_loop, _)) = loop_stack.last() {
1362 if sg_loop == Some(top_loop) || self.is_inside_loop(sg_loop, top_loop) {
1363 break;
1364 }
1365 let (closed_loop, child_body) = loop_stack.pop().unwrap();
1367 let target = if let Some((_, parent_body)) = loop_stack.last_mut() {
1368 parent_body
1369 } else {
1370 &mut current_output
1371 };
1372 self.emit_loop_gate(
1373 closed_loop,
1374 child_body,
1375 &loop_input_handoffs,
1376 &back_edge_hoffs_and_lazyness,
1377 &loop_swap_code,
1378 target,
1379 );
1380 }
1381
1382 if let Some(target_loop) = sg_loop
1384 && loop_stack.last().map(|&(l, _)| l) != Some(target_loop)
1385 {
1386 let mut path = Vec::new();
1388 let mut cur = Some(target_loop);
1389 while let Some(l) = cur {
1390 if loop_stack.last().map(|&(top, _)| top) == Some(l) {
1391 break;
1392 }
1393 path.push(l);
1394 cur = self.loop_parent(l);
1395 }
1396 for &loop_id in path.iter().rev() {
1399 if let Some(exit_hoffs) = loop_output_handoffs.get(loop_id) {
1401 let exit_hoff_decls = exit_hoffs.iter().map(|&hoff_id| {
1402 let span = self.nodes[hoff_id].span();
1403 let buf_ident = self.hoff_buf_ident(hoff_id, span);
1404 let GraphNode::Handoff { kind, .. } = self.node(hoff_id) else {
1405 panic!()
1406 };
1407 match kind {
1408 HandoffKind::Vec => quote_spanned! {span=>
1409 let mut #buf_ident = #root::bumpalo::collections::Vec::new_in(&#bump_ident);
1410 },
1411 HandoffKind::Singleton | HandoffKind::Optional => quote_spanned! {span=>
1412 let mut #buf_ident = ::std::option::Option::None;
1413 },
1414 }
1415 });
1416 let target = if let Some((_, body)) = loop_stack.last_mut() {
1417 body
1418 } else {
1419 &mut current_output
1420 };
1421 target.extend(quote! { #( #exit_hoff_decls )* });
1422 }
1423 loop_stack.push((loop_id, TokenStream::new()));
1424 }
1425 }
1426 let sg_metrics_ffi = subgraph_id.data().as_ffi();
1427 let (recv_hoffs, send_hoffs) = &subgraph_handoffs[subgraph_id];
1428
1429 let recv_port_idents: Vec<Ident> = recv_hoffs
1431 .iter()
1432 .map(|&hoff_id| self.node_as_ident(hoff_id, true))
1433 .collect();
1434 let send_port_idents: Vec<Ident> = send_hoffs
1435 .iter()
1436 .map(|&hoff_id| self.node_as_ident(hoff_id, false))
1437 .collect();
1438
1439 let recv_buf_idents: Vec<Ident> = recv_hoffs
1441 .iter()
1442 .map(|&hoff_id| self.hoff_buf_ident(hoff_id, self.nodes[hoff_id].span()))
1443 .collect();
1444 let send_buf_idents: Vec<Ident> = send_hoffs
1445 .iter()
1446 .map(|&hoff_id| self.hoff_buf_ident(hoff_id, self.nodes[hoff_id].span()))
1447 .collect();
1448
1449 let recv_kinds = recv_hoffs
1451 .iter()
1452 .map(|&hoff_id| {
1453 let GraphNode::Handoff { kind, .. } = self.node(hoff_id) else {
1454 panic!()
1455 };
1456 *kind
1457 })
1458 .collect::<Vec<_>>();
1459 let send_kinds = send_hoffs
1460 .iter()
1461 .map(|&hoff_id| {
1462 let GraphNode::Handoff { kind, .. } = self.node(hoff_id) else {
1463 panic!()
1464 };
1465 *kind
1466 })
1467 .collect::<Vec<_>>();
1468
1469 let recv_port_code: Vec<TokenStream> = recv_port_idents
1473 .iter()
1474 .zip(recv_buf_idents.iter())
1475 .zip(recv_kinds.iter())
1476 .zip(recv_hoffs.iter())
1477 .map(|(((port_ident, buf_ident), &kind), &hoff_id)| {
1478 let hoff_ffi = hoff_id.data().as_ffi();
1479 let work_done = Ident::new("__dfir_work_done", Span::call_site());
1483 let metrics = Ident::new("__dfir_metrics", Span::call_site());
1484
1485 let (len_expr, drain_expr) = match kind {
1487 HandoffKind::Singleton | HandoffKind::Optional => (
1488 quote! { if #buf_ident.is_some() { 1usize } else { 0usize } },
1489 quote! { #root::dfir_pipes::pull::iter(#buf_ident.take().into_iter()) },
1490 ),
1491 HandoffKind::Vec => {
1492 let drain_ident = if back_edge_hoffs_and_lazyness.contains_key(hoff_id) {
1496 &self.hoff_back_ident(hoff_id, buf_ident.span())
1497 } else {
1498 buf_ident
1499 };
1500 (
1501 quote! { #drain_ident.len() },
1502 quote! { #root::dfir_pipes::pull::iter(#drain_ident.drain(..)) },
1503 )
1504 }
1505 };
1506
1507 quote_spanned! {port_ident.span()=>
1508 {
1509 let hoff_len = #len_expr;
1510 if hoff_len > 0 {
1511 #work_done = true;
1512 }
1513 let hoff_metrics = &#metrics.handoffs[
1514 #root::slotmap::KeyData::from_ffi(#hoff_ffi).into()
1515 ];
1516 hoff_metrics.total_items_count.update(|x| x + hoff_len);
1517 hoff_metrics.curr_items_count.set(hoff_len);
1518 }
1519 let #port_ident = #drain_expr;
1520 }
1521 })
1522 .collect();
1523
1524 let send_port_code: Vec<TokenStream> = send_port_idents
1526 .iter()
1527 .zip(send_buf_idents.iter())
1528 .zip(send_kinds.iter())
1529 .map(|((port_ident, buf_ident), &kind)| {
1530 match kind {
1531 HandoffKind::Singleton => {
1532 quote_spanned! {port_ident.span()=>
1534 let #port_ident = #root::dfir_pipes::push::for_each(|__item| {
1535 if #buf_ident.replace(__item).is_some() {
1536 panic!("singleton() received more than one item");
1537 }
1538 });
1539 }
1540 }
1541 HandoffKind::Optional => {
1542 quote_spanned! {port_ident.span()=>
1544 let #port_ident = #root::dfir_pipes::push::for_each(|__item| {
1545 if #buf_ident.replace(__item).is_some() {
1546 panic!("optional() received more than one item");
1547 }
1548 });
1549 }
1550 }
1551 HandoffKind::Vec => {
1552 quote_spanned! {port_ident.span()=>
1553 let #port_ident = #root::dfir_pipes::push::for_each(|item| { #buf_ident.push(item); });
1555 }
1556 }
1557 }
1558 })
1559 .collect();
1560
1561 let loop_id = self.node_loop(subgraph_nodes[0]);
1563
1564 let mut subgraph_op_iter_code = Vec::new();
1565 let mut subgraph_op_iter_after_code = Vec::new();
1566 {
1567 let pull_to_push_idx = self.find_pull_to_push_idx(subgraph_nodes);
1568
1569 let (pull_half, push_half) = subgraph_nodes.split_at(pull_to_push_idx);
1570 let nodes_iter = pull_half.iter().chain(push_half.iter().rev());
1571
1572 for (idx, &node_id) in nodes_iter.enumerate() {
1573 let node = &self.nodes[node_id];
1574 assert!(
1575 matches!(node, GraphNode::Operator(_)),
1576 "Handoffs are not part of subgraphs."
1577 );
1578 let op_inst = &self.operator_instances[node_id];
1579
1580 let op_span = node.span();
1581 let op_name = op_inst.op_constraints.name;
1582 let root = change_spans(root.clone(), op_span);
1584 let op_constraints = OPERATORS
1585 .iter()
1586 .find(|op| op_name == op.name)
1587 .unwrap_or_else(|| panic!("Failed to find op: {}", op_name));
1588
1589 let ident = self.node_as_ident(node_id, false);
1590
1591 {
1592 let mut input_edges = self
1595 .graph
1596 .predecessor_edges(node_id)
1597 .map(|edge_id| (self.edge_ports(edge_id).1, edge_id))
1598 .collect::<Vec<_>>();
1599 input_edges.sort();
1601
1602 let inputs = input_edges
1603 .iter()
1604 .map(|&(_port, edge_id)| {
1605 let (pred, _) = self.edge(edge_id);
1606 self.node_as_ident(pred, true)
1607 })
1608 .collect::<Vec<_>>();
1609
1610 let mut output_edges = self
1612 .graph
1613 .successor_edges(node_id)
1614 .map(|edge_id| (&self.ports[edge_id].0, edge_id))
1615 .collect::<Vec<_>>();
1616 output_edges.sort();
1618
1619 let outputs = output_edges
1620 .iter()
1621 .map(|&(_port, edge_id)| {
1622 let (_, succ) = self.edge(edge_id);
1623 self.node_as_ident(succ, false)
1624 })
1625 .collect::<Vec<_>>();
1626
1627 let is_pull = idx < pull_to_push_idx;
1628
1629 let df_local = &Ident::new(GRAPH, op_span.resolved_at(df.span()));
1638 let context = &Ident::new(CONTEXT, op_span.resolved_at(context.span()));
1639
1640 let singletons_resolved =
1641 self.helper_resolve_singletons(node_id, op_span);
1642
1643 let arguments = &process_singletons::postprocess_singletons(
1644 op_inst.arguments_raw.clone(),
1645 singletons_resolved,
1646 );
1647
1648 let source_tag = 'a: {
1649 if let Some(tag) = self.operator_tag.get(node_id).cloned() {
1650 break 'a tag;
1651 }
1652
1653 if proc_macro::is_available() {
1654 let op_span = op_span.unwrap();
1655 break 'a format!(
1656 "loc_{}_{}_{}_{}_{}",
1657 crate::pretty_span::make_source_path_relative(
1658 &op_span.file()
1659 )
1660 .display()
1661 .to_string()
1662 .replace(|x: char| !x.is_ascii_alphanumeric(), "_"),
1663 op_span.start().line(),
1664 op_span.start().column(),
1665 op_span.end().line(),
1666 op_span.end().column(),
1667 );
1668 }
1669
1670 format!(
1671 "loc_nopath_{}_{}_{}_{}",
1672 op_span.start().line,
1673 op_span.start().column,
1674 op_span.end().line,
1675 op_span.end().column
1676 )
1677 };
1678
1679 let work_fn = format_ident!(
1680 "{}__{}__{}",
1681 ident,
1682 op_name,
1683 source_tag,
1684 span = op_span
1685 );
1686 let work_fn_async = format_ident!("{}__async", work_fn, span = op_span);
1687
1688 let context_args = WriteContextArgs {
1689 root: &root,
1690 df_ident: df_local,
1691 context,
1692 subgraph_id,
1693 node_id,
1694 loop_id,
1695 op_span,
1696 op_tag: self.operator_tag.get(node_id).cloned(),
1697 work_fn: &work_fn,
1698 work_fn_async: &work_fn_async,
1699 ident: &ident,
1700 is_pull,
1701 inputs: &inputs,
1702 outputs: &outputs,
1703 op_name,
1704 op_inst,
1705 arguments,
1706 };
1707
1708 let write_result =
1709 (op_constraints.write_fn)(&context_args, diagnostics);
1710 let OperatorWriteOutput {
1711 write_prologue,
1712 write_iterator,
1713 write_iterator_after,
1714 write_tick_end,
1715 } = write_result.unwrap_or_else(|()| {
1716 assert!(
1717 diagnostics.has_error(),
1718 "Operator `{}` returned `Err` but emitted no diagnostics, this is a bug.",
1719 op_name,
1720 );
1721 OperatorWriteOutput {
1722 write_iterator: null_write_iterator_fn(&context_args),
1723 ..Default::default()
1724 }
1725 });
1726
1727 op_prologue_code.push(syn::parse_quote! {
1728 #[allow(non_snake_case)]
1729 #[inline(always)]
1730 fn #work_fn<T>(thunk: impl ::std::ops::FnOnce() -> T) -> T {
1731 thunk()
1732 }
1733
1734 #[allow(non_snake_case)]
1735 #[inline(always)]
1736 async fn #work_fn_async<T>(
1737 thunk: impl ::std::future::Future<Output = T>,
1738 ) -> T {
1739 thunk.await
1740 }
1741 });
1742 op_prologue_code.push(write_prologue);
1743 op_tick_end_code.push(write_tick_end);
1744 subgraph_op_iter_code.push(write_iterator);
1745
1746 if include_type_guards {
1747 let type_guard = if is_pull {
1748 quote_spanned! {op_span=>
1749 let #ident = {
1750 #[allow(non_snake_case)]
1751 #[inline(always)]
1752 pub fn #work_fn<Item, Input>(input: Input)
1753 -> impl #root::dfir_pipes::pull::Pull<Item = Item, Meta = (), CanPend = Input::CanPend, CanEnd = Input::CanEnd>
1754 where
1755 Input: #root::dfir_pipes::pull::Pull<Item = Item, Meta = ()>,
1756 {
1757 #root::pin_project_lite::pin_project! {
1758 #[repr(transparent)]
1759 struct Pull<Item, Input: #root::dfir_pipes::pull::Pull<Item = Item>> {
1760 #[pin]
1761 inner: Input
1762 }
1763 }
1764
1765 impl<Item, Input> #root::dfir_pipes::pull::Pull for Pull<Item, Input>
1766 where
1767 Input: #root::dfir_pipes::pull::Pull<Item = Item>,
1768 {
1769 type Ctx<'ctx> = Input::Ctx<'ctx>;
1770
1771 type Item = Item;
1772 type Meta = Input::Meta;
1773 type CanPend = Input::CanPend;
1774 type CanEnd = Input::CanEnd;
1775
1776 #[inline(always)]
1777 fn pull(
1778 self: ::std::pin::Pin<&mut Self>,
1779 ctx: &mut Self::Ctx<'_>,
1780 ) -> #root::dfir_pipes::pull::PullStep<Self::Item, Self::Meta, Self::CanPend, Self::CanEnd> {
1781 #root::dfir_pipes::pull::Pull::pull(self.project().inner, ctx)
1782 }
1783
1784 #[inline(always)]
1785 fn size_hint(&self) -> (usize, Option<usize>) {
1786 #root::dfir_pipes::pull::Pull::size_hint(&self.inner)
1787 }
1788 }
1789
1790 Pull {
1791 inner: input
1792 }
1793 }
1794 #work_fn::<_, _>( #ident )
1795 };
1796 }
1797 } else {
1798 quote_spanned! {op_span=>
1799 let #ident = {
1800 #[allow(non_snake_case)]
1801 #[inline(always)]
1802 pub fn #work_fn<Item, Psh>(psh: Psh) -> impl #root::dfir_pipes::push::Push<Item, (), CanPend = Psh::CanPend>
1803 where
1804 Psh: #root::dfir_pipes::push::Push<Item, ()>
1805 {
1806 #root::pin_project_lite::pin_project! {
1807 #[repr(transparent)]
1808 struct PushGuard<Psh> {
1809 #[pin]
1810 inner: Psh,
1811 }
1812 }
1813
1814 impl<Item, Psh> #root::dfir_pipes::push::Push<Item, ()> for PushGuard<Psh>
1815 where
1816 Psh: #root::dfir_pipes::push::Push<Item, ()>,
1817 {
1818 type Ctx<'ctx> = Psh::Ctx<'ctx>;
1819
1820 type CanPend = Psh::CanPend;
1821
1822 #[inline(always)]
1823 fn poll_ready(
1824 self: ::std::pin::Pin<&mut Self>,
1825 ctx: &mut Self::Ctx<'_>,
1826 ) -> #root::dfir_pipes::push::PushStep<Self::CanPend> {
1827 #root::dfir_pipes::push::Push::poll_ready(self.project().inner, ctx)
1828 }
1829
1830 #[inline(always)]
1831 fn start_send(
1832 self: ::std::pin::Pin<&mut Self>,
1833 item: Item,
1834 meta: (),
1835 ) {
1836 #root::dfir_pipes::push::Push::start_send(self.project().inner, item, meta)
1837 }
1838
1839 #[inline(always)]
1840 fn poll_finalize(
1841 self: ::std::pin::Pin<&mut Self>,
1842 ctx: &mut Self::Ctx<'_>,
1843 ) -> #root::dfir_pipes::push::PushStep<Self::CanPend> {
1844 #root::dfir_pipes::push::Push::poll_finalize(self.project().inner, ctx)
1845 }
1846
1847 #[inline(always)]
1848 fn size_hint(
1849 self: ::std::pin::Pin<&mut Self>,
1850 hint: (usize, Option<usize>),
1851 ) {
1852 #root::dfir_pipes::push::Push::size_hint(self.project().inner, hint)
1853 }
1854 }
1855
1856 PushGuard {
1857 inner: psh
1858 }
1859 }
1860 #work_fn( #ident )
1861 };
1862 }
1863 };
1864 subgraph_op_iter_code.push(type_guard);
1865 }
1866 subgraph_op_iter_after_code.push(write_iterator_after);
1867 }
1868 }
1869
1870 {
1871 let pull_ident = if 0 < pull_to_push_idx {
1873 self.node_as_ident(subgraph_nodes[pull_to_push_idx - 1], false)
1874 } else {
1875 recv_port_idents[0].clone()
1877 };
1878
1879 #[rustfmt::skip]
1880 let push_ident = if let Some(&node_id) =
1881 subgraph_nodes.get(pull_to_push_idx)
1882 {
1883 self.node_as_ident(node_id, false)
1884 } else if 1 == send_port_idents.len() {
1885 send_port_idents[0].clone()
1887 } else {
1888 diagnostics.push(Diagnostic::spanned(
1889 pull_ident.span(),
1890 Level::Error,
1891 "Degenerate subgraph detected, is there a disconnected `null()` or other degenerate pipeline somewhere?",
1892 ));
1893 continue;
1894 };
1895
1896 let pivot_span = pull_ident
1898 .span()
1899 .join(push_ident.span())
1900 .unwrap_or_else(|| push_ident.span());
1901 let pivot_fn_ident = Ident::new(
1902 &format!("pivot_run_sg_{:?}", subgraph_id.data()),
1903 pivot_span,
1904 );
1905 let root = change_spans(root.clone(), pivot_span);
1906 subgraph_op_iter_code.push(quote_spanned! {pivot_span=>
1907 #[inline(always)]
1908 fn #pivot_fn_ident<Pul, Psh, Item>(pull: Pul, push: Psh)
1909 -> impl ::std::future::Future<Output = ()>
1910 where
1911 Pul: #root::dfir_pipes::pull::Pull<Item = Item>,
1912 Psh: #root::dfir_pipes::push::Push<Item, Pul::Meta>,
1913 {
1914 #root::dfir_pipes::pull::Pull::send_push(pull, push)
1915 }
1916 (#pivot_fn_ident)(#pull_ident, #push_ident).await;
1917 });
1918 }
1919 };
1920
1921 let sg_fut_ident = subgraph_id.as_ident(Span::call_site());
1925
1926 let send_metrics_code = send_hoffs
1928 .iter()
1929 .zip(send_buf_idents.iter())
1930 .zip(send_kinds.iter())
1931 .map(|((&hoff_id, buf_ident), &kind)| {
1932 let hoff_ffi = hoff_id.data().as_ffi();
1933 let len_expr = match kind {
1934 HandoffKind::Singleton | HandoffKind::Optional => {
1935 quote! { if #buf_ident.is_some() { 1 } else { 0 } }
1936 }
1937 HandoffKind::Vec => {
1938 quote! { #buf_ident.len() }
1939 }
1940 };
1941 quote! {
1942 __dfir_metrics.handoffs[
1943 #root::slotmap::KeyData::from_ffi(#hoff_ffi).into()
1944 ].curr_items_count.set(#len_expr);
1945 }
1946 })
1947 .collect::<Vec<_>>();
1948
1949 let send_hoff_make_code = send_buf_idents.iter()
1953 .zip(send_kinds.iter())
1954 .zip(send_hoffs.iter())
1955 .filter_map(|((buf_ident, &kind), &hoff_id)| {
1956 let span = buf_ident.span();
1957 if back_edge_hoffs_and_lazyness.contains_key(hoff_id) {
1958 Some(quote_spanned! {span=>
1961 #buf_ident.clear();
1962 })
1963 } else {
1964 let receiver_loop = self
1967 .node_successors(hoff_id)
1968 .next()
1969 .and_then(|(_, succ)| self.node_loop(succ));
1970 let is_exit = if let Some(sender_loop) = sg_loop {
1971 receiver_loop == self.loop_parent(sender_loop)
1972 } else {
1973 false
1974 };
1975 if is_exit {
1976 None
1978 } else {
1979 Some(match kind {
1980 HandoffKind::Vec => quote_spanned! {span=>
1981 let mut #buf_ident = #root::bumpalo::collections::Vec::new_in(&#bump_ident);
1982 },
1983 HandoffKind::Singleton | HandoffKind::Optional => quote_spanned! {span=>
1984 let mut #buf_ident = ::std::option::Option::None;
1985 },
1986 })
1987 }
1988 }
1989 })
1990 .collect::<Vec<_>>();
1991 let recv_hoff_drop_code = recv_buf_idents
1995 .iter()
1996 .zip(recv_hoffs.iter())
1997 .filter(|&(_, &hoff_id)| !back_edge_hoffs_and_lazyness.contains_key(hoff_id))
1998 .map(|(buf_ident, _)| {
1999 let span = buf_ident.span();
2000 quote_spanned! {span=>
2001 let _ = #buf_ident;
2002 }
2003 });
2004
2005 let sg_block = quote! {
2007 #( #send_hoff_make_code )*
2009
2010 let #sg_fut_ident = async {
2011 let #context = &#df;
2012 #( #recv_port_code )*
2013 #( #send_port_code )*
2014 #( #subgraph_op_iter_code )*
2015 #( #subgraph_op_iter_after_code )*
2016 };
2017 {
2018 let sg_metrics = &__dfir_metrics.subgraphs[
2020 #root::slotmap::KeyData::from_ffi(#sg_metrics_ffi).into()
2021 ];
2022 #root::scheduled::metrics::InstrumentSubgraph::new(
2023 #sg_fut_ident, sg_metrics
2024 ).await;
2025 sg_metrics.total_run_count.update(|x| x + 1);
2026
2027 #( #send_metrics_code )*
2029
2030 #( #recv_hoff_drop_code )*
2032 }
2033 };
2034 if let Some((_, body)) = loop_stack.last_mut() {
2035 body.extend(sg_block);
2036 } else {
2037 current_output.extend(sg_block);
2038 }
2039 }
2040 }
2041
2042 let gated_subgraph_code = {
2044 while let Some((closed_loop, child_body)) = loop_stack.pop() {
2045 let target = if let Some((_, parent_body)) = loop_stack.last_mut() {
2046 parent_body
2047 } else {
2048 &mut current_output
2049 };
2050 self.emit_loop_gate(
2051 closed_loop,
2052 child_body,
2053 &loop_input_handoffs,
2054 &back_edge_hoffs_and_lazyness,
2055 &loop_swap_code,
2056 target,
2057 );
2058 }
2059 current_output
2060 };
2061
2062 if diagnostics.has_error() {
2063 return Err(std::mem::take(diagnostics));
2064 }
2065 let _ = diagnostics; let (meta_graph_arg, diagnostics_arg) = if include_meta {
2068 let meta_graph_json = serde_json::to_string(&self).unwrap();
2069 let meta_graph_json = Literal::string(&meta_graph_json);
2070
2071 let serde_diagnostics: Vec<_> = diagnostics.iter().map(Diagnostic::to_serde).collect();
2072 let diagnostics_json = serde_json::to_string(&*serde_diagnostics).unwrap();
2073 let diagnostics_json = Literal::string(&diagnostics_json);
2074
2075 (
2076 quote! { Some(#meta_graph_json) },
2077 quote! { Some(#diagnostics_json) },
2078 )
2079 } else {
2080 (quote! { None }, quote! { None })
2081 };
2082
2083 let metrics_init_code = {
2085 let handoff_inits = handoff_nodes.iter().map(|&(node_id, _, _)| {
2086 let ffi = node_id.data().as_ffi();
2087 quote! {
2088 dfir_metrics.handoffs.insert(
2089 #root::slotmap::KeyData::from_ffi(#ffi).into(),
2090 ::std::default::Default::default(),
2091 );
2092 }
2093 });
2094 let subgraph_inits = all_subgraphs.iter().map(|&(sg_id, _)| {
2095 let ffi = sg_id.data().as_ffi();
2096 quote! {
2097 dfir_metrics.subgraphs.insert(
2098 #root::slotmap::KeyData::from_ffi(#ffi).into(),
2099 ::std::default::Default::default(),
2100 );
2101 }
2102 });
2103 handoff_inits.chain(subgraph_inits).collect::<Vec<_>>()
2104 };
2105
2106 let back_buffer_idents = back_buffer_idents_laziness
2108 .iter()
2109 .map(|(back_ident, _, _)| back_ident);
2110 let defer_tick_buf_idents = back_buffer_idents_laziness
2112 .iter()
2113 .map(|(_, buf_ident, _)| buf_ident);
2114 let non_lazy_schedule_idents: Vec<&Ident> = handoff_nodes
2119 .iter()
2120 .filter_map(|&(hoff_id, _, _)| {
2121 let delay_type = self.handoff_delay_type(hoff_id)?;
2122 if matches!(delay_type, DelayType::TickLazy | DelayType::LoopLazy) {
2124 return None;
2125 }
2126 let span = self.nodes[hoff_id].span();
2127 let expected_back_ident = self.hoff_back_ident(hoff_id, span);
2128 let entry = back_buffer_idents_laziness
2129 .iter()
2130 .find(|(back_ident, _, _)| *back_ident == expected_back_ident)?;
2131
2132 if delay_type == DelayType::Tick {
2134 let consumer_loop = self
2135 .node_successors(hoff_id)
2136 .next()
2137 .and_then(|(_, succ)| self.node_subgraph(succ))
2138 .and_then(|sg| self.subgraph_loop(sg));
2139 if consumer_loop.is_some_and(|lid| self.loop_parent(lid).is_none()) {
2140 return Some(&entry.0); }
2142 }
2143 Some(&entry.1) })
2145 .collect();
2146
2147 Ok(quote! {
2150 {
2151 #prefix
2152
2153 use #root::{var_expr, var_args};
2154
2155 let __dfir_wake_state = ::std::sync::Arc::new(
2156 #root::scheduled::context::WakeState::default()
2157 );
2158
2159 let __dfir_metrics = {
2160 let mut dfir_metrics = #root::scheduled::metrics::DfirMetrics::default();
2161 #( #metrics_init_code )*
2162 ::std::rc::Rc::new(dfir_metrics)
2163 };
2164
2165 #[allow(unused_mut)]
2166 let mut #df = #root::scheduled::context::Context::new(
2167 ::std::clone::Clone::clone(&__dfir_wake_state),
2168 __dfir_metrics,
2169 );
2170
2171 #( #op_prologue_code )*
2172
2173 #( let mut #back_buffer_idents = ::std::vec::Vec::new(); )*
2177 #( let mut #defer_tick_buf_idents = ::std::vec::Vec::new(); )*
2178
2179 let mut #bump_ident = #root::bumpalo::Bump::new();
2181
2182 let mut __dfir_work_done = true;
2187 #[allow(unused_qualifications, unused_mut, unused_variables, clippy::await_holding_refcell_ref, clippy::deref_addrof)]
2188 let __dfir_inline_tick = async move |#df: &mut #root::scheduled::context::Context| {
2189 #bump_ident.reset();
2191
2192 {
2193 let __dfir_metrics = #df.metrics();
2194
2195 #gated_subgraph_code
2196
2197 #[allow(clippy::nonminimal_bool, reason = "codegen")]
2200 if false #( || !#non_lazy_schedule_idents.is_empty() )* {
2201 #df.schedule_subgraph(true);
2202 }
2203
2204 #( #back_edge_swap_code )*
2207 }
2208
2209 #( #op_tick_end_code )*
2211
2212 #df.__end_tick();
2213
2214 ::std::mem::take(&mut __dfir_work_done)
2215 };
2216 #root::scheduled::context::Dfir::new(
2217 __dfir_inline_tick,
2218 #df,
2219 #meta_graph_arg,
2220 #diagnostics_arg,
2221 )
2222 }
2223 })
2224 }
2225
2226 pub fn node_color_map(&self) -> SparseSecondaryMap<GraphNodeId, Color> {
2229 let mut node_color_map: SparseSecondaryMap<GraphNodeId, Color> = self
2230 .node_ids()
2231 .filter_map(|node_id| {
2232 let op_color = self.node_color(node_id)?;
2233 Some((node_id, op_color))
2234 })
2235 .collect();
2236
2237 for sg_nodes in self.subgraph_nodes.values() {
2239 let pull_to_push_idx = self.find_pull_to_push_idx(sg_nodes);
2240
2241 for (idx, node_id) in sg_nodes.iter().copied().enumerate() {
2242 let is_pull = idx < pull_to_push_idx;
2243 node_color_map.insert(node_id, if is_pull { Color::Pull } else { Color::Push });
2244 }
2245 }
2246
2247 node_color_map
2248 }
2249
2250 pub fn to_mermaid(&self, write_config: &WriteConfig) -> String {
2252 let mut output = String::new();
2253 self.write_mermaid(&mut output, write_config).unwrap();
2254 output
2255 }
2256
2257 pub fn write_mermaid(
2259 &self,
2260 output: impl std::fmt::Write,
2261 write_config: &WriteConfig,
2262 ) -> std::fmt::Result {
2263 let mut graph_write = Mermaid::new(output);
2264 self.write_graph(&mut graph_write, write_config)
2265 }
2266
2267 pub fn to_dot(&self, write_config: &WriteConfig) -> String {
2269 let mut output = String::new();
2270 let mut graph_write = Dot::new(&mut output);
2271 self.write_graph(&mut graph_write, write_config).unwrap();
2272 output
2273 }
2274
2275 pub fn write_dot(
2277 &self,
2278 output: impl std::fmt::Write,
2279 write_config: &WriteConfig,
2280 ) -> std::fmt::Result {
2281 let mut graph_write = Dot::new(output);
2282 self.write_graph(&mut graph_write, write_config)
2283 }
2284
2285 pub(crate) fn write_graph<W>(
2287 &self,
2288 mut graph_write: W,
2289 write_config: &WriteConfig,
2290 ) -> Result<(), W::Err>
2291 where
2292 W: GraphWrite,
2293 {
2294 fn helper_edge_label(
2295 src_port: &PortIndexValue,
2296 dst_port: &PortIndexValue,
2297 ) -> Option<String> {
2298 let src_label = match src_port {
2299 PortIndexValue::Path(path) => Some(path.to_token_stream().to_string()),
2300 PortIndexValue::Int(index) => Some(index.value.to_string()),
2301 _ => None,
2302 };
2303 let dst_label = match dst_port {
2304 PortIndexValue::Path(path) => Some(path.to_token_stream().to_string()),
2305 PortIndexValue::Int(index) => Some(index.value.to_string()),
2306 _ => None,
2307 };
2308 let label = match (src_label, dst_label) {
2309 (Some(l1), Some(l2)) => Some(format!("{}\n{}", l1, l2)),
2310 (Some(l1), None) => Some(l1),
2311 (None, Some(l2)) => Some(l2),
2312 (None, None) => None,
2313 };
2314 label
2315 }
2316
2317 let node_color_map = self.node_color_map();
2319
2320 graph_write.write_prologue()?;
2322
2323 let mut skipped_handoffs = BTreeSet::new();
2325 for (node_id, node) in self.nodes() {
2326 if matches!(node, GraphNode::Handoff { .. }) && write_config.no_handoffs {
2327 skipped_handoffs.insert(node_id);
2328 continue;
2329 }
2330 graph_write.write_node_definition(
2331 node_id,
2332 &if write_config.op_short_text {
2333 node.to_name_string()
2334 } else if write_config.op_text_no_imports {
2335 let full_text = node.to_pretty_string();
2337 let mut output = String::new();
2338 for sentence in full_text.split('\n') {
2339 if sentence.trim().starts_with("use") {
2340 continue;
2341 }
2342 output.push('\n');
2343 output.push_str(sentence);
2344 }
2345 output.into()
2346 } else {
2347 node.to_pretty_string()
2348 },
2349 if write_config.no_pull_push {
2350 None
2351 } else {
2352 node_color_map.get(node_id).copied()
2353 },
2354 )?;
2355 }
2356
2357 for (edge_id, (src_id, mut dst_id)) in self.edges() {
2359 if skipped_handoffs.contains(&src_id) {
2361 continue;
2362 }
2363
2364 let (src_port, mut dst_port) = self.edge_ports(edge_id);
2365 if skipped_handoffs.contains(&dst_id) {
2366 let mut handoff_succs = self.node_successors(dst_id);
2370 if handoff_succs.len() == 0 {
2371 continue;
2372 }
2373 let (succ_edge, succ_node) = handoff_succs.next().unwrap();
2374 dst_id = succ_node;
2375 dst_port = self.edge_ports(succ_edge).1;
2376 }
2377
2378 let label = helper_edge_label(src_port, dst_port);
2379 let delay_type = self
2380 .node_op_inst(dst_id)
2381 .and_then(|op_inst| (op_inst.op_constraints.input_delaytype_fn)(dst_port));
2382 graph_write.write_edge(src_id, dst_id, delay_type, label.as_deref(), false)?;
2383 }
2384
2385 if !write_config.no_references {
2387 for dst_id in self.node_ids() {
2388 for src_ref_id in self
2389 .node_handoff_references(dst_id)
2390 .iter()
2391 .filter_map(|r| r.node_id)
2392 {
2393 let resolved_src = if skipped_handoffs.contains(&src_ref_id) {
2396 self.node_predecessor_nodes(src_ref_id).next()
2397 } else {
2398 Some(src_ref_id)
2399 };
2400 let Some(resolved_src) = resolved_src else {
2401 continue;
2402 };
2403 let label = None;
2404 graph_write.write_edge(resolved_src, dst_id, None, label, true)?;
2405 }
2406 }
2407 }
2408
2409 let loop_subgraphs = self.subgraph_ids().map(|sg_id| {
2417 let loop_id = if write_config.no_loops {
2418 None
2419 } else {
2420 self.subgraph_loop(sg_id)
2421 };
2422 (loop_id, sg_id)
2423 });
2424 let loop_subgraphs = into_group_map(loop_subgraphs);
2425 for (loop_id, subgraph_ids) in loop_subgraphs {
2426 if let Some(loop_id) = loop_id {
2427 graph_write.write_loop_start(loop_id)?;
2428 }
2429
2430 let subgraph_varnames_nodes = subgraph_ids.into_iter().flat_map(|sg_id| {
2432 self.subgraph(sg_id).iter().copied().map(move |node_id| {
2433 let opt_sg_id = if write_config.no_subgraphs {
2434 None
2435 } else {
2436 Some(sg_id)
2437 };
2438 (opt_sg_id, (self.node_varname(node_id), node_id))
2439 })
2440 });
2441 let subgraph_varnames_nodes = into_group_map(subgraph_varnames_nodes);
2442 for (sg_id, varnames) in subgraph_varnames_nodes {
2443 if let Some(sg_id) = sg_id {
2444 graph_write.write_subgraph_start(sg_id)?;
2445 }
2446
2447 let varname_nodes = varnames.into_iter().map(|(varname, node)| {
2449 let varname = if write_config.no_varnames {
2450 None
2451 } else {
2452 varname
2453 };
2454 (varname, node)
2455 });
2456 let varname_nodes = into_group_map(varname_nodes);
2457 for (varname, node_ids) in varname_nodes {
2458 if let Some(varname) = varname {
2459 graph_write.write_varname_start(&varname.0.to_string(), sg_id)?;
2460 }
2461
2462 for node_id in node_ids {
2464 graph_write.write_node(node_id)?;
2465 }
2466
2467 if varname.is_some() {
2468 graph_write.write_varname_end()?;
2469 }
2470 }
2471
2472 if sg_id.is_some() {
2473 graph_write.write_subgraph_end()?;
2474 }
2475 }
2476
2477 if loop_id.is_some() {
2478 graph_write.write_loop_end()?;
2479 }
2480 }
2481
2482 graph_write.write_epilogue()?;
2484
2485 Ok(())
2486 }
2487
2488 pub fn surface_syntax_string(&self) -> String {
2490 let mut string = String::new();
2491 self.write_surface_syntax(&mut string).unwrap();
2492 string
2493 }
2494
2495 pub fn write_surface_syntax(&self, write: &mut impl std::fmt::Write) -> std::fmt::Result {
2497 for (key, node) in self.nodes.iter() {
2498 match node {
2499 GraphNode::Operator(op) => {
2500 writeln!(write, "_{:?} = {};", key.data(), op.to_token_stream())?;
2501 }
2502 GraphNode::Handoff {
2503 kind: HandoffKind::Vec,
2504 ..
2505 } => {
2506 writeln!(write, "_{:?} = handoff();", key.data())?;
2507 }
2508 GraphNode::Handoff {
2509 kind: HandoffKind::Singleton,
2510 ..
2511 } => {
2512 writeln!(write, "_{:?} = singleton();", key.data())?;
2513 }
2514 GraphNode::Handoff {
2515 kind: HandoffKind::Optional,
2516 ..
2517 } => {
2518 writeln!(write, "_{:?} = optional();", key.data())?;
2519 }
2520 GraphNode::ModuleBoundary { .. } => panic!(),
2521 }
2522 }
2523 writeln!(write)?;
2524 for (e, (src_key, dst_key)) in self.graph.edges() {
2525 let (src_port, dst_port) = self.edge_ports(e);
2526 let src_port_str = if src_port.is_specified() {
2527 format!("[{}]", src_port)
2528 } else {
2529 String::new()
2530 };
2531 let dst_port_str = if dst_port.is_specified() {
2532 format!("[{}]", dst_port)
2533 } else {
2534 String::new()
2535 };
2536 writeln!(
2537 write,
2538 "_{:?}{} -> {}_{:?};",
2539 src_key.data(),
2540 src_port_str,
2541 dst_port_str,
2542 dst_key.data()
2543 )?;
2544 }
2545 Ok(())
2546 }
2547
2548 pub fn mermaid_string_flat(&self) -> String {
2550 let mut string = String::new();
2551 self.write_mermaid_flat(&mut string).unwrap();
2552 string
2553 }
2554
2555 pub fn write_mermaid_flat(&self, write: &mut impl std::fmt::Write) -> std::fmt::Result {
2557 writeln!(write, "flowchart TB")?;
2558 for (key, node) in self.nodes.iter() {
2559 match node {
2560 GraphNode::Operator(operator) => writeln!(
2561 write,
2562 " %% {span}\n {id:?}[\"{row_col} <tt>{code}</tt>\"]",
2563 span = PrettySpan(node.span()),
2564 id = key.data(),
2565 row_col = PrettyRowCol(node.span()),
2566 code = operator
2567 .to_token_stream()
2568 .to_string()
2569 .replace('&', "&")
2570 .replace('<', "<")
2571 .replace('>', ">")
2572 .replace('"', """)
2573 .replace('\n', "<br>"),
2574 ),
2575 GraphNode::Handoff {
2576 kind: HandoffKind::Vec,
2577 ..
2578 } => {
2579 writeln!(write, r#" {:?}{{"{}"}}"#, key.data(), HANDOFF_NODE_STR)
2580 }
2581 GraphNode::Handoff {
2582 kind: HandoffKind::Singleton | HandoffKind::Optional,
2583 ..
2584 } => {
2585 writeln!(
2586 write,
2587 r#" {:?}{{"{}"}}"#,
2588 key.data(),
2589 SINGLETON_SLOT_NODE_STR
2590 )
2591 }
2592 GraphNode::ModuleBoundary { .. } => {
2593 writeln!(
2594 write,
2595 r#" {:?}{{"{}"}}"#,
2596 key.data(),
2597 MODULE_BOUNDARY_NODE_STR
2598 )
2599 }
2600 }?;
2601 }
2602 writeln!(write)?;
2603 for (_e, (src_key, dst_key)) in self.graph.edges() {
2604 writeln!(write, " {:?}-->{:?}", src_key.data(), dst_key.data())?;
2605 }
2606 Ok(())
2607 }
2608}
2609
2610impl DfirGraph {
2612 pub fn loop_ids(&self) -> slotmap::basic::Keys<'_, GraphLoopId, Vec<GraphNodeId>> {
2614 self.loop_nodes.keys()
2615 }
2616
2617 pub fn loops(&self) -> slotmap::basic::Iter<'_, GraphLoopId, Vec<GraphNodeId>> {
2619 self.loop_nodes.iter()
2620 }
2621
2622 pub fn loop_nodes(&self, loop_id: GraphLoopId) -> &[GraphNodeId] {
2624 self.loop_nodes.get(loop_id).unwrap()
2625 }
2626
2627 pub fn insert_loop(&mut self, parent_loop: Option<GraphLoopId>) -> GraphLoopId {
2629 let loop_id = self.loop_nodes.insert(Vec::new());
2630 self.loop_children.insert(loop_id, Vec::new());
2631 if let Some(parent_loop) = parent_loop {
2632 self.loop_parent.insert(loop_id, parent_loop);
2633 self.loop_children
2634 .get_mut(parent_loop)
2635 .unwrap()
2636 .push(loop_id);
2637 } else {
2638 self.root_loops.push(loop_id);
2639 }
2640 loop_id
2641 }
2642
2643 pub fn node_loop(&self, node_id: GraphNodeId) -> Option<GraphLoopId> {
2645 self.node_loops.get(node_id).copied()
2646 }
2647
2648 pub fn subgraph_loop(&self, subgraph_id: GraphSubgraphId) -> Option<GraphLoopId> {
2650 let &node_id = self.subgraph(subgraph_id).first().unwrap();
2651 let out = self.node_loop(node_id);
2652 debug_assert!(
2653 self.subgraph(subgraph_id)
2654 .iter()
2655 .all(|&node_id| self.node_loop(node_id) == out),
2656 "Subgraph nodes should all have the same loop context."
2657 );
2658 out
2659 }
2660
2661 pub fn loop_parent(&self, loop_id: GraphLoopId) -> Option<GraphLoopId> {
2663 self.loop_parent.get(loop_id).copied()
2664 }
2665
2666 pub fn loop_children(&self, loop_id: GraphLoopId) -> &Vec<GraphLoopId> {
2668 self.loop_children.get(loop_id).unwrap()
2669 }
2670
2671 pub fn root_loops(&self) -> &[GraphLoopId] {
2673 &self.root_loops
2674 }
2675}
2676
2677#[derive(Clone, Debug, Default)]
2679#[cfg_attr(feature = "clap-derive", derive(clap::Args))]
2680pub struct WriteConfig {
2681 #[cfg_attr(feature = "clap-derive", arg(long))]
2683 pub no_subgraphs: bool,
2684 #[cfg_attr(feature = "clap-derive", arg(long))]
2686 pub no_varnames: bool,
2687 #[cfg_attr(feature = "clap-derive", arg(long))]
2689 pub no_pull_push: bool,
2690 #[cfg_attr(feature = "clap-derive", arg(long))]
2692 pub no_handoffs: bool,
2693 #[cfg_attr(feature = "clap-derive", arg(long))]
2695 pub no_references: bool,
2696 #[cfg_attr(feature = "clap-derive", arg(long))]
2698 pub no_loops: bool,
2699
2700 #[cfg_attr(feature = "clap-derive", arg(long))]
2702 pub op_short_text: bool,
2703 #[cfg_attr(feature = "clap-derive", arg(long))]
2705 pub op_text_no_imports: bool,
2706}
2707
2708#[derive(Copy, Clone, Debug)]
2710#[cfg_attr(feature = "clap-derive", derive(clap::Parser, clap::ValueEnum))]
2711pub enum WriteGraphType {
2712 Mermaid,
2714 Dot,
2716}
2717
2718fn into_group_map<K, V>(iter: impl IntoIterator<Item = (K, V)>) -> BTreeMap<K, Vec<V>>
2720where
2721 K: Ord,
2722{
2723 let mut out: BTreeMap<_, Vec<_>> = BTreeMap::new();
2724 for (k, v) in iter {
2725 out.entry(k).or_default().push(v);
2726 }
2727 out
2728}