Skip to main content

refaults_vmo/
atomic_vec.rs

1// Copyright 2026 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use std::cmp::min;
6use std::sync::atomic::{AtomicU64, Ordering};
7
8const BITS: u64 = u64::BITS as u64;
9
10/// An atomic bit-vector.
11pub struct AtomicBitVec {
12    storage: Box<[AtomicU64]>,
13    nbits: u64,
14}
15
16impl AtomicBitVec {
17    /// Creates a new `AtomicBitVec` with all bits set to false.
18    pub fn new(nbits: u64) -> Self {
19        let nwords = nbits.div_ceil(BITS);
20        let storage = (0..nwords).map(|_| AtomicU64::new(0)).collect();
21        Self { storage, nbits }
22    }
23
24    /// Sets the bits between `start_bit` (included) and `end_bit` (excluded), and returns the
25    /// number of bits that were already set.
26    pub fn test_and_set_range(&self, start_bit: u64, end_bit: u64) -> u64 {
27        assert!(start_bit < end_bit);
28        assert!(end_bit <= self.nbits);
29
30        let mut counter = 0;
31        let mut current_bit = start_bit;
32
33        while current_bit < end_bit {
34            let current_word_index = current_bit / BITS;
35            let current_word_start_bit = current_word_index * BITS;
36            let mask =
37                Self::get_mask(current_bit % BITS, min(end_bit - current_word_start_bit, BITS));
38            let old_word =
39                &self.storage[current_word_index as usize].fetch_or(mask, Ordering::Relaxed);
40            counter += (old_word & mask).count_ones();
41            current_bit = current_word_start_bit + BITS;
42        }
43
44        counter.into()
45    }
46
47    pub fn len(&self) -> u64 {
48        self.nbits
49    }
50
51    pub fn get(&self) -> Vec<bool> {
52        let fetch = |bit: u64| {
53            let word = bit / BITS;
54            let bit_mask = 1 << (bit % BITS);
55            self.storage[word as usize].fetch_or(0, Ordering::Relaxed) & bit_mask != 0
56        };
57        (0..self.nbits).map(fetch).collect()
58    }
59
60    fn get_mask(start_bit: u64, end_bit: u64) -> u64 {
61        let left_mask = u64::MAX << start_bit;
62        let right_mask = u64::MAX >> (BITS - end_bit);
63        left_mask & right_mask
64    }
65}
66
67impl Clone for AtomicBitVec {
68    fn clone(&self) -> Self {
69        let new_storage =
70            self.storage.iter().map(|a| AtomicU64::new(a.load(Ordering::Relaxed))).collect();
71        Self { storage: new_storage, nbits: self.nbits }
72    }
73}
74
75#[cfg(test)]
76mod tests {
77    use super::*;
78
79    #[test]
80    fn test_new() {
81        let vec = AtomicBitVec::new(0);
82        assert_eq!(vec.nbits, 0);
83        assert_eq!(vec.storage.len(), 0);
84
85        let vec = AtomicBitVec::new(1);
86        assert_eq!(vec.nbits, 1);
87        assert_eq!(vec.storage.len(), 1);
88        assert_eq!(vec.storage[0].load(Ordering::Relaxed), 0);
89
90        let vec = AtomicBitVec::new(64);
91        assert_eq!(vec.nbits, 64);
92        assert_eq!(vec.storage.len(), 1);
93        assert_eq!(vec.storage[0].load(Ordering::Relaxed), 0);
94
95        let vec = AtomicBitVec::new(65);
96        assert_eq!(vec.nbits, 65);
97        assert_eq!(vec.storage.len(), 2);
98        assert_eq!(vec.storage[0].load(Ordering::Relaxed), 0);
99        assert_eq!(vec.storage[1].load(Ordering::Relaxed), 0);
100    }
101
102    #[test]
103    fn test_test_and_set() {
104        let vec = AtomicBitVec::new(320);
105
106        // Set bits and check they were not set before.
107        assert_eq!(vec.test_and_set_range(10, 20), 0);
108        // Check they are set now.
109        assert_eq!(vec.test_and_set_range(10, 20), 10);
110
111        // Check another range, partially overlapping.
112        assert_eq!(vec.test_and_set_range(15, 25), 5);
113
114        // A range across two words.
115        assert_eq!(vec.test_and_set_range(60, 70), 0);
116        // Only 10 more.
117        assert_eq!(vec.test_and_set_range(55, 75), 10);
118
119        // Large range.
120        assert_eq!(vec.test_and_set_range(50, 300), 20);
121        assert_eq!(vec.test_and_set_range(50, 300), 250);
122    }
123
124    #[test]
125    fn test_test_and_set2() {
126        let vec = AtomicBitVec::new(320);
127        assert_eq!(vec.test_and_set_range(64, 128), 0);
128        assert_eq!(vec.test_and_set_range(64, 128), 64);
129    }
130
131    #[test]
132    fn test_test_and_set3() {
133        let vec = AtomicBitVec::new(150);
134
135        assert_eq!(vec.test_and_set_range(0, 1), 0);
136        assert_eq!(vec.test_and_set_range(63, 64), 0);
137        // 0 and 63 are already set.
138        assert_eq!(vec.test_and_set_range(0, 64), 2);
139        assert_eq!(vec.test_and_set_range(64, 65), 0);
140        // 63 and 64 are already set.
141        assert_eq!(vec.test_and_set_range(63, 65), 2);
142
143        // 63 and 64 are already set.
144        assert_eq!(vec.test_and_set_range(63, 128), 2);
145        // 63 to 127 (included) are already set.
146        assert_eq!(vec.test_and_set_range(63, 129), 65);
147    }
148
149    #[test]
150    #[should_panic]
151    fn test_and_set_out_of_bounds() {
152        let vec = AtomicBitVec::new(10);
153        vec.test_and_set_range(10, 20);
154    }
155
156    #[test]
157    fn test_clone() {
158        let vec = AtomicBitVec::new(100);
159        vec.test_and_set_range(10, 20);
160        vec.test_and_set_range(50, 60);
161
162        let clone = vec.clone();
163        assert_eq!(clone.nbits, vec.nbits);
164        assert_eq!(clone.test_and_set_range(10, 20), 10);
165        assert_eq!(clone.test_and_set_range(50, 60), 10);
166        assert_eq!(clone.test_and_set_range(20, 30), 0);
167
168        // Check the original is unaffected.
169        assert_eq!(vec.test_and_set_range(20, 30), 0);
170    }
171}