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::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
36impl TakeExecute for List {
40 #[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 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 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 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 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 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 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 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 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 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 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(), 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)); }
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 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 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 #[test]
1379 fn test_take_validity_length_mismatch_regression() {
1380 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 let idx = buffer![0u32, 1, 0, 1].into_array();
1391
1392 let result = list.take(idx).unwrap();
1394 assert_eq!(result.len(), 4);
1395 }
1396}