函数

添加用户自定义函数:标量/窗口/聚合/表函数

师成师成· 更新于 2026-09-28· 阅读 74 分钟· 0 次阅读

登录后可跨设备保存划线和私人笔记登录

添加用户自定义函数:标量/窗口/聚合/表函数

用户自定义函数(User Defined Functions,UDF)是在 DataFusion 执行上下文中可以使用的函数。

本页介绍如何向 DataFusion 添加 UDF。具体而言,涵盖如何添加标量、窗口和聚合 UDF。

UDF 类型描述示例
标量(Scalar)接收一行数据并返回单个值的函数。simple_udf.rs / advanced_udf.rs
窗口(Window)接收一行数据并返回单个值,同时还能访问其周围各行的函数。simple_udwf.rs / advanced_udwf.rs
聚合(Aggregate)接收一组行并返回单个值的函数。simple_udaf.rs / advanced_udaf.rs
表(Table)接收参数并返回一个 TableProvider,以便在查询计划中使用的函数。simple_udtf.rs
标量(异步)一种标量函数,用于在 UDF 中执行 async 操作(例如网络或 I/O 调用)。async_udf.rs

我们将先完整地讲解如何添加一个标量 UDF,然后再讨论不同类型 UDF 之间的差异。

添加标量 UDF

标量 UDF 是一种接收一行数据并返回单个值的函数。为了获得良好性能,这类函数在 DataFusion 中是「向量化」的:它们以一个或多个 Arrow Array 作为输入,并输出一个行数相同的 Arrow Array。

要创建标量 UDF,你需要:

  1. 实现 ScalarUDFImpl trait,向 DataFusion 描述你的函数,例如它接受哪些类型的参数以及如何计算结果。
  2. 创建一个 ScalarUDF,并通过 SessionContext::register_udf 注册它,以便按名称调用。

在下面的示例中,我们将添加一个函数,它接收一个 i64 并返回一个 i64,其值加 1:

为简洁起见,我们省略了一些错误处理。在生产代码中,你可能需要检查,例如 args.len() 是否与期望的参数数量一致。

通过 impl ScalarUDFImpl 添加

这是一个更底层的 API,功能更强大但也更复杂,相关文档见 advanced_udf.rs。

use std::sync::Arc;
use std::any::Any;
use std::sync::LazyLock;
use arrow::datatypes::DataType;
use datafusion_common::cast::as_int64_array;
use datafusion_common::{DataFusionError, plan_err, Result};
use datafusion_expr::{col, ColumnarValue, ScalarFunctionArgs, Signature, Volatility};
use datafusion::arrow::array::{ArrayRef, Int64Array};
use datafusion_expr::{ScalarUDFImpl, ScalarUDF};
use datafusion_macros::user_doc;
use datafusion_doc::Documentation;

/// This struct for a simple UDF that adds one to an int32
#[user_doc(
    doc_section(label = "Math Functions"),
    description = "Add one udf",
    syntax_example = "add_one(1)"
)]
#[derive(Debug, PartialEq, Eq, Hash)]
struct AddOne {
  signature: Signature,
}

impl AddOne {
  fn new() -> Self {
    Self {
      signature: Signature::uniform(1, vec![DataType::Int32], Volatility::Immutable),
     }
  }
}

/// Implement the ScalarUDFImpl trait for AddOne
impl ScalarUDFImpl for AddOne {
   fn name(&self) -> &str { "add_one" }
   fn signature(&self) -> &Signature { &self.signature }
   fn return_type(&self, args: &[DataType]) -> Result<DataType> {
     if !matches!(args.get(0), Some(&DataType::Int32)) {
       return plan_err!("add_one only accepts Int32 arguments");
     }
     Ok(DataType::Int32)
   }
   // The actual implementation would add one to the argument
   fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
        let args = ColumnarValue::values_to_arrays(&args.args)?;
        let i64s = as_int64_array(&args[0])?;

        let new_array = i64s
            .iter()
            .map(|array_elem| array_elem.map(|value| value + 1))
            .collect::<Int64Array>();

        Ok(ColumnarValue::from(Arc::new(new_array) as ArrayRef))
   }
   fn documentation(&self) -> Option<&Documentation> {
       self.doc()
    }
}

现在我们需要将该函数注册到 DataFusion 中,以便它能在查询上下文中使用。

use datafusion::execution::context::SessionContext;

// Create a new ScalarUDF from the implementation
let add_one = ScalarUDF::from(AddOne::new());

// Call the function `add_one(col)`
let expr = add_one.call(vec![col("a")]);

// register the UDF with the context so it can be invoked by name and from SQL
let mut ctx = SessionContext::new();
ctx.register_udf(add_one.clone());

