commonware_storage/merkle/
position.rs1use super::{Family, location::Location};
2use bytes::{Buf, BufMut};
3use commonware_codec::{ReadExt, varint::UInt};
4use core::{
5 fmt,
6 marker::PhantomData,
7 ops::{Add, AddAssign, Deref, Sub, SubAssign},
8};
9
10pub struct Position<F: Family>(u64, PhantomData<F>);
20
21#[cfg(feature = "arbitrary")]
22impl<F: Family> arbitrary::Arbitrary<'_> for Position<F> {
23 fn arbitrary(u: &mut arbitrary::Unstructured<'_>) -> arbitrary::Result<Self> {
24 let value = u.int_in_range(0..=F::MAX_NODES.as_u64())?;
25 Ok(Self::new(value))
26 }
27}
28
29impl<F: Family> Position<F> {
30 #[inline]
32 pub const fn new(pos: u64) -> Self {
33 Self(pos, PhantomData)
34 }
35
36 #[inline]
38 pub const fn as_u64(self) -> u64 {
39 self.0
40 }
41
42 #[inline]
44 pub const fn is_valid(self) -> bool {
45 self.0 <= F::MAX_NODES.as_u64()
46 }
47
48 #[inline]
50 pub const fn is_valid_index(self) -> bool {
51 self.0 < F::MAX_NODES.as_u64()
52 }
53
54 #[inline]
56 pub const fn checked_add(self, rhs: u64) -> Option<Self> {
57 match self.0.checked_add(rhs) {
58 Some(value) => {
59 if value <= F::MAX_NODES.as_u64() {
60 Some(Self::new(value))
61 } else {
62 None
63 }
64 }
65 None => None,
66 }
67 }
68
69 #[inline]
71 pub const fn checked_sub(self, rhs: u64) -> Option<Self> {
72 match self.0.checked_sub(rhs) {
73 Some(value) => Some(Self::new(value)),
74 None => None,
75 }
76 }
77
78 #[inline]
80 pub const fn saturating_add(self, rhs: u64) -> Self {
81 let result = self.0.saturating_add(rhs);
82 if result > F::MAX_NODES.as_u64() {
83 F::MAX_NODES
84 } else {
85 Self::new(result)
86 }
87 }
88
89 #[inline]
91 pub const fn saturating_sub(self, rhs: u64) -> Self {
92 Self::new(self.0.saturating_sub(rhs))
93 }
94
95 #[inline]
97 pub fn is_valid_size(self) -> bool {
98 F::is_valid_size(self)
99 }
100}
101
102impl<F: Family> Copy for Position<F> {}
105
106impl<F: Family> Clone for Position<F> {
107 #[inline]
108 fn clone(&self) -> Self {
109 *self
110 }
111}
112
113impl<F: Family> PartialEq for Position<F> {
114 #[inline]
115 fn eq(&self, other: &Self) -> bool {
116 self.0 == other.0
117 }
118}
119
120impl<F: Family> Eq for Position<F> {}
121
122impl<F: Family> PartialOrd for Position<F> {
123 #[inline]
124 fn partial_cmp(&self, other: &Self) -> Option<core::cmp::Ordering> {
125 Some(self.cmp(other))
126 }
127}
128
129impl<F: Family> Ord for Position<F> {
130 #[inline]
131 fn cmp(&self, other: &Self) -> core::cmp::Ordering {
132 self.0.cmp(&other.0)
133 }
134}
135
136impl<F: Family> core::hash::Hash for Position<F> {
137 #[inline]
138 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
139 self.0.hash(state);
140 }
141}
142
143impl<F: Family> Default for Position<F> {
144 #[inline]
145 fn default() -> Self {
146 Self::new(0)
147 }
148}
149
150impl<F: Family> fmt::Debug for Position<F> {
151 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
152 f.debug_tuple("Position").field(&self.0).finish()
153 }
154}
155
156impl<F: Family> fmt::Display for Position<F> {
157 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
158 write!(f, "Position({})", self.0)
159 }
160}
161
162impl<F: Family> Deref for Position<F> {
163 type Target = u64;
164 fn deref(&self) -> &Self::Target {
165 &self.0
166 }
167}
168
169impl<F: Family> AsRef<u64> for Position<F> {
170 fn as_ref(&self) -> &u64 {
171 &self.0
172 }
173}
174
175impl<F: Family> From<u64> for Position<F> {
176 #[inline]
177 fn from(value: u64) -> Self {
178 Self::new(value)
179 }
180}
181
182impl<F: Family> From<usize> for Position<F> {
183 #[inline]
184 fn from(value: usize) -> Self {
185 Self::new(value as u64)
186 }
187}
188
189impl<F: Family> From<Position<F>> for u64 {
190 #[inline]
191 fn from(position: Position<F>) -> Self {
192 *position
193 }
194}
195
196impl<F: Family> TryFrom<Location<F>> for Position<F> {
202 type Error = super::Error<F>;
203
204 #[inline]
205 fn try_from(loc: Location<F>) -> Result<Self, Self::Error> {
206 if !loc.is_valid() {
207 return Err(super::Error::LocationOverflow(loc));
208 }
209 Ok(F::location_to_position(loc))
210 }
211}
212
213impl<F: Family> Add for Position<F> {
221 type Output = Self;
222
223 #[inline]
224 fn add(self, rhs: Self) -> Self::Output {
225 Self::new(self.0 + rhs.0)
226 }
227}
228
229impl<F: Family> Add<u64> for Position<F> {
235 type Output = Self;
236
237 #[inline]
238 fn add(self, rhs: u64) -> Self::Output {
239 Self::new(self.0 + rhs)
240 }
241}
242
243impl<F: Family> Sub for Position<F> {
249 type Output = Self;
250
251 #[inline]
252 fn sub(self, rhs: Self) -> Self::Output {
253 Self::new(self.0 - rhs.0)
254 }
255}
256
257impl<F: Family> Sub<u64> for Position<F> {
263 type Output = Self;
264
265 #[inline]
266 fn sub(self, rhs: u64) -> Self::Output {
267 Self::new(*self - rhs)
268 }
269}
270
271impl<F: Family> PartialEq<u64> for Position<F> {
272 #[inline]
273 fn eq(&self, other: &u64) -> bool {
274 self.0 == *other
275 }
276}
277
278impl<F: Family> PartialOrd<u64> for Position<F> {
279 #[inline]
280 fn partial_cmp(&self, other: &u64) -> Option<core::cmp::Ordering> {
281 self.0.partial_cmp(other)
282 }
283}
284
285impl<F: Family> PartialEq<Position<F>> for u64 {
286 #[inline]
287 fn eq(&self, other: &Position<F>) -> bool {
288 *self == other.0
289 }
290}
291
292impl<F: Family> PartialOrd<Position<F>> for u64 {
293 #[inline]
294 fn partial_cmp(&self, other: &Position<F>) -> Option<core::cmp::Ordering> {
295 self.partial_cmp(&other.0)
296 }
297}
298
299impl<F: Family> AddAssign<u64> for Position<F> {
305 #[inline]
306 fn add_assign(&mut self, rhs: u64) {
307 self.0 += rhs;
308 }
309}
310
311impl<F: Family> SubAssign<u64> for Position<F> {
317 #[inline]
318 fn sub_assign(&mut self, rhs: u64) {
319 self.0 -= rhs;
320 }
321}
322
323impl<F: Family> commonware_codec::Write for Position<F> {
326 #[inline]
327 fn write(&self, buf: &mut impl BufMut) {
328 UInt(self.0).write(buf);
329 }
330}
331
332impl<F: Family> commonware_codec::EncodeSize for Position<F> {
333 #[inline]
334 fn encode_size(&self) -> usize {
335 UInt(self.0).encode_size()
336 }
337}
338
339impl<F: Family> commonware_codec::Read for Position<F> {
340 type Cfg = ();
341
342 #[inline]
343 fn read_cfg(buf: &mut impl Buf, _: &()) -> Result<Self, commonware_codec::Error> {
344 let pos = Self::new(UInt::read(buf)?.into());
345 if pos.is_valid() {
346 Ok(pos)
347 } else {
348 Err(commonware_codec::Error::Invalid(
349 "Position",
350 "value exceeds MAX_NODES",
351 ))
352 }
353 }
354}
355#[cfg(test)]
356mod tests {
357 use super::{Location as GenericLocation, Position as GenericPosition};
358 use crate::{
359 merkle::{Bagging::ForwardFold, Family as _},
360 mmr::{self, StandardHasher as Standard, mem::Mmr},
361 };
362 use commonware_cryptography::Sha256;
363
364 type Location = GenericLocation<mmr::Family>;
365 type Position = GenericPosition<mmr::Family>;
366
367 #[test]
369 fn test_from_location() {
370 const CASES: &[(Location, Position)] = &[
371 (Location::new(0), Position::new(0)),
372 (Location::new(1), Position::new(1)),
373 (Location::new(2), Position::new(3)),
374 (Location::new(3), Position::new(4)),
375 (Location::new(4), Position::new(7)),
376 (Location::new(5), Position::new(8)),
377 (Location::new(6), Position::new(10)),
378 (Location::new(7), Position::new(11)),
379 (Location::new(8), Position::new(15)),
380 (Location::new(9), Position::new(16)),
381 (Location::new(10), Position::new(18)),
382 (Location::new(11), Position::new(19)),
383 (Location::new(12), Position::new(22)),
384 (Location::new(13), Position::new(23)),
385 (Location::new(14), Position::new(25)),
386 (Location::new(15), Position::new(26)),
387 ];
388 for (loc, expected_pos) in CASES {
389 let pos = Position::try_from(*loc).unwrap();
390 assert_eq!(pos, *expected_pos);
391 }
392 }
393
394 #[test]
395 fn test_checked_add() {
396 let pos = Position::new(10);
397 assert_eq!(pos.checked_add(5).unwrap(), 15);
398
399 assert!(Position::new(u64::MAX).checked_add(1).is_none());
401
402 assert!(mmr::Family::MAX_NODES.checked_add(1).is_none());
404 assert!(
405 Position::new(*mmr::Family::MAX_NODES - 5)
406 .checked_add(10)
407 .is_none()
408 );
409 assert_eq!(
411 Position::new(*mmr::Family::MAX_NODES - 10)
412 .checked_add(10)
413 .unwrap(),
414 *mmr::Family::MAX_NODES
415 );
416
417 assert_eq!(
419 Position::new(*mmr::Family::MAX_NODES - 11)
420 .checked_add(10)
421 .unwrap(),
422 *mmr::Family::MAX_NODES - 1
423 );
424 }
425
426 #[test]
427 fn test_checked_sub() {
428 let pos = Position::new(10);
429 assert_eq!(pos.checked_sub(5).unwrap(), 5);
430 assert!(pos.checked_sub(11).is_none());
431 }
432
433 #[test]
434 fn test_saturating_add() {
435 let pos = Position::new(10);
436 assert_eq!(pos.saturating_add(5), 15);
437
438 assert_eq!(
440 Position::new(u64::MAX).saturating_add(1),
441 *mmr::Family::MAX_NODES
442 );
443 assert_eq!(
444 mmr::Family::MAX_NODES.saturating_add(1),
445 *mmr::Family::MAX_NODES
446 );
447 assert_eq!(
448 mmr::Family::MAX_NODES.saturating_add(1000),
449 *mmr::Family::MAX_NODES
450 );
451 assert_eq!(
452 Position::new(*mmr::Family::MAX_NODES - 5).saturating_add(10),
453 *mmr::Family::MAX_NODES
454 );
455 }
456
457 #[test]
458 fn test_saturating_sub() {
459 let pos = Position::new(10);
460 assert_eq!(pos.saturating_sub(5), 5);
461 assert_eq!(Position::new(0).saturating_sub(1), 0);
462 }
463
464 #[test]
465 fn test_display() {
466 let position = Position::new(42);
467 assert_eq!(position.to_string(), "Position(42)");
468 }
469
470 #[test]
471 fn test_add() {
472 let pos1 = Position::new(10);
473 let pos2 = Position::new(5);
474 assert_eq!((pos1 + pos2), 15);
475 }
476
477 #[test]
478 fn test_sub() {
479 let pos1 = Position::new(10);
480 let pos2 = Position::new(3);
481 assert_eq!((pos1 - pos2), 7);
482 }
483
484 #[test]
485 fn test_comparison_with_u64() {
486 let pos = Position::new(42);
487
488 assert_eq!(pos, 42u64);
490 assert_eq!(42u64, pos);
491 assert_ne!(pos, 43u64);
492 assert_ne!(43u64, pos);
493
494 assert!(pos < 43u64);
496 assert!(43u64 > pos);
497 assert!(pos > 41u64);
498 assert!(41u64 < pos);
499 assert!(pos <= 42u64);
500 assert!(42u64 >= pos);
501 }
502
503 #[test]
504 fn test_assignment_with_u64() {
505 let mut pos = Position::new(10);
506
507 pos += 5;
509 assert_eq!(pos, 15u64);
510
511 pos -= 3;
513 assert_eq!(pos, 12u64);
514 }
515
516 #[test]
517 fn test_max_position() {
518 let max_leaves = 1u64 << 62;
520 let max_size = 2 * max_leaves - 1; assert_eq!(*mmr::Family::MAX_NODES, max_size);
522 assert_eq!(*mmr::Family::MAX_NODES, (1u64 << 63) - 1);
523 assert_eq!(max_size.leading_zeros(), 1); let overflow_size = 2 * (max_leaves + 1) - 1;
527 assert_eq!(overflow_size.leading_zeros(), 0);
528
529 let pos = Position::try_from(mmr::Family::MAX_LEAVES).unwrap();
531 assert_eq!(pos, mmr::Family::MAX_NODES);
532 }
533
534 #[test]
535 fn test_is_valid_size() {
536 let mut size_to_check = Position::new(0);
539 let hasher = Standard::<Sha256>::new(ForwardFold);
540 let mut mmr = Mmr::new();
541 let digest = [1u8; 32];
542 for _i in 0..10000 {
543 while size_to_check != mmr.size() {
544 assert!(
545 !size_to_check.is_valid_size(),
546 "size_to_check: {} {}",
547 size_to_check,
548 mmr.size()
549 );
550 size_to_check += 1;
551 }
552 assert!(size_to_check.is_valid_size());
553 let batch = mmr
554 .new_batch()
555 .add(&hasher, &digest)
556 .merkleize(&mmr, &hasher);
557 mmr.apply_batch(&batch).unwrap();
558 size_to_check += 1;
559 }
560
561 assert!(!Position::new(u64::MAX).is_valid_size());
563 assert!(Position::new(u64::MAX >> 1).is_valid_size()); assert!(!Position::new((u64::MAX >> 1) + 1).is_valid_size());
565 assert!(mmr::Family::MAX_NODES.is_valid_size()); }
567
568 #[test]
569 fn test_read_cfg_valid_values() {
570 use commonware_codec::{Encode, ReadExt};
571
572 let pos = Position::new(0);
574 let encoded = pos.encode();
575 let decoded = Position::read(&mut encoded.as_ref()).unwrap();
576 assert_eq!(decoded, pos);
577
578 let pos = Position::new(12345);
580 let encoded = pos.encode();
581 let decoded = Position::read(&mut encoded.as_ref()).unwrap();
582 assert_eq!(decoded, pos);
583
584 let pos = mmr::Family::MAX_NODES;
586 let encoded = pos.encode();
587 let decoded = Position::read(&mut encoded.as_ref()).unwrap();
588 assert_eq!(decoded, pos);
589
590 let pos = mmr::Family::MAX_NODES - 1;
592 let encoded = pos.encode();
593 let decoded = Position::read(&mut encoded.as_ref()).unwrap();
594 assert_eq!(decoded, pos);
595 }
596
597 #[test]
598 fn test_read_cfg_invalid_values() {
599 use commonware_codec::{Encode, ReadExt, varint::UInt};
600
601 let invalid_value = *mmr::Family::MAX_NODES + 1;
603 let encoded = UInt(invalid_value).encode();
604 let result = Position::read(&mut encoded.as_ref());
605 assert!(result.is_err());
606 assert!(matches!(
607 result,
608 Err(commonware_codec::Error::Invalid("Position", _))
609 ));
610
611 let encoded = UInt(u64::MAX).encode();
613 let result = Position::read(&mut encoded.as_ref());
614 assert!(result.is_err());
615 assert!(matches!(
616 result,
617 Err(commonware_codec::Error::Invalid("Position", _))
618 ));
619 }
620}