@@ -20,8 +20,10 @@ use crate::arrays::Primitive;
2020use crate :: arrays:: PrimitiveArray ;
2121use crate :: arrays:: dict:: TakeExecute ;
2222use crate :: arrays:: list:: ListArrayExt ;
23+ use crate :: arrays:: piecewise_sequence:: UnitMultiplierLengths ;
2324use crate :: arrays:: piecewise_sequence:: execute_unit_multiplier_index_arrays;
2425use crate :: arrays:: piecewise_sequence:: validate_index_ranges;
26+ use crate :: arrays:: piecewise_sequence:: validate_index_ranges_constant;
2527use crate :: arrays:: primitive:: PrimitiveArrayExt ;
2628use crate :: builders:: ArrayBuilder ;
2729use crate :: builders:: PrimitiveBuilder ;
@@ -162,21 +164,174 @@ fn take_piecewise_sequence(
162164 } ;
163165 let offsets = array. offsets ( ) . clone ( ) . execute :: < PrimitiveArray > ( ctx) ?;
164166 let offsets = offsets. reinterpret_cast ( offsets. ptype ( ) . to_unsigned ( ) ) ;
167+ let output_len = indices_ref. len ( ) ;
168+
169+ let taken = match & lengths {
170+ UnitMultiplierLengths :: Constant ( length) => take_piecewise_sequence_constant_dispatch (
171+ array,
172+ & starts,
173+ * length,
174+ & offsets,
175+ indices_ref,
176+ output_len,
177+ ) ?,
178+ UnitMultiplierLengths :: Array ( lengths) => take_piecewise_sequence_lengths_dispatch (
179+ array,
180+ & starts,
181+ lengths,
182+ & offsets,
183+ indices_ref,
184+ output_len,
185+ ) ?,
186+ } ;
187+ Ok ( Some ( taken) )
188+ }
165189
190+ fn take_piecewise_sequence_constant_dispatch (
191+ array : ArrayView < ' _ , List > ,
192+ starts : & PrimitiveArray ,
193+ length : usize ,
194+ offsets : & PrimitiveArray ,
195+ indices_ref : & ArrayRef ,
196+ output_len : usize ,
197+ ) -> VortexResult < ArrayRef > {
166198 match_each_unsigned_integer_ptype ! ( starts. ptype( ) , |S | {
167- match_each_unsigned_integer_ptype!( lengths. ptype( ) , |L | {
168- match_each_unsigned_integer_ptype!( offsets. ptype( ) , |O | {
169- take_piecewise_sequence_typed:: <S , L , O >(
170- array,
171- starts. as_slice:: <S >( ) ,
172- lengths. as_slice:: <L >( ) ,
173- offsets. as_slice:: <O >( ) ,
174- indices_ref,
175- )
176- } )
177- } )
199+ take_piecewise_sequence_constant_start_dispatch:: <S >(
200+ array,
201+ starts,
202+ length,
203+ offsets,
204+ indices_ref,
205+ output_len,
206+ )
207+ } )
208+ }
209+
210+ fn take_piecewise_sequence_constant_start_dispatch < S > (
211+ array : ArrayView < ' _ , List > ,
212+ starts : & PrimitiveArray ,
213+ length : usize ,
214+ offsets : & PrimitiveArray ,
215+ indices_ref : & ArrayRef ,
216+ output_len : usize ,
217+ ) -> VortexResult < ArrayRef >
218+ where
219+ S : UnsignedPType ,
220+ {
221+ match_each_unsigned_integer_ptype ! ( offsets. ptype( ) , |O | {
222+ take_piecewise_sequence_constant_length:: <S , O >(
223+ array,
224+ starts. as_slice:: <S >( ) ,
225+ length,
226+ offsets. as_slice:: <O >( ) ,
227+ indices_ref,
228+ output_len,
229+ )
230+ } )
231+ }
232+
233+ fn take_piecewise_sequence_lengths_dispatch (
234+ array : ArrayView < ' _ , List > ,
235+ starts : & PrimitiveArray ,
236+ lengths : & PrimitiveArray ,
237+ offsets : & PrimitiveArray ,
238+ indices_ref : & ArrayRef ,
239+ output_len : usize ,
240+ ) -> VortexResult < ArrayRef > {
241+ match_each_unsigned_integer_ptype ! ( starts. ptype( ) , |S | {
242+ take_piecewise_sequence_lengths_start_dispatch:: <S >(
243+ array,
244+ starts,
245+ lengths,
246+ offsets,
247+ indices_ref,
248+ output_len,
249+ )
250+ } )
251+ }
252+
253+ fn take_piecewise_sequence_lengths_start_dispatch < S > (
254+ array : ArrayView < ' _ , List > ,
255+ starts : & PrimitiveArray ,
256+ lengths : & PrimitiveArray ,
257+ offsets : & PrimitiveArray ,
258+ indices_ref : & ArrayRef ,
259+ output_len : usize ,
260+ ) -> VortexResult < ArrayRef >
261+ where
262+ S : UnsignedPType ,
263+ {
264+ match_each_unsigned_integer_ptype ! ( lengths. ptype( ) , |L | {
265+ take_piecewise_sequence_lengths_start_length_dispatch:: <S , L >(
266+ array,
267+ starts,
268+ lengths,
269+ offsets,
270+ indices_ref,
271+ output_len,
272+ )
273+ } )
274+ }
275+
276+ fn take_piecewise_sequence_lengths_start_length_dispatch < S , L > (
277+ array : ArrayView < ' _ , List > ,
278+ starts : & PrimitiveArray ,
279+ lengths : & PrimitiveArray ,
280+ offsets : & PrimitiveArray ,
281+ indices_ref : & ArrayRef ,
282+ output_len : usize ,
283+ ) -> VortexResult < ArrayRef >
284+ where
285+ S : UnsignedPType ,
286+ L : UnsignedPType ,
287+ {
288+ match_each_unsigned_integer_ptype ! ( offsets. ptype( ) , |O | {
289+ take_piecewise_sequence_typed:: <S , L , O >(
290+ array,
291+ starts. as_slice:: <S >( ) ,
292+ lengths. as_slice:: <L >( ) ,
293+ offsets. as_slice:: <O >( ) ,
294+ indices_ref,
295+ output_len,
296+ )
297+ } )
298+ }
299+
300+ fn take_piecewise_sequence_constant_length < S , Offset > (
301+ array : ArrayView < ' _ , List > ,
302+ starts : & [ S ] ,
303+ length : usize ,
304+ offsets : & [ Offset ] ,
305+ indices_ref : & ArrayRef ,
306+ output_len : usize ,
307+ ) -> VortexResult < ArrayRef >
308+ where
309+ S : UnsignedPType ,
310+ Offset : UnsignedPType ,
311+ {
312+ validate_index_ranges_constant ( array. len ( ) , starts, length, output_len) ?;
313+ let total_elements =
314+ piecewise_list_elements_len_constant ( array. elements ( ) . len ( ) , offsets, starts, length) ?;
315+ let validity = array. validity ( ) ?. take ( indices_ref) ?;
316+
317+ match_smallest_offset_type ! ( total_elements, |OutputOffset | {
318+ let gathered = gather_piecewise_list_constant_length:: <S , Offset , OutputOffset >(
319+ array. elements( ) ,
320+ offsets,
321+ starts,
322+ length,
323+ output_len,
324+ total_elements,
325+ ) ?;
326+
327+ // SAFETY: output offsets are rebuilt from valid monotonic source offsets; output elements
328+ // are exactly the gathered child ranges referenced by those offsets; validity has one bit
329+ // per output row.
330+ Ok (
331+ unsafe { ListArray :: new_unchecked( gathered. elements, gathered. offsets, validity) }
332+ . into_array( ) ,
333+ )
178334 } )
179- . map ( Some )
180335}
181336
182337fn take_piecewise_sequence_typed < S , L , Offset > (
@@ -185,13 +340,14 @@ fn take_piecewise_sequence_typed<S, L, Offset>(
185340 lengths : & [ L ] ,
186341 offsets : & [ Offset ] ,
187342 indices_ref : & ArrayRef ,
343+ output_len : usize ,
188344) -> VortexResult < ArrayRef >
189345where
190346 S : UnsignedPType ,
191347 L : UnsignedPType ,
192348 Offset : UnsignedPType ,
193349{
194- validate_index_ranges ( array. len ( ) , starts, lengths, indices_ref . len ( ) ) ?;
350+ validate_index_ranges ( array. len ( ) , starts, lengths, output_len ) ?;
195351 let total_elements =
196352 piecewise_list_elements_len ( array. elements ( ) . len ( ) , offsets, starts, lengths) ?;
197353
@@ -201,7 +357,7 @@ where
201357 offsets,
202358 starts,
203359 lengths,
204- indices_ref . len ( ) ,
360+ output_len ,
205361 total_elements,
206362 ) ?;
207363 let validity = array. validity( ) ?. take( indices_ref) ?;
@@ -221,6 +377,37 @@ struct GatheredList {
221377 offsets : ArrayRef ,
222378}
223379
380+ fn piecewise_list_elements_len_constant < S , Offset > (
381+ elements_len : usize ,
382+ offsets : & [ Offset ] ,
383+ starts : & [ S ] ,
384+ length : usize ,
385+ ) -> VortexResult < usize >
386+ where
387+ S : UnsignedPType ,
388+ Offset : UnsignedPType ,
389+ {
390+ let mut total = 0usize ;
391+ for & start in starts {
392+ let start: usize = start. as_ ( ) ;
393+ let end = start + length;
394+ if length == 0 {
395+ continue ;
396+ }
397+
398+ let element_start: usize = offsets[ start] . as_ ( ) ;
399+ let element_end: usize = offsets[ end] . as_ ( ) ;
400+ vortex_ensure ! (
401+ element_start <= element_end && element_end <= elements_len,
402+ "List offsets range {element_start}..{element_end} exceeds elements length {elements_len}" ,
403+ ) ;
404+ total = total
405+ . checked_add ( element_end - element_start)
406+ . ok_or_else ( || vortex_err ! ( "List take output elements length overflow" ) ) ?;
407+ }
408+ Ok ( total)
409+ }
410+
224411fn piecewise_list_elements_len < S , L , Offset > (
225412 elements_len : usize ,
226413 offsets : & [ Offset ] ,
@@ -254,6 +441,74 @@ where
254441 Ok ( total)
255442}
256443
444+ fn gather_piecewise_list_constant_length < S , Offset , OutputOffset > (
445+ elements : & ArrayRef ,
446+ offsets : & [ Offset ] ,
447+ starts : & [ S ] ,
448+ length : usize ,
449+ output_len : usize ,
450+ total_elements : usize ,
451+ ) -> VortexResult < GatheredList >
452+ where
453+ S : UnsignedPType ,
454+ Offset : UnsignedPType ,
455+ OutputOffset : IntegerPType ,
456+ {
457+ let offsets_capacity = output_len
458+ . checked_add ( 1 )
459+ . ok_or_else ( || vortex_err ! ( "List take offsets length overflow" ) ) ?;
460+ let mut new_offsets = BufferMut :: < OutputOffset > :: with_capacity ( offsets_capacity) ;
461+ let mut element_starts = BufferMut :: < u64 > :: with_capacity ( starts. len ( ) ) ;
462+ let mut element_lengths = BufferMut :: < u64 > :: with_capacity ( starts. len ( ) ) ;
463+ let mut output_elements = 0usize ;
464+
465+ new_offsets. push ( OutputOffset :: zero ( ) ) ;
466+ for & start in starts {
467+ let start: usize = start. as_ ( ) ;
468+ let end = start + length;
469+ if length == 0 {
470+ continue ;
471+ }
472+
473+ let element_start: usize = offsets[ start] . as_ ( ) ;
474+ let element_end: usize = offsets[ end] . as_ ( ) ;
475+ for & offset in & offsets[ start + 1 ..=end] {
476+ let offset: usize = offset. as_ ( ) ;
477+ let relative = offset
478+ . checked_sub ( element_start)
479+ . ok_or_else ( || vortex_err ! ( "List offsets are not monotonic at offset {offset}" ) ) ?;
480+ let output_offset = output_elements
481+ . checked_add ( relative)
482+ . ok_or_else ( || vortex_err ! ( "List take output elements length overflow" ) ) ?;
483+ new_offsets. push ( new_offset_value :: < OutputOffset > ( output_offset) ?) ;
484+ }
485+
486+ let element_length = element_end - element_start;
487+ element_starts. push ( element_start as u64 ) ;
488+ element_lengths. push ( element_length as u64 ) ;
489+ output_elements = output_elements
490+ . checked_add ( element_length)
491+ . ok_or_else ( || vortex_err ! ( "List take output elements length overflow" ) ) ?;
492+ }
493+ debug_assert_eq ! ( output_elements, total_elements) ;
494+
495+ let offsets = PrimitiveArray :: new ( new_offsets. freeze ( ) , Validity :: NonNullable ) . into_array ( ) ;
496+ let multipliers = ConstantArray :: new ( 1u64 , element_starts. len ( ) ) . into_array ( ) ;
497+ // SAFETY: element ranges are derived from validated source list offsets, and total_elements is
498+ // the sum of the gathered element range lengths. Multiplier 1 preserves contiguous ranges.
499+ let element_indices = unsafe {
500+ PiecewiseSequenceArray :: new_unchecked (
501+ element_starts. into_array ( ) ,
502+ element_lengths. into_array ( ) ,
503+ multipliers,
504+ total_elements,
505+ )
506+ } ;
507+ let elements = elements. take ( element_indices. into_array ( ) ) ?;
508+
509+ Ok ( GatheredList { elements, offsets } )
510+ }
511+
257512fn gather_piecewise_list < S , L , Offset , OutputOffset > (
258513 elements : & ArrayRef ,
259514 offsets : & [ Offset ] ,
0 commit comments