Skip to main content

arrow_select/
merge.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
18//! [`merge`] and [`merge_n`]: Combine values from two or more arrays
19
20use crate::filter::{SlicesIterator, prep_null_mask_filter};
21use crate::zip::zip;
22use arrow_array::{Array, ArrayRef, BooleanArray, Datum, make_array, new_empty_array};
23use arrow_data::ArrayData;
24use arrow_data::transform::MutableArrayData;
25use arrow_schema::ArrowError;
26
27/// An index for the [merge_n] function.
28///
29/// This trait allows the indices argument for [merge_n] to be stored using a more
30/// compact representation than `usize` when the input arrays are small.
31/// If the number of input arrays is less than 256 for instance, the indices can be stored as `u8`.
32///
33/// Implementation must ensure that all values which return `None` from [MergeIndex::index] are
34/// considered equal by the [PartialEq] and [Eq] implementations.
35pub trait MergeIndex: PartialEq + Eq + Copy {
36    /// Returns the index value as an `Option<usize>`.
37    ///
38    /// `None` values returned by this function indicate holes in the index array and will result
39    /// in null values in the array created by [merge].
40    fn index(&self) -> Option<usize>;
41}
42
43impl MergeIndex for usize {
44    fn index(&self) -> Option<usize> {
45        Some(*self)
46    }
47}
48
49impl MergeIndex for Option<usize> {
50    fn index(&self) -> Option<usize> {
51        *self
52    }
53}
54
55/// Merges elements by index from a list of [`Array`], creating a new [`Array`] from
56/// those values.
57///
58/// Each element in `indices` is the index of an array in `values`. The `indices` array is processed
59/// sequentially. The first occurrence of index value `n` will be mapped to the first
60/// value of the array at index `n`. The second occurrence to the second value, and so on.
61/// An index value where `MergeIndex::index` returns `None` is interpreted as a null value.
62///
63/// # Implementation notes
64///
65/// This algorithm is similar in nature to both [zip] and
66/// [interleave](crate::interleave::interleave), but there are some important differences.
67///
68/// In contrast to [zip], this function supports multiple input arrays. Instead of
69/// a boolean selection vector, an index array is to take values from the input arrays, and a special
70/// marker values can be used to indicate null values.
71///
72/// In contrast to [interleave](crate::interleave::interleave), this function does not use pairs of
73/// indices. The values in `indices` serve the same purpose as the first value in the pairs passed
74/// to `interleave`.
75/// The index in the array is implicit and is derived from the number of times a particular array
76/// index occurs.
77/// The more constrained indexing mechanism used by this algorithm makes it easier to copy values
78/// in contiguous slices. In the example below, the two subsequent elements from array `2` can be
79/// copied in a single operation from the source array instead of copying them one by one.
80/// Long spans of null values are also especially cheap because they do not need to be represented
81/// in an input array.
82///
83/// # Panics
84///
85/// This function does not check that the number of occurrences of any particular array index matches
86/// the length of the corresponding input array. If an array contains more values than required, the
87/// spurious values will be ignored. If an array contains fewer values than necessary, this function
88/// will panic.
89///
90/// # Example
91///
92/// ```text
93/// ┌───────────┐  ┌─────────┐                             ┌─────────┐
94/// │┌─────────┐│  │   None  │                             │   NULL  │
95/// ││    A    ││  ├─────────┤                             ├─────────┤
96/// │└─────────┘│  │    1    │                             │    B    │
97/// │┌─────────┐│  ├─────────┤                             ├─────────┤
98/// ││    B    ││  │    0    │    merge(values, indices)   │    A    │
99/// │└─────────┘│  ├─────────┤  ─────────────────────────▶ ├─────────┤
100/// │┌─────────┐│  │   None  │                             │   NULL  │
101/// ││    C    ││  ├─────────┤                             ├─────────┤
102/// │├─────────┤│  │    2    │                             │    C    │
103/// ││    D    ││  ├─────────┤                             ├─────────┤
104/// │└─────────┘│  │    2    │                             │    D    │
105/// └───────────┘  └─────────┘                             └─────────┘
106///    values        indices                                  result
107///
108/// ```
109pub fn merge_n(values: &[&dyn Array], indices: &[impl MergeIndex]) -> Result<ArrayRef, ArrowError> {
110    if values.is_empty() {
111        return Err(ArrowError::InvalidArgumentError(
112            "merge_n requires at least one value array".to_string(),
113        ));
114    }
115
116    let data_type = values[0].data_type();
117
118    for array in values.iter().skip(1) {
119        if array.data_type() != data_type {
120            return Err(ArrowError::InvalidArgumentError(format!(
121                "It is not possible to merge arrays of different data types ({} and {})",
122                data_type,
123                array.data_type()
124            )));
125        }
126    }
127
128    if indices.is_empty() {
129        return Ok(new_empty_array(data_type));
130    }
131
132    #[cfg(debug_assertions)]
133    for ix in indices {
134        if let Some(index) = ix.index() {
135            assert!(
136                index < values.len(),
137                "Index out of bounds: {} >= {}",
138                index,
139                values.len()
140            );
141        }
142    }
143
144    let data: Vec<ArrayData> = values.iter().map(|a| a.to_data()).collect();
145    let data_refs = data.iter().collect();
146
147    let mut mutable = MutableArrayData::new(data_refs, true, indices.len());
148
149    // This loop extends the mutable array by taking slices from the partial results.
150    //
151    // take_offsets keeps track of how many values have been taken from each array.
152    let mut take_offsets = vec![0; values.len() + 1];
153    let mut start_row_ix = 0;
154    loop {
155        let array_ix = indices[start_row_ix];
156
157        // Determine the length of the slice to take.
158        let mut end_row_ix = start_row_ix + 1;
159        while end_row_ix < indices.len() && indices[end_row_ix] == array_ix {
160            end_row_ix += 1;
161        }
162        let slice_length = end_row_ix - start_row_ix;
163
164        // Extend mutable with either nulls or with values from the array.
165        match array_ix.index() {
166            None => mutable.try_extend_nulls(slice_length)?,
167            Some(index) => {
168                let start_offset = take_offsets[index];
169                let end_offset = start_offset + slice_length;
170                mutable.try_extend(index, start_offset, end_offset)?;
171                take_offsets[index] = end_offset;
172            }
173        }
174
175        if end_row_ix == indices.len() {
176            break;
177        }
178        // Set the start_row_ix for the next slice.
179        start_row_ix = end_row_ix;
180    }
181
182    Ok(make_array(mutable.freeze()))
183}
184
185/// Merges two arrays in the order specified by a boolean mask.
186///
187/// This algorithm is a variant of [zip] that does not require the truthy and
188/// falsy arrays to have the same length.
189///
190/// When truthy of falsy are [Scalar](arrow_array::Scalar), the single
191/// scalar value is repeated whenever the mask array contains true or false respectively.
192///
193/// # Example
194///
195/// ```text
196///  truthy
197/// ┌─────────┐  mask
198/// │    A    │  ┌─────────┐                             ┌─────────┐
199/// ├─────────┤  │  true   │                             │    A    │
200/// │    C    │  ├─────────┤                             ├─────────┤
201/// ├─────────┤  │  true   │                             │    C    │
202/// │   NULL  │  ├─────────┤                             ├─────────┤
203/// ├─────────┤  │  false  │  merge(mask, truthy, falsy) │    B    │
204/// │    D    │  ├─────────┤  ─────────────────────────▶ ├─────────┤
205/// └─────────┘  │  true   │                             │   NULL  │
206///  falsy       ├─────────┤                             ├─────────┤
207/// ┌─────────┐  │  false  │                             │    E    │
208/// │    B    │  ├─────────┤                             ├─────────┤
209/// ├─────────┤  │  true   │                             │    D    │
210/// │    E    │  └─────────┘                             └─────────┘
211/// └─────────┘
212/// ```
213pub fn merge(
214    mask: &BooleanArray,
215    truthy: &dyn Datum,
216    falsy: &dyn Datum,
217) -> Result<ArrayRef, ArrowError> {
218    let (truthy_array, truthy_is_scalar) = truthy.get();
219    let (falsy_array, falsy_is_scalar) = falsy.get();
220
221    if truthy_is_scalar && falsy_is_scalar {
222        // When both truthy and falsy are scalars, we can use `zip` since the result is the same
223        // and zip has optimized code for scalars.
224        return zip(mask, truthy, falsy);
225    }
226
227    if truthy_array.data_type() != falsy_array.data_type() {
228        return Err(ArrowError::InvalidArgumentError(
229            "arguments need to have the same data type".into(),
230        ));
231    }
232
233    if truthy_is_scalar && truthy_array.len() != 1 {
234        return Err(ArrowError::InvalidArgumentError(
235            "scalar arrays must have 1 element".into(),
236        ));
237    }
238    if falsy_is_scalar && falsy_array.len() != 1 {
239        return Err(ArrowError::InvalidArgumentError(
240            "scalar arrays must have 1 element".into(),
241        ));
242    }
243
244    let falsy = falsy_array.to_data();
245    let truthy = truthy_array.to_data();
246
247    let mut mutable = MutableArrayData::new(vec![&truthy, &falsy], false, mask.len());
248
249    // the SlicesIterator slices only the true values. So the gaps left by this iterator we need to
250    // fill with falsy values
251
252    // keep track of how much is filled
253    let mut filled = 0;
254    let mut falsy_offset = 0;
255    let mut truthy_offset = 0;
256
257    // Ensure nulls are treated as false
258    let mask_buffer = match mask.null_count() {
259        0 => mask.values().clone(),
260        _ => prep_null_mask_filter(mask).into_parts().0,
261    };
262
263    for (start, end) in SlicesIterator::from(&mask_buffer) {
264        // the gap needs to be filled with falsy values
265        if start > filled {
266            if falsy_is_scalar {
267                for _ in filled..start {
268                    // Copy the first item from the 'falsy' array into the output buffer.
269                    mutable.try_extend(1, 0, 1)?;
270                }
271            } else {
272                let falsy_length = start - filled;
273                let falsy_end = falsy_offset + falsy_length;
274                mutable.try_extend(1, falsy_offset, falsy_end)?;
275                falsy_offset = falsy_end;
276            }
277        }
278        // fill with truthy values
279        if truthy_is_scalar {
280            for _ in start..end {
281                // Copy the first item from the 'truthy' array into the output buffer.
282                mutable.try_extend(0, 0, 1)?;
283            }
284        } else {
285            let truthy_length = end - start;
286            let truthy_end = truthy_offset + truthy_length;
287            mutable.try_extend(0, truthy_offset, truthy_end)?;
288            truthy_offset = truthy_end;
289        }
290        filled = end;
291    }
292    // the remaining part is falsy
293    if filled < mask.len() {
294        if falsy_is_scalar {
295            for _ in filled..mask.len() {
296                // Copy the first item from the 'falsy' array into the output buffer.
297                mutable.try_extend(1, 0, 1)?;
298            }
299        } else {
300            let falsy_length = mask.len() - filled;
301            let falsy_end = falsy_offset + falsy_length;
302            mutable.try_extend(1, falsy_offset, falsy_end)?;
303        }
304    }
305
306    let data = mutable.freeze();
307    Ok(make_array(data))
308}
309
310#[cfg(test)]
311mod tests {
312    use crate::merge::{MergeIndex, merge, merge_n};
313    use arrow_array::cast::AsArray;
314    use arrow_array::{Array, BooleanArray, Datum, Int32Array, Scalar, StringArray, UInt64Array};
315    use arrow_schema::ArrowError::InvalidArgumentError;
316
317    #[derive(PartialEq, Eq, Copy, Clone)]
318    struct CompactMergeIndex {
319        index: u8,
320    }
321
322    impl MergeIndex for CompactMergeIndex {
323        fn index(&self) -> Option<usize> {
324            if self.index == u8::MAX {
325                None
326            } else {
327                Some(self.index as usize)
328            }
329        }
330    }
331
332    #[test]
333    fn test_merge() {
334        let a1 = StringArray::from(vec![Some("A"), Some("B"), Some("E"), None]);
335        let a2 = StringArray::from(vec![Some("C"), Some("D")]);
336
337        let indices = BooleanArray::from(vec![true, false, true, false, true, true]);
338
339        let merged = merge(&indices, &a1, &a2).unwrap();
340        let merged = merged.as_string::<i32>();
341
342        assert_eq!(merged.len(), indices.len());
343        assert!(merged.is_valid(0));
344        assert_eq!(merged.value(0), "A");
345        assert!(merged.is_valid(1));
346        assert_eq!(merged.value(1), "C");
347        assert!(merged.is_valid(2));
348        assert_eq!(merged.value(2), "B");
349        assert!(merged.is_valid(3));
350        assert_eq!(merged.value(3), "D");
351        assert!(merged.is_valid(4));
352        assert_eq!(merged.value(4), "E");
353        assert!(!merged.is_valid(5));
354    }
355
356    #[test]
357    fn test_merge_null_is_false() {
358        let a1 = StringArray::from(vec![Some("A"), Some("B"), Some("E"), None]);
359        let a2 = StringArray::from(vec![Some("C"), Some("D")]);
360
361        let indices = BooleanArray::from(vec![
362            Some(true),
363            None,
364            Some(true),
365            None,
366            Some(true),
367            Some(true),
368        ]);
369
370        let merged = merge(&indices, &a1, &a2).unwrap();
371        let merged = merged.as_string::<i32>();
372
373        assert_eq!(merged.len(), indices.len());
374        assert!(merged.is_valid(0));
375        assert_eq!(merged.value(0), "A");
376        assert!(merged.is_valid(1));
377        assert_eq!(merged.value(1), "C");
378        assert!(merged.is_valid(2));
379        assert_eq!(merged.value(2), "B");
380        assert!(merged.is_valid(3));
381        assert_eq!(merged.value(3), "D");
382        assert!(merged.is_valid(4));
383        assert_eq!(merged.value(4), "E");
384        assert!(!merged.is_valid(5));
385    }
386
387    #[test]
388    fn test_merge_false_tail() {
389        let a1 = StringArray::from(vec![Some("A"), Some("B"), Some("E"), None]);
390        let a2 = StringArray::from(vec![Some("C"), Some("D"), None, Some("F")]);
391
392        let indices = BooleanArray::from(vec![true, false, true, false, true, true, false, false]);
393
394        let merged = merge(&indices, &a1, &a2).unwrap();
395        let merged = merged.as_string::<i32>();
396
397        assert_eq!(merged.len(), indices.len());
398        assert!(merged.is_valid(0));
399        assert_eq!(merged.value(0), "A");
400        assert!(merged.is_valid(1));
401        assert_eq!(merged.value(1), "C");
402        assert!(merged.is_valid(2));
403        assert_eq!(merged.value(2), "B");
404        assert!(merged.is_valid(3));
405        assert_eq!(merged.value(3), "D");
406        assert!(merged.is_valid(4));
407        assert_eq!(merged.value(4), "E");
408        assert!(!merged.is_valid(5));
409        assert!(!merged.is_valid(6));
410        assert!(merged.is_valid(7));
411        assert_eq!(merged.value(7), "F");
412    }
413
414    #[test]
415    fn test_merge_scalars() {
416        let truthy = Scalar::new(StringArray::from(vec![Some("A")]));
417        let falsy = Scalar::new(StringArray::from(vec![Some("B")]));
418
419        let mask = BooleanArray::from(vec![true, false, false, true]);
420
421        let merged = merge(&mask, &truthy, &falsy).unwrap();
422        let merged = merged.as_string::<i32>();
423
424        assert_eq!(merged.len(), mask.len());
425        assert!(merged.is_valid(0));
426        assert_eq!(merged.value(0), "A");
427        assert!(merged.is_valid(1));
428        assert_eq!(merged.value(1), "B");
429        assert!(merged.is_valid(2));
430        assert_eq!(merged.value(2), "B");
431        assert!(merged.is_valid(3));
432        assert_eq!(merged.value(3), "A");
433    }
434
435    #[test]
436    fn test_merge_scalar_and_array() {
437        let truthy = Scalar::new(StringArray::from(vec![Some("A")]));
438        let falsy = StringArray::from(vec![Some("B"), Some("C")]);
439
440        let mask = BooleanArray::from(vec![true, false, false, true]);
441
442        let merged = merge(&mask, &truthy, &falsy).unwrap();
443        let merged = merged.as_string::<i32>();
444
445        assert_eq!(merged.len(), mask.len());
446        assert!(merged.is_valid(0));
447        assert_eq!(merged.value(0), "A");
448        assert!(merged.is_valid(1));
449        assert_eq!(merged.value(1), "B");
450        assert!(merged.is_valid(2));
451        assert_eq!(merged.value(2), "C");
452        assert!(merged.is_valid(3));
453        assert_eq!(merged.value(3), "A");
454    }
455
456    #[test]
457    fn test_merge_array_and_scalar() {
458        let truthy = StringArray::from(vec![Some("B"), Some("C")]);
459        let falsy = Scalar::new(StringArray::from(vec![Some("A")]));
460
461        let mask = BooleanArray::from(vec![true, false, false, true, false, false]);
462
463        let merged = merge(&mask, &truthy, &falsy).unwrap();
464        let merged = merged.as_string::<i32>();
465
466        assert_eq!(merged.len(), mask.len());
467        assert!(merged.is_valid(0));
468        assert_eq!(merged.value(0), "B");
469        assert!(merged.is_valid(1));
470        assert_eq!(merged.value(1), "A");
471        assert!(merged.is_valid(2));
472        assert_eq!(merged.value(2), "A");
473        assert!(merged.is_valid(3));
474        assert_eq!(merged.value(3), "C");
475        assert!(merged.is_valid(4));
476        assert_eq!(merged.value(4), "A");
477        assert!(merged.is_valid(5));
478        assert_eq!(merged.value(5), "A");
479    }
480
481    #[test]
482    fn test_merge_empty_mask() {
483        let a1 = StringArray::from(vec![Some("A")]);
484        let a2 = StringArray::from(vec![Some("B")]);
485        let mask: Vec<bool> = vec![];
486        let mask = BooleanArray::from(mask);
487        let result = merge(&mask, &a1, &a2).unwrap();
488        assert_eq!(result.len(), 0);
489    }
490
491    #[derive(Debug, Copy, Clone)]
492    pub struct UnsafeScalar<T: Array>(T);
493
494    impl<T: Array> Datum for UnsafeScalar<T> {
495        fn get(&self) -> (&dyn Array, bool) {
496            (&self.0, true)
497        }
498    }
499
500    #[test]
501    fn test_merge_invalid_truthy_scalar() {
502        let truthy = UnsafeScalar(StringArray::from(vec![Some("A"), Some("C")]));
503        let falsy = StringArray::from(vec![Some("B"), Some("D")]);
504        let mask = BooleanArray::from(vec![true, false, true, false]);
505        let merged = merge(&mask, &truthy, &falsy);
506        assert!(matches!(merged, Err(InvalidArgumentError { .. })));
507    }
508
509    #[test]
510    fn test_merge_invalid_falsy_scalar() {
511        let truthy = StringArray::from(vec![Some("A"), Some("C")]);
512        let falsy = UnsafeScalar(StringArray::from(vec![Some("B"), Some("D")]));
513        let mask = vec![true, false, true, false];
514        let mask = BooleanArray::from(mask);
515        let merged = merge(&mask, &truthy, &falsy);
516        assert!(matches!(merged, Err(InvalidArgumentError { .. })));
517    }
518
519    #[test]
520    fn test_merge_incompatible_arrays() {
521        let truthy = StringArray::from(vec![Some("A"), Some("B")]);
522        let falsy = Int32Array::from(vec![1, 2]);
523        let mask = BooleanArray::from(vec![true, false, true, false]);
524        let merged = merge(&mask, &truthy, &falsy);
525        assert!(matches!(merged, Err(InvalidArgumentError { .. })));
526    }
527
528    #[test]
529    fn test_merge_n() {
530        let a1 = StringArray::from(vec![Some("A")]);
531        let a2 = StringArray::from(vec![Some("B"), None, None]);
532        let a3 = StringArray::from(vec![Some("C"), Some("D")]);
533
534        let indices = vec![
535            CompactMergeIndex { index: u8::MAX },
536            CompactMergeIndex { index: 1 },
537            CompactMergeIndex { index: 0 },
538            CompactMergeIndex { index: u8::MAX },
539            CompactMergeIndex { index: 2 },
540            CompactMergeIndex { index: 2 },
541            CompactMergeIndex { index: 1 },
542            CompactMergeIndex { index: 1 },
543        ];
544
545        let arrays = [a1, a2, a3];
546        let array_refs = arrays.iter().map(|a| a as &dyn Array).collect::<Vec<_>>();
547        let merged = merge_n(&array_refs, &indices).unwrap();
548        let merged = merged.as_string::<i32>();
549
550        assert_eq!(merged.len(), indices.len());
551        assert!(!merged.is_valid(0));
552        assert!(merged.is_valid(1));
553        assert_eq!(merged.value(1), "B");
554        assert!(merged.is_valid(2));
555        assert_eq!(merged.value(2), "A");
556        assert!(!merged.is_valid(3));
557        assert!(merged.is_valid(4));
558        assert_eq!(merged.value(4), "C");
559        assert!(merged.is_valid(5));
560        assert_eq!(merged.value(5), "D");
561        assert!(!merged.is_valid(6));
562        assert!(!merged.is_valid(7));
563    }
564
565    #[test]
566    // The message differs between debug and release: in release the
567    // `cfg(debug_assertions)` bounds check in `merge_n` is compiled out and the
568    // slice index panics instead.
569    #[should_panic(expected = "out of bounds")]
570    fn test_merge_n_invalid_indices() {
571        let a1 = StringArray::from(vec![Some("A")]);
572
573        let indices = vec![CompactMergeIndex { index: 99 }];
574
575        let arrays = [a1];
576        let array_refs = arrays.iter().map(|a| a as &dyn Array).collect::<Vec<_>>();
577        let _ = merge_n(&array_refs, &indices);
578    }
579
580    #[test]
581    fn test_merge_n_empty_indices() {
582        let a1 = StringArray::from(vec![Some("A")]);
583        let a2 = StringArray::from(vec![Some("B"), None, None]);
584        let a3 = StringArray::from(vec![Some("C"), Some("D")]);
585
586        let indices: Vec<CompactMergeIndex> = vec![];
587
588        let arrays = [a1, a2, a3];
589        let array_refs = arrays.iter().map(|a| a as &dyn Array).collect::<Vec<_>>();
590        let merged = merge_n(&array_refs, &indices).unwrap();
591
592        assert_eq!(merged.len(), indices.len());
593    }
594
595    #[test]
596    fn test_merge_n_empty_values() {
597        let indices: Vec<CompactMergeIndex> = vec![];
598
599        let arrays: Vec<&dyn Array> = vec![];
600        let merged = merge_n(&arrays, &indices);
601
602        assert!(matches!(merged, Err(InvalidArgumentError { .. })));
603    }
604
605    #[test]
606    fn test_merge_n_incompatible_arrays() {
607        let a1: Box<dyn Array> = Box::new(StringArray::from(vec![Some("A")]));
608        let a2: Box<dyn Array> = Box::new(Int32Array::from(vec![1, 2, 3]));
609        let a3: Box<dyn Array> = Box::new(UInt64Array::from(vec![42, 314]));
610
611        let indices: Vec<CompactMergeIndex> = vec![];
612
613        let arrays = [a1.as_ref(), a2.as_ref(), a3.as_ref()];
614        let merged = merge_n(&arrays, &indices);
615
616        assert!(matches!(merged, Err(InvalidArgumentError { .. })));
617    }
618}