use openvm_stark_backend::{
hasher::MerkleHasher,
prover::{error::StackedPcsError, ColMajorMatrix, MatrixDimensions},
};
use p3_baby_bear::BabyBear;
use p3_dft::TwoAdicSubgroupDft;
use p3_field::TwoAdicField;
use p3_matrix::dense::RowMajorMatrix;
use p3_maybe_rayon::prelude::*;
use p3_util::log2_strict_usize;
use tracing::instrument;
use crate::{device::eval_to_coeff_cpu, two_adic::DftTwiddles};
#[derive(Clone, Debug)]
pub struct CpuMerkleTree<F, Digest> {
pub(crate) backing_matrix: RowMajorMatrix<F>,
pub(crate) digest_layers: Vec<Vec<Digest>>,
pub(crate) rows_per_query: usize,
}
impl<F, Digest> CpuMerkleTree<F, Digest> {
pub unsafe fn from_raw_parts(
backing_matrix: RowMajorMatrix<F>,
digest_layers: Vec<Vec<Digest>>,
rows_per_query: usize,
) -> Self {
Self {
backing_matrix,
digest_layers,
rows_per_query,
}
}
pub fn backing_matrix(&self) -> &RowMajorMatrix<F> {
&self.backing_matrix
}
pub fn digest_layers(&self) -> &Vec<Vec<Digest>> {
&self.digest_layers
}
pub fn rows_per_query(&self) -> usize {
self.rows_per_query
}
pub fn query_stride(&self) -> usize {
self.digest_layers[0].len()
}
pub fn proof_depth(&self) -> usize {
self.digest_layers.len() - 1
}
}
impl<F, Digest: Clone> CpuMerkleTree<F, Digest> {
pub fn root(&self) -> Result<Digest, StackedPcsError> {
Ok(self
.digest_layers
.last()
.ok_or(StackedPcsError::MerkleTreeNoRoot)?[0]
.clone())
}
pub fn query_merkle_proof(&self, query_idx: usize) -> Result<Vec<Digest>, StackedPcsError> {
let stride = self.query_stride();
if query_idx >= stride {
return Err(StackedPcsError::MerkleTreeQueryOutOfBounds {
query_idx,
query_stride: stride,
});
}
let mut idx = query_idx;
let mut proof = Vec::with_capacity(self.proof_depth());
for layer in self.digest_layers.iter().take(self.proof_depth()) {
let sibling = layer[idx ^ 1].clone();
proof.push(sibling);
idx >>= 1;
}
Ok(proof)
}
}
impl<F: Copy, Digest> CpuMerkleTree<F, Digest> {
pub fn get_opened_rows(&self, index: usize) -> Result<Vec<Vec<F>>, StackedPcsError> {
let query_stride = self.query_stride();
if index >= query_stride {
return Err(StackedPcsError::MerkleTreeOpenedRowsOutOfBounds {
index,
query_stride,
});
}
let width = self.backing_matrix.width;
let height = self.backing_matrix.values.len() / width;
let mut rows = Vec::with_capacity(self.rows_per_query);
for t in 0..self.rows_per_query {
let row_idx = t * query_stride + index;
if row_idx < height {
let start = row_idx * width;
rows.push(self.backing_matrix.values[start..start + width].to_vec());
} else {
rows.push(vec![]);
}
}
Ok(rows)
}
}
pub(crate) unsafe fn reinterpret_vec<A, B>(v: Vec<A>) -> Vec<B> {
debug_assert_eq!(std::mem::size_of::<A>(), std::mem::size_of::<B>());
debug_assert_eq!(std::mem::align_of::<A>(), std::mem::align_of::<B>());
let mut md = std::mem::ManuallyDrop::new(v);
Vec::from_raw_parts(md.as_mut_ptr().cast::<B>(), md.len(), md.capacity())
}
pub(crate) fn hash_rows_packed_babybear(
rm_vals: &[BabyBear],
width: usize,
codeword_height: usize,
num_leaves: usize,
) -> Vec<[BabyBear; 8]> {
use openvm_stark_backend::p3_symmetric::{CryptographicHasher, PaddingFreeSponge};
use p3_baby_bear::default_babybear_poseidon2_16;
use p3_field::{Field, PackedValue, PrimeCharacteristicRing};
type P = <BabyBear as Field>::Packing;
let perm = default_babybear_poseidon2_16();
let sponge = PaddingFreeSponge::<_, 16, 8, 8>::new(perm);
let pack_width = P::WIDTH;
let mut digests = vec![[BabyBear::ZERO; 8]; num_leaves];
digests
.par_chunks_mut(pack_width)
.enumerate()
.for_each(|(chunk_idx, digest_chunk)| {
let base_row = chunk_idx * pack_width;
if digest_chunk.len() == pack_width {
let packed_row: Vec<P> = (0..width)
.map(|col| {
P::from_fn(|lane| {
let row = base_row + lane;
if row < codeword_height {
rm_vals[row * width + col]
} else {
BabyBear::ZERO
}
})
})
.collect();
let packed_digest: [P; 8] = sponge.hash_slice(&packed_row);
for lane in 0..pack_width {
for d in 0..8 {
digest_chunk[lane][d] = packed_digest[d].as_slice()[lane];
}
}
} else {
for (lane, digest) in digest_chunk.iter_mut().enumerate() {
let row = base_row + lane;
if row < codeword_height {
*digest = sponge.hash_slice(&rm_vals[row * width..(row + 1) * width]);
}
}
}
});
digests
}
pub(crate) fn build_digest_layers<F, H>(
row_hashes: Vec<H::Digest>,
rows_per_query: usize,
hasher: &H,
) -> Vec<Vec<H::Digest>>
where
F: TwoAdicField + Ord + 'static,
H: MerkleHasher<F = F>,
{
use std::any::TypeId;
if TypeId::of::<F>() == TypeId::of::<BabyBear>()
&& TypeId::of::<H::Digest>() == TypeId::of::<[BabyBear; 8]>()
{
let bb_hashes: Vec<[BabyBear; 8]> = unsafe { reinterpret_vec(row_hashes) };
let bb_layers = build_digest_layers_packed_babybear(bb_hashes, rows_per_query);
bb_layers
.into_iter()
.map(|layer| unsafe { reinterpret_vec(layer) })
.collect()
} else {
build_digest_layers_scalar(row_hashes, rows_per_query, hasher)
}
}
pub(crate) fn hash_rows_with_padding<D, RowHashFn, PaddingHashFn>(
num_leaves: usize,
codeword_height: usize,
row_hash_fn: RowHashFn,
padding_hash_fn: PaddingHashFn,
) -> Vec<D>
where
D: Send,
RowHashFn: Fn(usize) -> D + Sync + Send,
PaddingHashFn: Fn() -> D + Sync + Send,
{
(0..num_leaves)
.into_par_iter()
.map(|r| {
if r < codeword_height {
row_hash_fn(r)
} else {
padding_hash_fn()
}
})
.collect()
}
#[instrument(name = "rs_encode_and_merkle_cpu", skip_all)]
pub(crate) fn rs_encode_and_merkle_cpu<F, H>(
hasher: &H,
l_skip: usize,
log_blowup: usize,
eval_matrix: &ColMajorMatrix<F>,
rows_per_query: usize,
) -> CpuMerkleTree<F, H::Digest>
where
F: TwoAdicField + Ord + 'static,
H: MerkleHasher<F = F>,
{
use p3_dft::Radix2DitParallel;
use p3_matrix::dense::RowMajorMatrix as P3RowMajorMatrix;
let height = eval_matrix.height();
let codeword_height = height.checked_shl(log_blowup as u32).unwrap();
let width = eval_matrix.width();
let twiddles = DftTwiddles::new(l_skip);
let coeff_vecs: Vec<Vec<F>> = tracing::info_span!("eval_to_coeff_phase").in_scope(|| {
eval_matrix
.values
.par_chunks_exact(height)
.map(|column_evals| {
let mut coeffs = eval_to_coeff_cpu(column_evals, &twiddles);
coeffs.resize(codeword_height, F::ZERO);
coeffs
})
.collect()
});
let rm_mat: P3RowMajorMatrix<F> = tracing::info_span!("transpose_to_rm").in_scope(|| {
let mut rm_values = F::zero_vec(codeword_height * width);
rm_values
.par_chunks_exact_mut(width)
.enumerate()
.for_each(|(i, row)| {
for (j, col) in coeff_vecs.iter().enumerate() {
row[j] = col[i];
}
});
P3RowMajorMatrix::new(rm_values, width)
});
drop(coeff_vecs);
let rm_result = tracing::info_span!("dft_batch").in_scope(|| {
use p3_matrix::Matrix as _;
Radix2DitParallel::default()
.dft_batch(rm_mat)
.to_row_major_matrix()
});
let num_leaves = codeword_height.next_power_of_two();
let rm_vals = &rm_result.values;
let row_hashes: Vec<H::Digest> = tracing::info_span!("row_hash").in_scope(|| {
use std::any::TypeId;
if TypeId::of::<F>() == TypeId::of::<BabyBear>()
&& TypeId::of::<H::Digest>() == TypeId::of::<[BabyBear; 8]>()
{
let bb_vals: &[BabyBear] = unsafe {
std::slice::from_raw_parts(rm_vals.as_ptr().cast::<BabyBear>(), rm_vals.len())
};
let bb_digests = hash_rows_packed_babybear(bb_vals, width, codeword_height, num_leaves);
unsafe { reinterpret_vec(bb_digests) }
} else {
let zero_row = vec![F::ZERO; width];
hash_rows_with_padding(
num_leaves,
codeword_height,
|r| hasher.hash_slice(&rm_vals[r * width..(r + 1) * width]),
|| hasher.hash_slice(&zero_row),
)
}
});
let digest_layers = tracing::info_span!("digest_layers")
.in_scope(|| build_digest_layers::<F, H>(row_hashes, rows_per_query, hasher));
unsafe { CpuMerkleTree::from_raw_parts(rm_result, digest_layers, rows_per_query) }
}
fn build_digest_layers_scalar<H: MerkleHasher>(
row_hashes: Vec<H::Digest>,
rows_per_query: usize,
hasher: &H,
) -> Vec<Vec<H::Digest>> {
let num_leaves = row_hashes.len();
let query_stride = num_leaves / rows_per_query;
let mut query_digest_layer = row_hashes;
for _ in 0..log2_strict_usize(rows_per_query) {
let prev_layer = query_digest_layer;
query_digest_layer = (0..prev_layer.len() / 2)
.into_par_iter()
.map(|i| {
let x = i / query_stride;
let y = i % query_stride;
let left = prev_layer[2 * x * query_stride + y];
let right = prev_layer[(2 * x + 1) * query_stride + y];
hasher.compress(left, right)
})
.collect();
}
let mut layers = vec![query_digest_layer];
while layers.last().unwrap().len() > 1 {
let prev = layers.last().unwrap();
let layer: Vec<_> = prev
.par_chunks_exact(2)
.map(|pair| hasher.compress(pair[0], pair[1]))
.collect();
layers.push(layer);
}
layers
}
fn build_digest_layers_packed_babybear(
row_hashes: Vec<[BabyBear; 8]>,
rows_per_query: usize,
) -> Vec<Vec<[BabyBear; 8]>> {
use openvm_stark_backend::p3_symmetric::{PseudoCompressionFunction, TruncatedPermutation};
use p3_baby_bear::default_babybear_poseidon2_16;
use p3_field::{Field, PackedValue, PrimeCharacteristicRing};
type P = <BabyBear as Field>::Packing;
let pack_width = P::WIDTH;
let perm = default_babybear_poseidon2_16();
let compressor = TruncatedPermutation::<_, 2, 8, 16>::new(perm);
let num_leaves = row_hashes.len();
let query_stride = num_leaves / rows_per_query;
let mut prev_layer = row_hashes;
for _ in 0..log2_strict_usize(rows_per_query) {
let n = prev_layer.len() / 2;
let qs = query_stride;
let mut next_layer = vec![[BabyBear::ZERO; 8]; n];
next_layer
.par_chunks_mut(pack_width)
.enumerate()
.for_each(|(chunk_idx, out_chunk)| {
let base = chunk_idx * pack_width;
let actual = out_chunk.len();
if actual == pack_width {
let mut packed_input: [[P; 8]; 2] = [[P::default(); 8]; 2];
for d in 0..8 {
packed_input[0][d] = P::from_fn(|lane| {
let i = base + lane;
let x = i / qs;
let y = i % qs;
prev_layer[2 * x * qs + y][d]
});
packed_input[1][d] = P::from_fn(|lane| {
let i = base + lane;
let x = i / qs;
let y = i % qs;
prev_layer[(2 * x + 1) * qs + y][d]
});
}
let packed_result: [P; 8] = compressor.compress(packed_input);
for lane in 0..pack_width {
for d in 0..8 {
out_chunk[lane][d] = packed_result[d].as_slice()[lane];
}
}
} else {
for lane in 0..actual {
let i = base + lane;
let x = i / qs;
let y = i % qs;
out_chunk[lane] = compressor.compress([
prev_layer[2 * x * qs + y],
prev_layer[(2 * x + 1) * qs + y],
]);
}
}
});
prev_layer = next_layer;
}
let mut layers = vec![prev_layer];
while layers.last().unwrap().len() > 1 {
let n = layers.last().unwrap().len() / 2;
let mut layer = vec![[BabyBear::ZERO; 8]; n];
{
let prev = layers.last().unwrap();
layer
.par_chunks_mut(pack_width)
.enumerate()
.for_each(|(chunk_idx, out_chunk)| {
let base = chunk_idx * pack_width;
let actual = out_chunk.len();
if actual == pack_width {
let mut packed_input: [[P; 8]; 2] = [[P::default(); 8]; 2];
for d in 0..8 {
packed_input[0][d] = P::from_fn(|lane| prev[2 * (base + lane)][d]);
packed_input[1][d] = P::from_fn(|lane| prev[2 * (base + lane) + 1][d]);
}
let packed_result: [P; 8] = compressor.compress(packed_input);
for lane in 0..pack_width {
for d in 0..8 {
out_chunk[lane][d] = packed_result[d].as_slice()[lane];
}
}
} else {
for lane in 0..actual {
let i = base + lane;
out_chunk[lane] = compressor.compress([prev[2 * i], prev[2 * i + 1]]);
}
}
});
}
layers.push(layer);
}
layers
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_from_raw_parts_and_accessors() {
let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4, 5, 6], 3);
let digest_layers = vec![vec![10u32, 20], vec![30]];
let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
assert_eq!(tree.rows_per_query(), 1);
assert_eq!(tree.backing_matrix().width, 3);
assert_eq!(tree.digest_layers().len(), 2);
assert_eq!(tree.query_stride(), 2);
assert_eq!(tree.proof_depth(), 1);
}
#[test]
fn test_root() {
let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4], 2);
let digest_layers = vec![vec![10u32, 20], vec![42]];
let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
assert_eq!(tree.root().unwrap(), 42);
}
#[test]
fn test_root_no_layers() {
let mat = RowMajorMatrix::new(vec![1u32, 2], 2);
let digest_layers: Vec<Vec<u32>> = vec![];
let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
assert!(tree.root().is_err());
}
#[test]
fn test_query_merkle_proof() {
let mat = RowMajorMatrix::new(vec![0u32; 8], 2);
let layer0 = vec![10u32, 20, 30, 40]; let layer1 = vec![100u32, 200]; let layer2 = vec![999u32]; let digest_layers = vec![layer0, layer1, layer2];
let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
assert_eq!(tree.query_stride(), 4);
assert_eq!(tree.proof_depth(), 2);
let proof = tree.query_merkle_proof(0).unwrap();
assert_eq!(proof, vec![20, 200]);
let proof = tree.query_merkle_proof(1).unwrap();
assert_eq!(proof, vec![10, 200]);
let proof = tree.query_merkle_proof(2).unwrap();
assert_eq!(proof, vec![40, 100]);
assert!(tree.query_merkle_proof(4).is_err());
}
#[test]
fn test_get_opened_rows_single() {
let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], 3);
let digest_layers = vec![vec![0u32; 4], vec![0u32; 2], vec![0u32; 1]];
let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
assert_eq!(tree.query_stride(), 4);
let rows = tree.get_opened_rows(0).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0], vec![1, 2, 3]);
let rows = tree.get_opened_rows(2).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0], vec![7, 8, 9]);
}
#[test]
fn test_get_opened_rows_batched() {
let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4, 5, 6, 7, 8], 2);
let digest_layers = vec![vec![0u32; 2], vec![0u32; 1]];
let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 2) };
assert_eq!(tree.query_stride(), 2);
let rows = tree.get_opened_rows(0).unwrap();
assert_eq!(rows.len(), 2);
assert_eq!(rows[0], vec![1, 2]);
assert_eq!(rows[1], vec![5, 6]);
let rows = tree.get_opened_rows(1).unwrap();
assert_eq!(rows.len(), 2);
assert_eq!(rows[0], vec![3, 4]);
assert_eq!(rows[1], vec![7, 8]);
assert!(tree.get_opened_rows(2).is_err());
}
}