Skip to content

Commit ec27c9c

Browse files
timsaucerclaude
andcommitted
feat: accept a bare PyCapsule in ScalarUDF and WindowUDF
AggregateUDF already accepted the capsule returned by __datafusion_aggregate_udf__ as well as an object exposing it (#1277). Extend the same to ScalarUDF and WindowUDF, with matching overloads on udf and udwf, so the three UDF kinds import the same way. Add FFI example tests for all three, including the previously untested AggregateUDF path. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 8ef34c0 commit ec27c9c

6 files changed

Lines changed: 92 additions & 23 deletions

File tree

‎crates/core/src/udf.rs‎

Lines changed: 18 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -209,6 +209,16 @@ impl ScalarUDFImpl for PythonFunctionScalarUDF {
209209
}
210210
}
211211

212+
fn scalar_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult<ScalarUDF> {
213+
let data: NonNull<FFI_ScalarUDF> = capsule
214+
.pointer_checked(Some(c"datafusion_scalar_udf"))?
215+
.cast();
216+
let udf = unsafe { data.as_ref() };
217+
let udf: Arc<dyn ScalarUDFImpl> = udf.into();
218+
219+
Ok(ScalarUDF::new_from_shared_impl(udf))
220+
}
221+
212222
/// Represents a PyScalarUDF
213223
#[pyclass(
214224
from_py_object,
@@ -247,22 +257,21 @@ impl PyScalarUDF {
247257

248258
#[staticmethod]
249259
pub fn from_pycapsule(func: Bound<'_, PyAny>) -> PyDataFusionResult<Self> {
260+
if func.is_instance_of::<PyCapsule>() {
261+
let capsule = func.cast::<PyCapsule>().map_err(to_datafusion_err)?;
262+
let function = scalar_udf_from_capsule(capsule)?;
263+
return Ok(Self { function });
264+
}
265+
250266
if func.hasattr("__datafusion_scalar_udf__")? {
251267
let capsule = call_capsule_getter(
252268
func.clone(),
253269
"__datafusion_scalar_udf__",
254270
CapsuleGetterArg::None,
255271
)?;
256272
let capsule = capsule.cast::<PyCapsule>().map_err(to_datafusion_err)?;
257-
let data: NonNull<FFI_ScalarUDF> = capsule
258-
.pointer_checked(Some(c"datafusion_scalar_udf"))?
259-
.cast();
260-
let udf = unsafe { data.as_ref() };
261-
let udf: Arc<dyn ScalarUDFImpl> = udf.into();
262-
263-
Ok(Self {
264-
function: ScalarUDF::new_from_shared_impl(udf),
265-
})
273+
let function = scalar_udf_from_capsule(capsule)?;
274+
Ok(Self { function })
266275
} else {
267276
Err(crate::errors::PyDataFusionError::Common(
268277
"__datafusion_scalar_udf__ does not exist on ScalarUDF object.".to_string(),

‎crates/core/src/udwf.rs‎

Lines changed: 18 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -216,6 +216,16 @@ pub fn to_rust_partition_evaluator(evaluator: Py<PyAny>) -> PartitionEvaluatorFa
216216
Arc::new(move || instantiate_partition_evaluator(&evaluator))
217217
}
218218

219+
fn window_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult<WindowUDF> {
220+
let data: NonNull<FFI_WindowUDF> = capsule
221+
.pointer_checked(Some(c"datafusion_window_udf"))?
222+
.cast();
223+
let udwf = unsafe { data.as_ref() };
224+
let udwf: Arc<dyn WindowUDFImpl> = udwf.into();
225+
226+
Ok(WindowUDF::new_from_shared_impl(udwf))
227+
}
228+
219229
/// Represents an WindowUDF
220230
#[pyclass(
221231
from_py_object,
@@ -262,19 +272,17 @@ impl PyWindowUDF {
262272

263273
#[staticmethod]
264274
pub fn from_pycapsule(func: Bound<'_, PyAny>) -> PyDataFusionResult<Self> {
275+
if func.is_instance_of::<PyCapsule>() {
276+
let capsule = func.cast::<PyCapsule>().map_err(to_datafusion_err)?;
277+
let function = window_udf_from_capsule(capsule)?;
278+
return Ok(Self { function });
279+
}
280+
265281
let capsule =
266282
call_capsule_getter(func, "__datafusion_window_udf__", CapsuleGetterArg::None)?;
267-
268283
let capsule = capsule.cast::<PyCapsule>().map_err(to_datafusion_err)?;
269-
let data: NonNull<FFI_WindowUDF> = capsule
270-
.pointer_checked(Some(c"datafusion_window_udf"))?
271-
.cast();
272-
let udwf = unsafe { data.as_ref() };
273-
let udwf: Arc<dyn WindowUDFImpl> = udwf.into();
274-
275-
Ok(Self {
276-
function: WindowUDF::new_from_shared_impl(udwf),
277-
})
284+
let function = window_udf_from_capsule(capsule)?;
285+
Ok(Self { function })
278286
}
279287

280288
fn __repr__(&self) -> PyResult<String> {

‎examples/datafusion-ffi-example/python/tests/_test_aggregate_udf.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -75,3 +75,11 @@ def test_ffi_aggregate_call_directly():
7575
]
7676

7777
assert result == expected
78+
79+
80+
def test_ffi_aggregate_from_bare_capsule():
81+
ctx = setup_context_with_table()
82+
my_udaf = udaf(MySumUDF().__datafusion_aggregate_udf__())
83+
84+
result = ctx.table("test_table").aggregate([], [my_udaf(col("a")).alias("r")])
85+
assert result.collect_column("r").to_pylist() == [6]

‎examples/datafusion-ffi-example/python/tests/_test_scalar_udf.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,3 +68,11 @@ def test_ffi_scalar_call_directly():
6868
]
6969

7070
assert result == expected
71+
72+
73+
def test_ffi_scalar_from_bare_capsule():
74+
ctx = setup_context_with_table()
75+
my_udf = udf(IsNullUDF().__datafusion_scalar_udf__())
76+
77+
result = ctx.table("test_table").select(my_udf(col("a")).alias("r"))
78+
assert result.collect_column("r").to_pylist() == [False, False, False, True]

‎examples/datafusion-ffi-example/python/tests/_test_window_udf.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,3 +87,15 @@ def test_ffi_window_call_directly():
8787
(40, 4),
8888
]
8989
assert results == expected
90+
91+
92+
def test_ffi_window_from_bare_capsule():
93+
ctx = setup_context_with_table()
94+
my_udwf = udwf(MyRankUDF().__datafusion_window_udf__())
95+
96+
result = (
97+
ctx.table("test_table")
98+
.select(col("a"), my_udwf().order_by(col("a")).build().alias("r"))
99+
.sort(col("a"))
100+
)
101+
assert result.collect_column("r").to_pylist() == [1, 2, 3, 4]

‎python/datafusion/user_defined.py‎

Lines changed: 28 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -283,6 +283,10 @@ def udf(
283283
@staticmethod
284284
def udf(func: ScalarUDFExportable) -> ScalarUDF: ...
285285

286+
@overload
287+
@staticmethod
288+
def udf(func: _PyCapsule) -> ScalarUDF: ...
289+
286290
@staticmethod
287291
def udf(*args: Any, **kwargs: Any): # noqa: D417
288292
"""Create a new User-Defined Function (UDF).
@@ -388,7 +392,7 @@ def wrapper(*args: Any, **kwargs: Any) -> Callable:
388392

389393
return decorator
390394

391-
if hasattr(args[0], "__datafusion_scalar_udf__"):
395+
if hasattr(args[0], "__datafusion_scalar_udf__") or _is_pycapsule(args[0]):
392396
return ScalarUDF.from_pycapsule(args[0])
393397

394398
if args and callable(args[0]):
@@ -398,12 +402,18 @@ def wrapper(*args: Any, **kwargs: Any) -> Callable:
398402
return _decorator(*args, **kwargs)
399403

400404
@staticmethod
401-
def from_pycapsule(func: ScalarUDFExportable) -> ScalarUDF:
405+
def from_pycapsule(func: ScalarUDFExportable | _PyCapsule) -> ScalarUDF:
402406
"""Create a Scalar UDF from ScalarUDF PyCapsule object.
403407
404408
This function will instantiate a Scalar UDF that uses a DataFusion
405409
ScalarUDF that is exported via the FFI bindings.
406410
"""
411+
if _is_pycapsule(func):
412+
scalar = cast("ScalarUDF", object.__new__(ScalarUDF))
413+
scalar._udf = df_internal.ScalarUDF.from_pycapsule(func)
414+
return scalar
415+
416+
func = cast("ScalarUDFExportable", func)
407417
name = str(func.__class__)
408418
return ScalarUDF(
409419
name=name,
@@ -1011,6 +1021,14 @@ def udwf(
10111021
name: str | None = None,
10121022
) -> WindowUDF: ...
10131023

1024+
@overload
1025+
@staticmethod
1026+
def udwf(func: WindowUDFExportable) -> WindowUDF: ...
1027+
1028+
@overload
1029+
@staticmethod
1030+
def udwf(func: _PyCapsule) -> WindowUDF: ...
1031+
10141032
@staticmethod
10151033
def udwf(*args: Any, **kwargs: Any): # noqa: D417
10161034
"""Create a new User-Defined Window Function (UDWF).
@@ -1075,7 +1093,7 @@ def udwf(*args: Any, **kwargs: Any): # noqa: D417
10751093
Returns:
10761094
A user-defined window function that can be used in window function calls.
10771095
"""
1078-
if hasattr(args[0], "__datafusion_window_udf__"):
1096+
if hasattr(args[0], "__datafusion_window_udf__") or _is_pycapsule(args[0]):
10791097
return WindowUDF.from_pycapsule(args[0])
10801098

10811099
if args and callable(args[0]):
@@ -1146,12 +1164,18 @@ def wrapper(*args: Any, **kwargs: Any) -> Expr:
11461164
return decorator
11471165

11481166
@staticmethod
1149-
def from_pycapsule(func: WindowUDFExportable) -> WindowUDF:
1167+
def from_pycapsule(func: WindowUDFExportable | _PyCapsule) -> WindowUDF:
11501168
"""Create a Window UDF from WindowUDF PyCapsule object.
11511169
11521170
This function will instantiate a Window UDF that uses a DataFusion
11531171
WindowUDF that is exported via the FFI bindings.
11541172
"""
1173+
if _is_pycapsule(func):
1174+
window = cast("WindowUDF", object.__new__(WindowUDF))
1175+
window._udwf = df_internal.WindowUDF.from_pycapsule(func)
1176+
return window
1177+
1178+
func = cast("WindowUDFExportable", func)
11551179
name = str(func.__class__)
11561180
return WindowUDF(
11571181
name=name,

0 commit comments

Comments
 (0)