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(air, config, periodic_columns, shared_period)
211}
212
213fn build_ace_dags_for_airs<A>(
215 airs: &[A],
216 config: AceConfig,
217) -> Result<Vec<AceArtifacts<QuadFelt>>, AceError>
218where
219 A: LiftedAir<Felt, QuadFelt>,
220{
221 let periodic_columns_by_air: Vec<_> =
222 airs.iter().map(BaseAir::<Felt>::periodic_columns).collect();
223 let shared_period = periodic_columns_by_air
224 .iter()
225 .map(|columns| max_period(columns))
226 .max()
227 .unwrap_or(1);
228
229 airs.iter()
230 .zip(periodic_columns_by_air)
231 .map(|(air, periodic_columns)| {
232 build_ace_dag_for_air_with_periodic_columns(
233 air,
234 config,
235 periodic_columns,
236 shared_period,
237 )
238 })
239 .collect()
240}
241
242fn build_ace_dag_for_air_with_periodic_columns<A>(
243 air: &A,
244 config: AceConfig,
245 periodic_columns: Vec<Vec<Felt>>,
246 shared_period: usize,
247) -> Result<AceArtifacts<QuadFelt>, AceError>
248where
249 A: LiftedAir<Felt, QuadFelt>,
250{
251 let counts = input_counts_for_air(air, config)?;
252 let layout = match (config.layout, config.num_airs >= 2) {
253 (LayoutKind::Native, false) => InputLayout::new(counts),
254 (LayoutKind::Masm, false) => InputLayout::new_masm(counts),
255 (LayoutKind::Native, true) => InputLayout::new_multi_air(counts, config.num_airs),
256 (LayoutKind::Masm, true) => InputLayout::new_masm_multi_air(counts, config.num_airs),
257 };
258 layout.validate();
259
260 let (graph, constraints) = capture(air);
261 let periodic_data = (!periodic_columns.is_empty())
262 .then(|| PeriodicColumnData::from_periodic_columns::<Felt>(periodic_columns));
263 let dag = build_verifier_dag_from_ir(
264 &graph,
265 &constraints,
266 &layout,
267 periodic_data.as_ref(),
268 shared_period,
269 );
270
271 Ok(AceArtifacts { layout, dag })
272}
273
274fn max_period<F>(periodic_columns: &[Vec<F>]) -> usize {
275 periodic_columns.iter().map(Vec::len).max().unwrap_or(1)
276}
277
278#[derive(Clone, Copy, Debug, Default)]
279struct TraceOffsets {
280 preprocessed: usize,
281 main: usize,
282 aux: usize,
283 boundary: usize,
284}
285
286fn reemit_air_root(
287 builder: &mut DagBuilder<QuadFelt>,
288 source: &AceDag<QuadFelt>,
289 air_index: usize,
290 offsets: TraceOffsets,
291) -> (NodeId, NodeId) {
292 debug_assert_eq!(source.root().index() + 1, source.nodes.len());
293 let NodeKind::Sub(accumulator, quotient_binding) = source.nodes[source.root().index()] else {
297 unreachable!("verifier DAGs always emit an accumulator - q*v root")
298 };
299
300 let mut translated = Vec::with_capacity(source.nodes.len() - 1);
301 for node in &source.nodes[..source.root().index()] {
302 let id = match *node {
303 NodeKind::Input(key) => {
304 let key = match key {
305 InputKey::Preprocessed { offset, index } => InputKey::Preprocessed {
306 offset,
307 index: index + offsets.preprocessed,
308 },
309 InputKey::Main { offset, index } => {
310 InputKey::Main { offset, index: index + offsets.main }
311 },
312 InputKey::AuxCoord { offset, index, coord } => InputKey::AuxCoord {
313 offset,
314 index: index + offsets.aux,
315 coord,
316 },
317 InputKey::AuxBusBoundary(index) => {
318 InputKey::AuxBusBoundary(index + offsets.boundary)
319 },
320 InputKey::IsFirst => InputKey::IsFirstAir(air_index),
321 InputKey::IsLast => InputKey::IsLastAir(air_index),
322 InputKey::IsTransition => InputKey::IsTransitionAir(air_index),
323 other => other,
324 };
325 builder.input(key)
326 },
327 NodeKind::Constant(value) => builder.constant(value),
328 NodeKind::Add(a, b) => builder.add(translated[a.index()], translated[b.index()]),
329 NodeKind::Sub(a, b) => builder.sub(translated[a.index()], translated[b.index()]),
330 NodeKind::Mul(a, b) => builder.mul(translated[a.index()], translated[b.index()]),
331 NodeKind::Neg(a) => builder.neg(translated[a.index()]),
332 };
333 translated.push(id);
334 }
335
336 (translated[accumulator.index()], translated[quotient_binding.index()])
337}
338
339fn input_counts_for_air<A>(air: &A, config: AceConfig) -> Result<InputCounts, AceError>
340where
341 A: LiftedAir<Felt, QuadFelt>,
342{
343 if config.num_quotient_chunks == 0 {
344 return Err(AceError::InvalidInputLayout {
345 message: "num_quotient_chunks must be > 0".into(),
346 });
347 }
348 let num_randomness = air.num_randomness();
349 if num_randomness != 2 {
350 return Err(AceError::InvalidInputLayout {
351 message: format!(
352 "AIR must declare exactly 2 randomness challenges (alpha, beta), got {num_randomness}"
353 ),
354 });
355 }
356
357 Ok(InputCounts {
358 preprocessed_width: air.preprocessed_width(),
359 width: air.width(),
360 aux_width: air.aux_width(),
361 num_aux_boundary: air.num_aux_values(),
362 num_public: air.num_public_values(),
363 num_randomness,
364 num_quotient_chunks: config.num_quotient_chunks,
365 })
366}
367
368#[derive(Debug, Clone)]
375pub struct FactoredMultiAirCircuit<EF> {
376 factored: FactoredAceCircuit<EF>,
377 blocks: Vec<TraceOffsets>,
380}
381
382impl<EF: Field> FactoredMultiAirCircuit<EF> {
383 pub fn layout(&self) -> &InputLayout {
385 self.factored.layout()
386 }
387
388 pub fn num_shuffle_ops(&self) -> usize {
390 self.factored.num_shuffle_ops()
391 }
392
393 pub fn num_airs(&self) -> usize {
395 self.blocks.len()
396 }
397
398 pub fn encode_shuffle_section_for_order<'a>(
406 &self,
407 proof_order: &[usize],
408 buffer: &'a mut ShuffleEncodeBuffer,
409 ) -> Result<&'a [Felt], AceError> {
410 let (srcs, exponents) = buffer.order_scratch();
411 self.shuffle_and_exponents(proof_order, srcs, exponents)?;
412 self.factored.encode_shuffle_section(buffer)
413 }
414
415 fn shuffle_and_exponents(
417 &self,
418 proof_order: &[usize],
419 srcs: &mut Vec<usize>,
420 exponents: &mut Vec<usize>,
421 ) -> Result<(), AceError> {
422 let num_airs = self.blocks.len();
423 if proof_order.len() != num_airs {
424 return Err(AceError::InvalidInputLayout {
425 message: format!("proof_order must be a permutation of 0..{num_airs}"),
426 });
427 }
428
429 exponents.clear();
430 exponents.resize(num_airs, usize::MAX);
431 for (position, &air_index) in proof_order.iter().enumerate() {
432 let Some(exponent) = exponents.get_mut(air_index) else {
433 return Err(AceError::InvalidInputLayout {
434 message: format!("proof_order must be a permutation of 0..{num_airs}"),
435 });
436 };
437 if *exponent != usize::MAX {
438 return Err(AceError::InvalidInputLayout {
439 message: format!("proof_order must be a permutation of 0..{num_airs}"),
440 });
441 }
442 *exponent = num_airs - 1 - position;
444 }
445
446 let proof_offsets = accumulate_block_offsets(&self.blocks, proof_order);
447 shuffled_slots(self.factored.layout(), &self.blocks, &proof_offsets, srcs)?;
448 Ok(())
449 }
450
451 pub fn circuit_for_order(&self, proof_order: &[usize]) -> Result<AceCircuit<EF>, AceError> {
456 let mut srcs = Vec::new();
457 let mut coeff_exponents = Vec::new();
458 self.shuffle_and_exponents(proof_order, &mut srcs, &mut coeff_exponents)?;
459 self.factored.assemble(&srcs, &coeff_exponents)
460 }
461}
462
463pub fn build_factored_multi_air_ace_circuit<A>(
471 airs: &[A],
472 config: AceConfig,
473 trace_width_alignment: usize,
474) -> Result<FactoredMultiAirCircuit<QuadFelt>, AceError>
475where
476 A: LiftedAir<Felt, QuadFelt>,
477{
478 let num_airs = airs.len();
479 if num_airs == 0 || config.num_airs != num_airs {
480 return Err(AceError::InvalidInputLayout {
481 message: format!(
482 "multi-AIR composition requires a nonempty airs slice and matching num_airs; got \
483 {} AIRs and num_airs {}",
484 num_airs, config.num_airs
485 ),
486 });
487 }
488 if trace_width_alignment == 0 {
489 return Err(AceError::InvalidInputLayout {
490 message: "trace width alignment must be nonzero".into(),
491 });
492 }
493
494 let sub_config = AceConfig { num_airs: 1, ..config };
495 let artifacts = build_ace_dags_for_airs(airs, sub_config)?;
496 let shared = artifacts[0].layout.counts;
497 if artifacts.iter().any(|air| air.layout.counts.num_public != shared.num_public) {
498 return Err(AceError::InvalidInputLayout {
499 message: "all AIRs must use the same public-value window".into(),
500 });
501 }
502
503 let mut blocks = Vec::with_capacity(num_airs);
504 for artifact in &artifacts {
505 let counts = artifact.layout.counts;
506 let aligned_aux = (counts.aux_width * EXT_DEGREE).next_multiple_of(trace_width_alignment);
507 if !aligned_aux.is_multiple_of(EXT_DEGREE) {
508 return Err(AceError::InvalidInputLayout {
509 message: "aligned auxiliary width must be divisible by the extension degree".into(),
510 });
511 }
512 blocks.push(TraceOffsets {
513 preprocessed: counts.preprocessed_width.next_multiple_of(trace_width_alignment),
514 main: counts.width.next_multiple_of(trace_width_alignment),
515 aux: aligned_aux / EXT_DEGREE,
516 boundary: counts.num_aux_boundary,
517 });
518 }
519
520 let canonical_order: Vec<usize> = (0..num_airs).collect();
521 let offsets = accumulate_block_offsets(&blocks, &canonical_order);
522 let totals = blocks.iter().fold(TraceOffsets::default(), |mut totals, block| {
523 totals.preprocessed += block.preprocessed;
524 totals.main += block.main;
525 totals.aux += block.aux;
526 totals.boundary += block.boundary;
527 totals
528 });
529
530 let counts = InputCounts {
531 preprocessed_width: totals.preprocessed,
532 width: totals.main,
533 aux_width: totals.aux,
534 num_aux_boundary: totals.boundary,
535 num_public: shared.num_public,
536 num_randomness: shared.num_randomness,
537 num_quotient_chunks: shared.num_quotient_chunks,
538 };
539 let layout = match config.layout {
540 LayoutKind::Native => InputLayout::new_multi_air(counts, num_airs),
541 LayoutKind::Masm => InputLayout::new_masm_multi_air(counts, num_airs),
542 };
543
544 let mut builder = DagBuilder::<QuadFelt>::new();
547 let mut roots = Vec::with_capacity(num_airs);
548 for (air_index, artifacts) in artifacts.iter().enumerate() {
549 roots.push(reemit_air_root(&mut builder, &artifacts.dag, air_index, offsets[air_index]));
550 }
551 let quotient_binding = roots[0].1;
552 if roots.iter().any(|&(_, binding)| binding != quotient_binding) {
553 return Err(AceError::InvalidInputLayout {
554 message: "all AIR quotient bindings must use the same q*v node".into(),
555 });
556 }
557
558 let mut accumulator = None;
559 for (air_index, &(acc, _)) in roots.iter().enumerate() {
560 let coeff = builder.input(InputKey::MultiAirFoldCoeff(air_index));
561 let scaled = builder.mul(acc, coeff);
562 accumulator = Some(match accumulator {
563 None => scaled,
564 Some(previous) => builder.add(previous, scaled),
565 });
566 }
567 let accumulator = accumulator.expect("multi-AIR composition is nonempty");
568
569 let root = builder.sub(accumulator, quotient_binding);
571 let mut dag = builder.build(root);
572 dag.compact();
573 let dag = normalize_dag(dag);
574
575 let mut shuffle_dsts = Vec::new();
576 shuffled_slots(&layout, &blocks, &offsets, &mut shuffle_dsts)?;
577 let factored = emit_factored_circuit(&dag, layout, shuffle_dsts, num_airs)?;
578 Ok(FactoredMultiAirCircuit { factored, blocks })
579}
580
581fn accumulate_block_offsets(blocks: &[TraceOffsets], order: &[usize]) -> Vec<TraceOffsets> {
583 let mut offsets = vec![TraceOffsets::default(); blocks.len()];
584 let mut totals = TraceOffsets::default();
585 for &air_index in order {
586 offsets[air_index] = totals;
587 let block = blocks[air_index];
588 totals.preprocessed += block.preprocessed;
589 totals.main += block.main;
590 totals.aux += block.aux;
591 totals.boundary += block.boundary;
592 }
593 offsets
594}
595
596fn shuffled_slots(
602 layout: &InputLayout,
603 blocks: &[TraceOffsets],
604 offsets: &[TraceOffsets],
605 slots: &mut Vec<usize>,
606) -> Result<(), AceError> {
607 slots.clear();
608 let mut push = |key: InputKey| -> Result<(), AceError> {
609 let index = layout.index(key).ok_or_else(|| AceError::InvalidInputLayout {
610 message: format!("shuffled slot {key:?} is missing from the layout"),
611 })?;
612 slots.push(index);
613 Ok(())
614 };
615
616 for row_offset in 0..2 {
617 for (block, air_offsets) in blocks.iter().zip(offsets) {
618 for column in 0..block.preprocessed {
619 push(InputKey::Preprocessed {
620 offset: row_offset,
621 index: air_offsets.preprocessed + column,
622 })?;
623 }
624 }
625 }
626 for row_offset in 0..2 {
627 for (block, air_offsets) in blocks.iter().zip(offsets) {
628 for column in 0..block.main {
629 push(InputKey::Main {
630 offset: row_offset,
631 index: air_offsets.main + column,
632 })?;
633 }
634 }
635 }
636 for row_offset in 0..2 {
637 for (block, air_offsets) in blocks.iter().zip(offsets) {
638 for column in 0..block.aux {
639 for coord in 0..EXT_DEGREE {
640 push(InputKey::AuxCoord {
641 offset: row_offset,
642 index: air_offsets.aux + column,
643 coord,
644 })?;
645 }
646 }
647 }
648 }
649 for (block, air_offsets) in blocks.iter().zip(offsets) {
650 for value in 0..block.boundary {
651 push(InputKey::AuxBusBoundary(air_offsets.boundary + value))?;
652 }
653 }
654
655 Ok(())
656}