diff --git a/common/utils/src/main/resources/error/error-conditions.json b/common/utils/src/main/resources/error/error-conditions.json index 94d6b897b18c8..d5208af6593ba 100644 --- a/common/utils/src/main/resources/error/error-conditions.json +++ b/common/utils/src/main/resources/error/error-conditions.json @@ -1248,6 +1248,60 @@ ], "sqlState" : "22003" }, + "COLUMN_UPDATE_DUPLICATE_REQUIRED_DATA_ATTRIBUTE" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` but declared duplicate column(s) in `requiredDataAttributes()`. Each column must be declared at most once." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_EMPTY_REQUIRED_DATA_ATTRIBUTES" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` but returned an empty array from `requiredDataAttributes()`. Connectors that opt into column-level updates must declare at least one required data attribute." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_METADATA_REQUIRED_DATA_ATTRIBUTE" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` but declared metadata column(s) in `requiredDataAttributes()`. Declare metadata columns in `requiredMetadataAttributes()` instead." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_NESTED_REQUIRED_DATA_ATTRIBUTE" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` but declared nested field(s) in `requiredDataAttributes()`. Declare the root struct column instead of a nested field." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_REQUIRED_DATA_ATTRIBUTES_MISSING_UPDATED_COLUMNS" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` but its `requiredDataAttributes()` does not cover every column being updated. Missing columns: . Connectors must include every column reported by `RowLevelOperationInfo.updatedColumns()` in `requiredDataAttributes()`." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_SPLIT_ROW_ID_NOT_DECLARED" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` and represents UPDATE as delete and insert, but row ID column(s) do not reach the reinserted row, which then has no identity for the connector to place it by. Declare each data row ID column in `requiredDataAttributes()`, and return each metadata row ID column from `requiredMetadataAttributes()` with `MetadataColumn.PRESERVE_ON_REINSERT` set." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_SPLIT_ROW_ID_REASSIGNMENT" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` and represents UPDATE as delete and insert, so UPDATE cannot assign row ID column(s) . The reinserted row carries only the columns in `requiredDataAttributes()`, and with a new row ID it cannot be matched to the row whose other columns the connector must preserve. Do not assign row ID columns, or include every table column in `requiredDataAttributes()`." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_UNDECLARED_WRITE_REQUIREMENT_COLUMNS" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates`, but its write requires a distribution or ordering by column(s) that the column-level UPDATE does not read. Declare each data column in `requiredDataAttributes()` and each metadata column in `requiredMetadataAttributes()`." + ], + "sqlState" : "42000" + }, + "COLUMN_UPDATE_UNKNOWN_REQUIRED_DATA_ATTRIBUTE" : { + "message" : [ + "Connector mixes in `SupportsColumnUpdates` but declared column(s) in `requiredDataAttributes()` that do not exist in the table." + ], + "sqlState" : "42703" + }, "COMPARATOR_RETURNS_NULL" : { "message" : [ "The comparator has returned a NULL for a comparison between and .", @@ -2397,6 +2451,12 @@ ], "sqlState" : "42K03" }, + "DATA_SOURCE_WRITE_COLUMN_UPDATE_NOT_IMPLEMENTED" : { + "message" : [ + " does not override `writeColumnUpdate(record)`. A data writer that receives rows in the `LogicalWriteInfo.columnUpdateSchema()` layout must override `writeColumnUpdate(record)`, or `writeColumnUpdate(metadata, record)` if the operation returns metadata columns from `requiredMetadataAttributes()`." + ], + "sqlState" : "0A000" + }, "DATETIME_FIELD_OUT_OF_BOUNDS" : { "message" : [ "." diff --git a/project/MimaExcludes.scala b/project/MimaExcludes.scala index 113757032b5c2..1ef63babe7f38 100644 --- a/project/MimaExcludes.scala +++ b/project/MimaExcludes.scala @@ -46,7 +46,10 @@ object MimaExcludes { "org.apache.spark.ml.regression.DecisionTreeRegressionModel.numLeave"), // [SPARK-59154] Remove unused prediction variance helper after inlining its implementation. ProblemFilters.exclude[DirectMissingMethodProblem]( - "org.apache.spark.ml.regression.DecisionTreeRegressionModel.predictVariance") + "org.apache.spark.ml.regression.DecisionTreeRegressionModel.predictVariance"), + // [SPARK-58111] Write schema narrowing for column-level UPDATE in DSv2 + ProblemFilters.exclude[ReversedMissingMethodProblem]( + "org.apache.spark.sql.connector.write.RowLevelOperationInfo.updatedColumns") ) // Exclude rules for 4.3.x from 4.2.0 (add 4.3-specific filters below as needed). diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DataWriter.java b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DataWriter.java index a4ec1abc9dd7d..963bc5e6c8e0d 100644 --- a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DataWriter.java +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DataWriter.java @@ -20,7 +20,9 @@ import java.io.Closeable; import java.io.IOException; import java.util.Iterator; +import java.util.Map; +import org.apache.spark.SparkUnsupportedOperationException; import org.apache.spark.annotation.Evolving; import org.apache.spark.sql.connector.metric.CustomTaskMetric; @@ -82,6 +84,60 @@ default void write(T metadata, T record) throws IOException { write(record); } + /** + * Writes one updated or copied record with metadata in the column update layout. + *

