Skip to content

Commit 09ca199

Browse files
committed
Rely on PiecewiseSequence validation in List take
Signed-off-by: Daniel King <dan@spiraldb.com>
1 parent f1880b1 commit 09ca199

1 file changed

Lines changed: 15 additions & 51 deletions

File tree

  • vortex-array/src/arrays/list/compute

vortex-array/src/arrays/list/compute/take.rs

Lines changed: 15 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,10 @@ use itertools::Itertools as _;
55
use vortex_buffer::BufferMut;
66
use vortex_error::VortexExpect;
77
use vortex_error::VortexResult;
8-
use vortex_error::vortex_bail;
98
use vortex_error::vortex_ensure;
109
use vortex_error::vortex_err;
1110

1211
use crate::ArrayRef;
13-
use crate::Canonical;
1412
use crate::Columnar;
1513
use crate::IntoArray;
1614
use crate::array::ArrayView;
@@ -169,7 +167,7 @@ fn take_piecewise_sequence(
169167

170168
let taken = match lengths {
171169
Columnar::Constant(lengths) => {
172-
let length = constant_unsigned_usize(&lengths)?;
170+
let length = constant_unsigned_usize(&lengths);
173171
take_piecewise_sequence_constant_dispatch(
174172
array,
175173
&starts,
@@ -179,7 +177,8 @@ fn take_piecewise_sequence(
179177
output_len,
180178
)?
181179
}
182-
Columnar::Canonical(Canonical::Primitive(lengths)) => {
180+
Columnar::Canonical(lengths) => {
181+
let lengths = lengths.into_primitive();
183182
take_piecewise_sequence_lengths_dispatch(
184183
array,
185184
&starts,
@@ -189,12 +188,6 @@ fn take_piecewise_sequence(
189188
output_len,
190189
)?
191190
}
192-
Columnar::Canonical(lengths) => {
193-
vortex_bail!(
194-
"PiecewiseSequenceArray lengths must be primitive or constant, got {}",
195-
lengths.dtype()
196-
)
197-
}
198191
};
199192
Ok(Some(taken))
200193
}
@@ -329,8 +322,7 @@ where
329322
computed_len == output_len,
330323
"PiecewiseSequenceArray expanded length {computed_len} does not match declared length {output_len}"
331324
);
332-
let total_elements =
333-
piecewise_list_elements_len_constant(array.elements().len(), offsets, starts, length)?;
325+
let total_elements = piecewise_list_elements_len_constant(offsets, starts, length)?;
334326
let validity = array.validity()?.take(indices_ref)?;
335327

336328
match_smallest_offset_type!(total_elements, |OutputOffset| {
@@ -377,8 +369,7 @@ where
377369
computed_len == output_len,
378370
"PiecewiseSequenceArray expanded length {computed_len} does not match declared length {output_len}"
379371
);
380-
let total_elements =
381-
piecewise_list_elements_len(array.elements().len(), offsets, starts, lengths)?;
372+
let total_elements = piecewise_list_elements_len(offsets, starts, lengths)?;
382373

383374
match_smallest_offset_type!(total_elements, |OutputOffset| {
384375
let gathered = gather_piecewise_list::<S, L, Offset, OutputOffset>(
@@ -407,7 +398,6 @@ struct GatheredList {
407398
}
408399

409400
fn piecewise_list_elements_len_constant<S, Offset>(
410-
elements_len: usize,
411401
offsets: &[Offset],
412402
starts: &[S],
413403
length: usize,
@@ -426,10 +416,6 @@ where
426416
let offset_range = &offsets[start..][..=length];
427417
let element_start: usize = offset_range[0].as_();
428418
let element_end: usize = offset_range[length].as_();
429-
vortex_ensure!(
430-
element_start <= element_end && element_end <= elements_len,
431-
"List offsets range {element_start}..{element_end} exceeds elements length {elements_len}",
432-
);
433419
total = total
434420
.checked_add(element_end - element_start)
435421
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
@@ -438,7 +424,6 @@ where
438424
}
439425

440426
fn piecewise_list_elements_len<S, L, Offset>(
441-
elements_len: usize,
442427
offsets: &[Offset],
443428
starts: &[S],
444429
lengths: &[L],
@@ -455,10 +440,6 @@ where
455440
let offset_range = &offsets[start..][..=length];
456441
let element_start: usize = offset_range[0].as_();
457442
let element_end: usize = offset_range[length].as_();
458-
vortex_ensure!(
459-
element_start <= element_end && element_end <= elements_len,
460-
"List offsets range {element_start}..{element_end} exceeds elements length {elements_len}",
461-
);
462443
total = total
463444
.checked_add(element_end - element_start)
464445
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
@@ -499,21 +480,15 @@ where
499480
let element_end: usize = offset_range[length].as_();
500481
for &offset in &offset_range[1..] {
501482
let offset: usize = offset.as_();
502-
let relative = offset
503-
.checked_sub(element_start)
504-
.ok_or_else(|| vortex_err!("List offsets are not monotonic at offset {offset}"))?;
505-
let output_offset = output_elements
506-
.checked_add(relative)
507-
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
508-
new_offsets.push(new_offset_value::<OutputOffset>(output_offset)?);
483+
let relative = offset - element_start;
484+
let output_offset = output_elements + relative;
485+
new_offsets.push(new_offset_value::<OutputOffset>(output_offset));
509486
}
510487

511488
let element_length = element_end - element_start;
512489
element_starts.push(element_start as u64);
513490
element_lengths.push(element_length as u64);
514-
output_elements = output_elements
515-
.checked_add(element_length)
516-
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
491+
output_elements += element_length;
517492
}
518493
debug_assert_eq!(output_elements, total_elements);
519494

@@ -569,21 +544,15 @@ where
569544
let element_end: usize = offset_range[length].as_();
570545
for &offset in &offset_range[1..] {
571546
let offset: usize = offset.as_();
572-
let relative = offset
573-
.checked_sub(element_start)
574-
.ok_or_else(|| vortex_err!("List offsets are not monotonic at offset {offset}"))?;
575-
let output_offset = output_elements
576-
.checked_add(relative)
577-
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
578-
new_offsets.push(new_offset_value::<OutputOffset>(output_offset)?);
547+
let relative = offset - element_start;
548+
let output_offset = output_elements + relative;
549+
new_offsets.push(new_offset_value::<OutputOffset>(output_offset));
579550
}
580551

581552
let element_length = element_end - element_start;
582553
element_starts.push(element_start as u64);
583554
element_lengths.push(element_length as u64);
584-
output_elements = output_elements
585-
.checked_add(element_length)
586-
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
555+
output_elements += element_length;
587556
}
588557
debug_assert_eq!(output_elements, total_elements);
589558

@@ -604,13 +573,8 @@ where
604573
Ok(GatheredList { elements, offsets })
605574
}
606575

607-
fn new_offset_value<T: IntegerPType>(value: usize) -> VortexResult<T> {
608-
T::from_usize(value).ok_or_else(|| {
609-
vortex_err!(
610-
"List take offset value {value} does not fit in {}",
611-
T::PTYPE
612-
)
613-
})
576+
fn new_offset_value<T: IntegerPType>(value: usize) -> T {
577+
T::from_usize(value).vortex_expect("output offset fits selected offset type")
614578
}
615579

616580
// Kept out-of-line: as a single-callsite generic helper it would otherwise be inlined into every

0 commit comments

Comments
 (0)