Skip to main content

arrow_select/coalesce/
primitive.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use 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/// InProgressArray for [`PrimitiveArray`]
28#[derive(Debug)]
29pub(crate) struct InProgressPrimitiveArray<T: ArrowPrimitiveType> {
30    /// Data type of the array
31    data_type: DataType,
32    /// The current source, if any
33    source: Option<ArrayRef>,
34    /// the target batch size (and thus size for views allocation)
35    batch_size: usize,
36    /// In progress nulls
37    nulls: NullBufferBuilder,
38    /// The currently in progress array
39    current: Vec<T::Native>,
40}
41
42impl<T: ArrowPrimitiveType> InProgressPrimitiveArray<T> {
43    /// Create a new `InProgressPrimitiveArray`
44    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    /// Allocate space for output values if necessary.
55    ///
56    /// This is done on write (when we know it is necessary) rather than
57    /// eagerly to avoid allocations that are not used.
58    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                // SAFETY: indices are derived from the filter predicate for this source.
81                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                // SAFETY: slices are derived from the filter predicate for this source.
110                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        // add nulls if necessary
159        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        // Copy the values
167        let values = s.values();
168        // SAFETY: copy_rows is called with ranges derived from the source array.
169        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            // Other selection shapes reuse the generic copy_rows path.
206            selection => self.copy_rows_by_selection(selection),
207        }
208    }
209
210    fn finish(&mut self) -> Result<ArrayRef, ArrowError> {
211        // take and reset the current values and nulls
212        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            // preserve timezone / precision+scale if applicable
218            .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        // A fresh in-progress array has allocated nothing yet
312        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        // Nothing copied yet, so size is exactly the source's memory
322        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        // Non-null source: the nulls builder stays empty (allocated_size == 0),
331        // so the only growth is the values buffer.
332        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        // finish() takes the buffered values/nulls but keeps the source, so the
347        // reported size drops back to exactly the source.
348        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        // All values valid: the nulls builder never allocates.
366        let all_valid = in_progress_bytes(Arc::new(Int32Array::from_iter_values(0..100)));
367        // Some values null: copying materializes a null buffer that must count.
368        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}