Skip to main content

defuse_bitmap/
lib.rs

1mod b256;
2
3pub use self::b256::*;
4
5use core::ops::{DerefMut, RangeInclusive, Shl};
6
7use defuse_map_utils::{IterableMap, Map, cleanup::DefaultMap};
8use num_traits::{One, PrimInt, Zero};
9
10/// Bitmap for primitive types
11#[cfg_attr(feature = "arbitrary", derive(::arbitrary::Arbitrary))]
12#[cfg_attr(
13    feature = "borsh",
14    derive(::borsh::BorshSerialize, ::borsh::BorshDeserialize),
15    cfg_attr(feature = "borsh-schema", derive(::borsh::BorshSchema))
16)]
17#[derive(Debug, Clone, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
18#[repr(transparent)]
19pub struct BitMap<M>(M);
20
21impl<M> BitMap<M>
22where
23    M: DefaultMap<K = <M as Map>::V>,
24    M::K: PrimInt + Shl<M::K, Output = M::K>,
25{
26    #[allow(clippy::as_conversions)]
27    const BITS_FOR_BIT_POS: usize = (size_of::<M::K>() * 8).ilog2() as usize;
28
29    #[inline]
30    pub const fn new(map: M) -> Self {
31        Self(map)
32    }
33
34    /// Get the bit `n`
35    ///
36    /// # Examples
37    ///
38    /// ```rust
39    /// # use std::collections::BTreeMap;
40    /// # use defuse_bitmap::BitMap;
41    /// let mut m = BitMap::<BTreeMap<u32, u32>>::default();
42    ///
43    /// assert!(!m.get_bit(42));
44    /// assert!(!m.set_bit(42));
45    /// assert!(m.get_bit(42));
46    /// ```
47    #[inline]
48    pub fn get_bit(&self, n: M::K) -> bool {
49        let (word, bit_mask) = Self::split_word_mask(n);
50        let Some(bitmap) = self.0.get(&word) else {
51            return false;
52        };
53        *bitmap & bit_mask != M::V::zero()
54    }
55
56    /// Set the bit `n` and return old value
57    ///
58    /// # Examples
59    ///
60    /// ```rust
61    /// # use std::collections::BTreeMap;
62    /// # use defuse_bitmap::BitMap;
63    /// let mut m = BitMap::<BTreeMap<u32, u32>>::default();
64    ///
65    /// assert!(!m.set_bit(42));
66    /// assert!(m.get_bit(42));
67    /// ```
68    #[inline]
69    pub fn set_bit(&mut self, n: M::K) -> bool {
70        let (mut bitmap, mask) = self.get_mut_with_mask(n);
71        let old = *bitmap & mask != M::V::zero();
72        *bitmap = *bitmap | mask;
73        old
74    }
75
76    /// Clear the bit `n` and return old value
77    ///
78    /// # Examples
79    ///
80    /// ```rust
81    /// # use std::collections::BTreeMap;
82    /// # use defuse_bitmap::BitMap;
83    /// let mut m = BitMap::<BTreeMap<u32, u32>>::default();
84    ///
85    /// assert!(!m.clear_bit(42));
86    /// assert!(!m.get_bit(42));
87    /// ```
88    #[inline]
89    pub fn clear_bit(&mut self, n: M::K) -> bool {
90        let (mut bitmap, mask) = self.get_mut_with_mask(n);
91        let old = *bitmap & mask != M::V::zero();
92        *bitmap = *bitmap & !mask;
93        old
94    }
95
96    /// Toggle the bit `n` and return old value
97    ///
98    /// # Examples
99    ///
100    /// ```rust
101    /// # use std::collections::BTreeMap;
102    /// # use defuse_bitmap::BitMap;
103    /// let mut m = BitMap::<BTreeMap<u32, u32>>::default();
104    ///
105    /// assert!(!m.toggle_bit(42));
106    /// assert!(m.toggle_bit(42));
107    /// assert!(!m.get_bit(42));
108    /// ```
109    #[inline]
110    pub fn toggle_bit(&mut self, n: M::K) -> bool {
111        let (mut bitmap, mask) = self.get_mut_with_mask(n);
112        let old = *bitmap & mask != M::V::zero();
113        *bitmap = *bitmap ^ mask;
114        old
115    }
116
117    /// Set bit `n` to given value and return old value
118    ///
119    /// # Examples
120    ///
121    /// ```rust
122    /// # use std::collections::BTreeMap;
123    /// # use defuse_bitmap::BitMap;
124    /// let mut m = BitMap::<BTreeMap<u32, u32>>::default();
125    ///
126    /// assert!(!m.set_bit_to(42, true));
127    /// assert!(m.set_bit_to(42, false));
128    /// assert!(!m.get_bit(42));
129    /// ```
130    #[inline]
131    pub fn set_bit_to(&mut self, n: M::K, v: bool) -> bool {
132        if v {
133            self.set_bit(n)
134        } else {
135            self.clear_bit(n)
136        }
137    }
138
139    /// Iterate over set bits
140    ///
141    /// # Examples
142    ///
143    /// ```rust
144    /// # use std::collections::BTreeMap;
145    /// # use defuse_bitmap::BitMap;
146    /// let mut m = BitMap::<BTreeMap<u32, u32>>::default();
147    /// for n in [100, 15, 1, 24, 0, 717, 999, u32::MAX] {
148    ///     assert!(!m.set_bit(n));
149    /// }
150    ///
151    /// assert_eq!(
152    ///     m.as_iter().collect::<Vec<_>>(),
153    ///     vec![0, 1, 15, 24, 100, 717, 999, u32::MAX],
154    /// );
155    /// ```
156    pub fn as_iter(&self) -> impl Iterator<Item = M::V> + '_
157    where
158        M: IterableMap,
159        RangeInclusive<M::V>: Iterator<Item = M::V>,
160    {
161        self.0.iter().flat_map(|(prefix, bitmap)| {
162            (M::V::zero()..=Self::bit_pos_mask())
163                .filter(|&bit_pos| {
164                    let bit_mask = M::V::one() << bit_pos;
165                    *bitmap & bit_mask != M::V::zero()
166                })
167                .map(|bit_pos| (*prefix << Self::BITS_FOR_BIT_POS) | bit_pos)
168        })
169    }
170
171    #[inline]
172    fn get_mut_with_mask(&mut self, n: M::K) -> (impl DerefMut<Target = M::V>, M::V) {
173        let (word, bit_mask) = Self::split_word_mask(n);
174        (self.0.entry_or_default(word), bit_mask)
175    }
176
177    /// Returns `(word, bit_pos_mask)`
178    #[inline]
179    fn split_word_mask(n: M::K) -> (M::K, M::V) {
180        let word = n >> Self::BITS_FOR_BIT_POS;
181        let bit_mask = M::V::one() << (n & Self::bit_pos_mask());
182        (word, bit_mask)
183    }
184
185    #[inline]
186    fn bit_pos_mask() -> M::V {
187        (M::V::one() << Self::BITS_FOR_BIT_POS) - M::V::one()
188    }
189}
190
191#[cfg(test)]
192mod tests {
193    use std::collections::BTreeMap;
194
195    use rstest::rstest;
196
197    use super::*;
198
199    #[allow(clippy::used_underscore_binding)]
200    #[rstest]
201    fn test<T>(#[values(0u8, 0u16, 0u32, 0u64, 0u128)] _n: T)
202    where
203        T: PrimInt + Shl<T, Output = T> + Default,
204    {
205        let mut m = BitMap::<BTreeMap<T, T>>::default();
206
207        for n in [
208            T::zero(),
209            T::one(),
210            T::max_value() - T::one(),
211            T::max_value(),
212        ] {
213            assert!(!m.get_bit(n));
214
215            assert!(!m.set_bit(n));
216            assert!(m.get_bit(n));
217            assert!(m.set_bit(n));
218            assert!(m.get_bit(n));
219
220            assert!(m.clear_bit(n));
221            assert!(!m.get_bit(n));
222            assert!(!m.clear_bit(n));
223            assert!(!m.get_bit(n));
224        }
225    }
226}