use alloc::vec::Vec;
use core::{borrow::Borrow, 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::{
ONE, ZERO,
chiplets::hasher,
crypto::merkle::{MerkleTree, NodeIndex},
field::PrimeCharacteristicRing,
mast::OpBatch,
};
use miden_utils_testing::rand::rand_array;
use super::{
ChipletTraceFragment, Digest, Felt, Hasher, HasherState, LINEAR_HASH, MP_VERIFY, MR_UPDATE_NEW,
MR_UPDATE_OLD, RETURN_HASH, RETURN_STATE, Selectors, absorb_into_state, get_digest, init_state,
init_state_from_words,
};
#[test]
fn hasher_permute() {
let mut hasher = Hasher::default();
let init_state: HasherState = rand_array();
let (addr, final_state) = hasher.permute(init_state);
assert_eq!(ONE, addr);
let expected_state = apply_permutation(init_state);
assert_eq!(expected_state, final_state);
let trace = build_trace(hasher);
assert_eq!(trace.controller.len(), controller_len(2));
assert_eq!(trace.poseidon2.len(), 2 * HASH_CYCLE_LEN);
check_controller_input(&trace.controller, 0, LINEAR_HASH, &init_state, ZERO, ONE, ZERO, ZERO);
check_controller_output(&trace.controller, 1, RETURN_STATE, &expected_state, ZERO, ONE, ZERO);
check_perm_segment(&trace.poseidon2, 0, &init_state, ONE);
}
#[test]
fn hasher_permute_two() {
let mut hasher = Hasher::default();
let init_state1: HasherState = rand_array();
let init_state2: HasherState = rand_array();
let (addr1, final_state1) = hasher.permute(init_state1);
let (addr2, final_state2) = hasher.permute(init_state2);
assert_eq!(ONE, addr1);
assert_eq!(Felt::from_u8(3), addr2);
assert_eq!(apply_permutation(init_state1), final_state1);
assert_eq!(apply_permutation(init_state2), final_state2);
let trace = build_trace(hasher);
assert_eq!(trace.controller.len(), controller_len(4));
assert_eq!(trace.poseidon2.len(), 3 * HASH_CYCLE_LEN);
check_controller_input(&trace.controller, 0, LINEAR_HASH, &init_state1, ZERO, ONE, ZERO, ZERO);
check_controller_output(&trace.controller, 1, RETURN_STATE, &final_state1, ZERO, ONE, ZERO);
check_controller_input(&trace.controller, 2, LINEAR_HASH, &init_state2, ZERO, ONE, ZERO, ZERO);
check_controller_output(&trace.controller, 3, RETURN_STATE, &final_state2, ZERO, ONE, ZERO);
}
#[test]
fn hasher_build_merkle_root_depth_1() {
let leaves = init_leaves(&[1, 2]);
let tree = MerkleTree::new(&leaves).unwrap();
let mut hasher = Hasher::default();
let path0 = tree.get_path(NodeIndex::new(1, 0).unwrap()).unwrap();
let (_, root) = hasher.build_merkle_root(leaves[0], &path0, ZERO);
assert_eq!(root, tree.root());
let trace = build_trace(hasher);
let init_state = init_state_from_words(&leaves[0], &path0[0]);
check_controller_input(&trace.controller, 0, MP_VERIFY, &init_state, ZERO, ONE, ZERO, ZERO);
check_controller_output(
&trace.controller,
1,
RETURN_HASH,
&apply_permutation(init_state),
ZERO,
ONE,
ZERO,
);
}
#[test]
fn hasher_build_merkle_root_depth_3() {
let leaves = init_leaves(&[1, 2, 3, 4, 5, 6, 7, 8]);
let tree = MerkleTree::new(&leaves).unwrap();
let mut hasher = Hasher::default();
let path = tree.get_path(NodeIndex::new(3, 5).unwrap()).unwrap();
let (_, root) = hasher.build_merkle_root(leaves[5], &path, Felt::from_u8(5));
assert_eq!(root, tree.root());
let trace = build_trace(hasher);
check_merkle_controller_pair(&trace.controller, 0, MP_VERIFY, 5, true, false, ZERO, ONE, ZERO);
check_merkle_controller_pair(&trace.controller, 2, MP_VERIFY, 2, false, false, ZERO, ZERO, ONE);
check_merkle_controller_pair(&trace.controller, 4, MP_VERIFY, 1, false, true, ZERO, ONE, ZERO);
for row in [0, 2, 4] {
for (i, value) in controller_row(&trace.controller, row).capacity().into_iter().enumerate()
{
assert_eq!(value, ZERO, "capacity[{i}] should be zero on tree input row {row}");
}
}
}
#[test]
fn hasher_update_merkle_root() {
let leaves = init_leaves(&[1, 2, 3, 4]);
let tree = MerkleTree::new(&leaves).unwrap();
let mut hasher = Hasher::default();
let index = 1u64;
let path = tree.get_path(NodeIndex::new(2, index).unwrap()).unwrap();
let new_leaf: Digest = [Felt::from_u8(100), ZERO, ZERO, ZERO].into();
let update = hasher.update_merkle_root(
leaves[index as usize],
new_leaf,
&path,
Felt::new_unchecked(index),
);
assert_eq!(update.get_old_root(), tree.root());
let trace = build_trace(hasher);
check_merkle_controller_pair(
&trace.controller,
0,
MR_UPDATE_OLD,
1,
true,
false,
ONE,
ONE,
ZERO,
);
check_merkle_controller_pair(
&trace.controller,
2,
MR_UPDATE_OLD,
0,
false,
true,
ONE,
ZERO,
ZERO,
);
check_merkle_controller_pair(
&trace.controller,
4,
MR_UPDATE_NEW,
1,
true,
false,
ONE,
ONE,
ZERO,
);
check_merkle_controller_pair(
&trace.controller,
6,
MR_UPDATE_NEW,
0,
false,
true,
ONE,
ZERO,
ZERO,
);
}
#[test]
fn poseidon2_trace_structure() {
let mut hasher = Hasher::default();
let init_state: HasherState = rand_array();
let (addr, result) = hasher.permute(init_state);
assert_eq!(addr, ONE, "first permutation should start at address 1");
assert_eq!(result, apply_permutation(init_state), "permuted state should match");
let trace = build_trace(hasher);
assert_eq!(trace.poseidon2.len(), 2 * HASH_CYCLE_LEN);
let perm_start = 0;
for offset in [1, 2, 3, 12, 13, 14] {
let row = perm_start + offset;
let cols = poseidon2_row(&trace.poseidon2, row);
assert_eq!(cols.witnesses[0], ZERO, "perm row {row}: witness 0 should be zero");
assert_eq!(cols.witnesses[1], ZERO, "perm row {row}: witness 1 should be zero");
assert_eq!(cols.witnesses[2], ZERO, "perm row {row}: witness 2 should be zero");
}
for offset in [0, 15] {
let row = perm_start + offset;
let cols = poseidon2_row(&trace.poseidon2, row);
assert_eq!(cols.witnesses[0], ONE, "perm row {row}: multiplicity mismatch");
assert_eq!(cols.witnesses[1], ZERO, "perm row {row}: witness 1 should be zero");
assert_eq!(cols.witnesses[2], ZERO, "perm row {row}: witness 2 should be zero");
}
let row_11 = perm_start + 11;
let row_11_cols = poseidon2_row(&trace.poseidon2, row_11);
assert_eq!(
row_11_cols.witnesses[1], ZERO,
"perm row {row_11}: witness 1 should be zero on int+ext row"
);
assert_eq!(
row_11_cols.witnesses[2], ZERO,
"perm row {row_11}: witness 2 should be zero on int+ext row"
);
assert_eq!(poseidon2_row(&trace.poseidon2, perm_start).witnesses[0], ONE);
assert_eq!(
poseidon2_row(&trace.poseidon2, HASH_CYCLE_LEN).witnesses[0],
ZERO,
"the final Poseidon2 cycle closes the accumulator"
);
assert_eq!(poseidon2_row(&trace.poseidon2, perm_start).perm_id, ZERO);
assert_eq!(poseidon2_row(&trace.poseidon2, HASH_CYCLE_LEN).perm_id, ONE);
}
#[test]
fn poseidon2_trace_deduplication() {
let mut hasher = Hasher::default();
let init_state: HasherState = rand_array();
let (addr1, result1) = hasher.permute(init_state);
let (addr2, result2) = hasher.permute(init_state);
assert_eq!(result1, result2, "same input should produce same output");
assert_ne!(addr1, addr2, "second call should have a different address");
let trace = build_trace(hasher);
assert_eq!(trace.controller.len(), controller_len(4));
assert_eq!(trace.poseidon2.len(), 2 * HASH_CYCLE_LEN);
assert_eq!(poseidon2_row(&trace.poseidon2, 0).witnesses[0], Felt::from_u8(2));
}
#[test]
fn hash_memoization_control_blocks() {
let h1: Digest = rand_array::<Felt, 4>().into();
let h2: Digest = rand_array::<Felt, 4>().into();
let domain = Felt::from_u8(7);
let state = super::init_state_from_words_with_domain(&h1, &h2, domain);
let permuted = apply_permutation(state);
let expected_hash: Digest = get_digest(&permuted);
let mut hasher = Hasher::default();
let (addr1, digest1) = hasher.hash_control_block(h1, h2, domain, expected_hash);
let (addr2, digest2) = hasher.hash_control_block(h1, h2, domain, expected_hash);
assert_eq!(digest1, digest2);
assert_eq!(digest1, expected_hash);
assert_ne!(addr1, addr2);
let trace = build_trace(hasher);
assert_eq!(trace.controller.len(), controller_len(4));
assert_eq!(trace.poseidon2.len(), 2 * HASH_CYCLE_LEN);
assert_eq!(poseidon2_row(&trace.poseidon2, 0).witnesses[0], Felt::from_u8(2));
}
#[test]
fn hash_memoization_basic_blocks_single_batch() {
let mut hasher = Hasher::default();
let batches = make_single_batch();
let expected_hash = compute_basic_block_hash(&batches);
let (addr1, digest1) = hasher.hash_basic_block(&batch_groups(&batches), expected_hash);
let (addr2, digest2) = hasher.hash_basic_block(&batch_groups(&batches), expected_hash);
assert_eq!(digest1, digest2, "memoized digest should match original");
assert_eq!(digest1, expected_hash);
assert_ne!(addr1, addr2, "memoized call should have a different address");
let trace = build_trace(hasher);
assert_eq!(trace.controller.len(), controller_len(4));
assert_eq!(trace.poseidon2.len(), 2 * HASH_CYCLE_LEN);
check_controller_input(
&trace.controller,
0,
LINEAR_HASH,
&init_state(batches[0].groups(), ZERO),
ZERO,
ONE,
ZERO,
ZERO,
);
check_controller_output(
&trace.controller,
1,
RETURN_HASH,
&apply_permutation(init_state(batches[0].groups(), ZERO)),
ZERO,
ONE,
ZERO,
);
check_memoized_trace(&trace.controller, 0..2, 2..4);
assert_eq!(poseidon2_row(&trace.poseidon2, 0).witnesses[0], Felt::from_u8(2));
}
#[test]
fn hash_memoization_basic_blocks_multi_batch() {
let mut hasher = Hasher::default();
let batches = make_multi_batch(3);
let expected_hash = compute_basic_block_hash(&batches);
let (addr1, digest1) = hasher.hash_basic_block(&batch_groups(&batches), expected_hash);
let (addr2, digest2) = hasher.hash_basic_block(&batch_groups(&batches), expected_hash);
assert_eq!(digest1, digest2);
assert_eq!(digest1, expected_hash);
assert_ne!(addr1, addr2);
let trace = build_trace(hasher);
assert_eq!(trace.controller.len(), controller_len(12));
assert_eq!(trace.poseidon2.len(), 4 * HASH_CYCLE_LEN);
assert_eq!(controller_row(&trace.controller, 0).is_boundary, ONE);
assert_eq!(controller_row(&trace.controller, 0).direction_bit, ZERO);
assert_eq!(controller_row(&trace.controller, 1).is_boundary, ZERO);
assert_eq!(controller_row(&trace.controller, 1).direction_bit, ZERO);
assert_eq!(controller_row(&trace.controller, 2).is_boundary, ZERO);
assert_eq!(controller_row(&trace.controller, 4).is_boundary, ZERO);
assert_eq!(controller_row(&trace.controller, 5).is_boundary, ONE);
check_memoized_trace(&trace.controller, 0..6, 6..12);
for i in 0..3 {
let cycle_start = i * HASH_CYCLE_LEN;
assert_eq!(
poseidon2_row(&trace.poseidon2, cycle_start).witnesses[0],
Felt::from_u8(2),
"perm cycle {i} should have multiplicity 2"
);
}
}
#[test]
fn hash_memoization_basic_blocks_check() {
let mut hasher = Hasher::default();
let batches = make_multi_batch(2);
let bb_hash = compute_basic_block_hash(&batches);
let loop_body_batches = make_single_batch();
let loop_body_hash = compute_basic_block_hash(&loop_body_batches);
let (bb1_addr, bb1_digest) = hasher.hash_basic_block(&batch_groups(&batches), bb_hash);
assert_eq!(bb1_digest, bb_hash);
let (_loop_addr, loop_digest) =
hasher.hash_basic_block(&batch_groups(&loop_body_batches), loop_body_hash);
assert_eq!(loop_digest, loop_body_hash);
let join2_state =
super::init_state_from_words_with_domain(&bb1_digest, &loop_digest, Felt::from_u8(7));
let join2_permuted = apply_permutation(join2_state);
let join2_hash = get_digest(&join2_permuted);
let (_join2_addr, join2_digest) =
hasher.hash_control_block(bb1_digest, loop_digest, Felt::from_u8(7), join2_hash);
assert_eq!(join2_digest, join2_hash);
let (bb2_addr, bb2_digest) = hasher.hash_basic_block(&batch_groups(&batches), bb_hash);
assert_eq!(bb2_digest, bb_hash);
assert_ne!(bb1_addr, bb2_addr, "memoized BB2 should have a different address");
let join1_state =
super::init_state_from_words_with_domain(&join2_digest, &bb2_digest, Felt::from_u8(7));
let join1_permuted = apply_permutation(join1_state);
let join1_hash = get_digest(&join1_permuted);
let (_join1_addr, join1_digest) =
hasher.hash_control_block(join2_digest, bb2_digest, Felt::from_u8(7), join1_hash);
assert_eq!(join1_digest, join1_hash);
let trace = build_trace(hasher);
let bb1_start = bb1_addr.as_canonical_u64() as usize - 1;
let bb2_start = bb2_addr.as_canonical_u64() as usize - 1;
check_memoized_trace(&trace.controller, bb1_start..bb1_start + 4, bb2_start..bb2_start + 4);
let num_perm_cycles = trace.poseidon2.len() / HASH_CYCLE_LEN;
assert!(num_perm_cycles >= 5, "expected at least 5 perm cycles, got {num_perm_cycles}");
let mut mult_2_count = 0;
let mut mult_1_count = 0;
for i in 0..num_perm_cycles {
let cycle_start = i * HASH_CYCLE_LEN;
let mult = poseidon2_row(&trace.poseidon2, cycle_start).witnesses[0];
if mult == Felt::from_u8(2) {
mult_2_count += 1;
} else if mult == ONE {
mult_1_count += 1;
}
}
assert_eq!(mult_2_count, 2, "expected 2 perm cycles with multiplicity 2 (BB1's states)");
assert_eq!(mult_1_count, 3, "expected 3 perm cycles with multiplicity 1");
}
struct HasherTestTrace {
controller: Vec<[Felt; TRACE_WIDTH]>,
poseidon2: Vec<[Felt; NUM_POSEIDON2_PERMUTATION_COLS]>,
}
fn controller_len(controller_rows: usize) -> usize {
controller_rows.next_multiple_of(CONTROLLER_TRACE_ALIGNMENT)
}
fn build_trace(hasher: Hasher) -> HasherTestTrace {
let trace_len = hasher.trace_len();
let mut band = Felt::zero_vec(TRACE_WIDTH * trace_len);
let mut fragment = ChipletTraceFragment::row_major(&mut band, TRACE_WIDTH, 0, TRACE_WIDTH);
let poseidon2_len = hasher.poseidon2_permutation_trace_len();
let mut poseidon2_band = Felt::zero_vec(NUM_POSEIDON2_PERMUTATION_COLS * poseidon2_len);
hasher.fill_trace(&mut fragment, &mut poseidon2_band);
let (controller, controller_remainder) = band.as_chunks::<TRACE_WIDTH>();
debug_assert!(controller_remainder.is_empty());
let (poseidon2, poseidon2_remainder) =
poseidon2_band.as_chunks::<NUM_POSEIDON2_PERMUTATION_COLS>();
debug_assert!(poseidon2_remainder.is_empty());
HasherTestTrace {
controller: controller.to_vec(),
poseidon2: poseidon2.to_vec(),
}
}
fn controller_row(trace: &[[Felt; TRACE_WIDTH]], row: usize) -> &ControllerCols<Felt> {
trace[row][..].borrow()
}
fn poseidon2_row(
trace: &[[Felt; NUM_POSEIDON2_PERMUTATION_COLS]],
row: usize,
) -> &Poseidon2PermutationCols<Felt> {
trace[row][..].borrow()
}
fn check_controller_input(
trace: &[[Felt; TRACE_WIDTH]],
row: usize,
selectors: Selectors,
state: &HasherState,
node_index: Felt,
is_boundary: Felt,
mrupdate_id: Felt,
direction_bit: Felt,
) {
let cols = controller_row(trace, row);
assert_eq!([cols.s0, cols.s1, cols.s2], selectors, "selectors at row {row}");
assert_eq!(cols.state, *state, "state at row {row}");
assert_eq!(cols.node_index, node_index, "node_index at row {row}");
assert_eq!(cols.is_boundary, is_boundary, "is_boundary at row {row}");
assert_eq!(cols.direction_bit, direction_bit, "direction_bit at row {row}");
assert_eq!(cols.mrupdate_id, mrupdate_id, "mrupdate_id at row {row}");
}
fn check_controller_output(
trace: &[[Felt; TRACE_WIDTH]],
row: usize,
selectors: Selectors,
state: &HasherState,
node_index: Felt,
is_boundary: Felt,
direction_bit: Felt,
) {
let cols = controller_row(trace, row);
let input_cols = controller_row(trace, row - 1);
assert_eq!([cols.s0, cols.s1, cols.s2], selectors, "selectors at row {row}");
assert_eq!(cols.state, *state, "state at row {row}");
assert_eq!(cols.node_index, node_index, "node_index at row {row}");
assert_eq!(cols.is_boundary, is_boundary, "is_boundary at row {row}");
assert_eq!(cols.direction_bit, direction_bit, "direction_bit at row {row}");
assert_eq!(
cols.perm_id,
input_cols.perm_id,
"perm_id mismatch between controller rows {} and {row}",
row - 1
);
}
fn check_merkle_controller_pair(
trace: &[[Felt; TRACE_WIDTH]],
input_row: usize,
input_selectors: Selectors,
node_index: u64,
is_boundary_input: bool,
is_boundary_output: bool,
mrupdate_id: Felt,
input_direction_bit: Felt,
output_direction_bit: Felt,
) {
let output_row = input_row + 1;
let is_boundary_input_felt = if is_boundary_input { ONE } else { ZERO };
let is_boundary_output_felt = if is_boundary_output { ONE } else { ZERO };
let input_cols = controller_row(trace, input_row);
let output_cols = controller_row(trace, output_row);
assert_eq!(
[input_cols.s0, input_cols.s1, input_cols.s2],
input_selectors,
"selectors at input row {input_row}"
);
assert_eq!(
input_cols.node_index,
Felt::new_unchecked(node_index),
"node_index at input row {input_row}"
);
assert_eq!(
input_cols.is_boundary, is_boundary_input_felt,
"is_boundary at input row {input_row}"
);
assert_eq!(
input_cols.direction_bit, input_direction_bit,
"direction_bit at input row {input_row}"
);
assert_eq!(input_cols.mrupdate_id, mrupdate_id, "mrupdate_id at input row {input_row}");
assert_eq!(
output_cols.node_index,
Felt::new_unchecked(node_index >> 1),
"node_index at output row {output_row}"
);
assert_eq!(
output_cols.is_boundary, is_boundary_output_felt,
"is_boundary at output row {output_row}"
);
assert_eq!(
output_cols.direction_bit, output_direction_bit,
"direction_bit at output row {output_row}"
);
assert_eq!(output_cols.mrupdate_id, mrupdate_id, "mrupdate_id at output row {output_row}");
assert_eq!(
output_cols.perm_id, input_cols.perm_id,
"perm_id mismatch between input row {input_row} and output row {output_row}"
);
}
fn check_perm_segment(
trace: &[[Felt; NUM_POSEIDON2_PERMUTATION_COLS]],
start_row: usize,
init_state: &HasherState,
expected_multiplicity: Felt,
) {
use miden_core::chiplets::hasher::Hasher;
let mut state = *init_state;
let first_row = poseidon2_row(trace, start_row);
assert_eq!(first_row.state, state, "state at perm row {CYCLE_INPUT_ROW} (row {start_row})");
assert_eq!(first_row.witnesses[0], expected_multiplicity);
assert_eq!(
poseidon2_row(trace, start_row + CYCLE_OUTPUT_ROW).witnesses[0],
expected_multiplicity
);
let expected_perm_id =
Felt::new_unchecked((start_row / HASH_CYCLE_LEN).try_into().expect("perm id exceeds u64"));
assert_eq!(first_row.perm_id, expected_perm_id);
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);
check_state_at_row(trace, start_row + INITIAL_EXTERNAL_ROUND_START, &state, "after init+ext1");
for round in INITIAL_EXTERNAL_ROUND_START..INITIAL_EXTERNAL_ROUND_END {
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_INITIAL[round]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
check_state_at_row(
trace,
start_row + round + 1,
&state,
&alloc::format!("after ext{}", round + 1),
);
}
for triple in 0..NUM_PACKED_INTERNAL_ROUND_ROWS {
let base = triple * NUM_SBOX_WITNESSES;
for k in 0..NUM_SBOX_WITNESSES {
state[0] += Hasher::ARK_INT[base + k];
state[0] = state[0].exp_const_u64::<7>();
Hasher::matmul_internal(&mut state, Hasher::MAT_DIAG);
}
check_state_at_row(
trace,
start_row + PACKED_INTERNAL_ROUND_START + 1 + triple,
&state,
&alloc::format!("after int triple {triple}"),
);
}
state[0] += Hasher::ARK_INT[LAST_INTERNAL_ROUND_ARK_IDX];
state[0] = state[0].exp_const_u64::<7>();
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);
check_state_at_row(
trace,
start_row + INTERNAL_PLUS_EXTERNAL_ROW + 1,
&state,
"after int22+ext5",
);
for round in 1..=NUM_TRAILING_EXTERNAL_ROUND_ROWS {
Hasher::add_rc(&mut state, &Hasher::ARK_EXT_TERMINAL[round]);
Hasher::apply_sbox(&mut state);
Hasher::apply_matmul_external(&mut state);
check_state_at_row(
trace,
start_row + INTERNAL_PLUS_EXTERNAL_ROW + 1 + round,
&state,
&alloc::format!("after ext{}", round + 5),
);
}
}
fn check_state_at_row(
trace: &[[Felt; NUM_POSEIDON2_PERMUTATION_COLS]],
row: usize,
state: &HasherState,
label: &str,
) {
assert_eq!(poseidon2_row(trace, row).state, *state, "state at row {row} ({label})");
}
fn apply_permutation(mut state: HasherState) -> HasherState {
hasher::apply_permutation(&mut state);
state
}
fn init_leaves(values: &[u64]) -> Vec<Digest> {
values.iter().map(|&v| init_leaf(v)).collect()
}
fn init_leaf(value: u64) -> Digest {
[Felt::new_unchecked(value), ZERO, ZERO, ZERO].into()
}
fn check_memoized_trace(
trace: &[[Felt; TRACE_WIDTH]],
original: Range<usize>,
copied: Range<usize>,
) {
assert_eq!(
original.len(),
copied.len(),
"original and copied ranges must have the same length"
);
for (orig_row, copy_row) in original.zip(copied) {
let original = controller_row(trace, orig_row);
let copied = controller_row(trace, copy_row);
assert_eq!(
[original.s0, original.s1, original.s2],
[copied.s0, copied.s1, copied.s2],
"selector mismatch: original row {orig_row} vs copied row {copy_row}"
);
assert_eq!(
original.state, copied.state,
"state mismatch: original row {orig_row} vs copied row {copy_row}"
);
assert_eq!(
original.node_index, copied.node_index,
"node_index mismatch: original row {orig_row} vs copied row {copy_row}"
);
assert_eq!(
original.is_boundary, copied.is_boundary,
"is_boundary mismatch: original row {orig_row} vs copied row {copy_row}"
);
assert_eq!(
original.direction_bit, copied.direction_bit,
"direction_bit mismatch: original row {orig_row} vs copied row {copy_row}"
);
assert_eq!(
original.perm_id, copied.perm_id,
"perm_id mismatch: original row {orig_row} vs copied row {copy_row}"
);
}
}
fn make_basic_block_batches(ops: Vec<miden_core::operations::Operation>) -> Vec<OpBatch> {
use miden_core::mast::BasicBlockNodeBuilder;
let node = BasicBlockNodeBuilder::new(ops).build().expect("failed to build basic block");
node.op_batches().to_vec()
}
fn batch_groups(batches: &[OpBatch]) -> Vec<[Felt; miden_air::trace::chiplets::hasher::RATE_LEN]> {
batches.iter().map(|batch| *batch.groups()).collect()
}
fn make_single_batch() -> Vec<OpBatch> {
use miden_core::operations::Operation;
make_basic_block_batches(vec![Operation::Pad])
}
fn make_multi_batch(n: usize) -> Vec<OpBatch> {
use miden_core::operations::Operation;
assert!(n >= 2, "use make_single_batch for n=1");
let num_ops = 72 * (n - 1) + 1;
let ops = vec![Operation::Noop; num_ops];
let batches = make_basic_block_batches(ops);
assert_eq!(batches.len(), n, "expected exactly {n} batches, got {}", batches.len());
batches
}
fn compute_basic_block_hash(batches: &[OpBatch]) -> Digest {
assert!(!batches.is_empty());
let mut state = init_state(batches[0].groups(), ZERO);
hasher::apply_permutation(&mut state);
for batch in batches.iter().skip(1) {
absorb_into_state(&mut state, batch.groups());
hasher::apply_permutation(&mut state);
}
get_digest(&state)
}