From 34cfbd983ee8ee7dac2ac48731a06ccd1b291628 Mon Sep 17 00:00:00 2001 From: Andrey Koshchiy Date: Tue, 4 Feb 2025 19:50:34 +0300 Subject: [PATCH 1/4] fix: rewrite order by on compound expr --- datafusion/expr/src/expr_rewriter/order_by.rs | 18 ++++++++++++++---- datafusion/sql/tests/sql_integration.rs | 16 ++++++++++++++++ datafusion/sqllogictest/test_files/order.slt | 13 +++++++++++++ 3 files changed, 43 insertions(+), 4 deletions(-) diff --git a/datafusion/expr/src/expr_rewriter/order_by.rs b/datafusion/expr/src/expr_rewriter/order_by.rs index 0044b6cf6f377..86a82f4478359 100644 --- a/datafusion/expr/src/expr_rewriter/order_by.rs +++ b/datafusion/expr/src/expr_rewriter/order_by.rs @@ -21,7 +21,7 @@ use crate::expr::Alias; use crate::expr_rewriter::normalize_col; use crate::{expr::Sort, Cast, Expr, LogicalPlan, TryCast}; -use datafusion_common::tree_node::{Transformed, TransformedResult, TreeNode}; +use datafusion_common::tree_node::{Transformed, TransformedResult, TreeNode, TreeNodeRecursion}; use datafusion_common::{Column, Result}; /// Rewrite sort on aggregate expressions to sort on the column of aggregate output @@ -74,7 +74,7 @@ fn rewrite_in_terms_of_projection( ) -> Result { // assumption is that each item in exprs, such as "b + c" is // available as an output column named "b + c" - expr.transform(|expr| { + expr.transform_down(|expr| { // search for unnormalized names first such as "c1" (such as aliases) if let Some(found) = proj_exprs.iter().find(|a| (**a) == expr) { let (qualifier, field_name) = found.qualified_name(); @@ -101,8 +101,18 @@ fn rewrite_in_terms_of_projection( let search_col = Expr::Column(Column::new_unqualified(name)); // look for the column named the same as this expr - if let Some(found) = proj_exprs.iter().find(|a| expr_match(&search_col, a)) { - let found = found.clone(); + let mut found = None; + for proj_expr in &proj_exprs { + proj_expr.apply(|e| { + if expr_match(e, &search_col) { + found = Some(e.clone()); + return Ok(TreeNodeRecursion::Stop); + } + Ok(TreeNodeRecursion::Continue) + })?; + } + + if let Some(found) = found { return Ok(Transformed::yes(match normalized_expr { Expr::Cast(Cast { expr: _, data_type }) => Expr::Cast(Cast { expr: Box::new(found), diff --git a/datafusion/sql/tests/sql_integration.rs b/datafusion/sql/tests/sql_integration.rs index b9502c1520049..88c6d3444f823 100644 --- a/datafusion/sql/tests/sql_integration.rs +++ b/datafusion/sql/tests/sql_integration.rs @@ -2511,6 +2511,22 @@ fn select_groupby_orderby() { FROM person GROUP BY person.birth_date ORDER BY birth_date; "#; quick_test(sql, expected); + + // Use columnized `avg(age)` in the order by + let sql = r#"SELECT + avg(age) + avg(age) + 1, + date_trunc('month', person.birth_date) AS "birth_date" + FROM person GROUP BY person.birth_date ORDER BY avg(age); +"#; + + let expected = + "Projection: avg(person.age) + avg(person.age) + Int64(1), birth_date\ + \n Sort: avg(person.age) ASC NULLS LAST\ + \n Projection: avg(person.age) + avg(person.age) + Int64(1), date_trunc(Utf8(\"month\"), person.birth_date) AS birth_date, avg(person.age)\ + \n Aggregate: groupBy=[[person.birth_date]], aggr=[[avg(person.age)]]\ + \n TableScan: person"; + + quick_test(sql, expected); } fn logical_plan(sql: &str) -> Result { diff --git a/datafusion/sqllogictest/test_files/order.slt b/datafusion/sqllogictest/test_files/order.slt index 0f84171697257..383be5ee9152a 100644 --- a/datafusion/sqllogictest/test_files/order.slt +++ b/datafusion/sqllogictest/test_files/order.slt @@ -384,6 +384,19 @@ ORDER BY time; 2 2022-01-01T01:00:00 3 2022-01-02T00:00:00 +# Tests for https://github.com/apache/datafusion/issues/14459 +query PI +select + date_trunc('minute',time) AS "time", + sum(value) + sum(value) +FROM t +GROUP BY time +ORDER BY sum(value) + sum(value); +---- +2022-01-01T00:00:00 2 +2022-01-01T01:00:00 4 +2022-01-02T00:00:00 6 + ## SORT BY is not supported statement error DataFusion error: This feature is not implemented: SORT BY select * from t SORT BY time; From 314ce84801c7923f38e8cad4d65eef013462d3d6 Mon Sep 17 00:00:00 2001 From: Andrey Koshchiy Date: Tue, 4 Feb 2025 19:52:19 +0300 Subject: [PATCH 2/4] fmt --- datafusion/expr/src/expr_rewriter/order_by.rs | 4 +++- datafusion/sql/tests/sql_integration.rs | 12 ++++++------ 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/datafusion/expr/src/expr_rewriter/order_by.rs b/datafusion/expr/src/expr_rewriter/order_by.rs index 86a82f4478359..72f679d296c27 100644 --- a/datafusion/expr/src/expr_rewriter/order_by.rs +++ b/datafusion/expr/src/expr_rewriter/order_by.rs @@ -21,7 +21,9 @@ use crate::expr::Alias; use crate::expr_rewriter::normalize_col; use crate::{expr::Sort, Cast, Expr, LogicalPlan, TryCast}; -use datafusion_common::tree_node::{Transformed, TransformedResult, TreeNode, TreeNodeRecursion}; +use datafusion_common::tree_node::{ + Transformed, TransformedResult, TreeNode, TreeNodeRecursion, +}; use datafusion_common::{Column, Result}; /// Rewrite sort on aggregate expressions to sort on the column of aggregate output diff --git a/datafusion/sql/tests/sql_integration.rs b/datafusion/sql/tests/sql_integration.rs index 88c6d3444f823..83c78d2159837 100644 --- a/datafusion/sql/tests/sql_integration.rs +++ b/datafusion/sql/tests/sql_integration.rs @@ -2519,12 +2519,12 @@ fn select_groupby_orderby() { FROM person GROUP BY person.birth_date ORDER BY avg(age); "#; - let expected = - "Projection: avg(person.age) + avg(person.age) + Int64(1), birth_date\ - \n Sort: avg(person.age) ASC NULLS LAST\ - \n Projection: avg(person.age) + avg(person.age) + Int64(1), date_trunc(Utf8(\"month\"), person.birth_date) AS birth_date, avg(person.age)\ - \n Aggregate: groupBy=[[person.birth_date]], aggr=[[avg(person.age)]]\ - \n TableScan: person"; + let expected = + "Projection: avg(person.age) + avg(person.age) + Int64(1), birth_date\ + \n Sort: avg(person.age) ASC NULLS LAST\ + \n Projection: avg(person.age) + avg(person.age) + Int64(1), date_trunc(Utf8(\"month\"), person.birth_date) AS birth_date, avg(person.age)\ + \n Aggregate: groupBy=[[person.birth_date]], aggr=[[avg(person.age)]]\ + \n TableScan: person"; quick_test(sql, expected); } From 941c1f19258d377e09ee864984acf690227376b7 Mon Sep 17 00:00:00 2001 From: Andrey Koshchiy Date: Tue, 4 Feb 2025 20:42:29 +0300 Subject: [PATCH 3/4] revert transform --- datafusion/expr/src/expr_rewriter/order_by.rs | 2 +- datafusion/sql/tests/sql_integration.rs | 13 ++++++------- 2 files changed, 7 insertions(+), 8 deletions(-) diff --git a/datafusion/expr/src/expr_rewriter/order_by.rs b/datafusion/expr/src/expr_rewriter/order_by.rs index 72f679d296c27..426c3205cfe3b 100644 --- a/datafusion/expr/src/expr_rewriter/order_by.rs +++ b/datafusion/expr/src/expr_rewriter/order_by.rs @@ -76,7 +76,7 @@ fn rewrite_in_terms_of_projection( ) -> Result { // assumption is that each item in exprs, such as "b + c" is // available as an output column named "b + c" - expr.transform_down(|expr| { + expr.transform(|expr| { // search for unnormalized names first such as "c1" (such as aliases) if let Some(found) = proj_exprs.iter().find(|a| (**a) == expr) { let (qualifier, field_name) = found.qualified_name(); diff --git a/datafusion/sql/tests/sql_integration.rs b/datafusion/sql/tests/sql_integration.rs index 83c78d2159837..74a5abf95c84a 100644 --- a/datafusion/sql/tests/sql_integration.rs +++ b/datafusion/sql/tests/sql_integration.rs @@ -2514,17 +2514,16 @@ fn select_groupby_orderby() { // Use columnized `avg(age)` in the order by let sql = r#"SELECT - avg(age) + avg(age) + 1, + avg(age) + avg(age), date_trunc('month', person.birth_date) AS "birth_date" - FROM person GROUP BY person.birth_date ORDER BY avg(age); + FROM person GROUP BY person.birth_date ORDER BY avg(age) + avg(age); "#; let expected = - "Projection: avg(person.age) + avg(person.age) + Int64(1), birth_date\ - \n Sort: avg(person.age) ASC NULLS LAST\ - \n Projection: avg(person.age) + avg(person.age) + Int64(1), date_trunc(Utf8(\"month\"), person.birth_date) AS birth_date, avg(person.age)\ - \n Aggregate: groupBy=[[person.birth_date]], aggr=[[avg(person.age)]]\ - \n TableScan: person"; + "Sort: avg(person.age) + avg(person.age) ASC NULLS LAST\ + \n Projection: avg(person.age) + avg(person.age), date_trunc(Utf8(\"month\"), person.birth_date) AS birth_date\ + \n Aggregate: groupBy=[[person.birth_date]], aggr=[[avg(person.age)]]\ + \n TableScan: person"; quick_test(sql, expected); } From c0ec7a843ca721b6581a1c4d562c406b84e2492c Mon Sep 17 00:00:00 2001 From: Andrey Koshchiy Date: Wed, 5 Feb 2025 21:24:08 +0300 Subject: [PATCH 4/4] test fix --- datafusion/expr/src/expr_rewriter/order_by.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/datafusion/expr/src/expr_rewriter/order_by.rs b/datafusion/expr/src/expr_rewriter/order_by.rs index 426c3205cfe3b..6db95555502da 100644 --- a/datafusion/expr/src/expr_rewriter/order_by.rs +++ b/datafusion/expr/src/expr_rewriter/order_by.rs @@ -106,7 +106,7 @@ fn rewrite_in_terms_of_projection( let mut found = None; for proj_expr in &proj_exprs { proj_expr.apply(|e| { - if expr_match(e, &search_col) { + if expr_match(&search_col, e) { found = Some(e.clone()); return Ok(TreeNodeRecursion::Stop); }