1use std::collections::HashMap;
2
3use miden_crypto::field::Field;
4
5use super::ir::{DagId, DagSnapshot, NodeId, NodeKind};
6use crate::layout::InputKey;
7
8#[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 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 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 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 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 pub fn into_nodes(self) -> Vec<NodeKind<EF>> {
83 self.nodes
84 }
85
86 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 pub fn input(&mut self, key: InputKey) -> NodeId {
95 self.intern(NodeKind::Input(key))
96 }
97
98 pub fn constant(&mut self, value: EF) -> NodeId {
100 self.intern(NodeKind::Constant(value))
101 }
102
103 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 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 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 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}