Skip to content

Commit 3b3d1c4

Browse files
authored
add support_literal_group_by_key for unparser dialect (apache#10)
1 parent afff2bb commit 3b3d1c4

3 files changed

Lines changed: 73 additions & 3 deletions

File tree

datafusion/sql/src/unparser/dialect.rs

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -245,6 +245,12 @@ pub trait Dialect: Send + Sync {
245245
fn to_unicode_string_literal(&self, _s: &str) -> Option<ast::Expr> {
246246
None
247247
}
248+
249+
/// Whether the dialect supports using literal values as GROUP BY keys.
250+
/// Some dialects like BigQuery do not support using literal values as GROUP BY keys.
251+
fn support_literal_group_by_key(&self) -> bool {
252+
true
253+
}
248254
}
249255

250256
/// `IntervalStyle` to use for unparsing
@@ -582,6 +588,10 @@ impl Dialect for BigQueryDialect {
582588
fn unnest_as_table_factor(&self) -> bool {
583589
true
584590
}
591+
592+
fn support_literal_group_by_key(&self) -> bool {
593+
false
594+
}
585595
}
586596

587597
impl BigQueryDialect {
@@ -786,6 +796,7 @@ pub struct CustomDialect {
786796
window_func_support_window_frame: bool,
787797
full_qualified_col: bool,
788798
unnest_as_table_factor: bool,
799+
support_literal_group_by_key: bool,
789800
}
790801

791802
impl Default for CustomDialect {
@@ -814,6 +825,7 @@ impl Default for CustomDialect {
814825
window_func_support_window_frame: true,
815826
full_qualified_col: false,
816827
unnest_as_table_factor: false,
828+
support_literal_group_by_key: true,
817829
}
818830
}
819831
}
@@ -935,6 +947,10 @@ impl Dialect for CustomDialect {
935947
fn unnest_as_table_factor(&self) -> bool {
936948
self.unnest_as_table_factor
937949
}
950+
951+
fn support_literal_group_by_key(&self) -> bool {
952+
self.support_literal_group_by_key
953+
}
938954
}
939955

940956
/// `CustomDialectBuilder` to build `CustomDialect` using builder pattern
@@ -973,6 +989,7 @@ pub struct CustomDialectBuilder {
973989
full_qualified_col: bool,
974990
unnest_as_table_factor: bool,
975991
unnest_to_flattened_table_factor: bool,
992+
support_literal_group_by_key: bool,
976993
}
977994

978995
impl Default for CustomDialectBuilder {
@@ -1008,6 +1025,7 @@ impl CustomDialectBuilder {
10081025
full_qualified_col: false,
10091026
unnest_as_table_factor: false,
10101027
unnest_to_flattened_table_factor: false,
1028+
support_literal_group_by_key: true,
10111029
}
10121030
}
10131031

@@ -1034,6 +1052,7 @@ impl CustomDialectBuilder {
10341052
window_func_support_window_frame: self.window_func_support_window_frame,
10351053
full_qualified_col: self.full_qualified_col,
10361054
unnest_as_table_factor: self.unnest_as_table_factor,
1055+
support_literal_group_by_key: self.support_literal_group_by_key,
10371056
}
10381057
}
10391058

@@ -1182,4 +1201,12 @@ impl CustomDialectBuilder {
11821201
self.unnest_to_flattened_table_factor = unnest_to_flattened_table_factor;
11831202
self
11841203
}
1204+
1205+
pub fn with_support_literal_group_by_key(
1206+
mut self,
1207+
support_literal_group_by_key: bool,
1208+
) -> Self {
1209+
self.support_literal_group_by_key = support_literal_group_by_key;
1210+
self
1211+
}
11851212
}

