Skip to content

Commit ae635f4

Browse files
committed
[SPARK-56599][SQL] Add scan narrowing for column-level UPDATEs in DSv2
1 parent df5c833 commit ae635f4

11 files changed

Lines changed: 1479 additions & 56 deletions

File tree

‎sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/RowLevelOperation.java‎

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -105,4 +105,45 @@ default String description() {
105105
default NamedReference[] requiredMetadataAttributes() {
106106
return new NamedReference[0];
107107
}
108+
109+
110+
/**
111+
* Controls whether to send only the required data columns to the connector rather than the
112+
* full row.
113+
* <p>
114+
* When true, Spark narrows the data column schema ({@link LogicalWriteInfo#schema()}) to only
115+
* the columns declared via {@link #requiredDataAttributes()}. Metadata columns (from
116+
* {@link #requiredMetadataAttributes()}) and row ID columns (from
117+
* {@link SupportsDelta#rowId()}) are unaffected and always projected separately.
118+
* <p>
119+
* If {@link #requiredDataAttributes()} returns a non-empty array, the write schema is exactly
120+
* those columns in declared order. The connector must include all columns it wants to receive,
121+
* including the columns being updated. If {@link #requiredDataAttributes()} returns an empty
122+
* array, Spark sends only the non-identity assigned columns (heuristic path).
123+
*
124+
* @since 4.2.0
125+
*/
126+
default boolean supportsColumnUpdates() {
127+
return false;
128+
}
129+
130+
/**
131+
* Returns data column references required to perform this row-level operation.
132+
* <p>
133+
* This method is only consulted by Spark when {@link #supportsColumnUpdates()} returns
134+
* {@code true}. If {@code supportsColumnUpdates()} returns {@code false}, the returned array
135+
* is ignored and the full table row is sent (the default behavior).
136+
* <p>
137+
* When non-empty, the returned columns become the write schema in declared order.
138+
* The connector must declare all columns it wants to receive, including the columns being
139+
* updated. Use {@link RowLevelOperationInfo#updatedColumns()} to learn which columns are being
140+
* assigned, then add any extra columns needed for row lookup or routing (e.g., primary key).
141+
* <p>
142+
* When empty (the default), Spark falls back to sending only the non-identity assigned columns.
143+
*
144+
* @since 4.2.0
145+
*/
146+
default NamedReference[] requiredDataAttributes() {
147+
return new NamedReference[0];
148+
}
108149
}

‎sql/catalyst/src/main/java/org/apache/spark/sql/connector/write/RowLevelOperationInfo.java‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
package org.apache.spark.sql.connector.write;
1919

2020
import org.apache.spark.annotation.Experimental;
21+
import org.apache.spark.sql.connector.expressions.NamedReference;
2122
import org.apache.spark.sql.connector.write.RowLevelOperation.Command;
2223
import org.apache.spark.sql.util.CaseInsensitiveStringMap;
2324

@@ -37,4 +38,20 @@ public interface RowLevelOperationInfo {
3738
* Returns the row-level SQL command (e.g. DELETE, UPDATE, MERGE).
3839
*/
3940
Command command();
41+
42+
/**
43+
* Returns the columns being updated in an UPDATE statement, as non-identity assignments.
44+
*
45+
* <p>For DELETE and MERGE, returns an empty array.
46+
*
47+
* <p>Connectors can use this to decide what {@link RowLevelOperation#requiredDataAttributes()}
48+
* to declare. For instance, a connector that needs its primary key for row lookup can check
49+
* whether pk is already in the updated columns list and, if not, add it to
50+
* requiredDataAttributes().
51+
*
52+
* @since 4.2.0
53+
*/
54+
default NamedReference[] updatedColumns() {
55+
return new NamedReference[0];
56+
}
4057
}

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/RewriteRowLevelCommand.scala‎