通过 create_udf 添加标量 UDF

此外,还有一个更早、更简洁但功能也更有限的 API create_udf 可供使用

添加标量 UDF

use std::sync::Arc;
use datafusion::arrow::array::{ArrayRef, Int64Array};
use datafusion::common::cast::as_int64_array;
use datafusion::common::Result;
use datafusion::logical_expr::ColumnarValue;

pub fn add_one(args: &[ColumnarValue]) -> Result<ColumnarValue> {
    // Error handling omitted for brevity
    let args = ColumnarValue::values_to_arrays(args)?;
    let i64s = as_int64_array(&args[0])?;

    let new_array = i64s
        .iter()
        .map(|array_elem| array_elem.map(|value| value + 1))
        .collect::<Int64Array>();

    Ok(ColumnarValue::from(Arc::new(new_array) as ArrayRef))
}

这在孤立情况下是“可行”的,也就是说,如果你有一个 ArrayRef 切片,你可以调用 add_one,它会返回一个新的 ArrayRef,其中每个值都加上了 1。

let input = vec![Some(1), None, Some(3)];
let input = ColumnarValue::from(Arc::new(Int64Array::from(input)) as ArrayRef);

let result = add_one(&[input]).unwrap();
let binding = result.into_array(1).unwrap();
let result = binding.as_any().downcast_ref::<Int64Array>().unwrap();

assert_eq!(result, &Int64Array::from(vec![Some(2), None, Some(4)]));

然而问题在于 DataFusion 并不知道这个函数。我们需要将它注册到 DataFusion 中,以便能在查询上下文中使用。

注册标量 UDF

要注册标量 UDF,你需要将函数实现封装到 ScalarUDF 结构体中,然后将其注册到 SessionContext。DataFusion 提供了 create_udf 及相关辅助函数,使这一过程更加简便。

use datafusion::logical_expr::{Volatility, create_udf};
use datafusion::arrow::datatypes::DataType;

let udf = create_udf(
    "add_one",
    vec![DataType::Int64],
    DataType::Int64,
    Volatility::Immutable,
    Arc::new(add_one),
);

关于 create_udf,有几点需要注意:

  • 第一个参数是函数的名称,即在 SQL 查询中使用的名称。
  • 第二个参数是一个 DataType 向量,即该函数接受的参数类型列表。也就是说,在本例中,函数接受单个 Int64 参数。
  • 第三个参数是函数的返回类型。也就是说,在本例中,函数返回一个 Int64。
  • 第四个参数是函数的 volatility(可变性)。简而言之,它用于判断在某些情况下函数的性能是否可以被优化。在本例中,该函数是 Immutable(不可变的),因为它对相同的输入总是返回相同的值。而随机数生成器则是 Volatile(可变的),因为它对相同的输入会返回不同的值。
  • 第五个参数是函数的实现,也就是我们上面定义的那个函数。

这样我们就得到了一个 ScalarUDF,可以将它注册到 SessionContext 中:

use datafusion::logical_expr::{Volatility, create_udf};
use datafusion::arrow::datatypes::DataType;
use datafusion::execution::context::SessionContext;

#[tokio::main]
async fn main() {
    let udf = create_udf(
        "add_one",
        vec![DataType::Int64],
        DataType::Int64,
        Volatility::Immutable,
        Arc::new(add_one),
    );

    let mut ctx = SessionContext::new();
    ctx.register_udf(udf);

    // At this point, you can use the `add_one` function in your query:
    let query = "SELECT add_one(1)";
    let df = ctx.sql(&query).await.unwrap();
}

添加异步标量 UDF

异步标量 UDF 允许你实现支持异步执行的用户自定义函数,例如在 UDF 内部执行网络或 I/O 操作。

要添加一个标量异步 UDF,你需要:

  1. 实现 AsyncScalarUDFImpl trait,以定义你的异步函数逻辑、签名和类型。
  2. 使用 AsyncScalarUDF::new 包装你的实现,并将其注册到 SessionContext 中。

通过 impl AsyncScalarUDFImpl 添加

#[derive(Debug, PartialEq, Eq, Hash)]
pub struct AsyncUpper {
    signature: Signature,
}

impl Default for AsyncUpper {
    fn default() -> Self {
        Self::new()
    }
}

impl AsyncUpper {
    pub fn new() -> Self {
        Self {
            signature: Signature::new(
                TypeSignature::Coercible(vec![Coercion::new_exact(
                    TypeSignatureClass::Native(logical_string()),
                )]),
                Volatility::Volatile,
            ),
        }
    }
}

/// Implement the normal ScalarUDFImpl trait for AsyncUpper
#[async_trait]
impl ScalarUDFImpl for AsyncUpper {
    fn name(&self) -> &str {
        "async_upper"
    }

