1use miden_constraint_compiler::ir::capture;
10use miden_core::{Felt, field::QuadFelt};
11use miden_crypto::stark::air::{BaseAir, LiftedAir};
12
13use crate::{
14 AceError, EXT_DEGREE,
15 circuit::{AceCircuit, emit_circuit},
16 dag::{AceDag, DagBuilder, NodeId, NodeKind, PeriodicColumnData, build_verifier_dag_from_ir},
17 layout::{InputCounts, InputKey, InputLayout},
18};
19
20#[derive(Debug, Clone, Copy)]
22pub enum LayoutKind {
23 Native,
25 Masm,
27}
28
29#[derive(Debug, Clone, Copy)]
31pub struct AceConfig {
32 pub num_quotient_chunks: usize,
34 pub layout: LayoutKind,
36 pub num_airs: usize,
41}
42
43#[derive(Debug)]
45pub struct AceArtifacts<EF> {
46 pub layout: InputLayout,
48 pub dag: AceDag<EF>,
50}
51
52pub fn build_ace_circuit_for_air<A>(
62 air: &A,
63 config: AceConfig,
64) -> Result<AceCircuit<QuadFelt>, AceError>
65where
66 A: LiftedAir<Felt, QuadFelt>,
67{
68 let artifacts = build_ace_dag_for_air(air, config)?;
69 emit_circuit(&artifacts.dag, artifacts.layout)
70}
71
72pub fn build_multi_air_ace_circuit<A>(
81 airs: &[A],
82 proof_order: &[usize],
83 config: AceConfig,
84 trace_width_alignment: usize,
85) -> Result<AceCircuit<QuadFelt>, AceError>
86where
87 A: LiftedAir<Felt, QuadFelt>,
88{
89 let num_airs = airs.len();
90 if num_airs == 0 || config.num_airs != num_airs {
91 return Err(AceError::InvalidInputLayout {
92 message: format!(
93 "multi-AIR composition requires a nonempty airs slice and matching num_airs; got \
94 {} AIRs and num_airs {}",
95 num_airs, config.num_airs
96 ),
97 });
98 }
99
100 let mut seen = vec![false; num_airs];
101 if proof_order.len() != num_airs
102 || proof_order
103 .iter()
104 .any(|&index| index >= num_airs || core::mem::replace(&mut seen[index], true))
105 {
106 return Err(AceError::InvalidInputLayout {
107 message: format!("proof_order must be a permutation of 0..{num_airs}"),
108 });
109 }
110 if trace_width_alignment == 0 {
111 return Err(AceError::InvalidInputLayout {
112 message: "trace width alignment must be nonzero".into(),
113 });
114 }
115
116 let sub_config = AceConfig { num_airs: 1, ..config };
117 let artifacts = build_ace_dags_for_airs(airs, sub_config)?;
118 let shared = artifacts[0].layout.counts;
119 if artifacts.iter().any(|air| air.layout.counts.num_public != shared.num_public) {
120 return Err(AceError::InvalidInputLayout {
121 message: "all AIRs must use the same public-value window".into(),
122 });
123 }
124
125 let mut offsets = vec![TraceOffsets::default(); num_airs];
126 let mut totals = TraceOffsets::default();
127 for &air_index in proof_order {
128 offsets[air_index] = totals;
129 let counts = artifacts[air_index].layout.counts;
130 totals.preprocessed += counts.preprocessed_width.next_multiple_of(trace_width_alignment);
131 totals.main += counts.width.next_multiple_of(trace_width_alignment);
132 let aligned_aux = (counts.aux_width * EXT_DEGREE).next_multiple_of(trace_width_alignment);
133 if !aligned_aux.is_multiple_of(EXT_DEGREE) {
134 return Err(AceError::InvalidInputLayout {
135 message: "aligned auxiliary width must be divisible by the extension degree".into(),
136 });
137 }
138 totals.aux += aligned_aux / EXT_DEGREE;
139 totals.boundary += counts.num_aux_boundary;
140 }
141
142 let counts = InputCounts {
143 preprocessed_width: totals.preprocessed,
144 width: totals.main,
145 aux_width: totals.aux,
146 num_aux_boundary: totals.boundary,
147 num_public: shared.num_public,
148 num_randomness: shared.num_randomness,
149 num_quotient_chunks: shared.num_quotient_chunks,
150 };
151 let layout = match config.layout {
152 LayoutKind::Native => InputLayout::new_multi_air(counts, num_airs),
153 LayoutKind::Masm => InputLayout::new_masm_multi_air(counts, num_airs),
154 };
155
156 let mut builder = DagBuilder::<QuadFelt>::new();
158 let mut roots = Vec::with_capacity(num_airs);
159 for (air_index, artifacts) in artifacts.iter().enumerate() {
160 roots.push(reemit_air_root(&mut builder, &artifacts.dag, air_index, offsets[air_index]));
161 }
162 let quotient_binding = roots[0].1;
163 if roots.iter().any(|&(_, binding)| binding != quotient_binding) {
164 return Err(AceError::InvalidInputLayout {
165 message: "all AIR quotient bindings must use the same q*v node".into(),
166 });
167 }
168
169 let beta = builder.input(InputKey::MultiAirFoldBeta);
170 let mut ordered = proof_order.iter().map(|&index| roots[index].0);
171 let mut accumulator = ordered.next().expect("multi-AIR composition is nonempty");
172 for next in ordered {
173 let scaled = builder.mul(accumulator, beta);
174 accumulator = builder.add(scaled, next);
175 }
176
177 let root = builder.sub(accumulator, quotient_binding);
179 let mut dag = builder.build(root);
180 dag.compact();
181 emit_circuit(&dag, layout)
182}
183
184pub fn build_ace_dag_for_air<A>(
188 air: &A,
189 config: AceConfig,
190) -> Result<AceArtifacts<QuadFelt>, AceError>
191where
192 A: LiftedAir<Felt, QuadFelt>,
193{
194 if config.num_airs == 0 {
195 return Err(AceError::InvalidInputLayout {
196 message: "num_airs must be at least 1".into(),
197 });
198 }
199
200 let periodic_columns = air.periodic_columns();
201 let shared_period = max_period(&periodic_columns);
202 build_ace_dag_for_air_with_periodic_columns(air, config, periodic_columns, shared_period)
203}
204
205fn build_ace_dags_for_airs<A>(
207 airs: &[A],
208 config: AceConfig,
209) -> Result<Vec<AceArtifacts<QuadFelt>>, AceError>
210where
211 A: LiftedAir<Felt, QuadFelt>,
212{
213 let periodic_columns_by_air: Vec<_> =
214 airs.iter().map(BaseAir::<Felt>::periodic_columns).collect();
215 let shared_period = periodic_columns_by_air
216 .iter()
217 .map(|columns| max_period(columns))
218 .max()
219 .unwrap_or(1);
220
221 airs.iter()
222 .zip(periodic_columns_by_air)
223 .map(|(air, periodic_columns)| {
224 build_ace_dag_for_air_with_periodic_columns(
225 air,
226 config,
227 periodic_columns,
228 shared_period,
229 )
230 })
231 .collect()
232}
233
234fn build_ace_dag_for_air_with_periodic_columns<A>(
235 air: &A,
236 config: AceConfig,
237 periodic_columns: Vec<Vec<Felt>>,
238 shared_period: usize,
239) -> Result<AceArtifacts<QuadFelt>, AceError>
240where
241 A: LiftedAir<Felt, QuadFelt>,
242{
243 let counts = input_counts_for_air(air, config)?;
244 let layout = match (config.layout, config.num_airs >= 2) {
245 (LayoutKind::Native, false) => InputLayout::new(counts),
246 (LayoutKind::Masm, false) => InputLayout::new_masm(counts),
247 (LayoutKind::Native, true) => InputLayout::new_multi_air(counts, config.num_airs),
248 (LayoutKind::Masm, true) => InputLayout::new_masm_multi_air(counts, config.num_airs),
249 };
250 layout.validate();
251
252 let (graph, constraints) = capture(air);
253 let periodic_data = (!periodic_columns.is_empty())
254 .then(|| PeriodicColumnData::from_periodic_columns::<Felt>(periodic_columns));
255 let dag = build_verifier_dag_from_ir(
256 &graph,
257 &constraints,
258 &layout,
259 periodic_data.as_ref(),
260 shared_period,
261 );
262
263 Ok(AceArtifacts { layout, dag })
264}
265
266fn max_period<F>(periodic_columns: &[Vec<F>]) -> usize {
267 periodic_columns.iter().map(Vec::len).max().unwrap_or(1)
268}
269
270#[derive(Clone, Copy, Default)]
271struct TraceOffsets {
272 preprocessed: usize,
273 main: usize,
274 aux: usize,
275 boundary: usize,
276}
277
278fn reemit_air_root(
279 builder: &mut DagBuilder<QuadFelt>,
280 source: &AceDag<QuadFelt>,
281 air_index: usize,
282 offsets: TraceOffsets,
283) -> (NodeId, NodeId) {
284 debug_assert_eq!(source.root().index() + 1, source.nodes.len());
285 let NodeKind::Sub(accumulator, quotient_binding) = source.nodes[source.root().index()] else {
286 unreachable!("verifier DAGs always emit an accumulator - q*v root")
287 };
288
289 let mut translated = Vec::with_capacity(source.nodes.len() - 1);
290 for node in &source.nodes[..source.root().index()] {
291 let id = match *node {
292 NodeKind::Input(key) => {
293 let key = match key {
294 InputKey::Preprocessed { offset, index } => InputKey::Preprocessed {
295 offset,
296 index: index + offsets.preprocessed,
297 },
298 InputKey::Main { offset, index } => {
299 InputKey::Main { offset, index: index + offsets.main }
300 },
301 InputKey::AuxCoord { offset, index, coord } => InputKey::AuxCoord {
302 offset,
303 index: index + offsets.aux,
304 coord,
305 },
306 InputKey::AuxBusBoundary(index) => {
307 InputKey::AuxBusBoundary(index + offsets.boundary)
308 },
309 InputKey::IsFirst => InputKey::IsFirstAir(air_index),
310 InputKey::IsLast => InputKey::IsLastAir(air_index),
311 InputKey::IsTransition => InputKey::IsTransitionAir(air_index),
312 other => other,
313 };
314 builder.input(key)
315 },
316 NodeKind::Constant(value) => builder.constant(value),
317 NodeKind::Add(a, b) => builder.add(translated[a.index()], translated[b.index()]),
318 NodeKind::Sub(a, b) => builder.sub(translated[a.index()], translated[b.index()]),
319 NodeKind::Mul(a, b) => builder.mul(translated[a.index()], translated[b.index()]),
320 NodeKind::Neg(a) => builder.neg(translated[a.index()]),
321 };
322 translated.push(id);
323 }
324
325 (translated[accumulator.index()], translated[quotient_binding.index()])
326}
327
328fn input_counts_for_air<A>(air: &A, config: AceConfig) -> Result<InputCounts, AceError>
329where
330 A: LiftedAir<Felt, QuadFelt>,
331{
332 if config.num_quotient_chunks == 0 {
333 return Err(AceError::InvalidInputLayout {
334 message: "num_quotient_chunks must be > 0".into(),
335 });
336 }
337 let num_randomness = air.num_randomness();
338 if num_randomness != 2 {
339 return Err(AceError::InvalidInputLayout {
340 message: format!(
341 "AIR must declare exactly 2 randomness challenges (alpha, beta), got {num_randomness}"
342 ),
343 });
344 }
345
346 Ok(InputCounts {
347 preprocessed_width: air.preprocessed_width(),
348 width: air.width(),
349 aux_width: air.aux_width(),
350 num_aux_boundary: air.num_aux_values(),
351 num_public: air.num_public_values(),
352 num_randomness,
353 num_quotient_chunks: config.num_quotient_chunks,
354 })
355}