use alloc::vec::Vec;
use core::{borrow::BorrowMut, ops::Range};
use miden_air::{
CYCLE_INPUT_ROW, CYCLE_OUTPUT_ROW, ControllerCols, INITIAL_EXTERNAL_ROUND_END,
INITIAL_EXTERNAL_ROUND_START, INTERNAL_PLUS_EXTERNAL_ROW, LAST_INTERNAL_ROUND_ARK_IDX,
NUM_PACKED_INTERNAL_ROUND_ROWS, NUM_SBOX_WITNESSES, NUM_TRAILING_EXTERNAL_ROUND_ROWS,
PACKED_INTERNAL_ROUND_START, Poseidon2PermutationCols,
trace::{
chiplets::hasher::{CONTROLLER_TRACE_ALIGNMENT, HASH_CYCLE_LEN, TRACE_WIDTH},
poseidon2_permutation::NUM_POSEIDON2_PERMUTATION_COLS,
},
};
use miden_core::chiplets::hasher::Hasher;
use rayon::prelude::*;
use super::{
ChipletTraceFragment, Felt, HasherState, ONE, PermRequest, STATE_WIDTH, Selectors, ZERO,
perm_id_felt,
};
#[derive(Debug, Clone)]
enum HasherOp {
Controller {
selectors: Selectors,
state: HasherState,
node_index: Felt,
mrupdate_id: Felt,
is_boundary: Felt,
direction_bit: Felt,
perm_id: Felt,
},
Padding { count: usize, mrupdate_id: Felt },
}
impl HasherOp {
fn row_count(&self) -> usize {
match self {
Self::Controller { .. } => 1,
Self::Padding { count, .. } => *count,
}
}
}
#[derive(Debug, Default)]
pub(super) struct HasherTrace {
ops: Vec<HasherOp>,
row_count: usize,
}
impl HasherTrace {
pub(super) fn trace_len(&self) -> usize {
self.row_count
}
pub(super) fn next_row_addr(&self) -> Felt {
Felt::new_unchecked(self.row_count as u64 + 1)
}
pub(super) fn next_op_index(&self) -> usize {
self.ops.len()
}
pub(super) fn append_controller_row(
&mut self,
selectors: Selectors,
state: &HasherState,
node_index: Felt,
mrupdate_id: Felt,
is_boundary: Felt,
direction_bit: Felt,
perm_id: Felt,
) {
self.ops.push(HasherOp::Controller {
selectors,
state: *state,
node_index,
mrupdate_id,
is_boundary,
direction_bit,
perm_id,
});
self.row_count += 1;
}
pub(super) fn pad_to_controller_boundary(&mut self, mrupdate_id: Felt) {
let remainder = self.row_count % CONTROLLER_TRACE_ALIGNMENT;
if remainder != 0 {
let count = CONTROLLER_TRACE_ALIGNMENT - remainder;
self.ops.push(HasherOp::Padding { count, mrupdate_id });
self.row_count += count;
}
}
pub(super) fn replay_ops_range(
&mut self,
range: Range<usize>,
new_mrupdate_id: Felt,
) -> (HasherState, Vec<HasherState>) {
let mut last_state = [ZERO; STATE_WIDTH];
let mut input_states = Vec::with_capacity(range.len() / 2);
for idx in range {
let mut op = self.ops[idx].clone();
match &mut op {
HasherOp::Controller { mrupdate_id, selectors, state, .. } => {
*mrupdate_id = new_mrupdate_id;
let [is_input, _, _] = *selectors;
if is_input == ONE {
input_states.push(*state);
}
last_state = *state;
},
HasherOp::Padding { mrupdate_id, .. } => {
*mrupdate_id = new_mrupdate_id;
},
}
self.row_count += op.row_count();
self.ops.push(op);
}
(last_state, input_states)
}
pub(super) fn fill_trace(self, trace: &mut ChipletTraceFragment) {
debug_assert_eq!(self.trace_len(), trace.len(), "inconsistent trace lengths");
debug_assert_eq!(TRACE_WIDTH, trace.width(), "inconsistent trace widths");
let mut chunk = [ZERO; TRACE_WIDTH * CONTROLLER_TRACE_ALIGNMENT];
let mut row_idx = 0usize;
for op in &self.ops {
let n = op.row_count();
debug_assert!(n <= CONTROLLER_TRACE_ALIGNMENT);
let (chunk_rows, _) = chunk.as_mut_slice().as_chunks_mut::<TRACE_WIDTH>();
match op {
HasherOp::Controller {
selectors,
state,
node_index,
mrupdate_id,
is_boundary,
direction_bit,
perm_id,
} => {
write_controller_row(
&mut chunk_rows[0],
*selectors,
state,
*node_index,
*mrupdate_id,
*is_boundary,
*direction_bit,
*perm_id,
);
},
HasherOp::Padding { count, mrupdate_id } => {
let padding_selectors = [ZERO, ONE, ZERO];
for row in &mut chunk_rows[..*count] {
write_controller_row(
row,
padding_selectors,
&[ZERO; STATE_WIDTH],
ZERO,
*mrupdate_id,
ZERO,
ZERO,
ZERO,
);
}
},
}
trace.copy_rows_into(row_idx, &chunk[..n * TRACE_WIDTH]);
row_idx += n;
}
debug_assert_eq!(row_idx, self.row_count);
}
}
fn write_controller_row(
row: &mut [Felt; TRACE_WIDTH],
selectors: Selectors,
state: &HasherState,
node_index: Felt,
mrupdate_id: Felt,
is_boundary: Felt,
direction_bit: Felt,
perm_id: Felt,
) {
let cols: &mut ControllerCols<Felt> = row.as_mut_slice().borrow_mut();
let [s0, s1, s2] = selectors;
cols.s0 = s0;
cols.s1 = s1;
cols.s2 = s2;
cols.state = *state;
cols.node_index = node_index;
cols.mrupdate_id = mrupdate_id;
cols.is_boundary = is_boundary;
cols.direction_bit = direction_bit;
cols.perm_id = perm_id;
}
pub(super) fn write_poseidon2_permutation_cycle(
rows: &mut [[Felt; NUM_POSEIDON2_PERMUTATION_COLS]],
init_state: &HasherState,
perm_id: Felt,
multiplicity: Felt,
) {
debug_assert_eq!(rows.len(), HASH_CYCLE_LEN);
let mut state = *init_state;
let zero_witnesses = [ZERO; NUM_SBOX_WITNESSES];
let multiplicity_witnesses = witnesses_with_first(multiplicity);
write_perm_row(&mut rows[CYCLE_INPUT_ROW], &state, perm_id, multiplicity_witnesses);
Hasher::apply_matmul_external(&mut state);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_INITIAL[0]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
for (offset, row) in rows[INITIAL_EXTERNAL_ROUND_START..INITIAL_EXTERNAL_ROUND_END]
.iter_mut()
.enumerate()
{
let round = INITIAL_EXTERNAL_ROUND_START + offset;
write_perm_row(row, &state, perm_id, zero_witnesses);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_INITIAL[round]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
}
for triple in 0..NUM_PACKED_INTERNAL_ROUND_ROWS {
let base = triple * NUM_SBOX_WITNESSES;
let pre_state = state;
let mut witnesses = zero_witnesses;
for (k, witness) in witnesses.iter_mut().enumerate() {
let sbox_out = (state[0] + Hasher::ARK_INT[base + k]).exp_const_u64::<7>();
*witness = sbox_out;
state[0] = sbox_out;
Hasher::matmul_internal(&mut state, Hasher::MAT_DIAG);
}
write_perm_row(
&mut rows[PACKED_INTERNAL_ROUND_START + triple],
&pre_state,
perm_id,
witnesses,
);
}
let pre_state = state;
let w0 = (state[0] + Hasher::ARK_INT[LAST_INTERNAL_ROUND_ARK_IDX]).exp_const_u64::<7>();
state[0] = w0;
Hasher::matmul_internal(&mut state, Hasher::MAT_DIAG);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_TERMINAL[0]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
let final_internal_witnesses = witnesses_with_first(w0);
write_perm_row(
&mut rows[INTERNAL_PLUS_EXTERNAL_ROW],
&pre_state,
perm_id,
final_internal_witnesses,
);
for round in 1..=NUM_TRAILING_EXTERNAL_ROUND_ROWS {
write_perm_row(
&mut rows[INTERNAL_PLUS_EXTERNAL_ROW + round],
&state,
perm_id,
zero_witnesses,
);
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_TERMINAL[round]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
}
write_perm_row(&mut rows[CYCLE_OUTPUT_ROW], &state, perm_id, multiplicity_witnesses);
}
pub(super) fn fill_poseidon2_permutation_trace(
perm_requests: Vec<PermRequest>,
trace: &mut [Felt],
) {
const W: usize = NUM_POSEIDON2_PERMUTATION_COLS;
assert_eq!(trace.len() % W, 0, "Poseidon2 trace buffer is not row-aligned");
let (rows, _) = trace.as_chunks_mut::<W>();
assert_eq!(rows.len() % HASH_CYCLE_LEN, 0, "Poseidon2 height must align to cycles");
assert!(
(perm_requests.len() + 1) * HASH_CYCLE_LEN <= rows.len(),
"Poseidon2 trace buffer is too short for permutation requests",
);
let request_count = perm_requests.len();
rows[..request_count * HASH_CYCLE_LEN]
.par_chunks_exact_mut(HASH_CYCLE_LEN)
.zip(perm_requests.par_iter())
.enumerate()
.for_each(|(perm_id, (cycle_rows, request))| {
let state = request.state.map(Felt::new_unchecked);
write_poseidon2_permutation_cycle(
cycle_rows,
&state,
perm_id_felt(perm_id),
Felt::new_unchecked(request.multiplicity),
);
});
let padding_start = request_count * HASH_CYCLE_LEN;
let zero_state = [ZERO; STATE_WIDTH];
if padding_start < rows.len() {
write_poseidon2_permutation_cycle(
&mut rows[padding_start..padding_start + HASH_CYCLE_LEN],
&zero_state,
perm_id_felt(request_count),
ZERO,
);
let (head, tail) = rows.split_at_mut(padding_start + HASH_CYCLE_LEN);
let template = &head[padding_start..];
tail.par_chunks_exact_mut(HASH_CYCLE_LEN)
.enumerate()
.for_each(|(cycle, cycle_rows)| {
cycle_rows.copy_from_slice(template);
set_perm_id(cycle_rows, perm_id_felt(request_count + 1 + cycle));
});
}
}
fn witnesses_with_first(value: Felt) -> [Felt; NUM_SBOX_WITNESSES] {
let mut witnesses = [ZERO; NUM_SBOX_WITNESSES];
witnesses[0] = value;
witnesses
}
fn set_perm_id(rows: &mut [[Felt; NUM_POSEIDON2_PERMUTATION_COLS]], perm_id: Felt) {
debug_assert_eq!(rows.len(), HASH_CYCLE_LEN);
for row in rows {
let cols: &mut Poseidon2PermutationCols<Felt> = row[..].borrow_mut();
cols.perm_id = perm_id;
}
}
fn write_perm_row(
row: &mut [Felt; NUM_POSEIDON2_PERMUTATION_COLS],
state: &HasherState,
perm_id: Felt,
witnesses: [Felt; NUM_SBOX_WITNESSES],
) {
let cols: &mut Poseidon2PermutationCols<Felt> = row[..].borrow_mut();
cols.witnesses = witnesses;
cols.state = *state;
cols.perm_id = perm_id;
}