    fn signature(&self) -> &Signature {
        &self.signature
    }

    fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
        Ok(DataType::Utf8)
    }

    // Note the normal invoke_with_args method is not called for Async UDFs
    fn invoke_with_args(
        &self,
        _args: ScalarFunctionArgs,
    ) -> Result<ColumnarValue> {
        not_impl_err!("AsyncUpper can only be called from async contexts")
    }
}

/// The actual implementation of the async UDF
#[async_trait]
impl AsyncScalarUDFImpl for AsyncUpper {
    fn ideal_batch_size(&self) -> Option<usize> {
        Some(10)
    }

    /// This method is called to execute the async UDF and is similar
    /// to the normal `invoke_with_args` except it is `async`.
    async fn invoke_async_with_args(
        &self,
        args: ScalarFunctionArgs,
    ) -> Result<ColumnarValue> {
        let value = &args.args[0];
        // This function simply implements a simple string to uppercase conversion
        // but can be used for any async operation such as network calls.
        let result = match value {
            ColumnarValue::Array(array) => {
                let string_array = array.as_string::<i32>();
                let iter = ArrayIter::new(string_array);
                let result = iter
                    .map(|string| string.map(|s| s.to_uppercase()))
                    .collect::<StringArray>();
                Arc::new(result) as ArrayRef
            }
            _ => return internal_err!("Expected a string argument, got {:?}", value),
        };
        Ok(ColumnarValue::from(result))
    }
}

现在我们可以使用 into_scalar_udf 将异步 UDF 转换为普通标量函数,并向 DataFusion 注册该函数,以便在查询上下文中使用它。

use datafusion::execution::context::SessionContext;
use datafusion::logical_expr::async_udf::AsyncScalarUDF;

let async_upper = AsyncUpper::new();
let udf = AsyncScalarUDF::new(Arc::new(async_upper));
let mut ctx = SessionContext::new();
ctx.register_udf(udf.into_scalar_udf());

注册完成后,你可以在 SQL 查询中直接使用这些异步 UDF,例如:

SELECT async_upper('datafusion');

有关异步 UDF 的实现细节,请参阅 async_udf.rs。

命名参数

DataFusion 支持标量、窗口和聚合 UDF 的命名参数,允许你按参数名称传递参数:

-- Scalar function
SELECT substr(str => 'hello', start_pos => 2, length => 3);

-- Window function
SELECT lead(expr => value, offset => 1) OVER (ORDER BY id) FROM table;

-- Aggregate function
SELECT corr(y => col1, x => col2) FROM table;

具名参数可以与位置参数混合使用,但位置参数必须放在最前面:

SELECT substr('hello', start_pos => 2, length => 3);  -- Valid

实现带命名参数的函数

要在 UDF 中支持命名参数,请使用 .with_parameter_names() 为函数签名添加参数名。这一方式对标量(Scalar)、窗口(Window)和聚合(Aggregate)UDF 均适用:

#[derive(Debug, PartialEq, Eq, Hash)]
struct PowerFunction {
    signature: Signature,
}

impl PowerFunction {
    fn new() -> Self {
        Self {
            signature: Signature::uniform(
                2,
                vec![DataType::Float64],
                Volatility::Immutable
            )
            .with_parameter_names(vec![
                "base".to_string(),
                "exponent".to_string()
            ])
            .expect("valid parameter names"),
        }
    }
}

impl ScalarUDFImpl for PowerFunction {
    fn name(&self) -> &str { "power" }
    fn signature(&self) -> &Signature { &self.signature }

    fn return_type(&self, _args: &[DataType]) -> Result<DataType> {
        Ok(DataType::Float64)
    }

    fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result<ColumnarValue> {
        // Your implementation - arguments are in correct positional order
        unimplemented!()
    }
}

参数名称应与函数签名中的参数顺序保持一致。DataFusion 会在调用你的函数之前,自动将命名参数解析为正确的参数位置。

注册之后,用户可以按任意顺序使用命名参数调用你的函数:

-- All equivalent
SELECT power(base => 2.0, exponent => 3.0);
SELECT power(exponent => 3.0, base => 2.0);
SELECT power(2.0, exponent => 3.0);

错误消息

当函数调用因参数不正确而失败时,DataFusion 会在错误消息中显示参数名称,以帮助用户:

No function matches the given name and argument types substr(Utf8).
    Candidate functions:
    substr(str: Any, start_pos: Any)
    substr(str: Any, start_pos: Any, length: Any)

添加窗口 UDF

标量 UDF 是接收一行数据并返回单个值的函数。窗口 UDF 与之类似,但它们还能访问其周围的行。能够访问相邻行非常有用,不过也会为实现带来一些复杂性。

