diff --git a/datafusion/sql/src/expr/subquery.rs b/datafusion/sql/src/expr/subquery.rs index 662c44f6f2620..096b9ff4a0955 100644 --- a/datafusion/sql/src/expr/subquery.rs +++ b/datafusion/sql/src/expr/subquery.rs @@ -16,9 +16,13 @@ // under the License. use crate::planner::{ContextProvider, PlannerContext, SqlToRel}; -use datafusion_common::{DFSchema, Diagnostic, Result, Span, Spans, plan_err}; +use datafusion_common::tree_node::{Transformed, TreeNode}; +use datafusion_common::{ + Column, DFSchema, Diagnostic, ExprSchema, Result, Span, Spans, plan_err, +}; use datafusion_expr::expr::{Exists, InSubquery, SetComparison, SetQuantifier}; -use datafusion_expr::{Expr, LogicalPlan, Subquery}; +use datafusion_expr::utils::conjunction; +use datafusion_expr::{Expr, LogicalPlan, LogicalPlanBuilder, Subquery}; use sqlparser::ast::Expr as SQLExpr; use sqlparser::ast::{BinaryOperator, Query, SelectItem, SetExpr}; use std::sync::Arc; @@ -67,6 +71,21 @@ impl SqlToRel<'_, S> { } let sub_plan = self.query_to_plan(subquery, planner_context)?; + let subquery_arity = sub_plan.schema().fields().len(); + if subquery_arity > 1 + && let SQLExpr::Tuple(values) = &expr + { + let result = self.build_tuple_in_subquery( + values.clone(), + sub_plan, + negated, + input_schema, + planner_context, + spans, + ); + planner_context.pop_outer_query_schema(); + return result; + } let outer_ref_columns = sub_plan.all_out_ref_exprs(); planner_context.pop_outer_query_schema(); @@ -90,6 +109,51 @@ impl SqlToRel<'_, S> { ))) } + fn build_tuple_in_subquery( + &self, + values: Vec, + sub_plan: LogicalPlan, + negated: bool, + input_schema: &DFSchema, + planner_context: &mut PlannerContext, + spans: Spans, + ) -> Result { + let tuple_arity = values.len(); + let subquery_arity = sub_plan.schema().fields().len(); + if tuple_arity != subquery_arity { + return plan_err!( + "IN subquery tuple has {tuple_arity} fields but the subquery returns {subquery_arity} columns" + ); + } + + let comparisons = values + .into_iter() + .enumerate() + .map(|(index, value)| { + let left = self.sql_to_expr(value, input_schema, planner_context)?; + let left = to_outer_reference(left, input_schema)?; + let right = + Expr::Column(Column::from(sub_plan.schema().qualified_field(index))); + Ok(Expr::eq(left, right)) + }) + .collect::>>()?; + let predicate = conjunction(comparisons) + .map_or(plan_err!("IN subquery tuple cannot be empty"), Ok)?; + let filtered = LogicalPlanBuilder::from(sub_plan) + .filter(predicate)? + .build()?; + let outer_ref_columns = filtered.all_out_ref_exprs(); + + Ok(Expr::Exists(Exists { + subquery: Subquery { + subquery: Arc::new(filtered), + outer_ref_columns, + spans, + }, + negated, + })) + } + pub(super) fn parse_scalar_subquery( &self, subquery: Query, @@ -206,3 +270,18 @@ impl SqlToRel<'_, S> { ))) } } + +fn to_outer_reference(expr: Expr, outer_schema: &DFSchema) -> Result { + expr.transform_up(|expr| match expr { + Expr::Column(col) => { + let field = outer_schema.field_from_column(&col)?; + Ok(Transformed::yes(Expr::OuterReferenceColumn( + Arc::clone(field), + col, + ))) + } + Expr::OuterReferenceColumn(_, _) => Ok(Transformed::no(expr)), + _ => Ok(Transformed::no(expr)), + }) + .map(|transformed| transformed.data) +} diff --git a/datafusion/sqllogictest/test_files/subquery_projection.slt b/datafusion/sqllogictest/test_files/subquery_projection.slt index 20de82531561a..9be629fb47946 100644 --- a/datafusion/sqllogictest/test_files/subquery_projection.slt +++ b/datafusion/sqllogictest/test_files/subquery_projection.slt @@ -55,3 +55,36 @@ FROM outer_values o; 2 false 3 NULL 4 false + +# Snowflake row-valued IN treats comparisons containing NULL as non-matches. +query BBBBB +WITH vals AS ( + SELECT * FROM (VALUES (1, 2), (4, 5), (4, NULL)) AS t(a, b) +), empty AS ( + SELECT a, b FROM vals WHERE false +) +SELECT + (1, 2) IN (SELECT a, b FROM vals) AS exact_match, + (9, 9) IN (SELECT a, b FROM vals) AS no_match, + (4, NULL) IN (SELECT a, b FROM vals) AS null_non_match, + (9, NULL) IN (SELECT a, b FROM vals) AS different_with_null, + (9, NULL) IN (SELECT a, b FROM empty) AS empty_rhs; +---- +true false false false false + +# Row-valued IN works as a top-level predicate as well as a projected value. +query I rowsort +WITH candidates AS ( + SELECT * FROM (VALUES (1, 2), (4, 5), (4, NULL)) AS t(a, b) +), probes AS ( + SELECT * FROM (VALUES (1, 1, 2), (2, 9, 9), (3, 4, NULL)) AS t(id, a, b) +) +SELECT id +FROM probes +WHERE (a, b) IN (SELECT a, b FROM candidates); +---- +1 + +# Tuple arity mismatches are rejected during SQL planning. +statement error IN subquery tuple has 2 fields but the subquery returns 3 columns +SELECT (1, 2) IN (SELECT column1, column2, column3 FROM VALUES (1, 2, 3));