From 988d0c839796a85ff947583a20e5b85895118650 Mon Sep 17 00:00:00 2001 From: Szehon Ho Date: Wed, 3 Sep 2025 13:11:30 -0700 Subject: [PATCH 1/4] [SPARK-53482][SQL] MERGE INTO support nested case where source has less fields than target --- .../sql/catalyst/analysis/Analyzer.scala | 16 ++-- .../catalyst/analysis/AssignmentUtils.scala | 83 +++++++++++-------- .../sql/catalyst/types/DataTypeUtils.scala | 20 +++++ .../connector/MergeIntoTableSuiteBase.scala | 39 +++++---- 4 files changed, 99 insertions(+), 59 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 1896a1c7ac279..4d05eb0f0c40e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -1695,10 +1695,10 @@ class Analyzer(override val catalogManager: CatalogManager) extends RuleExecutor case UpdateStarAction(updateCondition) => // Use only source columns. Missing columns in target will be handled in // ResolveRowLevelCommandAssignments. - val assignments = targetTable.output.flatMap{ targetAttr => - sourceTable.output.find( - sourceCol => conf.resolver(sourceCol.name, targetAttr.name)) - .map(Assignment(targetAttr, _))} + val sourceAttrs = DataTypeUtils.nestedAttributes(sourceTable.output) + val assignments = sourceAttrs.map{ sourceAttr => + Assignment(UnresolvedAttribute(sourceAttr.name), + UnresolvedAttribute(sourceAttr.name))} UpdateAction( updateCondition.map(resolveExpressionByPlanChildren(_, m)), // For UPDATE *, the value must be from source table. @@ -1721,10 +1721,10 @@ class Analyzer(override val catalogManager: CatalogManager) extends RuleExecutor resolveExpressionByPlanOutput(_, m.sourceTable)) // Use only source columns. Missing columns in target will be handled in // ResolveRowLevelCommandAssignments. - val assignments = targetTable.output.flatMap{ targetAttr => - sourceTable.output.find( - sourceCol => conf.resolver(sourceCol.name, targetAttr.name)) - .map(Assignment(targetAttr, _))} + val sourceAttrs = DataTypeUtils.nestedAttributes(sourceTable.output) + val assignments = sourceAttrs.map{ sourceAttr => + Assignment(UnresolvedAttribute(sourceAttr.name), + UnresolvedAttribute(sourceAttr.name))} InsertAction( resolvedInsertCondition, resolveAssignments(assignments, m, MergeResolvePolicy.SOURCE)) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/AssignmentUtils.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/AssignmentUtils.scala index 43631e1afc403..d42e243d561b8 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/AssignmentUtils.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/AssignmentUtils.scala @@ -79,8 +79,6 @@ object AssignmentUtils extends SQLConfHelper with CastSupport { * This method processes and reorders given assignments so that each target column gets * an expression it should be set to. There must be exactly one assignment for each top-level * attribute and its value must be compatible. - *