有关背景知识及其他注意事项,请参阅 DataFusion 中的用户自定义窗口函数 博客文章。

例如,我们将声明一个用于计算移动平均值的用户自定义窗口函数。

use datafusion::arrow::{array::{ArrayRef, Float64Array, AsArray}, datatypes::Float64Type};
use datafusion::logical_expr::{PartitionEvaluator};
use datafusion::common::ScalarValue;
use datafusion::error::Result;
/// This implements the lowest level evaluation for a window function
///
/// It handles calculating the value of the window function for each
/// distinct values of `PARTITION BY`
#[derive(Clone, Debug)]
struct MyPartitionEvaluator {}

impl MyPartitionEvaluator {
    fn new() -> Self {
        Self {}
    }
}

/// Different evaluation methods are called depending on the various
/// settings of WindowUDF. This example uses the simplest and most
/// general, `evaluate`. See `PartitionEvaluator` for the other more
/// advanced uses.
impl PartitionEvaluator for MyPartitionEvaluator {
    /// Tell DataFusion the window function varies based on the value
    /// of the window frame.
    fn uses_window_frame(&self) -> bool {
        true
    }

    /// This function is called once per input row.
    ///
    /// `range`specifies which indexes of `values` should be
    /// considered for the calculation.
    ///
    /// Note this is the SLOWEST, but simplest, way to evaluate a
    /// window function. It is much faster to implement
    /// evaluate_all or evaluate_all_with_rank, if possible
    fn evaluate(
        &mut self,
        values: &[ArrayRef],
        range: &std::ops::Range<usize>,
    ) -> Result<ScalarValue> {
        // Again, the input argument is an array of floating
        // point numbers to calculate a moving average
        let arr: &Float64Array = values[0].as_ref().as_primitive::<Float64Type>();

        let range_len = range.end - range.start;

        // our smoothing function will average all the values in the
        let output = if range_len > 0 {
            let sum: f64 = arr.values().iter().skip(range.start).take(range_len).sum();
            Some(sum / range_len as f64)
        } else {
            None
        };

        Ok(ScalarValue::Float64(output))
    }
}

/// Create a `PartitionEvaluator` to evaluate this function on a new
/// partition.
fn make_partition_evaluator() -> Result<Box<dyn PartitionEvaluator>> {
    Ok(Box::new(MyPartitionEvaluator::new()))
}

注册窗口 UDF

要注册窗口 UDF,你需要将函数实现包装进一个 WindowUDF 结构体,然后将其注册到 SessionContext 中。DataFusion 提供了 create_udwf 辅助函数来简化这一过程。另有一个功能更丰富但使用起来更复杂的底层 API,其文档位于 advanced_udwf.rs。

use datafusion::logical_expr::{Volatility, create_udwf};
use datafusion::arrow::datatypes::DataType;
use std::sync::Arc;

// here is where we define the UDWF. We also declare its signature:
let smooth_it = create_udwf(
    "smooth_it",
    DataType::Float64,
    Arc::new(DataType::Float64),
    Volatility::Immutable,
    Arc::new(make_partition_evaluator),
);

create_udwf 有五个参数需要检查:

  • 第一个参数是函数的名称。这是将在 SQL 查询中使用的名称。
  • 第二个参数是输入数组的 DataType(注意:这不是一个数组列表)。也就是说,在本例中,该函数接受 Float64 作为参数。
  • 第三个参数是函数的返回类型。也就是说,在本例中,该函数返回一个 Float64。
  • 第四个参数是函数的 volatility(易变性)。简而言之,这用于确定函数的执行结果在某些情况下是否可以被优化。在本例中,该函数是 Immutable(不可变的),因为对于相同的输入它总是返回相同的值。而随机数生成器则是 Volatile(可变的),因为对于相同的输入它会返回不同的值。
  • 第五个参数是函数的实现。也就是我们在上面定义的那个函数。

这样我们就得到了一个 WindowUDF,可以将其注册到 SessionContext 中:

use datafusion::execution::context::SessionContext;

let ctx = SessionContext::new();

ctx.register_udwf(smooth_it);

至此,你便可以在查询中使用 smooth_it 函数了:

例如,假设我们有一个 cars.csv,其内容如下:

car,speed,time
red,20.0,1996-04-12T12:05:03.000000000
red,20.3,1996-04-12T12:05:04.000000000
green,10.0,1996-04-12T12:05:03.000000000
green,10.3,1996-04-12T12:05:04.000000000
...

接着,我们可以像下面这样查询:

use datafusion::datasource::file_format::options::CsvReadOptions;

