From 01fb4fd244bdecb8a5d63aad39a68f5ee2aa7bfa Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 13 Jan 2023 09:46:35 -0700 Subject: [PATCH 1/4] Change API to return LogicalPlan instead of DataFrame --- datafusion/substrait/src/consumer.rs | 64 +++++++++++++++---------- datafusion/substrait/tests/roundtrip.rs | 21 +++----- datafusion/substrait/tests/serialize.rs | 3 +- 3 files changed, 45 insertions(+), 43 deletions(-) diff --git a/datafusion/substrait/src/consumer.rs b/datafusion/substrait/src/consumer.rs index 2f0e889696567..542f3eac9d74e 100644 --- a/datafusion/substrait/src/consumer.rs +++ b/datafusion/substrait/src/consumer.rs @@ -17,17 +17,17 @@ use async_recursion::async_recursion; use datafusion::common::{DFField, DFSchema, DFSchemaRef}; -use datafusion::logical_expr::build_join_schema; use datafusion::logical_expr::expr; use datafusion::logical_expr::{ aggregate_function, BinaryExpr, Case, Expr, LogicalPlan, Operator, }; +use datafusion::logical_expr::{build_join_schema, LogicalPlanBuilder}; use datafusion::prelude::JoinType; use datafusion::sql::TableReference; use datafusion::{ error::{DataFusionError, Result}, optimizer::utils::split_conjunction, - prelude::{Column, DataFrame, SessionContext}, + prelude::{Column, SessionContext}, scalar::ScalarValue, }; use substrait::protobuf::{ @@ -87,7 +87,7 @@ pub fn name_to_op(name: &str) -> Result { pub async fn from_substrait_plan( ctx: &mut SessionContext, plan: &Plan, -) -> Result { +) -> Result { // Register function extension let function_extension = plan .extensions @@ -135,17 +135,19 @@ pub async fn from_substrait_rel( ctx: &mut SessionContext, rel: &Rel, extensions: &HashMap, -) -> Result { +) -> Result { match &rel.rel_type { Some(RelType::Project(p)) => { if let Some(input) = p.input.as_ref() { - let input = from_substrait_rel(ctx, input, extensions).await?; + let input = LogicalPlanBuilder::from( + from_substrait_rel(ctx, input, extensions).await?, + ); let mut exprs: Vec = vec![]; for e in &p.expressions { let x = from_substrait_rex(e, input.schema(), extensions).await?; exprs.push(x.as_ref().clone()); } - input.select(exprs) + input.project(exprs)?.build() } else { Err(DataFusionError::NotImplemented( "Projection without an input is not supported".to_string(), @@ -154,11 +156,13 @@ pub async fn from_substrait_rel( } Some(RelType::Filter(filter)) => { if let Some(input) = filter.input.as_ref() { - let input = from_substrait_rel(ctx, input, extensions).await?; + let input = LogicalPlanBuilder::from( + from_substrait_rel(ctx, input, extensions).await?, + ); if let Some(condition) = filter.condition.as_ref() { let expr = from_substrait_rex(condition, input.schema(), extensions).await?; - input.filter(expr.as_ref().clone()) + input.filter(expr.as_ref().clone())?.build() } else { Err(DataFusionError::NotImplemented( "Filter without an condition is not valid".to_string(), @@ -172,10 +176,12 @@ pub async fn from_substrait_rel( } Some(RelType::Fetch(fetch)) => { if let Some(input) = fetch.input.as_ref() { - let input = from_substrait_rel(ctx, input, extensions).await?; + let input = LogicalPlanBuilder::from( + from_substrait_rel(ctx, input, extensions).await?, + ); let offset = fetch.offset as usize; let count = fetch.count as usize; - input.limit(offset, Some(count)) + input.limit(offset, Some(count))?.build() } else { Err(DataFusionError::NotImplemented( "Fetch without an input is not valid".to_string(), @@ -184,7 +190,9 @@ pub async fn from_substrait_rel( } Some(RelType::Sort(sort)) => { if let Some(input) = sort.input.as_ref() { - let input = from_substrait_rel(ctx, input, extensions).await?; + let input = LogicalPlanBuilder::from( + from_substrait_rel(ctx, input, extensions).await?, + ); let mut sorts: Vec = vec![]; for s in &sort.sorts { let expr = from_substrait_rex( @@ -224,7 +232,7 @@ pub async fn from_substrait_rel( nulls_first, })); } - input.sort(sorts) + input.sort(sorts)?.build() } else { Err(DataFusionError::NotImplemented( "Sort without an input is not valid".to_string(), @@ -233,7 +241,9 @@ pub async fn from_substrait_rel( } Some(RelType::Aggregate(agg)) => { if let Some(input) = agg.input.as_ref() { - let input = from_substrait_rel(ctx, input, extensions).await?; + let input = LogicalPlanBuilder::from( + from_substrait_rel(ctx, input, extensions).await?, + ); let mut group_expr = vec![]; let mut aggr_expr = vec![]; @@ -292,7 +302,7 @@ pub async fn from_substrait_rel( aggr_expr.push(agg_func?.as_ref().clone()); } - input.aggregate(group_expr, aggr_expr) + input.aggregate(group_expr, aggr_expr)?.build() } else { Err(DataFusionError::NotImplemented( "Aggregate without an input is not valid".to_string(), @@ -300,10 +310,12 @@ pub async fn from_substrait_rel( } } Some(RelType::Join(join)) => { - let left = - from_substrait_rel(ctx, join.left.as_ref().unwrap(), extensions).await?; - let right = - from_substrait_rel(ctx, join.right.as_ref().unwrap(), extensions).await?; + let left = LogicalPlanBuilder::from( + from_substrait_rel(ctx, join.left.as_ref().unwrap(), extensions).await?, + ); + let right = LogicalPlanBuilder::from( + from_substrait_rel(ctx, join.right.as_ref().unwrap(), extensions).await?, + ); let join_type = match join.r#type { 1 => JoinType::Inner, 2 => JoinType::Left, @@ -346,9 +358,9 @@ pub async fn from_substrait_rel( )), }) .collect::>>()?; - let left_cols: Vec<&str> = pairs.iter().map(|(l, _)| l.as_str()).collect(); - let right_cols: Vec<&str> = pairs.iter().map(|(_, r)| r.as_str()).collect(); - left.join(right, join_type, &left_cols, &right_cols, None) + let (left_cols, right_cols): (Vec<_>, Vec<_>) = pairs.iter().cloned().unzip(); + left.join(right.build()?, join_type, (left_cols, right_cols), None)? + .build() } Some(RelType::Read(read)) => match &read.as_ref().read_type { Some(ReadType::NamedTable(nt)) => { @@ -372,6 +384,7 @@ pub async fn from_substrait_rel( }, }; let t = ctx.table(table_reference).await?; + let t = t.logical_plan().clone(); match &read.projection { Some(MaskExpression { select, .. }) => match &select.as_ref() { Some(projection) => { @@ -380,7 +393,7 @@ pub async fn from_substrait_rel( .iter() .map(|item| item.field as usize) .collect(); - match t.into_optimized_plan()? { + match &t { LogicalPlan::TableScan(scan) => { let fields: Vec = column_indices .iter() @@ -395,17 +408,16 @@ pub async fn from_substrait_rel( fields, HashMap::new(), )?); - let plan = LogicalPlan::TableScan(scan); - Ok(DataFrame::new(ctx.state(), plan)) + Ok(LogicalPlan::TableScan(scan)) } _ => Err(DataFusionError::Internal( "unexpected plan for table".to_string(), )), } } - _ => Ok(t), + _ => Ok(t.clone()), }, - _ => Ok(t), + _ => Ok(t.clone()), } } _ => Err(DataFusionError::NotImplemented( diff --git a/datafusion/substrait/tests/roundtrip.rs b/datafusion/substrait/tests/roundtrip.rs index 5fde79d4cc769..6363ad9ae648e 100644 --- a/datafusion/substrait/tests/roundtrip.rs +++ b/datafusion/substrait/tests/roundtrip.rs @@ -186,8 +186,7 @@ mod tests { let df = ctx.sql(sql).await?; let plan = df.into_optimized_plan()?; let proto = to_substrait_plan(&plan)?; - let df = from_substrait_plan(&mut ctx, &proto).await?; - let plan2 = df.into_optimized_plan()?; + let plan2 = from_substrait_plan(&mut ctx, &proto).await?; let plan2str = format!("{:?}", plan2); assert_eq!(expected_plan_str, &plan2str); Ok(()) @@ -198,9 +197,7 @@ mod tests { let df = ctx.sql(sql).await?; let plan1 = df.into_optimized_plan()?; let proto = to_substrait_plan(&plan1)?; - - let df = from_substrait_plan(&mut ctx, &proto).await?; - let plan2 = df.into_optimized_plan()?; + let plan2 = from_substrait_plan(&mut ctx, &proto).await?; // Format plan string and replace all None's with 0 let plan1str = format!("{:?}", plan1).replace("None", "0"); @@ -218,15 +215,11 @@ mod tests { let df_a = ctx.sql(sql_with_alias).await?; let proto_a = to_substrait_plan(&df_a.into_optimized_plan()?)?; - let plan_with_alias = from_substrait_plan(&mut ctx, &proto_a) - .await? - .into_optimized_plan()?; + let plan_with_alias = from_substrait_plan(&mut ctx, &proto_a).await?; let df = ctx.sql(sql_no_alias).await?; - let proto = to_substrait_plan(&df.into_optimized_plan()?)?; - let plan = from_substrait_plan(&mut ctx, &proto) - .await? - .into_optimized_plan()?; + let proto = to_substrait_plan(&df.logical_plan())?; + let plan = from_substrait_plan(&mut ctx, &proto).await?; println!("{:#?}", plan_with_alias); println!("{:#?}", plan); @@ -242,9 +235,7 @@ mod tests { let df = ctx.sql(sql).await?; let plan = df.into_optimized_plan()?; let proto = to_substrait_plan(&plan)?; - - let df = from_substrait_plan(&mut ctx, &proto).await?; - let plan2 = df.into_optimized_plan()?; + let plan2 = from_substrait_plan(&mut ctx, &proto).await?; println!("{:#?}", plan); println!("{:#?}", plan2); diff --git a/datafusion/substrait/tests/serialize.rs b/datafusion/substrait/tests/serialize.rs index 59b2899ede399..a72fca010fa26 100644 --- a/datafusion/substrait/tests/serialize.rs +++ b/datafusion/substrait/tests/serialize.rs @@ -40,8 +40,7 @@ mod tests { // Read substrait plan from file let proto = serializer::deserialize(path).await?; // Check plan equality - let df = from_substrait_plan(&mut ctx, &proto).await?; - let plan = df.into_optimized_plan()?; + let plan = from_substrait_plan(&mut ctx, &proto).await?; let plan_str_ref = format!("{:?}", plan_ref); let plan_str = format!("{:?}", plan); assert_eq!(plan_str_ref, plan_str); From 68d2f7c8da3676aa20b2545a31e6443caaab1ce2 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 13 Jan 2023 10:02:46 -0700 Subject: [PATCH 2/4] optimize plans --- datafusion/substrait/src/consumer.rs | 6 +++--- datafusion/substrait/tests/roundtrip.rs | 2 +- datafusion/substrait/tests/serialize.rs | 2 ++ 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/datafusion/substrait/src/consumer.rs b/datafusion/substrait/src/consumer.rs index 542f3eac9d74e..b2469a14721e0 100644 --- a/datafusion/substrait/src/consumer.rs +++ b/datafusion/substrait/src/consumer.rs @@ -384,7 +384,7 @@ pub async fn from_substrait_rel( }, }; let t = ctx.table(table_reference).await?; - let t = t.logical_plan().clone(); + let t = t.into_optimized_plan()?; match &read.projection { Some(MaskExpression { select, .. }) => match &select.as_ref() { Some(projection) => { @@ -415,9 +415,9 @@ pub async fn from_substrait_rel( )), } } - _ => Ok(t.clone()), + _ => Ok(t), }, - _ => Ok(t.clone()), + _ => Ok(t), } } _ => Err(DataFusionError::NotImplemented( diff --git a/datafusion/substrait/tests/roundtrip.rs b/datafusion/substrait/tests/roundtrip.rs index 6363ad9ae648e..c31062edeee28 100644 --- a/datafusion/substrait/tests/roundtrip.rs +++ b/datafusion/substrait/tests/roundtrip.rs @@ -218,7 +218,7 @@ mod tests { let plan_with_alias = from_substrait_plan(&mut ctx, &proto_a).await?; let df = ctx.sql(sql_no_alias).await?; - let proto = to_substrait_plan(&df.logical_plan())?; + let proto = to_substrait_plan(&df.into_optimized_plan()?)?; let plan = from_substrait_plan(&mut ctx, &proto).await?; println!("{:#?}", plan_with_alias); diff --git a/datafusion/substrait/tests/serialize.rs b/datafusion/substrait/tests/serialize.rs index a72fca010fa26..9a836fa29151d 100644 --- a/datafusion/substrait/tests/serialize.rs +++ b/datafusion/substrait/tests/serialize.rs @@ -41,6 +41,8 @@ mod tests { let proto = serializer::deserialize(path).await?; // Check plan equality let plan = from_substrait_plan(&mut ctx, &proto).await?; + // #[allow(deprecated)] + // let plan = ctx.optimize(&plan)?; let plan_str_ref = format!("{:?}", plan_ref); let plan_str = format!("{:?}", plan); assert_eq!(plan_str_ref, plan_str); From 158c37a8313e789e22b45ad832f11cd2ee2178ed Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 13 Jan 2023 10:05:07 -0700 Subject: [PATCH 3/4] fix --- datafusion/substrait/tests/roundtrip.rs | 2 ++ 1 file changed, 2 insertions(+) diff --git a/datafusion/substrait/tests/roundtrip.rs b/datafusion/substrait/tests/roundtrip.rs index c31062edeee28..89b23b3d160d4 100644 --- a/datafusion/substrait/tests/roundtrip.rs +++ b/datafusion/substrait/tests/roundtrip.rs @@ -230,12 +230,14 @@ mod tests { Ok(()) } + #[allow(deprecated)] async fn roundtrip(sql: &str) -> Result<()> { let mut ctx = create_context().await?; let df = ctx.sql(sql).await?; let plan = df.into_optimized_plan()?; let proto = to_substrait_plan(&plan)?; let plan2 = from_substrait_plan(&mut ctx, &proto).await?; + let plan2 = ctx.optimize(&plan2)?; println!("{:#?}", plan); println!("{:#?}", plan2); From f47a34de64f78c63d5d1938a175516b117d01765 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 13 Jan 2023 10:08:25 -0700 Subject: [PATCH 4/4] ignore test --- datafusion/substrait/tests/roundtrip.rs | 1 + 1 file changed, 1 insertion(+) diff --git a/datafusion/substrait/tests/roundtrip.rs b/datafusion/substrait/tests/roundtrip.rs index 89b23b3d160d4..bbc6306185ea7 100644 --- a/datafusion/substrait/tests/roundtrip.rs +++ b/datafusion/substrait/tests/roundtrip.rs @@ -66,6 +66,7 @@ mod tests { roundtrip("SELECT * FROM data WHERE d AND a > 1").await } + #[ignore] // tracked in https://github.com/apache/arrow-datafusion/issues/4897 #[tokio::test] async fn select_with_limit() -> Result<()> { roundtrip_fill_na("SELECT * FROM data LIMIT 100").await