+ * When {@link LogicalWriteInfo#columnUpdateSchema()} is present for a row-level operation that + * does not mix in {@link SupportsDelta}, Spark passes updated and copied records to this method + * instead of {@link #write(Object, Object)} if the operation returns a non-empty + * {@link RowLevelOperation#requiredMetadataAttributes()}, and to + * {@link #writeColumnUpdate(Object)} otherwise. The record follows + * {@link LogicalWriteInfo#columnUpdateSchema()} and the metadata follows + * {@link LogicalWriteInfo#metadataSchema()}. Operations that mix in {@link SupportsDelta} receive + * such rows through {@link DeltaWriter#update} and {@link DeltaWriter#reinsert} instead. + *

+ * By default, delegates to {@link #writeColumnUpdate(Object)} and drops the metadata. + *

+ * If this method fails (by throwing an exception), {@link #abort()} will be called and this + * data writer is considered to have been failed. + * + * @throws IOException if failure happens during disk/network IO like writing files. + * @throws SparkUnsupportedOperationException if neither this method nor + * {@link #writeColumnUpdate(Object)} is overridden. + * + * @since 4.4.0 + */ + default void writeColumnUpdate(T metadata, T record) throws IOException { + writeColumnUpdate(record); + } + + /** + * Writes one updated or copied record without metadata in the column update layout. + *

+ * When {@link LogicalWriteInfo#columnUpdateSchema()} is present for a row-level operation that + * does not mix in {@link SupportsDelta}, Spark passes updated and copied records to this method + * instead of {@link #write(Object)} if the operation returns no + * {@link RowLevelOperation#requiredMetadataAttributes()}. The record follows + * {@link LogicalWriteInfo#columnUpdateSchema()}. + *

+ * A writer for such an operation must override this method, unless the operation returns a + * non-empty {@link RowLevelOperation#requiredMetadataAttributes()} and the writer overrides + * {@link #writeColumnUpdate(Object, Object)}. + *

+ * If this method fails (by throwing an exception), {@link #abort()} will be called and this + * data writer is considered to have been failed. + * + * @throws IOException if failure happens during disk/network IO like writing files. + * @throws SparkUnsupportedOperationException if this method is not overridden. + * + * @since 4.4.0 + */ + default void writeColumnUpdate(T record) throws IOException { + throw new SparkUnsupportedOperationException( + "DATA_SOURCE_WRITE_COLUMN_UPDATE_NOT_IMPLEMENTED", + Map.of("class", getClass().getName())); + } + /** * Writes one record. *

diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DeltaWriter.java b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DeltaWriter.java index a7ab0c162ddec..b973d55a68620 100644 --- a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DeltaWriter.java +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/DeltaWriter.java @@ -40,6 +40,9 @@ public interface DeltaWriter extends DataWriter { /** * Updates a row. + *

+ * When {@link LogicalWriteInfo#columnUpdateSchema()} is present, the {@code row} follows it; + * otherwise it follows {@link LogicalWriteInfo#schema()}. * * @param metadata values for metadata columns that were projected but are not part of the row ID * @param id a row ID to update @@ -52,6 +55,9 @@ public interface DeltaWriter extends DataWriter { * Reinserts a row with metadata. *

* This method handles the insert portion of updated rows split into deletes and inserts. + *

+ * When {@link LogicalWriteInfo#columnUpdateSchema()} is present, the {@code row} follows it; + * otherwise it follows {@link LogicalWriteInfo#schema()}. * * @param metadata values for metadata columns * @param row a row to reinsert diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/LogicalWriteInfo.java b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/LogicalWriteInfo.java index e7c2efc6d672d..221a2e4fffe04 100644 --- a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/LogicalWriteInfo.java +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/LogicalWriteInfo.java @@ -47,6 +47,10 @@ public interface LogicalWriteInfo { /** * the schema of the input data from Spark to data source. + *

+ * When {@link #columnUpdateSchema()} is present, this schema covers only newly inserted rows, as + * updated, copied, and reinserted rows follow {@link #columnUpdateSchema()}. It is then empty + * when the command inserts no new rows, such as UPDATE. */ StructType schema(); @@ -65,4 +69,18 @@ default Optional metadataSchema() { throw new SparkUnsupportedOperationException( "DATA_SOURCE_METADATA_SCHEMA_NOT_IMPLEMENTED", Map.of("class", getClass().getName())); } + + /** + * the schema of updated, copied, and reinserted rows from Spark to data source in a + * column-level update. Present when the operation mixes in {@link SupportsColumnUpdates} and + * Spark delivers these rows with the columns of + * {@link SupportsColumnUpdates#requiredDataAttributes()}, which currently happens only for + * UPDATE. It covers every table column if every column is declared. When present, + * {@link #schema()} covers only newly inserted rows. + * + * @since 4.4.0 + */ + default Optional columnUpdateSchema() { + return Optional.empty(); + } } diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/RowLevelOperationInfo.java b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/RowLevelOperationInfo.java index e3d7397aed91b..77525eaa906ba 100644 --- a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/RowLevelOperationInfo.java +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/RowLevelOperationInfo.java @@ -18,6 +18,7 @@ package org.apache.spark.sql.connector.write; import org.apache.spark.annotation.Experimental; +import org.apache.spark.sql.connector.expressions.NamedReference; import org.apache.spark.sql.connector.write.RowLevelOperation.Command; import org.apache.spark.sql.util.CaseInsensitiveStringMap; @@ -37,4 +38,16 @@ public interface RowLevelOperationInfo { * Returns the row-level SQL command (e.g. DELETE, UPDATE, MERGE). */ Command command(); + + /** + * Returns the columns being updated by this operation. Currently only UPDATE populates it; + * other commands report an empty array. + *

+ * A column is reported only if it is assigned a new value, so identity assignments such as + * {@code SET a = a} are excluded. Nested struct field updates are reported at root-column + * granularity (e.g. {@code SET s.c1 = -1} returns {@code s}). + * + * @since 4.4.0 + */ + NamedReference[] updatedColumns(); } diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/SupportsColumnUpdates.java b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/SupportsColumnUpdates.java new file mode 100644 index 0000000000000..a862d077b0472 --- /dev/null +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/SupportsColumnUpdates.java @@ -0,0 +1,94 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.connector.write; + +import org.apache.spark.annotation.Experimental; +import org.apache.spark.sql.connector.expressions.NamedReference; + +/** + * A mix-in interface for {@link RowLevelOperation}. Data sources can implement this interface to + * receive a narrow row containing only the columns declared via {@link #requiredDataAttributes()} + * for updated, copied, and reinserted records, instead of the full table row. + *

+ * When {@link LogicalWriteInfo#columnUpdateSchema()} is present, updated, copied, and reinserted + * rows follow it: group-based operations receive them through + * {@link DataWriter#writeColumnUpdate(Object, Object)} or + * {@link DataWriter#writeColumnUpdate(Object)}, and operations that mix in {@link SupportsDelta} + * through {@link DeltaWriter#update} and {@link DeltaWriter#reinsert}. Inserted rows always follow + * {@link LogicalWriteInfo#schema()} and arrive through {@link DataWriter#write(Object)} and + * {@link DeltaWriter#insert}. When {@link LogicalWriteInfo#columnUpdateSchema()} is absent, every + * row follows {@link LogicalWriteInfo#schema()}, as for any other operation. Connectors must + * decide which layout to expect by whether {@link LogicalWriteInfo#columnUpdateSchema()} is + * present, not by the command. A builder that wants full-width rows for some command can check + * {@link RowLevelOperationInfo#command()} and build an operation without this mix-in. Currently + * Spark narrows only UPDATE. + *

+ * The scan builder returned by {@link #newScanBuilder} should implement + * {@link org.apache.spark.sql.connector.read.SupportsPushDownRequiredColumns}, so that the scan + * does not read columns the command neither references nor declares. Otherwise the scan reads + * every column. + * + * @since 4.4.0 + */ +@Experimental +public interface SupportsColumnUpdates extends RowLevelOperation { + /** + * Returns the data column references required to perform this row-level operation. + *

+ * When Spark narrows the write, the returned columns become + * {@link LogicalWriteInfo#columnUpdateSchema()}, in declared order. Implementations must include + * every column they want to receive: every column reported by + * {@link RowLevelOperationInfo#updatedColumns()}, plus any columns needed for row lookup or + * routing, e.g. a primary key. Columns that are not declared are absent from the rows the + * connector receives, so the connector must preserve their values itself. + *

+ * Each entry must name a top-level data column of the table. For updates on nested fields such + * as {@code SET s.c1 = -1}, the connector must declare the root struct column {@code s}. Spark + * rejects with an analysis exception an empty array, a column that does not exist, a column + * declared more than once, a nested field, a metadata column (declare those through + * {@link #requiredMetadataAttributes()} instead), and an array that misses a column reported by + * {@link RowLevelOperationInfo#updatedColumns()}. + *

+ * If this operation also mixes in {@link SupportsDelta} and represents updates as deletes and + * inserts ({@link SupportsDelta#representUpdateAsDeleteAndInsert()} returns {@code true}), + * every row-ID column ({@link SupportsDelta#rowId()}) must reach the reinserted row: a data + * row-ID column must be declared here, and a metadata row-ID column must be returned by + * {@link #requiredMetadataAttributes()} with + * {@link org.apache.spark.sql.connector.catalog.MetadataColumn#PRESERVE_ON_REINSERT} set. In + * this mode, an UPDATE that assigns a new value to a row-ID column is also rejected, unless + * every table column is declared here. Both cases are rejected with an analysis exception. + *

+ * Data columns the write needs must be declared here too, even if the command does not + * reference them, such as the source columns of partition transforms and any columns used by + * {@link RequiresDistributionAndOrdering#requiredDistribution()} or + * {@link RequiresDistributionAndOrdering#requiredOrdering()}. Such columns appear in the rows + * the connector receives; a connector that does not want to persist them should project them + * away inside the connector before writing. A write whose distribution or ordering references + * any other column is rejected with an analysis exception, except for metadata columns returned + * by {@link #requiredMetadataAttributes()} and, for operations that mix in + * {@link SupportsDelta}, row-ID columns. + *

+ * Spark does not read an undeclared column that neither the command nor a table constraint + * references, so such a column is not available for runtime filtering. A scan that reports + * runtime filter attributes based on the columns it reads loses runtime group filtering on it, + * and a scan that still reports it as a runtime filter attribute fails with + * {@code DATA_SOURCE_INVALID_RUNTIME_FILTER_ATTRIBUTE}. Declare partition source columns used + * for runtime filtering to keep it. + */ + NamedReference[] requiredDataAttributes(); +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/ProjectingInternalRow.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/ProjectingInternalRow.scala index ccca512f080d9..051a2a60aebe1 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/ProjectingInternalRow.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/ProjectingInternalRow.scala @@ -38,6 +38,14 @@ case class ProjectingInternalRow(schema: StructType, this.row = row } + /** + * Returns a projection with the same schema whose field `i` reads input ordinal + * `ordinalMap(colOrdinals(i))` instead of `colOrdinals(i)`. + */ + def remapOrdinals(ordinalMap: Int => Int): ProjectingInternalRow = { + ProjectingInternalRow(schema, colOrdinals.map(ordinalMap)) + } + override def setNullAt(i: Int): Unit = throw SparkUnsupportedOperationException() override def update(i: Int, value: Any): Unit = throw SparkUnsupportedOperationException() diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteRowLevelCommand.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteRowLevelCommand.scala index a64bdb6cbf58f..d5003ebc734cd 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteRowLevelCommand.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteRowLevelCommand.scala @@ -17,6 +17,8 @@ package org.apache.spark.sql.catalyst.analysis +import java.util.Locale + import scala.collection.mutable import org.apache.spark.SparkException @@ -29,7 +31,7 @@ import org.apache.spark.sql.catalyst.util.{GeneratedColumn, ReplaceDataProjectio import org.apache.spark.sql.catalyst.util.RowDeltaUtils._ import org.apache.spark.sql.connector.catalog.SupportsRowLevelOperations import org.apache.spark.sql.connector.expressions.FieldReference -import org.apache.spark.sql.connector.write.{RowLevelOperation, RowLevelOperationInfoImpl, RowLevelOperationTable, SupportsDelta} +import org.apache.spark.sql.connector.write.{RowLevelOperation, RowLevelOperationInfoImpl, RowLevelOperationTable, SupportsColumnUpdates, SupportsDelta} import org.apache.spark.sql.connector.write.RowLevelOperation.Command import org.apache.spark.sql.errors.QueryCompilationErrors import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation @@ -72,8 +74,10 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] { protected def buildOperationTable( table: SupportsRowLevelOperations, command: Command, - options: CaseInsensitiveStringMap): RowLevelOperationTable = { - val info = RowLevelOperationInfoImpl(command, options) + options: CaseInsensitiveStringMap, + updatedAttrs: Seq[AttributeReference] = Nil): RowLevelOperationTable = { + val updatedColumns = updatedAttrs.map(attr => FieldReference(Seq(attr.name))) + val info = RowLevelOperationInfoImpl(command, options, updatedColumns) val operation = table.newRowLevelOperationBuilder(info).build() RowLevelOperationTable(table, operation) } @@ -109,6 +113,57 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] { relation) } + /** + * Resolves the connector-declared required data attributes for column-update writes against + * the given relation. An empty declaration, nested fields, duplicate columns, columns that do + * not exist in the relation and metadata columns are rejected with an `AnalysisException`. + */ + protected def resolveRequiredDataAttrs( + relation: DataSourceV2Relation, + operation: SupportsColumnUpdates): Seq[AttributeReference] = { + val refs = operation.requiredDataAttributes + if (refs.isEmpty) { + throw QueryCompilationErrors.emptyRequiredDataAttributesError(operation.getClass.getName) + } + val nested = refs.filter(_.fieldNames.length != 1).map(_.describe()).toImmutableArraySeq + if (nested.nonEmpty) { + throw QueryCompilationErrors.nestedRequiredDataAttributeError( + operation.getClass.getName, nested) + } + val normalizedNames = refs.map { ref => + val name = ref.fieldNames.head + if (conf.caseSensitiveAnalysis) name else name.toLowerCase(Locale.ROOT) + } + val duplicates = refs.zip(normalizedNames).groupBy(_._2).collect { + case (_, occurrences) if occurrences.length > 1 => occurrences.head._1.describe() + }.toSeq + if (duplicates.nonEmpty) { + throw QueryCompilationErrors.duplicateRequiredDataAttributeError( + operation.getClass.getName, duplicates) + } + val resolvedOpts = refs.toImmutableArraySeq.map(V2ExpressionUtils.resolveRefOpt(_, relation)) + val unknown = refs.toImmutableArraySeq.zip(resolvedOpts).collect { + case (ref, None) => ref.describe() + } + if (unknown.nonEmpty) { + throw QueryCompilationErrors.unknownRequiredDataAttributeError( + operation.getClass.getName, unknown) + } + val resolved = resolvedOpts.flatten.map(_.asInstanceOf[AttributeReference]) + val metadata = refs.toImmutableArraySeq.zip(resolved).collect { + case (ref, attr) if MetadataAttribute.isValid(attr.metadata) => ref.describe() + } + if (metadata.nonEmpty) { + throw QueryCompilationErrors.metadataRequiredDataAttributeError( + operation.getClass.getName, metadata) + } + // `resolveRefOpt` matches case-insensitively but keeps the declared spelling, so map each + // resolved attribute back to the relation's own (by exprId) to avoid a spelling mismatch + // between the connector's declaration and the table's actual column name. + val byExprId = relation.output.map(a => a.exprId -> a).toMap + resolved.map(a => byExprId.getOrElse(a.exprId, a).asInstanceOf[AttributeReference]) + } + protected def resolveRowIdAttrs( relation: DataSourceV2Relation, operation: SupportsDelta): Seq[AttributeReference] = { @@ -129,6 +184,24 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] { V2ExpressionUtils.resolveRef[AttributeReference](FieldReference(name), plan) } + protected def isIdentityAssignment(key: Attribute, value: Expression): Boolean = { + val valueWithoutAlias = value match { + case Alias(child, _) => child + case other => other + } + key.semanticEquals(valueWithoutAlias) + } + + /** + * Returns the table attributes that these assignments set to a new value, skipping identity + * assignments. + */ + protected def collectUpdatedAttrs(assignments: Seq[Assignment]): Seq[AttributeReference] = { + assignments.collect { + case Assignment(key: AttributeReference, value) if !isIdentityAssignment(key, value) => key + } + } + protected def deltaDeleteOutput( rowAttrs: Seq[Attribute], rowIdAttrs: Seq[Attribute], @@ -253,6 +326,24 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] { WriteDeltaProjections(rowProjection, rowIdProjection, metadataProjection) } + /** + * Moves the query columns that no projection reads after the ones that are read, and returns + * the reordered query with a map from each old ordinal to the new one. The optimizer prunes the + * query columns after the last one the write reads, which then never changes an ordinal. + */ + protected def orderReadColumnsFirst( + plan: LogicalPlan, + projections: Seq[ProjectingInternalRow]): (LogicalPlan, Int => Int) = { + val readOrdinals = (0 +: projections.flatMap(_.colOrdinals)).distinct.sorted + val readOrdinalSet = readOrdinals.toSet + val order = readOrdinals ++ plan.output.indices.filterNot(readOrdinalSet.contains) + if (order == plan.output.indices) { + (plan, identity) + } else { + (reorderOutputs(plan, order), order.zipWithIndex.toMap) + } + } + private def extractOutputs(plan: LogicalPlan): Seq[Seq[Expression]] = { plan match { case p: Project => Seq(p.projectList) @@ -263,6 +354,17 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] { } } + private def reorderOutputs(plan: LogicalPlan, order: Seq[Int]): LogicalPlan = { + plan match { + case p: Project => p.copy(projectList = order.map(p.projectList)) + case e: Expand => + e.copy(projections = e.projections.map(output => order.map(output)), + output = order.map(e.output)) + case u: Union => u.withNewChildren(u.children.map(reorderOutputs(_, order))) + case _ => throw SparkException.internalError("Can't reorder outputs of plan: " + plan) + } + } + private def filterOutputs( outputs: Seq[Seq[Expression]], operations: Set[Int]): Seq[Seq[Expression]] = { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteUpdateTable.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteUpdateTable.scala index 352ddc0b0acb1..e379cb7e761d3 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteUpdateTable.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteUpdateTable.scala @@ -17,14 +17,15 @@ package org.apache.spark.sql.catalyst.analysis -import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, AttributeReference, EqualNullSafe, Expression, If, Literal, MetadataAttribute, Not, SubqueryExpression} +import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, AttributeReference, AttributeSet, EqualNullSafe, Expression, If, Literal, MetadataAttribute, Not, SubqueryExpression} import org.apache.spark.sql.catalyst.expressions.Literal.TrueLiteral import org.apache.spark.sql.catalyst.plans.logical.{Assignment, Expand, Filter, LogicalPlan, Project, ReplaceData, Union, UpdateTable, WriteDelta} import org.apache.spark.sql.catalyst.trees.TreePattern.UPDATE_TABLE import org.apache.spark.sql.catalyst.util.RowDeltaUtils._ import org.apache.spark.sql.connector.catalog.SupportsRowLevelOperations -import org.apache.spark.sql.connector.write.{RowLevelOperationTable, SupportsDelta} +import org.apache.spark.sql.connector.write.{RowLevelOperation, RowLevelOperationTable, SupportsColumnUpdates, SupportsDelta} import org.apache.spark.sql.connector.write.RowLevelOperation.Command.UPDATE +import org.apache.spark.sql.errors.QueryCompilationErrors import org.apache.spark.sql.execution.datasources.v2.{DataSourceV2Relation, ExtractV2Table} import org.apache.spark.sql.types.IntegerType @@ -43,15 +44,17 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { EliminateSubqueryAliases(aliasedTable) match { case r @ ExtractV2Table(tbl: SupportsRowLevelOperations) => checkNoGeneratedColumns(r, UPDATE) - val table = buildOperationTable(tbl, UPDATE, r.options) + val updatedAttrs = collectUpdatedAttrs(assignments) + val table = buildOperationTable(tbl, UPDATE, r.options, updatedAttrs) val updateCond = cond.getOrElse(TrueLiteral) + val writeAttrs = resolveWriteAttrs(r, table.operation, assignments) table.operation match { case _: SupportsDelta => - buildWriteDeltaPlan(r, table, assignments, updateCond) + buildWriteDeltaPlan(r, table, assignments, updateCond, writeAttrs) case _ if SubqueryExpression.hasSubquery(updateCond) => - buildReplaceDataWithUnionPlan(r, table, assignments, updateCond) + buildReplaceDataWithUnionPlan(r, table, assignments, updateCond, writeAttrs) case _ => - buildReplaceDataPlan(r, table, assignments, updateCond) + buildReplaceDataPlan(r, table, assignments, updateCond, writeAttrs) } case _ => @@ -65,7 +68,8 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { relation: DataSourceV2Relation, operationTable: RowLevelOperationTable, assignments: Seq[Assignment], - cond: Expression): ReplaceData = { + cond: Expression, + writeAttrs: Seq[AttributeReference]): ReplaceData = { // resolve all required metadata attrs that may be used for grouping data on write val metadataAttrs = resolveRequiredMetadataAttrs(relation, operationTable.operation) @@ -77,10 +81,12 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { val query = buildReplaceDataUpdateProjection(readRelation, assignments, cond) // build a plan to replace read groups in the table - val writeRelation = relation.copy(table = operationTable) - val projections = buildReplaceDataProjections(query, relation.output, metadataAttrs) + val writeRelation = relation.copy(table = operationTable, output = writeAttrs) + val projections = buildReplaceDataProjections(query, writeRelation.output, metadataAttrs) + val (orderedQuery, ordinalMap) = orderReadColumnsFirst(query, projections.all) val groupFilterCond = if (groupFilterEnabled) Some(cond) else None - ReplaceData(writeRelation, cond, query, relation, projections, groupFilterCond) + ReplaceData(writeRelation, cond, orderedQuery, relation, projections.remapOrdinals(ordinalMap), + groupFilterCond) } // build a rewrite plan for sources that support replacing groups of data (e.g. files, partitions) @@ -89,7 +95,8 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { relation: DataSourceV2Relation, operationTable: RowLevelOperationTable, assignments: Seq[Assignment], - cond: Expression): ReplaceData = { + cond: Expression, + writeAttrs: Seq[AttributeReference]): ReplaceData = { // resolve all required metadata attrs that may be used for grouping data on write val metadataAttrs = resolveRequiredMetadataAttrs(relation, operationTable.operation) @@ -112,10 +119,12 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { val query = Union(updatedRowsPlan, remainingRowsPlan) // build a plan to replace read groups in the table - val writeRelation = relation.copy(table = operationTable) - val projections = buildReplaceDataProjections(query, relation.output, metadataAttrs) + val writeRelation = relation.copy(table = operationTable, output = writeAttrs) + val projections = buildReplaceDataProjections(query, writeRelation.output, metadataAttrs) + val (orderedQuery, ordinalMap) = orderReadColumnsFirst(query, projections.all) val groupFilterCond = if (groupFilterEnabled) Some(cond) else None - ReplaceData(writeRelation, cond, query, relation, projections, groupFilterCond) + ReplaceData(writeRelation, cond, orderedQuery, relation, projections.remapOrdinals(ordinalMap), + groupFilterCond) } // this method assumes the assignments have been already aligned before @@ -153,12 +162,12 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { relation: DataSourceV2Relation, operationTable: RowLevelOperationTable, assignments: Seq[Assignment], - cond: Expression): WriteDelta = { + cond: Expression, + writeAttrs: Seq[AttributeReference]): WriteDelta = { val operation = operationTable.operation.asInstanceOf[SupportsDelta] // resolve all needed attrs (e.g. row ID and any required metadata attrs) - val rowAttrs = relation.output val rowIdAttrs = resolveRowIdAttrs(relation, operation) val metadataAttrs = resolveRequiredMetadataAttrs(relation, operation) @@ -174,10 +183,13 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { } // build a plan to write the row delta to the table - val writeRelation = relation.copy(table = operationTable) - val projections = buildWriteDeltaProjections(rowDeltaPlan, rowAttrs, rowIdAttrs, metadataAttrs) + val writeRelation = relation.copy(table = operationTable, output = writeAttrs) + val projections = buildWriteDeltaProjections( + rowDeltaPlan, writeRelation.output, rowIdAttrs, metadataAttrs) + val (orderedPlan, ordinalMap) = orderReadColumnsFirst(rowDeltaPlan, projections.all) val groupFilterCond = if (groupFilterEnabled) Some(cond) else None - WriteDelta(writeRelation, cond, rowDeltaPlan, relation, projections, groupFilterCond) + WriteDelta(writeRelation, cond, orderedPlan, relation, projections.remapOrdinals(ordinalMap), + groupFilterCond) } // this method assumes the assignments have been already aligned before @@ -228,4 +240,114 @@ object RewriteUpdateTable extends RewriteRowLevelCommand { val expandOutput = generateExpandOutput(attrs, outputs) Expand(outputs, expandOutput, matchedRowsPlan) } + + /** + * Returns the output of the write relation, which is the layout of the data rows the connector + * receives. If the operation opts into column updates, this is the connector's + * `requiredDataAttributes()`, validated against the assignments. Otherwise it is the full + * relation output. + */ + private def resolveWriteAttrs( + relation: DataSourceV2Relation, + operation: RowLevelOperation, + assignments: Seq[Assignment]): Seq[AttributeReference] = operation match { + case op: SupportsColumnUpdates with SupportsDelta if op.representUpdateAsDeleteAndInsert => + val attrs = resolveDeclaredAttrs(relation, op, assignments) + val rowIdAttrs = resolveRowIdAttrs(relation, op) + validateNoRowIdReassignment(op, relation, attrs, assignments, rowIdAttrs) + validateRowIdDeclared(op, attrs, rowIdAttrs, resolveRequiredMetadataAttrs(relation, op)) + attrs + case op: SupportsColumnUpdates => + resolveDeclaredAttrs(relation, op, assignments) + case _ => + relation.output + } + + /** + * Resolves the connector's `requiredDataAttributes()` and checks that it covers every assigned + * column. + */ + private def resolveDeclaredAttrs( + relation: DataSourceV2Relation, + operation: SupportsColumnUpdates, + assignments: Seq[Assignment]): Seq[AttributeReference] = { + val attrs = resolveRequiredDataAttrs(relation, operation) + validateUpdatedColumnsSubset(operation, assignments, attrs) + attrs + } + + /** + * Enforces that every column being assigned (non-identity) is present in the connector-declared + * `requiredDataAttributes()`. Comparison is at root-column granularity. + */ + private def validateUpdatedColumnsSubset( + operation: RowLevelOperation, + assignments: Seq[Assignment], + connectorDataAttrs: Seq[AttributeReference]): Unit = { + val declaredIds = connectorDataAttrs.map(_.exprId).toSet + val missing = assignments.collect { + case Assignment(key: AttributeReference, value) + if !isIdentityAssignment(key, value) && !declaredIds.contains(key.exprId) => + key.name + }.distinct + if (missing.nonEmpty) { + throw QueryCompilationErrors.requiredDataAttributesMissingUpdatedColumnsError( + operation.getClass.getName, missing) + } + } + + /** + * For connectors that opt into narrow column updates AND represent UPDATE as delete + insert, + * reject reassignment of any row-ID column, unless `requiredDataAttributes()` covers every + * column in the relation. In that case the REINSERT payload already is the full row with the + * new row-ID value, and the DELETE half still carries the original row-ID via + * `newLazyRowIdProjection`, so reassignment is safe. Otherwise the REINSERT path has no + * row-ID channel to reconstruct columns outside `requiredDataAttributes()`. + */ + private def validateNoRowIdReassignment( + operation: RowLevelOperation, + relation: DataSourceV2Relation, + connectorDataAttrs: Seq[AttributeReference], + assignments: Seq[Assignment], + rowIdAttrs: Seq[Attribute]): Unit = { + val declaredIds = connectorDataAttrs.map(_.exprId).toSet + if (relation.output.forall(a => declaredIds.contains(a.exprId))) { + return + } + val rowIdAttrSet = AttributeSet(rowIdAttrs) + val reassigned = assignments.collect { + case Assignment(key: AttributeReference, value) + if rowIdAttrSet.contains(key) && !isIdentityAssignment(key, value) => + key.name + }.distinct + if (reassigned.nonEmpty) { + throw QueryCompilationErrors.splitUpdateRowIdReassignmentError( + operation.getClass.getName, reassigned) + } + } + + /** + * For connectors that opt into narrow column updates AND represent UPDATE as delete + insert, + * requires every row-ID column to reach the REINSERT row. A data column reaches it only if it + * is declared in `requiredDataAttributes()`. A metadata column reaches it only through the + * metadata projection, so it must be in `requiredMetadataAttributes()` and preserved on + * reinsert (see `deltaReinsertOutput`/`MetadataAttribute.isPreservedOnReinsert`). Otherwise the + * reinserted row would have no identity for the connector to place it by. + */ + private def validateRowIdDeclared( + operation: RowLevelOperation, + connectorDataAttrs: Seq[AttributeReference], + rowIdAttrs: Seq[Attribute], + metadataAttrs: Seq[Attribute]): Unit = { + val declaredIds = connectorDataAttrs.map(_.exprId).toSet + val reinsertedMetadataIds = + metadataAttrs.filter(MetadataAttribute.isPreservedOnReinsert).map(_.exprId).toSet + val undeclared = rowIdAttrs + .filterNot(a => declaredIds.contains(a.exprId) || reinsertedMetadataIds.contains(a.exprId)) + .map(_.name).distinct + if (undeclared.nonEmpty) { + throw QueryCompilationErrors.splitUpdateRowIdNotDeclaredError( + operation.getClass.getName, undeclared) + } + } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala index 7ba4493c0961e..799f6e29d98d5 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala @@ -1530,6 +1530,10 @@ object ColumnPruning extends Rule[LogicalPlan] { case e @ MergeRows(_, _, _, _, _, _, _, child) if !child.outputSet.subsetOf(e.references) => e.copy(child = prunedChild(child, e.references)) + // prune unused columns from the query of ReplaceData/WriteDelta for row-level operations + case w: RowLevelWrite if !w.query.outputSet.subsetOf(w.references) => + w.withNewQuery(prunedChild(w.query, w.references)) + // prune unrequired references case p @ Project(_, g: Generate) if p.references != g.outputSet => val requiredAttrs = p.references -- g.producedAttributes ++ g.generator.references diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/v2Commands.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/v2Commands.scala index fed130d0b316a..e142debe48030 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/v2Commands.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/v2Commands.scala @@ -19,6 +19,7 @@ package org.apache.spark.sql.catalyst.plans.logical import org.apache.spark.{SparkException, SparkIllegalArgumentException, SparkUnsupportedOperationException} import org.apache.spark.sql.AnalysisException +import org.apache.spark.sql.catalyst.ProjectingInternalRow import org.apache.spark.sql.catalyst.analysis.{AnalysisContext, AssignmentUtils, EliminateSubqueryAliases, FieldName, NamedRelation, PartitionSpec, ResolvedIdentifier, ResolvedProcedure, ResolveSchemaEvolution, TypeCheckResult, UnresolvedAttribute, UnresolvedException, UnresolvedProcedure, ViewSchemaMode} import org.apache.spark.sql.catalyst.analysis.TypeCheckResult.{DataTypeMismatch, TypeCheckSuccess} import org.apache.spark.sql.catalyst.catalog.{FunctionResource, RoutineLanguage} @@ -356,6 +357,18 @@ trait RowLevelWrite extends V2WriteCommand with SupportsSubquery { def condition: Expression def originalTable: NamedRelation + /** The projections that read query columns by ordinal. */ + protected def queryProjections: Seq[ProjectingInternalRow] + + /** + * The write reads the operation column at ordinal 0 and other columns only through projections. + * The query columns after the last one it reads can be pruned without changing its ordinals. + */ + override lazy val references: AttributeSet = { + val lastReadOrdinal = queryProjections.flatMap(_.colOrdinals).foldLeft(0)(math.max) + AttributeSet(query.output.take(lastReadOrdinal + 1)) + } + protected def operationResolved: Boolean = { val attr = query.output.head attr.name == RowDeltaUtils.OPERATION_COLUMN && attr.dataType == IntegerType && !attr.nullable @@ -397,8 +410,6 @@ case class ReplaceData( override def stringArgs: Iterator[Any] = Iterator(table, query, write) - override lazy val references: AttributeSet = query.outputSet - lazy val operation: RowLevelOperation = { EliminateSubqueryAliases(table) match { case ExtractV2Table(RowLevelOperationTable(_, operation)) => @@ -439,6 +450,8 @@ case class ReplaceData( operation.command == DELETE || MetadataAttribute.isPreservedOnUpdate(attr) } + override protected def queryProjections: Seq[ProjectingInternalRow] = projections.all + override def withNewQuery(newQuery: LogicalPlan): ReplaceData = copy(query = newQuery) override def withNewTable(newTable: NamedRelation): ReplaceData = copy(table = newTable) @@ -487,8 +500,6 @@ case class WriteDelta( override def stringArgs: Iterator[Any] = Iterator(table, query, write) - override lazy val references: AttributeSet = query.outputSet - lazy val operation: SupportsDelta = { EliminateSubqueryAliases(table) match { case ExtractV2Table(RowLevelOperationTable(_, operation)) => @@ -551,6 +562,8 @@ case class WriteDelta( } } + override protected def queryProjections: Seq[ProjectingInternalRow] = projections.all + override def withNewQuery(newQuery: LogicalPlan): V2WriteCommand = copy(query = newQuery) override def withNewTable(newTable: NamedRelation): V2WriteCommand = copy(table = newTable) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/ReplaceDataProjections.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/ReplaceDataProjections.scala index 99744faf2c749..7f843e225c370 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/ReplaceDataProjections.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/ReplaceDataProjections.scala @@ -21,4 +21,13 @@ import org.apache.spark.sql.catalyst.ProjectingInternalRow case class ReplaceDataProjections( rowProjection: ProjectingInternalRow, - metadataProjection: Option[ProjectingInternalRow]) + metadataProjection: Option[ProjectingInternalRow]) { + + def all: Seq[ProjectingInternalRow] = rowProjection +: metadataProjection.toSeq + + def remapOrdinals(ordinalMap: Int => Int): ReplaceDataProjections = { + ReplaceDataProjections( + rowProjection.remapOrdinals(ordinalMap), + metadataProjection.map(_.remapOrdinals(ordinalMap))) + } +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/WriteDeltaProjections.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/WriteDeltaProjections.scala index 90f0be60c5375..0b080f4798448 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/WriteDeltaProjections.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/WriteDeltaProjections.scala @@ -22,4 +22,15 @@ import org.apache.spark.sql.catalyst.ProjectingInternalRow case class WriteDeltaProjections( rowProjection: Option[ProjectingInternalRow], rowIdProjection: ProjectingInternalRow, - metadataProjection: Option[ProjectingInternalRow]) + metadataProjection: Option[ProjectingInternalRow]) { + + def all: Seq[ProjectingInternalRow] = + rowProjection.toSeq ++ (rowIdProjection +: metadataProjection.toSeq) + + def remapOrdinals(ordinalMap: Int => Int): WriteDeltaProjections = { + WriteDeltaProjections( + rowProjection.map(_.remapOrdinals(ordinalMap)), + rowIdProjection.remapOrdinals(ordinalMap), + metadataProjection.map(_.remapOrdinals(ordinalMap))) + } +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/LogicalWriteInfoImpl.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/LogicalWriteInfoImpl.scala index 1e4e1a5955f3c..81f00413e0c5e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/LogicalWriteInfoImpl.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/LogicalWriteInfoImpl.scala @@ -29,7 +29,8 @@ private[sql] case class LogicalWriteInfoImpl( schema: StructType, options: CaseInsensitiveStringMap, override val rowIdSchema: Optional[StructType] = Optional.empty[StructType], - override val metadataSchema: Optional[StructType] = Optional.empty[StructType]) + override val metadataSchema: Optional[StructType] = Optional.empty[StructType], + override val columnUpdateSchema: Optional[StructType] = Optional.empty[StructType]) extends LogicalWriteInfo object LogicalWriteInfoImpl { @@ -38,12 +39,14 @@ object LogicalWriteInfoImpl { schema: StructType, options: CaseInsensitiveStringMap, rowIdSchema: Option[StructType], - metadataSchema: Option[StructType]): LogicalWriteInfoImpl = { + metadataSchema: Option[StructType], + columnUpdateSchema: Option[StructType]): LogicalWriteInfoImpl = { LogicalWriteInfoImpl( queryId, schema, options, rowIdSchema.toJava, - metadataSchema.toJava) + metadataSchema.toJava, + columnUpdateSchema.toJava) } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/RowLevelOperationInfoImpl.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/RowLevelOperationInfoImpl.scala index 9d499cdef361b..eca42cf365dc2 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/RowLevelOperationInfoImpl.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/connector/write/RowLevelOperationInfoImpl.scala @@ -17,9 +17,15 @@ package org.apache.spark.sql.connector.write +import org.apache.spark.sql.connector.expressions.NamedReference import org.apache.spark.sql.connector.write.RowLevelOperation.Command import org.apache.spark.sql.util.CaseInsensitiveStringMap private[sql] case class RowLevelOperationInfoImpl( command: Command, - options: CaseInsensitiveStringMap) extends RowLevelOperationInfo + options: CaseInsensitiveStringMap, + private val updatedCols: Seq[NamedReference]) + extends RowLevelOperationInfo { + + override def updatedColumns(): Array[NamedReference] = updatedCols.toArray +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala index 2cc4fa51132f8..54c74f595caa7 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala @@ -4638,6 +4638,92 @@ private[sql] object QueryCompilationErrors extends QueryErrorsBase with Compilat messageParameters = Map("nullableRowIdAttrs" -> nullableRowIdAttrs.mkString(", "))) } + def emptyRequiredDataAttributesError(connectorClass: String): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_EMPTY_REQUIRED_DATA_ATTRIBUTES", + messageParameters = Map("connector" -> connectorClass)) + } + + def duplicateRequiredDataAttributeError( + connectorClass: String, + duplicateAttributes: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_DUPLICATE_REQUIRED_DATA_ATTRIBUTE", + messageParameters = Map( + "connector" -> connectorClass, + "duplicateAttributes" -> duplicateAttributes.mkString("[", ", ", "]"))) + } + + def metadataRequiredDataAttributeError( + connectorClass: String, + metadataAttributes: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_METADATA_REQUIRED_DATA_ATTRIBUTE", + messageParameters = Map( + "connector" -> connectorClass, + "metadataAttributes" -> metadataAttributes.mkString("[", ", ", "]"))) + } + + def nestedRequiredDataAttributeError( + connectorClass: String, + nestedAttributes: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_NESTED_REQUIRED_DATA_ATTRIBUTE", + messageParameters = Map( + "connector" -> connectorClass, + "nestedAttributes" -> nestedAttributes.mkString("[", ", ", "]"))) + } + + def requiredDataAttributesMissingUpdatedColumnsError( + connectorClass: String, + missingColumns: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_REQUIRED_DATA_ATTRIBUTES_MISSING_UPDATED_COLUMNS", + messageParameters = Map( + "connector" -> connectorClass, + "missingColumns" -> missingColumns.mkString("[", ", ", "]"))) + } + + def splitUpdateRowIdNotDeclaredError( + connectorClass: String, + rowIds: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_SPLIT_ROW_ID_NOT_DECLARED", + messageParameters = Map( + "connector" -> connectorClass, + "rowIds" -> rowIds.mkString("[", ", ", "]"))) + } + + def splitUpdateRowIdReassignmentError( + connectorClass: String, + rowIds: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_SPLIT_ROW_ID_REASSIGNMENT", + messageParameters = Map( + "connector" -> connectorClass, + "rowIds" -> rowIds.mkString("[", ", ", "]"))) + } + + def unknownRequiredDataAttributeError( + connectorClass: String, + unknownAttributes: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_UNKNOWN_REQUIRED_DATA_ATTRIBUTE", + messageParameters = Map( + "connector" -> connectorClass, + "unknownAttributes" -> unknownAttributes.mkString("[", ", ", "]"))) + } + + def undeclaredWriteRequirementColumnsError( + connectorClass: String, + columns: Seq[String]): Throwable = { + new AnalysisException( + errorClass = "COLUMN_UPDATE_UNDECLARED_WRITE_REQUIREMENT_COLUMNS", + messageParameters = Map( + "connector" -> connectorClass, + "columns" -> columns.mkString("[", ", ", "]"))) + } + def cannotRenameTableAcrossSchemaError(): Throwable = { new SparkUnsupportedOperationException( errorClass = "CANNOT_RENAME_ACROSS_SCHEMA", messageParameters = Map("type" -> "table") diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RowLevelWriteColumnPruningSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RowLevelWriteColumnPruningSuite.scala new file mode 100644 index 0000000000000..808d802f78d6f --- /dev/null +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/optimizer/RowLevelWriteColumnPruningSuite.scala @@ -0,0 +1,263 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.catalyst.optimizer + +import java.util + +import org.apache.spark.sql.catalyst.ProjectingInternalRow +import org.apache.spark.sql.catalyst.dsl.expressions._ +import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, AttributeReference, EqualNullSafe, Expression, If, Literal, Not} +import org.apache.spark.sql.catalyst.plans.{LeftOuter, PlanTest} +import org.apache.spark.sql.catalyst.plans.logical._ +import org.apache.spark.sql.catalyst.plans.logical.MergeRows.{Copy, Keep, Update} +import org.apache.spark.sql.catalyst.rules.RuleExecutor +import org.apache.spark.sql.catalyst.util.{ReplaceDataProjections, WriteDeltaProjections} +import org.apache.spark.sql.catalyst.util.RowDeltaUtils._ +import org.apache.spark.sql.connector.catalog.{Column, Table, TableCapability} +import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation +import org.apache.spark.sql.types.{IntegerType, StructField, StructType} +import org.apache.spark.sql.util.CaseInsensitiveStringMap + +/** + * Tests the [[ColumnPruning]] case that prunes the query columns of a [[RowLevelWrite]] after the + * last one its projections read. + */ +class RowLevelWriteColumnPruningSuite extends PlanTest { + + object Optimize extends RuleExecutor[LogicalPlan] { + val batches = Batch("Column Pruning", FixedPoint(10), + ColumnPruning, CollapseProject, RemoveNoopOperators) :: Nil + } + + // t(pk, id, dep, salary, bonus, extra) with the metadata column _partition + private val read = LocalRelation( + $"pk".int.notNull, $"id".int, $"dep".string, $"salary".int, $"bonus".int, $"extra".int, + $"_partition".string.notNull) + + // the write and original table relations; the pruning never inspects them + private val tableRelation: DataSourceV2Relation = { + val table = new Table { + override def name(): String = "t" + override def columns(): Array[Column] = Array(Column.create("pk", IntegerType)) + override def capabilities(): util.Set[TableCapability] = util.Set.of[TableCapability]() + } + DataSourceV2Relation.create(table, None, None, CaseInsensitiveStringMap.empty()) + } + + private val dataColumns = Seq("pk", "id", "dep", "salary", "bonus", "extra") + + // the data columns of a write that receives pk and salary, with the ones it reads first and + // _partition after them, as RewriteUpdateTable orders them + private val writeColumns = Seq("pk", "salary") + private val otherColumns = Seq("id", "dep", "bonus", "extra") + + private def column(plan: LogicalPlan, name: String): Attribute = { + plan.output.find(_.name == name).get + } + + private def cond(plan: LogicalPlan): Expression = column(plan, "dep") === "hr" + + // the value of every data column after `SET salary = salary + bonus` + private def updatedValue(plan: LogicalPlan, name: String): Expression = name match { + case "salary" => column(plan, "salary") + column(plan, "bonus") + case _ => column(plan, name) + } + + private def projection(query: LogicalPlan, names: String*): ProjectingInternalRow = { + val ordinals = names.map(name => query.output.indexWhere(_.name == name)).toIndexedSeq + val fields = ordinals.map { i => + val attr = query.output(i) + StructField(attr.name, attr.dataType, attr.nullable) + } + ProjectingInternalRow(StructType(fields), ordinals) + } + + // the copy-on-write UPDATE query built by RewriteUpdateTable.buildReplaceDataPlan + private def replaceDataQuery(plan: LogicalPlan, names: Seq[String]): LogicalPlan = { + val operation = If(cond(plan), Literal(UPDATE_OPERATION), Literal(COPY_OPERATION)) + val values = names.map { name => + if (name == "_partition") { + column(plan, name) + } else { + Alias(If(cond(plan), updatedValue(plan, name), column(plan, name)), name)() + } + } + Project(Alias(operation, OPERATION_COLUMN)() +: values, plan) + } + + private def readFirstQuery(plan: LogicalPlan): LogicalPlan = { + replaceDataQuery(plan, (writeColumns :+ "_partition") ++ otherColumns) + } + + private def replaceData(query: LogicalPlan, rowNames: Seq[String]): ReplaceData = { + val projections = ReplaceDataProjections( + projection(query, rowNames: _*), + Some(projection(query, "_partition"))) + ReplaceData(tableRelation, cond(read), query, tableRelation, projections) + } + + private def writeDelta(query: LogicalPlan): WriteDelta = { + val projections = WriteDeltaProjections( + Some(projection(query, writeColumns: _*)), + projection(query, "pk"), + Some(projection(query, "_partition"))) + WriteDelta(tableRelation, cond(read), query, tableRelation, projections) + } + + private def scannedColumns(plan: LogicalPlan, relation: LogicalPlan): Set[String] = { + plan.collect { case node => node.references }.flatten + .filter(relation.outputSet.contains).map(_.name).toSet + } + + private def checkPruned(optimized: RowLevelWrite, write: RowLevelWrite): Unit = { + assert(optimized.query.output.map(_.name) == + OPERATION_COLUMN +: writeColumns :+ "_partition") + assert(optimized.query.output == write.query.output.take(4)) + comparePlans(Optimize.execute(optimized), optimized, checkAnalysis = false) + } + + test("ReplaceData: prune the columns after the last one the write reads") { + val write = replaceData(readFirstQuery(read), writeColumns) + val optimized = Optimize.execute(write).asInstanceOf[ReplaceData] + + checkPruned(optimized, write) + assert(optimized.projections == write.projections) + assert(scannedColumns(optimized, read) == Set("pk", "dep", "salary", "bonus", "_partition")) + } + + test("ReplaceData: keep unread columns before the last one the write reads") { + val write = replaceData(replaceDataQuery(read, dataColumns :+ "_partition"), writeColumns) + comparePlans(Optimize.execute(write), write, checkAnalysis = false) + } + + test("ReplaceData: a CHECK filter on the query keeps its columns") { + val updated = readFirstQuery(read) + val write = replaceData(Filter(column(updated, "extra") > 0, updated), writeColumns) + val optimized = Optimize.execute(write).asInstanceOf[ReplaceData] + + checkPruned(optimized, write) + assert(optimized.query.exists(_.isInstanceOf[Filter])) + assert(optimized.projections == write.projections) + assert(scannedColumns(optimized, read) == + Set("pk", "dep", "salary", "bonus", "extra", "_partition")) + } + + test("ReplaceData: prune both children of the Union built for subquery conditions") { + val names = (writeColumns :+ "_partition") ++ otherColumns + val matched = Filter(cond(read), read) + val updated = Project( + Alias(Literal(UPDATE_OPERATION), OPERATION_COLUMN)() +: names.map { name => + if (name == "_partition") column(matched, name) + else Alias(updatedValue(matched, name), name)() + }, + matched) + val read2 = read.newInstance() + val remaining = Filter(Not(EqualNullSafe(cond(read2), Literal.TrueLiteral)), read2) + val copied = Project( + Alias(Literal(COPY_OPERATION), OPERATION_COLUMN)() +: names.map(column(remaining, _)), + remaining) + val write = replaceData(Union(updated, copied), writeColumns) + val optimized = Optimize.execute(write).asInstanceOf[ReplaceData] + + checkPruned(optimized, write) + val union = optimized.query.collectFirst { case u: Union => u }.get + union.children.foreach { child => + assert(child.output.map(_.name) == OPERATION_COLUMN +: writeColumns :+ "_partition") + } + assert(optimized.projections == write.projections) + assert(scannedColumns(optimized, read) == Set("pk", "dep", "salary", "bonus", "_partition")) + assert(scannedColumns(optimized, read2) == Set("pk", "dep", "salary", "_partition")) + } + + test("WriteDelta: prune the columns after the last one the write reads") { + val matched = Filter(cond(read), read) + val names = (writeColumns :+ "_partition") ++ otherColumns + val query = Project( + Alias(Literal(UPDATE_OPERATION), OPERATION_COLUMN)() +: names.map { name => + if (name == "_partition") column(matched, name) + else Alias(updatedValue(matched, name), name)() + }, + matched) + val write = writeDelta(query) + val optimized = Optimize.execute(write).asInstanceOf[WriteDelta] + + checkPruned(optimized, write) + assert(optimized.projections == write.projections) + assert(scannedColumns(optimized, read) == Set("pk", "dep", "salary", "bonus", "_partition")) + } + + test("WriteDelta: prune the Expand built for delete and reinsert") { + val matched = Filter(cond(read), read) + val names = (writeColumns :+ "_partition") ++ otherColumns + val deleteOutput = Literal(DELETE_OPERATION) +: names.map { name => + if (name == "pk" || name == "_partition") column(matched, name) + else Literal(null, column(matched, name).dataType) + } + val reinsertOutput = Literal(REINSERT_OPERATION) +: names.map(updatedValue(matched, _)) + val expandOutput = AttributeReference(OPERATION_COLUMN, IntegerType, nullable = false)() +: + names.map(column(matched, _).newInstance()) + val query = Expand(Seq(deleteOutput, reinsertOutput), expandOutput, matched) + val write = writeDelta(query) + val optimized = Optimize.execute(write).asInstanceOf[WriteDelta] + + checkPruned(optimized, write) + val expand = optimized.query.collectFirst { case e: Expand => e }.get + assert(expand.output.map(_.name) == OPERATION_COLUMN +: writeColumns :+ "_partition") + assert(expand.projections.forall(_.size == 4)) + assert(optimized.projections == write.projections) + assert(scannedColumns(optimized, read) == Set("pk", "dep", "salary", "bonus", "_partition")) + } + + test("no change when the projections read every query column") { + // DELETE and an UPDATE that is not a column update read every data column + val update = replaceData(replaceDataQuery(read, dataColumns :+ "_partition"), dataColumns) + comparePlans(Optimize.execute(update), update, checkAnalysis = false) + + val delete = replaceData( + Project(Alias(Literal(COPY_OPERATION), OPERATION_COLUMN)() +: read.output, + Filter(Not(EqualNullSafe(cond(read), Literal.TrueLiteral)), read)), + dataColumns) + comparePlans(Optimize.execute(delete), delete, checkAnalysis = false) + + // MERGE reads every target column through MergeRows + val source = LocalRelation($"spk".int, $"ssalary".int, $"present".boolean) + val Seq(spk, ssalary, present) = source.output + val joined = Join(read, source, LeftOuter, Some(column(read, "pk") === spk), JoinHint.NONE) + def target(name: String): Expression = column(read, name) + val updateOutput = Literal(UPDATE_OPERATION) +: + dataColumns.map(name => if (name == "salary") ssalary else target(name)) :+ + target("_partition") + val copyOutput = Literal(COPY_OPERATION) +: dataColumns.map(target) :+ target("_partition") + val mergeOutput = AttributeReference(OPERATION_COLUMN, IntegerType, nullable = false)() +: + read.output.map(_.newInstance()) + val mergeRows = MergeRows( + isSourceRowPresent = present, + isTargetRowPresent = Literal.TrueLiteral, + matchedInstructions = Seq(Keep(Update, Literal.TrueLiteral, updateOutput)), + notMatchedInstructions = Nil, + notMatchedBySourceInstructions = Seq(Keep(Copy, Literal.TrueLiteral, copyOutput)), + checkCardinality = false, + output = mergeOutput, + child = joined) + val merge = replaceData(mergeRows, dataColumns) + val optimizedMerge = Optimize.execute(merge).asInstanceOf[ReplaceData] + assert(optimizedMerge.query.isInstanceOf[MergeRows]) + assert(optimizedMerge.query.output == mergeRows.output) + assert(optimizedMerge.projections == merge.projections) + } +} diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala index 091a1a441f3bc..1e07435514260 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryBaseTable.scala @@ -1466,6 +1466,7 @@ private class BufferedRowsWriterFactory(schema: StructType) private class BufferWriter(schema: StructType) extends DataWriter[InternalRow] { private final val WRITE = UTF8String.fromString(Write.toString) + private final val WRITE_COLUMN_UPDATE = UTF8String.fromString(WriteColumnUpdate.toString) protected val buffer = new BufferedRows(Seq.empty, schema) @@ -1481,6 +1482,21 @@ private class BufferWriter(schema: StructType) extends DataWriter[InternalRow] { buffer.log.append(logEntry) } + // Tag UPDATE/COPY rows distinctly so tests can verify that Spark dispatches + // through writeColumnUpdate (the SupportsColumnUpdates path) rather than write. + override def writeColumnUpdate(metadata: InternalRow, row: InternalRow): Unit = { + buffer.rows.append(row.copy()) + val logEntry = new GenericInternalRow( + Array[Any](WRITE_COLUMN_UPDATE, null, metadata.copy(), row.copy())) + buffer.log.append(logEntry) + } + + override def writeColumnUpdate(row: InternalRow): Unit = { + buffer.rows.append(row.copy()) + val logEntry = new GenericInternalRow(Array[Any](WRITE_COLUMN_UPDATE, null, null, row.copy())) + buffer.log.append(logEntry) + } + override def commit(): WriterCommitMessage = buffer override def abort(): Unit = {} @@ -1524,6 +1540,7 @@ case class Commit(id: Long, writeSummary: Option[WriteSummary] = None) sealed trait Operation case object Write extends Operation +case object WriteColumnUpdate extends Operation case object Delete extends Operation case object Update extends Operation case object Reinsert extends Operation diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryRowLevelOperationTable.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryRowLevelOperationTable.scala index 3e04ab719b7bf..6b0691b971568 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryRowLevelOperationTable.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/InMemoryRowLevelOperationTable.scala @@ -18,15 +18,17 @@ package org.apache.spark.sql.connector.catalog import java.util +import java.util.concurrent.atomic.AtomicReference import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.catalyst.util.ResolveDefaultColumns import org.apache.spark.sql.connector.catalog.constraints.Constraint import org.apache.spark.sql.connector.distributions.{Distribution, Distributions} import org.apache.spark.sql.connector.expressions.{FieldReference, LogicalExpressions, NamedReference, SortDirection, SortOrder, Transform} import org.apache.spark.sql.connector.expressions.filter.Predicate import org.apache.spark.sql.connector.read.{InputPartition, Scan, ScanBuilder} -import org.apache.spark.sql.connector.write.{BatchWrite, DeltaBatchWrite, DeltaWrite, DeltaWriteBuilder, DeltaWriter, DeltaWriterFactory, LogicalWriteInfo, PhysicalWriteInfo, RequiresDistributionAndOrdering, RowLevelOperation, RowLevelOperationBuilder, RowLevelOperationInfo, SupportsDelta, Write, WriteBuilder, WriterCommitMessage} +import org.apache.spark.sql.connector.write.{BatchWrite, DataWriter, DataWriterFactory, DeltaBatchWrite, DeltaWrite, DeltaWriteBuilder, DeltaWriter, DeltaWriterFactory, LogicalWriteInfo, PhysicalWriteInfo, RequiresDistributionAndOrdering, RowLevelOperation, RowLevelOperationBuilder, RowLevelOperationInfo, SupportsColumnUpdates, SupportsDelta, Write, WriteBuilder, WriterCommitMessage} import org.apache.spark.sql.connector.write.RowLevelOperation.Command import org.apache.spark.sql.internal.connector.SchemaAlignmentConfig import org.apache.spark.sql.types.StructType @@ -87,11 +89,36 @@ class InMemoryRowLevelOperationTable private ( private final val noMetadata = properties.getOrDefault(NO_METADATA, "false") == "true" private final val useCatalystRuntimeFiltering = properties.getOrDefault(USE_CATALYST_RUNTIME_FILTERING, "false") == "true" + private final val COLUMN_UPDATE = "column-update" + private final val COLUMN_UPDATE_REQ_ATTRS = "column-update-req-attrs" + private final val COLUMN_UPDATE_COW = "column-update-cow" + private final val COLUMN_UPDATE_COW_REQ_ATTRS = "column-update-cow-req-attrs" + private final val COLUMN_UPDATE_COW_NO_WRITE_COLUMN_UPDATE = + "column-update-cow-no-write-column-update" + // makes the copy-on-write column-update writer override only writeColumnUpdate(record) + private final val COLUMN_UPDATE_COW_RECORD_ONLY_WRITER = "column-update-cow-record-only-writer" + // makes the copy-on-write column-update scan report every identity partition column from + // filterAttributes(), even one it doesn't read + private final val COLUMN_UPDATE_COW_UNREAD_FILTER_ATTRS = "column-update-cow-unread-filter-attrs" + private final val COLUMN_UPDATE_SPLIT = "column-update-split" + private final val COLUMN_UPDATE_SPLIT_REQ_ATTRS = "column-update-split-req-attrs" + private final val COLUMN_UPDATE_EMPTY_REQ_ATTRS = "column-update-empty-req-attrs" + private final val COLUMN_UPDATE_SPLIT_ROW_ID = "column-update-split-row-id" + private final val COLUMN_UPDATE_CLUSTER_BY = "column-update-cluster-by" // used in row-level operation tests to verify replaced partitions var replacedPartitions: Seq[Seq[Any]] = Seq.empty // used in row-level operation tests to verify reported write schema var lastWriteInfo: LogicalWriteInfo = _ + // used in column-update tests to verify that Spark passed the correct updated column list + // to the connector via RowLevelOperationInfo.updatedColumns() + var lastUpdatedColumns: Array[NamedReference] = Array.empty + // used in scan pruning tests to verify the schema Spark asked the connector to read. + // Routed through the companion object (see InMemoryRowLevelOperationTable.lastScanSchema) + // because Spark's planner and the test harness can hold references to different table + // instances for the same identifier -- a per-instance field would only be visible to one. + def lastScanSchema: StructType = + InMemoryRowLevelOperationTable.lastScanSchemaRef.get() // used in row-level operation tests to verify passed records // (operation, id, metadata, row) var lastWriteLog: Seq[InternalRow] = Seq.empty @@ -128,18 +155,84 @@ class InMemoryRowLevelOperationTable private ( copied.replacedPartitions = replacedPartitions copied.lastWriteInfo = lastWriteInfo copied.lastWriteLog = lastWriteLog + copied.lastUpdatedColumns = lastUpdatedColumns copied } override def newRowLevelOperationBuilder( info: RowLevelOperationInfo): RowLevelOperationBuilder = { - if (properties.getOrDefault(SUPPORTS_DELTAS, "false") == "true") { + lastUpdatedColumns = info.updatedColumns() + if (properties.getOrDefault(COLUMN_UPDATE, "false") == "true") { + () => new DeltaBasedColumnUpdateOperation( + info.command, info.updatedColumns().toSeq, info.options) + } else if (properties.containsKey(COLUMN_UPDATE_REQ_ATTRS)) { + val reqCols = properties.get(COLUMN_UPDATE_REQ_ATTRS).split(",").map(_.trim) + () => new DeltaBasedColumnUpdateOperationWithReqAttrs(info.command, reqCols, info.options) + } else if (properties.getOrDefault(COLUMN_UPDATE_EMPTY_REQ_ATTRS, "false") == "true") { + // Test-only: returns an empty requiredDataAttributes() so we can verify Spark rejects it. + () => new DeltaBasedColumnUpdateOperationWithReqAttrs( + info.command, Array.empty, info.options) + } else if (properties.getOrDefault(COLUMN_UPDATE_COW, "false") == "true") { + () => new PartitionBasedColumnUpdateOperation( + info.command, info.updatedColumns().toSeq, info.options) + } else if (properties.containsKey(COLUMN_UPDATE_COW_REQ_ATTRS)) { + val reqCols = properties.get(COLUMN_UPDATE_COW_REQ_ATTRS).split(",").map(_.trim) + () => new PartitionBasedColumnUpdateOperationWithReqAttrs( + info.command, reqCols, info.options) + } else if ( + properties.getOrDefault(COLUMN_UPDATE_COW_NO_WRITE_COLUMN_UPDATE, "false") == "true") { + () => new PartitionBasedColumnUpdateOperationNoWriteColumnUpdate( + info.command, info.updatedColumns().toSeq, info.options) + } else if (properties.getOrDefault(COLUMN_UPDATE_SPLIT, "false") == "true") { + () => new DeltaBasedColumnUpdateSplitOperation( + info.command, info.updatedColumns().toSeq, info.options) + } else if (properties.containsKey(COLUMN_UPDATE_SPLIT_REQ_ATTRS)) { + val reqCols = properties.get(COLUMN_UPDATE_SPLIT_REQ_ATTRS).split(",").map(_.trim) + () => new DeltaBasedColumnUpdateSplitOperationWithReqAttrs( + info.command, reqCols, info.options) + } else if (properties.containsKey(COLUMN_UPDATE_SPLIT_ROW_ID)) { + val rowIdCol = properties.get(COLUMN_UPDATE_SPLIT_ROW_ID) + () => new DeltaBasedColumnUpdateSplitOperationWithRowId( + info.command, rowIdCol, info.updatedColumns().toSeq, info.options) + } else if (properties.getOrDefault(SUPPORTS_DELTAS, "false") == "true") { () => DeltaBasedOperation(info.command, info.options) } else { () => PartitionBasedOperation(info.command, info.options) } } + private def currentRowByPk(pk: Int): Option[InternalRow] = { + dataMap.values.iterator.flatten.flatMap(_.rows) + .find(r => r.getInt(schema.fieldIndex("pk")) == pk) + } + + // Rebuilds a full table row for a narrow column-update row: the base row's values (or the + // existence default for a column added after the base row was written), overlaid with the + // narrow row's columns. + private def overlayNarrowRow( + baseRow: Option[InternalRow], + narrowRow: InternalRow, + narrowFieldIdx: Map[String, Int]): InternalRow = { + val fullRow = new GenericInternalRow(schema.length) + baseRow.foreach { base => + for (i <- schema.fields.indices) { + val field = schema.fields(i) + val value = if (i < base.numFields) { + base.get(i, field.dataType) + } else { + ResolveDefaultColumns.getExistenceDefaultValue(field) + } + fullRow.update(i, value) + } + } + schema.fields.zipWithIndex.foreach { case (field, i) => + narrowFieldIdx.get(field.name).foreach { j => + fullRow.update(i, narrowRow.get(j, field.dataType)) + } + } + fullRow + } + case class PartitionBasedOperation(command: Command, options: CaseInsensitiveStringMap) extends RowLevelOperation with RowLevelOperationWithOptions { var configuredScan: BatchScanBaseClass = _ @@ -154,6 +247,7 @@ class InMemoryRowLevelOperationTable private ( override def newScanBuilder(options: CaseInsensitiveStringMap): ScanBuilder = { newRowLevelScanBuilder(options) { scan => + InMemoryRowLevelOperationTable.recordLastScanSchema(scan.readSchema()) configuredScan = scan } } @@ -221,7 +315,9 @@ class InMemoryRowLevelOperationTable private ( override def rowId(): Array[NamedReference] = Array(PK_COLUMN_REF) override def newScanBuilder(options: CaseInsensitiveStringMap): ScanBuilder = { - newRowLevelScanBuilder(options)(_ => ()) + newRowLevelScanBuilder(options) { scan => + InMemoryRowLevelOperationTable.recordLastScanSchema(scan.readSchema()) + } } override def newWriteBuilder(info: LogicalWriteInfo): DeltaWriteBuilder = { @@ -258,6 +354,407 @@ class InMemoryRowLevelOperationTable private ( } } + // A delta-based operation that supports column-level updates: Spark sends only the declared + // columns in the row projection instead of the full row schema. It declares `pk` (the + // row-lookup key) and `dep` (an unconditionally declared base column) plus whatever columns + // Spark reports as being assigned via `RowLevelOperationInfo#updatedColumns()`. + class DeltaBasedColumnUpdateOperation( + command: Command, + updatedCols: Seq[NamedReference] = Nil, + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends DeltaBasedOperation(command, options) + with SupportsColumnUpdates { + override def representUpdateAsDeleteAndInsert(): Boolean = false + override def requiredDataAttributes(): Array[NamedReference] = { + val base = Seq(FieldReference("pk"), FieldReference("dep")) + val baseNames = base.map(_.describe()).toSet + (base ++ updatedCols.filterNot(r => baseNames.contains(r.describe()))).toArray + } + + private val clusterRef: NamedReference = Option(properties.get(COLUMN_UPDATE_CLUSTER_BY)) + .map(FieldReference(_)) + .getOrElse(PARTITION_COLUMN_REF) + + override def newWriteBuilder(info: LogicalWriteInfo): DeltaWriteBuilder = { + lastWriteInfo = info + // Capture info into a local val so nested writer/commit closures see a stable schema + // even if a subsequent newWriteBuilder call mutates lastWriteInfo. + val capturedInfo = info + val capturedWriteSchema = if (capturedInfo.columnUpdateSchema().isPresent) { + capturedInfo.columnUpdateSchema().get() + } else { + capturedInfo.schema() + } + new DeltaWriteBuilder { + override def build(): DeltaWrite = + new DeltaWrite with RequiresDistributionAndOrdering { + + override def requiredDistribution(): Distribution = { + Distributions.clustered(Array(clusterRef)) + } + + override def requiredOrdering(): Array[SortOrder] = { + Array[SortOrder]( + LogicalExpressions.sort( + clusterRef, + SortDirection.ASCENDING, + SortDirection.ASCENDING.defaultNullOrdering()) + ) + } + + override def toBatch: DeltaBatchWrite = + new TestBatchWrite with DeltaBatchWrite { + override def createBatchWriterFactory( + info: PhysicalWriteInfo): DeltaWriterFactory = { + new DeltaBufferedRowsWriterFactory(capturedWriteSchema) + } + + // For column-update writes, updated rows contain only the declared columns + // (columnUpdateSchema from LogicalWriteInfo). We expand each row to the full + // table schema by overlaying them on the base row found by pk. + override protected def doCommit(messages: Array[WriterCommitMessage]): Unit = + dataMap.synchronized { + val newData = messages.map(_.asInstanceOf[BufferedRows]) + val writeSchema = capturedWriteSchema + val writeFieldIdx = writeSchema.fieldNames.zipWithIndex.toMap + + val mergedData = newData.map { buf => + val merged = new BufferedRows(buf.key, schema) + val updateOpName = UTF8String.fromString(Update.toString) + val insertOpName = UTF8String.fromString(Insert.toString) + buf.log.foreach { logRow => + val opName = logRow.getUTF8String(0) + if (opName == updateOpName) { + val pk = logRow.getInt(1) + val narrowRow = logRow.get(3, writeSchema).asInstanceOf[InternalRow] + val baseRow = currentRowByPk(pk) + val fullRow = overlayNarrowRow(baseRow, narrowRow, writeFieldIdx) + merged.rows.append(fullRow) + } else if (opName == insertOpName) { + // INSERT rows arrive with the full table schema via writer.insert() + val insertRow = logRow.get(3, schema).asInstanceOf[InternalRow] + merged.rows.append(insertRow.copy()) + } + } + merged + } + + withDeletes(newData) + withData(mergedData) + lastWriteLog = newData.flatMap(buffer => buffer.log).toIndexedSeq + } + + override def abort(messages: Array[WriterCommitMessage]): Unit = {} + } + } + } + } + } + + class DeltaBasedColumnUpdateOperationWithReqAttrs( + command: Command, + reqCols: Array[String], + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends DeltaBasedColumnUpdateOperation(command, options = options) { + override def requiredDataAttributes(): Array[NamedReference] = reqCols.map(FieldReference(_)) + } + + class DeltaBasedColumnUpdateSplitOperation( + command: Command, + updatedCols: Seq[NamedReference] = Nil, + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends DeltaBasedColumnUpdateOperation(command, updatedCols, options) { + override def representUpdateAsDeleteAndInsert(): Boolean = true + + override def newWriteBuilder(info: LogicalWriteInfo): DeltaWriteBuilder = { + lastWriteInfo = info + // Capture info into a local val so nested writer/commit closures see a stable schema + // even if a subsequent newWriteBuilder call mutates lastWriteInfo. + val capturedInfo = info + val capturedWriteSchema = if (capturedInfo.columnUpdateSchema().isPresent) { + capturedInfo.columnUpdateSchema().get() + } else { + capturedInfo.schema() + } + new DeltaWriteBuilder { + override def build(): DeltaWrite = + new DeltaWrite with RequiresDistributionAndOrdering { + override def requiredDistribution(): Distribution = + Distributions.clustered(Array(PARTITION_COLUMN_REF)) + override def requiredOrdering(): Array[SortOrder] = Array[SortOrder]( + LogicalExpressions.sort( + PARTITION_COLUMN_REF, + SortDirection.ASCENDING, + SortDirection.ASCENDING.defaultNullOrdering())) + override def toBatch: DeltaBatchWrite = + new TestBatchWrite with DeltaBatchWrite { + override def createBatchWriterFactory( + info: PhysicalWriteInfo): DeltaWriterFactory = { + new DeltaBufferedRowsWriterFactory(capturedWriteSchema) + } + + // For delete+reinsert with narrow writes, the REINSERT row has only the + // connector-declared columns (requiredDataAttributes order). The base row is + // looked up by `pk`, so `pk` must be declared for the write to commit. + // Reconstruct the full row by overlaying the narrow row onto the original. + override protected def doCommit(messages: Array[WriterCommitMessage]): Unit = + dataMap.synchronized { + val newData = messages.map(_.asInstanceOf[BufferedRows]) + val writeSchema = capturedWriteSchema + val writeFieldIdx = writeSchema.fieldNames.zipWithIndex.toMap + val reinsertOpName = UTF8String.fromString(Reinsert.toString) + val insertOpName = UTF8String.fromString(Insert.toString) + val pkIdx = writeFieldIdx("pk") + + val expandedData = newData.map { buf => + val expanded = new BufferedRows(buf.key, schema) + buf.log.foreach { logRow => + val opName = logRow.getUTF8String(0) + if (opName == reinsertOpName) { + val narrowRow = logRow.get(3, writeSchema).asInstanceOf[InternalRow] + val pk = narrowRow.getInt(pkIdx) + val baseRow = currentRowByPk(pk) + val fullRow = overlayNarrowRow(baseRow, narrowRow, writeFieldIdx) + expanded.rows.append(fullRow) + } else if (opName == insertOpName) { + // INSERT rows arrive with the full table schema via writer.insert() + val insertRow = logRow.get(3, schema).asInstanceOf[InternalRow] + expanded.rows.append(insertRow.copy()) + } + } + expanded + } + + withDeletes(newData) + withData(expandedData) + lastWriteLog = newData.flatMap(buffer => buffer.log).toIndexedSeq + } + + override def abort(messages: Array[WriterCommitMessage]): Unit = {} + } + } + } + } + } + + // Test-only: a split-update fixture with a fully explicit requiredDataAttributes(), to + // exercise row-ID reassignment when the declaration covers every table column. + class DeltaBasedColumnUpdateSplitOperationWithReqAttrs( + command: Command, + reqCols: Array[String], + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends DeltaBasedColumnUpdateSplitOperation(command, options = options) { + override def requiredDataAttributes(): Array[NamedReference] = reqCols.map(FieldReference(_)) + } + + + // Test-only: a split-update fixture whose row ID is the given column and whose + // requiredDataAttributes() lists only the updated columns, to exercise which undeclared + // row-ID columns still reach the REINSERT row. + class DeltaBasedColumnUpdateSplitOperationWithRowId( + command: Command, + rowIdCol: String, + updatedCols: Seq[NamedReference] = Nil, + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends DeltaBasedColumnUpdateSplitOperation(command, updatedCols, options) { + override def rowId(): Array[NamedReference] = Array(FieldReference(rowIdCol)) + override def requiredDataAttributes(): Array[NamedReference] = updatedCols.toArray + } + + class PartitionBasedColumnUpdateOperation( + command: Command, + updatedCols: Seq[NamedReference] = Nil, + override val options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends RowLevelOperation with SupportsColumnUpdates with RowLevelOperationWithOptions { + var configuredScan: BatchScanBaseClass = _ + + override def command(): Command = command + + override def requiredDataAttributes(): Array[NamedReference] = { + val base = Seq(FieldReference("pk"), FieldReference("dep")) + val baseNames = base.map(_.describe()).toSet + (base ++ updatedCols.filterNot(r => baseNames.contains(r.describe()))).toArray + } + + override def requiredMetadataAttributes(): Array[NamedReference] = { + if (noMetadata) { + Array.empty + } else { + Array(PARTITION_COLUMN_REF, INDEX_COLUMN_REF) + } + } + + private def clusterColumnRef: NamedReference = + if (noMetadata) FieldReference("dep") else PARTITION_COLUMN_REF + + override def newScanBuilder(options: CaseInsensitiveStringMap): ScanBuilder = { + val onBuild: BatchScanBaseClass => Unit = { scan => + InMemoryRowLevelOperationTable.recordLastScanSchema(scan.readSchema()) + configuredScan = scan + } + if (properties.getOrDefault(COLUMN_UPDATE_COW_UNREAD_FILTER_ATTRS, "false") == "true") { + new InMemoryScanBuilder(schema, options) { + override protected def createScan( + partitions: Seq[InputPartition], + readSchema: StructType, + tableSchema: StructType, + options: CaseInsensitiveStringMap): BatchScanBaseClass = { + new InMemoryBatchScan(partitions, readSchema, tableSchema, options) { + override def filterAttributes(): Array[NamedReference] = identityPartitionReferences + } + } + + override def build(): Scan = { + val scan = super.build().asInstanceOf[BatchScanBaseClass] + onBuild(scan) + scan + } + } + } else { + newRowLevelScanBuilder(options)(onBuild) + } + } + + override def newWriteBuilder(info: LogicalWriteInfo): WriteBuilder = { + lastWriteInfo = info + new WriteBuilder { + override def build(): Write = new Write with RequiresDistributionAndOrdering { + override def requiredDistribution: Distribution = + Distributions.clustered(Array(clusterColumnRef)) + + override def requiredOrdering: Array[SortOrder] = Array[SortOrder]( + LogicalExpressions.sort( + clusterColumnRef, + SortDirection.ASCENDING, + SortDirection.ASCENDING.defaultNullOrdering())) + + override def toBatch: BatchWrite = { + val narrowSchema = if (info.columnUpdateSchema().isPresent) { + info.columnUpdateSchema().get() + } else { + info.schema() + } + // info.schema() is empty for column-update UPDATE writes (no INSERT rows); + // use the table schema to align any INSERT-tagged rows. + PartitionBasedNarrowReplaceData(configuredScan, narrowSchema, schema) + } + + override def description: String = "InMemoryNarrowCoWWrite" + } + } + } + + override def description(): String = "InMemoryPartitionColumnUpdateOperation" + } + + // Test-only: a copy-on-write column-update fixture with a fixed requiredDataAttributes(), so + // tests can declare columns independently of the UPDATE. The commit path looks rows up by `pk`, + // so `pk` must be declared. + class PartitionBasedColumnUpdateOperationWithReqAttrs( + command: Command, + reqCols: Array[String], + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends PartitionBasedColumnUpdateOperation(command, options = options) { + override def requiredDataAttributes(): Array[NamedReference] = reqCols.map(FieldReference(_)) + } + + // Test-only: mixes in SupportsColumnUpdates like PartitionBasedColumnUpdateOperation, but its + // DataWriter never overrides writeColumnUpdate(), to exercise the DataWriter#writeColumnUpdate + // default that throws DATA_SOURCE_WRITE_COLUMN_UPDATE_NOT_IMPLEMENTED. + class PartitionBasedColumnUpdateOperationNoWriteColumnUpdate( + command: Command, + updatedCols: Seq[NamedReference] = Nil, + options: CaseInsensitiveStringMap = CaseInsensitiveStringMap.empty()) + extends PartitionBasedColumnUpdateOperation(command, updatedCols, options) { + + override def newWriteBuilder(info: LogicalWriteInfo): WriteBuilder = { + lastWriteInfo = info + new WriteBuilder { + override def build(): Write = new Write with RequiresDistributionAndOrdering { + override def requiredDistribution: Distribution = + Distributions.clustered(Array(PARTITION_COLUMN_REF)) + + override def requiredOrdering: Array[SortOrder] = Array[SortOrder]( + LogicalExpressions.sort( + PARTITION_COLUMN_REF, + SortDirection.ASCENDING, + SortDirection.ASCENDING.defaultNullOrdering())) + + override def toBatch: BatchWrite = new TestBatchWrite { + override def createBatchWriterFactory( + info: PhysicalWriteInfo): DataWriterFactory = { + new NoWriteColumnUpdateOverrideWriterFactory + } + + override protected def doCommit(messages: Array[WriterCommitMessage]): Unit = {} + } + + override def description: String = "InMemoryColumnUpdateNoWriteColumnUpdateOverride" + } + } + } + } + + // CoW write handler for narrow column-update writes. + // Narrow rows (UPDATE/COPY) are sent via writeColumnUpdate, wide rows (INSERT) via write. + // Both arrive in the same buffer; rows are routed by the operation tag in the log entry, + // which lets tests assert that Spark dispatched through the correct writer method. + private case class PartitionBasedNarrowReplaceData( + scan: BatchScanBaseClass, + writeSchema: StructType, + fullSchema: StructType) extends TestBatchWrite { + + override def createBatchWriterFactory(info: PhysicalWriteInfo): DataWriterFactory = { + if (properties.getOrDefault(COLUMN_UPDATE_COW_RECORD_ONLY_WRITER, "false") == "true") { + new RecordOnlyColumnUpdateWriterFactory(CatalogV2Util.v2ColumnsToStructType(columns())) + } else { + super.createBatchWriterFactory(info) + } + } + + override protected def doCommit( + messages: Array[WriterCommitMessage]): Unit = dataMap.synchronized { + val newData = messages.map(_.asInstanceOf[BufferedRows]) + val readRows = scan.data.flatMap(_.asInstanceOf[BufferedRows].rows) + val readPartitions = readRows.map(r => getKey(r, schema)).distinct + dataMap --= readPartitions + replacedPartitions = readPartitions + + val writeFieldIdx = writeSchema.fieldNames.zipWithIndex.toMap + val pkIdxInWrite = writeFieldIdx("pk") + val pkIdxInFull = schema.fieldIndex("pk") + val writeColumnUpdateOpName = UTF8String.fromString(WriteColumnUpdate.toString) + + val expandedData = newData.map { buf => + val expanded = new BufferedRows(buf.key, schema) + // Walk the log so we can route on the operation tag (Write vs WriteColumnUpdate). + // buf.rows contains rows in the same order as the log; iterate together. + buf.log.zip(buf.rows).foreach { case (logEntry, row) => + val opName = logEntry.getUTF8String(0) + if (opName == writeColumnUpdateOpName) { + // UPDATE/COPY narrow row: look up base row by pk, overlay narrow values + val pk = row.getInt(pkIdxInWrite) + val origRow = readRows.find(r => r.getInt(pkIdxInFull) == pk) + val fullRow = overlayNarrowRow(origRow, row, writeFieldIdx) + expanded.rows.append(fullRow) + } else { + // INSERT row: full schema, append directly aligned to table schema + val fullRow = new GenericInternalRow(schema.length) + schema.fields.zipWithIndex.foreach { case (field, i) => + val srcIdx = fullSchema.fieldIndex(field.name) + fullRow.update(i, row.get(srcIdx, field.dataType)) + } + expanded.rows.append(fullRow) + } + } + expanded + } + + withData(expandedData, schema) + lastWriteLog = newData.flatMap(buffer => buffer.log).toImmutableArraySeq + } + } + private object TestDeltaBatchWrite extends TestBatchWrite with DeltaBatchWrite { override def createBatchWriterFactory(info: PhysicalWriteInfo): DeltaWriterFactory = { new DeltaBufferedRowsWriterFactory(CatalogV2Util.v2ColumnsToStructType(columns())) @@ -327,6 +824,40 @@ private class DeltaBufferedRowsWriterFactory(schema: StructType) extends DeltaWr } } +// Test-only: a writer factory for a connector that never overrides writeColumnUpdate(), to +// exercise the DataWriter#writeColumnUpdate default that throws +// DATA_SOURCE_WRITE_COLUMN_UPDATE_NOT_IMPLEMENTED. +// Declared top-level (no reference to the enclosing table) so it stays task-serializable. +private class NoWriteColumnUpdateOverrideWriterFactory extends DataWriterFactory { + override def createWriter(partitionId: Int, taskId: Long): DataWriter[InternalRow] = { + new NoWriteColumnUpdateOverrideWriter + } +} + +private class NoWriteColumnUpdateOverrideWriter extends DataWriter[InternalRow] { + override def write(record: InternalRow): Unit = {} + override def commit(): WriterCommitMessage = new BufferedRows(Seq.empty, StructType(Nil)) + override def abort(): Unit = {} + override def close(): Unit = {} +} + +// Test-only: a writer that overrides writeColumnUpdate(record) but not +// writeColumnUpdate(metadata, record), so rows with metadata reach it through the default. +private class RecordOnlyColumnUpdateWriterFactory(schema: StructType) extends DataWriterFactory { + override def createWriter(partitionId: Int, taskId: Long): DataWriter[InternalRow] = { + new RecordOnlyColumnUpdateWriter(schema) + } +} + +private class RecordOnlyColumnUpdateWriter(schema: StructType) extends DataWriter[InternalRow] { + private val delegate = new BufferWriter(schema) + override def write(record: InternalRow): Unit = delegate.write(record) + override def writeColumnUpdate(record: InternalRow): Unit = delegate.writeColumnUpdate(record) + override def commit(): WriterCommitMessage = delegate.commit() + override def abort(): Unit = delegate.abort() + override def close(): Unit = delegate.close() +} + private class DeltaBufferWriter(schema: StructType) extends BufferWriter(schema) with DeltaWriter[InternalRow] { @@ -373,6 +904,21 @@ private class DeltaBufferWriter(schema: StructType) extends BufferWriter(schema) } object InMemoryRowLevelOperationTable { + // Global holder for the scan schema of the last row-level operation. See the class-level + // `lastScanSchema` def for why this is a companion-object field rather than a per-instance one. + // Tests that assert on scan pruning should reset this before running the operation under test + // (see `RowLevelOperationSuiteBase.beforeEach`). + private[catalog] val lastScanSchemaRef = new AtomicReference[StructType]() + + /** Called by test-only scan builders on every scan build. Overwrites; tests reset between + * cases. */ + private[catalog] def recordLastScanSchema(schema: StructType): Unit = { + lastScanSchemaRef.set(schema) + } + + /** Called by test setup to clear the last recorded scan schema between test cases. */ + def resetLastScanSchema(): Unit = lastScanSchemaRef.set(null) + def withColumns( name: String, columns: Array[Column], diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/txns.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/txns.scala index 0a18f32e70295..f1b522a1ec27a 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/txns.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/connector/catalog/txns.scala @@ -160,6 +160,7 @@ class TxnTable( delegate.replacedPartitions = replacedPartitions delegate.lastWriteInfo = lastWriteInfo delegate.lastWriteLog = lastWriteLog + delegate.lastUpdatedColumns = lastUpdatedColumns delegate.commits ++= commits delegate.increaseVersion() } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/DataSourceV2Strategy.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/DataSourceV2Strategy.scala index 32232532bbd61..4b6a66cb223d1 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/DataSourceV2Strategy.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/DataSourceV2Strategy.scala @@ -590,7 +590,8 @@ class DataSourceV2Strategy(session: SparkSession) extends Strategy with Predicat projections, write, rd.operation.command, - r.name) :: Nil + r.name, + useWriteColumnUpdate = V2Writes.hasNarrowRows(rd)) :: Nil case wd @ WriteDelta(_: DataSourceV2Relation, _, query, r: DataSourceV2Relation, projections, _, Some(write)) => diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileTable.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileTable.scala index 70d40e32ac33a..33f32a5c546bc 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileTable.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/FileTable.scala @@ -212,7 +212,8 @@ abstract class FileTable( writeInfo.schema(), mergedOptions(writeInfo.options()), writeInfo.rowIdSchema(), - writeInfo.metadataSchema()) + writeInfo.metadataSchema(), + writeInfo.columnUpdateSchema()) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/GroupBasedRowLevelOperationScanPlanning.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/GroupBasedRowLevelOperationScanPlanning.scala index a41aad05d4351..8c6b4373d75b0 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/GroupBasedRowLevelOperationScanPlanning.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/GroupBasedRowLevelOperationScanPlanning.scala @@ -18,7 +18,7 @@ package org.apache.spark.sql.execution.datasources.v2 import org.apache.spark.internal.LogKeys -import org.apache.spark.sql.catalyst.expressions.{And, AttributeReference, AttributeSet, Expression, ExpressionSet, PredicateHelper, SubqueryExpression} +import org.apache.spark.sql.catalyst.expressions.{And, AttributeReference, AttributeSet, Expression, ExpressionSet, NamedExpression, PredicateHelper, SubqueryExpression} import org.apache.spark.sql.catalyst.expressions.Literal.TrueLiteral import org.apache.spark.sql.catalyst.planning.{GroupBasedRowLevelOperation, PhysicalOperation} import org.apache.spark.sql.catalyst.plans.logical.{Join, LogicalPlan, ReplaceData} @@ -27,6 +27,7 @@ import org.apache.spark.sql.catalyst.trees.TreePattern.REPLACE_DATA import org.apache.spark.sql.connector.expressions.filter.{Predicate => V2Filter} import org.apache.spark.sql.connector.read.ScanBuilder import org.apache.spark.sql.connector.write.RowLevelOperation.Command.MERGE +import org.apache.spark.sql.connector.write.RowLevelOperationTable import org.apache.spark.sql.execution.datasources.DataSourceStrategy import org.apache.spark.sql.internal.connector.PartitionPredicateField import org.apache.spark.sql.sources.Filter @@ -65,7 +66,9 @@ object GroupBasedRowLevelOperationScanPlanning extends Rule[LogicalPlan] with Pr .mkString(", ") } - val (scan, output) = PushDownUtils.pruneColumns(scanBuilder, relation, relation.output, Nil) + val requiredAttrs = + if (V2Writes.hasNarrowRows(rd)) readAttrs(rd.query, table, relation) else relation.output + val (scan, output) = PushDownUtils.pruneColumns(scanBuilder, relation, requiredAttrs, Nil) // scalastyle:off line.size.limit logInfo( @@ -92,6 +95,38 @@ object GroupBasedRowLevelOperationScanPlanning extends Rule[LogicalPlan] with Pr } } + /** + * Returns the columns of `relation` that the projections and filters directly above each read + * of the table use, as the rest of the query sees only their output. Only whole columns are + * pruned, as the query is not rewritten for a pruned nested schema. + * + * An UPDATE with a subquery in its condition reads the table twice, once per Union child, and + * the analyzer gives the second read new expr IDs. Both reads share one scan, so the columns + * are matched to `relation` by name. + */ + private def readAttrs( + plan: LogicalPlan, + table: RowLevelOperationTable, + relation: DataSourceV2Relation): Seq[AttributeReference] = { + val readNames = tableReads(plan, table).flatMap { case (projects, filters, read) => + val references = AttributeSet((projects ++ filters).flatMap(_.references)) + read.output.filter(references.contains).map(_.name) + } + relation.output.filter(attr => readNames.exists(conf.resolver(_, attr.name))) + } + + private def tableReads( + plan: LogicalPlan, + table: RowLevelOperationTable) + : Seq[(Seq[NamedExpression], Seq[Expression], DataSourceV2Relation)] = { + plan match { + case PhysicalOperation(projects, filters, r: DataSourceV2Relation) if r.table eq table => + Seq((projects, filters, r)) + case _ => + plan.children.flatMap(tableReads(_, table)) + } + } + // pushes down the operation condition and returns the following information: // - pushed down filters // - filter expressions that are fully evaluated on the data source side diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2Writes.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2Writes.scala index 5340937115d96..aaa4a4a1f38b5 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2Writes.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/V2Writes.scala @@ -22,13 +22,15 @@ import java.util.UUID import scala.jdk.CollectionConverters._ import org.apache.spark.sql.catalyst.expressions.PredicateHelper -import org.apache.spark.sql.catalyst.plans.logical.{AppendData, InsertOnlyMerge, LogicalPlan, OverwriteByExpression, OverwritePartitionsDynamic, ReplaceData, WriteDelta} +import org.apache.spark.sql.catalyst.plans.logical.{AppendData, InsertOnlyMerge, LogicalPlan, OverwriteByExpression, OverwritePartitionsDynamic, ReplaceData, RowLevelWrite, WriteDelta} import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.catalyst.streaming.InternalOutputModes._ import org.apache.spark.sql.catalyst.util.WriteDeltaProjections import org.apache.spark.sql.connector.catalog.Table +import org.apache.spark.sql.connector.distributions.{ClusteredDistribution, OrderedDistribution} import org.apache.spark.sql.connector.expressions.filter.Predicate -import org.apache.spark.sql.connector.write.{DeltaWriteBuilder, LogicalWriteInfoImpl, SupportsDynamicOverwrite, SupportsOverwriteV2, SupportsTruncate, Write, WriteBuilder} +import org.apache.spark.sql.connector.write.{DeltaWriteBuilder, LogicalWriteInfoImpl, RequiresDistributionAndOrdering, SupportsColumnUpdates, SupportsDynamicOverwrite, SupportsOverwriteV2, SupportsTruncate, Write, WriteBuilder} +import org.apache.spark.sql.connector.write.RowLevelOperation.Command.UPDATE import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors} import org.apache.spark.sql.execution.streaming.sources.{MicroBatchWrite, WriteToMicroBatchDataSource} import org.apache.spark.sql.internal.connector.SupportsStreamingUpdateAsAppend @@ -109,22 +111,75 @@ object V2Writes extends Rule[LogicalPlan] with PredicateHelper { WriteToDataSourceV2(Some(r), microBatchWrite, newQuery, customMetrics) case rd @ ReplaceData(r: DataSourceV2Relation, _, query, _, projections, _, None) => - val rowSchema = projections.rowProjection.schema + val (rowSchema, columnUpdateSchema) = + writeSchemas(projections.rowProjection.schema, hasNarrowRows(rd)) val metadataSchema = projections.metadataProjection.map(_.schema) val writeOptions = mergeOptions(Map.empty, r.options.asCaseSensitiveMap.asScala.toMap) - val writeBuilder = newWriteBuilder(r.table, writeOptions, rowSchema, metadataSchema) + val writeBuilder = newWriteBuilder(r.table, writeOptions, rowSchema, metadataSchema, + columnUpdateSchema) val write = writeBuilder.build() + validateColumnUpdateWrite(rd, write) val newQuery = DistributionAndOrderingUtils.prepareQuery(write, query, r.funCatalog) rd.copy(write = Some(write), query = newQuery) case wd @ WriteDelta(r: DataSourceV2Relation, _, query, _, projections, _, None) => val writeOptions = mergeOptions(Map.empty, r.options.asCaseSensitiveMap.asScala.toMap) - val deltaWriteBuilder = newDeltaWriteBuilder(r.table, writeOptions, projections) + val deltaWriteBuilder = + newDeltaWriteBuilder(r.table, writeOptions, projections, hasNarrowRows(wd)) val deltaWrite = deltaWriteBuilder.build() + validateColumnUpdateWrite(wd, deltaWrite) val newQuery = DistributionAndOrderingUtils.prepareQuery(deltaWrite, query, r.funCatalog) wd.copy(write = Some(deltaWrite), query = newQuery) } + /** + * Whether a row-level write delivers narrow rows for updated, copied, and reinserted records. + * This must match the writes whose write relation RewriteUpdateTable narrows, even when the + * declared columns cover the whole table. + */ + private[v2] def hasNarrowRows(write: RowLevelWrite): Boolean = { + write.operation.isInstanceOf[SupportsColumnUpdates] && write.operation.command == UPDATE + } + + /** + * Checks that the distribution and ordering of a write that delivers narrow rows reference only + * columns the write reads. Other columns stay in the query only until column pruning removes + * them, so they are rejected here whether or not pruning ran. + */ + private def validateColumnUpdateWrite(rowLevelWrite: RowLevelWrite, write: Write): Unit = { + write match { + case w: RequiresDistributionAndOrdering if hasNarrowRows(rowLevelWrite) => + val distribution = w.requiredDistribution match { + case d: ClusteredDistribution => d.clustering.toImmutableArraySeq + case d: OrderedDistribution => d.ordering.toImmutableArraySeq + case _ => Seq.empty + } + val readNames = rowLevelWrite.references.toSeq.map(_.name) + val unread = (distribution ++ w.requiredOrdering.toImmutableArraySeq) + .flatMap(_.references.toImmutableArraySeq) + .filterNot(_.fieldNames.headOption.exists(n => readNames.exists(conf.resolver(_, n)))) + .map(_.describe) + .distinct + .sorted + if (unread.nonEmpty) { + throw QueryCompilationErrors.undeclaredWriteRequirementColumnsError( + rowLevelWrite.operation.getClass.getName, unread) + } + case _ => + } + } + + /** + * Returns the `schema()` and `columnUpdateSchema()` of the write. Narrow rows are reported as + * `columnUpdateSchema()`, and `schema()` covers only newly inserted rows, which is empty as only + * UPDATE is narrowed and it inserts no new rows. + */ + private def writeSchemas( + rowSchema: StructType, + narrowRows: Boolean): (StructType, Option[StructType]) = { + if (narrowRows) (StructType(Nil), Some(rowSchema)) else (rowSchema, None) + } + private def mergeOptions( commandOptions: Map[String, String], dsOptions: Map[String, String]): Map[String, String] = { @@ -163,6 +218,7 @@ object V2Writes extends Rule[LogicalPlan] with PredicateHelper { writeOptions: Map[String, String], rowSchema: StructType, metadataSchema: Option[StructType] = None, + columnUpdateSchema: Option[StructType] = None, queryId: String = UUID.randomUUID().toString): WriteBuilder = { val info = LogicalWriteInfoImpl( @@ -170,7 +226,8 @@ object V2Writes extends Rule[LogicalPlan] with PredicateHelper { rowSchema, writeOptions.asOptions, rowIdSchema = None, - metadataSchema) + metadataSchema, + columnUpdateSchema) table.asWritable.newWriteBuilder(info) } @@ -178,9 +235,11 @@ object V2Writes extends Rule[LogicalPlan] with PredicateHelper { table: Table, writeOptions: Map[String, String], projections: WriteDeltaProjections, + narrowRows: Boolean, queryId: String = UUID.randomUUID().toString): DeltaWriteBuilder = { - val rowSchema = projections.rowProjection.map(_.schema).getOrElse(StructType(Nil)) + val (rowSchema, columnUpdateSchema) = writeSchemas( + projections.rowProjection.map(_.schema).getOrElse(StructType(Nil)), narrowRows) val rowIdSchema = Some(projections.rowIdProjection.schema) val metadataSchema = projections.metadataProjection.map(_.schema) @@ -189,7 +248,8 @@ object V2Writes extends Rule[LogicalPlan] with PredicateHelper { rowSchema, writeOptions.asOptions, rowIdSchema, - metadataSchema) + metadataSchema, + columnUpdateSchema) val writeBuilder = table.asWritable.newWriteBuilder(info) assert(writeBuilder.isInstanceOf[DeltaWriteBuilder], s"$writeBuilder must be DeltaWriteBuilder") diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/WriteToDataSourceV2Exec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/WriteToDataSourceV2Exec.scala index b9d01153e5ee4..d26e546c25e68 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/WriteToDataSourceV2Exec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/datasources/v2/WriteToDataSourceV2Exec.scala @@ -380,14 +380,17 @@ case class ReplaceDataExec( write: Write, rowLevelCommand: RowLevelOperation.Command, tableName: String, + useWriteColumnUpdate: Boolean, transaction: Option[Transaction] = None) extends RowLevelWriteExec { override def writingTask: WritingSparkTask[_] = { projections.metadataProjection match { case Some(metadataProj) => - DataAndMetadataWritingSparkTask(projections.rowProjection, metadataProj, sparkMetrics) + DataAndMetadataWritingSparkTask( + projections.rowProjection, metadataProj, useWriteColumnUpdate, sparkMetrics) case None => - DataWithProjectionWritingSparkTask(projections.rowProjection, sparkMetrics) + DataWithProjectionWritingSparkTask( + projections.rowProjection, useWriteColumnUpdate, sparkMetrics) } } @@ -772,6 +775,7 @@ trait WritingSparkTask[W <: DataWriter[InternalRow]] extends Logging with Serial case class DataAndMetadataWritingSparkTask( dataProj: ProjectingInternalRow, metadataProj: ProjectingInternalRow, + useWriteColumnUpdate: Boolean, sparkMetrics: Map[String, SQLMetric]) extends WritingSparkTask[DataWriter[InternalRow]] { @@ -779,6 +783,8 @@ case class DataAndMetadataWritingSparkTask( writer: DataWriter[InternalRow], iter: java.util.Iterator[InternalRow]): Unit = { var numUpdatedRows = 0L var numCopiedRows = 0L + val writeCopiedOrUpdatedRow: (InternalRow, InternalRow) => Unit = + if (useWriteColumnUpdate) writer.writeColumnUpdate(_, _) else writer.write(_, _) while (iter.hasNext) { val row = iter.next() @@ -789,13 +795,13 @@ case class DataAndMetadataWritingSparkTask( numUpdatedRows += 1L dataProj.project(row) metadataProj.project(row) - writer.write(metadataProj, dataProj) + writeCopiedOrUpdatedRow(metadataProj, dataProj) case COPY_OPERATION => numCopiedRows += 1L dataProj.project(row) metadataProj.project(row) - writer.write(metadataProj, dataProj) + writeCopiedOrUpdatedRow(metadataProj, dataProj) case INSERT_OPERATION => dataProj.project(row) @@ -813,6 +819,7 @@ case class DataAndMetadataWritingSparkTask( case class DataWithProjectionWritingSparkTask( dataProj: ProjectingInternalRow, + useWriteColumnUpdate: Boolean, sparkMetrics: Map[String, SQLMetric]) extends WritingSparkTask[DataWriter[InternalRow]] { @@ -820,6 +827,8 @@ case class DataWithProjectionWritingSparkTask( writer: DataWriter[InternalRow], iter: java.util.Iterator[InternalRow]): Unit = { var numUpdatedRows = 0L var numCopiedRows = 0L + val writeCopiedOrUpdatedRow: InternalRow => Unit = + if (useWriteColumnUpdate) writer.writeColumnUpdate(_) else writer.write(_) while (iter.hasNext) { val row = iter.next() @@ -829,12 +838,12 @@ case class DataWithProjectionWritingSparkTask( case UPDATE_OPERATION => numUpdatedRows += 1L dataProj.project(row) - writer.write(dataProj) + writeCopiedOrUpdatedRow(dataProj) case COPY_OPERATION => numCopiedRows += 1L dataProj.project(row) - writer.write(dataProj) + writeCopiedOrUpdatedRow(dataProj) case INSERT_OPERATION => dataProj.project(row) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/dynamicpruning/RowLevelOperationRuntimeGroupFiltering.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/dynamicpruning/RowLevelOperationRuntimeGroupFiltering.scala index a26a0cea74f1e..cf3715ebd804f 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/dynamicpruning/RowLevelOperationRuntimeGroupFiltering.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/dynamicpruning/RowLevelOperationRuntimeGroupFiltering.scala @@ -43,8 +43,9 @@ import org.apache.spark.util.ArrayImplicits._ * Note that this rule is also beneficial for operations that deal with deltas of rows. Even if * the data source is capable of handling specific changes, it is useful to first discard entire * groups that are not modified. The cost of the runtime query is small as it only projects columns - * required to evaluate the row level operation condition. The main scan, on the other hand, must - * project all columns, meaning the cost of reading unaffected groups can dominate the runtime. + * required to evaluate the row level operation condition. The main scan, on the other hand, also + * projects the columns the write needs, meaning the cost of reading unaffected groups can dominate + * the runtime. */ class RowLevelOperationRuntimeGroupFiltering(optimizeSubqueries: Rule[LogicalPlan]) extends Rule[LogicalPlan] with PredicateHelper { @@ -110,8 +111,10 @@ class RowLevelOperationRuntimeGroupFiltering(optimizeSubqueries: Rule[LogicalPla // this rule assigns runtime filters to both scan relations (will be shared at runtime) // and must transform the runtime filter condition to use correct expr IDs for each relation // note this only applies to group-based row-level operations (i.e. ReplaceData) + // the map is built from the original table as the write table may be narrower than the + // columns the condition references // see RewriteUpdateTable for more details - val attrMap = buildTableToScanAttrMap(write.table.output, relation.output) + val attrMap = buildTableToScanAttrMap(write.originalTable.output, relation.output) val transformedCond = cond transform { case attr: AttributeReference if attrMap.contains(attr) => attrMap(attr) } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedColumnUpdateTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedColumnUpdateTableSuite.scala new file mode 100644 index 0000000000000..4cb1518b3be7a --- /dev/null +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedColumnUpdateTableSuite.scala @@ -0,0 +1,999 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.connector + +import org.apache.spark.SparkRuntimeException +import org.apache.spark.sql.{AnalysisException, Row} +import org.apache.spark.sql.catalyst.optimizer.ColumnPruning +import org.apache.spark.sql.catalyst.plans.logical.WriteDelta +import org.apache.spark.sql.connector.catalog.{Delete, InMemoryBaseTable, InMemoryRowLevelOperationTable, Reinsert} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{IntegerType, StringType, StructField, StructType} + +/** + * Tests for UPDATE statements targeting connectors that mix in + * [[org.apache.spark.sql.connector.write.SupportsColumnUpdates]]. + * + * When a connector supports column updates, updated rows contain only the declared columns + * (LogicalWriteInfo.columnUpdateSchema()) rather than the full table row. + */ +class DeltaBasedColumnUpdateTableSuite extends DeltaBasedUpdateTableSuiteBase { + + override protected lazy val extraTableProps: java.util.Map[String, String] = { + val props = new java.util.HashMap[String, String]() + props.put("column-update", "true") + props + } + + private val reqAttrsOperationClass = + classOf[InMemoryRowLevelOperationTable#DeltaBasedColumnUpdateOperationWithReqAttrs].getName + + private val splitOperationClass = + classOf[InMemoryRowLevelOperationTable#DeltaBasedColumnUpdateSplitOperation].getName + + private val splitReqAttrsOperationClass = + classOf[InMemoryRowLevelOperationTable#DeltaBasedColumnUpdateSplitOperationWithReqAttrs] + .getName + + private val splitWithRowIdOperationClass = + classOf[InMemoryRowLevelOperationTable#DeltaBasedColumnUpdateSplitOperationWithRowId].getName + + test("column-update: column update schema contains a single assigned column") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + PK_FIELD, + DEP_FIELD, + StructField("id", IntegerType, nullable = false) + )))) + } + + test("column-update: column update schema contains multiple assigned columns") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1, dep = 'engineering' WHERE pk = 1") + + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + PK_FIELD, + StructField("dep", StringType, nullable = false), + StructField("id", IntegerType, nullable = false) + )))) + } + + test("column-update: column update schema is pk and dep for a full identity update") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = id, dep = dep WHERE pk = 1") + + // All assignments are identity, so updatedColumns is empty; pk and dep remain because the + // connector unconditionally declares them as base columns. + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Array(PK_FIELD, DEP_FIELD)))) + } + + test("column-update: row filter condition is orthogonal to column narrowing") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET dep = 'engineering' WHERE pk IN (1, 3)") + + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + PK_FIELD, + StructField("dep", StringType, nullable = false) + )))) + } + + test("column-update: update all rows (no WHERE clause)") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |{ "pk": 3, "salary": 300, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = salary * 2") + + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + PK_FIELD, + DEP_FIELD, + StructField("salary", IntegerType, nullable = true) + )))) + } + + test("column-update: column update schema excludes identity assignments in a mixed UPDATE") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = id, dep = 'engineering' WHERE pk = 1") + + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + PK_FIELD, + StructField("dep", StringType, nullable = false) + )))) + } + + + test("column-update: nested struct field update narrows to the root struct column") { + createAndInitTable("pk INT NOT NULL, s STRUCT, dep STRING", + """{ "pk": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET s.c1 = -1 WHERE pk = 1") + + checkLastUpdatedColumns("s") + + // `dep` is present because the connector unconditionally declares it as a base column. + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == Seq("pk", "dep", "s")) + } + + test("column-update: nested field identity update reports root struct as updated") { + createAndInitTable("pk INT NOT NULL, s STRUCT, dep STRING", + """{ "pk": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET s.c1 = s.c1 WHERE pk = 1") + + checkLastUpdatedColumns("s") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, Row(1, 2), "hr") :: Nil) + } + + test("column-update: updatedColumns is empty for DELETE") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |""".stripMargin) + + sql(s"DELETE FROM $tableNameAsString WHERE dep = 'hr'") + + checkLastUpdatedColumns() + } + + test("column-update: row-level DELETE falls back to full-width rows") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "hr" } + |""".stripMargin) + + sql(s"DELETE FROM $tableNameAsString WHERE pk < 2") + + assert(!table.lastWriteInfo.columnUpdateSchema().isPresent, + s"DELETE must not report a column update schema: " + + s"${table.lastWriteInfo.columnUpdateSchema()}") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(2, 200, "hr") :: Nil) + } + + test("column-update: MERGE falls back to full-width rows") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + checkMergeFallsBackToFullWidthRows() + } + + test("column-update split: MERGE falls back to full-width rows") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin, + java.util.Map.of("column-update-split", "true")) + checkMergeFallsBackToFullWidthRows() + } + + private def checkMergeFallsBackToFullWidthRows(): Unit = { + withTempView("source") { + import testImplicits._ + Seq((1, 150, "hr"), (3, 300, "software")).toDF("pk", "salary", "dep") + .createOrReplaceTempView("source") + + sql( + s"""MERGE INTO $tableNameAsString t + |USING source s + |ON t.pk = s.pk + |WHEN MATCHED THEN UPDATE SET salary = s.salary + |WHEN NOT MATCHED THEN INSERT * + |""".stripMargin) + + assert(!table.lastWriteInfo.columnUpdateSchema().isPresent, + s"MERGE must not report a column update schema: " + + s"${table.lastWriteInfo.columnUpdateSchema()}") + checkLastUpdatedColumns() + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 150, "hr") :: Row(2, 200, "software") :: Row(3, 300, "software") :: Nil) + } + } + + test("column-update: data correctness -- single column update") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr") :: Row(2, 2, "software") :: Row(3, 3, "hr") :: Nil) + checkUpdateMetrics(numUpdatedRows = 1, numCopiedRows = 0) + } + + test("column-update: data correctness -- update all rows") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |{ "pk": 3, "salary": 300, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = salary * 2") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 200, "hr") :: Row(2, 400, "software") :: Row(3, 600, "hr") :: Nil) + checkUpdateMetrics(numUpdatedRows = 3, numCopiedRows = 0) + } + + test("column-update: data correctness -- mixed identity and real assignments") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = id, dep = 'engineering' WHERE pk = 1") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "engineering") :: Row(2, 2, "software") :: Row(3, 3, "hr") :: Nil) + checkUpdateMetrics(numUpdatedRows = 1, numCopiedRows = 0) + } + + private def createAndInitTableWithReqAttrs( + reqAttrs: String, + schemaString: String, + jsonData: String): Unit = { + createAndInitTable(schemaString, jsonData, + java.util.Map.of("column-update-req-attrs", reqAttrs)) + } + + test("column-update: requiredDataAttributes - data correctness") { + createAndInitTableWithReqAttrs("dep,id", "pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr") :: Row(2, 2, "software") :: Row(3, 3, "hr") :: Nil) + } + + test("column-update: undeclared distribution column is rejected") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of( + "column-update-req-attrs", "pk,salary", + "column-update-cluster-by", "dep")) + + val excludeColumnPruning = SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> ColumnPruning.ruleName + Seq(Nil, Seq(excludeColumnPruning)).foreach { confs => + withSQLConf(confs: _*) { + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_UNDECLARED_WRITE_REQUIREMENT_COLUMNS", + parameters = Map("connector" -> reqAttrsOperationClass, "columns" -> "[dep]")) + } + } + } + + test("column-update: distribution by an undeclared row ID column is allowed") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "hr" } + |""".stripMargin, + java.util.Map.of( + "column-update-req-attrs", "salary,dep", + "column-update-cluster-by", "pk")) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr") :: Row(2, 200, "hr") :: Nil) + } + + test("column-update: partition source column not in requiredDataAttributes") { + createAndInitTableWithReqAttrs("pk,id", "pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkLastScanExcludes("dep") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr") :: Row(2, 2, "software") :: Row(3, 3, "hr") :: Nil) + } + + test("column-update: requiredDataAttributes resolves case-insensitively on the row-ID " + + "column without adopting the declared spelling") { + createAndInitTableWithReqAttrs("PK,salary,dep", "pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = salary + 1 WHERE pk = 1") + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == + Seq("pk", "salary", "dep")) + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 101, "hr") :: Row(2, 200, "software") :: Nil) + } + + + test("column-update: undeclared, differently-cased columns in the RHS and WHERE clause " + + "are read with the table's own spelling") { + createAndInitTableWithReqAttrs("pk,dep,salary", + "pk INT NOT NULL, salary INT, bonus INT, extra INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "extra": 5, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = salary + BONUS WHERE pk = 1 AND EXTRA > 3") + + checkLastScanSchema("pk INT, salary INT, bonus INT, extra INT, dep STRING, _partition STRING") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, 110, 10, 5, "hr") :: Nil) + } + + test("column-update: scan that returns unrequested columns still writes narrow rows") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |""".stripMargin, + java.util.Map.of( + "column-update", "true", + InMemoryBaseTable.SIMULATE_PARTIAL_COLUMN_PRUNING, "true")) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + checkLastScanIncludes("bonus") + val columnUpdateSchema = table.lastWriteInfo.columnUpdateSchema().get() + assert(!columnUpdateSchema.fieldNames.contains("bonus"), + s"bonus must not be in column update schema: $columnUpdateSchema") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, -1, 10, "hr") :: Nil) + } + + test("column-update: DELETE on a scan that returns unrequested columns is unaffected") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "software" } + |""".stripMargin, + java.util.Map.of( + "column-update", "true", + InMemoryBaseTable.SIMULATE_PARTIAL_COLUMN_PRUNING, "true")) + + sql(s"DELETE FROM $tableNameAsString WHERE pk < 2") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(2, 200, 20, "software") :: Nil) + } + + test("column-update: requiredDataAttributes resolves case-insensitively on an assigned " + + "column without adopting the declared spelling") { + createAndInitTableWithReqAttrs("pk,SALARY,dep", "pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = salary + 1 WHERE pk = 1") + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == + Seq("pk", "salary", "dep")) + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 101, "hr") :: Row(2, 200, "software") :: Nil) + } + + test("column-update: case-sensitive analysis accepts columns that differ only in case") { + withSQLConf(SQLConf.CASE_SENSITIVE.key -> "true") { + createAndInitTableWithReqAttrs("pk,salary,SALARY,dep", + "pk INT NOT NULL, salary INT, SALARY INT, dep STRING", + """{ "pk": 1, "salary": 100, "SALARY": 200, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == + Seq("pk", "salary", "SALARY", "dep")) + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), Row(1, -1, 200, "hr") :: Nil) + } + } + + test("column-update: empty requiredDataAttributes throws AnalysisException") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" }""".stripMargin, + java.util.Map.of("column-update-empty-req-attrs", "true")) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_EMPTY_REQUIRED_DATA_ATTRIBUTES", + parameters = Map("connector" -> reqAttrsOperationClass)) + } + + test("column-update: nested requiredDataAttributes throws AnalysisException") { + createAndInitTableWithReqAttrs("pk,s.c1,dep", + "pk INT NOT NULL, s STRUCT, dep STRING", + """{ "pk": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET s.c1 = -1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_NESTED_REQUIRED_DATA_ATTRIBUTE", + parameters = Map("connector" -> reqAttrsOperationClass, "nestedAttributes" -> "[s.c1]")) + } + + test("column-update: duplicate requiredDataAttributes throws AnalysisException") { + createAndInitTableWithReqAttrs("pk,pk,salary", "pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = salary + 1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_DUPLICATE_REQUIRED_DATA_ATTRIBUTE", + parameters = Map("connector" -> reqAttrsOperationClass, "duplicateAttributes" -> "[pk]")) + } + + test("column-update: case-only duplicate requiredDataAttributes throws AnalysisException") { + createAndInitTableWithReqAttrs("PK,pk,salary", "pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = salary + 1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_DUPLICATE_REQUIRED_DATA_ATTRIBUTE", + parameters = Map("connector" -> reqAttrsOperationClass, "duplicateAttributes" -> "[PK]")) + } + + test("column-update: metadata column in requiredDataAttributes throws AnalysisException") { + createAndInitTableWithReqAttrs("pk,_partition,dep", "pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET dep = 'x' WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_METADATA_REQUIRED_DATA_ATTRIBUTE", + parameters = Map( + "connector" -> reqAttrsOperationClass, + "metadataAttributes" -> "[_partition]")) + } + + test("column-update split: metadata column in requiredDataAttributes throws " + + "AnalysisException") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-split-req-attrs", "pk,_partition,dep")) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET dep = 'x' WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_METADATA_REQUIRED_DATA_ATTRIBUTE", + parameters = Map( + "connector" -> splitReqAttrsOperationClass, + "metadataAttributes" -> "[_partition]")) + } + + test("column-update: requiredDataAttributes rejects a column that does not exist") { + createAndInitTableWithReqAttrs("pk,nonexistent_col,id", "pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_UNKNOWN_REQUIRED_DATA_ATTRIBUTE", + parameters = Map( + "connector" -> reqAttrsOperationClass, + "unknownAttributes" -> "[nonexistent_col]")) + } + + test("column-update: column update schema excludes undeclared columns") { + createAndInitTable("pk INT NOT NULL, salary INT, id INT, dep STRING", + """{ "pk": 1, "salary": 100, "id": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "id": 20, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == + Seq("pk", "dep", "salary")) + } + + test("column-update: data correctness") { + createAndInitTable("pk INT NOT NULL, salary INT, id INT, dep STRING", + """{ "pk": 1, "salary": 100, "id": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "id": 20, "dep": "software" } + |{ "pk": 3, "salary": 300, "id": 30, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 10, "hr") :: + Row(2, 200, 20, "software") :: + Row(3, -1, 30, "hr") :: Nil) + } + + private def createAndInitTableSplit(schemaString: String, jsonData: String): Unit = { + createAndInitTable(schemaString, jsonData, + java.util.Map.of("column-update-split", "true")) + } + + test("column-update split: column update schema has the declared columns") { + createAndInitTableSplit("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkLastWriteInfo( + expectedRowIdSchema = Some(StructType(Array(PK_FIELD))), + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + PK_FIELD, + DEP_FIELD, + StructField("id", IntegerType, nullable = false) + )))) + } + + test("column-update split: data correctness") { + createAndInitTableSplit("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |{ "pk": 3, "id": 3, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE dep = 'hr'") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr") :: Row(2, 2, "software") :: Row(3, -1, "hr") :: Nil) + } + + test("column-update split: row-ID reassignment on narrow write is rejected") { + // `extra` is neither declared, updated, nor referenced, so requiredDataAttributes() + // ([pk, dep, salary]) doesn't cover every table column -- the write stays genuinely narrow + // and reassigning the row ID must be rejected. + createAndInitTableSplit("pk INT NOT NULL, salary INT, dep STRING, extra STRING", + """{ "pk": 1, "salary": 100, "dep": "hr", "extra": "x" } + |{ "pk": 2, "salary": 200, "dep": "software", "extra": "y" } + |""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET pk = pk + 10, salary = -1 WHERE dep = 'hr'") + }, + condition = "COLUMN_UPDATE_SPLIT_ROW_ID_REASSIGNMENT", + parameters = Map("connector" -> splitOperationClass, "rowIds" -> "[pk]")) + } + + test("column-update split: row-ID reassignment is allowed when requiredDataAttributes " + + "covers every table column") { + createAndInitTableSplitWithReqAttrs("pk,salary,dep", + "pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET pk = pk + 10, salary = -1 WHERE dep = 'hr'") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(2, 200, "software") :: Row(11, -1, "hr") :: Nil) + } + + private def createAndInitTableSplitWithReqAttrs( + reqAttrs: String, + schemaString: String, + jsonData: String): Unit = { + createAndInitTable(schemaString, jsonData, + java.util.Map.of("column-update-split-req-attrs", reqAttrs)) + } + + test("column-update split: undeclared row-ID column on narrow write is rejected") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin, + java.util.Map.of("column-update-split-row-id", "pk")) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + }, + condition = "COLUMN_UPDATE_SPLIT_ROW_ID_NOT_DECLARED", + parameters = Map("connector" -> splitWithRowIdOperationClass, "rowIds" -> "[pk]")) + } + + // analyzes without executing, as this fixture can't commit a write that omits `pk` + private def analyzeWriteDelta(sqlText: String): WriteDelta = { + val parsed = spark.sessionState.sqlParser.parsePlan(sqlText) + spark.sessionState.executePlan(parsed).analyzed.collectFirst { + case w: WriteDelta => w + }.getOrElse(fail("couldn't find WriteDelta in analyzed plan")) + } + + test("column-update split: metadata row ID preserved on reinsert need not be declared") { + // `_partition` is a required metadata attribute with PRESERVE_ON_REINSERT, so the REINSERT + // row carries it through the metadata projection. + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-split-row-id", "_partition")) + + val write = analyzeWriteDelta(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + + assert(write.table.output.map(_.name) == Seq("salary")) + assert(write.projections.metadataProjection.get.schema.fieldNames.contains("_partition")) + } + + test("column-update split: metadata row ID not preserved on reinsert must be declared") { + // `index` is a required metadata attribute, but its REINSERT value is null. + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-split-row-id", "index")) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + }, + condition = "COLUMN_UPDATE_SPLIT_ROW_ID_NOT_DECLARED", + parameters = Map("connector" -> splitWithRowIdOperationClass, "rowIds" -> "[index]")) + } + + test("column-update split: metadata row ID outside requiredMetadataAttributes must be " + + "declared") { + // Without required metadata attributes, nothing projects `_partition` into the REINSERT row. + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-split-row-id", "_partition", "no-metadata", "true")) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + }, + condition = "COLUMN_UPDATE_SPLIT_ROW_ID_NOT_DECLARED", + parameters = Map("connector" -> splitWithRowIdOperationClass, "rowIds" -> "[_partition]")) + } + + // --------------------------------------------------------------------------- + // Write relation tests: verify that only the write relation is narrowed to the + // declared columns, for in-place and delete+insert deltas. + // --------------------------------------------------------------------------- + + // t(pk, id, dep, salary, bonus, extra) with CHECK (extra > 0), partitioned by dep + private def createAndInitTableWithProps(props: java.util.Map[String, String]): Unit = { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, salary INT, bonus INT, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "salary": 100, "bonus": 10, "extra": 1 } + |{ "pk": 2, "id": 2, "dep": "hr", "salary": 200, "bonus": 20, "extra": 1 } + |{ "pk": 3, "id": 3, "dep": "software", "salary": 300, "bonus": 30, "extra": 1 } + |""".stripMargin, + props) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_extra CHECK (extra > 0)") + } + + test("column-update: write relation has the declared columns in declared order") { + createAndInitTableWithProps(java.util.Map.of("column-update-req-attrs", "salary,pk")) + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + checkColumnUpdateWrite(write, "salary", "pk") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + + test("column-update split: write relation has the declared columns") { + createAndInitTableWithProps(java.util.Map.of("column-update-split-req-attrs", "pk,salary")) + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + checkColumnUpdateWrite(write, "pk", "salary") + assert(table.lastWriteLog.map(_.getUTF8String(0).toString).toSet == + Set(Delete.toString, Reinsert.toString)) + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + + test("column-update: scan reads only the referenced columns") { + createAndInitTableWithProps(java.util.Map.of("column-update-req-attrs", "pk,salary")) + + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + + checkLastScanSchema("pk INT, dep STRING, salary INT, bonus INT, extra INT, _partition STRING") + } + + test("column-update split: scan reads only the referenced columns") { + createAndInitTableWithProps(java.util.Map.of("column-update-split-req-attrs", "pk,salary")) + + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + + checkLastScanSchema("pk INT, dep STRING, salary INT, bonus INT, extra INT, _partition STRING") + } + + test("column-update: excluding ColumnPruning keeps the write narrow") { + createAndInitTableWithProps(java.util.Map.of("column-update-req-attrs", "pk,salary")) + + withSQLConf(SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> ColumnPruning.ruleName) { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == Seq("pk", "salary")) + checkLastScanIncludes("id") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + + test("column-update: analyzed plan shows the declared columns as the write target") { + createAndInitTableWithProps(java.util.Map.of("column-update-req-attrs", "pk,salary")) + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + val header = write.simpleString(Int.MaxValue) + assert(header.matches("""WriteDelta RelationV2\[pk#\d+, salary#\d+\] .*"""), header) + } + + test("column-update: other operations keep a full-width write relation") { + createAndInitTableWithProps(java.util.Map.of("supports-deltas", "true")) + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + } + + assert(write.table.output == write.originalTable.output) + } + + test("column-update: empty requiredDataAttributes is rejected for an all-identity UPDATE") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" }""".stripMargin, + java.util.Map.of("column-update-empty-req-attrs", "true")) + + // Every assignment is identity, so only the non-empty check can reject the declaration. + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET id = id WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_EMPTY_REQUIRED_DATA_ATTRIBUTES", + parameters = Map("connector" -> reqAttrsOperationClass)) + } + + test("column-update: analysis fails when assignment key is outside requiredDataAttributes") { + createAndInitTableWithReqAttrs("pk", "pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" }""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_REQUIRED_DATA_ATTRIBUTES_MISSING_UPDATED_COLUMNS", + parameters = Map("connector" -> reqAttrsOperationClass, "missingColumns" -> "[id]")) + } + + test("column-update: a column assigned from another column must be declared") { + createAndInitTableWithReqAttrs("pk,dep", "pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 10, "dep": "hr" }""".stripMargin) + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET id = pk WHERE pk = 1") + }, + condition = "COLUMN_UPDATE_REQUIRED_DATA_ATTRIBUTES_MISSING_UPDATED_COLUMNS", + parameters = Map("connector" -> reqAttrsOperationClass, "missingColumns" -> "[id]")) + } + + test("column-update: CHECK constraint on an undeclared column referenced by the condition") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "extra": 5 } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_extra CHECK (extra > 0)") + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE extra > 3") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr", 5) :: Nil) + } + + test("column-update: CHECK constraint on an undeclared column referenced by the assignment") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "extra": 5 } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_extra CHECK (extra > 0)") + + sql(s"UPDATE $tableNameAsString SET id = extra + 1 WHERE pk = 1") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 6, "hr", 5) :: Nil) + } + + test("column-update: CHECK constraint is still enforced on a narrow delta write") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "extra": 5 } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_id CHECK (id > 0)") + + val ex = intercept[SparkRuntimeException] { + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + } + assert(ex.getCondition == "CHECK_CONSTRAINT_VIOLATION", + s"expected CHECK_CONSTRAINT_VIOLATION but got: ${ex.getCondition}") + } + + test("column-update: reassigning a row-ID column that has a CHECK constraint") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_pk CHECK (pk > 0)") + + sql(s"UPDATE $tableNameAsString SET pk = pk + 10, salary = -1 WHERE dep = 'hr'") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(2, 200, "software") :: Row(11, -1, "hr") :: Nil) + } + + test("column-update: CHECK constraint on a reassigned row-ID column is enforced") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_pk CHECK (pk > 0)") + + val ex = intercept[SparkRuntimeException] { + sql(s"UPDATE $tableNameAsString SET pk = -pk WHERE dep = 'hr'") + } + assert(ex.getCondition == "CHECK_CONSTRAINT_VIOLATION", + s"expected CHECK_CONSTRAINT_VIOLATION but got: ${ex.getCondition}") + } + + test("column-update: CHECK constraint column survives when a row-ID column is reassigned") { + createAndInitTable("pk INT NOT NULL, salary INT, extra INT, dep STRING", + """{ "pk": 1, "salary": 100, "extra": 5, "dep": "hr" } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_extra CHECK (extra > 0)") + + sql(s"UPDATE $tableNameAsString SET pk = pk + 10, salary = -1 WHERE dep = 'hr'") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(11, -1, 5, "hr") :: Nil) + } + + test("column-update: CHECK constraint on a nested field narrows correctly") { + createAndInitTable("pk INT NOT NULL, id INT, s STRUCT, dep STRING", + """{ "pk": 1, "id": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_c1 CHECK (s.c1 > 0)") + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, -1, Row(1, 2), "hr") :: Nil) + + val ex = intercept[SparkRuntimeException] { + sql(s"UPDATE $tableNameAsString SET s.c1 = -1 WHERE pk = 1") + } + assert(ex.getCondition == "CHECK_CONSTRAINT_VIOLATION", + s"expected CHECK_CONSTRAINT_VIOLATION but got: ${ex.getCondition}") + } + + test("column-update: dynamic options reach the connector") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + checkRowLevelOperationOptions( + sql(s"UPDATE $tableNameAsString WITH " + + s"(`load-option` = 'load-value', `write-option` = 'write-value') " + + s"SET id = -1 WHERE pk = 1"), + "load-option" -> "load-value", + "write-option" -> "write-value") + } +} diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedUpdateTableSuiteBase.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedUpdateTableSuiteBase.scala index 49e586535a0d0..a69fe3566424a 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedUpdateTableSuiteBase.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/DeltaBasedUpdateTableSuiteBase.scala @@ -23,6 +23,68 @@ abstract class DeltaBasedUpdateTableSuiteBase extends UpdateTableSuiteBase { override protected def deltaUpdate: Boolean = true + test("SPARK-58111: updatedColumns: single non-identity assignment") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkLastUpdatedColumns("id") + } + + test("SPARK-58111: updatedColumns: multiple non-identity assignments") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1, dep = 'eng' WHERE pk = 1") + + checkLastUpdatedColumns("id", "dep") + } + + test("SPARK-58111: updatedColumns: identity assignments are excluded") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = -1, dep = dep WHERE pk = 1") + + checkLastUpdatedColumns("id") + } + + test("SPARK-58111: updatedColumns: empty when all assignments are identity") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = id, dep = dep WHERE pk = 1") + + checkLastUpdatedColumns() + } + + test("SPARK-58111: updatedColumns: no WHERE clause still reports assigned columns") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |{ "pk": 2, "id": 2, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET dep = 'eng'") + + checkLastUpdatedColumns("dep") + } + + test("SPARK-58111: updatedColumns: assignment from another column is not identity") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 10, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET id = pk WHERE pk = 1") + + checkLastUpdatedColumns("id") + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), Row(1, 1, "hr") :: Nil) + } + test("nullable row ID attrs") { createAndInitTable("pk INT, salary INT, dep STRING", """{ "pk": 1, "salary": 300, "dep": 'hr' } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/GroupBasedColumnUpdateTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/GroupBasedColumnUpdateTableSuite.scala new file mode 100644 index 0000000000000..e90262a831618 --- /dev/null +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/GroupBasedColumnUpdateTableSuite.scala @@ -0,0 +1,754 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.connector + +import org.apache.spark.{SparkRuntimeException, SparkUnsupportedOperationException} +import org.apache.spark.sql.{AnalysisException, Row} +import org.apache.spark.sql.catalyst.optimizer.ColumnPruning +import org.apache.spark.sql.catalyst.plans.logical.Union +import org.apache.spark.sql.connector.catalog.{InMemoryBaseTable, InMemoryRowLevelOperationTable} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{IntegerType, StructField, StructType} + +class GroupBasedColumnUpdateTableSuite extends UpdateTableSuiteBase { + + override protected lazy val extraTableProps: java.util.Map[String, String] = { + val props = new java.util.HashMap[String, String]() + props.put("column-update-cow", "true") + props + } + + test("column-update ReplaceData: column update schema has only the declared columns") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == + Seq("pk", "dep", "salary")) + } + + test("column-update ReplaceData: data correctness -- bonus preserved, salary updated") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "software" } + |{ "pk": 3, "salary": 300, "bonus": 30, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE dep = 'hr'") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 10, "hr") :: + Row(2, 200, 20, "software") :: + Row(3, -1, 30, "hr") :: Nil) + // CoW: pk=1 and pk=3 match (2 updates); the 'hr' partition has no other rows so no copies. + checkUpdateMetrics(numUpdatedRows = 2, numCopiedRows = 0) + } + + test("column-update ReplaceData: subquery WHERE condition data correctness") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "software" } + |{ "pk": 3, "salary": 300, "bonus": 30, "dep": "hr" } + |""".stripMargin) + + import testImplicits._ + val subqueryDF = Seq("hr").toDF() + subqueryDF.createOrReplaceTempView("target_deps") + + sql( + s"""UPDATE $tableNameAsString + |SET salary = -1 + |WHERE dep IN (SELECT * FROM target_deps) + |""".stripMargin) + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 10, "hr") :: + Row(2, 200, 20, "software") :: + Row(3, -1, 30, "hr") :: Nil) + checkUpdateMetrics(numUpdatedRows = 2, numCopiedRows = 0) + } + + test("column-update ReplaceData: runtime group filtering data correctness") { + Seq(true, false).foreach { dppEnabled => + Seq(true, false).foreach { aqeEnabled => + withSQLConf( + SQLConf.DYNAMIC_PARTITION_PRUNING_ENABLED.key -> dppEnabled.toString, + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqeEnabled.toString) { + withTable(tableNameAsString) { + withTempView("matched_pk") { + createAndInitTable("pk INT, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "software" } + |{ "pk": 3, "salary": 300, "bonus": 30, "dep": "hr" } + |""".stripMargin) + + import testImplicits._ + val matchedPkDF = Seq(Some(1), None).toDF() + matchedPkDF.createOrReplaceTempView("matched_pk") + + executeAndCheckScans( + s"UPDATE $tableNameAsString SET salary = -1 " + + s"WHERE pk IN (SELECT * FROM matched_pk)", + primaryScanSchema = + "pk INT, salary INT, dep STRING, _partition STRING, index INT", + groupFilterScanSchema = Some("pk INT, dep STRING")) + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 10, "hr") :: + Row(2, 200, 20, "software") :: + Row(3, 300, 30, "hr") :: Nil) + + checkReplacedPartitions(Seq("hr")) + } + } + } + } + } + } + + test("column-update ReplaceData: updated and copied rows reach writeColumnUpdate with " + + "metadata") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "hr" } + |{ "pk": 3, "salary": 300, "bonus": 30, "dep": "software" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + checkLastColumnUpdateWriteLog( + writeColumnUpdateWithMetadataLogEntry(metadata = Row("hr", null), data = Row(1, "hr", -1)), + writeColumnUpdateWithMetadataLogEntry(metadata = Row("hr", 1), data = Row(2, "hr", 200))) + } + + test("column-update ReplaceData: without metadata, rows reach writeColumnUpdate(record)") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "hr" } + |{ "pk": 3, "salary": 300, "bonus": 30, "dep": "software" } + |""".stripMargin, + java.util.Map.of("column-update-cow", "true", "no-metadata", "true")) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + checkLastColumnUpdateWriteLog( + writeColumnUpdateLogEntry(data = Row(1, "hr", -1)), + writeColumnUpdateLogEntry(data = Row(2, "hr", 200))) + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 10, "hr") :: Row(2, 200, 20, "hr") :: Row(3, 300, 30, "software") :: Nil) + } + + test("column-update ReplaceData: a writer overriding only writeColumnUpdate(record) " + + "receives rows with metadata") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-cow", "true", "column-update-cow-record-only-writer", "true")) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + checkLastColumnUpdateWriteLog( + writeColumnUpdateLogEntry(data = Row(1, "hr", -1)), + writeColumnUpdateLogEntry(data = Row(2, "hr", 200))) + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 10, "hr") :: Row(2, 200, 20, "hr") :: Nil) + } + + test("column-update ReplaceData: without metadata, column update schema is the declared " + + "layout") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-cow", "true", "no-metadata", "true")) + + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE pk = 1") + + checkLastWriteInfo( + expectedColumnUpdateSchema = Some(StructType(Seq(PK_FIELD, DEP_FIELD, + StructField("salary", IntegerType, nullable = true))))) + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 110, 10, "hr") :: Row(2, 200, 20, "hr") :: Nil) + } + + test("column-update ReplaceData: connector missing writeColumnUpdate override is rejected") { + createAndInitTableReplaceDataNoWriteColumnUpdate("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + + checkError( + exception = intercept[SparkUnsupportedOperationException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + }, + condition = "DATA_SOURCE_WRITE_COLUMN_UPDATE_NOT_IMPLEMENTED", + parameters = Map( + "class" -> "org.apache.spark.sql.connector.catalog.NoWriteColumnUpdateOverrideWriter")) + } + + private def createAndInitTableReplaceDataNoWriteColumnUpdate( + schemaString: String, jsonData: String): Unit = { + createAndInitTable(schemaString, jsonData, + java.util.Map.of("column-update-cow-no-write-column-update", "true")) + } + + // --------------------------------------------------------------------------- + // Write relation tests: verify that only the write relation is narrowed to the + // declared columns, on both copy-on-write plan shapes. + // --------------------------------------------------------------------------- + + private val reqAttrsOperationClass = + classOf[InMemoryRowLevelOperationTable#PartitionBasedColumnUpdateOperationWithReqAttrs].getName + + // t(pk, id, dep, salary, bonus, extra) with CHECK (extra > 0), partitioned by dep + private def createAndInitTableWithReqAttrs(reqAttrs: String): Unit = { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, salary INT, bonus INT, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "salary": 100, "bonus": 10, "extra": 1 } + |{ "pk": 2, "id": 2, "dep": "hr", "salary": 200, "bonus": 20, "extra": 1 } + |{ "pk": 3, "id": 3, "dep": "software", "salary": 300, "bonus": 30, "extra": 1 } + |""".stripMargin, + java.util.Map.of("column-update-cow-req-attrs", reqAttrs)) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_extra CHECK (extra > 0)") + } + + test("column-update ReplaceData: write relation has the declared columns in declared order") { + createAndInitTableWithReqAttrs("salary,pk") + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + assert(!write.query.exists(_.isInstanceOf[Union])) + checkColumnUpdateWrite(write, "salary", "pk") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + + test("column-update ReplaceData: subquery write relation has the declared columns") { + withTempView("deps") { + createAndInitTableWithReqAttrs("pk,salary") + import testImplicits._ + Seq("hr").toDF("dep").createOrReplaceTempView("deps") + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus " + + s"WHERE dep IN (SELECT dep FROM deps)") + } + + assert(write.query.exists(_.isInstanceOf[Union])) + checkColumnUpdateWrite(write, "pk", "salary") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + } + + test("column-update ReplaceData: runtime group filtering on an undeclared condition column") { + withTempView("matched_ids") { + // `dep` is declared because the scan filters groups at runtime by the partition column + createAndInitTableWithReqAttrs("pk,dep,salary") + import testImplicits._ + Seq(1, 100).toDF().createOrReplaceTempView("matched_ids") + + // `id` is neither declared nor in the write relation; the group filter must still map it + // onto the second read of the table that the subquery rewrite adds + executeAndCheckScans( + s"UPDATE $tableNameAsString SET salary = salary + bonus " + + s"WHERE id IN (SELECT * FROM matched_ids)", + primaryScanSchema = "pk INT, id INT, dep STRING, salary INT, bonus INT, extra INT, " + + "_partition STRING, index INT", + groupFilterScanSchema = Some("id INT, dep STRING")) + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 200, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + checkReplacedPartitions(Seq("hr")) + } + } + + test("column-update ReplaceData: scan reads only the referenced columns") { + createAndInitTableWithReqAttrs("pk,salary") + + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + + // `id` is neither declared nor referenced; `dep`, `bonus` and `extra` are read for the + // condition, the assignment and the CHECK constraint + checkLastScanSchema( + "pk INT, dep STRING, salary INT, bonus INT, extra INT, _partition STRING, index INT") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + + test("column-update ReplaceData: subquery scan reads only the referenced columns") { + withTempView("deps") { + createAndInitTableWithReqAttrs("pk,salary") + import testImplicits._ + Seq("hr").toDF("dep").createOrReplaceTempView("deps") + + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus " + + s"WHERE dep IN (SELECT dep FROM deps)") + + checkLastScanSchema( + "pk INT, dep STRING, salary INT, bonus INT, extra INT, _partition STRING, index INT") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + } + + test("column-update ReplaceData: unread partition column skips runtime group filtering") { + withTempView("matched_ids") { + createAndInitTableWithReqAttrs("pk,salary") + import testImplicits._ + Seq(1, 100).toDF().createOrReplaceTempView("matched_ids") + + // `dep` is neither declared nor referenced, so the scan doesn't read it and doesn't + // report it from filterAttributes(); every group is read and replaced + executeAndCheckScans( + s"UPDATE $tableNameAsString SET salary = salary + bonus " + + s"WHERE id IN (SELECT * FROM matched_ids)", + primaryScanSchema = "pk INT, id INT, salary INT, bonus INT, extra INT, " + + "_partition STRING, index INT", + groupFilterScanSchema = None) + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 200, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + checkReplacedPartitions(Seq("hr", "software")) + } + } + + test("column-update ReplaceData: runtime filter attribute the scan doesn't read is rejected") { + withTempView("matched_ids") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, salary INT", + """{ "pk": 1, "id": 1, "dep": "hr", "salary": 100 } + |""".stripMargin, + java.util.Map.of( + "column-update-cow-req-attrs", "pk,salary", + "column-update-cow-unread-filter-attrs", "true")) + import testImplicits._ + Seq(1).toDF().createOrReplaceTempView("matched_ids") + + val ex = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE id IN (SELECT * FROM matched_ids)") + } + assert(ex.getCondition == "DATA_SOURCE_INVALID_RUNTIME_FILTER_ATTRIBUTE.CANNOT_RESOLVE") + assert(ex.getMessageParameters.get("attribute") == "`dep`") + assert(ex.getMessageParameters.get("method") == "filterAttributes()") + } + } + + test("column-update ReplaceData: a connector without column updates reads every column") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of[String, String]()) + + sql(s"UPDATE $tableNameAsString SET salary = 0") + + checkLastScanSchema("pk INT, salary INT, dep STRING, _partition STRING, index INT") + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), Row(1, 0, "hr") :: Nil) + } + + test("column-update ReplaceData: undeclared distribution column is rejected") { + // without metadata, the write clusters by `dep` + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-cow-req-attrs", "pk,salary", "no-metadata", "true")) + + val excludeColumnPruning = SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> ColumnPruning.ruleName + for { + cond <- Seq("pk = 1", "dep = 'hr'") + confs <- Seq(Nil, Seq(excludeColumnPruning)) + } withSQLConf(confs: _*) { + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE $cond") + }, + condition = "COLUMN_UPDATE_UNDECLARED_WRITE_REQUIREMENT_COLUMNS", + parameters = Map("connector" -> reqAttrsOperationClass, "columns" -> "[dep]")) + } + } + + test("column-update ReplaceData: excluding ColumnPruning keeps the write narrow") { + createAndInitTableWithReqAttrs("pk,salary") + + withSQLConf(SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> ColumnPruning.ruleName) { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == Seq("pk", "salary")) + checkLastScanSchema("pk INT, id INT, dep STRING, salary INT, bonus INT, extra INT, " + + "_partition STRING, index INT") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 1, "hr", 110, 10, 1) :: + Row(2, 2, "hr", 220, 20, 1) :: + Row(3, 3, "software", 300, 30, 1) :: Nil) + } + + test("column-update ReplaceData: other operations read every column") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "hr" } + |""".stripMargin, + java.util.Map.of()) + val fullScanSchema = "pk INT, salary INT, bonus INT, dep STRING, _partition STRING, index INT" + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + checkLastScanSchema(fullScanSchema) + + sql(s"DELETE FROM $tableNameAsString WHERE pk > 1") + checkLastScanSchema(fullScanSchema) + + withTempView("source") { + import testImplicits._ + Seq((1, 150)).toDF("pk", "salary").createOrReplaceTempView("source") + sql( + s"""MERGE INTO $tableNameAsString t + |USING source s + |ON t.pk = s.pk + |WHEN MATCHED THEN UPDATE SET salary = s.salary + |""".stripMargin) + checkLastScanSchema(fullScanSchema) + } + + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), Row(1, 150, 10, "hr") :: Nil) + } + + test("column-update ReplaceData: DELETE and MERGE on a column-update table read every column") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "dep": "hr" } + |""".stripMargin) + val fullScanSchema = "pk INT, salary INT, bonus INT, dep STRING, _partition STRING, index INT" + + sql(s"DELETE FROM $tableNameAsString WHERE pk > 1") + checkLastScanSchema(fullScanSchema) + + withTempView("source") { + import testImplicits._ + Seq((1, 150)).toDF("pk", "salary").createOrReplaceTempView("source") + sql( + s"""MERGE INTO $tableNameAsString t + |USING source s + |ON t.pk = s.pk + |WHEN MATCHED THEN UPDATE SET salary = s.salary + |""".stripMargin) + checkLastScanSchema(fullScanSchema) + } + + checkAnswer(sql(s"SELECT * FROM $tableNameAsString"), Row(1, 150, 10, "hr") :: Nil) + } + + test("column-update ReplaceData: assigned column outside requiredDataAttributes is rejected") { + createAndInitTableWithReqAttrs("pk,salary") + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET bonus = 0 WHERE dep = 'hr'") + }, + condition = "COLUMN_UPDATE_REQUIRED_DATA_ATTRIBUTES_MISSING_UPDATED_COLUMNS", + parameters = Map( + "connector" -> reqAttrsOperationClass, + "missingColumns" -> "[bonus]")) + } + + test("column-update ReplaceData: assigned column outside requiredDataAttributes is rejected " + + "on the subquery path") { + withTempView("deps") { + createAndInitTableWithReqAttrs("pk,salary") + import testImplicits._ + Seq("hr").toDF("dep").createOrReplaceTempView("deps") + + checkError( + exception = intercept[AnalysisException] { + sql(s"UPDATE $tableNameAsString SET bonus = 0 WHERE dep IN (SELECT dep FROM deps)") + }, + condition = "COLUMN_UPDATE_REQUIRED_DATA_ATTRIBUTES_MISSING_UPDATED_COLUMNS", + parameters = Map( + "connector" -> reqAttrsOperationClass, + "missingColumns" -> "[bonus]")) + } + } + + test("column-update ReplaceData: analyzed plan shows the declared columns as the write target") { + createAndInitTableWithReqAttrs("pk,salary") + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = salary + bonus WHERE dep = 'hr'") + } + + val header = write.simpleString(Int.MaxValue) + assert(header.matches("""ReplaceData RelationV2\[pk#\d+, salary#\d+\] .*"""), header) + } + + test("column-update ReplaceData: column update schema nullability covers UPDATE and COPY " + + "rows") { + createAndInitTable("pk INT NOT NULL, salary INT NOT NULL, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": null, "dep": "hr" } + |""".stripMargin, + java.util.Map.of("column-update-cow-req-attrs", "pk,salary,bonus")) + + sql(s"UPDATE $tableNameAsString SET salary = -1, bonus = 5 WHERE pk = 1") + + // UPDATE rows assign non-null literals, but the COPY row for pk = 2 carries a null bonus. + checkLastWriteInfo( + expectedMetadataSchema = Some(StructType(Array(PARTITION_FIELD, INDEX_FIELD_NULLABLE))), + expectedColumnUpdateSchema = Some(StructType(Seq( + StructField("pk", IntegerType, nullable = false), + StructField("salary", IntegerType, nullable = false), + StructField("bonus", IntegerType, nullable = true))))) + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, 5, "hr") :: Row(2, 200, null, "hr") :: Nil) + } + + test("column-update ReplaceData: other operations keep a full-width write relation") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |""".stripMargin, + java.util.Map.of()) + + val write = executeAndKeepAnalyzedWrite { + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + } + + assert(write.table.output == write.originalTable.output) + } + + + test("column-update ReplaceData: nested struct field update narrows to root struct column") { + createAndInitTable("pk INT NOT NULL, s STRUCT, dep STRING", + """{ "pk": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |{ "pk": 2, "s": { "c1": 3, "c2": 4 }, "dep": "hr" } + |""".stripMargin) + + // `SET s.c1 = -1` is aligned into a whole-struct assignment, so updatedColumns is [s] + // at root granularity. + sql(s"UPDATE $tableNameAsString SET s.c1 = -1 WHERE pk = 1") + + checkLastUpdatedColumns("s") + + val columnUpdateSchema = table.lastWriteInfo.columnUpdateSchema().get() + assert(columnUpdateSchema.fieldNames.contains("s"), + s"s must be in column update schema: $columnUpdateSchema") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, Row(-1, 2), "hr") :: Row(2, Row(3, 4), "hr") :: Nil) + } + + test("column-update ReplaceData: nested field identity update reports root struct as updated") { + createAndInitTable("pk INT NOT NULL, s STRUCT, dep STRING", + """{ "pk": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |""".stripMargin) + + // `SET s.c1 = s.c1` is semantically a no-op on the nested field, but the analyzer's + // AssignmentUtils rewrites nested field assignments into whole-struct rebuilds: + // Assignment(s, named_struct("c1", s.c1, "c2", s.c2)) + // `isIdentityAssignment` operates at root-column granularity and does NOT recognize the + // rebuilt struct as identity (its value is a `named_struct(...)`, not a plain Attribute + // matching the key). So `s` is reported in updatedColumns even though the values don't + // change. This documents the root-column granularity of `RowLevelOperationInfo` for the + // CoW path. + sql(s"UPDATE $tableNameAsString SET s.c1 = s.c1 WHERE pk = 1") + + checkLastUpdatedColumns("s") + + // Data correctness: the struct is rewritten but with equal values, so rows are unchanged. + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, Row(1, 2), "hr") :: Nil) + } + + test("column-update ReplaceData: scan reads an unreferenced CHECK constraint column") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "extra": 5 } + |{ "pk": 2, "id": 2, "dep": "software", "extra": 5 } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_extra CHECK (extra > 0)") + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + + checkLastScanIncludes("extra") + + val columnUpdateSchema = table.lastWriteInfo.columnUpdateSchema().get() + assert(!columnUpdateSchema.fieldNames.contains("extra"), + s"extra must not be in column update schema: $columnUpdateSchema") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, -1, "hr", 5) :: Row(2, 2, "software", 5) :: Nil) + } + + test("column-update ReplaceData: CHECK constraint is still enforced on a narrow write") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING, extra INT", + """{ "pk": 1, "id": 1, "dep": "hr", "extra": 5 } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_id CHECK (id > 0)") + + val ex = intercept[SparkRuntimeException] { + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + } + assert(ex.getCondition == "CHECK_CONSTRAINT_VIOLATION", + s"expected CHECK_CONSTRAINT_VIOLATION but got: ${ex.getCondition}") + } + + test("column-update ReplaceData: CHECK constraint on a nested field narrows correctly") { + createAndInitTable("pk INT NOT NULL, id INT, s STRUCT, dep STRING", + """{ "pk": 1, "id": 1, "s": { "c1": 1, "c2": 2 }, "dep": "hr" } + |""".stripMargin) + sql(s"ALTER TABLE $tableNameAsString ADD CONSTRAINT positive_c1 CHECK (s.c1 > 0)") + + sql(s"UPDATE $tableNameAsString SET id = -1 WHERE pk = 1") + checkLastScanIncludes("s") + + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, -1, Row(1, 2), "hr") :: Nil) + + val ex = intercept[SparkRuntimeException] { + sql(s"UPDATE $tableNameAsString SET s.c1 = -1 WHERE pk = 1") + } + assert(ex.getCondition == "CHECK_CONSTRAINT_VIOLATION", + s"expected CHECK_CONSTRAINT_VIOLATION but got: ${ex.getCondition}") + } + + + test("column-update ReplaceData: undeclared, differently-cased columns in the RHS and WHERE " + + "clause are read with the table's own spelling") { + createAndInitTable( + "pk INT NOT NULL, salary INT, bonus INT, extra INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "extra": 5, "dep": "hr" } + |{ "pk": 2, "salary": 200, "bonus": 20, "extra": 5, "dep": "hr" } + |""".stripMargin) + + sql(s"UPDATE $tableNameAsString SET salary = salary + BONUS WHERE pk = 1 AND EXTRA > 3") + + checkLastScanSchema( + "pk INT, salary INT, bonus INT, extra INT, dep STRING, _partition STRING, index INT") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 110, 10, 5, "hr") :: Row(2, 200, 20, 5, "hr") :: Nil) + } + + test("column-update ReplaceData: scan that returns unrequested columns still writes narrow " + + "rows") { + createAndInitTable("pk INT NOT NULL, salary INT, bonus INT, dep STRING", + """{ "pk": 1, "salary": 100, "bonus": 10, "dep": "hr" } + |""".stripMargin, + java.util.Map.of( + "column-update-cow", "true", + InMemoryBaseTable.SIMULATE_PARTIAL_COLUMN_PRUNING, "true")) + + sql(s"UPDATE $tableNameAsString SET salary = -1 WHERE pk = 1") + + checkLastScanIncludes("bonus") + val columnUpdateSchema = table.lastWriteInfo.columnUpdateSchema().get() + assert(!columnUpdateSchema.fieldNames.contains("bonus"), + s"bonus must not be in column update schema: $columnUpdateSchema") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(1, -1, 10, "hr") :: Nil) + } + + test("column-update ReplaceData: row-level DELETE falls back to full-width rows") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "hr" } + |""".stripMargin) + + sql(s"DELETE FROM $tableNameAsString WHERE pk < 2") + + assert(!table.lastWriteInfo.columnUpdateSchema().isPresent, + s"DELETE must not report a column update schema: " + + s"${table.lastWriteInfo.columnUpdateSchema()}") + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString"), + Row(2, 200, "hr") :: Nil) + } + + test("column-update ReplaceData: MERGE falls back to full-width rows") { + createAndInitTable("pk INT NOT NULL, salary INT, dep STRING", + """{ "pk": 1, "salary": 100, "dep": "hr" } + |{ "pk": 2, "salary": 200, "dep": "software" } + |""".stripMargin) + + withTempView("source") { + import testImplicits._ + Seq((1, 150, "hr"), (3, 300, "software")).toDF("pk", "salary", "dep") + .createOrReplaceTempView("source") + + sql( + s"""MERGE INTO $tableNameAsString t + |USING source s + |ON t.pk = s.pk + |WHEN MATCHED THEN UPDATE SET salary = s.salary + |WHEN NOT MATCHED THEN INSERT * + |""".stripMargin) + + assert(!table.lastWriteInfo.columnUpdateSchema().isPresent, + s"MERGE must not report a column update schema: " + + s"${table.lastWriteInfo.columnUpdateSchema()}") + checkLastUpdatedColumns() + checkAnswer( + sql(s"SELECT * FROM $tableNameAsString ORDER BY pk"), + Row(1, 150, "hr") :: Row(2, 200, "software") :: Row(3, 300, "software") :: Nil) + } + } + + test("column-update ReplaceData: dynamic options reach the connector") { + createAndInitTable("pk INT NOT NULL, id INT, dep STRING", + """{ "pk": 1, "id": 1, "dep": "hr" } + |""".stripMargin) + + checkRowLevelOperationOptions( + sql(s"UPDATE $tableNameAsString WITH " + + s"(`load-option` = 'load-value', `write-option` = 'write-value') " + + s"SET id = -1 WHERE pk = 1"), + "load-option" -> "load-value", + "write-option" -> "write-value") + } +} diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/RowLevelOperationSuiteBase.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/RowLevelOperationSuiteBase.scala index c62cc48cffa1c..c7482205f18c9 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/RowLevelOperationSuiteBase.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/RowLevelOperationSuiteBase.scala @@ -27,10 +27,10 @@ import org.apache.spark.sql.{DataFrame, Encoders, Row} import org.apache.spark.sql.QueryTest.{sameRows, withQueryExecutionsCaptured} import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder import org.apache.spark.sql.catalyst.expressions.{DynamicPruningExpression, Expression, GenericRowWithSchema} -import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, ReplaceData, WriteDelta} +import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, ReplaceData, RowLevelWrite, WriteDelta} import org.apache.spark.sql.catalyst.types.DataTypeUtils import org.apache.spark.sql.catalyst.util.METADATA_COL_ATTR_KEY -import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Column, Delete, Identifier, InMemoryRowLevelOperationTable, InMemoryRowLevelOperationTableCatalog, Insert, MetadataColumn, Operation, Reinsert, RowLevelOperationWithOptions, Table, TableInfo, Txn, TxnTable, Update, Write} +import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Column, Delete, Identifier, InMemoryRowLevelOperationTable, InMemoryRowLevelOperationTableCatalog, Insert, MetadataColumn, Operation, Reinsert, RowLevelOperationWithOptions, Table, TableInfo, Txn, TxnTable, Update, Write, WriteColumnUpdate} import org.apache.spark.sql.connector.expressions.LogicalExpressions.{identity, reference} import org.apache.spark.sql.connector.expressions.Transform import org.apache.spark.sql.connector.write.RowLevelOperationTable @@ -52,6 +52,7 @@ abstract class RowLevelOperationSuiteBase spark.conf.set("spark.sql.catalog.cat", classOf[InMemoryRowLevelOperationTableCatalog].getName) spark.conf.set( "spark.sql.catalog.cat.tableStateOptionKeys", "load-option,targetLoadOption") + InMemoryRowLevelOperationTable.resetLastScanSchema() } after { @@ -84,6 +85,7 @@ abstract class RowLevelOperationSuiteBase .putBoolean(MetadataColumn.PRESERVE_ON_UPDATE, value = false) .build()) protected final val INDEX_FIELD_NULLABLE = INDEX_FIELD.copy(nullable = true) + protected final val DEP_FIELD = StructField("dep", StringType, nullable = true) protected val namespace: Array[String] = Array("ns1") protected val ident: Identifier = Identifier.of(namespace, "test_table") @@ -141,6 +143,19 @@ abstract class RowLevelOperationSuiteBase append(schemaString, jsonData) } + protected def createAndInitTable( + schemaString: String, + jsonData: String, + tableProps: java.util.Map[String, String]): Unit = { + val tableInfo = new TableInfo.Builder() + .withColumns(CatalogV2Util.structTypeToV2Columns(StructType.fromDDL(schemaString))) + .withPartitions(Array[Transform](identity(reference(Seq("dep"))))) + .withProperties(tableProps) + .build() + catalog.createTable(ident, tableInfo) + append(schemaString, jsonData) + } + protected def append(schemaString: String, jsonData: String): Unit = { withSQLConf(SQLConf.LEGACY_RESPECT_NULLABILITY_IN_TEXT_DATASET_CONVERSION.key -> "true") { val df = toDF(jsonData, schemaString) @@ -193,6 +208,27 @@ abstract class RowLevelOperationSuiteBase }.getOrElse(fail("couldn't find row-level operation in optimized plan")) } + // executes an operation and extracts the analyzed ReplaceData or WriteDelta + protected def executeAndKeepAnalyzedWrite(func: => Unit): RowLevelWrite = { + val Seq(qe) = withQueryExecutionsCaptured(spark)(func) + qe.analyzed.collectFirst { + case w: RowLevelWrite => w + }.getOrElse(fail("couldn't find row-level operation in analyzed plan")) + } + + // checks that a column-update write narrows only the write relation to the declared columns + protected def checkColumnUpdateWrite(write: RowLevelWrite, declaredNames: String*): Unit = { + assert(write.table.output.map(_.name) == declaredNames, + s"write relation must have the declared columns: ${write.table.output}") + assert(write.originalTable.output.map(_.name) == table.schema.fieldNames.toSeq, + s"original table must have every column: ${write.originalTable.output}") + assert(table.lastWriteInfo.schema.isEmpty, + s"row schema must be empty: ${table.lastWriteInfo.schema}") + assert(table.lastWriteInfo.columnUpdateSchema.get.fieldNames.toSeq == declaredNames, + s"column update schema must have the declared columns: " + + s"${table.lastWriteInfo.columnUpdateSchema}") + } + // Asserts the given SQL options reached every V2 layer that should carry them: the target // catalog load, the rewritten DataSourceV2Relation, the RowLevelOperationInfo passed to the // operation builder, and the write builder's LogicalWriteInfo. @@ -306,16 +342,72 @@ abstract class RowLevelOperationSuiteBase protected def checkLastWriteInfo( expectedRowSchema: StructType = new StructType(), expectedRowIdSchema: Option[StructType] = None, - expectedMetadataSchema: Option[StructType] = None): Unit = { + expectedMetadataSchema: Option[StructType] = None, + expectedColumnUpdateSchema: Option[StructType] = None): Unit = { val info = table.lastWriteInfo assert(info.schema == expectedRowSchema, "row schema must match") val actualRowIdSchema = Option(info.rowIdSchema.orElse(null)) assert(actualRowIdSchema == expectedRowIdSchema, "row ID schema must match") val actualMetadataSchema = Option(info.metadataSchema.orElse(null)) assert(actualMetadataSchema == expectedMetadataSchema, "metadata schema must match") + val actualColumnUpdateSchema = Option(info.columnUpdateSchema.orElse(null)) + assert(actualColumnUpdateSchema == expectedColumnUpdateSchema, + "column update schema must match") + } + + + /** + * Asserts that the column names in RowLevelOperationInfo.updatedColumns() received by the + * last operation match exactly the expected set. Order is ignored. + */ + protected def checkLastUpdatedColumns(expectedNames: String*): Unit = { + val actual = table.lastUpdatedColumns.map(_.describe()).toSet + val expected = expectedNames.toSet + assert(actual == expected, + s"updatedColumns mismatch: expected ${expected.mkString("[", ", ", "]")} " + + s"but got ${actual.mkString("[", ", ", "]")}") + } + + /** Asserts that the last connector scan schema contains none of the given columns. */ + protected def checkLastScanExcludes(excludedNames: String*): Unit = { + val schema = table.lastScanSchema + assert(schema != null, "no scan schema was recorded") + val actual = schema.fieldNames.toSet + val leaked = excludedNames.toSet.intersect(actual) + assert(leaked.isEmpty, + s"scan should not include ${leaked.mkString("[", ", ", "]")} " + + s"but lastScanSchema=${schema.fieldNames.mkString("[", ", ", "]")}") + } + + /** Asserts that the last connector scan schema contains all of the given columns. */ + protected def checkLastScanIncludes(includedNames: String*): Unit = { + val schema = table.lastScanSchema + assert(schema != null, "no scan schema was recorded") + val actual = schema.fieldNames.toSet + val missing = includedNames.toSet.diff(actual) + assert(missing.isEmpty, + s"scan must include ${missing.mkString("[", ", ", "]")} " + + s"but lastScanSchema=${schema.fieldNames.mkString("[", ", ", "]")}") + } + + /** Asserts the exact schema of the last connector scan, ignoring nullability. */ + protected def checkLastScanSchema(expectedSchema: String): Unit = { + val schema = table.lastScanSchema + assert(schema != null, "no scan schema was recorded") + assert(DataTypeUtils.sameType(schema, StructType.fromDDL(expectedSchema)), + s"lastScanSchema=${schema.toDDL} but expected $expectedSchema") } protected def checkLastWriteLog(expectedEntries: WriteLogEntry*): Unit = { + checkWriteLog(table.schema, expectedEntries) + } + + /** Checks the last write log, reading row data in the layout of `columnUpdateSchema()`. */ + protected def checkLastColumnUpdateWriteLog(expectedEntries: WriteLogEntry*): Unit = { + checkWriteLog(table.lastWriteInfo.columnUpdateSchema.get, expectedEntries) + } + + private def checkWriteLog(dataSchema: StructType, expectedEntries: Seq[WriteLogEntry]): Unit = { val entryType = new StructType() .add(StructField("operation", StringType)) .add(StructField("id", IntegerType)) @@ -324,7 +416,7 @@ abstract class RowLevelOperationSuiteBase new StructType(Array( StructField("_partition", StringType), StructField("_index", IntegerType))))) - .add(StructField("data", table.schema)) + .add(StructField("data", dataSchema)) val expectedEntriesAsRows = expectedEntries.map { entry => new GenericRowWithSchema( @@ -354,6 +446,14 @@ abstract class RowLevelOperationSuiteBase WriteLogEntry(operation = Write, metadata = Some(metadata), data = Some(data)) } + protected def writeColumnUpdateLogEntry(data: Row): WriteLogEntry = { + WriteLogEntry(operation = WriteColumnUpdate, data = Some(data)) + } + + protected def writeColumnUpdateWithMetadataLogEntry(metadata: Row, data: Row): WriteLogEntry = { + WriteLogEntry(operation = WriteColumnUpdate, metadata = Some(metadata), data = Some(data)) + } + protected def deleteWriteLogEntry(id: Int, metadata: Row): WriteLogEntry = { WriteLogEntry(operation = Delete, id = Some(id), metadata = Some(metadata)) }