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