diff --git a/.ai/skills/make-pythonic/SKILL.md b/.ai/skills/make-pythonic/SKILL.md index 24c2bb817..4ef2224d8 100644 --- a/.ai/skills/make-pythonic/SKILL.md +++ b/.ai/skills/make-pythonic/SKILL.md @@ -321,12 +321,8 @@ def concat_ws(separator: str, *args: Expr) -> Expr: ### Category C: Arguments That Should Accept str as Column Name -In some contexts a string argument naturally refers to a column name rather than a literal. This is the pattern used by DataFrame methods. - **Type hint pattern:** `Expr | str` -**When to use:** Only when the string contextually means a column name (rare in `functions.py`, more common in DataFrame methods). - ```python # Use _to_raw_expr() from expr.py for this pattern from datafusion.expr import _to_raw_expr @@ -336,7 +332,21 @@ def some_function(column: Expr | str) -> Expr: return Expr(f.some_function(raw)) ``` -**IMPORTANT:** In `functions.py`, string arguments almost never mean column names. Functions operate on expressions, and column references should use `col()`. Category C applies mainly to DataFrame methods and context APIs, not to scalar/aggregate/window functions. Do NOT convert string arguments to column expressions in `functions.py` unless there is a very clear reason to do so. +#### Column name or literal? + +A bare string has to be read either as a column name or as a literal value. Decide by asking what a string *literal* would mean in that position: + +1. **The position only makes sense as column data.** A constant there is meaningless, so a string can only be a column name. Accept `Expr | str` and convert with `_to_raw_expr()`. This covers: + - the column inputs of aggregate functions (`sum("a")`, `corr("a", "b")`, `grouping("a")`), because aggregating a constant string is never what the caller wants + - the column arguments of DataFrame and context methods (`select`, `sort`, `aggregate`) +2. **The position is a parameter of the operation, not its data.** Delimiters, patterns, date parts, format strings and `string_agg`'s `delimiter` are examples. A string here is a literal, so use Category A or B. +3. **The position is the data of a scalar or array function.** A literal string is a valid, common input there (`upper("abc")`, `concat(col("a"), "-")`), so a bare string is ambiguous. Keep it `Expr` only and let the caller write `col()` or `lit()`. + +If you can't tell which rule applies, use rule 3. + +SQL's `count(*)` is a special value, not a column name: `count("*")` means `count()`. + +Window functions have not been converted to rule 1 yet. Their inputs (for example `lead`'s `arg`) are still `Expr` only. ## Implementation Steps @@ -430,7 +440,7 @@ from datafusion.expr import coerce_to_expr, coerce_to_expr_or_none ## What NOT to Change -- **Do not change arguments that represent data columns.** If an argument is the primary data being operated on (e.g., the `string` in `left(string, n)` or the `array` in `array_sort(array)`), it should remain `Expr` only. Users should use `col()` for column references. +- **Do not change the data arguments of scalar and array functions.** If an argument is the primary data being operated on (e.g., the `string` in `left(string, n)` or the `array` in `array_sort(array)`), a string there could be a literal, so it should remain `Expr` only (rule 3 under Category C). Aggregate column inputs are different: see rule 1. - **Do not change variadic `*args: Expr` parameters.** These represent multiple expressions and should stay as `Expr`. - **Do not change arguments where the coercion is ambiguous.** If it is unclear whether a string should be a column name or a literal, leave it as `Expr` and let the user be explicit. - **Do not add coercion logic to simple aliases.** If a function is just `return other_function(...)`, the primary function handles coercion. However, you **must update the alias's type hints** to match the primary function's signature so that type checkers and documentation accurately reflect what the alias accepts. diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 25762374a..114fede44 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -54,6 +54,7 @@ Expr, SortExpr, SortKey, + _to_raw_expr, coerce_to_expr, coerce_to_expr_list, coerce_to_expr_or_none, @@ -5395,7 +5396,7 @@ def element_at(map: Expr, key: Expr) -> Expr: # aggregate functions def approx_distinct( - expression: Expr, + expression: Expr | str, filter: Expr | None = None, ) -> Expr: """Returns the approximate number of distinct values. @@ -5408,7 +5409,7 @@ def approx_distinct( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: Values to check for distinct entries + expression: Values to check for distinct entries (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -5431,10 +5432,10 @@ def approx_distinct( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.approx_distinct(expression.expr, filter=filter_raw)) + return Expr(f.approx_distinct(_to_raw_expr(expression), filter=filter_raw)) -def approx_median(expression: Expr, filter: Expr | None = None) -> Expr: +def approx_median(expression: Expr | str, filter: Expr | None = None) -> Expr: """Returns the approximate median value. This aggregate function is similar to :py:func:`median`, but it will only @@ -5444,7 +5445,7 @@ def approx_median(expression: Expr, filter: Expr | None = None) -> Expr: the options ``order_by`` and ``null_treatment``, and ``distinct``. Args: - expression: Values to find the median for + expression: Values to find the median for (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -5466,11 +5467,11 @@ def approx_median(expression: Expr, filter: Expr | None = None) -> Expr: 2.5 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.approx_median(expression.expr, filter=filter_raw)) + return Expr(f.approx_median(_to_raw_expr(expression), filter=filter_raw)) def approx_percentile_cont( - sort_expression: Expr | SortExpr, + sort_expression: Expr | SortExpr | str, percentile: float, num_centroids: int | None = None, filter: Expr | None = None, @@ -5494,6 +5495,7 @@ def approx_percentile_cont( Args: sort_expression: Values for which to find the approximate percentile + (expression, sort expression, or column name) percentile: This must be between 0.0 and 1.0, inclusive num_centroids: Max bin size for the t-digest algorithm filter: If provided, only compute against rows for which the filter is True @@ -5517,6 +5519,8 @@ def approx_percentile_cont( >>> result.collect_column("v")[0].as_py() 3.5 """ + if isinstance(sort_expression, str): + sort_expression = Expr.column(sort_expression) sort_expr_raw = sort_or_default(sort_expression) filter_raw = filter.expr if filter is not None else None return Expr( @@ -5527,8 +5531,8 @@ def approx_percentile_cont( def approx_percentile_cont_with_weight( - sort_expression: Expr | SortExpr, - weight: Expr, + sort_expression: Expr | SortExpr | str, + weight: Expr | str, percentile: float, num_centroids: int | None = None, filter: Expr | None = None, @@ -5545,7 +5549,9 @@ def approx_percentile_cont_with_weight( Args: sort_expression: Values for which to find the approximate percentile + (expression, sort expression, or column name) weight: Relative weight for each of the values in ``expression`` + (expression or column name) percentile: This must be between 0.0 and 1.0, inclusive num_centroids: Max bin size for the t-digest algorithm filter: If provided, only compute against rows for which the filter is True @@ -5569,12 +5575,14 @@ def approx_percentile_cont_with_weight( >>> result.collect_column("v")[0].as_py() 2.5 """ + if isinstance(sort_expression, str): + sort_expression = Expr.column(sort_expression) sort_expr_raw = sort_or_default(sort_expression) filter_raw = filter.expr if filter is not None else None return Expr( f.approx_percentile_cont_with_weight( sort_expr_raw, - weight.expr, + _to_raw_expr(weight), percentile, num_centroids=num_centroids, filter=filter_raw, @@ -5606,7 +5614,7 @@ def _check_distinct(name: str, distinct: object, *shifted: str) -> None: def percentile_cont( - sort_expression: Expr | SortExpr, + sort_expression: Expr | SortExpr | str, percentile: float, distinct: bool = False, filter: Expr | None = None, @@ -5622,7 +5630,8 @@ def percentile_cont( the same expression. Args: - sort_expression: Values for which to find the percentile + sort_expression: Values for which to find the percentile (expression, sort + expression, or column name) 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 @@ -5654,6 +5663,8 @@ def percentile_cont( 2.5 """ _check_distinct("percentile_cont", distinct, "filter") + if isinstance(sort_expression, str): + sort_expression = Expr.column(sort_expression) sort_expr_raw = sort_or_default(sort_expression) filter_raw = filter.expr if filter is not None else None return Expr( @@ -5664,7 +5675,7 @@ def percentile_cont( def quantile_cont( - sort_expression: Expr | SortExpr, + sort_expression: Expr | SortExpr | str, percentile: float, distinct: bool = False, filter: Expr | None = None, @@ -5681,7 +5692,7 @@ def quantile_cont( def array_agg( - expression: Expr, + expression: Expr | str, distinct: bool = False, filter: Expr | None = None, order_by: list[SortKey] | SortKey | None = None, @@ -5696,7 +5707,7 @@ def array_agg( the option ``null_treatment``. Args: - expression: Values to combine into an array + expression: Values to combine into an array (expression or column name) distinct: If True, a single entry for each distinct value will be in the result filter: If provided, only compute against rows for which the filter is True order_by: Order the resultant array values. Accepts column names or expressions. @@ -5733,13 +5744,16 @@ def array_agg( return Expr( f.array_agg( - expression.expr, distinct=distinct, filter=filter_raw, order_by=order_by_raw + _to_raw_expr(expression), + distinct=distinct, + filter=filter_raw, + order_by=order_by_raw, ) ) def grouping( - expression: Expr, + expression: Expr | str, distinct: bool = False, filter: Expr | None = None, ) -> Expr: @@ -5769,7 +5783,8 @@ def grouping( on the result DataFrame instead. Args: - expression: The column to check grouping status for + expression: The column to check grouping status for (expression or column + name) distinct: If True, compute on distinct values only filter: If provided, only compute against rows for which the filter is True @@ -5794,11 +5809,13 @@ def grouping( :py:class:`~datafusion.expr.GroupingSet` """ filter_raw = filter.expr if filter is not None else None - return Expr(f.grouping(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr( + f.grouping(_to_raw_expr(expression), distinct=distinct, filter=filter_raw) + ) def avg( - expression: Expr, + expression: Expr | str, distinct: bool = False, filter: Expr | None = None, ) -> Expr: @@ -5810,7 +5827,7 @@ def avg( the options ``order_by`` and ``null_treatment``. Args: - expression: Values to combine into an array + expression: Values to combine into an array (expression or column name) distinct: If True, duplicate values are removed before averaging. filter: If provided, only compute against rows for which the filter is True @@ -5841,10 +5858,10 @@ def avg( 2.0 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.avg(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr(f.avg(_to_raw_expr(expression), distinct=distinct, filter=filter_raw)) -def corr(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: +def corr(value_y: Expr | str, value_x: Expr | str, filter: Expr | None = None) -> Expr: """Returns the correlation coefficient between ``value1`` and ``value2``. This aggregate function expects both values to be numeric and will return a float. @@ -5853,8 +5870,8 @@ def corr(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - value_y: The dependent variable for correlation - value_x: The independent variable for correlation + value_y: The dependent variable for correlation (expression or column name) + value_x: The independent variable for correlation (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -5876,11 +5893,11 @@ def corr(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: 1.0 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.corr(value_y.expr, value_x.expr, filter=filter_raw)) + return Expr(f.corr(_to_raw_expr(value_y), _to_raw_expr(value_x), filter=filter_raw)) def count( - expressions: Expr | list[Expr] | None = None, + expressions: Expr | str | list[Expr | str] | None = None, distinct: bool = False, filter: Expr | None = None, ) -> Expr: @@ -5892,7 +5909,8 @@ def count( the options ``order_by`` and ``null_treatment``. Args: - expressions: Argument to perform bitwise calculation on + expressions: Values to count (expression, column name, or a list of + either). ``"*"`` counts all rows, the same as ``count()``. distinct: If True, a single entry for each distinct value will be in the result filter: If provided, only compute against rows for which the filter is True @@ -5917,17 +5935,19 @@ def count( """ filter_raw = filter.expr if filter is not None else None - if expressions is None: + if expressions is None or (isinstance(expressions, str) and expressions == "*"): args = [Expr.literal(1).expr] elif isinstance(expressions, list): - args = [arg.expr for arg in expressions] + args = [_to_raw_expr(arg) for arg in expressions] else: - args = [expressions.expr] + args = [_to_raw_expr(expressions)] return Expr(f.count(*args, distinct=distinct, filter=filter_raw)) -def covar_pop(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: +def covar_pop( + value_y: Expr | str, value_x: Expr | str, filter: Expr | None = None +) -> Expr: """Computes the population covariance. This aggregate function expects both values to be numeric and will return a float. @@ -5936,8 +5956,8 @@ def covar_pop(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - value_y: The dependent variable for covariance - value_x: The independent variable for covariance + value_y: The dependent variable for covariance (expression or column name) + value_x: The independent variable for covariance (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -5965,10 +5985,14 @@ def covar_pop(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: 1.0 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.covar_pop(value_y.expr, value_x.expr, filter=filter_raw)) + return Expr( + f.covar_pop(_to_raw_expr(value_y), _to_raw_expr(value_x), filter=filter_raw) + ) -def covar_samp(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: +def covar_samp( + value_y: Expr | str, value_x: Expr | str, filter: Expr | None = None +) -> Expr: """Computes the sample covariance. This aggregate function expects both values to be numeric and will return a float. @@ -5977,8 +6001,8 @@ def covar_samp(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - value_y: The dependent variable for covariance - value_x: The independent variable for covariance + value_y: The dependent variable for covariance (expression or column name) + value_x: The independent variable for covariance (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6000,10 +6024,12 @@ def covar_samp(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr 0.5 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.covar_samp(value_y.expr, value_x.expr, filter=filter_raw)) + return Expr( + f.covar_samp(_to_raw_expr(value_y), _to_raw_expr(value_x), filter=filter_raw) + ) -def covar(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: +def covar(value_y: Expr | str, value_x: Expr | str, filter: Expr | None = None) -> Expr: """Computes the sample covariance. See Also: @@ -6012,14 +6038,14 @@ def covar(value_y: Expr, value_x: Expr, filter: Expr | None = None) -> Expr: return covar_samp(value_y, value_x, filter) -def max(expression: Expr, filter: Expr | None = None) -> Expr: +def max(expression: Expr | str, 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 the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: The value to find the maximum of + expression: The value to find the maximum of (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6041,11 +6067,11 @@ def max(expression: Expr, filter: Expr | None = None) -> Expr: 2 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.max(expression.expr, filter=filter_raw)) + return Expr(f.max(_to_raw_expr(expression), filter=filter_raw)) def mean( - expression: Expr, + expression: Expr | str, distinct: bool = False, filter: Expr | None = None, ) -> Expr: @@ -6059,7 +6085,7 @@ def mean( def median( - expression: Expr, distinct: bool = False, filter: Expr | None = None + expression: Expr | str, distinct: bool = False, filter: Expr | None = None ) -> Expr: """Computes the median of a set of numbers. @@ -6070,7 +6096,7 @@ def median( the options ``order_by`` and ``null_treatment``. Args: - expression: The value to compute the median of + expression: The value to compute the median of (expression or column name) distinct: If True, a single entry for each distinct value will be in the result filter: If provided, only compute against rows for which the filter is True @@ -6094,17 +6120,19 @@ def median( 1.5 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.median(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr( + f.median(_to_raw_expr(expression), distinct=distinct, filter=filter_raw) + ) -def min(expression: Expr, filter: Expr | None = None) -> Expr: +def min(expression: Expr | str, 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 the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: The value to find the minimum of + expression: The value to find the minimum of (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6126,11 +6154,11 @@ def min(expression: Expr, filter: Expr | None = None) -> Expr: 2 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.min(expression.expr, filter=filter_raw)) + return Expr(f.min(_to_raw_expr(expression), filter=filter_raw)) def sum( - expression: Expr, + expression: Expr | str, distinct: bool = False, filter: Expr | None = None, ) -> Expr: @@ -6142,7 +6170,7 @@ def sum( the options ``order_by`` and ``null_treatment``. Args: - expression: Values to combine into an array + expression: Values to combine into an array (expression or column name) distinct: If True, duplicate values are removed before summing. filter: If provided, only compute against rows for which the filter is True @@ -6156,6 +6184,12 @@ def sum( >>> result.collect_column("v")[0].as_py() 6 + A column name can be passed in place of ``col(...)``: + + >>> result = df.aggregate([], [dfn.functions.sum("a").alias("v")]) + >>> result.collect_column("v")[0].as_py() + 6 + >>> result = df.aggregate( ... [], [dfn.functions.sum( ... dfn.col("a"), @@ -6173,17 +6207,17 @@ def sum( 6 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.sum(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr(f.sum(_to_raw_expr(expression), distinct=distinct, filter=filter_raw)) -def stddev(expression: Expr, filter: Expr | None = None) -> Expr: +def stddev(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the standard deviation of the argument. If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: The value to find the minimum of + expression: The value to find the minimum of (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6205,17 +6239,17 @@ def stddev(expression: Expr, filter: Expr | None = None) -> Expr: 1.41... """ filter_raw = filter.expr if filter is not None else None - return Expr(f.stddev(expression.expr, filter=filter_raw)) + return Expr(f.stddev(_to_raw_expr(expression), filter=filter_raw)) -def stddev_pop(expression: Expr, filter: Expr | None = None) -> Expr: +def stddev_pop(expression: Expr | str, 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 the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: The value to find the minimum of + expression: The value to find the minimum of (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6240,10 +6274,10 @@ def stddev_pop(expression: Expr, filter: Expr | None = None) -> Expr: 1.0 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.stddev_pop(expression.expr, filter=filter_raw)) + return Expr(f.stddev_pop(_to_raw_expr(expression), filter=filter_raw)) -def stddev_samp(arg: Expr, filter: Expr | None = None) -> Expr: +def stddev_samp(arg: Expr | str, filter: Expr | None = None) -> Expr: """Computes the sample standard deviation of the argument. See Also: @@ -6252,7 +6286,7 @@ def stddev_samp(arg: Expr, filter: Expr | None = None) -> Expr: return stddev(arg, filter=filter) -def var(expression: Expr, filter: Expr | None = None) -> Expr: +def var(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the sample variance of the argument. See Also: @@ -6261,14 +6295,14 @@ def var(expression: Expr, filter: Expr | None = None) -> Expr: return var_samp(expression, filter) -def var_pop(expression: Expr, filter: Expr | None = None) -> Expr: +def var_pop(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the population variance of the argument. If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: The variable to compute the variance for + expression: The variable to compute the variance for (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6290,10 +6324,10 @@ def var_pop(expression: Expr, filter: Expr | None = None) -> Expr: 1.0 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.var_pop(expression.expr, filter=filter_raw)) + return Expr(f.var_pop(_to_raw_expr(expression), filter=filter_raw)) -def var_population(expression: Expr, filter: Expr | None = None) -> Expr: +def var_population(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the population variance of the argument. See Also: @@ -6302,14 +6336,14 @@ def var_population(expression: Expr, filter: Expr | None = None) -> Expr: return var_pop(expression, filter) -def var_samp(expression: Expr, filter: Expr | None = None) -> Expr: +def var_samp(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the sample variance of the argument. If using the builder functions described in :ref:`aggregation` this function ignores the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: The variable to compute the variance for + expression: The variable to compute the variance for (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6331,10 +6365,10 @@ def var_samp(expression: Expr, filter: Expr | None = None) -> Expr: 0.5 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.var_sample(expression.expr, filter=filter_raw)) + return Expr(f.var_sample(_to_raw_expr(expression), filter=filter_raw)) -def var_sample(expression: Expr, filter: Expr | None = None) -> Expr: +def var_sample(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the sample variance of the argument. See Also: @@ -6344,8 +6378,8 @@ def var_sample(expression: Expr, filter: Expr | None = None) -> Expr: def regr_avgx( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the average of the independent variable ``x``. @@ -6357,8 +6391,8 @@ def regr_avgx( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6381,12 +6415,12 @@ def regr_avgx( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_avgx(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_avgx(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_avgy( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the average of the dependent variable ``y``. @@ -6398,8 +6432,8 @@ def regr_avgy( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6422,12 +6456,12 @@ def regr_avgy( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_avgy(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_avgy(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_count( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Counts the number of rows in which both expressions are not null. @@ -6439,8 +6473,8 @@ def regr_count( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6463,12 +6497,12 @@ def regr_count( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_count(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_count(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_intercept( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the intercept from the linear regression. @@ -6480,8 +6514,8 @@ def regr_intercept( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6506,12 +6540,12 @@ def regr_intercept( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_intercept(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_intercept(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_r2( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the R-squared value from linear regression. @@ -6523,8 +6557,8 @@ def regr_r2( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6547,12 +6581,12 @@ def regr_r2( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_r2(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_r2(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_slope( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the slope from linear regression. @@ -6564,8 +6598,8 @@ def regr_slope( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6588,12 +6622,12 @@ def regr_slope( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_slope(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_slope(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_sxx( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the sum of squares of the independent variable ``x``. @@ -6605,8 +6639,8 @@ def regr_sxx( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6629,12 +6663,12 @@ def regr_sxx( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_sxx(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_sxx(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_sxy( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the sum of products of pairs of numbers. @@ -6646,8 +6680,8 @@ def regr_sxy( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6670,12 +6704,12 @@ def regr_sxy( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_sxy(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_sxy(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def regr_syy( - y: Expr, - x: Expr, + y: Expr | str, + x: Expr | str, filter: Expr | None = None, ) -> Expr: """Computes the sum of squares of the dependent variable ``y``. @@ -6687,8 +6721,8 @@ def regr_syy( the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - y: The linear regression dependent variable - x: The linear regression independent variable + y: The linear regression dependent variable (expression or column name) + x: The linear regression independent variable (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -6711,11 +6745,11 @@ def regr_syy( """ filter_raw = filter.expr if filter is not None else None - return Expr(f.regr_syy(y.expr, x.expr, filter=filter_raw)) + return Expr(f.regr_syy(_to_raw_expr(y), _to_raw_expr(x), filter=filter_raw)) def first_value( - expression: Expr, + expression: Expr | str, filter: Expr | None = None, order_by: list[SortKey] | SortKey | None = None, null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS, @@ -6729,6 +6763,7 @@ def first_value( Args: expression: Argument to perform bitwise calculation on + (expression or column name) 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. @@ -6762,7 +6797,7 @@ def first_value( return Expr( f.first_value( - expression.expr, + _to_raw_expr(expression), filter=filter_raw, order_by=order_by_raw, null_treatment=null_treatment.value, @@ -6771,7 +6806,7 @@ def first_value( def last_value( - expression: Expr, + expression: Expr | str, filter: Expr | None = None, order_by: list[SortKey] | SortKey | None = None, null_treatment: NullTreatment = NullTreatment.RESPECT_NULLS, @@ -6785,6 +6820,7 @@ def last_value( Args: expression: Argument to perform bitwise calculation on + (expression or column name) 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. @@ -6818,7 +6854,7 @@ def last_value( return Expr( f.last_value( - expression.expr, + _to_raw_expr(expression), filter=filter_raw, order_by=order_by_raw, null_treatment=null_treatment.value, @@ -6827,7 +6863,7 @@ def last_value( def nth_value( - expression: Expr, + expression: Expr | str, n: int, filter: Expr | None = None, order_by: list[SortKey] | SortKey | None = None, @@ -6842,6 +6878,7 @@ def nth_value( Args: expression: Argument to perform bitwise calculation on + (expression or column name) n: Index of value to return. Starts at 1. filter: If provided, only compute against rows for which the filter is True order_by: Set the ordering of the expression to evaluate. Accepts @@ -6875,7 +6912,7 @@ def nth_value( return Expr( f.nth_value( - expression.expr, + _to_raw_expr(expression), n, filter=filter_raw, order_by=order_by_raw, @@ -6884,7 +6921,7 @@ def nth_value( ) -def any_value(expression: Expr, filter: Expr | None = None) -> Expr: +def any_value(expression: Expr | str, 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 @@ -6894,7 +6931,7 @@ def any_value(expression: Expr, filter: Expr | None = None) -> Expr: the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: Argument to pick a value from + expression: Argument to pick a value from (expression or column name) filter: If provided, only consider rows for which the filter is True Examples: @@ -6917,11 +6954,11 @@ def any_value(expression: Expr, filter: Expr | None = None) -> Expr: 8 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.any_value(expression.expr, filter=filter_raw)) + return Expr(f.any_value(_to_raw_expr(expression), filter=filter_raw)) def bit_and( - expression: Expr, distinct: bool = False, filter: Expr | None = None + expression: Expr | str, distinct: bool = False, filter: Expr | None = None ) -> Expr: """Computes the bitwise AND of the argument. @@ -6932,6 +6969,7 @@ def bit_and( Args: expression: Argument to perform bitwise calculation on + (expression or column name) distinct: If True, evaluate each unique value of expression only once filter: If provided, only compute against rows for which the filter is True @@ -6956,11 +6994,13 @@ def bit_and( """ _check_distinct("bit_and", distinct, "filter") filter_raw = filter.expr if filter is not None else None - return Expr(f.bit_and(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr( + f.bit_and(_to_raw_expr(expression), distinct=distinct, filter=filter_raw) + ) def bit_or( - expression: Expr, distinct: bool = False, filter: Expr | None = None + expression: Expr | str, distinct: bool = False, filter: Expr | None = None ) -> Expr: """Computes the bitwise OR of the argument. @@ -6971,6 +7011,7 @@ def bit_or( Args: expression: Argument to perform bitwise calculation on + (expression or column name) distinct: If True, evaluate each unique value of expression only once filter: If provided, only compute against rows for which the filter is True @@ -6997,11 +7038,13 @@ def bit_or( """ _check_distinct("bit_or", distinct, "filter") filter_raw = filter.expr if filter is not None else None - return Expr(f.bit_or(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr( + f.bit_or(_to_raw_expr(expression), distinct=distinct, filter=filter_raw) + ) def bit_xor( - expression: Expr, distinct: bool = False, filter: Expr | None = None + expression: Expr | str, distinct: bool = False, filter: Expr | None = None ) -> Expr: """Computes the bitwise XOR of the argument. @@ -7012,6 +7055,7 @@ def bit_xor( Args: expression: Argument to perform bitwise calculation on + (expression or column name) distinct: If True, evaluate each unique value of expression only once filter: If provided, only compute against rows for which the filter is True @@ -7037,10 +7081,12 @@ def bit_xor( 5 """ filter_raw = filter.expr if filter is not None else None - return Expr(f.bit_xor(expression.expr, distinct=distinct, filter=filter_raw)) + return Expr( + f.bit_xor(_to_raw_expr(expression), distinct=distinct, filter=filter_raw) + ) -def bool_and(expression: Expr, filter: Expr | None = None) -> Expr: +def bool_and(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the boolean AND of the argument. This aggregate function will compare every value in the input partition. These are @@ -7050,7 +7096,7 @@ def bool_and(expression: Expr, filter: Expr | None = None) -> Expr: the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: Argument to perform calculation on + expression: Argument to perform calculation on (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -7076,10 +7122,10 @@ def bool_and(expression: Expr, filter: Expr | None = None) -> Expr: True """ filter_raw = filter.expr if filter is not None else None - return Expr(f.bool_and(expression.expr, filter=filter_raw)) + return Expr(f.bool_and(_to_raw_expr(expression), filter=filter_raw)) -def bool_or(expression: Expr, filter: Expr | None = None) -> Expr: +def bool_or(expression: Expr | str, filter: Expr | None = None) -> Expr: """Computes the boolean OR of the argument. This aggregate function will compare every value in the input partition. These are @@ -7089,7 +7135,7 @@ def bool_or(expression: Expr, filter: Expr | None = None) -> Expr: the options ``order_by``, ``null_treatment``, and ``distinct``. Args: - expression: Argument to perform calculation on + expression: Argument to perform calculation on (expression or column name) filter: If provided, only compute against rows for which the filter is True Examples: @@ -7115,7 +7161,7 @@ def bool_or(expression: Expr, filter: Expr | None = None) -> Expr: False """ filter_raw = filter.expr if filter is not None else None - return Expr(f.bool_or(expression.expr, filter=filter_raw)) + return Expr(f.bool_or(_to_raw_expr(expression), filter=filter_raw)) def lead( @@ -7652,7 +7698,7 @@ def ntile( def string_agg( - expression: Expr, + expression: Expr | str, delimiter: str, distinct: bool = False, filter: Expr | None = None, @@ -7669,6 +7715,7 @@ def string_agg( Args: expression: Argument to perform bitwise calculation on + (expression or column name) 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 @@ -7708,7 +7755,7 @@ def string_agg( return Expr( f.string_agg( - expression.expr, + _to_raw_expr(expression), delimiter, distinct=distinct, filter=filter_raw, diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index 38baef954..a0a62cfe1 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -102,6 +102,16 @@ def df_aggregate_100(): (f.var(column("a")), lambda a, b, c, d: np.array(np.var(a, ddof=1))), (f.var_pop(column("b")), lambda a, b, c, d: np.array(np.var(b, ddof=0))), (f.var_samp(column("c")), lambda a, b, c, d: np.array(np.var(c, ddof=1))), + # Column names in place of column(...) + (f.avg("a"), lambda a, b, c, d: np.array(np.average(a))), + (f.corr("a", "b"), lambda a, b, c, d: np.array(np.corrcoef(a, b)[0][1])), + (f.count("a"), lambda a, b, c, d: pa.array([len(a)])), + (f.covar("a", "b"), lambda a, b, c, d: np.array(np.cov(a, b, ddof=1)[0][1])), + (f.max("a"), lambda a, b, c, d: np.array(np.max(a))), + (f.mean("b"), lambda a, b, c, d: np.array(np.mean(b))), + (f.sum("b"), lambda a, b, c, d: np.array(np.sum(b.to_pylist()))), + (f.stddev_samp("c"), lambda a, b, c, d: np.array(np.std(c, ddof=1))), + (f.var("a"), lambda a, b, c, d: np.array(np.var(a, ddof=1))), ], ) def test_aggregation_stats(df, agg_expr, calc_expected): @@ -555,6 +565,54 @@ def test_string_agg(name, expr, result) -> None: assert df.collect()[0].to_pydict() == expected +@pytest.mark.parametrize( + ("by_name", "by_expr"), + [ + (f.count(["e"]), f.count([column("e")])), + (f.count("*"), f.count()), + (f.regr_slope("c", "b"), f.regr_slope(column("c"), column("b"))), + ( + f.approx_percentile_cont("b", 0.5), + f.approx_percentile_cont(column("b"), 0.5), + ), + ( + f.approx_percentile_cont_with_weight("b", "c", 0.5), + f.approx_percentile_cont_with_weight(column("b"), column("c"), 0.5), + ), + (f.percentile_cont("c", 0.5), f.percentile_cont(column("c"), 0.5)), + ( + f.first_value("b", order_by="c"), + f.first_value(column("b"), order_by="c"), + ), + (f.array_agg("b", order_by="b"), f.array_agg(column("b"), order_by="b")), + (f.bool_and("d"), f.bool_and(column("d"))), + ( + f.string_agg("s", ",", order_by="c"), + f.string_agg(column("s"), ",", order_by="c"), + ), + (f.grouping("a"), f.grouping(column("a"))), + ], +) +def test_aggregate_accepts_column_name(by_name, by_expr) -> None: + # grouping() needs its argument in the group by, so every case is grouped + # by "a". No alias: grouping() cannot be aliased (apache/datafusion#21411). + df = SessionContext().from_pydict( + { + "a": [1, 1, 2], + "b": [4, 4, 6], + "c": [9, 8, 5], + "d": [True, True, False], + "e": [1, None, 3], + "s": ["one", "two", "three"], + } + ) + + def run(expr): + return df.aggregate(["a"], [expr]).sort(column("a")).collect() + + assert run(by_name) == run(by_expr) + + _FILTER = column("b") > lit(1)