Skip to content

Commit d2ac890

Browse files
timsaucerclaude
andcommitted
fix: keep existing options when chaining aggregate and window builders
Upstream ExprFunctionExt methods on an Expr start from an empty builder, so build() resets every option not set again. The Python function wrappers already apply their keyword options, so calls such as string_agg(..., order_by=...).distinct().build() silently dropped the ordering, and .filter() dropped order_by, and so on. Seed the builder from the expression's existing params instead. A window frame equal to the default for its order_by is left unset so build() derives it again from the final order_by. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent b035c97 commit d2ac890

2 files changed

Lines changed: 138 additions & 8 deletions

File tree

‎crates/core/src/expr.rs‎

Lines changed: 63 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -625,34 +625,38 @@ impl PyExpr {
625625
// Expression Function Builder functions
626626

627627
pub fn order_by(&self, order_by: Vec<PySortExpr>) -> PyExprFuncBuilder {
628-
self.expr
629-
.clone()
628+
builder_from_expr(&self.expr)
630629
.order_by(to_sort_expressions(order_by))
631630
.into()
632631
}
633632

634633
pub fn filter(&self, filter: PyExpr) -> PyExprFuncBuilder {
635-
self.expr.clone().filter(filter.expr.clone()).into()
634+
builder_from_expr(&self.expr)
635+
.filter(filter.expr.clone())
636+
.into()
636637
}
637638

638639
pub fn distinct(&self) -> PyExprFuncBuilder {
639-
self.expr.clone().distinct().into()
640+
builder_from_expr(&self.expr).distinct().into()
640641
}
641642

642643
pub fn null_treatment(&self, null_treatment: NullTreatment) -> PyExprFuncBuilder {
643-
self.expr
644-
.clone()
644+
builder_from_expr(&self.expr)
645645
.null_treatment(Some(null_treatment.into()))
646646
.into()
647647
}
648648

649649
pub fn partition_by(&self, partition_by: Vec<PyExpr>) -> PyExprFuncBuilder {
650650
let partition_by = partition_by.iter().map(|e| e.expr.clone()).collect();
651-
self.expr.clone().partition_by(partition_by).into()
651+
builder_from_expr(&self.expr)
652+
.partition_by(partition_by)
653+
.into()
652654
}
653655

654656
pub fn window_frame(&self, window_frame: PyWindowFrame) -> PyExprFuncBuilder {
655-
self.expr.clone().window_frame(window_frame.into()).into()
657+
builder_from_expr(&self.expr)
658+
.window_frame(window_frame.into())
659+
.into()
656660
}
657661

658662
#[pyo3(signature = (partition_by=None, window_frame=None, order_by=None, null_treatment=None))]
@@ -743,6 +747,57 @@ impl PyExpr {
743747
}
744748
}
745749

750+
/// Start an [`ExprFuncBuilder`] that keeps the options already set on `expr`.
751+
///
752+
/// Upstream's `ExprFunctionExt` methods on an `Expr` start from an empty
753+
/// builder, so `build()` would reset every option not set again. The Python
754+
/// function wrappers already apply their keyword options, so chaining another
755+
/// builder method onto their result must not discard them.
756+
fn builder_from_expr(expr: &Expr) -> ExprFuncBuilder {
757+
match expr {
758+
Expr::AggregateFunction(agg) => {
759+
let params = &agg.params;
760+
let mut builder = expr.clone().null_treatment(params.null_treatment);
761+
if !params.order_by.is_empty() {
762+
builder = builder.order_by(params.order_by.clone());
763+
}
764+
if let Some(filter) = &params.filter {
765+
builder = builder.filter(filter.as_ref().clone());
766+
}
767+
if params.distinct {
768+
builder = builder.distinct();
769+
}
770+
builder
771+
}
772+
Expr::WindowFunction(window) => {
773+
let params = &window.params;
774+
let mut builder = expr.clone().null_treatment(params.null_treatment);
775+
if !params.partition_by.is_empty() {
776+
builder = builder.partition_by(params.partition_by.clone());
777+
}
778+
let has_order_by = !params.order_by.is_empty();
779+
if has_order_by {
780+
builder = builder.order_by(params.order_by.clone());
781+
}
782+
// A frame equal to the default `build()` derived from the order-by is
783+
// left unset, so it is derived again from the final order-by.
784+
if params.window_frame
785+
!= datafusion::logical_expr::WindowFrame::new(has_order_by.then_some(true))
786+
{
787+
builder = builder.window_frame(params.window_frame.clone());
788+
}
789+
if let Some(filter) = &params.filter {
790+
builder = builder.filter(filter.as_ref().clone());
791+
}
792+
if params.distinct {
793+
builder = builder.distinct();
794+
}
795+
builder
796+
}
797+
_ => expr.clone().null_treatment(None),
798+
}
799+
}
800+
746801
#[pyclass(
747802
from_py_object,
748803
frozen,