#[tokio::main]
async fn main() -> Result<()> {

    let ctx = SessionContext::new();

    let smooth_it = create_udwf(
        "smooth_it",
        DataType::Float64,
        Arc::new(DataType::Float64),
        Volatility::Immutable,
        Arc::new(make_partition_evaluator),
    );
    ctx.register_udwf(smooth_it);

    // register csv table first
    let csv_path = "../../datafusion/core/tests/data/cars.csv".to_string();
    ctx.register_csv("cars", &csv_path, CsvReadOptions::default().has_header(true)).await?;

    // do query with smooth_it
    let df = ctx
        .sql(r#"
            SELECT
                car,
                speed,
                smooth_it(speed) OVER (PARTITION BY car ORDER BY time) as smooth_speed,
                time
            FROM cars
            ORDER BY car
        "#)
        .await?;

    // print the results
    df.show().await?;
    Ok(())
}

输出结果如下:

+-------+-------+--------------------+---------------------+
| car   | speed | smooth_speed       | time                |
+-------+-------+--------------------+---------------------+
| green | 10.0  | 10.0               | 1996-04-12T12:05:03 |
| green | 10.3  | 10.15              | 1996-04-12T12:05:04 |
| green | 10.4  | 10.233333333333334 | 1996-04-12T12:05:05 |
| green | 10.5  | 10.3               | 1996-04-12T12:05:06 |
| green | 11.0  | 10.440000000000001 | 1996-04-12T12:05:07 |
| green | 12.0  | 10.700000000000001 | 1996-04-12T12:05:08 |
| green | 14.0  | 11.171428571428573 | 1996-04-12T12:05:09 |
| green | 15.0  | 11.65              | 1996-04-12T12:05:10 |
| green | 15.1  | 12.033333333333333 | 1996-04-12T12:05:11 |
| green | 15.2  | 12.35              | 1996-04-12T12:05:12 |
| green | 8.0   | 11.954545454545455 | 1996-04-12T12:05:13 |
| green | 2.0   | 11.125             | 1996-04-12T12:05:14 |
| red   | 20.0  | 20.0               | 1996-04-12T12:05:03 |
| red   | 20.3  | 20.15              | 1996-04-12T12:05:04 |
...

添加聚合 UDF

聚合 UDF 是接收一组行并返回单个值的函数,类似于 SQL 中的 SUM 或 COUNT 函数。

例如,我们将声明一个单输入类型、单返回类型的 UDAF,用于计算几何平均值。

use datafusion::arrow::array::ArrayRef;
use datafusion::scalar::ScalarValue;
use datafusion::{error::Result, physical_plan::Accumulator};

/// A UDAF has state across multiple rows, and thus we require a `struct` with that state.
#[derive(Debug)]
struct GeometricMean {
    n: u32,
    prod: f64,
}

impl GeometricMean {
    // how the struct is initialized
    pub fn new() -> Self {
        GeometricMean { n: 0, prod: 1.0 }
    }
}

// UDAFs are built using the trait `Accumulator`, that offers DataFusion the necessary functions
// to use them.
impl Accumulator for GeometricMean {
    // This function serializes our state to `ScalarValue`, which DataFusion uses
    // to pass this state between execution stages.
    // Note that this can be arbitrary data.
    fn state(&mut self) -> Result<Vec<ScalarValue>> {
        Ok(vec![
            ScalarValue::from(self.prod),
            ScalarValue::from(self.n),
        ])
    }

    // DataFusion expects this function to return the final value of this aggregator.
    // in this case, this is the formula of the geometric mean
    fn evaluate(&mut self) -> Result<ScalarValue> {
        let value = self.prod.powf(1.0 / self.n as f64);
        Ok(ScalarValue::from(value))
    }

    // DataFusion calls this function to update the accumulator's state for a batch
    // of inputs rows. In this case the product is updated with values from the first column
    // and the count is updated based on the row count
    fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
        if values.is_empty() {
            return Ok(());
        }
        let arr = &values[0];
        (0..arr.len()).try_for_each(|index| {
            let v = ScalarValue::try_from_array(arr, index)?;

            if let ScalarValue::Float64(Some(value)) = v {
                self.prod *= value;
                self.n += 1;
            } else {
                unreachable!("")
            }
            Ok(())
        })
    }

    // Optimization hint: this trait also supports `update_batch` and `merge_batch`,
    // that can be used to perform these operations on arrays instead of single values.
    fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
        if states.is_empty() {
            return Ok(());
        }
        let arr = &states[0];
        (0..arr.len()).try_for_each(|index| {
            let v = states
                .iter()
                .map(|array| ScalarValue::try_from_array(array, index))
                .collect::<Result<Vec<_>>>()?;
            if let (ScalarValue::Float64(Some(prod)), ScalarValue::UInt32(Some(n))) = (&v[0], &v[1])
            {
                self.prod *= prod;
                self.n += n;
            } else {
                unreachable!("")
            }
            Ok(())
        })
    }

    fn size(&self) -> usize {
        std::mem::size_of_val(self)
    }
}

