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                .is_some_and(|arr| arr.binary_search(&id).is_ok())
46        }
47    }
48
49    /// Add a snapshot ID to the set (returns a new set if modified).
50    pub fn set(&self, id: SnapshotId) -> Self {
51        if id < self.lower_bound {
52            if let Some(ref arr) = self.below_bound {
53                match arr.binary_search(&id) {
54                    Ok(_) => {
55                        return self.clone();
56                    }
57                    Err(insert_pos) => {
58                        let mut new_arr = Vec::with_capacity(arr.len() + 1);
59                        new_arr.extend_from_slice(&arr[..insert_pos]);
60                        new_arr.push(id);
61                        new_arr.extend_from_slice(&arr[insert_pos..]);
62                        return Self {
63                            upper_set: self.upper_set,
64                            lower_set: self.lower_set,
65                            lower_bound: self.lower_bound,
66                            below_bound: Some(new_arr.into_boxed_slice()),
67                        };
68                    }
69                }
70            } else {
71                return Self {
72                    upper_set: self.upper_set,
73                    lower_set: self.lower_set,
74                    lower_bound: self.lower_bound,
75                    below_bound: Some(vec![id].into_boxed_slice()),
76                };
77            }
78        }
79
80        let offset = id - self.lower_bound;
81
82        if offset < BITS_PER_SET {
83            let mask = 1u64 << offset;
84            if (self.lower_set & mask) == 0 {
85                return Self {
86                    upper_set: self.upper_set,
87                    lower_set: self.lower_set | mask,
88                    lower_bound: self.lower_bound,
89                    below_bound: self.below_bound.clone(),
90                };
91            }
92        } else if offset < BITS_PER_SET * 2 {
93            let mask = 1u64 << (offset - BITS_PER_SET);
94            if (self.upper_set & mask) == 0 {
95                return Self {
96                    upper_set: self.upper_set | mask,
97                    lower_set: self.lower_set,
98                    lower_bound: self.lower_bound,
99                    below_bound: self.below_bound.clone(),
100                };
101            }
102        } else if offset >= BITS_PER_SET * 2 && !self.get(id) {
103            return self.shift_and_set(id);
104        }
105
106        self.clone()
107    }
108
109    /// Remove a snapshot ID from the set (returns a new set if modified).
110    pub fn clear(&self, id: SnapshotId) -> Self {
111        let offset = id.wrapping_sub(self.lower_bound);
112
113        if offset < BITS_PER_SET {
114            let mask = 1u64 << offset;
115            if (self.lower_set & mask) != 0 {
116                return Self {
117                    upper_set: self.upper_set,
118                    lower_set: self.lower_set & !mask,
119                    lower_bound: self.lower_bound,
120                    below_bound: self.below_bound.clone(),
121                };
122            }
123        } else if offset < BITS_PER_SET * 2 {
124            let mask = 1u64 << (offset - BITS_PER_SET);
125            if (self.upper_set & mask) != 0 {
126                return Self {
127                    upper_set: self.upper_set & !mask,
128                    lower_set: self.lower_set,
129                    lower_bound: self.lower_bound,
130                    below_bound: self.below_bound.clone(),
131                };
132            }
133        } else if id < self.lower_bound
134            && let Some(ref arr) = self.below_bound
135            && let Ok(pos) = arr.binary_search(&id)
136        {
137            let mut new_arr = Vec::with_capacity(arr.len() - 1);
138            new_arr.extend_from_slice(&arr[..pos]);
139            new_arr.extend_from_slice(&arr[pos + 1..]);
140            return Self {
141                upper_set: self.upper_set,
142                lower_set: self.lower_set,
143                lower_bound: self.lower_bound,
144                below_bound: if new_arr.is_empty() {
145                    None
146                } else {
147                    Some(new_arr.into_boxed_slice())
148                },
149            };
150        }
151
152        self.clone()
153    }
154
155    /// Remove all IDs in `other` from this set (a & ~b).
156    pub fn and_not(&self, other: &Self) -> Self {
157        if other.is_empty() {
158            return self.clone();
159        }
160        if self.is_empty() {
161            return Self::EMPTY;
162        }
163
164        if self.lower_bound == other.lower_bound && self.below_bound_equals(&other.below_bound) {
165            return Self {
166                upper_set: self.upper_set & !other.upper_set,
167                lower_set: self.lower_set & !other.lower_set,
168                lower_bound: self.lower_bound,
169                below_bound: self.below_bound.clone(),
170            };
171        }
172
173        let mut result = self.clone();
174        for id in other.iter() {
175            result = result.clear(id);
176        }
177        result
178    }
179
180    /// Union this set with another (a | b).
181    pub fn or(&self, other: &Self) -> Self {
182        if other.is_empty() {
183            return self.clone();
184        }
185        if self.is_empty() {
186            return other.clone();
187        }
188
189        if self.lower_bound == other.lower_bound && self.below_bound_equals(&other.below_bound) {
190            return Self {
191                upper_set: self.upper_set | other.upper_set,
192                lower_set: self.lower_set | other.lower_set,
193                lower_bound: self.lower_bound,
194                below_bound: self.below_bound.clone(),
195            };
196        }
197
198        let mut result = self.clone();
199        for id in other.iter() {
200            result = result.set(id);
201        }
202        result
203    }
204
205    /// Find the lowest snapshot ID in the set that is <= upper.
206    pub fn lowest(&self, upper: SnapshotId) -> SnapshotId {
207        if let Some(ref arr) = self.below_bound
208            && let Some(&lowest) = arr.first()
209            && lowest <= upper
210        {
211            return lowest;
212        }
213
214        if self.lower_set != 0 {
215            let lowest_in_lower = self.lower_bound + self.lower_set.trailing_zeros() as usize;
216            if lowest_in_lower <= upper {
217                return lowest_in_lower;
218            }
219        }
220
221        if self.upper_set != 0 {
222            let lowest_in_upper =
223                self.lower_bound + BITS_PER_SET + self.upper_set.trailing_zeros() as usize;
224            if lowest_in_upper <= upper {
225                return lowest_in_upper;
226            }
227        }
228
229        upper
230    }
231
232    /// Check if the set is empty.
233    pub fn is_empty(&self) -> bool {
234        self.lower_set == 0 && self.upper_set == 0 && self.below_bound.is_none()
235    }
236
237    /// Iterate over all snapshot IDs in the set.
238    pub fn iter(&self) -> SnapshotIdSetIter<'_> {
239        SnapshotIdSetIter::new(self)
240    }
241
242    /// Convert to a Vec of snapshot IDs (for testing/debugging).
243    pub fn to_list(&self) -> Vec<SnapshotId> {
244        self.iter().collect()
245    }
246
247    /// Add a contiguous range of IDs [from, until) to the set.
248    /// Mirrors AndroidX SnapshotIdSet.addRange semantics used by Snapshot.kt.
249    pub fn add_range(&self, from: SnapshotId, until: SnapshotId) -> Self {
250        if from >= until {
251            return self.clone();
252        }
253        let mut result = self.clone();
254        let mut id = from;
255        while id < until {
256            result = result.set(id);
257            id += 1;
258        }
259        result
260    }
261
262    fn below_bound_equals(&self, other: &Option<Box<[SnapshotId]>>) -> bool {
263        match (&self.below_bound, other) {
264            (None, None) => true,
265            (Some(a), Some(b)) => a == b,
266            _ => false,
267        }
268    }
269
270    fn shift_and_set(&self, id: SnapshotId) -> Self {
271        let target_lower_bound = (id / SNAPSHOT_ID_SIZE) * SNAPSHOT_ID_SIZE;
272
273        let mut new_upper_set = self.upper_set;
274        let mut new_lower_set = self.lower_set;
275        let mut new_lower_bound = self.lower_bound;
276        let mut new_below_bound: Vec<SnapshotId> = if let Some(ref arr) = self.below_bound {
277            arr.to_vec()
278        } else {
279            Vec::new()
280        };
281
282        while new_lower_bound < target_lower_bound {
283            if new_lower_set != 0 {
284                for bit_offset in 0..BITS_PER_SET {
285                    if (new_lower_set & (1u64 << bit_offset)) != 0 {
286                        let id_to_add = new_lower_bound + bit_offset;
287                        match new_below_bound.binary_search(&id_to_add) {
288                            Ok(_) => {}
289                            Err(pos) => new_below_bound.insert(pos, id_to_add),
290                        }
291                    }
292                }
293            }
294
295            if new_upper_set == 0 {
296                new_lower_bound = target_lower_bound;
297                new_lower_set = 0;
298                break;
299            }
300
301            new_lower_set = new_upper_set;
302            new_upper_set = 0;
303            new_lower_bound += BITS_PER_SET;
304        }
305
306        let result = Self {
307            upper_set: new_upper_set,
308            lower_set: new_lower_set,
309            lower_bound: new_lower_bound,
310            below_bound: if new_below_bound.is_empty() {
311                None
312            } else {
313                Some(new_below_bound.into_boxed_slice())
314            },
315        };
316
317        result.set(id)
318    }
319}
320
321impl Default for SnapshotIdSet {
322    fn default() -> Self {
323        Self::EMPTY
324    }
325}
326
327impl fmt::Debug for SnapshotIdSet {
328    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
329        write!(f, "SnapshotIdSet{{")?;
330        let ids: Vec<_> = self.iter().collect();
331        for (i, id) in ids.iter().enumerate() {
332            if i > 0 {
333                write!(f, ", ")?;
334            }
335            write!(f, "{id}")?;
336        }
337        write!(f, "}}")
338    }
339}
340
341/// Iterator over snapshot IDs in a set.
342pub struct SnapshotIdSetIter<'a> {
343    set: &'a SnapshotIdSet,
344    below_index: usize,
345    lower_set: u64,
346    upper_set: u64,
347    current_offset: usize,
348}
349
350impl<'a> SnapshotIdSetIter<'a> {
351    fn new(set: &'a SnapshotIdSet) -> Self {
352        Self {
353            set,
354            below_index: 0,
355            lower_set: set.lower_set,
356            upper_set: set.upper_set,
357            current_offset: 0,
358        }
359    }
360}
361
362impl Iterator for SnapshotIdSetIter<'_> {
363    type Item = SnapshotId;
364
365    fn next(&mut self) -> Option<Self::Item> {
366        if let Some(ref arr) = self.set.below_bound
367            && self.below_index < arr.len()
368        {
369            let id = arr[self.below_index];
370            self.below_index += 1;
371            return Some(id);
372        }
373
374        while self.current_offset < BITS_PER_SET {
375            if (self.lower_set & (1u64 << self.current_offset)) != 0 {
376                let id = self.set.lower_bound + self.current_offset;
377                self.current_offset += 1;
378                return Some(id);
379            }
380            self.current_offset += 1;
381        }
382
383        while self.current_offset < BITS_PER_SET * 2 {
384            let bit_offset = self.current_offset - BITS_PER_SET;
385            if (self.upper_set & (1u64 << bit_offset)) != 0 {
386                let id = self.set.lower_bound + self.current_offset;
387                self.current_offset += 1;
388                return Some(id);
389            }
390            self.current_offset += 1;
391        }
392
393        None
394    }
395}
396
397#[cfg(test)]
398#[path = "tests/snapshot_id_set_tests.rs"]
399mod tests;