use std::collections::HashMap;
use miden_core::Felt;
use miden_crypto::field::Field;
use crate::{
AceError, EXT_DEGREE, InputLayout,
circuit::{AceCircuit, AceNode, AceOp, AceOpNode},
dag::{AceDag, NodeKind},
encode::{ADV_PIPE_BLOCK_FELTS, CONST_EF_ALIGN, StreamGeometry},
layout::InputKey,
};
const CONST_EF_BLOCK_ALIGN: usize = ADV_PIPE_BLOCK_FELTS / EXT_DEGREE;
const _: () = assert!(
CONST_EF_BLOCK_ALIGN.is_multiple_of(CONST_EF_ALIGN),
"constant block alignment must refine the encoder's READ-row alignment, or the two-segment split drifts off a block boundary"
);
const CONST_ZERO: usize = 0;
const CONST_ONE: usize = 1;
#[derive(Clone, Debug, Default)]
pub struct ShuffleEncodeBuffer {
srcs: Vec<usize>,
exponents: Vec<usize>,
seen_srcs: Vec<bool>,
seen_exponents: Vec<bool>,
ops: Vec<AceOpNode>,
felts: Vec<Felt>,
}
impl ShuffleEncodeBuffer {
pub fn new() -> Self {
Self::default()
}
pub(crate) fn order_scratch(&mut self) -> (&mut Vec<usize>, &mut Vec<usize>) {
(&mut self.srcs, &mut self.exponents)
}
}
#[derive(Debug, Clone)]
pub struct FactoredAceCircuit<EF> {
layout: InputLayout,
constants: Vec<EF>,
shuffle_dsts: Vec<usize>,
shuffle_dst_mask: Vec<bool>,
num_fold_coeffs: usize,
num_shuffle_ops: usize,
common_ops: Vec<AceOpNode>,
geometry: StreamGeometry,
}
impl<EF: Field> FactoredAceCircuit<EF> {
pub fn layout(&self) -> &InputLayout {
&self.layout
}
pub fn num_shuffle_ops(&self) -> usize {
self.num_shuffle_ops
}
fn emit_shuffle_ops(
&self,
shuffle_srcs: &[usize],
coeff_exponents: &[usize],
beta: Option<usize>,
out: &mut Vec<AceOpNode>,
) {
let start = out.len();
let zero = AceNode::Constant(CONST_ZERO);
let beta_node =
|| AceNode::Input(beta.expect("fold challenge is required beyond a single fold slot"));
let powers_start = start + self.shuffle_dsts.len();
let power_node = |e: usize| match e {
0 => AceNode::Constant(CONST_ONE),
1 => beta_node(),
_ => AceNode::Operation(powers_start + (e - 2)),
};
for &src in shuffle_srcs {
out.push(AceOpNode {
op: AceOp::Add,
lhs: AceNode::Input(src),
rhs: zero,
});
}
for e in 2..self.num_fold_coeffs {
out.push(AceOpNode {
op: AceOp::Mul,
lhs: power_node(e - 1),
rhs: beta_node(),
});
}
for &e in coeff_exponents {
out.push(AceOpNode {
op: AceOp::Add,
lhs: power_node(e),
rhs: zero,
});
}
debug_assert!(
out.len() - start <= self.num_shuffle_ops,
"shuffle emission overran its section and would displace the common ops"
);
while out.len() - start < self.num_shuffle_ops {
out.push(AceOpNode { op: AceOp::Add, lhs: zero, rhs: zero });
}
}
pub(crate) fn encode_shuffle_section<'a>(
&self,
buffer: &'a mut ShuffleEncodeBuffer,
) -> Result<&'a [Felt], AceError> {
self.geometry.validate()?;
let beta = self.validate_assembly_with_scratch(
&buffer.srcs,
&buffer.exponents,
&mut buffer.seen_srcs,
&mut buffer.seen_exponents,
)?;
let mut ops = core::mem::take(&mut buffer.ops);
ops.clear();
self.emit_shuffle_ops(&buffer.srcs, &buffer.exponents, beta, &mut ops);
buffer.ops = ops;
buffer.felts.clear();
buffer.felts.reserve(buffer.ops.len());
for op in &buffer.ops {
buffer.felts.push(self.geometry.encode_operation(op)?);
}
Ok(&buffer.felts)
}
fn validate_assembly(
&self,
shuffle_srcs: &[usize],
coeff_exponents: &[usize],
) -> Result<Option<usize>, AceError> {
let mut seen_srcs = Vec::new();
let mut seen_exponents = Vec::new();
self.validate_assembly_with_scratch(
shuffle_srcs,
coeff_exponents,
&mut seen_srcs,
&mut seen_exponents,
)
}
fn validate_assembly_with_scratch(
&self,
shuffle_srcs: &[usize],
coeff_exponents: &[usize],
seen_srcs: &mut Vec<bool>,
seen_exponents: &mut Vec<bool>,
) -> Result<Option<usize>, AceError> {
if shuffle_srcs.len() != self.shuffle_dsts.len() {
return Err(AceError::InvalidInputLayout {
message: format!(
"shuffle source count ({}) does not match destination count ({})",
shuffle_srcs.len(),
self.shuffle_dsts.len()
),
});
}
if !is_exact_permutation(
shuffle_srcs,
self.shuffle_dsts.len(),
&self.shuffle_dst_mask,
seen_srcs,
) {
return Err(AceError::InvalidInputLayout {
message: "shuffle sources must be a permutation of the shuffled slots".into(),
});
}
if coeff_exponents.len() != self.num_fold_coeffs {
return Err(AceError::InvalidInputLayout {
message: format!(
"fold coefficient count ({}) does not match AIR count ({})",
coeff_exponents.len(),
self.num_fold_coeffs
),
});
}
seen_exponents.resize(self.num_fold_coeffs, false);
seen_exponents.fill(false);
for &exponent in coeff_exponents {
let seen =
seen_exponents.get_mut(exponent).ok_or_else(|| AceError::InvalidInputLayout {
message: format!("fold coefficient exponent {exponent} out of range"),
})?;
if *seen {
return Err(AceError::InvalidInputLayout {
message: format!("fold coefficient exponent {exponent} is used twice"),
});
}
*seen = true;
}
let beta = match self.layout.index(InputKey::MultiAirFoldBeta) {
Some(beta) => Some(beta),
None if self.num_fold_coeffs == 1 => None,
None => {
return Err(AceError::InvalidInputLayout {
message: "factored circuit requires a MultiAirFoldBeta input slot".into(),
});
},
};
Ok(beta)
}
pub fn assemble(
&self,
shuffle_srcs: &[usize],
coeff_exponents: &[usize],
) -> Result<AceCircuit<EF>, AceError> {
let beta = self.validate_assembly(shuffle_srcs, coeff_exponents)?;
let mut operations = Vec::with_capacity(self.num_shuffle_ops + self.common_ops.len());
self.emit_shuffle_ops(shuffle_srcs, coeff_exponents, beta, &mut operations);
operations.extend_from_slice(&self.common_ops);
let root = AceNode::Operation(operations.len() - 1);
Ok(AceCircuit {
layout: self.layout.clone(),
constants: self.constants.clone(),
operations,
root,
})
}
}
fn is_exact_permutation(
values: &[usize],
expected_len: usize,
membership: &[bool],
seen: &mut Vec<bool>,
) -> bool {
if values.len() != expected_len {
return false;
}
seen.resize(membership.len(), false);
seen.fill(false);
values.iter().all(|&value| {
let Some(true) = membership.get(value).copied() else {
return false;
};
!core::mem::replace(&mut seen[value], true)
})
}
pub fn emit_factored_circuit<EF>(
dag: &AceDag<EF>,
layout: InputLayout,
shuffle_dsts: Vec<usize>,
num_fold_coeffs: usize,
) -> Result<FactoredAceCircuit<EF>, AceError>
where
EF: Field,
{
layout.validate();
if num_fold_coeffs == 0 {
return Err(AceError::InvalidInputLayout {
message: "factored circuit requires at least one fold coefficient".into(),
});
}
let mut copy_by_dst = HashMap::with_capacity(shuffle_dsts.len());
let mut shuffle_dst_mask = vec![false; layout.total_inputs];
for (copy_idx, &dst) in shuffle_dsts.iter().enumerate() {
if dst >= layout.total_inputs {
return Err(AceError::InvalidInputLayout {
message: format!("shuffle destination {dst} is outside the READ layout"),
});
}
if copy_by_dst.insert(dst, copy_idx).is_some() {
return Err(AceError::InvalidInputLayout {
message: format!("duplicate shuffle destination {dst}"),
});
}
shuffle_dst_mask[dst] = true;
}
let num_copies = shuffle_dsts.len();
let num_power_ops = num_fold_coeffs.saturating_sub(2);
let unpadded = num_copies + num_power_ops + num_fold_coeffs;
let num_shuffle_ops = unpadded.next_multiple_of(ADV_PIPE_BLOCK_FELTS);
let coeffs_start = num_copies + num_power_ops;
let mut constants = vec![EF::ZERO, EF::ONE];
let mut constant_map = HashMap::<EF, usize>::new();
constant_map.insert(EF::ZERO, CONST_ZERO);
constant_map.insert(EF::ONE, CONST_ONE);
let mut common_ops: Vec<AceOpNode> = Vec::new();
let mut node_map: Vec<Option<AceNode>> = vec![None; dag.nodes().len()];
let lookup = |map: &[Option<AceNode>], id: crate::dag::NodeId| -> AceNode {
map[id.index()].expect("ACE DAG nodes must be topologically ordered")
};
for (idx, node) in dag.nodes().iter().enumerate() {
let ace_node = match node {
NodeKind::Input(InputKey::MultiAirFoldCoeff(air)) => {
if *air >= num_fold_coeffs {
return Err(AceError::InvalidInputLayout {
message: format!("fold coefficient index {air} out of range"),
});
}
AceNode::Operation(coeffs_start + air)
},
NodeKind::Input(key) => {
let input_idx = layout.index(*key).ok_or_else(|| AceError::InvalidInputLayout {
message: format!("missing input key in layout: {key:?}"),
})?;
match copy_by_dst.get(&input_idx) {
Some(©_idx) => AceNode::Operation(copy_idx),
None => match *key {
InputKey::Public(_)
| InputKey::AuxRandAlpha
| InputKey::AuxRandBeta
| InputKey::MultiAirFoldBeta
| InputKey::Reserved
| InputKey::Alpha
| InputKey::ZPowN
| InputKey::ZK
| InputKey::IsFirst
| InputKey::IsLast
| InputKey::IsTransition
| InputKey::IsFirstAir(_)
| InputKey::IsLastAir(_)
| InputKey::IsTransitionAir(_)
| InputKey::Weight0
| InputKey::F
| InputKey::S0
| InputKey::QuotientChunkCoord { .. } => AceNode::Input(input_idx),
InputKey::Preprocessed { .. }
| InputKey::Main { .. }
| InputKey::AuxCoord { .. }
| InputKey::AuxBusBoundary(_) => {
return Err(AceError::InvalidInputLayout {
message: format!(
"shuffled input key {key:?} has no shuffle destination"
),
});
},
InputKey::MultiAirFoldCoeff(_) => unreachable!(),
},
}
},
NodeKind::Constant(value) => {
let const_idx = *constant_map.entry(*value).or_insert_with(|| {
constants.push(*value);
constants.len() - 1
});
AceNode::Constant(const_idx)
},
NodeKind::Add(a, b) => {
let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
common_ops.push(AceOpNode { op: AceOp::Add, lhs, rhs });
AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
},
NodeKind::Sub(a, b) => {
let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
common_ops.push(AceOpNode { op: AceOp::Sub, lhs, rhs });
AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
},
NodeKind::Mul(a, b) => {
let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
common_ops.push(AceOpNode { op: AceOp::Mul, lhs, rhs });
AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
},
NodeKind::Neg(a) => {
let rhs = lookup(&node_map, *a);
common_ops.push(AceOpNode {
op: AceOp::Sub,
lhs: AceNode::Constant(CONST_ZERO),
rhs,
});
AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
},
};
node_map[idx] = Some(ace_node);
}
match lookup(&node_map, dag.root()) {
AceNode::Operation(idx) if idx == num_shuffle_ops + common_ops.len() - 1 => {},
other => {
return Err(AceError::InvalidInputLayout {
message: format!("factored DAG root must be the last common op, got {other:?}"),
});
},
}
let padded_len = constants.len().next_multiple_of(CONST_EF_BLOCK_ALIGN);
constants.resize(padded_len, EF::ZERO);
let num_ops = num_shuffle_ops + common_ops.len();
let geometry = StreamGeometry::from_counts(layout.total_inputs, constants.len(), num_ops);
Ok(FactoredAceCircuit {
layout,
constants,
shuffle_dsts,
shuffle_dst_mask,
num_fold_coeffs,
num_shuffle_ops,
common_ops,
geometry,
})
}
#[cfg(test)]
mod tests {
use super::is_exact_permutation;
#[test]
fn exact_permutation_rejects_missing_duplicate_and_foreign_values() {
let membership = [false, true, false, true, true];
let mut seen = Vec::new();
assert!(is_exact_permutation(&[4, 1, 3], 3, &membership, &mut seen));
assert!(!is_exact_permutation(&[1, 3], 3, &membership, &mut seen));
assert!(!is_exact_permutation(&[1, 1, 4], 3, &membership, &mut seen));
assert!(!is_exact_permutation(&[1, 2, 4], 3, &membership, &mut seen));
assert!(!is_exact_permutation(&[1, 3, 5], 3, &membership, &mut seen));
}
}