Skip to main content

nxpu_analysis/
dataflow.rs

1//! Dataflow graph construction and analysis.
2//!
3//! Builds a dataflow graph (DFG) from an IR function by analyzing
4//! read/write sets of each statement and computing data dependencies
5//! (RAW, WAR, WAW) and control dependencies (barriers).
6//!
7//! Provides topological sorting and critical path analysis.
8
9use std::collections::{BTreeSet, HashSet, VecDeque};
10use std::fmt;
11
12use nxpu_ir::{Expression, Function, Handle, Statement};
13
14/// Errors during dataflow analysis.
15#[derive(Debug, thiserror::Error)]
16pub enum DataflowError {
17    /// The DFG contains a cycle, which should not occur in well-formed SSA IR.
18    #[error("cycle detected in dataflow graph ({visited} of {total} nodes visited)")]
19    CycleDetected { visited: usize, total: usize },
20}
21
22/// The kind of a DFG node, indicating what type of statement it represents.
23#[derive(Clone, Debug, PartialEq, Eq)]
24pub enum DfgNodeKind {
25    /// A `Store` statement (write through a pointer).
26    Store,
27    /// A function `Call` statement.
28    Call,
29    /// An `Atomic` operation statement.
30    Atomic,
31    /// A synchronization `Barrier` statement.
32    Barrier,
33    /// An `Emit` statement (makes expression results available).
34    Emit,
35    /// An `If` control flow statement.
36    If,
37    /// A `Loop` control flow statement.
38    Loop,
39    /// A `Return` statement.
40    Return,
41    /// A `Break` statement.
42    Break,
43    /// A `Continue` statement.
44    Continue,
45}
46
47impl fmt::Display for DfgNodeKind {
48    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
49        f.write_str(match self {
50            Self::Store => "Store",
51            Self::Call => "Call",
52            Self::Atomic => "Atomic",
53            Self::Barrier => "Barrier",
54            Self::Emit => "Emit",
55            Self::If => "If",
56            Self::Loop => "Loop",
57            Self::Return => "Return",
58            Self::Break => "Break",
59            Self::Continue => "Continue",
60        })
61    }
62}
63
64/// The kind of dependency between two DFG nodes.
65#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
66pub enum DependencyKind {
67    /// Read-after-write: consumer reads what producer wrote.
68    DataFlow,
69    /// Write-after-read: producer reads, consumer writes to same location.
70    AntiDependency,
71    /// Write-after-write: both nodes write to the same location.
72    OutputDependency,
73    /// Control dependency (barriers, control flow).
74    Control,
75}
76
77impl fmt::Display for DependencyKind {
78    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
79        f.write_str(match self {
80            Self::DataFlow => "RAW",
81            Self::AntiDependency => "WAR",
82            Self::OutputDependency => "WAW",
83            Self::Control => "Control",
84        })
85    }
86}
87
88/// A node in the dataflow graph representing a statement.
89#[derive(Clone, Debug)]
90pub struct DfgNode {
91    /// Unique node identifier (index into the DFG node array).
92    pub id: usize,
93    /// What kind of statement this node represents.
94    pub kind: DfgNodeKind,
95    /// Expression handles read by this statement.
96    pub reads: Vec<Handle<Expression>>,
97    /// Expression handles written by this statement.
98    pub writes: Vec<Handle<Expression>>,
99}
100
101/// An edge in the dataflow graph representing a dependency.
102#[derive(Clone, Debug)]
103pub struct DfgEdge {
104    /// Index of the producer (source) node.
105    pub from: usize,
106    /// Index of the consumer (destination) node.
107    pub to: usize,
108    /// Kind of dependency.
109    pub kind: DependencyKind,
110}
111
112/// A dataflow graph built from an IR function's statements.
113///
114/// Nodes correspond to statements and edges represent data/control
115/// dependencies between them.
116#[derive(Clone, Debug)]
117pub struct DataflowGraph {
118    /// Nodes in the graph, indexed by their `id`.
119    nodes: Vec<DfgNode>,
120    /// Edges representing dependencies.
121    edges: Vec<DfgEdge>,
122}
123
124impl DataflowGraph {
125    /// Build a dataflow graph from an IR function.
126    ///
127    /// Walks statements in order, computes read/write sets, and creates
128    /// dependency edges:
129    /// - RAW (DataFlow): a write followed by a read of the same expression
130    /// - WAR (AntiDependency): a read followed by a write of the same expression
131    /// - WAW (OutputDependency): two writes to the same expression
132    /// - Control: barriers create control edges to all subsequent memory operations
133    pub fn build(func: &Function) -> Self {
134        let mut nodes = Vec::new();
135        let mut edges = Vec::new();
136
137        // Build nodes from top-level statements.
138        for stmt in &func.body {
139            let id = nodes.len();
140            let (kind, reads, writes) = classify_statement(stmt, &func.expressions);
141            nodes.push(DfgNode {
142                id,
143                kind,
144                reads,
145                writes,
146            });
147        }
148
149        // Build dependency edges by comparing read/write sets.
150        for (i, node_i) in nodes.iter().enumerate() {
151            for (j, node_j) in nodes.iter().enumerate().skip(i + 1) {
152                // RAW: node i writes, node j reads the same handle.
153                for w in &node_i.writes {
154                    if node_j.reads.contains(w) {
155                        edges.push(DfgEdge {
156                            from: i,
157                            to: j,
158                            kind: DependencyKind::DataFlow,
159                        });
160                    }
161                }
162
163                // WAR: node i reads, node j writes the same handle.
164                for r in &node_i.reads {
165                    if node_j.writes.contains(r) {
166                        edges.push(DfgEdge {
167                            from: i,
168                            to: j,
169                            kind: DependencyKind::AntiDependency,
170                        });
171                    }
172                }
173
174                // WAW: node i writes, node j writes the same handle.
175                for w in &node_i.writes {
176                    if node_j.writes.contains(w) {
177                        edges.push(DfgEdge {
178                            from: i,
179                            to: j,
180                            kind: DependencyKind::OutputDependency,
181                        });
182                    }
183                }
184            }
185        }
186
187        // Barriers create control edges to all subsequent memory operations.
188        for (i, node_i) in nodes.iter().enumerate() {
189            if node_i.kind == DfgNodeKind::Barrier {
190                for (j, node_j) in nodes.iter().enumerate().skip(i + 1) {
191                    let has_memory_effect = !node_j.reads.is_empty()
192                        || !node_j.writes.is_empty()
193                        || node_j.kind == DfgNodeKind::Barrier;
194                    if has_memory_effect {
195                        edges.push(DfgEdge {
196                            from: i,
197                            to: j,
198                            kind: DependencyKind::Control,
199                        });
200                    }
201                }
202            }
203        }
204
205        Self { nodes, edges }
206    }
207
208    /// Returns the nodes in this graph.
209    pub fn nodes(&self) -> &[DfgNode] {
210        &self.nodes
211    }
212
213    /// Returns the edges in this graph.
214    pub fn edges(&self) -> &[DfgEdge] {
215        &self.edges
216    }
217
218    /// Returns the number of nodes.
219    pub fn node_count(&self) -> usize {
220        self.nodes.len()
221    }
222
223    /// Returns the number of edges.
224    pub fn edge_count(&self) -> usize {
225        self.edges.len()
226    }
227
228    /// Returns all edges originating from the given node.
229    pub fn successors(&self, node_id: usize) -> Vec<&DfgEdge> {
230        self.edges.iter().filter(|e| e.from == node_id).collect()
231    }
232
233    /// Returns all edges pointing to the given node.
234    pub fn predecessors(&self, node_id: usize) -> Vec<&DfgEdge> {
235        self.edges.iter().filter(|e| e.to == node_id).collect()
236    }
237
238    /// Perform topological sort using Kahn's algorithm.
239    ///
240    /// Returns nodes in a valid execution order respecting all dependencies.
241    ///
242    /// # Errors
243    ///
244    /// Returns [`DataflowError::CycleDetected`] if the graph contains a cycle.
245    pub fn topological_sort(&self) -> Result<Vec<usize>, DataflowError> {
246        let n = self.nodes.len();
247        if n == 0 {
248            return Ok(Vec::new());
249        }
250
251        // Build in-degree and adjacency.
252        let mut in_degree = vec![0usize; n];
253        let mut successors: Vec<Vec<usize>> = vec![Vec::new(); n];
254
255        for edge in &self.edges {
256            in_degree[edge.to] += 1;
257            successors[edge.from].push(edge.to);
258        }
259
260        // Use BTreeSet for deterministic ordering (lower id first).
261        let mut ready: BTreeSet<usize> = BTreeSet::new();
262        for (i, &deg) in in_degree.iter().enumerate() {
263            if deg == 0 {
264                ready.insert(i);
265            }
266        }
267
268        let mut result = Vec::with_capacity(n);
269
270        while let Some(&node_id) = ready.iter().next() {
271            ready.remove(&node_id);
272            result.push(node_id);
273
274            for &succ in &successors[node_id] {
275                in_degree[succ] -= 1;
276                if in_degree[succ] == 0 {
277                    ready.insert(succ);
278                }
279            }
280        }
281
282        if result.len() != n {
283            return Err(DataflowError::CycleDetected {
284                visited: result.len(),
285                total: n,
286            });
287        }
288
289        Ok(result)
290    }
291
292    /// Compute the critical path through the DFG.
293    ///
294    /// Assigns a unit cost of 1 to each node and computes the longest path
295    /// from any source node (no predecessors) to any sink node (no successors).
296    ///
297    /// Returns the cost for each node (longest path from any source to that node)
298    /// and the critical path length.
299    pub fn critical_path(&self) -> CriticalPathResult {
300        self.critical_path_with_costs(&vec![1usize; self.nodes.len()])
301    }
302
303    /// Compute the critical path with custom per-node costs.
304    ///
305    /// `costs[i]` is the execution cost of node `i`.
306    ///
307    /// Returns the cost for each node and the critical path length.
308    pub fn critical_path_with_costs(&self, costs: &[usize]) -> CriticalPathResult {
309        let n = self.nodes.len();
310        if n == 0 {
311            return CriticalPathResult {
312                node_distances: Vec::new(),
313                critical_path_length: 0,
314                critical_path: Vec::new(),
315            };
316        }
317
318        // Compute topological order (handle cycles gracefully).
319        let topo = match self.topological_sort() {
320            Ok(order) => order,
321            Err(_) => {
322                return CriticalPathResult {
323                    node_distances: vec![0; n],
324                    critical_path_length: 0,
325                    critical_path: Vec::new(),
326                };
327            }
328        };
329
330        // Build predecessor lists.
331        let mut preds: Vec<Vec<usize>> = vec![Vec::new(); n];
332        for edge in &self.edges {
333            preds[edge.to].push(edge.from);
334        }
335
336        // Compute longest distance from any source to each node.
337        let mut dist = vec![0usize; n];
338        // Track the predecessor on the critical path.
339        let mut prev: Vec<Option<usize>> = vec![None; n];
340
341        for &node_id in &topo {
342            let node_cost = costs.get(node_id).copied().unwrap_or(1);
343            let mut best_dist = 0;
344            let mut best_pred = None;
345
346            for &pred_id in &preds[node_id] {
347                if dist[pred_id] > best_dist {
348                    best_dist = dist[pred_id];
349                    best_pred = Some(pred_id);
350                }
351            }
352
353            dist[node_id] = best_dist + node_cost;
354            prev[node_id] = best_pred;
355        }
356
357        // Find the sink with the maximum distance (the critical path endpoint).
358        let mut max_dist = 0;
359        let mut max_node = 0;
360        for (i, &d) in dist.iter().enumerate() {
361            if d > max_dist {
362                max_dist = d;
363                max_node = i;
364            }
365        }
366
367        // Reconstruct the critical path by following prev pointers.
368        let mut path = VecDeque::new();
369        let mut current = Some(max_node);
370        while let Some(node_id) = current {
371            path.push_front(node_id);
372            current = prev[node_id];
373        }
374
375        CriticalPathResult {
376            node_distances: dist,
377            critical_path_length: max_dist,
378            critical_path: path.into(),
379        }
380    }
381
382    /// Identify groups of independent nodes that can execute concurrently.
383    ///
384    /// Returns a list of groups, where each group contains node IDs that have
385    /// no dependencies between them and can execute in parallel.
386    pub fn parallel_groups(&self) -> Vec<Vec<usize>> {
387        let n = self.nodes.len();
388        if n == 0 {
389            return Vec::new();
390        }
391
392        let topo = match self.topological_sort() {
393            Ok(order) => order,
394            Err(_) => return Vec::new(),
395        };
396
397        // Build predecessor/successor sets for dependency checks.
398        let mut has_pred: Vec<HashSet<usize>> = vec![HashSet::new(); n];
399        for edge in &self.edges {
400            has_pred[edge.to].insert(edge.from);
401        }
402
403        // Compute ASAP (As Soon As Possible) time for each node.
404        let mut asap = vec![0usize; n];
405        for &node_id in &topo {
406            let mut max_pred_time = 0;
407            for &pred_id in &has_pred[node_id] {
408                let pred_finish = asap[pred_id] + 1;
409                if pred_finish > max_pred_time {
410                    max_pred_time = pred_finish;
411                }
412            }
413            asap[node_id] = max_pred_time;
414        }
415
416        // Group nodes by their ASAP time.
417        let max_time = asap.iter().copied().max().unwrap_or(0);
418        let mut groups: Vec<Vec<usize>> = vec![Vec::new(); max_time + 1];
419        for &node_id in &topo {
420            groups[asap[node_id]].push(node_id);
421        }
422
423        // Filter out empty groups.
424        groups.retain(|g| !g.is_empty());
425        groups
426    }
427}
428
429/// Result of critical path analysis.
430#[derive(Clone, Debug)]
431pub struct CriticalPathResult {
432    /// The longest-path distance from any source to each node.
433    pub node_distances: Vec<usize>,
434    /// The length of the critical path (maximum distance).
435    pub critical_path_length: usize,
436    /// Node IDs on the critical path, from source to sink.
437    pub critical_path: Vec<usize>,
438}
439
440/// Classify a statement into its DFG node kind and compute read/write sets.
441fn classify_statement(
442    stmt: &Statement,
443    expressions: &nxpu_ir::Arena<Expression>,
444) -> (
445    DfgNodeKind,
446    Vec<Handle<Expression>>,
447    Vec<Handle<Expression>>,
448) {
449    match stmt {
450        Statement::Store { pointer, value } => {
451            let mut reads = Vec::new();
452            reads.push(*value);
453            collect_expr_reads(*value, expressions, &mut reads);
454            collect_expr_reads(*pointer, expressions, &mut reads);
455            let writes = vec![*pointer];
456            (DfgNodeKind::Store, reads, writes)
457        }
458        Statement::Call {
459            arguments, result, ..
460        } => {
461            let mut reads: Vec<Handle<Expression>> = arguments.clone();
462            for &arg in arguments {
463                collect_expr_reads(arg, expressions, &mut reads);
464            }
465            let writes = result.iter().copied().collect();
466            (DfgNodeKind::Call, reads, writes)
467        }
468        Statement::Atomic {
469            pointer,
470            value,
471            result,
472            fun,
473        } => {
474            let mut reads = vec![*pointer, *value];
475            collect_expr_reads(*pointer, expressions, &mut reads);
476            collect_expr_reads(*value, expressions, &mut reads);
477            if let nxpu_ir::AtomicFunction::Exchange {
478                compare: Some(cmp), ..
479            } = fun
480            {
481                reads.push(*cmp);
482                collect_expr_reads(*cmp, expressions, &mut reads);
483            }
484            let mut writes = vec![*pointer];
485            if let Some(r) = result {
486                writes.push(*r);
487            }
488            (DfgNodeKind::Atomic, reads, writes)
489        }
490        Statement::Barrier(_) => (DfgNodeKind::Barrier, Vec::new(), Vec::new()),
491        Statement::Emit(range) => {
492            let mut reads = Vec::new();
493            let mut writes = Vec::new();
494            let idx_range = range.index_range();
495            // Iterate expression arena and pick handles within the emit range.
496            for (handle, expr) in expressions.iter() {
497                let idx = handle.index() as u32;
498                if idx >= idx_range.start && idx < idx_range.end {
499                    writes.push(handle);
500                    for operand in expression_operands(expr) {
501                        if !reads.contains(&operand) {
502                            reads.push(operand);
503                        }
504                    }
505                }
506            }
507            (DfgNodeKind::Emit, reads, writes)
508        }
509        Statement::If {
510            condition,
511            accept,
512            reject,
513        } => {
514            let mut reads = vec![*condition];
515            collect_expr_reads(*condition, expressions, &mut reads);
516            // Conservatively include reads/writes from both branches.
517            for child in accept.iter().chain(reject.iter()) {
518                let (_, child_reads, _) = classify_statement(child, expressions);
519                for r in child_reads {
520                    if !reads.contains(&r) {
521                        reads.push(r);
522                    }
523                }
524            }
525            let mut writes = Vec::new();
526            for child in accept.iter().chain(reject.iter()) {
527                let (_, _, child_writes) = classify_statement(child, expressions);
528                for w in child_writes {
529                    if !writes.contains(&w) {
530                        writes.push(w);
531                    }
532                }
533            }
534            (DfgNodeKind::If, reads, writes)
535        }
536        Statement::Loop {
537            body,
538            continuing,
539            break_if,
540        } => {
541            let mut reads = Vec::new();
542            if let Some(brk) = break_if {
543                reads.push(*brk);
544                collect_expr_reads(*brk, expressions, &mut reads);
545            }
546            for child in body.iter().chain(continuing.iter()) {
547                let (_, child_reads, _) = classify_statement(child, expressions);
548                for r in child_reads {
549                    if !reads.contains(&r) {
550                        reads.push(r);
551                    }
552                }
553            }
554            let mut writes = Vec::new();
555            for child in body.iter().chain(continuing.iter()) {
556                let (_, _, child_writes) = classify_statement(child, expressions);
557                for w in child_writes {
558                    if !writes.contains(&w) {
559                        writes.push(w);
560                    }
561                }
562            }
563            (DfgNodeKind::Loop, reads, writes)
564        }
565        Statement::Return { value } => {
566            let mut reads = Vec::new();
567            if let Some(v) = value {
568                reads.push(*v);
569                collect_expr_reads(*v, expressions, &mut reads);
570            }
571            (DfgNodeKind::Return, reads, Vec::new())
572        }
573        Statement::Break => (DfgNodeKind::Break, Vec::new(), Vec::new()),
574        Statement::Continue => (DfgNodeKind::Continue, Vec::new(), Vec::new()),
575    }
576}
577
578/// Recursively collect expression handles that are read by an expression.
579fn collect_expr_reads(
580    handle: Handle<Expression>,
581    expressions: &nxpu_ir::Arena<Expression>,
582    reads: &mut Vec<Handle<Expression>>,
583) {
584    if let Some(expr) = expressions.try_get(handle) {
585        for operand in expression_operands(expr) {
586            if !reads.contains(&operand) {
587                reads.push(operand);
588                collect_expr_reads(operand, expressions, reads);
589            }
590        }
591    }
592}
593
594/// Returns all expression handles directly referenced by an expression.
595fn expression_operands(expr: &Expression) -> Vec<Handle<Expression>> {
596    match expr {
597        Expression::Literal(_)
598        | Expression::FunctionArgument(_)
599        | Expression::GlobalVariable(_)
600        | Expression::LocalVariable(_)
601        | Expression::CallResult(_)
602        | Expression::AtomicResult { .. }
603        | Expression::ZeroValue(_) => vec![],
604
605        Expression::Load { pointer } => vec![*pointer],
606        Expression::Unary { expr, .. } => vec![*expr],
607        Expression::ArrayLength(e) => vec![*e],
608        Expression::Splat { value, .. } => vec![*value],
609        Expression::As { expr, .. } => vec![*expr],
610
611        Expression::Binary { left, right, .. } => vec![*left, *right],
612        Expression::Access { base, index } => vec![*base, *index],
613        Expression::AccessIndex { base, .. } => vec![*base],
614        Expression::Select {
615            condition,
616            accept,
617            reject,
618        } => vec![*condition, *accept, *reject],
619        Expression::Swizzle { vector, .. } => vec![*vector],
620
621        Expression::Compose { components, .. } => components.clone(),
622        Expression::Math {
623            arg,
624            arg1,
625            arg2,
626            arg3,
627            ..
628        } => {
629            let mut ops = vec![*arg];
630            if let Some(a) = arg1 {
631                ops.push(*a);
632            }
633            if let Some(a) = arg2 {
634                ops.push(*a);
635            }
636            if let Some(a) = arg3 {
637                ops.push(*a);
638            }
639            ops
640        }
641    }
642}
643
644impl fmt::Display for DataflowGraph {
645    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
646        writeln!(
647            f,
648            "DataflowGraph ({} nodes, {} edges):",
649            self.nodes.len(),
650            self.edges.len()
651        )?;
652        for node in &self.nodes {
653            writeln!(
654                f,
655                "  node[{}]: {} (reads: {}, writes: {})",
656                node.id,
657                node.kind,
658                node.reads.len(),
659                node.writes.len()
660            )?;
661        }
662        for edge in &self.edges {
663            writeln!(f, "  edge: {} -> {} ({})", edge.from, edge.to, edge.kind)?;
664        }
665        Ok(())
666    }
667}
668
669#[cfg(test)]
670mod tests {
671    use super::*;
672    use nxpu_ir::{Expression, Function, Literal, Statement};
673
674    // Helper: build a simple function with two stores to different locations
675    // that share a common value (creating a RAW dependency).
676    fn make_store_chain() -> Function {
677        let mut func = Function::new("test");
678
679        // Expressions: gv0_ptr, gv1_ptr, literal, load(gv0_ptr), binary(load, lit)
680        let gv0 = {
681            let mut arena = nxpu_ir::Arena::new();
682            arena.append(nxpu_ir::GlobalVariable {
683                name: Some("a".into()),
684                space: nxpu_ir::AddressSpace::Storage {
685                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
686                },
687                binding: None,
688                ty: {
689                    let mut types = nxpu_ir::UniqueArena::new();
690                    types.insert(nxpu_ir::Type {
691                        name: None,
692                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
693                    })
694                },
695                init: None,
696                layout: None,
697            })
698        };
699        let gv1 = {
700            let mut arena = nxpu_ir::Arena::new();
701            let _ = arena.append(nxpu_ir::GlobalVariable {
702                name: None,
703                space: nxpu_ir::AddressSpace::Storage {
704                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
705                },
706                binding: None,
707                ty: {
708                    let mut types = nxpu_ir::UniqueArena::new();
709                    types.insert(nxpu_ir::Type {
710                        name: None,
711                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
712                    })
713                },
714                init: None,
715                layout: None,
716            });
717            arena.append(nxpu_ir::GlobalVariable {
718                name: Some("b".into()),
719                space: nxpu_ir::AddressSpace::Storage {
720                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
721                },
722                binding: None,
723                ty: {
724                    let mut types = nxpu_ir::UniqueArena::new();
725                    types.insert(nxpu_ir::Type {
726                        name: None,
727                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
728                    })
729                },
730                init: None,
731                layout: None,
732            })
733        };
734
735        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
736        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
737        let val = func
738            .expressions
739            .append(Expression::Literal(Literal::F32(42.0)));
740
741        // Store val to ptr_a
742        func.body.push(Statement::Store {
743            pointer: ptr_a,
744            value: val,
745        });
746
747        // Store val to ptr_b (no dependency on first store since different pointers)
748        func.body.push(Statement::Store {
749            pointer: ptr_b,
750            value: val,
751        });
752
753        func
754    }
755
756    #[test]
757    fn build_dfg_empty_function() {
758        let func = Function::new("empty");
759        let dfg = DataflowGraph::build(&func);
760        assert_eq!(dfg.node_count(), 0);
761        assert_eq!(dfg.edge_count(), 0);
762    }
763
764    #[test]
765    fn build_dfg_simple_stores() {
766        let func = make_store_chain();
767        let dfg = DataflowGraph::build(&func);
768        assert_eq!(dfg.node_count(), 2);
769        // Both stores read `val`, but write to different pointers.
770        // The shared read of `val` means no WAR/WAW between them
771        // unless the expressions overlap. Since val is read by both,
772        // and neither writes val, there is no dependency beyond shared reads.
773        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Store);
774        assert_eq!(dfg.nodes()[1].kind, DfgNodeKind::Store);
775    }
776
777    #[test]
778    fn build_dfg_raw_dependency() {
779        // Create: emit expr -> store -> load -> store
780        // This creates a RAW dependency.
781        let mut func = Function::new("test");
782
783        let gv = {
784            let mut arena = nxpu_ir::Arena::new();
785            arena.append(nxpu_ir::GlobalVariable {
786                name: Some("x".into()),
787                space: nxpu_ir::AddressSpace::Storage {
788                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
789                },
790                binding: None,
791                ty: {
792                    let mut types = nxpu_ir::UniqueArena::new();
793                    types.insert(nxpu_ir::Type {
794                        name: None,
795                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
796                    })
797                },
798                init: None,
799                layout: None,
800            })
801        };
802
803        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
804        let val = func
805            .expressions
806            .append(Expression::Literal(Literal::F32(1.0)));
807
808        // Statement 0: store val -> ptr
809        func.body.push(Statement::Store {
810            pointer: ptr,
811            value: val,
812        });
813        // Statement 1: store val -> ptr (WAW on ptr)
814        func.body.push(Statement::Store {
815            pointer: ptr,
816            value: val,
817        });
818
819        let dfg = DataflowGraph::build(&func);
820        assert_eq!(dfg.node_count(), 2);
821
822        // Should have WAW edge (both write to ptr).
823        let waw_edges: Vec<_> = dfg
824            .edges()
825            .iter()
826            .filter(|e| e.kind == DependencyKind::OutputDependency)
827            .collect();
828        assert!(
829            !waw_edges.is_empty(),
830            "expected WAW dependency for two stores to same pointer"
831        );
832    }
833
834    #[test]
835    fn build_dfg_war_dependency() {
836        // Create a read-then-write pattern.
837        let mut func = Function::new("test");
838
839        let gv = {
840            let mut arena = nxpu_ir::Arena::new();
841            arena.append(nxpu_ir::GlobalVariable {
842                name: Some("x".into()),
843                space: nxpu_ir::AddressSpace::Storage {
844                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
845                },
846                binding: None,
847                ty: {
848                    let mut types = nxpu_ir::UniqueArena::new();
849                    types.insert(nxpu_ir::Type {
850                        name: None,
851                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
852                    })
853                },
854                init: None,
855                layout: None,
856            })
857        };
858
859        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
860        let val = func
861            .expressions
862            .append(Expression::Literal(Literal::F32(1.0)));
863
864        // Statement 0: return (reads ptr via load expression contained within)
865        func.body.push(Statement::Return { value: Some(ptr) });
866        // Statement 1: store to ptr (writes ptr)
867        func.body.push(Statement::Store {
868            pointer: ptr,
869            value: val,
870        });
871
872        let dfg = DataflowGraph::build(&func);
873        assert_eq!(dfg.node_count(), 2);
874
875        // Should have WAR edge (node 0 reads ptr, node 1 writes ptr).
876        let war_edges: Vec<_> = dfg
877            .edges()
878            .iter()
879            .filter(|e| e.kind == DependencyKind::AntiDependency)
880            .collect();
881        assert!(!war_edges.is_empty(), "expected WAR dependency");
882    }
883
884    #[test]
885    fn topological_sort_linear_chain() {
886        // Three nodes: 0 -> 1 -> 2
887        let mut func = Function::new("test");
888
889        let gv = {
890            let mut arena = nxpu_ir::Arena::new();
891            arena.append(nxpu_ir::GlobalVariable {
892                name: Some("x".into()),
893                space: nxpu_ir::AddressSpace::Storage {
894                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
895                },
896                binding: None,
897                ty: {
898                    let mut types = nxpu_ir::UniqueArena::new();
899                    types.insert(nxpu_ir::Type {
900                        name: None,
901                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
902                    })
903                },
904                init: None,
905                layout: None,
906            })
907        };
908
909        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
910        let val = func
911            .expressions
912            .append(Expression::Literal(Literal::F32(1.0)));
913
914        // Three stores to the same pointer (creating WAW chain).
915        func.body.push(Statement::Store {
916            pointer: ptr,
917            value: val,
918        });
919        func.body.push(Statement::Store {
920            pointer: ptr,
921            value: val,
922        });
923        func.body.push(Statement::Store {
924            pointer: ptr,
925            value: val,
926        });
927
928        let dfg = DataflowGraph::build(&func);
929        let order = dfg.topological_sort().unwrap();
930        assert_eq!(order.len(), 3);
931        // Must be in order due to WAW dependencies.
932        assert_eq!(order[0], 0);
933        assert_eq!(order[1], 1);
934        assert_eq!(order[2], 2);
935    }
936
937    #[test]
938    fn topological_sort_empty() {
939        let func = Function::new("empty");
940        let dfg = DataflowGraph::build(&func);
941        let order = dfg.topological_sort().unwrap();
942        assert_eq!(order.len(), 0);
943    }
944
945    #[test]
946    fn topological_sort_independent_nodes() {
947        let mut func = Function::new("test");
948
949        // Two independent stores to different pointers with different values.
950        let gv0 = {
951            let mut arena = nxpu_ir::Arena::new();
952            arena.append(nxpu_ir::GlobalVariable {
953                name: Some("a".into()),
954                space: nxpu_ir::AddressSpace::Storage {
955                    access: nxpu_ir::StorageAccess::STORE,
956                },
957                binding: None,
958                ty: {
959                    let mut types = nxpu_ir::UniqueArena::new();
960                    types.insert(nxpu_ir::Type {
961                        name: None,
962                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
963                    })
964                },
965                init: None,
966                layout: None,
967            })
968        };
969        let gv1 = {
970            let mut arena = nxpu_ir::Arena::new();
971            let _ = arena.append(nxpu_ir::GlobalVariable {
972                name: None,
973                space: nxpu_ir::AddressSpace::Storage {
974                    access: nxpu_ir::StorageAccess::STORE,
975                },
976                binding: None,
977                ty: {
978                    let mut types = nxpu_ir::UniqueArena::new();
979                    types.insert(nxpu_ir::Type {
980                        name: None,
981                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
982                    })
983                },
984                init: None,
985                layout: None,
986            });
987            arena.append(nxpu_ir::GlobalVariable {
988                name: Some("b".into()),
989                space: nxpu_ir::AddressSpace::Storage {
990                    access: nxpu_ir::StorageAccess::STORE,
991                },
992                binding: None,
993                ty: {
994                    let mut types = nxpu_ir::UniqueArena::new();
995                    types.insert(nxpu_ir::Type {
996                        name: None,
997                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
998                    })
999                },
1000                init: None,
1001                layout: None,
1002            })
1003        };
1004
1005        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1006        let val_a = func
1007            .expressions
1008            .append(Expression::Literal(Literal::F32(1.0)));
1009        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1010        let val_b = func
1011            .expressions
1012            .append(Expression::Literal(Literal::F32(2.0)));
1013
1014        func.body.push(Statement::Store {
1015            pointer: ptr_a,
1016            value: val_a,
1017        });
1018        func.body.push(Statement::Store {
1019            pointer: ptr_b,
1020            value: val_b,
1021        });
1022
1023        let dfg = DataflowGraph::build(&func);
1024        let order = dfg.topological_sort().unwrap();
1025        assert_eq!(order.len(), 2);
1026        // Both are valid orderings since they are independent.
1027    }
1028
1029    #[test]
1030    fn critical_path_linear_chain() {
1031        let mut func = Function::new("test");
1032
1033        let gv = {
1034            let mut arena = nxpu_ir::Arena::new();
1035            arena.append(nxpu_ir::GlobalVariable {
1036                name: Some("x".into()),
1037                space: nxpu_ir::AddressSpace::Storage {
1038                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1039                },
1040                binding: None,
1041                ty: {
1042                    let mut types = nxpu_ir::UniqueArena::new();
1043                    types.insert(nxpu_ir::Type {
1044                        name: None,
1045                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1046                    })
1047                },
1048                init: None,
1049                layout: None,
1050            })
1051        };
1052
1053        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1054        let val = func
1055            .expressions
1056            .append(Expression::Literal(Literal::F32(1.0)));
1057
1058        // Three stores to the same pointer.
1059        func.body.push(Statement::Store {
1060            pointer: ptr,
1061            value: val,
1062        });
1063        func.body.push(Statement::Store {
1064            pointer: ptr,
1065            value: val,
1066        });
1067        func.body.push(Statement::Store {
1068            pointer: ptr,
1069            value: val,
1070        });
1071
1072        let dfg = DataflowGraph::build(&func);
1073        let cp = dfg.critical_path();
1074        // With unit cost and 3 chained nodes, the critical path length is 3.
1075        assert_eq!(cp.critical_path_length, 3);
1076        assert_eq!(cp.critical_path.len(), 3);
1077    }
1078
1079    #[test]
1080    fn critical_path_empty() {
1081        let func = Function::new("empty");
1082        let dfg = DataflowGraph::build(&func);
1083        let cp = dfg.critical_path();
1084        assert_eq!(cp.critical_path_length, 0);
1085        assert_eq!(cp.critical_path.len(), 0);
1086    }
1087
1088    #[test]
1089    fn barrier_creates_control_edges() {
1090        let mut func = Function::new("test");
1091
1092        let gv = {
1093            let mut arena = nxpu_ir::Arena::new();
1094            arena.append(nxpu_ir::GlobalVariable {
1095                name: Some("x".into()),
1096                space: nxpu_ir::AddressSpace::Storage {
1097                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1098                },
1099                binding: None,
1100                ty: {
1101                    let mut types = nxpu_ir::UniqueArena::new();
1102                    types.insert(nxpu_ir::Type {
1103                        name: None,
1104                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1105                    })
1106                },
1107                init: None,
1108                layout: None,
1109            })
1110        };
1111
1112        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1113        let val = func
1114            .expressions
1115            .append(Expression::Literal(Literal::F32(1.0)));
1116
1117        // Statement 0: store
1118        func.body.push(Statement::Store {
1119            pointer: ptr,
1120            value: val,
1121        });
1122        // Statement 1: barrier
1123        func.body
1124            .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1125        // Statement 2: store
1126        func.body.push(Statement::Store {
1127            pointer: ptr,
1128            value: val,
1129        });
1130
1131        let dfg = DataflowGraph::build(&func);
1132        assert_eq!(dfg.node_count(), 3);
1133
1134        // Barrier should have control edge to subsequent memory op.
1135        let control_edges: Vec<_> = dfg
1136            .edges()
1137            .iter()
1138            .filter(|e| e.kind == DependencyKind::Control)
1139            .collect();
1140        assert!(
1141            !control_edges.is_empty(),
1142            "expected control dependency from barrier"
1143        );
1144        // Barrier (node 1) -> store (node 2).
1145        assert!(control_edges.iter().any(|e| e.from == 1 && e.to == 2));
1146    }
1147
1148    #[test]
1149    fn parallel_groups_independent_ops() {
1150        let mut func = Function::new("test");
1151
1152        let gv0 = {
1153            let mut arena = nxpu_ir::Arena::new();
1154            arena.append(nxpu_ir::GlobalVariable {
1155                name: Some("a".into()),
1156                space: nxpu_ir::AddressSpace::Storage {
1157                    access: nxpu_ir::StorageAccess::STORE,
1158                },
1159                binding: None,
1160                ty: {
1161                    let mut types = nxpu_ir::UniqueArena::new();
1162                    types.insert(nxpu_ir::Type {
1163                        name: None,
1164                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1165                    })
1166                },
1167                init: None,
1168                layout: None,
1169            })
1170        };
1171        let gv1 = {
1172            let mut arena = nxpu_ir::Arena::new();
1173            let _ = arena.append(nxpu_ir::GlobalVariable {
1174                name: None,
1175                space: nxpu_ir::AddressSpace::Storage {
1176                    access: nxpu_ir::StorageAccess::STORE,
1177                },
1178                binding: None,
1179                ty: {
1180                    let mut types = nxpu_ir::UniqueArena::new();
1181                    types.insert(nxpu_ir::Type {
1182                        name: None,
1183                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1184                    })
1185                },
1186                init: None,
1187                layout: None,
1188            });
1189            arena.append(nxpu_ir::GlobalVariable {
1190                name: Some("b".into()),
1191                space: nxpu_ir::AddressSpace::Storage {
1192                    access: nxpu_ir::StorageAccess::STORE,
1193                },
1194                binding: None,
1195                ty: {
1196                    let mut types = nxpu_ir::UniqueArena::new();
1197                    types.insert(nxpu_ir::Type {
1198                        name: None,
1199                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1200                    })
1201                },
1202                init: None,
1203                layout: None,
1204            })
1205        };
1206
1207        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1208        let val_a = func
1209            .expressions
1210            .append(Expression::Literal(Literal::F32(1.0)));
1211        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1212        let val_b = func
1213            .expressions
1214            .append(Expression::Literal(Literal::F32(2.0)));
1215
1216        func.body.push(Statement::Store {
1217            pointer: ptr_a,
1218            value: val_a,
1219        });
1220        func.body.push(Statement::Store {
1221            pointer: ptr_b,
1222            value: val_b,
1223        });
1224
1225        let dfg = DataflowGraph::build(&func);
1226        let groups = dfg.parallel_groups();
1227
1228        // If they are independent, they should be in the same group.
1229        // Check that at least one group has 2 elements (parallel).
1230        let max_group_size = groups.iter().map(|g| g.len()).max().unwrap_or(0);
1231        assert!(
1232            max_group_size >= 2,
1233            "expected independent ops to be grouped together, got groups: {groups:?}"
1234        );
1235    }
1236
1237    #[test]
1238    fn display_dfg() {
1239        let func = Function::new("empty");
1240        let dfg = DataflowGraph::build(&func);
1241        let s = format!("{dfg}");
1242        assert!(s.contains("DataflowGraph"));
1243    }
1244
1245    #[test]
1246    fn critical_path_with_custom_costs() {
1247        let mut func = Function::new("test");
1248
1249        let gv = {
1250            let mut arena = nxpu_ir::Arena::new();
1251            arena.append(nxpu_ir::GlobalVariable {
1252                name: Some("x".into()),
1253                space: nxpu_ir::AddressSpace::Storage {
1254                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1255                },
1256                binding: None,
1257                ty: {
1258                    let mut types = nxpu_ir::UniqueArena::new();
1259                    types.insert(nxpu_ir::Type {
1260                        name: None,
1261                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1262                    })
1263                },
1264                init: None,
1265                layout: None,
1266            })
1267        };
1268
1269        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1270        let val = func
1271            .expressions
1272            .append(Expression::Literal(Literal::F32(1.0)));
1273
1274        func.body.push(Statement::Store {
1275            pointer: ptr,
1276            value: val,
1277        });
1278        func.body.push(Statement::Store {
1279            pointer: ptr,
1280            value: val,
1281        });
1282
1283        let dfg = DataflowGraph::build(&func);
1284        let cp = dfg.critical_path_with_costs(&[5, 3]);
1285        // With costs [5, 3] and chain 0->1, critical path = 5 + 3 = 8.
1286        assert_eq!(cp.critical_path_length, 8);
1287    }
1288
1289    // --- Helper to make a dummy global variable handle ---
1290    fn make_gv(index: usize) -> Handle<nxpu_ir::GlobalVariable> {
1291        let mut arena = nxpu_ir::Arena::new();
1292        // Append `index + 1` dummy variables so we get the handle at the desired index.
1293        let mut handle = None;
1294        for i in 0..=index {
1295            let h = arena.append(nxpu_ir::GlobalVariable {
1296                name: Some(format!("gv_{i}")),
1297                space: nxpu_ir::AddressSpace::Storage {
1298                    access: nxpu_ir::StorageAccess::LOAD | nxpu_ir::StorageAccess::STORE,
1299                },
1300                binding: None,
1301                ty: {
1302                    let mut types = nxpu_ir::UniqueArena::new();
1303                    types.insert(nxpu_ir::Type {
1304                        name: None,
1305                        inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::F32),
1306                    })
1307                },
1308                init: None,
1309                layout: None,
1310            });
1311            handle = Some(h);
1312        }
1313        handle.unwrap()
1314    }
1315
1316    // ===== Display tests =====
1317
1318    #[test]
1319    fn display_dfg_with_nodes_and_edges() {
1320        let mut func = Function::new("test");
1321        let gv = make_gv(0);
1322        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1323        let val = func
1324            .expressions
1325            .append(Expression::Literal(Literal::F32(1.0)));
1326
1327        // Two stores to the same pointer -> WAW edge.
1328        func.body.push(Statement::Store {
1329            pointer: ptr,
1330            value: val,
1331        });
1332        func.body.push(Statement::Store {
1333            pointer: ptr,
1334            value: val,
1335        });
1336
1337        let dfg = DataflowGraph::build(&func);
1338        let s = format!("{dfg}");
1339        assert!(s.contains("DataflowGraph (2 nodes, "));
1340        assert!(s.contains("node[0]: Store"));
1341        assert!(s.contains("node[1]: Store"));
1342        assert!(s.contains("edge:"));
1343        assert!(s.contains("WAW"));
1344    }
1345
1346    #[test]
1347    fn dfg_node_kind_display_all_variants() {
1348        assert_eq!(format!("{}", DfgNodeKind::Store), "Store");
1349        assert_eq!(format!("{}", DfgNodeKind::Call), "Call");
1350        assert_eq!(format!("{}", DfgNodeKind::Atomic), "Atomic");
1351        assert_eq!(format!("{}", DfgNodeKind::Barrier), "Barrier");
1352        assert_eq!(format!("{}", DfgNodeKind::Emit), "Emit");
1353        assert_eq!(format!("{}", DfgNodeKind::If), "If");
1354        assert_eq!(format!("{}", DfgNodeKind::Loop), "Loop");
1355        assert_eq!(format!("{}", DfgNodeKind::Return), "Return");
1356        assert_eq!(format!("{}", DfgNodeKind::Break), "Break");
1357        assert_eq!(format!("{}", DfgNodeKind::Continue), "Continue");
1358    }
1359
1360    #[test]
1361    fn dependency_kind_display_all_variants() {
1362        assert_eq!(format!("{}", DependencyKind::DataFlow), "RAW");
1363        assert_eq!(format!("{}", DependencyKind::AntiDependency), "WAR");
1364        assert_eq!(format!("{}", DependencyKind::OutputDependency), "WAW");
1365        assert_eq!(format!("{}", DependencyKind::Control), "Control");
1366    }
1367
1368    #[test]
1369    fn dataflow_error_display() {
1370        let err = DataflowError::CycleDetected {
1371            visited: 3,
1372            total: 5,
1373        };
1374        let msg = format!("{err}");
1375        assert!(msg.contains("cycle detected"));
1376        assert!(msg.contains("3 of 5"));
1377    }
1378
1379    // ===== Successors / Predecessors =====
1380
1381    #[test]
1382    fn successors_and_predecessors() {
1383        let mut func = Function::new("test");
1384        let gv = make_gv(0);
1385        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1386        let val = func
1387            .expressions
1388            .append(Expression::Literal(Literal::F32(1.0)));
1389
1390        // Three stores to same pointer: 0->1->2 (WAW chain).
1391        func.body.push(Statement::Store {
1392            pointer: ptr,
1393            value: val,
1394        });
1395        func.body.push(Statement::Store {
1396            pointer: ptr,
1397            value: val,
1398        });
1399        func.body.push(Statement::Store {
1400            pointer: ptr,
1401            value: val,
1402        });
1403
1404        let dfg = DataflowGraph::build(&func);
1405
1406        // Node 0 should have successors to 1 and 2.
1407        let succs_0 = dfg.successors(0);
1408        assert_ne!(succs_0.len(), 0);
1409
1410        // Node 2 should have predecessors from 0 and/or 1.
1411        let preds_2 = dfg.predecessors(2);
1412        assert_ne!(preds_2.len(), 0);
1413
1414        // Node 0 should have no predecessors.
1415        let preds_0 = dfg.predecessors(0);
1416        assert_eq!(preds_0.len(), 0);
1417    }
1418
1419    // ===== Break / Continue statement classification =====
1420
1421    #[test]
1422    fn break_and_continue_nodes() {
1423        let mut func = Function::new("test");
1424        func.body.push(Statement::Break);
1425        func.body.push(Statement::Continue);
1426
1427        let dfg = DataflowGraph::build(&func);
1428        assert_eq!(dfg.node_count(), 2);
1429        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Break);
1430        assert_eq!(dfg.nodes()[1].kind, DfgNodeKind::Continue);
1431        // Break and Continue have no reads/writes so no edges.
1432        assert_eq!(dfg.edge_count(), 0);
1433    }
1434
1435    // ===== Return with no value =====
1436
1437    #[test]
1438    fn return_no_value() {
1439        let mut func = Function::new("test");
1440        func.body.push(Statement::Return { value: None });
1441
1442        let dfg = DataflowGraph::build(&func);
1443        assert_eq!(dfg.node_count(), 1);
1444        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Return);
1445        assert_eq!(dfg.nodes()[0].reads.len(), 0);
1446        assert_eq!(dfg.nodes()[0].writes.len(), 0);
1447    }
1448
1449    // ===== Return with value =====
1450
1451    #[test]
1452    fn return_with_value() {
1453        let mut func = Function::new("test");
1454        let val = func
1455            .expressions
1456            .append(Expression::Literal(Literal::F32(42.0)));
1457        func.body.push(Statement::Return { value: Some(val) });
1458
1459        let dfg = DataflowGraph::build(&func);
1460        assert_eq!(dfg.node_count(), 1);
1461        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Return);
1462        assert!(dfg.nodes()[0].reads.contains(&val));
1463        assert_eq!(dfg.nodes()[0].writes.len(), 0);
1464    }
1465
1466    // ===== If statement classification =====
1467
1468    #[test]
1469    fn if_statement_classification() {
1470        let mut func = Function::new("test");
1471        let gv = make_gv(0);
1472        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1473        let cond = func
1474            .expressions
1475            .append(Expression::Literal(Literal::Bool(true)));
1476        let val = func
1477            .expressions
1478            .append(Expression::Literal(Literal::F32(1.0)));
1479
1480        // If statement with a store in the accept branch.
1481        func.body.push(Statement::If {
1482            condition: cond,
1483            accept: vec![Statement::Store {
1484                pointer: ptr,
1485                value: val,
1486            }],
1487            reject: vec![],
1488        });
1489
1490        let dfg = DataflowGraph::build(&func);
1491        assert_eq!(dfg.node_count(), 1);
1492        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::If);
1493        // The If node should read the condition and the value from the child store.
1494        assert!(dfg.nodes()[0].reads.contains(&cond));
1495        // The If node should write the pointer from the child store.
1496        assert!(dfg.nodes()[0].writes.contains(&ptr));
1497    }
1498
1499    // ===== If statement with reject branch =====
1500
1501    #[test]
1502    fn if_statement_with_reject_branch() {
1503        let mut func = Function::new("test");
1504        let gv0 = make_gv(0);
1505        let gv1 = make_gv(1);
1506        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1507        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1508        let cond = func
1509            .expressions
1510            .append(Expression::Literal(Literal::Bool(true)));
1511        let val = func
1512            .expressions
1513            .append(Expression::Literal(Literal::F32(1.0)));
1514
1515        func.body.push(Statement::If {
1516            condition: cond,
1517            accept: vec![Statement::Store {
1518                pointer: ptr_a,
1519                value: val,
1520            }],
1521            reject: vec![Statement::Store {
1522                pointer: ptr_b,
1523                value: val,
1524            }],
1525        });
1526
1527        let dfg = DataflowGraph::build(&func);
1528        assert_eq!(dfg.node_count(), 1);
1529        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::If);
1530        // Should include writes from both branches.
1531        assert!(dfg.nodes()[0].writes.contains(&ptr_a));
1532        assert!(dfg.nodes()[0].writes.contains(&ptr_b));
1533    }
1534
1535    // ===== Loop statement classification =====
1536
1537    #[test]
1538    fn loop_statement_classification() {
1539        let mut func = Function::new("test");
1540        let gv = make_gv(0);
1541        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1542        let val = func
1543            .expressions
1544            .append(Expression::Literal(Literal::F32(1.0)));
1545        let break_cond = func
1546            .expressions
1547            .append(Expression::Literal(Literal::Bool(false)));
1548
1549        func.body.push(Statement::Loop {
1550            body: vec![Statement::Store {
1551                pointer: ptr,
1552                value: val,
1553            }],
1554            continuing: vec![],
1555            break_if: Some(break_cond),
1556        });
1557
1558        let dfg = DataflowGraph::build(&func);
1559        assert_eq!(dfg.node_count(), 1);
1560        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Loop);
1561        // Should read break_if condition.
1562        assert!(dfg.nodes()[0].reads.contains(&break_cond));
1563        // Should write via the child store.
1564        assert!(dfg.nodes()[0].writes.contains(&ptr));
1565    }
1566
1567    // ===== Loop with continuing block =====
1568
1569    #[test]
1570    fn loop_with_continuing_block() {
1571        let mut func = Function::new("test");
1572        let gv0 = make_gv(0);
1573        let gv1 = make_gv(1);
1574        let ptr_body = func.expressions.append(Expression::GlobalVariable(gv0));
1575        let ptr_cont = func.expressions.append(Expression::GlobalVariable(gv1));
1576        let val = func
1577            .expressions
1578            .append(Expression::Literal(Literal::F32(1.0)));
1579
1580        func.body.push(Statement::Loop {
1581            body: vec![Statement::Store {
1582                pointer: ptr_body,
1583                value: val,
1584            }],
1585            continuing: vec![Statement::Store {
1586                pointer: ptr_cont,
1587                value: val,
1588            }],
1589            break_if: None,
1590        });
1591
1592        let dfg = DataflowGraph::build(&func);
1593        assert_eq!(dfg.node_count(), 1);
1594        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Loop);
1595        // Should include writes from both body and continuing blocks.
1596        assert!(dfg.nodes()[0].writes.contains(&ptr_body));
1597        assert!(dfg.nodes()[0].writes.contains(&ptr_cont));
1598    }
1599
1600    // ===== Call statement classification =====
1601
1602    #[test]
1603    fn call_statement_classification() {
1604        let mut func = Function::new("test");
1605        let arg_val = func
1606            .expressions
1607            .append(Expression::Literal(Literal::F32(1.0)));
1608        let result_val = func
1609            .expressions
1610            .append(Expression::Literal(Literal::F32(0.0)));
1611
1612        // Build a dummy function handle.
1613        let callee_handle = {
1614            let mut funcs = nxpu_ir::Arena::new();
1615            funcs.append(Function::new("callee"))
1616        };
1617
1618        func.body.push(Statement::Call {
1619            function: callee_handle,
1620            arguments: vec![arg_val],
1621            result: Some(result_val),
1622        });
1623
1624        let dfg = DataflowGraph::build(&func);
1625        assert_eq!(dfg.node_count(), 1);
1626        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Call);
1627        assert!(dfg.nodes()[0].reads.contains(&arg_val));
1628        assert!(dfg.nodes()[0].writes.contains(&result_val));
1629    }
1630
1631    // ===== Call with no result =====
1632
1633    #[test]
1634    fn call_no_result() {
1635        let mut func = Function::new("test");
1636        let arg_val = func
1637            .expressions
1638            .append(Expression::Literal(Literal::F32(1.0)));
1639
1640        let callee_handle = {
1641            let mut funcs = nxpu_ir::Arena::new();
1642            funcs.append(Function::new("callee"))
1643        };
1644
1645        func.body.push(Statement::Call {
1646            function: callee_handle,
1647            arguments: vec![arg_val],
1648            result: None,
1649        });
1650
1651        let dfg = DataflowGraph::build(&func);
1652        assert_eq!(dfg.node_count(), 1);
1653        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Call);
1654        assert_eq!(dfg.nodes()[0].writes.len(), 0);
1655    }
1656
1657    // ===== Atomic statement classification =====
1658
1659    #[test]
1660    fn atomic_statement_classification() {
1661        let mut func = Function::new("test");
1662        let gv = make_gv(0);
1663        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1664        let val = func
1665            .expressions
1666            .append(Expression::Literal(Literal::U32(1)));
1667        let result_expr = func
1668            .expressions
1669            .append(Expression::Literal(Literal::U32(0)));
1670
1671        func.body.push(Statement::Atomic {
1672            pointer: ptr,
1673            fun: nxpu_ir::AtomicFunction::Add,
1674            value: val,
1675            result: Some(result_expr),
1676        });
1677
1678        let dfg = DataflowGraph::build(&func);
1679        assert_eq!(dfg.node_count(), 1);
1680        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Atomic);
1681        assert!(dfg.nodes()[0].reads.contains(&ptr));
1682        assert!(dfg.nodes()[0].reads.contains(&val));
1683        assert!(dfg.nodes()[0].writes.contains(&ptr));
1684        assert!(dfg.nodes()[0].writes.contains(&result_expr));
1685    }
1686
1687    // ===== Atomic with Exchange + compare =====
1688
1689    #[test]
1690    fn atomic_exchange_with_compare() {
1691        let mut func = Function::new("test");
1692        let gv = make_gv(0);
1693        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1694        let val = func
1695            .expressions
1696            .append(Expression::Literal(Literal::U32(1)));
1697        let cmp = func
1698            .expressions
1699            .append(Expression::Literal(Literal::U32(0)));
1700
1701        func.body.push(Statement::Atomic {
1702            pointer: ptr,
1703            fun: nxpu_ir::AtomicFunction::Exchange { compare: Some(cmp) },
1704            value: val,
1705            result: None,
1706        });
1707
1708        let dfg = DataflowGraph::build(&func);
1709        assert_eq!(dfg.node_count(), 1);
1710        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Atomic);
1711        // Compare should appear in reads.
1712        assert!(dfg.nodes()[0].reads.contains(&cmp));
1713    }
1714
1715    // ===== Atomic with no result =====
1716
1717    #[test]
1718    fn atomic_no_result() {
1719        let mut func = Function::new("test");
1720        let gv = make_gv(0);
1721        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1722        let val = func
1723            .expressions
1724            .append(Expression::Literal(Literal::U32(1)));
1725
1726        func.body.push(Statement::Atomic {
1727            pointer: ptr,
1728            fun: nxpu_ir::AtomicFunction::Add,
1729            value: val,
1730            result: None,
1731        });
1732
1733        let dfg = DataflowGraph::build(&func);
1734        assert_eq!(dfg.nodes()[0].writes.len(), 1);
1735        assert!(dfg.nodes()[0].writes.contains(&ptr));
1736    }
1737
1738    // ===== Emit statement classification =====
1739
1740    #[test]
1741    fn emit_statement_classification() {
1742        let mut func = Function::new("test");
1743
1744        // Create some expressions to emit.
1745        let lit_a = func
1746            .expressions
1747            .append(Expression::Literal(Literal::F32(1.0)));
1748        let lit_b = func
1749            .expressions
1750            .append(Expression::Literal(Literal::F32(2.0)));
1751        let add = func.expressions.append(Expression::Binary {
1752            op: nxpu_ir::BinaryOp::Add,
1753            left: lit_a,
1754            right: lit_b,
1755        });
1756
1757        // Emit range covering all three expressions.
1758        let range = nxpu_ir::Range::from_index_range(0..3);
1759        func.body.push(Statement::Emit(range));
1760
1761        let dfg = DataflowGraph::build(&func);
1762        assert_eq!(dfg.node_count(), 1);
1763        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Emit);
1764        // The emit should write all 3 expressions.
1765        assert!(dfg.nodes()[0].writes.contains(&lit_a));
1766        assert!(dfg.nodes()[0].writes.contains(&lit_b));
1767        assert!(dfg.nodes()[0].writes.contains(&add));
1768    }
1769
1770    // ===== Barrier statement classification =====
1771
1772    #[test]
1773    fn barrier_statement_no_reads_writes() {
1774        let mut func = Function::new("test");
1775        func.body
1776            .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1777
1778        let dfg = DataflowGraph::build(&func);
1779        assert_eq!(dfg.node_count(), 1);
1780        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Barrier);
1781        assert_eq!(dfg.nodes()[0].reads.len(), 0);
1782        assert_eq!(dfg.nodes()[0].writes.len(), 0);
1783    }
1784
1785    // ===== Barrier to barrier control edge =====
1786
1787    #[test]
1788    fn barrier_to_barrier_control_edge() {
1789        let mut func = Function::new("test");
1790        func.body
1791            .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1792        func.body
1793            .push(Statement::Barrier(nxpu_ir::Barrier::WORKGROUP));
1794
1795        let dfg = DataflowGraph::build(&func);
1796        assert_eq!(dfg.node_count(), 2);
1797        // First barrier should have a control edge to the second barrier.
1798        let control_edges: Vec<_> = dfg
1799            .edges()
1800            .iter()
1801            .filter(|e| e.kind == DependencyKind::Control)
1802            .collect();
1803        assert!(
1804            control_edges.iter().any(|e| e.from == 0 && e.to == 1),
1805            "expected control edge from barrier[0] to barrier[1]"
1806        );
1807    }
1808
1809    // ===== Barrier does not create control edge to non-memory op =====
1810
1811    #[test]
1812    fn barrier_no_control_edge_to_break() {
1813        let mut func = Function::new("test");
1814        func.body
1815            .push(Statement::Barrier(nxpu_ir::Barrier::STORAGE));
1816        func.body.push(Statement::Break);
1817
1818        let dfg = DataflowGraph::build(&func);
1819        assert_eq!(dfg.node_count(), 2);
1820        // Break has no reads/writes and is not a barrier, so no control edge.
1821        let control_edges: Vec<_> = dfg
1822            .edges()
1823            .iter()
1824            .filter(|e| e.kind == DependencyKind::Control)
1825            .collect();
1826        assert!(
1827            control_edges.is_empty(),
1828            "should not have control edge from barrier to break"
1829        );
1830    }
1831
1832    // ===== RAW dependency (explicit read-after-write) =====
1833
1834    #[test]
1835    fn explicit_raw_dependency() {
1836        let mut func = Function::new("test");
1837        let gv = make_gv(0);
1838        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1839        let val = func
1840            .expressions
1841            .append(Expression::Literal(Literal::F32(1.0)));
1842        let load = func.expressions.append(Expression::Load { pointer: ptr });
1843
1844        // Statement 0: store val -> ptr (writes ptr)
1845        func.body.push(Statement::Store {
1846            pointer: ptr,
1847            value: val,
1848        });
1849        // Statement 1: return load(ptr) (reads ptr via load)
1850        func.body.push(Statement::Return { value: Some(load) });
1851
1852        let dfg = DataflowGraph::build(&func);
1853        assert_eq!(dfg.node_count(), 2);
1854
1855        // Should have RAW: node 0 writes ptr, node 1 reads ptr (via load's pointer).
1856        let raw_edges: Vec<_> = dfg
1857            .edges()
1858            .iter()
1859            .filter(|e| e.kind == DependencyKind::DataFlow)
1860            .collect();
1861        assert!(
1862            !raw_edges.is_empty(),
1863            "expected RAW dependency: store then read"
1864        );
1865    }
1866
1867    // ===== Combined WAR + WAW + RAW =====
1868
1869    #[test]
1870    fn combined_war_waw_raw_dependencies() {
1871        let mut func = Function::new("test");
1872        let gv = make_gv(0);
1873        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1874        let val = func
1875            .expressions
1876            .append(Expression::Literal(Literal::F32(1.0)));
1877        let load = func.expressions.append(Expression::Load { pointer: ptr });
1878
1879        // Statement 0: store val -> ptr (writes ptr, reads val)
1880        func.body.push(Statement::Store {
1881            pointer: ptr,
1882            value: val,
1883        });
1884        // Statement 1: store load(ptr) -> ptr (reads ptr via load, writes ptr)
1885        //   RAW: node 0 writes ptr, node 1 reads ptr (via load)
1886        //   WAW: both write ptr
1887        //   WAR: node 0 reads val (no WAR here since node 1 doesn't write val)
1888        func.body.push(Statement::Store {
1889            pointer: ptr,
1890            value: load,
1891        });
1892
1893        let dfg = DataflowGraph::build(&func);
1894        assert_eq!(dfg.node_count(), 2);
1895
1896        let has_raw = dfg
1897            .edges()
1898            .iter()
1899            .any(|e| e.kind == DependencyKind::DataFlow);
1900        let has_waw = dfg
1901            .edges()
1902            .iter()
1903            .any(|e| e.kind == DependencyKind::OutputDependency);
1904        assert!(has_raw, "expected RAW dependency");
1905        assert!(has_waw, "expected WAW dependency");
1906    }
1907
1908    // ===== Parallel groups with a dependency chain =====
1909
1910    #[test]
1911    fn parallel_groups_with_chain() {
1912        let mut func = Function::new("test");
1913        let gv = make_gv(0);
1914        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
1915        let val = func
1916            .expressions
1917            .append(Expression::Literal(Literal::F32(1.0)));
1918
1919        // Three stores to same pointer: all dependent.
1920        func.body.push(Statement::Store {
1921            pointer: ptr,
1922            value: val,
1923        });
1924        func.body.push(Statement::Store {
1925            pointer: ptr,
1926            value: val,
1927        });
1928        func.body.push(Statement::Store {
1929            pointer: ptr,
1930            value: val,
1931        });
1932
1933        let dfg = DataflowGraph::build(&func);
1934        let groups = dfg.parallel_groups();
1935
1936        // All dependent: each should be in its own group.
1937        assert_eq!(groups.len(), 3, "expected 3 sequential groups");
1938        for g in &groups {
1939            assert_eq!(g.len(), 1, "each group should have exactly 1 node");
1940        }
1941    }
1942
1943    // ===== Parallel groups empty graph =====
1944
1945    #[test]
1946    fn parallel_groups_empty() {
1947        let func = Function::new("empty");
1948        let dfg = DataflowGraph::build(&func);
1949        let groups = dfg.parallel_groups();
1950        assert_eq!(groups.len(), 0);
1951    }
1952
1953    // ===== Critical path multi-path diamond graph =====
1954
1955    #[test]
1956    fn critical_path_diamond_graph() {
1957        // Build a diamond: node 0 -> node 1 (independent), node 0 -> node 2,
1958        // node 1 -> node 3, node 2 -> node 3.
1959        // We simulate this with stores that create the right dependency pattern.
1960        let mut func = Function::new("test");
1961        let gv0 = make_gv(0);
1962        let gv1 = make_gv(1);
1963
1964        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
1965        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
1966        let val = func
1967            .expressions
1968            .append(Expression::Literal(Literal::F32(1.0)));
1969        let load_a = func.expressions.append(Expression::Load { pointer: ptr_a });
1970        let load_b = func.expressions.append(Expression::Load { pointer: ptr_b });
1971
1972        // stmt 0: store val -> ptr_a (writes a)
1973        func.body.push(Statement::Store {
1974            pointer: ptr_a,
1975            value: val,
1976        });
1977        // stmt 1: store val -> ptr_b (writes b, independent of stmt 0)
1978        func.body.push(Statement::Store {
1979            pointer: ptr_b,
1980            value: val,
1981        });
1982        // stmt 2: store load_a -> ptr_a (reads a -> RAW from 0, writes a -> WAW with 0)
1983        func.body.push(Statement::Store {
1984            pointer: ptr_a,
1985            value: load_a,
1986        });
1987        // stmt 3: store load_b -> ptr_b (reads b -> RAW from 1, writes b -> WAW with 1)
1988        func.body.push(Statement::Store {
1989            pointer: ptr_b,
1990            value: load_b,
1991        });
1992
1993        let dfg = DataflowGraph::build(&func);
1994        let cp = dfg.critical_path();
1995
1996        // Two parallel chains of length 2. Critical path = 2.
1997        assert_eq!(cp.critical_path_length, 2);
1998        assert_eq!(cp.critical_path.len(), 2);
1999        assert_eq!(cp.node_distances.len(), 4);
2000    }
2001
2002    // ===== Critical path result node distances =====
2003
2004    #[test]
2005    fn critical_path_node_distances() {
2006        let mut func = Function::new("test");
2007        let gv = make_gv(0);
2008        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2009        let val = func
2010            .expressions
2011            .append(Expression::Literal(Literal::F32(1.0)));
2012
2013        // Three stores: 0 -> 1 -> 2 (WAW chain).
2014        func.body.push(Statement::Store {
2015            pointer: ptr,
2016            value: val,
2017        });
2018        func.body.push(Statement::Store {
2019            pointer: ptr,
2020            value: val,
2021        });
2022        func.body.push(Statement::Store {
2023            pointer: ptr,
2024            value: val,
2025        });
2026
2027        let dfg = DataflowGraph::build(&func);
2028        let cp = dfg.critical_path();
2029        assert_eq!(cp.node_distances.len(), 3);
2030        assert_eq!(cp.node_distances[0], 1);
2031        assert_eq!(cp.node_distances[1], 2);
2032        assert_eq!(cp.node_distances[2], 3);
2033        assert_eq!(cp.critical_path, vec![0, 1, 2]);
2034    }
2035
2036    // ===== Expression operand coverage: Access, AccessIndex, Select, Swizzle =====
2037
2038    #[test]
2039    fn expression_operands_access_and_select() {
2040        let mut func = Function::new("test");
2041        let gv = make_gv(0);
2042        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2043        let idx = func
2044            .expressions
2045            .append(Expression::Literal(Literal::U32(0)));
2046        let access = func.expressions.append(Expression::Access {
2047            base: ptr,
2048            index: idx,
2049        });
2050        let val_a = func
2051            .expressions
2052            .append(Expression::Literal(Literal::F32(1.0)));
2053        let val_b = func
2054            .expressions
2055            .append(Expression::Literal(Literal::F32(2.0)));
2056        let cond = func
2057            .expressions
2058            .append(Expression::Literal(Literal::Bool(true)));
2059        let select = func.expressions.append(Expression::Select {
2060            condition: cond,
2061            accept: val_a,
2062            reject: val_b,
2063        });
2064
2065        // Emit a range covering all expressions.
2066        let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2067        func.body.push(Statement::Emit(range));
2068
2069        let dfg = DataflowGraph::build(&func);
2070        assert_eq!(dfg.node_count(), 1);
2071        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Emit);
2072        // All expressions should appear in writes.
2073        assert!(dfg.nodes()[0].writes.contains(&access));
2074        assert!(dfg.nodes()[0].writes.contains(&select));
2075    }
2076
2077    // ===== Expression operands: Splat, As, ArrayLength, Compose, Math =====
2078
2079    #[test]
2080    fn expression_operands_splat_as_arraylength_compose_math() {
2081        let mut func = Function::new("test");
2082        let lit = func
2083            .expressions
2084            .append(Expression::Literal(Literal::F32(1.0)));
2085        let splat = func.expressions.append(Expression::Splat {
2086            size: nxpu_ir::VectorSize::Quad,
2087            value: lit,
2088        });
2089        let cast = func.expressions.append(Expression::As {
2090            expr: lit,
2091            kind: nxpu_ir::ScalarKind::Uint,
2092            convert: Some(4),
2093        });
2094        let arr_len = func.expressions.append(Expression::ArrayLength(lit));
2095
2096        let ty_handle = {
2097            let mut types = nxpu_ir::UniqueArena::new();
2098            types.insert(nxpu_ir::Type {
2099                name: None,
2100                inner: nxpu_ir::TypeInner::Vector {
2101                    size: nxpu_ir::VectorSize::Quad,
2102                    scalar: nxpu_ir::Scalar::F32,
2103                },
2104            })
2105        };
2106        let compose = func.expressions.append(Expression::Compose {
2107            ty: ty_handle,
2108            components: vec![lit],
2109        });
2110
2111        let lit2 = func
2112            .expressions
2113            .append(Expression::Literal(Literal::F32(2.0)));
2114        let lit3 = func
2115            .expressions
2116            .append(Expression::Literal(Literal::F32(3.0)));
2117        let math = func.expressions.append(Expression::Math {
2118            fun: nxpu_ir::MathFunction::Clamp,
2119            arg: lit,
2120            arg1: Some(lit2),
2121            arg2: Some(lit3),
2122            arg3: None,
2123        });
2124
2125        let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2126        func.body.push(Statement::Emit(range));
2127
2128        let dfg = DataflowGraph::build(&func);
2129        assert_eq!(dfg.nodes()[0].kind, DfgNodeKind::Emit);
2130        assert!(dfg.nodes()[0].writes.contains(&splat));
2131        assert!(dfg.nodes()[0].writes.contains(&cast));
2132        assert!(dfg.nodes()[0].writes.contains(&arr_len));
2133        assert!(dfg.nodes()[0].writes.contains(&compose));
2134        assert!(dfg.nodes()[0].writes.contains(&math));
2135    }
2136
2137    // ===== Expression operands: Swizzle, Unary, AccessIndex =====
2138
2139    #[test]
2140    fn expression_operands_swizzle_unary_accessindex() {
2141        let mut func = Function::new("test");
2142        let lit = func
2143            .expressions
2144            .append(Expression::Literal(Literal::F32(1.0)));
2145        let neg = func.expressions.append(Expression::Unary {
2146            op: nxpu_ir::UnaryOp::Negate,
2147            expr: lit,
2148        });
2149        let gv = make_gv(0);
2150        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2151        let access_idx = func.expressions.append(Expression::AccessIndex {
2152            base: ptr,
2153            index: 0,
2154        });
2155
2156        let vec_lit = func
2157            .expressions
2158            .append(Expression::Literal(Literal::F32(0.0)));
2159        let swizzle = func.expressions.append(Expression::Swizzle {
2160            size: nxpu_ir::VectorSize::Bi,
2161            vector: vec_lit,
2162            pattern: [
2163                nxpu_ir::SwizzleComponent::X,
2164                nxpu_ir::SwizzleComponent::Y,
2165                nxpu_ir::SwizzleComponent::Z,
2166                nxpu_ir::SwizzleComponent::W,
2167            ],
2168        });
2169
2170        let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2171        func.body.push(Statement::Emit(range));
2172
2173        let dfg = DataflowGraph::build(&func);
2174        assert!(dfg.nodes()[0].writes.contains(&neg));
2175        assert!(dfg.nodes()[0].writes.contains(&access_idx));
2176        assert!(dfg.nodes()[0].writes.contains(&swizzle));
2177    }
2178
2179    // ===== Expression operands: leaf expressions (no operands) =====
2180
2181    #[test]
2182    fn expression_operands_leaf_nodes() {
2183        let mut func = Function::new("test");
2184        let _func_arg = func.expressions.append(Expression::FunctionArgument(0));
2185        let gv = make_gv(0);
2186        let _gv_expr = func.expressions.append(Expression::GlobalVariable(gv));
2187
2188        let callee_handle = {
2189            let mut funcs = nxpu_ir::Arena::new();
2190            funcs.append(Function::new("callee"))
2191        };
2192        let _call_result = func
2193            .expressions
2194            .append(Expression::CallResult(callee_handle));
2195
2196        let ty_handle = {
2197            let mut types = nxpu_ir::UniqueArena::new();
2198            types.insert(nxpu_ir::Type {
2199                name: None,
2200                inner: nxpu_ir::TypeInner::Scalar(nxpu_ir::Scalar::U32),
2201            })
2202        };
2203        let _atomic_result = func.expressions.append(Expression::AtomicResult {
2204            ty: ty_handle,
2205            comparison: false,
2206        });
2207        let _zero_val = func.expressions.append(Expression::ZeroValue(ty_handle));
2208
2209        let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2210        func.body.push(Statement::Emit(range));
2211
2212        let dfg = DataflowGraph::build(&func);
2213        assert_eq!(dfg.node_count(), 1);
2214        // Leaf expressions produce no reads (they have no operands).
2215        // FunctionArgument, GlobalVariable, CallResult, AtomicResult, ZeroValue.
2216    }
2217
2218    // ===== Topological sort with diamond shape =====
2219
2220    #[test]
2221    fn topological_sort_diamond_shape() {
2222        let mut func = Function::new("test");
2223        let gv0 = make_gv(0);
2224        let gv1 = make_gv(1);
2225        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
2226        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
2227        let val = func
2228            .expressions
2229            .append(Expression::Literal(Literal::F32(1.0)));
2230        let load_a = func.expressions.append(Expression::Load { pointer: ptr_a });
2231
2232        // stmt 0: store val -> ptr_a (writes a)
2233        func.body.push(Statement::Store {
2234            pointer: ptr_a,
2235            value: val,
2236        });
2237        // stmt 1: store val -> ptr_b (writes b, independent)
2238        func.body.push(Statement::Store {
2239            pointer: ptr_b,
2240            value: val,
2241        });
2242        // stmt 2: store load(a) -> ptr_b (reads a -> RAW from 0, writes b -> WAW with 1)
2243        func.body.push(Statement::Store {
2244            pointer: ptr_b,
2245            value: load_a,
2246        });
2247
2248        let dfg = DataflowGraph::build(&func);
2249        let order = dfg.topological_sort().unwrap();
2250        assert_eq!(order.len(), 3);
2251
2252        // Node 0 must come before node 2 (RAW dependency on ptr_a).
2253        // Node 1 must come before node 2 (WAW dependency on ptr_b).
2254        let pos_0 = order.iter().position(|&n| n == 0).unwrap();
2255        let pos_1 = order.iter().position(|&n| n == 1).unwrap();
2256        let pos_2 = order.iter().position(|&n| n == 2).unwrap();
2257        assert!(pos_0 < pos_2);
2258        assert!(pos_1 < pos_2);
2259    }
2260
2261    // ===== Parallel groups with mixed independent and dependent ops =====
2262
2263    #[test]
2264    fn parallel_groups_mixed() {
2265        let mut func = Function::new("test");
2266        let gv0 = make_gv(0);
2267        let gv1 = make_gv(1);
2268        let ptr_a = func.expressions.append(Expression::GlobalVariable(gv0));
2269        let ptr_b = func.expressions.append(Expression::GlobalVariable(gv1));
2270        let val_a = func
2271            .expressions
2272            .append(Expression::Literal(Literal::F32(1.0)));
2273        let val_b = func
2274            .expressions
2275            .append(Expression::Literal(Literal::F32(2.0)));
2276
2277        // Two independent stores (no shared pointers or values).
2278        func.body.push(Statement::Store {
2279            pointer: ptr_a,
2280            value: val_a,
2281        });
2282        func.body.push(Statement::Store {
2283            pointer: ptr_b,
2284            value: val_b,
2285        });
2286        // Third store depends on ptr_a (WAW with stmt 0).
2287        func.body.push(Statement::Store {
2288            pointer: ptr_a,
2289            value: val_b,
2290        });
2291
2292        let dfg = DataflowGraph::build(&func);
2293        let groups = dfg.parallel_groups();
2294
2295        // First group should have at least 2 independent ops.
2296        assert_ne!(groups.len(), 0);
2297        let first_group = &groups[0];
2298        assert!(
2299            first_group.len() >= 2,
2300            "expected first parallel group to contain independent ops"
2301        );
2302    }
2303
2304    // ===== Single node graph =====
2305
2306    #[test]
2307    fn single_node_graph() {
2308        let mut func = Function::new("test");
2309        func.body.push(Statement::Break);
2310
2311        let dfg = DataflowGraph::build(&func);
2312        assert_eq!(dfg.node_count(), 1);
2313        assert_eq!(dfg.edge_count(), 0);
2314
2315        let order = dfg.topological_sort().unwrap();
2316        assert_eq!(order, vec![0]);
2317
2318        let cp = dfg.critical_path();
2319        assert_eq!(cp.critical_path_length, 1);
2320        assert_eq!(cp.critical_path, vec![0]);
2321
2322        let groups = dfg.parallel_groups();
2323        assert_eq!(groups.len(), 1);
2324        assert_eq!(groups[0], vec![0]);
2325    }
2326
2327    // ===== Display with all edge kinds =====
2328
2329    #[test]
2330    fn display_all_edge_kinds() {
2331        let mut func = Function::new("test");
2332        let gv = make_gv(0);
2333        let ptr = func.expressions.append(Expression::GlobalVariable(gv));
2334        let val = func
2335            .expressions
2336            .append(Expression::Literal(Literal::F32(1.0)));
2337        let load = func.expressions.append(Expression::Load { pointer: ptr });
2338
2339        // stmt 0: store val -> ptr (writes ptr)
2340        func.body.push(Statement::Store {
2341            pointer: ptr,
2342            value: val,
2343        });
2344        // stmt 1: store load(ptr) -> ptr (reads ptr, writes ptr -> RAW + WAW)
2345        func.body.push(Statement::Store {
2346            pointer: ptr,
2347            value: load,
2348        });
2349
2350        let dfg = DataflowGraph::build(&func);
2351        let s = format!("{dfg}");
2352        assert!(s.contains("RAW") || s.contains("WAW") || s.contains("WAR"));
2353    }
2354
2355    // ===== Math expression with all 4 arguments =====
2356
2357    #[test]
2358    fn math_expression_four_args() {
2359        let mut func = Function::new("test");
2360        let a = func
2361            .expressions
2362            .append(Expression::Literal(Literal::F32(1.0)));
2363        let b = func
2364            .expressions
2365            .append(Expression::Literal(Literal::F32(2.0)));
2366        let c = func
2367            .expressions
2368            .append(Expression::Literal(Literal::F32(3.0)));
2369        let d = func
2370            .expressions
2371            .append(Expression::Literal(Literal::F32(4.0)));
2372        let _math = func.expressions.append(Expression::Math {
2373            fun: nxpu_ir::MathFunction::Fma,
2374            arg: a,
2375            arg1: Some(b),
2376            arg2: Some(c),
2377            arg3: Some(d),
2378        });
2379
2380        let range = nxpu_ir::Range::from_index_range(0..func.expressions.len() as u32);
2381        func.body.push(Statement::Emit(range));
2382
2383        let dfg = DataflowGraph::build(&func);
2384        // The math expression should read all four argument literals.
2385        let reads = &dfg.nodes()[0].reads;
2386        assert!(reads.contains(&a), "math should read arg");
2387        assert!(reads.contains(&b), "math should read arg1");
2388        assert!(reads.contains(&c), "math should read arg2");
2389        assert!(reads.contains(&d), "math should read arg3");
2390    }
2391}