Skip to main content

vortex_array/arrays/masked/
execute.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4//! Execution logic for MaskedArray - applies a validity mask to canonical arrays.
5
6use 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::ListView;
17use crate::arrays::ListViewArray;
18use crate::arrays::MapArray;
19use crate::arrays::MaskedArray;
20use crate::arrays::PrimitiveArray;
21use crate::arrays::StructArray;
22use crate::arrays::UnionArray;
23use crate::arrays::VarBinViewArray;
24use crate::arrays::VariantArray;
25use crate::arrays::bool::BoolArrayExt;
26use crate::arrays::extension::ExtensionArrayExt;
27use crate::arrays::fixed_size_list::FixedSizeListArrayExt;
28use crate::arrays::fixed_size_list::FixedSizeListArraySlotsExt;
29use crate::arrays::listview::ListViewArraySlotsExt;
30use crate::arrays::map::MapArrayExt;
31use crate::arrays::map::MapArraySlotsExt;
32use crate::arrays::struct_::StructArrayExt;
33use crate::arrays::union::UnionArrayExt;
34use crate::arrays::union::UnionArraySlotsExt;
35use crate::arrays::variant::VariantArraySlotsExt;
36use crate::builtins::ArrayBuiltins;
37use crate::executor::ExecutionCtx;
38use crate::validity::Validity;
39
40/// TODO: replace usage of compute fn.
41/// Apply a validity mask to a canonical array, ANDing with existing validity.
42///
43/// This is the core operation for MaskedArray execution - it intersects the child's
44/// validity with the provided mask, marking additional positions as invalid.
45pub fn mask_validity_canonical(
46    canonical: Canonical,
47    validity: Validity,
48    ctx: &mut ExecutionCtx,
49) -> VortexResult<Canonical> {
50    Ok(match canonical {
51        n @ Canonical::Null(_) => n,
52        Canonical::Bool(a) => Canonical::Bool(mask_validity_bool(a, validity)?),
53        Canonical::Primitive(a) => Canonical::Primitive(mask_validity_primitive(a, validity)?),
54        Canonical::Decimal(a) => Canonical::Decimal(mask_validity_decimal(a, validity)?),
55        Canonical::VarBinView(a) => Canonical::VarBinView(mask_validity_varbinview(a, validity)?),
56        Canonical::List(a) => Canonical::List(mask_validity_listview(a, validity)?),
57        Canonical::Map(a) => Canonical::Map(mask_validity_map(a, validity)?),
58        Canonical::FixedSizeList(a) => {
59            Canonical::FixedSizeList(mask_validity_fixed_size_list(a, validity)?)
60        }
61        Canonical::Struct(a) => Canonical::Struct(mask_validity_struct(a, validity)?),
62        Canonical::Union(a) => Canonical::Union(mask_validity_union(a, validity)?),
63        Canonical::Extension(a) => Canonical::Extension(mask_validity_extension(a, validity, ctx)?),
64        Canonical::Variant(a) => Canonical::Variant(mask_validity_variant(a, validity, ctx)?),
65    })
66}
67
68fn mask_validity_bool(array: BoolArray, mask: Validity) -> VortexResult<BoolArray> {
69    let new_validity = Validity::and(array.validity()?, mask)?;
70    Ok(BoolArray::new(array.to_bit_buffer(), new_validity))
71}
72
73fn mask_validity_primitive(
74    array: PrimitiveArray,
75    validity: Validity,
76) -> VortexResult<PrimitiveArray> {
77    let ptype = array.ptype();
78    let new_validity = Validity::and(array.validity()?, validity)?;
79    // SAFETY: We're only changing validity, not the data structure.
80    Ok(unsafe {
81        PrimitiveArray::new_unchecked_from_handle(
82            array.buffer_handle().clone(),
83            ptype,
84            new_validity,
85        )
86    })
87}
88
89fn mask_validity_decimal(array: DecimalArray, validity: Validity) -> VortexResult<DecimalArray> {
90    let new_validity = Validity::and(array.validity()?, validity)?;
91    // SAFETY: We're only changing validity, not the data structure.
92    Ok(unsafe {
93        DecimalArray::new_unchecked_handle(
94            array.buffer_handle().clone(),
95            array.values_type(),
96            array.decimal_dtype(),
97            new_validity,
98        )
99    })
100}
101
102/// Mask validity for VarBinViewArray.
103fn mask_validity_varbinview(
104    array: VarBinViewArray,
105    validity: Validity,
106) -> VortexResult<VarBinViewArray> {
107    let dtype = array.dtype().as_nullable();
108    let new_validity = Validity::and(array.validity()?, validity)?;
109    // SAFETY: We're only changing validity, not the data structure.
110    Ok(unsafe {
111        VarBinViewArray::new_handle_unchecked(
112            array.views_handle().clone(),
113            Arc::clone(array.data_buffers()),
114            dtype,
115            new_validity,
116        )
117    })
118}
119
120fn mask_validity_listview(array: ListViewArray, validity: Validity) -> VortexResult<ListViewArray> {
121    let new_validity = Validity::and(array.validity()?, validity)?;
122    // SAFETY: We're only changing validity, not the data structure.
123    let is_zctl = array.is_zero_copy_to_list();
124    Ok(unsafe {
125        ListViewArray::new_unchecked(
126            array.elements().clone(),
127            array.offsets().clone(),
128            array.sizes().clone(),
129            new_validity,
130        )
131        .with_zero_copy_to_list(is_zctl)
132    })
133}
134
135fn mask_validity_map(array: MapArray, validity: Validity) -> VortexResult<MapArray> {
136    let entries = mask_validity_listview(array.entries().as_::<ListView>().into_owned(), validity)?;
137    MapArray::try_new(array.map_dtype().clone(), entries)
138}
139
140fn mask_validity_fixed_size_list(
141    array: FixedSizeListArray,
142    validity: Validity,
143) -> VortexResult<FixedSizeListArray> {
144    let len = array.len();
145    let list_size = array.list_size();
146    let new_validity = Validity::and(array.validity()?, validity)?;
147    // SAFETY: We're only changing validity, not the data structure.
148    Ok(unsafe {
149        FixedSizeListArray::new_unchecked(array.elements().clone(), list_size, new_validity, len)
150    })
151}
152
153fn mask_validity_struct(array: StructArray, validity: Validity) -> VortexResult<StructArray> {
154    let len = array.len();
155    let new_validity = Validity::and(array.validity()?, validity)?;
156    // SAFETY: We're only changing validity, not the data structure.
157    Ok(unsafe {
158        StructArray::new_unchecked(
159            array.iter_unmasked_fields().cloned(),
160            array.struct_fields().clone(),
161            len,
162            new_validity,
163        )
164    })
165}
166
167fn mask_validity_union(array: UnionArray, validity: Validity) -> VortexResult<UnionArray> {
168    let type_ids = array
169        .type_ids()
170        .clone()
171        .mask(validity.to_array(array.len()))?;
172    let variants = array.variants().clone();
173
174    // SAFETY: We're only changing validity, not the data structure.
175    Ok(unsafe { UnionArray::new_unchecked(type_ids, variants, array.iter_children().cloned()) })
176}
177
178fn mask_validity_extension(
179    array: ExtensionArray,
180    validity: Validity,
181    ctx: &mut ExecutionCtx,
182) -> VortexResult<ExtensionArray> {
183    // For extension arrays, we need to mask the underlying storage.
184    let storage = array.storage_array().clone().execute::<Canonical>(ctx)?;
185    let masked_storage = mask_validity_canonical(storage, validity, ctx)?;
186    let masked_storage = masked_storage.into_array();
187    Ok(ExtensionArray::new(
188        array
189            .ext_dtype()
190            .with_nullability(masked_storage.dtype().nullability()),
191        masked_storage,
192    ))
193}
194
195fn mask_validity_variant(
196    array: VariantArray,
197    validity: Validity,
198    ctx: &mut ExecutionCtx,
199) -> VortexResult<VariantArray> {
200    let core_storage = array.core_storage().clone();
201    let len = core_storage.len();
202    let core_validity = core_storage.validity()?;
203    let shredded_validity = validity.clone();
204
205    let masked_core_storage = match core_validity {
206        Validity::NonNullable | Validity::AllValid => {
207            // Core storage has no nulls, so wrap it in MaskedArray to apply the mask.
208            MaskedArray::try_new(core_storage, validity)?.into_array()
209        }
210        Validity::AllInvalid => {
211            // Already all-null, ANDing with any mask is still all-null.
212            core_storage
213        }
214        Validity::Array(_) => {
215            // Core storage already has nulls, but its physical validity layout depends on the
216            // actual encoding. Use the mask operation instead of rewriting a presumed slot.
217            core_storage.mask(validity.to_array(len))?
218        }
219    };
220    let masked_shredded = if let Some(shredded) = array.shredded() {
221        let canonical = shredded.clone().execute::<Canonical>(ctx)?;
222        Some(mask_validity_canonical(canonical, shredded_validity, ctx)?.into_array())
223    } else {
224        None
225    };
226
227    VariantArray::try_new(masked_core_storage, masked_shredded)
228}