- * Insert assignments cannot refer to nested columns. * * @param attrs table attributes * @param assignments insert assignments to align @@ -92,51 +90,68 @@ object AssignmentUtils extends SQLConfHelper with CastSupport { val errors = new mutable.ArrayBuffer[String]() - val (topLevelAssignments, nestedAssignments) = assignments.partition { assignment => - assignment.key.isInstanceOf[Attribute] - } - - if (nestedAssignments.nonEmpty) { - val nestedAssignmentsStr = nestedAssignments.map(_.sql).mkString(", ") - errors += s"INSERT assignment keys cannot be nested fields: $nestedAssignmentsStr" - } - val alignedAssignments = attrs.map { attr => - val matchingAssignments = topLevelAssignments.collect { - case assignment if assignment.key.semanticEquals(attr) => assignment - } - val resolvedValue = if (matchingAssignments.isEmpty) { - val defaultExpr = getDefaultValueExprOrNullLit( - attr, conf.useNullsForMissingDefaultColumnValues) - if (defaultExpr.isEmpty) { - errors += s"No assignment for '${attr.name}'" - } - defaultExpr.getOrElse(attr) - } else if (matchingAssignments.length > 1) { - val conflictingValuesStr = matchingAssignments.map(_.value.sql).mkString(", ") - errors += s"Multiple assignments for '${attr.name}': $conflictingValuesStr" - attr - } else { - val colPath = Seq(attr.name) - val actualAttr = restoreActualType(attr) - val value = matchingAssignments.head.value - TableOutputResolver.resolveUpdate( - "", value, actualAttr, conf, err => errors += err, colPath) - } - Assignment(attr, resolvedValue) + applyInsertAssignment( + col = restoreActualType(attr), + colExpr = attr, + assignments, + addError = err => errors += err, + colPath = Seq(attr.name)) } if (errors.nonEmpty) { throw QueryCompilationErrors.invalidRowLevelOperationAssignments(assignments, errors.toSeq) } - alignedAssignments + attrs.zip(alignedAssignments).map { case (attr, expr) => Assignment(attr, expr) } } private def restoreActualType(attr: Attribute): Attribute = { attr.withDataType(CharVarcharUtils.getRawType(attr.metadata).getOrElse(attr.dataType)) } + private def applyInsertAssignment( + col: Attribute, + colExpr: Expression, + assignments: Seq[Assignment], + addError: String => Unit, + colPath: Seq[String]): Expression = { + val (exactAssignments, otherAssignments) = assignments.partition { assignment => + assignment.key.semanticEquals(colExpr) + } + + val fieldAssignments = otherAssignments.filter { assignment => + assignment.key.exists(_.semanticEquals(colExpr)) + } + + if (exactAssignments.isEmpty && fieldAssignments.isEmpty) { + val defaultExpr = getDefaultValueExprOrNullLit( + col, conf.useNullsForMissingDefaultColumnValues) + if (defaultExpr.isEmpty) { + addError(s"No assignment for '${col.name}'") + } + defaultExpr.getOrElse(col) + } else if (exactAssignments.length > 1) { + val conflictingValuesStr = exactAssignments.map(_.value.sql).mkString(", ") + addError(s"Multiple assignments for '${col.name}': $conflictingValuesStr") + col + } else if (exactAssignments.nonEmpty && fieldAssignments.nonEmpty) { + val conflictingAssignments = exactAssignments ++ fieldAssignments + val conflictingAssignmentsStr = conflictingAssignments.map(_.sql).mkString(", ") + addError(s"Conflicting assignments for '${col.name}': $conflictingAssignmentsStr") + col + } else if (exactAssignments.nonEmpty) { + val colPath = Seq(col.name) + val actualAttr = restoreActualType(col) + val value = exactAssignments.head.value + TableOutputResolver.resolveUpdate( + "", value, actualAttr, conf, addError, colPath) + } else { + applyFieldAssignments(restoreActualType(col), + col, fieldAssignments, addError, colPath) + } + } + private def applyAssignments( col: Attribute, colExpr: Expression, diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/types/DataTypeUtils.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/types/DataTypeUtils.scala index c6e51aab4584b..746b38db8108a 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/types/DataTypeUtils.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/types/DataTypeUtils.scala @@ -19,6 +19,7 @@ package org.apache.spark.sql.catalyst.types import org.apache.spark.sql.catalyst.analysis.Resolver import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, Cast, Literal} import org.apache.spark.sql.catalyst.util.TypeUtils.toSQLId +import org.apache.spark.sql.connector.catalog.CatalogV2Implicits.MultipartIdentifierHelper import org.apache.spark.sql.errors.QueryCompilationErrors import org.apache.spark.sql.internal.SQLConf.StoreAssignmentPolicy import org.apache.spark.sql.internal.SQLConf.StoreAssignmentPolicy.{ANSI, STRICT} @@ -237,6 +238,25 @@ object DataTypeUtils { schema.map(toAttribute) } + def nestedAttributes(attrs: Seq[Attribute]): Seq[AttributeReference] = { + nestedAttributes(fromAttributes(attrs)) + } + + private def nestedAttributes(schema: StructType, colPath: Seq[String] = Seq()) + : Seq[AttributeReference] = { + schema.flatMap { field => + field.dataType match { + case structType: StructType => + val newColPath = colPath :+ field.name + nestedAttributes(structType, newColPath) + case _ => Seq( + AttributeReference((colPath :+ field.name).quoted, + field.dataType, field.nullable, field.metadata)() + ) + } + } + } + def fromAttributes(attributes: Seq[Attribute]): StructType = StructType(attributes.map(a => StructField(a.name, a.dataType, a.nullable, a.metadata))) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala index 10586adab1f6e..b60b3d1b2f914 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala @@ -2546,8 +2546,7 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase |USING source src |ON t.pk = src.pk |WHEN MATCHED THEN - | UPDATE SET s.c1 = -1, s.c2.m = map('k', 'v'), s.c2.a = array(-1), - | s.c2.c3 = src.s.c2.c3 + | UPDATE SET s.c1 = -1, s.c2.m = map('k', 'v'), s.c2.a = array(-1) |WHEN NOT MATCHED THEN | INSERT (pk, s, dep) VALUES (src.pk, | named_struct('c1', src.s.c1, @@ -2572,7 +2571,6 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase } } - // TODO- support schema evolution for missing nested types using UPDATE SET * and INSERT * test("merge into schema evolution replace column with nested field and set all columns") { Seq(true, false).foreach { withSchemaEvolution => withTempView("source") { @@ -2602,21 +2600,28 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase .createOrReplaceTempView("source") val schemaEvolutionClause = if (withSchemaEvolution) "WITH SCHEMA EVOLUTION" else "" - val exception = intercept[org.apache.spark.sql.AnalysisException] { - sql( - s"""MERGE $schemaEvolutionClause - |INTO $tableNameAsString t - |USING source src - |ON t.pk = src.pk - |WHEN MATCHED THEN - | UPDATE SET * - |WHEN NOT MATCHED THEN - | INSERT * - |""".stripMargin) + val mergeStmt = s"""MERGE $schemaEvolutionClause + |INTO $tableNameAsString t + |USING source src + |ON t.pk = src.pk + |WHEN MATCHED THEN + | UPDATE SET * + |WHEN NOT MATCHED THEN + | INSERT * + |""".stripMargin + if (withSchemaEvolution) { + sql(mergeStmt) + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Seq(Row(1, Row(10, Row(Seq(1, 2), Map("c" -> "d"), false)), "sales"), + Row(2, Row(20, Row(null, Map("e" -> "f"), true)), "engineering"))) + } else { + val exception = intercept[org.apache.spark.sql.AnalysisException] { + sql(mergeStmt) + } + assert(exception.errorClass.get == "FIELD_NOT_FOUND") + assert(exception.getMessage.contains("No such struct field `c3` in `a`, `m`. ")) } - - assert(exception.errorClass.get == "INCOMPATIBLE_DATA_FOR_TABLE.CANNOT_FIND_DATA") - assert(exception.getMessage.contains("Cannot find data for the output column `s`.`c2`.`a`")) } sql(s"DROP TABLE IF EXISTS $tableNameAsString") } From 6b74861b32050d640763f222104a1b984918a987 Mon Sep 17 00:00:00 2001 From: Szehon Ho Date: Thu, 4 Sep 2025 13:36:57 -0700 Subject: [PATCH 2/4] Fix some tests --- .../sql/catalyst/analysis/Analyzer.scala | 20 +++--- .../connector/MergeIntoTableSuiteBase.scala | 62 ++++++++----------- 2 files changed, 39 insertions(+), 43 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 4d05eb0f0c40e..950138dd2ab7c 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -1693,12 +1693,14 @@ class Analyzer(override val catalogManager: CatalogManager) extends RuleExecutor // The update value can access columns from both target and source tables. resolveAssignments(assignments, m, MergeResolvePolicy.BOTH)) case UpdateStarAction(updateCondition) => - // Use only source columns. Missing columns in target will be handled in + // Use only common columns. Missing columns in target will be handled in // ResolveRowLevelCommandAssignments. val sourceAttrs = DataTypeUtils.nestedAttributes(sourceTable.output) - val assignments = sourceAttrs.map{ sourceAttr => - Assignment(UnresolvedAttribute(sourceAttr.name), - UnresolvedAttribute(sourceAttr.name))} + val targetAttrs = DataTypeUtils.nestedAttributes(targetTable.output) + val commonAttrs = sourceAttrs.filter(s => + targetAttrs.exists(t => conf.resolver(t.name, s.name))) + val assignments = commonAttrs.map{ a => + Assignment(UnresolvedAttribute(a.name), UnresolvedAttribute(a.name))} UpdateAction( updateCondition.map(resolveExpressionByPlanChildren(_, m)), // For UPDATE *, the value must be from source table. @@ -1719,12 +1721,14 @@ class Analyzer(override val catalogManager: CatalogManager) extends RuleExecutor // access columns from the source table. val resolvedInsertCondition = insertCondition.map( resolveExpressionByPlanOutput(_, m.sourceTable)) - // Use only source columns. Missing columns in target will be handled in + // Use only common columns. Missing columns in target will be handled in // ResolveRowLevelCommandAssignments. val sourceAttrs = DataTypeUtils.nestedAttributes(sourceTable.output) - val assignments = sourceAttrs.map{ sourceAttr => - Assignment(UnresolvedAttribute(sourceAttr.name), - UnresolvedAttribute(sourceAttr.name))} + val targetAttrs = DataTypeUtils.nestedAttributes(targetTable.output) + val commonAttrs = sourceAttrs.filter(s => + targetAttrs.exists(t => conf.resolver(t.name, s.name))) + val assignments = commonAttrs.map{ a => + Assignment(UnresolvedAttribute(a.name), UnresolvedAttribute(a.name))} InsertAction( resolvedInsertCondition, resolveAssignments(assignments, m, MergeResolvePolicy.SOURCE)) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala index b60b3d1b2f914..1f510bb2bb307 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/MergeIntoTableSuiteBase.scala @@ -2481,7 +2481,7 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase .createOrReplaceTempView("source") val schemaEvolutionClause = if (withSchemaEvolution) "WITH SCHEMA EVOLUTION" else "" - val mergeStmt = + sql( s"""MERGE $schemaEvolutionClause |INTO $tableNameAsString t |USING source src @@ -2490,22 +2490,16 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase | UPDATE SET * |WHEN NOT MATCHED THEN | INSERT * - |""".stripMargin + |""".stripMargin) - if (withSchemaEvolution) { - sql(mergeStmt) - checkAnswer( - sql(s"SELECT * FROM $tableNameAsString"), - Seq(Row(1, Row(10, Row(Seq(3, 4), Map("c" -> "d"), false)), "sales"), - Row(2, Row(20, Row(Seq(4, 5), Map("e" -> "f"), true)), "engineering"))) + val expectedAnswer = if (withSchemaEvolution) { + Seq(Row(1, Row(10, Row(Seq(3, 4), Map("c" -> "d"), false)), "sales"), + Row(2, Row(20, Row(Seq(4, 5), Map("e" -> "f"), true)), "engineering")) } else { - val exception = intercept[org.apache.spark.sql.AnalysisException] { - sql(mergeStmt) - } - assert(exception.errorClass.get == "INCOMPATIBLE_DATA_FOR_TABLE.EXTRA_STRUCT_FIELDS") - assert(exception.getMessage.contains( - "Cannot write extra fields `c3` to the struct `s`.`c2`")) + Seq(Row(1, Row(10, Row(Seq(3, 4), Map("c" -> "d"))), "sales"), + Row(2, Row(20, Row(Seq(4, 5), Map("e" -> "f"))), "engineering")) } + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), expectedAnswer) } sql(s"DROP TABLE IF EXISTS $tableNameAsString") } @@ -2546,7 +2540,8 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase |USING source src |ON t.pk = src.pk |WHEN MATCHED THEN - | UPDATE SET s.c1 = -1, s.c2.m = map('k', 'v'), s.c2.a = array(-1) + | UPDATE SET s.c1 = -1, s.c2.m = map('k', 'v'), s.c2.a = array(-1), + | s.c2.c3 = src.s.c2.c3 |WHEN NOT MATCHED THEN | INSERT (pk, s, dep) VALUES (src.pk, | named_struct('c1', src.s.c1, @@ -2600,28 +2595,25 @@ abstract class MergeIntoTableSuiteBase extends RowLevelOperationSuiteBase .createOrReplaceTempView("source") val schemaEvolutionClause = if (withSchemaEvolution) "WITH SCHEMA EVOLUTION" else "" - val mergeStmt = s"""MERGE $schemaEvolutionClause - |INTO $tableNameAsString t - |USING source src - |ON t.pk = src.pk - |WHEN MATCHED THEN - | UPDATE SET * - |WHEN NOT MATCHED THEN - | INSERT * - |""".stripMargin - if (withSchemaEvolution) { - sql(mergeStmt) - checkAnswer( - sql(s"SELECT * FROM $tableNameAsString"), - Seq(Row(1, Row(10, Row(Seq(1, 2), Map("c" -> "d"), false)), "sales"), - Row(2, Row(20, Row(null, Map("e" -> "f"), true)), "engineering"))) + sql( + s"""MERGE $schemaEvolutionClause + |INTO $tableNameAsString t + |USING source src + |ON t.pk = src.pk + |WHEN MATCHED THEN + | UPDATE SET * + |WHEN NOT MATCHED THEN + | INSERT * + |""".stripMargin) + + val expectedAnswer = if (withSchemaEvolution) { + Seq(Row(1, Row(10, Row(Seq(1, 2), Map("c" -> "d"), false)), "sales"), + Row(2, Row(20, Row(null, Map("e" -> "f"), true)), "engineering")) } else { - val exception = intercept[org.apache.spark.sql.AnalysisException] { - sql(mergeStmt) - } - assert(exception.errorClass.get == "FIELD_NOT_FOUND") - assert(exception.getMessage.contains("No such struct field `c3` in `a`, `m`. ")) + Seq(Row(1, Row(10, Row(Seq(1, 2), Map("c" -> "d"))), "sales"), + Row(2, Row(20, Row(null, Map("e" -> "f"))), "engineering")) } + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), expectedAnswer) } sql(s"DROP TABLE IF EXISTS $tableNameAsString") } From d37d3510912b6f9769bede0ec387d52257adc3f4 Mon Sep 17 00:00:00 2001 From: Szehon Ho Date: Thu, 4 Sep 2025 13:55:15 -0700 Subject: [PATCH 3/4] Fix PlanResolutionSuite --- .../command/PlanResolutionSuite.scala | 32 ++++++++++++------- 1 file changed, 21 insertions(+), 11 deletions(-) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/command/PlanResolutionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/command/PlanResolutionSuite.scala index ecc293a5acc2a..98643f0fb255e 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/command/PlanResolutionSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/command/PlanResolutionSuite.scala @@ -1624,10 +1624,10 @@ class PlanResolutionSuite extends SharedSparkSession with AnalysisTest { if (starInUpdate) { assert(updateAssigns.size == 2) - assert(updateAssigns(0).key.asInstanceOf[AttributeReference].sameRef(ti)) - assert(updateAssigns(0).value.asInstanceOf[AttributeReference].sameRef(si)) - assert(updateAssigns(1).key.asInstanceOf[AttributeReference].sameRef(ts)) - assert(updateAssigns(1).value.asInstanceOf[AttributeReference].sameRef(ss)) + assert(updateAssigns(0).key.asInstanceOf[AttributeReference].sameRef(ts)) + assert(updateAssigns(0).value.asInstanceOf[AttributeReference].sameRef(ss)) + assert(updateAssigns(1).key.asInstanceOf[AttributeReference].sameRef(ti)) + assert(updateAssigns(1).value.asInstanceOf[AttributeReference].sameRef(si)) } else { assert(updateAssigns.size == 1) assert(updateAssigns.head.key.asInstanceOf[AttributeReference].sameRef(ts)) @@ -1639,15 +1639,23 @@ class PlanResolutionSuite extends SharedSparkSession with AnalysisTest { target: LogicalPlan, source: LogicalPlan, insertCondAttr: Option[AttributeReference], - insertAssigns: Seq[Assignment]): Unit = { + insertAssigns: Seq[Assignment], + starInInsert: Boolean = false): Unit = { val (si, ss) = getAttributes(source) val (ti, ts) = getAttributes(target) insertCondAttr.foreach(a => assert(a.sameRef(ss))) assert(insertAssigns.size == 2) - assert(insertAssigns(0).key.asInstanceOf[AttributeReference].sameRef(ti)) - assert(insertAssigns(0).value.asInstanceOf[AttributeReference].sameRef(si)) - assert(insertAssigns(1).key.asInstanceOf[AttributeReference].sameRef(ts)) - assert(insertAssigns(1).value.asInstanceOf[AttributeReference].sameRef(ss)) + if (starInInsert) { + assert(insertAssigns(0).key.asInstanceOf[AttributeReference].sameRef(ts)) + assert(insertAssigns(0).value.asInstanceOf[AttributeReference].sameRef(ss)) + assert(insertAssigns(1).key.asInstanceOf[AttributeReference].sameRef(ti)) + assert(insertAssigns(1).value.asInstanceOf[AttributeReference].sameRef(si)) + } else { + assert(insertAssigns(0).key.asInstanceOf[AttributeReference].sameRef(ti)) + assert(insertAssigns(0).value.asInstanceOf[AttributeReference].sameRef(si)) + assert(insertAssigns(1).key.asInstanceOf[AttributeReference].sameRef(ts)) + assert(insertAssigns(1).value.asInstanceOf[AttributeReference].sameRef(ss)) + } } def checkNotMatchedBySourceClausesResolution( @@ -1726,7 +1734,8 @@ class PlanResolutionSuite extends SharedSparkSession with AnalysisTest { checkMergeConditionResolution(target, source, mergeCondition) checkMatchedClausesResolution(target, source, Some(dl), Some(ul), updateAssigns, starInUpdate = true) - checkNotMatchedClausesResolution(target, source, Some(il), insertAssigns) + checkNotMatchedClausesResolution(target, source, Some(il), insertAssigns, + starInInsert = true) assert(withSchemaEvolution === false) case other => fail("Expect MergeIntoTable, but got:\n" + other.treeString) @@ -1753,7 +1762,8 @@ class PlanResolutionSuite extends SharedSparkSession with AnalysisTest { checkMergeConditionResolution(target, source, mergeCondition) checkMatchedClausesResolution(target, source, None, None, updateAssigns, starInUpdate = true) - checkNotMatchedClausesResolution(target, source, None, insertAssigns) + checkNotMatchedClausesResolution(target, source, None, insertAssigns, + starInInsert = true) assert(withSchemaEvolution === false) case other => fail("Expect MergeIntoTable, but got:\n" + other.treeString) From d59d9b93c3cdc657b7eae9fa9e91c8c0ca9c06b2 Mon Sep 17 00:00:00 2001 From: Szehon Ho Date: Fri, 5 Sep 2025 11:18:24 -0700 Subject: [PATCH 4/4] Fix another test --- .../execution/command/AlignMergeAssignmentsSuite.scala | 8 -------- 1 file changed, 8 deletions(-) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/command/AlignMergeAssignmentsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/command/AlignMergeAssignmentsSuite.scala index cd099a2a94813..cb72918a1bb6e 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/command/AlignMergeAssignmentsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/command/AlignMergeAssignmentsSuite.scala @@ -631,14 +631,6 @@ class AlignMergeAssignmentsSuite extends AlignAssignmentsSuiteBase { | INSERT (i, l, txt, txt) VALUES (src.i, src.l, src.txt, src.txt) |""".stripMargin, "Multiple assignments for 'txt'") - - assertAnalysisException( - """MERGE INTO nested_struct_table t USING nested_struct_table src - |ON t.i = src.i - |WHEN NOT MATCHED THEN - | INSERT (s.n_i) VALUES (1) - |""".stripMargin, - "INSERT assignment keys cannot be nested fields: t.s.`n_i` = 1") } test("updates to nested structs in arrays") {