use std::{borrow::BorrowMut, sync::atomic::AtomicU32};
use openvm_cpu_backend::CpuBackend;
use openvm_stark_backend::{
p3_field::PrimeField32, p3_matrix::dense::RowMajorMatrix, prover::AirProvingContext,
StarkProtocolConfig, Val,
};
use tracing::instrument;
use crate::{
arch::{hasher::HasherChip, VmField},
system::{
memory::{
merkle::{tree::MerkleTree, FinalState, MemoryMerkleChip, MemoryMerkleCols},
Equipartition, MemoryImage,
},
poseidon2::{
Poseidon2PeripheryBaseChip, Poseidon2PeripheryChip, PERIPHERY_POSEIDON2_WIDTH,
},
},
};
impl<const CHUNK: usize, F: PrimeField32> MemoryMerkleChip<CHUNK, F> {
#[instrument(name = "merkle_finalize", level = "debug", skip_all)]
pub(crate) fn finalize(
&mut self,
initial_memory: &MemoryImage,
final_memory: &Equipartition<F, CHUNK>,
hasher: &impl HasherChip<CHUNK, F>,
) {
assert!(self.final_state.is_none(), "Merkle chip already finalized");
let memory_dimensions = &self.air.memory_dimensions;
let mut tree = MerkleTree::from_memory(initial_memory, memory_dimensions, hasher);
self.final_state = Some(tree.finalize(hasher, final_memory, memory_dimensions));
self.top_tree = tree.top_tree(memory_dimensions.addr_space_height);
}
}
impl<const CHUNK: usize, F> MemoryMerkleChip<CHUNK, F>
where
F: PrimeField32,
{
pub fn generate_proving_ctx<SC>(&mut self) -> AirProvingContext<CpuBackend<SC>>
where
SC: StarkProtocolConfig<F = F>,
{
assert!(
self.final_state.is_some(),
"Merkle chip must finalize before trace generation"
);
let FinalState {
mut rows,
init_root,
final_root,
} = self.final_state.take().unwrap();
rows.reverse();
rows.swap(0, 1);
#[cfg(feature = "metrics")]
{
self.current_height = rows.len();
}
let width = MemoryMerkleCols::<Val<SC>, CHUNK>::width();
let mut height = rows.len().next_power_of_two();
if let Some(mut oh) = self.overridden_height {
oh = oh.next_power_of_two();
assert!(
oh >= height,
"Overridden height {oh} is less than the required height {height}"
);
height = oh;
}
let mut trace = Val::<SC>::zero_vec(width * height);
for (trace_row, row) in trace.chunks_exact_mut(width).zip(rows) {
*trace_row.borrow_mut() = row;
}
let trace = RowMajorMatrix::new(trace, width);
let pvs = init_root.into_iter().chain(final_root).collect();
AirProvingContext::simple(trace, pvs)
}
}
pub trait SerialReceiver<T> {
fn receive(&self, msg: T);
}
impl<'a, F: VmField, const SBOX_REGISTERS: usize> SerialReceiver<&'a [F]>
for Poseidon2PeripheryBaseChip<F, SBOX_REGISTERS>
{
fn receive(&self, perm_preimage: &'a [F]) {
assert!(perm_preimage.len() <= PERIPHERY_POSEIDON2_WIDTH);
let mut state = [F::ZERO; PERIPHERY_POSEIDON2_WIDTH];
state[..perm_preimage.len()].copy_from_slice(perm_preimage);
let count = self.records.entry(state).or_insert(AtomicU32::new(0));
count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.nonempty
.store(true, std::sync::atomic::Ordering::Relaxed);
}
}
impl<'a, F: VmField> SerialReceiver<&'a [F]> for Poseidon2PeripheryChip<F> {
fn receive(&self, perm_preimage: &'a [F]) {
match self {
Poseidon2PeripheryChip::Register0(chip) => chip.receive(perm_preimage),
Poseidon2PeripheryChip::Register1(chip) => chip.receive(perm_preimage),
}
}
}