1use std::collections::HashMap;
23
24use miden_core::Felt;
25use miden_crypto::field::Field;
26
27use crate::{
28 AceError, EXT_DEGREE, InputLayout,
29 circuit::{AceCircuit, AceNode, AceOp, AceOpNode},
30 dag::{AceDag, NodeKind},
31 encode::{ADV_PIPE_BLOCK_FELTS, CONST_EF_ALIGN, StreamGeometry},
32 layout::InputKey,
33};
34
35const CONST_EF_BLOCK_ALIGN: usize = ADV_PIPE_BLOCK_FELTS / EXT_DEGREE;
37
38const _: () = assert!(
43 CONST_EF_BLOCK_ALIGN.is_multiple_of(CONST_EF_ALIGN),
44 "constant block alignment must refine the encoder's READ-row alignment, or the two-segment split drifts off a block boundary"
45);
46
47const CONST_ZERO: usize = 0;
49const CONST_ONE: usize = 1;
51
52#[derive(Clone, Debug, Default)]
57pub struct ShuffleEncodeBuffer {
58 srcs: Vec<usize>,
59 exponents: Vec<usize>,
60 seen_srcs: Vec<bool>,
61 seen_exponents: Vec<bool>,
62 ops: Vec<AceOpNode>,
63 felts: Vec<Felt>,
64}
65
66impl ShuffleEncodeBuffer {
67 pub fn new() -> Self {
69 Self::default()
70 }
71
72 pub(crate) fn order_scratch(&mut self) -> (&mut Vec<usize>, &mut Vec<usize>) {
74 (&mut self.srcs, &mut self.exponents)
75 }
76}
77
78#[derive(Debug, Clone)]
80pub struct FactoredAceCircuit<EF> {
81 layout: InputLayout,
82 constants: Vec<EF>,
84 shuffle_dsts: Vec<usize>,
86 shuffle_dst_mask: Vec<bool>,
88 num_fold_coeffs: usize,
90 num_shuffle_ops: usize,
92 common_ops: Vec<AceOpNode>,
94 geometry: StreamGeometry,
96}
97
98impl<EF: Field> FactoredAceCircuit<EF> {
99 pub fn layout(&self) -> &InputLayout {
101 &self.layout
102 }
103
104 pub fn num_shuffle_ops(&self) -> usize {
106 self.num_shuffle_ops
107 }
108
109 fn emit_shuffle_ops(
118 &self,
119 shuffle_srcs: &[usize],
120 coeff_exponents: &[usize],
121 beta: Option<usize>,
122 out: &mut Vec<AceOpNode>,
123 ) {
124 let start = out.len();
125 let zero = AceNode::Constant(CONST_ZERO);
126 let beta_node =
127 || AceNode::Input(beta.expect("fold challenge is required beyond a single fold slot"));
128 let powers_start = start + self.shuffle_dsts.len();
129 let power_node = |e: usize| match e {
130 0 => AceNode::Constant(CONST_ONE),
131 1 => beta_node(),
132 _ => AceNode::Operation(powers_start + (e - 2)),
133 };
134
135 for &src in shuffle_srcs {
136 out.push(AceOpNode {
137 op: AceOp::Add,
138 lhs: AceNode::Input(src),
139 rhs: zero,
140 });
141 }
142 for e in 2..self.num_fold_coeffs {
143 out.push(AceOpNode {
144 op: AceOp::Mul,
145 lhs: power_node(e - 1),
146 rhs: beta_node(),
147 });
148 }
149 for &e in coeff_exponents {
150 out.push(AceOpNode {
151 op: AceOp::Add,
152 lhs: power_node(e),
153 rhs: zero,
154 });
155 }
156 debug_assert!(
157 out.len() - start <= self.num_shuffle_ops,
158 "shuffle emission overran its section and would displace the common ops"
159 );
160 while out.len() - start < self.num_shuffle_ops {
161 out.push(AceOpNode { op: AceOp::Add, lhs: zero, rhs: zero });
162 }
163 }
164
165 pub(crate) fn encode_shuffle_section<'a>(
177 &self,
178 buffer: &'a mut ShuffleEncodeBuffer,
179 ) -> Result<&'a [Felt], AceError> {
180 self.geometry.validate()?;
181
182 let beta = self.validate_assembly_with_scratch(
183 &buffer.srcs,
184 &buffer.exponents,
185 &mut buffer.seen_srcs,
186 &mut buffer.seen_exponents,
187 )?;
188
189 let mut ops = core::mem::take(&mut buffer.ops);
190 ops.clear();
191 self.emit_shuffle_ops(&buffer.srcs, &buffer.exponents, beta, &mut ops);
192 buffer.ops = ops;
193
194 buffer.felts.clear();
195 buffer.felts.reserve(buffer.ops.len());
196 for op in &buffer.ops {
197 buffer.felts.push(self.geometry.encode_operation(op)?);
198 }
199 Ok(&buffer.felts)
200 }
201
202 fn validate_assembly(
210 &self,
211 shuffle_srcs: &[usize],
212 coeff_exponents: &[usize],
213 ) -> Result<Option<usize>, AceError> {
214 let mut seen_srcs = Vec::new();
215 let mut seen_exponents = Vec::new();
216 self.validate_assembly_with_scratch(
217 shuffle_srcs,
218 coeff_exponents,
219 &mut seen_srcs,
220 &mut seen_exponents,
221 )
222 }
223
224 fn validate_assembly_with_scratch(
225 &self,
226 shuffle_srcs: &[usize],
227 coeff_exponents: &[usize],
228 seen_srcs: &mut Vec<bool>,
229 seen_exponents: &mut Vec<bool>,
230 ) -> Result<Option<usize>, AceError> {
231 if shuffle_srcs.len() != self.shuffle_dsts.len() {
232 return Err(AceError::InvalidInputLayout {
233 message: format!(
234 "shuffle source count ({}) does not match destination count ({})",
235 shuffle_srcs.len(),
236 self.shuffle_dsts.len()
237 ),
238 });
239 }
240 if !is_exact_permutation(
241 shuffle_srcs,
242 self.shuffle_dsts.len(),
243 &self.shuffle_dst_mask,
244 seen_srcs,
245 ) {
246 return Err(AceError::InvalidInputLayout {
247 message: "shuffle sources must be a permutation of the shuffled slots".into(),
248 });
249 }
250 if coeff_exponents.len() != self.num_fold_coeffs {
251 return Err(AceError::InvalidInputLayout {
252 message: format!(
253 "fold coefficient count ({}) does not match AIR count ({})",
254 coeff_exponents.len(),
255 self.num_fold_coeffs
256 ),
257 });
258 }
259 seen_exponents.resize(self.num_fold_coeffs, false);
263 seen_exponents.fill(false);
264 for &exponent in coeff_exponents {
265 let seen =
266 seen_exponents.get_mut(exponent).ok_or_else(|| AceError::InvalidInputLayout {
267 message: format!("fold coefficient exponent {exponent} out of range"),
268 })?;
269 if *seen {
270 return Err(AceError::InvalidInputLayout {
271 message: format!("fold coefficient exponent {exponent} is used twice"),
272 });
273 }
274 *seen = true;
275 }
276
277 let beta = match self.layout.index(InputKey::MultiAirFoldBeta) {
280 Some(beta) => Some(beta),
281 None if self.num_fold_coeffs == 1 => None,
282 None => {
283 return Err(AceError::InvalidInputLayout {
284 message: "factored circuit requires a MultiAirFoldBeta input slot".into(),
285 });
286 },
287 };
288 Ok(beta)
289 }
290
291 pub fn assemble(
297 &self,
298 shuffle_srcs: &[usize],
299 coeff_exponents: &[usize],
300 ) -> Result<AceCircuit<EF>, AceError> {
301 let beta = self.validate_assembly(shuffle_srcs, coeff_exponents)?;
302
303 let mut operations = Vec::with_capacity(self.num_shuffle_ops + self.common_ops.len());
304 self.emit_shuffle_ops(shuffle_srcs, coeff_exponents, beta, &mut operations);
305 operations.extend_from_slice(&self.common_ops);
306
307 let root = AceNode::Operation(operations.len() - 1);
308 Ok(AceCircuit {
309 layout: self.layout.clone(),
310 constants: self.constants.clone(),
311 operations,
312 root,
313 })
314 }
315}
316
317fn is_exact_permutation(
319 values: &[usize],
320 expected_len: usize,
321 membership: &[bool],
322 seen: &mut Vec<bool>,
323) -> bool {
324 if values.len() != expected_len {
325 return false;
326 }
327 seen.resize(membership.len(), false);
328 seen.fill(false);
329 values.iter().all(|&value| {
330 let Some(true) = membership.get(value).copied() else {
331 return false;
332 };
333 !core::mem::replace(&mut seen[value], true)
334 })
335}
336
337pub fn emit_factored_circuit<EF>(
344 dag: &AceDag<EF>,
345 layout: InputLayout,
346 shuffle_dsts: Vec<usize>,
347 num_fold_coeffs: usize,
348) -> Result<FactoredAceCircuit<EF>, AceError>
349where
350 EF: Field,
351{
352 layout.validate();
353 if num_fold_coeffs == 0 {
354 return Err(AceError::InvalidInputLayout {
355 message: "factored circuit requires at least one fold coefficient".into(),
356 });
357 }
358
359 let mut copy_by_dst = HashMap::with_capacity(shuffle_dsts.len());
360 let mut shuffle_dst_mask = vec![false; layout.total_inputs];
361 for (copy_idx, &dst) in shuffle_dsts.iter().enumerate() {
362 if dst >= layout.total_inputs {
363 return Err(AceError::InvalidInputLayout {
364 message: format!("shuffle destination {dst} is outside the READ layout"),
365 });
366 }
367 if copy_by_dst.insert(dst, copy_idx).is_some() {
368 return Err(AceError::InvalidInputLayout {
369 message: format!("duplicate shuffle destination {dst}"),
370 });
371 }
372 shuffle_dst_mask[dst] = true;
373 }
374
375 let num_copies = shuffle_dsts.len();
376 let num_power_ops = num_fold_coeffs.saturating_sub(2);
377 let unpadded = num_copies + num_power_ops + num_fold_coeffs;
378 let num_shuffle_ops = unpadded.next_multiple_of(ADV_PIPE_BLOCK_FELTS);
379 let coeffs_start = num_copies + num_power_ops;
380
381 let mut constants = vec![EF::ZERO, EF::ONE];
382 let mut constant_map = HashMap::<EF, usize>::new();
383 constant_map.insert(EF::ZERO, CONST_ZERO);
384 constant_map.insert(EF::ONE, CONST_ONE);
385
386 let mut common_ops: Vec<AceOpNode> = Vec::new();
387 let mut node_map: Vec<Option<AceNode>> = vec![None; dag.nodes().len()];
388
389 let lookup = |map: &[Option<AceNode>], id: crate::dag::NodeId| -> AceNode {
390 map[id.index()].expect("ACE DAG nodes must be topologically ordered")
391 };
392
393 for (idx, node) in dag.nodes().iter().enumerate() {
394 let ace_node = match node {
395 NodeKind::Input(InputKey::MultiAirFoldCoeff(air)) => {
396 if *air >= num_fold_coeffs {
397 return Err(AceError::InvalidInputLayout {
398 message: format!("fold coefficient index {air} out of range"),
399 });
400 }
401 AceNode::Operation(coeffs_start + air)
402 },
403 NodeKind::Input(key) => {
404 let input_idx = layout.index(*key).ok_or_else(|| AceError::InvalidInputLayout {
405 message: format!("missing input key in layout: {key:?}"),
406 })?;
407 match copy_by_dst.get(&input_idx) {
408 Some(©_idx) => AceNode::Operation(copy_idx),
409 None => match *key {
414 InputKey::Public(_)
415 | InputKey::AuxRandAlpha
416 | InputKey::AuxRandBeta
417 | InputKey::MultiAirFoldBeta
418 | InputKey::Reserved
419 | InputKey::Alpha
420 | InputKey::ZPowN
421 | InputKey::ZK
422 | InputKey::IsFirst
423 | InputKey::IsLast
424 | InputKey::IsTransition
425 | InputKey::IsFirstAir(_)
426 | InputKey::IsLastAir(_)
427 | InputKey::IsTransitionAir(_)
428 | InputKey::Weight0
429 | InputKey::F
430 | InputKey::S0
431 | InputKey::QuotientChunkCoord { .. } => AceNode::Input(input_idx),
432 InputKey::Preprocessed { .. }
437 | InputKey::Main { .. }
438 | InputKey::AuxCoord { .. }
439 | InputKey::AuxBusBoundary(_) => {
440 return Err(AceError::InvalidInputLayout {
441 message: format!(
442 "shuffled input key {key:?} has no shuffle destination"
443 ),
444 });
445 },
446 InputKey::MultiAirFoldCoeff(_) => unreachable!(),
448 },
449 }
450 },
451 NodeKind::Constant(value) => {
452 let const_idx = *constant_map.entry(*value).or_insert_with(|| {
453 constants.push(*value);
454 constants.len() - 1
455 });
456 AceNode::Constant(const_idx)
457 },
458 NodeKind::Add(a, b) => {
459 let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
460 common_ops.push(AceOpNode { op: AceOp::Add, lhs, rhs });
461 AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
462 },
463 NodeKind::Sub(a, b) => {
464 let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
465 common_ops.push(AceOpNode { op: AceOp::Sub, lhs, rhs });
466 AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
467 },
468 NodeKind::Mul(a, b) => {
469 let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
470 common_ops.push(AceOpNode { op: AceOp::Mul, lhs, rhs });
471 AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
472 },
473 NodeKind::Neg(a) => {
474 let rhs = lookup(&node_map, *a);
475 common_ops.push(AceOpNode {
476 op: AceOp::Sub,
477 lhs: AceNode::Constant(CONST_ZERO),
478 rhs,
479 });
480 AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
481 },
482 };
483 node_map[idx] = Some(ace_node);
484 }
485
486 match lookup(&node_map, dag.root()) {
487 AceNode::Operation(idx) if idx == num_shuffle_ops + common_ops.len() - 1 => {},
488 other => {
489 return Err(AceError::InvalidInputLayout {
490 message: format!("factored DAG root must be the last common op, got {other:?}"),
491 });
492 },
493 }
494
495 let padded_len = constants.len().next_multiple_of(CONST_EF_BLOCK_ALIGN);
497 constants.resize(padded_len, EF::ZERO);
498
499 let num_ops = num_shuffle_ops + common_ops.len();
503 let geometry = StreamGeometry::from_counts(layout.total_inputs, constants.len(), num_ops);
504
505 Ok(FactoredAceCircuit {
506 layout,
507 constants,
508 shuffle_dsts,
509 shuffle_dst_mask,
510 num_fold_coeffs,
511 num_shuffle_ops,
512 common_ops,
513 geometry,
514 })
515}
516
517#[cfg(test)]
518mod tests {
519 use super::is_exact_permutation;
520
521 #[test]
522 fn exact_permutation_rejects_missing_duplicate_and_foreign_values() {
523 let membership = [false, true, false, true, true];
524 let mut seen = Vec::new();
525
526 assert!(is_exact_permutation(&[4, 1, 3], 3, &membership, &mut seen));
527 assert!(!is_exact_permutation(&[1, 3], 3, &membership, &mut seen));
528 assert!(!is_exact_permutation(&[1, 1, 4], 3, &membership, &mut seen));
529 assert!(!is_exact_permutation(&[1, 2, 4], 3, &membership, &mut seen));
530 assert!(!is_exact_permutation(&[1, 3, 5], 3, &membership, &mut seen));
531 }
532}