datafusion_spark/function/array/
repeat.rs1use arrow::datatypes::{DataType, Field};
19use datafusion_common::utils::take_function_args;
20use datafusion_common::{Result, ScalarValue, exec_err};
21use datafusion_expr::{
22 ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility,
23};
24use datafusion_functions_nested::repeat::ArrayRepeat;
25use std::sync::Arc;
26
27use crate::function::null_utils::{
28 NullMaskResolution, apply_null_mask, compute_null_mask,
29};
30
31#[derive(Debug, PartialEq, Eq, Hash)]
34pub struct SparkArrayRepeat {
35 signature: Signature,
36}
37
38impl Default for SparkArrayRepeat {
39 fn default() -> Self {
40 Self::new()
41 }
42}
43
44impl SparkArrayRepeat {
45 pub fn new() -> Self {
46 Self {
47 signature: Signature::user_defined(Volatility::Immutable),
48 }
49 }
50}
51
52impl ScalarUDFImpl for SparkArrayRepeat {
53 fn name(&self) -> &str {
54 "array_repeat"
55 }
56
57 fn signature(&self) -> &Signature {
58 &self.signature
59 }
60
61 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
62 Ok(DataType::List(Arc::new(Field::new_list_field(
63 arg_types[0].clone(),
64 true,
65 ))))
66 }
67
68 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
69 spark_array_repeat(args)
70 }
71
72 fn coerce_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
73 let [first_type, second_type] = take_function_args(self.name(), arg_types)?;
74
75 let second = match second_type {
77 DataType::Int8
78 | DataType::Int16
79 | DataType::Int32
80 | DataType::Int64
81 | DataType::Null => DataType::Int64,
82 DataType::UInt8 | DataType::UInt16 | DataType::UInt32 | DataType::UInt64 => {
83 DataType::UInt64
84 }
85 _ => return exec_err!("count must be an integer type"),
86 };
87
88 Ok(vec![first_type.clone(), second])
89 }
90}
91
92fn spark_array_repeat(args: ScalarFunctionArgs) -> Result<ColumnarValue> {
95 let ScalarFunctionArgs {
96 args: arg_values,
97 arg_fields,
98 number_rows,
99 return_field,
100 config_options,
101 } = args;
102 let return_type = return_field.data_type().clone();
103
104 let null_mask = compute_null_mask(&arg_values[1..]);
106
107 if matches!(null_mask, NullMaskResolution::ReturnNull) {
109 return Ok(ColumnarValue::Scalar(ScalarValue::try_from(return_type)?));
110 }
111
112 let array_repeat_func = ArrayRepeat::new();
113 let func_args = ScalarFunctionArgs {
114 args: arg_values,
115 arg_fields,
116 number_rows,
117 return_field,
118 config_options,
119 };
120 let result = array_repeat_func.invoke_with_args(func_args)?;
121
122 apply_null_mask(result, null_mask, &return_type)
123}