datafusion/sql/src/unparser/plan.rs

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -252,6 +252,10 @@ impl Unparser<'_> {
252252
select.group_by(ast::GroupByExpr::Expressions(
253253
agg.group_expr
254254
.iter()
255+
.filter(|expr| {
256+
self.dialect.support_literal_group_by_key()
257+
|| !matches!(expr, Expr::Literal(_, _))
258+
})
255259
.map(|expr| self.expr_to_sql(expr))
256260
.collect::<Result<Vec<_>>>()?,
257261
vec![],
@@ -556,6 +560,10 @@ impl Unparser<'_> {
556560
select.group_by(ast::GroupByExpr::Expressions(
557561
agg.group_expr
558562
.iter()
563+
.filter(|expr| {
564+
self.dialect.support_literal_group_by_key()
565+
|| !matches!(expr, Expr::Literal(_, _))
566+
})
559567
.map(|expr| self.expr_to_sql(expr))
560568
.collect::<Result<Vec<_>>>()?,
561569
vec![],

datafusion/sql/tests/cases/plan_to_sql.rs

Lines changed: 38 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -34,9 +34,10 @@ use datafusion_functions_nested::map::map_udf;
3434
use datafusion_functions_window::rank::rank_udwf;
3535
use datafusion_sql::planner::{ContextProvider, PlannerContext, SqlToRel};
3636
use datafusion_sql::unparser::dialect::{
37-
BigQueryDialect, CustomDialectBuilder, DefaultDialect as UnparserDefaultDialect,
38-
DefaultDialect, Dialect as UnparserDialect, MySqlDialect as UnparserMySqlDialect,
39-
PostgreSqlDialect as UnparserPostgreSqlDialect, SnowflakeDialect, SqliteDialect,
37+
BigQueryDialect, CustomDialect, CustomDialectBuilder,
38+
DefaultDialect as UnparserDefaultDialect, DefaultDialect, Dialect as UnparserDialect,
39+
MySqlDialect as UnparserMySqlDialect, PostgreSqlDialect as UnparserPostgreSqlDialect,
40+
SnowflakeDialect, SqliteDialect,
4041
};
4142
use datafusion_sql::unparser::{expr_to_sql, plan_to_sql, Unparser};
4243
use insta::assert_snapshot;
@@ -2708,3 +2709,37 @@ fn test_struct_expr3() {
27082709
@r#"SELECT test.c1."metadata".product."name" FROM (SELECT {"metadata": {product: {"name": 'Product Name'}}} AS c1) AS test"#
27092710
);
27102711
}
2712+
2713+
#[test]
2714+
fn test_literal_gbk() -> Result<(), DataFusionError> {
2715+
let sql = r#"
2716+
select 'ABC' as col1, first_name, min(id) from person group by 1, 2
2717+
2718+
"#;
2719+
let unparser = CustomDialectBuilder::new()
2720+
.with_support_literal_group_by_key(false)
2721+
.build();
2722+
roundtrip_statement_with_dialect_helper!(
2723+
sql: sql,
2724+
parser_dialect: GenericDialect {},
2725+
unparser_dialect: unparser,
2726+
expected: @r#"SELECT 'ABC' AS col1, person.first_name, min(person.id) FROM person GROUP BY person.first_name"#,
2727+
);
2728+
2729+
let unparser = BigQueryDialect {};
2730+
roundtrip_statement_with_dialect_helper!(
2731+
sql: sql,
2732+
parser_dialect: GenericDialect {},
2733+
unparser_dialect: unparser,
2734+
expected: @r#"SELECT 'ABC' AS `col1`, `person`.`first_name`, min(`person`.`id`) FROM `person` GROUP BY `person`.`first_name`"#,
2735+
);
2736+
2737+
let unparser = CustomDialect::default();
2738+
roundtrip_statement_with_dialect_helper!(
2739+
sql: sql,
2740+
parser_dialect: GenericDialect {},
2741+
unparser_dialect: unparser,
2742+
expected: @r#"SELECT 'ABC' AS col1, person.first_name, min(person.id) FROM person GROUP BY 'ABC', person.first_name"#,
2743+
);
2744+
Ok(())
2745+
}

0 commit comments

Comments
 (0)