Lines changed: 34 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -22,12 +22,13 @@ import scala.collection.mutable
2222
import org.apache.spark.SparkException
2323
import org.apache.spark.sql.catalyst.ProjectingInternalRow
2424
import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, AttributeReference, AttributeSet, Expression, ExprId, If, Literal, MetadataAttribute, NamedExpression, V2ExpressionUtils}
25+
import org.apache.spark.sql.catalyst.expressions.Literal.TrueLiteral
2526
import org.apache.spark.sql.catalyst.plans.logical.{Assignment, Expand, LogicalPlan, MergeRows, Project, Union}
2627
import org.apache.spark.sql.catalyst.rules.Rule
2728
import org.apache.spark.sql.catalyst.util.{ReplaceDataProjections, WriteDeltaProjections}
2829
import org.apache.spark.sql.catalyst.util.RowDeltaUtils._
2930
import org.apache.spark.sql.connector.catalog.SupportsRowLevelOperations
30-
import org.apache.spark.sql.connector.expressions.FieldReference
31+
import org.apache.spark.sql.connector.expressions.{FieldReference, NamedReference}
3132
import org.apache.spark.sql.connector.write.{RowLevelOperation, RowLevelOperationInfoImpl, RowLevelOperationTable, SupportsDelta}
3233
import org.apache.spark.sql.connector.write.RowLevelOperation.Command
3334
import org.apache.spark.sql.errors.QueryCompilationErrors
@@ -50,20 +51,35 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] {
5051
protected def buildOperationTable(
5152
table: SupportsRowLevelOperations,
5253
command: Command,
53-
options: CaseInsensitiveStringMap): RowLevelOperationTable = {
54-
val info = RowLevelOperationInfoImpl(command, options)
54+
options: CaseInsensitiveStringMap,
55+
updatedColumns: Seq[NamedReference] = Nil): RowLevelOperationTable = {
56+
val info = RowLevelOperationInfoImpl(command, options, updatedColumns)
5557
val operation = table.newRowLevelOperationBuilder(info).build()
5658
RowLevelOperationTable(table, operation)
5759
}
5860

61+
// Builds a DataSourceV2Relation for a row-level operation, optionally narrowing its output.
62+
//
63+
// When dataAttrs is non-empty, the relation output is narrowed to include only columns
64+
// required for a column-update write. When dataAttrs is empty, the full relation.output is
65+
// preserved.
5966
protected def buildRelationWithAttrs(
6067
relation: DataSourceV2Relation,
6168
table: RowLevelOperationTable,
6269
metadataAttrs: Seq[AttributeReference],
63-
rowIdAttrs: Seq[AttributeReference] = Nil): DataSourceV2Relation = {
64-
65-
val attrs = dedupAttrs(relation.output ++ rowIdAttrs ++ metadataAttrs)
66-
relation.copy(table = table, output = attrs)
70+
rowIdAttrs: Seq[AttributeReference] = Nil,
71+
dataAttrs: Seq[AttributeReference] = Nil,
72+
cond: Expression = TrueLiteral): DataSourceV2Relation = {
73+
74+
if (dataAttrs.nonEmpty) {
75+
val required =
76+
AttributeSet(dataAttrs) ++ AttributeSet(Seq(cond)) ++ AttributeSet(rowIdAttrs)
77+
val narrowOutput = relation.output.filter(required.contains)
78+
relation.copy(table = table, output = dedupAttrs(narrowOutput ++ rowIdAttrs ++ metadataAttrs))
79+
} else {
80+
val attrs = dedupAttrs(relation.output ++ rowIdAttrs ++ metadataAttrs)
81+
relation.copy(table = table, output = attrs)
82+
}
6783
}
6884

6985
protected def dedupAttrs(attrs: Seq[AttributeReference]): Seq[AttributeReference] = {
@@ -87,6 +103,14 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] {
87103
relation)
88104
}
89105

106+
protected def resolveRequiredDataAttrs(
107+
relation: DataSourceV2Relation,
108+
operation: RowLevelOperation): Seq[AttributeReference] = {
109+
V2ExpressionUtils.resolveRefs[AttributeReference](
110+
operation.requiredDataAttributes.toImmutableArraySeq,
111+
relation)
112+
}
113+
90114
protected def resolveRowIdAttrs(
91115
relation: DataSourceV2Relation,
92116
operation: SupportsDelta): Seq[AttributeReference] = {
@@ -211,11 +235,13 @@ trait RewriteRowLevelCommand extends Rule[LogicalPlan] {
211235
metadataAttrs: Seq[Attribute]): WriteDeltaProjections = {
212236
val outputs = extractOutputs(plan)
213237

238+
// Always produce Some(rowProjection) even for empty rowAttrs (identity-only column updates).
239+
// Physical execution calls rowProjection.project(row) unconditionally; None causes NPE.
214240
val rowProjection = if (rowAttrs.nonEmpty) {
215241
val outputsWithRow = filterOutputs(outputs, OPERATIONS_WITH_ROW)
216242
Some(newLazyProjection(plan, outputsWithRow, rowAttrs))
217243
} else {
218-
None
244+
Some(ProjectingInternalRow(StructType(Nil), Nil))
219245
}
220246

221247
val outputsWithRowId = filterOutputs(outputs, OPERATIONS_WITH_ROW_ID)

0 commit comments

Comments
 (0)