Skip to main content

miden_ace_codegen/dag/
builder.rs

1use std::collections::HashMap;
2
3use miden_crypto::field::Field;
4
5use super::ir::{DagId, DagSnapshot, NodeId, NodeKind};
6use crate::layout::InputKey;
7
8/// A hash-consed DAG builder.
9///
10/// The builder de-duplicates identical subexpressions to keep the circuit
11/// compact and deterministic.
12#[derive(Debug)]
13pub struct DagBuilder<EF> {
14    dag_id: DagId,
15    nodes: Vec<NodeKind<EF>>,
16    cache: HashMap<NodeKind<EF>, NodeId>,
17    imported_dag: Option<ImportedDag>,
18}
19
20impl<EF> DagBuilder<EF>
21where
22    EF: Field,
23{
24    /// Create an empty, hash-consed DAG builder.
25    pub fn new() -> Self {
26        Self {
27            dag_id: DagId::fresh(),
28            nodes: Vec::new(),
29            cache: HashMap::new(),
30            imported_dag: None,
31        }
32    }
33
34    /// Resume building from existing nodes using the published 0.23.0 API shape.
35    ///
36    /// Imported node ids are rebased onto the new builder, and ids from the source DAG
37    /// are accepted only when that provenance is encoded in the node graph itself.
38    pub fn from_nodes(nodes: Vec<NodeKind<EF>>) -> Self {
39        let imported_dag = infer_dag_id(&nodes)
40            .map(|source_dag_id| ImportedDag { source_dag_id, imported_len: nodes.len() });
41        let dag_id = DagId::fresh();
42        let nodes = rebase_nodes(nodes, dag_id);
43
44        Self::from_existing_nodes(dag_id, nodes, imported_dag)
45    }
46
47    /// Resume building from an exported snapshot.
48    ///
49    /// This preserves the original DAG id even when the imported nodes are all leaves.
50    pub fn from_snapshot(snapshot: DagSnapshot<EF>) -> Self {
51        let (source_dag_id, nodes, _) = snapshot.into_parts();
52        let dag_id = DagId::fresh();
53        let imported_dag = Some(ImportedDag { source_dag_id, imported_len: nodes.len() });
54        let nodes = rebase_nodes(nodes, dag_id);
55
56        Self::from_existing_nodes(dag_id, nodes, imported_dag)
57    }
58
59    /// Resume building from an existing DAG.
60    ///
61    /// Rebuilds the deduplication cache so that subsequent operations reuse
62    /// existing subexpressions.
63    pub fn from_dag(dag: super::AceDag<EF>) -> Self {
64        let dag_id = dag.dag_id();
65        Self::from_existing_nodes(dag_id, dag.into_nodes(), None)
66    }
67
68    fn from_existing_nodes(
69        dag_id: DagId,
70        nodes: Vec<NodeKind<EF>>,
71        imported_dag: Option<ImportedDag>,
72    ) -> Self {
73        let cache = nodes
74            .iter()
75            .enumerate()
76            .map(|(i, n)| (n.clone(), NodeId::in_dag(i, dag_id)))
77            .collect();
78        Self { dag_id, nodes, cache, imported_dag }
79    }
80
81    /// Consume the builder and return its node list.
82    pub fn into_nodes(self) -> Vec<NodeKind<EF>> {
83        self.nodes
84    }
85
86    /// Consume the builder and return a DAG with the provided root.
87    pub fn build(self, root: NodeId) -> super::AceDag<EF> {
88        let root = self.resolve_id(root, "DAG root must refer to a node built by this DagBuilder");
89
90        super::AceDag::from_parts(self.dag_id, self.nodes, root)
91    }
92
93    /// Add an input node.
94    pub fn input(&mut self, key: InputKey) -> NodeId {
95        self.intern(NodeKind::Input(key))
96    }
97
98    /// Add a constant node.
99    pub fn constant(&mut self, value: EF) -> NodeId {
100        self.intern(NodeKind::Constant(value))
101    }
102
103    /// Add an addition node with constant folding, add/sub cancellation, and negation
104    /// normalization.
105    pub fn add(&mut self, a: NodeId, b: NodeId) -> NodeId {
106        let a = self.resolve_node(a);
107        let b = self.resolve_node(b);
108        if let (Some(x), Some(y)) = (self.const_value(a), self.const_value(b)) {
109            return self.constant(x + y);
110        }
111        if self.is_zero(a) {
112            return b;
113        }
114        if self.is_zero(b) {
115            return a;
116        }
117        if let Some(result) = self.cancel_add(a, b) {
118            return result;
119        }
120        match (self.negated(a), self.negated(b)) {
121            (Some(a), None) => return self.sub(b, a),
122            (None, Some(b)) => return self.sub(a, b),
123            _ => {},
124        }
125        let (l, r) = if a <= b { (a, b) } else { (b, a) };
126        self.intern(NodeKind::Add(l, r))
127    }
128
129    /// Add a subtraction node with constant folding, add/sub cancellation, and negation
130    /// normalization.
131    pub fn sub(&mut self, a: NodeId, b: NodeId) -> NodeId {
132        let a = self.resolve_node(a);
133        let b = self.resolve_node(b);
134        if let (Some(x), Some(y)) = (self.const_value(a), self.const_value(b)) {
135            return self.constant(x - y);
136        }
137        if self.is_zero(b) {
138            return a;
139        }
140        if a == b {
141            return self.constant(EF::ZERO);
142        }
143        if let Some(result) = self.cancel_sub(a, b) {
144            return result;
145        }
146        if let Some(b) = self.negated(b) {
147            return self.add(a, b);
148        }
149        self.intern(NodeKind::Sub(a, b))
150    }
151
152    /// Add a multiplication node (with constant folding).
153    pub fn mul(&mut self, a: NodeId, b: NodeId) -> NodeId {
154        let a = self.resolve_node(a);
155        let b = self.resolve_node(b);
156        if let (Some(x), Some(y)) = (self.const_value(a), self.const_value(b)) {
157            return self.constant(x * y);
158        }
159        if self.is_zero(a) || self.is_zero(b) {
160            return self.constant(EF::ZERO);
161        }
162        if self.is_one(a) {
163            return b;
164        }
165        if self.is_one(b) {
166            return a;
167        }
168        let (l, r) = if a <= b { (a, b) } else { (b, a) };
169        self.intern(NodeKind::Mul(l, r))
170    }
171
172    /// Add a negation node (with constant folding).
173    pub fn neg(&mut self, a: NodeId) -> NodeId {
174        let a = self.resolve_node(a);
175        if let Some(x) = self.const_value(a) {
176            return self.constant(-x);
177        }
178        self.intern(NodeKind::Neg(a))
179    }
180
181    fn const_value(&self, id: NodeId) -> Option<EF> {
182        match self.nodes.get(id.index())? {
183            NodeKind::Constant(v) => Some(*v),
184            _ => None,
185        }
186    }
187
188    fn is_zero(&self, id: NodeId) -> bool {
189        self.const_value(id).is_some_and(|v| v == EF::ZERO)
190    }
191
192    fn is_one(&self, id: NodeId) -> bool {
193        self.const_value(id).is_some_and(|v| v == EF::ONE)
194    }
195
196    fn negated(&self, id: NodeId) -> Option<NodeId> {
197        match self.nodes.get(id.index())? {
198            NodeKind::Neg(inner) => Some(*inner),
199            _ => None,
200        }
201    }
202
203    fn cancel_add(&self, a: NodeId, b: NodeId) -> Option<NodeId> {
204        for (term, other) in [(a, b), (b, a)] {
205            if let NodeKind::Sub(lhs, rhs) = self.nodes[term.index()]
206                && rhs == other
207            {
208                return Some(lhs);
209            }
210        }
211        None
212    }
213
214    fn cancel_sub(&self, a: NodeId, b: NodeId) -> Option<NodeId> {
215        if let NodeKind::Add(lhs, rhs) = self.nodes[a.index()] {
216            if lhs == b {
217                return Some(rhs);
218            }
219            if rhs == b {
220                return Some(lhs);
221            }
222        }
223        if let NodeKind::Sub(lhs, rhs) = self.nodes[b.index()]
224            && lhs == a
225        {
226            return Some(rhs);
227        }
228        None
229    }
230
231    fn resolve_node(&self, id: NodeId) -> NodeId {
232        self.resolve_id(id, "DAG node must come from this DagBuilder")
233    }
234
235    fn intern(&mut self, node: NodeKind<EF>) -> NodeId {
236        if let Some(id) = self.cache.get(&node) {
237            return *id;
238        }
239        let id = NodeId::in_dag(self.nodes.len(), self.dag_id);
240        self.nodes.push(node.clone());
241        self.cache.insert(node, id);
242        id
243    }
244
245    fn resolve_id(&self, id: NodeId, message: &str) -> NodeId {
246        assert!(id.index() < self.nodes.len(), "{message}");
247
248        if id.dag_id == self.dag_id {
249            return id;
250        }
251
252        if let Some(imported) = &self.imported_dag
253            && imported.source_dag_id == id.dag_id
254            && id.index() < imported.imported_len
255        {
256            return NodeId::in_dag(id.index(), self.dag_id);
257        }
258
259        panic!("{message}");
260    }
261}
262
263fn infer_dag_id<EF>(nodes: &[NodeKind<EF>]) -> Option<DagId> {
264    nodes.iter().find_map(|node| match node {
265        NodeKind::Add(a, _) | NodeKind::Sub(a, _) | NodeKind::Mul(a, _) | NodeKind::Neg(a) => {
266            Some(a.dag_id)
267        },
268        NodeKind::Input(_) | NodeKind::Constant(_) => None,
269    })
270}
271
272fn rebase_nodes<EF>(nodes: Vec<NodeKind<EF>>, dag_id: DagId) -> Vec<NodeKind<EF>> {
273    nodes
274        .into_iter()
275        .map(|node| match node {
276            NodeKind::Input(key) => NodeKind::Input(key),
277            NodeKind::Constant(value) => NodeKind::Constant(value),
278            NodeKind::Add(a, b) => NodeKind::Add(rebase_node(a, dag_id), rebase_node(b, dag_id)),
279            NodeKind::Sub(a, b) => NodeKind::Sub(rebase_node(a, dag_id), rebase_node(b, dag_id)),
280            NodeKind::Mul(a, b) => NodeKind::Mul(rebase_node(a, dag_id), rebase_node(b, dag_id)),
281            NodeKind::Neg(a) => NodeKind::Neg(rebase_node(a, dag_id)),
282        })
283        .collect()
284}
285
286fn rebase_node(id: NodeId, dag_id: DagId) -> NodeId {
287    NodeId::in_dag(id.index(), dag_id)
288}
289
290#[derive(Debug, Clone)]
291struct ImportedDag {
292    source_dag_id: DagId,
293    imported_len: usize,
294}
295
296impl<EF> Default for DagBuilder<EF>
297where
298    EF: Field,
299{
300    fn default() -> Self {
301        Self::new()
302    }
303}
304
305#[cfg(test)]
306mod tests {
307    use miden_core::{Felt, field::QuadFelt};
308
309    use super::DagBuilder;
310    use crate::layout::InputKey;
311
312    fn felt(value: u64) -> QuadFelt {
313        QuadFelt::from(Felt::new_unchecked(value))
314    }
315
316    #[test]
317    #[should_panic(expected = "DAG root must refer to a node built by this DagBuilder")]
318    fn build_rejects_same_index_root_from_another_builder() {
319        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
320        let foreign_root = foreign_builder.constant(felt(1));
321
322        let mut builder = DagBuilder::<QuadFelt>::new();
323        builder.constant(felt(1));
324
325        let _ = builder.build(foreign_root);
326    }
327
328    #[test]
329    #[should_panic(expected = "DAG node must come from this DagBuilder")]
330    fn add_rejects_foreign_node() {
331        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
332        let foreign = foreign_builder.constant(felt(2));
333
334        let mut builder = DagBuilder::<QuadFelt>::new();
335        let local = builder.constant(felt(1));
336
337        let _ = builder.add(local, foreign);
338    }
339
340    #[test]
341    #[should_panic(expected = "DAG node must come from this DagBuilder")]
342    fn sub_rejects_foreign_node() {
343        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
344        let foreign = foreign_builder.constant(felt(2));
345
346        let mut builder = DagBuilder::<QuadFelt>::new();
347        let local = builder.constant(felt(1));
348
349        let _ = builder.sub(local, foreign);
350    }
351
352    #[test]
353    #[should_panic(expected = "DAG node must come from this DagBuilder")]
354    fn mul_rejects_foreign_node() {
355        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
356        let foreign = foreign_builder.constant(felt(2));
357
358        let mut builder = DagBuilder::<QuadFelt>::new();
359        let local = builder.constant(felt(1));
360
361        let _ = builder.mul(local, foreign);
362    }
363
364    #[test]
365    #[should_panic(expected = "DAG node must come from this DagBuilder")]
366    fn neg_rejects_foreign_node() {
367        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
368        let foreign = foreign_builder.constant(felt(2));
369
370        let mut builder = DagBuilder::<QuadFelt>::new();
371        let _ = builder.constant(felt(1));
372
373        let _ = builder.neg(foreign);
374    }
375
376    #[test]
377    fn from_dag_preserves_node_ownership() {
378        let mut builder = DagBuilder::<QuadFelt>::new();
379        let a = builder.constant(felt(1));
380        let dag = builder.build(a);
381        let root = dag.root();
382
383        let mut rebuilt = DagBuilder::from_dag(dag);
384        let b = rebuilt.constant(felt(2));
385        let sum = rebuilt.add(root, b);
386
387        let rebuilt_dag = rebuilt.build(sum);
388        assert_eq!(rebuilt_dag.root().index(), sum.index());
389    }
390
391    #[test]
392    fn from_nodes_accepts_published_root_shape() {
393        let mut builder = DagBuilder::<QuadFelt>::new();
394        let a = builder.input(InputKey::Reserved);
395        let b = builder.constant(felt(2));
396        let root = builder.add(a, b);
397        let dag = builder.build(root);
398
399        let mut rebuilt = DagBuilder::from_nodes(dag.nodes.clone());
400        let c = rebuilt.constant(felt(3));
401        let sum = rebuilt.add(dag.root, c);
402
403        let rebuilt_dag = rebuilt.build(sum);
404        assert_eq!(rebuilt_dag.root().index(), sum.index());
405    }
406
407    #[test]
408    fn from_nodes_accepts_leaf_only_root_shape() {
409        let mut builder = DagBuilder::<QuadFelt>::new();
410        let a = builder.constant(felt(1));
411        let dag = builder.build(a);
412
413        let root = dag.root();
414        let mut rebuilt = DagBuilder::from_snapshot(dag.into_snapshot());
415        let b = rebuilt.constant(felt(2));
416        let sum = rebuilt.add(root, b);
417
418        let rebuilt_dag = rebuilt.build(sum);
419        assert_eq!(rebuilt_dag.root().index(), sum.index());
420    }
421
422    #[test]
423    fn from_snapshot_accepts_leaf_only_root_after_source_dag_is_dropped() {
424        let mut builder = DagBuilder::<QuadFelt>::new();
425        let a = builder.constant(felt(1));
426        let snapshot = builder.build(a).into_snapshot();
427        let root = snapshot.root();
428
429        let mut rebuilt = DagBuilder::from_snapshot(snapshot);
430        let b = rebuilt.constant(felt(2));
431        let sum = rebuilt.add(root, b);
432
433        let rebuilt_dag = rebuilt.build(sum);
434        assert_eq!(rebuilt_dag.root().index(), sum.index());
435    }
436
437    #[test]
438    fn addition_absorbs_one_negated_operand() {
439        let mut builder = DagBuilder::<QuadFelt>::new();
440        let a = builder.input(InputKey::Public(0));
441        let b = builder.input(InputKey::Public(1));
442        let neg_a = builder.neg(a);
443        let neg_b = builder.neg(b);
444
445        let root = builder.add(a, neg_b);
446        assert_eq!(root, builder.sub(a, b));
447        assert_eq!(builder.add(neg_a, b), builder.sub(b, a));
448
449        let mut dag = builder.build(root);
450        dag.compact();
451        assert_eq!(dag.nodes.len(), 3, "the absorbed negation must become unreachable");
452    }
453
454    #[test]
455    fn subtraction_absorbs_a_negated_rhs() {
456        let mut builder = DagBuilder::<QuadFelt>::new();
457        let a = builder.input(InputKey::Public(0));
458        let b = builder.input(InputKey::Public(1));
459        let neg_b = builder.neg(b);
460
461        let root = builder.sub(a, neg_b);
462        assert_eq!(root, builder.add(a, b));
463
464        let mut dag = builder.build(root);
465        dag.compact();
466        assert_eq!(dag.nodes.len(), 3, "the absorbed negation must become unreachable");
467    }
468
469    #[test]
470    fn addition_and_subtraction_cancel_inverse_terms() {
471        let mut builder = DagBuilder::<QuadFelt>::new();
472        let a = builder.input(InputKey::Public(0));
473        let b = builder.input(InputKey::Public(1));
474        let difference = builder.sub(a, b);
475        let sum = builder.add(a, b);
476
477        assert_eq!(builder.add(difference, b), a);
478        assert_eq!(builder.add(b, difference), a);
479        assert_eq!(builder.sub(sum, a), b);
480        assert_eq!(builder.sub(sum, b), a);
481        assert_eq!(builder.sub(a, difference), b);
482        assert_eq!(builder.sub(a, a), builder.constant(felt(0)));
483    }
484
485    #[test]
486    #[should_panic(expected = "DAG node must come from this DagBuilder")]
487    fn from_nodes_rejects_foreign_node_from_another_builder() {
488        let mut source_builder = DagBuilder::<QuadFelt>::new();
489        let a = source_builder.input(InputKey::Reserved);
490        let b = source_builder.constant(felt(2));
491        let root = source_builder.add(a, b);
492        let dag = source_builder.build(root);
493
494        let mut rebuilt = DagBuilder::from_nodes(dag.nodes.clone());
495        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
496        let foreign = foreign_builder.constant(felt(3));
497
498        let _ = rebuilt.add(dag.root, foreign);
499    }
500
501    #[test]
502    #[should_panic(expected = "DAG root must refer to a node built by this DagBuilder")]
503    fn from_nodes_rejects_foreign_root_from_another_builder() {
504        let mut source_builder = DagBuilder::<QuadFelt>::new();
505        let a = source_builder.input(InputKey::Reserved);
506        let b = source_builder.constant(felt(2));
507        let root = source_builder.add(a, b);
508        let dag = source_builder.build(root);
509
510        let rebuilt = DagBuilder::from_nodes(dag.nodes);
511        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
512        let foreign = foreign_builder.constant(felt(3));
513
514        let _ = rebuilt.build(foreign);
515    }
516
517    #[test]
518    #[should_panic(expected = "DAG root must refer to a node built by this DagBuilder")]
519    fn from_nodes_leaf_only_rejects_foreign_root_before_any_imported_id() {
520        let mut source_builder = DagBuilder::<QuadFelt>::new();
521        let source = source_builder.constant(felt(1));
522        let dag = source_builder.build(source);
523
524        let rebuilt = DagBuilder::from_nodes(dag.nodes.clone());
525        let _ = rebuilt.build(dag.root);
526    }
527
528    #[test]
529    #[should_panic(expected = "DAG root must refer to a node built by this DagBuilder")]
530    fn from_snapshot_leaf_only_rejects_foreign_root() {
531        let mut source_builder = DagBuilder::<QuadFelt>::new();
532        let source = source_builder.constant(felt(1));
533        let snapshot = source_builder.build(source).into_snapshot();
534
535        let mut foreign_builder = DagBuilder::<QuadFelt>::new();
536        let foreign = foreign_builder.constant(felt(3));
537        let foreign_dag = foreign_builder.build(foreign);
538
539        let rebuilt = DagBuilder::from_snapshot(snapshot);
540        let _ = rebuilt.build(foreign_dag.root);
541    }
542}