use std::fmt::Display;
use std::mem::ManuallyDrop;
use std::sync::Arc;
use arrow_array::builder::UInt32Builder;
use arrow_array::cast::AsArray;
use arrow_array::types::*;
use arrow_array::*;
use arrow_buffer::{
ArrowNativeType, BooleanBuffer, Buffer, MutableBuffer, NullBuffer, NullBufferBuilder,
OffsetBuffer, RunEndBuffer, ScalarBuffer, bit_util,
};
use arrow_cmp::make_comparator;
use arrow_data::{ArrayData, transform::MutableArrayData};
use arrow_schema::{ArrowError, DataType, FieldRef, SortOptions, UnionFields, UnionMode};
use num_traits::{CheckedAdd, Zero};
pub fn take(
values: &dyn Array,
indices: &dyn Array,
options: Option<TakeOptions>,
) -> Result<ArrayRef, ArrowError> {
let options = options.unwrap_or_default();
downcast_integer_array!(
indices => {
let indices = indices.to_indices();
if options.check_bounds {
check_bounds(values.len(), &indices)?;
}
take_impl::<_, true>(values, &indices)
},
d => Err(ArrowError::InvalidArgumentError(format!("Take only supported for integers, got {d:?}")))
)
}
pub fn take_arrays(
arrays: &[ArrayRef],
indices: &dyn Array,
options: Option<TakeOptions>,
) -> Result<Vec<ArrayRef>, ArrowError> {
arrays
.iter()
.map(|array| take(array.as_ref(), indices, options.clone()))
.collect()
}
fn check_bounds<T: ArrowPrimitiveType>(
len: usize,
indices: &PrimitiveArray<T>,
) -> Result<(), ArrowError>
where
T::Native: Display,
{
let Some(len) = T::Native::from_usize(len) else {
return if T::DATA_TYPE.is_integer() {
Ok(())
} else {
Err(ArrowError::ComputeError("Cast to usize failed".to_string()))
};
};
if indices.null_count() > 0 {
indices.iter().flatten().try_for_each(|index| {
if index >= len {
return Err(ArrowError::ComputeError(format!(
"Array index out of bounds, cannot get item at index {index} from {len} entries"
)));
}
Ok(())
})
} else {
let in_bounds = indices.values().iter().fold(true, |in_bounds, &i| {
in_bounds & (i >= T::Native::ZERO) & (i < len)
});
if !in_bounds {
for &index in indices.values() {
if index < T::Native::ZERO || index >= len {
return Err(ArrowError::ComputeError(format!(
"Array index out of bounds, cannot get item at index {index} from {len} entries"
)));
}
}
}
Ok(())
}
}
#[inline(never)]
fn take_impl<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
values: &dyn Array,
indices: &PrimitiveArray<IndexType>,
) -> Result<ArrayRef, ArrowError> {
if indices.is_empty() {
if let DataType::Union(fields, _) = values.data_type()
&& fields.is_empty()
{
return Ok(values.slice(0, 0));
}
return Ok(new_empty_array(values.data_type()));
}
downcast_primitive_array! {
values => Ok(Arc::new(take_primitive::<_, _, CHECKED>(values, indices)?)),
DataType::Boolean => {
let values = values.as_any().downcast_ref::<BooleanArray>().unwrap();
Ok(Arc::new(take_boolean::<_, CHECKED>(values, indices)))
}
DataType::Utf8 => {
Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_string::<i32>(), indices)?))
}
DataType::LargeUtf8 => {
Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_string::<i64>(), indices)?))
}
DataType::Utf8View => {
Ok(Arc::new(take_byte_view::<_, _, CHECKED>(values.as_string_view(), indices)?))
}
DataType::List(_) => {
Ok(Arc::new(take_list::<_, Int32Type, CHECKED>(values.as_list(), indices)?))
}
DataType::LargeList(_) => {
Ok(Arc::new(take_list::<_, Int64Type, CHECKED>(values.as_list(), indices)?))
}
DataType::ListView(_) => {
Ok(Arc::new(take_list_view::<_, Int32Type, CHECKED>(values.as_list_view(), indices)?))
}
DataType::LargeListView(_) => {
Ok(Arc::new(take_list_view::<_, Int64Type, CHECKED>(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>(
values,
indices,
*length as u32,
)?))
}
DataType::Map(field, ordered) => {
let list_arr = ListArray::from(values.as_map().clone());
let list_data = take_list::<_, Int32Type, CHECKED>(&list_arr, indices)?;
let (_, offsets, entries, nulls) = list_data.into_parts();
let entries = entries.as_struct().clone();
Ok(Arc::new(MapArray::try_new(
field.clone(),
offsets,
entries,
nulls,
*ordered,
)?))
}
DataType::Struct(fields) => {
let array: &StructArray = values.as_struct();
let arrays = array
.columns()
.iter()
.map(|a| take_impl::<_, CHECKED>(a.as_ref(), indices))
.collect::<Result<Vec<ArrayRef>, _>>()?;
let fields: Vec<(FieldRef, ArrayRef)> =
fields.iter().cloned().zip(arrays).collect();
let is_valid: Buffer = indices
.iter()
.map(|index| {
if let Some(index) = index {
array.is_valid(index.to_usize().unwrap())
} else {
false
}
})
.collect();
if fields.is_empty() {
let nulls = NullBuffer::new(BooleanBuffer::new(is_valid, 0, indices.len()));
Ok(Arc::new(StructArray::new_empty_fields(indices.len(), Some(nulls))))
} else {
Ok(Arc::new(StructArray::from((fields, is_valid))) as ArrayRef)
}
}
DataType::Dictionary(_, _) => downcast_dictionary_array! {
values => Ok(Arc::new(take_dict::<_, _, CHECKED>(values, indices)?)),
t => unimplemented!("Take not supported for dictionary type {:?}", t)
}
DataType::RunEndEncoded(_, _) => downcast_run_array! {
values => Ok(Arc::new(take_run(values, indices)?)),
t => unimplemented!("Take not supported for run type {:?}", t)
}
DataType::Binary => {
Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_binary::<i32>(), indices)?))
}
DataType::LargeBinary => {
Ok(Arc::new(take_bytes::<_, _, CHECKED>(values.as_binary::<i64>(), indices)?))
}
DataType::BinaryView => {
Ok(Arc::new(take_byte_view::<_, _, CHECKED>(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)?))
}
DataType::Null => {
if values.len() >= indices.len() {
Ok(values.slice(0, indices.len()))
} else {
Ok(new_null_array(&DataType::Null, indices.len()))
}
}
DataType::Union(fields, UnionMode::Sparse) => {
let mut children = Vec::with_capacity(fields.len());
let values = values.as_any().downcast_ref::<UnionArray>().unwrap();
let type_ids = take_union_type_ids(fields, values.type_ids(), indices)?;
for (type_id, _field) in fields.iter() {
let values = values.child(type_id);
let values = take_impl::<_, CHECKED>(values, indices)?;
children.push(values);
}
let array = UnionArray::try_new(fields.clone(), type_ids, None, children)?;
Ok(Arc::new(array))
}
DataType::Union(fields, UnionMode::Dense) => {
let values = values.as_any().downcast_ref::<UnionArray>().unwrap();
let type_ids = PrimitiveArray::<Int8Type>::try_new(
take_union_type_ids(fields, values.type_ids(), indices)?,
None,
)?;
let offsets = <PrimitiveArray<Int32Type>>::try_new(
take_native(values.offsets().unwrap(), indices),
indices.nulls().cloned(),
)?;
let children = fields.iter()
.map(|(field_type_id, _)| {
let mask = BooleanArray::from_unary(&type_ids, |value_type_id| value_type_id == field_type_id);
let indices = crate::filter::filter(&offsets, &mask)?;
let values = values.child(field_type_id);
take_impl::<_, CHECKED>(values, indices.as_primitive::<Int32Type>())
})
.collect::<Result<_, _>>()?;
let mut child_offsets = [0; 128];
let offsets = type_ids.values()
.iter()
.map(|&i| {
let offset = child_offsets[i as usize];
child_offsets[i as usize] += 1;
offset
})
.collect();
let (_, type_ids, _) = type_ids.into_parts();
let array = UnionArray::try_new(fields.clone(), type_ids, Some(offsets), children)?;
Ok(Arc::new(array))
}
t => unimplemented!("Take not supported for data type {:?}", t)
}
}
fn take_union_type_ids<IndexType: ArrowPrimitiveType>(
fields: &UnionFields,
type_ids: &ScalarBuffer<i8>,
indices: &PrimitiveArray<IndexType>,
) -> Result<ScalarBuffer<i8>, ArrowError> {
if indices.null_count() == 0 {
return Ok(take_native(type_ids, indices));
}
let null_type_id = fields
.iter()
.next()
.map(|(type_id, _)| type_id)
.ok_or_else(|| {
ArrowError::ComputeError(
"Cannot take from a union with zero fields when indices contains nulls".into(),
)
})?;
let taken_type_ids = take_native(type_ids, indices);
let type_ids = indices
.iter()
.zip(&taken_type_ids)
.map(|(index, &type_id)| {
if index.is_some() {
type_id
} else {
null_type_id
}
})
.collect::<ScalarBuffer<_>>();
Ok(type_ids)
}
#[derive(Clone, Debug, Default)]
pub struct TakeOptions {
pub check_bounds: bool,
}
fn take_primitive<T, I, const CHECKED: bool>(
values: &PrimitiveArray<T>,
indices: &PrimitiveArray<I>,
) -> Result<PrimitiveArray<T>, ArrowError>
where
T: ArrowPrimitiveType,
I: ArrowPrimitiveType,
{
let values_buf = take_native(values.values(), indices);
let nulls = take_nulls::<_, CHECKED>(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>(
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(),
indices.len(),
),
None => indices.nulls().cloned(),
}
}
#[inline(never)]
fn take_native<T: ArrowNativeType, I: ArrowPrimitiveType>(
values: &[T],
indices: &PrimitiveArray<I>,
) -> ScalarBuffer<T> {
match indices.nulls().filter(|n| n.null_count() > 0) {
Some(n) => indices
.values()
.iter()
.enumerate()
.map(|(idx, index)| match values.get(index.as_usize()) {
Some(v) => *v,
None => match unsafe { n.inner().value_unchecked(idx) } {
false => T::default(),
true => panic!("Out-of-bounds index {index:?}"),
},
})
.collect(),
None => indices
.values()
.iter()
.map(|index| values[index.as_usize()])
.collect(),
}
}
#[inline(always)]
unsafe fn copy_bit_if_set(src: *const u8, src_bit_idx: usize, dst: *mut u8, dst_bit_idx: usize) {
unsafe {
if bit_util::get_bit_raw(src, src_bit_idx) {
bit_util::set_bit_raw(dst, dst_bit_idx);
}
}
}
#[inline(always)]
unsafe fn pack_bit(src: *const u8, bit_idx: usize, out_pos: usize) -> u8 {
let byte = unsafe { *src.add(bit_idx >> 3) }; ((byte >> (bit_idx & 7)) & 1) << out_pos }
#[inline(never)]
fn take_bits<I: ArrowPrimitiveType, const CHECKED: bool>(
values: &BooleanBuffer,
indices: &PrimitiveArray<I>,
) -> BooleanBuffer {
let len = indices.len();
let src_offset = values.offset();
let src_ptr = values.values().as_ptr();
let out_bytes = len.div_ceil(8);
match indices.nulls().filter(|nulls| nulls.null_count() > 0) {
Some(index_nulls) => {
let mut output = vec![0u8; out_bytes];
let out_ptr = output.as_mut_ptr();
index_nulls.valid_indices().for_each(|valid_idx| {
let index_val = unsafe { indices.value_unchecked(valid_idx) }.as_usize();
if CHECKED {
if values.value(index_val) {
unsafe { bit_util::set_bit_raw(out_ptr, valid_idx) };
}
} else {
unsafe { copy_bit_if_set(src_ptr, index_val + src_offset, out_ptr, valid_idx) };
}
});
BooleanBuffer::new(Buffer::from(output), 0, len)
}
None => {
let mut output = vec![0u8; out_bytes];
let out_slice = output.as_mut_slice();
let full_bytes = len / 8;
for (byte_idx, out_byte) in out_slice.iter_mut().enumerate().take(full_bytes) {
let base = byte_idx * 8;
let mut byte = 0u8;
for bit in 0..8usize {
let index_val = unsafe { indices.value_unchecked(base + bit) }.as_usize();
if CHECKED {
byte |= (values.value(index_val) as u8) << bit;
} else {
byte |= unsafe { pack_bit(src_ptr, index_val + src_offset, bit) };
}
}
*out_byte = byte;
}
if full_bytes < out_bytes {
let base = full_bytes * 8;
let mut byte = 0u8;
for bit in 0..(len - base) {
let index_val = unsafe { indices.value_unchecked(base + bit) }.as_usize();
if CHECKED {
byte |= (values.value(index_val) as u8) << bit;
} else {
byte |= unsafe { pack_bit(src_ptr, index_val + src_offset, bit) };
}
}
out_slice[full_bytes] = byte;
}
BooleanBuffer::new(Buffer::from(output), 0, len)
}
}
}
#[inline(never)]
fn take_bits_with_validity<I: ArrowPrimitiveType, const CHECKED: bool>(
values: &BooleanBuffer,
validity: &BooleanBuffer,
indices: &PrimitiveArray<I>,
) -> (BooleanBuffer, Option<NullBuffer>) {
let len = indices.len();
let value_bit_offset = values.offset();
let validity_bit_offset = validity.offset();
let value_data_ptr = values.values().as_ptr();
let validity_data_ptr = validity.values().as_ptr();
let out_bytes = len.div_ceil(8);
let mut value_out = vec![0u8; out_bytes];
let mut validity_out = vec![0u8; out_bytes];
match indices.nulls().filter(|nulls| nulls.null_count() > 0) {
Some(index_nulls) => {
let value_out_ptr = value_out.as_mut_ptr();
let validity_out_ptr = validity_out.as_mut_ptr();
for out_pos in index_nulls.valid_indices() {
let src_idx = unsafe { indices.value_unchecked(out_pos) }.as_usize();
if CHECKED {
if values.value(src_idx) {
unsafe { bit_util::set_bit_raw(value_out_ptr, out_pos) };
}
if validity.value(src_idx) {
unsafe { bit_util::set_bit_raw(validity_out_ptr, out_pos) };
}
} else {
unsafe {
copy_bit_if_set(
value_data_ptr,
src_idx + value_bit_offset,
value_out_ptr,
out_pos,
);
copy_bit_if_set(
validity_data_ptr,
src_idx + validity_bit_offset,
validity_out_ptr,
out_pos,
);
}
}
}
}
None => {
let value_out_slice = value_out.as_mut_slice();
let validity_out_slice = validity_out.as_mut_slice();
let full_bytes = len / 8;
for (byte_idx, (value_out_byte, validity_out_byte)) in value_out_slice
.iter_mut()
.zip(validity_out_slice.iter_mut())
.enumerate()
.take(full_bytes)
{
let bit_base = byte_idx * 8;
let mut packed_values = 0u8;
let mut packed_validity = 0u8;
for bit_pos in 0..8usize {
let src_idx = unsafe { indices.value_unchecked(bit_base + bit_pos) }.as_usize();
if CHECKED {
packed_values |= (values.value(src_idx) as u8) << bit_pos;
packed_validity |= (validity.value(src_idx) as u8) << bit_pos;
} else {
packed_values |= unsafe {
pack_bit(value_data_ptr, src_idx + value_bit_offset, bit_pos)
};
packed_validity |= unsafe {
pack_bit(validity_data_ptr, src_idx + validity_bit_offset, bit_pos)
};
}
}
*value_out_byte = packed_values;
*validity_out_byte = packed_validity;
}
if full_bytes < out_bytes {
let bit_base = full_bytes * 8;
let mut packed_values = 0u8;
let mut packed_validity = 0u8;
for bit_pos in 0..(len - bit_base) {
let src_idx = unsafe { indices.value_unchecked(bit_base + bit_pos) }.as_usize();
if CHECKED {
packed_values |= (values.value(src_idx) as u8) << bit_pos;
packed_validity |= (validity.value(src_idx) as u8) << bit_pos;
} else {
packed_values |= unsafe {
pack_bit(value_data_ptr, src_idx + value_bit_offset, bit_pos)
};
packed_validity |= unsafe {
pack_bit(validity_data_ptr, src_idx + validity_bit_offset, bit_pos)
};
}
}
value_out_slice[full_bytes] = packed_values;
validity_out_slice[full_bytes] = packed_validity;
}
}
}
let value_buf_out = BooleanBuffer::new(Buffer::from(value_out), 0, len);
let validity_buf_out = NullBuffer::from_unsliced_buffer(validity_out, len);
(value_buf_out, validity_buf_out)
}
fn take_boolean<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
array: &BooleanArray,
indices: &PrimitiveArray<IndexType>,
) -> BooleanArray {
let bits = array.values();
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);
BooleanArray::new(val_buf, null_buf)
}
None => {
let val_buf = take_bits::<_, CHECKED>(bits, indices);
let null_buf = take_nulls::<_, CHECKED>(None, indices);
BooleanArray::new(val_buf, null_buf)
}
}
}
fn take_bytes<T: ByteArrayType, IndexType: ArrowPrimitiveType, const CHECKED: bool>(
array: &GenericByteArray<T>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericByteArray<T>, ArrowError> {
let mut values: Vec<u8> = Vec::new();
let mut offsets = Vec::with_capacity(indices.len() + 1);
offsets.push(T::Offset::default());
let input_offsets = array.value_offsets();
let mut capacity = 0;
let nulls = take_nulls::<_, CHECKED>(array.nulls(), indices);
match nulls.as_ref().filter(|n| n.null_count() > 0) {
None => {
for index in indices.values() {
let index = index.as_usize();
let start = input_offsets[index].as_usize();
let end = input_offsets[index + 1].as_usize();
capacity += end - start;
offsets.push(
T::Offset::from_usize(capacity)
.ok_or_else(|| ArrowError::OffsetOverflowError(capacity))?,
);
}
values.reserve(capacity);
let dst = values.spare_capacity_mut();
debug_assert!(dst.len() >= capacity);
let mut offset = 0;
for index in indices.values() {
unsafe {
let data: &[u8] = array.value_unchecked(index.as_usize()).as_ref();
std::ptr::copy_nonoverlapping(
data.as_ptr(),
dst.get_unchecked_mut(offset..).as_mut_ptr().cast::<u8>(),
data.len(),
);
offset += data.len();
}
}
unsafe {
values.set_len(capacity);
}
}
Some(output_nulls) => {
let mut source_ranges = Vec::with_capacity(indices.len() - output_nulls.null_count());
let mut last_filled = 0;
offsets.resize(indices.len() + 1, T::Offset::default());
for i in output_nulls.valid_indices() {
let current_offset = T::Offset::from_usize(capacity)
.ok_or_else(|| ArrowError::OffsetOverflowError(capacity))?;
if last_filled < i {
offsets[last_filled + 1..=i].fill(current_offset);
}
let index = unsafe { indices.value_unchecked(i) }.as_usize();
let start = input_offsets[index].as_usize();
let end = input_offsets[index + 1].as_usize();
capacity += end - start;
offsets[i + 1] = T::Offset::from_usize(capacity)
.ok_or_else(|| ArrowError::OffsetOverflowError(capacity))?;
source_ranges.push((start, end));
last_filled = i + 1;
}
let final_offset = T::Offset::from_usize(capacity)
.ok_or_else(|| ArrowError::OffsetOverflowError(capacity))?;
offsets[last_filled + 1..].fill(final_offset);
values.reserve(capacity);
debug_assert_eq!(
source_ranges.iter().map(|(s, e)| e - s).sum::<usize>(),
capacity,
"capacity must equal total bytes across all ranges"
);
let src = array.value_data();
let src = src.as_ptr();
let dst = values.spare_capacity_mut();
debug_assert!(dst.len() >= capacity);
let mut offset = 0;
for (start, end) in source_ranges {
let value_len = end - start;
unsafe {
std::ptr::copy_nonoverlapping(
src.add(start),
dst.get_unchecked_mut(offset..).as_mut_ptr().cast::<u8>(),
value_len,
);
offset += value_len;
}
}
unsafe { values.set_len(capacity) };
}
}
let array = unsafe {
let offsets = OffsetBuffer::new_unchecked(offsets.into());
GenericByteArray::<T>::new_unchecked(offsets, values.into(), nulls)
};
Ok(array)
}
fn take_byte_view<T: ByteViewType, IndexType: ArrowPrimitiveType, const CHECKED: 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 buffers = Arc::clone(array.data_buffers());
Ok(unsafe { GenericByteViewArray::new_unchecked(new_views, buffers, new_nulls) })
}
fn take_list<IndexType, OffsetType, const CHECKED: bool>(
values: &GenericListArray<OffsetType::Native>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericListArray<OffsetType::Native>, ArrowError>
where
IndexType: ArrowPrimitiveType,
OffsetType: ArrowPrimitiveType,
OffsetType::Native: OffsetSizeTrait,
PrimitiveArray<OffsetType>: From<Vec<OffsetType::Native>>,
{
let src_offsets = values.value_offsets();
let child_data = values.values().to_data();
let nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
let mut dst_offsets = Vec::with_capacity(indices.len() + 1);
dst_offsets.push(OffsetType::Native::zero());
let field = values.value_field().clone();
if child_data.null_count() == 0
&& let Some(bytes_per_value) = child_data.data_type().primitive_width()
{
let values_buf = &child_data.buffers()[0];
let child_buf_offset = child_data.offset() * bytes_per_value;
let avg_row_len = child_data
.len()
.checked_div(values.len().max(1))
.unwrap_or(0);
let mut dst_buf = MutableBuffer::new(
avg_row_len
.saturating_mul(indices.len())
.saturating_mul(bytes_per_value),
);
let mut child_len = OffsetType::Native::zero();
match nulls.as_ref().filter(|n| n.null_count() > 0) {
None => {
for &idx in indices.values() {
let row = idx.as_usize();
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]))
.ok_or_else(|| ArrowError::OffsetOverflowError(child_len.as_usize()))?;
dst_offsets.push(child_len);
}
}
Some(valid) => {
let mut prev = 0;
for vidx in valid.valid_indices() {
if prev < vidx {
dst_offsets.extend(std::iter::repeat_n(child_len, vidx - prev));
}
let row = if CHECKED {
indices.value(vidx).as_usize()
} else {
unsafe { indices.value_unchecked(vidx) }.as_usize()
};
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]))
.ok_or_else(|| ArrowError::OffsetOverflowError(child_len.as_usize()))?;
dst_offsets.push(child_len);
prev = vidx + 1;
}
dst_offsets.extend(std::iter::repeat_n(child_len, indices.len() - prev));
}
}
debug_assert_eq!(
dst_offsets.len(),
indices.len() + 1,
"New offsets was filled under/over the expected capacity"
);
let child = make_array(unsafe {
ArrayData::builder(child_data.data_type().clone())
.len(child_len.as_usize())
.add_buffer(dst_buf.into())
.build_unchecked()
});
let offsets = unsafe { OffsetBuffer::new_unchecked(ScalarBuffer::from(dst_offsets)) };
return GenericListArray::<OffsetType::Native>::try_new(field, offsets, child, nulls);
}
let capacity = child_data
.len()
.checked_div(values.len())
.map(|avg| avg * indices.len())
.unwrap_or_default();
let mut mutable =
MutableArrayData::new(vec![&child_data], child_data.null_count() > 0, capacity);
match nulls.as_ref().filter(|n| n.null_count() > 0) {
None => {
for idx in indices.values() {
let row = idx.as_usize();
mutable.try_extend(
0,
src_offsets[row].as_usize(),
src_offsets[row + 1].as_usize(),
)?;
dst_offsets.push(
OffsetType::Native::from_usize(mutable.len())
.ok_or_else(|| ArrowError::OffsetOverflowError(mutable.len()))?,
);
}
}
Some(valid) => {
let mut last = 0;
for i in valid.valid_indices() {
let current = OffsetType::Native::from_usize(mutable.len())
.ok_or_else(|| ArrowError::OffsetOverflowError(mutable.len()))?;
if last < i {
dst_offsets.extend(std::iter::repeat_n(current, i - last));
}
let row = if CHECKED {
indices.value(i).as_usize()
} else {
unsafe { indices.value_unchecked(i) }.as_usize()
};
mutable.try_extend(
0,
src_offsets[row].as_usize(),
src_offsets[row + 1].as_usize(),
)?;
dst_offsets.push(
OffsetType::Native::from_usize(mutable.len())
.ok_or_else(|| ArrowError::OffsetOverflowError(mutable.len()))?,
);
last = i + 1;
}
let final_offset = OffsetType::Native::from_usize(mutable.len())
.ok_or_else(|| ArrowError::OffsetOverflowError(mutable.len()))?;
dst_offsets.extend(std::iter::repeat_n(final_offset, indices.len() - last));
}
}
debug_assert_eq!(dst_offsets.len(), indices.len() + 1);
let offsets = unsafe { OffsetBuffer::new_unchecked(ScalarBuffer::from(dst_offsets)) };
let child = make_array(mutable.freeze());
GenericListArray::<OffsetType::Native>::try_new(field, offsets, child, nulls)
}
fn take_list_view<IndexType, OffsetType, const CHECKED: bool>(
values: &GenericListViewArray<OffsetType::Native>,
indices: &PrimitiveArray<IndexType>,
) -> Result<GenericListViewArray<OffsetType::Native>, ArrowError>
where
IndexType: ArrowPrimitiveType,
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 field = match values.data_type() {
DataType::ListView(field) | DataType::LargeListView(field) => field.clone(),
d => unreachable!("take_list_view called with non-list-view data type {d}"),
};
Ok(unsafe {
GenericListViewArray::<OffsetType::Native>::new_unchecked(
field,
taken_offsets,
taken_sizes,
Arc::clone(values.values()),
nulls,
)
})
}
fn take_fixed_size_list<IndexType: ArrowPrimitiveType, const CHECKED: 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 taken_child = if child.null_count() == 0
&& let Some(element_size) = child.data_type().primitive_width()
{
take_fixed_size_list_primitive(
child,
indices,
length as usize,
element_size,
nulls.as_ref(),
)
} else {
let list_indices = take_value_indices_from_fixed_size_list(values, indices, length)?;
take_impl::<UInt32Type, CHECKED>(child.as_ref(), &list_indices)?
};
FixedSizeListArray::try_new_with_length(
field.clone(),
length as i32,
taken_child,
nulls,
indices.len(),
)
}
#[inline(never)]
fn take_fixed_size_list_primitive<IndexType: ArrowPrimitiveType>(
child: &ArrayRef,
indices: &PrimitiveArray<IndexType>,
list_size: usize,
element_size: usize,
taken_row_nulls: Option<&NullBuffer>,
) -> ArrayRef {
let row_bytes = list_size * element_size;
let child_data = child.to_data();
let src = child_data.buffers()[0].as_slice();
let child_byte_offset = child_data.offset() * element_size;
debug_assert!(
indices.len().checked_mul(row_bytes).is_some(),
"take_fixed_size_list_primitive: output buffer size overflows usize"
);
let out_len = indices.len() * list_size;
let mut out = MutableBuffer::from_len_zeroed(indices.len() * row_bytes);
let out_slice = out.as_slice_mut();
if indices.null_count() == 0 {
for (out_row, index) in indices.values().iter().enumerate() {
let src_start = child_byte_offset + index.as_usize() * row_bytes;
out_slice[out_row * row_bytes..(out_row + 1) * row_bytes]
.copy_from_slice(&src[src_start..src_start + row_bytes]);
}
} else {
for (out_row, index) in indices.values().iter().enumerate() {
if indices.is_valid(out_row) {
let src_start = child_byte_offset + index.as_usize() * row_bytes;
out_slice[out_row * row_bytes..(out_row + 1) * row_bytes]
.copy_from_slice(&src[src_start..src_start + row_bytes]);
}
}
}
let child_null_buf = taken_row_nulls.map(|n| n.expand(list_size).buffer().clone());
let array_data = unsafe {
ArrayData::builder(child.data_type().clone())
.len(out_len)
.add_buffer(out.into())
.null_bit_buffer(child_null_buf)
.build_unchecked()
};
make_array(array_data)
}
fn take_fixed_size_binary<IndexType: ArrowPrimitiveType, const CHECKED: bool>(
values: &FixedSizeBinaryArray,
indices: &PrimitiveArray<IndexType>,
size: i32,
) -> Result<FixedSizeBinaryArray, ArrowError> {
let size_usize = usize::try_from(size).map_err(|_| {
ArrowError::InvalidArgumentError(format!("Cannot convert size '{size}' to usize"))
})?;
let result_buffer = match size_usize {
1 => take_fixed_size::<IndexType, 1>(values.values(), indices),
2 => take_fixed_size::<IndexType, 2>(values.values(), indices),
4 => take_fixed_size::<IndexType, 4>(values.values(), indices),
8 => take_fixed_size::<IndexType, 8>(values.values(), indices),
16 => take_fixed_size::<IndexType, 16>(values.values(), indices),
_ => take_fixed_size_binary_buffer_dynamic_length(values, indices, size_usize),
};
let value_nulls = take_nulls::<_, CHECKED>(values.nulls(), indices);
let final_nulls = NullBuffer::union(value_nulls.as_ref(), indices.nulls());
return FixedSizeBinaryArray::try_new(size, result_buffer, final_nulls);
#[inline(never)]
fn take_fixed_size_binary_buffer_dynamic_length<IndexType: ArrowPrimitiveType>(
values: &FixedSizeBinaryArray,
indices: &PrimitiveArray<IndexType>,
size_usize: usize,
) -> Buffer {
let values_buffer = values.values().as_slice();
let mut output = Vec::with_capacity(indices.len() * size_usize);
if indices.null_count() == 0 {
let array_iter = indices.values().iter().map(|idx| {
let offset = idx.as_usize() * size_usize;
&values_buffer[offset..offset + size_usize]
});
for slice in array_iter {
output.extend_from_slice(slice);
}
} else {
let array_iter = indices.iter().map(|idx| {
idx.map(|idx| {
let offset = idx.as_usize() * size_usize;
&values_buffer[offset..offset + size_usize]
})
});
for slice in array_iter {
match slice {
None => output.resize(output.len() + size_usize, 0),
Some(slice) => output.extend_from_slice(slice),
}
}
}
output.into()
}
}
fn take_fixed_size<IndexType: ArrowPrimitiveType, const N: usize>(
buffer: &Buffer,
indices: &PrimitiveArray<IndexType>,
) -> Buffer {
assert_eq!(
buffer.len() % N,
0,
"Invalid array length in take_fixed_size"
);
let ptr = buffer.as_ptr();
let chunk_ptr = ptr.cast::<[u8; N]>();
let chunk_len = buffer.len() / N;
let buffer: &[[u8; N]] = unsafe {
std::slice::from_raw_parts(chunk_ptr, chunk_len)
};
let result_buffer = match indices.nulls().filter(|n| n.null_count() > 0) {
Some(n) => indices
.values()
.iter()
.enumerate()
.map(|(idx, index)| match buffer.get(index.as_usize()) {
Some(v) => *v,
None => match unsafe { n.inner().value_unchecked(idx) } {
false => [0u8; N],
true => panic!("Out-of-bounds index {index:?}"),
},
})
.collect::<Vec<_>>(),
None => indices
.values()
.iter()
.map(|index| buffer[index.as_usize()])
.collect::<Vec<_>>(),
};
let mut vec = ManuallyDrop::new(result_buffer); let ptr = vec.as_mut_ptr();
let len = vec.len();
let cap = vec.capacity();
let result_buffer = unsafe {
Vec::from_raw_parts(ptr.cast::<u8>(), len * N, cap * N)
};
Buffer::from_vec(result_buffer)
}
fn take_dict<T: ArrowDictionaryKeyType, I: ArrowPrimitiveType, const CHECKED: bool>(
values: &DictionaryArray<T>,
indices: &PrimitiveArray<I>,
) -> Result<DictionaryArray<T>, ArrowError> {
let new_keys = take_primitive::<_, _, CHECKED>(values.keys(), indices)?;
Ok(unsafe { DictionaryArray::new_unchecked(new_keys, values.values().clone()) })
}
fn take_run<T: RunEndIndexType, I: ArrowPrimitiveType>(
run_array: &RunArray<T>,
logical_indices: &PrimitiveArray<I>,
) -> Result<RunArray<T>, ArrowError> {
let physical_indices = physical_indices_for_take(run_array, logical_indices)?;
let mut new_run_ends = Vec::with_capacity(1);
let mut take_value_indices = Vec::with_capacity(1);
let mut take_value_is_valid = NullBufferBuilder::new(1);
let values_cmp = make_comparator(
run_array.values().as_ref(),
run_array.values().as_ref(),
SortOptions::default(),
)?;
for ix in 1..physical_indices.len() {
let prev_idx = physical_indices[ix - 1];
let cur_idx = physical_indices[ix];
if is_new_run_take(run_array.values().as_ref(), prev_idx, cur_idx, &values_cmp) {
let index = I::Native::from_usize(prev_idx.unwrap_or_default()).unwrap();
take_value_indices.push(index);
take_value_is_valid
.append(prev_idx.is_some_and(|idx| run_array.values().is_valid(idx)));
new_run_ends.push(T::Native::from_usize(ix).unwrap());
}
}
let last = physical_indices[physical_indices.len() - 1];
let index = I::Native::from_usize(last.unwrap_or_default()).unwrap();
take_value_indices.push(index);
take_value_is_valid.append(last.is_some_and(|idx| run_array.values().is_valid(idx)));
new_run_ends.push(T::Native::from_usize(physical_indices.len()).unwrap());
let run_ends = unsafe {
RunEndBuffer::new_unchecked(ScalarBuffer::from(new_run_ends), 0, physical_indices.len())
};
let nulls = take_value_is_valid.finish();
let take_value_indices =
PrimitiveArray::<I>::new(ScalarBuffer::from(take_value_indices), nulls);
let new_values = take(run_array.values(), &take_value_indices, None)?;
Ok(
unsafe {
RunArray::<T>::new_unchecked(run_array.data_type().clone(), run_ends, new_values)
},
)
}
fn physical_indices_for_take<T: RunEndIndexType, I: ArrowPrimitiveType>(
run_array: &RunArray<T>,
logical_indices: &PrimitiveArray<I>,
) -> Result<Vec<Option<usize>>, ArrowError> {
if logical_indices.null_count() == 0 {
return Ok(run_array
.get_physical_indices(logical_indices.values())?
.into_iter()
.map(Some)
.collect());
}
let valid_logical: Vec<_> = logical_indices.iter().flatten().collect();
let valid_physical = if valid_logical.is_empty() {
Vec::new()
} else {
run_array.get_physical_indices(&valid_logical)?
};
let mut valid_physical = valid_physical.into_iter();
Ok(logical_indices
.iter()
.map(|index| index.map(|_| valid_physical.next().unwrap()))
.collect())
}
fn is_new_run_take(
values: &dyn Array,
prev_idx: Option<usize>,
cur_idx: Option<usize>,
values_cmp: &arrow_cmp::DynComparator,
) -> bool {
let prev_valid = prev_idx.is_some_and(|idx| values.is_valid(idx));
let cur_valid = cur_idx.is_some_and(|idx| values.is_valid(idx));
match (prev_valid, cur_valid) {
(false, false) => false,
(true, true) => {
let prev = prev_idx.unwrap();
let cur = cur_idx.unwrap();
prev != cur && values_cmp(cur, prev).is_ne()
}
_ => true,
}
}
fn take_value_indices_from_fixed_size_list<IndexType>(
list: &FixedSizeListArray,
indices: &PrimitiveArray<IndexType>,
length: <UInt32Type as ArrowPrimitiveType>::Native,
) -> Result<PrimitiveArray<UInt32Type>, ArrowError>
where
IndexType: ArrowPrimitiveType,
{
let mut values = UInt32Builder::with_capacity(length as usize * indices.len());
for i in 0..indices.len() {
if indices.is_valid(i) {
let index = indices
.value(i)
.to_usize()
.ok_or_else(|| ArrowError::ComputeError("Cast to usize failed".to_string()))?;
let start = list.value_offset(index) as <UInt32Type as ArrowPrimitiveType>::Native;
unsafe {
values.append_trusted_len_iter(start..start + length);
}
} else {
values.append_nulls(length as usize);
}
}
Ok(values.finish())
}
trait ToIndices {
type T: ArrowPrimitiveType;
fn to_indices(&self) -> PrimitiveArray<Self::T>;
}
macro_rules! to_indices_reinterpret {
($t:ty, $o:ty) => {
impl ToIndices for PrimitiveArray<$t> {
type T = $o;
fn to_indices(&self) -> PrimitiveArray<$o> {
let cast = ScalarBuffer::new(self.values().inner().clone(), 0, self.len());
PrimitiveArray::new(cast, self.nulls().cloned())
}
}
};
}
macro_rules! to_indices_identity {
($t:ty) => {
impl ToIndices for PrimitiveArray<$t> {
type T = $t;
fn to_indices(&self) -> PrimitiveArray<$t> {
self.clone()
}
}
};
}
macro_rules! to_indices_widening {
($t:ty, $o:ty) => {
impl ToIndices for PrimitiveArray<$t> {
type T = UInt32Type;
fn to_indices(&self) -> PrimitiveArray<$o> {
let cast = self.values().iter().copied().map(|x| x as _).collect();
PrimitiveArray::new(cast, self.nulls().cloned())
}
}
};
}
to_indices_widening!(UInt8Type, UInt32Type);
to_indices_widening!(Int8Type, UInt32Type);
to_indices_widening!(UInt16Type, UInt32Type);
to_indices_widening!(Int16Type, UInt32Type);
to_indices_identity!(UInt32Type);
to_indices_reinterpret!(Int32Type, UInt32Type);
to_indices_identity!(UInt64Type);
to_indices_reinterpret!(Int64Type, UInt64Type);
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)
}
#[cfg(test)]
mod tests {
use super::*;
use arrow_array::builder::*;
use arrow_buffer::{IntervalDayTime, IntervalMonthDayNano};
use arrow_data::ArrayData;
use arrow_schema::{Field, Fields, TimeUnit, UnionFields};
use num_traits::ToPrimitive;
fn test_take_decimal_arrays(
data: Vec<Option<i128>>,
index: &UInt32Array,
options: Option<TakeOptions>,
expected_data: Vec<Option<i128>>,
precision: &u8,
scale: &i8,
) -> Result<(), ArrowError> {
let output = data
.into_iter()
.collect::<Decimal128Array>()
.with_precision_and_scale(*precision, *scale)
.unwrap();
let expected = expected_data
.into_iter()
.collect::<Decimal128Array>()
.with_precision_and_scale(*precision, *scale)
.unwrap();
let expected = Arc::new(expected) as ArrayRef;
let output = take(&output, index, options).unwrap();
assert_eq!(&output, &expected);
Ok(())
}
fn test_take_boolean_arrays(
data: Vec<Option<bool>>,
index: &UInt32Array,
options: Option<TakeOptions>,
expected_data: Vec<Option<bool>>,
) {
let output = BooleanArray::from(data);
let expected = Arc::new(BooleanArray::from(expected_data)) as ArrayRef;
let output = take(&output, index, options).unwrap();
assert_eq!(&output, &expected)
}
fn test_take_primitive_arrays<T>(
data: Vec<Option<T::Native>>,
index: &UInt32Array,
options: Option<TakeOptions>,
expected_data: Vec<Option<T::Native>>,
) -> Result<(), ArrowError>
where
T: ArrowPrimitiveType,
PrimitiveArray<T>: From<Vec<Option<T::Native>>>,
{
let output = PrimitiveArray::<T>::from(data);
let expected = Arc::new(PrimitiveArray::<T>::from(expected_data)) as ArrayRef;
let output = take(&output, index, options)?;
assert_eq!(&output, &expected);
Ok(())
}
fn test_take_primitive_arrays_non_null<T>(
data: Vec<T::Native>,
index: &UInt32Array,
options: Option<TakeOptions>,
expected_data: Vec<Option<T::Native>>,
) -> Result<(), ArrowError>
where
T: ArrowPrimitiveType,
PrimitiveArray<T>: From<Vec<T::Native>>,
PrimitiveArray<T>: From<Vec<Option<T::Native>>>,
{
let output = PrimitiveArray::<T>::from(data);
let expected = Arc::new(PrimitiveArray::<T>::from(expected_data)) as ArrayRef;
let output = take(&output, index, options)?;
assert_eq!(&output, &expected);
Ok(())
}
fn test_take_impl_primitive_arrays<T, I>(
data: Vec<Option<T::Native>>,
index: &PrimitiveArray<I>,
options: Option<TakeOptions>,
expected_data: Vec<Option<T::Native>>,
) where
T: ArrowPrimitiveType,
PrimitiveArray<T>: From<Vec<Option<T::Native>>>,
I: ArrowPrimitiveType,
{
let output = PrimitiveArray::<T>::from(data);
let expected = PrimitiveArray::<T>::from(expected_data);
let output = take(&output, index, options).unwrap();
let output = output.as_any().downcast_ref::<PrimitiveArray<T>>().unwrap();
assert_eq!(output, &expected)
}
fn create_test_struct(values: Vec<Option<(Option<bool>, Option<i32>)>>) -> StructArray {
let mut struct_builder = StructBuilder::new(
Fields::from(vec![
Field::new("a", DataType::Boolean, true),
Field::new("b", DataType::Int32, true),
]),
vec![
Box::new(BooleanBuilder::with_capacity(values.len())),
Box::new(Int32Builder::with_capacity(values.len())),
],
);
for value in values {
struct_builder
.field_builder::<BooleanBuilder>(0)
.unwrap()
.append_option(value.and_then(|v| v.0));
struct_builder
.field_builder::<Int32Builder>(1)
.unwrap()
.append_option(value.and_then(|v| v.1));
struct_builder.append(value.is_some());
}
struct_builder.finish()
}
#[test]
fn test_take_decimal128_non_null_indices() {
let index = UInt32Array::from(vec![0, 5, 3, 1, 4, 2]);
let precision: u8 = 10;
let scale: i8 = 5;
test_take_decimal_arrays(
vec![None, Some(3), Some(5), Some(2), Some(3), None],
&index,
None,
vec![None, None, Some(2), Some(3), Some(3), Some(5)],
&precision,
&scale,
)
.unwrap();
}
#[test]
fn test_take_decimal128() {
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(2)]);
let precision: u8 = 10;
let scale: i8 = 5;
test_take_decimal_arrays(
vec![Some(0), Some(1), Some(2), Some(3), Some(4)],
&index,
None,
vec![Some(3), None, Some(1), Some(3), Some(2)],
&precision,
&scale,
)
.unwrap();
}
#[test]
fn test_take_primitive_non_null_indices() {
let index = UInt32Array::from(vec![0, 5, 3, 1, 4, 2]);
test_take_primitive_arrays::<Int8Type>(
vec![None, Some(3), Some(5), Some(2), Some(3), None],
&index,
None,
vec![None, None, Some(2), Some(3), Some(3), Some(5)],
)
.unwrap();
}
#[test]
fn test_take_primitive_non_null_values() {
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(2)]);
test_take_primitive_arrays::<Int8Type>(
vec![Some(0), Some(1), Some(2), Some(3), Some(4)],
&index,
None,
vec![Some(3), None, Some(1), Some(3), Some(2)],
)
.unwrap();
}
#[test]
fn test_take_primitive_non_null() {
let index = UInt32Array::from(vec![0, 5, 3, 1, 4, 2]);
test_take_primitive_arrays::<Int8Type>(
vec![Some(0), Some(3), Some(5), Some(2), Some(3), Some(1)],
&index,
None,
vec![Some(0), Some(1), Some(2), Some(3), Some(3), Some(5)],
)
.unwrap();
}
#[test]
fn test_take_primitive_nullable_indices_non_null_values_with_offset() {
let index = UInt32Array::from(vec![Some(0), Some(1), Some(2), Some(3), None, None]);
let index = index.slice(2, 4);
let index = index.as_any().downcast_ref::<UInt32Array>().unwrap();
assert_eq!(
index,
&UInt32Array::from(vec![Some(2), Some(3), None, None])
);
test_take_primitive_arrays_non_null::<Int64Type>(
vec![0, 10, 20, 30, 40, 50],
index,
None,
vec![Some(20), Some(30), None, None],
)
.unwrap();
}
#[test]
fn test_take_primitive_nullable_indices_nullable_values_with_offset() {
let index = UInt32Array::from(vec![Some(0), Some(1), Some(2), Some(3), None, None]);
let index = index.slice(2, 4);
let index = index.as_any().downcast_ref::<UInt32Array>().unwrap();
assert_eq!(
index,
&UInt32Array::from(vec![Some(2), Some(3), None, None])
);
test_take_primitive_arrays::<Int64Type>(
vec![None, None, Some(20), Some(30), Some(40), Some(50)],
index,
None,
vec![Some(20), Some(30), None, None],
)
.unwrap();
}
#[test]
fn test_take_primitive() {
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(2)]);
test_take_primitive_arrays::<Int8Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<Int16Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<Int32Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<Int64Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<UInt8Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<UInt16Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<UInt32Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<Int64Type>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<IntervalYearMonthType>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
)
.unwrap();
let v1 = IntervalDayTime::new(0, 0);
let v2 = IntervalDayTime::new(2, 0);
let v3 = IntervalDayTime::new(-15, 0);
test_take_primitive_arrays::<IntervalDayTimeType>(
vec![Some(v1), None, Some(v2), Some(v3), None],
&index,
None,
vec![Some(v3), None, None, Some(v3), Some(v2)],
)
.unwrap();
let v1 = IntervalMonthDayNano::new(0, 0, 0);
let v2 = IntervalMonthDayNano::new(2, 0, 0);
let v3 = IntervalMonthDayNano::new(-15, 0, 0);
test_take_primitive_arrays::<IntervalMonthDayNanoType>(
vec![Some(v1), None, Some(v2), Some(v3), None],
&index,
None,
vec![Some(v3), None, None, Some(v3), Some(v2)],
)
.unwrap();
test_take_primitive_arrays::<DurationSecondType>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<DurationMillisecondType>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<DurationMicrosecondType>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<DurationNanosecondType>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
)
.unwrap();
test_take_primitive_arrays::<Float32Type>(
vec![Some(0.0), None, Some(2.21), Some(-3.1), None],
&index,
None,
vec![Some(-3.1), None, None, Some(-3.1), Some(2.21)],
)
.unwrap();
test_take_primitive_arrays::<Float64Type>(
vec![Some(0.0), None, Some(2.21), Some(-3.1), None],
&index,
None,
vec![Some(-3.1), None, None, Some(-3.1), Some(2.21)],
)
.unwrap();
}
#[test]
fn test_take_preserve_timezone() {
let index = Int64Array::from(vec![Some(0), None]);
let input = TimestampNanosecondArray::from(vec![
1_639_715_368_000_000_000,
1_639_715_368_000_000_000,
])
.with_timezone("UTC".to_string());
let result = take(&input, &index, None).unwrap();
match result.data_type() {
DataType::Timestamp(TimeUnit::Nanosecond, tz) => {
assert_eq!(tz.clone(), Some("UTC".into()))
}
_ => panic!(),
}
}
#[test]
fn test_take_impl_primitive_with_int64_indices() {
let index = Int64Array::from(vec![Some(3), None, Some(1), Some(3), Some(2)]);
test_take_impl_primitive_arrays::<Int16Type, Int64Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
);
test_take_impl_primitive_arrays::<Int64Type, Int64Type>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
);
test_take_impl_primitive_arrays::<UInt64Type, Int64Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
);
test_take_impl_primitive_arrays::<DurationMillisecondType, Int64Type>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
);
test_take_impl_primitive_arrays::<Float32Type, Int64Type>(
vec![Some(0.0), None, Some(2.21), Some(-3.1), None],
&index,
None,
vec![Some(-3.1), None, None, Some(-3.1), Some(2.21)],
);
}
#[test]
fn test_take_impl_primitive_with_uint8_indices() {
let index = UInt8Array::from(vec![Some(3), None, Some(1), Some(3), Some(2)]);
test_take_impl_primitive_arrays::<Int16Type, UInt8Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
None,
vec![Some(3), None, None, Some(3), Some(2)],
);
test_take_impl_primitive_arrays::<DurationMillisecondType, UInt8Type>(
vec![Some(0), None, Some(2), Some(-15), None],
&index,
None,
vec![Some(-15), None, None, Some(-15), Some(2)],
);
test_take_impl_primitive_arrays::<Float32Type, UInt8Type>(
vec![Some(0.0), None, Some(2.21), Some(-3.1), None],
&index,
None,
vec![Some(-3.1), None, None, Some(-3.1), Some(2.21)],
);
}
#[test]
fn test_take_bool() {
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(2)]);
test_take_boolean_arrays(
vec![Some(false), None, Some(true), Some(false), None],
&index,
None,
vec![Some(false), None, None, Some(false), Some(true)],
);
}
#[test]
fn test_take_bool_nullable_index() {
let index_data = ArrayData::try_new(
DataType::UInt32,
6,
Some(Buffer::from_iter(vec![
false, true, false, true, false, true,
])),
0,
vec![Buffer::from_iter(vec![99, 0, 999, 1, 9999, 2])],
vec![],
)
.unwrap();
let index = UInt32Array::from(index_data);
test_take_boolean_arrays(
vec![Some(true), None, Some(false)],
&index,
None,
vec![None, Some(true), None, None, None, Some(false)],
);
}
#[test]
fn test_take_bool_nullable_index_nonnull_values() {
let index_data = ArrayData::try_new(
DataType::UInt32,
6,
Some(Buffer::from_iter(vec![
false, true, false, true, false, true,
])),
0,
vec![Buffer::from_iter(vec![99, 0, 999, 1, 9999, 2])],
vec![],
)
.unwrap();
let index = UInt32Array::from(index_data);
test_take_boolean_arrays(
vec![Some(true), Some(true), Some(false)],
&index,
None,
vec![None, Some(true), None, Some(true), None, Some(false)],
);
}
#[test]
fn test_take_bool_with_offset() {
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(2), None]);
let index = index.slice(2, 4);
let index = index
.as_any()
.downcast_ref::<PrimitiveArray<UInt32Type>>()
.unwrap();
test_take_boolean_arrays(
vec![Some(false), None, Some(true), Some(false), None],
index,
None,
vec![None, Some(false), Some(true), None],
);
}
#[test]
fn test_take_bool_nullable_index_sliced_source() {
let source = BooleanArray::from(vec![Some(true), Some(false), Some(true), Some(false)]);
let source = source.slice(1, 3); let source = source;
let indices = UInt32Array::from(vec![Some(2), None, Some(0)]);
let result = take(&source, &indices, None).unwrap();
let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
let expected = BooleanArray::from(vec![Some(false), None, Some(false)]);
assert_eq!(result, &expected);
}
#[test]
fn test_take_bool_no_nulls_multi_byte() {
let source = BooleanArray::from(vec![
true, false, true, true, false, false, true, false, true, true,
]);
let indices = UInt32Array::from(vec![0, 2, 4, 6, 8, 1, 3, 5, 7, 9]);
let result = take(&source, &indices, None).unwrap();
let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
let expected = BooleanArray::from(vec![
true, true, false, true, true, false, true, false, false, true,
]);
assert_eq!(result, &expected);
}
#[test]
fn test_take_bool_no_nulls_sliced_source() {
let source = BooleanArray::from(vec![true, false, true, false, true]);
let source = source.slice(2, 3); let source = source.as_any().downcast_ref::<BooleanArray>().unwrap();
let indices = UInt32Array::from(vec![2, 0, 1]);
let result = take(source, &indices, None).unwrap();
let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
let expected = BooleanArray::from(vec![true, true, false]);
assert_eq!(result, &expected);
}
#[test]
fn test_take_bool_nullable_values_multi_byte() {
let source = BooleanArray::from(vec![
Some(true),
None,
Some(false),
Some(true),
None,
Some(false),
Some(true),
Some(false),
Some(true),
None,
]);
let indices = UInt32Array::from(vec![0, 2, 4, 6, 8, 1, 3, 5, 7, 9]);
let result = take(&source, &indices, None).unwrap();
let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
let expected = BooleanArray::from(vec![
Some(true),
Some(false),
None,
Some(true),
Some(true),
None,
Some(true),
Some(false),
Some(false),
None,
]);
assert_eq!(result, &expected);
}
#[test]
fn test_take_bool_nullable_values_sliced_source_null_indices() {
let source =
BooleanArray::from(vec![Some(true), Some(false), None, Some(true), Some(false)]);
let source = source.slice(1, 4); let source = source.as_any().downcast_ref::<BooleanArray>().unwrap();
let indices = UInt32Array::from(vec![Some(3), None, Some(1), Some(0)]);
let result = take(source, &indices, None).unwrap();
let result = result.as_any().downcast_ref::<BooleanArray>().unwrap();
let expected = BooleanArray::from(vec![Some(false), None, None, Some(false)]);
assert_eq!(result, &expected);
}
#[test]
#[should_panic(expected = "assertion failed: idx < self.bit_len")]
fn test_take_bool_oob_no_check_bounds_panics() {
let array = BooleanArray::from(vec![true, false, true]);
let indices = Int32Array::from(vec![0, 1, 10]);
take(&array, &indices, None).unwrap();
}
fn _test_take_string<'a, K>()
where
K: Array + PartialEq + From<Vec<Option<&'a str>>> + 'static,
{
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(4)]);
let array = K::from(vec![
Some("one"),
None,
Some("three"),
Some("four"),
Some("five"),
]);
let actual = take(&array, &index, None).unwrap();
assert_eq!(actual.len(), index.len());
let actual = actual.as_any().downcast_ref::<K>().unwrap();
let expected = K::from(vec![Some("four"), None, None, Some("four"), Some("five")]);
assert_eq!(actual, &expected);
}
#[test]
fn test_take_string() {
_test_take_string::<StringArray>()
}
#[test]
fn test_take_large_string() {
_test_take_string::<LargeStringArray>()
}
#[test]
fn test_take_slice_string() {
let strings = StringArray::from(vec![Some("hello"), None, Some("world"), None, Some("hi")]);
let indices = Int32Array::from(vec![Some(0), Some(1), None, Some(0), Some(2)]);
let indices_slice = indices.slice(1, 4);
let expected = StringArray::from(vec![None, None, Some("hello"), Some("world")]);
let result = take(&strings, &indices_slice, None).unwrap();
assert_eq!(result.as_ref(), &expected);
}
#[test]
fn test_take_bytes_sliced_values() {
let values = StringArray::from(vec![
Some("aaa"),
Some("bbb"),
None,
Some("ccccc"),
Some("dd"),
None,
Some("eeee"),
]);
let sliced = values.slice(2, 5);
let indices = Int32Array::from(vec![1, 2, 4, 1]);
let result = take(&sliced, &indices, None).unwrap();
let expected =
StringArray::from(vec![Some("ccccc"), Some("dd"), Some("eeee"), Some("ccccc")]);
assert_eq!(result.as_string::<i32>(), &expected);
let indices = Int32Array::from(vec![Some(1), None, Some(0), Some(4), Some(3)]);
let result = take(&sliced, &indices, None).unwrap();
let expected = StringArray::from(vec![Some("ccccc"), None, None, Some("eeee"), None]);
assert_eq!(result.as_string::<i32>(), &expected);
}
fn _test_byte_view<T>()
where
T: ByteViewType,
str: AsRef<T::Native>,
T::Native: PartialEq,
{
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(4), Some(2)]);
let array = {
let mut builder = GenericByteViewBuilder::<T>::new();
builder.append_value("hello");
builder.append_value("world");
builder.append_null();
builder.append_value("large payload over 12 bytes");
builder.append_value("lulu");
builder.finish()
};
let actual = take(&array, &index, None).unwrap();
assert_eq!(actual.len(), index.len());
let actual_buffers = actual.as_byte_view::<T>().data_buffers();
let input_buffers = array.data_buffers();
assert!(Arc::ptr_eq(actual_buffers, input_buffers));
let expected = {
let mut builder = GenericByteViewBuilder::<T>::new();
builder.append_value("large payload over 12 bytes");
builder.append_null();
builder.append_value("world");
builder.append_value("large payload over 12 bytes");
builder.append_value("lulu");
builder.append_null();
builder.finish()
};
assert_eq!(actual.as_ref(), &expected);
}
#[test]
fn test_take_string_view() {
_test_byte_view::<StringViewType>()
}
#[test]
fn test_take_binary_view() {
_test_byte_view::<BinaryViewType>()
}
macro_rules! test_take_list {
($offset_type:ty, $list_data_type:ident, $list_array_type:ident) => {{
let value_data = Int32Array::from(vec![0, 0, 0, -1, -2, -1, 2, 3]).into_data();
let value_offsets: [$offset_type; 5] = [0, 3, 6, 6, 8];
let value_offsets = Buffer::from_slice_ref(&value_offsets);
let list_data_type =
DataType::$list_data_type(Arc::new(Field::new_list_field(DataType::Int32, false)));
let list_data = ArrayData::builder(list_data_type.clone())
.len(4)
.add_buffer(value_offsets)
.add_child_data(value_data)
.build()
.unwrap();
let list_array = $list_array_type::from(list_data);
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(2), Some(0)]);
let a = take(&list_array, &index, None).unwrap();
let a: &$list_array_type = a.as_any().downcast_ref::<$list_array_type>().unwrap();
let expected_data = Int32Array::from(vec![
Some(2),
Some(3),
Some(-1),
Some(-2),
Some(-1),
Some(0),
Some(0),
Some(0),
])
.into_data();
let expected_offsets: [$offset_type; 6] = [0, 2, 2, 5, 5, 8];
let expected_offsets = Buffer::from_slice_ref(&expected_offsets);
let expected_list_data = ArrayData::builder(list_data_type)
.len(5)
.nulls(index.nulls().cloned())
.add_buffer(expected_offsets)
.add_child_data(expected_data)
.build()
.unwrap();
let expected_list_array = $list_array_type::from(expected_list_data);
assert_eq!(a, &expected_list_array);
}};
}
macro_rules! test_take_list_with_value_nulls {
($offset_type:ty, $list_data_type:ident, $list_array_type:ident) => {{
let value_data = Int32Array::from(vec![
Some(0),
None,
Some(0),
Some(-1),
Some(-2),
Some(3),
None,
Some(5),
None,
])
.into_data();
let value_offsets: [$offset_type; 5] = [0, 3, 6, 7, 9];
let value_offsets = Buffer::from_slice_ref(&value_offsets);
let list_data_type =
DataType::$list_data_type(Arc::new(Field::new_list_field(DataType::Int32, true)));
let list_data = ArrayData::builder(list_data_type.clone())
.len(4)
.add_buffer(value_offsets)
.null_bit_buffer(Some(Buffer::from([0b11111111])))
.add_child_data(value_data)
.build()
.unwrap();
let list_array = $list_array_type::from(list_data);
let index = UInt32Array::from(vec![Some(2), None, Some(1), Some(3), Some(0)]);
let a = take(&list_array, &index, None).unwrap();
let a: &$list_array_type = a.as_any().downcast_ref::<$list_array_type>().unwrap();
let expected_data = Int32Array::from(vec![
None,
Some(-1),
Some(-2),
Some(3),
Some(5),
None,
Some(0),
None,
Some(0),
])
.into_data();
let expected_offsets: [$offset_type; 6] = [0, 1, 1, 4, 6, 9];
let expected_offsets = Buffer::from_slice_ref(&expected_offsets);
let expected_list_data = ArrayData::builder(list_data_type)
.len(5)
.nulls(index.nulls().cloned())
.add_buffer(expected_offsets)
.add_child_data(expected_data)
.build()
.unwrap();
let expected_list_array = $list_array_type::from(expected_list_data);
assert_eq!(a, &expected_list_array);
}};
}
macro_rules! test_take_list_with_nulls {
($offset_type:ty, $list_data_type:ident, $list_array_type:ident) => {{
let value_data = Int32Array::from(vec![
Some(0),
None,
Some(0),
Some(-1),
Some(-2),
Some(3),
Some(5),
None,
])
.into_data();
let value_offsets: [$offset_type; 5] = [0, 3, 6, 6, 8];
let value_offsets = Buffer::from_slice_ref(&value_offsets);
let list_data_type =
DataType::$list_data_type(Arc::new(Field::new_list_field(DataType::Int32, true)));
let list_data = ArrayData::builder(list_data_type.clone())
.len(4)
.add_buffer(value_offsets)
.null_bit_buffer(Some(Buffer::from([0b11111011])))
.add_child_data(value_data)
.build()
.unwrap();
let list_array = $list_array_type::from(list_data);
let index = UInt32Array::from(vec![Some(2), None, Some(1), Some(3), Some(0)]);
let a = take(&list_array, &index, None).unwrap();
let a: &$list_array_type = a.as_any().downcast_ref::<$list_array_type>().unwrap();
let expected_data = Int32Array::from(vec![
Some(-1),
Some(-2),
Some(3),
Some(5),
None,
Some(0),
None,
Some(0),
])
.into_data();
let expected_offsets: [$offset_type; 6] = [0, 0, 0, 3, 5, 8];
let expected_offsets = Buffer::from_slice_ref(&expected_offsets);
let mut null_bits: [u8; 1] = [0; 1];
bit_util::set_bit(&mut null_bits, 2);
bit_util::set_bit(&mut null_bits, 3);
bit_util::set_bit(&mut null_bits, 4);
let expected_list_data = ArrayData::builder(list_data_type)
.len(5)
.null_bit_buffer(Some(Buffer::from(null_bits)))
.add_buffer(expected_offsets)
.add_child_data(expected_data)
.build()
.unwrap();
let expected_list_array = $list_array_type::from(expected_list_data);
assert_eq!(a, &expected_list_array);
}};
}
fn test_take_list_view_generic<OffsetType: OffsetSizeTrait, ValuesType: ArrowPrimitiveType, F>(
values: Vec<Option<Vec<Option<ValuesType::Native>>>>,
take_indices: Vec<Option<usize>>,
expected: Vec<Option<Vec<Option<ValuesType::Native>>>>,
mapper: F,
) where
F: Fn(GenericListViewArray<OffsetType>) -> GenericListViewArray<OffsetType>,
{
let mut list_view_array =
GenericListViewBuilder::<OffsetType, _>::new(PrimitiveBuilder::<ValuesType>::new());
for value in values {
list_view_array.append_option(value);
}
let list_view_array = list_view_array.finish();
let list_view_array = mapper(list_view_array);
let mut indices = UInt64Builder::new();
for idx in take_indices {
indices.append_option(idx.map(|i| i.to_u64().unwrap()));
}
let indices = indices.finish();
let taken = take(&list_view_array, &indices, None)
.unwrap()
.as_list_view()
.clone();
let mut expected_array =
GenericListViewBuilder::<OffsetType, _>::new(PrimitiveBuilder::<ValuesType>::new());
for value in expected {
expected_array.append_option(value);
}
let expected_array = expected_array.finish();
assert_eq!(taken, expected_array);
}
macro_rules! list_view_test_case {
(values: $values:expr, indices: $indices:expr, expected: $expected: expr) => {{
test_take_list_view_generic::<i32, Int8Type, _>($values, $indices, $expected, |x| x);
test_take_list_view_generic::<i64, Int8Type, _>($values, $indices, $expected, |x| x);
}};
(values: $values:expr, transform: $fn:expr, indices: $indices:expr, expected: $expected: expr) => {{
test_take_list_view_generic::<i32, Int8Type, _>($values, $indices, $expected, $fn);
test_take_list_view_generic::<i64, Int8Type, _>($values, $indices, $expected, $fn);
}};
}
fn do_take_fixed_size_list_test<T>(
length: <Int32Type as ArrowPrimitiveType>::Native,
input_data: Vec<Option<Vec<Option<T::Native>>>>,
indices: Vec<<UInt32Type as ArrowPrimitiveType>::Native>,
expected_data: Vec<Option<Vec<Option<T::Native>>>>,
) where
T: ArrowPrimitiveType,
PrimitiveArray<T>: From<Vec<Option<T::Native>>>,
{
let indices = UInt32Array::from(indices);
let input_array = FixedSizeListArray::from_iter_primitive::<T, _, _>(input_data, length);
let output =
take_fixed_size_list::<_, true>(&input_array, &indices, length as u32).unwrap();
let expected = FixedSizeListArray::from_iter_primitive::<T, _, _>(expected_data, length);
assert_eq!(&output, &expected)
}
#[test]
fn test_take_list_primitive_child_null_indices() {
let list = ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(1)]),
Some(vec![Some(2), Some(3), Some(4)]),
Some(vec![Some(5), Some(6)]),
]);
let indices = Int32Array::from(vec![Some(2), None, Some(0), Some(1)]);
let result = take(&list, &indices, None).unwrap();
let result = result.as_any().downcast_ref::<ListArray>().unwrap();
let expected = ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(5), Some(6)]),
None,
Some(vec![Some(1)]),
Some(vec![Some(2), Some(3), Some(4)]),
]);
assert_eq!(result, &expected);
}
#[test]
fn test_take_list() {
test_take_list!(i32, List, ListArray);
}
#[test]
fn test_take_large_list() {
test_take_list!(i64, LargeList, LargeListArray);
}
#[test]
fn test_take_list_with_value_nulls() {
test_take_list_with_value_nulls!(i32, List, ListArray);
}
#[test]
fn test_take_large_list_with_value_nulls() {
test_take_list_with_value_nulls!(i64, LargeList, LargeListArray);
}
#[test]
fn test_test_take_list_with_nulls() {
test_take_list_with_nulls!(i32, List, ListArray);
}
#[test]
fn test_test_take_large_list_with_nulls() {
test_take_list_with_nulls!(i64, LargeList, LargeListArray);
}
#[test]
fn test_test_take_list_view_reversed() {
list_view_test_case! {
values: vec![
Some(vec![Some(1), None, Some(3)]),
None,
Some(vec![Some(7), Some(8), None]),
],
indices: vec![Some(2), Some(1), Some(0)],
expected: vec![
Some(vec![Some(7), Some(8), None]),
None,
Some(vec![Some(1), None, Some(3)]),
]
}
}
#[test]
fn test_take_list_view_null_indices() {
list_view_test_case! {
values: vec![
Some(vec![Some(1), None, Some(3)]),
None,
Some(vec![Some(7), Some(8), None]),
],
indices: vec![None, Some(0), None],
expected: vec![None, Some(vec![Some(1), None, Some(3)]), None]
}
}
#[test]
fn test_take_list_view_null_values() {
list_view_test_case! {
values: vec![
Some(vec![Some(1), None, Some(3)]),
None,
Some(vec![Some(7), Some(8), None]),
],
indices: vec![Some(1), Some(1), Some(1), None, None],
expected: vec![None; 5]
}
}
#[test]
fn test_take_list_view_sliced() {
list_view_test_case! {
values: vec![
Some(vec![Some(1)]),
None,
None,
Some(vec![Some(2), Some(3)]),
Some(vec![Some(4), Some(5)]),
None,
],
transform: |l| l.slice(2, 4),
indices: vec![Some(0), Some(3), None, Some(1), Some(2)],
expected: vec![
None, None, None, Some(vec![Some(2), Some(3)]), Some(vec![Some(4), Some(5)])
]
}
}
#[test]
fn test_take_fixed_size_list() {
do_take_fixed_size_list_test::<Int32Type>(
3,
vec![
Some(vec![None, Some(1), Some(2)]),
Some(vec![Some(3), Some(4), None]),
Some(vec![Some(6), Some(7), Some(8)]),
],
vec![2, 1, 0],
vec![
Some(vec![Some(6), Some(7), Some(8)]),
Some(vec![Some(3), Some(4), None]),
Some(vec![None, Some(1), Some(2)]),
],
);
do_take_fixed_size_list_test::<UInt8Type>(
1,
vec![
Some(vec![Some(1)]),
Some(vec![Some(2)]),
Some(vec![Some(3)]),
Some(vec![Some(4)]),
Some(vec![Some(5)]),
Some(vec![Some(6)]),
Some(vec![Some(7)]),
Some(vec![Some(8)]),
],
vec![2, 7, 0],
vec![
Some(vec![Some(3)]),
Some(vec![Some(8)]),
Some(vec![Some(1)]),
],
);
do_take_fixed_size_list_test::<UInt64Type>(
3,
vec![
Some(vec![Some(10), Some(11), Some(12)]),
Some(vec![Some(13), Some(14), Some(15)]),
None,
Some(vec![Some(16), Some(17), Some(18)]),
],
vec![3, 2, 1, 2, 0],
vec![
Some(vec![Some(16), Some(17), Some(18)]),
None,
Some(vec![Some(13), Some(14), Some(15)]),
None,
Some(vec![Some(10), Some(11), Some(12)]),
],
);
}
#[test]
fn test_take_fixed_size_binary_with_nulls_indices() {
let fsb = FixedSizeBinaryArray::try_from_sparse_iter_with_size(
[
Some(vec![0x01, 0x01, 0x01, 0x01]),
Some(vec![0x02, 0x02, 0x02, 0x02]),
Some(vec![0x03, 0x03, 0x03, 0x03]),
Some(vec![0x04, 0x04, 0x04, 0x04]),
]
.into_iter(),
4,
)
.unwrap();
let indices = UInt32Array::from(vec![Some(0), None, None, Some(3)]);
let result = take_fixed_size_binary::<_, true>(&fsb, &indices, 4).unwrap();
assert_eq!(result.len(), 4);
assert_eq!(result.null_count(), 2);
assert_eq!(
result.nulls().unwrap().iter().collect::<Vec<_>>(),
vec![true, false, false, true]
);
}
#[test]
fn test_take_fixed_size_binary_with_nulls_indices_not_optimized_length() {
let fsb = FixedSizeBinaryArray::try_from_sparse_iter_with_size(
[
Some(vec![0x01, 0x01, 0x01, 0x01, 0x01]),
Some(vec![0x02, 0x02, 0x02, 0x02, 0x01]),
Some(vec![0x03, 0x03, 0x03, 0x03, 0x01]),
Some(vec![0x04, 0x04, 0x04, 0x04, 0x01]),
]
.into_iter(),
5,
)
.unwrap();
let indices = UInt32Array::from(vec![Some(0), None, None, Some(3)]);
let result = take_fixed_size_binary::<_, true>(&fsb, &indices, 5).unwrap();
assert_eq!(result.len(), 4);
assert_eq!(result.null_count(), 2);
assert_eq!(
result.nulls().unwrap().iter().collect::<Vec<_>>(),
vec![true, false, false, true]
);
}
#[test]
#[should_panic(expected = "index out of bounds: the len is 4 but the index is 1000")]
fn test_take_list_out_of_bounds() {
let value_data = Int32Array::from(vec![0, 0, 0, -1, -2, -1, 2, 3]).into_data();
let value_offsets = Buffer::from_slice_ref([0, 3, 6, 8]);
let list_data_type =
DataType::List(Arc::new(Field::new_list_field(DataType::Int32, false)));
let list_data = ArrayData::builder(list_data_type)
.len(3)
.add_buffer(value_offsets)
.add_child_data(value_data)
.build()
.unwrap();
let list_array = ListArray::from(list_data);
let index = UInt32Array::from(vec![1000]);
take(&list_array, &index, None).unwrap();
}
#[test]
fn test_take_map() {
let values = Int32Array::from(vec![1, 2, 3, 4]);
let array =
MapArray::new_from_strings(vec!["a", "b", "c", "a"].into_iter(), &values, &[0, 3, 4])
.unwrap();
let index = UInt32Array::from(vec![0]);
let result = take(&array, &index, None).unwrap();
let expected: ArrayRef = Arc::new(
MapArray::new_from_strings(
vec!["a", "b", "c"].into_iter(),
&values.slice(0, 3),
&[0, 3],
)
.unwrap(),
);
assert_eq!(&expected, &result);
}
#[test]
fn test_take_struct() {
let array = create_test_struct(vec![
Some((Some(true), Some(42))),
Some((Some(false), Some(28))),
Some((Some(false), Some(19))),
Some((Some(true), Some(31))),
None,
]);
let index = UInt32Array::from(vec![0, 3, 1, 0, 2, 4]);
let actual = take(&array, &index, None).unwrap();
let actual: &StructArray = actual.as_any().downcast_ref::<StructArray>().unwrap();
assert_eq!(index.len(), actual.len());
assert_eq!(1, actual.null_count());
let expected = create_test_struct(vec![
Some((Some(true), Some(42))),
Some((Some(true), Some(31))),
Some((Some(false), Some(28))),
Some((Some(true), Some(42))),
Some((Some(false), Some(19))),
None,
]);
assert_eq!(&expected, actual);
let nulls = NullBuffer::from(&[false, true, false, true, false, true]);
let empty_struct_arr = StructArray::new_empty_fields(6, Some(nulls));
let index = UInt32Array::from(vec![0, 2, 1, 4]);
let actual = take(&empty_struct_arr, &index, None).unwrap();
let expected_nulls = NullBuffer::from(&[false, false, true, false]);
let expected_struct_arr = StructArray::new_empty_fields(4, Some(expected_nulls));
assert_eq!(&expected_struct_arr, actual.as_struct());
}
#[test]
fn test_take_struct_with_null_indices() {
let array = create_test_struct(vec![
Some((Some(true), Some(42))),
Some((Some(false), Some(28))),
Some((Some(false), Some(19))),
Some((Some(true), Some(31))),
None,
]);
let index = UInt32Array::from(vec![None, Some(3), Some(1), None, Some(0), Some(4)]);
let actual = take(&array, &index, None).unwrap();
let actual: &StructArray = actual.as_any().downcast_ref::<StructArray>().unwrap();
assert_eq!(index.len(), actual.len());
assert_eq!(3, actual.null_count());
let expected = create_test_struct(vec![
None,
Some((Some(true), Some(31))),
Some((Some(false), Some(28))),
None,
Some((Some(true), Some(42))),
None,
]);
assert_eq!(&expected, actual);
}
#[test]
fn test_take_out_of_bounds() {
let index = UInt32Array::from(vec![Some(3), None, Some(1), Some(3), Some(6)]);
let take_opt = TakeOptions { check_bounds: true };
let result = test_take_primitive_arrays::<Int64Type>(
vec![Some(0), None, Some(2), Some(3), None],
&index,
Some(take_opt),
vec![None],
);
assert!(result.is_err());
}
#[test]
#[should_panic(expected = "index out of bounds: the len is 4 but the index is 1000")]
fn test_take_out_of_bounds_panic() {
let index = UInt32Array::from(vec![Some(1000)]);
test_take_primitive_arrays::<Int64Type>(
vec![Some(0), Some(1), Some(2), Some(3)],
&index,
None,
vec![None],
)
.unwrap();
}
#[test]
fn test_null_array_smaller_than_indices() {
let values = NullArray::new(2);
let indices = UInt32Array::from(vec![Some(0), None, Some(15)]);
let result = take(&values, &indices, None).unwrap();
let expected: ArrayRef = Arc::new(NullArray::new(3));
assert_eq!(&result, &expected);
}
#[test]
fn test_null_array_larger_than_indices() {
let values = NullArray::new(5);
let indices = UInt32Array::from(vec![Some(0), None, Some(15)]);
let result = take(&values, &indices, None).unwrap();
let expected: ArrayRef = Arc::new(NullArray::new(3));
assert_eq!(&result, &expected);
}
#[test]
fn test_null_array_indices_out_of_bounds() {
let values = NullArray::new(5);
let indices = UInt32Array::from(vec![Some(0), None, Some(15)]);
let result = take(&values, &indices, Some(TakeOptions { check_bounds: true }));
assert_eq!(
result.unwrap_err().to_string(),
"Compute error: Array index out of bounds, cannot get item at index 15 from 5 entries"
);
}
#[test]
fn test_take_dict() {
let mut dict_builder = StringDictionaryBuilder::<Int16Type>::new();
dict_builder.append("foo").unwrap();
dict_builder.append("bar").unwrap();
dict_builder.append("").unwrap();
dict_builder.append_null();
dict_builder.append("foo").unwrap();
dict_builder.append("bar").unwrap();
dict_builder.append("bar").unwrap();
dict_builder.append("foo").unwrap();
let array = dict_builder.finish();
let dict_values = array.values().clone();
let dict_values = dict_values.as_any().downcast_ref::<StringArray>().unwrap();
let indices = UInt32Array::from(vec![
Some(0), Some(7), None, Some(5), Some(6), Some(2), Some(3), ]);
let result = take(&array, &indices, None).unwrap();
let result = result
.as_any()
.downcast_ref::<DictionaryArray<Int16Type>>()
.unwrap();
let result_values: StringArray = result.values().to_data().into();
let expected_values = StringArray::from(vec!["foo", "bar", ""]);
assert_eq!(&expected_values, dict_values);
assert_eq!(&expected_values, &result_values);
let expected_keys = Int16Array::from(vec![
Some(0),
Some(0),
None,
Some(1),
Some(1),
Some(2),
None,
]);
assert_eq!(result.keys(), &expected_keys);
}
fn build_generic_list<S, T>(data: Vec<Option<Vec<T::Native>>>) -> GenericListArray<S>
where
S: OffsetSizeTrait + 'static,
T: ArrowPrimitiveType,
PrimitiveArray<T>: From<Vec<Option<T::Native>>>,
{
GenericListArray::from_iter_primitive::<T, _, _>(
data.iter()
.map(|x| x.as_ref().map(|x| x.iter().map(|x| Some(*x)))),
)
}
fn test_take_sliced_list_generic<S: OffsetSizeTrait + 'static>() {
let list = build_generic_list::<S, Int32Type>(vec![
Some(vec![0, 1]),
Some(vec![2, 3, 4]),
None,
Some(vec![]),
Some(vec![5, 6]),
Some(vec![7]),
]);
let sliced = list.slice(1, 4);
let indices = UInt32Array::from(vec![Some(3), Some(0), None, Some(2), Some(1)]);
let taken = take(&sliced, &indices, None).unwrap();
let taken = taken.as_list::<S>();
let expected = build_generic_list::<S, Int32Type>(vec![
Some(vec![5, 6]),
Some(vec![2, 3, 4]),
None,
Some(vec![]),
None,
]);
assert_eq!(taken, &expected);
}
fn test_take_sliced_list_with_value_nulls_generic<S: OffsetSizeTrait + 'static>() {
let list = GenericListArray::<S>::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(10)]),
Some(vec![None, Some(1)]),
None,
Some(vec![Some(2), None]),
Some(vec![]),
Some(vec![Some(3)]),
]);
let sliced = list.slice(1, 4);
let indices = UInt32Array::from(vec![Some(2), Some(0), None, Some(3), Some(1)]);
let taken = take(&sliced, &indices, None).unwrap();
let taken = taken.as_list::<S>();
let expected = GenericListArray::<S>::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(2), None]),
Some(vec![None, Some(1)]),
None,
Some(vec![]),
None,
]);
assert_eq!(taken, &expected);
}
#[test]
fn test_take_sliced_list() {
test_take_sliced_list_generic::<i32>();
}
#[test]
fn test_take_sliced_large_list() {
test_take_sliced_list_generic::<i64>();
}
#[test]
fn test_take_sliced_list_with_value_nulls() {
test_take_sliced_list_with_value_nulls_generic::<i32>();
}
#[test]
fn test_take_sliced_large_list_with_value_nulls() {
test_take_sliced_list_with_value_nulls_generic::<i64>();
}
#[test]
fn test_take_runs() {
let logical_array: Vec<i32> = vec![1_i32, 1, 2, 2, 1, 1, 1, 2, 2, 1, 1, 2, 2];
let mut builder = PrimitiveRunBuilder::<Int32Type, Int32Type>::new();
builder.extend(logical_array.into_iter().map(Some));
let run_array = builder.finish();
let take_indices: PrimitiveArray<Int32Type> =
vec![7, 2, 3, 7, 11, 4, 6].into_iter().collect();
let take_out = take_run(&run_array, &take_indices).unwrap();
assert_eq!(take_out.len(), 7);
assert_eq!(
take_out.run_ends().values().len(),
2,
"expected two physical runs"
);
assert_eq!(take_out.run_ends().values(), &[5_i32, 7]);
let take_out_values = take_out.values().as_primitive::<Int32Type>();
assert_eq!(take_out_values.values(), &[2, 1]);
}
#[test]
fn test_take_runs_null_indices() {
let mut builder = PrimitiveRunBuilder::<Int32Type, Int32Type>::new();
builder.extend([Some(10), Some(10), None, None, Some(99)]);
let run_array = builder.finish();
let indices = Int32Array::from(vec![Some(0), None, Some(2), Some(3), Some(4)]);
let taken = take(&run_array, &indices, None).unwrap();
let run = taken.as_run::<Int32Type>();
let logical: Vec<Option<i32>> = run.downcast::<Int32Array>().unwrap().into_iter().collect();
assert_eq!(logical, vec![Some(10), None, None, None, Some(99)]);
assert_eq!(run.run_ends().values(), &[1_i32, 4, 5]);
}
#[test]
fn test_take_runs_sliced() {
let logical_array: Vec<i32> = vec![1, 1, 2, 2, 3, 3, 3, 4, 4, 5, 5, 6, 6];
let mut builder = PrimitiveRunBuilder::<Int32Type, Int32Type>::new();
builder.extend(logical_array.into_iter().map(Some));
let run_array = builder.finish();
let run_array = run_array.slice(4, 6);
let take_indices: PrimitiveArray<Int32Type> = vec![0, 5, 5, 1, 4].into_iter().collect();
let result = take_run(&run_array, &take_indices).unwrap();
let result = result.downcast::<Int32Array>().unwrap();
assert_eq!(
result.run_ends().values().len(),
4,
"expected four physical runs"
);
assert_eq!(result.run_ends().values(), &[1_i32, 3, 4, 5]);
let expected = vec![3, 5, 5, 3, 4];
let actual = result.into_iter().flatten().collect::<Vec<_>>();
assert_eq!(expected, actual);
}
#[test]
fn test_take_value_index_from_fixed_list() {
let list = FixedSizeListArray::from_iter_primitive::<Int32Type, _, _>(
vec![
Some(vec![Some(1), Some(2), None]),
Some(vec![Some(4), None, Some(6)]),
None,
Some(vec![None, Some(8), Some(9)]),
],
3,
);
let indices = UInt32Array::from(vec![2, 1, 0]);
let indexed = take_value_indices_from_fixed_size_list(&list, &indices, 3).unwrap();
assert_eq!(indexed, UInt32Array::from(vec![6, 7, 8, 3, 4, 5, 0, 1, 2]));
let indices = UInt32Array::from(vec![3, 2, 1, 2, 0]);
let indexed = take_value_indices_from_fixed_size_list(&list, &indices, 3).unwrap();
assert_eq!(
indexed,
UInt32Array::from(vec![9, 10, 11, 6, 7, 8, 3, 4, 5, 6, 7, 8, 0, 1, 2])
);
}
#[test]
fn test_take_null_indices() {
let indices = Int32Array::new(
vec![1, 2, 400, 400].into(),
Some(NullBuffer::from(vec![true, true, false, false])),
);
let values = Int32Array::from(vec![1, 23, 4, 5]);
let r = take(&values, &indices, None).unwrap();
let values = r
.as_primitive::<Int32Type>()
.into_iter()
.collect::<Vec<_>>();
assert_eq!(&values, &[Some(23), Some(4), None, None])
}
#[test]
fn test_take_fixed_size_list_null_indices() {
let indices = Int32Array::from_iter([Some(0), None]);
let values = Arc::new(Int32Array::from(vec![0, 1, 2, 3]));
let arr_field = Arc::new(Field::new_list_field(values.data_type().clone(), true));
let values = FixedSizeListArray::try_new(arr_field, 2, values, None).unwrap();
let r = take(&values, &indices, None).unwrap();
let values = r
.as_fixed_size_list()
.values()
.as_primitive::<Int32Type>()
.into_iter()
.collect::<Vec<_>>();
assert_eq!(values, &[Some(0), Some(1), None, None])
}
#[test]
fn test_take_bytes_null_indices() {
let indices = Int32Array::new(
vec![0, 1, 400, 400].into(),
Some(NullBuffer::from_iter(vec![true, true, false, false])),
);
let values = StringArray::from(vec![Some("foo"), None]);
let r = take(&values, &indices, None).unwrap();
let values = r.as_string::<i32>().iter().collect::<Vec<_>>();
assert_eq!(&values, &[Some("foo"), None, None, None])
}
#[test]
fn test_take_union_sparse() {
let structs = create_test_struct(vec![
Some((Some(true), Some(42))),
Some((Some(false), Some(28))),
Some((Some(false), Some(19))),
Some((Some(true), Some(31))),
None,
]);
let strings = StringArray::from(vec![Some("a"), None, Some("c"), None, Some("d")]);
let type_ids = [1; 5].into_iter().collect::<ScalarBuffer<i8>>();
let union_fields = [
(
0,
Arc::new(Field::new("f1", structs.data_type().clone(), true)),
),
(
1,
Arc::new(Field::new("f2", strings.data_type().clone(), true)),
),
]
.into_iter()
.collect();
let children = vec![Arc::new(structs) as Arc<dyn Array>, Arc::new(strings)];
let array = UnionArray::try_new(union_fields, type_ids, None, children).unwrap();
let indices = vec![0, 3, 1, 0, 2, 4];
let index = UInt32Array::from(indices.clone());
let actual = take(&array, &index, None).unwrap();
let actual = actual.as_any().downcast_ref::<UnionArray>().unwrap();
let strings = actual.child(1);
let strings = strings.as_any().downcast_ref::<StringArray>().unwrap();
let actual = strings.iter().collect::<Vec<_>>();
let expected = vec![Some("a"), None, None, Some("a"), Some("c"), Some("d")];
assert_eq!(expected, actual);
}
#[test]
fn test_take_union_dense() {
let type_ids = vec![0, 1, 1, 0, 0, 1, 0];
let offsets = vec![0, 0, 1, 1, 2, 2, 3];
let ints = vec![10, 20, 30, 40];
let strings = vec![Some("a"), None, Some("c"), Some("d")];
let indices = vec![0, 3, 1, 0, 2, 4];
let taken_type_ids = vec![0, 0, 1, 0, 1, 0];
let taken_offsets = vec![0, 1, 0, 2, 1, 3];
let taken_ints = vec![10, 20, 10, 30];
let taken_strings = vec![Some("a"), None];
let type_ids = <ScalarBuffer<i8>>::from(type_ids);
let offsets = <ScalarBuffer<i32>>::from(offsets);
let ints = UInt32Array::from(ints);
let strings = StringArray::from(strings);
let union_fields = [
(
0,
Arc::new(Field::new("f1", ints.data_type().clone(), true)),
),
(
1,
Arc::new(Field::new("f2", strings.data_type().clone(), true)),
),
]
.into_iter()
.collect();
let array = UnionArray::try_new(
union_fields,
type_ids,
Some(offsets),
vec![Arc::new(ints), Arc::new(strings)],
)
.unwrap();
let index = UInt32Array::from(indices);
let actual = take(&array, &index, None).unwrap();
let actual = actual.as_any().downcast_ref::<UnionArray>().unwrap();
assert_eq!(actual.offsets(), Some(&ScalarBuffer::from(taken_offsets)));
assert_eq!(actual.type_ids(), &ScalarBuffer::from(taken_type_ids));
assert_eq!(
UInt32Array::from(actual.child(0).to_data()),
UInt32Array::from(taken_ints)
);
assert_eq!(
StringArray::from(actual.child(1).to_data()),
StringArray::from(taken_strings)
);
}
fn union_i32_logical(array: &UnionArray) -> Vec<Option<i32>> {
(0..array.len())
.map(|i| {
let child = array.child(array.type_id(i)).as_primitive::<Int32Type>();
let offset = array.value_offset(i);
if child.is_null(offset) {
None
} else {
Some(child.value(offset))
}
})
.collect()
}
#[test]
fn test_take_union_dense_null_indices() {
let fields =
UnionFields::try_new(vec![5], vec![Field::new("i", DataType::Int32, true)]).unwrap();
let dense = UnionArray::try_new(
fields.clone(),
ScalarBuffer::from(vec![5_i8, 5, 5]),
Some(ScalarBuffer::from(vec![0_i32, 1, 2])),
vec![Arc::new(Int32Array::from(vec![1, 2, 3]))],
)
.unwrap();
let sparse = UnionArray::try_new(
fields,
ScalarBuffer::from(vec![5_i8, 5, 5]),
None,
vec![Arc::new(Int32Array::from(vec![1, 2, 3]))],
)
.unwrap();
let indices = UInt32Array::new(
ScalarBuffer::from(vec![0_u32, 99, 2]),
Some(NullBuffer::from(vec![true, false, true])),
);
let dense_taken = take(&dense, &indices, None).unwrap();
let sparse_taken = take(&sparse, &indices, None).unwrap();
let dense_logical = union_i32_logical(dense_taken.as_any().downcast_ref().unwrap());
let sparse_logical = union_i32_logical(sparse_taken.as_any().downcast_ref().unwrap());
assert_eq!(dense_logical, vec![Some(1), None, Some(3)]);
assert_eq!(dense_logical, sparse_logical);
}
#[test]
fn test_take_empty_union_without_null_indices() {
let fields = UnionFields::try_new(vec![], Vec::<Field>::new()).unwrap();
let indices = UInt32Array::from(Vec::<u32>::new());
let sparse = UnionArray::try_new(
fields.clone(),
ScalarBuffer::<i8>::from(vec![]),
None,
vec![],
)
.unwrap();
let dense = UnionArray::try_new(
fields,
ScalarBuffer::<i8>::from(vec![]),
Some(ScalarBuffer::<i32>::from(vec![])),
vec![],
)
.unwrap();
for values in [&sparse, &dense] {
let taken = take(values, &indices, None).unwrap();
assert_eq!(taken.len(), 0);
assert_eq!(taken.data_type(), values.data_type());
}
}
#[test]
fn test_take_empty_union_with_null_indices() {
let fields = UnionFields::try_new(vec![], Vec::<Field>::new()).unwrap();
let indices = UInt32Array::from(vec![None]);
let sparse = UnionArray::try_new(
fields.clone(),
ScalarBuffer::<i8>::from(vec![]),
None,
vec![],
)
.unwrap();
let dense = UnionArray::try_new(
fields,
ScalarBuffer::<i8>::from(vec![]),
Some(ScalarBuffer::<i32>::from(vec![])),
vec![],
)
.unwrap();
for values in [&sparse, &dense] {
let error = take(values, &indices, None).unwrap_err();
assert_eq!(
error.to_string(),
"Compute error: Cannot take from a union with zero fields when indices contains nulls"
);
}
}
#[test]
fn test_take_union_dense_using_builder() {
let mut builder = UnionBuilder::new_dense();
builder.append::<Int32Type>("a", 1).unwrap();
builder.append::<Float64Type>("b", 3.0).unwrap();
builder.append::<Int32Type>("a", 4).unwrap();
builder.append::<Int32Type>("a", 5).unwrap();
builder.append::<Float64Type>("b", 2.0).unwrap();
let union = builder.build().unwrap();
let indices = UInt32Array::from(vec![2, 0, 1, 2]);
let mut builder = UnionBuilder::new_dense();
builder.append::<Int32Type>("a", 4).unwrap();
builder.append::<Int32Type>("a", 1).unwrap();
builder.append::<Float64Type>("b", 3.0).unwrap();
builder.append::<Int32Type>("a", 4).unwrap();
let taken = builder.build().unwrap();
assert_eq!(
taken.to_data(),
take(&union, &indices, None).unwrap().to_data()
);
}
#[test]
fn test_take_union_dense_all_match_issue_6206() {
let fields = UnionFields::from_fields(vec![Field::new("a", DataType::Int64, false)]);
let ints = Arc::new(Int64Array::from(vec![1, 2, 3, 4, 5]));
let array = UnionArray::try_new(
fields,
ScalarBuffer::from(vec![0_i8, 0, 0, 0, 0]),
Some(ScalarBuffer::from_iter(0_i32..5)),
vec![ints],
)
.unwrap();
let indices = Int64Array::from(vec![0, 2, 4]);
let array = take(&array, &indices, None).unwrap();
assert_eq!(array.len(), 3);
}
fn offset_overflow_fixture() -> (StringArray, usize) {
let value_len = 1_000_000;
let values = StringArray::from(vec![Some("a".repeat(value_len))]);
let n = i32::MAX as usize / value_len + 1;
(values, n)
}
#[test]
fn test_take_bytes_offset_overflow() {
let (values, n) = offset_overflow_fixture();
let indices = Int32Array::from(vec![0; n]);
assert!(matches!(
take(&values, &indices, None),
Err(ArrowError::OffsetOverflowError(_))
));
}
#[test]
fn test_take_bytes_offset_overflow_nullable() {
let (values, n) = offset_overflow_fixture();
let validity =
NullBuffer::from_iter(std::iter::once(false).chain(std::iter::repeat_n(true, n)));
let indices = Int32Array::new(vec![0i32; n + 1].into(), Some(validity));
assert!(matches!(
take(&values, &indices, None),
Err(ArrowError::OffsetOverflowError(_))
));
}
#[test]
fn test_take_run_empty_indices() {
let mut builder = PrimitiveRunBuilder::<Int32Type, Int32Type>::new();
builder.extend([Some(1), Some(1), Some(2), Some(2)]);
let run_array = builder.finish();
let logical_indices: PrimitiveArray<Int32Type> = PrimitiveArray::from(Vec::<i32>::new());
let result = take_impl::<_, true>(&run_array, &logical_indices)
.expect("take_run with empty indices");
assert_eq!(result.len(), 0);
assert_eq!(result.null_count(), 0);
let run_result = result
.as_any()
.downcast_ref::<RunArray<Int32Type>>()
.expect("result should be a RunArray");
assert_eq!(run_result.run_ends().len(), 0);
assert_eq!(run_result.values().len(), 0);
}
#[test]
fn test_take_run_end_encoded_merges_identical_runs() {
let mut builder = PrimitiveRunBuilder::<Int32Type, Int32Type>::new();
builder.extend([1, 1, 0, 0, 1, 1].into_iter().map(Some));
let ree = builder.finish();
let indexes = Int32Array::from_iter_values(vec![0, 1, 4, 5]);
let result = take(&ree, &indexes, None).unwrap();
let result = result
.as_run::<Int32Type>()
.downcast::<Int32Array>()
.unwrap();
assert_eq!(
result.run_ends().values().len(),
1,
"expected a single physical run"
);
assert_eq!(result.run_ends().values(), &[4_i32]);
let actual = result.into_iter().flatten().collect::<Vec<_>>();
assert_eq!(actual, vec![1, 1, 1, 1]);
}
#[test]
fn test_take_run_end_encoded_merges_identical_string_runs() {
let mut builder = StringRunBuilder::<Int32Type>::new();
builder.extend(
["bob", "bob", "alice", "alice", "bob", "bob"]
.into_iter()
.map(Some),
);
let ree = builder.finish();
let indexes = Int32Array::from_iter_values(vec![0, 1, 4, 5]);
let result = take(&ree, &indexes, None).unwrap();
let result = result
.as_run::<Int32Type>()
.downcast::<StringArray>()
.unwrap();
assert_eq!(
result.run_ends().values().len(),
1,
"expected a single physical run"
);
assert_eq!(result.run_ends().values(), &[4_i32]);
let actual = result.into_iter().flatten().collect::<Vec<_>>();
assert_eq!(actual, vec!["bob", "bob", "bob", "bob"]);
}
#[test]
fn test_take_run_end_encoded_mixed_runs() {
let mut builder = StringRunBuilder::<Int32Type>::new();
builder.extend(
["bob", "bob", "alice", "alice", "bob", "bob", "eve", "eve"]
.into_iter()
.map(Some),
);
let ree = builder.finish();
let indexes = Int32Array::from_iter_values(vec![0, 0, 1, 4, 5, 2, 3, 2, 6, 7, 6]);
let result = take(&ree, &indexes, None).unwrap();
let result = result
.as_run::<Int32Type>()
.downcast::<StringArray>()
.unwrap();
println!("run_ends_raw: {:?}", result.run_ends());
println!("run_ends: {:?}", result.run_ends().values());
println!("values : {:?}", result.values());
assert_eq!(
result.run_ends().values().len(),
3,
"expected three physical runs"
);
assert_eq!(result.run_ends().values(), &[5_i32, 8, 11]);
let actual = result.into_iter().flatten().collect::<Vec<_>>();
assert_eq!(
actual,
vec![
"bob", "bob", "bob", "bob", "bob", "alice", "alice", "alice", "eve", "eve", "eve"
]
);
}
#[test]
fn test_take_fixed_size_list_parent_nulls() {
let list = FixedSizeListArray::from_iter_primitive::<Int32Type, _, _>(
vec![
Some(vec![Some(1), Some(2)]),
None,
Some(vec![Some(5), Some(6)]),
],
2,
);
let indices = UInt32Array::from(vec![2, 1, 0]);
let result = take(&list, &indices, None).unwrap();
let result = result.as_fixed_size_list();
assert_eq!(result.len(), 3);
assert!(result.is_valid(0));
assert!(result.is_null(1));
assert!(result.is_valid(2));
let child = result.values().as_primitive::<Int32Type>();
assert_eq!(child.value(0), 5);
assert_eq!(child.value(1), 6);
assert!(child.is_null(2));
assert!(child.is_null(3));
assert_eq!(child.value(4), 1);
assert_eq!(child.value(5), 2);
}
#[test]
fn test_take_zero_sized_fixed_size_list() {
let input = FixedSizeListArray::try_new_with_length(
Field::new_list_field(DataType::Int32, true).into(),
0,
Arc::new(Int32Array::new_null(0)),
None,
3,
)
.unwrap();
let indices = UInt32Array::from(vec![2, 0]);
let result = take(&input, &indices, None).unwrap();
assert_eq!(result.len(), 2);
}
}