use arrow::array::{
Array, ArrayRef, AsArray, Capacities, GenericListArray, MutableArrayData,
NullBufferBuilder, OffsetSizeTrait, Scalar, new_null_array,
};
use arrow::buffer::OffsetBuffer;
use arrow::datatypes::{DataType, Field, FieldRef};
use datafusion_common::cast::as_int64_array;
use datafusion_common::utils::ListCoercion;
use datafusion_common::{
Result, ScalarValue, exec_err, internal_err, utils::take_function_args,
};
use datafusion_expr::{
ArrayFunctionArgument, ArrayFunctionSignature, ColumnarValue, Documentation,
ReturnFieldArgs, ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignature,
Volatility,
};
use datafusion_macros::user_doc;
use crate::utils::{compare_element_to_list, list_inner_field, list_type_with_element};
use std::sync::Arc;
make_udf_expr_and_func!(ArrayReplace,
array_replace,
array from to,
"replaces the first occurrence of the specified element with another specified element.",
array_replace_udf
);
make_udf_expr_and_func!(ArrayReplaceN,
array_replace_n,
array from to max,
"replaces the first `max` occurrences of the specified element with another specified element.",
array_replace_n_udf
);
make_udf_expr_and_func!(ArrayReplaceAll,
array_replace_all,
array from to,
"replaces all occurrences of the specified element with another specified element.",
array_replace_all_udf
);
#[user_doc(
doc_section(label = "Array Functions"),
description = "Replaces the first occurrence of the specified element with another specified element.",
syntax_example = "array_replace(array, from, to)",
sql_example = r#"```sql
> select array_replace([1, 2, 2, 3, 2, 1, 4], 2, 5);
+--------------------------------------------------------+
| array_replace(List([1,2,2,3,2,1,4]),Int64(2),Int64(5)) |
+--------------------------------------------------------+
| [1, 5, 2, 3, 2, 1, 4] |
+--------------------------------------------------------+
```"#,
argument(
name = "array",
description = "Array expression. Can be a constant, column, or function, and any combination of array operators."
),
argument(name = "from", description = "Initial element."),
argument(name = "to", description = "Final element.")
)]
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct ArrayReplace {
signature: Signature,
aliases: Vec<String>,
}
impl Default for ArrayReplace {
fn default() -> Self {
Self::new()
}
}
impl ArrayReplace {
pub fn new() -> Self {
Self {
signature: Signature {
type_signature: TypeSignature::ArraySignature(
ArrayFunctionSignature::Array {
arguments: vec![
ArrayFunctionArgument::Array,
ArrayFunctionArgument::Element,
ArrayFunctionArgument::Element,
],
array_coercion: Some(ListCoercion::FixedSizedListToList),
},
),
volatility: Volatility::Immutable,
parameter_names: None,
},
aliases: vec![String::from("list_replace")],
}
}
}
impl ScalarUDFImpl for ArrayReplace {
fn name(&self) -> &str {
"array_replace"
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
internal_err!("return_field_from_args should be used instead")
}
fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result<FieldRef> {
replace_return_field(self.name(), args.arg_fields)
}
fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let return_type = args.return_field.data_type().clone();
let [list_arg, from_arg, to_arg] = take_function_args(self.name(), &args.args)?;
let num_rows = args.number_rows;
let list_array = list_arg.to_array(num_rows)?;
match (from_arg, to_arg) {
(ColumnarValue::Scalar(scalar_from), ColumnarValue::Scalar(scalar_to)) => {
let result = array_replace_with_scalar_args(
self.name(),
&list_array,
scalar_from,
scalar_to,
1i64,
&return_type,
)?;
Ok(ColumnarValue::Array(result))
}
(from_arg, to_arg) => {
let from_array = from_arg.to_array(num_rows)?;
let to_array = to_arg.to_array(num_rows)?;
let result = array_replace_internal(
self.name(),
&list_array,
&from_array,
&to_array,
&[Some(1)],
&return_type,
)?;
Ok(ColumnarValue::Array(result))
}
}
}
fn aliases(&self) -> &[String] {
&self.aliases
}
fn documentation(&self) -> Option<&Documentation> {
self.doc()
}
}
#[user_doc(
doc_section(label = "Array Functions"),
description = "Replaces the first `max` occurrences of the specified element with another specified element.",
syntax_example = "array_replace_n(array, from, to, max)",
sql_example = r#"```sql
> select array_replace_n([1, 2, 2, 3, 2, 1, 4], 2, 5, 2);
+-------------------------------------------------------------------+
| array_replace_n(List([1,2,2,3,2,1,4]),Int64(2),Int64(5),Int64(2)) |
+-------------------------------------------------------------------+
| [1, 5, 5, 3, 2, 1, 4] |
+-------------------------------------------------------------------+
```"#,
argument(
name = "array",
description = "Array expression. Can be a constant, column, or function, and any combination of array operators."
),
argument(name = "from", description = "Initial element."),
argument(name = "to", description = "Final element."),
argument(name = "max", description = "Number of first occurrences to replace.")
)]
#[derive(Debug, PartialEq, Eq, Hash)]
pub(super) struct ArrayReplaceN {
signature: Signature,
aliases: Vec<String>,
}
impl ArrayReplaceN {
pub fn new() -> Self {
Self {
signature: Signature {
type_signature: TypeSignature::ArraySignature(
ArrayFunctionSignature::Array {
arguments: vec![
ArrayFunctionArgument::Array,
ArrayFunctionArgument::Element,
ArrayFunctionArgument::Element,
ArrayFunctionArgument::Index,
],
array_coercion: Some(ListCoercion::FixedSizedListToList),
},
),
volatility: Volatility::Immutable,
parameter_names: None,
},
aliases: vec![String::from("list_replace_n")],
}
}
}
impl ScalarUDFImpl for ArrayReplaceN {
fn name(&self) -> &str {
"array_replace_n"
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
internal_err!("return_field_from_args should be used instead")
}
fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result<FieldRef> {
replace_return_field(self.name(), args.arg_fields)
}
fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let return_type = args.return_field.data_type().clone();
let [list_arg, from_arg, to_arg, max_arg] =
take_function_args(self.name(), &args.args)?;
let num_rows = args.number_rows;
let list_array = list_arg.to_array(num_rows)?;
match (from_arg, to_arg, max_arg) {
(
ColumnarValue::Scalar(scalar_from),
ColumnarValue::Scalar(scalar_to),
ColumnarValue::Scalar(scalar_max),
) => {
let ScalarValue::Int64(Some(n)) = scalar_max else {
return Ok(ColumnarValue::Array(new_null_array(
&return_type,
num_rows,
)));
};
let result = array_replace_with_scalar_args(
self.name(),
&list_array,
scalar_from,
scalar_to,
*n,
&return_type,
)?;
Ok(ColumnarValue::Array(result))
}
(from_arg, to_arg, max_arg) => {
let from_array = from_arg.to_array(num_rows)?;
let to_array = to_arg.to_array(num_rows)?;
let max_array = max_arg.to_array(num_rows)?;
let result = array_replace_n_inner(
self.name(),
&list_array,
&from_array,
&to_array,
&max_array,
&return_type,
)?;
Ok(ColumnarValue::Array(result))
}
}
}
fn aliases(&self) -> &[String] {
&self.aliases
}
fn documentation(&self) -> Option<&Documentation> {
self.doc()
}
}
#[user_doc(
doc_section(label = "Array Functions"),
description = "Replaces all occurrences of the specified element with another specified element.",
syntax_example = "array_replace_all(array, from, to)",
sql_example = r#"```sql
> select array_replace_all([1, 2, 2, 3, 2, 1, 4], 2, 5);
+------------------------------------------------------------+
| array_replace_all(List([1,2,2,3,2,1,4]),Int64(2),Int64(5)) |
+------------------------------------------------------------+
| [1, 5, 5, 3, 5, 1, 4] |
+------------------------------------------------------------+
```"#,
argument(
name = "array",
description = "Array expression. Can be a constant, column, or function, and any combination of array operators."
),
argument(name = "from", description = "Initial element."),
argument(name = "to", description = "Final element.")
)]
#[derive(Debug, PartialEq, Eq, Hash)]
pub(super) struct ArrayReplaceAll {
signature: Signature,
aliases: Vec<String>,
}
impl ArrayReplaceAll {
pub fn new() -> Self {
Self {
signature: Signature {
type_signature: TypeSignature::ArraySignature(
ArrayFunctionSignature::Array {
arguments: vec![
ArrayFunctionArgument::Array,
ArrayFunctionArgument::Element,
ArrayFunctionArgument::Element,
],
array_coercion: Some(ListCoercion::FixedSizedListToList),
},
),
volatility: Volatility::Immutable,
parameter_names: None,
},
aliases: vec![String::from("list_replace_all")],
}
}
}
impl ScalarUDFImpl for ArrayReplaceAll {
fn name(&self) -> &str {
"array_replace_all"
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
internal_err!("return_field_from_args should be used instead")
}
fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result<FieldRef> {
replace_return_field(self.name(), args.arg_fields)
}
fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let return_type = args.return_field.data_type().clone();
let [list_arg, from_arg, to_arg] = take_function_args(self.name(), &args.args)?;
let num_rows = args.number_rows;
let list_array = list_arg.to_array(num_rows)?;
match (from_arg, to_arg) {
(ColumnarValue::Scalar(scalar_from), ColumnarValue::Scalar(scalar_to)) => {
let result = array_replace_with_scalar_args(
self.name(),
&list_array,
scalar_from,
scalar_to,
i64::MAX,
&return_type,
)?;
Ok(ColumnarValue::Array(result))
}
(from_arg, to_arg) => {
let from_array = from_arg.to_array(num_rows)?;
let to_array = to_arg.to_array(num_rows)?;
let result = array_replace_internal(
self.name(),
&list_array,
&from_array,
&to_array,
&[Some(i64::MAX)],
&return_type,
)?;
Ok(ColumnarValue::Array(result))
}
}
}
fn aliases(&self) -> &[String] {
&self.aliases
}
fn documentation(&self) -> Option<&Documentation> {
self.doc()
}
}
fn replace_return_field(name: &str, arg_fields: &[FieldRef]) -> Result<FieldRef> {
let [array_field, _from_field, to_field, ..] = arg_fields else {
return exec_err!(
"{name} expects at least 3 arguments, got {}",
arg_fields.len()
);
};
let data_type =
list_type_with_element(array_field.data_type(), to_field.is_nullable());
Ok(Arc::new(Field::new(name, data_type, true)))
}
fn general_replace<O: OffsetSizeTrait>(
list_array: &GenericListArray<O>,
from_array: &ArrayRef,
to_array: &ArrayRef,
arr_n: &[Option<i64>],
field: FieldRef,
) -> Result<ArrayRef> {
let mut offsets: Vec<O> = Vec::with_capacity(list_array.len() + 1);
offsets.push(O::usize_as(0));
let values = list_array.values();
let original_data = values.to_data();
let to_data = to_array.to_data();
let capacity = Capacities::Array(original_data.len());
let mut mutable = MutableArrayData::with_capacities(
vec![&original_data, &to_data],
false,
capacity,
);
let mut valid = NullBufferBuilder::new(list_array.len());
for (row_index, offset_window) in list_array.offsets().windows(2).enumerate() {
if list_array.is_null(row_index) {
offsets.push(offsets[row_index]);
valid.append_null();
continue;
}
let n = if arr_n.len() == 1 {
arr_n[0]
} else {
arr_n[row_index]
};
let Some(n) = n else {
offsets.push(offsets[row_index]);
valid.append_null();
continue;
};
let start = offset_window[0];
let end = offset_window[1];
let list_array_row = list_array.value(row_index);
let eq_array =
compare_element_to_list(&list_array_row, &from_array, row_index, true)?;
let original_idx = O::usize_as(0);
let replace_idx = O::usize_as(1);
let mut counter = 0;
if n <= 0 || !eq_array.has_true() {
mutable.try_extend(
original_idx.to_usize().unwrap(),
start.to_usize().unwrap(),
end.to_usize().unwrap(),
)?;
offsets.push(offsets[row_index] + (end - start));
valid.append_non_null();
continue;
}
let mut pending_retain: Option<O> = None;
for (i, to_replace) in eq_array.iter().enumerate() {
let i = O::usize_as(i);
if to_replace == Some(true) && counter < n {
if let Some(rs) = pending_retain.take() {
mutable.try_extend(
original_idx.to_usize().unwrap(),
(start + rs).to_usize().unwrap(),
(start + i).to_usize().unwrap(),
)?;
}
mutable.try_extend(
replace_idx.to_usize().unwrap(),
row_index,
row_index + 1,
)?;
counter += 1;
if counter == n {
mutable.try_extend(
original_idx.to_usize().unwrap(),
(start + i).to_usize().unwrap() + 1,
end.to_usize().unwrap(),
)?;
break;
}
} else if pending_retain.is_none() {
pending_retain = Some(i);
}
}
if counter < n
&& let Some(rs) = pending_retain
{
mutable.try_extend(
original_idx.to_usize().unwrap(),
(start + rs).to_usize().unwrap(),
end.to_usize().unwrap(),
)?;
}
offsets.push(offsets[row_index] + (end - start));
valid.append_non_null();
}
let data = mutable.freeze();
Ok(Arc::new(GenericListArray::<O>::try_new(
field,
OffsetBuffer::<O>::new(offsets.into()),
arrow::array::make_array(data),
valid.finish(),
)?))
}
fn general_replace_with_scalar<O: OffsetSizeTrait>(
list_array: &GenericListArray<O>,
needle: &Scalar<ArrayRef>,
scalar_to: &ScalarValue,
max_replacements: i64,
field: FieldRef,
) -> Result<ArrayRef> {
if max_replacements <= 0 {
return Ok(Arc::new(GenericListArray::<O>::try_new(
field,
list_array.offsets().clone(),
Arc::clone(list_array.values()),
list_array.nulls().cloned(),
)?));
}
let first_offset = list_array.offsets()[0].to_usize().unwrap();
let last_offset = list_array.offsets()[list_array.len()].to_usize().unwrap();
let visible_values = list_array
.values()
.slice(first_offset, last_offset - first_offset);
let to_array = scalar_to.to_array_of_size(1)?;
let original_data = visible_values.to_data();
let to_data = to_array.to_data();
let capacity = Capacities::Array(original_data.len());
let mut mutable = MutableArrayData::with_capacities(
vec![&original_data, &to_data],
false,
capacity,
);
let mut offsets = Vec::<O>::with_capacity(list_array.len() + 1);
offsets.push(O::zero());
let match_bitmap = arrow_ord::cmp::not_distinct(&visible_values, needle)?;
let match_bits = match_bitmap.values();
for (row_index, offset_window) in list_array.offsets().windows(2).enumerate() {
let start = offset_window[0].to_usize().unwrap() - first_offset;
let end = offset_window[1].to_usize().unwrap() - first_offset;
let row_len = end - start;
if list_array.is_null(row_index) {
offsets.push(offsets[row_index]);
continue;
}
let row_bits = match_bits.slice(start, row_len);
let mut match_positions = row_bits
.set_indices()
.take(max_replacements as usize)
.peekable();
if match_positions.peek().is_none() {
mutable.try_extend(0, start, end)?;
offsets.push(offsets[row_index] + O::usize_as(row_len));
continue;
}
let mut prev_end = 0usize;
for match_pos in match_positions {
if match_pos > prev_end {
mutable.try_extend(0, start + prev_end, start + match_pos)?;
}
mutable.try_extend(1, 0, 1)?;
prev_end = match_pos + 1;
}
if prev_end < row_len {
mutable.try_extend(0, start + prev_end, end)?;
}
offsets.push(offsets[row_index] + O::usize_as(row_len));
}
let data = mutable.freeze();
Ok(Arc::new(GenericListArray::<O>::try_new(
field,
OffsetBuffer::new(offsets.into()),
arrow::array::make_array(data),
list_array.nulls().cloned(),
)?))
}
fn array_replace_with_scalar_args(
name: &str,
list_array: &ArrayRef,
scalar_from: &ScalarValue,
scalar_to: &ScalarValue,
max_replacements: i64,
return_type: &DataType,
) -> Result<ArrayRef> {
if scalar_from.data_type().is_nested() {
let num_rows = list_array.len();
let from_array = scalar_from.to_array_of_size(num_rows)?;
let to_array = scalar_to.to_array_of_size(num_rows)?;
return array_replace_internal(
name,
list_array,
&from_array,
&to_array,
&vec![Some(max_replacements); num_rows],
return_type,
);
}
let needle = Scalar::new(scalar_from.to_array_of_size(1)?);
match list_array.data_type() {
DataType::List(_) => general_replace_with_scalar::<i32>(
list_array.as_list::<i32>(),
&needle,
scalar_to,
max_replacements,
list_inner_field(name, return_type)?,
),
DataType::LargeList(_) => general_replace_with_scalar::<i64>(
list_array.as_list::<i64>(),
&needle,
scalar_to,
max_replacements,
list_inner_field(name, return_type)?,
),
DataType::Null => Ok(new_null_array(return_type, list_array.len())),
array_type => exec_err!("{name} does not support type '{array_type}'."),
}
}
fn array_replace_internal(
name: &str,
array: &ArrayRef,
from: &ArrayRef,
to: &ArrayRef,
arr_n: &[Option<i64>],
return_type: &DataType,
) -> Result<ArrayRef> {
match array.data_type() {
DataType::List(_) => general_replace::<i32>(
array.as_list::<i32>(),
from,
to,
arr_n,
list_inner_field(name, return_type)?,
),
DataType::LargeList(_) => general_replace::<i64>(
array.as_list::<i64>(),
from,
to,
arr_n,
list_inner_field(name, return_type)?,
),
DataType::Null => Ok(new_null_array(return_type, array.len())),
array_type => exec_err!("{name} does not support type '{array_type}'."),
}
}
fn array_replace_n_inner(
name: &str,
array: &ArrayRef,
from: &ArrayRef,
to: &ArrayRef,
max: &ArrayRef,
return_type: &DataType,
) -> Result<ArrayRef> {
let arr_n = as_int64_array(max)?.iter().collect::<Vec<_>>();
array_replace_internal(name, array, from, to, &arr_n, return_type)
}
#[cfg(test)]
mod tests {
use super::{ArrayReplaceN, array_replace_n_inner};
use arrow::array::{ArrayRef, AsArray, Int32Array, Int64Array, ListArray};
use arrow::buffer::{NullBuffer, ScalarBuffer};
use arrow::datatypes::{DataType, Field, Int32Type};
use datafusion_common::{Result, ScalarValue, config::ConfigOptions};
use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl};
use std::sync::Arc;
#[test]
fn test_array_replace_n_null_max_returns_null() -> Result<()> {
let array: ArrayRef =
Arc::new(ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(1), Some(2), Some(3)]),
Some(vec![Some(4), Some(2)]),
]));
let from: ArrayRef = Arc::new(Int32Array::from(vec![2, 2]));
let to: ArrayRef = Arc::new(Int32Array::from(vec![9, 9]));
let max: ArrayRef = Arc::new(Int64Array::new(
ScalarBuffer::from(vec![1, 1]),
Some(NullBuffer::from(vec![true, false])),
));
let result = array_replace_n_inner(
"array_replace_n",
&array,
&from,
&to,
&max,
array.data_type(),
)?;
let expected = ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(1), Some(9), Some(3)]),
None,
]);
assert_eq!(result.as_list::<i32>(), &expected);
Ok(())
}
#[test]
fn test_array_replace_n_scalar_null_max_returns_null() -> Result<()> {
let array: ArrayRef =
Arc::new(ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
Some(vec![Some(1), Some(2), Some(3)]),
Some(vec![Some(4), Some(2)]),
]));
let array_field = Arc::new(Field::new("array", array.data_type().clone(), true));
let result = ArrayReplaceN::new().invoke_with_args(ScalarFunctionArgs {
args: vec![
ColumnarValue::Array(Arc::clone(&array)),
ColumnarValue::Scalar(ScalarValue::Int32(Some(2))),
ColumnarValue::Scalar(ScalarValue::Int32(Some(9))),
ColumnarValue::Scalar(ScalarValue::Int64(None)),
],
arg_fields: vec![
Arc::clone(&array_field),
Arc::new(Field::new("from", DataType::Int32, false)),
Arc::new(Field::new("to", DataType::Int32, false)),
Arc::new(Field::new("max", DataType::Int64, true)),
],
number_rows: array.len(),
return_field: Arc::clone(&array_field),
config_options: Arc::new(ConfigOptions::default()),
})?;
let result = result.into_array(array.len())?;
let expected = ListArray::from_iter_primitive::<Int32Type, _, _>(vec![
Option::<Vec<Option<i32>>>::None,
Option::<Vec<Option<i32>>>::None,
]);
assert_eq!(result.as_list::<i32>(), &expected);
Ok(())
}
}