use miden_constraint_compiler::ir::capture;
use miden_core::{Felt, field::QuadFelt};
use miden_crypto::{
field::Field,
stark::air::{BaseAir, LiftedAir},
};
use crate::{
AceError, EXT_DEGREE,
circuit::{AceCircuit, emit_circuit},
dag::{
AceDag, DagBuilder, NodeId, NodeKind, PeriodicColumnData, build_verifier_dag_from_ir,
normalize_dag,
},
factored::{FactoredAceCircuit, ShuffleEncodeBuffer, emit_factored_circuit},
layout::{InputCounts, InputKey, InputLayout},
};
#[derive(Debug, Clone, Copy)]
pub enum LayoutKind {
Native,
Masm,
}
#[derive(Debug, Clone, Copy)]
pub struct AceConfig {
pub num_quotient_chunks: usize,
pub layout: LayoutKind,
pub num_airs: usize,
}
#[derive(Debug)]
pub struct AceArtifacts<EF> {
pub layout: InputLayout,
pub dag: AceDag<EF>,
}
pub fn build_ace_circuit_for_air<A>(
air: &A,
config: AceConfig,
) -> Result<AceCircuit<QuadFelt>, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
let artifacts = build_ace_dag_for_air(air, config)?;
emit_circuit(&artifacts.dag, artifacts.layout)
}
pub fn build_multi_air_ace_circuit<A>(
airs: &[A],
proof_order: &[usize],
config: AceConfig,
trace_width_alignment: usize,
) -> Result<AceCircuit<QuadFelt>, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
let num_airs = airs.len();
if num_airs == 0 || config.num_airs != num_airs {
return Err(AceError::InvalidInputLayout {
message: format!(
"multi-AIR composition requires a nonempty airs slice and matching num_airs; got \
{} AIRs and num_airs {}",
num_airs, config.num_airs
),
});
}
let mut seen = vec![false; num_airs];
if proof_order.len() != num_airs
|| proof_order
.iter()
.any(|&index| index >= num_airs || core::mem::replace(&mut seen[index], true))
{
return Err(AceError::InvalidInputLayout {
message: format!("proof_order must be a permutation of 0..{num_airs}"),
});
}
if trace_width_alignment == 0 {
return Err(AceError::InvalidInputLayout {
message: "trace width alignment must be nonzero".into(),
});
}
let sub_config = AceConfig { num_airs: 1, ..config };
let artifacts = build_ace_dags_for_airs(airs, sub_config)?;
let shared = artifacts[0].layout.counts;
if artifacts.iter().any(|air| air.layout.counts.num_public != shared.num_public) {
return Err(AceError::InvalidInputLayout {
message: "all AIRs must use the same public-value window".into(),
});
}
let mut offsets = vec![TraceOffsets::default(); num_airs];
let mut totals = TraceOffsets::default();
for &air_index in proof_order {
offsets[air_index] = totals;
let counts = artifacts[air_index].layout.counts;
totals.preprocessed += counts.preprocessed_width.next_multiple_of(trace_width_alignment);
totals.main += counts.width.next_multiple_of(trace_width_alignment);
let aligned_aux = (counts.aux_width * EXT_DEGREE).next_multiple_of(trace_width_alignment);
if !aligned_aux.is_multiple_of(EXT_DEGREE) {
return Err(AceError::InvalidInputLayout {
message: "aligned auxiliary width must be divisible by the extension degree".into(),
});
}
totals.aux += aligned_aux / EXT_DEGREE;
totals.boundary += counts.num_aux_boundary;
}
let counts = InputCounts {
preprocessed_width: totals.preprocessed,
width: totals.main,
aux_width: totals.aux,
num_aux_boundary: totals.boundary,
num_public: shared.num_public,
num_randomness: shared.num_randomness,
num_quotient_chunks: shared.num_quotient_chunks,
};
let layout = match config.layout {
LayoutKind::Native => InputLayout::new_multi_air(counts, num_airs),
LayoutKind::Masm => InputLayout::new_masm_multi_air(counts, num_airs),
};
let mut builder = DagBuilder::<QuadFelt>::new();
let mut roots = Vec::with_capacity(num_airs);
for (air_index, artifacts) in artifacts.iter().enumerate() {
roots.push(reemit_air_root(&mut builder, &artifacts.dag, air_index, offsets[air_index]));
}
let quotient_binding = roots[0].1;
if roots.iter().any(|&(_, binding)| binding != quotient_binding) {
return Err(AceError::InvalidInputLayout {
message: "all AIR quotient bindings must use the same q*v node".into(),
});
}
let beta = builder.input(InputKey::MultiAirFoldBeta);
let mut ordered = proof_order.iter().map(|&index| roots[index].0);
let mut accumulator = ordered.next().expect("multi-AIR composition is nonempty");
for next in ordered {
let scaled = builder.mul(accumulator, beta);
accumulator = builder.add(scaled, next);
}
let root = builder.sub(accumulator, quotient_binding);
let mut dag = builder.build(root);
dag.compact();
let dag = normalize_dag(dag);
emit_circuit(&dag, layout)
}
pub fn build_ace_dag_for_air<A>(
air: &A,
config: AceConfig,
) -> Result<AceArtifacts<QuadFelt>, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
if config.num_airs == 0 {
return Err(AceError::InvalidInputLayout {
message: "num_airs must be at least 1".into(),
});
}
let periodic_columns = air.periodic_columns();
let shared_period = max_period(&periodic_columns);
build_ace_dag_for_air_with_periodic_columns(air, config, periodic_columns, shared_period)
}
fn build_ace_dags_for_airs<A>(
airs: &[A],
config: AceConfig,
) -> Result<Vec<AceArtifacts<QuadFelt>>, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
let periodic_columns_by_air: Vec<_> =
airs.iter().map(BaseAir::<Felt>::periodic_columns).collect();
let shared_period = periodic_columns_by_air
.iter()
.map(|columns| max_period(columns))
.max()
.unwrap_or(1);
airs.iter()
.zip(periodic_columns_by_air)
.map(|(air, periodic_columns)| {
build_ace_dag_for_air_with_periodic_columns(
air,
config,
periodic_columns,
shared_period,
)
})
.collect()
}
fn build_ace_dag_for_air_with_periodic_columns<A>(
air: &A,
config: AceConfig,
periodic_columns: Vec<Vec<Felt>>,
shared_period: usize,
) -> Result<AceArtifacts<QuadFelt>, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
let counts = input_counts_for_air(air, config)?;
let layout = match (config.layout, config.num_airs >= 2) {
(LayoutKind::Native, false) => InputLayout::new(counts),
(LayoutKind::Masm, false) => InputLayout::new_masm(counts),
(LayoutKind::Native, true) => InputLayout::new_multi_air(counts, config.num_airs),
(LayoutKind::Masm, true) => InputLayout::new_masm_multi_air(counts, config.num_airs),
};
layout.validate();
let (graph, constraints) = capture(air);
let periodic_data = (!periodic_columns.is_empty())
.then(|| PeriodicColumnData::from_periodic_columns::<Felt>(periodic_columns));
let dag = build_verifier_dag_from_ir(
&graph,
&constraints,
&layout,
periodic_data.as_ref(),
shared_period,
);
Ok(AceArtifacts { layout, dag })
}
fn max_period<F>(periodic_columns: &[Vec<F>]) -> usize {
periodic_columns.iter().map(Vec::len).max().unwrap_or(1)
}
#[derive(Clone, Copy, Debug, Default)]
struct TraceOffsets {
preprocessed: usize,
main: usize,
aux: usize,
boundary: usize,
}
fn reemit_air_root(
builder: &mut DagBuilder<QuadFelt>,
source: &AceDag<QuadFelt>,
air_index: usize,
offsets: TraceOffsets,
) -> (NodeId, NodeId) {
debug_assert_eq!(source.root().index() + 1, source.nodes.len());
let NodeKind::Sub(accumulator, quotient_binding) = source.nodes[source.root().index()] else {
unreachable!("verifier DAGs always emit an accumulator - q*v root")
};
let mut translated = Vec::with_capacity(source.nodes.len() - 1);
for node in &source.nodes[..source.root().index()] {
let id = match *node {
NodeKind::Input(key) => {
let key = match key {
InputKey::Preprocessed { offset, index } => InputKey::Preprocessed {
offset,
index: index + offsets.preprocessed,
},
InputKey::Main { offset, index } => {
InputKey::Main { offset, index: index + offsets.main }
},
InputKey::AuxCoord { offset, index, coord } => InputKey::AuxCoord {
offset,
index: index + offsets.aux,
coord,
},
InputKey::AuxBusBoundary(index) => {
InputKey::AuxBusBoundary(index + offsets.boundary)
},
InputKey::IsFirst => InputKey::IsFirstAir(air_index),
InputKey::IsLast => InputKey::IsLastAir(air_index),
InputKey::IsTransition => InputKey::IsTransitionAir(air_index),
other => other,
};
builder.input(key)
},
NodeKind::Constant(value) => builder.constant(value),
NodeKind::Add(a, b) => builder.add(translated[a.index()], translated[b.index()]),
NodeKind::Sub(a, b) => builder.sub(translated[a.index()], translated[b.index()]),
NodeKind::Mul(a, b) => builder.mul(translated[a.index()], translated[b.index()]),
NodeKind::Neg(a) => builder.neg(translated[a.index()]),
};
translated.push(id);
}
(translated[accumulator.index()], translated[quotient_binding.index()])
}
fn input_counts_for_air<A>(air: &A, config: AceConfig) -> Result<InputCounts, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
if config.num_quotient_chunks == 0 {
return Err(AceError::InvalidInputLayout {
message: "num_quotient_chunks must be > 0".into(),
});
}
let num_randomness = air.num_randomness();
if num_randomness != 2 {
return Err(AceError::InvalidInputLayout {
message: format!(
"AIR must declare exactly 2 randomness challenges (alpha, beta), got {num_randomness}"
),
});
}
Ok(InputCounts {
preprocessed_width: air.preprocessed_width(),
width: air.width(),
aux_width: air.aux_width(),
num_aux_boundary: air.num_aux_values(),
num_public: air.num_public_values(),
num_randomness,
num_quotient_chunks: config.num_quotient_chunks,
})
}
#[derive(Debug, Clone)]
pub struct FactoredMultiAirCircuit<EF> {
factored: FactoredAceCircuit<EF>,
blocks: Vec<TraceOffsets>,
}
impl<EF: Field> FactoredMultiAirCircuit<EF> {
pub fn layout(&self) -> &InputLayout {
self.factored.layout()
}
pub fn num_shuffle_ops(&self) -> usize {
self.factored.num_shuffle_ops()
}
pub fn num_airs(&self) -> usize {
self.blocks.len()
}
pub fn encode_shuffle_section_for_order<'a>(
&self,
proof_order: &[usize],
buffer: &'a mut ShuffleEncodeBuffer,
) -> Result<&'a [Felt], AceError> {
let (srcs, exponents) = buffer.order_scratch();
self.shuffle_and_exponents(proof_order, srcs, exponents)?;
self.factored.encode_shuffle_section(buffer)
}
fn shuffle_and_exponents(
&self,
proof_order: &[usize],
srcs: &mut Vec<usize>,
exponents: &mut Vec<usize>,
) -> Result<(), AceError> {
let num_airs = self.blocks.len();
if proof_order.len() != num_airs {
return Err(AceError::InvalidInputLayout {
message: format!("proof_order must be a permutation of 0..{num_airs}"),
});
}
exponents.clear();
exponents.resize(num_airs, usize::MAX);
for (position, &air_index) in proof_order.iter().enumerate() {
let Some(exponent) = exponents.get_mut(air_index) else {
return Err(AceError::InvalidInputLayout {
message: format!("proof_order must be a permutation of 0..{num_airs}"),
});
};
if *exponent != usize::MAX {
return Err(AceError::InvalidInputLayout {
message: format!("proof_order must be a permutation of 0..{num_airs}"),
});
}
*exponent = num_airs - 1 - position;
}
let proof_offsets = accumulate_block_offsets(&self.blocks, proof_order);
shuffled_slots(self.factored.layout(), &self.blocks, &proof_offsets, srcs)?;
Ok(())
}
pub fn circuit_for_order(&self, proof_order: &[usize]) -> Result<AceCircuit<EF>, AceError> {
let mut srcs = Vec::new();
let mut coeff_exponents = Vec::new();
self.shuffle_and_exponents(proof_order, &mut srcs, &mut coeff_exponents)?;
self.factored.assemble(&srcs, &coeff_exponents)
}
}
pub fn build_factored_multi_air_ace_circuit<A>(
airs: &[A],
config: AceConfig,
trace_width_alignment: usize,
) -> Result<FactoredMultiAirCircuit<QuadFelt>, AceError>
where
A: LiftedAir<Felt, QuadFelt>,
{
let num_airs = airs.len();
if num_airs == 0 || config.num_airs != num_airs {
return Err(AceError::InvalidInputLayout {
message: format!(
"multi-AIR composition requires a nonempty airs slice and matching num_airs; got \
{} AIRs and num_airs {}",
num_airs, config.num_airs
),
});
}
if trace_width_alignment == 0 {
return Err(AceError::InvalidInputLayout {
message: "trace width alignment must be nonzero".into(),
});
}
let sub_config = AceConfig { num_airs: 1, ..config };
let artifacts = build_ace_dags_for_airs(airs, sub_config)?;
let shared = artifacts[0].layout.counts;
if artifacts.iter().any(|air| air.layout.counts.num_public != shared.num_public) {
return Err(AceError::InvalidInputLayout {
message: "all AIRs must use the same public-value window".into(),
});
}
let mut blocks = Vec::with_capacity(num_airs);
for artifact in &artifacts {
let counts = artifact.layout.counts;
let aligned_aux = (counts.aux_width * EXT_DEGREE).next_multiple_of(trace_width_alignment);
if !aligned_aux.is_multiple_of(EXT_DEGREE) {
return Err(AceError::InvalidInputLayout {
message: "aligned auxiliary width must be divisible by the extension degree".into(),
});
}
blocks.push(TraceOffsets {
preprocessed: counts.preprocessed_width.next_multiple_of(trace_width_alignment),
main: counts.width.next_multiple_of(trace_width_alignment),
aux: aligned_aux / EXT_DEGREE,
boundary: counts.num_aux_boundary,
});
}
let canonical_order: Vec<usize> = (0..num_airs).collect();
let offsets = accumulate_block_offsets(&blocks, &canonical_order);
let totals = blocks.iter().fold(TraceOffsets::default(), |mut totals, block| {
totals.preprocessed += block.preprocessed;
totals.main += block.main;
totals.aux += block.aux;
totals.boundary += block.boundary;
totals
});
let counts = InputCounts {
preprocessed_width: totals.preprocessed,
width: totals.main,
aux_width: totals.aux,
num_aux_boundary: totals.boundary,
num_public: shared.num_public,
num_randomness: shared.num_randomness,
num_quotient_chunks: shared.num_quotient_chunks,
};
let layout = match config.layout {
LayoutKind::Native => InputLayout::new_multi_air(counts, num_airs),
LayoutKind::Masm => InputLayout::new_masm_multi_air(counts, num_airs),
};
let mut builder = DagBuilder::<QuadFelt>::new();
let mut roots = Vec::with_capacity(num_airs);
for (air_index, artifacts) in artifacts.iter().enumerate() {
roots.push(reemit_air_root(&mut builder, &artifacts.dag, air_index, offsets[air_index]));
}
let quotient_binding = roots[0].1;
if roots.iter().any(|&(_, binding)| binding != quotient_binding) {
return Err(AceError::InvalidInputLayout {
message: "all AIR quotient bindings must use the same q*v node".into(),
});
}
let mut accumulator = None;
for (air_index, &(acc, _)) in roots.iter().enumerate() {
let coeff = builder.input(InputKey::MultiAirFoldCoeff(air_index));
let scaled = builder.mul(acc, coeff);
accumulator = Some(match accumulator {
None => scaled,
Some(previous) => builder.add(previous, scaled),
});
}
let accumulator = accumulator.expect("multi-AIR composition is nonempty");
let root = builder.sub(accumulator, quotient_binding);
let mut dag = builder.build(root);
dag.compact();
let dag = normalize_dag(dag);
let mut shuffle_dsts = Vec::new();
shuffled_slots(&layout, &blocks, &offsets, &mut shuffle_dsts)?;
let factored = emit_factored_circuit(&dag, layout, shuffle_dsts, num_airs)?;
Ok(FactoredMultiAirCircuit { factored, blocks })
}
fn accumulate_block_offsets(blocks: &[TraceOffsets], order: &[usize]) -> Vec<TraceOffsets> {
let mut offsets = vec![TraceOffsets::default(); blocks.len()];
let mut totals = TraceOffsets::default();
for &air_index in order {
offsets[air_index] = totals;
let block = blocks[air_index];
totals.preprocessed += block.preprocessed;
totals.main += block.main;
totals.aux += block.aux;
totals.boundary += block.boundary;
}
offsets
}
fn shuffled_slots(
layout: &InputLayout,
blocks: &[TraceOffsets],
offsets: &[TraceOffsets],
slots: &mut Vec<usize>,
) -> Result<(), AceError> {
slots.clear();
let mut push = |key: InputKey| -> Result<(), AceError> {
let index = layout.index(key).ok_or_else(|| AceError::InvalidInputLayout {
message: format!("shuffled slot {key:?} is missing from the layout"),
})?;
slots.push(index);
Ok(())
};
for row_offset in 0..2 {
for (block, air_offsets) in blocks.iter().zip(offsets) {
for column in 0..block.preprocessed {
push(InputKey::Preprocessed {
offset: row_offset,
index: air_offsets.preprocessed + column,
})?;
}
}
}
for row_offset in 0..2 {
for (block, air_offsets) in blocks.iter().zip(offsets) {
for column in 0..block.main {
push(InputKey::Main {
offset: row_offset,
index: air_offsets.main + column,
})?;
}
}
}
for row_offset in 0..2 {
for (block, air_offsets) in blocks.iter().zip(offsets) {
for column in 0..block.aux {
for coord in 0..EXT_DEGREE {
push(InputKey::AuxCoord {
offset: row_offset,
index: air_offsets.aux + column,
coord,
})?;
}
}
}
}
for (block, air_offsets) in blocks.iter().zip(offsets) {
for value in 0..block.boundary {
push(InputKey::AuxBusBoundary(air_offsets.boundary + value))?;
}
}
Ok(())
}