1use std::collections::HashMap;
7
8use miden_crypto::field::Field;
9
10use crate::{
11 AceError, InputLayout,
12 dag::{AceDag, NodeId, NodeKind},
13};
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub(crate) enum AceOp {
18 Add,
19 Sub,
20 Mul,
21}
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
25pub(crate) enum AceNode {
26 Input(usize),
27 Constant(usize),
28 Operation(usize),
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub(crate) struct AceOpNode {
34 pub op: AceOp,
35 pub lhs: AceNode,
36 pub rhs: AceNode,
37}
38
39#[derive(Debug, Clone)]
43pub struct AceCircuit<EF> {
44 pub(crate) layout: InputLayout,
45 pub(crate) constants: Vec<EF>,
46 pub(crate) operations: Vec<AceOpNode>,
47 pub(crate) root: AceNode,
48}
49
50impl<EF: Field> AceCircuit<EF> {
51 pub fn layout(&self) -> &InputLayout {
53 &self.layout
54 }
55
56 pub fn eval(&self, inputs: &[EF]) -> Result<EF, AceError> {
58 if inputs.len() != self.layout.total_inputs {
59 return Err(AceError::InvalidInputLength {
60 expected: self.layout.total_inputs,
61 got: inputs.len(),
62 });
63 }
64 let mut op_values = vec![EF::ZERO; self.operations.len()];
65 for (idx, op) in self.operations.iter().enumerate() {
66 let lhs = self.node_value(op.lhs, inputs, &op_values);
67 let rhs = self.node_value(op.rhs, inputs, &op_values);
68 op_values[idx] = match op.op {
69 AceOp::Add => lhs + rhs,
70 AceOp::Sub => lhs - rhs,
71 AceOp::Mul => lhs * rhs,
72 };
73 }
74 Ok(self.node_value(self.root, inputs, &op_values))
75 }
76
77 pub fn num_nodes(&self) -> usize {
79 self.layout.total_inputs + self.constants.len() + self.operations.len()
80 }
81
82 fn node_value(&self, node: AceNode, inputs: &[EF], op_values: &[EF]) -> EF {
83 match node {
84 AceNode::Input(index) => inputs[index],
85 AceNode::Constant(index) => self.constants[index],
86 AceNode::Operation(index) => op_values[index],
87 }
88 }
89}
90
91pub fn emit_circuit<EF>(dag: &AceDag<EF>, layout: InputLayout) -> Result<AceCircuit<EF>, AceError>
93where
94 EF: Field,
95{
96 layout.validate();
97
98 let mut constants = Vec::new();
99 let mut constant_map = HashMap::<EF, usize>::new();
100 let mut operations = Vec::new();
101 let mut node_map: Vec<Option<AceNode>> = vec![None; dag.nodes().len()];
102
103 for (idx, node) in dag.nodes().iter().enumerate() {
104 let ace_node = match node {
105 NodeKind::Input(key) => {
106 let input_idx = layout.index(*key).ok_or_else(|| AceError::InvalidInputLayout {
107 message: format!("missing input key in layout: {key:?}"),
108 })?;
109 AceNode::Input(input_idx)
110 },
111 NodeKind::Constant(value) => {
112 let const_idx = *constant_map.entry(*value).or_insert_with(|| {
113 constants.push(*value);
114 constants.len() - 1
115 });
116 AceNode::Constant(const_idx)
117 },
118 NodeKind::Add(a, b) => {
119 let lhs = lookup_node(&node_map, *a);
120 let rhs = lookup_node(&node_map, *b);
121 let op_idx = operations.len();
122 operations.push(AceOpNode { op: AceOp::Add, lhs, rhs });
123 AceNode::Operation(op_idx)
124 },
125 NodeKind::Sub(a, b) => {
126 let lhs = lookup_node(&node_map, *a);
127 let rhs = lookup_node(&node_map, *b);
128 let op_idx = operations.len();
129 operations.push(AceOpNode { op: AceOp::Sub, lhs, rhs });
130 AceNode::Operation(op_idx)
131 },
132 NodeKind::Mul(a, b) => {
133 let lhs = lookup_node(&node_map, *a);
134 let rhs = lookup_node(&node_map, *b);
135 let op_idx = operations.len();
136 operations.push(AceOpNode { op: AceOp::Mul, lhs, rhs });
137 AceNode::Operation(op_idx)
138 },
139 NodeKind::Neg(a) => {
140 let rhs = lookup_node(&node_map, *a);
141 let zero = *constant_map.entry(EF::ZERO).or_insert_with(|| {
142 constants.push(EF::ZERO);
143 constants.len() - 1
144 });
145 let op_idx = operations.len();
146 operations.push(AceOpNode {
147 op: AceOp::Sub,
148 lhs: AceNode::Constant(zero),
149 rhs,
150 });
151 AceNode::Operation(op_idx)
152 },
153 };
154 node_map[idx] = Some(ace_node);
155 }
156
157 let root = lookup_node(&node_map, dag.root());
158 Ok(AceCircuit { layout, constants, operations, root })
159}
160
161fn lookup_node(map: &[Option<AceNode>], id: NodeId) -> AceNode {
162 map[id.index()].expect("ACE DAG nodes must be topologically ordered")
163}