diff --git a/rust/cubesql/cubesql/src/compile/engine/df/wrapper.rs b/rust/cubesql/cubesql/src/compile/engine/df/wrapper.rs index 643f1f3b560cf..520be0bdf3d7f 100644 --- a/rust/cubesql/cubesql/src/compile/engine/df/wrapper.rs +++ b/rust/cubesql/cubesql/src/compile/engine/df/wrapper.rs @@ -3101,16 +3101,34 @@ impl WrappedSelectNode { sql_query, ), ScalarValue::Float32(f) => ( - f.map(|f| format!("{f}")).map_or_else( - || Self::generate_null_for_literal(sql_generator, &literal), - Ok, + f.map_or_else( + || Self::generate_null_for_literal(sql_generator.clone(), &literal), + |f| { + let data_type = + Self::generate_sql_type(sql_generator.clone(), DataType::Float32)?; + Self::generate_sql_cast_expr( + sql_generator.clone(), + format!("{f}"), + data_type, + ) + }, )?, sql_query, ), ScalarValue::Float64(f) => ( - f.map(|f| format!("{f}")).map_or_else( - || Self::generate_null_for_literal(sql_generator, &literal), - Ok, + // Display formats integral floats without a decimal point. Keep the + // scalar type so the source does not infer integer arithmetic. + f.map_or_else( + || Self::generate_null_for_literal(sql_generator.clone(), &literal), + |f| { + let data_type = + Self::generate_sql_type(sql_generator.clone(), DataType::Float64)?; + Self::generate_sql_cast_expr( + sql_generator.clone(), + format!("{f}"), + data_type, + ) + }, )?, sql_query, ), @@ -4762,6 +4780,38 @@ impl<'ctx, 'mem> ExpressionVisitor for CollectMembersVisitor<'ctx, 'mem> { #[cfg(test)] mod tests { use super::*; + #[test] + fn test_float_literal_preserves_sql_type() { + let generator = crate::compile::test::sql_generator(vec![ + ("types/float".into(), "FLOAT(24)".into()), + ("types/double".into(), "FLOAT(53)".into()), + ]); + for (literal, expected) in [ + (ScalarValue::Float32(Some(100.0)), "CAST(100 AS FLOAT(24))"), + (ScalarValue::Float64(Some(100.0)), "CAST(100 AS FLOAT(53))"), + ( + ScalarValue::Float64(Some(100.1)), + "CAST(100.1 AS FLOAT(53))", + ), + (ScalarValue::Float64(Some(0.0)), "CAST(0 AS FLOAT(53))"), + ( + ScalarValue::Float64(Some(-100.0)), + "CAST(-100 AS FLOAT(53))", + ), + (ScalarValue::Float32(None), "CAST(NULL AS FLOAT(24))"), + (ScalarValue::Float64(None), "CAST(NULL AS FLOAT(53))"), + (ScalarValue::Int64(Some(100)), "100"), + ] { + let (sql, _) = WrappedSelectNode::generate_sql_for_literal( + SqlQuery::new(String::new(), vec![]), + generator.clone(), + literal, + ) + .unwrap(); + assert_eq!(sql, expected); + } + } + use crate::{ compile::engine::df::scan::CubeScanOptions, sql::HttpAuthContext, diff --git a/rust/cubesql/cubesql/src/compile/mod.rs b/rust/cubesql/cubesql/src/compile/mod.rs index bda2ae3768937..616e78d103d72 100644 --- a/rust/cubesql/cubesql/src/compile/mod.rs +++ b/rust/cubesql/cubesql/src/compile/mod.rs @@ -7284,7 +7284,7 @@ ORDER BY "expr": { "type": "SqlFunction", "cubeParams": [], - "sql": "0", + "sql": "CAST(0 AS DOUBLE)", }, "groupingSet": null, }) @@ -7615,7 +7615,7 @@ ORDER BY "source"."str0" ASC assert_eq!( member_expression_sql(&request.dimensions), [ - "((FLOOR(((${KibanaSampleDataEcommerce.taxful_total_price} - 1.1) / 0.025)) * 0.025) + 1.1)", + "((FLOOR(((${KibanaSampleDataEcommerce.taxful_total_price} - CAST(1.1 AS DOUBLE)) / CAST(0.025 AS DOUBLE))) * CAST(0.025 AS DOUBLE)) + CAST(1.1 AS DOUBLE))", ] ); } @@ -7827,7 +7827,7 @@ ORDER BY "source"."str0" ASC assert_eq!( member_expression_sql(&request.dimensions), [ - "CEIL((CAST(EXTRACT(doy FROM CAST(${KibanaSampleDataEcommerce.order_date.week} AS TIMESTAMP)) AS INTEGER) / 7))", + "CEIL((CAST(EXTRACT(doy FROM CAST(${KibanaSampleDataEcommerce.order_date.week} AS TIMESTAMP)) AS INTEGER) / CAST(7 AS DOUBLE)))", ] ); } @@ -12064,7 +12064,7 @@ ORDER BY "source"."str0" ASC assert!(member_expression_sql(&request.measures).is_empty()); assert_eq!( member_expression_sql(&request.dimensions), - ["(EXTRACT(day FROM ${KibanaSampleDataEcommerce.order_date}) = 15)",] + ["(EXTRACT(day FROM ${KibanaSampleDataEcommerce.order_date}) = CAST(15 AS DOUBLE))",] ); } @@ -12125,7 +12125,7 @@ ORDER BY "source"."str0" ASC assert_eq!( member_expression_sql(&request.dimensions), [ - "(EXTRACT(month FROM ${KibanaSampleDataEcommerce.order_date}) < (EXTRACT(month FROM ${KibanaSampleDataEcommerce.last_mod}) + 1))", + "(EXTRACT(month FROM ${KibanaSampleDataEcommerce.order_date}) < (EXTRACT(month FROM ${KibanaSampleDataEcommerce.last_mod}) + CAST(1 AS DOUBLE)))", ] ); } @@ -12270,7 +12270,7 @@ ORDER BY "source"."str0" ASC assert!(member_expression_sql(&request.measures).is_empty()); assert_eq!( member_expression_sql(&request.dimensions), - ["(${KibanaSampleDataEcommerce.taxful_total_price} > 10)",] + ["(${KibanaSampleDataEcommerce.taxful_total_price} > CAST(10 AS DOUBLE))",] ); } @@ -15534,6 +15534,37 @@ ORDER BY "source"."str0" ASC ); } + #[tokio::test] + async fn test_float_literal_percentage_pushdown() { + for (constant, rendered) in [ + ("100.0", "CAST(100 AS FLOAT(53))"), + ("CAST(100 AS DOUBLE)", "CAST(100 AS FLOAT(53))"), + ("100.1", "CAST(100.1 AS FLOAT(53))"), + ] { + let query_plan = convert_select_to_query_plan_customized( + format!( + "SELECT customer_gender, {constant} * COUNT(*) / NULLIF(COUNT(DISTINCT notes), 0) AS ratio + FROM KibanaSampleDataEcommerce + WHERE LOWER(customer_gender) = 'test' + GROUP BY 1 ORDER BY 2 DESC LIMIT 100" + ), + DatabaseProtocol::PostgreSQL, + vec![ + ("types/double".into(), "FLOAT(53)".into()), + ("expressions/int_division".into(), "UNEXPECTED_INT_DIVISION({{ left }}, {{ right }})".into()), + ], + ).await; + let sql = query_plan + .as_logical_plan() + .find_cube_scan_wrapped_sql() + .wrapped_sql + .sql; + assert!(sql.contains(rendered), "{}: {}", constant, sql); + assert!(!sql.contains("UNEXPECTED_INT_DIVISION"), "{}", sql); + assert!(sql.contains("NULLIF("), "{}", sql); + } + } + #[tokio::test] async fn test_timestamp_literal_from_date_only_string() { if !Rewriter::sql_push_down_enabled() { diff --git a/rust/cubesql/cubesql/src/compile/test/test_wrapper.rs b/rust/cubesql/cubesql/src/compile/test/test_wrapper.rs index 21a002f84c03b..f623888c46d8d 100644 --- a/rust/cubesql/cubesql/src/compile/test/test_wrapper.rs +++ b/rust/cubesql/cubesql/src/compile/test/test_wrapper.rs @@ -4094,7 +4094,7 @@ async fn test_wrapper_multi_arg_aggregate_function() { .request .measures ), - vec!["APPROX_PERCENTILE(${KibanaSampleDataEcommerce.taxful_total_price}, 0.5)"], + vec!["APPROX_PERCENTILE(${KibanaSampleDataEcommerce.taxful_total_price}, CAST(0.5 AS DOUBLE))"], "{} is not pushed down", call );