1use miden_constraint_compiler::ir::capture;
10use miden_core::{Felt, field::QuadFelt};
11use miden_crypto::{
12 field::Field,
13 stark::air::{BaseAir, LiftedAir},
14};
15
16use crate::{
17 AceError, EXT_DEGREE,
18 circuit::{AceCircuit, emit_circuit},
19 dag::{
20 AceDag, DagBuilder, NodeId, NodeKind, PeriodicColumnData, build_verifier_dag_from_ir,
21 normalize_dag,
22 },
23 factored::{FactoredAceCircuit, ShuffleEncodeBuffer, emit_factored_circuit},
24 layout::{InputCounts, InputKey, InputLayout},
25};
26
27#[derive(Debug, Clone, Copy)]
29pub enum LayoutKind {
30 Native,
32 Masm,
34}
35
36#[derive(Debug, Clone, Copy)]
38pub struct AceConfig {
39 pub num_quotient_chunks: usize,
41 pub layout: LayoutKind,
43 pub num_airs: usize,
48}
49
50#[derive(Debug)]
52pub struct AceArtifacts<EF> {
53 pub layout: InputLayout,
55 pub dag: AceDag<EF>,
57}
58
59pub fn build_ace_circuit_for_air<A>(
69 air: &A,
70 config: AceConfig,
71) -> Result<AceCircuit<QuadFelt>, AceError>
72where
73 A: LiftedAir<Felt, QuadFelt>,
74{
75 let artifacts = build_ace_dag_for_air(air, config)?;
76 emit_circuit(&artifacts.dag, artifacts.layout)
77}
78
79pub fn build_multi_air_ace_circuit<A>(
88 airs: &[A],
89 proof_order: &[usize],
90 config: AceConfig,
91 trace_width_alignment: usize,
92) -> Result<AceCircuit<QuadFelt>, AceError>
93where
94 A: LiftedAir<Felt, QuadFelt>,
95{
96 let num_airs = airs.len();
97 if num_airs == 0 || config.num_airs != num_airs {
98 return Err(AceError::InvalidInputLayout {
99 message: format!(
100 "multi-AIR composition requires a nonempty airs slice and matching num_airs; got \
101 {} AIRs and num_airs {}",
102 num_airs, config.num_airs
103 ),
104 });
105 }
106
107 let mut seen = vec![false; num_airs];
108 if proof_order.len() != num_airs
109 || proof_order
110 .iter()
111 .any(|&index| index >= num_airs || core::mem::replace(&mut seen[index], true))
112 {
113 return Err(AceError::InvalidInputLayout {
114 message: format!("proof_order must be a permutation of 0..{num_airs}"),
115 });
116 }
117 if trace_width_alignment == 0 {
118 return Err(AceError::InvalidInputLayout {
119 message: "trace width alignment must be nonzero".into(),
120 });
121 }
122
123 let sub_config = AceConfig { num_airs: 1, ..config };
124 let artifacts = build_ace_dags_for_airs(airs, sub_config)?;
125 let shared = artifacts[0].layout.counts;
126 if artifacts.iter().any(|air| air.layout.counts.num_public != shared.num_public) {
127 return Err(AceError::InvalidInputLayout {
128 message: "all AIRs must use the same public-value window".into(),
129 });
130 }
131
132 let mut offsets = vec![TraceOffsets::default(); num_airs];
133 let mut totals = TraceOffsets::default();
134 for &air_index in proof_order {
135 offsets[air_index] = totals;
136 let counts = artifacts[air_index].layout.counts;
137 totals.preprocessed += counts.preprocessed_width.next_multiple_of(trace_width_alignment);
138 totals.main += counts.width.next_multiple_of(trace_width_alignment);
139 let aligned_aux = (counts.aux_width * EXT_DEGREE).next_multiple_of(trace_width_alignment);
140 if !aligned_aux.is_multiple_of(EXT_DEGREE) {
141 return Err(AceError::InvalidInputLayout {
142 message: "aligned auxiliary width must be divisible by the extension degree".into(),
143 });
144 }
145 totals.aux += aligned_aux / EXT_DEGREE;
146 totals.boundary += counts.num_aux_boundary;
147 }
148
149 let counts = InputCounts {
150 preprocessed_width: totals.preprocessed,
151 width: totals.main,
152 aux_width: totals.aux,
153 num_aux_boundary: totals.boundary,
154 num_public: shared.num_public,
155 num_randomness: shared.num_randomness,
156 num_quotient_chunks: shared.num_quotient_chunks,
157 };
158 let layout = match config.layout {
159 LayoutKind::Native => InputLayout::new_multi_air(counts, num_airs),
160 LayoutKind::Masm => InputLayout::new_masm_multi_air(counts, num_airs),
161 };
162
163 let mut builder = DagBuilder::<QuadFelt>::new();
165 let mut roots = Vec::with_capacity(num_airs);
166 for (air_index, artifacts) in artifacts.iter().enumerate() {
167 roots.push(reemit_air_root(&mut builder, &artifacts.dag, air_index, offsets[air_index]));
168 }
169 let quotient_binding = roots[0].1;
170 if roots.iter().any(|&(_, binding)| binding != quotient_binding) {
171 return Err(AceError::InvalidInputLayout {
172 message: "all AIR quotient bindings must use the same q*v node".into(),
173 });
174 }
175
176 let beta = builder.input(InputKey::MultiAirFoldBeta);
177 let mut ordered = proof_order.iter().map(|&index| roots[index].0);
178 let mut accumulator = ordered.next().expect("multi-AIR composition is nonempty");
179 for next in ordered {
180 let scaled = builder.mul(accumulator, beta);
181 accumulator = builder.add(scaled, next);
182 }
183
184 let root = builder.sub(accumulator, quotient_binding);
186 let mut dag = builder.build(root);
187 dag.compact();
188 let dag = normalize_dag(dag);
189 emit_circuit(&dag, layout)
190}
191
192pub fn build_ace_dag_for_air<A>(
196 air: &A,
197 config: AceConfig,
198) -> Result<AceArtifacts<QuadFelt>, AceError>
199where
200 A: LiftedAir<Felt, QuadFelt>,
201{
202 if config.num_airs == 0 {
203 return Err(AceError::InvalidInputLayout {
204 message: "num_airs must be at least 1".into(),
205 });
206 }
207
208 let periodic_columns = air.periodic_columns();
209 let shared_period = max_period(&periodic_columns);
210 build_ace_dag_for_air_with_periodic_columns(
211 air,
212 config,
213 periodic_columns.into_owned(),
214 shared_period,
215 )
216}
217
218fn build_ace_dags_for_airs<A>(
220 airs: &[A],
221 config: AceConfig,
222) -> Result<Vec<AceArtifacts<QuadFelt>>, AceError>
223where
224 A: LiftedAir<Felt, QuadFelt>,
225{
226 let periodic_columns_by_air: Vec<_> =
227 airs.iter().map(BaseAir::<Felt>::periodic_columns).collect();
228 let shared_period = periodic_columns_by_air
229 .iter()
230 .map(|columns| max_period(columns))
231 .max()
232 .unwrap_or(1);
233
234 airs.iter()
235 .zip(periodic_columns_by_air)
236 .map(|(air, periodic_columns)| {
237 build_ace_dag_for_air_with_periodic_columns(
238 air,
239 config,
240 periodic_columns.into_owned(),
241 shared_period,
242 )
243 })
244 .collect()
245}
246
247fn build_ace_dag_for_air_with_periodic_columns<A>(
248 air: &A,
249 config: AceConfig,
250 periodic_columns: Vec<Vec<Felt>>,
251 shared_period: usize,
252) -> Result<AceArtifacts<QuadFelt>, AceError>
253where
254 A: LiftedAir<Felt, QuadFelt>,
255{
256 let counts = input_counts_for_air(air, config)?;
257 let layout = match (config.layout, config.num_airs >= 2) {
258 (LayoutKind::Native, false) => InputLayout::new(counts),
259 (LayoutKind::Masm, false) => InputLayout::new_masm(counts),
260 (LayoutKind::Native, true) => InputLayout::new_multi_air(counts, config.num_airs),
261 (LayoutKind::Masm, true) => InputLayout::new_masm_multi_air(counts, config.num_airs),
262 };
263 layout.validate();
264
265 let (graph, constraints) = capture(air);
266 let periodic_data = (!periodic_columns.is_empty())
267 .then(|| PeriodicColumnData::from_periodic_columns::<Felt>(periodic_columns));
268 let dag = build_verifier_dag_from_ir(
269 &graph,
270 &constraints,
271 &layout,
272 periodic_data.as_ref(),
273 shared_period,
274 );
275
276 Ok(AceArtifacts { layout, dag })
277}
278
279fn max_period<F>(periodic_columns: &[Vec<F>]) -> usize {
280 periodic_columns.iter().map(Vec::len).max().unwrap_or(1)
281}
282
283#[derive(Clone, Copy, Debug, Default)]
284struct TraceOffsets {
285 preprocessed: usize,
286 main: usize,
287 aux: usize,
288 boundary: usize,
289}
290
291fn reemit_air_root(
292 builder: &mut DagBuilder<QuadFelt>,
293 source: &AceDag<QuadFelt>,
294 air_index: usize,
295 offsets: TraceOffsets,
296) -> (NodeId, NodeId) {
297 debug_assert_eq!(source.root().index() + 1, source.nodes.len());
298 let NodeKind::Sub(accumulator, quotient_binding) = source.nodes[source.root().index()] else {
302 unreachable!("verifier DAGs always emit an accumulator - q*v root")
303 };
304
305 let mut translated = Vec::with_capacity(source.nodes.len() - 1);
306 for node in &source.nodes[..source.root().index()] {
307 let id = match *node {
308 NodeKind::Input(key) => {
309 let key = match key {
310 InputKey::Preprocessed { offset, index } => InputKey::Preprocessed {
311 offset,
312 index: index + offsets.preprocessed,
313 },
314 InputKey::Main { offset, index } => {
315 InputKey::Main { offset, index: index + offsets.main }
316 },
317 InputKey::AuxCoord { offset, index, coord } => InputKey::AuxCoord {
318 offset,
319 index: index + offsets.aux,
320 coord,
321 },
322 InputKey::AuxBusBoundary(index) => {
323 InputKey::AuxBusBoundary(index + offsets.boundary)
324 },
325 InputKey::IsFirst => InputKey::IsFirstAir(air_index),
326 InputKey::IsLast => InputKey::IsLastAir(air_index),
327 InputKey::IsTransition => InputKey::IsTransitionAir(air_index),
328 other => other,
329 };
330 builder.input(key)
331 },
332 NodeKind::Constant(value) => builder.constant(value),
333 NodeKind::Add(a, b) => builder.add(translated[a.index()], translated[b.index()]),
334 NodeKind::Sub(a, b) => builder.sub(translated[a.index()], translated[b.index()]),
335 NodeKind::Mul(a, b) => builder.mul(translated[a.index()], translated[b.index()]),
336 NodeKind::Neg(a) => builder.neg(translated[a.index()]),
337 };
338 translated.push(id);
339 }
340
341 (translated[accumulator.index()], translated[quotient_binding.index()])
342}
343
344fn input_counts_for_air<A>(air: &A, config: AceConfig) -> Result<InputCounts, AceError>
345where
346 A: LiftedAir<Felt, QuadFelt>,
347{
348 if config.num_quotient_chunks == 0 {
349 return Err(AceError::InvalidInputLayout {
350 message: "num_quotient_chunks must be > 0".into(),
351 });
352 }
353 let num_randomness = air.num_randomness();
354 if num_randomness != 2 {
355 return Err(AceError::InvalidInputLayout {
356 message: format!(
357 "AIR must declare exactly 2 randomness challenges (alpha, beta), got {num_randomness}"
358 ),
359 });
360 }
361
362 Ok(InputCounts {
363 preprocessed_width: air.preprocessed_width(),
364 width: air.width(),
365 aux_width: air.aux_width(),
366 num_aux_boundary: air.num_aux_values(),
367 num_public: air.num_public_values(),
368 num_randomness,
369 num_quotient_chunks: config.num_quotient_chunks,
370 })
371}
372
373#[derive(Debug, Clone)]
380pub struct FactoredMultiAirCircuit<EF> {
381 factored: FactoredAceCircuit<EF>,
382 blocks: Vec<TraceOffsets>,
385}
386
387impl<EF: Field> FactoredMultiAirCircuit<EF> {
388 pub fn layout(&self) -> &InputLayout {
390 self.factored.layout()
391 }
392
393 pub fn num_shuffle_ops(&self) -> usize {
395 self.factored.num_shuffle_ops()
396 }
397
398 pub fn num_airs(&self) -> usize {
400 self.blocks.len()
401 }
402
403 pub fn encode_shuffle_section_for_order<'a>(
411 &self,
412 proof_order: &[usize],
413 buffer: &'a mut ShuffleEncodeBuffer,
414 ) -> Result<&'a [Felt], AceError> {
415 let (srcs, exponents) = buffer.order_scratch();
416 self.shuffle_and_exponents(proof_order, srcs, exponents)?;
417 self.factored.encode_shuffle_section(buffer)
418 }
419
420 fn shuffle_and_exponents(
422 &self,
423 proof_order: &[usize],
424 srcs: &mut Vec<usize>,
425 exponents: &mut Vec<usize>,
426 ) -> Result<(), AceError> {
427 let num_airs = self.blocks.len();
428 if proof_order.len() != num_airs {
429 return Err(AceError::InvalidInputLayout {
430 message: format!("proof_order must be a permutation of 0..{num_airs}"),
431 });
432 }
433
434 exponents.clear();
435 exponents.resize(num_airs, usize::MAX);
436 for (position, &air_index) in proof_order.iter().enumerate() {
437 let Some(exponent) = exponents.get_mut(air_index) else {
438 return Err(AceError::InvalidInputLayout {
439 message: format!("proof_order must be a permutation of 0..{num_airs}"),
440 });
441 };
442 if *exponent != usize::MAX {
443 return Err(AceError::InvalidInputLayout {
444 message: format!("proof_order must be a permutation of 0..{num_airs}"),
445 });
446 }
447 *exponent = num_airs - 1 - position;
449 }
450
451 let proof_offsets = accumulate_block_offsets(&self.blocks, proof_order);
452 shuffled_slots(self.factored.layout(), &self.blocks, &proof_offsets, srcs)?;
453 Ok(())
454 }
455
456 pub fn circuit_for_order(&self, proof_order: &[usize]) -> Result<AceCircuit<EF>, AceError> {
461 let mut srcs = Vec::new();
462 let mut coeff_exponents = Vec::new();
463 self.shuffle_and_exponents(proof_order, &mut srcs, &mut coeff_exponents)?;
464 self.factored.assemble(&srcs, &coeff_exponents)
465 }
466}
467
468pub fn build_factored_multi_air_ace_circuit<A>(
476 airs: &[A],
477 config: AceConfig,
478 trace_width_alignment: usize,
479) -> Result<FactoredMultiAirCircuit<QuadFelt>, AceError>
480where
481 A: LiftedAir<Felt, QuadFelt>,
482{
483 let num_airs = airs.len();
484 if num_airs == 0 || config.num_airs != num_airs {
485 return Err(AceError::InvalidInputLayout {
486 message: format!(
487 "multi-AIR composition requires a nonempty airs slice and matching num_airs; got \
488 {} AIRs and num_airs {}",
489 num_airs, config.num_airs
490 ),
491 });
492 }
493 if trace_width_alignment == 0 {
494 return Err(AceError::InvalidInputLayout {
495 message: "trace width alignment must be nonzero".into(),
496 });
497 }
498
499 let sub_config = AceConfig { num_airs: 1, ..config };
500 let artifacts = build_ace_dags_for_airs(airs, sub_config)?;
501 let shared = artifacts[0].layout.counts;
502 if artifacts.iter().any(|air| air.layout.counts.num_public != shared.num_public) {
503 return Err(AceError::InvalidInputLayout {
504 message: "all AIRs must use the same public-value window".into(),
505 });
506 }
507
508 let mut blocks = Vec::with_capacity(num_airs);
509 for artifact in &artifacts {
510 let counts = artifact.layout.counts;
511 let aligned_aux = (counts.aux_width * EXT_DEGREE).next_multiple_of(trace_width_alignment);
512 if !aligned_aux.is_multiple_of(EXT_DEGREE) {
513 return Err(AceError::InvalidInputLayout {
514 message: "aligned auxiliary width must be divisible by the extension degree".into(),
515 });
516 }
517 blocks.push(TraceOffsets {
518 preprocessed: counts.preprocessed_width.next_multiple_of(trace_width_alignment),
519 main: counts.width.next_multiple_of(trace_width_alignment),
520 aux: aligned_aux / EXT_DEGREE,
521 boundary: counts.num_aux_boundary,
522 });
523 }
524
525 let canonical_order: Vec<usize> = (0..num_airs).collect();
526 let offsets = accumulate_block_offsets(&blocks, &canonical_order);
527 let totals = blocks.iter().fold(TraceOffsets::default(), |mut totals, block| {
528 totals.preprocessed += block.preprocessed;
529 totals.main += block.main;
530 totals.aux += block.aux;
531 totals.boundary += block.boundary;
532 totals
533 });
534
535 let counts = InputCounts {
536 preprocessed_width: totals.preprocessed,
537 width: totals.main,
538 aux_width: totals.aux,
539 num_aux_boundary: totals.boundary,
540 num_public: shared.num_public,
541 num_randomness: shared.num_randomness,
542 num_quotient_chunks: shared.num_quotient_chunks,
543 };
544 let layout = match config.layout {
545 LayoutKind::Native => InputLayout::new_multi_air(counts, num_airs),
546 LayoutKind::Masm => InputLayout::new_masm_multi_air(counts, num_airs),
547 };
548
549 let mut builder = DagBuilder::<QuadFelt>::new();
552 let mut roots = Vec::with_capacity(num_airs);
553 for (air_index, artifacts) in artifacts.iter().enumerate() {
554 roots.push(reemit_air_root(&mut builder, &artifacts.dag, air_index, offsets[air_index]));
555 }
556 let quotient_binding = roots[0].1;
557 if roots.iter().any(|&(_, binding)| binding != quotient_binding) {
558 return Err(AceError::InvalidInputLayout {
559 message: "all AIR quotient bindings must use the same q*v node".into(),
560 });
561 }
562
563 let mut accumulator = None;
564 for (air_index, &(acc, _)) in roots.iter().enumerate() {
565 let coeff = builder.input(InputKey::MultiAirFoldCoeff(air_index));
566 let scaled = builder.mul(acc, coeff);
567 accumulator = Some(match accumulator {
568 None => scaled,
569 Some(previous) => builder.add(previous, scaled),
570 });
571 }
572 let accumulator = accumulator.expect("multi-AIR composition is nonempty");
573
574 let root = builder.sub(accumulator, quotient_binding);
576 let mut dag = builder.build(root);
577 dag.compact();
578 let dag = normalize_dag(dag);
579
580 let mut shuffle_dsts = Vec::new();
581 shuffled_slots(&layout, &blocks, &offsets, &mut shuffle_dsts)?;
582 let factored = emit_factored_circuit(&dag, layout, shuffle_dsts, num_airs)?;
583 Ok(FactoredMultiAirCircuit { factored, blocks })
584}
585
586fn accumulate_block_offsets(blocks: &[TraceOffsets], order: &[usize]) -> Vec<TraceOffsets> {
588 let mut offsets = vec![TraceOffsets::default(); blocks.len()];
589 let mut totals = TraceOffsets::default();
590 for &air_index in order {
591 offsets[air_index] = totals;
592 let block = blocks[air_index];
593 totals.preprocessed += block.preprocessed;
594 totals.main += block.main;
595 totals.aux += block.aux;
596 totals.boundary += block.boundary;
597 }
598 offsets
599}
600
601fn shuffled_slots(
607 layout: &InputLayout,
608 blocks: &[TraceOffsets],
609 offsets: &[TraceOffsets],
610 slots: &mut Vec<usize>,
611) -> Result<(), AceError> {
612 slots.clear();
613 let mut push = |key: InputKey| -> Result<(), AceError> {
614 let index = layout.index(key).ok_or_else(|| AceError::InvalidInputLayout {
615 message: format!("shuffled slot {key:?} is missing from the layout"),
616 })?;
617 slots.push(index);
618 Ok(())
619 };
620
621 for row_offset in 0..2 {
622 for (block, air_offsets) in blocks.iter().zip(offsets) {
623 for column in 0..block.preprocessed {
624 push(InputKey::Preprocessed {
625 offset: row_offset,
626 index: air_offsets.preprocessed + column,
627 })?;
628 }
629 }
630 }
631 for row_offset in 0..2 {
632 for (block, air_offsets) in blocks.iter().zip(offsets) {
633 for column in 0..block.main {
634 push(InputKey::Main {
635 offset: row_offset,
636 index: air_offsets.main + column,
637 })?;
638 }
639 }
640 }
641 for row_offset in 0..2 {
642 for (block, air_offsets) in blocks.iter().zip(offsets) {
643 for column in 0..block.aux {
644 for coord in 0..EXT_DEGREE {
645 push(InputKey::AuxCoord {
646 offset: row_offset,
647 index: air_offsets.aux + column,
648 coord,
649 })?;
650 }
651 }
652 }
653 }
654 for (block, air_offsets) in blocks.iter().zip(offsets) {
655 for value in 0..block.boundary {
656 push(InputKey::AuxBusBoundary(air_offsets.boundary + value))?;
657 }
658 }
659
660 Ok(())
661}