1use arrow_array::builder::BufferBuilder;
21use arrow_array::*;
22use arrow_buffer::ArrowNativeType;
23use arrow_buffer::MutableBuffer;
24use arrow_buffer::buffer::NullBuffer;
25use arrow_data::ArrayData;
26use arrow_schema::ArrowError;
27
28pub fn unary<I, F, O>(array: &PrimitiveArray<I>, op: F) -> PrimitiveArray<O>
30where
31 I: ArrowPrimitiveType,
32 O: ArrowPrimitiveType,
33 F: Fn(I::Native) -> O::Native,
34{
35 array.unary(op)
36}
37
38pub fn unary_mut<I, F>(
40 array: PrimitiveArray<I>,
41 op: F,
42) -> Result<PrimitiveArray<I>, PrimitiveArray<I>>
43where
44 I: ArrowPrimitiveType,
45 F: Fn(I::Native) -> I::Native,
46{
47 array.unary_mut(op)
48}
49
50pub fn try_unary<I, F, O>(array: &PrimitiveArray<I>, op: F) -> Result<PrimitiveArray<O>, ArrowError>
52where
53 I: ArrowPrimitiveType,
54 O: ArrowPrimitiveType,
55 F: Fn(I::Native) -> Result<O::Native, ArrowError>,
56{
57 array.try_unary(op)
58}
59
60pub fn try_unary_mut<I, F>(
62 array: PrimitiveArray<I>,
63 op: F,
64) -> Result<Result<PrimitiveArray<I>, ArrowError>, PrimitiveArray<I>>
65where
66 I: ArrowPrimitiveType,
67 F: Fn(I::Native) -> Result<I::Native, ArrowError>,
68{
69 array.try_unary_mut(op)
70}
71
72pub fn binary<A, B, F, O>(
105 a: &PrimitiveArray<A>,
106 b: &PrimitiveArray<B>,
107 op: F,
108) -> Result<PrimitiveArray<O>, ArrowError>
109where
110 A: ArrowPrimitiveType,
111 B: ArrowPrimitiveType,
112 O: ArrowPrimitiveType,
113 F: Fn(A::Native, B::Native) -> O::Native,
114{
115 if a.len() != b.len() {
116 return Err(ArrowError::ComputeError(
117 "Cannot perform binary operation on arrays of different length".to_string(),
118 ));
119 }
120
121 if a.is_empty() {
122 return Ok(PrimitiveArray::from(ArrayData::new_empty(&O::DATA_TYPE)));
123 }
124
125 let nulls = NullBuffer::union(a.logical_nulls().as_ref(), b.logical_nulls().as_ref());
126
127 let values = a
128 .values()
129 .into_iter()
130 .zip(b.values())
131 .map(|(l, r)| op(*l, *r));
132
133 let buffer: Vec<_> = values.collect();
134 Ok(PrimitiveArray::new(buffer.into(), nulls))
135}
136
137pub fn binary_mut<T, U, F>(
202 a: PrimitiveArray<T>,
203 b: &PrimitiveArray<U>,
204 op: F,
205) -> Result<Result<PrimitiveArray<T>, ArrowError>, PrimitiveArray<T>>
206where
207 T: ArrowPrimitiveType,
208 U: ArrowPrimitiveType,
209 F: Fn(T::Native, U::Native) -> T::Native,
210{
211 if a.len() != b.len() {
212 return Ok(Err(ArrowError::ComputeError(
213 "Cannot perform binary operation on arrays of different length".to_string(),
214 )));
215 }
216
217 if a.is_empty() {
218 return Ok(Ok(PrimitiveArray::from(ArrayData::new_empty(
219 &T::DATA_TYPE,
220 ))));
221 }
222
223 let mut builder = a.into_builder()?;
224
225 builder
226 .values_slice_mut()
227 .iter_mut()
228 .zip(b.values())
229 .for_each(|(l, r)| *l = op(*l, *r));
230
231 let array = builder.finish();
232
233 let nulls = NullBuffer::union(array.logical_nulls().as_ref(), b.logical_nulls().as_ref());
235
236 let array_builder = array.into_data().into_builder().nulls(nulls);
237
238 let array_data = unsafe { array_builder.build_unchecked() };
239 Ok(Ok(PrimitiveArray::<T>::from(array_data)))
240}
241
242pub fn try_binary<A: ArrayAccessor, B: ArrayAccessor, F, O>(
255 a: A,
256 b: B,
257 op: F,
258) -> Result<PrimitiveArray<O>, ArrowError>
259where
260 O: ArrowPrimitiveType,
261 F: Fn(A::Item, B::Item) -> Result<O::Native, ArrowError>,
262{
263 if a.len() != b.len() {
264 return Err(ArrowError::ComputeError(
265 "Cannot perform a binary operation on arrays of different length".to_string(),
266 ));
267 }
268 if a.is_empty() {
269 return Ok(PrimitiveArray::from(ArrayData::new_empty(&O::DATA_TYPE)));
270 }
271 let len = a.len();
272
273 if !a.is_nullable() && !b.is_nullable() {
278 try_binary_no_nulls(len, a, b, op)
279 } else {
280 let Some(nulls) = NullBuffer::union(a.logical_nulls().as_ref(), b.logical_nulls().as_ref())
281 else {
282 return try_binary_no_nulls(len, a, b, op);
283 };
284
285 let mut buffer = BufferBuilder::<O::Native>::new(len);
286 buffer.append_n_zeroed(len);
287 let slice = buffer.as_slice_mut();
288
289 nulls.try_for_each_valid_idx(|idx| {
290 unsafe {
291 *slice.get_unchecked_mut(idx) = op(a.value_unchecked(idx), b.value_unchecked(idx))?
292 };
293 Ok::<_, ArrowError>(())
294 })?;
295
296 let values = buffer.finish().into();
297 Ok(PrimitiveArray::new(values, Some(nulls)))
298 }
299}
300
301pub fn try_binary_mut<T, F>(
312 a: PrimitiveArray<T>,
313 b: &PrimitiveArray<T>,
314 op: F,
315) -> Result<Result<PrimitiveArray<T>, ArrowError>, PrimitiveArray<T>>
316where
317 T: ArrowPrimitiveType,
318 F: Fn(T::Native, T::Native) -> Result<T::Native, ArrowError>,
319{
320 if a.len() != b.len() {
321 return Ok(Err(ArrowError::ComputeError(
322 "Cannot perform binary operation on arrays of different length".to_string(),
323 )));
324 }
325 let len = a.len();
326
327 if a.is_empty() {
328 return Ok(Ok(PrimitiveArray::from(ArrayData::new_empty(
329 &T::DATA_TYPE,
330 ))));
331 }
332
333 if !a.is_nullable() && !b.is_nullable() {
337 try_binary_no_nulls_mut(len, a, b, op)
338 } else {
339 let Some(nulls) =
340 create_union_null_buffer(a.logical_nulls().as_ref(), b.logical_nulls().as_ref())
341 else {
342 return try_binary_no_nulls_mut(len, a, b, op);
343 };
344
345 let mut builder = a.into_builder()?;
346
347 let slice = builder.values_slice_mut();
348
349 let r = nulls.try_for_each_valid_idx(|idx| {
350 unsafe {
351 *slice.get_unchecked_mut(idx) =
352 op(*slice.get_unchecked(idx), b.value_unchecked(idx))?
353 };
354 Ok::<_, ArrowError>(())
355 });
356 if let Err(err) = r {
357 return Ok(Err(err));
358 }
359 let array_builder = builder.finish().into_data().into_builder();
360 let array_data = unsafe { array_builder.nulls(Some(nulls)).build_unchecked() };
361 Ok(Ok(PrimitiveArray::<T>::from(array_data)))
362 }
363}
364
365fn create_union_null_buffer(
371 lhs: Option<&NullBuffer>,
372 rhs: Option<&NullBuffer>,
373) -> Option<NullBuffer> {
374 match (lhs, rhs) {
375 (Some(lhs), Some(rhs)) => Some(NullBuffer::new(lhs.inner() & rhs.inner())),
376 (Some(n), None) | (None, Some(n)) => Some(NullBuffer::new(n.inner() & n.inner())),
377 (None, None) => None,
378 }
379}
380
381#[inline(never)]
383fn try_binary_no_nulls<A: ArrayAccessor, B: ArrayAccessor, F, O>(
384 len: usize,
385 a: A,
386 b: B,
387 op: F,
388) -> Result<PrimitiveArray<O>, ArrowError>
389where
390 O: ArrowPrimitiveType,
391 F: Fn(A::Item, B::Item) -> Result<O::Native, ArrowError>,
392{
393 let mut buffer = MutableBuffer::new(len * O::Native::get_byte_width());
394 for idx in 0..len {
395 unsafe {
396 buffer.push_unchecked(op(a.value_unchecked(idx), b.value_unchecked(idx))?);
397 };
398 }
399 Ok(PrimitiveArray::new(buffer.into(), None))
400}
401
402#[inline(never)]
404fn try_binary_no_nulls_mut<T, F>(
405 len: usize,
406 a: PrimitiveArray<T>,
407 b: &PrimitiveArray<T>,
408 op: F,
409) -> Result<Result<PrimitiveArray<T>, ArrowError>, PrimitiveArray<T>>
410where
411 T: ArrowPrimitiveType,
412 F: Fn(T::Native, T::Native) -> Result<T::Native, ArrowError>,
413{
414 let mut builder = a.into_builder()?;
415 let slice = builder.values_slice_mut();
416
417 for idx in 0..len {
418 unsafe {
419 match op(*slice.get_unchecked(idx), b.value_unchecked(idx)) {
420 Ok(value) => *slice.get_unchecked_mut(idx) = value,
421 Err(err) => return Ok(Err(err)),
422 }
423 };
424 }
425 Ok(Ok(builder.finish()))
426}
427
428#[cfg(test)]
429mod tests {
430 use super::*;
431 use arrow_array::types::*;
432 use std::sync::Arc;
433
434 #[test]
435 fn test_unary_f64_slice() {
436 let input = Float64Array::from(vec![Some(5.1f64), None, Some(6.8), None, Some(7.2)]);
437 let input_slice = input.slice(1, 4);
438 let result = unary(&input_slice, |n| n.round());
439 assert_eq!(
440 result,
441 Float64Array::from(vec![None, Some(7.0), None, Some(7.0)])
442 );
443 }
444
445 #[test]
446 fn test_try_binary_run_array_logical_nulls() {
447 let run_ends = Int32Array::from(vec![1, 2, 3]);
449 let values = Int32Array::from(vec![Some(10), None, Some(30)]);
450 let run = RunArray::<Int32Type>::try_new(&run_ends, &values).expect("valid run array");
451 assert_eq!(run.null_count(), 0);
452 assert_eq!(run.logical_null_count(), 1);
453
454 let typed = run.downcast::<Int32Array>().expect("Int32 values");
455 let other = Int32Array::from(vec![1, 1, 1]);
456 let result =
457 try_binary::<_, _, _, Int32Type>(typed, &other, |a, b| Ok(a + b)).expect("no overflow");
458 assert_eq!(result, Int32Array::from(vec![Some(11), None, Some(31)]));
459 }
460
461 #[test]
462 fn test_binary_mut() {
463 let a = Int32Array::from(vec![15, 14, 9, 8, 1]);
464 let b = Int32Array::from(vec![Some(1), None, Some(3), None, Some(5)]);
465 let c = binary_mut(a, &b, |l, r| l + r).unwrap().unwrap();
466
467 let expected = Int32Array::from(vec![Some(16), None, Some(12), None, Some(6)]);
468 assert_eq!(c, expected);
469 }
470
471 #[test]
472 fn test_binary_mut_null_buffer() {
473 let a = Int32Array::from(vec![Some(3), Some(4), Some(5), Some(6), None]);
474
475 let b = Int32Array::from(vec![Some(10), Some(11), Some(12), Some(13), Some(14)]);
476
477 let r1 = binary_mut(a, &b, |a, b| a + b).unwrap();
478
479 let a = Int32Array::from(vec![Some(3), Some(4), Some(5), Some(6), None]);
480 let b = Int32Array::new(
481 vec![10, 11, 12, 13, 14].into(),
482 Some(vec![true, true, true, true, true].into()),
483 );
484
485 let r2 = binary_mut(a, &b, |a, b| a + b).unwrap();
487 assert_eq!(r1.unwrap(), r2.unwrap());
488 }
489
490 #[test]
491 fn test_try_binary_mut() {
492 let a = Int32Array::from(vec![15, 14, 9, 8, 1]);
493 let b = Int32Array::from(vec![Some(1), None, Some(3), None, Some(5)]);
494 let c = try_binary_mut(a, &b, |l, r| Ok(l + r)).unwrap().unwrap();
495
496 let expected = Int32Array::from(vec![Some(16), None, Some(12), None, Some(6)]);
497 assert_eq!(c, expected);
498
499 let a = Int32Array::from(vec![15, 14, 9, 8, 1]);
500 let b = Int32Array::from(vec![1, 2, 3, 4, 5]);
501 let c = try_binary_mut(a, &b, |l, r| Ok(l + r)).unwrap().unwrap();
502 let expected = Int32Array::from(vec![16, 16, 12, 12, 6]);
503 assert_eq!(c, expected);
504
505 let a = Int32Array::from(vec![15, 14, 9, 8, 1]);
506 let b = Int32Array::from(vec![Some(1), None, Some(3), None, Some(5)]);
507 let _ = try_binary_mut(a, &b, |l, r| {
508 if l == 1 {
509 Err(ArrowError::InvalidArgumentError(
510 "got error".parse().unwrap(),
511 ))
512 } else {
513 Ok(l + r)
514 }
515 })
516 .unwrap()
517 .expect_err("should got error");
518 }
519
520 #[test]
521 fn test_try_binary_mut_null_buffer() {
522 let a = Int32Array::from(vec![Some(3), Some(4), Some(5), Some(6), None]);
523
524 let b = Int32Array::from(vec![Some(10), Some(11), Some(12), Some(13), Some(14)]);
525
526 let r1 = try_binary_mut(a, &b, |a, b| Ok(a + b)).unwrap();
527
528 let a = Int32Array::from(vec![Some(3), Some(4), Some(5), Some(6), None]);
529 let b = Int32Array::new(
530 vec![10, 11, 12, 13, 14].into(),
531 Some(vec![true, true, true, true, true].into()),
532 );
533
534 let r2 = try_binary_mut(a, &b, |a, b| Ok(a + b)).unwrap();
536 assert_eq!(r1.unwrap(), r2.unwrap());
537 }
538
539 #[test]
540 fn test_try_binary_mut_all_valid_null_buffers() {
541 let a = Int32Array::new(vec![1, 2].into(), Some(vec![true, true].into()));
543 let b = Int32Array::new(vec![10, 20].into(), Some(vec![true, true].into()));
544 let c = try_binary_mut(a, &b, |a, b| Ok(a + b))
545 .expect("not shared")
546 .expect("no overflow");
547 assert_eq!(c, Int32Array::from(vec![11, 22]));
548 assert_eq!(c.logical_null_count(), 0);
549 }
550
551 #[test]
552 fn test_unary_dict_mut() {
553 let values = Int32Array::from(vec![Some(10), Some(20), None]);
554 let keys = Int8Array::from_iter_values([0, 0, 1, 2]);
555 let dictionary = DictionaryArray::new(keys, Arc::new(values));
556
557 let updated = dictionary.unary_mut::<_, Int32Type>(|x| x + 1).unwrap();
558 let typed = updated.downcast_dict::<Int32Array>().unwrap();
559 assert_eq!(typed.value(0), 11);
560 assert_eq!(typed.value(1), 11);
561 assert_eq!(typed.value(2), 21);
562
563 let values = updated.values();
564 assert!(values.is_null(2));
565 }
566}