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        .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        // add nulls if necessary
158        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        // Copy the values
166        let values = s.values();
167        // SAFETY: copy_rows is called with ranges derived from the source array.
168        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            // Other selection shapes reuse the generic copy_rows path.
205            selection => self.copy_rows_by_selection(selection),
206        }
207    }
208
209    fn finish(&mut self) -> Result<ArrayRef, ArrowError> {
210        // take and reset the current values and nulls
211        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            // preserve timezone / precision+scale if applicable
217            .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        // A fresh in-progress array has allocated nothing yet
311        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        // Nothing copied yet, so size is exactly the source's memory
321        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        // Non-null source: the nulls builder stays empty (allocated_size == 0),
330        // so the only growth is the values buffer.
331        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        // finish() takes the buffered values/nulls but keeps the source, so the
346        // reported size drops back to exactly the source.
347        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        // All values valid: the nulls builder never allocates.
365        let all_valid = in_progress_bytes(Arc::new(Int32Array::from_iter_values(0..100)));
366        // Some values null: copying materializes a null buffer that must count.
367        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}