Skip to main content

competitive/data_structure/
bitset.rs

1#[cfg(target_arch = "x86_64")]
2use super::avx512_enabled;
3use std::{
4    cmp::Ordering,
5    ops::{
6        BitAnd, BitAndAssign, BitOr, BitOrAssign, BitXor, BitXorAssign, Not, Shl, ShlAssign, Shr,
7        ShrAssign,
8    },
9};
10
11const BIT_AND: u8 = 0;
12const BIT_OR: u8 = 1;
13const BIT_XOR: u8 = 2;
14#[cfg(target_arch = "x86_64")]
15const SIMD_MIN_BLOCKS: usize = 8;
16
17#[repr(C, align(64))]
18#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
19struct Block([u64; 8]);
20
21#[derive(Clone, Debug, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
22pub struct BitSet {
23    size: usize,
24    bits: Vec<Block>,
25}
26
27impl BitSet {
28    pub fn new(size: usize) -> Self {
29        Self {
30            size,
31            bits: vec![Block::default(); size.div_ceil(512)],
32        }
33    }
34
35    pub fn len(&self) -> usize {
36        self.size
37    }
38
39    pub fn is_empty(&self) -> bool {
40        self.size == 0
41    }
42
43    /// Parses ASCII `0` and `1`, with the first character at bit index zero.
44    /// Returns `None` if any other character occurs.
45    pub fn from_binary(s: &str) -> Option<Self> {
46        let bytes = s.as_bytes();
47        let mut bits = Self::new(bytes.len());
48        let end = bytes.len() / 64 * 64;
49        #[cfg(target_arch = "x86_64")]
50        let parsed = if avx512_enabled() && is_x86_feature_detected!("avx512bw") {
51            // SAFETY: AVX-512BW is available; each complete chunk contains 64 bytes.
52            unsafe { simd::parse_binary_avx512(&bytes[..end], bits.words_mut()) }
53        } else if is_x86_feature_detected!("avx2") {
54            // SAFETY: AVX2 is available; each complete chunk contains 64 bytes.
55            unsafe { simd::parse_binary_avx2(&bytes[..end], bits.words_mut()) }
56        } else {
57            Self::parse_binary_scalar(&bytes[..end], bits.words_mut())
58        };
59        #[cfg(not(target_arch = "x86_64"))]
60        let parsed = Self::parse_binary_scalar(&bytes[..end], bits.words_mut());
61        if !parsed {
62            return None;
63        }
64        if end != bytes.len() {
65            let mut word = 0;
66            for (i, &b) in bytes[end..].iter().enumerate() {
67                if b != b'0' && b != b'1' {
68                    return None;
69                }
70                word |= u64::from(b & 1) << i;
71            }
72            bits.words_mut()[end / 64] = word;
73        }
74        Some(bits)
75    }
76
77    fn parse_binary_scalar(bytes: &[u8], words: &mut [u64]) -> bool {
78        for (chunk, word) in bytes.as_chunks::<64>().0.iter().zip(words) {
79            for (i, byte) in chunk.as_chunks::<8>().0.iter().enumerate() {
80                let x = u64::from_le_bytes(*byte);
81                if x & 0xfefe_fefe_fefe_fefe != 0x3030_3030_3030_3030 {
82                    return false;
83                }
84                *word |= ((x & 0x0101_0101_0101_0101).wrapping_mul(0x0102_0408_1020_4080) >> 56)
85                    << (i * 8);
86            }
87        }
88        true
89    }
90
91    /// Returns ASCII `0` and `1` in increasing bit-index order.
92    pub fn to_binary(&self) -> String {
93        const TABLE: [[u8; 8]; 256] = {
94            let mut table = [[b'0'; 8]; 256];
95            let mut i = 0;
96            while i < 256 {
97                let mut j = 0;
98                while j < 8 {
99                    table[i][j] |= ((i >> j) & 1) as u8;
100                    j += 1;
101                }
102                i += 1;
103            }
104            table
105        };
106        let mut bytes = vec![b'0'; self.size.div_ceil(8) * 8];
107        #[cfg(target_arch = "x86_64")]
108        let end = if self.size >= 64 && avx512_enabled() && is_x86_feature_detected!("avx512bw") {
109            let end = self.size / 64 * 64;
110            // SAFETY: AVX-512BW is available; each output chunk holds one 64-bit word.
111            unsafe {
112                simd::write_binary_avx512(&mut bytes[..end], self.words());
113            }
114            end
115        } else if self.size >= 64 && is_x86_feature_detected!("avx2") {
116            let end = self.size / 64 * 64;
117            // SAFETY: AVX2 is available; each output chunk holds one complete 64-bit word.
118            unsafe {
119                simd::write_binary_avx2(&mut bytes[..end], self.words());
120            }
121            end
122        } else {
123            0
124        };
125        #[cfg(not(target_arch = "x86_64"))]
126        let end = 0;
127        for (chunk, &word) in bytes[end..].chunks_mut(64).zip(&self.words()[end / 64..]) {
128            for (i, byte) in chunk.as_chunks_mut::<8>().0.iter_mut().enumerate() {
129                byte.copy_from_slice(&TABLE[(word >> (i * 8) & 255) as usize]);
130            }
131        }
132        bytes.truncate(self.size);
133        // SAFETY: every output byte is ASCII `0` or `1`.
134        unsafe { String::from_utf8_unchecked(bytes) }
135    }
136
137    pub fn ones(size: usize) -> Self {
138        let mut self_ = Self {
139            size,
140            bits: vec![Block([u64::MAX; 8]); size.div_ceil(512)],
141        };
142        self_.trim();
143        self_
144    }
145
146    pub fn get(&self, i: usize) -> bool {
147        self.bits[i >> 9].0[i >> 6 & 7] & (1 << (i & 63)) != 0
148    }
149
150    pub fn set(&mut self, i: usize, b: bool) {
151        let word = &mut self.bits[i >> 9].0[i >> 6 & 7];
152        if b {
153            *word |= 1 << (i & 63);
154        } else {
155            *word &= !(1 << (i & 63));
156        }
157    }
158
159    /// Clears all bits.
160    pub fn reset(&mut self) {
161        self.bits.fill(Block::default());
162    }
163
164    /// Sets all bits to `value`.
165    pub fn fill(&mut self, value: bool) {
166        self.bits.fill(Block([if value { u64::MAX } else { 0 }; 8]));
167        self.trim();
168    }
169
170    /// Tests whether any bit is set.
171    #[inline]
172    pub fn any(&self) -> bool {
173        !self.none()
174    }
175
176    /// Tests whether all bits are unset.
177    #[inline]
178    pub fn none(&self) -> bool {
179        #[cfg(target_arch = "x86_64")]
180        if self.bits.len() >= SIMD_MIN_BLOCKS {
181            if self.bits[0].0[0] != 0 {
182                return false;
183            }
184            if avx512_enabled() && is_x86_feature_detected!("avx512f") {
185                // SAFETY: blocks are 64-byte aligned and feature detection checked AVX-512F.
186                return unsafe { simd::none_avx512(&self.bits) };
187            }
188            if is_x86_feature_detected!("avx2") {
189                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
190                return unsafe { simd::none_avx2(&self.bits) };
191            }
192        }
193        self.words().iter().all(|&word| word == 0)
194    }
195
196    /// Tests whether all bits are set.
197    #[inline(always)]
198    pub fn all(&self) -> bool {
199        let words = self.words();
200        let full_words = self.size >> 6;
201        #[cfg(target_arch = "x86_64")]
202        if self.size >> 9 >= SIMD_MIN_BLOCKS {
203            let full_blocks = self.size >> 9;
204            if words[0] != u64::MAX {
205                return false;
206            }
207            let full_blocks_are_set = if avx512_enabled() && is_x86_feature_detected!("avx512f") {
208                // SAFETY: blocks are 64-byte aligned and feature detection checked AVX-512F.
209                Some(unsafe { simd::all_avx512(&self.bits[..full_blocks]) })
210            } else if is_x86_feature_detected!("avx2") {
211                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
212                Some(unsafe { simd::all_avx2(&self.bits[..full_blocks]) })
213            } else {
214                None
215            };
216            if let Some(full_blocks_are_set) = full_blocks_are_set {
217                return full_blocks_are_set
218                    && words[full_blocks * 8..full_words]
219                        .iter()
220                        .all(|&word| word == u64::MAX)
221                    && (self.size & 63 == 0
222                        || words[full_words] == u64::MAX >> (64 - (self.size & 63)));
223            }
224        }
225        if self.size & 63 == 0 {
226            return words.iter().all(|&word| word == u64::MAX);
227        }
228        words[..full_words].iter().all(|&word| word == u64::MAX)
229            && words[full_words] == u64::MAX >> (64 - (self.size & 63))
230    }
231
232    /// Iterates over set-bit indices in ascending order.
233    pub fn iter_ones(&self) -> impl Iterator<Item = usize> + '_ {
234        self.words()
235            .iter()
236            .copied()
237            .enumerate()
238            .flat_map(|(word_index, mut word)| {
239                std::iter::from_fn(move || {
240                    if word == 0 {
241                        None
242                    } else {
243                        let bit = word.trailing_zeros() as usize;
244                        word &= word - 1;
245                        Some((word_index << 6) | bit)
246                    }
247                })
248            })
249    }
250
251    /// Counts set bits.
252    #[inline]
253    pub fn count_ones(&self) -> u64 {
254        let words = self.words();
255        #[cfg(target_arch = "x86_64")]
256        if words.len() >= 8 {
257            if avx512_enabled()
258                && is_x86_feature_detected!("avx512f")
259                && is_x86_feature_detected!("avx512vpopcntdq")
260            {
261                // SAFETY: blocks are aligned and feature detection checked both requirements.
262                return unsafe { simd::count_ones_avx512(&self.bits) };
263            }
264            if is_x86_feature_detected!("avx2") {
265                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
266                return unsafe { simd::count_ones_avx2(&self.bits) };
267            }
268        }
269        words.iter().map(|word| word.count_ones() as u64).sum()
270    }
271
272    /// Counts unset bits.
273    #[inline]
274    pub fn count_zeros(&self) -> u64 {
275        self.size as u64 - self.count_ones()
276    }
277
278    pub fn push(&mut self, b: bool) {
279        if self.size & 511 == 0 {
280            self.bits.push(Block::default());
281        }
282        if b {
283            self.bits[self.size >> 9].0[self.size >> 6 & 7] |= 1 << (self.size & 63);
284        }
285        self.size += 1;
286    }
287
288    pub fn resize(&mut self, new_size: usize) {
289        match self.size.cmp(&new_size) {
290            Ordering::Less => self.bits.resize(new_size.div_ceil(512), Block::default()),
291            Ordering::Equal => {}
292            Ordering::Greater => self.bits.truncate(new_size.div_ceil(512)),
293        }
294        self.size = new_size;
295        self.trim();
296    }
297
298    /// Assigns `self | (self << rhs)` to `self`.
299    #[inline]
300    pub fn shl_bitor_assign(&mut self, rhs: usize) {
301        self.shift_left::<true>(rhs);
302    }
303
304    /// Assigns `self | (self >> rhs)` to `self`.
305    #[inline]
306    pub fn shr_bitor_assign(&mut self, rhs: usize) {
307        self.shift_right::<true>(rhs);
308    }
309
310    /// Returns words in increasing bit-index order; unused high bits in the last word are zero.
311    pub fn words(&self) -> &[u64] {
312        // SAFETY: `Block` is exactly eight contiguous `u64`s, with no trailing padding.
313        unsafe { std::slice::from_raw_parts(self.bits.as_ptr().cast(), self.size.div_ceil(64)) }
314    }
315
316    /// Returns mutable words in increasing bit-index order.
317    /// Callers must keep unused high bits in the last word zero.
318    pub fn words_mut(&mut self) -> &mut [u64] {
319        let len = self.size.div_ceil(64);
320        // SAFETY: `Block` is exactly eight contiguous `u64`s, with no trailing padding.
321        unsafe { std::slice::from_raw_parts_mut(self.bits.as_mut_ptr().cast(), len) }
322    }
323
324    fn trim(&mut self) {
325        let used_words = self.size.div_ceil(64) & 7;
326        if let Some(last) = self.bits.last_mut() {
327            if self.size & 63 != 0 {
328                last.0[(self.size >> 6) & 7] &= u64::MAX >> (64 - (self.size & 63));
329            }
330            if used_words != 0 {
331                last.0[used_words..].fill(0);
332            }
333        }
334    }
335
336    #[inline]
337    fn bitop_assign<const OP: u8>(&mut self, rhs: &Self) {
338        assert_eq!(
339            self.size, rhs.size,
340            "bitwise operations require equal lengths"
341        );
342        #[cfg(target_arch = "x86_64")]
343        if self.bits.len() >= SIMD_MIN_BLOCKS {
344            if avx512_enabled() && is_x86_feature_detected!("avx512f") {
345                // SAFETY: blocks are 64-byte aligned and feature detection checked AVX-512F.
346                unsafe { simd::bitop_avx512::<OP>(&mut self.bits, &rhs.bits) };
347                return;
348            }
349            if is_x86_feature_detected!("avx2") {
350                // SAFETY: 64-byte alignment also satisfies AVX2 and feature detection checked it.
351                unsafe { simd::bitop_avx2::<OP>(&mut self.bits, &rhs.bits) };
352                return;
353            }
354        }
355        for (lhs, &rhs) in self.words_mut().iter_mut().zip(rhs.words()) {
356            *lhs = match OP {
357                BIT_AND => *lhs & rhs,
358                BIT_OR => *lhs | rhs,
359                BIT_XOR => *lhs ^ rhs,
360                _ => unreachable!(),
361            };
362        }
363    }
364
365    #[inline]
366    fn shift_left<const OR_ASSIGN: bool>(&mut self, rhs: usize) {
367        if rhs == 0 {
368            return;
369        }
370        if rhs >= self.size {
371            if !OR_ASSIGN {
372                self.reset();
373            }
374            return;
375        }
376        #[cfg(target_arch = "x86_64")]
377        if self.bits.len() >= SIMD_MIN_BLOCKS {
378            if avx512_enabled()
379                && is_x86_feature_detected!("avx512f")
380                && is_x86_feature_detected!("avx512vbmi2")
381            {
382                // SAFETY: blocks are aligned and feature detection checked AVX-512F and VBMI2.
383                unsafe { simd::shift_left_avx512::<OR_ASSIGN>(&mut self.bits, rhs) };
384                if self.size & 511 != 0 {
385                    self.trim();
386                }
387                return;
388            }
389            if is_x86_feature_detected!("avx2") {
390                // SAFETY: feature detection checked AVX2 support.
391                unsafe { simd::shift_left_avx2::<OR_ASSIGN>(self.words_mut(), rhs) };
392                if self.size & 63 != 0 {
393                    self.trim();
394                }
395                return;
396            }
397        }
398
399        let bits = self.words_mut();
400        let word_shift = rhs >> 6;
401        let bit_shift = rhs & 63;
402        if bit_shift == 0 {
403            for i in (0..bits.len() - word_shift).rev() {
404                if OR_ASSIGN {
405                    bits[i + word_shift] |= bits[i];
406                } else {
407                    bits[i + word_shift] = bits[i];
408                }
409            }
410        } else {
411            for i in (1..bits.len() - word_shift).rev() {
412                let value = (bits[i] << bit_shift) | (bits[i - 1] >> (64 - bit_shift));
413                if OR_ASSIGN {
414                    bits[i + word_shift] |= value;
415                } else {
416                    bits[i + word_shift] = value;
417                }
418            }
419            if OR_ASSIGN {
420                bits[word_shift] |= bits[0] << bit_shift;
421            } else {
422                bits[word_shift] = bits[0] << bit_shift;
423            }
424        }
425        if !OR_ASSIGN {
426            bits[..word_shift].fill(0);
427        }
428        if self.size & 63 != 0 {
429            self.trim();
430        }
431    }
432
433    #[inline]
434    fn shift_right<const OR_ASSIGN: bool>(&mut self, rhs: usize) {
435        if rhs == 0 {
436            return;
437        }
438        if rhs >= self.size {
439            if !OR_ASSIGN {
440                self.reset();
441            }
442            return;
443        }
444        #[cfg(target_arch = "x86_64")]
445        if self.bits.len() >= SIMD_MIN_BLOCKS {
446            if avx512_enabled()
447                && is_x86_feature_detected!("avx512f")
448                && is_x86_feature_detected!("avx512vbmi2")
449            {
450                // SAFETY: blocks are aligned and feature detection checked AVX-512F and VBMI2.
451                unsafe { simd::shift_right_avx512::<OR_ASSIGN>(&mut self.bits, rhs) };
452                return;
453            }
454            if is_x86_feature_detected!("avx2") {
455                // SAFETY: feature detection checked AVX2 support.
456                unsafe { simd::shift_right_avx2::<OR_ASSIGN>(self.words_mut(), rhs) };
457                return;
458            }
459        }
460
461        let bits = self.words_mut();
462        let word_shift = rhs >> 6;
463        let bit_shift = rhs & 63;
464        if bit_shift == 0 {
465            for i in word_shift..bits.len() {
466                if OR_ASSIGN {
467                    bits[i - word_shift] |= bits[i];
468                } else {
469                    bits[i - word_shift] = bits[i];
470                }
471            }
472        } else {
473            for i in word_shift..bits.len() - 1 {
474                let value = (bits[i] >> bit_shift) | (bits[i + 1] << (64 - bit_shift));
475                if OR_ASSIGN {
476                    bits[i - word_shift] |= value;
477                } else {
478                    bits[i - word_shift] = value;
479                }
480            }
481            if OR_ASSIGN {
482                bits[bits.len() - word_shift - 1] |= bits[bits.len() - 1] >> bit_shift;
483            } else {
484                bits[bits.len() - word_shift - 1] = bits[bits.len() - 1] >> bit_shift;
485            }
486        }
487        if !OR_ASSIGN {
488            let end = bits.len() - word_shift;
489            bits[end..].fill(0);
490        }
491    }
492}
493
494impl Extend<bool> for BitSet {
495    fn extend<T: IntoIterator<Item = bool>>(&mut self, iter: T) {
496        for bit in iter {
497            self.push(bit);
498        }
499    }
500}
501
502impl FromIterator<bool> for BitSet {
503    fn from_iter<T: IntoIterator<Item = bool>>(iter: T) -> Self {
504        let mut set = BitSet::new(0);
505        set.extend(iter);
506        set
507    }
508}
509
510impl ShlAssign<usize> for BitSet {
511    #[inline]
512    fn shl_assign(&mut self, rhs: usize) {
513        self.shift_left::<false>(rhs);
514    }
515}
516
517impl Shl<usize> for BitSet {
518    type Output = Self;
519    fn shl(mut self, rhs: usize) -> Self::Output {
520        self <<= rhs;
521        self
522    }
523}
524
525impl ShrAssign<usize> for BitSet {
526    #[inline]
527    fn shr_assign(&mut self, rhs: usize) {
528        self.shift_right::<false>(rhs);
529    }
530}
531
532impl Shr<usize> for BitSet {
533    type Output = Self;
534    fn shr(mut self, rhs: usize) -> Self::Output {
535        self >>= rhs;
536        self
537    }
538}
539
540impl BitOrAssign<&BitSet> for BitSet {
541    #[inline]
542    fn bitor_assign(&mut self, rhs: &Self) {
543        self.bitop_assign::<BIT_OR>(rhs);
544    }
545}
546
547impl BitOr<&BitSet> for BitSet {
548    type Output = Self;
549    fn bitor(mut self, rhs: &Self) -> Self::Output {
550        self |= rhs;
551        self
552    }
553}
554
555impl BitOr<&BitSet> for &BitSet {
556    type Output = BitSet;
557    fn bitor(self, rhs: &BitSet) -> Self::Output {
558        let mut res = self.clone();
559        res |= rhs;
560        res
561    }
562}
563
564impl BitAndAssign<&BitSet> for BitSet {
565    #[inline]
566    fn bitand_assign(&mut self, rhs: &Self) {
567        self.bitop_assign::<BIT_AND>(rhs);
568    }
569}
570
571impl BitAnd<&BitSet> for BitSet {
572    type Output = Self;
573    fn bitand(mut self, rhs: &Self) -> Self::Output {
574        self &= rhs;
575        self
576    }
577}
578
579impl BitAnd<&BitSet> for &BitSet {
580    type Output = BitSet;
581    fn bitand(self, rhs: &BitSet) -> Self::Output {
582        let mut res = self.clone();
583        res &= rhs;
584        res
585    }
586}
587
588impl BitXorAssign<&BitSet> for BitSet {
589    #[inline]
590    fn bitxor_assign(&mut self, rhs: &Self) {
591        self.bitop_assign::<BIT_XOR>(rhs);
592    }
593}
594
595impl BitXor<&BitSet> for BitSet {
596    type Output = Self;
597    fn bitxor(mut self, rhs: &Self) -> Self::Output {
598        self ^= rhs;
599        self
600    }
601}
602
603impl BitXor<&BitSet> for &BitSet {
604    type Output = BitSet;
605    fn bitxor(self, rhs: &BitSet) -> Self::Output {
606        let mut res = self.clone();
607        res ^= rhs;
608        res
609    }
610}
611
612impl Not for BitSet {
613    type Output = Self;
614    fn not(mut self) -> Self::Output {
615        for word in self.words_mut() {
616            *word = !*word;
617        }
618        self.trim();
619        self
620    }
621}
622
623impl Not for &BitSet {
624    type Output = BitSet;
625    fn not(self) -> Self::Output {
626        !self.clone()
627    }
628}
629
630#[cfg(target_arch = "x86_64")]
631#[allow(unsafe_op_in_unsafe_fn)] // SIMD intrinsics and raw pointers are confined here
632mod simd {
633    use super::{BIT_AND, BIT_OR, BIT_XOR, Block};
634    use std::arch::x86_64::*;
635
636    #[target_feature(enable = "avx512bw")]
637    pub unsafe fn write_binary_avx512(bytes: &mut [u8], words: &[u64]) {
638        let one = _mm512_set1_epi8(1);
639        let zero = _mm512_set1_epi8(b'0' as i8);
640        for (chunk, &word) in bytes.as_chunks_mut::<64>().0.iter_mut().zip(words) {
641            let value = _mm512_mask_add_epi8(zero, word, zero, one);
642            _mm512_storeu_si512(chunk.as_mut_ptr().cast(), value);
643        }
644    }
645    #[target_feature(enable = "avx2")]
646    pub unsafe fn write_binary_avx2(bytes: &mut [u8], words: &[u64]) {
647        let indices = _mm256_setr_epi64x(
648            0,
649            0x0101_0101_0101_0101,
650            0x0202_0202_0202_0202,
651            0x0303_0303_0303_0303,
652        );
653        let mask = _mm256_set1_epi64x(0x8040_2010_0804_0201u64 as i64);
654        let one = _mm256_set1_epi8(b'1' as i8);
655        let zero = _mm256_setzero_si256();
656        for (chunk, &word) in bytes.as_chunks_mut::<64>().0.iter_mut().zip(words) {
657            let lo = _mm256_shuffle_epi8(_mm256_set1_epi32(word as i32), indices);
658            let hi = _mm256_shuffle_epi8(_mm256_set1_epi32((word >> 32) as i32), indices);
659            let lo = _mm256_add_epi8(one, _mm256_cmpeq_epi8(_mm256_and_si256(lo, mask), zero));
660            let hi = _mm256_add_epi8(one, _mm256_cmpeq_epi8(_mm256_and_si256(hi, mask), zero));
661            _mm256_storeu_si256(chunk.as_mut_ptr().cast(), lo);
662            _mm256_storeu_si256(chunk.as_mut_ptr().add(32).cast(), hi);
663        }
664    }
665    #[target_feature(enable = "avx2")]
666    pub unsafe fn parse_binary_avx2(bytes: &[u8], words: &mut [u64]) -> bool {
667        let one = _mm256_set1_epi8(b'1' as i8);
668        let mask = _mm256_set1_epi8(!1);
669        let zero = _mm256_set1_epi8(b'0' as i8);
670        for (chunk, word) in bytes.as_chunks::<64>().0.iter().zip(words) {
671            let a = _mm256_loadu_si256(chunk.as_ptr().cast());
672            let b = _mm256_loadu_si256(chunk.as_ptr().add(32).cast());
673            let valid_a = _mm256_cmpeq_epi8(_mm256_and_si256(a, mask), zero);
674            let valid_b = _mm256_cmpeq_epi8(_mm256_and_si256(b, mask), zero);
675            if _mm256_movemask_epi8(_mm256_and_si256(valid_a, valid_b)) != -1 {
676                return false;
677            }
678            *word = _mm256_movemask_epi8(_mm256_cmpeq_epi8(a, one)) as u32 as u64
679                | ((_mm256_movemask_epi8(_mm256_cmpeq_epi8(b, one)) as u32 as u64) << 32);
680        }
681        true
682    }
683
684    #[target_feature(enable = "avx512bw")]
685    pub unsafe fn parse_binary_avx512(bytes: &[u8], words: &mut [u64]) -> bool {
686        let one = _mm512_set1_epi8(b'1' as i8);
687        let mask = _mm512_set1_epi8(!1);
688        let zero = _mm512_set1_epi8(b'0' as i8);
689        for (chunk, word) in bytes.as_chunks::<64>().0.iter().zip(words) {
690            let x = _mm512_loadu_si512(chunk.as_ptr().cast());
691            if _mm512_cmpeq_epi8_mask(_mm512_and_si512(x, mask), zero) != u64::MAX {
692                return false;
693            }
694            *word = _mm512_cmpeq_epi8_mask(x, one);
695        }
696        true
697    }
698
699    #[target_feature(enable = "avx2")]
700    pub unsafe fn bitop_avx2<const OP: u8>(lhs: &mut [Block], rhs: &[Block]) {
701        let lhs_ptr = lhs.as_mut_ptr().cast::<__m256i>();
702        let rhs_ptr = rhs.as_ptr().cast::<__m256i>();
703        for i in 0..lhs.len() * 2 {
704            let lhs_value = _mm256_load_si256(lhs_ptr.add(i));
705            let rhs_value = _mm256_load_si256(rhs_ptr.add(i));
706            let value = match OP {
707                BIT_AND => _mm256_and_si256(lhs_value, rhs_value),
708                BIT_OR => _mm256_or_si256(lhs_value, rhs_value),
709                BIT_XOR => _mm256_xor_si256(lhs_value, rhs_value),
710                _ => unreachable!(),
711            };
712            _mm256_store_si256(lhs_ptr.add(i), value);
713        }
714    }
715
716    #[target_feature(enable = "avx512f")]
717    pub unsafe fn bitop_avx512<const OP: u8>(lhs: &mut [Block], rhs: &[Block]) {
718        for i in 0..lhs.len() {
719            let lhs_value = _mm512_load_si512(lhs.as_ptr().add(i).cast());
720            let rhs_value = _mm512_load_si512(rhs.as_ptr().add(i).cast());
721            let value = match OP {
722                BIT_AND => _mm512_and_si512(lhs_value, rhs_value),
723                BIT_OR => _mm512_or_si512(lhs_value, rhs_value),
724                BIT_XOR => _mm512_xor_si512(lhs_value, rhs_value),
725                _ => unreachable!(),
726            };
727            _mm512_store_si512(lhs.as_mut_ptr().add(i).cast(), value);
728        }
729    }
730
731    #[target_feature(enable = "avx2")]
732    pub unsafe fn count_ones_avx2(bits: &[Block]) -> u64 {
733        let table = _mm256_set_epi64x(
734            0x0403_0302_0302_0201,
735            0x0302_0201_0201_0100,
736            0x0403_0302_0302_0201,
737            0x0302_0201_0201_0100,
738        );
739        let low_mask = _mm256_set1_epi8(0x0f);
740        let zero = _mm256_setzero_si256();
741        let mut sum = zero;
742        let ptr = bits.as_ptr().cast::<__m256i>();
743        let mut i = 0;
744        while i + 16 <= bits.len() * 2 {
745            // Sixteen vectors contribute at most 128 set bits to each byte.
746            let mut counts = zero;
747            for offset in 0..16 {
748                let value = _mm256_load_si256(ptr.add(i + offset));
749                let low = _mm256_shuffle_epi8(table, _mm256_and_si256(value, low_mask));
750                let high = _mm256_shuffle_epi8(
751                    table,
752                    _mm256_and_si256(_mm256_srli_epi16::<4>(value), low_mask),
753                );
754                counts = _mm256_add_epi8(counts, _mm256_add_epi8(low, high));
755            }
756            sum = _mm256_add_epi64(sum, _mm256_sad_epu8(counts, zero));
757            i += 16;
758        }
759        while i < bits.len() * 2 {
760            let value = _mm256_load_si256(ptr.add(i));
761            let low = _mm256_shuffle_epi8(table, _mm256_and_si256(value, low_mask));
762            let high = _mm256_shuffle_epi8(
763                table,
764                _mm256_and_si256(_mm256_srli_epi16::<4>(value), low_mask),
765            );
766            sum = _mm256_add_epi64(sum, _mm256_sad_epu8(_mm256_add_epi8(low, high), zero));
767            i += 1;
768        }
769        let mut lanes = [0; 4];
770        _mm256_storeu_si256(lanes.as_mut_ptr().cast(), sum);
771        lanes.into_iter().sum()
772    }
773
774    #[target_feature(enable = "avx512f,avx512vpopcntdq")]
775    pub unsafe fn count_ones_avx512(bits: &[Block]) -> u64 {
776        let mut sum = _mm512_setzero_si512();
777        for i in 0..bits.len() {
778            let value = _mm512_load_si512(bits.as_ptr().add(i).cast());
779            sum = _mm512_add_epi64(sum, _mm512_popcnt_epi64(value));
780        }
781        let mut lanes = [0; 8];
782        _mm512_storeu_si512(lanes.as_mut_ptr().cast(), sum);
783        lanes.into_iter().sum()
784    }
785
786    #[target_feature(enable = "avx2")]
787    pub unsafe fn none_avx2(bits: &[Block]) -> bool {
788        let ptr = bits.as_ptr().cast::<__m256i>();
789        for i in 0..bits.len() * 2 {
790            let value = _mm256_load_si256(ptr.add(i));
791            if _mm256_testz_si256(value, value) == 0 {
792                return false;
793            }
794        }
795        true
796    }
797
798    #[target_feature(enable = "avx512f")]
799    pub unsafe fn none_avx512(bits: &[Block]) -> bool {
800        let mut i = 0;
801        while i + 4 <= bits.len() {
802            let value = _mm512_or_si512(
803                _mm512_or_si512(
804                    _mm512_load_si512(bits.as_ptr().add(i).cast()),
805                    _mm512_load_si512(bits.as_ptr().add(i + 1).cast()),
806                ),
807                _mm512_or_si512(
808                    _mm512_load_si512(bits.as_ptr().add(i + 2).cast()),
809                    _mm512_load_si512(bits.as_ptr().add(i + 3).cast()),
810                ),
811            );
812            if _mm512_test_epi64_mask(value, value) != 0 {
813                return false;
814            }
815            i += 4;
816        }
817        while i < bits.len() {
818            let value = _mm512_load_si512(bits.as_ptr().add(i).cast());
819            if _mm512_test_epi64_mask(value, value) != 0 {
820                return false;
821            }
822            i += 1;
823        }
824        true
825    }
826
827    #[target_feature(enable = "avx2")]
828    pub unsafe fn all_avx2(bits: &[Block]) -> bool {
829        let ones = _mm256_set1_epi64x(-1);
830        let ptr = bits.as_ptr().cast::<__m256i>();
831        for i in 0..bits.len() {
832            let value = _mm256_and_si256(
833                _mm256_load_si256(ptr.add(i * 2)),
834                _mm256_load_si256(ptr.add(i * 2 + 1)),
835            );
836            if _mm256_movemask_epi8(_mm256_cmpeq_epi64(value, ones)) != -1 {
837                return false;
838            }
839        }
840        true
841    }
842
843    #[target_feature(enable = "avx512f")]
844    pub unsafe fn all_avx512(bits: &[Block]) -> bool {
845        let ones = _mm512_set1_epi64(-1);
846        let mut i = 0;
847        while i + 4 <= bits.len() {
848            let value = _mm512_and_si512(
849                _mm512_and_si512(
850                    _mm512_load_si512(bits.as_ptr().add(i).cast()),
851                    _mm512_load_si512(bits.as_ptr().add(i + 1).cast()),
852                ),
853                _mm512_and_si512(
854                    _mm512_load_si512(bits.as_ptr().add(i + 2).cast()),
855                    _mm512_load_si512(bits.as_ptr().add(i + 3).cast()),
856                ),
857            );
858            if _mm512_cmpeq_epi64_mask(value, ones) != u8::MAX {
859                return false;
860            }
861            i += 4;
862        }
863        while i < bits.len() {
864            if _mm512_cmpeq_epi64_mask(_mm512_load_si512(bits.as_ptr().add(i).cast()), ones)
865                != u8::MAX
866            {
867                return false;
868            }
869            i += 1;
870        }
871        true
872    }
873
874    #[target_feature(enable = "avx2")]
875    pub unsafe fn shift_left_avx2<const OR_ASSIGN: bool>(bits: &mut [u64], rhs: usize) {
876        let word_shift = rhs >> 6;
877        let bit_shift = rhs & 63;
878        let lower = word_shift + usize::from(bit_shift != 0);
879        let mut end = bits.len();
880        let count = _mm_cvtsi64_si128(bit_shift as i64);
881        while end >= lower + 4 {
882            let start = end - 4;
883            let value = _mm256_loadu_si256(bits.as_ptr().add(start - word_shift).cast());
884            let mut value = if bit_shift == 0 {
885                value
886            } else {
887                _mm256_or_si256(
888                    _mm256_sll_epi64(value, count),
889                    _mm256_srl_epi64(
890                        _mm256_loadu_si256(bits.as_ptr().add(start - word_shift - 1).cast()),
891                        _mm_cvtsi64_si128((64 - bit_shift) as i64),
892                    ),
893                )
894            };
895            if OR_ASSIGN {
896                value = _mm256_or_si256(value, _mm256_loadu_si256(bits.as_ptr().add(start).cast()));
897            }
898            _mm256_storeu_si256(bits.as_mut_ptr().add(start).cast(), value);
899            end = start;
900        }
901        for i in (lower..end).rev() {
902            let source = i - word_shift;
903            let value = if bit_shift == 0 {
904                bits[source]
905            } else {
906                (bits[source] << bit_shift) | (bits[source - 1] >> (64 - bit_shift))
907            };
908            if OR_ASSIGN {
909                bits[i] |= value;
910            } else {
911                bits[i] = value;
912            }
913        }
914        if bit_shift != 0 {
915            if OR_ASSIGN {
916                bits[word_shift] |= bits[0] << bit_shift;
917            } else {
918                bits[word_shift] = bits[0] << bit_shift;
919            }
920        }
921        if !OR_ASSIGN {
922            bits[..word_shift].fill(0);
923        }
924    }
925
926    #[target_feature(enable = "avx2")]
927    pub unsafe fn shift_right_avx2<const OR_ASSIGN: bool>(bits: &mut [u64], rhs: usize) {
928        let word_shift = rhs >> 6;
929        let bit_shift = rhs & 63;
930        let upper = bits.len() - word_shift - usize::from(bit_shift != 0);
931        let mut start = 0;
932        let count = _mm_cvtsi64_si128(bit_shift as i64);
933        while start + 4 <= upper {
934            let value = _mm256_loadu_si256(bits.as_ptr().add(start + word_shift).cast());
935            let mut value = if bit_shift == 0 {
936                value
937            } else {
938                _mm256_or_si256(
939                    _mm256_srl_epi64(value, count),
940                    _mm256_sll_epi64(
941                        _mm256_loadu_si256(bits.as_ptr().add(start + word_shift + 1).cast()),
942                        _mm_cvtsi64_si128((64 - bit_shift) as i64),
943                    ),
944                )
945            };
946            if OR_ASSIGN {
947                value = _mm256_or_si256(value, _mm256_loadu_si256(bits.as_ptr().add(start).cast()));
948            }
949            _mm256_storeu_si256(bits.as_mut_ptr().add(start).cast(), value);
950            start += 4;
951        }
952        for i in start..upper {
953            let source = i + word_shift;
954            let value = if bit_shift == 0 {
955                bits[source]
956            } else {
957                (bits[source] >> bit_shift) | (bits[source + 1] << (64 - bit_shift))
958            };
959            if OR_ASSIGN {
960                bits[i] |= value;
961            } else {
962                bits[i] = value;
963            }
964        }
965        if bit_shift != 0 {
966            if OR_ASSIGN {
967                bits[bits.len() - word_shift - 1] |= bits[bits.len() - 1] >> bit_shift;
968            } else {
969                bits[bits.len() - word_shift - 1] = bits[bits.len() - 1] >> bit_shift;
970            }
971        }
972        if !OR_ASSIGN {
973            let end = bits.len() - word_shift;
974            bits[end..].fill(0);
975        }
976    }
977
978    #[target_feature(enable = "avx512f,avx512vbmi2")]
979    pub unsafe fn shift_left_avx512<const OR_ASSIGN: bool>(bits: &mut [Block], rhs: usize) {
980        let block_shift = rhs >> 9;
981        let word_shift = rhs >> 6 & 7;
982        let bit_shift = rhs & 63;
983        if word_shift | bit_shift == 0 {
984            for i in (block_shift..bits.len()).rev() {
985                let mut value = _mm512_load_si512(bits.as_ptr().add(i - block_shift).cast());
986                if OR_ASSIGN {
987                    value = _mm512_or_si512(value, _mm512_load_si512(bits.as_ptr().add(i).cast()));
988                }
989                _mm512_store_si512(bits.as_mut_ptr().add(i).cast(), value);
990            }
991        } else {
992            let zero = _mm512_setzero_si512();
993            let indices = _mm512_setr_epi64(0, 1, 2, 3, 4, 5, 6, 7);
994            let current_indices =
995                _mm512_add_epi64(indices, _mm512_set1_epi64((8 - word_shift) as i64));
996            let previous_indices = _mm512_sub_epi64(current_indices, _mm512_set1_epi64(1));
997            let count = _mm512_set1_epi64(bit_shift as i64);
998            for i in (block_shift..bits.len()).rev() {
999                let source = i - block_shift;
1000                let previous = if source == 0 {
1001                    zero
1002                } else {
1003                    _mm512_load_si512(bits.as_ptr().add(source - 1).cast())
1004                };
1005                let current = _mm512_load_si512(bits.as_ptr().add(source).cast());
1006                let value = if word_shift == 0 {
1007                    current
1008                } else {
1009                    _mm512_permutex2var_epi64(previous, current_indices, current)
1010                };
1011                let mut value = if bit_shift == 0 {
1012                    value
1013                } else {
1014                    _mm512_shldv_epi64(
1015                        value,
1016                        _mm512_permutex2var_epi64(previous, previous_indices, current),
1017                        count,
1018                    )
1019                };
1020                if OR_ASSIGN {
1021                    value = _mm512_or_si512(value, _mm512_load_si512(bits.as_ptr().add(i).cast()));
1022                }
1023                _mm512_store_si512(bits.as_mut_ptr().add(i).cast(), value);
1024            }
1025        }
1026        if !OR_ASSIGN {
1027            bits[..block_shift].fill(Block::default());
1028        }
1029    }
1030
1031    #[target_feature(enable = "avx512f,avx512vbmi2")]
1032    pub unsafe fn shift_right_avx512<const OR_ASSIGN: bool>(bits: &mut [Block], rhs: usize) {
1033        let block_shift = rhs >> 9;
1034        let word_shift = rhs >> 6 & 7;
1035        let bit_shift = rhs & 63;
1036        let remaining = bits.len() - block_shift;
1037        if word_shift | bit_shift == 0 {
1038            for i in 0..remaining {
1039                let mut value = _mm512_load_si512(bits.as_ptr().add(i + block_shift).cast());
1040                if OR_ASSIGN {
1041                    value = _mm512_or_si512(value, _mm512_load_si512(bits.as_ptr().add(i).cast()));
1042                }
1043                _mm512_store_si512(bits.as_mut_ptr().add(i).cast(), value);
1044            }
1045        } else {
1046            let zero = _mm512_setzero_si512();
1047            let indices = _mm512_setr_epi64(0, 1, 2, 3, 4, 5, 6, 7);
1048            let current_indices = _mm512_add_epi64(indices, _mm512_set1_epi64(word_shift as i64));
1049            let next_indices = _mm512_add_epi64(current_indices, _mm512_set1_epi64(1));
1050            let count = _mm512_set1_epi64(bit_shift as i64);
1051            for i in 0..remaining {
1052                let source = i + block_shift;
1053                let current = _mm512_load_si512(bits.as_ptr().add(source).cast());
1054                let next = if source + 1 == bits.len() {
1055                    zero
1056                } else {
1057                    _mm512_load_si512(bits.as_ptr().add(source + 1).cast())
1058                };
1059                let value = if word_shift == 0 {
1060                    current
1061                } else {
1062                    _mm512_permutex2var_epi64(current, current_indices, next)
1063                };
1064                let mut value = if bit_shift == 0 {
1065                    value
1066                } else {
1067                    _mm512_shrdv_epi64(
1068                        value,
1069                        _mm512_permutex2var_epi64(current, next_indices, next),
1070                        count,
1071                    )
1072                };
1073                if OR_ASSIGN {
1074                    value = _mm512_or_si512(value, _mm512_load_si512(bits.as_ptr().add(i).cast()));
1075                }
1076                _mm512_store_si512(bits.as_mut_ptr().add(i).cast(), value);
1077            }
1078        }
1079        if !OR_ASSIGN {
1080            bits[remaining..].fill(Block::default());
1081        }
1082    }
1083}
1084
1085#[cfg(test)]
1086mod tests {
1087    use super::*;
1088    use crate::tools::Xorshift;
1089    use crate::tools::testutil::sample_usize;
1090    use std::{
1091        mem::size_of,
1092        panic::{AssertUnwindSafe, catch_unwind},
1093    };
1094
1095    #[test]
1096    fn test_binary_conversion() {
1097        let mut rng = Xorshift::new_with_seed(41513);
1098        for case in 0..512 {
1099            let size = match case % 3 {
1100                0 => rng.random(0..32),
1101                1 => rng.random(0..512),
1102                _ => rng.random(0..20000),
1103            };
1104            let density = rng.rand(1010);
1105            let model: Vec<bool> = (0..size).map(|_| rng.rand(1009) < density).collect();
1106            let s: String = model.iter().map(|&b| if b { '1' } else { '0' }).collect();
1107            let expected: BitSet = model.iter().copied().collect();
1108            let parsed = BitSet::from_binary(&s).unwrap();
1109            assert_model(&parsed, &model);
1110            assert_eq!(parsed, expected);
1111            assert_eq!(expected.to_binary(), s);
1112            let words: Vec<u64> = model
1113                .chunks(64)
1114                .map(|chunk| {
1115                    chunk
1116                        .iter()
1117                        .enumerate()
1118                        .fold(0, |w, (i, &b)| w | (u64::from(b) << i))
1119                })
1120                .collect();
1121            assert_eq!(expected.words(), words);
1122            let mut packed = BitSet::new(size);
1123            packed.words_mut().copy_from_slice(&words);
1124            assert_model(&packed, &model);
1125
1126            let mut modified = s.as_bytes().to_vec();
1127            if size != 0 {
1128                for _ in 0..rng.random(1..=size.min(16)) {
1129                    modified[rng.random(0..size)] = rng.random(0..128);
1130                }
1131            }
1132            let valid = modified.iter().all(|b| matches!(b, b'0' | b'1'));
1133            let modified = String::from_utf8(modified).unwrap();
1134            assert_eq!(BitSet::from_binary(&modified).is_some(), valid);
1135            let mut unicode = s.clone();
1136            unicode.insert(
1137                rng.random(0..=size),
1138                char::from_u32(rng.random(0xe000..0x110000)).unwrap(),
1139            );
1140            assert!(BitSet::from_binary(&unicode).is_none());
1141
1142            let end = size / 64 * 64;
1143            for invalid in [false, true] {
1144                let mut input = s.as_bytes()[..end].to_vec();
1145                if invalid && end != 0 {
1146                    for _ in 0..rng.random(1..=end.min(16)) {
1147                        input[rng.random(0..end)] = rng.random(0..=255);
1148                    }
1149                }
1150                let valid = input.iter().all(|b| matches!(b, b'0' | b'1'));
1151                let mut scalar = BitSet::new(end);
1152                assert_eq!(
1153                    BitSet::parse_binary_scalar(&input, scalar.words_mut()),
1154                    valid
1155                );
1156                if valid {
1157                    assert_model(
1158                        &scalar,
1159                        &input.iter().map(|&b| b == b'1').collect::<Vec<_>>(),
1160                    );
1161                }
1162                #[cfg(target_arch = "x86_64")]
1163                {
1164                    if is_x86_feature_detected!("avx2") {
1165                        let mut parsed = BitSet::new(end);
1166                        // SAFETY: AVX2 support was checked; input has full 64-byte chunks.
1167                        assert_eq!(
1168                            unsafe { simd::parse_binary_avx2(&input, parsed.words_mut()) },
1169                            valid
1170                        );
1171                        if valid {
1172                            assert_eq!(parsed, scalar);
1173                            let mut output = vec![0; end];
1174                            // SAFETY: AVX2 support was checked; output has full 64-byte chunks.
1175                            unsafe { simd::write_binary_avx2(&mut output, parsed.words()) };
1176                            assert_eq!(output, input);
1177                        }
1178                    }
1179                    if is_x86_feature_detected!("avx512bw") {
1180                        let mut parsed = BitSet::new(end);
1181                        // SAFETY: AVX-512BW support was checked; input has full 64-byte chunks.
1182                        assert_eq!(
1183                            unsafe { simd::parse_binary_avx512(&input, parsed.words_mut()) },
1184                            valid
1185                        );
1186                        if valid {
1187                            assert_eq!(parsed, scalar);
1188                            let mut output = vec![0; end];
1189                            // SAFETY: AVX-512BW support was checked; output has full chunks.
1190                            unsafe { simd::write_binary_avx512(&mut output, parsed.words()) };
1191                            assert_eq!(output, input);
1192                        }
1193                    }
1194                }
1195            }
1196        }
1197    }
1198
1199    fn bitset(model: &[bool]) -> BitSet {
1200        model.iter().copied().collect()
1201    }
1202
1203    fn assert_model(actual: &BitSet, expected: &[bool]) {
1204        assert_eq!(actual.len(), expected.len());
1205        assert_eq!(
1206            actual.iter_ones().collect::<Vec<_>>(),
1207            expected
1208                .iter()
1209                .enumerate()
1210                .filter_map(|(i, &bit)| bit.then_some(i))
1211                .collect::<Vec<_>>()
1212        );
1213        assert_eq!(
1214            actual.count_ones(),
1215            expected.iter().filter(|&&bit| bit).count() as u64
1216        );
1217        assert_eq!(
1218            actual.count_zeros(),
1219            expected.iter().filter(|&&bit| !bit).count() as u64
1220        );
1221        assert_eq!(actual.any(), expected.iter().any(|&bit| bit));
1222        assert_eq!(actual.none(), expected.iter().all(|&bit| !bit));
1223        assert_eq!(actual.all(), expected.iter().all(|&bit| bit));
1224        assert!(
1225            actual
1226                .bits
1227                .iter()
1228                .flat_map(|block| block.0)
1229                .skip(expected.len().div_ceil(64))
1230                .all(|word| word == 0)
1231        );
1232        if expected.len() & 63 != 0 {
1233            assert_eq!(actual.words().last().unwrap() >> (expected.len() & 63), 0);
1234        }
1235        for (i, &bit) in expected.iter().enumerate() {
1236            assert_eq!(actual.get(i), bit, "bit {i}");
1237        }
1238    }
1239
1240    fn random_model(rng: &mut Xorshift, size: usize) -> Vec<bool> {
1241        (0..size)
1242            .map(|_| rng.random::<u64, _>(..) & 1 != 0)
1243            .collect()
1244    }
1245
1246    fn bitset_sizes(rng: &mut Xorshift) -> Vec<usize> {
1247        let block_bits = size_of::<Block>() * u8::BITS as usize;
1248        let max_size = 16 * block_bits + 1;
1249        let mut sizes = sample_usize(rng, u64::BITS as usize, 0..=max_size, 30);
1250        // Cover every block boundary, including both sides of SIMD dispatch.
1251        sizes.extend(
1252            (0..max_size)
1253                .step_by(block_bits)
1254                .flat_map(|boundary| boundary.saturating_sub(1)..=boundary + 1),
1255        );
1256        sizes.sort_unstable();
1257        sizes.dedup();
1258        sizes
1259    }
1260
1261    #[test]
1262    fn access_fill_reset_push_extend_resize() {
1263        assert_eq!(size_of::<Block>(), 64);
1264        assert_eq!(std::mem::align_of::<Block>(), 64);
1265        let aligned = BitSet::new(1);
1266        assert_eq!(aligned.bits.as_ptr() as usize & 63, 0);
1267
1268        let mut rng = Xorshift::default();
1269        for size in bitset_sizes(&mut rng) {
1270            let model = random_model(&mut rng, size);
1271            let mut actual = bitset(&model);
1272            assert_model(&actual, &model);
1273
1274            actual.fill(true);
1275            assert_model(&actual, &vec![true; size]);
1276            actual.fill(false);
1277            assert_model(&actual, &vec![false; size]);
1278            actual.fill(true);
1279            actual.reset();
1280            assert_model(&actual, &vec![false; size]);
1281
1282            for (i, &value) in model.iter().enumerate() {
1283                actual.set(i, value);
1284            }
1285            assert_model(&actual, &model);
1286
1287            let extra_len = rng.random(64..=128);
1288            let extra = random_model(&mut rng, extra_len);
1289            actual.extend(extra.iter().copied());
1290            let mut extended = model.clone();
1291            extended.extend(extra);
1292            assert_model(&actual, &extended);
1293
1294            actual.resize(size / 2);
1295            assert_model(&actual, &model[..size / 2]);
1296            actual.resize(size + extra_len);
1297            let mut resized = model[..size / 2].to_vec();
1298            resized.resize(size + extra_len, false);
1299            assert_model(&actual, &resized);
1300
1301            let value = rng.random::<u64, _>(..) & 1 != 0;
1302            actual.push(value);
1303            resized.push(value);
1304            assert_model(&actual, &resized);
1305        }
1306    }
1307
1308    #[test]
1309    fn bitwise_operations_match_boolean_model() {
1310        let mut rng = Xorshift::default();
1311        for size in bitset_sizes(&mut rng) {
1312            let lhs = random_model(&mut rng, size);
1313            let rhs = random_model(&mut rng, size);
1314            let lhs_set = bitset(&lhs);
1315            let rhs_set = bitset(&rhs);
1316
1317            assert_model(
1318                &(&lhs_set & &rhs_set),
1319                &lhs.iter()
1320                    .zip(&rhs)
1321                    .map(|(&x, &y)| x & y)
1322                    .collect::<Vec<_>>(),
1323            );
1324            assert_model(
1325                &(&lhs_set | &rhs_set),
1326                &lhs.iter()
1327                    .zip(&rhs)
1328                    .map(|(&x, &y)| x | y)
1329                    .collect::<Vec<_>>(),
1330            );
1331            assert_model(
1332                &(&lhs_set ^ &rhs_set),
1333                &lhs.iter()
1334                    .zip(&rhs)
1335                    .map(|(&x, &y)| x ^ y)
1336                    .collect::<Vec<_>>(),
1337            );
1338            assert_model(&!&lhs_set, &lhs.iter().map(|&x| !x).collect::<Vec<_>>());
1339        }
1340    }
1341
1342    #[test]
1343    fn shifts_match_boolean_model() {
1344        let mut rng = Xorshift::default();
1345        for size in bitset_sizes(&mut rng) {
1346            let model = random_model(&mut rng, size);
1347            for shift in sample_usize(&mut rng, 16, 0..=size + 512, 10)
1348                .into_iter()
1349                .chain([size, size + 1])
1350            {
1351                let mut expected_left = vec![false; size];
1352                let mut expected_right = vec![false; size];
1353                for (i, &value) in model.iter().enumerate() {
1354                    if let Some(i) = i.checked_add(shift)
1355                        && i < size
1356                    {
1357                        expected_left[i] = value;
1358                    }
1359                    if i >= shift {
1360                        expected_right[i - shift] = value;
1361                    }
1362                }
1363
1364                let mut actual = bitset(&model);
1365                actual <<= shift;
1366                assert_model(&actual, &expected_left);
1367                let mut actual = bitset(&model);
1368                actual >>= shift;
1369                assert_model(&actual, &expected_right);
1370
1371                let mut expected = model.clone();
1372                for (value, shifted) in expected.iter_mut().zip(&expected_left) {
1373                    *value |= *shifted;
1374                }
1375                let mut actual = bitset(&model);
1376                actual.shl_bitor_assign(shift);
1377                assert_model(&actual, &expected);
1378
1379                let mut expected = model.clone();
1380                for (value, shifted) in expected.iter_mut().zip(&expected_right) {
1381                    *value |= *shifted;
1382                }
1383                let mut actual = bitset(&model);
1384                actual.shr_bitor_assign(shift);
1385                assert_model(&actual, &expected);
1386            }
1387        }
1388    }
1389
1390    #[test]
1391    fn bitwise_operations_reject_different_lengths_without_mutation() {
1392        let mut rng = Xorshift::default();
1393        for _ in 0..100 {
1394            let size = rng.random(0..=1024);
1395            let other_size = size + rng.random(1..=1024usize);
1396            let mut lhs = bitset(&random_model(&mut rng, size));
1397            let rhs = bitset(&random_model(&mut rng, other_size));
1398            let before = lhs.clone();
1399            let op = rng.random(0..3);
1400            let result = catch_unwind(AssertUnwindSafe(|| match op {
1401                0 => lhs &= &rhs,
1402                1 => lhs |= &rhs,
1403                _ => lhs ^= &rhs,
1404            }));
1405            assert!(result.is_err());
1406            assert_eq!(lhs, before);
1407        }
1408    }
1409}