Skip to main content

vortex_array/arrays/list/compute/
take.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use itertools::Itertools as _;
5use vortex_buffer::BufferMut;
6use vortex_error::VortexExpect;
7use vortex_error::VortexResult;
8use vortex_error::vortex_ensure;
9use vortex_error::vortex_err;
10use vortex_mask::Mask;
11
12use crate::ArrayRef;
13use crate::Columnar;
14use crate::IntoArray;
15use crate::array::ArrayView;
16use crate::arrays::ConstantArray;
17use crate::arrays::List;
18use crate::arrays::ListArray;
19use crate::arrays::PiecewiseSequence;
20use crate::arrays::PiecewiseSequenceArray;
21use crate::arrays::Primitive;
22use crate::arrays::PrimitiveArray;
23use crate::arrays::dict::TakeExecute;
24use crate::arrays::list::ListArrayExt;
25use crate::arrays::piecewise_sequence::constant_unsigned_usize;
26use crate::arrays::piecewise_sequence::maybe_contiguous_slices;
27use crate::arrays::primitive::PrimitiveArrayExt;
28use crate::dtype::IntegerPType;
29use crate::dtype::UnsignedPType;
30use crate::executor::ExecutionCtx;
31use crate::match_each_unsigned_integer_ptype;
32use crate::match_smallest_offset_type;
33use crate::validity::Validity;
34
35// TODO(connor)[ListView]: Re-revert to the version where we simply convert to a `ListView` and call
36// the `ListView::take` compute function once `ListView` is more stable.
37
38impl TakeExecute for List {
39    /// Take implementation for [`ListArray`].
40    ///
41    /// Unlike `ListView`, `ListArray` must rebuild the elements array to maintain its invariant
42    /// that lists are stored contiguously and in-order (`offset[i+1] >= offset[i]`). Taking
43    /// non-contiguous indices would violate this requirement.
44    #[expect(clippy::cognitive_complexity)]
45    fn take(
46        array: ArrayView<'_, List>,
47        indices: &ArrayRef,
48        ctx: &mut ExecutionCtx,
49    ) -> VortexResult<Option<ArrayRef>> {
50        if let Some(piecewise_indices) = indices.as_opt::<PiecewiseSequence>()
51            && let Some(taken) = take_slices(array, piecewise_indices, indices, ctx)?
52        {
53            return Ok(Some(taken));
54        }
55
56        let new_validity = array.validity()?.take(indices)?;
57        let indices = indices.clone().execute::<PrimitiveArray>(ctx)?;
58        let indices = indices.reinterpret_cast(indices.ptype().to_unsigned());
59        let offsets = array.offsets().clone().execute::<PrimitiveArray>(ctx)?;
60        let offsets = offsets.reinterpret_cast(offsets.ptype().to_unsigned());
61        let validity_mask = new_validity.execute_mask(indices.len(), ctx)?;
62        // This is an over-approximation of the total number of elements in the resulting array.
63        let total_approx = array.elements().len().saturating_mul(indices.len());
64
65        match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
66            match_each_unsigned_integer_ptype!(indices.ptype(), |I| {
67                match_smallest_offset_type!(total_approx, |OutOffset| {
68                    take_with_piecewise_elements::<I, O, OutOffset>(
69                        array,
70                        offsets.as_view(),
71                        indices.as_view(),
72                        new_validity,
73                        &validity_mask,
74                    )
75                    .map(Some)
76                })
77            })
78        })
79    }
80}
81
82fn take_with_piecewise_elements<I: IntegerPType, O: IntegerPType, OutOffset: IntegerPType>(
83    array: ArrayView<'_, List>,
84    offsets_array: ArrayView<'_, Primitive>,
85    indices_array: ArrayView<'_, Primitive>,
86    new_validity: Validity,
87    validity_mask: &Mask,
88) -> VortexResult<ArrayRef> {
89    let offsets: &[O] = offsets_array.as_slice();
90    let indices: &[I] = indices_array.as_slice();
91
92    let offsets_capacity = indices
93        .len()
94        .checked_add(1)
95        .ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
96    let mut new_offsets = BufferMut::<OutOffset>::with_capacity(offsets_capacity);
97    let mut element_starts = BufferMut::<u64>::with_capacity(indices.len());
98    let mut element_lengths = BufferMut::<u64>::with_capacity(indices.len());
99
100    let mut current_offset = 0usize;
101    new_offsets.push(OutOffset::zero());
102
103    for (&data_idx, is_valid) in indices.iter().zip_eq(validity_mask.iter()) {
104        if !is_valid {
105            new_offsets.push(new_offset_value::<OutOffset>(current_offset));
106            element_starts.push(0);
107            element_lengths.push(0);
108            continue;
109        }
110
111        let data_idx: usize = data_idx.as_();
112
113        let start = offsets[data_idx];
114        let stop = offsets[data_idx + 1];
115        let start: usize = start.as_();
116        let stop: usize = stop.as_();
117        let length = stop - start;
118
119        current_offset = current_offset
120            .checked_add(length)
121            .ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
122        new_offsets.push(new_offset_value::<OutOffset>(current_offset));
123        element_starts.push(start as u64);
124        element_lengths.push(length as u64);
125    }
126
127    let new_offsets = PrimitiveArray::new(new_offsets.freeze(), Validity::NonNullable).into_array();
128    let multipliers = ConstantArray::new(1u64, element_starts.len()).into_array();
129
130    // SAFETY: valid source rows contribute ranges derived from list offsets; null index/source
131    // rows contribute zero-length placeholder ranges. `current_offset` is the sum of all generated
132    // element lengths, and multiplier 1 preserves contiguous ranges.
133    let element_indices = unsafe {
134        PiecewiseSequenceArray::new_unchecked(
135            element_starts.into_array(),
136            element_lengths.into_array(),
137            multipliers,
138            current_offset,
139        )
140    };
141    let new_elements = array.elements().take(element_indices.into_array())?;
142
143    // SAFETY: offsets are rebuilt from the gathered element ranges and have one entry per output
144    // row plus the leading zero; validity is produced by the usual take-validity path.
145    Ok(unsafe { ListArray::new_unchecked(new_elements, new_offsets, new_validity) }.into_array())
146}
147
148fn take_slices(
149    array: ArrayView<'_, List>,
150    indices: ArrayView<'_, PiecewiseSequence>,
151    indices_ref: &ArrayRef,
152    ctx: &mut ExecutionCtx,
153) -> VortexResult<Option<ArrayRef>> {
154    let Some((starts, lengths)) = maybe_contiguous_slices(indices, ctx)? else {
155        return Ok(None);
156    };
157    let data_validity = array
158        .list_validity()
159        .execute_mask(array.as_ref().len(), ctx)?;
160    let offsets = array.offsets().clone().execute::<PrimitiveArray>(ctx)?;
161    let offsets = offsets.reinterpret_cast(offsets.ptype().to_unsigned());
162    let output_len = indices_ref.len();
163
164    let taken = match lengths {
165        Columnar::Constant(lengths) => {
166            let length = constant_unsigned_usize(&lengths);
167            take_slices_constant_start_dispatch(
168                array,
169                &starts,
170                length,
171                &offsets,
172                indices_ref,
173                output_len,
174                &data_validity,
175            )?
176        }
177        Columnar::Canonical(lengths) => {
178            let lengths = lengths.into_primitive();
179            take_slices_start_dispatch(
180                array,
181                &starts,
182                &lengths,
183                &offsets,
184                indices_ref,
185                output_len,
186                &data_validity,
187            )?
188        }
189    };
190    Ok(Some(taken))
191}
192
193fn take_slices_constant_start_dispatch(
194    array: ArrayView<'_, List>,
195    starts: &PrimitiveArray,
196    length: usize,
197    offsets: &PrimitiveArray,
198    indices_ref: &ArrayRef,
199    output_len: usize,
200    data_validity: &Mask,
201) -> VortexResult<ArrayRef> {
202    match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
203        take_slices_constant_offset_dispatch::<S>(
204            array,
205            starts,
206            length,
207            offsets,
208            indices_ref,
209            output_len,
210            data_validity,
211        )
212    })
213}
214
215fn take_slices_constant_offset_dispatch<S>(
216    array: ArrayView<'_, List>,
217    starts: &PrimitiveArray,
218    length: usize,
219    offsets: &PrimitiveArray,
220    indices_ref: &ArrayRef,
221    output_len: usize,
222    data_validity: &Mask,
223) -> VortexResult<ArrayRef>
224where
225    S: UnsignedPType,
226{
227    match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
228        take_slices_constant_length::<S, O>(
229            array,
230            starts.as_slice::<S>(),
231            length,
232            offsets.as_slice::<O>(),
233            indices_ref,
234            output_len,
235            data_validity,
236        )
237    })
238}
239
240fn take_slices_start_dispatch(
241    array: ArrayView<'_, List>,
242    starts: &PrimitiveArray,
243    lengths: &PrimitiveArray,
244    offsets: &PrimitiveArray,
245    indices_ref: &ArrayRef,
246    output_len: usize,
247    data_validity: &Mask,
248) -> VortexResult<ArrayRef> {
249    match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
250        take_slices_length_dispatch::<S>(
251            array,
252            starts,
253            lengths,
254            offsets,
255            indices_ref,
256            output_len,
257            data_validity,
258        )
259    })
260}
261
262fn take_slices_length_dispatch<S>(
263    array: ArrayView<'_, List>,
264    starts: &PrimitiveArray,
265    lengths: &PrimitiveArray,
266    offsets: &PrimitiveArray,
267    indices_ref: &ArrayRef,
268    output_len: usize,
269    data_validity: &Mask,
270) -> VortexResult<ArrayRef>
271where
272    S: UnsignedPType,
273{
274    match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
275        take_slices_offset_dispatch::<S, L>(
276            array,
277            starts,
278            lengths,
279            offsets,
280            indices_ref,
281            output_len,
282            data_validity,
283        )
284    })
285}
286
287fn take_slices_offset_dispatch<S, L>(
288    array: ArrayView<'_, List>,
289    starts: &PrimitiveArray,
290    lengths: &PrimitiveArray,
291    offsets: &PrimitiveArray,
292    indices_ref: &ArrayRef,
293    output_len: usize,
294    data_validity: &Mask,
295) -> VortexResult<ArrayRef>
296where
297    S: UnsignedPType,
298    L: UnsignedPType,
299{
300    match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
301        take_slices_typed::<S, L, O>(
302            array,
303            starts.as_slice::<S>(),
304            lengths.as_slice::<L>(),
305            offsets.as_slice::<O>(),
306            indices_ref,
307            output_len,
308            data_validity,
309        )
310    })
311}
312
313fn take_slices_constant_length<S, Offset>(
314    array: ArrayView<'_, List>,
315    starts: &[S],
316    length: usize,
317    offsets: &[Offset],
318    indices_ref: &ArrayRef,
319    output_len: usize,
320    data_validity: &Mask,
321) -> VortexResult<ArrayRef>
322where
323    S: UnsignedPType,
324    Offset: UnsignedPType,
325{
326    let computed_len = starts
327        .len()
328        .checked_mul(length)
329        .ok_or_else(|| vortex_err!("PiecewiseSequenceArray output length overflows usize"))?;
330    vortex_ensure!(
331        computed_len == output_len,
332        "PiecewiseSequenceArray expanded length {computed_len} does not match declared length {output_len}"
333    );
334    let all_valid = data_validity.all_true();
335    let total_elements = if all_valid {
336        piecewise_list_elements_len_constant(offsets, starts, length)?
337    } else {
338        piecewise_list_elements_len_constant_validity(offsets, starts, length, data_validity)?
339    };
340    let validity = array.validity()?.take(indices_ref)?;
341
342    match_smallest_offset_type!(total_elements, |OutOffset| {
343        let gathered = if all_valid {
344            gather_piecewise_list_constant_length::<S, Offset, OutOffset>(
345                array.elements(),
346                offsets,
347                starts,
348                length,
349                output_len,
350                total_elements,
351            )?
352        } else {
353            gather_piecewise_list_constant_length_validity::<S, Offset, OutOffset>(
354                array.elements(),
355                offsets,
356                starts,
357                length,
358                output_len,
359                total_elements,
360                data_validity,
361            )?
362        };
363
364        // SAFETY: output offsets are rebuilt from valid monotonic source offsets; output elements
365        // are exactly the gathered child ranges referenced by those offsets; validity has one bit
366        // per output row.
367        Ok(
368            unsafe { ListArray::new_unchecked(gathered.elements, gathered.offsets, validity) }
369                .into_array(),
370        )
371    })
372}
373
374fn take_slices_typed<S, L, Offset>(
375    array: ArrayView<'_, List>,
376    starts: &[S],
377    lengths: &[L],
378    offsets: &[Offset],
379    indices_ref: &ArrayRef,
380    output_len: usize,
381    data_validity: &Mask,
382) -> VortexResult<ArrayRef>
383where
384    S: UnsignedPType,
385    L: UnsignedPType,
386    Offset: UnsignedPType,
387{
388    let mut computed_len = 0usize;
389    for &length in lengths {
390        let length: usize = length.as_();
391        computed_len = computed_len
392            .checked_add(length)
393            .ok_or_else(|| vortex_err!("PiecewiseSequenceArray output length overflows usize"))?;
394    }
395    vortex_ensure!(
396        computed_len == output_len,
397        "PiecewiseSequenceArray expanded length {computed_len} does not match declared length {output_len}"
398    );
399    let all_valid = data_validity.all_true();
400    let total_elements = if all_valid {
401        piecewise_list_elements_len(offsets, starts, lengths)?
402    } else {
403        piecewise_list_elements_len_validity(offsets, starts, lengths, data_validity)?
404    };
405
406    match_smallest_offset_type!(total_elements, |OutOffset| {
407        let gathered = if all_valid {
408            gather_piecewise_list::<S, L, Offset, OutOffset>(
409                array.elements(),
410                offsets,
411                starts,
412                lengths,
413                output_len,
414                total_elements,
415            )?
416        } else {
417            gather_piecewise_list_validity::<S, L, Offset, OutOffset>(
418                array.elements(),
419                offsets,
420                starts,
421                lengths,
422                output_len,
423                total_elements,
424                data_validity,
425            )?
426        };
427        let validity = array.validity()?.take(indices_ref)?;
428
429        // SAFETY: output offsets are rebuilt from valid monotonic source offsets; output elements
430        // are exactly the gathered child ranges referenced by those offsets; validity has one bit
431        // per output row.
432        Ok(
433            unsafe { ListArray::new_unchecked(gathered.elements, gathered.offsets, validity) }
434                .into_array(),
435        )
436    })
437}
438
439struct GatheredList {
440    elements: ArrayRef,
441    offsets: ArrayRef,
442}
443
444struct ValidPieceGather<OutOffset> {
445    new_offsets: BufferMut<OutOffset>,
446    element_starts: BufferMut<u64>,
447    element_lengths: BufferMut<u64>,
448    output_elements: usize,
449}
450
451fn piecewise_list_elements_len_constant<S, Offset>(
452    offsets: &[Offset],
453    starts: &[S],
454    length: usize,
455) -> VortexResult<usize>
456where
457    S: UnsignedPType,
458    Offset: UnsignedPType,
459{
460    if length == 0 {
461        return Ok(0);
462    }
463
464    let mut total = 0usize;
465    for start in starts {
466        let start: usize = start.as_();
467        let offset_range = &offsets[start..][..=length];
468        let element_start: usize = offset_range[0].as_();
469        let element_end: usize = offset_range[length].as_();
470        total = total
471            .checked_add(element_end - element_start)
472            .ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
473    }
474    Ok(total)
475}
476
477fn piecewise_list_elements_len_constant_validity<S, Offset>(
478    offsets: &[Offset],
479    starts: &[S],
480    length: usize,
481    data_validity: &Mask,
482) -> VortexResult<usize>
483where
484    S: UnsignedPType,
485    Offset: UnsignedPType,
486{
487    if length == 0 {
488        return Ok(0);
489    }
490
491    let mut total = 0usize;
492    for start in starts {
493        let start: usize = start.as_();
494        let additional = valid_piece_elements_len(offsets, data_validity, start, length)?;
495        total = total
496            .checked_add(additional)
497            .ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
498    }
499    Ok(total)
500}
501
502fn piecewise_list_elements_len<S, L, Offset>(
503    offsets: &[Offset],
504    starts: &[S],
505    lengths: &[L],
506) -> VortexResult<usize>
507where
508    S: UnsignedPType,
509    L: UnsignedPType,
510    Offset: UnsignedPType,
511{
512    let mut total = 0usize;
513    for (&start, &length) in starts.iter().zip_eq(lengths) {
514        let start: usize = start.as_();
515        let length: usize = length.as_();
516        let offset_range = &offsets[start..][..=length];
517        let element_start: usize = offset_range[0].as_();
518        let element_end: usize = offset_range[length].as_();
519        total = total
520            .checked_add(element_end - element_start)
521            .ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
522    }
523    Ok(total)
524}
525
526fn piecewise_list_elements_len_validity<S, L, Offset>(
527    offsets: &[Offset],
528    starts: &[S],
529    lengths: &[L],
530    data_validity: &Mask,
531) -> VortexResult<usize>
532where
533    S: UnsignedPType,
534    L: UnsignedPType,
535    Offset: UnsignedPType,
536{
537    let mut total = 0usize;
538    for (&start, &length) in starts.iter().zip_eq(lengths) {
539        let start: usize = start.as_();
540        let length: usize = length.as_();
541        let additional = valid_piece_elements_len(offsets, data_validity, start, length)?;
542        total = total
543            .checked_add(additional)
544            .ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
545    }
546    Ok(total)
547}
548
549fn valid_piece_elements_len<Offset>(
550    offsets: &[Offset],
551    data_validity: &Mask,
552    start: usize,
553    length: usize,
554) -> VortexResult<usize>
555where
556    Offset: UnsignedPType,
557{
558    let offset_range = &offsets[start..][..=length];
559    let mut total = 0usize;
560    for (data_idx, window) in (start..).zip(offset_range.windows(2)) {
561        if !data_validity.value(data_idx) {
562            continue;
563        }
564        let element_start: usize = window[0].as_();
565        let element_end: usize = window[1].as_();
566        total = total
567            .checked_add(element_end - element_start)
568            .ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
569    }
570    Ok(total)
571}
572
573fn gather_piecewise_list_constant_length<S, Offset, OutOffset>(
574    elements: &ArrayRef,
575    offsets: &[Offset],
576    starts: &[S],
577    length: usize,
578    output_len: usize,
579    total_elements: usize,
580) -> VortexResult<GatheredList>
581where
582    S: UnsignedPType,
583    Offset: UnsignedPType,
584    OutOffset: IntegerPType,
585{
586    let offsets_capacity = output_len
587        .checked_add(1)
588        .ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
589    let mut new_offsets = BufferMut::<OutOffset>::with_capacity(offsets_capacity);
590    let mut element_starts = BufferMut::<u64>::with_capacity(starts.len());
591    let mut element_lengths = BufferMut::<u64>::with_capacity(starts.len());
592    let mut output_elements = 0usize;
593
594    new_offsets.push(OutOffset::zero());
595    for start in starts {
596        let start: usize = start.as_();
597        if length == 0 {
598            continue;
599        }
600
601        let offset_range = &offsets[start..][..=length];
602        let element_start: usize = offset_range[0].as_();
603        let element_end: usize = offset_range[length].as_();
604        for &offset in &offset_range[1..] {
605            let offset: usize = offset.as_();
606            let relative = offset - element_start;
607            let output_offset = output_elements + relative;
608            new_offsets.push(new_offset_value::<OutOffset>(output_offset));
609        }
610
611        let element_length = element_end - element_start;
612        element_starts.push(element_start as u64);
613        element_lengths.push(element_length as u64);
614        output_elements += element_length;
615    }
616    debug_assert_eq!(output_elements, total_elements);
617
618    let offsets = PrimitiveArray::new(new_offsets.freeze(), Validity::NonNullable).into_array();
619    let multipliers = ConstantArray::new(1u64, element_starts.len()).into_array();
620    // SAFETY: element ranges are derived from validated source list offsets, and total_elements is
621    // the sum of the gathered element range lengths. Multiplier 1 preserves contiguous ranges.
622    let element_indices = unsafe {
623        PiecewiseSequenceArray::new_unchecked(
624            element_starts.into_array(),
625            element_lengths.into_array(),
626            multipliers,
627            total_elements,
628        )
629    };
630    let elements = elements.take(element_indices.into_array())?;
631
632    Ok(GatheredList { elements, offsets })
633}
634
635fn gather_piecewise_list_constant_length_validity<S, Offset, OutOffset>(
636    elements: &ArrayRef,
637    offsets: &[Offset],
638    starts: &[S],
639    length: usize,
640    output_len: usize,
641    total_elements: usize,
642    data_validity: &Mask,
643) -> VortexResult<GatheredList>
644where
645    S: UnsignedPType,
646    Offset: UnsignedPType,
647    OutOffset: IntegerPType,
648{
649    let offsets_capacity = output_len
650        .checked_add(1)
651        .ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
652    let mut gather = ValidPieceGather {
653        new_offsets: BufferMut::<OutOffset>::with_capacity(offsets_capacity),
654        element_starts: BufferMut::<u64>::with_capacity(output_len),
655        element_lengths: BufferMut::<u64>::with_capacity(output_len),
656        output_elements: 0,
657    };
658
659    gather.new_offsets.push(OutOffset::zero());
660    for start in starts {
661        let start: usize = start.as_();
662        if length == 0 {
663            continue;
664        }
665
666        gather_valid_piece(offsets, data_validity, start, length, &mut gather);
667    }
668    debug_assert_eq!(gather.output_elements, total_elements);
669
670    let offsets =
671        PrimitiveArray::new(gather.new_offsets.freeze(), Validity::NonNullable).into_array();
672    let multipliers = ConstantArray::new(1u64, gather.element_starts.len()).into_array();
673    // SAFETY: element ranges come only from valid source list rows. Source list construction
674    // validated those row offsets, and null source rows produce no element range.
675    let element_indices = unsafe {
676        PiecewiseSequenceArray::new_unchecked(
677            gather.element_starts.into_array(),
678            gather.element_lengths.into_array(),
679            multipliers,
680            total_elements,
681        )
682    };
683    let elements = elements.take(element_indices.into_array())?;
684
685    Ok(GatheredList { elements, offsets })
686}
687
688fn gather_piecewise_list<S, L, Offset, OutOffset>(
689    elements: &ArrayRef,
690    offsets: &[Offset],
691    starts: &[S],
692    lengths: &[L],
693    output_len: usize,
694    total_elements: usize,
695) -> VortexResult<GatheredList>
696where
697    S: UnsignedPType,
698    L: UnsignedPType,
699    Offset: UnsignedPType,
700    OutOffset: IntegerPType,
701{
702    let offsets_capacity = output_len
703        .checked_add(1)
704        .ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
705    let mut new_offsets = BufferMut::<OutOffset>::with_capacity(offsets_capacity);
706    let mut element_starts = BufferMut::<u64>::with_capacity(starts.len());
707    let mut element_lengths = BufferMut::<u64>::with_capacity(lengths.len());
708    let mut output_elements = 0usize;
709
710    new_offsets.push(OutOffset::zero());
711    for (&start, &length) in starts.iter().zip_eq(lengths) {
712        let start: usize = start.as_();
713        let length: usize = length.as_();
714        if length == 0 {
715            continue;
716        }
717
718        let offset_range = &offsets[start..][..=length];
719        let element_start: usize = offset_range[0].as_();
720        let element_end: usize = offset_range[length].as_();
721        for &offset in &offset_range[1..] {
722            let offset: usize = offset.as_();
723            let relative = offset - element_start;
724            let output_offset = output_elements + relative;
725            new_offsets.push(new_offset_value::<OutOffset>(output_offset));
726        }
727
728        let element_length = element_end - element_start;
729        element_starts.push(element_start as u64);
730        element_lengths.push(element_length as u64);
731        output_elements += element_length;
732    }
733    debug_assert_eq!(output_elements, total_elements);
734
735    let offsets = PrimitiveArray::new(new_offsets.freeze(), Validity::NonNullable).into_array();
736    let multipliers = ConstantArray::new(1u64, element_starts.len()).into_array();
737    // SAFETY: element ranges are derived from validated source list offsets, and total_elements is
738    // the sum of the gathered element range lengths. Multiplier 1 preserves contiguous ranges.
739    let element_indices = unsafe {
740        PiecewiseSequenceArray::new_unchecked(
741            element_starts.into_array(),
742            element_lengths.into_array(),
743            multipliers,
744            total_elements,
745        )
746    };
747    let elements = elements.take(element_indices.into_array())?;
748
749    Ok(GatheredList { elements, offsets })
750}
751
752fn gather_piecewise_list_validity<S, L, Offset, OutOffset>(
753    elements: &ArrayRef,
754    offsets: &[Offset],
755    starts: &[S],
756    lengths: &[L],
757    output_len: usize,
758    total_elements: usize,
759    data_validity: &Mask,
760) -> VortexResult<GatheredList>
761where
762    S: UnsignedPType,
763    L: UnsignedPType,
764    Offset: UnsignedPType,
765    OutOffset: IntegerPType,
766{
767    let offsets_capacity = output_len
768        .checked_add(1)
769        .ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
770    let mut gather = ValidPieceGather {
771        new_offsets: BufferMut::<OutOffset>::with_capacity(offsets_capacity),
772        element_starts: BufferMut::<u64>::with_capacity(output_len),
773        element_lengths: BufferMut::<u64>::with_capacity(output_len),
774        output_elements: 0,
775    };
776
777    gather.new_offsets.push(OutOffset::zero());
778    for (&start, &length) in starts.iter().zip_eq(lengths) {
779        let start: usize = start.as_();
780        let length: usize = length.as_();
781        if length == 0 {
782            continue;
783        }
784
785        gather_valid_piece(offsets, data_validity, start, length, &mut gather);
786    }
787    debug_assert_eq!(gather.output_elements, total_elements);
788
789    let offsets =
790        PrimitiveArray::new(gather.new_offsets.freeze(), Validity::NonNullable).into_array();
791    let multipliers = ConstantArray::new(1u64, gather.element_starts.len()).into_array();
792    // SAFETY: element ranges come only from valid source list rows. Source list construction
793    // validated those row offsets, and null source rows produce no element range.
794    let element_indices = unsafe {
795        PiecewiseSequenceArray::new_unchecked(
796            gather.element_starts.into_array(),
797            gather.element_lengths.into_array(),
798            multipliers,
799            total_elements,
800        )
801    };
802    let elements = elements.take(element_indices.into_array())?;
803
804    Ok(GatheredList { elements, offsets })
805}
806
807fn gather_valid_piece<Offset, OutOffset>(
808    offsets: &[Offset],
809    data_validity: &Mask,
810    start: usize,
811    length: usize,
812    gather: &mut ValidPieceGather<OutOffset>,
813) where
814    Offset: UnsignedPType,
815    OutOffset: IntegerPType,
816{
817    let offset_range = &offsets[start..][..=length];
818    for (data_idx, window) in (start..).zip(offset_range.windows(2)) {
819        if !data_validity.value(data_idx) {
820            gather
821                .new_offsets
822                .push(new_offset_value::<OutOffset>(gather.output_elements));
823            continue;
824        }
825
826        let element_start: usize = window[0].as_();
827        let element_end: usize = window[1].as_();
828        let element_length = element_end - element_start;
829        if element_length != 0 {
830            gather.element_starts.push(element_start as u64);
831            gather.element_lengths.push(element_length as u64);
832            gather.output_elements += element_length;
833        }
834        gather
835            .new_offsets
836            .push(new_offset_value::<OutOffset>(gather.output_elements));
837    }
838}
839
840fn new_offset_value<T: IntegerPType>(value: usize) -> T {
841    T::from_usize(value).vortex_expect("output offset fits selected offset type")
842}
843
844#[cfg(test)]
845mod test {
846    use std::sync::Arc;
847
848    use rstest::rstest;
849    use vortex_buffer::buffer;
850    use vortex_error::VortexResult;
851
852    use crate::IntoArray as _;
853    use crate::VortexSessionExecute;
854    use crate::array_session;
855    use crate::arrays::BoolArray;
856    use crate::arrays::ConstantArray;
857    use crate::arrays::ListArray;
858    use crate::arrays::ListViewArray;
859    use crate::arrays::PiecewiseSequenceArray;
860    use crate::arrays::PrimitiveArray;
861    use crate::arrays::listview::ListViewArrayExt;
862    use crate::assert_arrays_eq;
863    use crate::compute::conformance::take::test_take_conformance;
864    use crate::dtype::DType;
865    use crate::dtype::Nullability;
866    use crate::dtype::PType::I32;
867    use crate::scalar::Scalar;
868    use crate::validity::Validity;
869
870    #[test]
871    fn nullable_take() {
872        let mut ctx = array_session().create_execution_ctx();
873        let list = ListArray::try_new(
874            buffer![0i32, 5, 3, 4].into_array(),
875            buffer![0, 2, 3, 4, 4].into_array(),
876            Validity::Array(BoolArray::from_iter(vec![true, true, false, true]).into_array()),
877        )
878        .unwrap()
879        .into_array();
880
881        let idx =
882            PrimitiveArray::from_option_iter(vec![Some(0), None, Some(1), Some(3)]).into_array();
883
884        let result = list.take(idx).unwrap();
885
886        assert_eq!(
887            result.dtype(),
888            &DType::List(
889                Arc::new(DType::Primitive(I32, Nullability::NonNullable)),
890                Nullability::Nullable
891            )
892        );
893
894        let result = result.execute::<ListViewArray>(&mut ctx).unwrap();
895
896        assert_eq!(result.len(), 4);
897
898        let element_dtype: Arc<DType> = Arc::new(I32.into());
899
900        assert!(
901            result
902                .is_valid(0, &mut array_session().create_execution_ctx())
903                .unwrap()
904        );
905        assert_eq!(
906            result
907                .execute_scalar(0, &mut array_session().create_execution_ctx())
908                .unwrap(),
909            Scalar::list(
910                Arc::clone(&element_dtype),
911                vec![0i32.into(), 5.into()],
912                Nullability::Nullable
913            )
914        );
915
916        assert!(
917            result
918                .is_invalid(1, &mut array_session().create_execution_ctx())
919                .unwrap()
920        );
921
922        assert!(
923            result
924                .is_valid(2, &mut array_session().create_execution_ctx())
925                .unwrap()
926        );
927        assert_eq!(
928            result
929                .execute_scalar(2, &mut array_session().create_execution_ctx())
930                .unwrap(),
931            Scalar::list(
932                Arc::clone(&element_dtype),
933                vec![3i32.into()],
934                Nullability::Nullable
935            )
936        );
937
938        assert!(
939            result
940                .is_valid(3, &mut array_session().create_execution_ctx())
941                .unwrap()
942        );
943        assert_eq!(
944            result
945                .execute_scalar(3, &mut array_session().create_execution_ctx())
946                .unwrap(),
947            Scalar::list(element_dtype, vec![], Nullability::Nullable)
948        );
949    }
950
951    #[test]
952    fn null_index_ignores_out_of_bounds_payload() {
953        let mut ctx = array_session().create_execution_ctx();
954        let list = ListArray::try_new(
955            buffer![1i32, 2, 3, 4].into_array(),
956            buffer![0u32, 2, 4].into_array(),
957            Validity::NonNullable,
958        )
959        .unwrap()
960        .into_array();
961
962        let idx = PrimitiveArray::new(
963            buffer![1u32, 99, 0],
964            Validity::from_iter([true, false, true]),
965        )
966        .into_array();
967        let result = list.take(idx).unwrap();
968
969        let expected = ListArray::new(
970            buffer![3i32, 4, 1, 2].into_array(),
971            buffer![0u32, 2, 2, 4].into_array(),
972            Validity::from_iter([true, false, true]),
973        );
974        assert_arrays_eq!(expected, result, &mut ctx);
975    }
976
977    #[test]
978    fn null_source_row_uses_valid_empty_output_range() {
979        let mut ctx = array_session().create_execution_ctx();
980        let list = ListArray::new(
981            buffer![1i32, 2, 7, 8].into_array(),
982            buffer![0u32, 2, 4].into_array(),
983            Validity::from_iter([true, false]),
984        )
985        .into_array();
986
987        let idx = buffer![0u32, 1].into_array();
988        let result = list.take(idx).unwrap();
989
990        let expected = ListArray::new(
991            buffer![1i32, 2].into_array(),
992            buffer![0u32, 2, 2].into_array(),
993            Validity::from_iter([true, false]),
994        );
995        assert_arrays_eq!(expected, result, &mut ctx);
996    }
997
998    #[test]
999    fn change_validity() {
1000        let list = ListArray::try_new(
1001            buffer![0i32, 5, 3, 4].into_array(),
1002            buffer![0, 2, 3].into_array(),
1003            Validity::NonNullable,
1004        )
1005        .unwrap()
1006        .into_array();
1007
1008        let idx = PrimitiveArray::from_option_iter(vec![Some(0), Some(1), None]).into_array();
1009        // since idx is nullable, the final list will also be nullable
1010
1011        let result = list.take(idx).unwrap();
1012        assert_eq!(
1013            result.dtype(),
1014            &DType::List(
1015                Arc::new(DType::Primitive(I32, Nullability::NonNullable)),
1016                Nullability::Nullable
1017            )
1018        );
1019    }
1020
1021    #[test]
1022    fn non_nullable_take() {
1023        let mut ctx = array_session().create_execution_ctx();
1024        let list = ListArray::try_new(
1025            buffer![0i32, 5, 3, 4].into_array(),
1026            buffer![0, 2, 3, 3, 4].into_array(),
1027            Validity::NonNullable,
1028        )
1029        .unwrap()
1030        .into_array();
1031
1032        let idx = buffer![1, 0, 2].into_array();
1033
1034        let result = list.take(idx).unwrap();
1035
1036        assert_eq!(
1037            result.dtype(),
1038            &DType::List(
1039                Arc::new(DType::Primitive(I32, Nullability::NonNullable)),
1040                Nullability::NonNullable
1041            )
1042        );
1043
1044        let result = result.execute::<ListViewArray>(&mut ctx).unwrap();
1045
1046        assert_eq!(result.len(), 3);
1047
1048        let element_dtype: Arc<DType> = Arc::new(I32.into());
1049
1050        assert!(
1051            result
1052                .is_valid(0, &mut array_session().create_execution_ctx())
1053                .unwrap()
1054        );
1055        assert_eq!(
1056            result
1057                .execute_scalar(0, &mut array_session().create_execution_ctx())
1058                .unwrap(),
1059            Scalar::list(
1060                Arc::clone(&element_dtype),
1061                vec![3i32.into()],
1062                Nullability::NonNullable
1063            )
1064        );
1065
1066        assert!(
1067            result
1068                .is_valid(1, &mut array_session().create_execution_ctx())
1069                .unwrap()
1070        );
1071        assert_eq!(
1072            result
1073                .execute_scalar(1, &mut array_session().create_execution_ctx())
1074                .unwrap(),
1075            Scalar::list(
1076                Arc::clone(&element_dtype),
1077                vec![0i32.into(), 5.into()],
1078                Nullability::NonNullable
1079            )
1080        );
1081
1082        assert!(
1083            result
1084                .is_valid(2, &mut array_session().create_execution_ctx())
1085                .unwrap()
1086        );
1087        assert_eq!(
1088            result
1089                .execute_scalar(2, &mut array_session().create_execution_ctx())
1090                .unwrap(),
1091            Scalar::list(element_dtype, vec![], Nullability::NonNullable)
1092        );
1093    }
1094
1095    #[test]
1096    fn piecewise_sequence_take() {
1097        let mut ctx = array_session().create_execution_ctx();
1098        let list = ListArray::try_new(
1099            buffer![0i32, 1, 2, 3, 4, 5, 6].into_array(),
1100            buffer![0u32, 2, 5, 5, 7].into_array(),
1101            Validity::NonNullable,
1102        )
1103        .unwrap()
1104        .into_array();
1105        let idx = PiecewiseSequenceArray::try_new(
1106            buffer![1u64, 0].into_array(),
1107            buffer![2u64, 1].into_array(),
1108            ConstantArray::new(1u64, 2).into_array(),
1109            3,
1110        )
1111        .unwrap()
1112        .into_array();
1113
1114        let result = list
1115            .take(idx)
1116            .unwrap()
1117            .execute::<ListViewArray>(&mut ctx)
1118            .unwrap();
1119
1120        let element_dtype: Arc<DType> = Arc::new(I32.into());
1121        assert_eq!(
1122            result.execute_scalar(0, &mut ctx).unwrap(),
1123            Scalar::list(
1124                Arc::clone(&element_dtype),
1125                vec![2i32.into(), 3.into(), 4.into()],
1126                Nullability::NonNullable
1127            )
1128        );
1129        assert_eq!(
1130            result.execute_scalar(1, &mut ctx).unwrap(),
1131            Scalar::list(Arc::clone(&element_dtype), vec![], Nullability::NonNullable)
1132        );
1133        assert_eq!(
1134            result.execute_scalar(2, &mut ctx).unwrap(),
1135            Scalar::list(
1136                element_dtype,
1137                vec![0i32.into(), 1.into()],
1138                Nullability::NonNullable
1139            )
1140        );
1141    }
1142
1143    #[test]
1144    fn piecewise_sequence_take_nullable_list_constant_lengths() -> VortexResult<()> {
1145        let mut ctx = array_session().create_execution_ctx();
1146        let list = ListArray::try_new(
1147            buffer![0i32, 1, 99, 100, 2, 3, 4, 5].into_array(),
1148            buffer![0u32, 2, 4, 7, 8].into_array(),
1149            Validity::Array(BoolArray::from_iter([true, false, true, true]).into_array()),
1150        )?
1151        .into_array();
1152        let idx = PiecewiseSequenceArray::try_new(
1153            buffer![0u64].into_array(),
1154            ConstantArray::new(4u64, 1).into_array(),
1155            ConstantArray::new(1u64, 1).into_array(),
1156            4,
1157        )?
1158        .into_array();
1159
1160        let result = list.take(idx)?.execute::<ListViewArray>(&mut ctx)?;
1161        assert_eq!(result.offset_at(0), 0);
1162        assert_eq!(result.size_at(0), 2);
1163        assert_eq!(result.offset_at(1), 2);
1164        assert_eq!(result.size_at(1), 0);
1165        assert_eq!(result.offset_at(2), 2);
1166        assert_eq!(result.size_at(2), 3);
1167        assert_eq!(result.offset_at(3), 5);
1168        assert_eq!(result.size_at(3), 1);
1169
1170        let element_dtype: Arc<DType> = Arc::new(I32.into());
1171        assert_eq!(
1172            result.execute_scalar(0, &mut ctx)?,
1173            Scalar::list(
1174                Arc::clone(&element_dtype),
1175                vec![0i32.into(), 1.into()],
1176                Nullability::Nullable
1177            )
1178        );
1179        assert!(result.is_invalid(1, &mut ctx)?);
1180        assert_eq!(
1181            result.execute_scalar(2, &mut ctx)?,
1182            Scalar::list(
1183                Arc::clone(&element_dtype),
1184                vec![2i32.into(), 3.into(), 4.into()],
1185                Nullability::Nullable
1186            )
1187        );
1188        assert_eq!(
1189            result.execute_scalar(3, &mut ctx)?,
1190            Scalar::list(element_dtype, vec![5i32.into()], Nullability::Nullable)
1191        );
1192        Ok(())
1193    }
1194
1195    #[test]
1196    fn piecewise_sequence_take_nullable_list_array_lengths() -> VortexResult<()> {
1197        let mut ctx = array_session().create_execution_ctx();
1198        let list = ListArray::try_new(
1199            buffer![0i32, 1, 99, 100, 2, 3, 4, 5].into_array(),
1200            buffer![0u32, 2, 4, 7, 8].into_array(),
1201            Validity::Array(BoolArray::from_iter([true, false, true, true]).into_array()),
1202        )?
1203        .into_array();
1204        let idx = PiecewiseSequenceArray::try_new(
1205            buffer![1u64, 0].into_array(),
1206            buffer![2u64, 1].into_array(),
1207            ConstantArray::new(1u64, 2).into_array(),
1208            3,
1209        )?
1210        .into_array();
1211
1212        let result = list.take(idx)?.execute::<ListViewArray>(&mut ctx)?;
1213        assert_eq!(result.offset_at(0), 0);
1214        assert_eq!(result.size_at(0), 0);
1215        assert_eq!(result.offset_at(1), 0);
1216        assert_eq!(result.size_at(1), 3);
1217        assert_eq!(result.offset_at(2), 3);
1218        assert_eq!(result.size_at(2), 2);
1219
1220        let element_dtype: Arc<DType> = Arc::new(I32.into());
1221        assert!(result.is_invalid(0, &mut ctx)?);
1222        assert_eq!(
1223            result.execute_scalar(1, &mut ctx)?,
1224            Scalar::list(
1225                Arc::clone(&element_dtype),
1226                vec![2i32.into(), 3.into(), 4.into()],
1227                Nullability::Nullable
1228            )
1229        );
1230        assert_eq!(
1231            result.execute_scalar(2, &mut ctx)?,
1232            Scalar::list(
1233                element_dtype,
1234                vec![0i32.into(), 1.into()],
1235                Nullability::Nullable
1236            )
1237        );
1238        Ok(())
1239    }
1240
1241    #[test]
1242    fn test_take_empty_array() {
1243        let list = ListArray::try_new(
1244            buffer![0i32, 5, 3, 4].into_array(),
1245            buffer![0].into_array(),
1246            Validity::NonNullable,
1247        )
1248        .unwrap()
1249        .into_array();
1250
1251        let idx = PrimitiveArray::empty::<i32>(Nullability::Nullable).into_array();
1252
1253        let result = list.take(idx).unwrap();
1254        assert_eq!(
1255            result.dtype(),
1256            &DType::List(
1257                Arc::new(DType::Primitive(I32, Nullability::NonNullable)),
1258                Nullability::Nullable
1259            )
1260        );
1261        assert_eq!(result.len(), 0,);
1262    }
1263
1264    #[rstest]
1265    #[case(ListArray::try_new(
1266        buffer![0i32, 1, 2, 3, 4, 5].into_array(),
1267        buffer![0, 2, 3, 5, 5, 6].into_array(),
1268        Validity::NonNullable,
1269    ).unwrap())]
1270    #[case(ListArray::try_new(
1271        buffer![10i32, 20, 30, 40, 50].into_array(),
1272        buffer![0, 2, 3, 4, 5].into_array(),
1273        Validity::Array(BoolArray::from_iter(vec![true, false, true, true]).into_array()),
1274    ).unwrap())]
1275    #[case(ListArray::try_new(
1276        buffer![1i32, 2, 3].into_array(),
1277        buffer![0, 0, 2, 2, 3].into_array(), // First and third are empty
1278        Validity::NonNullable,
1279    ).unwrap())]
1280    #[case(ListArray::try_new(
1281        buffer![42i32, 43].into_array(),
1282        buffer![0, 2].into_array(),
1283        Validity::NonNullable,
1284    ).unwrap())]
1285    #[case({
1286        let elements = buffer![0i32..200].into_array();
1287        let mut offsets = vec![0u64];
1288        for i in 1..=50 {
1289            offsets.push(offsets[i - 1] + (i as u64 % 5)); // Variable length lists
1290        }
1291        ListArray::try_new(
1292            elements,
1293            PrimitiveArray::from_iter(offsets).into_array(),
1294            Validity::NonNullable,
1295        ).unwrap()
1296    })]
1297    #[case(ListArray::try_new(
1298        PrimitiveArray::from_option_iter([Some(1i32), None, Some(3), Some(4), None]).into_array(),
1299        buffer![0, 2, 3, 5].into_array(),
1300        Validity::NonNullable,
1301    ).unwrap())]
1302    fn test_take_list_conformance(#[case] list: ListArray) {
1303        test_take_conformance(
1304            &list.into_array(),
1305            &mut array_session().create_execution_ctx(),
1306        );
1307    }
1308
1309    #[test]
1310    fn test_u64_offset_accumulation_non_nullable() {
1311        let mut ctx = array_session().create_execution_ctx();
1312        let elements = buffer![0i32; 200].into_array();
1313        let offsets = buffer![0u8, 200].into_array();
1314        let list = ListArray::try_new(elements, offsets, Validity::NonNullable)
1315            .unwrap()
1316            .into_array();
1317
1318        // Take the same large list twice - would overflow u8 but works with u64.
1319        let idx = buffer![0u8, 0].into_array();
1320        let result = list.take(idx).unwrap();
1321
1322        assert_eq!(result.len(), 2);
1323
1324        let result_view = result.execute::<ListViewArray>(&mut ctx).unwrap();
1325        assert_eq!(result_view.len(), 2);
1326        assert!(
1327            result_view
1328                .is_valid(0, &mut array_session().create_execution_ctx())
1329                .unwrap()
1330        );
1331        assert!(
1332            result_view
1333                .is_valid(1, &mut array_session().create_execution_ctx())
1334                .unwrap()
1335        );
1336    }
1337
1338    #[test]
1339    fn test_u64_offset_accumulation_nullable() {
1340        let mut ctx = array_session().create_execution_ctx();
1341        let elements = buffer![0i32; 150].into_array();
1342        let offsets = buffer![0u8, 150, 150].into_array();
1343        let validity = BoolArray::from_iter(vec![true, false]).into_array();
1344        let list = ListArray::try_new(elements, offsets, Validity::Array(validity))
1345            .unwrap()
1346            .into_array();
1347
1348        // Take the same large list twice - would overflow u8 but works with u64.
1349        let idx = PrimitiveArray::from_option_iter(vec![Some(0u8), None, Some(0u8)]).into_array();
1350        let result = list.take(idx).unwrap();
1351
1352        assert_eq!(result.len(), 3);
1353
1354        let result_view = result.execute::<ListViewArray>(&mut ctx).unwrap();
1355        assert_eq!(result_view.len(), 3);
1356        assert!(
1357            result_view
1358                .is_valid(0, &mut array_session().create_execution_ctx())
1359                .unwrap()
1360        );
1361        assert!(
1362            result_view
1363                .is_invalid(1, &mut array_session().create_execution_ctx())
1364                .unwrap()
1365        );
1366        assert!(
1367            result_view
1368                .is_valid(2, &mut array_session().create_execution_ctx())
1369                .unwrap()
1370        );
1371    }
1372
1373    /// Regression test for validity length mismatch bug.
1374    ///
1375    /// When source array has `Validity::Array(...)` and indices are non-nullable,
1376    /// the result validity must have length equal to indices.len(), not source.len().
1377    #[test]
1378    fn test_take_validity_length_mismatch_regression() {
1379        // Source array with explicit validity array (length 2).
1380        let list = ListArray::try_new(
1381            buffer![1i32, 2, 3, 4].into_array(),
1382            buffer![0, 2, 4].into_array(),
1383            Validity::Array(BoolArray::from_iter(vec![true, true]).into_array()),
1384        )
1385        .unwrap()
1386        .into_array();
1387
1388        // Take more indices than source length (4 vs 2) with non-nullable indices.
1389        let idx = buffer![0u32, 1, 0, 1].into_array();
1390
1391        // This should not panic - result should have length 4.
1392        let result = list.take(idx).unwrap();
1393        assert_eq!(result.len(), 4);
1394    }
1395}