1use openvm_stark_backend::{
8 hasher::MerkleHasher,
9 prover::{error::StackedPcsError, ColMajorMatrix, MatrixDimensions},
10};
11use p3_baby_bear::BabyBear;
12use p3_dft::TwoAdicSubgroupDft;
13use p3_field::TwoAdicField;
14use p3_matrix::dense::RowMajorMatrix;
15use p3_maybe_rayon::prelude::*;
16use p3_util::log2_strict_usize;
17use tracing::instrument;
18
19use crate::{device::eval_to_coeff_cpu, two_adic::DftTwiddles};
20
21#[derive(Clone, Debug)]
28pub struct CpuMerkleTree<F, Digest> {
29 pub(crate) backing_matrix: RowMajorMatrix<F>,
30 pub(crate) digest_layers: Vec<Vec<Digest>>,
31 pub(crate) rows_per_query: usize,
32}
33
34impl<F, Digest> CpuMerkleTree<F, Digest> {
35 pub unsafe fn from_raw_parts(
51 backing_matrix: RowMajorMatrix<F>,
52 digest_layers: Vec<Vec<Digest>>,
53 rows_per_query: usize,
54 ) -> Self {
55 Self {
56 backing_matrix,
57 digest_layers,
58 rows_per_query,
59 }
60 }
61
62 pub fn backing_matrix(&self) -> &RowMajorMatrix<F> {
64 &self.backing_matrix
65 }
66
67 pub fn digest_layers(&self) -> &Vec<Vec<Digest>> {
69 &self.digest_layers
70 }
71
72 pub fn rows_per_query(&self) -> usize {
74 self.rows_per_query
75 }
76
77 pub fn query_stride(&self) -> usize {
79 self.digest_layers[0].len()
80 }
81
82 pub fn proof_depth(&self) -> usize {
84 self.digest_layers.len() - 1
85 }
86}
87
88impl<F, Digest: Clone> CpuMerkleTree<F, Digest> {
89 pub fn root(&self) -> Result<Digest, StackedPcsError> {
91 Ok(self
92 .digest_layers
93 .last()
94 .ok_or(StackedPcsError::MerkleTreeNoRoot)?[0]
95 .clone())
96 }
97
98 pub fn query_merkle_proof(&self, query_idx: usize) -> Result<Vec<Digest>, StackedPcsError> {
100 let stride = self.query_stride();
101 if query_idx >= stride {
102 return Err(StackedPcsError::MerkleTreeQueryOutOfBounds {
103 query_idx,
104 query_stride: stride,
105 });
106 }
107
108 let mut idx = query_idx;
109 let mut proof = Vec::with_capacity(self.proof_depth());
110 for layer in self.digest_layers.iter().take(self.proof_depth()) {
111 let sibling = layer[idx ^ 1].clone();
112 proof.push(sibling);
113 idx >>= 1;
114 }
115 Ok(proof)
116 }
117}
118
119impl<F: Copy, Digest> CpuMerkleTree<F, Digest> {
120 pub fn get_opened_rows(&self, index: usize) -> Result<Vec<Vec<F>>, StackedPcsError> {
127 let query_stride = self.query_stride();
128 if index >= query_stride {
129 return Err(StackedPcsError::MerkleTreeOpenedRowsOutOfBounds {
130 index,
131 query_stride,
132 });
133 }
134
135 let width = self.backing_matrix.width;
136 let height = self.backing_matrix.values.len() / width;
137 let mut rows = Vec::with_capacity(self.rows_per_query);
138 for t in 0..self.rows_per_query {
139 let row_idx = t * query_stride + index;
140 if row_idx < height {
141 let start = row_idx * width;
142 rows.push(self.backing_matrix.values[start..start + width].to_vec());
143 } else {
144 rows.push(vec![]);
146 }
147 }
148 Ok(rows)
149 }
150}
151
152pub(crate) unsafe fn reinterpret_vec<A, B>(v: Vec<A>) -> Vec<B> {
158 debug_assert_eq!(std::mem::size_of::<A>(), std::mem::size_of::<B>());
159 debug_assert_eq!(std::mem::align_of::<A>(), std::mem::align_of::<B>());
160 let mut md = std::mem::ManuallyDrop::new(v);
161 Vec::from_raw_parts(md.as_mut_ptr().cast::<B>(), md.len(), md.capacity())
162}
163
164pub(crate) fn hash_rows_packed_babybear(
172 rm_vals: &[BabyBear],
173 width: usize,
174 codeword_height: usize,
175 num_leaves: usize,
176) -> Vec<[BabyBear; 8]> {
177 use openvm_stark_backend::p3_symmetric::{CryptographicHasher, PaddingFreeSponge};
178 use p3_baby_bear::default_babybear_poseidon2_16;
179 use p3_field::{Field, PackedValue, PrimeCharacteristicRing};
180
181 type P = <BabyBear as Field>::Packing;
182
183 let perm = default_babybear_poseidon2_16();
184 let sponge = PaddingFreeSponge::<_, 16, 8, 8>::new(perm);
185 let pack_width = P::WIDTH;
186
187 let mut digests = vec![[BabyBear::ZERO; 8]; num_leaves];
188
189 digests
190 .par_chunks_mut(pack_width)
191 .enumerate()
192 .for_each(|(chunk_idx, digest_chunk)| {
193 let base_row = chunk_idx * pack_width;
194
195 if digest_chunk.len() == pack_width {
196 let packed_row: Vec<P> = (0..width)
198 .map(|col| {
199 P::from_fn(|lane| {
200 let row = base_row + lane;
201 if row < codeword_height {
202 rm_vals[row * width + col]
203 } else {
204 BabyBear::ZERO
205 }
206 })
207 })
208 .collect();
209
210 let packed_digest: [P; 8] = sponge.hash_slice(&packed_row);
212
213 for lane in 0..pack_width {
215 for d in 0..8 {
216 digest_chunk[lane][d] = packed_digest[d].as_slice()[lane];
217 }
218 }
219 } else {
220 for (lane, digest) in digest_chunk.iter_mut().enumerate() {
222 let row = base_row + lane;
223 if row < codeword_height {
224 *digest = sponge.hash_slice(&rm_vals[row * width..(row + 1) * width]);
225 }
226 }
227 }
228 });
229
230 digests
231}
232
233pub(crate) fn build_digest_layers<F, H>(
239 row_hashes: Vec<H::Digest>,
240 rows_per_query: usize,
241 hasher: &H,
242) -> Vec<Vec<H::Digest>>
243where
244 F: TwoAdicField + Ord + 'static,
245 H: MerkleHasher<F = F>,
246{
247 use std::any::TypeId;
248 if TypeId::of::<F>() == TypeId::of::<BabyBear>()
249 && TypeId::of::<H::Digest>() == TypeId::of::<[BabyBear; 8]>()
250 {
251 let bb_hashes: Vec<[BabyBear; 8]> = unsafe { reinterpret_vec(row_hashes) };
253 let bb_layers = build_digest_layers_packed_babybear(bb_hashes, rows_per_query);
254 bb_layers
255 .into_iter()
256 .map(|layer| unsafe { reinterpret_vec(layer) })
257 .collect()
258 } else {
259 build_digest_layers_scalar(row_hashes, rows_per_query, hasher)
260 }
261}
262
263pub(crate) fn hash_rows_with_padding<D, RowHashFn, PaddingHashFn>(
264 num_leaves: usize,
265 codeword_height: usize,
266 row_hash_fn: RowHashFn,
267 padding_hash_fn: PaddingHashFn,
268) -> Vec<D>
269where
270 D: Send,
271 RowHashFn: Fn(usize) -> D + Sync + Send,
272 PaddingHashFn: Fn() -> D + Sync + Send,
273{
274 (0..num_leaves)
275 .into_par_iter()
276 .map(|r| {
277 if r < codeword_height {
278 row_hash_fn(r)
279 } else {
280 padding_hash_fn()
281 }
282 })
283 .collect()
284}
285
286#[instrument(name = "rs_encode_and_merkle_cpu", skip_all)]
291pub(crate) fn rs_encode_and_merkle_cpu<F, H>(
292 hasher: &H,
293 l_skip: usize,
294 log_blowup: usize,
295 eval_matrix: &ColMajorMatrix<F>,
296 rows_per_query: usize,
297) -> CpuMerkleTree<F, H::Digest>
298where
299 F: TwoAdicField + Ord + 'static,
300 H: MerkleHasher<F = F>,
301{
302 use p3_dft::Radix2DitParallel;
303 use p3_matrix::dense::RowMajorMatrix as P3RowMajorMatrix;
304
305 let height = eval_matrix.height();
306 let codeword_height = height.checked_shl(log_blowup as u32).unwrap();
307 let width = eval_matrix.width();
308 let twiddles = DftTwiddles::new(l_skip);
309
310 let coeff_vecs: Vec<Vec<F>> = tracing::info_span!("eval_to_coeff_phase").in_scope(|| {
312 eval_matrix
313 .values
314 .par_chunks_exact(height)
315 .map(|column_evals| {
316 let mut coeffs = eval_to_coeff_cpu(column_evals, &twiddles);
317 coeffs.resize(codeword_height, F::ZERO);
318 coeffs
319 })
320 .collect()
321 });
322
323 let rm_mat: P3RowMajorMatrix<F> = tracing::info_span!("transpose_to_rm").in_scope(|| {
325 let mut rm_values = F::zero_vec(codeword_height * width);
326 rm_values
327 .par_chunks_exact_mut(width)
328 .enumerate()
329 .for_each(|(i, row)| {
330 for (j, col) in coeff_vecs.iter().enumerate() {
331 row[j] = col[i];
332 }
333 });
334 P3RowMajorMatrix::new(rm_values, width)
335 });
336 drop(coeff_vecs);
337
338 let rm_result = tracing::info_span!("dft_batch").in_scope(|| {
340 use p3_matrix::Matrix as _;
341 Radix2DitParallel::default()
342 .dft_batch(rm_mat)
343 .to_row_major_matrix()
344 });
345
346 let num_leaves = codeword_height.next_power_of_two();
348 let rm_vals = &rm_result.values;
349 let row_hashes: Vec<H::Digest> = tracing::info_span!("row_hash").in_scope(|| {
350 use std::any::TypeId;
351 if TypeId::of::<F>() == TypeId::of::<BabyBear>()
352 && TypeId::of::<H::Digest>() == TypeId::of::<[BabyBear; 8]>()
353 {
354 let bb_vals: &[BabyBear] = unsafe {
355 std::slice::from_raw_parts(rm_vals.as_ptr().cast::<BabyBear>(), rm_vals.len())
356 };
357 let bb_digests = hash_rows_packed_babybear(bb_vals, width, codeword_height, num_leaves);
358 unsafe { reinterpret_vec(bb_digests) }
359 } else {
360 let zero_row = vec![F::ZERO; width];
361 hash_rows_with_padding(
362 num_leaves,
363 codeword_height,
364 |r| hasher.hash_slice(&rm_vals[r * width..(r + 1) * width]),
365 || hasher.hash_slice(&zero_row),
366 )
367 }
368 });
369
370 let digest_layers = tracing::info_span!("digest_layers")
372 .in_scope(|| build_digest_layers::<F, H>(row_hashes, rows_per_query, hasher));
373
374 unsafe { CpuMerkleTree::from_raw_parts(rm_result, digest_layers, rows_per_query) }
381}
382
383fn build_digest_layers_scalar<H: MerkleHasher>(
385 row_hashes: Vec<H::Digest>,
386 rows_per_query: usize,
387 hasher: &H,
388) -> Vec<Vec<H::Digest>> {
389 let num_leaves = row_hashes.len();
390 let query_stride = num_leaves / rows_per_query;
391 let mut query_digest_layer = row_hashes;
392 for _ in 0..log2_strict_usize(rows_per_query) {
393 let prev_layer = query_digest_layer;
394 query_digest_layer = (0..prev_layer.len() / 2)
395 .into_par_iter()
396 .map(|i| {
397 let x = i / query_stride;
398 let y = i % query_stride;
399 let left = prev_layer[2 * x * query_stride + y];
400 let right = prev_layer[(2 * x + 1) * query_stride + y];
401 hasher.compress(left, right)
402 })
403 .collect();
404 }
405 let mut layers = vec![query_digest_layer];
406 while layers.last().unwrap().len() > 1 {
407 let prev = layers.last().unwrap();
408 let layer: Vec<_> = prev
409 .par_chunks_exact(2)
410 .map(|pair| hasher.compress(pair[0], pair[1]))
411 .collect();
412 layers.push(layer);
413 }
414 layers
415}
416
417fn build_digest_layers_packed_babybear(
419 row_hashes: Vec<[BabyBear; 8]>,
420 rows_per_query: usize,
421) -> Vec<Vec<[BabyBear; 8]>> {
422 use openvm_stark_backend::p3_symmetric::{PseudoCompressionFunction, TruncatedPermutation};
423 use p3_baby_bear::default_babybear_poseidon2_16;
424 use p3_field::{Field, PackedValue, PrimeCharacteristicRing};
425
426 type P = <BabyBear as Field>::Packing;
427 let pack_width = P::WIDTH;
428
429 let perm = default_babybear_poseidon2_16();
430 let compressor = TruncatedPermutation::<_, 2, 8, 16>::new(perm);
431
432 let num_leaves = row_hashes.len();
433 let query_stride = num_leaves / rows_per_query;
434
435 let mut prev_layer = row_hashes;
437 for _ in 0..log2_strict_usize(rows_per_query) {
438 let n = prev_layer.len() / 2;
439 let qs = query_stride;
440 let mut next_layer = vec![[BabyBear::ZERO; 8]; n];
441
442 next_layer
443 .par_chunks_mut(pack_width)
444 .enumerate()
445 .for_each(|(chunk_idx, out_chunk)| {
446 let base = chunk_idx * pack_width;
447 let actual = out_chunk.len();
448
449 if actual == pack_width {
450 let mut packed_input: [[P; 8]; 2] = [[P::default(); 8]; 2];
451 for d in 0..8 {
452 packed_input[0][d] = P::from_fn(|lane| {
453 let i = base + lane;
454 let x = i / qs;
455 let y = i % qs;
456 prev_layer[2 * x * qs + y][d]
457 });
458 packed_input[1][d] = P::from_fn(|lane| {
459 let i = base + lane;
460 let x = i / qs;
461 let y = i % qs;
462 prev_layer[(2 * x + 1) * qs + y][d]
463 });
464 }
465 let packed_result: [P; 8] = compressor.compress(packed_input);
466 for lane in 0..pack_width {
467 for d in 0..8 {
468 out_chunk[lane][d] = packed_result[d].as_slice()[lane];
469 }
470 }
471 } else {
472 for lane in 0..actual {
473 let i = base + lane;
474 let x = i / qs;
475 let y = i % qs;
476 out_chunk[lane] = compressor.compress([
477 prev_layer[2 * x * qs + y],
478 prev_layer[(2 * x + 1) * qs + y],
479 ]);
480 }
481 }
482 });
483
484 prev_layer = next_layer;
485 }
486
487 let mut layers = vec![prev_layer];
489 while layers.last().unwrap().len() > 1 {
490 let n = layers.last().unwrap().len() / 2;
491 let mut layer = vec![[BabyBear::ZERO; 8]; n];
492 {
493 let prev = layers.last().unwrap();
494 layer
495 .par_chunks_mut(pack_width)
496 .enumerate()
497 .for_each(|(chunk_idx, out_chunk)| {
498 let base = chunk_idx * pack_width;
499 let actual = out_chunk.len();
500
501 if actual == pack_width {
502 let mut packed_input: [[P; 8]; 2] = [[P::default(); 8]; 2];
503 for d in 0..8 {
504 packed_input[0][d] = P::from_fn(|lane| prev[2 * (base + lane)][d]);
505 packed_input[1][d] = P::from_fn(|lane| prev[2 * (base + lane) + 1][d]);
506 }
507 let packed_result: [P; 8] = compressor.compress(packed_input);
508 for lane in 0..pack_width {
509 for d in 0..8 {
510 out_chunk[lane][d] = packed_result[d].as_slice()[lane];
511 }
512 }
513 } else {
514 for lane in 0..actual {
515 let i = base + lane;
516 out_chunk[lane] = compressor.compress([prev[2 * i], prev[2 * i + 1]]);
517 }
518 }
519 });
520 }
521 layers.push(layer);
522 }
523
524 layers
525}
526
527#[cfg(test)]
528mod tests {
529 use super::*;
530
531 #[test]
532 fn test_from_raw_parts_and_accessors() {
533 let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4, 5, 6], 3);
534 let digest_layers = vec![vec![10u32, 20], vec![30]];
535
536 let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
537
538 assert_eq!(tree.rows_per_query(), 1);
539 assert_eq!(tree.backing_matrix().width, 3);
540 assert_eq!(tree.digest_layers().len(), 2);
541 assert_eq!(tree.query_stride(), 2);
542 assert_eq!(tree.proof_depth(), 1);
543 }
544
545 #[test]
546 fn test_root() {
547 let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4], 2);
548 let digest_layers = vec![vec![10u32, 20], vec![42]];
549
550 let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
551 assert_eq!(tree.root().unwrap(), 42);
552 }
553
554 #[test]
555 fn test_root_no_layers() {
556 let mat = RowMajorMatrix::new(vec![1u32, 2], 2);
557 let digest_layers: Vec<Vec<u32>> = vec![];
558
559 let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
560 assert!(tree.root().is_err());
561 }
562
563 #[test]
564 fn test_query_merkle_proof() {
565 let mat = RowMajorMatrix::new(vec![0u32; 8], 2);
567 let layer0 = vec![10u32, 20, 30, 40]; let layer1 = vec![100u32, 200]; let layer2 = vec![999u32]; let digest_layers = vec![layer0, layer1, layer2];
571
572 let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
573 assert_eq!(tree.query_stride(), 4);
574 assert_eq!(tree.proof_depth(), 2);
575
576 let proof = tree.query_merkle_proof(0).unwrap();
578 assert_eq!(proof, vec![20, 200]);
579
580 let proof = tree.query_merkle_proof(1).unwrap();
582 assert_eq!(proof, vec![10, 200]);
583
584 let proof = tree.query_merkle_proof(2).unwrap();
586 assert_eq!(proof, vec![40, 100]);
587
588 assert!(tree.query_merkle_proof(4).is_err());
590 }
591
592 #[test]
593 fn test_get_opened_rows_single() {
594 let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12], 3);
596 let digest_layers = vec![vec![0u32; 4], vec![0u32; 2], vec![0u32; 1]];
597
598 let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 1) };
599 assert_eq!(tree.query_stride(), 4);
600
601 let rows = tree.get_opened_rows(0).unwrap();
602 assert_eq!(rows.len(), 1);
603 assert_eq!(rows[0], vec![1, 2, 3]);
604
605 let rows = tree.get_opened_rows(2).unwrap();
606 assert_eq!(rows.len(), 1);
607 assert_eq!(rows[0], vec![7, 8, 9]);
608 }
609
610 #[test]
611 fn test_get_opened_rows_batched() {
612 let mat = RowMajorMatrix::new(vec![1u32, 2, 3, 4, 5, 6, 7, 8], 2);
615 let digest_layers = vec![vec![0u32; 2], vec![0u32; 1]];
616
617 let tree = unsafe { CpuMerkleTree::from_raw_parts(mat, digest_layers, 2) };
618 assert_eq!(tree.query_stride(), 2);
619
620 let rows = tree.get_opened_rows(0).unwrap();
622 assert_eq!(rows.len(), 2);
623 assert_eq!(rows[0], vec![1, 2]);
624 assert_eq!(rows[1], vec![5, 6]);
625
626 let rows = tree.get_opened_rows(1).unwrap();
628 assert_eq!(rows.len(), 2);
629 assert_eq!(rows[0], vec![3, 4]);
630 assert_eq!(rows[1], vec![7, 8]);
631
632 assert!(tree.get_opened_rows(2).is_err());
634 }
635}