Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 7 additions & 9 deletions datafusion/core/src/physical_planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -869,11 +869,9 @@ impl DefaultPhysicalPlanner {
e.context(format!("MERGE INTO operation on table '{table_name}'"))
})?;
let input_exec = children.one()?;
let target_schema = DFSchema::try_from_qualified_schema(
table_name.clone(),
&target.schema(),
)?;
let merge_schema = Arc::new(target_schema.join(input.schema())?);
let merge_schema = Arc::new(
merge_op.expression_schema(&target.schema(), input.schema())?,
);
provider
.merge_into(
session_state,
Expand Down Expand Up @@ -3571,8 +3569,8 @@ mod tests {
ctx.register_table("source", source)?;

ctx.sql(
"MERGE INTO target AS t USING source AS s ON t.id = s.id \
WHEN MATCHED AND t.id > s.id THEN DELETE",
"MERGE INTO target AS t USING source AS target ON t.id = target.id \
WHEN MATCHED AND t.id > target.id THEN DELETE",
)
.await?
.create_physical_plan()
Expand All @@ -3583,11 +3581,11 @@ mod tests {
captured.as_ref().expect("merge_into should be called");
assert_eq!(*clause_count, 1);
assert_eq!(
merge_schema.index_of_column(&Column::new(Some("target"), "id"))?,
merge_schema.index_of_column(&Column::new(Some("t"), "id"))?,
0
);
assert_eq!(
merge_schema.index_of_column(&Column::new(Some("s"), "id"))?,
merge_schema.index_of_column(&Column::new(Some("target"), "id"))?,
1
);
assert_contains!(physical_on, "index: 0");
Expand Down
141 changes: 93 additions & 48 deletions datafusion/core/tests/sql/sql_api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,10 @@
// under the License.

use datafusion::prelude::*;
use datafusion_common::assert_contains;
use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion};
use datafusion_common::{TableReference, assert_contains};
use datafusion_expr::dml::MergeIntoOp;
use datafusion_expr::{Expr, LogicalPlan, WriteOp};

use tempfile::TempDir;

Expand Down Expand Up @@ -215,69 +218,111 @@ async fn merge_into_context() -> SessionContext {
ctx
}

async fn assert_merge_sql_error(ctx: &SessionContext, sql: &str, expected: &str) {
let err = ctx.sql(sql).await.unwrap_err();
assert_contains!(err.strip_backtrace(), expected);
}

async fn assert_merge_physical_error(ctx: &SessionContext, sql: &str, expected: &str) {
let err = ctx
let result = ctx
.sql(sql)
.await
.unwrap()
.unwrap_or_else(|error| panic!("failed to plan MERGE SQL:\n{sql}\n{error}"))
.create_physical_plan()
.await
.unwrap_err();
assert_contains!(err.strip_backtrace(), expected);
.await;
let err = match result {
Ok(_) => panic!("expected physical planning to fail:\n{sql}"),
Err(error) => error,
};
let actual = err.strip_backtrace();
assert!(
actual.contains(expected),
"MERGE SQL:\n{sql}\n\nExpected:\n{expected}\n\nActual:\n{actual}"
);
}

#[tokio::test]
async fn merge_into_rejects_source_alias_colliding_with_target_name() {
// Canonicalizing `t.id` to `target.id` must not collapse it onto a source
// that also uses `target` as its qualifier.
let ctx = merge_into_context().await;
async fn merge_operation(ctx: &SessionContext, sql: &str) -> Box<MergeIntoOp> {
let plan = ctx.state().create_logical_plan(sql).await.unwrap();
let LogicalPlan::Dml(dml) = plan else {
panic!("expected MERGE DML")
};
let WriteOp::MergeInto(merge_op) = dml.op else {
panic!("expected MERGE operation")
};
merge_op
}

for target_ref in ["target", "public.target", "datafusion.public.target"] {
assert_merge_sql_error(
&ctx,
&format!(
"MERGE INTO {target_ref} AS t USING source AS target \
ON t.id = target.id WHEN MATCHED THEN DELETE"
),
&format!(
"MERGE source may not use the target table name '{target_ref}' \
as a qualifier"
),
)
.await;
}
fn has_outer_reference_to(expr: &Expr, qualifier: &TableReference) -> bool {
let mut found = false;
expr.apply(|expr| {
let outer_refs = match expr {
Expr::Exists(exists) => Some(&exists.subquery.outer_ref_columns),
Expr::InSubquery(in_subquery) => {
Some(&in_subquery.subquery.outer_ref_columns)
}
Expr::SetComparison(set_comparison) => {
Some(&set_comparison.subquery.outer_ref_columns)
}
Expr::ScalarSubquery(subquery) => Some(&subquery.outer_ref_columns),
_ => None,
};
found = outer_refs.is_some_and(|outer_refs| {
outer_refs.iter().any(|expr| {
matches!(
expr,
Expr::OuterReferenceColumn(_, column)
if column.relation.as_ref() == Some(qualifier)
)
})
});
Ok(if found {
TreeNodeRecursion::Stop
} else {
TreeNodeRecursion::Continue
})
})
.unwrap();
found
}

#[tokio::test]
async fn merge_into_rejects_subqueries_correlated_to_target_alias() {
async fn merge_into_preserves_target_alias_in_correlated_subquery() {
let ctx = merge_into_context().await;
assert_merge_sql_error(
&ctx,
"MERGE INTO target AS t USING source AS s \
let direct_exists = "MERGE INTO target AS t USING source AS s \
ON EXISTS (SELECT 1 FROM source AS x WHERE x.id = t.id) \
WHEN MATCHED THEN DELETE",
"MERGE subqueries correlated to target alias 't' are not supported",
)
.await;
WHEN MATCHED THEN DELETE";
let direct_in = "MERGE INTO target AS t USING source AS s \
ON t.id IN (SELECT x.id FROM source AS x WHERE x.id = t.id) \
WHEN MATCHED THEN DELETE";
let direct_any = "MERGE INTO target AS t USING source AS s \
ON t.id = ANY (SELECT x.id FROM source AS x WHERE x.id = t.id) \
WHEN MATCHED THEN DELETE";
let direct_all = "MERGE INTO target AS t USING source AS s \
ON t.id = ALL (SELECT x.id FROM source AS x WHERE x.id = t.id) \
WHEN MATCHED THEN DELETE";
let direct_scalar = "MERGE INTO target AS t USING source AS s \
ON t.id = (SELECT max(x.id) FROM source AS x WHERE x.id = t.id) \
WHEN MATCHED THEN DELETE";

// Source-correlated and uncorrelated subqueries remain supported through
// logical optimization.
for sql in [
"MERGE INTO target AS t USING source AS s \
ON EXISTS (SELECT 1 FROM source AS x WHERE x.id = s.id) \
WHEN MATCHED THEN DELETE",
"MERGE INTO target AS t USING source AS s \
ON t.id = ANY (SELECT id FROM source) \
WHEN MATCHED THEN DELETE",
direct_exists,
direct_in,
direct_any,
direct_all,
direct_scalar,
] {
assert_merge_physical_error(&ctx, sql, "MERGE INTO not supported for Base table")
.await;
let merge_op = merge_operation(&ctx, sql).await;
assert_eq!(merge_op.target_qualifier(), &TableReference::bare("t"));
assert!(has_outer_reference_to(
&merge_op.on,
&TableReference::bare("t")
));
}

let shadowed_correlation = "MERGE INTO target AS t USING source AS s \
ON EXISTS (SELECT 1 FROM source AS t \
WHERE EXISTS (SELECT 1 FROM source AS x WHERE x.id = t.id)) \
WHEN MATCHED THEN DELETE";
let merge_op = merge_operation(&ctx, shadowed_correlation).await;
assert!(!has_outer_reference_to(
&merge_op.on,
&TableReference::bare("t")
));
}

#[tokio::test]
Expand Down
Loading
Loading