声明聚合 UDF 如何处理 DISTINCT

默认情况下,DataFusion 认为聚合函数对 DISTINCT 修饰符敏感,即累加器需要读取 AccumulatorArgs::is_distinct 并对输入进行去重。如果你的函数并非如此,请重写 AggregateUDFImpl::distinct_handling:

  • 当重复项不会改变结果时,返回 DistinctHandling::Insensitive,也就是说,当合并一个累加器已经见过的值是无操作的。min、max、bool_and 和 bit_or 都属于这一类。优化器随后会将 f(DISTINCT x) 规划为 f(x),从而跳过每组的哈希集合,也跳过 SingleDistinctToGroupBy 本来会引入的额外分组阶段。
  • 当累加器未实现 DISTINCT 时,返回 DistinctHandling::Unsupported:它不读取 is_distinct,或者遇到 DISTINCT 会报错。此时规划器必须先对输入去重,或者拒绝该查询。目前这只是一个声明;在规划阶段拒绝这类查询是后续的改进项。
  • 如果累加器会读取 AccumulatorArgs::is_distinct 并自行对输入去重,则保留默认的 DistinctHandling::Sensitive。

声明错误会改变查询结果,因此只有在你的合并操作确实满足幂等时,才应声明为 Insensitive。

注册聚合 UDF

要注册聚合 UDF,需要先将函数实现包装进 AggregateUDF 结构体,然后将其注册到 SessionContext 中。DataFusion 提供了 create_udaf 辅助函数来简化这一过程。另有一套功能更丰富但也更复杂的底层 API,相关文档见 advanced_udaf.rs。

use datafusion::logical_expr::{Volatility, create_udaf};
use datafusion::arrow::datatypes::DataType;
use std::sync::Arc;

// here is where we define the UDAF. We also declare its signature:
let geometric_mean = create_udaf(
    // the name; used to represent it in plan descriptions and in the registry, to use in SQL.
    "geo_mean",
    // the input type; DataFusion guarantees that the first entry of `values` in `update` has this type.
    vec![DataType::Float64],
    // the return type; DataFusion expects this to match the type returned by `evaluate`.
    Arc::new(DataType::Float64),
    Volatility::Immutable,
    // This is the accumulator factory; DataFusion uses it to create new accumulators.
    Arc::new( | _ | Ok(Box::new(GeometricMean::new()))),
    // This is the description of the state. `state()` must match the types here.
    Arc::new(vec![DataType::Float64, DataType::UInt32]),
);

create_udaf 有六个参数需要检查:

  • 第一个参数是函数的名称,也就是 SQL 查询中使用的名称。
  • 第二个参数是一个 DataType 向量,即该函数接受的参数类型列表。例如在本例中,该函数接受单个 Float64 参数。
  • 第三个参数是函数的返回类型。例如在本例中,该函数返回 Int64。
  • 第四个参数是函数的易变性(volatility)。简而言之,它用于判断在某些情况下函数的性能是否可以被优化。本例中该函数是 Immutable,因为它对相同的输入总是返回相同的值;而随机数生成器则是 Volatile,因为相同的输入会返回不同的值。
  • 第五个参数是函数实现,即我们上面定义的那个函数。
  • 第六个参数是状态的描述,该状态会在各个执行阶段之间传递。

从聚合 UDF 返回多个值

当一个聚合结果需要携带多个值时,聚合 UDF 可以返回 DataType::Struct。这对于需要同时返回窗口起始时间、窗口结束时间以及聚合值等元数据的时间窗口扩展非常有用。

请将相关的输入列传递给聚合,使累加器在多阶段聚合计划中拥有足够的信息来正常更新和合并状态。例如,可以使用内置的 date_bin 函数将行分组到时间桶中,同时由返回结构体的聚合计算出每个桶的值并携带相应的元数据:

SELECT
  augmented_avg(time, value)['window_start'] AS window_start,
  augmented_avg(time, value)['window_end'] AS window_end,
  augmented_avg(time, value)['window_duration'] AS window_duration,
  augmented_avg(time, value)['avg_value'] AS avg_value
FROM t
GROUP BY date_bin(INTERVAL '30 seconds', time)
ORDER BY window_start;

