From 5a0e79ddd77e632d466ce47d58a5a94fa00276c5 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:19:42 -0400 Subject: [PATCH 01/31] feat: expose any_value, array math, array_first, and missing Spark functions Close gaps found by auditing the Python API against upstream DataFusion 55.1.0. - Add the any_value aggregate. - Add array_add, array_subtract, array_scale, array_sum, array_avg, array_product, and array_first, each with its list_* alias. - Add Spark monthname, weekday, atan2, hypot, pow/power, quote, and concat_ws. concat_ws calls the UDF directly so *cols stays variadic. - Fix the aggregations guide, which listed regr_slope twice and omitted regr_sxy. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/functions.rs | 22 ++ crates/core/src/spark_functions.rs | 23 ++ .../common-operations/aggregations.md | 3 +- python/datafusion/functions/__init__.py | 261 ++++++++++++++++++ python/datafusion/functions/spark.py | 127 +++++++++ python/tests/test_aggregation.py | 14 + python/tests/test_functions.py | 70 ++++- python/tests/test_lambda.py | 17 ++ python/tests/test_spark_functions.py | 27 ++ 9 files changed, 562 insertions(+), 2 deletions(-) diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index e57c7702d..c8fd77c51 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -197,6 +197,13 @@ fn array_filter(array: PyExpr, predicate: PyExpr) -> PyExpr { datafusion::functions_nested::expr_fn::array_filter(array.into(), predicate.into()).into() } +/// Higher-order function: return the first element of `array` for which +/// `predicate` (a lambda returning a boolean) is true, or null if none match. +#[pyfunction] +fn array_first(array: PyExpr, predicate: PyExpr) -> PyExpr { + datafusion::functions_nested::expr_fn::array_first(array.into(), predicate.into()).into() +} + /// Computes a binary hash of the given data. type is the algorithm to use. /// Standard algorithms are md5, sha224, sha256, sha384, sha512, blake2s, blake2b, and blake3. // #[pyfunction(value, method)] @@ -663,6 +670,12 @@ array_fn!(array_compact, array); array_fn!(array_normalize, array); array_fn!(cosine_distance, array1 array2); array_fn!(inner_product, array1 array2); +array_fn!(array_add, array1 array2); +array_fn!(array_subtract, array1 array2); +array_fn!(array_scale, array scalar); +array_fn!(array_sum, array); +array_fn!(array_avg, array); +array_fn!(array_product, array); array_fn!(array_intersect, first_array second_array); array_fn!(array_union, array1 array2); array_fn!(array_except, first_array second_array); @@ -688,6 +701,7 @@ aggregate_function!(avg); aggregate_function!(sum); aggregate_function!(bit_and); aggregate_function!(bit_or); +aggregate_function!(any_value); aggregate_function!(bit_xor); aggregate_function!(bool_and); aggregate_function!(bool_or); @@ -1126,6 +1140,7 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(nth_value))?; m.add_wrapped(wrap_pyfunction!(bit_and))?; m.add_wrapped(wrap_pyfunction!(bit_or))?; + m.add_wrapped(wrap_pyfunction!(any_value))?; m.add_wrapped(wrap_pyfunction!(bit_xor))?; m.add_wrapped(wrap_pyfunction!(bool_and))?; m.add_wrapped(wrap_pyfunction!(bool_or))?; @@ -1140,6 +1155,7 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(array_transform))?; m.add_wrapped(wrap_pyfunction!(array_any_match))?; m.add_wrapped(wrap_pyfunction!(array_filter))?; + m.add_wrapped(wrap_pyfunction!(array_first))?; // Array Functions m.add_wrapped(wrap_pyfunction!(array_append))?; @@ -1151,6 +1167,12 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(array_normalize))?; m.add_wrapped(wrap_pyfunction!(cosine_distance))?; m.add_wrapped(wrap_pyfunction!(inner_product))?; + m.add_wrapped(wrap_pyfunction!(array_add))?; + m.add_wrapped(wrap_pyfunction!(array_subtract))?; + m.add_wrapped(wrap_pyfunction!(array_scale))?; + m.add_wrapped(wrap_pyfunction!(array_sum))?; + m.add_wrapped(wrap_pyfunction!(array_avg))?; + m.add_wrapped(wrap_pyfunction!(array_product))?; m.add_wrapped(wrap_pyfunction!(array_element))?; m.add_wrapped(wrap_pyfunction!(array_empty))?; m.add_wrapped(wrap_pyfunction!(array_length))?; diff --git a/crates/core/src/spark_functions.rs b/crates/core/src/spark_functions.rs index e7cb94f8c..3ff2af250 100644 --- a/crates/core/src/spark_functions.rs +++ b/crates/core/src/spark_functions.rs @@ -155,6 +155,7 @@ spark_expr_fn!(date_add, start_date days); spark_expr_fn!(date_sub, start_date days); spark_expr_fn!(hour, arg1); spark_expr_fn!(minute, arg1); +spark_expr_fn!(monthname, arg1); spark_expr_fn!(second, arg1); spark_expr_fn!(last_day, arg1); spark_expr_fn!(make_dt_interval, days hours mins secs); @@ -171,6 +172,7 @@ spark_expr_fn!(unix_date, dt); spark_expr_fn!(unix_micros, ts); spark_expr_fn!(unix_millis, ts); spark_expr_fn!(unix_seconds, ts); +spark_expr_fn!(weekday, arg1); // --------------------------------------------------------------------------- // Hash functions @@ -200,13 +202,16 @@ spark_expr_fn!(str_to_map, text pair_delim key_value_delim); // --------------------------------------------------------------------------- spark_expr_fn!(abs, arg1); +spark_expr_fn!(atan2, arg1 arg2); spark_expr_fn!(ceil, arg1); spark_expr_fn!(expm1, arg1); spark_expr_fn!(factorial, arg1); spark_expr_fn!(floor, arg1); spark_expr_fn!(hex, arg1); +spark_expr_fn!(hypot, arg1 arg2); spark_expr_fn!(modulus, dividend divisor); spark_expr_fn!(pmod, dividend divisor); +spark_expr_fn!(pow, arg1 arg2); spark_expr_fn!(rint, arg1); spark_expr_fn!(round, value scale); spark_expr_fn!(unhex, arg1); @@ -230,12 +235,23 @@ fn char_fn(arg1: PyExpr) -> PyExpr { expr_fn::char(arg1.into()).into() } spark_udf_vec!(concat, udf::string::concat); +/// `concat_ws(sep, *cols)`. The upstream `expr_fn::concat_ws` takes a single +/// `Expr` for the values, so call the UDF directly to keep `*cols` variadic. +#[pyfunction] +#[pyo3(signature = (sep, *cols))] +fn concat_ws(sep: PyExpr, cols: Vec) -> PyExpr { + let args: Vec = std::iter::once(sep.into()) + .chain(cols.into_iter().map(Into::into)) + .collect(); + Expr::ScalarFunction(ScalarFunction::new_udf(udf::string::concat_ws(), args)).into() +} spark_udf_vec!(elt, udf::string::elt); spark_expr_fn!(ilike, str pattern); spark_expr_fn!(length, arg1); spark_expr_fn!(like, str pattern); spark_expr_fn!(luhn_check, arg1); spark_udf_vec!(format_string, udf::string::format_string); +spark_expr_fn!(quote, arg1); spark_expr_fn!(space, arg1); spark_expr_fn!(substring, str pos length); spark_expr_fn!(unbase64, str); @@ -292,6 +308,7 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(date_sub))?; m.add_wrapped(wrap_pyfunction!(hour))?; m.add_wrapped(wrap_pyfunction!(minute))?; + m.add_wrapped(wrap_pyfunction!(monthname))?; m.add_wrapped(wrap_pyfunction!(second))?; m.add_wrapped(wrap_pyfunction!(last_day))?; m.add_wrapped(wrap_pyfunction!(make_dt_interval))?; @@ -308,6 +325,7 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(unix_micros))?; m.add_wrapped(wrap_pyfunction!(unix_millis))?; m.add_wrapped(wrap_pyfunction!(unix_seconds))?; + m.add_wrapped(wrap_pyfunction!(weekday))?; // Hash m.add_wrapped(wrap_pyfunction!(crc32))?; m.add_wrapped(wrap_pyfunction!(sha1))?; @@ -321,13 +339,16 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(str_to_map))?; // Math m.add_wrapped(wrap_pyfunction!(abs))?; + m.add_wrapped(wrap_pyfunction!(atan2))?; m.add_wrapped(wrap_pyfunction!(ceil))?; m.add_wrapped(wrap_pyfunction!(expm1))?; m.add_wrapped(wrap_pyfunction!(factorial))?; m.add_wrapped(wrap_pyfunction!(floor))?; m.add_wrapped(wrap_pyfunction!(hex))?; + m.add_wrapped(wrap_pyfunction!(hypot))?; m.add_wrapped(wrap_pyfunction!(modulus))?; m.add_wrapped(wrap_pyfunction!(pmod))?; + m.add_wrapped(wrap_pyfunction!(pow))?; m.add_wrapped(wrap_pyfunction!(rint))?; m.add_wrapped(wrap_pyfunction!(round))?; m.add_wrapped(wrap_pyfunction!(unhex))?; @@ -341,11 +362,13 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(base64))?; m.add_wrapped(wrap_pyfunction!(char_fn))?; m.add_wrapped(wrap_pyfunction!(concat))?; + m.add_wrapped(wrap_pyfunction!(concat_ws))?; m.add_wrapped(wrap_pyfunction!(elt))?; m.add_wrapped(wrap_pyfunction!(ilike))?; m.add_wrapped(wrap_pyfunction!(length))?; m.add_wrapped(wrap_pyfunction!(like))?; m.add_wrapped(wrap_pyfunction!(luhn_check))?; + m.add_wrapped(wrap_pyfunction!(quote))?; m.add_wrapped(wrap_pyfunction!(format_string))?; m.add_wrapped(wrap_pyfunction!(space))?; m.add_wrapped(wrap_pyfunction!(substring))?; diff --git a/docs/source/user-guide/common-operations/aggregations.md b/docs/source/user-guide/common-operations/aggregations.md index 9c2d58e3c..0e073185e 100644 --- a/docs/source/user-guide/common-operations/aggregations.md +++ b/docs/source/user-guide/common-operations/aggregations.md @@ -393,11 +393,12 @@ The available aggregate functions are: - {py:func}`datafusion.functions.regr_avgy` - {py:func}`datafusion.functions.regr_sxx` - {py:func}`datafusion.functions.regr_syy` - - {py:func}`datafusion.functions.regr_slope` + - {py:func}`datafusion.functions.regr_sxy` 07. Positional Functions : - {py:func}`datafusion.functions.first_value` - {py:func}`datafusion.functions.last_value` - {py:func}`datafusion.functions.nth_value` + - {py:func}`datafusion.functions.any_value` 08. String Functions : - {py:func}`datafusion.functions.string_agg` 09. Percentile Functions diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 291957490..39a31e00d 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -82,15 +82,18 @@ def _warn_if_expr_for_literal_arg( "acosh", "alias", "any_match", + "any_value", "approx_distinct", "approx_median", "approx_percentile_cont", "approx_percentile_cont_with_weight", "array", + "array_add", "array_agg", "array_any_match", "array_any_value", "array_append", + "array_avg", "array_cat", "array_compact", "array_concat", @@ -103,6 +106,7 @@ def _warn_if_expr_for_literal_arg( "array_except", "array_extract", "array_filter", + "array_first", "array_has", "array_has_all", "array_has_any", @@ -119,6 +123,7 @@ def _warn_if_expr_for_literal_arg( "array_position", "array_positions", "array_prepend", + "array_product", "array_push_back", "array_push_front", "array_remove", @@ -130,8 +135,11 @@ def _warn_if_expr_for_literal_arg( "array_replace_n", "array_resize", "array_reverse", + "array_scale", "array_slice", "array_sort", + "array_subtract", + "array_sum", "array_to_string", "array_transform", "array_union", @@ -230,9 +238,11 @@ def _warn_if_expr_for_literal_arg( "left", "length", "levenshtein", + "list_add", "list_any_match", "list_any_value", "list_append", + "list_avg", "list_cat", "list_compact", "list_concat", @@ -245,6 +255,7 @@ def _warn_if_expr_for_literal_arg( "list_except", "list_extract", "list_filter", + "list_first", "list_has", "list_has_all", "list_has_any", @@ -262,6 +273,7 @@ def _warn_if_expr_for_literal_arg( "list_position", "list_positions", "list_prepend", + "list_product", "list_push_back", "list_push_front", "list_remove", @@ -273,8 +285,11 @@ def _warn_if_expr_for_literal_arg( "list_replace_n", "list_resize", "list_reverse", + "list_scale", "list_slice", "list_sort", + "list_subtract", + "list_sum", "list_to_string", "list_transform", "list_union", @@ -712,6 +727,46 @@ def list_filter(array: Expr, predicate: Expr | Callable[..., Any]) -> Expr: return array_filter(array, predicate) +def array_first(array: Expr, predicate: Expr | Callable[..., Any]) -> Expr: + """Return the first element of ``array`` for which ``predicate`` is ``True``. + + ``predicate`` may be a Python callable, converted to a lambda + automatically, or an explicit lambda built with :py:func:`lambda_`. It must + return a boolean expression. Returns NULL if no element matches. + + Examples: + Using a Python callable: + + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": [[1, 2, 3, 4]]}) + >>> df.select( + ... F.array_first(col("a"), lambda v: v > 2).alias("f") + ... ).collect_column("f")[0].as_py() + 3 + + Using an explicit lambda built with :py:func:`lambda_`: + + >>> predicate = F.lambda_(["v"], F.lambda_var("v") > lit(2)) + >>> df.select( + ... F.array_first(col("a"), predicate).alias("f") + ... ).collect_column("f")[0].as_py() + 3 + + See Also: + :py:func:`array_filter`, :py:func:`array_any_match`, :py:func:`lambda_`. + """ + return Expr(f.array_first(array.expr, _to_lambda(predicate).expr)) + + +def list_first(array: Expr, predicate: Expr | Callable[..., Any]) -> Expr: + """Return the first element of a list for which a predicate is ``True``. + + See Also: + This is an alias for :py:func:`array_first`. + """ + return array_first(array, predicate) + + def in_list(arg: Expr, values: list[Expr], negated: bool = False) -> Expr: """Returns whether the argument is contained within the list ``values``. @@ -3687,6 +3742,122 @@ def dot_product(array1: Expr, array2: Expr) -> Expr: return inner_product(array1, array2) +def array_add(array1: Expr, array2: Expr) -> Expr: + """Returns the element-wise sum of two numeric arrays of equal length. + + A NULL element in either input produces a NULL at that position. Execution + fails if the arrays in a row have different lengths. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict( + ... {"a": [[1.0, 2.0, 3.0]], "b": [[10.0, 20.0, 30.0]]} + ... ) + >>> result = df.select( + ... dfn.functions.array_add(dfn.col("a"), dfn.col("b")).alias("result") + ... ) + >>> result.collect_column("result")[0].as_py() + [11.0, 22.0, 33.0] + """ + return Expr(f.array_add(array1.expr, array2.expr)) + + +def array_subtract(array1: Expr, array2: Expr) -> Expr: + """Returns the element-wise difference of two numeric arrays of equal length. + + Computes ``array1[i] - array2[i]``. A NULL element in either input produces + a NULL at that position. Execution fails if the arrays in a row have + different lengths. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict( + ... {"a": [[10.0, 20.0, 30.0]], "b": [[1.0, 2.0, 3.0]]} + ... ) + >>> result = df.select( + ... dfn.functions.array_subtract( + ... dfn.col("a"), dfn.col("b") + ... ).alias("result") + ... ) + >>> result.collect_column("result")[0].as_py() + [9.0, 18.0, 27.0] + """ + return Expr(f.array_subtract(array1.expr, array2.expr)) + + +def array_scale(array: Expr, scalar: Expr | float) -> Expr: + """Multiplies each element of a numeric array by ``scalar``. + + A NULL element produces a NULL at that position. Returns NULL if ``scalar`` + is NULL. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": [[1.0, 2.0, 3.0]]}) + >>> result = df.select( + ... dfn.functions.array_scale(dfn.col("a"), 2.0).alias("result") + ... ) + >>> result.collect_column("result")[0].as_py() + [2.0, 4.0, 6.0] + """ + scalar = coerce_to_expr(scalar) + return Expr(f.array_scale(array.expr, scalar.expr)) + + +def array_sum(array: Expr) -> Expr: + """Returns the sum of the elements of a numeric array. + + NULL elements are skipped. Returns NULL if the array is NULL, empty, or + contains only NULL elements. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": [[1.0, None, 3.0]]}) + >>> result = df.select( + ... dfn.functions.array_sum(dfn.col("a")).alias("result") + ... ) + >>> result.collect_column("result")[0].as_py() + 4.0 + """ + return Expr(f.array_sum(array.expr)) + + +def array_avg(array: Expr) -> Expr: + """Returns the arithmetic mean of the elements of a numeric array. + + NULL elements are skipped and excluded from the count. Returns NULL if the + array is NULL, empty, or contains only NULL elements. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": [[1.0, None, 3.0]]}) + >>> result = df.select( + ... dfn.functions.array_avg(dfn.col("a")).alias("result") + ... ) + >>> result.collect_column("result")[0].as_py() + 2.0 + """ + return Expr(f.array_avg(array.expr)) + + +def array_product(array: Expr) -> Expr: + """Returns the product of the elements of a numeric array. + + NULL elements are skipped. Returns NULL if the array is NULL, empty, or + contains only NULL elements. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": [[2.0, None, 3.0]]}) + >>> result = df.select( + ... dfn.functions.array_product(dfn.col("a")).alias("result") + ... ) + >>> result.collect_column("result")[0].as_py() + 6.0 + """ + return Expr(f.array_product(array.expr)) + + def list_cat(*args: Expr) -> Expr: """Concatenates the input arrays. @@ -3732,6 +3903,60 @@ def list_normalize(array: Expr) -> Expr: return array_normalize(array) +def list_add(array1: Expr, array2: Expr) -> Expr: + """Returns the element-wise sum of two numeric lists of equal length. + + See Also: + This is an alias for :py:func:`array_add`. + """ + return array_add(array1, array2) + + +def list_subtract(array1: Expr, array2: Expr) -> Expr: + """Returns the element-wise difference of two numeric lists of equal length. + + See Also: + This is an alias for :py:func:`array_subtract`. + """ + return array_subtract(array1, array2) + + +def list_scale(array: Expr, scalar: Expr | float) -> Expr: + """Multiplies each element of a numeric list by a scalar. + + See Also: + This is an alias for :py:func:`array_scale`. + """ + return array_scale(array, scalar) + + +def list_sum(array: Expr) -> Expr: + """Returns the sum of the elements of a numeric list. + + See Also: + This is an alias for :py:func:`array_sum`. + """ + return array_sum(array) + + +def list_avg(array: Expr) -> Expr: + """Returns the arithmetic mean of the elements of a numeric list. + + See Also: + This is an alias for :py:func:`array_avg`. + """ + return array_avg(array) + + +def list_product(array: Expr) -> Expr: + """Returns the product of the elements of a numeric list. + + See Also: + This is an alias for :py:func:`array_product`. + """ + return array_product(array) + + def list_dims(array: Expr) -> Expr: """Returns an array of the array's dimensions. @@ -6360,6 +6585,42 @@ def nth_value( ) +def any_value(expression: Expr, filter: Expr | None = None) -> Expr: + """Returns an arbitrary non-null value from each group. + + Returns NULL if every value in the group is NULL. Which value is returned + is not specified and may differ between runs. + + If using the builder functions described in ref:`_aggregation` this function ignores + the options ``order_by``, ``null_treatment``, and ``distinct``. + + Args: + expression: Argument to pick a value from + filter: If provided, only consider rows for which the filter is True + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": [None, 7, None]}) + >>> result = df.aggregate( + ... [], [dfn.functions.any_value(dfn.col("a")).alias("v")] + ... ) + >>> result.collect_column("v")[0].as_py() + 7 + + >>> df = ctx.from_pydict({"a": [None, 7, 8], "b": [1, 2, 3]}) + >>> result = df.aggregate( + ... [], [dfn.functions.any_value( + ... dfn.col("a"), + ... filter=dfn.col("b") > dfn.lit(2) + ... ).alias("v")] + ... ) + >>> result.collect_column("v")[0].as_py() + 8 + """ + filter_raw = filter.expr if filter is not None else None + return Expr(f.any_value(expression.expr, filter=filter_raw)) + + def bit_and(expression: Expr, filter: Expr | None = None) -> Expr: """Computes the bitwise AND of the argument. diff --git a/python/datafusion/functions/spark.py b/python/datafusion/functions/spark.py index 0a0f41400..7d84583d2 100644 --- a/python/datafusion/functions/spark.py +++ b/python/datafusion/functions/spark.py @@ -593,6 +593,22 @@ def minute(col: Expr) -> Expr: return Expr(_f.minute(col.expr)) +def monthname(col: Expr) -> Expr: + """Spark ``monthname``: three-letter abbreviated month name. + + Examples: + >>> import pyarrow as pa + >>> from datetime import date + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> d = dfn.lit(pa.scalar(date(2024, 3, 15))) + >>> r = df.select(dfn.functions.spark.monthname(d).alias("v")) + >>> r.collect_column("v")[0].as_py() + 'Mar' + """ + return Expr(_f.monthname(col.expr)) + + def second(col: Expr) -> Expr: """Spark ``second``: extract second component of a timestamp. @@ -933,6 +949,22 @@ def unix_seconds(col: Expr) -> Expr: # --------------------------------------------------------------------------- +def weekday(col: Expr) -> Expr: + """Spark ``weekday``: day of the week, Monday = 0 through Sunday = 6. + + Examples: + >>> import pyarrow as pa + >>> from datetime import date + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> d = dfn.lit(pa.scalar(date(2024, 3, 15))) + >>> r = df.select(dfn.functions.spark.weekday(d).alias("v")) + >>> r.collect_column("v")[0].as_py() + 4 + """ + return Expr(_f.weekday(col.expr)) + + def crc32(col: Expr) -> Expr: """Spark ``crc32``: cyclic redundancy check value as a bigint. @@ -1114,6 +1146,22 @@ def abs(col: Expr) -> Expr: return Expr(_f.abs(col.expr)) +def atan2(col1: Expr | float, col2: Expr | float) -> Expr: + """Spark ``atan2``: angle in radians of the point ``(col2, col1)``. + + ``col1`` is the y coordinate and ``col2`` the x coordinate. Both accept + native numbers or :class:`Expr`. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> r = df.select(dfn.functions.spark.atan2(1.0, 0.0).alias("v")) + >>> r.collect_column("v")[0].as_py() + 1.5707963267948966 + """ + return Expr(_f.atan2(coerce_to_expr(col1).expr, coerce_to_expr(col2).expr)) + + def ceil(col: Expr) -> Expr: """Spark ``ceil``: smallest integer ≥ arg. @@ -1184,6 +1232,21 @@ def hex(col: Expr) -> Expr: return Expr(_f.hex(col.expr)) +def hypot(col1: Expr | float, col2: Expr | float) -> Expr: + """Spark ``hypot``: ``sqrt(col1^2 + col2^2)`` without intermediate overflow. + + Both arguments accept native numbers or :class:`Expr`. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> r = df.select(dfn.functions.spark.hypot(3.0, 4.0).alias("v")) + >>> r.collect_column("v")[0].as_py() + 5.0 + """ + return Expr(_f.hypot(coerce_to_expr(col1).expr, coerce_to_expr(col2).expr)) + + def modulus(dividend: Expr | float, divisor: Expr | float) -> Expr: """Spark ``mod``: remainder of ``dividend / divisor`` (sign follows dividend). @@ -1216,6 +1279,30 @@ def pmod(dividend: Expr | float, divisor: Expr | float) -> Expr: return Expr(_f.pmod(coerce_to_expr(dividend).expr, coerce_to_expr(divisor).expr)) +def pow(col1: Expr | float, col2: Expr | float) -> Expr: + """Spark ``pow``: ``col1`` raised to the power ``col2``, as a double. + + Both arguments accept native numbers or :class:`Expr`. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> r = df.select(dfn.functions.spark.pow(2, 10).alias("v")) + >>> r.collect_column("v")[0].as_py() + 1024.0 + """ + return Expr(_f.pow(coerce_to_expr(col1).expr, coerce_to_expr(col2).expr)) + + +def power(col1: Expr | float, col2: Expr | float) -> Expr: + """Spark ``power``: ``col1`` raised to the power ``col2``. + + See Also: + This is an alias for :py:func:`pow`. + """ + return pow(col1, col2) + + def rint(col: Expr) -> Expr: """Spark ``rint``: round to nearest mathematical integer (as double). @@ -1400,6 +1487,25 @@ def concat(*cols: Expr) -> Expr: return Expr(_f.concat(*[c.expr for c in cols])) +def concat_ws(sep: Expr | str, *cols: Expr) -> Expr: + """Spark ``concat_ws``: joins strings and arrays of strings with ``sep``. + + NULL inputs are skipped rather than propagated. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": ["x"], "b": [None], "c": ["z"]}) + >>> r = df.select( + ... dfn.functions.spark.concat_ws( + ... "-", dfn.col("a"), dfn.col("b"), dfn.col("c") + ... ).alias("v") + ... ) + >>> r.collect_column("v")[0].as_py() + 'x-z' + """ + return Expr(_f.concat_ws(coerce_to_expr(sep).expr, *[c.expr for c in cols])) + + def elt(*inputs: Expr) -> Expr: """Spark ``elt``: returns the n-th input (1-indexed). @@ -1499,6 +1605,19 @@ def luhn_check(col: Expr) -> Expr: return Expr(_f.luhn_check(col.expr)) +def quote(col: Expr) -> Expr: + r"""Spark ``quote``: wraps a string in single quotes, escaping inner quotes. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> r = df.select(dfn.functions.spark.quote(dfn.lit("it's")).alias("v")) + >>> print(r.collect_column("v")[0].as_py()) + 'it\'s' + """ + return Expr(_f.quote(col.expr)) + + def format_string(format: str | Expr, *cols: Expr) -> Expr: """Spark ``format_string``: printf-style format string. @@ -1735,6 +1854,7 @@ def url_encode(str: Expr) -> Expr: # String "ascii", # Aggregate + "atan2", "avg", "base64", "bin", @@ -1752,6 +1872,7 @@ def url_encode(str: Expr) -> Expr: "collect_set", "concat", # Hash + "concat_ws", "crc32", "csc", "date_add", @@ -1767,6 +1888,7 @@ def url_encode(str: Expr) -> Expr: "from_utc_timestamp", "hex", "hour", + "hypot", "if_", "ilike", "is_valid_utf8", @@ -1784,11 +1906,15 @@ def url_encode(str: Expr) -> Expr: "map_from_entries", "minute", "modulus", + "monthname", "negative", "next_day", # URL "parse_url", "pmod", + "pow", + "power", + "quote", "rint", "round", "sec", @@ -1821,6 +1947,7 @@ def url_encode(str: Expr) -> Expr: "unix_seconds", "url_decode", "url_encode", + "weekday", "width_bucket", "xxhash64", ] diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index ef51343aa..735fef35d 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -317,7 +317,21 @@ def test_aggregate_100(df_aggregate_100, name, expr, expected): assert df.collect()[0].to_pydict() == expected_dict +def test_any_value_skips_nulls_per_group(): + ctx = SessionContext() + df = ctx.from_pydict( + {"g": ["x", "x", "y", "y", "z"], "v": [None, 7, 8, None, None]} + ) + result = ( + df.aggregate([column("g")], [f.any_value(column("v")).alias("v")]) + .sort(column("g").sort()) + .to_pydict() + ) + assert result == {"g": ["x", "y", "z"], "v": [7, 8, None]} + + data_test_bitwise_and_boolean_functions = [ + ("any_value_filter", f.any_value(column("a"), filter=column("a") == lit(2)), [2]), ("bit_and", f.bit_and(column("a")), [0]), ("bit_and_filter", f.bit_and(column("a"), filter=column("a") != lit(2)), [1]), ("bit_or", f.bit_or(column("b")), [6]), diff --git a/python/tests/test_functions.py b/python/tests/test_functions.py index fabe6dff2..49d9b33fb 100644 --- a/python/tests/test_functions.py +++ b/python/tests/test_functions.py @@ -745,6 +745,12 @@ def test_array_function_obj_tests(stmt, py_expr): f.inner_product, {"a": [[1.0, 2.0, 3.0]], "b": [[4.0, 5.0, 6.0]]}, ), + (f.list_add, f.array_add, {"a": [[1.0, 2.0]], "b": [[3.0, 4.0]]}), + (f.list_subtract, f.array_subtract, {"a": [[1.0, 2.0]], "b": [[3.0, 4.0]]}), + (f.list_scale, f.array_scale, {"a": [[1.0, 2.0]], "b": [3.0]}), + (f.list_sum, f.array_sum, {"a": [[1.0, 2.0, 3.0]]}), + (f.list_avg, f.array_avg, {"a": [[1.0, 2.0, 3.0]]}), + (f.list_product, f.array_product, {"a": [[1.0, 2.0, 3.0]]}), ], ) def test_array_function_aliases(alias_fn, primary_fn, data): @@ -759,7 +765,69 @@ def test_array_function_aliases(alias_fn, primary_fn, data): ) -@pytest.mark.parametrize("fn", [f.cosine_distance, f.inner_product, f.dot_product]) +@pytest.mark.parametrize( + ("build_expr", "expected"), + [ + pytest.param( + lambda: f.array_add(column("a"), column("b")), + [[11.0, None, 33.0], [], None], + id="array_add", + ), + pytest.param( + lambda: f.array_subtract(column("b"), column("a")), + [[9.0, None, 27.0], [], None], + id="array_subtract", + ), + pytest.param( + lambda: f.array_scale(column("a"), 2), + [[2.0, 4.0, 6.0], [], [None, None]], + id="array_scale_native_scalar", + ), + pytest.param( + lambda: f.array_scale(column("a"), literal(None).cast(pa.float64())), + [None, None, None], + id="array_scale_null_scalar", + ), + pytest.param( + lambda: f.array_sum(column("a")), + [6.0, None, None], + id="array_sum", + ), + pytest.param( + lambda: f.array_avg(column("a")), + [2.0, None, None], + id="array_avg", + ), + pytest.param( + lambda: f.array_product(column("a")), + [6.0, None, None], + id="array_product", + ), + ], +) +def test_array_arithmetic_functions(build_expr, expected): + """Element-wise and reducing array math, including NULL and empty rows.""" + ctx = SessionContext() + df = ctx.from_pydict( + { + "a": [[1.0, 2.0, 3.0], [], [None, None]], + "b": [[10.0, None, 30.0], [], None], + } + ) + result = df.select(build_expr().alias("r")).collect_column("r").to_pylist() + assert result == expected + + +@pytest.mark.parametrize( + "fn", + [ + f.cosine_distance, + f.inner_product, + f.dot_product, + f.array_add, + f.array_subtract, + ], +) def test_array_distance_length_mismatch_raises(fn): """Length-mismatched inputs to vector distance fns should raise at execute.""" ctx = SessionContext() diff --git a/python/tests/test_lambda.py b/python/tests/test_lambda.py index ce37546be..71f7b947a 100644 --- a/python/tests/test_lambda.py +++ b/python/tests/test_lambda.py @@ -97,6 +97,23 @@ def _column(df, expr, name): [[3], [4, 5]], id="list_filter_alias", ), + pytest.param( + lambda: f.array_first(col("a"), lambda v: v > 2), + [3, 4], + id="array_first_callable", + ), + pytest.param( + lambda: f.array_first( + col("a"), f.lambda_(["v"], f.lambda_var("v") > lit(3)) + ), + [None, 4], + id="array_first_explicit_lambda_no_match_is_null", + ), + pytest.param( + lambda: f.list_first(col("a"), lambda v: v > 4), + [None, 5], + id="list_first_alias", + ), ], ) def test_higher_order_function_results(df, build_expr, expected): diff --git a/python/tests/test_spark_functions.py b/python/tests/test_spark_functions.py index 39735a543..5add70732 100644 --- a/python/tests/test_spark_functions.py +++ b/python/tests/test_spark_functions.py @@ -17,6 +17,8 @@ """Tests for the Spark-compatible function bindings.""" +import math + import pyarrow as pa import pytest from datafusion import SessionContext, col, lit @@ -75,12 +77,30 @@ def _dt(*args): (lambda: spark.rint(lit(2.5)), 2.0), (lambda: spark.round(lit(2.5), lit(0)), 3.0), (lambda: spark.negative(lit(3)), -3), + (lambda: spark.atan2(lit(1.0), lit(1.0)), 0.7853981633974483), + (lambda: spark.atan2(0.0, -1.0), 3.141592653589793), + (lambda: spark.hypot(lit(3.0), lit(4.0)), 5.0), + (lambda: spark.hypot(1e200, 1e200), math.hypot(1e200, 1e200)), + (lambda: spark.pow(lit(2), lit(10)), 1024.0), + (lambda: spark.pow(0.0, -1.0), float("inf")), + (lambda: spark.power(2, 3), 8.0), ], ) def test_math(df, expr_factory, expected): assert _val(df, expr_factory()) == expected +@pytest.mark.parametrize( + ("expr_factory", "expected"), + [ + (lambda: spark.monthname(_ts()), "Jan"), + (lambda: spark.weekday(_ts()), 2), + ], +) +def test_monthname_weekday(df, expr_factory, expected): + assert _val(df, expr_factory()) == expected + + def test_factorial(df): # factorial wants Int32; lit(int) is Int64 by default. expr = spark.factorial(lit(pa.scalar(5, type=pa.int32()))) @@ -105,6 +125,13 @@ def test_factorial(df): (lambda: spark.is_valid_utf8(lit("hi")), True), (lambda: spark.concat(lit("a"), lit("b")), "ab"), (lambda: spark.elt(lit(2), lit("a"), lit("b")), "b"), + (lambda: spark.quote(lit("it's")), "'it\\'s'"), + (lambda: spark.concat_ws(",", lit("a"), lit("b")), "a,b"), + (lambda: spark.concat_ws(lit("-"), lit("a"), lit(None), lit("b")), "a-b"), + ( + lambda: spark.concat_ws(",", f.make_array(lit("a"), lit("b")), lit("c")), + "a,b,c", + ), ], ) def test_string(df, expr_factory, expected): From 8027fcb46cf36098b9733ebf2c3ec65f69a09b96 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:23:25 -0400 Subject: [PATCH 02/31] feat: add rand, substring_index, file metadata functions, and distinct kwargs - Add rand and substring_index as aliases of random and substr_index. - Add input_file_name and file_row_index, which report the source file and row offset during a file scan. - Add a distinct argument to bit_and, bit_or, mean, percentile_cont, quantile_cont, and string_agg. distinct goes before filter, matching sum and avg; the upgrade guide covers positional callers. - Fix mean, which passed filter into avg's distinct slot and raised a TypeError whenever filter was given. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/functions.rs | 9 +- docs/source/user-guide/upgrade-guides.md | 20 ++++ python/datafusion/functions/__init__.py | 133 +++++++++++++++++++++-- python/tests/test_aggregation.py | 35 ++++++ python/tests/test_functions.py | 39 +++++++ 5 files changed, 222 insertions(+), 14 deletions(-) diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index c8fd77c51..eb37d64c0 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -622,6 +622,8 @@ expr_fn_vec!(arrow_metadata); expr_fn_vec!(with_metadata); expr_fn!(union_tag, arg1); expr_fn!(random); +expr_fn!(input_file_name); +expr_fn!(file_row_index); #[pyfunction] fn get_field(expr: PyExpr, names: Vec) -> PyExpr { @@ -768,16 +770,17 @@ pub fn approx_percentile_cont_with_weight( } #[pyfunction] -#[pyo3(signature = (sort_expression, percentile, filter=None))] +#[pyo3(signature = (sort_expression, percentile, distinct=None, filter=None))] pub fn percentile_cont( sort_expression: PySortExpr, percentile: f64, + distinct: Option, filter: Option, ) -> PyDataFusionResult { let agg_fn = functions_aggregate::expr_fn::percentile_cont(sort_expression.sort, lit(percentile)); - add_builder_fns_to_aggregate(agg_fn, None, filter, None, None) + add_builder_fns_to_aggregate(agg_fn, distinct, filter, None, None) } // We handle last_value explicitly because the signature expects an order_by @@ -1070,6 +1073,8 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(power))?; m.add_wrapped(wrap_pyfunction!(radians))?; m.add_wrapped(wrap_pyfunction!(random))?; + m.add_wrapped(wrap_pyfunction!(input_file_name))?; + m.add_wrapped(wrap_pyfunction!(file_row_index))?; m.add_wrapped(wrap_pyfunction!(regexp_count))?; m.add_wrapped(wrap_pyfunction!(regexp_instr))?; m.add_wrapped(wrap_pyfunction!(regexp_like))?; diff --git a/docs/source/user-guide/upgrade-guides.md b/docs/source/user-guide/upgrade-guides.md index 0b8e94bb5..5b16683fc 100644 --- a/docs/source/user-guide/upgrade-guides.md +++ b/docs/source/user-guide/upgrade-guides.md @@ -198,6 +198,26 @@ ctx.execute(plan, partitions=0) # before ctx.execute(plan, partition=0) # after ``` +### More aggregate functions accept `distinct` + +{py:func}`~datafusion.functions.bit_and`, +{py:func}`~datafusion.functions.bit_or`, +{py:func}`~datafusion.functions.mean`, +{py:func}`~datafusion.functions.percentile_cont`, +{py:func}`~datafusion.functions.quantile_cont`, and +{py:func}`~datafusion.functions.string_agg` now accept a `distinct` argument. +As with `sum` and `avg` in 54.0.0, `distinct` is inserted *before* `filter`, so +code that passed `filter` (or, for `string_agg`, `order_by`) positionally must +pass it by keyword. + +```python +f.bit_and(column("a"), my_filter) # before +f.bit_and(column("a"), filter=my_filter) # after +``` + +Passing `filter` to `mean` previously raised a `TypeError`, whether passed +positionally or by keyword; it now works. + ### Changes to the `datafusion-python-util` crate Extension libraries written in Rust usually depend on the diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 39a31e00d..f2f2b16ef 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -209,6 +209,7 @@ def _warn_if_expr_for_literal_arg( "exp", "extract", "factorial", + "file_row_index", "find_in_set", "first_value", "flatten", @@ -224,6 +225,7 @@ def _warn_if_expr_for_literal_arg( "in_list", "initcap", "inner_product", + "input_file_name", "instr", "is_nan", "isnan", @@ -334,6 +336,7 @@ def _warn_if_expr_for_literal_arg( "power", "quantile_cont", "radians", + "rand", "random", "range", "rank", @@ -382,6 +385,7 @@ def _warn_if_expr_for_literal_arg( "substr", "substr_index", "substring", + "substring_index", "sum", "tan", "tanh", @@ -2342,6 +2346,15 @@ def substr_index(string: Expr, delimiter: Expr | str, count: Expr | int) -> Expr return Expr(f.substr_index(string.expr, delimiter.expr, count.expr)) +def substring_index(string: Expr, delimiter: Expr | str, count: Expr | int) -> Expr: + """Returns an indexed substring. + + See Also: + This is an alias for :py:func:`substr_index`. + """ + return substr_index(string, delimiter, count) + + def substring(string: Expr, position: Expr | int, length: Expr | int) -> Expr: """Substring from the ``position`` with ``length`` characters. @@ -3487,6 +3500,64 @@ def random() -> Expr: return Expr(f.random()) +def rand() -> Expr: + """Returns a random value in the range ``0.0 <= x < 1.0``. + + See Also: + This is an alias for :py:func:`random`. + """ + return random() + + +def input_file_name() -> Expr: + """Returns the path of the file that produced the current row. + + Only valid inside a scan of a file-backed table; evaluating it anywhere + else raises an error. + + Examples: + >>> import tempfile, os + >>> import pyarrow as pa, pyarrow.parquet as pq + >>> tmp = tempfile.mkdtemp() + >>> path = os.path.join(tmp, "data.parquet") + >>> pq.write_table(pa.table({"a": [1, 2]}), path) + >>> ctx = dfn.SessionContext() + >>> df = ctx.read_parquet(path) + >>> result = df.select(dfn.functions.input_file_name().alias("f")) + >>> result.collect_column("f")[0].as_py().endswith("data.parquet") + True + + See Also: + :py:func:`file_row_index`. + """ + return Expr(f.input_file_name()) + + +def file_row_index() -> Expr: + """Returns the zero-based position of the current row within its source file. + + The index restarts at zero for each file, so rows from different files in one + scan can share a value. Only valid inside a scan of a Parquet table; + evaluating it anywhere else raises an error. + + Examples: + >>> import tempfile, os + >>> import pyarrow as pa, pyarrow.parquet as pq + >>> tmp = tempfile.mkdtemp() + >>> path = os.path.join(tmp, "data.parquet") + >>> pq.write_table(pa.table({"a": [10, 20, 30]}), path) + >>> ctx = dfn.SessionContext() + >>> df = ctx.read_parquet(path).filter(dfn.col("a") > dfn.lit(10)) + >>> result = df.select(dfn.functions.file_row_index().alias("i")) + >>> result.collect_column("i").to_pylist() + [1, 2] + + See Also: + :py:func:`input_file_name`. + """ + return Expr(f.file_row_index()) + + def array_append(array: Expr, element: Expr) -> Expr: """Appends an element to the end of an array. @@ -5335,6 +5406,7 @@ def approx_percentile_cont_with_weight( def percentile_cont( sort_expression: Expr | SortExpr, percentile: float, + distinct: bool = False, filter: Expr | None = None, ) -> Expr: """Computes the exact percentile of input values using continuous interpolation. @@ -5343,11 +5415,12 @@ def percentile_cont( percentile value rather than an approximation. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``order_by``, ``null_treatment``, and ``distinct``. + the options ``order_by`` and ``null_treatment``. Args: sort_expression: Values for which to find the percentile percentile: This must be between 0.0 and 1.0, inclusive + distinct: If True, duplicate values are removed before computing filter: If provided, only compute against rows for which the filter is True Examples: @@ -5367,15 +5440,28 @@ def percentile_cont( ... ).alias("v")]) >>> result.collect_column("v")[0].as_py() 3.5 + + >>> df = ctx.from_pydict({"a": [1.0, 1.0, 1.0, 4.0]}) + >>> result = df.aggregate( + ... [], [dfn.functions.percentile_cont( + ... dfn.col("a"), 0.5, distinct=True, + ... ).alias("v")]) + >>> result.collect_column("v")[0].as_py() + 2.5 """ sort_expr_raw = sort_or_default(sort_expression) filter_raw = filter.expr if filter is not None else None - return Expr(f.percentile_cont(sort_expr_raw, percentile, filter=filter_raw)) + return Expr( + f.percentile_cont( + sort_expr_raw, percentile, distinct=distinct, filter=filter_raw + ) + ) def quantile_cont( sort_expression: Expr | SortExpr, percentile: float, + distinct: bool = False, filter: Expr | None = None, ) -> Expr: """Computes the exact percentile of input values using continuous interpolation. @@ -5383,7 +5469,9 @@ def quantile_cont( See Also: This is an alias for :py:func:`percentile_cont`. """ - return percentile_cont(sort_expression, percentile, filter) + return percentile_cont( + sort_expression, percentile, distinct=distinct, filter=filter + ) def array_agg( @@ -5750,13 +5838,17 @@ def max(expression: Expr, filter: Expr | None = None) -> Expr: return Expr(f.max(expression.expr, filter=filter_raw)) -def mean(expression: Expr, filter: Expr | None = None) -> Expr: +def mean( + expression: Expr, + distinct: bool = False, + filter: Expr | None = None, +) -> Expr: """Returns the average (mean) value of the argument. See Also: This is an alias for :py:func:`avg`. """ - return avg(expression, filter) + return avg(expression, distinct=distinct, filter=filter) def median( @@ -6621,16 +6713,19 @@ def any_value(expression: Expr, filter: Expr | None = None) -> Expr: return Expr(f.any_value(expression.expr, filter=filter_raw)) -def bit_and(expression: Expr, filter: Expr | None = None) -> Expr: +def bit_and( + expression: Expr, distinct: bool = False, filter: Expr | None = None +) -> Expr: """Computes the bitwise AND of the argument. This aggregate function will bitwise compare every value in the input partition. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``order_by``, ``null_treatment``, and ``distinct``. + the options ``order_by`` and ``null_treatment``. Args: expression: Argument to perform bitwise calculation on + distinct: If True, evaluate each unique value of expression only once filter: If provided, only compute against rows for which the filter is True Examples: @@ -6653,19 +6748,22 @@ def bit_and(expression: Expr, filter: Expr | None = None) -> Expr: 5 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.bit_and(expression.expr, filter=filter_raw)) + return Expr(f.bit_and(expression.expr, distinct=distinct, filter=filter_raw)) -def bit_or(expression: Expr, filter: Expr | None = None) -> Expr: +def bit_or( + expression: Expr, distinct: bool = False, filter: Expr | None = None +) -> Expr: """Computes the bitwise OR of the argument. This aggregate function will bitwise compare every value in the input partition. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``order_by``, ``null_treatment``, and ``distinct``. + the options ``order_by`` and ``null_treatment``. Args: expression: Argument to perform bitwise calculation on + distinct: If True, evaluate each unique value of expression only once filter: If provided, only compute against rows for which the filter is True Examples: @@ -6690,7 +6788,7 @@ def bit_or(expression: Expr, filter: Expr | None = None) -> Expr: 6 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.bit_or(expression.expr, filter=filter_raw)) + return Expr(f.bit_or(expression.expr, distinct=distinct, filter=filter_raw)) def bit_xor( @@ -7315,6 +7413,7 @@ def ntile( def string_agg( expression: Expr, delimiter: str, + distinct: bool = False, filter: Expr | None = None, order_by: list[SortKey] | SortKey | None = None, ) -> Expr: @@ -7325,11 +7424,12 @@ def string_agg( their string equivalents. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``distinct`` and ``null_treatment``. + the option ``null_treatment``. Args: expression: Argument to perform bitwise calculation on delimiter: Text to place between each value of expression + distinct: If True, each unique value of expression is included only once filter: If provided, only compute against rows for which the filter is True order_by: Set the ordering of the expression to evaluate. Accepts column names or expressions. @@ -7352,6 +7452,14 @@ def string_agg( ... ).alias("s")]) >>> result.collect_column("s")[0].as_py() 'y,z' + + >>> df = ctx.from_pydict({"a": ["y", "x", "y"]}) + >>> result = df.aggregate( + ... [], [dfn.functions.string_agg( + ... dfn.col("a"), ",", distinct=True, order_by="a", + ... ).alias("s")]) + >>> result.collect_column("s")[0].as_py() + 'x,y' """ order_by_raw = sort_list_to_raw_sort_list(order_by) filter_raw = filter.expr if filter is not None else None @@ -7360,6 +7468,7 @@ def string_agg( f.string_agg( expression.expr, delimiter, + distinct=distinct, filter=filter_raw, order_by=order_by_raw, ) diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index 735fef35d..7847c089c 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -330,12 +330,42 @@ def test_any_value_skips_nulls_per_group(): assert result == {"g": ["x", "y", "z"], "v": [7, 8, None]} +@pytest.mark.parametrize( + ("expr", "expected"), + [ + pytest.param(f.mean(column("v")), 2.0, id="mean"), + pytest.param(f.mean(column("v"), distinct=True), 3.0, id="mean_distinct"), + pytest.param( + f.mean(column("v"), filter=column("v") > lit(1.0)), 5.0, id="mean_filter" + ), + pytest.param(f.percentile_cont(column("v"), 0.5), 1.0, id="percentile_cont"), + pytest.param( + f.percentile_cont(column("v"), 0.5, distinct=True), + 3.0, + id="percentile_cont_distinct", + ), + pytest.param( + f.quantile_cont(column("v"), 0.5, distinct=True), + 3.0, + id="quantile_cont_distinct", + ), + ], +) +def test_distinct_numeric_aggregates(expr, expected): + ctx = SessionContext() + df = ctx.from_pydict({"v": [1.0, 1.0, 1.0, 5.0]}) + result = df.aggregate([], [expr.alias("r")]).collect_column("r")[0].as_py() + assert result == expected + + data_test_bitwise_and_boolean_functions = [ ("any_value_filter", f.any_value(column("a"), filter=column("a") == lit(2)), [2]), ("bit_and", f.bit_and(column("a")), [0]), ("bit_and_filter", f.bit_and(column("a"), filter=column("a") != lit(2)), [1]), ("bit_or", f.bit_or(column("b")), [6]), ("bit_or_filter", f.bit_or(column("b"), filter=column("a") != lit(3)), [4]), + ("bit_and_distinct", f.bit_and(column("b"), distinct=True), [4]), + ("bit_or_distinct", f.bit_or(column("b"), distinct=True), [6]), ("bit_xor", f.bit_xor(column("c")), [4]), ("bit_xor_distinct", f.bit_xor(column("b"), distinct=True), [2]), ("bit_xor_filter", f.bit_xor(column("b"), filter=column("a") != lit(3)), [0]), @@ -491,6 +521,11 @@ def test_first_last_value(df_partitioned, name, expr, result) -> None: f.string_agg(column("a"), ",", order_by=column("b")), "one,three,two,two", ), + ( + "string_agg", + f.string_agg(column("a"), ",", distinct=True, order_by=column("a")), + "one,three,two", + ), ], ) def test_string_agg(name, expr, result) -> None: diff --git a/python/tests/test_functions.py b/python/tests/test_functions.py index 49d9b33fb..ae1dcc33d 100644 --- a/python/tests/test_functions.py +++ b/python/tests/test_functions.py @@ -20,6 +20,7 @@ import numpy as np import pyarrow as pa +import pyarrow.parquet as pq import pytest from datafusion import SessionContext, column, literal from datafusion import functions as f @@ -765,6 +766,44 @@ def test_array_function_aliases(alias_fn, primary_fn, data): ) +@pytest.mark.parametrize( + ("fn", "expected"), + [ + pytest.param(f.input_file_name, ["data.parquet"] * 2, id="input_file_name"), + pytest.param(f.file_row_index, [1, 2], id="file_row_index"), + ], +) +def test_file_metadata_functions(tmp_path, fn, expected): + path = tmp_path / "data.parquet" + pq.write_table(pa.table({"a": [10, 20, 30]}), path) + ctx = SessionContext() + df = ctx.read_parquet(str(path)).filter(column("a") > literal(10)) + result = df.select(fn().alias("r")).collect_column("r").to_pylist() + if fn is f.input_file_name: + result = [r.rsplit("/", 1)[-1] for r in result] + assert result == expected + + +@pytest.mark.parametrize("fn", [f.input_file_name, f.file_row_index]) +def test_file_metadata_functions_outside_scan_raise(fn): + ctx = SessionContext() + df = ctx.from_pydict({"a": [1]}) + with pytest.raises(Exception, match="source dependent"): + df.select(fn().alias("r")).collect() + + +def test_rand_and_substring_index_aliases(): + ctx = SessionContext() + df = ctx.from_pydict({"s": ["a.b.c"]}) + r = df.select( + f.rand().alias("r"), + f.substring_index(column("s"), ".", 2).alias("si"), + f.substr_index(column("s"), ".", 2).alias("sp"), + ).to_pydict() + assert 0.0 <= r["r"][0] < 1.0 + assert r["si"] == r["sp"] == ["a.b"] + + @pytest.mark.parametrize( ("build_expr", "expected"), [ From 97e6031f7333d4ea530485776f70a2bdc461ce5d Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:25:04 -0400 Subject: [PATCH 03/31] feat(spark): add pyspark aliases and optional-length substr - Add getbit, dateadd, datediff, datepart, sha, ceiling, printf, char_length, and character_length as aliases of their Spark primaries. - Add substr, whose len argument is optional as in pyspark. It calls the UDF directly because upstream expr_fn::substring always takes a length. - Rename the spark.last_day parameter from col to date to match pyspark, with an upgrade-guide note for keyword callers. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/spark_functions.rs | 13 +++ docs/source/user-guide/upgrade-guides.md | 11 +++ python/datafusion/functions/spark.py | 117 ++++++++++++++++++++++- python/tests/test_spark_functions.py | 36 +++++++ 4 files changed, 175 insertions(+), 2 deletions(-) diff --git a/crates/core/src/spark_functions.rs b/crates/core/src/spark_functions.rs index 3ff2af250..868fc6934 100644 --- a/crates/core/src/spark_functions.rs +++ b/crates/core/src/spark_functions.rs @@ -254,6 +254,18 @@ spark_udf_vec!(format_string, udf::string::format_string); spark_expr_fn!(quote, arg1); spark_expr_fn!(space, arg1); spark_expr_fn!(substring, str pos length); +/// `substr(str, pos, len=None)`. Upstream `expr_fn::substring` always takes a +/// length, so call the UDF directly to allow the two-argument form. +#[pyfunction] +#[pyo3(signature = (str, pos, len=None))] +fn substr(str: PyExpr, pos: PyExpr, len: Option) -> PyExpr { + let args: Vec = [Some(str), Some(pos), len] + .into_iter() + .flatten() + .map(Into::into) + .collect(); + Expr::ScalarFunction(ScalarFunction::new_udf(udf::string::substring(), args)).into() +} spark_expr_fn!(unbase64, str); spark_expr_fn!(soundex, str); spark_expr_fn!(is_valid_utf8, str); @@ -372,6 +384,7 @@ pub(crate) fn init_module(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_wrapped(wrap_pyfunction!(format_string))?; m.add_wrapped(wrap_pyfunction!(space))?; m.add_wrapped(wrap_pyfunction!(substring))?; + m.add_wrapped(wrap_pyfunction!(substr))?; m.add_wrapped(wrap_pyfunction!(unbase64))?; m.add_wrapped(wrap_pyfunction!(soundex))?; m.add_wrapped(wrap_pyfunction!(is_valid_utf8))?; diff --git a/docs/source/user-guide/upgrade-guides.md b/docs/source/user-guide/upgrade-guides.md index 5b16683fc..a9f4e3d85 100644 --- a/docs/source/user-guide/upgrade-guides.md +++ b/docs/source/user-guide/upgrade-guides.md @@ -218,6 +218,17 @@ f.bit_and(column("a"), filter=my_filter) # after Passing `filter` to `mean` previously raised a `TypeError`, whether passed positionally or by keyword; it now works. +### `spark.last_day` renamed its parameter + +The parameter of {py:func}`datafusion.functions.spark.last_day` is now named +`date`, matching `pyspark.sql.functions.last_day`. Positional calls are +unaffected; update any call passing it by keyword. + +```python +spark.last_day(col=d) # before +spark.last_day(date=d) # after +``` + ### Changes to the `datafusion-python-util` crate Extension libraries written in Rust usually depend on the diff --git a/python/datafusion/functions/spark.py b/python/datafusion/functions/spark.py index 7d84583d2..25ff8190e 100644 --- a/python/datafusion/functions/spark.py +++ b/python/datafusion/functions/spark.py @@ -356,6 +356,15 @@ def bit_get(col: Expr, pos: Expr | str) -> Expr: return Expr(_f.bit_get(col.expr, _to_raw_expr(pos))) +def getbit(col: Expr, pos: Expr | str) -> Expr: + """Spark ``getbit``: returns the bit (0 or 1) at ``pos``. + + See Also: + This is an alias for :py:func:`bit_get`. + """ + return bit_get(col, pos) + + def bit_count(col: Expr) -> Expr: """Spark ``bit_count``: number of bits set in the integer's binary form. @@ -539,6 +548,15 @@ def date_add(start: Expr, days: Expr | int) -> Expr: return Expr(_f.date_add(start.expr, _coerce_i32(days).expr)) +def dateadd(start: Expr, days: Expr | int) -> Expr: + """Spark ``dateadd``: date + N days. + + See Also: + This is an alias for :py:func:`date_add`. + """ + return date_add(start, days) + + def date_sub(start: Expr, days: Expr | int) -> Expr: """Spark ``date_sub``: date - N days. @@ -627,7 +645,7 @@ def second(col: Expr) -> Expr: return Expr(_f.second(col.expr)) -def last_day(col: Expr) -> Expr: +def last_day(date: Expr) -> Expr: """Spark ``last_day``: last day of the month containing the date. Examples: @@ -640,7 +658,7 @@ def last_day(col: Expr) -> Expr: >>> r.collect_column("v")[0].as_py() datetime.date(2020, 1, 31) """ - return Expr(_f.last_day(col.expr)) + return Expr(_f.last_day(date.expr)) def make_dt_interval( @@ -754,6 +772,15 @@ def date_diff(end: Expr, start: Expr) -> Expr: return Expr(_f.date_diff(end.expr, start.expr)) +def datediff(end: Expr, start: Expr) -> Expr: + """Spark ``datediff``: number of days from ``start`` to ``end``. + + See Also: + This is an alias for :py:func:`date_diff`. + """ + return date_diff(end, start) + + def date_trunc(format: Expr | str, timestamp: Expr) -> Expr: """Spark ``date_trunc``: truncate timestamp to unit ``fmt``. @@ -832,6 +859,15 @@ def date_part(field: Expr | str, source: Expr) -> Expr: return Expr(_f.date_part(coerce_to_expr(field).expr, source.expr)) +def datepart(field: Expr | str, source: Expr) -> Expr: + """Spark ``datepart``: extract ``field`` from a date/time/timestamp. + + See Also: + This is an alias for :py:func:`date_part`. + """ + return date_part(field, source) + + def from_utc_timestamp(timestamp: Expr, tz: Expr | str) -> Expr: """Spark ``from_utc_timestamp``: interpret ``ts`` as UTC, convert to ``tz``. @@ -991,6 +1027,15 @@ def sha1(col: Expr) -> Expr: return Expr(_f.sha1(col.expr)) +def sha(col: Expr) -> Expr: + """Spark ``sha``: SHA-1 hash as a hex string. + + See Also: + This is an alias for :py:func:`sha1`. + """ + return sha1(col) + + def sha2(col: Expr, numBits: Expr | int) -> Expr: # noqa: N803 """Spark ``sha2``: SHA-2 family hash (224, 256, 384, 512). Bit length 0 = 256. @@ -1175,6 +1220,15 @@ def ceil(col: Expr) -> Expr: return Expr(_f.ceil(col.expr)) +def ceiling(col: Expr) -> Expr: + """Spark ``ceiling``: smallest integer ≥ arg. + + See Also: + This is an alias for :py:func:`ceil`. + """ + return ceil(col) + + def expm1(col: Expr) -> Expr: """Spark ``expm1``: exp(arg) - 1. @@ -1562,6 +1616,24 @@ def length(col: Expr) -> Expr: return Expr(_f.length(col.expr)) +def character_length(col: Expr) -> Expr: + """Spark ``character_length``: character length of a string, or bytes of binary. + + See Also: + This is an alias for :py:func:`length`. + """ + return length(col) + + +def char_length(col: Expr) -> Expr: + """Spark ``char_length``: character length of a string, or bytes of binary. + + See Also: + This is an alias for :py:func:`length`. + """ + return length(col) + + def like( str: Expr, pattern: Expr | str, @@ -1639,6 +1711,15 @@ def format_string(format: str | Expr, *cols: Expr) -> Expr: return Expr(_f.format_string(fmt_expr.expr, *[c.expr for c in cols])) +def printf(format: str | Expr, *cols: Expr) -> Expr: + """Spark ``printf``: printf-style format string. + + See Also: + This is an alias for :py:func:`format_string`. + """ + return format_string(format, *cols) + + def space(col: Expr | int) -> Expr: """Spark ``space``: string of n spaces. @@ -1673,6 +1754,28 @@ def substring(str: Expr, pos: Expr | int, len: Expr | int) -> Expr: ) +def substr(str: Expr, pos: Expr | int, len: Expr | int | None = None) -> Expr: + """Spark ``substr``: 1-indexed substring, to the end when ``len`` is omitted. + + Same as :py:func:`substring` except that ``len`` is optional. ``pos`` and + ``len`` accept native ``int`` values or :class:`Expr`. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"x": [1]}) + >>> r = df.select(dfn.functions.spark.substr(dfn.lit("hello"), 2).alias("v")) + >>> r.collect_column("v")[0].as_py() + 'ello' + + >>> r = df.select( + ... dfn.functions.spark.substr(dfn.lit("hello"), 2, len=3).alias("v")) + >>> r.collect_column("v")[0].as_py() + 'ell' + """ + len_raw = coerce_to_expr(len).expr if len is not None else None + return Expr(_f.substr(str.expr, coerce_to_expr(pos).expr, len_raw)) + + def unbase64(col: Expr) -> Expr: """Spark ``unbase64``: decode a base64 string to binary. @@ -1867,7 +1970,10 @@ def url_encode(str: Expr) -> Expr: "bitmap_count", "bitwise_not", "ceil", + "ceiling", "char", + "char_length", + "character_length", "collect_list", "collect_set", "concat", @@ -1880,12 +1986,16 @@ def url_encode(str: Expr) -> Expr: "date_part", "date_sub", "date_trunc", + "dateadd", + "datediff", + "datepart", "elt", "expm1", "factorial", "floor", "format_string", "from_utc_timestamp", + "getbit", "hex", "hour", "hypot", @@ -1914,11 +2024,13 @@ def url_encode(str: Expr) -> Expr: "pmod", "pow", "power", + "printf", "quote", "rint", "round", "sec", "second", + "sha", "sha1", "sha2", "shiftleft", @@ -1932,6 +2044,7 @@ def url_encode(str: Expr) -> Expr: "space", "spark_cast", "str_to_map", + "substr", "substring", "time_trunc", "to_utc_timestamp", diff --git a/python/tests/test_spark_functions.py b/python/tests/test_spark_functions.py index 5add70732..67487d0da 100644 --- a/python/tests/test_spark_functions.py +++ b/python/tests/test_spark_functions.py @@ -500,3 +500,39 @@ def test_sql_concat_semantics_override(): ctx2.sql("SELECT concat('a', NULL, 'b') AS c").collect_column("c")[0].as_py() ) assert spark_out is None + + +@pytest.mark.parametrize( + ("alias_fn", "primary_fn", "args"), + [ + (spark.getbit, spark.bit_get, lambda: (lit(5), lit(0))), + (spark.dateadd, spark.date_add, lambda: (_ts().cast(pa.date32()), 3)), + ( + spark.datediff, + spark.date_diff, + lambda: (_ts().cast(pa.date32()), lit("2020-01-01").cast(pa.date32())), + ), + (spark.datepart, spark.date_part, lambda: ("YEAR", _ts())), + (spark.sha, spark.sha1, lambda: (lit("abc"),)), + (spark.ceiling, spark.ceil, lambda: (lit(1.2),)), + (spark.printf, spark.format_string, lambda: ("%d-%s", lit(42), lit("hi"))), + (spark.char_length, spark.length, lambda: (lit("hello"),)), + (spark.character_length, spark.length, lambda: (lit("hello"),)), + (spark.power, spark.pow, lambda: (lit(2), lit(3))), + (spark.substr, spark.substring, lambda: (lit("hello"), 2, 3)), + ], +) +def test_aliases_match_primary(df, alias_fn, primary_fn, args): + assert _val(df, alias_fn(*args())) == _val(df, primary_fn(*args())) + + +def test_substr_without_len(df): + assert _val(df, spark.substr(lit("hello"), 2)) == "ello" + assert _val(df, spark.substr(lit("hello"), -3)) == "llo" + + +def test_last_day_date_keyword(df): + import datetime as dt + + d = lit(pa.scalar(dt.date(2024, 2, 10), type=pa.date32())) + assert _val(df, spark.last_day(date=d)) == dt.date(2024, 2, 29) From e8c219286f07aaa740917490abbf4346a49a1357 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:26:55 -0400 Subject: [PATCH 04/31] feat: accept optional arguments upstream supports on trim, array_to_string, substr - btrim, ltrim, rtrim, and trim take an optional characters argument naming the set to strip. - array_to_string and its aliases take an optional null_string that is written in place of NULL elements. - substr takes an optional length. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/functions.rs | 15 ++- python/datafusion/functions/__init__.py | 133 +++++++++++++++++++----- python/tests/test_functions.py | 43 ++++++++ 3 files changed, 166 insertions(+), 25 deletions(-) diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index eb37d64c0..a88b8de72 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -117,6 +117,20 @@ fn string_to_array(string: PyExpr, delimiter: PyExpr, null_string: Option) -> PyExpr { + let mut args = vec![array.into(), delimiter.into()]; + if let Some(null_string) = null_string { + args.push(null_string.into()); + } + Expr::ScalarFunction(datafusion::logical_expr::expr::ScalarFunction::new_udf( + datafusion::functions_nested::string::array_to_string_udf(), + args, + )) + .into() +} + #[pyfunction] #[pyo3(signature = (start, stop, step=None))] fn gen_series(start: PyExpr, stop: PyExpr, step: Option) -> PyExpr { @@ -646,7 +660,6 @@ fn version() -> PyExpr { // Array Functions array_fn!(array_append, array element); -array_fn!(array_to_string, array delimiter); array_fn!(array_dims, array); array_fn!(array_distinct, array); array_fn!(array_element, array element); diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index f2f2b16ef..9ccacdccb 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -486,46 +486,71 @@ def decode(expr: Expr, encoding: Expr | str) -> Expr: return Expr(f.decode(expr.expr, encoding.expr)) -def array_to_string(expr: Expr, delimiter: Expr | str) -> Expr: +def array_to_string( + expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None +) -> Expr: """Converts each element to its text representation. + NULL elements are omitted unless ``null_string`` is given, in which case it + is written in their place. + Examples: >>> ctx = dfn.SessionContext() - >>> df = ctx.from_pydict({"a": [[1, 2, 3]]}) + >>> df = ctx.from_pydict({"a": [[1, None, 3]]}) >>> result = df.select( ... dfn.functions.array_to_string(dfn.col("a"), ",").alias("s")) >>> result.collect_column("s")[0].as_py() - '1,2,3' + '1,3' + + >>> result = df.select( + ... dfn.functions.array_to_string( + ... dfn.col("a"), ",", null_string="*" + ... ).alias("s")) + >>> result.collect_column("s")[0].as_py() + '1,*,3' """ delimiter = coerce_to_expr(delimiter) - return Expr(f.array_to_string(expr.expr, delimiter.expr.cast(pa.string()))) + null_string = coerce_to_expr_or_none(null_string) + return Expr( + f.array_to_string( + expr.expr, + delimiter.expr.cast(pa.string()), + null_string.expr if null_string is not None else None, + ) + ) -def array_join(expr: Expr, delimiter: Expr | str) -> Expr: +def array_join( + expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None +) -> Expr: """Converts each element to its text representation. See Also: This is an alias for :py:func:`array_to_string`. """ - return array_to_string(expr, delimiter) + return array_to_string(expr, delimiter, null_string=null_string) -def list_to_string(expr: Expr, delimiter: Expr | str) -> Expr: +def list_to_string( + expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None +) -> Expr: """Converts each element to its text representation. See Also: This is an alias for :py:func:`array_to_string`. """ - return array_to_string(expr, delimiter) + return array_to_string(expr, delimiter, null_string=null_string) -def list_join(expr: Expr, delimiter: Expr | str) -> Expr: +def list_join( + expr: Expr, delimiter: Expr | str, null_string: Expr | str | None = None +) -> Expr: """Converts each element to its text representation. See Also: This is an alias for :py:func:`array_to_string`. """ - return array_to_string(expr, delimiter) + return array_to_string(expr, delimiter, null_string=null_string) def lambda_var(name: str) -> Expr: @@ -1128,8 +1153,8 @@ def bit_length(arg: Expr) -> Expr: return Expr(f.bit_length(arg.expr)) -def btrim(arg: Expr) -> Expr: - """Removes all characters, spaces by default, from both sides of a string. +def btrim(arg: Expr, characters: Expr | str | None = None) -> Expr: + """Removes ``characters``, spaces by default, from both sides of a string. Examples: >>> ctx = dfn.SessionContext() @@ -1137,8 +1162,20 @@ def btrim(arg: Expr) -> Expr: >>> trim_df = df.select(dfn.functions.btrim(dfn.col("a")).alias("trimmed")) >>> trim_df.collect_column("trimmed")[0].as_py() 'a' + + Trim a different set of characters: + + >>> df = ctx.from_pydict({"a": ["xxaxx"]}) + >>> trim_df = df.select( + ... dfn.functions.btrim(dfn.col("a"), characters="x").alias("trimmed") + ... ) + >>> trim_df.collect_column("trimmed")[0].as_py() + 'a' """ - return Expr(f.btrim(arg.expr)) + args = [arg.expr] + if characters is not None: + args.append(coerce_to_expr(characters).expr) + return Expr(f.btrim(*args)) def cbrt(arg: Expr) -> Expr: @@ -1616,8 +1653,8 @@ def lpad(string: Expr, count: Expr | int, characters: Expr | str | None = None) return Expr(f.lpad(string.expr, count.expr, characters.expr)) -def ltrim(arg: Expr) -> Expr: - """Removes all characters, spaces by default, from the beginning of a string. +def ltrim(arg: Expr, characters: Expr | str | None = None) -> Expr: + """Removes ``characters``, spaces by default, from the beginning of a string. Examples: >>> ctx = dfn.SessionContext() @@ -1625,8 +1662,20 @@ def ltrim(arg: Expr) -> Expr: >>> trim_df = df.select(dfn.functions.ltrim(dfn.col("a")).alias("trimmed")) >>> trim_df.collect_column("trimmed")[0].as_py() 'a ' + + Trim a different set of characters: + + >>> df = ctx.from_pydict({"a": ["xxaxx"]}) + >>> trim_df = df.select( + ... dfn.functions.ltrim(dfn.col("a"), characters="x").alias("trimmed") + ... ) + >>> trim_df.collect_column("trimmed")[0].as_py() + 'axx' """ - return Expr(f.ltrim(arg.expr)) + args = [arg.expr] + if characters is not None: + args.append(coerce_to_expr(characters).expr) + return Expr(f.ltrim(*args)) def md5(arg: Expr) -> Expr: @@ -2138,8 +2187,8 @@ def rpad(string: Expr, count: Expr | int, characters: Expr | str | None = None) return Expr(f.rpad(string.expr, count.expr, characters.expr)) -def rtrim(arg: Expr) -> Expr: - """Removes all characters, spaces by default, from the end of a string. +def rtrim(arg: Expr, characters: Expr | str | None = None) -> Expr: + """Removes ``characters``, spaces by default, from the end of a string. Examples: >>> ctx = dfn.SessionContext() @@ -2147,8 +2196,20 @@ def rtrim(arg: Expr) -> Expr: >>> trim_df = df.select(dfn.functions.rtrim(dfn.col("a")).alias("trimmed")) >>> trim_df.collect_column("trimmed")[0].as_py() ' a' + + Trim a different set of characters: + + >>> df = ctx.from_pydict({"a": ["xxaxx"]}) + >>> trim_df = df.select( + ... dfn.functions.rtrim(dfn.col("a"), characters="x").alias("trimmed") + ... ) + >>> trim_df.collect_column("trimmed")[0].as_py() + 'xxa' """ - return Expr(f.rtrim(arg.expr)) + args = [arg.expr] + if characters is not None: + args.append(coerce_to_expr(characters).expr) + return Expr(f.rtrim(*args)) def sha224(arg: Expr) -> Expr: @@ -2312,8 +2373,10 @@ def strpos(string: Expr, substring: Expr | str) -> Expr: return Expr(f.strpos(string.expr, substring.expr)) -def substr(string: Expr, position: Expr | int) -> Expr: - """Substring from the ``position`` to the end. +def substr( + string: Expr, position: Expr | int, length: Expr | int | None = None +) -> Expr: + """Substring from the ``position``, to the end or for ``length`` characters. Examples: >>> ctx = dfn.SessionContext() @@ -2322,7 +2385,17 @@ def substr(string: Expr, position: Expr | int) -> Expr: ... dfn.functions.substr(dfn.col("a"), 3).alias("s")) >>> result.collect_column("s")[0].as_py() 'llo' + + >>> result = df.select( + ... dfn.functions.substr(dfn.col("a"), 2, length=3).alias("s")) + >>> result.collect_column("s")[0].as_py() + 'ell' + + See Also: + :py:func:`substring`. """ + if length is not None: + return substring(string, position, length) position = coerce_to_expr(position) return Expr(f.substr(string.expr, position.expr)) @@ -2958,8 +3031,8 @@ def translate(string: Expr, from_val: Expr | str, to_val: Expr | str) -> Expr: return Expr(f.translate(string.expr, from_val.expr, to_val.expr)) -def trim(arg: Expr) -> Expr: - """Removes all characters, spaces by default, from both sides of a string. +def trim(arg: Expr, characters: Expr | str | None = None) -> Expr: + """Removes ``characters``, spaces by default, from both sides of a string. Examples: >>> ctx = dfn.SessionContext() @@ -2967,8 +3040,20 @@ def trim(arg: Expr) -> Expr: >>> result = df.select(dfn.functions.trim(dfn.col("a")).alias("t")) >>> result.collect_column("t")[0].as_py() 'hello' + + Trim a different set of characters: + + >>> df = ctx.from_pydict({"a": ["xxhelloxx"]}) + >>> result = df.select( + ... dfn.functions.trim(dfn.col("a"), characters="x").alias("t") + ... ) + >>> result.collect_column("t")[0].as_py() + 'hello' """ - return Expr(f.trim(arg.expr)) + args = [arg.expr] + if characters is not None: + args.append(coerce_to_expr(characters).expr) + return Expr(f.trim(*args)) def trunc(num: Expr, precision: Expr | int | None = None) -> Expr: diff --git a/python/tests/test_functions.py b/python/tests/test_functions.py index ae1dcc33d..f5f1921f8 100644 --- a/python/tests/test_functions.py +++ b/python/tests/test_functions.py @@ -2547,3 +2547,46 @@ def test_backward_compat_with_lit(self): f.split_part(column("a"), literal(","), literal(2)).alias("s") ).collect() assert result[0].column(0)[0].as_py() == "b" + + +@pytest.mark.parametrize( + ("fn", "expected"), + [ + (f.btrim, "hi"), + (f.trim, "hi"), + (f.ltrim, "hixyx"), + (f.rtrim, "xyxhi"), + ], +) +def test_trim_characters(fn, expected): + ctx = SessionContext() + df = ctx.from_pydict({"a": ["xyxhixyx"]}) + assert df.select(fn(column("a"), characters="xy").alias("r")).collect_column( + "r" + ).to_pylist() == [expected] + assert df.select( + fn(column("a"), characters=literal("xy")).alias("r") + ).collect_column("r").to_pylist() == [expected] + + +@pytest.mark.parametrize( + "fn", [f.array_to_string, f.array_join, f.list_to_string, f.list_join] +) +def test_array_to_string_null_string(fn): + ctx = SessionContext() + df = ctx.from_pydict({"a": [[1, None, 3]]}) + without = df.select(fn(column("a"), "-").alias("r")).collect_column("r") + with_null = df.select(fn(column("a"), "-", null_string="NA").alias("r")) + assert without.to_pylist() == ["1-3"] + assert with_null.collect_column("r").to_pylist() == ["1-NA-3"] + + +def test_substr_length(): + ctx = SessionContext() + df = ctx.from_pydict({"a": ["hello"]}) + r = df.select( + f.substr(column("a"), 2).alias("tail"), + f.substr(column("a"), 2, length=3).alias("mid"), + f.substr(column("a"), 2, length=literal(3)).alias("mid_expr"), + ).to_pydict() + assert r == {"tail": ["ello"], "mid": ["ell"], "mid_expr": ["ell"]} From 2747573aa5816fb3cf876df8230759c5f11af72c Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:28:26 -0400 Subject: [PATCH 05/31] feat: add DataFrame.fill_nan Replace NaN in floating-point columns, optionally limited to a subset. Mirrors fill_null and wraps upstream DataFrame::fill_nan. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/dataframe.rs | 20 ++++++++++++++++ python/datafusion/dataframe.py | 28 ++++++++++++++++++++++ python/tests/test_dataframe.py | 44 ++++++++++++++++++++++++++++++++++ 3 files changed, 92 insertions(+) diff --git a/crates/core/src/dataframe.rs b/crates/core/src/dataframe.rs index b1f305551..61b480645 100644 --- a/crates/core/src/dataframe.rs +++ b/crates/core/src/dataframe.rs @@ -1320,6 +1320,26 @@ impl PyDataFrame { let df = self.df.as_ref().fill_null(&scalar_value.0, &cols)?; Ok(Self::new(df)) } + + /// Fill NaN values with a specified value for specific floating-point columns + #[pyo3(signature = (value, columns=None))] + fn fill_nan( + &self, + value: Py, + columns: Option>, + py: Python, + ) -> PyDataFusionResult { + let scalar_value: PyScalarValue = value.extract(py)?; + + let cols = match columns { + Some(col_names) => col_names.iter().map(|c| c.to_string()).collect(), + None => Vec::new(), // Empty vector means fill NaN for all columns + }; + + let cols = cols.iter().map(String::as_str).collect::>(); + let df = self.df.as_ref().fill_nan(&scalar_value.0, &cols)?; + Ok(Self::new(df)) + } } #[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)] diff --git a/python/datafusion/dataframe.py b/python/datafusion/dataframe.py index de00ff474..5d46b7fc0 100644 --- a/python/datafusion/dataframe.py +++ b/python/datafusion/dataframe.py @@ -1875,6 +1875,34 @@ def fill_null(self, value: Any, subset: list[str] | None = None) -> DataFrame: """ return DataFrame(self.df.fill_null(value, subset)) + def fill_nan(self, value: float, subset: list[str] | None = None) -> DataFrame: + """Fill NaN values in floating-point columns with a value. + + Only floating-point columns are changed; others are kept unchanged, as is + any column ``value`` cannot be cast to. NaN is distinct from null, which + :py:meth:`fill_null` handles. + + Args: + value: Value to replace NaN with. Will be cast to match column type. + subset: Optional list of column names to fill. If None, fills all + floating-point columns. + + Returns: + DataFrame with NaN values replaced. + + Examples: + >>> from datafusion import SessionContext + >>> ctx = SessionContext() + >>> nan = float("nan") + >>> df = ctx.from_pydict({"a": [1.0, nan, None], "b": [nan, 2.0, 3.0]}) + >>> df.fill_nan(0.0).to_pydict() + {'a': [1.0, 0.0, None], 'b': [0.0, 2.0, 3.0]} + + >>> df.fill_nan(0.0, subset=["a"]).collect_column("b")[0].as_py() + nan + """ + return DataFrame(self.df.fill_nan(value, subset)) + class InsertOp(Enum): """Insert operation mode. diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index bb21a3974..56023e460 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -3450,6 +3450,50 @@ def test_fill_null_all_null_column(ctx): assert result.column(1).to_pylist() == ["filled", "filled", "filled"] +def _nan_df(ctx): + nan = float("nan") + batch = pa.RecordBatch.from_arrays( + [ + pa.array([1.0, nan, None], type=pa.float64()), + pa.array([nan, 2.0, 3.0], type=pa.float32()), + pa.array([1, 2, 3]), + pa.array(["x", "nan", None]), + ], + names=["f64", "f32", "i", "s"], + ) + return ctx.create_dataframe([[batch]]) + + +def _is_nan(v): + return v is not None and v != v # noqa: PLR0124 + + +def test_fill_nan_all_columns(ctx): + result = _nan_df(ctx).fill_nan(0.0).to_pydict() + # NaN replaced in both float widths; null is not NaN and stays null. + assert result["f64"] == [1.0, 0.0, None] + assert result["f32"] == [0.0, 2.0, 3.0] + # Non-float columns are untouched. + assert result["i"] == [1, 2, 3] + assert result["s"] == ["x", "nan", None] + + +def test_fill_nan_subset(ctx): + result = _nan_df(ctx).fill_nan(-1.0, subset=["f32"]).to_pydict() + assert result["f32"] == [-1.0, 2.0, 3.0] + assert _is_nan(result["f64"][1]) + + +def test_fill_nan_preserves_schema(ctx): + df = _nan_df(ctx) + assert df.fill_nan(0.0).schema() == df.schema() + + +def test_fill_nan_unknown_column_raises(ctx): + with pytest.raises(Exception, match="missing"): + _nan_df(ctx).fill_nan(0.0, subset=["missing"]).collect() + + _slow_udf_started = threading.Event() From 24764bddce0bfb9011886079bb228d230bb60b30 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:42:21 -0400 Subject: [PATCH 06/31] feat: add null_treatment to lead and lag IGNORE_NULLS skips null values when counting shift_offset rows, matching SQL LEAD/LAG ... IGNORE NULLS. The argument is appended after order_by, so existing calls are unaffected. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/functions.rs | 10 +++++---- python/datafusion/functions/__init__.py | 28 +++++++++++++++++++++++++ python/tests/test_dataframe.py | 21 +++++++++++++++++++ 3 files changed, 55 insertions(+), 4 deletions(-) diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index a88b8de72..af1de455f 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -891,33 +891,35 @@ pub(crate) fn add_builder_fns_to_window( } #[pyfunction] -#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None))] +#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None, null_treatment=None))] pub fn lead( arg: PyExpr, shift_offset: i64, default_value: Option, partition_by: Option>, order_by: Option>, + null_treatment: Option, ) -> PyDataFusionResult { let default_value = default_value.map(|v| v.into()); let window_fn = functions_window::expr_fn::lead(arg.expr, Some(shift_offset), default_value); - add_builder_fns_to_window(window_fn, partition_by, None, order_by, None) + add_builder_fns_to_window(window_fn, partition_by, None, order_by, null_treatment) } #[pyfunction] -#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None))] +#[pyo3(signature = (arg, shift_offset, default_value=None, partition_by=None, order_by=None, null_treatment=None))] pub fn lag( arg: PyExpr, shift_offset: i64, default_value: Option, partition_by: Option>, order_by: Option>, + null_treatment: Option, ) -> PyDataFusionResult { let default_value = default_value.map(|v| v.into()); let window_fn = functions_window::expr_fn::lag(arg.expr, Some(shift_offset), default_value); - add_builder_fns_to_window(window_fn, partition_by, None, order_by, None) + add_builder_fns_to_window(window_fn, partition_by, None, order_by, null_treatment) } #[pyfunction] diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 9ccacdccb..7b4f24a62 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -7000,6 +7000,7 @@ def lead( default_value: Any | None = None, partition_by: list[Expr] | Expr | None = None, order_by: list[SortKey] | SortKey | None = None, + null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS, ) -> Expr: """Create a lead window function. @@ -7030,6 +7031,8 @@ def lead( partition_by: Expressions to partition the window frame on. order_by: Set ordering within the window frame. Accepts column names or expressions. + null_treatment: Set to ``IGNORE_NULLS`` to skip null values when + counting ``shift_offset`` rows. Examples: >>> ctx = dfn.SessionContext() @@ -7052,6 +7055,16 @@ def lead( ... ).alias("lead")) >>> result.sort(dfn.col("g"), dfn.col("v")).collect_column("lead").to_pylist() [2, 0, 0] + + >>> df = ctx.from_pydict({"i": [1, 2, 3, 4], "v": [1, None, None, 4]}) + >>> result = df.select( + ... dfn.col("i"), + ... dfn.functions.lead( + ... dfn.col("v"), order_by="i", + ... null_treatment=dfn.common.NullTreatment.IGNORE_NULLS, + ... ).alias("lead")) + >>> result.sort(dfn.col("i")).collect_column("lead").to_pylist() + [4, 4, 4, None] """ if not isinstance(default_value, pa.Scalar) and default_value is not None: default_value = pa.scalar(default_value) @@ -7066,6 +7079,7 @@ def lead( default_value, partition_by=partition_by_raw, order_by=order_by_raw, + null_treatment=null_treatment.value, ) ) @@ -7076,6 +7090,7 @@ def lag( default_value: Any | None = None, partition_by: list[Expr] | Expr | None = None, order_by: list[SortKey] | SortKey | None = None, + null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS, ) -> Expr: """Create a lag window function. @@ -7103,6 +7118,8 @@ def lag( partition_by: Expressions to partition the window frame on. order_by: Set ordering within the window frame. Accepts column names or expressions. + null_treatment: Set to ``IGNORE_NULLS`` to skip null values when + counting ``shift_offset`` rows. Examples: >>> ctx = dfn.SessionContext() @@ -7125,6 +7142,16 @@ def lag( ... ).alias("lag")) >>> result.sort(dfn.col("g"), dfn.col("v")).collect_column("lag").to_pylist() [0, 1, 0] + + >>> df = ctx.from_pydict({"i": [1, 2, 3, 4], "v": [1, None, None, 4]}) + >>> result = df.select( + ... dfn.col("i"), + ... dfn.functions.lag( + ... dfn.col("v"), order_by="i", + ... null_treatment=dfn.common.NullTreatment.IGNORE_NULLS, + ... ).alias("lag")) + >>> result.sort(dfn.col("i")).collect_column("lag").to_pylist() + [None, 1, 1, 1] """ if not isinstance(default_value, pa.Scalar): default_value = pa.scalar(default_value) @@ -7139,6 +7166,7 @@ def lag( default_value, partition_by=partition_by_raw, order_by=order_by_raw, + null_treatment=null_treatment.value, ) ) diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index 56023e460..e44afce98 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -46,6 +46,7 @@ from datafusion import ( functions as f, ) +from datafusion.common import NullTreatment from datafusion.dataframe import DataFrameWriteOptions from datafusion.dataframe_formatter import ( DataFrameHtmlFormatter, @@ -1077,6 +1078,26 @@ def test_distinct(): ), [-1, -1, None, 7, -1, -1, None], ), + ( + "lead_ignore_nulls", + f.lead( + column("b"), + order_by=column("a"), + partition_by=column("c"), + null_treatment=NullTreatment.IGNORE_NULLS, + ), + [7, 7, 8, None, 9, 9, None], + ), + ( + "lag_ignore_nulls", + f.lag( + column("b"), + order_by=column("a"), + partition_by=column("c"), + null_treatment=NullTreatment.IGNORE_NULLS, + ), + [None, 7, 7, 7, None, 9, 9], + ), ( "first_value", f.first_value(column("a")).over( From b035c97fab0fb519923f352bae9e82673f58d542 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:43:54 -0400 Subject: [PATCH 07/31] feat: add show_statistics, analyze_level, and analyze_categories to explain Expose the remaining upstream ExplainOption fields as keywords on DataFrame.explain. Each defaults to None, which falls back to the matching datafusion.explain.* session setting, so existing calls are unaffected. New ExplainAnalyzeLevel and ExplainMetricCategory enums sit beside ExplainFormat. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/dataframe.rs | 29 +++++++++++++++-- python/datafusion/dataframe.py | 51 +++++++++++++++++++++++++++++- python/tests/test_dataframe.py | 57 +++++++++++++++++++++++++++++++++- 3 files changed, 133 insertions(+), 4 deletions(-) diff --git a/crates/core/src/dataframe.rs b/crates/core/src/dataframe.rs index 61b480645..c8d907bde 100644 --- a/crates/core/src/dataframe.rs +++ b/crates/core/src/dataframe.rs @@ -844,13 +844,24 @@ impl PyDataFrame { } /// Print the query plan - #[pyo3(signature = (verbose=false, analyze=false, format=None))] + #[pyo3(signature = ( + verbose=false, + analyze=false, + format=None, + show_statistics=None, + analyze_level=None, + analyze_categories=None + ))] + #[allow(clippy::too_many_arguments)] fn explain( &self, py: Python, verbose: bool, analyze: bool, format: Option<&str>, + show_statistics: Option, + analyze_level: Option<&str>, + analyze_categories: Option>, ) -> PyDataFusionResult<()> { let explain_format = match format { Some(f) => f @@ -860,10 +871,24 @@ impl PyDataFrame { })?, None => datafusion::common::format::ExplainFormat::Indent, }; + let analyze_level = analyze_level + .map(|l| l.parse::()) + .transpose()?; + let analyze_categories = analyze_categories + .map(|cats| { + cats.iter() + .map(|c| c.parse::()) + .collect::>>() + .map(datafusion::common::format::ExplainAnalyzeCategories::Only) + }) + .transpose()?; let opts = datafusion::logical_expr::ExplainOption::default() .with_verbose(verbose) .with_analyze(analyze) - .with_format(explain_format); + .with_format(explain_format) + .with_show_statistics(show_statistics) + .with_analyze_level(analyze_level) + .with_analyze_categories(analyze_categories); let df = self.df.as_ref().clone().explain_with_options(opts)?; print_dataframe(py, df) } diff --git a/python/datafusion/dataframe.py b/python/datafusion/dataframe.py index 5d46b7fc0..b724935a7 100644 --- a/python/datafusion/dataframe.py +++ b/python/datafusion/dataframe.py @@ -108,6 +108,32 @@ class ExplainFormat(Enum): """Graphviz DOT format for graph rendering.""" +class ExplainAnalyzeLevel(Enum): + """Which metrics :py:meth:`DataFrame.explain` reports when ``analyze=True``.""" + + SUMMARY = "summary" + """Common metrics for finding which operator is slow.""" + + DEV = "dev" + """All metrics, including those for deep operator-level introspection.""" + + +class ExplainMetricCategory(Enum): + """Category of metric reported by :py:meth:`DataFrame.explain` with ``analyze``.""" + + ROWS = "rows" + """Row counts, such as ``output_rows``.""" + + BYTES = "bytes" + """Byte sizes, such as ``output_bytes``.""" + + TIMING = "timing" + """Elapsed times, such as ``elapsed_compute``.""" + + UNCATEGORIZED = "uncategorized" + """Metrics that declare no category.""" + + # excerpt from deltalake # https://github.com/apache/datafusion-python/pull/981#discussion_r1905619163 class Compression(Enum): @@ -1207,6 +1233,9 @@ def explain( verbose: bool = False, analyze: bool = False, format: ExplainFormat | None = None, + show_statistics: bool | None = None, + analyze_level: ExplainAnalyzeLevel | None = None, + analyze_categories: Iterable[ExplainMetricCategory] | None = None, ) -> None: """Print an explanation of the DataFrame's plan so far. @@ -1217,6 +1246,13 @@ def explain( analyze: If ``True``, the plan will run and metrics reported. format: Output format for the plan. Defaults to :py:attr:`ExplainFormat.INDENT`. + show_statistics: If ``True``, include each operator's statistics. + ``None`` uses the ``datafusion.explain.show_statistics`` setting. + analyze_level: Which metrics to report with ``analyze``. ``None`` + uses the ``datafusion.explain.analyze_level`` setting. + analyze_categories: Report only metrics in these categories with + ``analyze``; an empty iterable reports none. ``None`` uses the + ``datafusion.explain.analyze_categories`` setting. Examples: Show the plan in tree format: @@ -1229,9 +1265,22 @@ def explain( Show plan with runtime metrics: >>> df.explain(analyze=True) # doctest: +SKIP + + Show only row-count metrics: + + >>> from datafusion.dataframe import ExplainMetricCategory + >>> df.explain( + ... analyze=True, analyze_categories=[ExplainMetricCategory.ROWS] + ... ) # doctest: +SKIP """ fmt = format.value if format is not None else None - self.df.explain(verbose, analyze, fmt) + level = analyze_level.value if analyze_level is not None else None + categories = ( + [c.value for c in analyze_categories] + if analyze_categories is not None + else None + ) + self.df.explain(verbose, analyze, fmt, show_statistics, level, categories) def logical_plan(self) -> LogicalPlan: """Return the unoptimized ``LogicalPlan``. diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index e44afce98..4e1f48167 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -47,7 +47,11 @@ functions as f, ) from datafusion.common import NullTreatment -from datafusion.dataframe import DataFrameWriteOptions +from datafusion.dataframe import ( + DataFrameWriteOptions, + ExplainAnalyzeLevel, + ExplainMetricCategory, +) from datafusion.dataframe_formatter import ( DataFrameHtmlFormatter, configure_formatter, @@ -3871,6 +3875,57 @@ def test_explain_with_format(capsys, fmt, verbose, analyze, expected_substring): assert expected_substring in captured.out +def _explain_output(capsys, **kwargs): + ctx = SessionContext() + df = ctx.from_pydict({"a": [1, 2]}).filter(column("a") > literal(1)) + df.explain(**kwargs) + return capsys.readouterr().out + + +@pytest.mark.parametrize( + ("kwargs", "present", "absent"), + [ + pytest.param({}, [], ["statistics="], id="default_no_statistics"), + pytest.param( + {"show_statistics": True}, ["statistics=[Rows="], [], id="show_statistics" + ), + pytest.param( + {"analyze": True, "analyze_level": ExplainAnalyzeLevel.DEV}, + ["output_rows=", "output_batches="], + [], + id="analyze_level_dev", + ), + pytest.param( + {"analyze": True, "analyze_level": ExplainAnalyzeLevel.SUMMARY}, + ["output_rows="], + ["output_batches="], + id="analyze_level_summary", + ), + pytest.param( + { + "analyze": True, + "analyze_categories": [ExplainMetricCategory.ROWS], + }, + ["output_rows="], + ["elapsed_compute=", "output_bytes="], + id="analyze_categories_rows", + ), + pytest.param( + {"analyze": True, "analyze_categories": []}, + ["FilterExec: a@0 > 1, metrics=[]"], + ["output_rows="], + id="analyze_categories_empty_suppresses_metrics", + ), + ], +) +def test_explain_options(capsys, kwargs, present, absent): + out = _explain_output(capsys, **kwargs) + for text in present: + assert text in out + for text in absent: + assert text not in out + + @pytest.mark.parametrize( ("window_exprs", "expected_columns"), [ From d2ac89033899479842b209039cccf45422ec15e1 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:47:42 -0400 Subject: [PATCH 08/31] fix: keep existing options when chaining aggregate and window builders Upstream ExprFunctionExt methods on an Expr start from an empty builder, so build() resets every option not set again. The Python function wrappers already apply their keyword options, so calls such as string_agg(..., order_by=...).distinct().build() silently dropped the ordering, and .filter() dropped order_by, and so on. Seed the builder from the expression's existing params instead. A window frame equal to the default for its order_by is left unset so build() derives it again from the final order_by. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/expr.rs | 71 +++++++++++++++++++++++++++++++----- python/tests/test_expr.py | 75 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 138 insertions(+), 8 deletions(-) diff --git a/crates/core/src/expr.rs b/crates/core/src/expr.rs index cab997d7a..2227d0910 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -625,34 +625,38 @@ impl PyExpr { // Expression Function Builder functions pub fn order_by(&self, order_by: Vec) -> PyExprFuncBuilder { - self.expr - .clone() + builder_from_expr(&self.expr) .order_by(to_sort_expressions(order_by)) .into() } pub fn filter(&self, filter: PyExpr) -> PyExprFuncBuilder { - self.expr.clone().filter(filter.expr.clone()).into() + builder_from_expr(&self.expr) + .filter(filter.expr.clone()) + .into() } pub fn distinct(&self) -> PyExprFuncBuilder { - self.expr.clone().distinct().into() + builder_from_expr(&self.expr).distinct().into() } pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder { - self.expr - .clone() + builder_from_expr(&self.expr) .null_treatment(Some(null_treatment.into())) .into() } pub fn partition_by(&self, partition_by: Vec) -> PyExprFuncBuilder { let partition_by = partition_by.iter().map(|e| e.expr.clone()).collect(); - self.expr.clone().partition_by(partition_by).into() + builder_from_expr(&self.expr) + .partition_by(partition_by) + .into() } pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder { - self.expr.clone().window_frame(window_frame.into()).into() + builder_from_expr(&self.expr) + .window_frame(window_frame.into()) + .into() } #[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None))] @@ -743,6 +747,57 @@ impl PyExpr { } } +/// Start an [`ExprFuncBuilder`] that keeps the options already set on `expr`. +/// +/// Upstream's `ExprFunctionExt` methods on an `Expr` start from an empty +/// builder, so `build()` would reset every option not set again. The Python +/// function wrappers already apply their keyword options, so chaining another +/// builder method onto their result must not discard them. +fn builder_from_expr(expr: &Expr) -> ExprFuncBuilder { + match expr { + Expr::AggregateFunction(agg) => { + let params = &agg.params; + let mut builder = expr.clone().null_treatment(params.null_treatment); + if !params.order_by.is_empty() { + builder = builder.order_by(params.order_by.clone()); + } + if let Some(filter) = ¶ms.filter { + builder = builder.filter(filter.as_ref().clone()); + } + if params.distinct { + builder = builder.distinct(); + } + builder + } + Expr::WindowFunction(window) => { + let params = &window.params; + let mut builder = expr.clone().null_treatment(params.null_treatment); + if !params.partition_by.is_empty() { + builder = builder.partition_by(params.partition_by.clone()); + } + let has_order_by = !params.order_by.is_empty(); + if has_order_by { + builder = builder.order_by(params.order_by.clone()); + } + // A frame equal to the default `build()` derived from the order-by is + // left unset, so it is derived again from the final order-by. + if params.window_frame + != datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true)) + { + builder = builder.window_frame(params.window_frame.clone()); + } + if let Some(filter) = ¶ms.filter { + builder = builder.filter(filter.as_ref().clone()); + } + if params.distinct { + builder = builder.distinct(); + } + builder + } + _ => expr.clone().null_treatment(None), + } +} + #[pyclass( from_py_object, frozen, diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index ef006bd91..4e03e56e4 100644 --- a/python/tests/test_expr.py +++ b/python/tests/test_expr.py @@ -53,6 +53,8 @@ TransactionEnd, TransactionStart, Values, + Window, + WindowFrame, coerce_to_expr, coerce_to_expr_list, coerce_to_expr_or_none, @@ -1251,3 +1253,76 @@ def test_expr_to_bytes_no_ctx_default_codec() -> None: restored = Expr.from_bytes(blob, ctx=fresh) assert restored.canonical_name() == original.canonical_name() + + +@pytest.fixture +def builder_df(): + ctx = SessionContext() + return ctx.from_pydict( + {"g": [1, 1, 1, 2], "s": ["y", "x", "z", "w"], "v": [3, 1, 2, 4]} + ) + + +@pytest.mark.parametrize( + ("build_expr", "expected"), + [ + pytest.param( + lambda: ( + functions.array_agg(col("s"), order_by="s") + .filter(col("v") > lit(1)) + .build() + ), + ["w", "y", "z"], + id="order_by_kept_after_filter", + ), + pytest.param( + lambda: ( + functions.array_agg(col("s"), filter=col("v") > lit(1)) + .order_by(col("s").sort(ascending=False)) + .build() + ), + ["z", "y", "w"], + id="filter_kept_after_order_by", + ), + pytest.param( + lambda: ( + functions.string_agg(col("s"), ",", order_by="s").distinct().build() + ), + "w,x,y,z", + id="order_by_kept_after_distinct", + ), + pytest.param( + lambda: ( + functions.first_value(col("s"), order_by="v") + .filter(col("v") > lit(1)) + .build() + ), + "z", + id="first_value_order_by_kept_after_filter", + ), + ], +) +def test_aggregate_builder_keeps_existing_options(builder_df, build_expr, expected): + result = builder_df.aggregate([], [build_expr().alias("r")]) + assert result.collect_column("r")[0].as_py() == expected + + +def test_window_builder_keeps_existing_options(builder_df): + expr = functions.lead(col("v"), order_by="v").partition_by(col("g")).build() + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [2, 3, None, None] + + +def test_window_builder_keeps_explicit_frame(builder_df): + window = Window(order_by=col("v"), window_frame=WindowFrame("rows", 1, 0)) + expr = functions.sum(col("v")).over(window).partition_by(col("g")).build() + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [1, 3, 5, 4] + + +def test_window_builder_rederives_default_frame(builder_df): + # No order_by means a whole-partition frame; adding one later must switch + # to the running frame rather than keep the whole-partition default. + expr = functions.sum(col("v")).over(Window()).order_by(col("v")).build() + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [1, 3, 6, 10] From 97e5905f8a2c9851f84d91e91cf55561975228c8 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:51:44 -0400 Subject: [PATCH 09/31] 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 --- crates/core/src/udf.rs | 27 ++++++++++------ crates/core/src/udwf.rs | 28 ++++++++++------ .../python/tests/_test_aggregate_udf.py | 8 +++++ .../python/tests/_test_scalar_udf.py | 8 +++++ .../python/tests/_test_window_udf.py | 12 +++++++ python/datafusion/user_defined.py | 32 ++++++++++++++++--- 6 files changed, 92 insertions(+), 23 deletions(-) diff --git a/crates/core/src/udf.rs b/crates/core/src/udf.rs index 6376c81a8..5093a564f 100644 --- a/crates/core/src/udf.rs +++ b/crates/core/src/udf.rs @@ -209,6 +209,16 @@ impl ScalarUDFImpl for PythonFunctionScalarUDF { } } +fn scalar_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult { + let data: NonNull = capsule + .pointer_checked(Some(c"datafusion_scalar_udf"))? + .cast(); + let udf = unsafe { data.as_ref() }; + let udf: Arc = udf.into(); + + Ok(ScalarUDF::new_from_shared_impl(udf)) +} + /// Represents a PyScalarUDF #[pyclass( from_py_object, @@ -247,6 +257,12 @@ impl PyScalarUDF { #[staticmethod] pub fn from_pycapsule(func: Bound<'_, PyAny>) -> PyDataFusionResult { + if func.is_instance_of::() { + let capsule = func.cast::().map_err(to_datafusion_err)?; + let function = scalar_udf_from_capsule(capsule)?; + return Ok(Self { function }); + } + if func.hasattr("__datafusion_scalar_udf__")? { let capsule = call_capsule_getter( func.clone(), @@ -254,15 +270,8 @@ impl PyScalarUDF { CapsuleGetterArg::None, )?; let capsule = capsule.cast::().map_err(to_datafusion_err)?; - let data: NonNull = capsule - .pointer_checked(Some(c"datafusion_scalar_udf"))? - .cast(); - let udf = unsafe { data.as_ref() }; - let udf: Arc = udf.into(); - - Ok(Self { - function: ScalarUDF::new_from_shared_impl(udf), - }) + let function = scalar_udf_from_capsule(capsule)?; + Ok(Self { function }) } else { Err(crate::errors::PyDataFusionError::Common( "__datafusion_scalar_udf__ does not exist on ScalarUDF object.".to_string(), diff --git a/crates/core/src/udwf.rs b/crates/core/src/udwf.rs index 8935c9ba8..af3871bb6 100644 --- a/crates/core/src/udwf.rs +++ b/crates/core/src/udwf.rs @@ -216,6 +216,16 @@ pub fn to_rust_partition_evaluator(evaluator: Py) -> PartitionEvaluatorFa Arc::new(move || instantiate_partition_evaluator(&evaluator)) } +fn window_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult { + let data: NonNull = capsule + .pointer_checked(Some(c"datafusion_window_udf"))? + .cast(); + let udwf = unsafe { data.as_ref() }; + let udwf: Arc = udwf.into(); + + Ok(WindowUDF::new_from_shared_impl(udwf)) +} + /// Represents an WindowUDF #[pyclass( from_py_object, @@ -262,19 +272,17 @@ impl PyWindowUDF { #[staticmethod] pub fn from_pycapsule(func: Bound<'_, PyAny>) -> PyDataFusionResult { + if func.is_instance_of::() { + let capsule = func.cast::().map_err(to_datafusion_err)?; + let function = window_udf_from_capsule(capsule)?; + return Ok(Self { function }); + } + let capsule = call_capsule_getter(func, "__datafusion_window_udf__", CapsuleGetterArg::None)?; - let capsule = capsule.cast::().map_err(to_datafusion_err)?; - let data: NonNull = capsule - .pointer_checked(Some(c"datafusion_window_udf"))? - .cast(); - let udwf = unsafe { data.as_ref() }; - let udwf: Arc = udwf.into(); - - Ok(Self { - function: WindowUDF::new_from_shared_impl(udwf), - }) + let function = window_udf_from_capsule(capsule)?; + Ok(Self { function }) } fn __repr__(&self) -> PyResult { diff --git a/examples/datafusion-ffi-example/python/tests/_test_aggregate_udf.py b/examples/datafusion-ffi-example/python/tests/_test_aggregate_udf.py index 7ea6b295c..2df549de2 100644 --- a/examples/datafusion-ffi-example/python/tests/_test_aggregate_udf.py +++ b/examples/datafusion-ffi-example/python/tests/_test_aggregate_udf.py @@ -75,3 +75,11 @@ def test_ffi_aggregate_call_directly(): ] assert result == expected + + +def test_ffi_aggregate_from_bare_capsule(): + ctx = setup_context_with_table() + my_udaf = udaf(MySumUDF().__datafusion_aggregate_udf__()) + + result = ctx.table("test_table").aggregate([], [my_udaf(col("a")).alias("r")]) + assert result.collect_column("r").to_pylist() == [6] diff --git a/examples/datafusion-ffi-example/python/tests/_test_scalar_udf.py b/examples/datafusion-ffi-example/python/tests/_test_scalar_udf.py index 0c949c34a..ecb51b6c5 100644 --- a/examples/datafusion-ffi-example/python/tests/_test_scalar_udf.py +++ b/examples/datafusion-ffi-example/python/tests/_test_scalar_udf.py @@ -68,3 +68,11 @@ def test_ffi_scalar_call_directly(): ] assert result == expected + + +def test_ffi_scalar_from_bare_capsule(): + ctx = setup_context_with_table() + my_udf = udf(IsNullUDF().__datafusion_scalar_udf__()) + + result = ctx.table("test_table").select(my_udf(col("a")).alias("r")) + assert result.collect_column("r").to_pylist() == [False, False, False, True] diff --git a/examples/datafusion-ffi-example/python/tests/_test_window_udf.py b/examples/datafusion-ffi-example/python/tests/_test_window_udf.py index 7d96994b9..9fad7772e 100644 --- a/examples/datafusion-ffi-example/python/tests/_test_window_udf.py +++ b/examples/datafusion-ffi-example/python/tests/_test_window_udf.py @@ -87,3 +87,15 @@ def test_ffi_window_call_directly(): (40, 4), ] assert results == expected + + +def test_ffi_window_from_bare_capsule(): + ctx = setup_context_with_table() + my_udwf = udwf(MyRankUDF().__datafusion_window_udf__()) + + result = ( + ctx.table("test_table") + .select(col("a"), my_udwf().order_by(col("a")).build().alias("r")) + .sort(col("a")) + ) + assert result.collect_column("r").to_pylist() == [1, 2, 3, 4] diff --git a/python/datafusion/user_defined.py b/python/datafusion/user_defined.py index f8d273177..cc17758b1 100644 --- a/python/datafusion/user_defined.py +++ b/python/datafusion/user_defined.py @@ -283,6 +283,10 @@ def udf( @staticmethod def udf(func: ScalarUDFExportable) -> ScalarUDF: ... + @overload + @staticmethod + def udf(func: _PyCapsule) -> ScalarUDF: ... + @staticmethod def udf(*args: Any, **kwargs: Any): # noqa: D417 """Create a new User-Defined Function (UDF). @@ -388,7 +392,7 @@ def wrapper(*args: Any, **kwargs: Any) -> Callable: return decorator - if hasattr(args[0], "__datafusion_scalar_udf__"): + if hasattr(args[0], "__datafusion_scalar_udf__") or _is_pycapsule(args[0]): return ScalarUDF.from_pycapsule(args[0]) if args and callable(args[0]): @@ -398,12 +402,18 @@ def wrapper(*args: Any, **kwargs: Any) -> Callable: return _decorator(*args, **kwargs) @staticmethod - def from_pycapsule(func: ScalarUDFExportable) -> ScalarUDF: + def from_pycapsule(func: ScalarUDFExportable | _PyCapsule) -> ScalarUDF: """Create a Scalar UDF from ScalarUDF PyCapsule object. This function will instantiate a Scalar UDF that uses a DataFusion ScalarUDF that is exported via the FFI bindings. """ + if _is_pycapsule(func): + scalar = cast("ScalarUDF", object.__new__(ScalarUDF)) + scalar._udf = df_internal.ScalarUDF.from_pycapsule(func) + return scalar + + func = cast("ScalarUDFExportable", func) name = str(func.__class__) return ScalarUDF( name=name, @@ -1011,6 +1021,14 @@ def udwf( name: str | None = None, ) -> WindowUDF: ... + @overload + @staticmethod + def udwf(func: WindowUDFExportable) -> WindowUDF: ... + + @overload + @staticmethod + def udwf(func: _PyCapsule) -> WindowUDF: ... + @staticmethod def udwf(*args: Any, **kwargs: Any): # noqa: D417 """Create a new User-Defined Window Function (UDWF). @@ -1075,7 +1093,7 @@ def udwf(*args: Any, **kwargs: Any): # noqa: D417 Returns: A user-defined window function that can be used in window function calls. """ - if hasattr(args[0], "__datafusion_window_udf__"): + if hasattr(args[0], "__datafusion_window_udf__") or _is_pycapsule(args[0]): return WindowUDF.from_pycapsule(args[0]) if args and callable(args[0]): @@ -1146,12 +1164,18 @@ def wrapper(*args: Any, **kwargs: Any) -> Expr: return decorator @staticmethod - def from_pycapsule(func: WindowUDFExportable) -> WindowUDF: + def from_pycapsule(func: WindowUDFExportable | _PyCapsule) -> WindowUDF: """Create a Window UDF from WindowUDF PyCapsule object. This function will instantiate a Window UDF that uses a DataFusion WindowUDF that is exported via the FFI bindings. """ + if _is_pycapsule(func): + window = cast("WindowUDF", object.__new__(WindowUDF)) + window._udwf = df_internal.WindowUDF.from_pycapsule(func) + return window + + func = cast("WindowUDFExportable", func) name = str(func.__class__) return WindowUDF( name=name, From d53a255f3eb99dc3d10c0aed091587c1ce09383e Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 08:52:31 -0400 Subject: [PATCH 10/31] chore: export TableProviderFactory from datafusion.catalog The ABC is defined in datafusion.catalog beside CatalogProvider and SchemaProvider, which are in its __all__, but was only listed in the package root's __all__. Co-Authored-By: Claude Opus 5.5 --- python/datafusion/catalog.py | 1 + 1 file changed, 1 insertion(+) diff --git a/python/datafusion/catalog.py b/python/datafusion/catalog.py index 20da5e671..d1d501172 100644 --- a/python/datafusion/catalog.py +++ b/python/datafusion/catalog.py @@ -45,6 +45,7 @@ "Schema", "SchemaProvider", "Table", + "TableProviderFactory", ] From 5d43576fb21395c18c9975a280284aee8d00308a Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 09:22:24 -0400 Subject: [PATCH 11/31] docs(skill): record DataFrame.to_string as not needing exposure str(df) is the Pythonic way to get a string and already uses __repr__ and the configurable formatter, so check-upstream should not flag the upstream to_string as a gap. Co-Authored-By: Claude Opus 5.5 --- .ai/skills/check-upstream/SKILL.md | 1 + 1 file changed, 1 insertion(+) diff --git a/.ai/skills/check-upstream/SKILL.md b/.ai/skills/check-upstream/SKILL.md index 828f227d8..0c80d31ae 100644 --- a/.ai/skills/check-upstream/SKILL.md +++ b/.ai/skills/check-upstream/SKILL.md @@ -146,6 +146,7 @@ The user may specify an area via `$ARGUMENTS`. If no area is specified or "all" - `show_limit` — already covered by `DataFrame.show()`, which provides the same functionality with a simpler API - `with_param_values` — already covered by the `param_values` argument on `SessionContext.sql()`, which accomplishes the same thing more robustly - `union_by_name_distinct` — already covered by `DataFrame.union_by_name(distinct=True)`, which provides a more Pythonic API +- `to_string` — `str(df)` is the Pythonic way to get a string and already goes through `__repr__` and the configurable formatter. A separate `to_string()` would either duplicate `str(df)` or render every row through a different path (session `datafusion.format.*` options, no formatter), giving a third text rendering alongside `repr` and `show()` **How to check:** 1. Fetch the upstream DataFrame documentation page listing all methods From f6d047889999d6abc80fb7982651544f29272fce Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 10:53:10 -0400 Subject: [PATCH 12/31] fix: resolve the _PyCapsule type alias so capsule overloads type-check user_defined.py imported CapsuleType from _typeshed, which does not define it. Pyright reports `"CapsuleType" is unknown import symbol`, so _PyCapsule resolved to Unknown and every overload and parameter typed with it accepted any argument. This affected the udaf/from_pycapsule hints added in #1277 as well as the new udf/udwf ones. Import it from types on Python 3.13+ and from typing_extensions (already a dependency below 3.13) otherwise. Checked with pyright at --pythonversion 3.10 and 3.13: udf/udaf/udwf and each from_pycapsule accept a CapsuleType and return the right wrapper, and WindowUDF.from_pycapsule(1) is now rejected where it previously passed. Co-Authored-By: Claude Opus 5.5 --- python/datafusion/user_defined.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/python/datafusion/user_defined.py b/python/datafusion/user_defined.py index cc17758b1..92d541591 100644 --- a/python/datafusion/user_defined.py +++ b/python/datafusion/user_defined.py @@ -31,7 +31,12 @@ from datafusion.expr import Expr if TYPE_CHECKING: - from _typeshed import CapsuleType as _PyCapsule + import sys + + if sys.version_info >= (3, 13): + from types import CapsuleType as _PyCapsule + else: + from typing_extensions import CapsuleType as _PyCapsule _R = TypeVar("_R", bound=pa.Array) from collections.abc import Callable, Sequence From 6345b17f9a0d66a3d169a280070ce2d2c317f189 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 11:19:11 -0400 Subject: [PATCH 13/31] fix: keep an explicit window frame that equals the default when chaining A built window function always stores a concrete frame, so the builder could not tell a frame the user chose from the default. A frame equal to the no-order_by default was treated as unset and re-derived as the running frame once order_by was chained, silently changing results. Record on the Python Expr whether over() or window_frame() set the frame explicitly, and pass that to builder_from_expr so the frame is kept. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/expr.rs | 53 ++++++++++++++++++++++-------- python/datafusion/expr.py | 68 ++++++++++++++++++++++++++++----------- python/tests/test_expr.py | 22 +++++++++++++ 3 files changed, 110 insertions(+), 33 deletions(-) diff --git a/crates/core/src/expr.rs b/crates/core/src/expr.rs index 2227d0910..1b427c9f8 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -624,37 +624,56 @@ impl PyExpr { // Expression Function Builder functions - pub fn order_by(&self, order_by: Vec) -> PyExprFuncBuilder { - builder_from_expr(&self.expr) + #[pyo3(signature = (order_by, keep_window_frame=false))] + pub fn order_by( + &self, + order_by: Vec, + keep_window_frame: bool, + ) -> PyExprFuncBuilder { + builder_from_expr(&self.expr, keep_window_frame) .order_by(to_sort_expressions(order_by)) .into() } - pub fn filter(&self, filter: PyExpr) -> PyExprFuncBuilder { - builder_from_expr(&self.expr) + #[pyo3(signature = (filter, keep_window_frame=false))] + pub fn filter(&self, filter: PyExpr, keep_window_frame: bool) -> PyExprFuncBuilder { + builder_from_expr(&self.expr, keep_window_frame) .filter(filter.expr.clone()) .into() } - pub fn distinct(&self) -> PyExprFuncBuilder { - builder_from_expr(&self.expr).distinct().into() + #[pyo3(signature = (keep_window_frame=false))] + pub fn distinct(&self, keep_window_frame: bool) -> PyExprFuncBuilder { + builder_from_expr(&self.expr, keep_window_frame) + .distinct() + .into() } - pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder { - builder_from_expr(&self.expr) + #[pyo3(signature = (null_treatment, keep_window_frame=false))] + pub fn null_treatment( + &self, + null_treatment: NullTreatment, + keep_window_frame: bool, + ) -> PyExprFuncBuilder { + builder_from_expr(&self.expr, keep_window_frame) .null_treatment(Some(null_treatment.into())) .into() } - pub fn partition_by(&self, partition_by: Vec) -> PyExprFuncBuilder { + #[pyo3(signature = (partition_by, keep_window_frame=false))] + pub fn partition_by( + &self, + partition_by: Vec, + keep_window_frame: bool, + ) -> PyExprFuncBuilder { let partition_by = partition_by.iter().map(|e| e.expr.clone()).collect(); - builder_from_expr(&self.expr) + builder_from_expr(&self.expr, keep_window_frame) .partition_by(partition_by) .into() } pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder { - builder_from_expr(&self.expr) + builder_from_expr(&self.expr, false) .window_frame(window_frame.into()) .into() } @@ -753,7 +772,12 @@ impl PyExpr { /// builder, so `build()` would reset every option not set again. The Python /// function wrappers already apply their keyword options, so chaining another /// builder method onto their result must not discard them. -fn builder_from_expr(expr: &Expr) -> ExprFuncBuilder { +/// +/// A built window function always stores a concrete frame, so whether the user +/// chose it is lost. `keep_window_frame` carries that from the Python side; when +/// false, a frame equal to the default for the current order-by is treated as +/// unset. +fn builder_from_expr(expr: &Expr, keep_window_frame: bool) -> ExprFuncBuilder { match expr { Expr::AggregateFunction(agg) => { let params = &agg.params; @@ -781,8 +805,9 @@ fn builder_from_expr(expr: &Expr) -> ExprFuncBuilder { } // A frame equal to the default `build()` derived from the order-by is // left unset, so it is derived again from the final order-by. - if params.window_frame - != datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true)) + if keep_window_frame + || params.window_frame + != datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true)) { builder = builder.window_frame(params.window_frame.clone()); } diff --git a/python/datafusion/expr.py b/python/datafusion/expr.py index 18ce3554d..458c5dc87 100644 --- a/python/datafusion/expr.py +++ b/python/datafusion/expr.py @@ -453,6 +453,10 @@ class Expr: # noqa: PLW1641 :ref:`Expressions` in the online documentation for more information. """ + # Set by ``over()`` when the window frame was given explicitly, so chaining a + # builder method keeps it even if it equals the default frame. + _explicit_window_frame = False + def __init__(self, expr: expr_internal.RawExpr) -> None: """This constructor should not be called by the end user.""" self.expr = expr @@ -1031,7 +1035,7 @@ def order_by(self, *exprs: Expr | SortExpr) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder(self.expr.order_by([sort_or_default(e) for e in exprs])) + return self._builder(self.expr.order_by, [sort_or_default(e) for e in exprs]) def filter(self, filter: Expr) -> ExprFuncBuilder: """Filter an aggregate function. @@ -1040,7 +1044,7 @@ def filter(self, filter: Expr) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder(self.expr.filter(filter.expr)) + return self._builder(self.expr.filter, filter.expr) def distinct(self) -> ExprFuncBuilder: """Only evaluate distinct values for an aggregate function. @@ -1049,7 +1053,7 @@ def distinct(self) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder(self.expr.distinct()) + return self._builder(self.expr.distinct) def null_treatment(self, null_treatment: NullTreatment) -> ExprFuncBuilder: """Set the treatment for ``null`` values for a window or aggregate function. @@ -1058,7 +1062,7 @@ def null_treatment(self, null_treatment: NullTreatment) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder(self.expr.null_treatment(null_treatment.value)) + return self._builder(self.expr.null_treatment, null_treatment.value) def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: """Set the partitioning for a window function. @@ -1067,7 +1071,7 @@ def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder(self.expr.partition_by([e.expr for e in partition_by])) + return self._builder(self.expr.partition_by, [e.expr for e in partition_by]) def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: """Set the frame fora window function. @@ -1076,7 +1080,16 @@ def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder(self.expr.window_frame(window_frame.window_frame)) + return ExprFuncBuilder( + self.expr.window_frame(window_frame.window_frame), + explicit_window_frame=True, + ) + + def _builder(self, method: Any, *args: Any) -> ExprFuncBuilder: + keep = self._explicit_window_frame + return ExprFuncBuilder( + method(*args, keep_window_frame=keep), explicit_window_frame=keep + ) def over(self, window: Window) -> Expr: """Turn an aggregate function into a window function. @@ -1099,7 +1112,7 @@ def over(self, window: Window) -> Expr: window._null_treatment.value if window._null_treatment is not None else None ) - return Expr( + result = Expr( self.expr.over( partition_by=partition_by_raw, order_by=order_by_raw, @@ -1107,6 +1120,8 @@ def over(self, window: Window) -> Expr: null_treatment=null_treatment_raw, ) ) + result._explicit_window_frame = window_frame_raw is not None + return result def asin(self) -> Expr: """Returns the arc sine or inverse sine of a number.""" @@ -1534,8 +1549,22 @@ def cot(self) -> Expr: class ExprFuncBuilder: - def __init__(self, builder: expr_internal.ExprFuncBuilder) -> None: + def __init__( + self, + builder: expr_internal.ExprFuncBuilder, + explicit_window_frame: bool = False, + ) -> None: self.builder = builder + self._explicit_window_frame = explicit_window_frame + + def _wrap( + self, + builder: expr_internal.ExprFuncBuilder, + explicit_window_frame: bool = False, + ) -> ExprFuncBuilder: + return ExprFuncBuilder( + builder, self._explicit_window_frame or explicit_window_frame + ) def order_by(self, *exprs: Expr) -> ExprFuncBuilder: """Set the ordering for a window or aggregate function. @@ -1543,35 +1572,36 @@ def order_by(self, *exprs: Expr) -> ExprFuncBuilder: Values given in ``exprs`` must be sort expressions. You can convert any other expression to a sort expression using `.sort()`. """ - return ExprFuncBuilder( - self.builder.order_by([sort_or_default(e) for e in exprs]) - ) + return self._wrap(self.builder.order_by([sort_or_default(e) for e in exprs])) def filter(self, filter: Expr) -> ExprFuncBuilder: """Filter values during aggregation.""" - return ExprFuncBuilder(self.builder.filter(filter.expr)) + return self._wrap(self.builder.filter(filter.expr)) def distinct(self) -> ExprFuncBuilder: """Only evaluate distinct values during aggregation.""" - return ExprFuncBuilder(self.builder.distinct()) + return self._wrap(self.builder.distinct()) def null_treatment(self, null_treatment: NullTreatment) -> ExprFuncBuilder: """Set how nulls are treated for either window or aggregate functions.""" - return ExprFuncBuilder(self.builder.null_treatment(null_treatment.value)) + return self._wrap(self.builder.null_treatment(null_treatment.value)) def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: """Set partitioning for window functions.""" - return ExprFuncBuilder( - self.builder.partition_by([e.expr for e in partition_by]) - ) + return self._wrap(self.builder.partition_by([e.expr for e in partition_by])) def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: """Set window frame for window functions.""" - return ExprFuncBuilder(self.builder.window_frame(window_frame.window_frame)) + return self._wrap( + self.builder.window_frame(window_frame.window_frame), + explicit_window_frame=True, + ) def build(self) -> Expr: """Create an expression from a Function Builder.""" - return Expr(self.builder.build()) + result = Expr(self.builder.build()) + result._explicit_window_frame = self._explicit_window_frame + return result class Window: diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index 4e03e56e4..eeb7ef16f 100644 --- a/python/tests/test_expr.py +++ b/python/tests/test_expr.py @@ -1320,6 +1320,28 @@ def test_window_builder_keeps_explicit_frame(builder_df): assert result.collect_column("r").to_pylist() == [1, 3, 5, 4] +def test_window_builder_keeps_explicit_default_frame(builder_df): + # An explicit frame equal to the no-order_by default must survive a later + # order_by instead of being re-derived as the running frame. + window = Window(window_frame=WindowFrame("rows", None, None)) + expr = functions.sum(col("v")).over(window).order_by(col("v")).build() + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] + + +def test_window_builder_keeps_frame_set_on_builder(builder_df): + expr = ( + functions.sum(col("v")) + .over(Window()) + .window_frame(WindowFrame("rows", None, None)) + .build() + .order_by(col("v")) + .build() + ) + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] + + def test_window_builder_rederives_default_frame(builder_df): # No order_by means a whole-partition frame; adding one later must switch # to the running frame rather than keep the whole-partition default. From 73b052d76cac3e06763ae6c77b9c92bc21a3308b Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 11:20:04 -0400 Subject: [PATCH 14/31] fix: leave lead and lag null_treatment unset by default Defaulting to NullTreatment.RESPECT_NULLS passed Some(RespectNulls) to Rust, which adds "RESPECT NULLS" to the generated column name and breaks code that refers to an un-aliased lead/lag output by name. Default to None instead; respecting nulls is already the behavior when unset. Co-Authored-By: Claude Opus 5.5 --- python/datafusion/functions/__init__.py | 12 ++++++++---- python/tests/test_dataframe.py | 8 ++++++++ 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 7b4f24a62..2b189902e 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -7000,7 +7000,7 @@ def lead( default_value: Any | None = None, partition_by: list[Expr] | Expr | None = None, order_by: list[SortKey] | SortKey | None = None, - null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS, + null_treatment: NullTreatment | None = None, ) -> Expr: """Create a lead window function. @@ -7079,7 +7079,9 @@ def lead( default_value, partition_by=partition_by_raw, order_by=order_by_raw, - null_treatment=null_treatment.value, + null_treatment=( + null_treatment.value if null_treatment is not None else None + ), ) ) @@ -7090,7 +7092,7 @@ def lag( default_value: Any | None = None, partition_by: list[Expr] | Expr | None = None, order_by: list[SortKey] | SortKey | None = None, - null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS, + null_treatment: NullTreatment | None = None, ) -> Expr: """Create a lag window function. @@ -7166,7 +7168,9 @@ def lag( default_value, partition_by=partition_by_raw, order_by=order_by_raw, - null_treatment=null_treatment.value, + null_treatment=( + null_treatment.value if null_treatment is not None else None + ), ) ) diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index 4e1f48167..590c5cd6a 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -1186,6 +1186,14 @@ def test_window_partition_by_accepts_string(partitioned_df, partition): assert table.column("fv").to_pylist() == [1, 1, 1, 1, 5, 5, 5] +@pytest.mark.parametrize("func", [f.lead, f.lag]) +def test_lead_lag_default_null_treatment_keeps_column_name(partitioned_df, func): + """Omitting null_treatment must not add RESPECT NULLS to the output name.""" + df = partitioned_df.select(func(column("b"), order_by=column("a"))) + name = df.schema().names[0] + assert "RESPECT NULLS" not in name + + @pytest.mark.parametrize( ("units", "start_bound", "end_bound"), [ From 81735542b9c7db5c52d866591af6a64b39b4d3f2 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 12:13:08 -0400 Subject: [PATCH 15/31] fix: keep existing window function options in over() Calling over() on an expression that was already a window function rebuilt it from an empty builder, silently dropping its order_by, null_treatment, and explicit window frame. For example, lead(v, order_by="i", null_treatment=IGNORE_NULLS).over(Window(partition_by=[g])) lost both the ordering and IGNORE NULLS. Start from builder_from_expr instead, and apply only the options the Window sets. Pass the Python-side explicit-frame flag through so a frame equal to the default is not re-derived when over() adds an order_by. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/expr.rs | 9 +++++---- crates/core/src/functions.rs | 25 ++++++++++++++++++++++--- python/datafusion/expr.py | 5 ++++- python/tests/test_expr.py | 23 +++++++++++++++++++++++ 4 files changed, 54 insertions(+), 8 deletions(-) diff --git a/crates/core/src/expr.rs b/crates/core/src/expr.rs index 1b427c9f8..74a7576b2 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -55,7 +55,7 @@ use crate::expr::aggregate_expr::PyAggregateFunction; use crate::expr::binary_expr::PyBinaryExpr; use crate::expr::column::PyColumn; use crate::expr::literal::PyLiteral; -use crate::functions::add_builder_fns_to_window; +use crate::functions::{add_builder_fns_to_window, apply_window_options}; use crate::pyarrow_util::scalar_to_pyarrow; use crate::sql::logical::PyLogicalPlan; @@ -678,13 +678,14 @@ impl PyExpr { .into() } - #[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None))] + #[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None, keep_window_frame=false))] pub fn over( &self, partition_by: Option>, window_frame: Option, order_by: Option>, null_treatment: Option, + keep_window_frame: bool, ) -> PyDataFusionResult { match &self.expr { Expr::AggregateFunction(agg_fn) => { @@ -701,8 +702,8 @@ impl PyExpr { null_treatment, ) } - Expr::WindowFunction(_) => add_builder_fns_to_window( - self.expr.clone(), + Expr::WindowFunction(_) => apply_window_options( + builder_from_expr(&self.expr, keep_window_frame), partition_by, window_frame, order_by, diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index af1de455f..f9c936ebb 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -19,7 +19,7 @@ use std::collections::HashMap; use datafusion::common::{Column, ScalarValue, TableReference}; use datafusion::logical_expr::expr::{Alias, FieldMetadata, NullTreatment as DFNullTreatment}; -use datafusion::logical_expr::{Expr, ExprFunctionExt, lit}; +use datafusion::logical_expr::{Expr, ExprFuncBuilder, ExprFunctionExt, lit}; use datafusion::{functions, functions_aggregate, functions_window}; use pyo3::prelude::*; use pyo3::wrap_pyfunction; @@ -866,8 +866,27 @@ pub(crate) fn add_builder_fns_to_window( order_by: Option>, null_treatment: Option, ) -> PyDataFusionResult { - let null_treatment = null_treatment.map(|n| n.into()); - let mut builder = window_fn.null_treatment(null_treatment); + apply_window_options( + window_fn.null_treatment(None), + partition_by, + window_frame, + order_by, + null_treatment, + ) +} + +/// Applies the options that are `Some` to `builder` and builds it. Options +/// that are `None` keep whatever `builder` already holds. +pub(crate) fn apply_window_options( + mut builder: ExprFuncBuilder, + partition_by: Option>, + window_frame: Option, + order_by: Option>, + null_treatment: Option, +) -> PyDataFusionResult { + if let Some(null_treatment) = null_treatment { + builder = builder.null_treatment(Some(null_treatment.into())); + } if let Some(partition_cols) = partition_by { builder = builder.partition_by( diff --git a/python/datafusion/expr.py b/python/datafusion/expr.py index 458c5dc87..7bd72942b 100644 --- a/python/datafusion/expr.py +++ b/python/datafusion/expr.py @@ -1118,9 +1118,12 @@ def over(self, window: Window) -> Expr: order_by=order_by_raw, window_frame=window_frame_raw, null_treatment=null_treatment_raw, + keep_window_frame=self._explicit_window_frame, ) ) - result._explicit_window_frame = window_frame_raw is not None + result._explicit_window_frame = ( + window_frame_raw is not None or self._explicit_window_frame + ) return result def asin(self) -> Expr: diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index eeb7ef16f..504436a5e 100644 --- a/python/tests/test_expr.py +++ b/python/tests/test_expr.py @@ -32,6 +32,7 @@ lit_with_metadata, literal_with_metadata, ) +from datafusion.common import NullTreatment from datafusion.expr import ( EXPR_TYPE_ERROR, Aggregate, @@ -1342,6 +1343,28 @@ def test_window_builder_keeps_frame_set_on_builder(builder_df): assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] +def test_over_keeps_window_function_options(): + ctx = SessionContext() + df = ctx.from_pydict({"g": [1, 1, 1, 2], "i": [1, 2, 3, 4], "v": [1, None, 3, 4]}) + expr = functions.lead( + col("v"), 1, order_by="i", null_treatment=NullTreatment.IGNORE_NULLS + ).over(Window(partition_by=[col("g")])) + result = df.select(col("i"), expr.alias("r")).sort(col("i")) + assert result.collect_column("r").to_pylist() == [3, 3, None, None] + + +def test_over_keeps_explicit_default_frame(builder_df): + # A frame equal to the no-order_by default, set by an earlier over(), must + # survive a later over() that adds an order_by. + expr = ( + functions.sum(col("v")) + .over(Window(window_frame=WindowFrame("rows", None, None))) + .over(Window(order_by=col("v"))) + ) + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] + + def test_window_builder_rederives_default_frame(builder_df): # No order_by means a whole-partition frame; adding one later must switch # to the running frame rather than keep the whole-partition default. From 96f990b4d0ae1ac29bc742aff8170ddb016b51df Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Fri, 25 Sep 2026 12:26:40 -0400 Subject: [PATCH 16/31] feat: accept native ints and a single argument in range and gen_series range() required all three of start, stop, and step as Expr, although upstream also accepts range(stop) and range(start, stop). Make stop and step optional in the bindings for both range and gen_series, so a single argument is the upper bound starting at 0, like Python's built-in range. Accept plain ints for start, stop, and step, coerced to literals, so callers no longer need lit(). Passing step without stop raises ValueError. Co-Authored-By: Claude Opus 5.5 --- crates/core/src/functions.rs | 45 ++++++++++--- python/datafusion/functions/__init__.py | 90 +++++++++++++++++++------ python/tests/test_functions.py | 17 +++++ 3 files changed, 120 insertions(+), 32 deletions(-) diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index f9c936ebb..79c0524e7 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -131,20 +131,46 @@ fn array_to_string(array: PyExpr, delimiter: PyExpr, null_string: Option .into() } -#[pyfunction] -#[pyo3(signature = (start, stop, step=None))] -fn gen_series(start: PyExpr, stop: PyExpr, step: Option) -> PyExpr { - let mut args = vec![start.into(), stop.into()]; - if let Some(step) = step { - args.push(step.into()); - } +/// Builds `range` or `gen_series` from its one, two, or three arguments. +fn series_expr( + udf: std::sync::Arc, + start: PyExpr, + stop: Option, + step: Option, +) -> PyExpr { + let args = std::iter::once(start) + .chain(stop) + .chain(step) + .map(Into::into) + .collect(); Expr::ScalarFunction(datafusion::logical_expr::expr::ScalarFunction::new_udf( - datafusion::functions_nested::range::gen_series_udf(), - args, + udf, args, )) .into() } +#[pyfunction] +#[pyo3(signature = (start, stop=None, step=None))] +fn range(start: PyExpr, stop: Option, step: Option) -> PyExpr { + series_expr( + datafusion::functions_nested::range::range_udf(), + start, + stop, + step, + ) +} + +#[pyfunction] +#[pyo3(signature = (start, stop=None, step=None))] +fn gen_series(start: PyExpr, stop: Option, step: Option) -> PyExpr { + series_expr( + datafusion::functions_nested::range::gen_series_udf(), + start, + stop, + step, + ) +} + #[pyfunction] fn make_map(keys: Vec, values: Vec) -> PyExpr { let keys = keys.into_iter().map(|x| x.into()).collect(); @@ -701,7 +727,6 @@ array_fn!(array_min, array); array_fn!(array_reverse, array); array_fn!(cardinality, array); array_fn!(flatten, array); -array_fn!(range, start stop step); // Map Functions array_fn!(map_keys, map); diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 2b189902e..af255c120 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -3126,18 +3126,57 @@ def array(*args: Expr) -> Expr: return make_array(*args) -def range(start: Expr, stop: Expr, step: Expr) -> Expr: - """Create a list of values in the range between start and stop. +def _series( + fn: Callable[..., Any], + name: str, + start: Expr | int, + stop: Expr | int | None, + step: Expr | int | None, +) -> Expr: + if stop is None and step is not None: + msg = f"{name}() requires stop when step is given" + raise ValueError(msg) + stop = coerce_to_expr_or_none(stop) + step = coerce_to_expr_or_none(step) + return Expr( + fn( + coerce_to_expr(start).expr, + stop.expr if stop is not None else None, + step.expr if step is not None else None, + ) + ) + + +def range( + start: Expr | int, + stop: Expr | int | None = None, + step: Expr | int | None = None, +) -> Expr: + """Create a list of values from ``start`` up to, but excluding, ``stop``. + + With a single argument, it is the upper bound and the range starts at 0, + like Python's built-in :py:class:`range`. Examples: >>> ctx = dfn.SessionContext() >>> df = ctx.from_pydict({"a": [1]}) - >>> result = df.select( - ... dfn.functions.range(dfn.lit(0), dfn.lit(5), dfn.lit(2)).alias("r")) + >>> result = df.select(dfn.functions.range(5).alias("r")) + >>> result.collect_column("r")[0].as_py() + [0, 1, 2, 3, 4] + + Specify a ``stop``: + + >>> result = df.select(dfn.functions.range(1, stop=5).alias("r")) + >>> result.collect_column("r")[0].as_py() + [1, 2, 3, 4] + + Specify a ``step``: + + >>> result = df.select(dfn.functions.range(0, stop=5, step=2).alias("r")) >>> result.collect_column("r")[0].as_py() [0, 2, 4] """ - return Expr(f.range(start.expr, stop.expr, step.expr)) + return _series(f.range, "range", start, stop, step) def uuid() -> Expr: @@ -5077,38 +5116,45 @@ def string_to_list( return string_to_array(string, delimiter, null_string) -def gen_series(start: Expr, stop: Expr, step: Expr | None = None) -> Expr: - """Creates a list of values in the range between start and stop. +def gen_series( + start: Expr | int, + stop: Expr | int | None = None, + step: Expr | int | None = None, +) -> Expr: + """Creates a list of values from ``start`` up to and including ``stop``. - Unlike :py:func:`range`, this includes the upper bound. + Unlike :py:func:`range`, this includes the upper bound. With a single + argument, it is the upper bound and the series starts at 0. Examples: >>> ctx = dfn.SessionContext() >>> df = ctx.from_pydict({"a": [0]}) - >>> result = df.select( - ... dfn.functions.gen_series( - ... dfn.lit(1), dfn.lit(5), - ... ).alias("result")) + >>> result = df.select(dfn.functions.gen_series(3).alias("result")) + >>> result.collect_column("result")[0].as_py() + [0, 1, 2, 3] + + Specify a ``stop``: + + >>> result = df.select(dfn.functions.gen_series(1, stop=5).alias("result")) >>> result.collect_column("result")[0].as_py() [1, 2, 3, 4, 5] - Specify a custom ``step``: + Specify a ``step``: >>> result = df.select( - ... dfn.functions.gen_series( - ... dfn.lit(1), dfn.lit(10), step=dfn.lit(3), - ... ).alias("result")) + ... dfn.functions.gen_series(1, stop=10, step=3).alias("result")) >>> result.collect_column("result")[0].as_py() [1, 4, 7, 10] """ - step_expr = step.expr if step is not None else None - return Expr(f.gen_series(start.expr, stop.expr, step_expr)) + return _series(f.gen_series, "gen_series", start, stop, step) -def generate_series(start: Expr, stop: Expr, step: Expr | None = None) -> Expr: - """Creates a list of values in the range between start and stop. - - Unlike :py:func:`range`, this includes the upper bound. +def generate_series( + start: Expr | int, + stop: Expr | int | None = None, + step: Expr | int | None = None, +) -> Expr: + """Creates a list of values from ``start`` up to and including ``stop``. See Also: This is an alias for :py:func:`gen_series`. diff --git a/python/tests/test_functions.py b/python/tests/test_functions.py index f5f1921f8..3305fdc71 100644 --- a/python/tests/test_functions.py +++ b/python/tests/test_functions.py @@ -2359,6 +2359,23 @@ def test_gen_series_with_step(): assert result[0].column(0)[0].as_py() == [1, 4, 7, 10] +@pytest.mark.parametrize( + ("func", "expected"), + [(f.range, [[0], [0, 1]]), (f.gen_series, [[0, 1], [0, 1, 2]])], +) +def test_series_single_arg_accepts_column(func, expected): + ctx = SessionContext() + df = ctx.from_pydict({"n": [1, 2]}) + result = df.select(func(column("n")).alias("v")) + assert result.collect_column("v").to_pylist() == expected + + +@pytest.mark.parametrize("func", [f.range, f.gen_series, f.generate_series]) +def test_series_step_requires_stop(func): + with pytest.raises(ValueError, match="requires stop"): + func(0, step=2) + + class TestPythonicNativeTypes: """Tests for accepting native Python types instead of requiring lit().""" From d7312799f6ff9a54e1640256abcfc67fa421c669 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:02:37 -0400 Subject: [PATCH 17/31] fix: decide window frame re-derivation from the expression alone The explicit-frame flag added in 6345b17f lived only on the Python Expr, so copy, pickle, from_bytes and SQL all lost it. The same builder chain then gave different results for identical expressions, including between a driver and the workers it pickles expressions to. Drop the flag. A frame equal to the default for the current order-by is treated as unset when chaining, which matches main. Document the rule and how to keep such a frame (set it after order_by, or in the same Window) in the window functions guide. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/core/src/expr.rs | 59 +++++---------- .../user-guide/common-operations/windows.md | 27 +++++++ python/datafusion/expr.py | 71 +++++-------------- python/tests/test_expr.py | 39 +++++----- 4 files changed, 83 insertions(+), 113 deletions(-) diff --git a/crates/core/src/expr.rs b/crates/core/src/expr.rs index 74a7576b2..eeb01cbeb 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -624,68 +624,48 @@ impl PyExpr { // Expression Function Builder functions - #[pyo3(signature = (order_by, keep_window_frame=false))] - pub fn order_by( - &self, - order_by: Vec, - keep_window_frame: bool, - ) -> PyExprFuncBuilder { - builder_from_expr(&self.expr, keep_window_frame) + pub fn order_by(&self, order_by: Vec) -> PyExprFuncBuilder { + builder_from_expr(&self.expr) .order_by(to_sort_expressions(order_by)) .into() } - #[pyo3(signature = (filter, keep_window_frame=false))] - pub fn filter(&self, filter: PyExpr, keep_window_frame: bool) -> PyExprFuncBuilder { - builder_from_expr(&self.expr, keep_window_frame) + pub fn filter(&self, filter: PyExpr) -> PyExprFuncBuilder { + builder_from_expr(&self.expr) .filter(filter.expr.clone()) .into() } - #[pyo3(signature = (keep_window_frame=false))] - pub fn distinct(&self, keep_window_frame: bool) -> PyExprFuncBuilder { - builder_from_expr(&self.expr, keep_window_frame) - .distinct() - .into() + pub fn distinct(&self) -> PyExprFuncBuilder { + builder_from_expr(&self.expr).distinct().into() } - #[pyo3(signature = (null_treatment, keep_window_frame=false))] - pub fn null_treatment( - &self, - null_treatment: NullTreatment, - keep_window_frame: bool, - ) -> PyExprFuncBuilder { - builder_from_expr(&self.expr, keep_window_frame) + pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder { + builder_from_expr(&self.expr) .null_treatment(Some(null_treatment.into())) .into() } - #[pyo3(signature = (partition_by, keep_window_frame=false))] - pub fn partition_by( - &self, - partition_by: Vec, - keep_window_frame: bool, - ) -> PyExprFuncBuilder { + pub fn partition_by(&self, partition_by: Vec) -> PyExprFuncBuilder { let partition_by = partition_by.iter().map(|e| e.expr.clone()).collect(); - builder_from_expr(&self.expr, keep_window_frame) + builder_from_expr(&self.expr) .partition_by(partition_by) .into() } pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder { - builder_from_expr(&self.expr, false) + builder_from_expr(&self.expr) .window_frame(window_frame.into()) .into() } - #[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None, keep_window_frame=false))] + #[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None))] pub fn over( &self, partition_by: Option>, window_frame: Option, order_by: Option>, null_treatment: Option, - keep_window_frame: bool, ) -> PyDataFusionResult { match &self.expr { Expr::AggregateFunction(agg_fn) => { @@ -703,7 +683,7 @@ impl PyExpr { ) } Expr::WindowFunction(_) => apply_window_options( - builder_from_expr(&self.expr, keep_window_frame), + builder_from_expr(&self.expr), partition_by, window_frame, order_by, @@ -775,10 +755,10 @@ impl PyExpr { /// builder method onto their result must not discard them. /// /// A built window function always stores a concrete frame, so whether the user -/// chose it is lost. `keep_window_frame` carries that from the Python side; when -/// false, a frame equal to the default for the current order-by is treated as -/// unset. -fn builder_from_expr(expr: &Expr, keep_window_frame: bool) -> ExprFuncBuilder { +/// chose it is lost. A frame equal to the default for the current order-by is +/// treated as unset, which depends only on the expression and so behaves the +/// same after a copy, pickle, or round trip through protobuf or SQL. +fn builder_from_expr(expr: &Expr) -> ExprFuncBuilder { match expr { Expr::AggregateFunction(agg) => { let params = &agg.params; @@ -806,9 +786,8 @@ fn builder_from_expr(expr: &Expr, keep_window_frame: bool) -> ExprFuncBuilder { } // A frame equal to the default `build()` derived from the order-by is // left unset, so it is derived again from the final order-by. - if keep_window_frame - || params.window_frame - != datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true)) + if params.window_frame + != datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true)) { builder = builder.window_frame(params.window_frame.clone()); } diff --git a/docs/source/user-guide/common-operations/windows.md b/docs/source/user-guide/common-operations/windows.md index bee96c820..dfab8f958 100644 --- a/docs/source/user-guide/common-operations/windows.md +++ b/docs/source/user-guide/common-operations/windows.md @@ -135,6 +135,33 @@ df.select( ) ``` +(window_frame_chaining)= + +#### Chaining onto a window function + +A built window function stores a concrete frame, with no record of whether you +chose it. When you chain another builder method or `over()` onto one, a frame +equal to the default for its current `order_by` is treated as unset and derived +again from the final `order_by`. Any other frame is kept. + +This rule depends only on the expression, so a copy, a pickled expression sent +to a worker, or one parsed from SQL all chain the same way. The cost is that an +explicit frame that matches the default is indistinguishable from no frame: + +```python +whole = WindowFrame("rows", None, None) # same as the no-order_by default + +# The frame is re-derived, giving a running sum. +f.sum(col("v")).over(Window(window_frame=whole)).order_by(col("v")).build() +``` + +To keep such a frame, set it after the `order_by`, or pass both in one `Window`: + +```python +f.sum(col("v")).over(Window()).order_by(col("v")).window_frame(whole).build() +f.sum(col("v")).over(Window(order_by=col("v"), window_frame=whole)) +``` + ### Null Treatment When using aggregate functions as window functions, it is often useful to specify how null values diff --git a/python/datafusion/expr.py b/python/datafusion/expr.py index 7bd72942b..18ce3554d 100644 --- a/python/datafusion/expr.py +++ b/python/datafusion/expr.py @@ -453,10 +453,6 @@ class Expr: # noqa: PLW1641 :ref:`Expressions` in the online documentation for more information. """ - # Set by ``over()`` when the window frame was given explicitly, so chaining a - # builder method keeps it even if it equals the default frame. - _explicit_window_frame = False - def __init__(self, expr: expr_internal.RawExpr) -> None: """This constructor should not be called by the end user.""" self.expr = expr @@ -1035,7 +1031,7 @@ def order_by(self, *exprs: Expr | SortExpr) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return self._builder(self.expr.order_by, [sort_or_default(e) for e in exprs]) + return ExprFuncBuilder(self.expr.order_by([sort_or_default(e) for e in exprs])) def filter(self, filter: Expr) -> ExprFuncBuilder: """Filter an aggregate function. @@ -1044,7 +1040,7 @@ def filter(self, filter: Expr) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return self._builder(self.expr.filter, filter.expr) + return ExprFuncBuilder(self.expr.filter(filter.expr)) def distinct(self) -> ExprFuncBuilder: """Only evaluate distinct values for an aggregate function. @@ -1053,7 +1049,7 @@ def distinct(self) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return self._builder(self.expr.distinct) + return ExprFuncBuilder(self.expr.distinct()) def null_treatment(self, null_treatment: NullTreatment) -> ExprFuncBuilder: """Set the treatment for ``null`` values for a window or aggregate function. @@ -1062,7 +1058,7 @@ def null_treatment(self, null_treatment: NullTreatment) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return self._builder(self.expr.null_treatment, null_treatment.value) + return ExprFuncBuilder(self.expr.null_treatment(null_treatment.value)) def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: """Set the partitioning for a window function. @@ -1071,7 +1067,7 @@ def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return self._builder(self.expr.partition_by, [e.expr for e in partition_by]) + return ExprFuncBuilder(self.expr.partition_by([e.expr for e in partition_by])) def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: """Set the frame fora window function. @@ -1080,16 +1076,7 @@ def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: set parameters for either window or aggregate functions. If used on any other type of expression, an error will be generated when ``build()`` is called. """ - return ExprFuncBuilder( - self.expr.window_frame(window_frame.window_frame), - explicit_window_frame=True, - ) - - def _builder(self, method: Any, *args: Any) -> ExprFuncBuilder: - keep = self._explicit_window_frame - return ExprFuncBuilder( - method(*args, keep_window_frame=keep), explicit_window_frame=keep - ) + return ExprFuncBuilder(self.expr.window_frame(window_frame.window_frame)) def over(self, window: Window) -> Expr: """Turn an aggregate function into a window function. @@ -1112,19 +1099,14 @@ def over(self, window: Window) -> Expr: window._null_treatment.value if window._null_treatment is not None else None ) - result = Expr( + return Expr( self.expr.over( partition_by=partition_by_raw, order_by=order_by_raw, window_frame=window_frame_raw, null_treatment=null_treatment_raw, - keep_window_frame=self._explicit_window_frame, ) ) - result._explicit_window_frame = ( - window_frame_raw is not None or self._explicit_window_frame - ) - return result def asin(self) -> Expr: """Returns the arc sine or inverse sine of a number.""" @@ -1552,22 +1534,8 @@ def cot(self) -> Expr: class ExprFuncBuilder: - def __init__( - self, - builder: expr_internal.ExprFuncBuilder, - explicit_window_frame: bool = False, - ) -> None: + def __init__(self, builder: expr_internal.ExprFuncBuilder) -> None: self.builder = builder - self._explicit_window_frame = explicit_window_frame - - def _wrap( - self, - builder: expr_internal.ExprFuncBuilder, - explicit_window_frame: bool = False, - ) -> ExprFuncBuilder: - return ExprFuncBuilder( - builder, self._explicit_window_frame or explicit_window_frame - ) def order_by(self, *exprs: Expr) -> ExprFuncBuilder: """Set the ordering for a window or aggregate function. @@ -1575,36 +1543,35 @@ def order_by(self, *exprs: Expr) -> ExprFuncBuilder: Values given in ``exprs`` must be sort expressions. You can convert any other expression to a sort expression using `.sort()`. """ - return self._wrap(self.builder.order_by([sort_or_default(e) for e in exprs])) + return ExprFuncBuilder( + self.builder.order_by([sort_or_default(e) for e in exprs]) + ) def filter(self, filter: Expr) -> ExprFuncBuilder: """Filter values during aggregation.""" - return self._wrap(self.builder.filter(filter.expr)) + return ExprFuncBuilder(self.builder.filter(filter.expr)) def distinct(self) -> ExprFuncBuilder: """Only evaluate distinct values during aggregation.""" - return self._wrap(self.builder.distinct()) + return ExprFuncBuilder(self.builder.distinct()) def null_treatment(self, null_treatment: NullTreatment) -> ExprFuncBuilder: """Set how nulls are treated for either window or aggregate functions.""" - return self._wrap(self.builder.null_treatment(null_treatment.value)) + return ExprFuncBuilder(self.builder.null_treatment(null_treatment.value)) def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: """Set partitioning for window functions.""" - return self._wrap(self.builder.partition_by([e.expr for e in partition_by])) + return ExprFuncBuilder( + self.builder.partition_by([e.expr for e in partition_by]) + ) def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: """Set window frame for window functions.""" - return self._wrap( - self.builder.window_frame(window_frame.window_frame), - explicit_window_frame=True, - ) + return ExprFuncBuilder(self.builder.window_frame(window_frame.window_frame)) def build(self) -> Expr: """Create an expression from a Function Builder.""" - result = Expr(self.builder.build()) - result._explicit_window_frame = self._explicit_window_frame - return result + return Expr(self.builder.build()) class Window: diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index 504436a5e..f146d60ed 100644 --- a/python/tests/test_expr.py +++ b/python/tests/test_expr.py @@ -15,6 +15,8 @@ # specific language governing permissions and limitations # under the License. +import copy +import pickle import re from concurrent.futures import ThreadPoolExecutor from datetime import date, datetime, time, timezone @@ -1321,22 +1323,29 @@ def test_window_builder_keeps_explicit_frame(builder_df): assert result.collect_column("r").to_pylist() == [1, 3, 5, 4] -def test_window_builder_keeps_explicit_default_frame(builder_df): - # An explicit frame equal to the no-order_by default must survive a later - # order_by instead of being re-derived as the running frame. +def test_window_builder_default_frame_same_after_copy_and_pickle(builder_df): + # Whether a frame is re-derived depends only on the expression, so a copy + # or a pickled round trip chains to the same result as the original. window = Window(window_frame=WindowFrame("rows", None, None)) - expr = functions.sum(col("v")).over(window).order_by(col("v")).build() - result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) - assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] + expr = functions.sum(col("v")).over(window) + + def chained(e): + r = e.order_by(col("v")).build().alias("r") + return builder_df.select(col("v"), r).sort(col("v")).collect_column("r") + expected = [1, 3, 6, 10] + assert chained(expr).to_pylist() == expected + assert chained(copy.copy(expr)).to_pylist() == expected + assert chained(pickle.loads(pickle.dumps(expr))).to_pylist() == expected # noqa: S301 -def test_window_builder_keeps_frame_set_on_builder(builder_df): + +def test_window_builder_keeps_default_frame_set_last(builder_df): + frame = WindowFrame("rows", None, None) expr = ( functions.sum(col("v")) .over(Window()) - .window_frame(WindowFrame("rows", None, None)) - .build() .order_by(col("v")) + .window_frame(frame) .build() ) result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) @@ -1353,18 +1362,6 @@ def test_over_keeps_window_function_options(): assert result.collect_column("r").to_pylist() == [3, 3, None, None] -def test_over_keeps_explicit_default_frame(builder_df): - # A frame equal to the no-order_by default, set by an earlier over(), must - # survive a later over() that adds an order_by. - expr = ( - functions.sum(col("v")) - .over(Window(window_frame=WindowFrame("rows", None, None))) - .over(Window(order_by=col("v"))) - ) - result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) - assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] - - def test_window_builder_rederives_default_frame(builder_df): # No order_by means a whole-partition frame; adding one later must switch # to the running frame rather than keep the whole-partition default. From b4cad2d2c30ee6bf9ebf1bf5e8fd8423dc0e0dce Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:05:46 -0400 Subject: [PATCH 18/31] fix: treat the frame an empty order_by derives as a default when chaining build() derives RANGE UNBOUNDED PRECEDING .. CURRENT ROW from an empty order_by list and the whole-partition frame from an absent one, but both are stored as an empty order_by. Only the whole-partition frame was recognised as a default, so a frame derived from order_by=[] looked explicit and survived a later order_by, summing tied rows together. Treat either frame as the default when the stored order_by is empty. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/core/src/expr.rs | 15 +++++++++++---- .../user-guide/common-operations/windows.md | 4 +++- python/tests/test_expr.py | 12 ++++++++++++ 3 files changed, 26 insertions(+), 5 deletions(-) diff --git a/crates/core/src/expr.rs b/crates/core/src/expr.rs index eeb01cbeb..6781e1e03 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -785,10 +785,17 @@ fn builder_from_expr(expr: &Expr) -> ExprFuncBuilder { builder = builder.order_by(params.order_by.clone()); } // A frame equal to the default `build()` derived from the order-by is - // left unset, so it is derived again from the final order-by. - if params.window_frame - != datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true)) - { + // left unset, so it is derived again from the final order-by. An + // absent and an empty order-by derive different frames but are both + // stored as empty, so either frame counts as the default here. + let is_default_frame = if has_order_by { + params.window_frame == datafusion::logical_expr::WindowFrame::new(Some(true)) + } else { + params.window_frame == datafusion::logical_expr::WindowFrame::new(None) + || params.window_frame + == datafusion::logical_expr::WindowFrame::new(Some(false)) + }; + if !is_default_frame { builder = builder.window_frame(params.window_frame.clone()); } if let Some(filter) = ¶ms.filter { diff --git a/docs/source/user-guide/common-operations/windows.md b/docs/source/user-guide/common-operations/windows.md index dfab8f958..2fde84665 100644 --- a/docs/source/user-guide/common-operations/windows.md +++ b/docs/source/user-guide/common-operations/windows.md @@ -142,7 +142,9 @@ df.select( A built window function stores a concrete frame, with no record of whether you chose it. When you chain another builder method or `over()` onto one, a frame equal to the default for its current `order_by` is treated as unset and derived -again from the final `order_by`. Any other frame is kept. +again from the final `order_by`. Any other frame is kept. With no `order_by`, +both the whole-partition frame and `RANGE UNBOUNDED PRECEDING .. CURRENT ROW` +(the frame an empty `order_by` list derives) count as the default. This rule depends only on the expression, so a copy, a pickled expression sent to a worker, or one parsed from SQL all chain the same way. The cost is that an diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index f146d60ed..81b29a479 100644 --- a/python/tests/test_expr.py +++ b/python/tests/test_expr.py @@ -1368,3 +1368,15 @@ def test_window_builder_rederives_default_frame(builder_df): expr = functions.sum(col("v")).over(Window()).order_by(col("v")).build() result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) assert result.collect_column("r").to_pylist() == [1, 3, 6, 10] + + +@pytest.mark.parametrize("first", [Window(), Window(order_by=[])]) +def test_window_builder_rederives_default_frame_from_empty_order_by(first): + # An empty order_by derives a RANGE frame rather than the whole partition; + # both are defaults, so a later order_by gets the ROWS running frame, which + # does not sum the tied ``i`` values together. + ctx = SessionContext() + df = ctx.from_pydict({"i": [1, 1, 2, 3], "v": [1, 2, 3, 4]}) + expr = functions.sum(col("v")).over(first).over(Window(order_by="i")) + result = df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [1, 3, 6, 10] From 92ace16b6ae1323d82b3e153ff7b253ccce3e640 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:10:13 -0400 Subject: [PATCH 19/31] fix: raise when a builder option does not apply to the function kind Seeding the builder from the existing expression always gives it a function kind, so upstream's per-method kind checks no longer applied. partition_by and window_frame on a plain aggregate, and distinct on a non-aggregate window function, built successfully and silently dropped the option instead of raising at build(). Fall back to upstream's ExprFunctionExt method for those cases so build() raises again. distinct on an aggregate run as a window function still keeps the existing options. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/core/src/expr.rs | 25 ++++++++++++++++++++++++- python/tests/test_expr.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 53 insertions(+), 1 deletion(-) diff --git a/crates/core/src/expr.rs b/crates/core/src/expr.rs index 6781e1e03..8c55aca58 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -637,7 +637,21 @@ impl PyExpr { } pub fn distinct(&self) -> PyExprFuncBuilder { - builder_from_expr(&self.expr).distinct().into() + // Only aggregates support DISTINCT, including an aggregate run as a window + // function. For anything else, upstream's empty builder makes `build()` + // raise instead of dropping the option. + let supports_distinct = match &self.expr { + Expr::AggregateFunction(_) => true, + Expr::WindowFunction(window) => { + matches!(window.fun, WindowFunctionDefinition::AggregateUDF(_)) + } + _ => false, + }; + if supports_distinct { + builder_from_expr(&self.expr).distinct().into() + } else { + self.expr.clone().distinct().into() + } } pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder { @@ -648,12 +662,21 @@ impl PyExpr { pub fn partition_by(&self, partition_by: Vec) -> PyExprFuncBuilder { let partition_by = partition_by.iter().map(|e| e.expr.clone()).collect(); + // Window-only option: on an aggregate, upstream's empty builder makes + // `build()` raise instead of dropping it. + if matches!(self.expr, Expr::AggregateFunction(_)) { + return self.expr.clone().partition_by(partition_by).into(); + } builder_from_expr(&self.expr) .partition_by(partition_by) .into() } pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder { + // Window-only option; see `partition_by`. + if matches!(self.expr, Expr::AggregateFunction(_)) { + return self.expr.clone().window_frame(window_frame.into()).into(); + } builder_from_expr(&self.expr) .window_frame(window_frame.into()) .into() diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index 81b29a479..39a02a3c7 100644 --- a/python/tests/test_expr.py +++ b/python/tests/test_expr.py @@ -1380,3 +1380,32 @@ def test_window_builder_rederives_default_frame_from_empty_order_by(first): expr = functions.sum(col("v")).over(first).over(Window(order_by="i")) result = df.select(col("v"), expr.alias("r")).sort(col("v")) assert result.collect_column("r").to_pylist() == [1, 3, 6, 10] + + +@pytest.mark.parametrize( + "builder", + [ + pytest.param( + lambda: functions.sum(col("v")).partition_by(col("g")), + id="partition_by on aggregate", + ), + pytest.param( + lambda: functions.sum(col("v")).window_frame(WindowFrame("rows", 1, 0)), + id="window_frame on aggregate", + ), + pytest.param( + lambda: functions.lead(col("v"), order_by="v").distinct(), + id="distinct on window function", + ), + ], +) +def test_builder_rejects_option_for_other_function_kind(builder): + with pytest.raises(Exception, match="ExprFunctionExt can only be used with"): + builder().build() + + +def test_window_builder_distinct_on_aggregate_window(builder_df): + expr = functions.sum(col("v")).over(Window()).distinct().build() + df = builder_df.select(col("g"), (col("v") % lit(2)).alias("v")) + result = df.select(expr.alias("r")) + assert result.collect_column("r").to_pylist() == [1, 1, 1, 1] From 8e21c8ad2c9a7ab20a184fda01fe9722f3bdee7d Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:13:13 -0400 Subject: [PATCH 20/31] docs: note that over() drops options set on an aggregate Expr.over() on an aggregate drops its order_by, null_treatment, filter, and distinct options (#1764). The new distinct arguments on mean, percentile_cont, quantile_cont, and string_agg run straight into it. Document the limitation and the workaround (pass order_by and null_treatment in the Window, chain filter() or distinct() after over()) in the window functions guide, and point to it from Expr.over. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../user-guide/common-operations/windows.md | 22 +++++++++++++++++++ python/datafusion/expr.py | 4 ++++ 2 files changed, 26 insertions(+) diff --git a/docs/source/user-guide/common-operations/windows.md b/docs/source/user-guide/common-operations/windows.md index 2fde84665..c71930555 100644 --- a/docs/source/user-guide/common-operations/windows.md +++ b/docs/source/user-guide/common-operations/windows.md @@ -225,6 +225,28 @@ df.select( ) ``` +(aggregate_over_options)= + +### Options set on the aggregate + +`over()` currently drops the `order_by`, `null_treatment`, `filter`, and +`distinct` options an aggregate was built with +([apache/datafusion-python#1764](https://github.com/apache/datafusion-python/issues/1764)), +so the window runs as if they were never set: + +```python +# DISTINCT is dropped: averages 1.0, 1.0, 4.0 rather than 1.0, 4.0. +f.avg(col("v"), distinct=True).over(Window()) +``` + +Until that is fixed, pass `order_by` and `null_treatment` in the `Window`, and +chain `filter()` or `distinct()` after `over()`: + +```python +f.avg(col("v")).over(Window()).distinct().build() +f.sum(col("v")).over(Window()).filter(col("v") > lit(1)).build() +``` + ## Available Functions The possible window functions are: diff --git a/python/datafusion/expr.py b/python/datafusion/expr.py index 18ce3554d..92f79ddbd 100644 --- a/python/datafusion/expr.py +++ b/python/datafusion/expr.py @@ -1085,6 +1085,10 @@ def over(self, window: Window) -> Expr: exception of ``partition_by``, how each of the parameters is used is determined by the underlying aggregate function. + On an aggregate, the ``order_by``, ``null_treatment``, ``filter``, and + ``distinct`` options it was built with are dropped; see + :ref:`aggregate_over_options`. + Args: window: Window definition """ From 6b52225a54763b7ee63440c6cc98613f77507129 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:14:10 -0400 Subject: [PATCH 21/31] docs: document that chaining now keeps options already set Add an upgrade-guide section for the behavior change from keeping existing options when chaining builder methods or over(): results can change (lead keeps its partition), chains whose options used to be dropped can now raise (array_agg with distinct and order_by), and first_value, last_value, and nth_value gain RESPECT NULLS in their generated names. Rewrite the Expr.over docstring, which described only aggregates, to cover window functions and point to the frame-chaining rule. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/source/user-guide/upgrade-guides.md | 40 ++++++++++++++++++++++++ python/datafusion/expr.py | 8 +++-- 2 files changed, 46 insertions(+), 2 deletions(-) diff --git a/docs/source/user-guide/upgrade-guides.md b/docs/source/user-guide/upgrade-guides.md index a9f4e3d85..c884bc667 100644 --- a/docs/source/user-guide/upgrade-guides.md +++ b/docs/source/user-guide/upgrade-guides.md @@ -218,6 +218,46 @@ f.bit_and(column("a"), filter=my_filter) # after Passing `filter` to `mean` previously raised a `TypeError`, whether passed positionally or by keyword; it now works. +### Chaining keeps options already set + +Chaining a builder method (`order_by`, `filter`, `distinct`, `null_treatment`, +`partition_by`, `window_frame`) or `over()` onto a function used to start from +an empty builder, so options set by the function's keyword arguments were +silently reset. They are now kept, which can change results: + +```python +e = f.lead(col("v"), 1, partition_by=[col("g")], order_by="t") +e.over(Window(order_by="t")) # before: partition dropped; after: kept +``` + +Options that used to be dropped now take effect, so a chain that ran before +may now raise. For example, DISTINCT requires the ORDER BY expressions to be +among the arguments: + +```python +f.array_agg(col("s"), distinct=True).order_by(col("v")).build() +# before: ran without DISTINCT +# after: Execution error: In an aggregate with DISTINCT, ORDER BY expressions +# must appear in argument list +``` + +Drop `distinct`, or order by the aggregated column, to get either of the +results the chain can actually produce. + +The default `RESPECT NULLS` set by `first_value`, `last_value`, and `nth_value` +is also kept, so their generated column names change: + +```python +f.first_value(col("a")).order_by(col("b")).build() +# before: first_value(a) ORDER BY [b ASC NULLS FIRST] +# after: first_value(a) RESPECT NULLS ORDER BY [b ASC NULLS FIRST] +``` + +This now matches the name from `f.first_value(col("a"), order_by=col("b"))`. +Code that selects the result by its generated name should `alias()` it instead. +How a window frame is handled when chaining is described in +{ref}`window_frame_chaining`. + ### `spark.last_day` renamed its parameter The parameter of {py:func}`datafusion.functions.spark.last_day` is now named diff --git a/python/datafusion/expr.py b/python/datafusion/expr.py index 92f79ddbd..faab23eb3 100644 --- a/python/datafusion/expr.py +++ b/python/datafusion/expr.py @@ -1079,9 +1079,9 @@ def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: return ExprFuncBuilder(self.expr.window_frame(window_frame.window_frame)) def over(self, window: Window) -> Expr: - """Turn an aggregate function into a window function. + """Evaluate an aggregate or window function over a window. - This function turns any aggregate function into a window function. With the + On an aggregate function this turns it into a window function. With the exception of ``partition_by``, how each of the parameters is used is determined by the underlying aggregate function. @@ -1089,6 +1089,10 @@ def over(self, window: Window) -> Expr: ``distinct`` options it was built with are dropped; see :ref:`aggregate_over_options`. + On a window function, each option set in ``window`` replaces the one already + set and the others are kept; see :ref:`window_frame_chaining` for how the + frame is handled. + Args: window: Window definition """ From 0ac494e6058bb70927aacd62a44f6ff5ec47897c Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:20:20 -0400 Subject: [PATCH 22/31] fix: reject a non-bool distinct in string_agg Inserting distinct before filter shifts positional arguments. Most shifted calls fail loudly because an Expr is not a bool, but the old way to pass order_by positionally, string_agg(expr, ",", None, order), still ran: the Rust binding takes Option, so None became distinct=False and the ordering column became the filter. Raise a TypeError for a non-bool distinct that says to pass filter and order_by by keyword. Co-Authored-By: Claude Opus 5.5 (1M context) --- python/datafusion/functions/__init__.py | 6 ++++++ python/tests/test_aggregation.py | 9 +++++++++ 2 files changed, 15 insertions(+) diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index af255c120..a05e331c2 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -7624,6 +7624,12 @@ def string_agg( >>> result.collect_column("s")[0].as_py() 'x,y' """ + if not isinstance(distinct, bool): + msg = ( + f"distinct must be a bool, got {type(distinct).__name__}; " + "pass filter and order_by by keyword" + ) + raise TypeError(msg) order_by_raw = sort_list_to_raw_sort_list(order_by) filter_raw = filter.expr if filter is not None else None diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index 7847c089c..118de3415 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -545,3 +545,12 @@ def test_string_agg(name, expr, result) -> None: } df.show() assert df.collect()[0].to_pydict() == expected + + +@pytest.mark.parametrize("distinct", [None, column("b")]) +def test_string_agg_rejects_non_bool_distinct(distinct) -> None: + # A positional ``None`` placeholder for the old ``filter`` slot now lands + # in ``distinct`` and must raise rather than shift ``order_by`` into + # ``filter``. + with pytest.raises(TypeError, match="distinct must be a bool"): + f.string_agg(column("a"), ",", distinct, column("b")) From 607dde49917134f639ce2568efed8d586c5170a6 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:50:56 -0400 Subject: [PATCH 23/31] fix: reject show_statistics combined with analyze in explain Upstream's LogicalPlan::Analyze has statement-level overrides for the analyze level and categories but none for statistics, so explain(analyze=True, show_statistics=True) silently printed no statistics. SQL rejects the same combination ("EXPLAIN option COSTS cannot be combined with ANALYZE"); raise a ValueError to match. Co-Authored-By: Claude Opus 5.5 (1M context) --- python/datafusion/dataframe.py | 9 ++++++++- python/tests/test_dataframe.py | 7 +++++++ 2 files changed, 15 insertions(+), 1 deletion(-) diff --git a/python/datafusion/dataframe.py b/python/datafusion/dataframe.py index b724935a7..edcaae0e3 100644 --- a/python/datafusion/dataframe.py +++ b/python/datafusion/dataframe.py @@ -1247,13 +1247,17 @@ def explain( format: Output format for the plan. Defaults to :py:attr:`ExplainFormat.INDENT`. show_statistics: If ``True``, include each operator's statistics. - ``None`` uses the ``datafusion.explain.show_statistics`` setting. + Cannot be combined with ``analyze``. ``None`` uses the + ``datafusion.explain.show_statistics`` setting. analyze_level: Which metrics to report with ``analyze``. ``None`` uses the ``datafusion.explain.analyze_level`` setting. analyze_categories: Report only metrics in these categories with ``analyze``; an empty iterable reports none. ``None`` uses the ``datafusion.explain.analyze_categories`` setting. + Raises: + ValueError: If both ``analyze`` and ``show_statistics`` are ``True``. + Examples: Show the plan in tree format: @@ -1273,6 +1277,9 @@ def explain( ... analyze=True, analyze_categories=[ExplainMetricCategory.ROWS] ... ) # doctest: +SKIP """ + if analyze and show_statistics: + msg = "show_statistics cannot be combined with analyze" + raise ValueError(msg) fmt = format.value if format is not None else None level = analyze_level.value if analyze_level is not None else None categories = ( diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index 590c5cd6a..9ff5e78ff 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -3934,6 +3934,13 @@ def test_explain_options(capsys, kwargs, present, absent): assert text not in out +def test_explain_rejects_show_statistics_with_analyze(capsys): + # Upstream's Analyze plan has no statistics override, so the flag would be + # silently ignored; SQL rejects the same combination. + with pytest.raises(ValueError, match="cannot be combined with analyze"): + _explain_output(capsys, analyze=True, show_statistics=True) + + @pytest.mark.parametrize( ("window_exprs", "expected_columns"), [ From 4933fa28e8f952b42bbbc5b5e64836ec7ae0c999 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 08:52:44 -0400 Subject: [PATCH 24/31] fix: treat a bare str as a column name in spark.printf pyspark's printf(format: ColumnOrName, *cols: ColumnOrName) reads a bare str as a column name, unlike format_string, whose format is a plain str template. As an alias of format_string, spark.printf("a", ...) returned the literal text of the column name instead of formatting with it. Give printf its own body that resolves str arguments as column names, as bit_get already does for pos. Co-Authored-By: Claude Opus 5.5 (1M context) --- python/datafusion/functions/spark.py | 23 +++++++++++++++++++---- python/tests/test_spark_functions.py | 9 ++++++++- 2 files changed, 27 insertions(+), 5 deletions(-) diff --git a/python/datafusion/functions/spark.py b/python/datafusion/functions/spark.py index 25ff8190e..c4002670d 100644 --- a/python/datafusion/functions/spark.py +++ b/python/datafusion/functions/spark.py @@ -1711,13 +1711,28 @@ def format_string(format: str | Expr, *cols: Expr) -> Expr: return Expr(_f.format_string(fmt_expr.expr, *[c.expr for c in cols])) -def printf(format: str | Expr, *cols: Expr) -> Expr: +def printf(format: Expr | str, *cols: Expr | str) -> Expr: """Spark ``printf``: printf-style format string. - See Also: - This is an alias for :py:func:`format_string`. + Unlike :py:func:`format_string`, a bare ``str`` ``format`` or ``cols`` entry + is treated as a column name (matching pyspark), not a literal; pass + :func:`~datafusion.lit` for a literal format. + + Examples: + >>> ctx = dfn.SessionContext() + >>> df = ctx.from_pydict({"a": ["aa%d%s"], "b": [123], "c": ["cc"]}) + >>> r = df.select(dfn.functions.spark.printf("a", "b", "c").alias("v")) + >>> r.collect_column("v")[0].as_py() + 'aa123cc' + + >>> r = df.select( + ... dfn.functions.spark.printf(dfn.lit("%d-%s"), "b", "c").alias("v")) + >>> r.collect_column("v")[0].as_py() + '123-cc' """ - return format_string(format, *cols) + return Expr( + _f.format_string(_to_raw_expr(format), *[_to_raw_expr(c) for c in cols]) + ) def space(col: Expr | int) -> Expr: diff --git a/python/tests/test_spark_functions.py b/python/tests/test_spark_functions.py index 67487d0da..c1189cbf1 100644 --- a/python/tests/test_spark_functions.py +++ b/python/tests/test_spark_functions.py @@ -515,7 +515,6 @@ def test_sql_concat_semantics_override(): (spark.datepart, spark.date_part, lambda: ("YEAR", _ts())), (spark.sha, spark.sha1, lambda: (lit("abc"),)), (spark.ceiling, spark.ceil, lambda: (lit(1.2),)), - (spark.printf, spark.format_string, lambda: ("%d-%s", lit(42), lit("hi"))), (spark.char_length, spark.length, lambda: (lit("hello"),)), (spark.character_length, spark.length, lambda: (lit("hello"),)), (spark.power, spark.pow, lambda: (lit(2), lit(3))), @@ -526,6 +525,14 @@ def test_aliases_match_primary(df, alias_fn, primary_fn, args): assert _val(df, alias_fn(*args())) == _val(df, primary_fn(*args())) +def test_printf_str_is_column_name(): + ctx = SessionContext() + df = ctx.from_pydict({"a": ["aa%d%s"], "b": [123], "c": ["cc"]}) + assert _val(df, spark.printf("a", "b", "c")) == "aa123cc" + assert _val(df, spark.printf(col("a"), col("b"), col("c"))) == "aa123cc" + assert _val(df, spark.printf(lit("%d-%s"), lit(42), lit("hi"))) == "42-hi" + + def test_substr_without_len(df): assert _val(df, spark.substr(lit("hello"), 2)) == "ello" assert _val(df, spark.substr(lit("hello"), -3)) == "llo" From 23ed533dba87b0ae0bcd5c42913b77e73a45f7bc Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 09:11:44 -0400 Subject: [PATCH 25/31] fix: name the expected and found capsule when importing a UDF Now that udf, udaf, and udwf accept a bare PyCapsule, passing the wrong kind is easy, and pointer_checked alone surfaces CPython's fixed "PyCapsule_GetPointer called with incorrect name". Call validate_pycapsule first in the scalar, window, and aggregate capsule imports, as its doc comment asks of every extraction site, so the error names both the expected and the actual capsule. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/core/src/udaf.rs | 5 ++++- crates/core/src/udf.rs | 5 ++++- crates/core/src/udwf.rs | 5 ++++- python/tests/test_udf.py | 20 +++++++++++++++++++- 4 files changed, 31 insertions(+), 4 deletions(-) diff --git a/crates/core/src/udaf.rs b/crates/core/src/udaf.rs index 6a2675193..96583e694 100644 --- a/crates/core/src/udaf.rs +++ b/crates/core/src/udaf.rs @@ -28,7 +28,9 @@ use datafusion::logical_expr::{ Accumulator, AccumulatorFactoryFunction, AggregateUDF, AggregateUDFImpl, Signature, Volatility, }; use datafusion_ffi::udaf::FFI_AggregateUDF; -use datafusion_python_util::{CapsuleGetterArg, call_capsule_getter, parse_volatility}; +use datafusion_python_util::{ + CapsuleGetterArg, call_capsule_getter, parse_volatility, validate_pycapsule, +}; use pyo3::prelude::*; use pyo3::types::{PyCapsule, PyTuple}; @@ -310,6 +312,7 @@ impl AggregateUDFImpl for PythonFunctionAggregateUDF { } fn aggregate_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult { + validate_pycapsule(capsule, "datafusion_aggregate_udf")?; let data: NonNull = capsule .pointer_checked(Some(c"datafusion_aggregate_udf"))? .cast(); diff --git a/crates/core/src/udf.rs b/crates/core/src/udf.rs index 5093a564f..1b44bffac 100644 --- a/crates/core/src/udf.rs +++ b/crates/core/src/udf.rs @@ -31,7 +31,9 @@ use datafusion::logical_expr::{ Volatility, }; use datafusion_ffi::udf::FFI_ScalarUDF; -use datafusion_python_util::{CapsuleGetterArg, call_capsule_getter, parse_volatility}; +use datafusion_python_util::{ + CapsuleGetterArg, call_capsule_getter, parse_volatility, validate_pycapsule, +}; use pyo3::prelude::*; use pyo3::types::{PyCapsule, PyTuple}; @@ -210,6 +212,7 @@ impl ScalarUDFImpl for PythonFunctionScalarUDF { } fn scalar_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult { + validate_pycapsule(capsule, "datafusion_scalar_udf")?; let data: NonNull = capsule .pointer_checked(Some(c"datafusion_scalar_udf"))? .cast(); diff --git a/crates/core/src/udwf.rs b/crates/core/src/udwf.rs index af3871bb6..b9b7089f2 100644 --- a/crates/core/src/udwf.rs +++ b/crates/core/src/udwf.rs @@ -30,7 +30,9 @@ use datafusion::logical_expr::{ }; use datafusion::scalar::ScalarValue; use datafusion_ffi::udwf::FFI_WindowUDF; -use datafusion_python_util::{CapsuleGetterArg, call_capsule_getter, parse_volatility}; +use datafusion_python_util::{ + CapsuleGetterArg, call_capsule_getter, parse_volatility, validate_pycapsule, +}; use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use pyo3::types::{PyCapsule, PyList, PyTuple}; @@ -217,6 +219,7 @@ pub fn to_rust_partition_evaluator(evaluator: Py) -> PartitionEvaluatorFa } fn window_udf_from_capsule(capsule: &Bound<'_, PyCapsule>) -> PyDataFusionResult { + validate_pycapsule(capsule, "datafusion_window_udf")?; let data: NonNull = capsule .pointer_checked(Some(c"datafusion_window_udf"))? .cast(); diff --git a/python/tests/test_udf.py b/python/tests/test_udf.py index 3a41fa6e1..c2af1907b 100644 --- a/python/tests/test_udf.py +++ b/python/tests/test_udf.py @@ -20,7 +20,7 @@ import pyarrow as pa import pyarrow.compute as pc import pytest -from datafusion import SessionContext, column, udf +from datafusion import SessionContext, column, udaf, udf, udwf from datafusion import functions as f @@ -278,3 +278,21 @@ def non_nullable_abs(input_col): with pytest.raises(Exception) as e_info: _results = df_result.collect() assert "Invalid argument error" in str(e_info) + + +@pytest.mark.parametrize( + ("decorator", "expected"), + [ + (udf, "datafusion_scalar_udf"), + (udaf, "datafusion_aggregate_udf"), + (udwf, "datafusion_window_udf"), + ], +) +def test_wrong_capsule_kind_names_expected_and_found(decorator, expected): + capsule = SessionContext().__datafusion_logical_extension_codec__() + with pytest.raises( + ValueError, + match=f"Expected name '{expected}' in PyCapsule, " + "instead got 'datafusion_logical_extension_codec'", + ): + decorator(capsule) From 8454a0998fe1d7445d9f124f5eeb361f637faa29 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 09:18:48 -0400 Subject: [PATCH 26/31] fix: treat an empty fill_null or fill_nan subset as no columns Both bindings collect subset into a Vec, and upstream reads an empty column list as "all columns", so subset=[] behaved like subset=None. A subset computed from the schema that happened to match nothing then rewrote every column (every float column for fill_nan). Return the DataFrame unchanged for an empty subset, matching pyspark's na.fill, and note the fill_null change in the upgrade guide. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/source/user-guide/upgrade-guides.md | 13 +++++++++++++ python/datafusion/dataframe.py | 15 +++++++++++++-- python/tests/test_dataframe.py | 10 ++++++++++ 3 files changed, 36 insertions(+), 2 deletions(-) diff --git a/docs/source/user-guide/upgrade-guides.md b/docs/source/user-guide/upgrade-guides.md index c884bc667..749b31484 100644 --- a/docs/source/user-guide/upgrade-guides.md +++ b/docs/source/user-guide/upgrade-guides.md @@ -258,6 +258,19 @@ Code that selects the result by its generated name should `alias()` it instead. How a window frame is handled when chaining is described in {ref}`window_frame_chaining`. +### `fill_null(subset=[])` fills no columns + +{py:meth}`~datafusion.dataframe.DataFrame.fill_null` with an empty `subset` +list used to fill every column, the same as `subset=None`. It now returns the +DataFrame unchanged, so a subset computed from the schema that matches nothing +no longer rewrites every column. The new +{py:meth}`~datafusion.dataframe.DataFrame.fill_nan` behaves the same way. + +```python +df.fill_null(0, subset=[]) # before: fills all columns; after: fills none +df.fill_null(0) # fills all columns, before and after +``` + ### `spark.last_day` renamed its parameter The parameter of {py:func}`datafusion.functions.spark.last_day` is now named diff --git a/python/datafusion/dataframe.py b/python/datafusion/dataframe.py index edcaae0e3..c41c66434 100644 --- a/python/datafusion/dataframe.py +++ b/python/datafusion/dataframe.py @@ -1911,7 +1911,8 @@ def fill_null(self, value: Any, subset: list[str] | None = None) -> DataFrame: Args: value: Value to replace nulls with. Will be cast to match column type. - subset: Optional list of column names to fill. If None, fills all columns. + subset: Optional list of column names to fill. If None, fills all columns; + an empty list fills none. Returns: DataFrame with null values replaced where type casting is possible @@ -1924,11 +1925,16 @@ def fill_null(self, value: Any, subset: list[str] | None = None) -> DataFrame: >>> filled.sort(col("a")).collect()[0].column("a").to_pylist() [0, 1, 3] + >>> df.fill_null(0, subset=[]).to_pydict() + {'a': [1, None, 3], 'b': [None, 5, 6]} + Notes: - Only fills nulls in columns where the value can be cast to the column type - For columns where casting fails, the original column is kept unchanged - For columns not in subset, the original column is kept unchanged """ + if subset is not None and not subset: + return self return DataFrame(self.df.fill_null(value, subset)) def fill_nan(self, value: float, subset: list[str] | None = None) -> DataFrame: @@ -1941,7 +1947,7 @@ def fill_nan(self, value: float, subset: list[str] | None = None) -> DataFrame: Args: value: Value to replace NaN with. Will be cast to match column type. subset: Optional list of column names to fill. If None, fills all - floating-point columns. + floating-point columns; an empty list fills none. Returns: DataFrame with NaN values replaced. @@ -1956,7 +1962,12 @@ def fill_nan(self, value: float, subset: list[str] | None = None) -> DataFrame: >>> df.fill_nan(0.0, subset=["a"]).collect_column("b")[0].as_py() nan + + >>> df.fill_nan(0.0, subset=[]).collect_column("b")[0].as_py() + nan """ + if subset is not None and not subset: + return self return DataFrame(self.df.fill_nan(value, subset)) diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index 9ff5e78ff..7b56461b6 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -3429,6 +3429,10 @@ def test_fill_null_specific_types(null_df): ] +def test_fill_null_empty_subset_fills_nothing(null_df): + assert null_df.fill_null(0, subset=[]).to_pydict() == null_df.to_pydict() + + def test_fill_null_immutability(null_df): """Test that original DataFrame is unchanged after fill_null.""" # Get original values with nulls @@ -3517,6 +3521,12 @@ def test_fill_nan_subset(ctx): assert _is_nan(result["f64"][1]) +def test_fill_nan_empty_subset_fills_nothing(ctx): + result = _nan_df(ctx).fill_nan(0.0, subset=[]).to_pydict() + assert _is_nan(result["f64"][1]) + assert _is_nan(result["f32"][0]) + + def test_fill_nan_preserves_schema(ctx): df = _nan_df(ctx) assert df.fill_nan(0.0).schema() == df.schema() From 1b0184562ceec7c2c3d6a95c36aca246f0fd058f Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 09:23:42 -0400 Subject: [PATCH 27/31] fix: keep the sort direction in percentile_cont and friends Upstream's percentile_cont(sort, p) stores the sort as the aggregate's WITHIN GROUP ordering, but add_builder_fns_to_aggregate starts from an empty builder, so build() overwrote it. A descending sort was silently ignored in percentile_cont, quantile_cont, approx_percentile_cont, and approx_percentile_cont_with_weight (2.0 instead of SQL's 4.0 for the 0.25 percentile of 1..5). Pass the sort through as order_by. Generated column names now include the ordering, as they do from SQL; note that in the upgrade guide. Correct the docstrings, which said these functions ignore order_by: a chained order_by replaces the ordering of sort_expression. Co-Authored-By: Claude Opus 5.5 (1M context) --- crates/core/src/functions.rs | 20 ++++++--- docs/source/user-guide/upgrade-guides.md | 23 ++++++++++ python/datafusion/functions/__init__.py | 9 ++-- python/tests/test_aggregation.py | 54 ++++++++++++++++++++++++ 4 files changed, 96 insertions(+), 10 deletions(-) diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index 79c0524e7..29c4aa377 100644 --- a/crates/core/src/functions.rs +++ b/crates/core/src/functions.rs @@ -780,12 +780,14 @@ pub fn approx_percentile_cont( filter: Option, ) -> PyDataFusionResult { let agg_fn = functions_aggregate::expr_fn::approx_percentile_cont( - sort_expression.sort, + sort_expression.sort.clone(), lit(percentile), num_centroids.map(lit), ); - add_builder_fns_to_aggregate(agg_fn, None, filter, None, None) + // The builder starts empty, so the WITHIN GROUP ordering upstream stored + // must be passed again or `build()` drops its direction. + add_builder_fns_to_aggregate(agg_fn, None, filter, Some(vec![sort_expression]), None) } #[pyfunction] @@ -798,13 +800,14 @@ pub fn approx_percentile_cont_with_weight( filter: Option, ) -> PyDataFusionResult { let agg_fn = functions_aggregate::expr_fn::approx_percentile_cont_with_weight( - sort_expression.sort, + sort_expression.sort.clone(), weight.expr, lit(percentile), num_centroids.map(lit), ); - add_builder_fns_to_aggregate(agg_fn, None, filter, None, None) + // See `approx_percentile_cont`. + add_builder_fns_to_aggregate(agg_fn, None, filter, Some(vec![sort_expression]), None) } #[pyfunction] @@ -815,10 +818,13 @@ pub fn percentile_cont( distinct: Option, filter: Option, ) -> PyDataFusionResult { - let agg_fn = - functions_aggregate::expr_fn::percentile_cont(sort_expression.sort, lit(percentile)); + let agg_fn = functions_aggregate::expr_fn::percentile_cont( + sort_expression.sort.clone(), + lit(percentile), + ); - add_builder_fns_to_aggregate(agg_fn, distinct, filter, None, None) + // See `approx_percentile_cont`. + add_builder_fns_to_aggregate(agg_fn, distinct, filter, Some(vec![sort_expression]), None) } // We handle last_value explicitly because the signature expects an order_by diff --git a/docs/source/user-guide/upgrade-guides.md b/docs/source/user-guide/upgrade-guides.md index 749b31484..e2faec5b8 100644 --- a/docs/source/user-guide/upgrade-guides.md +++ b/docs/source/user-guide/upgrade-guides.md @@ -258,6 +258,29 @@ Code that selects the result by its generated name should `alias()` it instead. How a window frame is handled when chaining is described in {ref}`window_frame_chaining`. +### Percentile functions keep the sort direction + +{py:func}`~datafusion.functions.percentile_cont`, +{py:func}`~datafusion.functions.quantile_cont`, +{py:func}`~datafusion.functions.approx_percentile_cont`, and +{py:func}`~datafusion.functions.approx_percentile_cont_with_weight` ignored the +direction of `sort_expression`, so a descending sort gave the ascending result. +They now match `WITHIN GROUP (ORDER BY ... DESC)` in SQL: + +```python +f.percentile_cont(col("a").sort(ascending=False), 0.25) +# a = [1, 2, 3, 4, 5]; before: 2.0, after: 4.0 +``` + +Their generated column names now include the ordering, in the same form as SQL. +Code that selects the result by its generated name should `alias()` it instead. + +```python +f.percentile_cont(col("a"), 0.25) +# before: percentile_cont(t.a,Float64(0.25)) +# after: percentile_cont(Float64(0.25)) WITHIN GROUP [t.a ASC NULLS FIRST] +``` + ### `fill_null(subset=[])` fills no columns {py:meth}`~datafusion.dataframe.DataFrame.fill_null` with an empty `subset` diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index a05e331c2..3dee48b48 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -5444,7 +5444,8 @@ def approx_percentile_cont( setting the ``num_centroids`` parameter. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``order_by``, ``null_treatment``, and ``distinct``. + the options ``null_treatment`` and ``distinct``, and ``order_by`` replaces the + ordering of ``sort_expression``. Args: sort_expression: Values for which to find the approximate percentile @@ -5493,7 +5494,8 @@ def approx_percentile_cont_with_weight( it uses the associated associated weights. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``order_by``, ``null_treatment``, and ``distinct``. + the options ``null_treatment`` and ``distinct``, and ``order_by`` replaces the + ordering of ``sort_expression``. Args: sort_expression: Values for which to find the approximate percentile @@ -5546,7 +5548,8 @@ def percentile_cont( percentile value rather than an approximation. If using the builder functions described in ref:`_aggregation` this function ignores - the options ``order_by`` and ``null_treatment``. + the option ``null_treatment``, and ``order_by`` replaces the ordering of + ``sort_expression``. Args: sort_expression: Values for which to find the percentile diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index 118de3415..34545ed75 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -554,3 +554,57 @@ def test_string_agg_rejects_non_bool_distinct(distinct) -> None: # ``filter``. with pytest.raises(TypeError, match="distinct must be a bool"): f.string_agg(column("a"), ",", distinct, column("b")) + + +@pytest.mark.parametrize( + ("expr", "sql"), + [ + pytest.param( + lambda s: f.percentile_cont(s, 0.25), + "percentile_cont(0.25)", + id="percentile_cont", + ), + pytest.param( + lambda s: f.quantile_cont(s, 0.25), + "quantile_cont(0.25)", + id="quantile_cont", + ), + pytest.param( + lambda s: f.approx_percentile_cont(s, 0.25), + "approx_percentile_cont(0.25)", + id="approx_percentile_cont", + ), + pytest.param( + lambda s: f.approx_percentile_cont_with_weight(s, lit(1.0), 0.25), + "approx_percentile_cont_with_weight(1.0, 0.25)", + id="approx_percentile_cont_with_weight", + ), + ], +) +def test_percentile_keeps_sort_direction(expr, sql) -> None: + ctx = SessionContext() + df = ctx.from_pydict({"a": [1.0, 2.0, 3.0, 4.0, 5.0]}, name="t") + desc = column("a").sort(ascending=False) + + result = df.aggregate([], [expr(desc).alias("p")]).collect_column("p") + expected = ctx.sql( + f"SELECT {sql} WITHIN GROUP (ORDER BY a DESC) AS p FROM t" + ).collect_column("p") + assert result.to_pylist() == expected.to_pylist() + assert ( + result.to_pylist() + != df.aggregate([], [expr(column("a")).alias("p")]) + .collect_column("p") + .to_pylist() + ) + + +def test_percentile_cont_order_by_replaces_sort() -> None: + ctx = SessionContext() + df = ctx.from_pydict({"a": [1.0, 2.0, 3.0, 4.0, 5.0]}) + expr = ( + f.percentile_cont(column("a"), 0.25) + .order_by(column("a").sort(ascending=False)) + .build() + ) + assert df.aggregate([], [expr.alias("p")]).collect_column("p")[0].as_py() == 4.0 From f8d7b70fe14448a4ce37494b37f8c1d3943585ca Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 09:25:31 -0400 Subject: [PATCH 28/31] fix: allow the keyword-only form of the udf, udaf, and udwf decorators The capsule check read args[0] before the "if args" guard that follows it, so calling a decorator with only keyword arguments, such as udf(input_fields=..., return_field=..., volatility=...), raised IndexError instead of returning the decorator. Guard the capsule check on args being non-empty, as udtf already does. Co-Authored-By: Claude Opus 5.5 (1M context) --- python/datafusion/user_defined.py | 12 +++++++++--- python/tests/test_udaf.py | 14 ++++++++++++++ python/tests/test_udf.py | 9 +++++++++ python/tests/test_udwf.py | 10 ++++++++++ 4 files changed, 42 insertions(+), 3 deletions(-) diff --git a/python/datafusion/user_defined.py b/python/datafusion/user_defined.py index 92d541591..763553a2e 100644 --- a/python/datafusion/user_defined.py +++ b/python/datafusion/user_defined.py @@ -397,7 +397,9 @@ def wrapper(*args: Any, **kwargs: Any) -> Callable: return decorator - if hasattr(args[0], "__datafusion_scalar_udf__") or _is_pycapsule(args[0]): + if args and ( + hasattr(args[0], "__datafusion_scalar_udf__") or _is_pycapsule(args[0]) + ): return ScalarUDF.from_pycapsule(args[0]) if args and callable(args[0]): @@ -747,7 +749,9 @@ def wrapper(*args: Any, **kwargs: Any) -> Expr: return decorator - if hasattr(args[0], "__datafusion_aggregate_udf__") or _is_pycapsule(args[0]): + if args and ( + hasattr(args[0], "__datafusion_aggregate_udf__") or _is_pycapsule(args[0]) + ): return AggregateUDF.from_pycapsule(args[0]) if args and callable(args[0]): @@ -1098,7 +1102,9 @@ def udwf(*args: Any, **kwargs: Any): # noqa: D417 Returns: A user-defined window function that can be used in window function calls. """ - if hasattr(args[0], "__datafusion_window_udf__") or _is_pycapsule(args[0]): + if args and ( + hasattr(args[0], "__datafusion_window_udf__") or _is_pycapsule(args[0]) + ): return WindowUDF.from_pycapsule(args[0]) if args and callable(args[0]): diff --git a/python/tests/test_udaf.py b/python/tests/test_udaf.py index 8cd480e37..167720312 100644 --- a/python/tests/test_udaf.py +++ b/python/tests/test_udaf.py @@ -168,6 +168,20 @@ def summarize(): assert result.column(0) == pa.array([1.0 + 2.0 + 3.0]) +def test_udaf_decorator_keyword_arguments(df): + @udaf( + input_types=pa.float64(), + return_type=pa.float64(), + state_type=[pa.float64()], + volatility="immutable", + ) + def summarize(): + return Summarize() + + result = df.aggregate([], [summarize(column("a"))]).collect()[0] + assert result.column(0) == pa.array([1.0 + 2.0 + 3.0]) + + @pytest.mark.parametrize("as_scalar", [True, False]) def test_udaf_aggregate_with_arguments(df, as_scalar): bias = 10.0 diff --git a/python/tests/test_udf.py b/python/tests/test_udf.py index c2af1907b..e67f70613 100644 --- a/python/tests/test_udf.py +++ b/python/tests/test_udf.py @@ -59,6 +59,15 @@ def is_null(x: pa.Array) -> pa.Array: assert result == pa.array([False, False, True]) +def test_udf_decorator_keyword_arguments(df): + @udf(input_fields=[pa.int64()], return_field=pa.bool_(), volatility="immutable") + def is_null(x: pa.Array) -> pa.Array: + return x.is_null() + + result = df.select(is_null(column("b"))).collect()[0].column(0) + assert result == pa.array([False, False, True]) + + def test_register_udf(ctx, df) -> None: is_null = udf( lambda x: x.is_null(), diff --git a/python/tests/test_udwf.py b/python/tests/test_udwf.py index 38b935b7e..eb66f2ada 100644 --- a/python/tests/test_udwf.py +++ b/python/tests/test_udwf.py @@ -466,3 +466,13 @@ def test_udwf_named_function(ctx, count_window_df): FOLLOWING) FROM test_table""" ).collect()[0] assert result.column(0) == pa.array([0, 1, 2]) + + +def test_udwf_decorator_keyword_arguments(ctx): + @udwf(input_types=[pa.int64()], return_type=pa.int64(), volatility="immutable") + def window_count() -> WindowEvaluator: + return SimpleWindowCount() + + df = ctx.from_pydict({"a": [1, 2, 3]}) + result = df.select(window_count(column("a")).alias("c")).collect_column("c") + assert result.to_pylist() == [0, 1, 2] From 72598193d3cdf40d14a016d2a7c103537ccc00eb Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 12:04:22 -0400 Subject: [PATCH 29/31] docs: fix the broken guide links in aggregate docstrings The aggregate docstrings linked the aggregation guide as ref:`_aggregation`, with no leading colon and a leading underscore on the label, so it rendered as plain italic text instead of a link. The pattern came in with the builder parameters and was copied into each new aggregate since; the NullTreatment docstring and the window function note had the same mistake with _window_functions. Use :ref:`aggregation` and :ref:`window_functions`, matching the labels in the guide and the module docstring that already links them. Co-Authored-By: Claude Opus 5.5 (1M context) --- python/datafusion/common.py | 2 +- python/datafusion/functions/__init__.py | 80 ++++++++++++------------- 2 files changed, 41 insertions(+), 41 deletions(-) diff --git a/python/datafusion/common.py b/python/datafusion/common.py index c689a816d..a3c6b3960 100644 --- a/python/datafusion/common.py +++ b/python/datafusion/common.py @@ -61,7 +61,7 @@ class NullTreatment(Enum): This is used primarily by aggregate and window functions. It can be set on these functions using the builder approach described in - ref:`_window_functions` and ref:`_aggregation` in the online documentation. + :ref:`window_functions` and :ref:`aggregation` in the online documentation. """ diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 3dee48b48..1f8e17791 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -958,7 +958,7 @@ def count_star(filter: Expr | None = None) -> Expr: This aggregate function will count all of the rows in the partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``distinct``, and ``null_treatment``. Args: @@ -5360,7 +5360,7 @@ def approx_distinct( will approximate the number of distinct entries. It may return significantly faster than :py:func:`count` for some DataFrames. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -5396,7 +5396,7 @@ def approx_median(expression: Expr, filter: Expr | None = None) -> Expr: This aggregate function is similar to :py:func:`median`, but it will only approximate the median. It may return significantly faster for some DataFrames. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``, and ``distinct``. Args: @@ -5443,7 +5443,7 @@ def approx_percentile_cont( compute the percentile. You can limit the number of bins used in this algorithm by setting the ``num_centroids`` parameter. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``null_treatment`` and ``distinct``, and ``order_by`` replaces the ordering of ``sort_expression``. @@ -5493,7 +5493,7 @@ def approx_percentile_cont_with_weight( This aggregate function is similar to :py:func:`approx_percentile_cont` except that it uses the associated associated weights. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``null_treatment`` and ``distinct``, and ``order_by`` replaces the ordering of ``sort_expression``. @@ -5547,7 +5547,7 @@ def percentile_cont( Unlike :py:func:`approx_percentile_cont`, this function computes the exact percentile value rather than an approximation. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the option ``null_treatment``, and ``order_by`` replaces the ordering of ``sort_expression``. @@ -5620,7 +5620,7 @@ def array_agg( consider :py:func:`array_sort` after aggregation. [Issue Tracker](https://github.com/apache/datafusion/issues/12371) - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the option ``null_treatment``. Args: @@ -5734,7 +5734,7 @@ def avg( This aggregate function expects a numeric expression and will return a float. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -5777,7 +5777,7 @@ def corr(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: This aggregate function expects both values to be numeric and will return a float. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -5816,7 +5816,7 @@ def count( This aggregate function will count the non-null rows provided in the expression. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -5860,7 +5860,7 @@ def covar_pop(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: This aggregate function expects both values to be numeric and will return a float. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -5901,7 +5901,7 @@ def covar_samp(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr This aggregate function expects both values to be numeric and will return a float. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -5943,7 +5943,7 @@ def covar(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: def max(expression: Expr, filter: Expr | None = None) -> Expr: """Aggregate function that returns the maximum value of the argument. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -5993,7 +5993,7 @@ def median( This aggregate function returns the median value of the expression for the given aggregate function. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -6027,7 +6027,7 @@ def median( def min(expression: Expr, filter: Expr | None = None) -> Expr: """Aggregate function that returns the minimum value of the argument. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6065,7 +6065,7 @@ def sum( This aggregate function expects a numeric expression. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -6106,7 +6106,7 @@ def sum( def stddev(expression: Expr, filter: Expr | None = None) -> Expr: """Computes the standard deviation of the argument. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6138,7 +6138,7 @@ def stddev(expression: Expr, filter: Expr | None = None) -> Expr: def stddev_pop(expression: Expr, filter: Expr | None = None) -> Expr: """Computes the population standard deviation of the argument. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6191,7 +6191,7 @@ def var(expression: Expr, filter: Expr | None = None) -> Expr: def var_pop(expression: Expr, filter: Expr | None = None) -> Expr: """Computes the population variance of the argument. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6232,7 +6232,7 @@ def var_population(expression: Expr, filter: Expr | None = None) -> Expr: def var_samp(expression: Expr, filter: Expr | None = None) -> Expr: """Computes the sample variance of the argument. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6280,7 +6280,7 @@ def regr_avgx( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6321,7 +6321,7 @@ def regr_avgy( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6362,7 +6362,7 @@ def regr_count( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6403,7 +6403,7 @@ def regr_intercept( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6446,7 +6446,7 @@ def regr_r2( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6487,7 +6487,7 @@ def regr_slope( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6528,7 +6528,7 @@ def regr_sxx( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6569,7 +6569,7 @@ def regr_sxy( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6610,7 +6610,7 @@ def regr_syy( This is a linear regression aggregate function. Only non-null pairs of the inputs are evaluated. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6651,7 +6651,7 @@ def first_value( This aggregate function will return the first value in the partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the option ``distinct``. Args: @@ -6707,7 +6707,7 @@ def last_value( This aggregate function will return the last value in the partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the option ``distinct``. Args: @@ -6764,7 +6764,7 @@ def nth_value( This aggregate function will return the n-th value in the partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the option ``distinct``. Args: @@ -6817,7 +6817,7 @@ def any_value(expression: Expr, filter: Expr | None = None) -> Expr: Returns NULL if every value in the group is NULL. Which value is returned is not specified and may differ between runs. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -6854,7 +6854,7 @@ def bit_and( This aggregate function will bitwise compare every value in the input partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -6892,7 +6892,7 @@ def bit_or( This aggregate function will bitwise compare every value in the input partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -6932,7 +6932,7 @@ def bit_xor( This aggregate function will bitwise compare every value in the input partition. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by`` and ``null_treatment``. Args: @@ -6971,7 +6971,7 @@ def bool_and(expression: Expr, filter: Expr | None = None) -> Expr: This aggregate function will compare every value in the input partition. These are expected to be boolean values. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -7010,7 +7010,7 @@ def bool_or(expression: Expr, filter: Expr | None = None) -> Expr: This aggregate function will compare every value in the input partition. These are expected to be boolean values. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: @@ -7071,7 +7071,7 @@ def lead( +--------+------+-----+ To set window function parameters use the window builder approach described in the - ref:`_window_functions` online documentation. + :ref:`window_functions` online documentation. Args: arg: Value to return @@ -7589,7 +7589,7 @@ def string_agg( separating them with the specified delimiter. Non-string values will be converted to their string equivalents. - If using the builder functions described in ref:`_aggregation` this function ignores + If using the builder functions described in :ref:`aggregation` this function ignores the option ``null_treatment``. Args: From d7795029b756c1e8d84352b58e6c01455e136a6e Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 12:08:31 -0400 Subject: [PATCH 30/31] test: check that bit_and and bit_or keep distinct in the expression AND and OR ignore duplicate inputs, so the existing result checks for bit_and(distinct=True) and bit_or(distinct=True) pass whether or not the argument reaches the plan. Assert on the built expression's canonical name as well. Co-Authored-By: Claude Opus 5.5 (1M context) --- python/tests/test_aggregation.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index 34545ed75..bf418b142 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -394,6 +394,14 @@ def test_bit_and_bool_fns(df, name, expr, result): assert df.collect()[0].to_pydict() == expected +@pytest.mark.parametrize("fn", [f.bit_and, f.bit_or]) +def test_bitwise_distinct_is_kept(fn): + # AND and OR ignore duplicates, so the result alone cannot show whether + # ``distinct`` reached the plan. + assert "DISTINCT" in fn(column("b"), distinct=True).canonical_name() + assert "DISTINCT" not in fn(column("b")).canonical_name() + + @pytest.mark.parametrize( ("name", "expr", "result"), [ From db5b7cc199487f962cf61dd39c16f348107c6955 Mon Sep 17 00:00:00 2001 From: Tim Saucer Date: Mon, 28 Sep 2026 12:18:37 -0400 Subject: [PATCH 31/31] docs(skill): show substr's length argument as valid The skill said F.substr takes only two arguments and labelled F.substr(col("c_phone"), lit(1), lit(2)) as wrong. substr now accepts an optional length, so that call returns the first two characters, and the skill was steering agents away from a valid form. Describe both forms, mark length as requiring datafusion-python 55, and keep the TypeError note for older versions. Co-Authored-By: Claude Opus 5.5 (1M context) --- skills/datafusion_python/SKILL.md | 25 +++++++++++++------------ 1 file changed, 13 insertions(+), 12 deletions(-) diff --git a/skills/datafusion_python/SKILL.md b/skills/datafusion_python/SKILL.md index 296b6fdd5..34a4b6bf3 100644 --- a/skills/datafusion_python/SKILL.md +++ b/skills/datafusion_python/SKILL.md @@ -744,18 +744,19 @@ The `functions` module (imported as `F`) provides 290+ functions. Key categories `initcap`, `ascii`, `chr`, `left`, `right`, `strpos`, `translate`, `overlay`, `levenshtein` -`F.substr(str, start)` takes **only two arguments** and returns the tail of -the string from `start` onward — passing a third length argument raises -`TypeError: substr() takes 2 positional arguments but 3 were given`. For the -SQL-style 3-arg form (`SUBSTRING(str FROM start FOR length)`), use -`F.substring(col("s"), lit(start), lit(length))`. For a fixed-length prefix, -`F.left(col("s"), lit(n))` is cleanest. - -```python -# WRONG — substr does not accept a length argument -F.substr(col("c_phone"), lit(1), lit(2)) -# CORRECT -F.substring(col("c_phone"), lit(1), lit(2)) # explicit length +`F.substr(str, start)` returns the tail of the string from `start` onward; +`F.substr(str, start, length=n)` returns `n` characters from `start`, like +SQL's `SUBSTRING(str FROM start FOR length)`. `start` and `length` accept +native ints. `F.substring(str, start, length)` is the same with `length` +required. For a fixed-length prefix, `F.left(col("s"), lit(n))` is cleanest. + +*The `length` argument to `substr` requires datafusion-python 55 or newer. +On earlier versions a third argument raises `TypeError: substr() takes 2 +positional arguments but 3 were given`; use `F.substring` there.* + +```python +F.substr(col("c_phone"), 1, length=2) # first 2 characters +F.substring(col("c_phone"), lit(1), lit(2)) # same, on any version F.left(col("c_phone"), lit(2)) # prefix shortcut ```