1use crate::coalesce::InProgressArray;
19use crate::filter::{FilterIndices, FilterPredicate, FilterSelection, FilterSlices};
20use arrow_array::cast::AsArray;
21use arrow_array::{Array, ArrayRef, ArrowPrimitiveType, PrimitiveArray};
22use arrow_buffer::{NullBuffer, NullBufferBuilder, ScalarBuffer};
23use arrow_schema::{ArrowError, DataType};
24use std::fmt::Debug;
25use std::sync::Arc;
26
27#[derive(Debug)]
29pub(crate) struct InProgressPrimitiveArray<T: ArrowPrimitiveType> {
30 data_type: DataType,
32 source: Option<ArrayRef>,
34 batch_size: usize,
36 nulls: NullBufferBuilder,
38 current: Vec<T::Native>,
40}
41
42impl<T: ArrowPrimitiveType> InProgressPrimitiveArray<T> {
43 pub(crate) fn new(batch_size: usize, data_type: DataType) -> Self {
45 Self {
46 data_type,
47 batch_size,
48 source: None,
49 nulls: NullBufferBuilder::new(batch_size),
50 current: vec![],
51 }
52 }
53
54 fn ensure_capacity(&mut self) {
59 if self.current.capacity() == 0 {
60 self.current.reserve(self.batch_size);
61 }
62 }
63
64 fn append_values_by_indices(
65 current: &mut Vec<T::Native>,
66 values: &[T::Native],
67 indices: FilterIndices<'_>,
68 selected_count: usize,
69 ) {
70 let current_len = current.len();
71 let mut written = 0;
72
73 unsafe {
74 let mut out = current
75 .spare_capacity_mut()
76 .as_mut_ptr()
77 .cast::<T::Native>();
78
79 indices.for_each(|idx| {
80 out.write(*values.get_unchecked(idx));
82 out = out.add(1);
83 written += 1;
84 });
85
86 current.set_len(current_len + written);
87 }
88
89 debug_assert_eq!(written, selected_count);
90 }
91
92 fn append_values_by_slices(
93 current: &mut Vec<T::Native>,
94 values: &[T::Native],
95 slices: FilterSlices<'_>,
96 selected_count: usize,
97 ) {
98 let current_len = current.len();
99 let mut written = 0;
100
101 unsafe {
102 let mut out = current
103 .spare_capacity_mut()
104 .as_mut_ptr()
105 .cast::<T::Native>();
106
107 slices.for_each(|(start, end)| {
108 let len = end - start;
109 std::ptr::copy_nonoverlapping(values.as_ptr().add(start), out, len);
111 out = out.add(len);
112 written += len;
113 });
114
115 current.set_len(current_len + written);
116 }
117
118 debug_assert_eq!(written, selected_count);
119 }
120}
121
122#[inline]
123fn primitive_source<T: ArrowPrimitiveType>(
124 source: &Option<ArrayRef>,
125) -> Result<&PrimitiveArray<T>, ArrowError> {
126 Ok(source
127 .as_ref()
128 .ok_or_else(|| {
129 ArrowError::InvalidArgumentError(
130 "Internal Error: InProgressPrimitiveArray: source not set".to_string(),
131 )
132 })?
133 .as_primitive::<T>())
134}
135
136fn append_filtered_nulls(
137 nulls: &mut NullBufferBuilder,
138 source_nulls: Option<&NullBuffer>,
139 filter: &FilterPredicate,
140) {
141 if let Some(filtered_nulls) = filter.filter_nulls(source_nulls) {
142 nulls.append_buffer(&filtered_nulls);
143 } else {
144 nulls.append_n_non_nulls(filter.count());
145 }
146}
147
148impl<T: ArrowPrimitiveType + Debug> InProgressArray for InProgressPrimitiveArray<T> {
149 fn set_source(&mut self, source: Option<ArrayRef>) {
150 self.source = source;
151 }
152
153 fn copy_rows(&mut self, offset: usize, len: usize) -> Result<(), ArrowError> {
154 self.ensure_capacity();
155
156 let s = primitive_source::<T>(&self.source)?;
157
158 if let Some(nulls) = s.nulls().as_ref() {
160 let nulls = nulls.slice(offset, len);
161 self.nulls.append_buffer(&nulls);
162 } else {
163 self.nulls.append_n_non_nulls(len);
164 };
165
166 let values = s.values();
168 self.current
170 .extend_from_slice(unsafe { values.get_unchecked(offset..offset + len) });
171
172 Ok(())
173 }
174
175 fn copy_rows_by_filter(&mut self, filter: &FilterPredicate) -> Result<(), ArrowError> {
176 match filter.selection() {
177 FilterSelection::Indices(indices) => {
178 self.ensure_capacity();
179 let s = primitive_source::<T>(&self.source)?;
180
181 append_filtered_nulls(&mut self.nulls, s.nulls(), filter);
182 self.current.reserve(filter.count());
183 Self::append_values_by_indices(
184 &mut self.current,
185 s.values(),
186 indices,
187 filter.count(),
188 );
189 Ok(())
190 }
191 FilterSelection::Slices(slices) => {
192 self.ensure_capacity();
193 let s = primitive_source::<T>(&self.source)?;
194
195 append_filtered_nulls(&mut self.nulls, s.nulls(), filter);
196 self.current.reserve(filter.count());
197 Self::append_values_by_slices(
198 &mut self.current,
199 s.values(),
200 slices,
201 filter.count(),
202 );
203 Ok(())
204 }
205 selection => self.copy_rows_by_selection(selection),
207 }
208 }
209
210 fn finish(&mut self) -> Result<ArrayRef, ArrowError> {
211 let values = std::mem::take(&mut self.current);
213 let nulls = self.nulls.finish();
214 self.nulls = NullBufferBuilder::new(self.batch_size);
215
216 let array = PrimitiveArray::<T>::try_new(ScalarBuffer::from(values), nulls)?
217 .with_data_type(self.data_type.clone());
219 Ok(Arc::new(array))
220 }
221
222 fn size(&self) -> usize {
223 self.source
224 .as_ref()
225 .map_or(0, |source| source.get_array_memory_size())
226 + self.current.capacity() * std::mem::size_of::<T::Native>()
227 + self.nulls.allocated_size()
228 }
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234 use crate::filter::FilterBuilder;
235 use arrow_array::types::Int32Type;
236 use arrow_array::{BooleanArray, Int32Array};
237
238 #[test]
239 fn test_copy_rows_by_filter_index_iterator() {
240 let source =
241 Int32Array::from_iter((0..21).map(|idx| if idx % 5 == 0 { None } else { Some(idx) }));
242 let filter = BooleanArray::from_iter(
243 (0..21).map(|idx| Some(matches!(idx, 0 | 1 | 2 | 3 | 5 | 8 | 13))),
244 );
245 let predicate = FilterBuilder::new(&filter).build();
246 let FilterSelection::Indices(indices) = predicate.selection() else {
247 panic!("expected index iterator selection");
248 };
249 let mut selected_indices = Vec::new();
250 indices.for_each(|idx| selected_indices.push(idx));
251 assert_eq!(selected_indices, vec![0, 1, 2, 3, 5, 8, 13]);
252
253 let mut in_progress = InProgressPrimitiveArray::<Int32Type>::new(7, DataType::Int32);
254 in_progress.set_source(Some(Arc::new(source)));
255 in_progress.copy_rows_by_filter(&predicate).unwrap();
256
257 let result = in_progress.finish().unwrap();
258 let result = result.as_primitive::<Int32Type>();
259 let expected = Int32Array::from(vec![
260 None,
261 Some(1),
262 Some(2),
263 Some(3),
264 None,
265 Some(8),
266 Some(13),
267 ]);
268 assert_eq!(result, &expected);
269 }
270
271 #[test]
272 fn test_copy_rows_by_filter_slice_iterator() {
273 let source =
274 Int32Array::from_iter((0..16).map(|idx| if idx % 5 == 0 { None } else { Some(idx) }));
275 let filter = BooleanArray::from_iter((0..16).map(|idx| Some(!matches!(idx, 3 | 9))));
276 let predicate = FilterBuilder::new(&filter).build();
277 let FilterSelection::Slices(slices) = predicate.selection() else {
278 panic!("expected slice iterator selection");
279 };
280 let mut selected_slices = Vec::new();
281 slices.for_each(|slice| selected_slices.push(slice));
282 assert_eq!(selected_slices, vec![(0, 3), (4, 9), (10, 16)]);
283
284 let mut in_progress = InProgressPrimitiveArray::<Int32Type>::new(14, DataType::Int32);
285 in_progress.set_source(Some(Arc::new(source)));
286 in_progress.copy_rows_by_filter(&predicate).unwrap();
287
288 let result = in_progress.finish().unwrap();
289 let result = result.as_primitive::<Int32Type>();
290 let expected = Int32Array::from(vec![
291 None,
292 Some(1),
293 Some(2),
294 Some(4),
295 None,
296 Some(6),
297 Some(7),
298 Some(8),
299 None,
300 Some(11),
301 Some(12),
302 Some(13),
303 Some(14),
304 None,
305 ]);
306 assert_eq!(result, &expected);
307 }
308
309 #[test]
310 fn test_size_empty() {
311 let in_progress = InProgressPrimitiveArray::<Int32Type>::new(64, DataType::Int32);
313 assert_eq!(in_progress.size(), 0);
314 }
315
316 #[test]
317 fn test_size_counts_source() {
318 let mut in_progress = InProgressPrimitiveArray::<Int32Type>::new(64, DataType::Int32);
319 let source: ArrayRef = Arc::new(Int32Array::from_iter_values(0..100));
320 in_progress.set_source(Some(Arc::clone(&source)));
321 assert_eq!(in_progress.size(), source.get_array_memory_size());
323 }
324
325 #[test]
326 fn test_size_counts_values_buffer_and_resets_on_finish() {
327 const BATCH_SIZE: usize = 64;
328 let mut in_progress =
329 InProgressPrimitiveArray::<Int32Type>::new(BATCH_SIZE, DataType::Int32);
330 let source: ArrayRef = Arc::new(Int32Array::from_iter_values(0..100));
333 let source_size = source.get_array_memory_size();
334 in_progress.set_source(Some(Arc::clone(&source)));
335
336 in_progress.copy_rows(0, 50).unwrap();
337 assert!(
338 in_progress.size() >= source_size + 50 * size_of::<i32>(),
339 "values buffer under-counted: {} < {} + {} * {}",
340 in_progress.size(),
341 source_size,
342 50,
343 size_of::<i32>(),
344 );
345
346 in_progress.finish().unwrap();
349 assert_eq!(in_progress.size(), source_size);
350 }
351
352 #[test]
353 fn test_size_counts_null_buffer() {
354 const BATCH_SIZE: usize = 64;
355
356 let in_progress_bytes = |source: ArrayRef| {
357 let mut in_progress =
358 InProgressPrimitiveArray::<Int32Type>::new(BATCH_SIZE, DataType::Int32);
359 let source_len = source.len();
360 in_progress.set_source(Some(source));
361 in_progress.copy_rows(0, source_len / 2).unwrap();
362 in_progress.size()
363 };
364
365 let all_valid = in_progress_bytes(Arc::new(Int32Array::from_iter_values(0..100)));
367 let with_nulls = in_progress_bytes(Arc::new(Int32Array::from_iter(
369 (0..100).map(|i| (i % 2 == 0).then_some(i)),
370 )));
371
372 assert!(
373 with_nulls > all_valid,
374 "null buffer must be included in size(): with_nulls={with_nulls} all_valid={all_valid}"
375 );
376 }
377}