use std::sync::Arc;
use arrow::array::{Array, ArrayRef, Float64Array};
use arrow::datatypes::DataType;
use datafusion_common::utils::take_function_args;
use datafusion_common::{Result, ScalarValue};
use datafusion_expr::{
ColumnarValue, Documentation, ScalarFunctionArgs, ScalarUDFImpl, Signature,
};
use datafusion_functions::math::power::PowerFunc;
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct SparkPow {
inner: PowerFunc,
aliases: Vec<String>,
}
impl Default for SparkPow {
fn default() -> Self {
Self::new()
}
}
impl SparkPow {
pub fn new() -> Self {
Self {
inner: PowerFunc::new(),
aliases: vec!["power".to_string()],
}
}
}
impl ScalarUDFImpl for SparkPow {
fn name(&self) -> &str {
"pow"
}
fn aliases(&self) -> &[String] {
&self.aliases
}
fn signature(&self) -> &Signature {
self.inner.signature()
}
fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
self.inner.return_type(arg_types)
}
fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
match args.args.as_slice() {
[base, exponent]
if matches!(base.data_type(), DataType::Float64)
&& matches!(exponent.data_type(), DataType::Float64) => {}
_ => return self.inner.invoke_with_args(args),
}
let num_rows = args.number_rows;
if let [
ColumnarValue::Scalar(ScalarValue::Float64(base)),
ColumnarValue::Scalar(ScalarValue::Float64(exp)),
] = args.args.as_slice()
{
let result = (*base).zip(*exp).map(|(base, exp)| {
if base == 0.0 && exp < 0.0 {
f64::INFINITY
} else {
base.powf(exp)
}
});
return Ok(ColumnarValue::Scalar(ScalarValue::Float64(result)));
}
let [base, exponent] = take_function_args(self.name(), &args.args)?;
let base_arr: ArrayRef = base.to_array(num_rows)?;
let exp_arr: ArrayRef = exponent.to_array(num_rows)?;
let base_f64 = base_arr
.as_any()
.downcast_ref::<Float64Array>()
.expect("base must be Float64Array");
let exp_f64 = exp_arr
.as_any()
.downcast_ref::<Float64Array>()
.expect("exponent must be Float64Array");
let result: Float64Array = base_f64
.iter()
.zip(exp_f64.iter())
.map(|(base, exp)| match (base, exp) {
(Some(base), Some(exp)) => {
if base == 0.0 && exp < 0.0 {
Some(f64::INFINITY)
} else {
Some(base.powf(exp))
}
}
_ => None,
})
.collect();
Ok(ColumnarValue::Array(Arc::new(result)))
}
fn documentation(&self) -> Option<&Documentation> {
self.inner.documentation()
}
}