Skip to main content

oxicuda_graph/
graph.rs

1//! The `ComputeGraph` — a directed acyclic graph of GPU operations.
2//!
3//! `ComputeGraph` is the central data structure of `oxicuda-graph`. It
4//! stores nodes (GPU operations), directed dependency edges, and buffer
5//! descriptors. Analysis and optimisation passes operate on this structure
6//! before it is lowered to an `ExecutionPlan` by the executor.
7//!
8//! # Design
9//!
10//! * Nodes are stored in a `Vec` indexed by a dense `u32` slot.
11//! * Adjacency is maintained as two maps: `successors` and `predecessors`,
12//!   each mapping `NodeId → Vec<NodeId>`. This makes both forward and
13//!   backward traversals O(degree).
14//! * Cycle detection happens eagerly on `add_edge` using DFS, so the
15//!   invariant "this graph is a DAG" always holds after successful mutation.
16
17use std::collections::{HashMap, HashSet, VecDeque};
18
19use crate::error::{GraphError, GraphResult};
20use crate::node::{BufferDescriptor, BufferId, GraphNode, NodeId};
21
22// ---------------------------------------------------------------------------
23// ComputeGraph
24// ---------------------------------------------------------------------------
25
26/// A directed acyclic graph (DAG) of GPU operations.
27///
28/// Nodes represent individual GPU operations ([`GraphNode`]). Directed edges
29/// express execution-order dependencies: an edge `a → b` means `b` cannot
30/// begin until `a` has completed. Buffers referenced by nodes are registered
31/// separately via [`add_buffer`](ComputeGraph::add_buffer).
32///
33/// # Guarantees
34///
35/// * The graph is always a DAG — [`add_edge`](ComputeGraph::add_edge) returns
36///   [`GraphError::CycleDetected`] if adding the edge would create a cycle.
37/// * Node IDs are unique and stable: nodes are never moved or re-indexed.
38/// * Buffer IDs are unique and stable.
39#[derive(Debug, Clone)]
40pub struct ComputeGraph {
41    /// Nodes stored in insertion order (indexed by NodeId.0).
42    nodes: Vec<GraphNode>,
43    /// Successor adjacency: node → nodes that depend on it.
44    successors: Vec<Vec<NodeId>>,
45    /// Predecessor adjacency: node → nodes it depends on.
46    predecessors: Vec<Vec<NodeId>>,
47    /// Buffer metadata (indexed by BufferId.0).
48    buffers: Vec<BufferDescriptor>,
49    /// Next node id to allocate.
50    next_node: u32,
51    /// Next buffer id to allocate.
52    next_buf: u32,
53}
54
55impl Default for ComputeGraph {
56    fn default() -> Self {
57        Self::new()
58    }
59}
60
61impl ComputeGraph {
62    /// Creates an empty computation graph.
63    #[must_use]
64    pub fn new() -> Self {
65        Self {
66            nodes: Vec::new(),
67            successors: Vec::new(),
68            predecessors: Vec::new(),
69            buffers: Vec::new(),
70            next_node: 0,
71            next_buf: 0,
72        }
73    }
74
75    // -----------------------------------------------------------------------
76    // Node management
77    // -----------------------------------------------------------------------
78
79    /// Adds a node to the graph and returns its assigned `NodeId`.
80    ///
81    /// The node's `id` field will be overwritten with the allocated ID.
82    pub fn add_node(&mut self, mut node: GraphNode) -> NodeId {
83        let id = NodeId(self.next_node);
84        node.id = id;
85        self.next_node += 1;
86        self.nodes.push(node);
87        self.successors.push(Vec::new());
88        self.predecessors.push(Vec::new());
89        id
90    }
91
92    /// Returns a reference to the node with the given ID, or an error.
93    pub fn node(&self, id: NodeId) -> GraphResult<&GraphNode> {
94        self.nodes
95            .get(id.0 as usize)
96            .ok_or(GraphError::NodeNotFound(id))
97    }
98
99    /// Returns a mutable reference to the node with the given ID, or an error.
100    pub fn node_mut(&mut self, id: NodeId) -> GraphResult<&mut GraphNode> {
101        self.nodes
102            .get_mut(id.0 as usize)
103            .ok_or(GraphError::NodeNotFound(id))
104    }
105
106    /// Returns a slice of all nodes in insertion order.
107    #[inline]
108    pub fn nodes(&self) -> &[GraphNode] {
109        &self.nodes
110    }
111
112    /// Returns the total number of nodes.
113    #[inline]
114    pub fn node_count(&self) -> usize {
115        self.nodes.len()
116    }
117
118    // -----------------------------------------------------------------------
119    // Buffer management
120    // -----------------------------------------------------------------------
121
122    /// Registers a buffer and returns its assigned `BufferId`.
123    ///
124    /// The buffer's `id` field will be overwritten with the allocated ID.
125    pub fn add_buffer(&mut self, mut buf: BufferDescriptor) -> BufferId {
126        let id = BufferId(self.next_buf);
127        buf.id = id;
128        self.next_buf += 1;
129        self.buffers.push(buf);
130        id
131    }
132
133    /// Returns a reference to the buffer descriptor, or an error.
134    pub fn buffer(&self, id: BufferId) -> GraphResult<&BufferDescriptor> {
135        self.buffers
136            .get(id.0 as usize)
137            .ok_or(GraphError::NodeNotFound(NodeId(id.0)))
138    }
139
140    /// Returns a slice of all buffer descriptors.
141    #[inline]
142    pub fn buffers(&self) -> &[BufferDescriptor] {
143        &self.buffers
144    }
145
146    /// Returns the number of registered buffers.
147    #[inline]
148    pub fn buffer_count(&self) -> usize {
149        self.buffers.len()
150    }
151
152    // -----------------------------------------------------------------------
153    // Edge management
154    // -----------------------------------------------------------------------
155
156    /// Adds a directed dependency edge `from → to`.
157    ///
158    /// This means `to` will not begin execution until `from` has completed.
159    ///
160    /// # Errors
161    ///
162    /// * [`GraphError::NodeNotFound`] if either ID is invalid.
163    /// * [`GraphError::CycleDetected`] if the edge would create a cycle.
164    pub fn add_edge(&mut self, from: NodeId, to: NodeId) -> GraphResult<()> {
165        let n = self.nodes.len();
166        if from.0 as usize >= n {
167            return Err(GraphError::NodeNotFound(from));
168        }
169        if to.0 as usize >= n {
170            return Err(GraphError::NodeNotFound(to));
171        }
172        if from == to {
173            return Err(GraphError::CycleDetected { from, to });
174        }
175        // Check whether adding from→to would create a cycle: it does iff
176        // `from` is reachable from `to` in the current graph.
177        if self.is_reachable(to, from) {
178            return Err(GraphError::CycleDetected { from, to });
179        }
180        // Avoid duplicate edges.
181        if !self.successors[from.0 as usize].contains(&to) {
182            self.successors[from.0 as usize].push(to);
183            self.predecessors[to.0 as usize].push(from);
184        }
185        Ok(())
186    }
187
188    /// Returns whether node `src` can reach node `dst` via directed edges.
189    ///
190    /// Uses iterative BFS to avoid stack overflow on deep graphs.
191    pub fn is_reachable(&self, src: NodeId, dst: NodeId) -> bool {
192        if src == dst {
193            return true;
194        }
195        let n = self.nodes.len();
196        let mut visited = vec![false; n];
197        let mut queue = VecDeque::new();
198        queue.push_back(src);
199        visited[src.0 as usize] = true;
200        while let Some(curr) = queue.pop_front() {
201            for &next in &self.successors[curr.0 as usize] {
202                if next == dst {
203                    return true;
204                }
205                if !visited[next.0 as usize] {
206                    visited[next.0 as usize] = true;
207                    queue.push_back(next);
208                }
209            }
210        }
211        false
212    }
213
214    /// Returns successors of `id` (nodes that depend on it).
215    pub fn successors(&self, id: NodeId) -> GraphResult<&[NodeId]> {
216        if id.0 as usize >= self.nodes.len() {
217            return Err(GraphError::NodeNotFound(id));
218        }
219        Ok(&self.successors[id.0 as usize])
220    }
221
222    /// Returns predecessors of `id` (nodes it depends on).
223    pub fn predecessors(&self, id: NodeId) -> GraphResult<&[NodeId]> {
224        if id.0 as usize >= self.nodes.len() {
225            return Err(GraphError::NodeNotFound(id));
226        }
227        Ok(&self.predecessors[id.0 as usize])
228    }
229
230    /// Returns the total number of directed edges in the graph.
231    pub fn edge_count(&self) -> usize {
232        self.successors.iter().map(|v| v.len()).sum()
233    }
234
235    /// Returns all edges as `(from, to)` pairs.
236    pub fn edges(&self) -> Vec<(NodeId, NodeId)> {
237        let mut edges = Vec::new();
238        for (i, succs) in self.successors.iter().enumerate() {
239            for &to in succs {
240                edges.push((NodeId(i as u32), to));
241            }
242        }
243        edges
244    }
245
246    // -----------------------------------------------------------------------
247    // Source / sink queries
248    // -----------------------------------------------------------------------
249
250    /// Returns all nodes with no predecessors (entry points of the graph).
251    pub fn sources(&self) -> Vec<NodeId> {
252        self.predecessors
253            .iter()
254            .enumerate()
255            .filter(|(_, preds)| preds.is_empty())
256            .map(|(i, _)| NodeId(i as u32))
257            .collect()
258    }
259
260    /// Returns all nodes with no successors (terminal nodes of the graph).
261    pub fn sinks(&self) -> Vec<NodeId> {
262        self.successors
263            .iter()
264            .enumerate()
265            .filter(|(_, succs)| succs.is_empty())
266            .map(|(i, _)| NodeId(i as u32))
267            .collect()
268    }
269
270    // -----------------------------------------------------------------------
271    // Topological sort (Kahn's algorithm)
272    // -----------------------------------------------------------------------
273
274    /// Returns nodes in topological order (all predecessors before successors).
275    ///
276    /// Uses Kahn's BFS algorithm. Because the graph is guaranteed to be a DAG
277    /// after successful `add_edge` calls, this will always succeed unless the
278    /// graph is empty.
279    ///
280    /// # Errors
281    ///
282    /// Returns [`GraphError::EmptyGraph`] if there are no nodes.
283    pub fn topological_order(&self) -> GraphResult<Vec<NodeId>> {
284        if self.nodes.is_empty() {
285            return Err(GraphError::EmptyGraph);
286        }
287        let n = self.nodes.len();
288        let mut in_degree: Vec<u32> = self.predecessors.iter().map(|p| p.len() as u32).collect();
289        let mut queue: VecDeque<NodeId> = (0..n)
290            .filter(|&i| in_degree[i] == 0)
291            .map(|i| NodeId(i as u32))
292            .collect();
293        let mut order = Vec::with_capacity(n);
294        while let Some(id) = queue.pop_front() {
295            order.push(id);
296            for &succ in &self.successors[id.0 as usize] {
297                let d = &mut in_degree[succ.0 as usize];
298                *d -= 1;
299                if *d == 0 {
300                    queue.push_back(succ);
301                }
302            }
303        }
304        // Safety: the DAG invariant guarantees all nodes are reachable.
305        debug_assert_eq!(
306            order.len(),
307            n,
308            "topological sort incomplete — internal invariant broken"
309        );
310        Ok(order)
311    }
312
313    // -----------------------------------------------------------------------
314    // Buffer-derived data-flow edges
315    // -----------------------------------------------------------------------
316
317    /// Infers and adds control edges from buffer data-flow.
318    ///
319    /// For every buffer `b`, if node `a` writes `b` (has `b` in outputs) and
320    /// node `c` reads `b` (has `b` in inputs), adds the dependency edge `a → c`.
321    ///
322    /// # Errors
323    ///
324    /// Returns [`GraphError::CycleDetected`] if any inferred edge creates a cycle.
325    pub fn infer_data_edges(&mut self) -> GraphResult<()> {
326        // Build writer map: buffer → list of nodes that write it.
327        let mut writers: HashMap<BufferId, Vec<NodeId>> = HashMap::new();
328        for node in &self.nodes {
329            for &buf in &node.outputs {
330                writers.entry(buf).or_default().push(node.id);
331            }
332        }
333        // For each node that reads a buffer, add edge from all writers → this node.
334        let reader_data: Vec<(NodeId, Vec<BufferId>)> = self
335            .nodes
336            .iter()
337            .map(|n| (n.id, n.inputs.clone()))
338            .collect();
339        for (reader_id, inputs) in reader_data {
340            for buf in inputs {
341                if let Some(node_writers) = writers.get(&buf) {
342                    for &writer_id in node_writers {
343                        if writer_id != reader_id {
344                            // Ignore duplicate-edge or cycle errors here: cycles
345                            // introduced by explicit buffer write-after-write are
346                            // a user error caught by add_edge.
347                            self.add_edge(writer_id, reader_id)?;
348                        }
349                    }
350                }
351            }
352        }
353        Ok(())
354    }
355
356    // -----------------------------------------------------------------------
357    // Subgraph extraction
358    // -----------------------------------------------------------------------
359
360    /// Returns the set of all nodes reachable (forward) from the given roots.
361    pub fn reachable_from(&self, roots: &[NodeId]) -> HashSet<NodeId> {
362        let mut visited = HashSet::new();
363        let mut stack: Vec<NodeId> = roots.to_vec();
364        while let Some(id) = stack.pop() {
365            if visited.insert(id) {
366                for &s in &self.successors[id.0 as usize] {
367                    if !visited.contains(&s) {
368                        stack.push(s);
369                    }
370                }
371            }
372        }
373        visited
374    }
375
376    /// Returns the set of all nodes that can reach any of the given targets (reverse reachability).
377    pub fn reaching(&self, targets: &[NodeId]) -> HashSet<NodeId> {
378        let mut visited = HashSet::new();
379        let mut stack: Vec<NodeId> = targets.to_vec();
380        while let Some(id) = stack.pop() {
381            if visited.insert(id) {
382                for &p in &self.predecessors[id.0 as usize] {
383                    if !visited.contains(&p) {
384                        stack.push(p);
385                    }
386                }
387            }
388        }
389        visited
390    }
391
392    // -----------------------------------------------------------------------
393    // Graph properties
394    // -----------------------------------------------------------------------
395
396    /// Returns `true` if the graph has no nodes.
397    #[inline]
398    pub fn is_empty(&self) -> bool {
399        self.nodes.is_empty()
400    }
401
402    /// Returns the longest path length (in edges) from any source to any sink.
403    ///
404    /// This is the critical path length — a lower bound on the sequential
405    /// execution depth of the graph.
406    ///
407    /// # Errors
408    ///
409    /// Returns [`GraphError::EmptyGraph`] if the graph is empty.
410    pub fn critical_path_length(&self) -> GraphResult<usize> {
411        let order = self.topological_order()?;
412        let n = self.nodes.len();
413        let mut dist = vec![0usize; n];
414        let mut max_len = 0usize;
415        for id in &order {
416            let d = dist[id.0 as usize];
417            max_len = max_len.max(d);
418            for &succ in &self.successors[id.0 as usize] {
419                let nd = d + 1;
420                if nd > dist[succ.0 as usize] {
421                    dist[succ.0 as usize] = nd;
422                }
423            }
424        }
425        Ok(max_len)
426    }
427
428    /// Returns the maximum fan-in (in-degree) among all nodes.
429    pub fn max_in_degree(&self) -> usize {
430        self.predecessors.iter().map(|p| p.len()).max().unwrap_or(0)
431    }
432
433    /// Returns the maximum fan-out (out-degree) among all nodes.
434    pub fn max_out_degree(&self) -> usize {
435        self.successors.iter().map(|s| s.len()).max().unwrap_or(0)
436    }
437
438    /// Returns the number of parallel chains (independent execution paths).
439    ///
440    /// This is an upper bound on the number of streams that can be usefully
441    /// exploited for concurrent execution.
442    pub fn parallelism_width(&self) -> GraphResult<usize> {
443        // Width = maximum number of nodes at the same topological "level"
444        // (BFS layer from sources).
445        if self.nodes.is_empty() {
446            return Ok(0);
447        }
448        let n = self.nodes.len();
449        let mut level = vec![0usize; n];
450        let mut in_degree: Vec<u32> = self.predecessors.iter().map(|p| p.len() as u32).collect();
451        let mut queue: VecDeque<NodeId> = (0..n)
452            .filter(|&i| in_degree[i] == 0)
453            .map(|i| NodeId(i as u32))
454            .collect();
455        let mut max_width = queue.len();
456        while let Some(id) = queue.pop_front() {
457            for &succ in &self.successors[id.0 as usize] {
458                let nl = level[id.0 as usize] + 1;
459                if nl > level[succ.0 as usize] {
460                    level[succ.0 as usize] = nl;
461                }
462                let d = &mut in_degree[succ.0 as usize];
463                *d -= 1;
464                if *d == 0 {
465                    queue.push_back(succ);
466                }
467            }
468            // Compute width at each level by scanning all level values.
469        }
470        // Count nodes per level.
471        let max_level = *level.iter().max().unwrap_or(&0);
472        let mut width_at_level = vec![0usize; max_level + 1];
473        for &lv in &level {
474            width_at_level[lv] += 1;
475        }
476        max_width = max_width.max(*width_at_level.iter().max().unwrap_or(&0));
477        Ok(max_width)
478    }
479
480    // -----------------------------------------------------------------------
481    // Kernel-type queries
482    // -----------------------------------------------------------------------
483
484    /// Returns all nodes that are kernel launches (compute nodes).
485    pub fn kernel_nodes(&self) -> Vec<NodeId> {
486        self.nodes
487            .iter()
488            .filter(|n| n.kind.is_compute())
489            .map(|n| n.id)
490            .collect()
491    }
492
493    /// Returns all nodes that are fusible kernel launches.
494    pub fn fusible_nodes(&self) -> Vec<NodeId> {
495        self.nodes
496            .iter()
497            .filter(|n| n.kind.is_fusible())
498            .map(|n| n.id)
499            .collect()
500    }
501
502    // -----------------------------------------------------------------------
503    // DOT format serialisation
504    // -----------------------------------------------------------------------
505
506    /// Renders the graph in Graphviz DOT format for visualisation.
507    pub fn to_dot(&self) -> String {
508        let mut s = String::from("digraph ComputeGraph {\n  rankdir=TB;\n");
509        for node in &self.nodes {
510            let label = node.display_name();
511            let shape = if node.kind.is_compute() {
512                "box"
513            } else if node.kind.is_memory_op() {
514                "parallelogram"
515            } else {
516                "ellipse"
517            };
518            s.push_str(&format!(
519                "  {} [label=\"{} ({})\", shape={shape}];\n",
520                node.id.0,
521                label,
522                node.kind.tag()
523            ));
524        }
525        for (i, succs) in self.successors.iter().enumerate() {
526            for &to in succs {
527                s.push_str(&format!("  {} -> {};\n", i, to.0));
528            }
529        }
530        s.push('}');
531        s
532    }
533}
534
535impl std::fmt::Display for ComputeGraph {
536    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
537        write!(
538            f,
539            "ComputeGraph({} nodes, {} edges, {} buffers)",
540            self.node_count(),
541            self.edge_count(),
542            self.buffer_count()
543        )
544    }
545}
546
547// ---------------------------------------------------------------------------
548// Tests
549// ---------------------------------------------------------------------------
550
551#[cfg(test)]
552mod tests {
553    use super::*;
554    use crate::node::{BufferDescriptor, KernelConfig, MemcpyDir, NodeKind};
555
556    fn kernel_node(name: &str) -> GraphNode {
557        GraphNode::new(
558            NodeId(0),
559            NodeKind::KernelLaunch {
560                function_name: name.into(),
561                config: KernelConfig::linear(1, 32, 0),
562                fusible: true,
563            },
564        )
565    }
566
567    fn barrier_node() -> GraphNode {
568        GraphNode::new(NodeId(0), NodeKind::Barrier)
569    }
570
571    fn memcpy_node(dir: MemcpyDir, size: usize) -> GraphNode {
572        GraphNode::new(
573            NodeId(0),
574            NodeKind::Memcpy {
575                dir,
576                size_bytes: size,
577            },
578        )
579    }
580
581    // --- Basic construction ---
582
583    #[test]
584    fn new_graph_is_empty() {
585        let g = ComputeGraph::new();
586        assert!(g.is_empty());
587        assert_eq!(g.node_count(), 0);
588        assert_eq!(g.edge_count(), 0);
589        assert_eq!(g.buffer_count(), 0);
590    }
591
592    #[test]
593    fn default_is_empty() {
594        let g = ComputeGraph::default();
595        assert!(g.is_empty());
596    }
597
598    #[test]
599    fn add_node_assigns_sequential_ids() {
600        let mut g = ComputeGraph::new();
601        let a = g.add_node(barrier_node());
602        let b = g.add_node(barrier_node());
603        let c = g.add_node(barrier_node());
604        assert_eq!(a, NodeId(0));
605        assert_eq!(b, NodeId(1));
606        assert_eq!(c, NodeId(2));
607        assert_eq!(g.node_count(), 3);
608    }
609
610    #[test]
611    fn node_lookup_valid() {
612        let mut g = ComputeGraph::new();
613        let id = g.add_node(kernel_node("add"));
614        assert!(g.node(id).is_ok());
615        assert_eq!(
616            g.node(id)
617                .expect("node registered in graph")
618                .kind
619                .function_name(),
620            Some("add")
621        );
622    }
623
624    #[test]
625    fn node_lookup_invalid() {
626        let g = ComputeGraph::new();
627        assert!(matches!(
628            g.node(NodeId(0)),
629            Err(GraphError::NodeNotFound(_))
630        ));
631    }
632
633    #[test]
634    fn node_mut_allows_modification() {
635        let mut g = ComputeGraph::new();
636        let id = g.add_node(barrier_node());
637        g.node_mut(id).expect("node registered in graph").cost_hint = 42;
638        assert_eq!(g.node(id).expect("node registered in graph").cost_hint, 42);
639    }
640
641    // --- Buffer management ---
642
643    #[test]
644    fn add_buffer_assigns_ids() {
645        let mut g = ComputeGraph::new();
646        let b0 = g.add_buffer(BufferDescriptor::new(BufferId(0), 1024));
647        let b1 = g.add_buffer(BufferDescriptor::new(BufferId(0), 2048));
648        assert_eq!(b0, BufferId(0));
649        assert_eq!(b1, BufferId(1));
650        assert_eq!(g.buffer_count(), 2);
651    }
652
653    #[test]
654    fn buffer_lookup_invalid() {
655        let g = ComputeGraph::new();
656        assert!(g.buffer(BufferId(0)).is_err());
657    }
658
659    // --- Edge management ---
660
661    #[test]
662    fn add_edge_valid() {
663        let mut g = ComputeGraph::new();
664        let a = g.add_node(kernel_node("a"));
665        let b = g.add_node(kernel_node("b"));
666        assert!(g.add_edge(a, b).is_ok());
667        assert_eq!(g.edge_count(), 1);
668        assert_eq!(g.successors(a).expect("node a registered in graph"), &[b]);
669        assert_eq!(g.predecessors(b).expect("node b registered in graph"), &[a]);
670    }
671
672    #[test]
673    fn add_edge_self_loop_rejected() {
674        let mut g = ComputeGraph::new();
675        let a = g.add_node(barrier_node());
676        assert!(matches!(
677            g.add_edge(a, a),
678            Err(GraphError::CycleDetected { .. })
679        ));
680    }
681
682    #[test]
683    fn add_edge_cycle_rejected() {
684        let mut g = ComputeGraph::new();
685        let a = g.add_node(barrier_node());
686        let b = g.add_node(barrier_node());
687        g.add_edge(a, b).expect("valid DAG edge from a to b");
688        assert!(matches!(
689            g.add_edge(b, a),
690            Err(GraphError::CycleDetected { .. })
691        ));
692    }
693
694    #[test]
695    fn add_edge_invalid_node_rejected() {
696        let mut g = ComputeGraph::new();
697        let a = g.add_node(barrier_node());
698        assert!(matches!(
699            g.add_edge(a, NodeId(99)),
700            Err(GraphError::NodeNotFound(_))
701        ));
702        assert!(matches!(
703            g.add_edge(NodeId(99), a),
704            Err(GraphError::NodeNotFound(_))
705        ));
706    }
707
708    #[test]
709    fn add_edge_duplicate_is_idempotent() {
710        let mut g = ComputeGraph::new();
711        let a = g.add_node(barrier_node());
712        let b = g.add_node(barrier_node());
713        g.add_edge(a, b).expect("valid DAG edge from a to b");
714        g.add_edge(a, b).expect("valid DAG edge from a to b"); // second call must succeed and NOT duplicate
715        assert_eq!(g.edge_count(), 1);
716    }
717
718    #[test]
719    fn edges_returns_all() {
720        let mut g = ComputeGraph::new();
721        let a = g.add_node(barrier_node());
722        let b = g.add_node(barrier_node());
723        let c = g.add_node(barrier_node());
724        g.add_edge(a, b).expect("valid DAG edge from a to b");
725        g.add_edge(a, c).expect("valid DAG edge from a to c");
726        let mut edges = g.edges();
727        edges.sort();
728        assert_eq!(edges, vec![(a, b), (a, c)]);
729    }
730
731    // --- Reachability ---
732
733    #[test]
734    fn is_reachable_direct() {
735        let mut g = ComputeGraph::new();
736        let a = g.add_node(barrier_node());
737        let b = g.add_node(barrier_node());
738        g.add_edge(a, b).expect("valid DAG edge from a to b");
739        assert!(g.is_reachable(a, b));
740        assert!(!g.is_reachable(b, a));
741    }
742
743    #[test]
744    fn is_reachable_transitive() {
745        let mut g = ComputeGraph::new();
746        let a = g.add_node(barrier_node());
747        let b = g.add_node(barrier_node());
748        let c = g.add_node(barrier_node());
749        g.add_edge(a, b).expect("valid DAG edge from a to b");
750        g.add_edge(b, c).expect("valid DAG edge from b to c");
751        assert!(g.is_reachable(a, c));
752        assert!(!g.is_reachable(c, a));
753    }
754
755    #[test]
756    fn is_reachable_disconnected() {
757        let mut g = ComputeGraph::new();
758        let a = g.add_node(barrier_node());
759        let b = g.add_node(barrier_node());
760        assert!(!g.is_reachable(a, b));
761        assert!(!g.is_reachable(b, a));
762    }
763
764    // --- Sources / sinks ---
765
766    #[test]
767    fn sources_and_sinks_linear_chain() {
768        let mut g = ComputeGraph::new();
769        let a = g.add_node(barrier_node());
770        let b = g.add_node(barrier_node());
771        let c = g.add_node(barrier_node());
772        g.add_edge(a, b).expect("valid DAG edge from a to b");
773        g.add_edge(b, c).expect("valid DAG edge from b to c");
774        let sources = g.sources();
775        let sinks = g.sinks();
776        assert_eq!(sources, vec![a]);
777        assert_eq!(sinks, vec![c]);
778    }
779
780    #[test]
781    fn sources_and_sinks_diamond() {
782        let mut g = ComputeGraph::new();
783        let a = g.add_node(barrier_node()); // source
784        let b = g.add_node(barrier_node());
785        let c = g.add_node(barrier_node());
786        let d = g.add_node(barrier_node()); // sink
787        g.add_edge(a, b).expect("valid DAG edge from a to b");
788        g.add_edge(a, c).expect("valid DAG edge from a to c");
789        g.add_edge(b, d).expect("valid DAG edge from b to d");
790        g.add_edge(c, d).expect("valid DAG edge from c to d");
791        assert_eq!(g.sources(), vec![a]);
792        assert_eq!(g.sinks(), vec![d]);
793    }
794
795    // --- Topological sort ---
796
797    #[test]
798    fn topological_order_linear() {
799        let mut g = ComputeGraph::new();
800        let a = g.add_node(barrier_node());
801        let b = g.add_node(barrier_node());
802        let c = g.add_node(barrier_node());
803        g.add_edge(a, b).expect("valid DAG edge from a to b");
804        g.add_edge(b, c).expect("valid DAG edge from b to c");
805        let order = g
806            .topological_order()
807            .expect("topological sort of valid DAG");
808        let pos_a = order
809            .iter()
810            .position(|&x| x == a)
811            .expect("node a present in topological order");
812        let pos_b = order
813            .iter()
814            .position(|&x| x == b)
815            .expect("node b present in topological order");
816        let pos_c = order
817            .iter()
818            .position(|&x| x == c)
819            .expect("node c present in topological order");
820        assert!(pos_a < pos_b && pos_b < pos_c);
821    }
822
823    #[test]
824    fn topological_order_diamond() {
825        let mut g = ComputeGraph::new();
826        let a = g.add_node(barrier_node());
827        let b = g.add_node(barrier_node());
828        let c = g.add_node(barrier_node());
829        let d = g.add_node(barrier_node());
830        g.add_edge(a, b).expect("valid DAG edge from a to b");
831        g.add_edge(a, c).expect("valid DAG edge from a to c");
832        g.add_edge(b, d).expect("valid DAG edge from b to d");
833        g.add_edge(c, d).expect("valid DAG edge from c to d");
834        let order = g
835            .topological_order()
836            .expect("topological sort of valid DAG");
837        assert_eq!(order.len(), 4);
838        let pos = |n: NodeId| {
839            order
840                .iter()
841                .position(|&x| x == n)
842                .expect("node present in topological order")
843        };
844        assert!(pos(a) < pos(b));
845        assert!(pos(a) < pos(c));
846        assert!(pos(b) < pos(d));
847        assert!(pos(c) < pos(d));
848    }
849
850    #[test]
851    fn topological_order_empty_graph() {
852        let g = ComputeGraph::new();
853        assert!(matches!(g.topological_order(), Err(GraphError::EmptyGraph)));
854    }
855
856    #[test]
857    fn topological_order_isolated_nodes() {
858        let mut g = ComputeGraph::new();
859        g.add_node(barrier_node());
860        g.add_node(barrier_node());
861        g.add_node(barrier_node());
862        let order = g
863            .topological_order()
864            .expect("topological sort of valid DAG");
865        assert_eq!(order.len(), 3);
866    }
867
868    // --- Data-flow edge inference ---
869
870    #[test]
871    fn infer_data_edges_connects_writer_to_reader() {
872        let mut g = ComputeGraph::new();
873        let buf = g.add_buffer(BufferDescriptor::new(BufferId(0), 1024));
874        let writer = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_outputs([buf]));
875        let reader = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_inputs([buf]));
876        g.infer_data_edges()
877            .expect("data edge inference on valid graph");
878        assert!(g.is_reachable(writer, reader));
879    }
880
881    #[test]
882    fn infer_data_edges_multiple_readers() {
883        let mut g = ComputeGraph::new();
884        let buf = g.add_buffer(BufferDescriptor::new(BufferId(0), 1024));
885        let writer = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_outputs([buf]));
886        let r1 = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_inputs([buf]));
887        let r2 = g.add_node(GraphNode::new(NodeId(0), NodeKind::Barrier).with_inputs([buf]));
888        g.infer_data_edges()
889            .expect("data edge inference on valid graph");
890        assert!(g.is_reachable(writer, r1));
891        assert!(g.is_reachable(writer, r2));
892    }
893
894    // --- Critical path, degrees ---
895
896    #[test]
897    fn critical_path_linear_chain() {
898        let mut g = ComputeGraph::new();
899        let a = g.add_node(barrier_node());
900        let b = g.add_node(barrier_node());
901        let c = g.add_node(barrier_node());
902        let d = g.add_node(barrier_node());
903        g.add_edge(a, b).expect("valid DAG edge from a to b");
904        g.add_edge(b, c).expect("valid DAG edge from b to c");
905        g.add_edge(c, d).expect("valid DAG edge from c to d");
906        assert_eq!(
907            g.critical_path_length()
908                .expect("critical path length of valid DAG"),
909            3
910        );
911    }
912
913    #[test]
914    fn critical_path_diamond() {
915        let mut g = ComputeGraph::new();
916        let a = g.add_node(barrier_node());
917        let b = g.add_node(barrier_node());
918        let c = g.add_node(barrier_node());
919        let d = g.add_node(barrier_node());
920        g.add_edge(a, b).expect("valid DAG edge from a to b");
921        g.add_edge(a, c).expect("valid DAG edge from a to c");
922        g.add_edge(b, d).expect("valid DAG edge from b to d");
923        g.add_edge(c, d).expect("valid DAG edge from c to d");
924        assert_eq!(
925            g.critical_path_length()
926                .expect("critical path length of valid DAG"),
927            2
928        );
929    }
930
931    #[test]
932    fn max_degrees_computed_correctly() {
933        let mut g = ComputeGraph::new();
934        let a = g.add_node(barrier_node());
935        let b = g.add_node(barrier_node());
936        let c = g.add_node(barrier_node());
937        let d = g.add_node(barrier_node());
938        g.add_edge(a, b).expect("valid DAG edge from a to b");
939        g.add_edge(a, c).expect("valid DAG edge from a to c");
940        g.add_edge(a, d).expect("valid DAG edge from a to d");
941        assert_eq!(g.max_out_degree(), 3);
942        assert_eq!(g.max_in_degree(), 1);
943    }
944
945    // --- Subgraph reachability ---
946
947    #[test]
948    fn reachable_from_set() {
949        let mut g = ComputeGraph::new();
950        let a = g.add_node(barrier_node());
951        let b = g.add_node(barrier_node());
952        let c = g.add_node(barrier_node());
953        let d = g.add_node(barrier_node());
954        g.add_edge(a, b).expect("valid DAG edge from a to b");
955        g.add_edge(b, c).expect("valid DAG edge from b to c");
956        // d is isolated
957        let reach = g.reachable_from(&[a]);
958        assert!(reach.contains(&a));
959        assert!(reach.contains(&b));
960        assert!(reach.contains(&c));
961        assert!(!reach.contains(&d));
962    }
963
964    #[test]
965    fn reaching_set() {
966        let mut g = ComputeGraph::new();
967        let a = g.add_node(barrier_node());
968        let b = g.add_node(barrier_node());
969        let c = g.add_node(barrier_node());
970        g.add_edge(a, b).expect("valid DAG edge from a to b");
971        g.add_edge(b, c).expect("valid DAG edge from b to c");
972        let reaching = g.reaching(&[c]);
973        assert!(reaching.contains(&a));
974        assert!(reaching.contains(&b));
975        assert!(reaching.contains(&c));
976    }
977
978    // --- Kernel / fusible queries ---
979
980    #[test]
981    fn kernel_nodes_returns_compute_only() {
982        let mut g = ComputeGraph::new();
983        g.add_node(kernel_node("k0"));
984        g.add_node(barrier_node());
985        g.add_node(memcpy_node(MemcpyDir::HostToDevice, 1024));
986        g.add_node(kernel_node("k1"));
987        let kernels = g.kernel_nodes();
988        assert_eq!(kernels.len(), 2);
989    }
990
991    #[test]
992    fn fusible_nodes_returns_fusible_only() {
993        let mut g = ComputeGraph::new();
994        g.add_node(kernel_node("fusible")); // fusible=true in helper
995        g.add_node(GraphNode::new(
996            NodeId(0),
997            NodeKind::KernelLaunch {
998                function_name: "custom".into(),
999                config: KernelConfig::linear(1, 32, 0),
1000                fusible: false,
1001            },
1002        ));
1003        assert_eq!(g.fusible_nodes().len(), 1);
1004    }
1005
1006    // --- Parallelism width ---
1007
1008    #[test]
1009    fn parallelism_width_linear_is_one() {
1010        let mut g = ComputeGraph::new();
1011        let a = g.add_node(barrier_node());
1012        let b = g.add_node(barrier_node());
1013        let c = g.add_node(barrier_node());
1014        g.add_edge(a, b).expect("valid DAG edge from a to b");
1015        g.add_edge(b, c).expect("valid DAG edge from b to c");
1016        assert_eq!(
1017            g.parallelism_width()
1018                .expect("parallelism width of valid DAG"),
1019            1
1020        );
1021    }
1022
1023    #[test]
1024    fn parallelism_width_fork_join() {
1025        let mut g = ComputeGraph::new();
1026        let src = g.add_node(barrier_node());
1027        let b = g.add_node(barrier_node());
1028        let c = g.add_node(barrier_node());
1029        let d = g.add_node(barrier_node());
1030        let sink = g.add_node(barrier_node());
1031        g.add_edge(src, b).expect("valid DAG edge from src to b");
1032        g.add_edge(src, c).expect("valid DAG edge from src to c");
1033        g.add_edge(src, d).expect("valid DAG edge from src to d");
1034        g.add_edge(b, sink).expect("valid DAG edge from b to sink");
1035        g.add_edge(c, sink).expect("valid DAG edge from c to sink");
1036        g.add_edge(d, sink).expect("valid DAG edge from d to sink");
1037        assert_eq!(
1038            g.parallelism_width()
1039                .expect("parallelism width of valid DAG"),
1040            3
1041        );
1042    }
1043
1044    // --- DOT output ---
1045
1046    #[test]
1047    fn to_dot_contains_node_labels() {
1048        let mut g = ComputeGraph::new();
1049        let a = g.add_node(kernel_node("my_kernel").with_name("k0"));
1050        let b = g.add_node(barrier_node());
1051        g.add_edge(a, b).expect("valid DAG edge from a to b");
1052        let dot = g.to_dot();
1053        assert!(dot.contains("digraph"));
1054        assert!(dot.contains("k0"));
1055        assert!(dot.contains("->"));
1056    }
1057
1058    // --- Display ---
1059
1060    #[test]
1061    fn display_shows_counts() {
1062        let mut g = ComputeGraph::new();
1063        g.add_node(barrier_node());
1064        g.add_node(barrier_node());
1065        g.add_buffer(BufferDescriptor::new(BufferId(0), 1));
1066        let s = g.to_string();
1067        assert!(s.contains("2 nodes"));
1068        assert!(s.contains("1 buffers"));
1069    }
1070}