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 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 unsafe { simd::parse_binary_avx512(&bytes[..end], bits.words_mut()) }
53 } else if is_x86_feature_detected!("avx2") {
54 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 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 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 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 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 pub fn reset(&mut self) {
161 self.bits.fill(Block::default());
162 }
163
164 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 #[inline]
172 pub fn any(&self) -> bool {
173 !self.none()
174 }
175
176 #[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 return unsafe { simd::none_avx512(&self.bits) };
187 }
188 if is_x86_feature_detected!("avx2") {
189 return unsafe { simd::none_avx2(&self.bits) };
191 }
192 }
193 self.words().iter().all(|&word| word == 0)
194 }
195
196 #[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 Some(unsafe { simd::all_avx512(&self.bits[..full_blocks]) })
210 } else if is_x86_feature_detected!("avx2") {
211 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 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 #[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 return unsafe { simd::count_ones_avx512(&self.bits) };
263 }
264 if is_x86_feature_detected!("avx2") {
265 return unsafe { simd::count_ones_avx2(&self.bits) };
267 }
268 }
269 words.iter().map(|word| word.count_ones() as u64).sum()
270 }
271
272 #[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 #[inline]
300 pub fn shl_bitor_assign(&mut self, rhs: usize) {
301 self.shift_left::<true>(rhs);
302 }
303
304 #[inline]
306 pub fn shr_bitor_assign(&mut self, rhs: usize) {
307 self.shift_right::<true>(rhs);
308 }
309
310 pub fn words(&self) -> &[u64] {
312 unsafe { std::slice::from_raw_parts(self.bits.as_ptr().cast(), self.size.div_ceil(64)) }
314 }
315
316 pub fn words_mut(&mut self) -> &mut [u64] {
319 let len = self.size.div_ceil(64);
320 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 unsafe { simd::bitop_avx512::<OP>(&mut self.bits, &rhs.bits) };
347 return;
348 }
349 if is_x86_feature_detected!("avx2") {
350 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 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 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 unsafe { simd::shift_right_avx512::<OR_ASSIGN>(&mut self.bits, rhs) };
452 return;
453 }
454 if is_x86_feature_detected!("avx2") {
455 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)] mod 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 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 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 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 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 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 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}