1#![deny(missing_docs)]
6
7mod bitops;
8mod eq;
9mod intersect_by_rank;
10
11#[cfg(test)]
12mod tests;
13
14use std::cmp::Ordering;
15use std::fmt::Debug;
16use std::fmt::Formatter;
17use std::ops::Bound;
18use std::ops::RangeBounds;
19use std::sync::Arc;
20use std::sync::OnceLock;
21
22use itertools::Itertools;
23use vortex_buffer::BitBuffer;
24use vortex_buffer::BitBufferMut;
25use vortex_buffer::BitIterator;
26use vortex_error::VortexResult;
27use vortex_error::vortex_panic;
28
29pub enum AllOr<T> {
31 All,
33 None,
35 Some(T),
37}
38
39impl<T> AllOr<T> {
40 #[inline]
42 pub fn unwrap_or_else<F, G>(self, all_true: F, all_false: G) -> T
43 where
44 F: FnOnce() -> T,
45 G: FnOnce() -> T,
46 {
47 match self {
48 Self::Some(v) => v,
49 AllOr::All => all_true(),
50 AllOr::None => all_false(),
51 }
52 }
53}
54
55impl<T> AllOr<&T> {
56 #[inline]
58 pub fn cloned(self) -> AllOr<T>
59 where
60 T: Clone,
61 {
62 match self {
63 Self::All => AllOr::All,
64 Self::None => AllOr::None,
65 Self::Some(v) => AllOr::Some(v.clone()),
66 }
67 }
68}
69
70impl<T> Debug for AllOr<T>
71where
72 T: Debug,
73{
74 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
75 match self {
76 Self::All => f.write_str("All"),
77 Self::None => f.write_str("None"),
78 Self::Some(v) => f.debug_tuple("Some").field(v).finish(),
79 }
80 }
81}
82
83impl<T> PartialEq for AllOr<T>
84where
85 T: PartialEq,
86{
87 fn eq(&self, other: &Self) -> bool {
88 match (self, other) {
89 (Self::All, Self::All) => true,
90 (Self::None, Self::None) => true,
91 (Self::Some(lhs), Self::Some(rhs)) => lhs == rhs,
92 _ => false,
93 }
94 }
95}
96
97impl<T> Eq for AllOr<T> where T: Eq {}
98
99#[derive(Clone)]
105#[cfg_attr(feature = "serde", derive(::serde::Serialize, ::serde::Deserialize))]
106pub enum Mask {
107 AllTrue(usize),
109 AllFalse(usize),
111 Values(Arc<MaskValues>),
113}
114
115impl Debug for Mask {
116 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
117 match self {
118 Self::AllTrue(len) => write!(f, "All true({len})"),
119 Self::AllFalse(len) => write!(f, "All false({len})"),
120 Self::Values(mask) => write!(f, "{mask:?}"),
121 }
122 }
123}
124
125impl Default for Mask {
126 fn default() -> Self {
127 Self::new_true(0)
128 }
129}
130
131#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
133pub struct MaskValues {
134 buffer: BitBuffer,
135
136 #[cfg_attr(feature = "serde", serde(skip))]
139 indices: OnceLock<Vec<usize>>,
140 #[cfg_attr(feature = "serde", serde(skip))]
141 slices: OnceLock<Vec<(usize, usize)>>,
142
143 true_count: usize,
145 density: f64,
147}
148
149impl Debug for MaskValues {
150 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
151 write!(f, "true_count={}, ", self.true_count)?;
152 write!(f, "density={}, ", self.density)?;
153 if let Some(v) = self.indices.get() {
154 write!(f, "indices={v:?}, ")?;
155 }
156 if let Some(v) = self.slices.get() {
157 write!(f, "slices={v:?}, ")?;
158 }
159 if f.alternate() {
160 f.write_str("\n")?;
161 }
162 write!(f, "{}", self.buffer)
163 }
164}
165
166impl Mask {
167 pub fn new(length: usize, value: bool) -> Self {
169 if value {
170 Self::AllTrue(length)
171 } else {
172 Self::AllFalse(length)
173 }
174 }
175
176 #[inline]
178 pub fn new_true(length: usize) -> Self {
179 Self::AllTrue(length)
180 }
181
182 #[inline]
184 pub fn new_false(length: usize) -> Self {
185 Self::AllFalse(length)
186 }
187
188 pub fn from_buffer(buffer: BitBuffer) -> Self {
190 let len = buffer.len();
191 let true_count = buffer.true_count();
192
193 if true_count == 0 {
194 return Self::AllFalse(len);
195 }
196 if true_count == len {
197 return Self::AllTrue(len);
198 }
199
200 Self::Values(Arc::new(MaskValues {
201 buffer,
202 indices: Default::default(),
203 slices: Default::default(),
204 true_count,
205 density: true_count as f64 / len as f64,
206 }))
207 }
208
209 pub fn from_indices(len: usize, indices: impl IntoIterator<Item = usize>) -> Self {
211 let indices = indices.into_iter().collect::<Vec<_>>();
212 assert!(indices.is_sorted(), "Mask indices must be sorted");
213 assert!(
214 indices.windows(2).all(|w| w[0] != w[1]),
215 "Mask indices must be unique"
216 );
217 let buffer = BitBuffer::from_indices(len, indices.iter().copied());
218 debug_assert_eq!(buffer.len(), len);
219 let true_count = buffer.true_count();
220
221 if true_count == 0 {
222 return Self::AllFalse(len);
223 }
224 if true_count == len {
225 return Self::AllTrue(len);
226 }
227
228 Self::Values(Arc::new(MaskValues {
229 buffer,
230 indices: OnceLock::from(indices),
231 slices: Default::default(),
232 true_count,
233 density: true_count as f64 / len as f64,
234 }))
235 }
236
237 pub fn from_excluded_indices(len: usize, indices: impl IntoIterator<Item = usize>) -> Self {
239 let mut buf = BitBufferMut::new_set(len);
240
241 let mut false_count: usize = 0;
242 indices.into_iter().for_each(|idx| {
243 buf.unset(idx);
244 false_count += 1;
245 });
246 debug_assert_eq!(buf.len(), len);
247 let true_count = len - false_count;
248
249 if false_count == 0 {
251 return Self::AllTrue(len);
252 }
253 if false_count == len {
254 return Self::AllFalse(len);
255 }
256
257 Self::Values(Arc::new(MaskValues {
258 buffer: buf.freeze(),
259 indices: Default::default(),
260 slices: Default::default(),
261 true_count,
262 density: true_count as f64 / len as f64,
263 }))
264 }
265
266 pub fn from_slices(len: usize, vec: Vec<(usize, usize)>) -> Self {
269 Self::check_slices(len, &vec);
270 Self::from_slices_unchecked(len, vec)
271 }
272
273 fn from_slices_unchecked(len: usize, slices: Vec<(usize, usize)>) -> Self {
274 #[cfg(debug_assertions)]
275 Self::check_slices(len, &slices);
276
277 let true_count = slices.iter().map(|(b, e)| e - b).sum();
278 if true_count == 0 {
279 return Self::AllFalse(len);
280 }
281 if true_count == len {
282 return Self::AllTrue(len);
283 }
284
285 let mut buf = BitBufferMut::with_capacity(len);
286 let mut cursor = 0;
287 for (start, end) in slices.iter().copied() {
288 buf.append_n(false, start - cursor);
289 buf.append_n(true, end - start);
290 cursor = end;
291 }
292 buf.append_n(false, len - cursor);
293 debug_assert_eq!(buf.len(), len);
294
295 Self::Values(Arc::new(MaskValues {
296 buffer: buf.freeze(),
297 indices: Default::default(),
298 slices: OnceLock::from(slices),
299 true_count,
300 density: true_count as f64 / len as f64,
301 }))
302 }
303
304 #[inline(always)]
305 fn check_slices(len: usize, vec: &[(usize, usize)]) {
306 assert!(vec.iter().all(|&(b, e)| b < e && e <= len));
307 for (first, second) in vec.iter().tuple_windows() {
308 assert!(
309 first.0 < second.0,
310 "Slices must be sorted, got {first:?} and {second:?}"
311 );
312 assert!(
313 first.1 <= second.0,
314 "Slices must be non-overlapping, got {first:?} and {second:?}"
315 );
316 }
317 }
318
319 pub fn from_intersection_indices(
321 len: usize,
322 lhs: impl Iterator<Item = usize>,
323 rhs: impl Iterator<Item = usize>,
324 ) -> Self {
325 let mut intersection = Vec::with_capacity(len);
326 let mut lhs = lhs.peekable();
327 let mut rhs = rhs.peekable();
328 while let (Some(&l), Some(&r)) = (lhs.peek(), rhs.peek()) {
329 match l.cmp(&r) {
330 Ordering::Less => {
331 lhs.next();
332 }
333 Ordering::Greater => {
334 rhs.next();
335 }
336 Ordering::Equal => {
337 intersection.push(l);
338 lhs.next();
339 rhs.next();
340 }
341 }
342 }
343 Self::from_indices(len, intersection)
344 }
345
346 pub fn clear(&mut self) {
348 *self = Self::new_false(0);
349 }
350
351 #[inline]
353 pub fn len(&self) -> usize {
354 match self {
355 Self::AllTrue(len) => *len,
356 Self::AllFalse(len) => *len,
357 Self::Values(values) => values.len(),
358 }
359 }
360
361 #[inline]
363 pub fn is_empty(&self) -> bool {
364 match self {
365 Self::AllTrue(len) => *len == 0,
366 Self::AllFalse(len) => *len == 0,
367 Self::Values(values) => values.is_empty(),
368 }
369 }
370
371 #[inline]
373 pub fn true_count(&self) -> usize {
374 match &self {
375 Self::AllTrue(len) => *len,
376 Self::AllFalse(_) => 0,
377 Self::Values(values) => values.true_count,
378 }
379 }
380
381 #[inline]
383 pub fn false_count(&self) -> usize {
384 match &self {
385 Self::AllTrue(_) => 0,
386 Self::AllFalse(len) => *len,
387 Self::Values(values) => values.buffer.len() - values.true_count,
388 }
389 }
390
391 #[inline]
393 pub fn all_true(&self) -> bool {
394 match &self {
395 Self::AllTrue(_) => true,
396 Self::AllFalse(0) => true,
397 Self::AllFalse(_) => false,
398 Self::Values(values) => values.buffer.len() == values.true_count,
399 }
400 }
401
402 #[inline]
404 pub fn all_false(&self) -> bool {
405 self.true_count() == 0
406 }
407
408 #[inline]
410 pub fn density(&self) -> f64 {
411 match &self {
412 Self::AllTrue(_) => 1.0,
413 Self::AllFalse(_) => 0.0,
414 Self::Values(values) => values.density,
415 }
416 }
417
418 #[inline]
424 pub fn value(&self, idx: usize) -> bool {
425 match self {
426 Mask::AllTrue(_) => true,
427 Mask::AllFalse(_) => false,
428 Mask::Values(values) => values.buffer.value(idx),
429 }
430 }
431
432 #[inline]
438 pub fn iter(&self) -> MaskBoolIter<'_> {
439 match self {
440 Mask::AllTrue(len) => MaskBoolIter::Repeat {
441 value: true,
442 remaining: *len,
443 },
444 Mask::AllFalse(len) => MaskBoolIter::Repeat {
445 value: false,
446 remaining: *len,
447 },
448 Mask::Values(values) => MaskBoolIter::Bits(values.bit_buffer().iter()),
449 }
450 }
451
452 pub fn first(&self) -> Option<usize> {
454 match &self {
455 Self::AllTrue(len) => (*len > 0).then_some(0),
456 Self::AllFalse(_) => None,
457 Self::Values(values) => {
458 if let Some(indices) = values.indices.get() {
459 return indices.first().copied();
460 }
461 if let Some(slices) = values.slices.get() {
462 return slices.first().map(|(start, _)| *start);
463 }
464 values.buffer.set_indices().next()
465 }
466 }
467 }
468
469 pub fn last(&self) -> Option<usize> {
471 match &self {
472 Self::AllTrue(len) => (*len > 0).then_some(*len - 1),
473 Self::AllFalse(_) => None,
474 Self::Values(values) => {
475 if let Some(indices) = values.indices.get() {
476 return indices.last().copied();
477 }
478 if let Some(slices) = values.slices.get() {
479 return slices.last().map(|(_, end)| end - 1);
480 }
481
482 if values.true_count == 0 {
483 return None;
484 }
485
486 Some(
487 values
488 .buffer
489 .select(values.true_count - 1)
490 .unwrap_or_else(|| {
491 vortex_panic!(
492 "Rank {} out of bounds for mask with true count {}",
493 values.true_count - 1,
494 values.true_count
495 )
496 }),
497 )
498 }
499 }
500 }
501
502 pub fn rank(&self, n: usize) -> usize {
504 if n >= self.true_count() {
505 vortex_panic!(
506 "Rank {n} out of bounds for mask with true count {}",
507 self.true_count()
508 );
509 }
510 match &self {
511 Self::AllTrue(_) => n,
512 Self::AllFalse(_) => unreachable!("no true values in all-false mask"),
513 Self::Values(values) => {
514 if let Some(indices) = values.indices.get() {
515 return indices[n];
516 }
517
518 values.buffer.select(n).unwrap_or_else(|| {
519 vortex_panic!(
520 "Rank {} out of bounds for mask with true count {}",
521 values.true_count - 1,
522 values.true_count
523 )
524 })
525 }
526 }
527 }
528
529 pub fn slice(&self, range: impl RangeBounds<usize>) -> Self {
531 let start = match range.start_bound() {
532 Bound::Included(&s) => s,
533 Bound::Excluded(&s) => s + 1,
534 Bound::Unbounded => 0,
535 };
536 let end = match range.end_bound() {
537 Bound::Included(&e) => e + 1,
538 Bound::Excluded(&e) => e,
539 Bound::Unbounded => self.len(),
540 };
541
542 assert!(start <= end);
543 assert!(start <= self.len());
544 assert!(end <= self.len());
545 let len = end - start;
546
547 if len == self.len() {
550 return self.clone();
551 }
552
553 match &self {
554 Self::AllTrue(_) => Self::new_true(len),
555 Self::AllFalse(_) => Self::new_false(len),
556 Self::Values(values) => Self::from_buffer(values.buffer.slice(range)),
557 }
558 }
559
560 #[inline]
562 pub fn bit_buffer(&self) -> AllOr<&BitBuffer> {
563 match &self {
564 Self::AllTrue(_) => AllOr::All,
565 Self::AllFalse(_) => AllOr::None,
566 Self::Values(values) => AllOr::Some(&values.buffer),
567 }
568 }
569
570 #[inline]
573 pub fn to_bit_buffer(&self) -> BitBuffer {
574 match self {
575 Self::AllTrue(l) => BitBuffer::new_set(*l),
576 Self::AllFalse(l) => BitBuffer::new_unset(*l),
577 Self::Values(values) => values.bit_buffer().clone(),
578 }
579 }
580
581 #[inline]
584 pub fn into_bit_buffer(self) -> BitBuffer {
585 match self {
586 Self::AllTrue(l) => BitBuffer::new_set(l),
587 Self::AllFalse(l) => BitBuffer::new_unset(l),
588 Self::Values(values) => Arc::try_unwrap(values)
589 .map(|v| v.into_bit_buffer())
590 .unwrap_or_else(|v| v.bit_buffer().clone()),
591 }
592 }
593
594 #[inline]
596 pub fn indices(&self) -> AllOr<&[usize]> {
597 match &self {
598 Self::AllTrue(_) => AllOr::All,
599 Self::AllFalse(_) => AllOr::None,
600 Self::Values(values) => AllOr::Some(values.indices()),
601 }
602 }
603
604 #[inline]
606 pub fn slices(&self) -> AllOr<&[(usize, usize)]> {
607 match &self {
608 Self::AllTrue(_) => AllOr::All,
609 Self::AllFalse(_) => AllOr::None,
610 Self::Values(values) => AllOr::Some(values.slices()),
611 }
612 }
613
614 #[inline]
616 pub fn threshold_iter(&self, threshold: f64) -> AllOr<MaskIter<'_>> {
617 match &self {
618 Self::AllTrue(_) => AllOr::All,
619 Self::AllFalse(_) => AllOr::None,
620 Self::Values(values) => AllOr::Some(values.threshold_iter(threshold)),
621 }
622 }
623
624 #[inline]
626 pub fn values(&self) -> Option<&MaskValues> {
627 if let Self::Values(values) = self {
628 Some(values)
629 } else {
630 None
631 }
632 }
633
634 pub fn valid_counts_for_indices(&self, indices: &[usize]) -> Vec<usize> {
640 match self {
641 Self::AllTrue(_) => indices.to_vec(),
642 Self::AllFalse(_) => vec![0; indices.len()],
643 Self::Values(values) => {
644 let buffer = values.bit_buffer();
645 let mut valid_counts = Vec::with_capacity(indices.len());
646 let mut valid_count = 0;
647 let mut prev = 0;
648 for &next_idx in indices {
649 assert!(next_idx <= buffer.len(), "Row indices exceed array length");
650 if next_idx > prev {
653 valid_count += buffer.count_range(prev, next_idx);
654 prev = next_idx;
655 }
656 valid_counts.push(valid_count);
657 }
658
659 valid_counts
660 }
661 }
662 }
663
664 pub fn limit(self, limit: usize) -> Self {
666 if self.len() <= limit {
670 return self;
671 }
672
673 match &self {
674 Mask::AllTrue(len) => {
675 Self::from_iter([Self::new_true(limit), Self::new_false(len - limit)])
676 }
677 Mask::AllFalse(_) => self,
678 Mask::Values(mask_values) => {
679 if limit >= mask_values.true_count() {
680 return self;
681 }
682
683 let existing_buffer = mask_values.bit_buffer();
684
685 let mut new_buffer_builder = BitBufferMut::new_unset(mask_values.len());
686 debug_assert!(limit < mask_values.len());
687
688 for index in existing_buffer.set_indices().take(limit) {
689 unsafe { new_buffer_builder.set_unchecked(index) }
692 }
693
694 Self::from(new_buffer_builder.freeze())
695 }
696 }
697 }
698
699 pub fn concat<'a>(masks: impl Iterator<Item = &'a Self>) -> VortexResult<Self> {
701 let masks: Vec<_> = masks.collect();
702 let len = masks.iter().map(|t| t.len()).sum();
703
704 if masks.iter().all(|t| t.all_true()) {
705 return Ok(Mask::AllTrue(len));
706 }
707
708 if masks.iter().all(|t| t.all_false()) {
709 return Ok(Mask::AllFalse(len));
710 }
711
712 let mut builder = BitBufferMut::with_capacity(len);
713
714 for mask in masks {
715 match mask {
716 Mask::AllTrue(n) => builder.append_n(true, *n),
717 Mask::AllFalse(n) => builder.append_n(false, *n),
718 Mask::Values(v) => builder.append_buffer(v.bit_buffer()),
719 }
720 }
721
722 Ok(Mask::from_buffer(builder.freeze()))
723 }
724}
725
726impl MaskValues {
727 #[inline]
729 pub fn len(&self) -> usize {
730 self.buffer.len()
731 }
732
733 #[inline]
735 pub fn is_empty(&self) -> bool {
736 self.buffer.is_empty()
737 }
738
739 #[inline]
741 pub fn density(&self) -> f64 {
742 self.density
743 }
744
745 #[inline]
747 pub fn true_count(&self) -> usize {
748 self.true_count
749 }
750
751 #[inline]
753 pub fn bit_buffer(&self) -> &BitBuffer {
754 &self.buffer
755 }
756
757 #[inline]
759 pub fn into_bit_buffer(self) -> BitBuffer {
760 self.buffer
761 }
762
763 #[inline]
765 pub fn value(&self, index: usize) -> bool {
766 self.buffer.value(index)
767 }
768
769 pub fn indices(&self) -> &[usize] {
771 self.indices.get_or_init(|| {
772 if self.true_count == 0 {
773 return vec![];
774 }
775
776 if self.true_count == self.len() {
777 return (0..self.len()).collect();
778 }
779
780 if let Some(slices) = self.slices.get() {
781 let mut indices = Vec::with_capacity(self.true_count);
782 indices.extend(slices.iter().flat_map(|(start, end)| *start..*end));
783 debug_assert!(indices.is_sorted());
784 assert_eq!(indices.len(), self.true_count);
785 return indices;
786 }
787
788 let mut indices = Vec::with_capacity(self.true_count);
789 self.buffer.for_each_set_index(|i| indices.push(i));
793 debug_assert!(indices.is_sorted());
794 assert_eq!(indices.len(), self.true_count);
795 indices
796 })
797 }
798
799 #[inline]
804 pub fn cached_indices(&self) -> Option<&[usize]> {
805 self.indices.get().map(Vec::as_slice)
806 }
807
808 #[inline]
810 pub fn slices(&self) -> &[(usize, usize)] {
811 self.slices.get_or_init(|| {
812 if self.true_count == self.len() {
813 return vec![(0, self.len())];
814 }
815
816 self.buffer.set_slices().collect()
817 })
818 }
819
820 #[inline]
825 pub fn cached_slices(&self) -> Option<&[(usize, usize)]> {
826 self.slices.get().map(Vec::as_slice)
827 }
828
829 #[inline]
831 pub fn threshold_iter(&self, threshold: f64) -> MaskIter<'_> {
832 if self.density >= threshold {
833 MaskIter::Slices(self.slices())
834 } else {
835 MaskIter::Indices(self.indices())
836 }
837 }
838}
839
840pub enum MaskIter<'a> {
842 Indices(&'a [usize]),
844 Slices(&'a [(usize, usize)]),
846}
847
848pub enum MaskBoolIter<'a> {
852 Repeat {
854 value: bool,
856 remaining: usize,
858 },
859 Bits(BitIterator<'a>),
861}
862
863impl Iterator for MaskBoolIter<'_> {
864 type Item = bool;
865
866 #[inline]
867 fn next(&mut self) -> Option<Self::Item> {
868 match self {
869 Self::Repeat { remaining: 0, .. } => None,
870 Self::Repeat { value, remaining } => {
871 *remaining -= 1;
872 Some(*value)
873 }
874 Self::Bits(bits) => bits.next(),
875 }
876 }
877
878 #[inline]
879 fn size_hint(&self) -> (usize, Option<usize>) {
880 let remaining = match self {
881 Self::Repeat { remaining, .. } => *remaining,
882 Self::Bits(bits) => bits.len(),
883 };
884 (remaining, Some(remaining))
885 }
886}
887
888impl ExactSizeIterator for MaskBoolIter<'_> {}
889
890impl From<BitBuffer> for Mask {
891 fn from(value: BitBuffer) -> Self {
892 Self::from_buffer(value)
893 }
894}
895
896impl FromIterator<bool> for Mask {
897 #[inline]
898 fn from_iter<T: IntoIterator<Item = bool>>(iter: T) -> Self {
899 Self::from_buffer(BitBuffer::from_iter(iter))
900 }
901}
902
903impl FromIterator<Mask> for Mask {
904 fn from_iter<T: IntoIterator<Item = Mask>>(iter: T) -> Self {
905 let masks = iter
906 .into_iter()
907 .filter(|m| !m.is_empty())
908 .collect::<Vec<_>>();
909 let total_length = masks.iter().map(|v| v.len()).sum();
910
911 if masks.iter().all(|v| v.all_true()) {
913 return Self::AllTrue(total_length);
914 }
915 if masks.iter().all(|v| v.all_false()) {
917 return Self::AllFalse(total_length);
918 }
919
920 let mut buffer = BitBufferMut::with_capacity(total_length);
922 for mask in masks {
923 match mask {
924 Mask::AllTrue(count) => buffer.append_n(true, count),
925 Mask::AllFalse(count) => buffer.append_n(false, count),
926 Mask::Values(values) => {
927 buffer.append_buffer(values.bit_buffer());
928 }
929 };
930 }
931 Self::from_buffer(buffer.freeze())
932 }
933}