Skip to main content

cranpose_core/
snapshot_id_set.rs

1use std::fmt;
2
3pub type SnapshotId = usize;
4
5const BITS_PER_SET: usize = 64;
6const SNAPSHOT_ID_SIZE: usize = 64;
7
8#[derive(Clone, PartialEq, Eq)]
9pub struct SnapshotIdSet {
10    upper_set: u64,
11    lower_set: u64,
12    lower_bound: SnapshotId,
13    below_bound: Option<Box<[SnapshotId]>>,
14}
15
16impl SnapshotIdSet {
17    /// Empty snapshot ID set.
18    pub const EMPTY: SnapshotIdSet = SnapshotIdSet {
19        upper_set: 0,
20        lower_set: 0,
21        lower_bound: 0,
22        below_bound: None,
23    };
24
25    /// Create a new empty snapshot ID set.
26    pub fn new() -> Self {
27        Self::EMPTY
28    }
29
30    /// Check if a snapshot ID is in the set.
31    pub fn get(&self, id: SnapshotId) -> bool {
32        let offset = id.wrapping_sub(self.lower_bound);
33
34        if offset < BITS_PER_SET {
35            let mask = 1u64 << offset;
36            (self.lower_set & mask) != 0
37        } else if offset < BITS_PER_SET * 2 {
38            let mask = 1u64 << (offset - BITS_PER_SET);
39            (self.upper_set & mask) != 0
40        } else if id > self.lower_bound {
41            false
42        } else {
43            self.below_bound
44                .as_ref()
45                .map(|arr| arr.binary_search(&id).is_ok())
46                .unwrap_or(false)
47        }
48    }
49
50    /// Add a snapshot ID to the set (returns a new set if modified).
51    pub fn set(&self, id: SnapshotId) -> Self {
52        if id < self.lower_bound {
53            if let Some(ref arr) = self.below_bound {
54                match arr.binary_search(&id) {
55                    Ok(_) => {
56                        return self.clone();
57                    }
58                    Err(insert_pos) => {
59                        let mut new_arr = Vec::with_capacity(arr.len() + 1);
60                        new_arr.extend_from_slice(&arr[..insert_pos]);
61                        new_arr.push(id);
62                        new_arr.extend_from_slice(&arr[insert_pos..]);
63                        return Self {
64                            upper_set: self.upper_set,
65                            lower_set: self.lower_set,
66                            lower_bound: self.lower_bound,
67                            below_bound: Some(new_arr.into_boxed_slice()),
68                        };
69                    }
70                }
71            } else {
72                return Self {
73                    upper_set: self.upper_set,
74                    lower_set: self.lower_set,
75                    lower_bound: self.lower_bound,
76                    below_bound: Some(vec![id].into_boxed_slice()),
77                };
78            }
79        }
80
81        let offset = id - self.lower_bound;
82
83        if offset < BITS_PER_SET {
84            let mask = 1u64 << offset;
85            if (self.lower_set & mask) == 0 {
86                return Self {
87                    upper_set: self.upper_set,
88                    lower_set: self.lower_set | mask,
89                    lower_bound: self.lower_bound,
90                    below_bound: self.below_bound.clone(),
91                };
92            }
93        } else if offset < BITS_PER_SET * 2 {
94            let mask = 1u64 << (offset - BITS_PER_SET);
95            if (self.upper_set & mask) == 0 {
96                return Self {
97                    upper_set: self.upper_set | mask,
98                    lower_set: self.lower_set,
99                    lower_bound: self.lower_bound,
100                    below_bound: self.below_bound.clone(),
101                };
102            }
103        } else if offset >= BITS_PER_SET * 2 && !self.get(id) {
104            return self.shift_and_set(id);
105        }
106
107        self.clone()
108    }
109
110    /// Remove a snapshot ID from the set (returns a new set if modified).
111    pub fn clear(&self, id: SnapshotId) -> Self {
112        let offset = id.wrapping_sub(self.lower_bound);
113
114        if offset < BITS_PER_SET {
115            let mask = 1u64 << offset;
116            if (self.lower_set & mask) != 0 {
117                return Self {
118                    upper_set: self.upper_set,
119                    lower_set: self.lower_set & !mask,
120                    lower_bound: self.lower_bound,
121                    below_bound: self.below_bound.clone(),
122                };
123            }
124        } else if offset < BITS_PER_SET * 2 {
125            let mask = 1u64 << (offset - BITS_PER_SET);
126            if (self.upper_set & mask) != 0 {
127                return Self {
128                    upper_set: self.upper_set & !mask,
129                    lower_set: self.lower_set,
130                    lower_bound: self.lower_bound,
131                    below_bound: self.below_bound.clone(),
132                };
133            }
134        } else if id < self.lower_bound
135            && let Some(ref arr) = self.below_bound
136            && let Ok(pos) = arr.binary_search(&id)
137        {
138            let mut new_arr = Vec::with_capacity(arr.len() - 1);
139            new_arr.extend_from_slice(&arr[..pos]);
140            new_arr.extend_from_slice(&arr[pos + 1..]);
141            return Self {
142                upper_set: self.upper_set,
143                lower_set: self.lower_set,
144                lower_bound: self.lower_bound,
145                below_bound: if new_arr.is_empty() {
146                    None
147                } else {
148                    Some(new_arr.into_boxed_slice())
149                },
150            };
151        }
152
153        self.clone()
154    }
155
156    /// Remove all IDs in `other` from this set (a & ~b).
157    pub fn and_not(&self, other: &Self) -> Self {
158        if other.is_empty() {
159            return self.clone();
160        }
161        if self.is_empty() {
162            return Self::EMPTY;
163        }
164
165        if self.lower_bound == other.lower_bound && self.below_bound_equals(&other.below_bound) {
166            return Self {
167                upper_set: self.upper_set & !other.upper_set,
168                lower_set: self.lower_set & !other.lower_set,
169                lower_bound: self.lower_bound,
170                below_bound: self.below_bound.clone(),
171            };
172        }
173
174        let mut result = self.clone();
175        for id in other.iter() {
176            result = result.clear(id);
177        }
178        result
179    }
180
181    /// Union this set with another (a | b).
182    pub fn or(&self, other: &Self) -> Self {
183        if other.is_empty() {
184            return self.clone();
185        }
186        if self.is_empty() {
187            return other.clone();
188        }
189
190        if self.lower_bound == other.lower_bound && self.below_bound_equals(&other.below_bound) {
191            return Self {
192                upper_set: self.upper_set | other.upper_set,
193                lower_set: self.lower_set | other.lower_set,
194                lower_bound: self.lower_bound,
195                below_bound: self.below_bound.clone(),
196            };
197        }
198
199        let mut result = self.clone();
200        for id in other.iter() {
201            result = result.set(id);
202        }
203        result
204    }
205
206    /// Find the lowest snapshot ID in the set that is <= upper.
207    pub fn lowest(&self, upper: SnapshotId) -> SnapshotId {
208        if let Some(ref arr) = self.below_bound
209            && let Some(&lowest) = arr.first()
210            && lowest <= upper
211        {
212            return lowest;
213        }
214
215        if self.lower_set != 0 {
216            let lowest_in_lower = self.lower_bound + self.lower_set.trailing_zeros() as usize;
217            if lowest_in_lower <= upper {
218                return lowest_in_lower;
219            }
220        }
221
222        if self.upper_set != 0 {
223            let lowest_in_upper =
224                self.lower_bound + BITS_PER_SET + self.upper_set.trailing_zeros() as usize;
225            if lowest_in_upper <= upper {
226                return lowest_in_upper;
227            }
228        }
229
230        upper
231    }
232
233    /// Check if the set is empty.
234    pub fn is_empty(&self) -> bool {
235        self.lower_set == 0 && self.upper_set == 0 && self.below_bound.is_none()
236    }
237
238    /// Iterate over all snapshot IDs in the set.
239    pub fn iter(&self) -> SnapshotIdSetIter<'_> {
240        SnapshotIdSetIter::new(self)
241    }
242
243    /// Convert to a Vec of snapshot IDs (for testing/debugging).
244    pub fn to_list(&self) -> Vec<SnapshotId> {
245        self.iter().collect()
246    }
247
248    /// Add a contiguous range of IDs [from, until) to the set.
249    /// Mirrors AndroidX SnapshotIdSet.addRange semantics used by Snapshot.kt.
250    pub fn add_range(&self, from: SnapshotId, until: SnapshotId) -> Self {
251        if from >= until {
252            return self.clone();
253        }
254        let mut result = self.clone();
255        let mut id = from;
256        while id < until {
257            result = result.set(id);
258            id += 1;
259        }
260        result
261    }
262
263    fn below_bound_equals(&self, other: &Option<Box<[SnapshotId]>>) -> bool {
264        match (&self.below_bound, other) {
265            (None, None) => true,
266            (Some(a), Some(b)) => a == b,
267            _ => false,
268        }
269    }
270
271    fn shift_and_set(&self, id: SnapshotId) -> Self {
272        let target_lower_bound = (id / SNAPSHOT_ID_SIZE) * SNAPSHOT_ID_SIZE;
273
274        let mut new_upper_set = self.upper_set;
275        let mut new_lower_set = self.lower_set;
276        let mut new_lower_bound = self.lower_bound;
277        let mut new_below_bound: Vec<SnapshotId> = if let Some(ref arr) = self.below_bound {
278            arr.to_vec()
279        } else {
280            Vec::new()
281        };
282
283        while new_lower_bound < target_lower_bound {
284            if new_lower_set != 0 {
285                for bit_offset in 0..BITS_PER_SET {
286                    if (new_lower_set & (1u64 << bit_offset)) != 0 {
287                        let id_to_add = new_lower_bound + bit_offset;
288                        match new_below_bound.binary_search(&id_to_add) {
289                            Ok(_) => {}
290                            Err(pos) => new_below_bound.insert(pos, id_to_add),
291                        }
292                    }
293                }
294            }
295
296            if new_upper_set == 0 {
297                new_lower_bound = target_lower_bound;
298                new_lower_set = 0;
299                break;
300            }
301
302            new_lower_set = new_upper_set;
303            new_upper_set = 0;
304            new_lower_bound += BITS_PER_SET;
305        }
306
307        let result = Self {
308            upper_set: new_upper_set,
309            lower_set: new_lower_set,
310            lower_bound: new_lower_bound,
311            below_bound: if new_below_bound.is_empty() {
312                None
313            } else {
314                Some(new_below_bound.into_boxed_slice())
315            },
316        };
317
318        result.set(id)
319    }
320}
321
322impl Default for SnapshotIdSet {
323    fn default() -> Self {
324        Self::EMPTY
325    }
326}
327
328impl fmt::Debug for SnapshotIdSet {
329    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
330        write!(f, "SnapshotIdSet{{")?;
331        let ids: Vec<_> = self.iter().collect();
332        for (i, id) in ids.iter().enumerate() {
333            if i > 0 {
334                write!(f, ", ")?;
335            }
336            write!(f, "{}", id)?;
337        }
338        write!(f, "}}")
339    }
340}
341
342/// Iterator over snapshot IDs in a set.
343pub struct SnapshotIdSetIter<'a> {
344    set: &'a SnapshotIdSet,
345    below_index: usize,
346    lower_set: u64,
347    upper_set: u64,
348    current_offset: usize,
349}
350
351impl<'a> SnapshotIdSetIter<'a> {
352    fn new(set: &'a SnapshotIdSet) -> Self {
353        Self {
354            set,
355            below_index: 0,
356            lower_set: set.lower_set,
357            upper_set: set.upper_set,
358            current_offset: 0,
359        }
360    }
361}
362
363impl<'a> Iterator for SnapshotIdSetIter<'a> {
364    type Item = SnapshotId;
365
366    fn next(&mut self) -> Option<Self::Item> {
367        if let Some(ref arr) = self.set.below_bound
368            && self.below_index < arr.len()
369        {
370            let id = arr[self.below_index];
371            self.below_index += 1;
372            return Some(id);
373        }
374
375        while self.current_offset < BITS_PER_SET {
376            if (self.lower_set & (1u64 << self.current_offset)) != 0 {
377                let id = self.set.lower_bound + self.current_offset;
378                self.current_offset += 1;
379                return Some(id);
380            }
381            self.current_offset += 1;
382        }
383
384        while self.current_offset < BITS_PER_SET * 2 {
385            let bit_offset = self.current_offset - BITS_PER_SET;
386            if (self.upper_set & (1u64 << bit_offset)) != 0 {
387                let id = self.set.lower_bound + self.current_offset;
388                self.current_offset += 1;
389                return Some(id);
390            }
391            self.current_offset += 1;
392        }
393
394        None
395    }
396}
397
398#[cfg(test)]
399mod tests {
400    use super::*;
401
402    #[test]
403    fn test_empty_set() {
404        let set = SnapshotIdSet::EMPTY;
405        assert!(set.is_empty());
406        assert!(!set.get(0));
407        assert!(!set.get(100));
408    }
409
410    #[test]
411    fn test_set_and_get_lower_range() {
412        let set = SnapshotIdSet::new();
413        let set = set.set(0);
414        assert!(set.get(0));
415        assert!(!set.get(1));
416
417        let set = set.set(63);
418        assert!(set.get(0));
419        assert!(set.get(63));
420        assert!(!set.get(64));
421    }
422
423    #[test]
424    fn test_set_and_get_upper_range() {
425        let set = SnapshotIdSet::new();
426        let set = set.set(64);
427        assert!(set.get(64));
428        assert!(!set.get(63));
429        assert!(!set.get(128));
430
431        let set = set.set(127);
432        assert!(set.get(64));
433        assert!(set.get(127));
434        assert!(!set.get(128));
435    }
436
437    #[test]
438    fn test_set_idempotent() {
439        let set = SnapshotIdSet::new();
440        let set1 = set.set(10);
441        let set2 = set1.set(10);
442        assert_eq!(set1, set2);
443    }
444
445    #[test]
446    fn test_clear() {
447        let set = SnapshotIdSet::new().set(10).set(20).set(30);
448        assert!(set.get(10));
449        assert!(set.get(20));
450        assert!(set.get(30));
451
452        let set = set.clear(20);
453        assert!(set.get(10));
454        assert!(!set.get(20));
455        assert!(set.get(30));
456    }
457
458    #[test]
459    fn test_clear_idempotent() {
460        let set = SnapshotIdSet::new().set(10);
461        let set1 = set.clear(10);
462        let set2 = set1.clear(10);
463        assert_eq!(set1, set2);
464    }
465
466    #[test]
467    fn test_below_bound_insertion() {
468        let mut set = SnapshotIdSet::new();
469        set = set.set(100);
470        assert_eq!(set.lower_bound, 0);
471
472        set = set.set(50);
473        assert!(set.get(50));
474        assert!(set.get(100));
475
476        set = set.set(25);
477        set = set.set(75);
478        assert!(set.get(25));
479        assert!(set.get(50));
480        assert!(set.get(75));
481        assert!(set.get(100));
482
483        let list = set.to_list();
484        assert_eq!(list, vec![25, 50, 75, 100]);
485    }
486
487    #[test]
488    fn test_below_bound_removal() {
489        let set = SnapshotIdSet::new();
490        let set = set.set(25);
491        let set = set.set(50);
492        let set = set.set(75);
493        let set = set.set(200);
494
495        let set = set.clear(50);
496        assert!(set.get(25));
497        assert!(!set.get(50));
498        assert!(set.get(75));
499        assert!(set.get(200));
500
501        let list = set.to_list();
502        assert_eq!(list, vec![25, 75, 200]);
503    }
504
505    #[test]
506    fn test_shift_and_set() {
507        let set = SnapshotIdSet::new();
508        let set = set.set(10);
509        assert_eq!(set.lower_bound, 0);
510
511        let set = set.set(200);
512        assert!(set.get(10));
513        assert!(set.get(200));
514
515        assert!(set.below_bound.is_some());
516    }
517
518    #[test]
519    fn test_shift_and_set_boundary_values() {
520        let mut set = SnapshotIdSet::new();
521        let boundary = SNAPSHOT_ID_SIZE * 12 - 1;
522        set = set.set(boundary);
523        assert!(set.get(boundary));
524
525        set = set.set(boundary + 1);
526        assert!(set.get(boundary));
527        assert!(set.get(boundary + 1));
528    }
529
530    #[test]
531    fn test_set_below_lower_bound_inserts() {
532        let set = SnapshotIdSet::new().set(200);
533        let lower_bound = set.lower_bound;
534        assert!(lower_bound > 0);
535
536        let below = lower_bound - 1;
537        let set = set.set(below);
538        assert!(set.get(below));
539        assert!(set.get(200));
540    }
541
542    #[test]
543    fn test_and_not_fast_path() {
544        let set1 = SnapshotIdSet::new().set(10).set(20).set(30);
545        let set2 = SnapshotIdSet::new().set(20).set(40);
546
547        let result = set1.and_not(&set2);
548        assert!(result.get(10));
549        assert!(!result.get(20));
550        assert!(result.get(30));
551        assert!(!result.get(40));
552    }
553
554    #[test]
555    fn test_and_not_slow_path() {
556        let set1 = SnapshotIdSet::new().set(10).set(20).set(30);
557        let set2 = SnapshotIdSet::new().set(100).set(20);
558
559        let result = set1.and_not(&set2);
560        assert!(result.get(10));
561        assert!(!result.get(20));
562        assert!(result.get(30));
563    }
564
565    #[test]
566    fn test_or_fast_path() {
567        let set1 = SnapshotIdSet::new().set(10).set(20);
568        let set2 = SnapshotIdSet::new().set(20).set(30);
569
570        let result = set1.or(&set2);
571        assert!(result.get(10));
572        assert!(result.get(20));
573        assert!(result.get(30));
574    }
575
576    #[test]
577    fn test_or_slow_path() {
578        let set1 = SnapshotIdSet::new().set(10).set(20);
579        let set2 = SnapshotIdSet::new().set(100).set(30);
580
581        let result = set1.or(&set2);
582        assert!(result.get(10));
583        assert!(result.get(20));
584        assert!(result.get(30));
585        assert!(result.get(100));
586    }
587
588    #[test]
589    fn test_lowest_in_below_bound() {
590        let set = SnapshotIdSet::new();
591        let set = set.set(25);
592        let set = set.set(50);
593        let set = set.set(200);
594        assert_eq!(set.lowest(1000), 25);
595        assert_eq!(set.lowest(100), 25);
596        assert_eq!(set.lowest(30), 25);
597    }
598
599    #[test]
600    fn test_lowest_in_lower_set() {
601        let set = SnapshotIdSet::new().set(10).set(20).set(30);
602        assert_eq!(set.lowest(1000), 10);
603        assert_eq!(set.lowest(25), 10);
604    }
605
606    #[test]
607    fn test_lowest_in_upper_set() {
608        let set = SnapshotIdSet::new().set(70).set(80).set(90);
609        assert_eq!(set.lowest(1000), 70);
610    }
611
612    #[test]
613    fn test_lowest_returns_upper_if_none_found() {
614        let set = SnapshotIdSet::new().set(100);
615        assert_eq!(set.lowest(50), 50);
616    }
617
618    #[test]
619    fn test_iterator() {
620        let set = SnapshotIdSet::new().set(10).set(20).set(5).set(30);
621        let list: Vec<_> = set.iter().collect();
622        assert_eq!(list, vec![5, 10, 20, 30]);
623    }
624
625    #[test]
626    fn test_iterator_empty() {
627        let set = SnapshotIdSet::new();
628        let list: Vec<_> = set.iter().collect();
629        assert_eq!(list, Vec::<SnapshotId>::new());
630    }
631
632    #[test]
633    fn test_iterator_all_ranges() {
634        let set = SnapshotIdSet::new().set(5).set(10).set(70).set(200);
635
636        let list: Vec<_> = set.iter().collect();
637        assert_eq!(list, vec![5, 10, 70, 200]);
638    }
639
640    #[test]
641    fn test_to_list() {
642        let set = SnapshotIdSet::new().set(10).set(20).set(30);
643        assert_eq!(set.to_list(), vec![10, 20, 30]);
644    }
645
646    #[test]
647    fn test_debug_format() {
648        let set = SnapshotIdSet::new().set(10).set(20);
649        let debug_str = format!("{:?}", set);
650        assert_eq!(debug_str, "SnapshotIdSet{10, 20}");
651    }
652
653    #[test]
654    fn test_large_snapshot_ids() {
655        let set = SnapshotIdSet::new();
656        let set = set.set(500);
657        let set = set.set(1000);
658        let set = set.set(2000);
659
660        assert!(set.get(500));
661        assert!(set.get(1000));
662        assert!(set.get(2000));
663        assert!(!set.get(1500));
664    }
665
666    #[test]
667    fn test_boundary_transitions() {
668        let set = SnapshotIdSet::new();
669
670        let set = set.set(63);
671        let set = set.set(64);
672        assert!(set.get(63));
673        assert!(set.get(64));
674
675        let set = set.set(127);
676        let set = set.set(128);
677        assert!(set.get(127));
678        assert!(set.get(128));
679    }
680}