Skip to main content

vortex_array/arrays/varbinview/compute/
take.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use std::iter;
5use std::ptr;
6use std::sync::Arc;
7
8use itertools::Itertools as _;
9use num_traits::AsPrimitive;
10use vortex_buffer::Buffer;
11use vortex_buffer::BufferMut;
12use vortex_error::VortexResult;
13use vortex_error::vortex_ensure;
14use vortex_error::vortex_err;
15use vortex_mask::AllOr;
16use vortex_mask::Mask;
17
18use crate::ArrayRef;
19use crate::Columnar;
20use crate::IntoArray;
21use crate::array::ArrayView;
22use crate::arrays::PiecewiseSequence;
23use crate::arrays::PrimitiveArray;
24use crate::arrays::VarBinView;
25use crate::arrays::VarBinViewArray;
26use crate::arrays::dict::TakeExecute;
27use crate::arrays::piecewise_sequence::constant_unsigned_usize;
28use crate::arrays::piecewise_sequence::maybe_contiguous_slices;
29use crate::arrays::varbinview::BinaryView;
30use crate::buffer::BufferHandle;
31use crate::dtype::UnsignedPType;
32use crate::executor::ExecutionCtx;
33use crate::match_each_integer_ptype;
34use crate::match_each_unsigned_integer_ptype;
35
36impl TakeExecute for VarBinView {
37    /// Take involves creating a new array that references the old array, just with the given set of views.
38    fn take(
39        array: ArrayView<'_, VarBinView>,
40        indices: &ArrayRef,
41        ctx: &mut ExecutionCtx,
42    ) -> VortexResult<Option<ArrayRef>> {
43        if let Some(piecewise_indices) = indices.as_opt::<PiecewiseSequence>()
44            && let Some(taken) = take_contiguous_ranges(array, piecewise_indices, indices, ctx)?
45        {
46            return Ok(Some(taken));
47        }
48
49        let validity = array.validity()?.take(indices)?;
50        let indices = indices.clone().execute::<PrimitiveArray>(ctx)?;
51
52        let indices_mask = indices
53            .as_ref()
54            .validity()?
55            .execute_mask(indices.as_ref().len(), ctx)?;
56        let views_buffer = match_each_integer_ptype!(indices.ptype(), |I| {
57            take_views(array.views(), indices.as_slice::<I>(), &indices_mask)
58        });
59
60        // SAFETY: taking all components at same indices maintains invariants
61        unsafe {
62            Ok(Some(
63                VarBinViewArray::new_handle_unchecked(
64                    BufferHandle::new_host(views_buffer.into_byte_buffer()),
65                    Arc::clone(array.data_buffers()),
66                    array
67                        .dtype()
68                        .union_nullability(indices.dtype().nullability()),
69                    validity,
70                )
71                .into_array(),
72            ))
73        }
74    }
75}
76
77fn take_contiguous_ranges(
78    array: ArrayView<'_, VarBinView>,
79    indices: ArrayView<'_, PiecewiseSequence>,
80    indices_ref: &ArrayRef,
81    ctx: &mut ExecutionCtx,
82) -> VortexResult<Option<ArrayRef>> {
83    let Some((starts, lengths)) = maybe_contiguous_slices(indices, ctx)? else {
84        return Ok(None);
85    };
86    let source = array.views();
87    let output_len = indices_ref.len();
88    let views = match lengths {
89        Columnar::Constant(lengths) => {
90            let length = constant_unsigned_usize(&lengths);
91            match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
92                gather_view_slices_constant_length(
93                    source,
94                    starts.as_slice::<S>(),
95                    length,
96                    output_len,
97                )?
98            })
99        }
100        Columnar::Canonical(lengths) => {
101            let lengths = lengths.into_primitive();
102            match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
103                match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
104                    gather_view_slices(
105                        source,
106                        starts.as_slice::<S>(),
107                        lengths.as_slice::<L>(),
108                        output_len,
109                    )?
110                })
111            })
112        }
113    };
114    let validity = array.validity()?.take(indices_ref)?;
115
116    // SAFETY: ranges were validated against the source views, and copied views still reference the
117    // same backing data buffers.
118    unsafe {
119        Ok(Some(
120            VarBinViewArray::new_handle_unchecked(
121                BufferHandle::new_host(views.into_byte_buffer()),
122                Arc::clone(array.data_buffers()),
123                array.dtype().clone(),
124                validity,
125            )
126            .into_array(),
127        ))
128    }
129}
130
131fn take_views<I: AsPrimitive<usize>>(
132    views_ref: &[BinaryView],
133    indices: &[I],
134    mask: &Mask,
135) -> Buffer<BinaryView> {
136    // NOTE(ngates): this deref is not actually trivial, so we run it once.
137    // We do not use iter_bools directly, since the resulting dyn iterator cannot
138    // implement TrustedLen.
139    match mask.bit_buffer() {
140        AllOr::All => {
141            Buffer::<BinaryView>::from_trusted_len_iter(indices.iter().map(|i| views_ref[i.as_()]))
142        }
143        AllOr::None => Buffer::<BinaryView>::from_trusted_len_iter(iter::repeat_n(
144            BinaryView::default(),
145            indices.len(),
146        )),
147        AllOr::Some(buffer) => Buffer::<BinaryView>::from_trusted_len_iter(
148            buffer.iter().zip(indices.iter()).map(|(valid, idx)| {
149                if valid {
150                    views_ref[idx.as_()]
151                } else {
152                    BinaryView::default()
153                }
154            }),
155        ),
156    }
157}
158
159fn gather_view_slices_constant_length<S>(
160    source: &[BinaryView],
161    starts: &[S],
162    length: usize,
163    output_len: usize,
164) -> VortexResult<Buffer<BinaryView>>
165where
166    S: UnsignedPType,
167{
168    let computed_len = starts
169        .len()
170        .checked_mul(length)
171        .ok_or_else(|| vortex_err!("PiecewiseSequenceArray output length overflows usize"))?;
172    vortex_ensure!(
173        computed_len == output_len,
174        "PiecewiseSequenceArray expanded length {computed_len} does not match declared length {output_len}"
175    );
176
177    let mut views = BufferMut::<BinaryView>::with_capacity(output_len);
178    let spare = &mut views.spare_capacity_mut()[..output_len];
179    let mut cursor = 0usize;
180    for &start in starts {
181        let start = start.as_();
182        let src = &source[start..][..length];
183        // SAFETY: `src` and the checked `spare` range have equal lengths and cannot overlap.
184        unsafe {
185            ptr::copy_nonoverlapping(
186                src.as_ptr(),
187                spare[cursor..][..src.len()]
188                    .as_mut_ptr()
189                    .cast::<BinaryView>(),
190                src.len(),
191            );
192        }
193        cursor += src.len();
194    }
195    // SAFETY: the loop initialized the prefix `0..cursor` of the spare capacity.
196    unsafe { views.set_len(cursor) };
197    vortex_ensure!(
198        views.len() == output_len,
199        "PiecewiseSequenceArray expanded length {} does not match declared length {output_len}",
200        views.len()
201    );
202    Ok(views.freeze())
203}
204
205fn gather_view_slices<S, L>(
206    source: &[BinaryView],
207    starts: &[S],
208    lengths: &[L],
209    output_len: usize,
210) -> VortexResult<Buffer<BinaryView>>
211where
212    S: UnsignedPType,
213    L: UnsignedPType,
214{
215    let mut views = BufferMut::<BinaryView>::with_capacity(output_len);
216    let spare = &mut views.spare_capacity_mut()[..output_len];
217    let mut cursor = 0usize;
218    for (&start, &length) in starts.iter().zip_eq(lengths) {
219        let start = start.as_();
220        let length = length.as_();
221        let src = &source[start..][..length];
222        // SAFETY: `src` and the checked `spare` range have equal lengths and cannot overlap.
223        unsafe {
224            ptr::copy_nonoverlapping(
225                src.as_ptr(),
226                spare[cursor..][..src.len()]
227                    .as_mut_ptr()
228                    .cast::<BinaryView>(),
229                src.len(),
230            );
231        }
232        cursor += src.len();
233    }
234    // SAFETY: the loop initialized the prefix `0..cursor` of the spare capacity.
235    unsafe { views.set_len(cursor) };
236    vortex_ensure!(
237        views.len() == output_len,
238        "PiecewiseSequenceArray expanded length {} does not match declared length {output_len}",
239        views.len()
240    );
241    Ok(views.freeze())
242}
243
244#[cfg(test)]
245mod tests {
246    use rstest::rstest;
247    use vortex_buffer::BitBuffer;
248    use vortex_buffer::buffer;
249    use vortex_error::VortexResult;
250
251    use crate::IntoArray;
252    use crate::VortexSessionExecute;
253    use crate::array_session;
254    use crate::arrays::VarBinViewArray;
255    use crate::arrays::varbinview::compute::take::PrimitiveArray;
256    use crate::compute::conformance::take::test_take_conformance;
257    use crate::dtype::DType;
258    use crate::dtype::Nullability::NonNullable;
259    use crate::validity::Validity;
260
261    #[test]
262    fn take_nullable() -> VortexResult<()> {
263        let arr = VarBinViewArray::from_iter_nullable_str([
264            Some("one"),
265            None,
266            Some("three"),
267            Some("four"),
268            None,
269            Some("six"),
270        ]);
271
272        let taken = arr.take(buffer![0, 3].into_array())?;
273
274        assert!(taken.dtype().is_nullable());
275        let mut ctx = array_session().create_execution_ctx();
276        let taken = taken.execute::<VarBinViewArray>(&mut ctx)?;
277        let mask = taken.validity()?.execute_mask(taken.len(), &mut ctx)?;
278        let result = (0..taken.len())
279            .map(|i| {
280                mask.value(i)
281                    .then(|| unsafe { String::from_utf8_unchecked(taken.bytes_at(i).to_vec()) })
282            })
283            .collect::<Vec<_>>();
284        assert_eq!(result, [Some("one".to_string()), Some("four".to_string())]);
285        Ok(())
286    }
287
288    #[test]
289    fn take_nullable_indices() -> VortexResult<()> {
290        let arr = VarBinViewArray::from_iter(["one", "two"].map(Some), DType::Utf8(NonNullable));
291
292        let indices = PrimitiveArray::new(
293            // Verify that garbage values at NULL indices are ignored.
294            buffer![1u64, 999],
295            Validity::from(BitBuffer::from(vec![true, false])),
296        );
297
298        let taken = arr.take(indices.into_array())?;
299
300        assert!(taken.dtype().is_nullable());
301        let mut ctx = array_session().create_execution_ctx();
302        let taken = taken.execute::<VarBinViewArray>(&mut ctx)?;
303        let mask = taken.validity()?.execute_mask(taken.len(), &mut ctx)?;
304        let result = (0..taken.len())
305            .map(|i| {
306                mask.value(i)
307                    .then(|| unsafe { String::from_utf8_unchecked(taken.bytes_at(i).to_vec()) })
308            })
309            .collect::<Vec<_>>();
310        assert_eq!(result, [Some("two".to_string()), None]);
311        Ok(())
312    }
313
314    #[rstest]
315    #[case(VarBinViewArray::from_iter(
316        ["hello", "world", "test", "data", "array"].map(Some),
317        DType::Utf8(NonNullable),
318    ))]
319    #[case(VarBinViewArray::from_iter_nullable_str([
320        Some("hello"),
321        None,
322        Some("test"),
323        Some("data"),
324        None,
325    ]))]
326    #[case(VarBinViewArray::from_iter(
327        [b"hello".as_slice(), b"world", b"test", b"data", b"array"].map(Some),
328        DType::Binary(NonNullable),
329    ))]
330    #[case(VarBinViewArray::from_iter(["single"].map(Some), DType::Utf8(NonNullable)))]
331    fn test_take_varbinview_conformance(#[case] array: VarBinViewArray) {
332        test_take_conformance(
333            &array.into_array(),
334            &mut array_session().create_execution_ctx(),
335        );
336    }
337}