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::UnionArray;
21use crate::arrays::VarBinViewArray;
22use crate::arrays::VariantArray;
23use crate::arrays::bool::BoolArrayExt;
24use crate::arrays::extension::ExtensionArrayExt;
25use crate::arrays::fixed_size_list::FixedSizeListArrayExt;
26use crate::arrays::fixed_size_list::FixedSizeListArraySlotsExt;
27use crate::arrays::listview::ListViewArraySlotsExt;
28use crate::arrays::struct_::StructArrayExt;
29use crate::arrays::union::UnionArrayExt;
30use crate::arrays::variant::VariantArraySlotsExt;
31use crate::builtins::ArrayBuiltins;
32use crate::executor::ExecutionCtx;
33use crate::validity::Validity;
34
35pub fn mask_validity_canonical(
41 canonical: Canonical,
42 validity: Validity,
43 ctx: &mut ExecutionCtx,
44) -> VortexResult<Canonical> {
45 Ok(match canonical {
46 n @ Canonical::Null(_) => n,
47 Canonical::Bool(a) => Canonical::Bool(mask_validity_bool(a, validity)?),
48 Canonical::Primitive(a) => Canonical::Primitive(mask_validity_primitive(a, validity)?),
49 Canonical::Decimal(a) => Canonical::Decimal(mask_validity_decimal(a, validity)?),
50 Canonical::VarBinView(a) => Canonical::VarBinView(mask_validity_varbinview(a, validity)?),
51 Canonical::List(a) => Canonical::List(mask_validity_listview(a, validity)?),
52 Canonical::FixedSizeList(a) => {
53 Canonical::FixedSizeList(mask_validity_fixed_size_list(a, validity)?)
54 }
55 Canonical::Struct(a) => Canonical::Struct(mask_validity_struct(a, validity)?),
56 Canonical::Union(a) => Canonical::Union(mask_validity_union(a, validity)?),
57 Canonical::Extension(a) => Canonical::Extension(mask_validity_extension(a, validity, ctx)?),
58 Canonical::Variant(a) => Canonical::Variant(mask_validity_variant(a, validity, ctx)?),
59 })
60}
61
62fn mask_validity_bool(array: BoolArray, mask: Validity) -> VortexResult<BoolArray> {
63 let new_validity = Validity::and(array.validity()?, mask)?;
64 Ok(BoolArray::new(array.to_bit_buffer(), new_validity))
65}
66
67fn mask_validity_primitive(
68 array: PrimitiveArray,
69 validity: Validity,
70) -> VortexResult<PrimitiveArray> {
71 let ptype = array.ptype();
72 let new_validity = Validity::and(array.validity()?, validity)?;
73 Ok(unsafe {
75 PrimitiveArray::new_unchecked_from_handle(
76 array.buffer_handle().clone(),
77 ptype,
78 new_validity,
79 )
80 })
81}
82
83fn mask_validity_decimal(array: DecimalArray, validity: Validity) -> VortexResult<DecimalArray> {
84 let new_validity = Validity::and(array.validity()?, validity)?;
85 Ok(unsafe {
87 DecimalArray::new_unchecked_handle(
88 array.buffer_handle().clone(),
89 array.values_type(),
90 array.decimal_dtype(),
91 new_validity,
92 )
93 })
94}
95
96fn mask_validity_varbinview(
98 array: VarBinViewArray,
99 validity: Validity,
100) -> VortexResult<VarBinViewArray> {
101 let dtype = array.dtype().as_nullable();
102 let new_validity = Validity::and(array.validity()?, validity)?;
103 Ok(unsafe {
105 VarBinViewArray::new_handle_unchecked(
106 array.views_handle().clone(),
107 Arc::clone(array.data_buffers()),
108 dtype,
109 new_validity,
110 )
111 })
112}
113
114fn mask_validity_listview(array: ListViewArray, validity: Validity) -> VortexResult<ListViewArray> {
115 let new_validity = Validity::and(array.validity()?, validity)?;
116 let is_zctl = array.is_zero_copy_to_list();
118 Ok(unsafe {
119 ListViewArray::new_unchecked(
120 array.elements().clone(),
121 array.offsets().clone(),
122 array.sizes().clone(),
123 new_validity,
124 )
125 .with_zero_copy_to_list(is_zctl)
126 })
127}
128
129fn mask_validity_fixed_size_list(
130 array: FixedSizeListArray,
131 validity: Validity,
132) -> VortexResult<FixedSizeListArray> {
133 let len = array.len();
134 let list_size = array.list_size();
135 let new_validity = Validity::and(array.validity()?, validity)?;
136 Ok(unsafe {
138 FixedSizeListArray::new_unchecked(array.elements().clone(), list_size, new_validity, len)
139 })
140}
141
142fn mask_validity_struct(array: StructArray, validity: Validity) -> VortexResult<StructArray> {
143 let len = array.len();
144 let new_validity = Validity::and(array.validity()?, validity)?;
145 let fields = array.unmasked_fields();
146 let struct_fields = array.struct_fields();
147 Ok(unsafe { StructArray::new_unchecked(fields, struct_fields.clone(), len, new_validity) })
149}
150
151fn mask_validity_union(array: UnionArray, validity: Validity) -> VortexResult<UnionArray> {
152 let type_ids = array
153 .type_ids()
154 .clone()
155 .mask(validity.to_array(array.len()))?;
156 let variants = array.variants().clone();
157 let children = array.children();
158
159 Ok(unsafe { UnionArray::new_unchecked(type_ids, variants, children) })
161}
162
163fn mask_validity_extension(
164 array: ExtensionArray,
165 validity: Validity,
166 ctx: &mut ExecutionCtx,
167) -> VortexResult<ExtensionArray> {
168 let storage = array.storage_array().clone().execute::<Canonical>(ctx)?;
170 let masked_storage = mask_validity_canonical(storage, validity, ctx)?;
171 let masked_storage = masked_storage.into_array();
172 Ok(ExtensionArray::new(
173 array
174 .ext_dtype()
175 .with_nullability(masked_storage.dtype().nullability()),
176 masked_storage,
177 ))
178}
179
180fn mask_validity_variant(
181 array: VariantArray,
182 validity: Validity,
183 ctx: &mut ExecutionCtx,
184) -> VortexResult<VariantArray> {
185 let core_storage = array.core_storage().clone();
186 let len = core_storage.len();
187 let core_validity = core_storage.validity()?;
188 let shredded_validity = validity.clone();
189
190 let masked_core_storage = match core_validity {
191 Validity::NonNullable | Validity::AllValid => {
192 MaskedArray::try_new(core_storage, validity)?.into_array()
194 }
195 Validity::AllInvalid => {
196 core_storage
198 }
199 Validity::Array(_) => {
200 core_storage.mask(validity.to_array(len))?
203 }
204 };
205 let masked_shredded = if let Some(shredded) = array.shredded() {
206 let canonical = shredded.clone().execute::<Canonical>(ctx)?;
207 Some(mask_validity_canonical(canonical, shredded_validity, ctx)?.into_array())
208 } else {
209 None
210 };
211
212 VariantArray::try_new(masked_core_storage, masked_shredded)
213}