vortex_array/arrays/masked/
execute.rs1use std::sync::Arc;
7
8use vortex_error::VortexResult;
9
10use crate::Canonical;
11use crate::IntoArray;
12use crate::arrays::BoolArray;
13use crate::arrays::DecimalArray;
14use crate::arrays::ExtensionArray;
15use crate::arrays::FixedSizeListArray;
16use crate::arrays::ListViewArray;
17use crate::arrays::MaskedArray;
18use crate::arrays::PrimitiveArray;
19use crate::arrays::StructArray;
20use crate::arrays::VarBinViewArray;
21use crate::arrays::VariantArray;
22use crate::arrays::bool::BoolArrayExt;
23use crate::arrays::extension::ExtensionArrayExt;
24use crate::arrays::fixed_size_list::FixedSizeListArrayExt;
25use crate::arrays::listview::ListViewArrayExt;
26use crate::arrays::struct_::StructArrayExt;
27use crate::arrays::variant::VariantArrayExt;
28use crate::builtins::ArrayBuiltins;
29use crate::executor::ExecutionCtx;
30use crate::validity::Validity;
31
32pub fn mask_validity_canonical(
38 canonical: Canonical,
39 validity: Validity,
40 ctx: &mut ExecutionCtx,
41) -> VortexResult<Canonical> {
42 Ok(match canonical {
43 n @ Canonical::Null(_) => n,
44 Canonical::Bool(a) => Canonical::Bool(mask_validity_bool(a, validity)?),
45 Canonical::Primitive(a) => Canonical::Primitive(mask_validity_primitive(a, validity)?),
46 Canonical::Decimal(a) => Canonical::Decimal(mask_validity_decimal(a, validity)?),
47 Canonical::VarBinView(a) => Canonical::VarBinView(mask_validity_varbinview(a, validity)?),
48 Canonical::List(a) => Canonical::List(mask_validity_listview(a, validity)?),
49 Canonical::FixedSizeList(a) => {
50 Canonical::FixedSizeList(mask_validity_fixed_size_list(a, validity)?)
51 }
52 Canonical::Struct(a) => Canonical::Struct(mask_validity_struct(a, validity)?),
53 Canonical::Union(_) => {
54 todo!("TODO(connor)[Union]: implement masking for Union arrays")
55 }
56 Canonical::Extension(a) => Canonical::Extension(mask_validity_extension(a, validity, ctx)?),
57 Canonical::Variant(a) => Canonical::Variant(mask_validity_variant(a, validity, ctx)?),
58 })
59}
60
61fn mask_validity_bool(array: BoolArray, mask: Validity) -> VortexResult<BoolArray> {
62 let new_validity = Validity::and(array.validity()?, mask)?;
63 Ok(BoolArray::new(array.to_bit_buffer(), new_validity))
64}
65
66fn mask_validity_primitive(
67 array: PrimitiveArray,
68 validity: Validity,
69) -> VortexResult<PrimitiveArray> {
70 let ptype = array.ptype();
71 let new_validity = Validity::and(array.validity()?, validity)?;
72 Ok(unsafe {
74 PrimitiveArray::new_unchecked_from_handle(
75 array.buffer_handle().clone(),
76 ptype,
77 new_validity,
78 )
79 })
80}
81
82fn mask_validity_decimal(array: DecimalArray, validity: Validity) -> VortexResult<DecimalArray> {
83 let new_validity = Validity::and(array.validity()?, validity)?;
84 Ok(unsafe {
86 DecimalArray::new_unchecked_handle(
87 array.buffer_handle().clone(),
88 array.values_type(),
89 array.decimal_dtype(),
90 new_validity,
91 )
92 })
93}
94
95fn mask_validity_varbinview(
97 array: VarBinViewArray,
98 validity: Validity,
99) -> VortexResult<VarBinViewArray> {
100 let dtype = array.dtype().as_nullable();
101 let new_validity = Validity::and(array.validity()?, validity)?;
102 Ok(unsafe {
104 VarBinViewArray::new_handle_unchecked(
105 array.views_handle().clone(),
106 Arc::clone(array.data_buffers()),
107 dtype,
108 new_validity,
109 )
110 })
111}
112
113fn mask_validity_listview(array: ListViewArray, validity: Validity) -> VortexResult<ListViewArray> {
114 let new_validity = Validity::and(array.validity()?, validity)?;
115 let is_zctl = array.is_zero_copy_to_list();
117 Ok(unsafe {
118 ListViewArray::new_unchecked(
119 array.elements().clone(),
120 array.offsets().clone(),
121 array.sizes().clone(),
122 new_validity,
123 )
124 .with_zero_copy_to_list(is_zctl)
125 })
126}
127
128fn mask_validity_fixed_size_list(
129 array: FixedSizeListArray,
130 validity: Validity,
131) -> VortexResult<FixedSizeListArray> {
132 let len = array.len();
133 let list_size = array.list_size();
134 let new_validity = Validity::and(array.validity()?, validity)?;
135 Ok(unsafe {
137 FixedSizeListArray::new_unchecked(array.elements().clone(), list_size, new_validity, len)
138 })
139}
140
141fn mask_validity_struct(array: StructArray, validity: Validity) -> VortexResult<StructArray> {
142 let len = array.len();
143 let new_validity = Validity::and(array.validity()?, validity)?;
144 let fields = array.unmasked_fields();
145 let struct_fields = array.struct_fields();
146 Ok(unsafe { StructArray::new_unchecked(fields, struct_fields.clone(), len, new_validity) })
148}
149
150fn mask_validity_extension(
151 array: ExtensionArray,
152 validity: Validity,
153 ctx: &mut ExecutionCtx,
154) -> VortexResult<ExtensionArray> {
155 let storage = array.storage_array().clone().execute::<Canonical>(ctx)?;
157 let masked_storage = mask_validity_canonical(storage, validity, ctx)?;
158 let masked_storage = masked_storage.into_array();
159 Ok(ExtensionArray::new(
160 array
161 .ext_dtype()
162 .with_nullability(masked_storage.dtype().nullability()),
163 masked_storage,
164 ))
165}
166
167fn mask_validity_variant(
168 array: VariantArray,
169 validity: Validity,
170 ctx: &mut ExecutionCtx,
171) -> VortexResult<VariantArray> {
172 let core_storage = array.core_storage().clone();
173 let len = core_storage.len();
174 let core_validity = core_storage.validity()?;
175 let shredded_validity = validity.clone();
176
177 let masked_core_storage = match core_validity {
178 Validity::NonNullable | Validity::AllValid => {
179 MaskedArray::try_new(core_storage, validity)?.into_array()
181 }
182 Validity::AllInvalid => {
183 core_storage
185 }
186 Validity::Array(_) => {
187 core_storage.mask(validity.to_array(len))?
190 }
191 };
192 let masked_shredded = if let Some(shredded) = array.shredded() {
193 let canonical = shredded.clone().execute::<Canonical>(ctx)?;
194 Some(mask_validity_canonical(canonical, shredded_validity, ctx)?.into_array())
195 } else {
196 None
197 };
198
199 VariantArray::try_new(masked_core_storage, masked_shredded)
200}