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::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::listview::ListViewArrayExt;
27use crate::arrays::struct_::StructArrayExt;
28use crate::arrays::union::UnionArrayExt;
29use crate::arrays::variant::VariantArrayExt;
30use crate::builtins::ArrayBuiltins;
31use crate::executor::ExecutionCtx;
32use crate::validity::Validity;
33
34/// TODO: replace usage of compute fn.
35/// Apply a validity mask to a canonical array, ANDing with existing validity.
36///
37/// This is the core operation for MaskedArray execution - it intersects the child's
38/// validity with the provided mask, marking additional positions as invalid.
39pub fn mask_validity_canonical(
40    canonical: Canonical,
41    validity: Validity,
42    ctx: &mut ExecutionCtx,
43) -> VortexResult<Canonical> {
44    Ok(match canonical {
45        n @ Canonical::Null(_) => n,
46        Canonical::Bool(a) => Canonical::Bool(mask_validity_bool(a, validity)?),
47        Canonical::Primitive(a) => Canonical::Primitive(mask_validity_primitive(a, validity)?),
48        Canonical::Decimal(a) => Canonical::Decimal(mask_validity_decimal(a, validity)?),
49        Canonical::VarBinView(a) => Canonical::VarBinView(mask_validity_varbinview(a, validity)?),
50        Canonical::List(a) => Canonical::List(mask_validity_listview(a, validity)?),
51        Canonical::FixedSizeList(a) => {
52            Canonical::FixedSizeList(mask_validity_fixed_size_list(a, validity)?)
53        }
54        Canonical::Struct(a) => Canonical::Struct(mask_validity_struct(a, validity)?),
55        Canonical::Union(a) => Canonical::Union(mask_validity_union(a, validity)?),
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    // SAFETY: We're only changing validity, not the data structure.
73    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    // SAFETY: We're only changing validity, not the data structure.
85    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
95/// Mask validity for VarBinViewArray.
96fn 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    // SAFETY: We're only changing validity, not the data structure.
103    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    // SAFETY: We're only changing validity, not the data structure.
116    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    // SAFETY: We're only changing validity, not the data structure.
136    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    // SAFETY: We're only changing validity, not the data structure.
147    Ok(unsafe { StructArray::new_unchecked(fields, struct_fields.clone(), len, new_validity) })
148}
149
150fn mask_validity_union(array: UnionArray, validity: Validity) -> VortexResult<UnionArray> {
151    let type_ids = array
152        .type_ids()
153        .clone()
154        .mask(validity.to_array(array.len()))?;
155    let variants = array.variants().clone();
156    let children = array.children();
157
158    // SAFETY: We're only changing validity, not the data structure.
159    Ok(unsafe { UnionArray::new_unchecked(type_ids, variants, children) })
160}
161
162fn mask_validity_extension(
163    array: ExtensionArray,
164    validity: Validity,
165    ctx: &mut ExecutionCtx,
166) -> VortexResult<ExtensionArray> {
167    // For extension arrays, we need to mask the underlying storage.
168    let storage = array.storage_array().clone().execute::<Canonical>(ctx)?;
169    let masked_storage = mask_validity_canonical(storage, validity, ctx)?;
170    let masked_storage = masked_storage.into_array();
171    Ok(ExtensionArray::new(
172        array
173            .ext_dtype()
174            .with_nullability(masked_storage.dtype().nullability()),
175        masked_storage,
176    ))
177}
178
179fn mask_validity_variant(
180    array: VariantArray,
181    validity: Validity,
182    ctx: &mut ExecutionCtx,
183) -> VortexResult<VariantArray> {
184    let core_storage = array.core_storage().clone();
185    let len = core_storage.len();
186    let core_validity = core_storage.validity()?;
187    let shredded_validity = validity.clone();
188
189    let masked_core_storage = match core_validity {
190        Validity::NonNullable | Validity::AllValid => {
191            // Core storage has no nulls, so wrap it in MaskedArray to apply the mask.
192            MaskedArray::try_new(core_storage, validity)?.into_array()
193        }
194        Validity::AllInvalid => {
195            // Already all-null, ANDing with any mask is still all-null.
196            core_storage
197        }
198        Validity::Array(_) => {
199            // Core storage already has nulls, but its physical validity layout depends on the
200            // actual encoding. Use the mask operation instead of rewriting a presumed slot.
201            core_storage.mask(validity.to_array(len))?
202        }
203    };
204    let masked_shredded = if let Some(shredded) = array.shredded() {
205        let canonical = shredded.clone().execute::<Canonical>(ctx)?;
206        Some(mask_validity_canonical(canonical, shredded_validity, ctx)?.into_array())
207    } else {
208        None
209    };
210
211    VariantArray::try_new(masked_core_storage, masked_shredded)
212}