use std::sync::Arc;
use arrow::array::{ArrayRef, AsArray, Float64Array};
use arrow::compute::kernels::arity::binary;
use arrow::datatypes::{DataType, Float64Type};
use datafusion_common::Result;
use datafusion_common::utils::take_function_args;
use datafusion_expr::{
ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility,
};
use datafusion_functions::utils::make_scalar_function;
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct SparkAtan2 {
signature: Signature,
}
impl Default for SparkAtan2 {
fn default() -> Self {
Self::new()
}
}
impl SparkAtan2 {
pub fn new() -> Self {
Self {
signature: Signature::exact(
vec![DataType::Float64, DataType::Float64],
Volatility::Immutable,
),
}
}
}
impl ScalarUDFImpl for SparkAtan2 {
fn name(&self) -> &str {
"atan2"
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
Ok(DataType::Float64)
}
fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
make_scalar_function(spark_atan2, vec![])(&args.args)
}
}
fn spark_atan2(args: &[ArrayRef]) -> Result<ArrayRef> {
let [y, x] = take_function_args("atan2", args)?;
let y = y.as_primitive::<Float64Type>();
let x = x.as_primitive::<Float64Type>();
let result: Float64Array = binary(y, x, |y, x| y.atan2(x))?;
Ok(Arc::new(result))
}