在这种模式中,date_bin(...) 将各行分配到时间桶,而 augmented_avg(time, value) 是一个普通的聚合 UDF,其累加器保存可合并的状态,例如 window_start、window_end、sum 和 count。该聚合的 evaluate 方法返回一个 ScalarValue::Struct,调用方可以从该结构体中投影出单个字段。

use datafusion::logical_expr::{Volatility, create_udaf};
use datafusion::arrow::datatypes::DataType;
use std::sync::Arc;
use datafusion::execution::context::SessionContext;
use datafusion::datasource::file_format::options::CsvReadOptions;

#[tokio::main]
async fn main() -> Result<()> {
    let geometric_mean = create_udaf(
        "geo_mean",
        vec![DataType::Float64],
        Arc::new(DataType::Float64),
        Volatility::Immutable,
        Arc::new( | _ | Ok(Box::new(GeometricMean::new()))),
        Arc::new(vec![DataType::Float64, DataType::UInt32]),
    );

    // That gives us a `AggregateUDF` that we can register with the `SessionContext`:
    use datafusion::execution::context::SessionContext;

    let ctx = SessionContext::new();
    ctx.register_udaf(geometric_mean);

    // register csv table first
    let csv_path = "../../datafusion/core/tests/data/cars.csv".to_string();
    ctx.register_csv("cars", &csv_path, CsvReadOptions::default().has_header(true)).await?;

    // Then, we can query like below:
    let df = ctx.sql("SELECT geo_mean(speed) FROM cars").await?;
    Ok(())
}

添加表 UDF

用户自定义表函数(UDTF)是一种接受参数并返回 TableProvider 的函数。

由于我们返回的是 TableProvider,在本示例中将使用 MemTable 数据源来表示一张表。这是一个简单的结构体,它在内存中持有一组 RecordBatch,并将其作为表来处理。在你的场景中,这会被替换为你自己的、实现了 TableProvider 的结构体。

虽然这只是一个用于说明问题的简单示例,但 UDTF 有许多潜在的应用场景,尤其适用于从外部数据源读取数据以及进行交互式分析。请参阅这个可运行的示例,它从 CSV 文件中读取数据。再举一个例子,你可以在 CLI 中使用内置的 UDTF parquet_metadata 来读取 Parquet 文件的元数据。

> select filename, row_group_id, row_group_num_rows, row_group_bytes, stats_min, stats_max from parquet_metadata('./benchmarks/data/hits.parquet') where  column_id = 17 limit 10;
+--------------------------------+--------------+--------------------+-----------------+-----------+-----------+
| filename                       | row_group_id | row_group_num_rows | row_group_bytes | stats_min | stats_max |
+--------------------------------+--------------+--------------------+-----------------+-----------+-----------+
| ./benchmarks/data/hits.parquet | 0            | 450560             | 188921521       | 0         | 73256     |
| ./benchmarks/data/hits.parquet | 1            | 612174             | 210338885       | 0         | 109827    |
| ./benchmarks/data/hits.parquet | 2            | 344064             | 161242466       | 0         | 122484    |
| ./benchmarks/data/hits.parquet | 3            | 606208             | 235549898       | 0         | 121073    |
| ./benchmarks/data/hits.parquet | 4            | 335872             | 137103898       | 0         | 108996    |
| ./benchmarks/data/hits.parquet | 5            | 311296             | 145453612       | 0         | 108996    |
| ./benchmarks/data/hits.parquet | 6            | 303104             | 138833963       | 0         | 108996    |
| ./benchmarks/data/hits.parquet | 7            | 303104             | 191140113       | 0         | 73256     |
| ./benchmarks/data/hits.parquet | 8            | 573440             | 208038598       | 0         | 95823     |
| ./benchmarks/data/hits.parquet | 9            | 344064             | 147838157       | 0         | 73256     |
+--------------------------------+--------------+--------------------+-----------------+-----------+-----------+

编写 UDTF

这里使用的简单 UDTF 接受一个 Int64 参数,并返回一张只有一列的表,该列的值即为参数值。要在 DataFusion 中创建函数,你需要实现 TableFunctionImpl trait。该 trait 只有一个方法 call_with_args,它接受一个 TableFunctionArgs 结构体并返回 Result<Arc<dyn TableProvider>>。传入的结构体中包含以 Expr 切片形式表示的函数参数。

在 call_with_args 方法中,你需要解析输入的 Expr 并返回一个 TableProvider。你可能还想对输入的 Expr 做一些校验,例如检查参数的数量是否正确。

use std::sync::Arc;
use datafusion::common::{plan_err, ScalarValue, Result};
use datafusion::catalog::{TableFunctionArgs, TableFunctionImpl, TableProvider};
use datafusion::arrow::array::{ArrayRef, Int64Array};
use datafusion::datasource::memory::MemTable;
use arrow::record_batch::RecordBatch;
use arrow::datatypes::{DataType, Field, Schema};
use datafusion_expr::Expr;

