1use crate::hash::Hash;
18use crate::write_auth::PartCommitment;
19use blake3::hazmat::{self, HasherExt, Mode};
20
21pub const MIN_PART_SIZE: u64 = 8 * 1024 * 1024;
23
24pub type ChainingValue = [u8; 32];
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
29#[non_exhaustive]
30pub enum PartError {
31 #[error("part size is not a power of two")]
33 PartSizeNotPowerOfTwo,
34 #[error("part size is below the minimum")]
36 PartSizeTooSmall,
37 #[error("pack fits one part; upload it whole")]
39 NotMultipart,
40 #[error("too many parts")]
42 TooManyParts,
43 #[error("part index out of range")]
45 IndexOutOfRange,
46 #[error("part data exceeds the part length")]
48 Overrun,
49 #[error("part data is shorter than the part length")]
51 LengthMismatch,
52 #[error("wrong number of part chaining values")]
54 WrongPartCount,
55}
56
57#[derive(Clone, Copy, Debug, PartialEq, Eq)]
59pub struct PartPlan {
60 total: u64,
61 part_size: u64,
62 count: u32,
63}
64
65impl PartPlan {
66 pub fn new(total: u64, part_size: u64, max_parts: u32) -> Result<Self, PartError> {
74 if !part_size.is_power_of_two() {
75 return Err(PartError::PartSizeNotPowerOfTwo);
76 }
77 if part_size < MIN_PART_SIZE {
78 return Err(PartError::PartSizeTooSmall);
79 }
80 Self::build(total, part_size, max_parts)
81 }
82
83 #[cfg(test)]
86 pub(crate) fn new_small(total: u64, part_size: u64) -> Result<Self, PartError> {
87 if !part_size.is_power_of_two() {
88 return Err(PartError::PartSizeNotPowerOfTwo);
89 }
90 if part_size < blake3::CHUNK_LEN as u64 {
91 return Err(PartError::PartSizeTooSmall);
92 }
93 Self::build(total, part_size, u32::MAX)
94 }
95
96 fn build(total: u64, part_size: u64, max_parts: u32) -> Result<Self, PartError> {
97 if total <= part_size {
98 return Err(PartError::NotMultipart);
99 }
100 let count =
101 u32::try_from(total.div_ceil(part_size)).map_err(|_| PartError::TooManyParts)?;
102 if count > max_parts {
103 return Err(PartError::TooManyParts);
104 }
105 u64::from(count)
108 .checked_mul(part_size)
109 .ok_or(PartError::TooManyParts)?;
110 Ok(Self {
111 total,
112 part_size,
113 count,
114 })
115 }
116
117 #[must_use]
119 pub fn count(&self) -> u32 {
120 self.count
121 }
122
123 #[must_use]
125 pub fn part_size(&self) -> u64 {
126 self.part_size
127 }
128
129 #[must_use]
131 pub fn total(&self) -> u64 {
132 self.total
133 }
134
135 pub fn offset(&self, index: u32) -> Result<u64, PartError> {
140 if index >= self.count {
141 return Err(PartError::IndexOutOfRange);
142 }
143 u64::from(index)
145 .checked_mul(self.part_size)
146 .ok_or(PartError::IndexOutOfRange)
147 }
148
149 pub fn expected_len(&self, index: u32) -> Result<u64, PartError> {
155 Ok((self.total - self.offset(index)?).min(self.part_size))
156 }
157
158 pub fn check(&self, commitment: &PartCommitment) -> Result<(), PartError> {
165 if commitment.len == self.expected_len(commitment.index)? {
166 Ok(())
167 } else {
168 Err(PartError::LengthMismatch)
169 }
170 }
171
172 fn range_len(&self, lo: u32, hi: u32) -> u64 {
174 let end = (u64::from(hi) * self.part_size).min(self.total);
175 end - u64::from(lo) * self.part_size
176 }
177}
178
179#[derive(Debug, Clone)]
184pub struct PartHasher {
185 hasher: blake3::Hasher,
186 expected: u64,
187 seen: u64,
188}
189
190impl PartHasher {
191 pub fn new(plan: &PartPlan, index: u32) -> Result<Self, PartError> {
196 let offset = plan.offset(index)?;
197 let expected = plan.expected_len(index)?;
198 let mut hasher = blake3::Hasher::new();
199 hasher.set_input_offset(offset);
202 Ok(Self {
203 hasher,
204 expected,
205 seen: 0,
206 })
207 }
208
209 pub fn update(&mut self, data: &[u8]) -> Result<(), PartError> {
216 let len = u64::try_from(data.len()).map_err(|_| PartError::Overrun)?;
217 if len > self.expected - self.seen {
218 return Err(PartError::Overrun);
219 }
220 self.hasher.update(data);
223 self.seen += len;
224 Ok(())
225 }
226
227 pub fn finalize(self) -> Result<ChainingValue, PartError> {
233 if self.seen != self.expected {
234 return Err(PartError::LengthMismatch);
235 }
236 Ok(self.hasher.finalize_non_root())
238 }
239}
240
241pub fn part_subtree_cv(
247 plan: &PartPlan,
248 index: u32,
249 bytes: &[u8],
250) -> Result<ChainingValue, PartError> {
251 let mut hasher = PartHasher::new(plan, index)?;
252 hasher.update(bytes)?;
253 hasher.finalize()
254}
255
256pub fn merge_to_root(plan: &PartPlan, cvs: &[ChainingValue]) -> Result<Hash, PartError> {
266 if u32::try_from(cvs.len()) != Ok(plan.count) {
267 return Err(PartError::WrongPartCount);
268 }
269 let mid = split(plan, 0, plan.count);
271 let left = merge_range(plan, cvs, 0, mid);
272 let right = merge_range(plan, cvs, mid, plan.count);
273 Ok(*hazmat::merge_subtrees_root(&left, &right, Mode::Hash).as_bytes())
274}
275
276fn split(plan: &PartPlan, lo: u32, hi: u32) -> u32 {
278 let left_parts = hazmat::left_subtree_len(plan.range_len(lo, hi)) / plan.part_size;
282 lo + u32::try_from(left_parts).unwrap_or(hi - lo - 1)
284}
285
286fn merge_range(plan: &PartPlan, cvs: &[ChainingValue], lo: u32, hi: u32) -> ChainingValue {
287 if hi - lo == 1 {
288 return cvs[lo as usize];
289 }
290 let mid = split(plan, lo, hi);
291 hazmat::merge_subtrees_non_root(
292 &merge_range(plan, cvs, lo, mid),
293 &merge_range(plan, cvs, mid, hi),
294 Mode::Hash,
295 )
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301 use proptest::prelude::*;
302
303 const MIB: u64 = 1024 * 1024;
304
305 fn us(n: u64) -> usize {
306 usize::try_from(n).unwrap()
307 }
308
309 fn input(len: u64) -> Vec<u8> {
311 (0..len).map(|i| (i % 251) as u8).collect()
312 }
313
314 fn cvs(plan: &PartPlan, data: &[u8]) -> Vec<ChainingValue> {
315 (0..plan.count())
316 .map(|i| {
317 let start = us(plan.offset(i).unwrap());
318 let end = start + us(plan.expected_len(i).unwrap());
319 part_subtree_cv(plan, i, &data[start..end]).unwrap()
320 })
321 .collect()
322 }
323
324 #[test]
325 fn plan_rejects_part_size_not_power_of_two() {
326 for size in [0, 9 * MIB, 12 * MIB, MIN_PART_SIZE + 1, u64::MAX] {
327 assert_eq!(
328 PartPlan::new(u64::MAX, size, u32::MAX),
329 Err(PartError::PartSizeNotPowerOfTwo),
330 "{size}"
331 );
332 }
333 }
334
335 #[test]
336 fn plan_rejects_small_part_size() {
337 for size in [1, 1024, 4 * MIB, MIN_PART_SIZE / 2] {
338 assert_eq!(
339 PartPlan::new(64 * MIB, size, u32::MAX),
340 Err(PartError::PartSizeTooSmall),
341 "{size}"
342 );
343 }
344 }
345
346 #[test]
347 fn plan_rejects_single_part_packs() {
348 for total in [0, 1, MIN_PART_SIZE - 1, MIN_PART_SIZE] {
349 assert_eq!(
350 PartPlan::new(total, MIN_PART_SIZE, u32::MAX),
351 Err(PartError::NotMultipart),
352 "{total}"
353 );
354 }
355 let plan = PartPlan::new(MIN_PART_SIZE + 1, MIN_PART_SIZE, 2).unwrap();
356 assert_eq!((plan.count(), plan.total()), (2, MIN_PART_SIZE + 1));
357 assert_eq!(plan.part_size(), MIN_PART_SIZE);
358 }
359
360 #[test]
361 fn plan_rejects_too_many_parts() {
362 assert_eq!(
363 PartPlan::new(3 * MIN_PART_SIZE, MIN_PART_SIZE, 2),
364 Err(PartError::TooManyParts)
365 );
366 assert_eq!(
367 PartPlan::new(MIN_PART_SIZE + 1, MIN_PART_SIZE, 0),
368 Err(PartError::TooManyParts)
369 );
370 PartPlan::new(3 * MIN_PART_SIZE, MIN_PART_SIZE, 3).unwrap();
371 assert_eq!(
373 PartPlan::new(u64::MAX, MIN_PART_SIZE, u32::MAX),
374 Err(PartError::TooManyParts)
375 );
376 }
377
378 #[test]
379 fn plan_arithmetic_near_u64_max() {
380 for (total, part_size, max_parts) in [
383 (u64::MAX, 1 << 63, u32::MAX),
384 (u64::MAX, 1 << 51, 10_000),
385 (u64::MAX - (1 << 51) + 2, 1 << 51, 10_000),
386 ] {
387 assert_eq!(
388 PartPlan::new(total, part_size, max_parts),
389 Err(PartError::TooManyParts),
390 "{total} {part_size}"
391 );
392 }
393 let top = 1 << 62;
395 let plan = PartPlan::new(u64::MAX - top, top, u32::MAX).unwrap();
396 assert_eq!(plan.count(), 3);
397 assert_eq!(plan.offset(2), Ok(2 * top));
398 assert_eq!(plan.expected_len(2), Ok(top - 1));
399 assert_eq!(plan.offset(3), Err(PartError::IndexOutOfRange));
400 let hasher = PartHasher::new(&plan, 2).unwrap();
401 assert_eq!(hasher.finalize(), Err(PartError::LengthMismatch));
402 assert!(merge_to_root(&plan, &[[0; 32]; 3]).is_ok());
403 let plan = PartPlan::new(u64::MAX - (1 << 51) + 1, 1 << 51, 10_000).unwrap();
404 assert_eq!(plan.count(), 8191);
405 assert!(merge_to_root(&plan, &vec![[0; 32]; 8191]).is_ok());
406 }
407
408 #[test]
409 fn plan_checks_part_commitments() {
410 let plan = PartPlan::new(2 * MIN_PART_SIZE + 5, MIN_PART_SIZE, 3).unwrap();
411 let part = |index, len| PartCommitment {
412 ticket: [0x5a; 32],
413 index,
414 subtree: [0xcd; 32],
415 len,
416 };
417 assert_eq!(plan.check(&part(0, MIN_PART_SIZE)), Ok(()));
418 assert_eq!(plan.check(&part(1, MIN_PART_SIZE)), Ok(()));
419 assert_eq!(plan.check(&part(2, 5)), Ok(()));
420 for (index, len) in [
421 (0, MIN_PART_SIZE - 1),
422 (1, MIN_PART_SIZE + 1),
423 (1, 5),
424 (2, 4),
425 (2, 6),
426 (2, MIN_PART_SIZE),
427 ] {
428 assert_eq!(
429 plan.check(&part(index, len)),
430 Err(PartError::LengthMismatch),
431 "{index} {len}"
432 );
433 }
434 for index in [3, u32::MAX] {
435 assert_eq!(plan.check(&part(index, 5)), Err(PartError::IndexOutOfRange));
436 }
437 }
438
439 #[test]
440 fn expected_len_last_part_is_remainder() {
441 let plan = PartPlan::new(4 * MIN_PART_SIZE + 5, MIN_PART_SIZE, 5).unwrap();
442 assert_eq!(plan.count(), 5);
443 for i in 0..4 {
444 assert_eq!(plan.expected_len(i), Ok(MIN_PART_SIZE));
445 assert_eq!(plan.offset(i), Ok(u64::from(i) * MIN_PART_SIZE));
446 }
447 assert_eq!(plan.expected_len(4), Ok(5));
448 let exact = PartPlan::new(2 * MIN_PART_SIZE, MIN_PART_SIZE, 2).unwrap();
449 assert_eq!(exact.expected_len(1), Ok(MIN_PART_SIZE));
450 }
451
452 #[test]
453 fn index_out_of_range() {
454 let plan = PartPlan::new(2 * MIN_PART_SIZE, MIN_PART_SIZE, 2).unwrap();
455 for index in [2, 3, u32::MAX] {
456 assert_eq!(plan.offset(index), Err(PartError::IndexOutOfRange));
457 assert_eq!(plan.expected_len(index), Err(PartError::IndexOutOfRange));
458 assert_eq!(
459 PartHasher::new(&plan, index).unwrap_err(),
460 PartError::IndexOutOfRange
461 );
462 }
463 }
464
465 #[test]
466 fn hasher_overrun_is_error_not_panic() {
467 let plan = PartPlan::new_small(5 * 1024 + 7, 1024).unwrap();
468 for index in 0..plan.count() {
469 let expected = us(plan.expected_len(index).unwrap());
470 let good = part_subtree_cv(&plan, index, &vec![7; expected]).unwrap();
471 assert_eq!(
473 part_subtree_cv(&plan, index, &vec![7; expected + 1]),
474 Err(PartError::Overrun)
475 );
476 let mut hasher = PartHasher::new(&plan, index).unwrap();
478 hasher.update(&vec![7; expected]).unwrap();
479 assert_eq!(hasher.update(&[7]), Err(PartError::Overrun));
480 assert_eq!(hasher.finalize(), Ok(good));
481 }
482 let plan = PartPlan::new(3 * MIN_PART_SIZE, MIN_PART_SIZE, 3).unwrap();
484 let mut hasher = PartHasher::new(&plan, 1).unwrap();
485 hasher.update(&vec![0; us(MIN_PART_SIZE)]).unwrap();
486 assert_eq!(hasher.update(&[0]), Err(PartError::Overrun));
487 }
488
489 #[test]
490 fn hasher_short_is_length_mismatch() {
491 let plan = PartPlan::new_small(3 * 1024 + 1, 1024).unwrap();
492 for index in 0..plan.count() {
493 let expected = us(plan.expected_len(index).unwrap());
494 assert_eq!(
495 part_subtree_cv(&plan, index, &vec![1; expected - 1]),
496 Err(PartError::LengthMismatch)
497 );
498 assert_eq!(
499 PartHasher::new(&plan, index).unwrap().finalize(),
500 Err(PartError::LengthMismatch)
501 );
502 }
503 }
504
505 #[test]
506 fn merge_rejects_wrong_count() {
507 let plan = PartPlan::new_small(3 * 1024, 1024).unwrap();
508 let data = input(plan.total());
509 let mut all = cvs(&plan, &data);
510 assert_eq!(
511 merge_to_root(&plan, &all[..2]),
512 Err(PartError::WrongPartCount)
513 );
514 assert_eq!(merge_to_root(&plan, &[]), Err(PartError::WrongPartCount));
515 all.push(all[0]);
516 assert_eq!(merge_to_root(&plan, &all), Err(PartError::WrongPartCount));
517 }
518
519 #[test]
520 fn swapped_cvs_do_not_match_root() {
521 let plan = PartPlan::new_small(4 * 1024, 1024).unwrap();
522 let data = input(plan.total());
523 let mut all = cvs(&plan, &data);
524 assert_eq!(merge_to_root(&plan, &all), Ok(crate::hash::hash(&data)));
525 all.swap(1, 2);
526 assert_ne!(merge_to_root(&plan, &all), Ok(crate::hash::hash(&data)));
527 }
528
529 #[test]
530 fn part_cv_is_offset_bound() {
531 let plan = PartPlan::new_small(3 * 1024, 1024).unwrap();
533 let bytes = [9; 1024];
534 let a = part_subtree_cv(&plan, 0, &bytes).unwrap();
535 let b = part_subtree_cv(&plan, 1, &bytes).unwrap();
536 assert_ne!(a, b);
537 }
538
539 #[test]
542 fn small_geometry_goldens_match_module() {
543 let fixture: serde_json::Value = serde_json::from_str(include_str!(
544 "../../../tests/golden/uploads/subtree-merge.json"
545 ))
546 .unwrap();
547 let mut checked = 0;
548 for vector in fixture["vectors"].as_array().unwrap() {
549 if vector["test_geometry"] != serde_json::Value::Bool(true) {
550 continue;
551 }
552 let plan = PartPlan::new_small(
553 vector["total"].as_u64().unwrap(),
554 vector["part_size"].as_u64().unwrap(),
555 )
556 .unwrap();
557 let data = input(plan.total());
558 let all = cvs(&plan, &data);
559 let parts = vector["parts"].as_array().unwrap();
560 assert_eq!(parts.len(), all.len());
561 for (part, cv) in parts.iter().zip(&all) {
562 assert_eq!(part["cv"].as_str().unwrap(), crate::hash::to_hex(cv));
563 }
564 let root = merge_to_root(&plan, &all).unwrap();
565 assert_eq!(vector["root"].as_str().unwrap(), crate::hash::to_hex(&root));
566 checked += 1;
567 }
568 assert!(checked >= 6, "expected the small-geometry vectors");
569 }
570
571 proptest! {
572 #![proptest_config(ProptestConfig::with_cases(256))]
573
574 #[test]
575 fn streaming_split_equals_one_shot(
576 exp in 10u32..=14,
577 extra in 1u64..=40_000,
578 index_seed in any::<u32>(),
579 cuts in proptest::collection::vec(any::<u16>(), 0..8),
580 ) {
581 let part_size = 1u64 << exp;
582 let plan = PartPlan::new_small(part_size + extra, part_size).unwrap();
583 let index = index_seed % plan.count();
584 let data = input(plan.total());
585 let start = us(plan.offset(index).unwrap());
586 let part = &data[start..start + us(plan.expected_len(index).unwrap())];
587 let one_shot = part_subtree_cv(&plan, index, part).unwrap();
588 let mut points: Vec<usize> = cuts.iter().map(|c| usize::from(*c) % (part.len() + 1)).collect();
589 points.sort_unstable();
590 let mut hasher = PartHasher::new(&plan, index).unwrap();
591 let mut last = 0;
592 for point in points.into_iter().chain([part.len()]) {
593 hasher.update(&part[last..point]).unwrap();
594 last = point;
595 }
596 prop_assert_eq!(hasher.finalize().unwrap(), one_shot);
597 }
598
599 #[test]
600 fn merge_matches_blake3_hash_small_geometry(
601 exp in 10u32..=16,
602 parts in 2u64..=20,
603 short in any::<u64>(),
604 ) {
605 let part_size = 1u64 << exp;
606 let total = (parts - 1) * part_size + 1 + short % part_size;
608 let plan = PartPlan::new_small(total, part_size).unwrap();
609 prop_assert_eq!(u64::from(plan.count()), parts);
610 let data = input(total);
611 let root = merge_to_root(&plan, &cvs(&plan, &data)).unwrap();
612 prop_assert_eq!(root, crate::hash::hash(&data));
613 }
614
615 #[test]
616 fn arbitrary_geometry_never_panics(
617 total in any::<u64>(),
618 part_size in prop_oneof![any::<u64>(), (0u32..64).prop_map(|e| 1u64 << e)],
619 max_parts in any::<u32>(),
620 index in any::<u32>(),
621 ) {
622 if let Ok(plan) = PartPlan::new(total, part_size, max_parts) {
623 prop_assert!(plan.count() >= 2 && plan.count() <= max_parts);
624 let _ = plan.offset(index);
625 let _ = plan.expected_len(index);
626 if let Ok(mut hasher) = PartHasher::new(&plan, index) {
627 let _ = hasher.update(&[0; 3]);
628 let _ = hasher.finalize();
629 }
630 let _ = plan.check(&PartCommitment {
631 ticket: [0; 32],
632 index,
633 subtree: [0; 32],
634 len: total % 97,
635 });
636 if plan.count() <= 64 {
638 let cvs = vec![[0; 32]; plan.count() as usize];
639 prop_assert!(merge_to_root(&plan, &cvs).is_ok());
640 }
641 let _ = merge_to_root(&plan, &[[0; 32]; 3]);
642 }
643 }
644 }
645}