This is an automated email from the ASF dual-hosted git repository.
alamb pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-rs.git
The following commit(s) were added to refs/heads/main by this push:
new 1c2c390875 perf(arrow-select): add take_record_batch_unchecked to skip
redundant bounds checks (#10945)
1c2c390875 is described below
commit 1c2c390875ee7f9edc6cecd6c8dfe044d5334e58
Author: RIchard Baah <[email protected]>
AuthorDate: Mon Sep 28 10:25:53 2026 -0400
perf(arrow-select): add take_record_batch_unchecked to skip redundant
bounds checks (#10945)
#10945 <- here, this PR is stacked on top of #10944
#10944
# Which issue does this PR close?
<!--
We generally require a GitHub issue to be filed for all bug fixes and
enhancements and this helps us generate change logs for our releases.
You can link an issue to this PR using the GitHub syntax.
-->
- Closes #8879.
- the only unique commit compared to #10944 is
47e2de03ef98a65f39efa01dbc5ba2dc576ca962
# Rationale for this change
`take_record_batch` previously called take per column, which always runs
a bounds check on every index even when the caller already knows the
indices are valid. For workloads doing repeated record batch takes with
pre-validated indices — such as sort, merge, or filter pipelines — this
check is redundant and measurable overhead. Benchmarks show ~13% speedup
on primitive-column batches at 1024 rows when the check is skipped.
<img width="643" height="161" alt="Image 9-1-26 at 2 28 PM"
src="https://github.com/user-attachments/assets/5170db51-ff1c-4062-b4fe-c69de894ddbc"
/>
<!--
Why are you proposing this change? If this is already explained clearly
in the issue then this section is not needed.
Explaining clearly why changes are proposed helps reviewers understand
your changes and offer better suggestions for fixes.
-->
# What changes are included in this PR?
- Introduces pub unsafe fn take_record_batch_unchecked which bypasses
index bounds checking by dispatching through `take_impl::<_, false>`
directly
<!--
There is no need to duplicate the description in the issue here but it
is sometimes worth providing a summary of the individual changes in this
PR.
-->
# Are these changes tested?
yes, existing test + miri
<!--
We typically require tests for all PRs in order to:
1. Prevent the code from being accidentally broken by subsequent changes
2. Serve as another way to document the expected behavior of the code
If tests are not included in your PR, please explain why (for example,
are they covered by existing tests)?
If this PR claims a performance improvement, please include evidence
such as benchmark results.
-->
# Are there any user-facing changes?
new `take_record_batch_unchecked()` method.
<!--
If there are user-facing changes then we may require documentation to be
updated before approving the PR.
If there are any breaking changes to public APIs, please call them out.
-->
---
arrow-select/src/take.rs | 242 +++++++++++++++++++++++++++++------------------
1 file changed, 152 insertions(+), 90 deletions(-)
diff --git a/arrow-select/src/take.rs b/arrow-select/src/take.rs
index 4ed8d856d4..ea505762e8 100644
--- a/arrow-select/src/take.rs
+++ b/arrow-select/src/take.rs
@@ -163,6 +163,33 @@ pub fn take_arrays(
.collect()
}
+/// For each [ArrayRef] in the [`Vec<ArrayRef>`], take elements by index and
create a new
+/// [`Vec<ArrayRef>`] from those indices, without bounds checking.
+///
+/// # Safety
+///
+/// Caller must ensure all values in `indices` are valid indices into each
array
+/// (i.e. `< array.len()`). Out-of-bounds indices will cause undefined
behavior.
+///
+/// # Errors
+///
+/// Returns an error if an index value cannot be cast to `usize`.
+pub unsafe fn take_arrays_unchecked(
+ arrays: &[ArrayRef],
+ indices: &dyn Array,
+) -> Result<Vec<ArrayRef>, ArrowError> {
+ downcast_integer_array!(
+ indices => {
+ let indices = indices.to_indices();
+ arrays
+ .iter()
+ .map(|array| take_impl::<_, false>(array.as_ref(), &indices))
+ .collect()
+ },
+ d => Err(ArrowError::InvalidArgumentError(format!("Take only supported
for integers, got {d:?}")))
+ )
+}
+
/// Verifies that the non-null values of `indices` are all `< len`
fn check_bounds<T: ArrowPrimitiveType>(
len: usize,
@@ -209,7 +236,7 @@ where
}
#[inline(never)]
-fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_impl<IndexType: ArrowPrimitiveType, const VALIDATE_INDICES: bool>(
values: &dyn Array,
indices: &PrimitiveArray<IndexType>,
) -> Result<ArrayRef, ArrowError> {
@@ -224,38 +251,38 @@ fn take_impl<IndexType: ArrowPrimitiveType, const
CHECKED: bool>(
return Ok(new_empty_array(values.data_type()));
}
downcast_primitive_array! {
- values => Ok(Arc::new(take_primitive::<_, _, CHECKED>(values,
indices)?)),
+ values => Ok(Arc::new(take_primitive::<_, _, VALIDATE_INDICES>(values,
indices)?)),
DataType::Boolean => {
let values =
values.as_any().downcast_ref::<BooleanArray>().unwrap();
- Ok(Arc::new(take_boolean::<_, CHECKED>(values, indices)))
+ Ok(Arc::new(take_boolean::<_, VALIDATE_INDICES>(values, indices)))
}
DataType::Utf8 => {
- Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_string::<i32>(),
indices)?))
+ Ok(Arc::new(take_bytes::<_, _,
VALIDATE_INDICES>(values.as_string::<i32>(), indices)?))
}
DataType::LargeUtf8 => {
- Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_string::<i64>(),
indices)?))
+ Ok(Arc::new(take_bytes::<_, _,
VALIDATE_INDICES>(values.as_string::<i64>(), indices)?))
}
DataType::Utf8View => {
- Ok(Arc::new(take_byte_view::<_, _,
CHECKED>(values.as_string_view(), indices)?))
+ Ok(Arc::new(take_byte_view::<_, _,
VALIDATE_INDICES>(values.as_string_view(), indices)?))
}
DataType::List(_) => {
- Ok(Arc::new(take_list::<_, Int32Type, CHECKED>(values.as_list(),
indices)?))
+ Ok(Arc::new(take_list::<_, Int32Type,
VALIDATE_INDICES>(values.as_list(), indices)?))
}
DataType::LargeList(_) => {
- Ok(Arc::new(take_list::<_, Int64Type, CHECKED>(values.as_list(),
indices)?))
+ Ok(Arc::new(take_list::<_, Int64Type,
VALIDATE_INDICES>(values.as_list(), indices)?))
}
DataType::ListView(_) => {
- Ok(Arc::new(take_list_view::<_, Int32Type,
CHECKED>(values.as_list_view(), indices)?))
+ Ok(Arc::new(take_list_view::<_, Int32Type,
VALIDATE_INDICES>(values.as_list_view(), indices)?))
}
DataType::LargeListView(_) => {
- Ok(Arc::new(take_list_view::<_, Int64Type,
CHECKED>(values.as_list_view(), indices)?))
+ Ok(Arc::new(take_list_view::<_, Int64Type,
VALIDATE_INDICES>(values.as_list_view(), indices)?))
}
DataType::FixedSizeList(_, length) => {
let values = values
.as_any()
.downcast_ref::<FixedSizeListArray>()
.unwrap();
- Ok(Arc::new(take_fixed_size_list::<_, CHECKED>(
+ Ok(Arc::new(take_fixed_size_list::<_, VALIDATE_INDICES>(
values,
indices,
*length as u32,
@@ -263,7 +290,7 @@ fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
}
DataType::Map(field, ordered) => {
let list_arr = ListArray::from(values.as_map().clone());
- let list_data = take_list::<_, Int32Type, CHECKED>(&list_arr,
indices)?;
+ let list_data = take_list::<_, Int32Type,
VALIDATE_INDICES>(&list_arr, indices)?;
let (_, offsets, entries, nulls) = list_data.into_parts();
let entries = entries.as_struct().clone();
Ok(Arc::new(MapArray::try_new(
@@ -279,7 +306,7 @@ fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
let arrays = array
.columns()
.iter()
- .map(|a| take_impl::<_, CHECKED>(a.as_ref(), indices))
+ .map(|a| take_impl::<_, VALIDATE_INDICES>(a.as_ref(), indices))
.collect::<Result<Vec<ArrayRef>, _>>()?;
let fields: Vec<(FieldRef, ArrayRef)> =
fields.iter().cloned().zip(arrays).collect();
@@ -304,7 +331,7 @@ fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
}
}
DataType::Dictionary(_, _) => downcast_dictionary_array! {
- values => Ok(Arc::new(take_dict::<_, _, CHECKED>(values,
indices)?)),
+ values => Ok(Arc::new(take_dict::<_, _, VALIDATE_INDICES>(values,
indices)?)),
t => unimplemented!("Take not supported for dictionary type {:?}",
t)
}
DataType::RunEndEncoded(_, _) => downcast_run_array! {
@@ -312,20 +339,20 @@ fn take_impl<IndexType: ArrowPrimitiveType, const
CHECKED: bool>(
t => unimplemented!("Take not supported for run type {:?}", t)
}
DataType::Binary => {
- Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_binary::<i32>(),
indices)?))
+ Ok(Arc::new(take_bytes::<_, _,
VALIDATE_INDICES>(values.as_binary::<i32>(), indices)?))
}
DataType::LargeBinary => {
- Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_binary::<i64>(),
indices)?))
+ Ok(Arc::new(take_bytes::<_, _,
VALIDATE_INDICES>(values.as_binary::<i64>(), indices)?))
}
DataType::BinaryView => {
- Ok(Arc::new(take_byte_view::<_, _,
CHECKED>(values.as_binary_view(), indices)?))
+ Ok(Arc::new(take_byte_view::<_, _,
VALIDATE_INDICES>(values.as_binary_view(), indices)?))
}
DataType::FixedSizeBinary(size) => {
let values = values
.as_any()
.downcast_ref::<FixedSizeBinaryArray>()
.unwrap();
- Ok(Arc::new(take_fixed_size_binary::<_, CHECKED>(values, indices,
*size)?))
+ Ok(Arc::new(take_fixed_size_binary::<_, VALIDATE_INDICES>(values,
indices, *size)?))
}
DataType::Null => {
// Take applied to a null array produces a null array.
@@ -345,7 +372,7 @@ fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
let children = fields
.iter()
.map(|(type_id, _)| {
- take_sparse_union_child::<_, CHECKED>(
+ take_sparse_union_child::<_, VALIDATE_INDICES>(
values.child(type_id),
null_type_id == Some(type_id),
indices,
@@ -363,7 +390,7 @@ fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
// Keep index nulls so `take` of each child writes a null instead
of
// reading child offset 0 (the default `take_native` fills in).
let offsets = <PrimitiveArray<Int32Type>>::try_new(
- take_native(values.offsets().unwrap(), indices),
+ take_native::<_, _,
VALIDATE_INDICES>(values.offsets().unwrap(), indices),
indices.nulls().cloned(),
)?;
@@ -375,7 +402,7 @@ fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
let values = values.child(field_type_id);
- take_impl::<_, CHECKED>(values,
indices.as_primitive::<Int32Type>())
+ take_impl::<_, VALIDATE_INDICES>(values,
indices.as_primitive::<Int32Type>())
})
.collect::<Result<_, _>>()?;
@@ -418,11 +445,11 @@ fn take_union_type_ids<IndexType: ArrowPrimitiveType>(
indices: &PrimitiveArray<IndexType>,
) -> Result<(ScalarBuffer<i8>, Option<i8>), ArrowError> {
if indices.null_count() == 0 {
- return Ok((take_native(type_ids, indices), None));
+ return Ok((take_native::<_, _, true>(type_ids, indices), None));
}
let null_type_id = union_null_type_id(fields)?;
- let taken_type_ids = take_native(type_ids, indices);
+ let taken_type_ids = take_native::<_, _, true>(type_ids, indices);
let type_ids = indices
.iter()
.zip(&taken_type_ids)
@@ -474,13 +501,13 @@ fn field_can_represent_take_null(field: &FieldRef) ->
bool {
/// ([`take_union_type_ids`]). Values in the other children at those positions
/// are unspecified, so they are taken with dummy indices instead of
introducing
/// nulls that would contradict field metadata.
-fn take_sparse_union_child<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_sparse_union_child<IndexType: ArrowPrimitiveType, const
VALIDATE_INDICES: bool>(
values: &dyn Array,
represent_nulls: bool,
indices: &PrimitiveArray<IndexType>,
) -> Result<ArrayRef, ArrowError> {
if represent_nulls || indices.null_count() == 0 {
- return take_impl::<_, CHECKED>(values, indices);
+ return take_impl::<_, VALIDATE_INDICES>(values, indices);
}
if values.is_empty() {
@@ -496,7 +523,7 @@ fn take_sparse_union_child<IndexType: ArrowPrimitiveType,
const CHECKED: bool>(
));
}
- take_impl::<_, CHECKED>(values, &indices_without_nulls(indices))
+ take_impl::<_, VALIDATE_INDICES>(values, &indices_without_nulls(indices))
}
/// Replaces null take indices with `0` and drops the null bitmap.
@@ -526,7 +553,7 @@ pub struct TakeOptions {
/// values: [1, 2, 3, null, 5]
/// indices: [0, null, 4, 3]
/// The result is: [1 (slot 0), null (null slot), 5 (slot 4), null (slot 3)]
-fn take_primitive<T, I, const CHECKED: bool>(
+fn take_primitive<T, I, const VALIDATE_INDICES: bool>(
values: &PrimitiveArray<T>,
indices: &PrimitiveArray<I>,
) -> Result<PrimitiveArray<T>, ArrowError>
@@ -534,19 +561,19 @@ where
T: ArrowPrimitiveType,
I: ArrowPrimitiveType,
{
- let values_buf = take_native(values.values(), indices);
- let nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
+ let values_buf = take_native::<_, _, VALIDATE_INDICES>(values.values(),
indices);
+ let nulls = take_nulls::<_, VALIDATE_INDICES>(values.nulls(), indices);
Ok(PrimitiveArray::try_new(values_buf,
nulls)?.with_data_type(values.data_type().clone()))
}
#[inline(never)]
-fn take_nulls<I: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_nulls<I: ArrowPrimitiveType, const VALIDATE_INDICES: bool>(
values: Option<&NullBuffer>,
indices: &PrimitiveArray<I>,
) -> Option<NullBuffer> {
match values.filter(|n| n.null_count() > 0) {
Some(n) => NullBuffer::from_unsliced_buffer(
- take_bits::<_, CHECKED>(n.inner(), indices).into_inner(),
+ take_bits::<_, VALIDATE_INDICES>(n.inner(), indices).into_inner(),
indices.len(),
),
None => indices.nulls().cloned(),
@@ -554,7 +581,7 @@ fn take_nulls<I: ArrowPrimitiveType, const CHECKED: bool>(
}
#[inline(never)]
-fn take_native<T: ArrowNativeType, I: ArrowPrimitiveType>(
+fn take_native<T: ArrowNativeType, I: ArrowPrimitiveType, const
VALIDATE_INDICES: bool>(
values: &[T],
indices: &PrimitiveArray<I>,
) -> ScalarBuffer<T> {
@@ -572,11 +599,24 @@ fn take_native<T: ArrowNativeType, I: ArrowPrimitiveType>(
},
})
.collect(),
- None => indices
- .values()
- .iter()
- .map(|index| values[index.as_usize()])
- .collect(),
+ None => {
+ if VALIDATE_INDICES {
+ indices
+ .values()
+ .iter()
+ .map(|index| values[index.as_usize()])
+ .collect()
+ } else {
+ indices
+ .values()
+ .iter()
+ .map(|index| {
+ // SAFETY: caller guarantees all indices are in-bounds.
+ unsafe { *values.get_unchecked(index.as_usize()) }
+ })
+ .collect()
+ }
+ }
}
}
@@ -617,7 +657,7 @@ unsafe fn pack_bit(src: *const u8, bit_idx: usize, out_pos:
usize) -> u8 {
}
#[inline(never)]
-fn take_bits<I: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_bits<I: ArrowPrimitiveType, const VALIDATE_INDICES: bool>(
values: &BooleanBuffer,
indices: &PrimitiveArray<I>,
) -> BooleanBuffer {
@@ -633,7 +673,7 @@ fn take_bits<I: ArrowPrimitiveType, const CHECKED: bool>(
index_nulls.valid_indices().for_each(|valid_idx| {
// SAFETY: valid_idx < indices.len(), guaranteed by
valid_indices().
let index_val = unsafe { indices.value_unchecked(valid_idx)
}.as_usize();
- if CHECKED {
+ if VALIDATE_INDICES {
if values.value(index_val) {
// SAFETY: valid_idx < indices.len() = len, output
buffer holds len bits.
unsafe { bit_util::set_bit_raw(out_ptr, valid_idx) };
@@ -659,7 +699,7 @@ fn take_bits<I: ArrowPrimitiveType, const CHECKED: bool>(
// SAFETY: base + bit < full_bytes * 8 <= len, so base +
bit is a valid
// position in the indices array.
let index_val = unsafe { indices.value_unchecked(base +
bit) }.as_usize();
- if CHECKED {
+ if VALIDATE_INDICES {
byte |= (values.value(index_val) as u8) << bit;
} else {
// SAFETY: caller guarantees index_val < values.len().
@@ -676,7 +716,7 @@ fn take_bits<I: ArrowPrimitiveType, const CHECKED: bool>(
// SAFETY: base + bit < len (remainder loop bound), so
base + bit is a
// valid position in the indices array.
let index_val = unsafe { indices.value_unchecked(base +
bit) }.as_usize();
- if CHECKED {
+ if VALIDATE_INDICES {
byte |= (values.value(index_val) as u8) << bit;
} else {
// SAFETY: caller guarantees index_val < values.len().
@@ -693,7 +733,7 @@ fn take_bits<I: ArrowPrimitiveType, const CHECKED: bool>(
/// Gather value bits and validity bits from two boolean buffers in a single
pass.
/// Used when the values array itself has nulls, avoiding two separate
`take_bits` calls.
#[inline(never)]
-fn take_bits_with_validity<I: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_bits_with_validity<I: ArrowPrimitiveType, const VALIDATE_INDICES:
bool>(
values: &BooleanBuffer,
validity: &BooleanBuffer,
indices: &PrimitiveArray<I>,
@@ -715,7 +755,7 @@ fn take_bits_with_validity<I: ArrowPrimitiveType, const
CHECKED: bool>(
for out_pos in index_nulls.valid_indices() {
// SAFETY: out_pos < indices.len(), guaranteed by
valid_indices().
let src_idx = unsafe { indices.value_unchecked(out_pos)
}.as_usize();
- if CHECKED {
+ if VALIDATE_INDICES {
if values.value(src_idx) {
// SAFETY: out_pos < indices.len() = len, output
buffer holds len bits.
unsafe { bit_util::set_bit_raw(value_out_ptr, out_pos)
};
@@ -760,7 +800,7 @@ fn take_bits_with_validity<I: ArrowPrimitiveType, const
CHECKED: bool>(
for bit_pos in 0..8usize {
// SAFETY: bit_base + bit_pos < full_bytes * 8 <= len.
let src_idx = unsafe { indices.value_unchecked(bit_base +
bit_pos) }.as_usize();
- if CHECKED {
+ if VALIDATE_INDICES {
packed_values |= (values.value(src_idx) as u8) <<
bit_pos;
packed_validity |= (validity.value(src_idx) as u8) <<
bit_pos;
} else {
@@ -784,7 +824,7 @@ fn take_bits_with_validity<I: ArrowPrimitiveType, const
CHECKED: bool>(
for bit_pos in 0..(len - bit_base) {
// SAFETY: bit_base + bit_pos < len (remainder loop bound).
let src_idx = unsafe { indices.value_unchecked(bit_base +
bit_pos) }.as_usize();
- if CHECKED {
+ if VALIDATE_INDICES {
packed_values |= (values.value(src_idx) as u8) <<
bit_pos;
packed_validity |= (validity.value(src_idx) as u8) <<
bit_pos;
} else {
@@ -809,7 +849,7 @@ fn take_bits_with_validity<I: ArrowPrimitiveType, const
CHECKED: bool>(
}
/// `take` implementation for boolean arrays
-fn take_boolean<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_boolean<IndexType: ArrowPrimitiveType, const VALIDATE_INDICES: bool>(
array: &BooleanArray,
indices: &PrimitiveArray<IndexType>,
) -> BooleanArray {
@@ -817,19 +857,19 @@ fn take_boolean<IndexType: ArrowPrimitiveType, const
CHECKED: bool>(
match array.nulls().filter(|n| n.null_count() > 0) {
Some(array_nulls) => {
let (val_buf, null_buf) =
- take_bits_with_validity::<_, CHECKED>(bits,
array_nulls.inner(), indices);
+ take_bits_with_validity::<_, VALIDATE_INDICES>(bits,
array_nulls.inner(), indices);
BooleanArray::new(val_buf, null_buf)
}
None => {
- let val_buf = take_bits::<_, CHECKED>(bits, indices);
- let null_buf = take_nulls::<_, CHECKED>(None, indices);
+ let val_buf = take_bits::<_, VALIDATE_INDICES>(bits, indices);
+ let null_buf = take_nulls::<_, VALIDATE_INDICES>(None, indices);
BooleanArray::new(val_buf, null_buf)
}
}
}
/// `take` implementation for string arrays
-fn take_bytes<T: ByteArrayType, IndexType: ArrowPrimitiveType, const CHECKED:
bool>(
+fn take_bytes<T: ByteArrayType, IndexType: ArrowPrimitiveType, const
VALIDATE_INDICES: bool>(
array: &GenericByteArray<T>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericByteArray<T>, ArrowError> {
@@ -839,7 +879,7 @@ fn take_bytes<T: ByteArrayType, IndexType:
ArrowPrimitiveType, const CHECKED: bo
let input_offsets = array.value_offsets();
let mut capacity = 0;
- let nulls = take_nulls::<_, CHECKED>(array.nulls(), indices);
+ let nulls = take_nulls::<_, VALIDATE_INDICES>(array.nulls(), indices);
// Branch on output nulls — `None` means every output slot is valid.
match nulls.as_ref().filter(|n| n.null_count() > 0) {
@@ -960,12 +1000,12 @@ fn take_bytes<T: ByteArrayType, IndexType:
ArrowPrimitiveType, const CHECKED: bo
}
/// `take` implementation for byte view arrays
-fn take_byte_view<T: ByteViewType, IndexType: ArrowPrimitiveType, const
CHECKED: bool>(
+fn take_byte_view<T: ByteViewType, IndexType: ArrowPrimitiveType, const
VALIDATE_INDICES: bool>(
array: &GenericByteViewArray<T>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericByteViewArray<T>, ArrowError> {
- let new_views = take_native(array.views(), indices);
- let new_nulls = take_nulls::<_, CHECKED>(array.nulls(), indices);
+ let new_views = take_native::<_, _, VALIDATE_INDICES>(array.views(),
indices);
+ let new_nulls = take_nulls::<_, VALIDATE_INDICES>(array.nulls(), indices);
let buffers = Arc::clone(array.data_buffers());
// Safety: array.views was valid, and take_native copies only valid
values, and verifies bounds
Ok(unsafe { GenericByteViewArray::new_unchecked(new_views, buffers,
new_nulls) })
@@ -975,7 +1015,7 @@ fn take_byte_view<T: ByteViewType, IndexType:
ArrowPrimitiveType, const CHECKED:
///
/// Copies the selected list entries' child slices into a new child array
/// via `MutableArrayData`, then reconstructs a list array with new offsets
-fn take_list<IndexType, OffsetType, const CHECKED: bool>(
+fn take_list<IndexType, OffsetType, const VALIDATE_INDICES: bool>(
values: &GenericListArray<OffsetType::Native>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericListArray<OffsetType::Native>, ArrowError>
@@ -987,7 +1027,7 @@ where
{
let src_offsets = values.value_offsets();
let child_data = values.values().to_data();
- let nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
+ let nulls = take_nulls::<_, VALIDATE_INDICES>(values.nulls(), indices);
let mut dst_offsets = Vec::with_capacity(indices.len() + 1);
dst_offsets.push(OffsetType::Native::zero());
@@ -1032,15 +1072,25 @@ where
if prev < vidx {
dst_offsets.extend(std::iter::repeat_n(child_len, vidx
- prev));
}
- let row = if CHECKED {
- indices.value(vidx).as_usize()
+ // SAFETY: vidx < indices.len(), guaranteed by
valid_indices().
+ let row = unsafe { indices.value_unchecked(vidx)
}.as_usize();
+ let (start, end) = if VALIDATE_INDICES {
+ (
+ child_buf_offset + src_offsets[row].as_usize() *
bytes_per_value,
+ child_buf_offset + src_offsets[row + 1].as_usize()
* bytes_per_value,
+ )
} else {
- // SAFETY: !CHECKED means the caller guarantees all
indices are valid;
- // `vidx` is further bounded by the validity bitmap of
`indices`.
- unsafe { indices.value_unchecked(vidx) }.as_usize()
+ // SAFETY: caller guarantees row < values.len(), so
row+1 is also in bounds.
+ unsafe {
+ (
+ child_buf_offset
+ +
src_offsets.get_unchecked(row).as_usize() * bytes_per_value,
+ child_buf_offset
+ + src_offsets.get_unchecked(row +
1).as_usize()
+ * bytes_per_value,
+ )
+ }
};
- let start = child_buf_offset + src_offsets[row].as_usize()
* bytes_per_value;
- let end = child_buf_offset + src_offsets[row +
1].as_usize() * bytes_per_value;
dst_buf.extend_from_slice(&values_buf[start..end]);
child_len = child_len
.checked_add(&(src_offsets[row + 1] -
src_offsets[row]))
@@ -1102,18 +1152,20 @@ where
if last < i {
dst_offsets.extend(std::iter::repeat_n(current, i - last));
}
- let row = if CHECKED {
- indices.value(i).as_usize()
+ // SAFETY: i < indices.len(), guaranteed by valid_indices().
+ let row = unsafe { indices.value_unchecked(i) }.as_usize();
+ let (start, end) = if VALIDATE_INDICES {
+ (src_offsets[row].as_usize(), src_offsets[row +
1].as_usize())
} else {
- // SAFETY: !CHECKED means the caller guarantees all
indices are valid;
- // `i` is further bounded by the validity bitmap of
`indices`.
- unsafe { indices.value_unchecked(i) }.as_usize()
+ // SAFETY: caller guarantees row < values.len(), so row+1
is also in bounds.
+ unsafe {
+ (
+ src_offsets.get_unchecked(row).as_usize(),
+ src_offsets.get_unchecked(row + 1).as_usize(),
+ )
+ }
};
- mutable.try_extend(
- 0,
- src_offsets[row].as_usize(),
- src_offsets[row + 1].as_usize(),
- )?;
+ mutable.try_extend(0, start, end)?;
dst_offsets.push(
OffsetType::Native::from_usize(mutable.len())
.ok_or_else(||
ArrowError::OffsetOverflowError(mutable.len()))?,
@@ -1135,7 +1187,7 @@ where
GenericListArray::<OffsetType::Native>::try_new(field, offsets, child,
nulls)
}
-fn take_list_view<IndexType, OffsetType, const CHECKED: bool>(
+fn take_list_view<IndexType, OffsetType, const VALIDATE_INDICES: bool>(
values: &GenericListViewArray<OffsetType::Native>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericListViewArray<OffsetType::Native>, ArrowError>
@@ -1144,9 +1196,9 @@ where
OffsetType: ArrowPrimitiveType,
OffsetType::Native: OffsetSizeTrait,
{
- let taken_offsets = take_native(values.offsets(), indices);
- let taken_sizes = take_native(values.sizes(), indices);
- let nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
+ let taken_offsets = take_native::<_, _,
VALIDATE_INDICES>(values.offsets(), indices);
+ let taken_sizes = take_native::<_, _, VALIDATE_INDICES>(values.sizes(),
indices);
+ let nulls = take_nulls::<_, VALIDATE_INDICES>(values.nulls(), indices);
let field = match values.data_type() {
DataType::ListView(field) | DataType::LargeListView(field) =>
field.clone(),
@@ -1171,14 +1223,14 @@ where
/// Calculates the index and indexed offset for the inner array,
/// applying `take` on the inner array, then reconstructing a list array
/// with the indexed offsets
-fn take_fixed_size_list<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_fixed_size_list<IndexType: ArrowPrimitiveType, const VALIDATE_INDICES:
bool>(
values: &FixedSizeListArray,
indices: &PrimitiveArray<IndexType>,
length: <UInt32Type as ArrowPrimitiveType>::Native,
) -> Result<FixedSizeListArray, ArrowError> {
let field = values.value_field();
let child = values.values();
- let nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
+ let nulls = take_nulls::<_, VALIDATE_INDICES>(values.nulls(), indices);
// Fast path: primitive child with no nulls copy row-sized byte blocks
directly,
let taken_child = if child.null_count() == 0
@@ -1193,7 +1245,7 @@ fn take_fixed_size_list<IndexType: ArrowPrimitiveType,
const CHECKED: bool>(
)
} else {
let list_indices = take_value_indices_from_fixed_size_list(values,
indices, length)?;
- take_impl::<UInt32Type, CHECKED>(child.as_ref(), &list_indices)?
+ take_impl::<UInt32Type, VALIDATE_INDICES>(child.as_ref(),
&list_indices)?
};
FixedSizeListArray::try_new_with_length(
@@ -1265,7 +1317,7 @@ fn take_fixed_size_list_primitive<IndexType:
ArrowPrimitiveType>(
/// The computation is done in two steps:
/// - Compute the values buffer
/// - Compute the null buffer
-fn take_fixed_size_binary<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
+fn take_fixed_size_binary<IndexType: ArrowPrimitiveType, const
VALIDATE_INDICES: bool>(
values: &FixedSizeBinaryArray,
indices: &PrimitiveArray<IndexType>,
size: i32,
@@ -1283,7 +1335,7 @@ fn take_fixed_size_binary<IndexType: ArrowPrimitiveType,
const CHECKED: bool>(
_ => take_fixed_size_binary_buffer_dynamic_length(values, indices,
size_usize),
};
- let value_nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
+ let value_nulls = take_nulls::<_, VALIDATE_INDICES>(values.nulls(),
indices);
let final_nulls = NullBuffer::union(value_nulls.as_ref(), indices.nulls());
return FixedSizeBinaryArray::try_new(size, result_buffer, final_nulls);
@@ -1395,11 +1447,11 @@ fn take_fixed_size<IndexType: ArrowPrimitiveType, const
N: usize>(
///
/// applies `take` to the keys of the dictionary array and returns a new
dictionary array
/// with the same dictionary values and reordered keys
-fn take_dict<T: ArrowDictionaryKeyType, I: ArrowPrimitiveType, const CHECKED:
bool>(
+fn take_dict<T: ArrowDictionaryKeyType, I: ArrowPrimitiveType, const
VALIDATE_INDICES: bool>(
values: &DictionaryArray<T>,
indices: &PrimitiveArray<I>,
) -> Result<DictionaryArray<T>, ArrowError> {
- let new_keys = take_primitive::<_, _, CHECKED>(values.keys(), indices)?;
+ let new_keys = take_primitive::<_, _, VALIDATE_INDICES>(values.keys(),
indices)?;
Ok(unsafe { DictionaryArray::new_unchecked(new_keys,
values.values().clone()) })
}
@@ -1661,12 +1713,22 @@ pub fn take_record_batch(
record_batch: &RecordBatch,
indices: &dyn Array,
) -> Result<RecordBatch, ArrowError> {
- let columns = record_batch
- .columns()
- .iter()
- .map(|c| take(c, indices, None))
- .collect::<Result<Vec<_>, _>>()?;
- RecordBatch::try_new(record_batch.schema(), columns)
+ downcast_integer_array!(
+ indices => {
+ let indices = indices.to_indices();
+ let cols = record_batch.columns();
+ let mut columns = Vec::with_capacity(cols.len());
+ if let Some(first) = cols.first() {
+ columns.push(take_impl::<_, true>(first.as_ref(), &indices)?);
+ for col in &cols[1..] {
+ columns.push(take_impl::<_, false>(col.as_ref(),
&indices)?);
+ }
+ }
+ // Safety: indices were validated by the first take_impl call
+ Ok(unsafe { RecordBatch::new_unchecked(record_batch.schema(),
columns, indices.len()) })
+ },
+ d => Err(ArrowError::InvalidArgumentError(format!("Take only supported
for integers, got {d:?}")))
+ )
}
#[cfg(test)]