@@ -22,12 +22,13 @@ import scala.collection.mutable
2222import org .apache .spark .SparkException
2323import org .apache .spark .sql .catalyst .ProjectingInternalRow
2424import 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
2526import org .apache .spark .sql .catalyst .plans .logical .{Assignment , Expand , LogicalPlan , MergeRows , Project , Union }
2627import org .apache .spark .sql .catalyst .rules .Rule
2728import org .apache .spark .sql .catalyst .util .{ReplaceDataProjections , WriteDeltaProjections }
2829import org .apache .spark .sql .catalyst .util .RowDeltaUtils ._
2930import 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 }
3132import org .apache .spark .sql .connector .write .{RowLevelOperation , RowLevelOperationInfoImpl , RowLevelOperationTable , SupportsDelta }
3233import org .apache .spark .sql .connector .write .RowLevelOperation .Command
3334import 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