1use 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
27pub trait MergeIndex: PartialEq + Eq + Copy {
36 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
55pub 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 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 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 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 start_row_ix = end_row_ix;
180 }
181
182 Ok(make_array(mutable.freeze()))
183}
184
185pub 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 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 let mut filled = 0;
254 let mut falsy_offset = 0;
255 let mut truthy_offset = 0;
256
257 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 if start > filled {
266 if falsy_is_scalar {
267 for _ in filled..start {
268 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 if truthy_is_scalar {
280 for _ in start..end {
281 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 if filled < mask.len() {
294 if falsy_is_scalar {
295 for _ in filled..mask.len() {
296 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 #[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}