vortex_array/arrays/varbinview/compute/
take.rs1use std::iter;
5use std::ptr;
6use std::sync::Arc;
7
8use itertools::Itertools as _;
9use num_traits::AsPrimitive;
10use vortex_buffer::Buffer;
11use vortex_buffer::BufferMut;
12use vortex_error::VortexResult;
13use vortex_error::vortex_ensure;
14use vortex_error::vortex_err;
15use vortex_mask::AllOr;
16use vortex_mask::Mask;
17
18use crate::ArrayRef;
19use crate::Columnar;
20use crate::IntoArray;
21use crate::array::ArrayView;
22use crate::arrays::PiecewiseSequence;
23use crate::arrays::PrimitiveArray;
24use crate::arrays::VarBinView;
25use crate::arrays::VarBinViewArray;
26use crate::arrays::dict::TakeExecute;
27use crate::arrays::piecewise_sequence::constant_unsigned_usize;
28use crate::arrays::piecewise_sequence::maybe_contiguous_slices;
29use crate::arrays::varbinview::BinaryView;
30use crate::buffer::BufferHandle;
31use crate::dtype::UnsignedPType;
32use crate::executor::ExecutionCtx;
33use crate::match_each_integer_ptype;
34use crate::match_each_unsigned_integer_ptype;
35
36impl TakeExecute for VarBinView {
37 fn take(
39 array: ArrayView<'_, VarBinView>,
40 indices: &ArrayRef,
41 ctx: &mut ExecutionCtx,
42 ) -> VortexResult<Option<ArrayRef>> {
43 if let Some(piecewise_indices) = indices.as_opt::<PiecewiseSequence>()
44 && let Some(taken) = take_contiguous_ranges(array, piecewise_indices, indices, ctx)?
45 {
46 return Ok(Some(taken));
47 }
48
49 let validity = array.validity()?.take(indices)?;
50 let indices = indices.clone().execute::<PrimitiveArray>(ctx)?;
51
52 let indices_mask = indices
53 .as_ref()
54 .validity()?
55 .execute_mask(indices.as_ref().len(), ctx)?;
56 let views_buffer = match_each_integer_ptype!(indices.ptype(), |I| {
57 take_views(array.views(), indices.as_slice::<I>(), &indices_mask)
58 });
59
60 unsafe {
62 Ok(Some(
63 VarBinViewArray::new_handle_unchecked(
64 BufferHandle::new_host(views_buffer.into_byte_buffer()),
65 Arc::clone(array.data_buffers()),
66 array
67 .dtype()
68 .union_nullability(indices.dtype().nullability()),
69 validity,
70 )
71 .into_array(),
72 ))
73 }
74 }
75}
76
77fn take_contiguous_ranges(
78 array: ArrayView<'_, VarBinView>,
79 indices: ArrayView<'_, PiecewiseSequence>,
80 indices_ref: &ArrayRef,
81 ctx: &mut ExecutionCtx,
82) -> VortexResult<Option<ArrayRef>> {
83 let Some((starts, lengths)) = maybe_contiguous_slices(indices, ctx)? else {
84 return Ok(None);
85 };
86 let source = array.views();
87 let output_len = indices_ref.len();
88 let views = match lengths {
89 Columnar::Constant(lengths) => {
90 let length = constant_unsigned_usize(&lengths);
91 match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
92 gather_view_slices_constant_length(
93 source,
94 starts.as_slice::<S>(),
95 length,
96 output_len,
97 )?
98 })
99 }
100 Columnar::Canonical(lengths) => {
101 let lengths = lengths.into_primitive();
102 match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
103 match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
104 gather_view_slices(
105 source,
106 starts.as_slice::<S>(),
107 lengths.as_slice::<L>(),
108 output_len,
109 )?
110 })
111 })
112 }
113 };
114 let validity = array.validity()?.take(indices_ref)?;
115
116 unsafe {
119 Ok(Some(
120 VarBinViewArray::new_handle_unchecked(
121 BufferHandle::new_host(views.into_byte_buffer()),
122 Arc::clone(array.data_buffers()),
123 array.dtype().clone(),
124 validity,
125 )
126 .into_array(),
127 ))
128 }
129}
130
131fn take_views<I: AsPrimitive<usize>>(
132 views_ref: &[BinaryView],
133 indices: &[I],
134 mask: &Mask,
135) -> Buffer<BinaryView> {
136 match mask.bit_buffer() {
140 AllOr::All => {
141 Buffer::<BinaryView>::from_trusted_len_iter(indices.iter().map(|i| views_ref[i.as_()]))
142 }
143 AllOr::None => Buffer::<BinaryView>::from_trusted_len_iter(iter::repeat_n(
144 BinaryView::default(),
145 indices.len(),
146 )),
147 AllOr::Some(buffer) => Buffer::<BinaryView>::from_trusted_len_iter(
148 buffer.iter().zip(indices.iter()).map(|(valid, idx)| {
149 if valid {
150 views_ref[idx.as_()]
151 } else {
152 BinaryView::default()
153 }
154 }),
155 ),
156 }
157}
158
159fn gather_view_slices_constant_length<S>(
160 source: &[BinaryView],
161 starts: &[S],
162 length: usize,
163 output_len: usize,
164) -> VortexResult<Buffer<BinaryView>>
165where
166 S: UnsignedPType,
167{
168 let computed_len = starts
169 .len()
170 .checked_mul(length)
171 .ok_or_else(|| vortex_err!("PiecewiseSequenceArray output length overflows usize"))?;
172 vortex_ensure!(
173 computed_len == output_len,
174 "PiecewiseSequenceArray expanded length {computed_len} does not match declared length {output_len}"
175 );
176
177 let mut views = BufferMut::<BinaryView>::with_capacity(output_len);
178 let spare = &mut views.spare_capacity_mut()[..output_len];
179 let mut cursor = 0usize;
180 for &start in starts {
181 let start = start.as_();
182 let src = &source[start..][..length];
183 unsafe {
185 ptr::copy_nonoverlapping(
186 src.as_ptr(),
187 spare[cursor..][..src.len()]
188 .as_mut_ptr()
189 .cast::<BinaryView>(),
190 src.len(),
191 );
192 }
193 cursor += src.len();
194 }
195 unsafe { views.set_len(cursor) };
197 vortex_ensure!(
198 views.len() == output_len,
199 "PiecewiseSequenceArray expanded length {} does not match declared length {output_len}",
200 views.len()
201 );
202 Ok(views.freeze())
203}
204
205fn gather_view_slices<S, L>(
206 source: &[BinaryView],
207 starts: &[S],
208 lengths: &[L],
209 output_len: usize,
210) -> VortexResult<Buffer<BinaryView>>
211where
212 S: UnsignedPType,
213 L: UnsignedPType,
214{
215 let mut views = BufferMut::<BinaryView>::with_capacity(output_len);
216 let spare = &mut views.spare_capacity_mut()[..output_len];
217 let mut cursor = 0usize;
218 for (&start, &length) in starts.iter().zip_eq(lengths) {
219 let start = start.as_();
220 let length = length.as_();
221 let src = &source[start..][..length];
222 unsafe {
224 ptr::copy_nonoverlapping(
225 src.as_ptr(),
226 spare[cursor..][..src.len()]
227 .as_mut_ptr()
228 .cast::<BinaryView>(),
229 src.len(),
230 );
231 }
232 cursor += src.len();
233 }
234 unsafe { views.set_len(cursor) };
236 vortex_ensure!(
237 views.len() == output_len,
238 "PiecewiseSequenceArray expanded length {} does not match declared length {output_len}",
239 views.len()
240 );
241 Ok(views.freeze())
242}
243
244#[cfg(test)]
245mod tests {
246 use rstest::rstest;
247 use vortex_buffer::BitBuffer;
248 use vortex_buffer::buffer;
249 use vortex_error::VortexResult;
250
251 use crate::IntoArray;
252 use crate::VortexSessionExecute;
253 use crate::array_session;
254 use crate::arrays::VarBinViewArray;
255 use crate::arrays::varbinview::compute::take::PrimitiveArray;
256 use crate::compute::conformance::take::test_take_conformance;
257 use crate::dtype::DType;
258 use crate::dtype::Nullability::NonNullable;
259 use crate::validity::Validity;
260
261 #[test]
262 fn take_nullable() -> VortexResult<()> {
263 let arr = VarBinViewArray::from_iter_nullable_str([
264 Some("one"),
265 None,
266 Some("three"),
267 Some("four"),
268 None,
269 Some("six"),
270 ]);
271
272 let taken = arr.take(buffer![0, 3].into_array())?;
273
274 assert!(taken.dtype().is_nullable());
275 let mut ctx = array_session().create_execution_ctx();
276 let taken = taken.execute::<VarBinViewArray>(&mut ctx)?;
277 let mask = taken.validity()?.execute_mask(taken.len(), &mut ctx)?;
278 let result = (0..taken.len())
279 .map(|i| {
280 mask.value(i)
281 .then(|| unsafe { String::from_utf8_unchecked(taken.bytes_at(i).to_vec()) })
282 })
283 .collect::<Vec<_>>();
284 assert_eq!(result, [Some("one".to_string()), Some("four".to_string())]);
285 Ok(())
286 }
287
288 #[test]
289 fn take_nullable_indices() -> VortexResult<()> {
290 let arr = VarBinViewArray::from_iter(["one", "two"].map(Some), DType::Utf8(NonNullable));
291
292 let indices = PrimitiveArray::new(
293 buffer![1u64, 999],
295 Validity::from(BitBuffer::from(vec![true, false])),
296 );
297
298 let taken = arr.take(indices.into_array())?;
299
300 assert!(taken.dtype().is_nullable());
301 let mut ctx = array_session().create_execution_ctx();
302 let taken = taken.execute::<VarBinViewArray>(&mut ctx)?;
303 let mask = taken.validity()?.execute_mask(taken.len(), &mut ctx)?;
304 let result = (0..taken.len())
305 .map(|i| {
306 mask.value(i)
307 .then(|| unsafe { String::from_utf8_unchecked(taken.bytes_at(i).to_vec()) })
308 })
309 .collect::<Vec<_>>();
310 assert_eq!(result, [Some("two".to_string()), None]);
311 Ok(())
312 }
313
314 #[rstest]
315 #[case(VarBinViewArray::from_iter(
316 ["hello", "world", "test", "data", "array"].map(Some),
317 DType::Utf8(NonNullable),
318 ))]
319 #[case(VarBinViewArray::from_iter_nullable_str([
320 Some("hello"),
321 None,
322 Some("test"),
323 Some("data"),
324 None,
325 ]))]
326 #[case(VarBinViewArray::from_iter(
327 [b"hello".as_slice(), b"world", b"test", b"data", b"array"].map(Some),
328 DType::Binary(NonNullable),
329 ))]
330 #[case(VarBinViewArray::from_iter(["single"].map(Some), DType::Utf8(NonNullable)))]
331 fn test_take_varbinview_conformance(#[case] array: VarBinViewArray) {
332 test_take_conformance(
333 &array.into_array(),
334 &mut array_session().create_execution_ctx(),
335 );
336 }
337}