use miden_core::{Felt, Word, crypto::hash::Poseidon2};
use miden_crypto::{
field::{ExtensionField, Field},
hash::poseidon2::Poseidon2Permutation256,
stark::symmetric::Permutation,
};
use miden_field::{PackedValue, PrimeCharacteristicRing};
use crate::{
AceError, EXT_DEGREE, encode::EncodedCircuit, factored::ShuffleEncodeBuffer,
pipeline::FactoredMultiAirCircuit,
};
const RATE_WIDTH: usize = Poseidon2::RATE_RANGE.end - Poseidon2::RATE_RANGE.start;
type PackedFelt = <Felt as Field>::Packing;
pub const LEAF_LANES: usize = <PackedFelt as PackedValue>::WIDTH;
#[derive(Default)]
pub struct PackedLeafScratch {
buffer: ShuffleEncodeBuffer,
streams: Vec<Vec<Felt>>,
}
impl PackedLeafScratch {
pub fn new() -> Self {
Self::default()
}
}
#[derive(Clone, Debug)]
pub struct FactoredEncodedCircuit {
pub encoded: EncodedCircuit,
pub shuffle_prefix_len: usize,
pub shuffle_commitment: Word,
pub common_commitment: Word,
pub commitment: Word,
}
pub struct FactoredCircuitFactory<EF> {
factored: FactoredMultiAirCircuit<EF>,
constants_state: [Felt; Poseidon2::STATE_WIDTH],
const_felts: usize,
common_commitment: Word,
}
impl<EF> FactoredCircuitFactory<EF>
where
EF: ExtensionField<Felt>,
{
pub fn new(factored: FactoredMultiAirCircuit<EF>) -> Result<Self, AceError> {
let canonical: Vec<usize> = (0..factored.num_airs()).collect();
let circuit = factored.circuit_for_order(&canonical)?;
let encoded = circuit.to_ace()?;
let instructions = encoded.instructions();
let const_felts = encoded.num_constants() * EXT_DEGREE;
let prefix_len = const_felts + factored.num_shuffle_ops();
if !const_felts.is_multiple_of(RATE_WIDTH)
|| !prefix_len.is_multiple_of(RATE_WIDTH)
|| prefix_len >= instructions.len()
{
return Err(AceError::InvalidInputLayout {
message: "ACE stream sections must be rate-aligned for prefix resumption".into(),
});
}
let mut constants_state = [<Felt as PrimeCharacteristicRing>::ZERO; Poseidon2::STATE_WIDTH];
absorb_rate_blocks(&mut constants_state, &instructions[..const_felts]);
let common_commitment = Poseidon2::hash_elements(&instructions[prefix_len..]);
let mut buffer = ShuffleEncodeBuffer::new();
let fast = factored.encode_shuffle_section_for_order(&canonical, &mut buffer)?;
if fast != &instructions[const_felts..prefix_len] {
return Err(AceError::InvalidInputLayout {
message: "encode-only shuffle section diverges from the assembled stream".into(),
});
}
let mut resumed = constants_state;
absorb_rate_blocks(&mut resumed, fast);
let resumed_prefix =
Word::new(resumed[Poseidon2::RATE0_RANGE].try_into().expect("digest is one word"));
if resumed_prefix != Poseidon2::hash_elements(&instructions[..prefix_len]) {
return Err(AceError::InvalidInputLayout {
message: "resumed prefix hash diverges from hashing the full prefix".into(),
});
}
Ok(Self {
factored,
constants_state,
const_felts,
common_commitment,
})
}
pub fn factored(&self) -> &FactoredMultiAirCircuit<EF> {
&self.factored
}
pub fn const_felts(&self) -> usize {
self.const_felts
}
pub fn leaf_for_order(
&self,
proof_order: &[usize],
buffer: &mut ShuffleEncodeBuffer,
) -> Result<Word, AceError> {
let shuffle = self.factored.encode_shuffle_section_for_order(proof_order, buffer)?;
let mut state = self.constants_state;
absorb_rate_blocks(&mut state, shuffle);
let shuffle_commitment =
Word::new(state[Poseidon2::RATE0_RANGE].try_into().expect("digest is one word"));
Ok(Poseidon2::merge(&[shuffle_commitment, self.common_commitment]))
}
pub fn leaves_for_orders(
&self,
orders: &[&[usize]],
scratch: &mut PackedLeafScratch,
out: &mut Vec<Word>,
) -> Result<(), AceError> {
scratch.streams.resize_with(LEAF_LANES, Vec::new);
for chunk in orders.chunks(LEAF_LANES) {
for lane in 0..LEAF_LANES {
let order = chunk.get(lane).copied().unwrap_or(chunk[chunk.len() - 1]);
let shuffle =
self.factored.encode_shuffle_section_for_order(order, &mut scratch.buffer)?;
scratch.streams[lane].clear();
scratch.streams[lane].extend_from_slice(shuffle);
}
let mut state: [PackedFelt; Poseidon2::STATE_WIDTH] = core::array::from_fn(|e| {
let mut packed = <PackedFelt as PrimeCharacteristicRing>::ZERO;
packed.as_slice_mut().fill(self.constants_state[e]);
packed
});
assert!(
scratch.streams[0].len().is_multiple_of(RATE_WIDTH),
"shuffle streams must be rate-aligned"
);
let blocks = scratch.streams[0].len() / RATE_WIDTH;
for block in 0..blocks {
for i in 0..RATE_WIDTH {
let elem = &mut state[Poseidon2::RATE_RANGE.start + i];
for lane in 0..LEAF_LANES {
elem.as_slice_mut()[lane] = scratch.streams[lane][block * RATE_WIDTH + i];
}
}
Poseidon2Permutation256.permute_mut(&mut state);
}
let mut merge_state: [PackedFelt; Poseidon2::STATE_WIDTH] =
core::array::from_fn(|_| <PackedFelt as PrimeCharacteristicRing>::ZERO);
for i in 0..4 {
merge_state[i] = state[Poseidon2::RATE0_RANGE.start + i];
let mut common = <PackedFelt as PrimeCharacteristicRing>::ZERO;
common.as_slice_mut().fill(self.common_commitment[i]);
merge_state[4 + i] = common;
}
Poseidon2Permutation256.permute_mut(&mut merge_state);
for lane in 0..chunk.len() {
let leaf: [Felt; 4] = core::array::from_fn(|i| merge_state[i].as_slice()[lane]);
out.push(Word::new(leaf));
}
}
Ok(())
}
pub fn circuit_for_order(
&self,
proof_order: &[usize],
) -> Result<FactoredEncodedCircuit, AceError> {
let circuit = self.factored.circuit_for_order(proof_order)?;
let encoded = circuit.to_ace()?;
let instructions = encoded.instructions();
let stream_len = encoded.size_in_felt();
if stream_len != instructions.len() {
return Err(AceError::InvalidInputLayout {
message: format!(
"ACE circuit stream length ({stream_len}) does not match instruction count \
({})",
instructions.len()
),
});
}
let shuffle_prefix_len = self.const_felts + self.factored.num_shuffle_ops();
if encoded.num_constants() * EXT_DEGREE != self.const_felts
|| !stream_len.is_multiple_of(RATE_WIDTH)
|| shuffle_prefix_len >= stream_len
{
return Err(AceError::InvalidInputLayout {
message: "assembled ACE stream does not match the factored section layout".into(),
});
}
let mut state = self.constants_state;
absorb_rate_blocks(&mut state, &instructions[self.const_felts..shuffle_prefix_len]);
let shuffle_commitment =
Word::new(state[Poseidon2::RATE0_RANGE].try_into().expect("digest is one word"));
let common_commitment = self.common_commitment;
let commitment = Poseidon2::merge(&[shuffle_commitment, common_commitment]);
Ok(FactoredEncodedCircuit {
encoded,
shuffle_prefix_len,
shuffle_commitment,
common_commitment,
commitment,
})
}
}
fn absorb_rate_blocks(state: &mut [Felt; Poseidon2::STATE_WIDTH], elements: &[Felt]) {
assert!(
elements.len().is_multiple_of(RATE_WIDTH),
"sponge absorption requires whole rate blocks"
);
for block in elements.as_chunks::<RATE_WIDTH>().0 {
state[Poseidon2::RATE_RANGE].copy_from_slice(block);
Poseidon2::apply_permutation(state);
}
}