‎python/tests/test_expr.py‎

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,8 @@
5353
TransactionEnd,
5454
TransactionStart,
5555
Values,
56+
Window,
57+
WindowFrame,
5658
coerce_to_expr,
5759
coerce_to_expr_list,
5860
coerce_to_expr_or_none,
@@ -1251,3 +1253,76 @@ def test_expr_to_bytes_no_ctx_default_codec() -> None:
12511253
restored = Expr.from_bytes(blob, ctx=fresh)
12521254

12531255
assert restored.canonical_name() == original.canonical_name()
1256+
1257+
1258+
@pytest.fixture
1259+
def builder_df():
1260+
ctx = SessionContext()
1261+
return ctx.from_pydict(
1262+
{"g": [1, 1, 1, 2], "s": ["y", "x", "z", "w"], "v": [3, 1, 2, 4]}
1263+
)
1264+
1265+
1266+
@pytest.mark.parametrize(
1267+
("build_expr", "expected"),
1268+
[
1269+
pytest.param(
1270+
lambda: (
1271+
functions.array_agg(col("s"), order_by="s")
1272+
.filter(col("v") > lit(1))
1273+
.build()
1274+
),
1275+
["w", "y", "z"],
1276+
id="order_by_kept_after_filter",
1277+
),
1278+
pytest.param(
1279+
lambda: (
1280+
functions.array_agg(col("s"), filter=col("v") > lit(1))
1281+
.order_by(col("s").sort(ascending=False))
1282+
.build()
1283+
),
1284+
["z", "y", "w"],
1285+
id="filter_kept_after_order_by",
1286+
),
1287+
pytest.param(
1288+
lambda: (
1289+
functions.string_agg(col("s"), ",", order_by="s").distinct().build()
1290+
),
1291+
"w,x,y,z",
1292+
id="order_by_kept_after_distinct",
1293+
),
1294+
pytest.param(
1295+
lambda: (
1296+
functions.first_value(col("s"), order_by="v")
1297+
.filter(col("v") > lit(1))
1298+
.build()
1299+
),
1300+
"z",
1301+
id="first_value_order_by_kept_after_filter",
1302+
),
1303+
],
1304+
)
1305+
def test_aggregate_builder_keeps_existing_options(builder_df, build_expr, expected):
1306+
result = builder_df.aggregate([], [build_expr().alias("r")])
1307+
assert result.collect_column("r")[0].as_py() == expected
1308+
1309+
1310+
def test_window_builder_keeps_existing_options(builder_df):
1311+
expr = functions.lead(col("v"), order_by="v").partition_by(col("g")).build()
1312+
result = builder_df.select(col("v"), expr.alias("r")).sort(col("v"))
1313+
assert result.collect_column("r").to_pylist() == [2, 3, None, None]
1314+
1315+
1316+
def test_window_builder_keeps_explicit_frame(builder_df):
1317+
window = Window(order_by=col("v"), window_frame=WindowFrame("rows", 1, 0))
1318+
expr = functions.sum(col("v")).over(window).partition_by(col("g")).build()
1319+
result = builder_df.select(col("v"), expr.alias("r")).sort(col("v"))
1320+
assert result.collect_column("r").to_pylist() == [1, 3, 5, 4]
1321+
1322+
1323+
def test_window_builder_rederives_default_frame(builder_df):
1324+
# No order_by means a whole-partition frame; adding one later must switch
1325+
# to the running frame rather than keep the whole-partition default.
1326+
expr = functions.sum(col("v")).over(Window()).order_by(col("v")).build()
1327+
result = builder_df.select(col("v"), expr.alias("r")).sort(col("v"))
1328+
assert result.collect_column("r").to_pylist() == [1, 3, 6, 10]

0 commit comments

Comments
 (0)