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 .ok_or_else(|| {
128 ArrowError::InvalidArgumentError(
129 "Internal Error: InProgressPrimitiveArray: source not set".to_string(),
130 )
131 })?
132 .as_primitive::<T>())
133}
134
135fn append_filtered_nulls(
136 nulls: &mut NullBufferBuilder,
137 source_nulls: Option<&NullBuffer>,
138 filter: &FilterPredicate,
139) {
140 if let Some(filtered_nulls) = filter.filter_nulls(source_nulls) {
141 nulls.append_buffer(&filtered_nulls);
142 } else {
143 nulls.append_n_non_nulls(filter.count());
144 }
145}
146
147impl<T: ArrowPrimitiveType + Debug> InProgressArray for InProgressPrimitiveArray<T> {
148 fn set_source(&mut self, source: Option<ArrayRef>) {
149 self.source = source;
150 }
151
152 fn copy_rows(&mut self, offset: usize, len: usize) -> Result<(), ArrowError> {
153 self.ensure_capacity();
154
155 let s = primitive_source::<T>(self.source.as_ref())?;
156
157 if let Some(nulls) = s.nulls().as_ref() {
159 let nulls = nulls.slice(offset, len);
160 self.nulls.append_buffer(&nulls);
161 } else {
162 self.nulls.append_n_non_nulls(len);
163 }
164
165 let values = s.values();
167 self.current
169 .extend_from_slice(unsafe { values.get_unchecked(offset..offset + len) });
170
171 Ok(())
172 }
173
174 fn copy_rows_by_filter(&mut self, filter: &FilterPredicate) -> Result<(), ArrowError> {
175 match filter.selection() {
176 FilterSelection::Indices(indices) => {
177 self.ensure_capacity();
178 let s = primitive_source::<T>(self.source.as_ref())?;
179
180 append_filtered_nulls(&mut self.nulls, s.nulls(), filter);
181 self.current.reserve(filter.count());
182 Self::append_values_by_indices(
183 &mut self.current,
184 s.values(),
185 indices,
186 filter.count(),
187 );
188 Ok(())
189 }
190 FilterSelection::Slices(slices) => {
191 self.ensure_capacity();
192 let s = primitive_source::<T>(self.source.as_ref())?;
193
194 append_filtered_nulls(&mut self.nulls, s.nulls(), filter);
195 self.current.reserve(filter.count());
196 Self::append_values_by_slices(
197 &mut self.current,
198 s.values(),
199 slices,
200 filter.count(),
201 );
202 Ok(())
203 }
204 selection => self.copy_rows_by_selection(selection),
206 }
207 }
208
209 fn finish(&mut self) -> Result<ArrayRef, ArrowError> {
210 let values = std::mem::take(&mut self.current);
212 let nulls = self.nulls.finish();
213 self.nulls = NullBufferBuilder::new(self.batch_size);
214
215 let array = PrimitiveArray::<T>::try_new(ScalarBuffer::from(values), nulls)?
216 .with_data_type(self.data_type.clone());
218 Ok(Arc::new(array))
219 }
220
221 fn size(&self) -> usize {
222 self.source
223 .as_ref()
224 .map_or(0, |source| source.get_array_memory_size())
225 + self.current.capacity() * std::mem::size_of::<T::Native>()
226 + self.nulls.allocated_size()
227 }
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233 use crate::filter::FilterBuilder;
234 use arrow_array::types::Int32Type;
235 use arrow_array::{BooleanArray, Int32Array};
236
237 #[test]
238 fn test_copy_rows_by_filter_index_iterator() {
239 let source =
240 Int32Array::from_iter((0..21).map(|idx| if idx % 5 == 0 { None } else { Some(idx) }));
241 let filter = BooleanArray::from_iter(
242 (0..21).map(|idx| Some(matches!(idx, 0 | 1 | 2 | 3 | 5 | 8 | 13))),
243 );
244 let predicate = FilterBuilder::new(&filter).build();
245 let FilterSelection::Indices(indices) = predicate.selection() else {
246 panic!("expected index iterator selection");
247 };
248 let mut selected_indices = Vec::new();
249 indices.for_each(|idx| selected_indices.push(idx));
250 assert_eq!(selected_indices, vec![0, 1, 2, 3, 5, 8, 13]);
251
252 let mut in_progress = InProgressPrimitiveArray::<Int32Type>::new(7, DataType::Int32);
253 in_progress.set_source(Some(Arc::new(source)));
254 in_progress.copy_rows_by_filter(&predicate).unwrap();
255
256 let result = in_progress.finish().unwrap();
257 let result = result.as_primitive::<Int32Type>();
258 let expected = Int32Array::from(vec![
259 None,
260 Some(1),
261 Some(2),
262 Some(3),
263 None,
264 Some(8),
265 Some(13),
266 ]);
267 assert_eq!(result, &expected);
268 }
269
270 #[test]
271 fn test_copy_rows_by_filter_slice_iterator() {
272 let source =
273 Int32Array::from_iter((0..16).map(|idx| if idx % 5 == 0 { None } else { Some(idx) }));
274 let filter = BooleanArray::from_iter((0..16).map(|idx| Some(!matches!(idx, 3 | 9))));
275 let predicate = FilterBuilder::new(&filter).build();
276 let FilterSelection::Slices(slices) = predicate.selection() else {
277 panic!("expected slice iterator selection");
278 };
279 let mut selected_slices = Vec::new();
280 slices.for_each(|slice| selected_slices.push(slice));
281 assert_eq!(selected_slices, vec![(0, 3), (4, 9), (10, 16)]);
282
283 let mut in_progress = InProgressPrimitiveArray::<Int32Type>::new(14, DataType::Int32);
284 in_progress.set_source(Some(Arc::new(source)));
285 in_progress.copy_rows_by_filter(&predicate).unwrap();
286
287 let result = in_progress.finish().unwrap();
288 let result = result.as_primitive::<Int32Type>();
289 let expected = Int32Array::from(vec![
290 None,
291 Some(1),
292 Some(2),
293 Some(4),
294 None,
295 Some(6),
296 Some(7),
297 Some(8),
298 None,
299 Some(11),
300 Some(12),
301 Some(13),
302 Some(14),
303 None,
304 ]);
305 assert_eq!(result, &expected);
306 }
307
308 #[test]
309 fn test_size_empty() {
310 let in_progress = InProgressPrimitiveArray::<Int32Type>::new(64, DataType::Int32);
312 assert_eq!(in_progress.size(), 0);
313 }
314
315 #[test]
316 fn test_size_counts_source() {
317 let mut in_progress = InProgressPrimitiveArray::<Int32Type>::new(64, DataType::Int32);
318 let source: ArrayRef = Arc::new(Int32Array::from_iter_values(0..100));
319 in_progress.set_source(Some(Arc::clone(&source)));
320 assert_eq!(in_progress.size(), source.get_array_memory_size());
322 }
323
324 #[test]
325 fn test_size_counts_values_buffer_and_resets_on_finish() {
326 const BATCH_SIZE: usize = 64;
327 let mut in_progress =
328 InProgressPrimitiveArray::<Int32Type>::new(BATCH_SIZE, DataType::Int32);
329 let source: ArrayRef = Arc::new(Int32Array::from_iter_values(0..100));
332 let source_size = source.get_array_memory_size();
333 in_progress.set_source(Some(Arc::clone(&source)));
334
335 in_progress.copy_rows(0, 50).unwrap();
336 assert!(
337 in_progress.size() >= source_size + 50 * size_of::<i32>(),
338 "values buffer under-counted: {} < {} + {} * {}",
339 in_progress.size(),
340 source_size,
341 50,
342 size_of::<i32>(),
343 );
344
345 in_progress.finish().unwrap();
348 assert_eq!(in_progress.size(), source_size);
349 }
350
351 #[test]
352 fn test_size_counts_null_buffer() {
353 const BATCH_SIZE: usize = 64;
354
355 let in_progress_bytes = |source: ArrayRef| {
356 let mut in_progress =
357 InProgressPrimitiveArray::<Int32Type>::new(BATCH_SIZE, DataType::Int32);
358 let source_len = source.len();
359 in_progress.set_source(Some(source));
360 in_progress.copy_rows(0, source_len / 2).unwrap();
361 in_progress.size()
362 };
363
364 let all_valid = in_progress_bytes(Arc::new(Int32Array::from_iter_values(0..100)));
366 let with_nulls = in_progress_bytes(Arc::new(Int32Array::from_iter(
368 (0..100).map(|i| (i % 2 == 0).then_some(i)),
369 )));
370
371 assert!(
372 with_nulls > all_valid,
373 "null buffer must be included in size(): with_nulls={with_nulls} all_valid={all_valid}"
374 );
375 }
376}