Skip to main content

openvm_cpu_backend/
merkle.rs

1//! [CpuMerkleTree] — a Merkle tree backed by a [RowMajorMatrix].
2//!
3//! This mirrors the interface of [`openvm_stark_backend::prover::stacked_pcs::MerkleTree`]
4//! but stores its codeword matrix in row-major layout for cache-friendly row hashing
5//! and contiguous row access in query answering.
6
7use 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/// Merkle tree with row-major codeword backing.
22///
23/// Each leaf corresponds to a row of `backing_matrix` (or a batch of rows when
24/// `rows_per_query > 1`).  The `digest_layers` are built exactly as in
25/// [`MerkleTree`](openvm_stark_backend::prover::stacked_pcs::MerkleTree), so
26/// proof construction and verification are identical.
27#[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    /// Construct a `CpuMerkleTree` from pre-computed parts without validation.
36    ///
37    /// # Safety
38    ///
39    /// The caller must guarantee:
40    /// - `digest_layers` form a valid Merkle tree over `backing_matrix`: the leaf layer contains
41    ///   correct hashes of the matrix rows and each subsequent layer contains correct compressions
42    ///   of consecutive pairs from the previous layer, terminating in a single root digest.
43    /// - `rows_per_query` is a power of two and does not exceed the number of leaves (i.e.,
44    ///   `backing_matrix.height().next_power_of_two()`).
45    /// - The leaf layer length equals `backing_matrix.height().next_power_of_two() /
46    ///   rows_per_query`.
47    ///
48    /// Violating these invariants will produce incorrect Merkle proofs or panics
49    /// in downstream query/verification code.
50    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    /// Returns a reference to the row-major codeword matrix.
63    pub fn backing_matrix(&self) -> &RowMajorMatrix<F> {
64        &self.backing_matrix
65    }
66
67    /// Returns a reference to the digest layers.
68    pub fn digest_layers(&self) -> &Vec<Vec<Digest>> {
69        &self.digest_layers
70    }
71
72    /// Number of rows batched into each Merkle leaf.
73    pub fn rows_per_query(&self) -> usize {
74        self.rows_per_query
75    }
76
77    /// Number of distinct query indices (= number of entries in the bottom digest layer).
78    pub fn query_stride(&self) -> usize {
79        self.digest_layers[0].len()
80    }
81
82    /// Depth of a Merkle proof (number of sibling digests).
83    pub fn proof_depth(&self) -> usize {
84        self.digest_layers.len() - 1
85    }
86}
87
88impl<F, Digest: Clone> CpuMerkleTree<F, Digest> {
89    /// Returns the Merkle root (the single element of the topmost digest layer).
90    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    /// Returns the Merkle authentication path for the given `query_idx`.
99    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    /// Returns the opened rows for the given query index.
121    ///
122    /// The rows are `{ index + t * query_stride() }` for `t` in `0..rows_per_query`.
123    ///
124    /// Because the backing matrix is row-major, each row is a contiguous slice —
125    /// no scatter-gather is needed (unlike the column-major variant).
126    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                // Padding row for matrices whose height is not a power of two.
145                rows.push(vec![]);
146            }
147        }
148        Ok(rows)
149    }
150}
151
152/// Reinterpret a `Vec<A>` as `Vec<B>` when both types have identical layout.
153///
154/// # Safety
155/// The caller must guarantee that `A` and `B` have the same size, alignment,
156/// and compatible memory representations (e.g. verified via `TypeId` checks).
157pub(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
164/// Packed SIMD row hashing for BabyBear Poseidon2.
165///
166/// Hashes `F::Packing::WIDTH` rows simultaneously using packed field arithmetic.
167/// On aarch64 NEON: 4 rows/hash, on x86 AVX2: 8 rows/hash, scalar fallback: 1 row/hash.
168///
169/// Constructs a fresh `PaddingFreeSponge` from `default_babybear_poseidon2_16()`,
170/// which is deterministic and identical to the one in `BabyBearPoseidon2Config`.
171pub(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                // SIMD: pack `pack_width` rows into packed field elements
197                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                // Hash produces [P; 8] — pack_width digests interleaved across lanes
211                let packed_digest: [P; 8] = sponge.hash_slice(&packed_row);
212
213                // Unpack individual digests from SIMD lanes
214                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                // Scalar fallback for partial final chunk
221                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
233/// Build Merkle digest layers from row hashes, dispatching to packed SIMD for
234/// BabyBear or scalar fallback otherwise.
235///
236/// This is the shared implementation used by both `rs_encode_and_merkle_cpu` and
237/// `build_ef_merkle_tree_packed` in the WHIR module.
238pub(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        // SAFETY: TypeId checks guarantee H::Digest = [BabyBear; 8].
252        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/// Fused RS encoding + Merkle tree construction with RowMajor backing.
287///
288/// Eliminates the Phase 6 transpose to col-major that the reference implementation requires,
289/// since `CpuMerkleTree` stores the codeword matrix in row-major layout directly.
290#[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    // Phase 1: Convert PLE evaluations to coefficients (parallel per column).
311    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    // Phase 2: Transpose column vectors into a RowMajorMatrix for batch DFT.
324    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    // Phase 3: Batch DFT — single level of rayon parallelism + SIMD butterflies.
339    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    // Phase 4: Hash rows — use packed SIMD for BabyBear, scalar fallback otherwise.
347    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    // Phase 5: Build Merkle digest layers.
371    let digest_layers = tracing::info_span!("digest_layers")
372        .in_scope(|| build_digest_layers::<F, H>(row_hashes, rows_per_query, hasher));
373
374    // No Phase 6: RowMajor DFT result is stored directly as the backing matrix.
375    // This eliminates the O(n*m) transpose to col-major.
376
377    // SAFETY: digest_layers were just computed as correct Merkle hashes over rm_result
378    // by hash_rows_packed_babybear and build_digest_layers_packed_babybear above.
379    // rows_per_query is forwarded from the validated SystemParams.
380    unsafe { CpuMerkleTree::from_raw_parts(rm_result, digest_layers, rows_per_query) }
381}
382
383/// Scalar fallback for building Merkle digest layers.
384fn 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
417/// Packed SIMD Merkle tree digest layer compression for BabyBear Poseidon2.
418fn 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    // Phase 1: Query-stride interleaved layers.
436    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    // Phase 2: Standard binary tree layers (adjacent pairs).
488    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        // 4 leaves -> 2 layers: [a, b, c, d] -> [ab, cd] -> [abcd]
566        let mat = RowMajorMatrix::new(vec![0u32; 8], 2);
567        let layer0 = vec![10u32, 20, 30, 40]; // 4 entries
568        let layer1 = vec![100u32, 200]; // 2 entries
569        let layer2 = vec![999u32]; // root
570        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        // Query index 0: siblings are layer0[1], layer1[1]
577        let proof = tree.query_merkle_proof(0).unwrap();
578        assert_eq!(proof, vec![20, 200]);
579
580        // Query index 1: siblings are layer0[0], layer1[1]
581        let proof = tree.query_merkle_proof(1).unwrap();
582        assert_eq!(proof, vec![10, 200]);
583
584        // Query index 2: siblings are layer0[3], layer1[0]
585        let proof = tree.query_merkle_proof(2).unwrap();
586        assert_eq!(proof, vec![40, 100]);
587
588        // Out of bounds
589        assert!(tree.query_merkle_proof(4).is_err());
590    }
591
592    #[test]
593    fn test_get_opened_rows_single() {
594        // 4 rows x 3 cols, rows_per_query = 1
595        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        // 4 rows x 2 cols, rows_per_query = 2
613        // query_stride = 4 / 2 = 2 (from digest_layers[0].len())
614        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        // index=0: rows at 0 and 0 + 2 = 2
621        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        // index=1: rows at 1 and 1 + 2 = 3
627        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        // Out of bounds
633        assert!(tree.get_opened_rows(2).is_err());
634    }
635}