/// A table function that returns a table provider with the value as a single column
#[derive(Debug)]
pub struct EchoFunction {}

impl TableFunctionImpl for EchoFunction {
    fn call_with_args(&self, args: TableFunctionArgs) -> Result<Arc<dyn TableProvider>> {
        let exprs = args.exprs();
        let Some(Expr::Literal(ScalarValue::Int64(Some(value)), _)) = exprs.get(0) else {
            return plan_err!("First argument must be an integer");
        };

        // Create the schema for the table
        let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)]));

        // Create a single RecordBatch with the value as a single column
        let batch = RecordBatch::try_new(
            schema.clone(),
            vec![Arc::new(Int64Array::from(vec![*value]))],
        )?;

        // Create a MemTable plan that returns the RecordBatch
        let provider = MemTable::try_new(schema, vec![vec![batch]])?;

        Ok(Arc::new(provider))
    }
}

注册和使用 UDTF

UDTF 实现完成后,你可以将其注册到 SessionContext 中:

use datafusion::execution::context::SessionContext;
use datafusion::arrow::util::pretty;

#[tokio::main]
async fn main() -> Result<()> {
    let ctx = SessionContext::new();

    ctx.register_udtf("echo", Arc::new(EchoFunction::default()));

    // And if all goes well, you can use it in your query:

    let results = ctx.sql("SELECT * FROM echo(1)").await?.collect().await?;
    pretty::print_batches(&results)?;
    Ok(())
}

// +---+
// | a |
// +---+
// | 1 |
// +---+

自定义表达式规划

DataFusion 默认对常见的 SQL 运算符和语法结构提供原生支持,例如 +、-、||。但它不支持一些其他运算符,例如 @>,也不支持 TABLESAMPLE 这类在不同 SQL 方言之间差异较大、使用较少的语法结构。为了覆盖 DataFusion 的默认处理方式或支持这些不受支持的特性,开发者可以通过实现自定义表达式规划来扩展 DataFusion,这是 DataFusion 的核心特性之一。

有关扩展 SQL 语法(包括 ExprPlanner、TypePlanner 和 RelationPlanner)的完整指南,请参阅扩展 DataFusion 的 SQL 语法。

实现自定义表达式规划

要扩展 DataFusion 以支持其原生不提供的自定义运算符,你需要:

  1. 实现 ExprPlanner trait:这使你能够为 DataFusion 原生无法识别的表达式定义自定义规划逻辑。该 trait 提供了将 SQL AST 节点转换为逻辑 Expr 所需的接口。

    详细文档请参阅:Trait ExprPlanner

  2. 注册你的自定义规划器:将你的实现与 DataFusion 的 SessionContext 集成,以确保在查询优化和执行规划阶段调用你的自定义规划逻辑。

    详细文档请参阅:fn register_expr_planner

示例如下:

// Implement ExprPlanner to add support for the `->` custom operator
impl ExprPlanner for MyCustomPlanner {
    fn plan_binary_op(
        &self,
        expr: RawBinaryExpr,
        _schema: &DFSchema,
    ) -> Result<PlannerResult<RawBinaryExpr>> {
        match &expr.op {
            // Map `->` to string concatenation
            BinaryOperator::Arrow => {
                // Rewrite `->` as a string concatenation operation
                // - `left` and `right` are the operands (e.g., 'hello' and 'world')
                // - `Operator::StringConcat` tells DataFusion to concatenate them
                Ok(PlannerResult::Planned(Expr::BinaryExpr(BinaryExpr {
                    left: Box::new(expr.left.clone()),
                    right: Box::new(expr.right.clone()),
                    op: Operator::StringConcat,
                })))
            }
            _ => Ok(PlannerResult::Original(expr)),
        }
    }
}

use datafusion::execution::context::SessionContext;
use datafusion::arrow::util::pretty;

#[tokio::main]
async fn main() -> Result<()> {
    let config = SessionConfig::new().set_str("datafusion.sql_parser.dialect", "postgres");
    let mut ctx = SessionContext::new_with_config(config);
    ctx.register_expr_planner(Arc::new(MyCustomPlanner))?;
    let results = ctx.sql("select 'foo'->'bar';").await?.collect().await?;

    let expected = [
         "+----------------------------+",
         "| Utf8(\"foo\") || Utf8(\"bar\") |",
         "+----------------------------+",
         "| foobar                     |",
         "+----------------------------+",
     ];
    assert_batches_eq!(&expected, &results);

    pretty::print_batches(&results)?;
    Ok(())
}

评论

登录后参与评论

正在加载评论…