1use 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
35impl TakeExecute for List {
39 #[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 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 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 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 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 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 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 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 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 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 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(), 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)); }
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 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 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 #[test]
1378 fn test_take_validity_length_mismatch_regression() {
1379 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 let idx = buffer![0u32, 1, 0, 1].into_array();
1390
1391 let result = list.take(idx).unwrap();
1393 assert_eq!(result.len(), 4);
1394 }
1395}