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 diff --git a/crates/core/src/dataframe.rs b/crates/core/src/dataframe.rs index b1f305551..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) } @@ -1320,6 +1345,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/crates/core/src/expr.rs b/crates/core/src/expr.rs index cab997d7a..11961521f 100644 --- a/crates/core/src/expr.rs +++ b/crates/core/src/expr.rs @@ -29,7 +29,7 @@ use datafusion::logical_expr::expr::{ use datafusion::logical_expr::utils::exprlist_to_fields; use datafusion::logical_expr::{ Between, BinaryExpr, Case, Cast, Expr, ExprFuncBuilder, ExprFunctionExt, Like, LogicalPlan, - Operator, TryCast, WindowFunctionDefinition, col, lit, lit_with_metadata, + Operator, TryCast, WindowFrame, WindowFunctionDefinition, col, lit, lit_with_metadata, }; use datafusion_proto::logical_plan::{from_proto, to_proto}; use prost::Message; @@ -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::apply_window_options; use crate::pyarrow_util::scalar_to_pyarrow; use crate::sql::logical::PyLogicalPlan; @@ -624,35 +624,34 @@ impl PyExpr { // Expression Function Builder functions - pub fn order_by(&self, order_by: Vec) -> PyExprFuncBuilder { - self.expr - .clone() - .order_by(to_sort_expressions(order_by)) - .into() + pub fn order_by(&self, order_by: Vec) -> PyDataFusionResult { + Ok(PyExprFuncBuilder::from_expr(&self.expr)?.order_by(order_by)) } - pub fn filter(&self, filter: PyExpr) -> PyExprFuncBuilder { - self.expr.clone().filter(filter.expr.clone()).into() + pub fn filter(&self, filter: PyExpr) -> PyDataFusionResult { + PyExprFuncBuilder::from_expr(&self.expr)?.filter(filter) } - pub fn distinct(&self) -> PyExprFuncBuilder { - self.expr.clone().distinct().into() + pub fn distinct(&self) -> PyDataFusionResult { + PyExprFuncBuilder::from_expr(&self.expr)?.distinct() } - pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder { - self.expr - .clone() - .null_treatment(Some(null_treatment.into())) - .into() + pub fn null_treatment( + &self, + null_treatment: NullTreatment, + ) -> PyDataFusionResult { + Ok(PyExprFuncBuilder::from_expr(&self.expr)?.null_treatment(null_treatment)) } - 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() + pub fn partition_by(&self, partition_by: Vec) -> PyDataFusionResult { + PyExprFuncBuilder::from_expr(&self.expr)?.partition_by(partition_by) } - pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder { - self.expr.clone().window_frame(window_frame.into()).into() + pub fn window_frame( + &self, + window_frame: PyWindowFrame, + ) -> PyDataFusionResult { + PyExprFuncBuilder::from_expr(&self.expr)?.window_frame(window_frame) } #[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None))] @@ -665,21 +664,41 @@ impl PyExpr { ) -> PyDataFusionResult { match &self.expr { Expr::AggregateFunction(agg_fn) => { - let window_fn = Expr::WindowFunction(Box::new(WindowFunction::new( - WindowFunctionDefinition::AggregateUDF(agg_fn.func.clone()), - agg_fn.params.args.clone(), - ))); + let params = &agg_fn.params; + // A window never passes an ordering to the accumulator, so an + // aggregate's order_by cannot be kept. A WITHIN GROUP function + // runs ascending as a window, so only an ascending ordering can + // be dropped without changing the result. + let order_by_is_redundant = agg_fn.func.supports_within_group_clause() + && params.order_by.iter().all(|sort| sort.asc); + if !params.order_by.is_empty() && !order_by_is_redundant { + return Err(datafusion::error::DataFusionError::Plan(format!( + "Aggregate order_by is not supported when {} is used as a window \ + function", + agg_fn.func.name() + )) + .into()); + } - add_builder_fns_to_window( - window_fn, + let mut window_fn = WindowFunction::new( + WindowFunctionDefinition::AggregateUDF(agg_fn.func.clone()), + params.args.clone(), + ); + window_fn.params.filter = params.filter.clone(); + window_fn.params.distinct = params.distinct; + window_fn.params.null_treatment = params.null_treatment; + + apply_window_options( + PyExprFuncBuilder::from_expr(&Expr::WindowFunction(Box::new(window_fn)))? + .builder, partition_by, window_frame, order_by, null_treatment, ) } - Expr::WindowFunction(_) => add_builder_fns_to_window( - self.expr.clone(), + Expr::WindowFunction(_) => apply_window_options( + PyExprFuncBuilder::from_expr(&self.expr)?.builder, partition_by, window_frame, order_by, @@ -753,48 +772,177 @@ impl PyExpr { #[derive(Debug, Clone)] pub struct PyExprFuncBuilder { pub builder: ExprFuncBuilder, + kind: FuncKind, + name: String, +} + +/// The kind of function a builder was started from, which decides the +/// options it accepts. Upstream checks the kind only when a builder is first +/// created from an `Expr` and silently drops options at `build()`, so it is +/// checked here on every call instead. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum FuncKind { + Aggregate, + AggregateWindow, + Window, + /// Not a function; upstream's empty builder makes `build()` raise. + Other, } -impl From for PyExprFuncBuilder { - fn from(builder: ExprFuncBuilder) -> Self { - Self { builder } +impl PyExprFuncBuilder { + /// Start a builder 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. + /// + /// A built window function always stores a concrete frame, so whether the user + /// chose it is lost and its options cannot be merged without guessing. A window + /// function that already has a partition, an order-by, or a frame other than + /// the whole-partition default raises instead; see apache/datafusion#25934. + fn from_expr(expr: &Expr) -> PyDataFusionResult { + 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(); + } + Ok(Self { + builder, + kind: FuncKind::Aggregate, + name: agg.func.name().to_string(), + }) + } + Expr::WindowFunction(window) => { + let params = &window.params; + let name = window.fun.name().to_string(); + // Any frame but the whole-partition default was set explicitly or + // derived from an order-by. A derived one is reported as the order-by. + let has_order_by = !params.order_by.is_empty(); + let derived = |strict| params.window_frame == WindowFrame::new(Some(strict)); + let has_frame = params.window_frame != WindowFrame::new(None); + let set: Vec<&str> = [ + ("partition_by", !params.partition_by.is_empty()), + ("order_by", has_order_by), + ( + "window_frame", + has_frame && !(has_order_by && (derived(true) || derived(false))), + ), + ] + .into_iter() + .filter_map(|(option, is_set)| is_set.then_some(option)) + .collect(); + if !set.is_empty() { + return Err(datafusion::error::DataFusionError::Plan(format!( + "{name} already has window options ({}); set partition_by, order_by, \ + and window_frame in one place: the function's keyword arguments, a \ + single over(Window(...)), or one builder chain", + set.join(", ") + )) + .into()); + } + let mut builder = expr.clone().null_treatment(params.null_treatment); + if let Some(filter) = ¶ms.filter { + builder = builder.filter(filter.as_ref().clone()); + } + if params.distinct { + builder = builder.distinct(); + } + let kind = match window.fun { + WindowFunctionDefinition::AggregateUDF(_) => FuncKind::AggregateWindow, + WindowFunctionDefinition::WindowUDF(_) => FuncKind::Window, + }; + Ok(Self { + builder, + kind, + name, + }) + } + _ => Ok(Self { + builder: expr.clone().null_treatment(None), + kind: FuncKind::Other, + name: String::new(), + }), + } + } + + fn with_builder(&self, builder: ExprFuncBuilder) -> Self { + Self { + builder, + kind: self.kind, + name: self.name.clone(), + } + } + + fn require_aggregate(&self, option: &str) -> PyDataFusionResult<()> { + if self.kind == FuncKind::Window { + return Err(datafusion::error::DataFusionError::Plan(format!( + "{option}() applies only to aggregate functions, including one used as a \ + window function; {} is a window function", + self.name + )) + .into()); + } + Ok(()) + } + + fn require_window(&self, option: &str) -> PyDataFusionResult<()> { + if self.kind == FuncKind::Aggregate { + return Err(datafusion::error::DataFusionError::Plan(format!( + "{option}() applies only to window functions; {} is an aggregate function, \ + use over() to run it as a window function", + self.name + )) + .into()); + } + Ok(()) } } #[pymethods] impl PyExprFuncBuilder { pub fn order_by(&self, order_by: Vec) -> PyExprFuncBuilder { - self.builder - .clone() - .order_by(to_sort_expressions(order_by)) - .into() + self.with_builder(self.builder.clone().order_by(to_sort_expressions(order_by))) } - pub fn filter(&self, filter: PyExpr) -> PyExprFuncBuilder { - self.builder.clone().filter(filter.expr.clone()).into() + pub fn filter(&self, filter: PyExpr) -> PyDataFusionResult { + self.require_aggregate("filter")?; + Ok(self.with_builder(self.builder.clone().filter(filter.expr))) } - pub fn distinct(&self) -> PyExprFuncBuilder { - self.builder.clone().distinct().into() + pub fn distinct(&self) -> PyDataFusionResult { + self.require_aggregate("distinct")?; + Ok(self.with_builder(self.builder.clone().distinct())) } pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder { - self.builder - .clone() - .null_treatment(Some(null_treatment.into())) - .into() + self.with_builder( + self.builder + .clone() + .null_treatment(Some(null_treatment.into())), + ) } - pub fn partition_by(&self, partition_by: Vec) -> PyExprFuncBuilder { - let partition_by = partition_by.iter().map(|e| e.expr.clone()).collect(); - self.builder.clone().partition_by(partition_by).into() + pub fn partition_by(&self, partition_by: Vec) -> PyDataFusionResult { + self.require_window("partition_by")?; + let partition_by = partition_by.into_iter().map(|e| e.expr).collect(); + Ok(self.with_builder(self.builder.clone().partition_by(partition_by))) } - pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder { - self.builder - .clone() - .window_frame(window_frame.into()) - .into() + pub fn window_frame( + &self, + window_frame: PyWindowFrame, + ) -> PyDataFusionResult { + self.require_window("window_frame")?; + Ok(self.with_builder(self.builder.clone().window_frame(window_frame.into()))) } pub fn build(&self) -> PyDataFusionResult { diff --git a/crates/core/src/functions.rs b/crates/core/src/functions.rs index e57c7702d..915ae8719 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; @@ -118,19 +118,69 @@ fn string_to_array(string: PyExpr, delimiter: PyExpr, null_string: Option) -> PyExpr { - let mut args = vec![start.into(), stop.into()]; - if let Some(step) = step { - args.push(step.into()); +#[pyo3(signature = (array, delimiter, null_string=None))] +fn array_to_string(array: 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::range::gen_series_udf(), + datafusion::functions_nested::string::array_to_string_udf(), args, )) .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, +) -> PyResult { + // Upstream reads the arguments by position, so a step without a stop + // would be taken as the stop. + if stop.is_none() && step.is_some() { + return Err(pyo3::exceptions::PyValueError::new_err(format!( + "{}() requires stop when step is given", + udf.name() + ))); + } + let args = std::iter::once(start) + .chain(stop) + .chain(step) + .map(Into::into) + .collect(); + Ok( + Expr::ScalarFunction(datafusion::logical_expr::expr::ScalarFunction::new_udf( + udf, args, + )) + .into(), + ) +} + +#[pyfunction] +#[pyo3(signature = (start, stop=None, step=None))] +fn range(start: PyExpr, stop: Option, step: Option) -> PyResult { + 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) -> PyResult { + 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(); @@ -197,6 +247,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)] @@ -615,6 +672,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 { @@ -637,7 +696,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); @@ -663,6 +721,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); @@ -673,7 +737,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); @@ -688,6 +751,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); @@ -726,12 +790,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] @@ -744,26 +810,31 @@ 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] -#[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)); + let agg_fn = functions_aggregate::expr_fn::percentile_cont( + sort_expression.sort.clone(), + lit(percentile), + ); - add_builder_fns_to_aggregate(agg_fn, None, 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 @@ -836,8 +907,30 @@ 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. +/// +/// An empty `order_by` is treated as `None`. Passing it through would make +/// `build()` derive a RANGE frame with no sort key, which cannot execute. +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( @@ -848,7 +941,7 @@ pub(crate) fn add_builder_fns_to_window( ); } - if let Some(order_by_cols) = order_by { + if let Some(order_by_cols) = order_by.filter(|cols| !cols.is_empty()) { let order_by_cols = to_sort_expressions(order_by_cols); builder = builder.order_by(order_by_cols); } @@ -861,33 +954,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] @@ -1056,6 +1151,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))?; @@ -1126,6 +1223,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 +1238,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 +1250,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..868fc6934 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,14 +235,37 @@ 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); +/// `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); @@ -292,6 +320,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 +337,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 +351,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,14 +374,17 @@ 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))?; + 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/crates/core/src/udaf.rs b/crates/core/src/udaf.rs index 6a2675193..3b16180bb 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}; @@ -300,7 +302,17 @@ impl AggregateUDFImpl for PythonFunctionAggregateUDF { Ok(self.return_type.clone()) } - fn accumulator(&self, _acc_args: AccumulatorArgs) -> Result> { + fn accumulator(&self, acc_args: AccumulatorArgs) -> Result> { + // The Python accumulator cannot see the flag, so it would count every + // row. A query with a single DISTINCT argument is rewritten by the + // optimizer to group by those values; it arrives here without the + // flag and still runs. + if acc_args.is_distinct { + return datafusion::common::not_impl_err!( + "DISTINCT is not supported for the Python aggregate UDF {}", + self.name + ); + } instantiate_accumulator(&self.accumulator) } @@ -310,6 +322,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 6376c81a8..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}; @@ -209,6 +211,17 @@ 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(); + 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 +260,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 +273,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..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}; @@ -216,6 +218,17 @@ 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 { + validate_pycapsule(capsule, "datafusion_window_udf")?; + 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 +275,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/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/docs/source/user-guide/common-operations/expressions.md b/docs/source/user-guide/common-operations/expressions.md index 7985705ef..ab42172ad 100644 --- a/docs/source/user-guide/common-operations/expressions.md +++ b/docs/source/user-guide/common-operations/expressions.md @@ -157,7 +157,9 @@ In this example, the `repeated_array` column will contain `[[1, 2, 3], [1, 2, 3] Some array functions take a *lambda function*: a small function that runs once per element. {py:func}`~datafusion.functions.array_transform` maps a lambda over every element, {py:func}`~datafusion.functions.array_filter` keeps the elements -for which a predicate lambda is true, and +for which a predicate lambda is true, +{py:func}`~datafusion.functions.array_first` returns the first element that +satisfies a predicate lambda, and {py:func}`~datafusion.functions.array_any_match` returns whether any element satisfies a predicate lambda. (Functions that take another function as an argument are sometimes called *higher-order* functions.) @@ -173,6 +175,7 @@ ctx = SessionContext() df = ctx.from_pydict({"a": [[1, 2, 3], [4, 5]]}) df.select(f.array_transform(col("a"), lambda v: v * 2).alias("doubled")) df.select(f.array_filter(col("a"), lambda v: v > 2).alias("big_only")) +df.select(f.array_first(col("a"), lambda v: v > 2).alias("first_big")) df.select(f.array_any_match(col("a"), lambda v: v > 3).alias("has_big")) ``` diff --git a/docs/source/user-guide/common-operations/functions.md b/docs/source/user-guide/common-operations/functions.md index 9939e2f00..d5ddffb76 100644 --- a/docs/source/user-guide/common-operations/functions.md +++ b/docs/source/user-guide/common-operations/functions.md @@ -158,3 +158,38 @@ df = df.fill_null("missing", subset=["name", "category"]) ``` The fill value will be cast to match each column's type. If casting fails for a column, that column remains unchanged. + +## fill_nan + +The `fill_nan()` method replaces NaN values in floating-point columns. NaN is +distinct from NULL, which `fill_null()` handles: + +```python +# Replace NaN with 0.0 in every floating-point column +df = df.fill_nan(0.0) + +# Replace NaN only in specific columns +df = df.fill_nan(0.0, subset=["price"]) +``` + +(fill_column_names)= + +## Column names with uppercase letters or dots + +`fill_null()` and `fill_nan()` fail on a DataFrame that has any column whose +name contains an uppercase letter or a `.`, even when that column is not in +`subset` +([apache/datafusion#25829](https://github.com/apache/datafusion/issues/25829)): + +```python +df = ctx.from_pydict({"Price": [float("nan"), 2.0], "qty": [float("nan"), 1.0]}) +df.fill_nan(0.0, subset=["qty"]) +# Schema error: No field named price. Did you mean '..."Price"'? +``` + +Until that is fixed, rename such columns first. Quote the old name so it is +not normalized: + +```python +df.with_column_renamed('"Price"', "price").fill_nan(0.0, subset=["qty"]) +``` diff --git a/docs/source/user-guide/common-operations/windows.md b/docs/source/user-guide/common-operations/windows.md index bee96c820..31fcdbe01 100644 --- a/docs/source/user-guide/common-operations/windows.md +++ b/docs/source/user-guide/common-operations/windows.md @@ -135,11 +135,54 @@ df.select( ) ``` +(window_function_chaining)= + +#### Chaining onto a window function + +Set a window function's `partition_by`, `order_by`, and `window_frame` in one +place: its keyword arguments, a single `over()`, or one builder chain ending in +`build()`. Chaining a builder method or `over()` onto a window function that +already has any of them raises: + +```python +# Raises: lead already has window options (order_by) +f.lead(col("v"), order_by="t").over(Window(partition_by=[col("g")])) + +# Set them together instead. +f.lead(col("v")).over(Window(partition_by=[col("g")], order_by="t")) +``` + +A built window function stores a concrete frame with no record of whether you +chose it or it was derived from `order_by`, so the options cannot be merged +without guessing. Merging may become possible once +[apache/datafusion#25934](https://github.com/apache/datafusion/issues/25934) +is resolved. + +The `null_treatment` already set is kept, as are `filter` and `distinct` on an +aggregate used as a window function. The whole-partition frame counts as no +frame, so adding an `order_by` derives the running frame, even when you passed +that frame explicitly: + +```python +whole = WindowFrame("rows", None, None) # same as the no-order_by default + +# Both give a running sum. +f.sum(col("v")).over(Window()).order_by(col("v")).build() +f.sum(col("v")).over(Window(window_frame=whole)).order_by(col("v")).build() +``` + +To keep that 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 -should be treated. In order to do this you need to use the builder function. In future releases -we expect this to be simplified in the interface. +should be treated. Pass `null_treatment` in the `Window`, or set it on the aggregate itself, which +`over()` keeps (see {ref}`aggregate_over_options`). One common usage for handling nulls is the case where you want to find the last value up to the current row. In the following example we demonstrate how setting the null treatment to ignore @@ -196,6 +239,30 @@ df.select( ) ``` +(aggregate_over_options)= + +### Options set on the aggregate + +`over()` keeps the `filter`, `distinct`, and `null_treatment` options an +aggregate was built with: + +```python +# Averages the distinct values 1.0 and 4.0. +f.avg(col("v"), distinct=True).over(Window()) +``` + +An aggregate's `order_by` raises, as `ORDER BY` inside an aggregate call does +with `OVER` in SQL. A window does not pass an ordering to the aggregate, so +the `order_by` in the `Window` only sets the frame and the order of rows. The +one exception is a `WITHIN GROUP` function such as +{py:func}`~datafusion.functions.percentile_cont`, which computes ascending as a +window: an ascending `sort_expression` is accepted, and a descending one raises. + +```python +f.percentile_cont(col("v"), 0.25).over(Window()) # ascending, accepted +f.percentile_cont(col("v").sort(ascending=False), 0.25).over(Window()) # raises +``` + ## Available Functions The possible window functions are: diff --git a/docs/source/user-guide/upgrade-guides.md b/docs/source/user-guide/upgrade-guides.md index 0b8e94bb5..3194eb61f 100644 --- a/docs/source/user-guide/upgrade-guides.md +++ b/docs/source/user-guide/upgrade-guides.md @@ -198,6 +198,179 @@ 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 when passed by keyword. + +### 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. On an aggregate they are now kept, which can change results: + +```python +e = f.string_agg(col("s"), ",", order_by="s") +e.distinct().build() # before: order_by dropped; after: kept +``` + +On a window function that already has a `partition_by`, `order_by`, or +`window_frame`, chaining now raises instead of dropping them. Set them in one +place, as described in {ref}`window_function_chaining`: + +```python +e = f.lead(col("v"), 1, partition_by=[col("g")], order_by="t") +e.over(Window(order_by="t")) # before: partition dropped; after: raises +f.lead(col("v"), 1).over(Window(partition_by=[col("g")], order_by="t")) # after +``` + +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. + +An option that does not apply to the function now raises as soon as it is +set, anywhere in the chain. `filter` and `distinct` need an aggregate, +including one used as a window function, and `partition_by` and +`window_frame` need a window function. Later in a chain these were silently +dropped: + +```python +f.sum(col("v")).filter(col("v") > lit(1)).partition_by(col("g")) +# before: partition_by dropped; after: raises +``` + +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. + +### `over()` keeps options set on an aggregate + +`Expr.over()` on an aggregate used to drop the `filter`, `distinct`, +`null_treatment`, and `order_by` options it was built with. The first three are +now kept, which can change results: + +```python +f.avg(col("v"), distinct=True).over(Window()) +# v = [1, 1, 4]; before: 2.0, after: 2.5 +``` + +An `order_by` on the aggregate now raises instead of being dropped, as it does +with `OVER` in SQL. Remove it, or move it into the `Window` if it was meant to +order the rows: + +```python +# before: order_by dropped; after: raises +f.first_value(col("v"), order_by=col("i").sort(ascending=False)).over( + Window(partition_by=[col("g")]) +) + +# after: the Window orders the rows the aggregate sees +f.first_value(col("v")).over( + Window(partition_by=[col("g")], order_by=[col("i").sort(ascending=False)]) +) +``` + +A `WITHIN GROUP` function such as `percentile_cont` still accepts +an ascending `sort_expression`, and raises on a descending one, which used to +give the ascending result. See {ref}`aggregate_over_options`. + +### Python aggregate UDFs reject `DISTINCT` + +A Python {py:class}`~datafusion.user_defined.Accumulator` cannot deduplicate its +input, so a Python aggregate UDF with `DISTINCT` counted every row. It now +raises instead of returning that result: + +```python +my_sum(col("v")).over(Window()).distinct().build() +# v = [1, 1, 1, 5]; before: 8.0; after: DISTINCT is not supported ... +``` + +When the optimizer rewrites the query to group by the distinct values first, +as it does for SQL's `SELECT my_sum(DISTINCT v) FROM t`, the query still runs +and gives the distinct result. + +### 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. +Used as aggregates, they now match `WITHIN GROUP (ORDER BY ... DESC)` in SQL +(see {ref}`aggregate_over_options` for their use in a window): + +```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` +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 +`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/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/pyproject.toml b/pyproject.toml index dae84dd44..604d3f847 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,7 +54,7 @@ dependencies = [ "cloudpickle>=2.0", "pyarrow>=16.0.0;python_version<'3.14'", "pyarrow>=22.0.0;python_version>='3.14'", - "typing-extensions;python_version<'3.13'", + "typing-extensions>=4.12;python_version<'3.13'", ] dynamic = ["version"] @@ -197,7 +197,7 @@ ignore-words-list = ["IST", "ans"] # Keep typing-extensions on a single version across all supported Pythons. # Adding the docs `myst-nb` stack pulled a newer typing-extensions only on # Python < 3.11, splitting it into two locked versions. Our own -# `typing-extensions; python_full_version < '3.13'` dependency then spanned +# `typing-extensions>=4.12; python_full_version < '3.13'` dependency then spanned # that split, which uv records as an under-specified lock entry that older # uv versions (such as the one in CI) refuse to parse. Pinning to one # version removes the fork. 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", ] 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/context.py b/python/datafusion/context.py index 4043e357b..c663dea2b 100644 --- a/python/datafusion/context.py +++ b/python/datafusion/context.py @@ -85,11 +85,16 @@ if TYPE_CHECKING: import pathlib + import sys from collections.abc import Iterable, Sequence import pandas as pd import polars as pl # type: ignore[import] - from _typeshed import CapsuleType as _PyCapsule + + if sys.version_info >= (3, 13): + from types import CapsuleType as _PyCapsule + else: + from typing_extensions import CapsuleType as _PyCapsule from datafusion.catalog import CatalogProvider, Table from datafusion.common import DFSchema diff --git a/python/datafusion/dataframe.py b/python/datafusion/dataframe.py index de00ff474..9c6feeb66 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,18 @@ 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. + + Raises: + ValueError: If ``show_statistics`` is set with ``analyze``, or + ``analyze_level`` or ``analyze_categories`` is set without it. Examples: Show the plan in tree format: @@ -1229,9 +1270,31 @@ 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 + """ + if analyze and show_statistics is not None: + msg = "show_statistics cannot be combined with analyze" + raise ValueError(msg) + if not analyze and analyze_level is not None: + msg = "analyze_level requires analyze" + raise ValueError(msg) + if not analyze and analyze_categories is not None: + msg = "analyze_categories requires analyze" + raise ValueError(msg) 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``. @@ -1855,7 +1918,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 @@ -1868,13 +1932,54 @@ 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 + - Fails on a DataFrame with an uppercase or dotted column name; see + :ref:`fill_column_names` """ + if subset is not None and len(subset) == 0: + return self 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. Fails on a DataFrame with an uppercase or + dotted column name; see :ref:`fill_column_names`. + + 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; an empty list fills none. + + 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 + + >>> df.fill_nan(0.0, subset=[]).collect_column("b")[0].as_py() + nan + """ + if subset is not None and len(subset) == 0: + return self + return DataFrame(self.df.fill_nan(value, subset)) + class InsertOp(Enum): """Insert operation mode. diff --git a/python/datafusion/expr.py b/python/datafusion/expr.py index 18ce3554d..fb005896f 100644 --- a/python/datafusion/expr.py +++ b/python/datafusion/expr.py @@ -1037,8 +1037,9 @@ def filter(self, filter: Expr) -> ExprFuncBuilder: """Filter an aggregate function. This function will create an :py:class:`ExprFuncBuilder` that can be used to - 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. + set parameters for either window or aggregate functions. It raises on a window + function unless that is an aggregate used as one. If used on any other type of + expression, an error will be generated when ``build()`` is called. """ return ExprFuncBuilder(self.expr.filter(filter.expr)) @@ -1046,8 +1047,9 @@ def distinct(self) -> ExprFuncBuilder: """Only evaluate distinct values for an aggregate function. This function will create an :py:class:`ExprFuncBuilder` that can be used to - 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. + set parameters for either window or aggregate functions. It raises on a window + function unless that is an aggregate used as one. If used on any other type of + expression, an error will be generated when ``build()`` is called. """ return ExprFuncBuilder(self.expr.distinct()) @@ -1064,27 +1066,38 @@ def partition_by(self, *partition_by: Expr) -> ExprFuncBuilder: """Set the partitioning for a window function. This function will create an :py:class:`ExprFuncBuilder` that can be used to - 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. + set parameters for either window or aggregate functions. It raises on an + aggregate function that has not been turned into a window function with + :py:meth:`over`. 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])) def window_frame(self, window_frame: WindowFrame) -> ExprFuncBuilder: - """Set the frame fora window function. + """Set the frame for a window function. This function will create an :py:class:`ExprFuncBuilder` that can be used to - 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. + set parameters for either window or aggregate functions. It raises on an + aggregate function that has not been turned into a window function with + :py:meth:`over`. 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)) 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. + On an aggregate, the ``null_treatment``, ``filter``, and ``distinct`` options + it was built with are kept, and an ``order_by`` raises; see + :ref:`aggregate_over_options`. + + On a window function that already has a ``partition_by``, ``order_by``, or + ``window_frame``, this raises; see :ref:`window_function_chaining`. + Args: window: Window definition """ @@ -1534,6 +1547,15 @@ def cot(self) -> Expr: class ExprFuncBuilder: + """Sets the options of an aggregate or window function. + + ``filter`` and ``distinct`` apply to an aggregate, including one used as a + window function; ``partition_by`` and ``window_frame`` apply to a window + function. Starting a builder from a window function that already has a + ``partition_by``, ``order_by``, or ``window_frame`` raises; see + :ref:`window_function_chaining`. + """ + def __init__(self, builder: expr_internal.ExprFuncBuilder) -> None: self.builder = builder diff --git a/python/datafusion/extensions.py b/python/datafusion/extensions.py index edb9a513e..cbcc1a9ac 100644 --- a/python/datafusion/extensions.py +++ b/python/datafusion/extensions.py @@ -49,7 +49,12 @@ from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable 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 from datafusion.context import SessionContext from datafusion.user_defined import ( diff --git a/python/datafusion/functions/__init__.py b/python/datafusion/functions/__init__.py index 291957490..25762374a 100644 --- a/python/datafusion/functions/__init__.py +++ b/python/datafusion/functions/__init__.py @@ -40,7 +40,7 @@ import inspect import warnings -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, overload import pyarrow as pa @@ -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", @@ -201,6 +209,7 @@ def _warn_if_expr_for_literal_arg( "exp", "extract", "factorial", + "file_row_index", "find_in_set", "first_value", "flatten", @@ -216,6 +225,7 @@ def _warn_if_expr_for_literal_arg( "in_list", "initcap", "inner_product", + "input_file_name", "instr", "is_nan", "isnan", @@ -230,9 +240,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 +257,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 +275,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 +287,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", @@ -319,6 +336,7 @@ def _warn_if_expr_for_literal_arg( "power", "quantile_cont", "radians", + "rand", "random", "range", "rank", @@ -367,6 +385,7 @@ def _warn_if_expr_for_literal_arg( "substr", "substr_index", "substring", + "substring_index", "sum", "tan", "tanh", @@ -467,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: @@ -712,6 +756,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``. @@ -874,8 +958,9 @@ 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 - the options ``order_by``, ``distinct``, and ``null_treatment``. + If using the builder functions described in :ref:`aggregation` this function ignores + the options ``order_by`` and ``null_treatment``. ``distinct`` counts the distinct + values of the constant ``1``, so the result is 1 for any non-empty input. Args: filter: If provided, only count rows for which the filter is True @@ -1069,8 +1154,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() @@ -1078,8 +1163,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: @@ -1557,8 +1654,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() @@ -1566,8 +1663,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: @@ -2079,8 +2188,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() @@ -2088,8 +2197,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: @@ -2253,8 +2374,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() @@ -2263,7 +2386,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)) @@ -2287,6 +2420,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. @@ -2890,8 +3032,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() @@ -2899,8 +3041,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: @@ -2973,18 +3127,84 @@ 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. +_SERIES_SIGNATURE = inspect.Signature( + [ + inspect.Parameter("start", inspect.Parameter.POSITIONAL_OR_KEYWORD), + inspect.Parameter("stop", inspect.Parameter.POSITIONAL_OR_KEYWORD), + inspect.Parameter( + "step", inspect.Parameter.POSITIONAL_OR_KEYWORD, default=None + ), + ] +) + + +def _series( + fn: Callable[..., Any], name: str, args: tuple[Any, ...], kwargs: dict[str, Any] +) -> Expr: + # ``stop`` is required and ``start`` defaults to 0, as in numpy.arange: a + # lone positional argument is ``stop``. ``start=`` alone is then a missing + # ``stop``, so it cannot silently become the upper bound. + if len(args) == 1 and not kwargs: + args = (0, *args) + elif not args and "start" not in kwargs: + kwargs = {"start": 0, **kwargs} + try: + bound = _SERIES_SIGNATURE.bind(*args, **kwargs) + except TypeError as e: + msg = f"{name}() {e}" + raise TypeError(msg) from None + start, stop = bound.arguments["start"], bound.arguments["stop"] + if stop is None: + msg = f"{name}() stop cannot be None" + raise TypeError(msg) + step = coerce_to_expr_or_none(bound.arguments.get("step")) + return Expr( + fn( + coerce_to_expr(start).expr, + coerce_to_expr(stop).expr, + step.expr if step is not None else None, + ) + ) + + +@overload +def range(stop: Expr | int) -> Expr: ... + + +@overload +def range( + start: Expr | int, + stop: Expr | int, + step: Expr | int | None = None, +) -> Expr: ... + + +def range(*args: Any, **kwargs: Any) -> Expr: + """Create a list of values from ``start`` up to, but excluding, ``stop``. + + With a single argument it is ``stop`` 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", args, kwargs) def uuid() -> Expr: @@ -3432,6 +3652,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. @@ -3687,6 +3965,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 +4126,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. @@ -4696,43 +5144,66 @@ 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. +@overload +def gen_series(stop: Expr | int) -> Expr: ... + - Unlike :py:func:`range`, this includes the upper bound. +@overload +def gen_series( + start: Expr | int, + stop: Expr | int, + step: Expr | int | None = None, +) -> Expr: ... + + +def gen_series(*args: Any, **kwargs: Any) -> Expr: + """Creates a list of values from ``start`` up to and including ``stop``. + + Unlike :py:func:`range`, this includes the upper bound. With a single + argument it is ``stop`` 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", args, kwargs) -def generate_series(start: Expr, stop: Expr, step: Expr | None = None) -> Expr: - """Creates a list of values in the range between start and stop. +@overload +def generate_series(stop: Expr | int) -> Expr: ... + + +@overload +def generate_series( + start: Expr | int, + stop: Expr | int, + step: Expr | int | None = None, +) -> Expr: ... + - Unlike :py:func:`range`, this includes the upper bound. +def generate_series(*args: Any, **kwargs: Any) -> Expr: + """Creates a list of values from ``start`` up to and including ``stop``. See Also: This is an alias for :py:func:`gen_series`. """ - return gen_series(start, stop, step) + return _series(f.gen_series, "generate_series", args, kwargs) def flatten(array: Expr) -> Expr: @@ -4933,7 +5404,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: @@ -4969,7 +5440,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: @@ -5016,8 +5487,10 @@ 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 - the options ``order_by``, ``null_treatment``, and ``distinct``. + If using the builder functions described in :ref:`aggregation` this function ignores + the options ``null_treatment`` and ``distinct``. A chained ``order_by`` sets only + the sort direction; the percentile is still computed over ``sort_expression``, + so pass the same expression. Args: sort_expression: Values for which to find the approximate percentile @@ -5065,8 +5538,10 @@ 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 - the options ``order_by``, ``null_treatment``, and ``distinct``. + If using the builder functions described in :ref:`aggregation` this function ignores + the option ``null_treatment`` and rejects ``distinct``. A chained ``order_by`` + sets only the sort direction; the percentile is still computed over + ``sort_expression``, so pass the same expression. Args: sort_expression: Values for which to find the approximate percentile @@ -5107,9 +5582,33 @@ def approx_percentile_cont_with_weight( ) +def _check_distinct(name: str, distinct: object, *shifted: str) -> None: + """Raise when a positional ``filter`` from before ``distinct`` landed here. + + Args: + name: Function name for the message. + distinct: The value received for ``distinct``. + shifted: Names of the arguments that follow ``distinct``. + + Examples: + >>> dfn.functions._check_distinct("mean", True, "filter") + >>> dfn.functions._check_distinct("mean", dfn.col("a"), "filter") + Traceback (most recent call last): + ... + TypeError: mean() distinct must be a bool, got Expr; pass filter by keyword + """ + if distinct is None or isinstance(distinct, Expr): + msg = ( + f"{name}() distinct must be a bool, got {type(distinct).__name__}; " + f"pass {' and '.join(shifted)} by keyword" + ) + raise TypeError(msg) + + 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. @@ -5117,12 +5616,15 @@ 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 - the options ``order_by``, ``null_treatment``, and ``distinct``. + If using the builder functions described in :ref:`aggregation` this function ignores + the option ``null_treatment``. A chained ``order_by`` sets only the sort + direction; the percentile is still computed over ``sort_expression``, so pass + the same expression. 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: @@ -5142,15 +5644,29 @@ 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 """ + _check_distinct("percentile_cont", distinct, "filter") 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. @@ -5158,7 +5674,10 @@ def quantile_cont( See Also: This is an alias for :py:func:`percentile_cont`. """ - return percentile_cont(sort_expression, percentile, filter) + _check_distinct("quantile_cont", distinct, "filter") + return percentile_cont( + sort_expression, percentile, distinct=distinct, filter=filter + ) def array_agg( @@ -5173,7 +5692,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: @@ -5287,7 +5806,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: @@ -5330,7 +5849,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: @@ -5369,7 +5888,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: @@ -5413,7 +5932,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: @@ -5454,7 +5973,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: @@ -5496,7 +6015,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: @@ -5525,13 +6044,18 @@ 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) + _check_distinct("mean", distinct, "filter") + return avg(expression, distinct=distinct, filter=filter) def median( @@ -5542,7 +6066,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: @@ -5576,7 +6100,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: @@ -5614,7 +6138,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: @@ -5655,7 +6179,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: @@ -5687,7 +6211,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: @@ -5740,7 +6264,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: @@ -5781,7 +6305,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: @@ -5829,7 +6353,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: @@ -5870,7 +6394,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: @@ -5911,7 +6435,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: @@ -5952,7 +6476,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: @@ -5995,7 +6519,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: @@ -6036,7 +6560,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: @@ -6077,7 +6601,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: @@ -6118,7 +6642,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: @@ -6159,7 +6683,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: @@ -6200,7 +6724,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: @@ -6256,7 +6780,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: @@ -6313,7 +6837,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: @@ -6360,16 +6884,55 @@ def nth_value( ) -def bit_and(expression: Expr, filter: Expr | None = None) -> Expr: +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, 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``. + If using the builder functions described in :ref:`aggregation` this function ignores + 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: @@ -6391,20 +6954,24 @@ def bit_and(expression: Expr, filter: Expr | None = None) -> Expr: >>> result.collect_column("v")[0].as_py() 5 """ + _check_distinct("bit_and", distinct, "filter") 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``. + If using the builder functions described in :ref:`aggregation` this function ignores + 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: @@ -6428,8 +6995,9 @@ def bit_or(expression: Expr, filter: Expr | None = None) -> Expr: >>> result.collect_column("v")[0].as_py() 6 """ + _check_distinct("bit_or", distinct, "filter") 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( @@ -6439,7 +7007,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: @@ -6478,7 +7046,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: @@ -6517,7 +7085,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: @@ -6556,6 +7124,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 | None = None, ) -> Expr: """Create a lead window function. @@ -6577,7 +7146,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 @@ -6586,6 +7155,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() @@ -6608,6 +7179,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) @@ -6622,6 +7203,9 @@ def lead( default_value, partition_by=partition_by_raw, order_by=order_by_raw, + null_treatment=( + null_treatment.value if null_treatment is not None else None + ), ) ) @@ -6632,6 +7216,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 | None = None, ) -> Expr: """Create a lag window function. @@ -6659,6 +7244,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() @@ -6681,6 +7268,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) @@ -6695,6 +7292,9 @@ def lag( default_value, partition_by=partition_by_raw, order_by=order_by_raw, + null_treatment=( + null_treatment.value if null_treatment is not None else None + ), ) ) @@ -7054,6 +7654,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: @@ -7063,12 +7664,13 @@ 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 - the options ``distinct`` and ``null_treatment``. + If using the builder functions described in :ref:`aggregation` this function ignores + 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. @@ -7091,7 +7693,16 @@ 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' """ + _check_distinct("string_agg", distinct, "filter", "order_by") order_by_raw = sort_list_to_raw_sort_list(order_by) filter_raw = filter.expr if filter is not None else None @@ -7099,6 +7710,7 @@ def string_agg( f.string_agg( expression.expr, delimiter, + distinct=distinct, filter=filter_raw, order_by=order_by_raw, ) diff --git a/python/datafusion/functions/spark.py b/python/datafusion/functions/spark.py index 0a0f41400..c4002670d 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. @@ -593,6 +611,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. @@ -611,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: @@ -624,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( @@ -738,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``. @@ -816,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``. @@ -933,6 +985,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. @@ -959,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. @@ -1114,6 +1191,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. @@ -1127,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. @@ -1184,6 +1286,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 +1333,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 +1541,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). @@ -1456,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, @@ -1499,6 +1677,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. @@ -1520,6 +1711,30 @@ 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: Expr | str, *cols: Expr | str) -> Expr: + """Spark ``printf``: printf-style 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 Expr( + _f.format_string(_to_raw_expr(format), *[_to_raw_expr(c) for c in cols]) + ) + + def space(col: Expr | int) -> Expr: """Spark ``space``: string of n spaces. @@ -1554,6 +1769,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. @@ -1735,6 +1972,7 @@ def url_encode(str: Expr) -> Expr: # String "ascii", # Aggregate + "atan2", "avg", "base64", "bin", @@ -1747,11 +1985,15 @@ def url_encode(str: Expr) -> Expr: "bitmap_count", "bitwise_not", "ceil", + "ceiling", "char", + "char_length", + "character_length", "collect_list", "collect_set", "concat", # Hash + "concat_ws", "crc32", "csc", "date_add", @@ -1759,14 +2001,19 @@ 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", "if_", "ilike", "is_valid_utf8", @@ -1784,15 +2031,21 @@ def url_encode(str: Expr) -> Expr: "map_from_entries", "minute", "modulus", + "monthname", "negative", "next_day", # URL "parse_url", "pmod", + "pow", + "power", + "printf", + "quote", "rint", "round", "sec", "second", + "sha", "sha1", "sha2", "shiftleft", @@ -1806,6 +2059,7 @@ def url_encode(str: Expr) -> Expr: "space", "spark_cast", "str_to_map", + "substr", "substring", "time_trunc", "to_utc_timestamp", @@ -1821,6 +2075,7 @@ def url_encode(str: Expr) -> Expr: "unix_seconds", "url_decode", "url_encode", + "weekday", "width_bucket", "xxhash64", ] diff --git a/python/datafusion/user_defined.py b/python/datafusion/user_defined.py index f8d273177..1d8ccbe3f 100644 --- a/python/datafusion/user_defined.py +++ b/python/datafusion/user_defined.py @@ -20,6 +20,7 @@ from __future__ import annotations import functools +import sys from abc import ABCMeta, abstractmethod from enum import Enum from typing import TYPE_CHECKING, Any, Protocol, TypeGuard, TypeVar, cast, overload @@ -30,9 +31,14 @@ from datafusion import SessionContext from datafusion.expr import Expr -if TYPE_CHECKING: - from _typeshed import CapsuleType as _PyCapsule +# Imported at runtime so ``typing.get_type_hints`` resolves the capsule +# overloads; typing_extensions is a runtime dependency below 3.13. +if sys.version_info >= (3, 13): + from types import CapsuleType as _PyCapsule +else: + from typing_extensions import CapsuleType as _PyCapsule +if TYPE_CHECKING: _R = TypeVar("_R", bound=pa.Array) from collections.abc import Callable, Sequence @@ -216,9 +222,9 @@ def __init__( def _from_internal(cls, internal: df_internal.ScalarUDF) -> ScalarUDF: """Wrap an already-constructed internal ``ScalarUDF`` handle. - Used by :py:meth:`SessionContext.udf` to surface a function looked - up from the session's function registry without re-running - :py:meth:`__init__`. + Used by :py:meth:`SessionContext.udf` and :py:meth:`from_pycapsule` + to wrap a handle from the session's function registry or an FFI + capsule without re-running :py:meth:`__init__`. """ wrapper = cls.__new__(cls) wrapper._udf = internal @@ -283,8 +289,12 @@ def udf( @staticmethod def udf(func: ScalarUDFExportable) -> ScalarUDF: ... + @overload + @staticmethod + def udf(func: _PyCapsule) -> ScalarUDF: ... + @staticmethod - def udf(*args: Any, **kwargs: Any): # noqa: D417 + def udf(*args: Any, **kwargs: Any): # noqa: D417, C901 """Create a new User-Defined Function (UDF). This class can be used both as either a function or a decorator. @@ -388,7 +398,11 @@ def wrapper(*args: Any, **kwargs: Any) -> Callable: return decorator - if hasattr(args[0], "__datafusion_scalar_udf__"): + if not args and "func" in kwargs: + args = (kwargs.pop("func"),) + 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]): @@ -398,12 +412,16 @@ 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): + return ScalarUDF._from_internal(df_internal.ScalarUDF.from_pycapsule(func)) + + func = cast("ScalarUDFExportable", func) name = str(func.__class__) return ScalarUDF( name=name, @@ -529,9 +547,9 @@ def __init__( def _from_internal(cls, internal: df_internal.AggregateUDF) -> AggregateUDF: """Wrap an already-constructed internal ``AggregateUDF`` handle. - Used by :py:meth:`SessionContext.udaf` to surface a function looked - up from the session's function registry without re-running - :py:meth:`__init__`. + Used by :py:meth:`SessionContext.udaf` and :py:meth:`from_pycapsule` + to wrap a handle from the session's function registry or an FFI + capsule without re-running :py:meth:`__init__`. """ wrapper = cls.__new__(cls) wrapper._udaf = internal @@ -732,7 +750,11 @@ def wrapper(*args: Any, **kwargs: Any) -> Expr: return decorator - if hasattr(args[0], "__datafusion_aggregate_udf__") or _is_pycapsule(args[0]): + if not args and "accum" in kwargs: + args = (kwargs.pop("accum"),) + 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]): @@ -749,9 +771,9 @@ def from_pycapsule(func: AggregateUDFExportable | _PyCapsule) -> AggregateUDF: AggregateUDF that is exported via the FFI bindings. """ if _is_pycapsule(func): - aggregate = cast("AggregateUDF", object.__new__(AggregateUDF)) - aggregate._udaf = df_internal.AggregateUDF.from_pycapsule(func) - return aggregate + return AggregateUDF._from_internal( + df_internal.AggregateUDF.from_pycapsule(func) + ) capsule = cast("AggregateUDFExportable", func) name = str(capsule.__class__) @@ -961,9 +983,9 @@ def __init__( def _from_internal(cls, internal: df_internal.WindowUDF) -> WindowUDF: """Wrap an already-constructed internal ``WindowUDF`` handle. - Used by :py:meth:`SessionContext.udwf` to surface a function looked - up from the session's function registry without re-running - :py:meth:`__init__`. + Used by :py:meth:`SessionContext.udwf` and :py:meth:`from_pycapsule` + to wrap a handle from the session's function registry or an FFI + capsule without re-running :py:meth:`__init__`. """ wrapper = cls.__new__(cls) wrapper._udwf = internal @@ -1011,6 +1033,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 +1105,11 @@ 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 not args and "func" in kwargs: + args = (kwargs.pop("func"),) + 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]): @@ -1146,12 +1180,16 @@ 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): + return WindowUDF._from_internal(df_internal.WindowUDF.from_pycapsule(func)) + + func = cast("WindowUDFExportable", func) name = str(func.__class__) return WindowUDF( name=name, @@ -1257,6 +1295,8 @@ def udtf(*args: Any, with_session: bool = False, **kwargs: Any): :class:`SessionContext` injected as a ``session`` keyword argument on each invocation. """ + if not args and "func" in kwargs: + args = (kwargs.pop("func"),) if args and callable(args[0]): # Case 1: Used as a function, require the first parameter to be callable return TableFunction._create_table_udf( diff --git a/python/tests/test_aggregation.py b/python/tests/test_aggregation.py index ef51343aa..38baef954 100644 --- a/python/tests/test_aggregation.py +++ b/python/tests/test_aggregation.py @@ -317,11 +317,55 @@ 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]} + + +@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]), @@ -350,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"), [ @@ -477,6 +529,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: @@ -496,3 +553,116 @@ def test_string_agg(name, expr, result) -> None: } df.show() assert df.collect()[0].to_pydict() == expected + + +_FILTER = column("b") > lit(1) + + +@pytest.mark.parametrize( + ("call", "match"), + [ + pytest.param( + lambda: f.percentile_cont(column("a"), 0.5, _FILTER), + r"percentile_cont\(\).*pass filter by keyword", + id="percentile_cont", + ), + pytest.param( + lambda: f.quantile_cont(column("a"), 0.5, _FILTER), + r"quantile_cont\(\).*pass filter by keyword", + id="quantile_cont", + ), + pytest.param( + lambda: f.mean(column("a"), _FILTER), + r"mean\(\).*pass filter by keyword", + id="mean", + ), + pytest.param( + lambda: f.bit_and(column("a"), _FILTER), + r"bit_and\(\).*pass filter by keyword", + id="bit_and", + ), + pytest.param( + lambda: f.bit_or(column("a"), _FILTER), + r"bit_or\(\).*pass filter by keyword", + id="bit_or", + ), + pytest.param( + lambda: f.string_agg(column("a"), ",", _FILTER, column("b")), + r"string_agg\(\).*pass filter and order_by by keyword", + id="string_agg filter", + ), + pytest.param( + lambda: f.string_agg(column("a"), ",", None, column("b")), + r"string_agg\(\).*pass filter and order_by by keyword", + id="string_agg None placeholder", + ), + ], +) +def test_positional_filter_names_the_function(call, match) -> None: + # ``distinct`` was inserted before ``filter``. A call written for the old + # signature must name the function and the fix, not fail inside PyO3. The + # ``None`` placeholder would otherwise run with order_by shifted into + # filter. + with pytest.raises(TypeError, match=match): + call() + + +def test_string_agg_accepts_numpy_bool_distinct() -> None: + np = pytest.importorskip("numpy") + df = SessionContext().from_pydict({"a": ["x", "y", "x"]}) + expr = f.string_agg(column("a"), ",", distinct=np.True_, order_by="a") + assert df.aggregate([], [expr.alias("s")]).collect_column("s")[0].as_py() == "x,y" + + +@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 diff --git a/python/tests/test_dataframe.py b/python/tests/test_dataframe.py index bb21a3974..30443b3cf 100644 --- a/python/tests/test_dataframe.py +++ b/python/tests/test_dataframe.py @@ -24,6 +24,7 @@ from pathlib import Path from typing import Any +import numpy as np import pyarrow as pa import pyarrow.parquet as pq import pytest @@ -46,7 +47,12 @@ from datafusion import ( functions as f, ) -from datafusion.dataframe import DataFrameWriteOptions +from datafusion.common import NullTreatment +from datafusion.dataframe import ( + DataFrameWriteOptions, + ExplainAnalyzeLevel, + ExplainMetricCategory, +) from datafusion.dataframe_formatter import ( DataFrameHtmlFormatter, configure_formatter, @@ -1077,6 +1083,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( @@ -1161,6 +1187,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"), [ @@ -3396,6 +3430,26 @@ def test_fill_null_specific_types(null_df): ] +@pytest.mark.parametrize( + ("subset", "expected_cols"), + [ + pytest.param([], [], id="empty list"), + pytest.param(np.array([], dtype=str), [], id="empty numpy"), + pytest.param(np.array(["int_col"]), ["int_col"], id="numpy"), + ], +) +def test_fill_null_subset_forms(null_df, subset, expected_cols): + # An empty subset fills nothing rather than everything, and a numpy array + # is accepted without being tested for truthiness. + result = null_df.fill_null(0, subset=subset).to_pydict() + original = null_df.to_pydict() + for name, values in result.items(): + if name in expected_cols: + assert None not in values + else: + assert values == original[name] + + def test_fill_null_immutability(null_df): """Test that original DataFrame is unchanged after fill_null.""" # Get original values with nulls @@ -3450,6 +3504,58 @@ 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): + df = _nan_df(ctx) + assert df.fill_nan(0.0).schema() == df.schema() + result = df.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]) + + +@pytest.mark.parametrize( + ("subset", "f64_filled", "f32_filled"), + [ + pytest.param([], False, False, id="empty list"), + pytest.param(np.array([], dtype=str), False, False, id="empty numpy"), + pytest.param(np.array(["f32"]), False, True, id="numpy"), + ], +) +def test_fill_nan_subset_forms(ctx, subset, f64_filled, f32_filled): + # An empty subset fills nothing rather than everything, and a numpy array + # is accepted without being tested for truthiness. + result = _nan_df(ctx).fill_nan(0.0, subset=subset).to_pydict() + assert _is_nan(result["f64"][1]) != f64_filled + assert _is_nan(result["f32"][0]) != f32_filled + + _slow_udf_started = threading.Event() @@ -3806,6 +3912,88 @@ 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( + ("kwargs", "message"), + [ + pytest.param( + {"analyze": True, "show_statistics": True}, + "show_statistics cannot be combined with analyze", + id="show_statistics_true_with_analyze", + ), + pytest.param( + {"analyze": True, "show_statistics": False}, + "show_statistics cannot be combined with analyze", + id="show_statistics_false_with_analyze", + ), + pytest.param( + {"analyze_level": ExplainAnalyzeLevel.DEV}, + "analyze_level requires analyze", + id="analyze_level_without_analyze", + ), + pytest.param( + {"analyze_categories": [ExplainMetricCategory.ROWS]}, + "analyze_categories requires analyze", + id="analyze_categories_without_analyze", + ), + ], +) +def test_explain_rejects_options_that_would_be_ignored(capsys, kwargs, message): + # Upstream would silently ignore these; SQL's EXPLAIN rejects them. + with pytest.raises(ValueError, match=message): + _explain_output(capsys, **kwargs) + + @pytest.mark.parametrize( ("window_exprs", "expected_columns"), [ diff --git a/python/tests/test_expr.py b/python/tests/test_expr.py index ef006bd91..aca01f9b3 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 @@ -25,6 +27,7 @@ import pyarrow as pa import pytest from datafusion import ( + Expr, SessionContext, col, functions, @@ -32,6 +35,7 @@ lit_with_metadata, literal_with_metadata, ) +from datafusion.common import NullTreatment from datafusion.expr import ( EXPR_TYPE_ERROR, Aggregate, @@ -53,6 +57,8 @@ TransactionEnd, TransactionStart, Values, + Window, + WindowFrame, coerce_to_expr, coerce_to_expr_list, coerce_to_expr_or_none, @@ -1251,3 +1257,276 @@ 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", + ), + ], +) +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 + + +@pytest.mark.parametrize( + ("chain", "options"), + [ + pytest.param( + lambda: functions.lead(col("v"), order_by="v").partition_by(col("g")), + "order_by", + id="builder on keyword order_by", + ), + pytest.param( + lambda: functions.lead(col("v"), partition_by=[col("g")]).over( + Window(order_by="v") + ), + "partition_by", + id="over on keyword partition_by", + ), + pytest.param( + lambda: functions.lead(col("v"), order_by="v").null_treatment( + NullTreatment.IGNORE_NULLS + ), + "order_by", + id="null_treatment on keyword order_by", + ), + pytest.param( + lambda: ( + functions.sum(col("v")) + .over(Window(window_frame=WindowFrame("rows", 1, 0))) + .partition_by(col("g")) + ), + "window_frame", + id="builder on explicit frame", + ), + pytest.param( + lambda: pickle.loads( # noqa: S301 + pickle.dumps( + functions.sum(col("v")).over( + Window(window_frame=WindowFrame("range", None, None)) + ) + ) + ).partition_by(col("g")), + "order_by, window_frame", + id="builder on decoded range frame", + ), + ], +) +def test_window_chain_rejects_existing_window_options(chain, options): + with pytest.raises(Exception, match=f"already has window options \\({options}\\)"): + chain() + + +def test_over_keeps_window_function_null_treatment(): + 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, null_treatment=NullTreatment.IGNORE_NULLS).over( + Window(partition_by=[col("g")], order_by="i") + ) + result = df.select(col("i"), expr.alias("r")).sort(col("i")) + assert result.collect_column("r").to_pylist() == [3, 3, None, None] + + +@pytest.mark.parametrize( + "window", + [ + pytest.param(Window(), id="no order_by"), + pytest.param( + Window(window_frame=WindowFrame("rows", None, None)), id="rows frame" + ), + ], +) +@pytest.mark.parametrize( + "round_trip", + [ + pytest.param(lambda e: e, id="original"), + pytest.param(copy.copy, id="copy"), + pytest.param(lambda e: pickle.loads(pickle.dumps(e)), id="pickle"), # noqa: S301 + pytest.param(lambda e: Expr.from_bytes(e.to_bytes()), id="from_bytes"), + ], +) +def test_window_builder_default_frame_same_after_round_trip(window, round_trip): + # The whole-partition frame counts as no frame, so a copy or a decoded round + # trip chains to the same result as the original. + ctx = SessionContext() + df = ctx.from_pydict({"i": [1, 1, 2, 3], "v": [1, 2, 3, 4]}) + expr = round_trip(functions.sum(col("v")).over(window)) + r = expr.over(Window(order_by="i")).alias("r") + result = df.select(col("v"), r).sort(col("v")) + assert result.collect_column("r").to_pylist() == [1, 3, 6, 10] + + +def test_window_builder_keeps_default_frame_set_last(builder_df): + frame = WindowFrame("rows", None, None) + expr = ( + functions.sum(col("v")) + .over(Window()) + .order_by(col("v")) + .window_frame(frame) + .build() + ) + result = builder_df.select(col("v"), expr.alias("r")).sort(col("v")) + assert result.collect_column("r").to_pylist() == [10, 10, 10, 10] + + +@pytest.mark.parametrize( + ("aggregate", "expected"), + [ + pytest.param( + functions.avg(col("v"), distinct=True), [2.5, 2.5, 2.5], id="distinct" + ), + pytest.param( + functions.sum(col("v"), filter=col("v") > lit(1.0)), + [4.0, 4.0, 4.0], + id="filter", + ), + pytest.param( + functions.first_value(col("n"), null_treatment=NullTreatment.IGNORE_NULLS), + [2.0, 2.0, 2.0], + id="null_treatment", + ), + pytest.param( + functions.percentile_cont(col("v"), 0.5), + [1.0, 1.0, 1.0], + id="ascending within group", + ), + ], +) +def test_over_keeps_aggregate_options(aggregate, expected): + ctx = SessionContext() + df = ctx.from_pydict({"v": [1.0, 1.0, 4.0], "n": [None, 2.0, 3.0]}) + result = df.select(aggregate.over(Window()).alias("r")) + assert result.collect_column("r").to_pylist() == expected + + +@pytest.mark.parametrize( + "aggregate", + [ + pytest.param( + functions.percentile_cont(col("v").sort(ascending=False), 0.25), + id="descending within group", + ), + pytest.param(functions.array_agg(col("v"), order_by="v"), id="order_by"), + ], +) +def test_over_rejects_aggregate_order_by(aggregate): + with pytest.raises(Exception, match="Aggregate order_by is not supported"): + aggregate.over(Window()) + + +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] + + +@pytest.mark.parametrize( + ("expr", "expected"), + [ + pytest.param( + functions.sum(col("v")).over(Window(order_by=[])), + [10, 10, 10, 10], + id="over", + ), + pytest.param(functions.row_number(order_by=[]), [1, 2, 3, 4], id="keyword"), + ], +) +def test_window_empty_order_by_executes_as_no_order_by(expr, expected): + # An empty order_by must not derive a RANGE frame with no sort key, which + # fails at execution with "ORDER BY column cannot be empty". + ctx = SessionContext() + df = ctx.from_pydict({"v": [1, 2, 3, 4]}) + result = df.select(expr.alias("r")).collect_column("r").to_pylist() + assert result == expected + + +@pytest.mark.parametrize( + ("builder", "message"), + [ + pytest.param( + lambda: functions.sum(col("v")).partition_by(col("g")), + "partition_by\\(\\) applies only to window functions; sum is an aggregate", + id="partition_by on aggregate", + ), + pytest.param( + lambda: ( + functions.sum(col("v")) + .order_by(col("v")) + .window_frame(WindowFrame("rows", 1, 0)) + ), + "window_frame\\(\\) applies only to window functions; sum is an aggregate", + id="window_frame chained on aggregate", + ), + pytest.param( + lambda: functions.lead(col("v")).distinct(), + "distinct\\(\\) applies only to aggregate functions.*lead is a window", + id="distinct on window function", + ), + pytest.param( + lambda: ( + functions.lead(col("v")) + .partition_by(col("g")) + .filter(col("v") > lit(2)) + ), + "filter\\(\\) applies only to aggregate functions.*lead is a window", + id="filter chained on window function", + ), + ], +) +def test_builder_rejects_option_for_other_function_kind(builder, message): + # The chained rows cover the kind being checked on every call, not only + # the first one made from the expression. + with pytest.raises(Exception, match=message): + builder() + + +@pytest.mark.parametrize( + ("chain", "expected"), + [ + pytest.param(lambda b: b.filter(col("v") > lit(0)), [2, 2, 2, 2], id="filter"), + pytest.param(lambda b: b.distinct(), [1, 1, 1, 1], id="distinct"), + ], +) +def test_window_builder_aggregate_options_on_aggregate_window( + builder_df, chain, expected +): + # An aggregate run as a window function still takes aggregate options. + df = builder_df.select(col("g"), (col("v") % lit(2)).alias("v")) + expr = chain(functions.sum(col("v")).over(Window())).build() + assert df.select(expr.alias("r")).collect_column("r").to_pylist() == expected diff --git a/python/tests/test_functions.py b/python/tests/test_functions.py index fabe6dff2..9cec531f9 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 @@ -745,6 +746,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 +766,99 @@ 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( + ("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 + + +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"), + [ + 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() @@ -2252,6 +2351,75 @@ 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", "args", "kwargs", "expected"), + [ + pytest.param(f.range, (5,), {}, [0, 1, 2, 3, 4], id="range stop"), + pytest.param(f.range, (1, 5), {}, [1, 2, 3, 4], id="range start stop"), + pytest.param(f.range, (1,), {"stop": 5}, [1, 2, 3, 4], id="range stop="), + pytest.param(f.range, (), {"stop": 5}, [0, 1, 2, 3, 4], id="range stop= only"), + pytest.param( + f.range, (), {"stop": 5, "step": 2}, [0, 2, 4], id="range stop= step=" + ), + pytest.param( + f.range, + (), + {"start": 0, "stop": 5, "step": 2}, + [0, 2, 4], + id="range all keywords", + ), + pytest.param(f.gen_series, (3,), {}, [0, 1, 2, 3], id="gen_series stop"), + pytest.param( + f.generate_series, + (), + {"start": 1, "stop": 3}, + [1, 2, 3], + id="generate_series keywords", + ), + ], +) +def test_series_argument_forms(func, args, kwargs, expected): + df = SessionContext().from_pydict({"a": [0]}) + result = df.select(func(*args, **kwargs).alias("v")) + assert result.collect_column("v")[0].as_py() == expected + + +@pytest.mark.parametrize( + ("args", "kwargs", "match"), + [ + pytest.param( + (), {"start": 5}, "missing a required argument: 'stop'", id="start only" + ), + pytest.param( + (1,), + {"step": 2}, + "missing a required argument: 'stop'", + id="step without stop", + ), + pytest.param((1, None), {}, "stop cannot be None", id="positional None stop"), + pytest.param( + (1,), {"stop": None}, "stop cannot be None", id="keyword None stop" + ), + ], +) +def test_series_invalid_arguments(args, kwargs, match): + # Each of these would otherwise silently promote the given value to stop. + # gen_series and generate_series share _series, so range stands for all. + with pytest.raises(TypeError, match=match): + f.range(*args, **kwargs) + + class TestPythonicNativeTypes: """Tests for accepting native Python types instead of requiring lit().""" @@ -2440,3 +2608,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"]} 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..c1189cbf1 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): @@ -473,3 +500,46 @@ 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.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_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" + + +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) diff --git a/python/tests/test_udaf.py b/python/tests/test_udaf.py index 8cd480e37..2a3466eec 100644 --- a/python/tests/test_udaf.py +++ b/python/tests/test_udaf.py @@ -22,7 +22,8 @@ import pyarrow as pa import pyarrow.compute as pc import pytest -from datafusion import Accumulator, column, udaf +from datafusion import Accumulator, SessionContext, column, udaf +from datafusion.expr import Window class Summarize(Accumulator): @@ -168,6 +169,32 @@ 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]) + + +def test_udaf_function_keyword_arguments(df): + summarize = udaf( + accum=Summarize, + input_types=pa.float64(), + return_type=pa.float64(), + state_type=[pa.float64()], + volatility="immutable", + ) + 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 @@ -251,6 +278,59 @@ def test_register_udaf(ctx, df) -> None: assert df_result.collect()[0][0][0].as_py() == 14.0 +@pytest.fixture +def distinct_ctx(): + ctx = SessionContext() + ctx.register_udaf( + udaf( + Summarize, + pa.float64(), + pa.float64(), + [pa.float64()], + volatility="immutable", + ) + ) + ctx.from_pydict({"v": [1.0, 1.0, 1.0, 5.0]}, name="t") + return ctx + + +@pytest.mark.parametrize( + "run", + [ + pytest.param( + lambda ctx, summarize: ctx.table("t").aggregate( + [], [summarize(column("v")).distinct().build().alias("r")] + ), + id="distinct aggregate", + ), + pytest.param( + lambda ctx, summarize: ctx.table("t").select( + summarize(column("v")).over(Window()).distinct().build().alias("r") + ), + id="distinct window", + ), + ], +) +def test_udaf_distinct_raises(distinct_ctx, run): + # The Python accumulator cannot deduplicate, so it would count every row. + summarize = udaf( + Summarize, + pa.float64(), + pa.float64(), + [pa.float64()], + volatility="immutable", + ) + with pytest.raises(Exception, match="DISTINCT is not supported"): + run(distinct_ctx, summarize).collect() + + +def test_udaf_distinct_rewritten_to_group_by_runs(distinct_ctx): + # The optimizer groups by the distinct values, so the accumulator never + # sees DISTINCT. + result = distinct_ctx.sql("select summarize(distinct v) as r from t") + assert result.collect_column("r").to_pylist() == [6.0] + + @pytest.mark.parametrize("wrap_in_scalar", [True, False]) def test_udaf_list_timestamp_return(ctx, wrap_in_scalar) -> None: timestamps1 = [ diff --git a/python/tests/test_udf.py b/python/tests/test_udf.py index 3a41fa6e1..8bed040fd 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 @@ -59,6 +59,26 @@ 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_udf_function_keyword_arguments(df): + is_null = udf( + func=lambda x: x.is_null(), + input_fields=[pa.int64()], + return_field=pa.bool_(), + volatility="immutable", + ) + 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(), @@ -278,3 +298,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) diff --git a/python/tests/test_udtf.py b/python/tests/test_udtf.py index aa0599ffa..ebddae93e 100644 --- a/python/tests/test_udtf.py +++ b/python/tests/test_udtf.py @@ -99,6 +99,17 @@ def static_table_func() -> Table: assert list(result[0].column(1).to_pylist()) == [0, 1, 2] +def test_python_table_function_keyword_arguments() -> None: + ctx = SessionContext() + static_func = udtf( + func=lambda: python_table_function_inner(2, 3, 1), name="static_func" + ) + ctx.register_udtf(static_func) + + result = ctx.sql("SELECT * FROM static_func()").collect() + assert result[0].column(0).to_pylist() == [0, 1, 2] + + def test_python_table_function_single_arg() -> None: """Test Python TableFunction with a single argument.""" ctx = SessionContext() diff --git a/python/tests/test_udwf.py b/python/tests/test_udwf.py index 38b935b7e..e9a35f0cf 100644 --- a/python/tests/test_udwf.py +++ b/python/tests/test_udwf.py @@ -466,3 +466,25 @@ 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] + + +def test_udwf_function_keyword_arguments(ctx): + window_count = udwf( + func=SimpleWindowCount, + input_types=[pa.int64()], + return_type=pa.int64(), + volatility="immutable", + ) + 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] 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 ``` diff --git a/uv.lock b/uv.lock index e79b871b5..499829710 100644 --- a/uv.lock +++ b/uv.lock @@ -530,7 +530,7 @@ requires-dist = [ { name = "cloudpickle", specifier = ">=2.0" }, { name = "pyarrow", marker = "python_full_version < '3.14'", specifier = ">=16.0.0" }, { name = "pyarrow", marker = "python_full_version >= '3.14'", specifier = ">=22.0.0" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'", specifier = ">=4.12" }, ] [package.metadata.requires-dev]