binius_iop/merkle_tree/
scheme.rs1use std::{
5 fmt::{self, Debug, Formatter},
6 marker::PhantomData,
7};
8
9use binius_hash::{CompressionFunction, HashSuite, hash_serialize};
10use binius_transcript::{Buf, TranscriptReader};
11use binius_utils::{
12 FixedSizeSerializeBytes,
13 checked_arithmetics::{checked_log_2, log2_ceil_usize},
14};
15use digest::{Digest, Output};
16
17use super::{
18 error::{Error, VerificationError},
19 merkle_tree_vcs::MerkleTreeScheme,
20};
21
22pub struct BinaryMerkleTreeScheme<T, H: HashSuite> {
28 compression: H::Compression,
30 _phantom: PhantomData<fn() -> T>,
33}
34
35impl<T, H: HashSuite> Default for BinaryMerkleTreeScheme<T, H> {
36 fn default() -> Self {
37 Self::new()
38 }
39}
40
41impl<T, H: HashSuite> BinaryMerkleTreeScheme<T, H> {
42 pub fn new() -> Self {
43 Self {
44 compression: H::Compression::default(),
46 _phantom: PhantomData,
47 }
48 }
49
50 fn fold_to_root(&self, digests: &[Output<H::LeafHash>]) -> Output<H::LeafHash> {
68 assert!(
71 digests.len().is_power_of_two(),
72 "precondition: the number of digests must be a non-zero power of two"
73 );
74
75 if let [root] = digests {
77 return root.clone();
78 }
79
80 let mut layer = digests
83 .chunks_exact(2)
84 .map(|pair| {
85 self.compression
86 .compress([pair[0].clone(), pair[1].clone()])
87 })
88 .collect::<Vec<_>>();
89
90 while layer.len() > 1 {
93 let half = layer.len() / 2;
94 for i in 0..half {
95 layer[i] = self
96 .compression
97 .compress([layer[2 * i].clone(), layer[2 * i + 1].clone()]);
98 }
99 layer.truncate(half);
101 }
102
103 layer
104 .pop()
105 .expect("a non-empty layer folds down to exactly one digest")
106 }
107}
108
109impl<T, H: HashSuite> Clone for BinaryMerkleTreeScheme<T, H> {
110 fn clone(&self) -> Self {
111 Self {
115 compression: self.compression.clone(),
116 _phantom: PhantomData,
117 }
118 }
119}
120
121impl<T, H: HashSuite> Debug for BinaryMerkleTreeScheme<T, H> {
122 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
123 f.debug_struct("BinaryMerkleTreeScheme").finish()
127 }
128}
129
130impl<T, H> BinaryMerkleTreeScheme<T, H>
131where
132 T: FixedSizeSerializeBytes,
133 H: HashSuite,
134{
135 fn compute_leaf_digest(&self, values: &[T]) -> Result<Output<H::LeafHash>, Error> {
137 Ok(hash_serialize::<T, H::LeafHash>(values)?)
138 }
139}
140
141impl<T, H> MerkleTreeScheme<T> for BinaryMerkleTreeScheme<T, H>
142where
143 T: FixedSizeSerializeBytes,
144 H: HashSuite,
145{
146 type Digest = Output<H::LeafHash>;
147
148 fn optimal_verify_layer(&self, n_queries: usize, tree_depth: usize) -> usize {
149 log2_ceil_usize(n_queries).min(tree_depth)
154 }
155
156 fn proof_size(&self, len: usize, n_queries: usize, layer_depth: usize) -> usize {
157 assert!(len.is_power_of_two(), "precondition: len must be a power of two");
158
159 let log_len = checked_log_2(len);
161
162 assert!(layer_depth <= log_len, "precondition: layer_depth must be at most log2(len)");
163
164 ((log_len - layer_depth) * n_queries + (1 << layer_depth))
170 * <H::LeafHash as Digest>::output_size()
171 }
172
173 fn verify_vector(
174 &self,
175 root: &Self::Digest,
176 data: &[T],
177 batch_size: usize,
178 ) -> Result<(), Error> {
179 assert_ne!(batch_size, 0, "precondition: batch_size must be non-zero");
181 assert!(
183 data.len().is_multiple_of(batch_size),
184 "precondition: data length must be a multiple of batch_size"
185 );
186 assert!(
188 (data.len() / batch_size).is_power_of_two(),
189 "precondition: data.len() / batch_size must be a non-zero power of two"
190 );
191
192 let digests = data
194 .chunks(batch_size)
195 .map(|chunk| self.compute_leaf_digest(chunk))
196 .collect::<Result<Vec<_>, _>>()?;
197
198 if self.fold_to_root(&digests) != *root {
200 return Err(VerificationError::InvalidProof.into());
201 }
202 Ok(())
203 }
204
205 fn verify_layer(
206 &self,
207 root: &Self::Digest,
208 layer_depth: usize,
209 layer_digests: &[Self::Digest],
210 ) -> Result<(), Error> {
211 assert_eq!(
213 layer_digests.len(),
214 1 << layer_depth,
215 "precondition: layer_digests must have 2^layer_depth entries"
216 );
217
218 if self.fold_to_root(layer_digests) != *root {
221 return Err(VerificationError::InvalidProof.into());
222 }
223 Ok(())
224 }
225
226 fn verify_opening<B: Buf>(
227 &self,
228 mut index: usize,
229 values: &[T],
230 layer_depth: usize,
231 tree_depth: usize,
232 layer_digests: &[Self::Digest],
233 proof: &mut TranscriptReader<'_, B>,
234 ) -> Result<(), Error> {
235 assert_eq!(
237 layer_digests.len(),
238 1 << layer_depth,
239 "precondition: layer_digests must have 2^layer_depth entries"
240 );
241 assert!(layer_depth <= tree_depth, "precondition: layer_depth must be at most tree_depth");
243 assert!(index < (1 << tree_depth), "precondition: index must be less than 2^tree_depth");
245
246 let mut digest = self.compute_leaf_digest(values)?;
248
249 for _ in layer_depth..tree_depth {
255 let sibling = proof.read::<Self::Digest>()?;
256 digest = self.compression.compress(if index & 1 == 0 {
258 [digest, sibling]
259 } else {
260 [sibling, digest]
261 });
262 index >>= 1;
264 }
265
266 if digest != layer_digests[index] {
269 return Err(VerificationError::InvalidProof.into());
270 }
271 Ok(())
272 }
273}