1use std::fmt::Debug;
12use std::ops::Range;
13
14use itertools::Itertools as _;
15use vortex_buffer::BitBuffer;
16use vortex_error::VortexExpect as _;
17use vortex_error::VortexResult;
18use vortex_error::vortex_bail;
19use vortex_error::vortex_err;
20use vortex_mask::Mask;
21use vortex_mask::MaskValues;
22
23use crate::ArrayRef;
24use crate::Canonical;
25use crate::ExecutionCtx;
26use crate::IntoArray;
27use crate::VortexSessionExecute;
28use crate::arrays::BoolArray;
29use crate::arrays::ChunkedArray;
30use crate::arrays::ConstantArray;
31use crate::arrays::scalar_fn::ScalarFnFactoryExt;
32use crate::builtins::ArrayBuiltins;
33use crate::dtype::DType;
34use crate::dtype::Nullability;
35use crate::legacy_session;
36use crate::optimizer::ArrayOptimizer;
37use crate::patches::Patches;
38use crate::scalar::Scalar;
39use crate::scalar_fn::fns::binary::Binary;
40use crate::scalar_fn::fns::operators::Operator;
41
42#[derive(Clone)]
44pub enum Validity {
45 NonNullable,
47 AllValid,
49 AllInvalid,
51 Array(ArrayRef),
55}
56
57impl Debug for Validity {
58 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
59 match self {
60 Self::NonNullable => write!(f, "NonNullable"),
61 Self::AllValid => write!(f, "AllValid"),
62 Self::AllInvalid => write!(f, "AllInvalid"),
63 Self::Array(arr) => write!(f, "SomeValid({})", arr.display_values()),
64 }
65 }
66}
67
68impl Validity {
69 pub fn execute(self, ctx: &mut ExecutionCtx) -> VortexResult<Validity> {
71 match self {
72 v @ Validity::NonNullable | v @ Validity::AllValid | v @ Validity::AllInvalid => Ok(v),
73 Validity::Array(a) => Ok(Validity::Array(a.execute::<Canonical>(ctx)?.into_array())),
74 }
75 }
76}
77
78impl Validity {
79 pub const DTYPE: DType = DType::Bool(Nullability::NonNullable);
81
82 pub fn to_array(&self, len: usize) -> ArrayRef {
84 match self {
85 Self::NonNullable | Self::AllValid => ConstantArray::new(true, len).into_array(),
86 Self::AllInvalid => ConstantArray::new(false, len).into_array(),
87 Self::Array(a) => a.clone(),
88 }
89 }
90
91 #[inline]
93 pub fn into_array(self) -> Option<ArrayRef> {
94 if let Self::Array(a) = self {
95 Some(a)
96 } else {
97 None
98 }
99 }
100
101 #[inline]
103 pub fn as_array(&self) -> Option<&ArrayRef> {
104 if let Self::Array(a) = self {
105 Some(a)
106 } else {
107 None
108 }
109 }
110
111 #[inline]
112 pub fn nullability(&self) -> Nullability {
113 if matches!(self, Self::NonNullable) {
114 Nullability::NonNullable
115 } else {
116 Nullability::Nullable
117 }
118 }
119
120 #[inline]
127 pub fn definitely_no_nulls(&self) -> bool {
128 matches!(self, Self::NonNullable | Self::AllValid)
129 }
130
131 #[inline]
140 pub fn definitely_all_null(&self) -> bool {
141 matches!(self, Self::AllInvalid)
142 }
143
144 pub fn execute_no_nulls(&self, length: usize, ctx: &mut ExecutionCtx) -> VortexResult<bool> {
150 match self {
151 Self::NonNullable | Self::AllValid => Ok(true),
152 Self::AllInvalid => Ok(length == 0),
153 Self::Array(_) => Ok(self.execute_mask(length, ctx)?.all_true()),
154 }
155 }
156
157 #[inline]
159 pub fn union_nullability(self, nullability: Nullability) -> Self {
160 match nullability {
161 Nullability::NonNullable => self,
162 Nullability::Nullable => self.into_nullable(),
163 }
164 }
165
166 #[inline]
168 pub fn execute_is_valid(&self, index: usize, ctx: &mut ExecutionCtx) -> VortexResult<bool> {
169 Ok(match self {
170 Self::NonNullable | Self::AllValid => true,
171 Self::AllInvalid => false,
172 Self::Array(a) => a
173 .execute_scalar(index, ctx)?
174 .as_bool()
175 .value()
176 .ok_or_else(|| vortex_err!("validity value at index {index} is null"))?,
177 })
178 }
179
180 #[inline]
182 pub fn execute_is_null(&self, index: usize, ctx: &mut ExecutionCtx) -> VortexResult<bool> {
183 Ok(!self.execute_is_valid(index, ctx)?)
184 }
185
186 #[deprecated(note = "use `execute_is_valid` with an explicit `ExecutionCtx`")]
188 #[inline]
189 #[allow(clippy::disallowed_methods)]
190 pub fn is_valid(&self, index: usize) -> VortexResult<bool> {
191 self.execute_is_valid(index, &mut legacy_session().create_execution_ctx())
192 }
193
194 #[deprecated(note = "use `execute_is_null` with an explicit `ExecutionCtx`")]
196 #[inline]
197 #[allow(clippy::disallowed_methods)]
198 pub fn is_null(&self, index: usize) -> VortexResult<bool> {
199 self.execute_is_null(index, &mut legacy_session().create_execution_ctx())
200 }
201
202 #[inline]
203 pub fn slice(&self, range: Range<usize>) -> VortexResult<Self> {
204 match self {
205 Self::Array(a) => Ok(Self::Array(a.slice(range)?)),
206 Self::NonNullable | Self::AllValid | Self::AllInvalid => Ok(self.clone()),
207 }
208 }
209
210 pub fn take(&self, indices: &ArrayRef) -> VortexResult<Self> {
211 match self {
212 Self::NonNullable => indices.validity(),
213 Self::AllValid => Ok(match indices.validity()? {
214 Self::NonNullable => Self::AllValid,
215 v => v,
216 }),
217 Self::AllInvalid => Ok(Self::AllInvalid),
218 Self::Array(is_valid) => {
219 let maybe_is_valid = is_valid.take(indices.clone())?;
220 let is_valid = maybe_is_valid.fill_null(Scalar::from(false))?;
222 Ok(Self::Array(is_valid))
223 }
224 }
225 }
226
227 pub fn not(&self) -> VortexResult<Self> {
229 match self {
230 Validity::NonNullable => Ok(Validity::NonNullable),
231 Validity::AllValid => Ok(Validity::AllInvalid),
232 Validity::AllInvalid => Ok(Validity::AllValid),
233 Validity::Array(arr) => Ok(Validity::Array(arr.not()?)),
234 }
235 }
236
237 pub fn filter(&self, mask: &Mask) -> VortexResult<Self> {
245 match self {
248 v @ (Validity::NonNullable | Validity::AllValid | Validity::AllInvalid) => {
249 Ok(v.clone())
250 }
251 Validity::Array(arr) => Ok(Validity::Array(arr.filter(mask.clone())?)),
252 }
253 }
254
255 #[deprecated(note = "Use execute_mask")]
259 pub fn to_mask(&self, length: usize, ctx: &mut ExecutionCtx) -> VortexResult<Mask> {
260 match self {
261 Self::NonNullable | Self::AllValid => Ok(Mask::new_true(length)),
262 Self::AllInvalid => Ok(Mask::new_false(length)),
263 Self::Array(arr) => arr.clone().execute::<Mask>(ctx),
264 }
265 }
266
267 #[inline]
268 pub fn execute_mask(&self, length: usize, ctx: &mut ExecutionCtx) -> VortexResult<Mask> {
269 match self {
270 Self::NonNullable | Self::AllValid => Ok(Mask::AllTrue(length)),
271 Self::AllInvalid => Ok(Mask::AllFalse(length)),
272 Self::Array(arr) => {
273 assert_eq!(
274 arr.len(),
275 length,
276 "Validity::Array length must equal to_logical's argument: {}, {}.",
277 arr.len(),
278 length,
279 );
280 arr.clone().execute::<Mask>(ctx)
283 }
284 }
285 }
286
287 pub fn mask_eq(
290 &self,
291 other: &Validity,
292 length: usize,
293 ctx: &mut ExecutionCtx,
294 ) -> VortexResult<bool> {
295 match (self, other) {
296 (
298 Validity::NonNullable | Validity::AllValid,
299 Validity::NonNullable | Validity::AllValid,
300 )
301 | (Validity::AllInvalid, Validity::AllInvalid) => Ok(true),
302 _ => Ok(self.execute_mask(length, ctx)? == other.execute_mask(length, ctx)?),
303 }
304 }
305
306 #[inline]
308 pub fn and(self, rhs: Validity) -> VortexResult<Validity> {
309 Ok(match (self, rhs) {
310 (Validity::NonNullable, Validity::NonNullable) => Validity::NonNullable,
312 (Validity::AllInvalid, _) | (_, Validity::AllInvalid) => Validity::AllInvalid,
314 (Validity::Array(a), Validity::AllValid)
316 | (Validity::Array(a), Validity::NonNullable)
317 | (Validity::NonNullable, Validity::Array(a))
318 | (Validity::AllValid, Validity::Array(a)) => Validity::Array(a),
319 (Validity::NonNullable, Validity::AllValid)
321 | (Validity::AllValid, Validity::NonNullable)
322 | (Validity::AllValid, Validity::AllValid) => Validity::AllValid,
323 (Validity::Array(lhs), Validity::Array(rhs)) => Validity::Array(
325 Binary
326 .try_new_array(lhs.len(), Operator::And, [lhs, rhs])?
327 .optimize()?,
328 ),
329 })
330 }
331
332 pub fn patch(
333 self,
334 len: usize,
335 indices_offset: usize,
336 indices: &ArrayRef,
337 patches: &Validity,
338 ctx: &mut ExecutionCtx,
339 ) -> VortexResult<Self> {
340 match (&self, patches) {
341 (Validity::NonNullable, Validity::NonNullable) => return Ok(Validity::NonNullable),
342 (Validity::NonNullable, _) => {
343 vortex_bail!("Can't patch a non-nullable validity with nullable validity")
344 }
345 (_, Validity::NonNullable) => {
346 vortex_bail!("Can't patch a nullable validity with non-nullable validity")
347 }
348 (Validity::AllValid, Validity::AllValid) => return Ok(Validity::AllValid),
349 (Validity::AllInvalid, Validity::AllInvalid) => return Ok(Validity::AllInvalid),
350 _ => {}
351 };
352
353 if matches!(self, Validity::NonNullable) {
354 return Ok(Self::NonNullable);
355 }
356
357 let source = match self {
359 Validity::NonNullable => BoolArray::from(BitBuffer::new_set(len)),
360 Validity::AllValid => BoolArray::from(BitBuffer::new_set(len)),
361 Validity::AllInvalid => BoolArray::from(BitBuffer::new_unset(len)),
362 Validity::Array(a) => a.execute::<BoolArray>(ctx)?,
363 };
364
365 let patch_values = match patches {
366 Validity::NonNullable => BoolArray::from(BitBuffer::new_set(indices.len())),
367 Validity::AllValid => BoolArray::from(BitBuffer::new_set(indices.len())),
368 Validity::AllInvalid => BoolArray::from(BitBuffer::new_unset(indices.len())),
369 Validity::Array(a) => a.clone().execute::<BoolArray>(ctx)?,
370 };
371
372 let patches = Patches::new(
373 len,
374 indices_offset,
375 indices.clone(),
376 patch_values.into_array(),
377 None,
379 )?;
380
381 Ok(Self::Array(source.patch(&patches, ctx)?.into_array()))
382 }
383
384 #[inline]
386 pub fn into_nullable(self) -> Validity {
387 match self {
388 Self::NonNullable => Self::AllValid,
389 Self::AllValid | Self::AllInvalid | Self::Array(_) => self,
390 }
391 }
392
393 #[inline]
399 pub fn into_non_nullable(self, len: usize, ctx: &mut ExecutionCtx) -> Option<Validity> {
400 match self {
401 _ if len == 0 => Some(Validity::NonNullable),
402 Self::NonNullable => Some(Self::NonNullable),
403 Self::AllValid => Some(Self::NonNullable),
404 Self::AllInvalid => None,
405 Self::Array(is_valid) => {
406 is_valid
407 .statistics()
408 .compute_min::<bool>(ctx)
409 .vortex_expect("validity array must support min")
410 .then(|| {
411 Self::NonNullable
413 })
414 }
415 }
416 }
417
418 #[inline]
429 pub fn trivial_into_non_nullable(self, len: usize) -> VortexResult<Option<Validity>> {
430 match self {
431 _ if len == 0 => Ok(Some(Validity::NonNullable)),
432 Self::NonNullable => Ok(Some(Self::NonNullable)),
433 Self::AllValid => Ok(Some(Self::NonNullable)),
434 Self::AllInvalid => {
435 Err(vortex_err!(InvalidArgument: "Cannot cast AllInvalid to NonNullable"))
436 }
437 Self::Array(_) => Ok(None),
438 }
439 }
440
441 #[inline]
457 pub fn cast_nullability(
458 self,
459 nullability: Nullability,
460 len: usize,
461 ctx: &mut ExecutionCtx,
462 ) -> VortexResult<Validity> {
463 match nullability {
464 Nullability::NonNullable => self.into_non_nullable(len, ctx).ok_or_else(|| {
465 vortex_err!(InvalidArgument: "Cannot cast array with invalid values to non-nullable type.")
466 }),
467 Nullability::Nullable => Ok(self.into_nullable()),
468 }
469 }
470
471 #[inline]
496 pub fn trivially_cast_nullability(
497 self,
498 nullability: Nullability,
499 len: usize,
500 ) -> VortexResult<Option<Validity>> {
501 match nullability {
502 Nullability::NonNullable => self.trivial_into_non_nullable(len),
503 Nullability::Nullable => Ok(Some(self.into_nullable())),
504 }
505 }
506
507 #[inline]
509 pub fn maybe_len(&self) -> Option<usize> {
510 match self {
511 Self::NonNullable | Self::AllValid | Self::AllInvalid => None,
512 Self::Array(a) => Some(a.len()),
513 }
514 }
515}
516
517impl From<BitBuffer> for Validity {
518 #[inline]
519 fn from(value: BitBuffer) -> Self {
520 let true_count = value.true_count();
521 if true_count == value.len() {
522 Self::AllValid
523 } else if true_count == 0 {
524 Self::AllInvalid
525 } else {
526 Self::Array(BoolArray::from(value).into_array())
527 }
528 }
529}
530
531impl FromIterator<Mask> for Validity {
532 #[inline]
533 fn from_iter<T: IntoIterator<Item = Mask>>(iter: T) -> Self {
534 Validity::from_mask(iter.into_iter().collect(), Nullability::Nullable)
535 }
536}
537
538impl FromIterator<bool> for Validity {
539 #[inline]
540 fn from_iter<T: IntoIterator<Item = bool>>(iter: T) -> Self {
541 Validity::from(BitBuffer::from_iter(iter))
542 }
543}
544
545impl From<Nullability> for Validity {
546 #[inline]
547 fn from(value: Nullability) -> Self {
548 Validity::from(&value)
549 }
550}
551
552impl From<&Nullability> for Validity {
553 #[inline]
554 fn from(value: &Nullability) -> Self {
555 match *value {
556 Nullability::NonNullable => Validity::NonNullable,
557 Nullability::Nullable => Validity::AllValid,
558 }
559 }
560}
561
562impl Validity {
563 pub fn concat(validities: Vec<(Validity, usize)>) -> Option<Self> {
567 let mut validity_kinds = validities
568 .iter()
569 .map(|(v, _)| std::mem::discriminant(v))
570 .unique();
571 let validity_kind = validity_kinds.next()?;
572 if validity_kinds.next().is_none() {
573 if validity_kind == std::mem::discriminant(&Validity::AllValid) {
576 return Some(Validity::AllValid);
577 }
578 if validity_kind == std::mem::discriminant(&Validity::AllInvalid) {
579 return Some(Validity::AllInvalid);
580 }
581 if validity_kind == std::mem::discriminant(&Validity::NonNullable) {
582 return Some(Validity::NonNullable);
583 }
584 }
585
586 Some(Validity::Array(
587 unsafe {
588 ChunkedArray::new_unchecked(
589 validities.into_iter().map(|(v, len)| v.to_array(len)),
590 DType::Bool(Nullability::NonNullable),
591 )
592 }
593 .into_array(),
594 ))
595 }
596}
597
598impl Validity {
599 pub fn from_bit_buffer(buffer: BitBuffer, nullability: Nullability) -> Self {
600 if buffer.true_count() == buffer.len() {
601 nullability.into()
602 } else if buffer.true_count() == 0 {
603 Validity::AllInvalid
604 } else {
605 Validity::Array(BoolArray::new(buffer, Validity::NonNullable).into_array())
606 }
607 }
608
609 pub fn from_mask(mask: Mask, nullability: Nullability) -> Self {
610 assert!(
611 nullability == Nullability::Nullable || matches!(mask, Mask::AllTrue(_)),
612 "NonNullable validity must be AllValid",
613 );
614 match mask {
615 Mask::AllTrue(_) => match nullability {
616 Nullability::NonNullable => Validity::NonNullable,
617 Nullability::Nullable => Validity::AllValid,
618 },
619 Mask::AllFalse(_) => Validity::AllInvalid,
620 Mask::Values(values) => Validity::Array(values.into_array()),
621 }
622 }
623}
624
625impl IntoArray for Mask {
626 #[inline]
627 fn into_array(self) -> ArrayRef {
628 match self {
629 Self::AllTrue(len) => ConstantArray::new(true, len).into_array(),
630 Self::AllFalse(len) => ConstantArray::new(false, len).into_array(),
631 Self::Values(a) => a.into_array(),
632 }
633 }
634}
635
636impl IntoArray for &MaskValues {
637 #[inline]
638 fn into_array(self) -> ArrayRef {
639 BoolArray::new(self.bit_buffer().clone(), Validity::NonNullable).into_array()
640 }
641}
642
643#[cfg(test)]
644mod tests {
645 use rstest::rstest;
646 use vortex_buffer::Buffer;
647 use vortex_buffer::buffer;
648 use vortex_mask::Mask;
649
650 use crate::ArrayRef;
651 use crate::IntoArray;
652 use crate::VortexSessionExecute;
653 use crate::array_session;
654 use crate::arrays::PrimitiveArray;
655 use crate::dtype::Nullability;
656 use crate::validity::BoolArray;
657 use crate::validity::Validity;
658
659 #[rstest]
660 #[case(Validity::AllValid, 5, &[2, 4], Validity::AllValid, Validity::AllValid)]
661 #[case(
662 Validity::AllValid,
663 5,
664 &[2, 4],
665 Validity::AllInvalid,
666 Validity::Array(BoolArray::from_iter([true, true, false, true, false]).into_array())
667 )]
668 #[case(
669 Validity::AllValid,
670 5,
671 &[2, 4],
672 Validity::Array(BoolArray::from_iter([true, false]).into_array()),
673 Validity::Array(BoolArray::from_iter([true, true, true, true, false]).into_array())
674 )]
675 #[case(
676 Validity::AllInvalid,
677 5,
678 &[2, 4],
679 Validity::AllValid,
680 Validity::Array(BoolArray::from_iter([false, false, true, false, true]).into_array())
681 )]
682 #[case(Validity::AllInvalid, 5, &[2, 4], Validity::AllInvalid, Validity::AllInvalid)]
683 #[case(
684 Validity::AllInvalid,
685 5,
686 &[2, 4],
687 Validity::Array(BoolArray::from_iter([true, false]).into_array()),
688 Validity::Array(BoolArray::from_iter([false, false, true, false, false]).into_array())
689 )]
690 #[case(
691 Validity::Array(BoolArray::from_iter([false, true, false, true, false]).into_array()),
692 5,
693 &[2, 4],
694 Validity::AllValid,
695 Validity::Array(BoolArray::from_iter([false, true, true, true, true]).into_array())
696 )]
697 #[case(
698 Validity::Array(BoolArray::from_iter([false, true, false, true, false]).into_array()),
699 5,
700 &[2, 4],
701 Validity::AllInvalid,
702 Validity::Array(BoolArray::from_iter([false, true, false, true, false]).into_array())
703 )]
704 #[case(
705 Validity::Array(BoolArray::from_iter([false, true, false, true, false]).into_array()),
706 5,
707 &[2, 4],
708 Validity::Array(BoolArray::from_iter([true, false]).into_array()),
709 Validity::Array(BoolArray::from_iter([false, true, true, true, false]).into_array())
710 )]
711
712 fn patch_validity(
713 #[case] validity: Validity,
714 #[case] len: usize,
715 #[case] positions: &[u64],
716 #[case] patches: Validity,
717 #[case] expected: Validity,
718 ) {
719 let indices =
720 PrimitiveArray::new(Buffer::copy_from(positions), Validity::NonNullable).into_array();
721
722 let mut ctx = array_session().create_execution_ctx();
723
724 assert!(
725 validity
726 .patch(len, 0, &indices, &patches, &mut ctx,)
727 .unwrap()
728 .mask_eq(&expected, len, &mut ctx)
729 .unwrap()
730 );
731 }
732
733 #[test]
734 #[should_panic]
735 fn out_of_bounds_patch() {
736 let mut ctx = array_session().create_execution_ctx();
737 Validity::NonNullable
738 .patch(
739 2,
740 0,
741 &buffer![4].into_array(),
742 &Validity::AllInvalid,
743 &mut ctx,
744 )
745 .unwrap();
746 }
747
748 #[test]
749 #[should_panic]
750 fn into_validity_nullable() {
751 Validity::from_mask(Mask::AllFalse(10), Nullability::NonNullable);
752 }
753
754 #[test]
755 #[should_panic]
756 fn into_validity_nullable_array() {
757 Validity::from_mask(Mask::from_iter(vec![true, false]), Nullability::NonNullable);
758 }
759
760 #[rstest]
761 #[case(
762 Validity::AllValid,
763 PrimitiveArray::new(buffer![0, 1], Validity::from_iter(vec![true, false])).into_array(),
764 Validity::from_iter(vec![true, false])
765 )]
766 #[case(Validity::AllValid, buffer![0, 1].into_array(), Validity::AllValid)]
767 #[case(
768 Validity::AllValid,
769 PrimitiveArray::new(buffer![0, 1], Validity::AllInvalid).into_array(),
770 Validity::AllInvalid
771 )]
772 #[case(
773 Validity::NonNullable,
774 PrimitiveArray::new(buffer![0, 1], Validity::from_iter(vec![true, false])).into_array(),
775 Validity::from_iter(vec![true, false])
776 )]
777 #[case(Validity::NonNullable, buffer![0, 1].into_array(), Validity::NonNullable)]
778 #[case(
779 Validity::NonNullable,
780 PrimitiveArray::new(buffer![0, 1], Validity::AllInvalid).into_array(),
781 Validity::AllInvalid
782 )]
783 fn validity_take(
784 #[case] validity: Validity,
785 #[case] indices: ArrayRef,
786 #[case] expected: Validity,
787 ) {
788 let mut ctx = array_session().create_execution_ctx();
789 assert!(
790 validity
791 .take(&indices)
792 .unwrap()
793 .mask_eq(&expected, indices.len(), &mut ctx)
794 .unwrap()
795 );
796 }
797
798 #[rstest]
799 #[case(Validity::NonNullable, Validity::AllValid, true)]
801 #[case(Validity::AllValid, Validity::NonNullable, true)]
802 #[case(Validity::AllValid, Validity::AllInvalid, false)]
803 #[case(Validity::NonNullable, Validity::AllInvalid, false)]
804 #[case(
806 Validity::Array(BoolArray::from_iter([true, true, true]).into_array()),
807 Validity::AllValid,
808 true
809 )]
810 #[case(
811 Validity::NonNullable,
812 Validity::Array(BoolArray::from_iter([true, true, true]).into_array()),
813 true
814 )]
815 #[case(
816 Validity::Array(BoolArray::from_iter([false, false, false]).into_array()),
817 Validity::AllInvalid,
818 true
819 )]
820 #[case(
821 Validity::Array(BoolArray::from_iter([true, false, true]).into_array()),
822 Validity::AllValid,
823 false
824 )]
825 #[case(
826 Validity::Array(BoolArray::from_iter([true, false, true]).into_array()),
827 Validity::AllInvalid,
828 false
829 )]
830 fn mask_eq_mixed_variants(
831 #[case] lhs: Validity,
832 #[case] rhs: Validity,
833 #[case] expected: bool,
834 ) -> vortex_error::VortexResult<()> {
835 let mut ctx = array_session().create_execution_ctx();
836 assert_eq!(lhs.mask_eq(&rhs, 3, &mut ctx)?, expected);
837 Ok(())
838 }
839}