更新时间:2026-10-08 GMT+08:00
分享

Scalar UDF

Scalar UDF是最基本的行处理函数, 它接受零个或多个输入参数,并对这一行数据进行操作,最终返回一行结果。

通常Scalar UDF适用于数学运算、复杂类型转换和自定义格式化等行处理的场景。

Scalar UDF的行为与内置函数非常相似,但其具体实现由用户自定义。

示例

import os
import ibis
from aura_frame.multimodal import ai_lake
target_database = "test"
#连接AuraJobV2端点
con = ai_lake.connect(
    aura_endpoint=os.getenv("aura_endpoint"),         # AuraJobV2端点所在的终端节点(Endpoint)
    aura_endpoint_id=os.getenv("aura_endpoint_id"),    # AuraJobV2端点ID
    aura_endpoint_name=os.getenv("aura_endpoint_name"),# AuraJobV2端点名称
    aura_workspace_id=os.getenv("aura_workspace_id"),  # 工作空间ID
    lf_catalog_name=os.getenv("lf_catalog_name"),      # LakeFormation Catalog名称
    lf_instance_id=os.getenv("lf_instance_id"),        # LakeFormation实例ID
    project_id=os.getenv("project_id"),                # 项目ID
    access_key=os.getenv("access_key"),                # 访问密钥ID(AK)
    secret_key=os.getenv("secret_key"),                # 访问密钥(SK)
    default_database=target_database,                  # 默认数据库
    use_single_cn_mode=True,                           # 是否开启单CN模式
)
#配置UDF部署工作区
con.set_function_staging_workspace(
    obs_directory_base=os.getenv("obs_directory_base"),  # 并行文件系统中UDF的存储路径
    obs_bucket_name=os.getenv("obs_bucket_name"),        # 并行文件系统名称
    obs_server=os.getenv("obs_server"),                  # OBS终端节点(Endpoint)
    access_key=os.getenv("access_key"),                  # 访问密钥ID(AK)
    secret_key=os.getenv("secret_key"),                  # 访问密钥(SK)
)
# 定义 Scalar UDF:两数求和
def add_udf(lhs: float, rhs: int) -> float:
    return lhs + rhs
try:
    # 创建测试表并写入数据
    table_name = "product_orders"
    table_schema = ibis.schema({"price": "float64", "quantity": "int32"})
    con.sql(f'DROP TABLE IF EXISTS "{target_database}"."{table_name}";')
    con.create_table(
        table_name,
        schema=table_schema,
        database=target_database,
        table_format="orc",
        if_not_exists=True,
    )
    con.sql(f'''
    INSERT INTO "{target_database}"."{table_name}" (price, quantity) VALUES
        (100.0, 10),
        (200.0, 20),
        (150.0, 15)
    ''')
    # 删除已存在的同名UDF(支持重复执行)
    con.delete_function("add_udf", database=target_database, if_it_exists=True)
    # 注册 Scalar UDF
    udf = con.create_scalar_function(
        add_udf,
        database=target_database,
        comment="To test create_scalar_function",
    )
    # 使用UDF处理product_orders表数据
    ds = con.load_dataset("product_orders", database=target_database)
    ds = ds.map(fn=udf, on=[ds.price, ds.quantity], as_col="sum_column")
    ds = ds.select_columns(ds.price, ds.quantity, ds.sum_column)
    print(ds.execute())
finally:
    con.close()

相关文档