andygrove commented on code in PR #5051:
URL: https://github.com/apache/datafusion-comet/pull/5051#discussion_r3872504786


##########
spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala:
##########
@@ -398,33 +430,69 @@ object Utils extends CometTypeShim with Logging {
     }
   }
 
+  /**
+   * Whether every column in `batch` is an Arrow-backed `CometVector`, so 
[[getBatchFieldVectors]]
+   * can hand out its vectors directly. Callers that may receive batches from 
a plan they did not
+   * build (e.g. Comet's cache serializer, which Spark hands the cached plan's 
columnar output)
+   * use this to convert foreign vectors to Arrow instead of tripping the 
exception below.
+   *
+   * Stricter than what [[getBatchFieldVectors]] accepts: a 
`ConstantColumnVector` is rejected
+   * here even though that method materializes one, so such a batch takes the 
conversion path
+   * rather than being materialized column by column.
+   */
+  def isArrowBacked(batch: ColumnarBatch): Boolean =
+    (0 until batch.numCols()).forall { i =>
+      batch.column(i) match {
+        // Not every CometVector can be handed to getFieldVector: a 
CometPlainVector can wrap a
+        // LargeVarCharVector or LargeVarBinaryVector (an accelerated 
mapInArrow returning
+        // pa.large_string(), for instance), which it rejects. Answering true 
for those would
+        // send a batch down the direct write path that then fails, so check 
the vector itself
+        // and let the caller convert instead.
+        case v: CometVector => isSupportedFieldVector(v.getValueVector)
+        case _ => false
+      }
+    }
+
   def getBatchFieldVectors(
       batch: ColumnarBatch): (Seq[FieldVector], Option[DictionaryProvider]) = {
-    var provider: Option[DictionaryProvider] = None
+    val columns = getBatchFieldVectorsWithProviders(batch)
+    (columns.map(_._1), columns.flatMap(_._2).headOption)

Review Comment:
   Fixed in 70f046abf. `getBatchFieldVectors` now combines the providers 
instead of taking the first: each dictionary-encoded column's dictionary is 
looked up under the provider that column was decoded with, and the result is 
one `MapDictionaryProvider` covering the whole batch.
   
   Your repro is the regression test, `Comet in-memory cache broadcasts a batch 
whose columns have separate dictionaries`. With the fix reverted it fails with 
exactly `IllegalArgumentException: Could not find dictionary with ID 1`.
   
   On normalize versus combine: combine. Renumbering would mean rewriting each 
vector's `Field` dictionary ID, which is not settable without copying the 
vector, so a same-ID/different-dictionary clash now raises rather than 
resolving one column against another's dictionary. That case is unreachable 
today, since every provider reaching this path descends either from a single 
`NativeUtil` importer or from a stream whose IDs that importer wrote, but 
silently decoding one column with the wrong dictionary was the worse of the two 
failure modes to leave available.



##########
spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala:
##########
@@ -0,0 +1,539 @@
+/*
+ * 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.comet.execution.arrow
+
+import scala.collection.JavaConverters._
+import scala.util.control.NonFatal
+
+import org.apache.spark.TaskContext
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, 
GenericInternalRow, IsNotNull, IsNull, UnsafeProjection}
+import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch, 
SimpleMetricsCachedBatchSerializer}
+import org.apache.spark.sql.comet.util.Utils
+import org.apache.spark.sql.execution.columnar.DefaultCachedBatchSerializer
+import org.apache.spark.sql.internal.SQLConf
+import org.apache.spark.sql.types._
+import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector}
+import org.apache.spark.storage.StorageLevel
+import org.apache.spark.unsafe.types.{ByteArray, UTF8String}
+import org.apache.spark.util.io.ChunkedByteBuffer
+
+import org.apache.comet.CometArrowAllocator
+
+/**
+ * Cached batch format used when Comet writes Spark in-memory cache data.
+ *
+ * `columns` holds one compressed Arrow stream per cached column, in 
cache-schema order, produced
+ * by `Utils.serializeBatchColumns`. Storing columns separately is what lets a 
scan decode only
+ * the ones it projected; a single stream covering the whole batch would have 
to be inflated in
+ * full before any projection could be applied. The cache manager still owns 
storage and eviction;
+ * this class only changes the cached payload.
+ */
+private case class CometCachedBatch(
+    override val numRows: Int,
+    override val sizeInBytes: Long,
+    override val stats: InternalRow,
+    columns: Array[ChunkedByteBuffer])
+    extends SimpleMetricsCachedBatch
+
+/**
+ * Cache serializer that stores Comet-compatible Arrow batches in Spark's 
in-memory cache.
+ *
+ * The cached payload format is decided by the schema alone. A relation whose 
schema Comet's Arrow
+ * writer supports is stored as `CometCachedBatch`, and every other relation 
is delegated in full
+ * to Spark's `DefaultCachedBatchSerializer`. The format deliberately does not 
depend on any
+ * runtime config: `spark.sql.cache.serializer` is a static conf, so 
installing this serializer is
+ * already a per-application decision, and a relation whose format could flip 
mid-session cannot
+ * be read back reliably. `spark.comet.exec.inMemoryCache.enabled` still 
governs whether a scan
+ * over the cache runs natively, and its value at startup is what makes 
`CometDriverPlugin`
+ * install this serializer in the first place.
+ *
+ * Reads of `CometCachedBatch` keep working when the native scan is disabled, 
because Spark then
+ * reads the same cached data through the SparkToColumnar fallback path.
+ */
+class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer {
+
+  import ArrowCachedBatchSerializer.supportsSchema
+
+  private val fallback = new DefaultCachedBatchSerializer()
+
+  // Bounds and null counts per column, gathered before the batch is 
serialized: serializing
+  // clears the batch's vectors, and the per-column byte sizes that complete 
the statistics row
+  // are only known afterwards. See statsRow.
+  private def gatherColumnStats(
+      batch: ColumnarBatch,
+      attrs: Seq[Attribute]): (Array[Any], Array[Any], Array[Int]) = {
+    val numCols = attrs.length
+    val lower = new Array[Any](numCols)
+    val upper = new Array[Any](numCols)
+    val nulls = Array.fill[Int](numCols)(0)
+    val numRows = batch.numRows()
+
+    var c = 0
+    while (c < numCols) {
+      val dt = attrs(c).dataType
+      val col = batch.column(c)
+      var r = 0
+      while (r < numRows) {
+        if (col.isNullAt(r)) {
+          nulls(c) += 1
+        } else if (tracksBounds(dt)) {
+          val value = readValue(col, dt, r)
+          if (lower(c) == null || compare(dt, value, lower(c)) < 0) {
+            lower(c) = value
+          }
+          if (upper(c) == null || compare(dt, value, upper(c)) > 0) {
+            upper(c) = value
+          }
+        }
+        r += 1
+      }
+      c += 1
+    }
+
+    (lower, upper, nulls)
+  }
+
+  // Build the statistics row expected by SimpleMetricsCachedBatchSerializer.
+  // For each cached column Spark expects five values in this order:
+  // lower bound, upper bound, null count, row count, and size in bytes.
+  private def statsRow(
+      lower: Array[Any],
+      upper: Array[Any],
+      nulls: Array[Int],
+      numRows: Int,
+      columnSizes: Array[Long]): InternalRow = {
+    val numCols = lower.length
+    val values = new Array[Any](numCols * 5)
+    var c = 0
+    while (c < numCols) {
+      val base = c * 5
+      values(base) = lower(c)
+      values(base + 1) = upper(c)
+      values(base + 2) = nulls(c)
+      values(base + 3) = numRows
+      // Each column is its own compressed stream, so its size is known 
exactly. Cache pruning
+      // uses bounds/null-count/row-count rather than this field, but Spark 
reserves it and
+      // reports it, so record the real value.
+      values(base + 4) = columnSizes(c)
+      c += 1
+    }
+
+    new GenericInternalRow(values)
+  }
+
+  // Spark can prune cache batches only for types whose bounds can be compared.
+  // Other types still report null count and row count but leave bounds as 
null.
+  private def tracksBounds(dt: DataType): Boolean = dt match {
+    case BooleanType | ByteType | ShortType | IntegerType | LongType | 
FloatType | DoubleType |
+        _: DecimalType | StringType | DateType | TimestampType | 
TimestampNTZType =>
+      true
+    case _ => false
+  }
+
+  // Read a non-null value from a ColumnVector using Spark's internal value 
type
+  // for the corresponding DataType.
+  private def readValue(col: ColumnVector, dt: DataType, rowId: Int): Any = dt 
match {
+    case BooleanType => col.getBoolean(rowId)
+    case ByteType => col.getByte(rowId)
+    case ShortType => col.getShort(rowId)
+    case IntegerType | DateType => col.getInt(rowId)
+    case LongType | TimestampType | TimestampNTZType => col.getLong(rowId)
+    case FloatType => col.getFloat(rowId)
+    case DoubleType => col.getDouble(rowId)
+    case d: DecimalType => col.getDecimal(rowId, d.precision, d.scale)
+    case StringType => col.getUTF8String(rowId).copy()
+    case _ => null
+  }
+
+  // Compare values using the same physical representation used in the stats 
row.
+  private def compare(dt: DataType, left: Any, right: Any): Int = dt match {
+    case BooleanType =>
+      java.lang.Boolean.compare(left.asInstanceOf[Boolean], 
right.asInstanceOf[Boolean])
+    case ByteType =>
+      java.lang.Byte.compare(left.asInstanceOf[Byte], right.asInstanceOf[Byte])
+    case ShortType =>
+      java.lang.Short.compare(left.asInstanceOf[Short], 
right.asInstanceOf[Short])
+    case IntegerType | DateType =>
+      java.lang.Integer.compare(left.asInstanceOf[Int], 
right.asInstanceOf[Int])
+    case LongType | TimestampType | TimestampNTZType =>
+      java.lang.Long.compare(left.asInstanceOf[Long], right.asInstanceOf[Long])
+    case FloatType =>
+      java.lang.Float.compare(left.asInstanceOf[Float], 
right.asInstanceOf[Float])
+    case DoubleType =>
+      java.lang.Double.compare(left.asInstanceOf[Double], 
right.asInstanceOf[Double])
+    case _: DecimalType =>
+      left.asInstanceOf[Decimal].compare(right.asInstanceOf[Decimal])
+    case StringType =>
+      ByteArray.compareBinary(
+        left.asInstanceOf[UTF8String].getBytes,
+        right.asInstanceOf[UTF8String].getBytes)
+    case other =>
+      throw new IllegalStateException(s"compare called for unsupported type 
$other")
+  }
+
+  // Compute Spark-compatible cache stats before serializing each batch to 
Arrow.
+  // The stats are stored beside the Arrow bytes so Spark's cache filter can 
prune
+  // CometCachedBatch without decoding the batch first.
+  //
+  // A columnar input batch is not guaranteed to be Arrow-backed; see 
supportsColumnarInput for
+  // why. Batches that are not get copied into Arrow first, since 
Utils.serializeBatches only
+  // writes CometVector columns.
+  private def encodeBatches(
+      batches: Iterator[ColumnarBatch],
+      attrs: Seq[Attribute]): Iterator[CachedBatch] = {
+    val arrowSchema =
+      Utils.toArrowSchema(Utils.fromAttributes(attrs), 
CometArrowStream.NATIVE_TIMEZONE)
+
+    batches.map { batch =>
+      // Bounds and null counts are read from the input batch, which 
serializing then clears, so
+      // they have to be gathered first. The row is only assembled once the 
per-column sizes are
+      // known.
+      val (lower, upper, nulls) = gatherColumnStats(batch, attrs)
+      val numRows = batch.numRows()
+
+      val columns = if (Utils.isArrowBacked(batch)) {
+        Utils.serializeBatchColumns(batch)
+      } else {
+        val arrowBatch =
+          CometArrowConverters.columnarBatchToArrowBatch(batch, arrowSchema, 
CometArrowAllocator)
+        try Utils.serializeBatchColumns(arrowBatch)
+        finally arrowBatch.close()
+      }
+
+      val columnSizes = columns.map(_.size)
+      CometCachedBatch(
+        numRows = numRows,
+        sizeInBytes = columnSizes.sum,
+        stats = statsRow(lower, upper, nulls, numRows, columnSizes),
+        columns = columns)
+    }
+  }
+
+  // Resolve requested columns by exprId, not by name, because aliases may 
reuse names.
+  //
+  // An empty selection stays empty rather than expanding to every column. 
Spark asks for no
+  // columns when the query only needs the row count (SELECT count(*)), and 
since projection now
+  // decides what gets decoded, expanding it would turn the cheapest possible 
read into the most
+  // expensive one.
+  private def selectedIndices(
+      cacheAttributes: Seq[Attribute],
+      selectedAttributes: Seq[Attribute]): Array[Int] = {
+    val byExprId = cacheAttributes.zipWithIndex.map { case (attr, idx) =>
+      attr.exprId -> idx
+    }.toMap
+
+    selectedAttributes.map { attr =>
+      byExprId.getOrElse(
+        attr.exprId,
+        throw new IllegalStateException(
+          s"Could not resolve selected attribute ${attr.name} from cache 
attributes"))
+    }.toArray
+  }
+
+  // Spark's SimpleMetricsCachedBatchSerializer prunes a batch when the 
generated partition filter
+  // does not evaluate to true against the stats row. Bounds are only computed 
for the types
+  // tracksBounds accepts, and for every other column the lower and upper 
bounds stay null, which
+  // makes a comparison against them evaluate to null and therefore prune the 
batch. That would
+  // silently drop rows, so predicates over columns without bounds are not 
pushed down at all.
+  // Null counts and row counts are recorded for every column, so IsNull and 
IsNotNull stay safe.
+  override def buildFilter(
+      predicates: Seq[Expression],
+      cachedAttributes: Seq[Attribute]): (Int, Iterator[CachedBatch]) => 
Iterator[CachedBatch] = {
+    val prunable = cachedAttributes.collect {
+      case a if tracksBounds(a.dataType) => a.exprId
+    }.toSet
+
+    val prunablePredicates = predicates.filter {
+      case _: IsNull | _: IsNotNull => true
+      case p => p.references.forall(a => prunable.contains(a.exprId))
+    }
+
+    super.buildFilter(prunablePredicates, cachedAttributes)
+  }
+
+  // Comet's Arrow writer only handles the types listed in supportsSchema. 
Reporting false here
+  // sends the relation down the row path, where it is delegated to Spark's 
default serializer,
+  // instead of failing at cache materialization inside Utils.serializeBatches.
+  //
+  // This answer is schema-only, because attributes are all Spark gives us; it 
says nothing about
+  // the vectors. Returning true also makes InMemoryRelation strip the 
ColumnarToRow above the
+  // cached plan, so convertColumnarBatchToCachedBatch then receives whatever 
that plan produces:
+  // a Comet scan's CometVectors, but equally Spark's vectorized Parquet/ORC 
reader or a
+  // connector's own vectors. encodeBatches converts the non-Arrow ones; that 
conversion is load
+  // bearing, not defensive.
+  override def supportsColumnarInput(schema: Seq[Attribute]): Boolean = 
supportsSchema(schema)
+
+  // A relation Comet stores is always readable as columnar Arrow. Anything 
else holds
+  // DefaultCachedBatch, so defer to Spark, which only claims columnar output 
for the primitive
+  // types its ColumnAccessor.decompress path can actually decode.
+  override def supportsColumnarOutput(schema: StructType): Boolean = {
+    if (schema.fields.forall(f => 
ArrowCachedBatchSerializer.supportsType(f.dataType))) {
+      true
+    } else {
+      fallback.supportsColumnarOutput(schema)
+    }
+  }
+
+  // Columnar Comet output is stored as compressed Arrow stream bytes. Spark 
only calls this when
+  // supportsColumnarInput returned true, so the schema is known to be 
Comet-writable here.
+  override def convertColumnarBatchToCachedBatch(
+      input: RDD[ColumnarBatch],
+      schema: Seq[Attribute],
+      storageLevel: StorageLevel,
+      conf: SQLConf): RDD[CachedBatch] = {
+
+    input.mapPartitions { batches =>
+      encodeBatches(batches, schema)
+    }
+  }
+
+  override def convertCachedBatchToColumnarBatch(
+      input: RDD[CachedBatch],
+      cacheAttributes: Seq[Attribute],
+      selectedAttributes: Seq[Attribute],
+      conf: SQLConf): RDD[ColumnarBatch] = {
+    if (!supportsSchema(cacheAttributes)) {
+      return fallback.convertCachedBatchToColumnarBatch(
+        input,
+        cacheAttributes,
+        selectedAttributes,
+        conf)
+    }
+
+    val indices = selectedIndices(cacheAttributes, selectedAttributes)
+
+    input.mapPartitions { it =>
+      // A ColumnReaders closes its readers (releasing the vectors they are 
holding) only when the
+      // batch it produced has been consumed. A consumer that stops early -- 
LIMIT, take(), or a
+      // cancelled task -- leaves the readers for the batch in flight open, so 
close them on task
+      // completion. Spark's own ArrowCachedBatchSerializer registers a 
listener for the same
+      // reason.
+      //
+      // flatMap consumes each inner iterator fully before building the next, 
so at most one batch
+      // is open at a time and tracking the current one is enough. close() is 
idempotent, so
+      // closing one that already released itself is a no-op.
+      @volatile var current: ColumnReaders = null
+      Option(TaskContext.get()).foreach { tc =>
+        tc.addTaskCompletionListener[Unit] { _ =>
+          val readers = current
+          current = null
+          if (readers != null) {
+            readers.close()
+          }
+        }
+      }
+
+      it.flatMap {
+        case cb: CometCachedBatch =>
+          if (indices.isEmpty) {
+            // Nothing to decode: the row count is the whole answer, and it is 
already here.
+            Iterator.single(new ColumnarBatch(Array.empty[ColumnVector], 
cb.numRows))
+          } else {
+            val readers = new ColumnReaders(indices.map(i => cb.columns(i)), 
cb.numRows)
+            current = readers
+            readers.batches
+          }
+
+        case other =>
+          throw new IllegalStateException(
+            s"Unsupported cached batch type ${other.getClass.getName}")
+      }
+    }
+  }
+
+  // Decodes one selected column stream apiece and stitches the results back 
into a single batch.
+  //
+  // Each stream is self-contained, so the columns a scan did not select are 
never inflated. The
+  // decoded vectors stay owned by their readers: closing them releases the 
batch, which is why
+  // this yields a single-element iterator that closes on exhaustion, matching 
what
+  // ArrowReaderIterator did when the payload was one stream.
+  private class ColumnReaders(buffers: Array[ChunkedByteBuffer], numRows: Int) 
{
+    // decodeBatches opens a reader and eagerly decodes its first batch, so it 
allocates. If a
+    // later column throws, the readers already opened here are unreachable: 
the task-completion
+    // listener cannot release them because `current` is only assigned once 
this constructor
+    // returns, so they would leak off-heap for the life of the executor.
+    private val readers: Array[Iterator[ColumnarBatch]] = {
+      val opened = new Array[Iterator[ColumnarBatch]](buffers.length)
+      var i = 0
+      try {
+        while (i < buffers.length) {
+          opened(i) = Utils.decodeBatches(buffers(i), "CometCache")

Review Comment:
   Fixed in 70f046abf. `ArrowReaderIterator` now closes the reader if its eager 
first decode throws, so a column that loads its dictionary and then fails on 
the record batch releases it. I also guarded `StreamReader`'s own 
`getVectorSchemaRoot`, which allocates the root's vectors under the same 
condition: nothing holds the reader until the constructor returns, so nothing 
else can close it.
   
   Regression test: `Comet in-memory cache releases a reader whose own first 
batch fails to decode`. Instead of an allocator limit it truncates the decoded 
Arrow stream, dropping the end-of-stream marker plus part of the record batch 
body, so the dictionary loads and the read that follows runs out of input. 
Reverted, it leaks 64 bytes per attempt, matching the growth you measured; with 
the fix the allocator returns to baseline.
   
   The existing `corruptColumnStream` helper was not usable for this. A small 
column compresses to a single LZ4 block, so truncating the compressed bytes 
fails the decompressor before Arrow reads anything at all, which is the case 
that was already covered. The new helper cuts the decoded bytes and 
re-compresses.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to