From d7ae3bb2438e97c5d0ea70bf47b77413867a6c7d Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 28 Aug 2026 17:51:57 -0600 Subject: [PATCH 01/24] perf: project cached batches by buffer selection, prune on collated strings Follow-up to #5051, applying items from #5487. Replace the per-column Arrow IPC stream layout of `CometCachedBatch` with a single encapsulated IPC record batch message per cached batch, carrying no Schema message and no end-of-stream marker. The reader rebuilds the schema from the cached relation's attributes, so a wide relation no longer repeats the same schema bytes once per cached batch. Compression moves from a whole-payload Spark codec to Arrow's per-buffer IPC compression. That is what makes projection cheap: the message metadata records every buffer's offset and length in the body, so `CachedBatchIpc.readProjected` copies out only the byte ranges of the columns a scan selected and decompresses just those. This subsumes the separate "drop the schema message" item, since there is no longer a per-column stream to frame. Dictionary-encoded columns are decoded before being stored: a payload with no schema message cannot describe a dictionary encoding. The codec defaults to zstd, and lz4 is deliberately not offered. Arrow's lz4 is commons-compress's pure-Java implementation, unrelated to the JNI-accelerated lz4-java behind `spark.io.compression.codec`. Over a 200k-row six-column relation it measured 205s to write against 347ms for zstd, while also producing larger output, so no workload prefers it. zstd also beats storing batches uncompressed on both axes (347ms and 2 MiB against 1743ms and 13 MiB), because the bytes it saves cost more to copy and store than compressing them costs. Decompression is done here rather than left to `VectorLoader`, which leaks: `VectorLoader.loadBuffers` collects a field's decompressed buffers into a local list and releases them only after the whole field loads, so a buffer that fails to decompress strands every buffer of that field decompressed before it. A string column reaches this, its offsets buffer decompressing before its data buffer throws. Also track statistics bounds for collated string columns, comparing with the collation's own ordering through a new `CometTypeShim.compareStrings`. Matching the bare `StringType` object excluded collated columns, which then got null bounds and no pruning. Benchmark over a 5M-row six-column relation, keeping the cached scan native against falling back to a Spark cache scan and converting: 1.3x on a repeated scan, 1.3x on a narrow projection and 2.3x on a full projection. --- .../user-guide/latest/in-memory-cache.md | 148 +++++++ docs/source/user-guide/latest/index.rst | 1 + pom.xml | 26 ++ spark/pom.xml | 4 + .../scala/org/apache/comet/CometConf.scala | 55 ++- .../arrow/ArrowCachedBatchSerializer.scala | 245 ++++++----- .../execution/arrow/CachedBatchIpc.scala | 409 ++++++++++++++++++ .../apache/spark/sql/comet/util/Utils.scala | 39 +- .../apache/comet/shims/CometTypeShim.scala | 12 + .../apache/comet/shims/CometTypeShim.scala | 11 + .../comet/exec/CometInMemoryCacheSuite.scala | 303 ++++++++++--- .../CometInMemoryCacheBenchmark.scala | 7 +- .../arrow/CometCachedBatchHelper.scala | 259 ++++++++--- 13 files changed, 1216 insertions(+), 303 deletions(-) create mode 100644 docs/source/user-guide/latest/in-memory-cache.md create mode 100644 spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md new file mode 100644 index 00000000000..67f0b0e9755 --- /dev/null +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -0,0 +1,148 @@ + + +# In-Memory Cache + +Comet can store Spark's in-memory cache (`CACHE TABLE`, `df.cache()`, `df.persist()`) in an Arrow +format that Comet operators read directly. Without it, a cached table is stored in Spark's own +format and every scan of it has to convert each batch before Comet can continue, which shows up in +the plan as a `CometSparkColumnarToColumnar` above the cache scan. + +This feature is **experimental and disabled by default**. + +```scala +spark.conf.set("spark.comet.exec.inMemoryCache.enabled", "true") +``` + +## What changes when it is enabled + +`spark.comet.exec.inMemoryCache.enabled` is read at startup, and its value then decides whether +Comet installs its cache serializer as `spark.sql.cache.serializer`. When it is installed: + +- Cached data is stored as `CometCachedBatch` rather than Spark's `DefaultCachedBatch`. +- Cached tables are scanned by `CometInMemoryTableScan`, which feeds Comet operators directly. +- Per-batch column statistics are recorded in the layout Spark's `SimpleMetricsCachedBatchSerializer` + expects, so Spark can prune whole cached batches on a predicate before any of them is decoded. + +Relations whose schema Comet's Arrow writer cannot store — interval types, most notably — are +delegated in full to Spark's default cache format, per relation. Nothing about the format depends +on a runtime config, because `spark.sql.cache.serializer` is a static setting and a relation whose +format could change mid-session could not be read back reliably. Turning +`spark.comet.exec.inMemoryCache.enabled` off at runtime only sends cached scans back to Spark's +execution path; the cached data stays readable either way. + +## Storage format + +Each cached batch is stored as a single Arrow IPC record batch message and its body. + +The message carries **no Arrow schema**. The reader already has one: `InMemoryRelation` knows the +cached relation's attributes, and Comet maps them to exactly the Arrow fields the writer produced. +Storing a schema in every batch would repeat the same bytes once per cached batch — for a wide +relation cached in many batches, a large share of a payload that is not data. + +Compression is applied by Arrow to **each buffer separately**, rather than by wrapping the whole +payload in a Spark compression codec. That is what makes a projected read cheap: the message +metadata records every buffer's offset and length within the body, so a scan copies out only the +byte ranges belonging to the columns it selected, and only those are decompressed. A read of one +column out of six does roughly a sixth of the decompression work, and a `SELECT count(*)`, which +selects no columns at all, answers from the row count stored beside the payload without touching +it. + +Compression defaults to `zstd`, which is faster than storing cached batches uncompressed: the +bytes it saves cost more to copy and store than compressing them costs. Measured over a 200k-row, +six-column relation: + +| Codec | Materialize | Footprint | Read 1 of 6 | Read 6 of 6 | +| ------ | ----------: | --------: | ----------: | ----------: | +| `zstd` | 347 ms | 2 MiB | 52 ms | 63 ms | +| `none` | 1743 ms | 13 MiB | 74 ms | 79 ms | + +Arrow's other IPC codec, LZ4, is deliberately not offered. It is commons-compress's pure-Java +implementation and is unrelated to the JNI-accelerated lz4-java behind `spark.io.compression.codec`; +it measured three orders of magnitude slower to write than `zstd` while also producing larger +output, so no workload prefers it. + +Dictionary-encoded columns are decoded before they are stored. A payload with no schema message has +nowhere to record either that a column is dictionary encoded or the dictionary itself. + +## Configuration + +| Config | Default | Description | +| ------------------------------------------------------- | ------- | ---------------------------------------------------------------------------------------------------------------------------------------------- | +| `spark.comet.exec.inMemoryCache.enabled` | `false` | Whether to store and scan Spark's in-memory cache in Comet's format. Read at startup. | +| `spark.comet.exec.inMemoryCache.compression.codec` | `zstd` | Arrow IPC compression codec for cached data: `zstd` or `none`. Affects newly cached data only — a batch records the codec it was written with. | +| `spark.comet.exec.inMemoryCache.compression.zstd.level` | `1` | Compression level when the codec is `zstd`. Ignored otherwise. | + +## Performance + +Measured with `CometInMemoryCacheBenchmark` on a 5M-row, six-column relation (Apple M3 Ultra, +JDK 17, Spark 4.1, release build). Regenerate with: + +```sh +SPARK_GENERATE_BENCHMARK_FILES=1 \ + make benchmark-org.apache.spark.sql.benchmark.CometInMemoryCacheBenchmark +``` + +| Query shape | Spark cache scan + convert | `CometInMemoryTableScan` | Relative | +| ------------------------------ | -------------------------: | -----------------------: | -------: | +| Repeated scan (3 of 6 columns) | 156 ms | 118 ms | 1.3x | +| Selective filter | 44 ms | 39 ms | 1.1x | +| Row count only (0 of 6) | 32 ms | 28 ms | 1.1x | +| Narrow projection (1 of 6) | 50 ms | 39 ms | 1.3x | +| Full projection (6 of 6) | 316 ms | 135 ms | 2.3x | + +Read what this compares carefully. Comet execution is on in both columns, so the aggregation runs +on Comet either way and only the cache-scan boundary moves: on the left, Spark's +`InMemoryTableScanExec` feeds those same Comet operators through a `CometSparkColumnarToColumnar` +bridge; on the right, `CometInMemoryTableScan` feeds them directly. Both columns read the same +Comet-written `CometCachedBatch` — `spark.sql.cache.serializer` is static, so one session cannot +also materialize Spark's format to compare against. These numbers are therefore "keep the cached +scan native" against "fall back to a Spark cache scan and convert", not Comet against Spark +execution, and not a comparison with Spark's own cache format. + +## Kryo + +Spark serializes a cached batch with `spark.serializer` whenever the block leaves the heap: the +`_SER` storage levels, replication, cross-executor fetches, and the disk half of the default +`MEMORY_AND_DISK`. So an ordinary `df.cache()` that spills is enough to reach it. + +If you run with `spark.kryo.registrationRequired=true`, register Comet's classes: + +``` +spark.serializer=org.apache.spark.serializer.KryoSerializer +spark.kryo.registrationRequired=true +spark.kryo.registrator=org.apache.comet.CometKryoRegistrator +``` + +Comet cannot set `spark.kryo.registrator` for you the way it sets `spark.sql.cache.serializer`: +`KryoSerializer` reads it when `SparkEnv` builds the serializer, which happens before any plugin +runs. Without it, caching fails with a "Class is not registered" error that does not name this +feature. Comet's driver plugin warns at startup when it sees Kryo, `registrationRequired`, and no +registrator. + +## Limitations + +Reads that feed **Spark** operators rather than Comet ones are still slower than Spark's own cache +format, by roughly 1.7x to 2.5x depending on how wide the projection is. Those reads pay a row +conversion that Spark's format avoids with generated code over its own layout. This is why the +feature is off by default. + +Comet's serializer exists because Spark's own Arrow cache format +([SPARK-57268](https://issues.apache.org/jira/browse/SPARK-57268)) is only available from Spark +4.3, which Comet does not yet support. diff --git a/docs/source/user-guide/latest/index.rst b/docs/source/user-guide/latest/index.rst index 815e12289c7..063a5581a04 100644 --- a/docs/source/user-guide/latest/index.rst +++ b/docs/source/user-guide/latest/index.rst @@ -74,6 +74,7 @@ to read more. Understanding Comet Plans Tuning Guide Metrics Guide + In-Memory Cache PyArrow UDF Acceleration .. toctree:: diff --git a/pom.xml b/pom.xml index 11494665755..2fa63d71103 100644 --- a/pom.xml +++ b/pom.xml @@ -225,6 +225,32 @@ under the License. arrow-c-data ${arrow.version} + + org.apache.arrow + arrow-compression + ${arrow.version} + + + + org.apache.commons + commons-compress + + + com.github.luben + zstd-jni + + + io.netty + netty-common + + + com.google.code.findbugs + jsr305 + + + diff --git a/spark/pom.xml b/spark/pom.xml index a257415dd3f..d30afc75914 100644 --- a/spark/pom.xml +++ b/spark/pom.xml @@ -60,6 +60,10 @@ under the License. org.apache.arrow arrow-vector + + org.apache.arrow + arrow-compression + org.scala-lang.modules scala-collection-compat_${scala.binary.version} diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index f5b7c7c09d0..4183ffef457 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -261,23 +261,50 @@ object CometConf extends ShimCometConf { val COMET_EXEC_IN_MEMORY_CACHE_ENABLED: ConfigEntry[Boolean] = conf("spark.comet.exec.inMemoryCache.enabled") .category(CATEGORY_EXEC) - .doc( - "Whether to enable Comet native execution for in-memory cached tables. Its value at " + - "startup also decides whether CometDriverPlugin installs Comet's cache serializer, " + - "which stores cached data in Arrow format. Because spark.sql.cache.serializer is a " + - "static config, the cached format is fixed for the application, and disabling this " + - "at runtime only sends cached scans back to Spark's execution path. Relations whose " + - "schema Comet's Arrow writer does not support are always cached in Spark's default " + - "format. Each cached column is stored as its own compressed Arrow IPC stream, so a " + - "scan decodes only the columns it projected. Reads that feed Spark operators rather " + - "than Comet ones still pay a row conversion the default format avoids, and can be " + - "slower than Spark's cache. With spark.kryo.registrationRequired=true, also set " + - "spark.kryo.registrator=org.apache.comet.CometKryoRegistrator before creating the " + - "SparkContext, otherwise caching fails as soon as a block is serialized, including " + - "the disk half of the default MEMORY_AND_DISK storage level.") + .doc("Whether to enable Comet native execution for in-memory cached tables. Its value at " + + "startup also decides whether CometDriverPlugin installs Comet's cache serializer, " + + "which stores cached data in Arrow format. Because spark.sql.cache.serializer is a " + + "static config, the cached format is fixed for the application, and disabling this " + + "at runtime only sends cached scans back to Spark's execution path. Relations whose " + + "schema Comet's Arrow writer does not support are always cached in Spark's default " + + "format. Each cached batch is stored as one Arrow IPC record batch with per-buffer " + + "zstd compression, and a scan copies out only the buffers of the columns it projected, " + + "so the unselected ones are never decompressed. Reads that feed Spark operators rather " + + "than Comet ones still pay a row conversion the default format avoids, and can be " + + "slower than Spark's cache. With spark.kryo.registrationRequired=true, also set " + + "spark.kryo.registrator=org.apache.comet.CometKryoRegistrator before creating the " + + "SparkContext, otherwise caching fails as soon as a block is serialized, including " + + "the disk half of the default MEMORY_AND_DISK storage level.") .booleanConf .createWithDefault(false) + val COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC: ConfigEntry[String] = + conf("spark.comet.exec.inMemoryCache.compression.codec") + .category(CATEGORY_EXEC) + .doc( + "The Arrow IPC compression codec used when Comet's cache serializer writes cached " + + "data. Unlike spark.io.compression.codec, this compresses each Arrow buffer " + + "separately rather than the batch as a whole, which is what lets a projected scan " + + "decompress only the columns it selected. Set to none to store cached batches " + + "uncompressed, which is both slower to write and larger than zstd because the extra " + + "bytes cost more to move and store than compressing them costs. Only affects newly " + + "cached data; the codec a batch was written with is recorded in the batch itself and " + + "is what the read path uses. Arrow's lz4 is deliberately not offered: it is a " + + "pure-Java implementation, unrelated to the JNI-accelerated lz4 behind " + + "spark.io.compression.codec, and is orders of magnitude slower to write than zstd " + + "while also producing larger output.") + .stringConf + .checkValues(Set("none", "zstd")) + .createWithDefault("zstd") + + val COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL: ConfigEntry[Int] = + conf("spark.comet.exec.inMemoryCache.compression.zstd.level") + .category(CATEGORY_EXEC) + .doc("The compression level to use when Comet's cache serializer compresses cached data " + + "with zstd. Ignored for other codecs.") + .intConf + .createWithDefault(1) + val COMET_NATIVE_COLUMNAR_TO_ROW_ENABLED: ConfigEntry[Boolean] = conf(s"$COMET_EXEC_CONFIG_PREFIX.columnarToRow.native.enabled") .category(CATEGORY_EXEC) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 46a51ad8775..24eba56f737 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -22,6 +22,8 @@ package org.apache.spark.sql.comet.execution.arrow import scala.collection.JavaConverters._ import scala.util.control.NonFatal +import org.apache.arrow.vector.VectorSchemaRoot +import org.apache.arrow.vector.types.pojo.{Field, Schema} import org.apache.spark.TaskContext import org.apache.spark.rdd.RDD import org.apache.spark.sql.catalyst.InternalRow @@ -33,25 +35,27 @@ 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.spark.unsafe.types.UTF8String -import org.apache.comet.CometArrowAllocator +import org.apache.comet.{CometArrowAllocator, CometConf} +import org.apache.comet.shims.CometTypeShim +import org.apache.comet.vector.NativeUtil /** * 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. + * `bytes` is one encapsulated Arrow IPC RecordBatch message and its body, with no Schema message + * and no end-of-stream marker, produced by `CachedBatchIpc.serialize`. Compression is applied per + * Arrow buffer rather than over the payload as a whole, which is what lets a scan decompress only + * the columns it projected: the message records every buffer's offset and length, so + * `CachedBatchIpc.readProjected` copies out just the selected columns' byte ranges. 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]) + bytes: Array[Byte]) extends SimpleMetricsCachedBatch /** @@ -69,7 +73,7 @@ private case class CometCachedBatch( * 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 { +class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with CometTypeShim { import ArrowCachedBatchSerializer.supportsSchema @@ -130,9 +134,11 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { 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. + // The stored size of the column's own Arrow buffers, taken from the message's buffer + // layout, so it is exact rather than an estimate. Cache pruning uses + // bounds/null-count/row-count rather than this field, but Spark reserves it and reports it, + // so record the real value. The per-batch message framing is not attributed to any column, + // so these sum to slightly less than sizeInBytes. values(base + 4) = columnSizes(c) c += 1 } @@ -142,9 +148,15 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { // 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. + // + // Every StringType qualifies, collated ones included: bounds are recorded with the collation's + // own comparison, which is the same ordering the predicate Spark generates over that column + // uses. Matching the bare `StringType` object instead would exclude collated columns, since a + // collated StringType is not equal to the default one, and they would then get null bounds and + // no pruning at all. private def tracksBounds(dt: DataType): Boolean = dt match { case BooleanType | ByteType | ShortType | IntegerType | LongType | FloatType | DoubleType | - _: DecimalType | StringType | DateType | TimestampType | TimestampNTZType => + _: DecimalType | _: StringType | DateType | TimestampType | TimestampNTZType => true case _ => false } @@ -160,7 +172,7 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { 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 _: StringType => col.getUTF8String(rowId).copy() case _ => null } @@ -182,10 +194,8 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { 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 st: StringType => + compareStrings(left.asInstanceOf[UTF8String], right.asInstanceOf[UTF8String], st) case other => throw new IllegalStateException(s"compare called for unsupported type $other") } @@ -199,32 +209,33 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { // writes CometVector columns. private def encodeBatches( batches: Iterator[ColumnarBatch], - attrs: Seq[Attribute]): Iterator[CachedBatch] = { + attrs: Seq[Attribute], + codecName: String, + zstdLevel: Int): Iterator[CachedBatch] = { val arrowSchema = Utils.toArrowSchema(Utils.fromAttributes(attrs), CometArrowStream.NATIVE_TIMEZONE) + val codec = CachedBatchIpc.compressionCodec(codecName, zstdLevel) 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. + // Bounds and null counts are read from the input batch before it is serialized, and the row + // is only assembled once the per-column sizes the message reports are known. val (lower, upper, nulls) = gatherColumnStats(batch, attrs) val numRows = batch.numRows() - val columns = if (Utils.isArrowBacked(batch)) { - Utils.serializeBatchColumns(batch) + val (bytes, columnSizes) = if (Utils.isArrowBacked(batch)) { + CachedBatchIpc.serialize(batch, codec, CometArrowAllocator) } else { val arrowBatch = CometArrowConverters.columnarBatchToArrowBatch(batch, arrowSchema, CometArrowAllocator) - try Utils.serializeBatchColumns(arrowBatch) + try CachedBatchIpc.serialize(arrowBatch, codec, CometArrowAllocator) finally arrowBatch.close() } - val columnSizes = columns.map(_.size) CometCachedBatch( numRows = numRows, - sizeInBytes = columnSizes.sum, + sizeInBytes = bytes.length.toLong, stats = statsRow(lower, upper, nulls, numRows, columnSizes), - columns = columns) + bytes = bytes) } } @@ -293,16 +304,22 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { } } - // 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. + // Columnar Comet output is stored as one Arrow IPC record batch message per cached batch. 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] = { + // Read on the driver: the closure ships to the executors, where CometConf would resolve + // against whatever SQLConf happens to be current on that thread rather than this session's. + val codecName = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.get(conf) + val zstdLevel = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL.get(conf) + input.mapPartitions { batches => - encodeBatches(batches, schema) + encodeBatches(batches, schema, codecName, zstdLevel) } } @@ -320,24 +337,33 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { } val indices = selectedIndices(cacheAttributes, selectedAttributes) + // Captured as a StructType rather than the attributes themselves: this closure ships to the + // executors, and the Arrow schema is rebuilt there from the same mapping the writer used. + val cacheSchema = Utils.fromAttributes(cacheAttributes) 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. + val arrowFields = + Utils + .toArrowSchema(cacheSchema, CometArrowStream.NATIVE_TIMEZONE) + .getFields + .asScala + .toSeq + + // A ProjectedBatch owns the vectors of the batch it produced, and releases them only when + // that batch has been consumed. A consumer that stops early -- LIMIT, take(), or a + // cancelled task -- leaves the batch in flight open, so close it 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 + @volatile var current: ProjectedBatch = null Option(TaskContext.get()).foreach { tc => tc.addTaskCompletionListener[Unit] { _ => - val readers = current + val open = current current = null - if (readers != null) { - readers.close() + if (open != null) { + open.close() } } } @@ -348,9 +374,9 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { // 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 + val projected = new ProjectedBatch(cb, arrowFields, indices) + current = projected + projected.batches } case other => @@ -360,78 +386,56 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { } } - // 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 + /** + * Loads the projected columns of one cached batch into Arrow vectors. + * + * The schema is rebuilt from the cached relation's attributes rather than read from the + * payload, which stores none. `CachedBatchIpc.readProjected` then materializes only the + * selected columns' buffers, so the rest are never copied out of the cached bytes or + * decompressed. + * + * The decoded vectors stay owned by this object: closing it releases the batch, which is why + * this yields a single-element iterator that closes on exhaustion. + */ + private class ProjectedBatch( + cached: CometCachedBatch, + arrowFields: Seq[Field], + indices: Array[Int]) { + + // Allocated before anything can throw, so that a failure below has a root to release. + private val root = VectorSchemaRoot.create( + new Schema(indices.map(arrowFields).toSeq.asJava), + CometArrowAllocator) + private var closed = false + + // Loading happens during construction, so `batches` below can hand out the root directly. + try { + val recordBatch = + CachedBatchIpc.readProjected(cached.bytes, arrowFields, indices, CometArrowAllocator) try { - while (i < buffers.length) { - opened(i) = Utils.decodeBatches(buffers(i), "CometCache") - i += 1 - } - } catch { - case NonFatal(e) => - var j = 0 - while (j < i) { - opened(j) match { - case reader: ArrowReaderIterator => - try reader.close() - catch { case NonFatal(closeError) => e.addSuppressed(closeError) } - case _ => () - } - j += 1 - } - throw e + CachedBatchIpc.loaderFor(root).load(recordBatch) + } finally { + recordBatch.close() + } + // A cached batch's columns all cover the same rows. Check rather than trust: a mismatch + // would otherwise build a batch whose columns disagree with the row count recorded beside + // them, which reads as corrupt data far from here. + if (root.getRowCount != cached.numRows) { + throw new IllegalStateException( + s"Cached batch decoded ${root.getRowCount} rows, expected ${cached.numRows}") } - opened + } catch { + case NonFatal(e) => + try root.close() + catch { case NonFatal(closeError) => e.addSuppressed(closeError) } + throw e } - private var closed = false def close(): Unit = synchronized { if (!closed) { closed = true - readers.foreach { - case reader: ArrowReaderIterator => reader.close() - case _ => () - } - } - } - - private def assemble(): ColumnarBatch = { - val columns = new Array[ColumnVector](readers.length) - var i = 0 - while (i < readers.length) { - val reader = readers(i) - if (!reader.hasNext) { - throw new IllegalStateException( - s"Cached column stream $i of ${readers.length} decoded to no batch") - } - val decoded = reader.next() - // Each stream holds exactly one single-column record batch, and every column of a cached - // batch covers the same rows. Check rather than trust: a mismatch would otherwise build a - // batch whose columns disagree on length, which reads as corrupt data far from here. - if (decoded.numCols() != 1) { - throw new IllegalStateException( - s"Cached column stream $i decoded to ${decoded.numCols()} columns, expected 1") - } - if (decoded.numRows() != numRows) { - throw new IllegalStateException( - s"Cached column stream $i decoded ${decoded.numRows()} rows, expected $numRows") - } - columns(i) = decoded.column(0) - i += 1 + root.close() } - new ColumnarBatch(columns, numRows) } def batches: Iterator[ColumnarBatch] = new Iterator[ColumnarBatch] { @@ -451,7 +455,7 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { throw new NoSuchElementException } emitted = true - assemble() + NativeUtil.rootAsBatch(root) } } } @@ -467,24 +471,26 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { fallback.convertInternalRowToCachedBatch(input, schema, storageLevel, conf) } else { val batchSize = conf.columnBatchSize + val codecName = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.get(conf) + val zstdLevel = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL.get(conf) input.mapPartitions { rows => val iter = CometArrowConverters.rowToArrowBatchIter( rows, Utils.fromAttributes(schema), batchSize, - // NATIVE_TIMEZONE ("UTC"), not conf.sessionLocalTimeZone, so both write paths produce - // the same physical format: the columnar path above already encodes with - // NATIVE_TIMEZONE. Unlike Spark's Arrow cache, whose RecordBatch is deliberately - // schema-less, CometCachedBatch stores a full IPC stream including the schema, so a - // session-local label would persist the writing session's mutable timezone into cached - // data. This is a label only: Spark's internal timestamp representation is micros since - // the Unix epoch regardless of session timezone, so no values are converted. It also - // matches Comet's native schema, avoiding a cast at the native boundary. + // NATIVE_TIMEZONE ("UTC"), not conf.sessionLocalTimeZone. The payload stores no schema, + // so the read path rebuilds one with toArrowSchema(cacheSchema, NATIVE_TIMEZONE); a + // write that labelled its timestamps with the writing session's timezone would be read + // back under a different label. Both write paths therefore have to agree on this, and + // the columnar path above encodes with NATIVE_TIMEZONE too. This is a label only: + // Spark's internal timestamp representation is micros since the Unix epoch regardless + // of session timezone, so no values are converted. It also matches Comet's native + // schema, avoiding a cast at the native boundary. CometArrowStream.NATIVE_TIMEZONE, CometArrowAllocator) - encodeBatches(iter, schema) + encodeBatches(iter, schema, codecName, zstdLevel) } } } @@ -551,6 +557,9 @@ object ArrowCachedBatchSerializer { */ def kryoClasses: Seq[Class[_]] = Seq( classOf[CometCachedBatch], + // The payload itself. Kryo registers Array[Byte] by default, but registering it here is what + // keeps that true if the payload type ever changes again. + classOf[Array[Byte]], // The statistics row, whose values are bounds in Spark's internal representation: boxed // primitives, which Kryo registers by default, plus UTF8String and Decimal, which it does not. // A Decimal above Long precision holds a scala.math.BigDecimal, which Chill's Scala registrar diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala new file mode 100644 index 00000000000..55215ef548a --- /dev/null +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala @@ -0,0 +1,409 @@ +/* + * 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 java.io.{ByteArrayInputStream, ByteArrayOutputStream} +import java.nio.channels.Channels + +import scala.collection.mutable +import scala.jdk.CollectionConverters._ +import scala.util.control.NonFatal + +import org.apache.arrow.compression.{CommonsCompressionFactory, ZstdCompressionCodec} +import org.apache.arrow.flatbuf.{RecordBatch => FlatBufRecordBatch} +import org.apache.arrow.memory.{ArrowBuf, BufferAllocator} +import org.apache.arrow.vector.{FieldVector, TypeLayout, ValueVector, VectorSchemaRoot, VectorUnloader} +import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil, NoCompressionCodec} +import org.apache.arrow.vector.dictionary.DictionaryEncoder +import org.apache.arrow.vector.ipc.{ReadChannel, WriteChannel} +import org.apache.arrow.vector.ipc.message.{ArrowBodyCompression, ArrowFieldNode, ArrowRecordBatch, MessageSerializer} +import org.apache.arrow.vector.types.pojo.{ArrowType, Field} +import org.apache.spark.SparkException +import org.apache.spark.sql.comet.util.Utils +import org.apache.spark.sql.vectorized.ColumnarBatch + +/** + * The on-disk shape of a `CometCachedBatch` payload, and the two operations over it. + * + * A cached batch is one encapsulated Arrow IPC RecordBatch message followed by its body, with no + * Schema message and no end-of-stream marker. The schema is not stored because the reader already + * has it: `InMemoryRelation` knows the cached relation's attributes, and `Utils.toArrowSchema` + * maps them to exactly the fields the writer unloaded. Leaving it out saves a schema message per + * cached batch, which for a wide relation cached in many batches is a large share of the payload + * that is not data. + * + * Compression is applied by Arrow per buffer rather than by wrapping the whole payload in a Spark + * `CompressionCodec`. That is what makes projection cheap: the message metadata records every + * buffer's offset and length within the body, so [[readProjected]] can copy out only the buffers + * of the columns a scan selected and let `VectorLoader` decompress just those. A whole-payload + * codec would have to inflate everything before any column could be read. + */ +private[comet] object CachedBatchIpc { + + /** + * The Arrow compression codec named by `spark.comet.exec.inMemoryCache.compression.codec`. + * + * Only the write path consults the config. A batch records which codec compressed it, so the + * read path looks the codec up from the batch itself and keeps reading data cached before the + * config changed. + */ + def compressionCodec(codecName: String, zstdLevel: Int): CompressionCodec = codecName match { + case "none" => NoCompressionCodec.INSTANCE + // Constructed directly rather than through CompressionCodec.Factory, which ignores the level + // and always builds a codec at zstd's default. + case "zstd" => new ZstdCompressionCodec(zstdLevel) + // Arrow's other codec, LZ4_FRAME, is not offered. It is commons-compress's pure-Java LZ4 -- + // no relation to the JNI-accelerated lz4-java behind spark.io.compression.codec -- and + // measures three orders of magnitude slower to write than zstd while also producing larger + // output, so nothing prefers it. Reads still accept it, since the factory the read path uses + // handles whatever codec a batch records. + case other => + throw new SparkException( + s"Unsupported Arrow compression codec for Comet's cache: $other. " + + "Supported values: none, zstd") + } + + /** + * Serialize `batch` into one encapsulated IPC RecordBatch message. + * + * Returns the message bytes and the on-body compressed size of each top-level column, which the + * caller records in the statistics row. The sizes come from the message's own buffer layout, so + * they are the real stored sizes rather than an estimate. + * + * Dictionary-encoded columns are decoded to their plain form first. A payload with no Schema + * message cannot describe a dictionary encoding, and the schema the reader rebuilds from Spark + * attributes never carries one, so a dictionary-encoded column has nowhere to record either its + * index type or the dictionary itself. Comet's native scans do produce such columns, so this is + * a real path, not a defensive one. + * + * As in `Utils.serializeBatches`, `batch`'s vectors are cleared once written, so callers gather + * anything they need from the batch (statistics, for instance) before calling this. + */ + def serialize( + batch: ColumnarBatch, + codec: CompressionCodec, + allocator: BufferAllocator): (Array[Byte], Array[Long]) = { + val (vectors, hydrated) = hydrateDictionaries(batch, allocator) + try { + val root = new VectorSchemaRoot(vectors.asJava) + // A batch of zero columns carries only a row count, which a VectorSchemaRoot cannot infer + // without vectors to measure. + if (vectors.isEmpty) { + root.setRowCount(batch.numRows()) + } + + // alignBuffers=true matches the 8-byte buffer alignment readProjected reproduces when it + // repacks the selected buffers. + val unloader = new VectorUnloader(root, true, codec, true) + val recordBatch = unloader.getRecordBatch + try { + val fields = vectors.map(_.getField) + // Serializing consumes the batch, as it does in Utils.serializeBatches. The record batch + // holds its own buffers by now -- compressed copies, or retained references when the codec + // is none -- so releasing the vectors here does not touch it. getField still answers + // afterwards: clearing releases buffers, not the schema. + // + // Not load bearing for memory: the plan that produced the batch releases its vectors + // either way, and dropping this line leaks nothing. It is here because serializeBatches + // does the same, so both writers leave a batch they were handed in the same state. + root.clear() + + val out = new ByteArrayOutputStream() + val channel = new WriteChannel(Channels.newChannel(out)) + MessageSerializer.serialize(channel, recordBatch) + (out.toByteArray, columnSizes(fields, recordBatch)) + } finally { + recordBatch.close() + } + } finally { + // Only the vectors this method allocated. The rest belong to the input batch. + hydrated.foreach(v => + try v.close() + catch { case NonFatal(_) => () }) + } + } + + /** + * Read an encapsulated IPC RecordBatch message, materializing off-heap only the buffers of the + * requested top-level columns. + * + * The body is a flat, depth-first sequence of buffers in schema order, so each top-level column + * owns a contiguous run of buffers whose length is [[fieldBufferCount]]; field nodes and + * variadic buffer counts run in the same order. The selected columns' bytes are copied into a + * single off-heap allocation, each buffer 8-byte aligned exactly as Arrow's IPC body lays them + * out, and the returned batch's buffers are windows into it -- one allocation, no per-buffer + * bookkeeping. + * + * Only the selected buffers are ever decompressed. A buffer's recorded (offset, length) covers + * its on-body bytes including the uncompressed-length prefix, so a copied window is exactly + * what the writer emitted; the columns that were not selected are never read, let alone + * inflated. The copied windows are then decompressed in one pass -- see [[decompressed]] for + * why that is not left to `VectorLoader` -- so what comes back is an uncompressed batch. + * + * The returned batch owns its buffers; the caller closes it. + */ + def readProjected( + data: Array[Byte], + schemaFields: Seq[Field], + selectedIndices: Array[Int], + allocator: BufferAllocator): ArrowRecordBatch = { + val readChannel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) + // Reads the message metadata only. The body stays in `data` and is copied selectively below. + val metadata = MessageSerializer.readMessage(readChannel) + if (metadata == null) { + throw new SparkException("Unexpected end of input reading a Comet cached batch") + } + val batch = + metadata.getMessage.header(new FlatBufRecordBatch()).asInstanceOf[FlatBufRecordBatch] + // serialize writes exactly [encapsulated message][body] and nothing after it, so the body is + // the tail of `data`. + val bodyStart = data.length - metadata.getMessageBodyLength.toInt + + val compression = + if (batch.compression() == null) NoCompressionCodec.DEFAULT_BODY_COMPRESSION + else new ArrowBodyCompression(batch.compression().codec(), batch.compression().method()) + + val nodeStarts = schemaFields.scanLeft(0)(_ + fieldNodeCount(_)).toArray + val bufferStarts = schemaFields.scanLeft(0)(_ + fieldBufferCount(_)).toArray + val variadicStarts = schemaFields.scanLeft(0)(_ + fieldVariadicCount(_)).toArray + val hasVariadic = batch.variadicBufferCountsLength() > 0 + + // The selected columns' field nodes, buffer indices and variadic counts, in output order. + val nodes = new java.util.ArrayList[ArrowFieldNode]() + val bufferIndices = mutable.ArrayBuffer.empty[Int] + val variadicCounts = new java.util.ArrayList[java.lang.Long]() + selectedIndices.foreach { i => + val field = schemaFields(i) + val nodeStart = nodeStarts(i) + (nodeStart until nodeStart + fieldNodeCount(field)).foreach { j => + val node = batch.nodes(j) + nodes.add(new ArrowFieldNode(node.length(), node.nullCount())) + } + val bufferStart = bufferStarts(i) + (bufferStart until bufferStart + fieldBufferCount(field)).foreach(bufferIndices += _) + if (hasVariadic) { + val variadicStart = variadicStarts(i) + (variadicStart until variadicStart + fieldVariadicCount(field)) + .foreach(j => variadicCounts.add(batch.variadicBufferCounts(j))) + } + } + + val layout = bufferIndices.map { j => + val buffer = batch.buffers(j) + (buffer.offset(), buffer.length()) + } + val alignedSizes = layout.map { case (_, length) => ((length + 7) / 8) * 8 } + // allocator.buffer(0) is legal but yields a buffer no window can be sliced from, and an + // all-empty projection (every selected column a NullVector, say) would ask for exactly that. + val body = allocator.buffer(math.max(alignedSizes.sum, 1L)) + val compressedBatch = + try { + val buffers = new java.util.ArrayList[ArrowBuf]() + var position = 0L + layout.indices.foreach { k => + val (sourceOffset, length) = layout(k) + if (length > 0) { + body.setBytes(position, data, bodyStart + sourceOffset.toInt, length.toInt) + } + val window = body.slice(position, length) + window.writerIndex(length) + buffers.add(window) + position += alignedSizes(k) + } + new ArrowRecordBatch( + batch.length().toInt, + nodes, + buffers, + compression, + variadicCounts, + false) + } catch { + case NonFatal(e) => + body.close() + throw e + } + + // The constructor retained each window; slice() alone does not. Dropping `body`'s own + // reference leaves the batch as sole owner of the one allocation, so closing the batch below + // is what frees it -- and closing `body` again here would drive its reference count negative. + body.close() + try decompressed(compressedBatch, allocator) + finally compressedBatch.close() + } + + /** + * The same record batch with every buffer decompressed, as a new batch the caller owns. + * + * `VectorLoader` would do this itself, but arrow-java 18.3.0 leaks on the failure path: + * `VectorLoader.loadBuffers` decompresses a field's buffers into a local list and only releases + * them after the whole field has loaded, so if one buffer of a field fails to decompress, every + * buffer of that field decompressed before it is unreachable and never freed. A string column + * is enough to reach it -- its offsets buffer decompresses, then its data buffer throws -- so a + * single corrupt cached batch leaks off-heap for the life of the executor. Doing the + * decompression here keeps every allocation reachable from this method's own error path. + * + * Buffers are retained before decompressing rather than after, which is the other half of the + * difference. `decompress` consumes a reference to its input on the paths where it allocates, + * so retaining afterwards leaves the reference stranded if it throws -- and, when a batch has a + * single buffer, drops the shared body to zero references and frees it before the retain that + * was meant to protect it. + */ + private def decompressed( + batch: ArrowRecordBatch, + allocator: BufferAllocator): ArrowRecordBatch = { + // getCodec is the raw IPC byte; the factory keys off the enum. Both sides of the comparison + // below have to be CodecType: NoCompressionCodec.COMPRESSION_TYPE is the byte -1, and Scala + // compares a CodecType against it by universal equality, which is quietly always unequal. + val codecType = + CompressionUtil.CodecType.fromCompressionType(batch.getBodyCompression.getCodec) + val compressed = codecType != CompressionUtil.CodecType.NO_COMPRESSION + val codec: CompressionCodec = + if (compressed) CommonsCompressionFactory.INSTANCE.createCodec(codecType) + else NoCompressionCodec.INSTANCE + + val buffers = new java.util.ArrayList[ArrowBuf]() + try { + batch.getBuffers.asScala.foreach { buffer => + buffer.getReferenceManager.retain() + val plain = + try { + // An empty buffer carries no compressed length prefix to read. + if (compressed && buffer.writerIndex() > 0) codec.decompress(allocator, buffer) + else buffer + } catch { + case NonFatal(e) => + buffer.getReferenceManager.release() + throw e + } + buffers.add(plain) + } + + val result = new ArrowRecordBatch( + batch.getLength, + batch.getNodes, + buffers, + NoCompressionCodec.DEFAULT_BODY_COMPRESSION, + batch.getVariadicBufferCounts, + false) + // The constructor retained each buffer, so drop the references held here. + buffers.asScala.foreach(_.close()) + result + } catch { + case NonFatal(e) => + buffers.asScala.foreach { buffer => + try buffer.close() + catch { case NonFatal(closeError) => e.addSuppressed(closeError) } + } + throw e + } + } + + /** + * A `VectorLoader` for what [[readProjected]] returns. + * + * No compression factory: [[readProjected]] has already decompressed every buffer, so the + * loader only ever sees a batch marked uncompressed. + */ + def loaderFor(root: VectorSchemaRoot): org.apache.arrow.vector.VectorLoader = + new org.apache.arrow.vector.VectorLoader(root) + + /** + * The on-body compressed size of each top-level column. + * + * Each column owns the run of buffers its subtree occupies, so its stored size is the sum of + * those buffers' recorded lengths. With one payload per batch these are the only per-column + * sizes available -- there is no separate stream to measure -- and they are exact. + */ + private def columnSizes(fields: Seq[Field], recordBatch: ArrowRecordBatch): Array[Long] = { + val buffers = recordBatch.getBuffersLayout + val starts = fields.scanLeft(0)(_ + fieldBufferCount(_)).toArray + fields.indices.map { i => + (starts(i) until starts(i) + fieldBufferCount(fields(i))) + .map(j => buffers.get(j).getSize) + .sum + }.toArray + } + + /** + * Replace every dictionary-encoded column of `batch` with its decoded form. + * + * Returns the vectors to write and, separately, the ones allocated here so the caller can close + * exactly those. Columns that needed no decoding are returned as they are and stay owned by + * `batch`. + */ + private def hydrateDictionaries( + batch: ColumnarBatch, + allocator: BufferAllocator): (Seq[FieldVector], Seq[ValueVector]) = { + val hydrated = mutable.ArrayBuffer.empty[ValueVector] + try { + val vectors = + Utils.getBatchFieldVectorsWithProviders(batch).map { case (vector, providerOpt) => + val encoding = vector.getField.getDictionary + if (encoding == null) { + vector + } else { + val dictionary = providerOpt.map(_.lookup(encoding.getId)).orNull + if (dictionary == null) { + throw new SparkException( + s"Column ${vector.getField.getName} is dictionary encoded with ID " + + s"${encoding.getId}, but no dictionary with that ID was provided") + } + val decoded = DictionaryEncoder.decode(vector, dictionary, allocator) + hydrated += decoded + decoded.asInstanceOf[FieldVector] + } + } + (vectors, hydrated.toSeq) + } catch { + case NonFatal(e) => + hydrated.foreach(v => + try v.close() + catch { case NonFatal(closeError) => e.addSuppressed(closeError) }) + throw e + } + } + + /** + * Number of Arrow buffers a field occupies in a RecordBatch body, including every descendant, + * in the depth-first order `VectorLoader` consumes them. The type's own count covers its + * validity and offset/data buffers; each child contributes its whole subtree. + */ + private def fieldBufferCount(field: Field): Int = + TypeLayout.getTypeBufferCount(field.getType) + + field.getChildren.asScala.map(fieldBufferCount).sum + + /** Number of field nodes a field occupies: itself plus every descendant. */ + private def fieldNodeCount(field: Field): Int = + 1 + field.getChildren.asScala.map(fieldNodeCount).sum + + /** + * Number of variadic buffer counts a field contributes, one per view-type buffer, recursively. + * + * Only Utf8View and BinaryView carry one. Comet's cache never writes view vectors today, but + * the span arithmetic above has to stay correct if that changes. + */ + private def fieldVariadicCount(field: Field): Int = { + val own = field.getType match { + case _: ArrowType.Utf8View | _: ArrowType.BinaryView => 1 + case _ => 0 + } + own + field.getChildren.asScala.map(fieldVariadicCount).sum + } +} diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala index 769d8058de5..f418c31a626 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala @@ -270,45 +270,12 @@ object Utils extends CometTypeShim with Logging { } /** - * Serializes each column of `batch` into its own compressed Arrow IPC stream, in column order. - * - * [[serializeBatches]] writes one stream covering every column, so a reader has to inflate all - * of them before it can project. Comet's in-memory cache stores columns separately instead, so - * a scan decodes only the ones it selected. Each stream is self-contained, including its schema - * and any dictionaries the column needs. - * - * The row count is not recoverable from the result when `batch` has no columns, so callers keep - * it alongside. As with [[serializeBatches]], the batch's vectors are cleared once written. - */ - def serializeBatchColumns(batch: ColumnarBatch): Array[ChunkedByteBuffer] = { - val codec = CompressionCodec.createCodec(SparkEnv.get.conf) - - // Each column is written with the provider it was decoded with, not the batch's first one: - // columns decoded from separate streams have independent dictionary ID namespaces. - getBatchFieldVectorsWithProviders(batch).map { case (fieldVector, providerOpt) => - val provider = providerOpt.getOrElse(new CDataDictionaryProvider) - val cbbos = new ChunkedByteBufferOutputStream(1024 * 1024, ByteBuffer.allocate) - val out = new DataOutputStream(codec.compressedOutputStream(cbbos)) - - val root = new VectorSchemaRoot(Seq(fieldVector).asJava) - val writer = new ArrowStreamWriter(root, provider, Channels.newChannel(out)) - writer.start() - writer.writeBatch() - root.clear() - writer.close() - - cbbos.toChunkedByteBuffer - }.toArray - } - - /** - * The classes that carry the output of [[serializeBatches]] and [[serializeBatchColumns]] out - * of Comet, for Kryo registration by [[org.apache.comet.CometKryoRegistrator]]. + * The classes that carry the output of [[serializeBatches]] out of Comet, for Kryo registration + * by [[org.apache.comet.CometKryoRegistrator]]. * * Spark registers `ChunkedByteBuffer` itself but not an array of them, and * `CometBroadcastExchangeExec` broadcasts exactly that array, so a native broadcast fails under - * `spark.kryo.registrationRequired=true` whichever Comet features are enabled. Comet's cache - * format stores one buffer per column and so needs the same registrations. + * `spark.kryo.registrationRequired=true` whichever Comet features are enabled. */ def arrowBytesKryoClasses: Seq[Class[_]] = Seq( classOf[ChunkedByteBuffer], diff --git a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala index b71476c3dd1..a6b34b74e7d 100644 --- a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala @@ -22,6 +22,7 @@ package org.apache.comet.shims import scala.annotation.nowarn import org.apache.spark.sql.types.{DataType, StructType} +import org.apache.spark.unsafe.types.{ByteArray, UTF8String} trait CometTypeShim { @nowarn // Spark 4 feature; stubbed to false in Spark 3.x for compatibility. @@ -41,4 +42,15 @@ trait CometTypeShim { @nowarn // Spark 4.1 feature; TimeType doesn't exist in Spark 3.x. def isTimeType(dt: DataType): Boolean = false + + /** + * Compare two strings under the collation of `dt`, which must be a `StringType`. + * + * Spark 3.x has no collations, so every string comparison is byte order. Callers that record + * comparable bounds (Comet's cache statistics, for instance) use this so the ordering they + * store is the one Spark's own comparison would produce. + */ + @nowarn // Collation is a Spark 4 feature; on 3.x every StringType compares as bytes. + def compareStrings(left: UTF8String, right: UTF8String, dt: DataType): Int = + ByteArray.compareBinary(left.getBytes, right.getBytes) } diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala index f48955a7da5..71e72dd2de4 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala @@ -21,6 +21,7 @@ package org.apache.comet.shims import org.apache.spark.sql.execution.datasources.VariantMetadata import org.apache.spark.sql.types.{ArrayType, DataType, MapType, StringType, StructType, VariantType} +import org.apache.spark.unsafe.types.UTF8String trait CometTypeShim { // A `StringType` carries collation metadata in Spark 4.0. Only non-default (non-UTF8_BINARY) @@ -64,4 +65,14 @@ trait CometTypeShim { dt.getClass.getSimpleName.startsWith("TimeType") def hasCollationSupport: Boolean = true + + /** + * Compare two strings under the collation of `dt`, which must be a `StringType`. + * + * `semanticCompare` is the comparison Spark's own expressions use for the type, so bounds + * recorded with it order the same way a predicate over the column does. For the default + * UTF8_BINARY collation it is byte order, which is what Spark 3.x always does. + */ + def compareStrings(left: UTF8String, right: UTF8String, dt: DataType): Int = + left.semanticCompare(right, dt.asInstanceOf[StringType].collationId) } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index a7a9a6c8590..3464e5e3ddc 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -29,6 +29,7 @@ import org.apache.spark.sql.catalyst.expressions.{And, Attribute, Expression, Gr import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} import org.apache.spark.sql.comet.CometInMemoryTableScanExec import org.apache.spark.sql.comet.execution.arrow.CometCachedBatchHelper +import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation} import org.apache.spark.sql.execution.exchange.{Exchange, ReusedExchangeExec} import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} @@ -724,13 +725,14 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache prunes only on columns that have bounds") { + test("Comet in-memory cache prunes on collated string columns") { assume(isSpark40Plus, "collated string types require Spark 4.0+") withNativeCache { - // A collated StringType does not match `case StringType` in the serializer's bounds - // tracking, so its lower and upper bounds stay null. Spark still builds a partition filter - // for it because a collated string literal is an AtomicType, and comparing against null - // bounds prunes every batch. Without the buildFilter guard this query returns no rows. + // Bounds for a collated column are recorded with that collation's own comparison, which is + // the same ordering the partition filter Spark generates over the column uses. Tracking + // bounds only for the bare `StringType` object would leave a collated column's bounds null, + // and a comparison against null bounds prunes every batch, so the query below would return + // no rows at all rather than merely losing the pruning. spark .sql("SELECT id, CAST(id AS STRING) COLLATE UTF8_LCASE AS s FROM range(100)") .createOrReplaceTempView("collated_cache") @@ -747,12 +749,57 @@ class CometInMemoryCacheSuite extends CometTestBase { assert( spark.sql("SELECT id FROM collated_cache WHERE s >= '5'").collect().length == expected) + // UTF8_LCASE compares case-insensitively, so bounds recorded under it have to as well: a + // batch whose values all sort above 'A' under byte order still contains matches for a + // predicate that is looking for lower-case letters. + spark + .sql( + "SELECT id, CAST(concat('X', cast(id as string)) AS STRING) COLLATE UTF8_LCASE AS s " + + "FROM range(100)") + .createOrReplaceTempView("collated_case_cache") + spark.catalog.cacheTable("collated_case_cache") + spark.table("collated_case_cache").count() + assert( + spark.sql("SELECT id FROM collated_case_cache WHERE s = 'x1'").collect().length == 1, + "a case-insensitive match must survive pruning") + // Null-count based pruning stays available for columns without bounds. assert( spark.sql("SELECT id FROM collated_cache WHERE s IS NOT NULL").collect().length == 100) } } + test("Comet in-memory cache prunes only on columns that have bounds") { + withNativeCache { + // Binary has no bounds recorded, so its lower and upper stay null. Spark would still build + // a partition filter for it, and comparing against null bounds prunes every batch, so + // without the buildFilter guard this query returns no rows. + spark + .sql("SELECT id, CAST(CAST(id AS STRING) AS BINARY) AS b FROM range(100)") + .createOrReplaceTempView("binary_cache") + spark.catalog.cacheTable("binary_cache") + spark.table("binary_cache").count() + + assert( + cachedBatchTypes("binary_cache").sameElements( + Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) + + assert( + spark + .sql("SELECT id FROM binary_cache WHERE b >= CAST('5' AS BINARY)") + .collect() + .length == + spark + .sql("SELECT id FROM range(100) WHERE CAST(CAST(id AS STRING) AS BINARY) >= " + + "CAST('5' AS BINARY)") + .collect() + .length) + + // Null-count based pruning stays available for columns without bounds. + assert(spark.sql("SELECT id FROM binary_cache WHERE b IS NOT NULL").collect().length == 100) + } + } + test("Comet in-memory cache is readable when Comet is disabled") { // spark.sql.cache.serializer is static, so the cached format cannot depend on a runtime // config. Disabling Comet must still leave the cached relation readable, including for @@ -986,8 +1033,12 @@ class CometInMemoryCacheSuite extends CometTestBase { SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> "true") { spark.catalog.clearCache() + // Wide enough that every column's buffers are big enough for Arrow to actually compress + // them. Arrow stores a buffer verbatim when compressing it would not make it smaller, and a + // boolean column of a few hundred rows is a few dozen bytes, which takes that fallback -- + // leaving the corruption the projection tests rely on with nothing to corrupt. spark - .range(0, 500, 1, 2) + .range(0, 8000, 1, 2) .selectExpr( "id", "id % 100 AS k", @@ -997,7 +1048,7 @@ class CometInMemoryCacheSuite extends CometTestBase { "cast(id % 2 = 0 as boolean) AS flag") .createOrReplaceTempView("projection_cache") spark.catalog.cacheTable("projection_cache") - assert(spark.table("projection_cache").count() == 500) + assert(spark.table("projection_cache").count() == 8000) assert( cachedBatchTypes("projection_cache").sameElements( Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) @@ -1032,67 +1083,149 @@ class CometInMemoryCacheSuite extends CometTestBase { .sum } - test("Comet in-memory cache stores one stream per column") { + /** + * Run `f`, require it to fail, and require the failure to be the decode error itself. + * + * The read path allocates an off-heap body, hands it to a record batch that takes its own + * references, and drops its own. A cleanup path that then releases the body a second time + * drives its reference count negative, and the reference-count error replaces the decode + * failure that caused it -- leaving a plain `intercept[Exception]` green while the user sees a + * error that says nothing about their corrupt cache. + */ + private def interceptDecodeFailure(f: => Unit): Throwable = { + val thrown = intercept[Exception](f) + val chain = + Iterator.iterate(thrown: Throwable)(_.getCause).takeWhile(_ != null).take(20).toSeq + assert( + !chain.exists { t => + t.getClass.getName.contains("IllegalReferenceCount") || + Option(t.getMessage).exists(m => m.contains("RefCnt") || m.contains("refCnt")) + }, + s"the decode failure must surface as itself, not as a reference-count error: $thrown") + thrown + } + + test("Comet in-memory cache round-trips under every compression codec") { + // Every codec the config accepts, not just the default. `none` takes a different path on read + // -- the payload records no codec, so nothing is decompressed -- and shipped broken for a + // while because the only tests that ran were on the default codec. + Seq("none", "zstd").foreach { codec => + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.key -> codec, + "spark.comet.sparkToColumnar.enabled" -> "true") { + + spark.catalog.clearCache() + val view = s"codec_cache_$codec" + spark + .range(0, 4000, 1, 2) + .selectExpr( + "id", + "cast(id as double) / 3 AS d", + "concat('s_', cast(id as string)) AS s", + "cast(id % 2 = 0 as boolean) AS flag") + .createOrReplaceTempView(view) + spark.catalog.cacheTable(view) + + assert( + cachedBatchTypes(view).sameElements( + Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch")), + s"codec $codec should still store CometCachedBatch") + + // A full read, a projected read (the buffer-selection path), and a row count that decodes + // nothing -- the three shapes the read path distinguishes. + checkSparkAnswer(spark.sql(s"SELECT * FROM $view")) + checkSparkAnswer(spark.sql(s"SELECT s FROM $view WHERE id >= 3990")) + assert(spark.sql(s"SELECT count(*) FROM $view").collect()(0).getLong(0) == 4000) + // Pruning reads the statistics rather than the payload, so exercise it too. + assert(spark.sql(s"SELECT id FROM $view WHERE id >= 3990").collect().length == 10) + + spark.catalog.clearCache() + } + } + } + + test("Comet in-memory cache stores no schema message per cached batch") { + // The reader rebuilds the schema from the cached relation's attributes, so storing one in + // every batch would repeat the same bytes for as many batches as the relation was cached in. withProjectionCache { (relation, batches) => assert(batches.nonEmpty) + val cacheSchema = Utils.fromAttributes(relation.output) batches.foreach { batch => assert( - CometCachedBatchHelper.numColumnStreams(batch) == relation.output.length, - "a cached batch must hold one independently decodable stream per cached column") + !CometCachedBatchHelper.hasSchemaMessage(batch), + "a cached batch must begin with its record batch, not a schema message") + val sizes = CometCachedBatchHelper.columnSizes(batch, cacheSchema) assert( - CometCachedBatchHelper.columnStreamSizes(batch).forall(_ > 0), - "every column stream must carry data") + sizes.length == relation.output.length, + "every cached column must own a run of buffers in the payload") + assert(sizes.forall(_ > 0), "every cached column must carry data") } } } test("Comet in-memory cache decodes only the projected columns") { - // Timings would be a weak assertion here, so this corrupts the streams the read must not - // touch. Reading still has to succeed, which it only can if those streams were never - // inflated. The second half checks the corruption is detectable at all, so that the first - // half cannot pass just because the bad bytes decode silently to nothing. + // Timings would be a weak assertion here, so this scrambles the compressed bytes of the + // columns the read must not touch, leaving every other byte of the payload identical. + // Reading still has to succeed, which it only can if those columns' buffers were never copied + // out of the payload and handed to the decompressor. The second half checks the corruption is + // detectable at all, so the first half cannot pass just because the bad bytes decode silently + // to nothing. withProjectionCache { (relation, batches) => + val cacheSchema = Utils.fromAttributes(relation.output) val selectedIdx = 1 val selected = Seq(relation.output(selectedIdx)) + relation.output.indices.foreach { i => + assert( + batches.forall(b => CometCachedBatchHelper.columnIsCompressed(b, cacheSchema, i)), + s"column $i is not stored compressed, so corrupting it would prove nothing") + } + relation.output.indices.filter(_ != selectedIdx).foreach { i => - batches.foreach(b => CometCachedBatchHelper.corruptColumnStream(b, i)) + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) } assert( - decodedRowCount(relation, batches, selected) == 500, - "reading one column must not decode the other five") + decodedRowCount(relation, batches, selected) == 8000, + "reading one column must not decompress the other five") - batches.foreach(b => CometCachedBatchHelper.corruptColumnStream(b, selectedIdx)) - intercept[Exception] { + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, selectedIdx)) + interceptDecodeFailure { decodedRowCount(relation, batches, selected) } } } test("Comet in-memory cache decodes no columns for a row-count-only read") { - // SELECT count(*) selects no columns. Every stream is corrupted, so the read can only succeed - // by decoding none of them and answering from the row count the cached batch already carries. + // SELECT count(*) selects no columns. Every column's bytes are corrupted, so the read can + // only succeed by touching none of them and answering from the row count the cached batch + // already carries beside the payload. withProjectionCache { (relation, batches) => + val cacheSchema = Utils.fromAttributes(relation.output) relation.output.indices.foreach { i => - batches.foreach(b => CometCachedBatchHelper.corruptColumnStream(b, i)) + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) } - assert(decodedRowCount(relation, batches, Seq.empty) == 500) + assert(decodedRowCount(relation, batches, Seq.empty) == 8000) } } test("Comet in-memory cache records per-column sizes in its statistics") { - // SimpleMetricsCachedBatch reserves a fifth field per column for its size. Each column is now - // its own stream, so the real size is known and must be reported rather than left at zero. - withProjectionCache { (_, batches) => + // SimpleMetricsCachedBatch reserves a fifth field per column for its size. A column owns a + // known run of buffers in the payload, so the real stored size is known and must be reported + // rather than left at zero. + withProjectionCache { (relation, batches) => + val cacheSchema = Utils.fromAttributes(relation.output) batches.foreach { batch => - val sizes = CometCachedBatchHelper.columnStreamSizes(batch) + val sizes = CometCachedBatchHelper.columnSizes(batch, cacheSchema) val stats = batch.asInstanceOf[SimpleMetricsCachedBatch].stats sizes.zipWithIndex.foreach { case (size, i) => assert( stats.getLong(i * 5 + 4) == size, - s"column $i should report its own stream size in the statistics row") + s"column $i should report the stored size of its own buffers in the statistics row") } } } @@ -1169,11 +1302,12 @@ class CometInMemoryCacheSuite extends CometTestBase { test( "Comet in-memory cache re-encodes a decoded batch whose columns have separate dictionaries") { - // Each cached column is decoded from its own stream, so dictionary-backed columns come back - // with independent providers whose IDs collide. Re-encoding such a batch with only the first - // column's provider cannot resolve the later columns' dictionary IDs. Spark's columnar Union - // hands decoded cached batches straight back to this serializer, so caching a union of a - // cached relation exercises exactly that. + // Spark's columnar Union hands decoded cached batches straight back to this serializer, so + // caching a union of a cached relation re-encodes batches that came out of the cache. The + // cache no longer stores dictionary-encoded columns -- the writer decodes them first -- but + // batches reaching serializeBatches from a shuffle or broadcast still carry independent + // dictionary providers whose IDs collide, and re-encoding one with only the first column's + // provider cannot resolve the later columns' IDs. withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", @@ -1248,23 +1382,26 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache releases opened readers when a later column fails to decode") { - // A cached batch is several independent Arrow streams and decodeBatches opens each eagerly. - // If a later column throws, the readers already opened are unreachable: the task-completion - // listener cannot release them, because the holder is only published once its constructor - // returns. The failure would then leak off-heap for the life of the executor. + test("Comet in-memory cache releases its vectors when a column fails to decode") { + // Reading a batch allocates twice before anything can go wrong: the root that receives the + // projected columns, and the off-heap body the selected buffers are copied into. A column + // that fails to decompress throws between the two, and neither is reachable from anywhere + // else -- the holder is published to the task-completion listener only once its constructor + // returns -- so a failure that does not release them leaks off-heap for the life of the + // executor. withProjectionCache { (relation, batches) => - // Corrupt the second selected column, so the first is opened successfully first. + val cacheSchema = Utils.fromAttributes(relation.output) + // Corrupt the second selected column, so the first is copied out successfully first. val selected = Seq(relation.output(0), relation.output(1)) - batches.foreach(b => CometCachedBatchHelper.corruptColumnStream(b, 1)) + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, 1)) val before = CometArrowAllocator.getAllocatedMemory - intercept[Exception] { + interceptDecodeFailure { decodedRowCount(relation, batches, selected) } assert( CometArrowAllocator.getAllocatedMemory == before, - "readers opened before the failure must be released") + "everything allocated before the failure must be released") } } @@ -1272,7 +1409,7 @@ class CometInMemoryCacheSuite extends CometTestBase { * Cache two low-cardinality string columns and hand the test the cached payload. * * The shuffle is what makes this worth its own fixture: its reader hands the cache writer - * dictionary-encoded columns, so each cached column stream carries a dictionary of its own. + * dictionary-encoded columns, which the writer has to decode before storing them. */ private def withDictionaryCache(f: (InMemoryRelation, Array[CachedBatch]) => Unit): Unit = { withSQLConf( @@ -1308,15 +1445,51 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache broadcasts a batch whose columns have separate dictionaries") { - // A broadcast of a cache scan re-serializes each decoded batch as one stream covering every - // column, and the writer resolves all of their dictionary IDs against the single provider it - // is handed. The columns were decoded from separate streams, so they arrive carrying separate - // providers: passing any one of them cannot resolve the others. - withDictionaryCache { (relation, batches) => + test("Comet in-memory cache releases its vectors when a column fails after a partial decode") { + // Tighter than the two cases above, and the one that actually catches a leak. A string column + // stores its offsets and its data as separate compressed buffers, so corrupting only the + // second makes the decoder decompress one buffer of the column into a fresh allocation and + // then throw on the next, with the first reachable from nothing the failure path can see. + withProjectionCache { (relation, batches) => + val cacheSchema = Utils.fromAttributes(relation.output) + val stringIdx = 3 + assert(relation.output(stringIdx).dataType.typeName == "string") + batches.foreach(b => + CometCachedBatchHelper.corruptTrailingBuffer(b, cacheSchema, stringIdx)) + + val before = CometArrowAllocator.getAllocatedMemory + interceptDecodeFailure { + decodedRowCount(relation, batches, Seq(relation.output(stringIdx))) + } assert( - CometCachedBatchHelper.columnsAreDictionaryEncoded(batches.head).forall(identity), - "this test is only meaningful over dictionary-encoded cached columns") + CometArrowAllocator.getAllocatedMemory == before, + "a buffer decoded before the failure must be released") + } + } + + test("Comet in-memory cache decodes dictionary-encoded columns before storing them") { + // The payload carries no schema, so it has nowhere to record that a column is dictionary + // encoded, nor the dictionary itself. The reader rebuilds a plain Utf8 field for a string + // column either way, so a writer that stored the index vector as-is would hand the loader + // integer indices to read as strings. Reading the values back correctly is what proves the + // writer decoded them first; a row count alone would not. + withDictionaryCache { (relation, _) => + assert(relation.output.length == 2) + + val df = spark.sql("SELECT s1, s2 FROM dictionary_cache") + checkSparkAnswer(df) + + val distinct = + spark.sql("SELECT DISTINCT s1, s2 FROM dictionary_cache ORDER BY s1, s2").collect() + assert(distinct.length == 12, "3 distinct s1 values by 4 distinct s2 values") + assert(distinct.head.getString(0) == "a_0" && distinct.head.getString(1) == "b_0") + } + } + + test("Comet in-memory cache broadcasts a batch read back from the cache") { + // A broadcast of a cache scan re-serializes each decoded batch through serializeBatches, + // which is a different writer from the one that produced the cached payload. + withDictionaryCache { (relation, _) => assert(relation.output.length == 2) val df = spark.sql( @@ -1326,27 +1499,23 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache releases a reader whose own first batch fails to decode") { - // A dictionary-encoded column loads its dictionary before the record batch that indexes into - // it, so a reader can allocate and then fail while opening. Nothing else can release what it - // took: its constructor never returns, so no caller holds the reader it would close, and the - // task-completion listener has not been told about it either. + test("Comet in-memory cache releases its vectors when a column fails part way through") { + // Distinct from the corrupted-column case above: there the compressed bytes are wrong from + // their first byte, so the decompressor rejects them outright. Here the bytes start out + // genuine and only the tail is destroyed, so the failure lands after the loader has already + // begun filling vectors. Nothing else can release them -- the holder's constructor never + // returns, so no caller holds it and the task-completion listener has not been told about it. withDictionaryCache { (relation, batches) => - assert( - CometCachedBatchHelper.columnsAreDictionaryEncoded(batches.head).forall(identity), - "this test is only meaningful over dictionary-encoded cached columns") - - // Enough to take out the end-of-stream marker and bite into the record batch body, so the - // read fails after the dictionary has been loaded rather than before. - batches.foreach(b => CometCachedBatchHelper.truncateColumnStream(b, 0, 64)) + val cacheSchema = Utils.fromAttributes(relation.output) + batches.foreach(b => CometCachedBatchHelper.truncateColumn(b, cacheSchema, 0)) val before = CometArrowAllocator.getAllocatedMemory - intercept[Exception] { + interceptDecodeFailure { decodedRowCount(relation, batches, Seq(relation.output.head)) } assert( CometArrowAllocator.getAllocatedMemory == before, - "a reader that fails while opening must release what it already allocated") + "a read that fails while loading must release what it already allocated") } } diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index 959b5590b34..c269dc220c6 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -85,9 +85,10 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { |WHERE id >= 4500000 AND id < 4750000 """.stripMargin) - // A CometCachedBatch stores each column as its own stream, so a scan decodes only what it - // projected and cost tracks the width of the projection. These three cases span that range - // over one cached relation: no columns, one column, and all six. + // A CometCachedBatch records where each column's buffers sit in its payload, so a scan + // copies out and decompresses only what it projected and cost tracks the width of the + // projection. These three cases span that range over one cached relation: no columns, one + // column, and all six. runCacheBenchmark( "in-memory cache row count only (0 of 6 columns)", s"SELECT count(*) FROM $cacheTable") diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala index 5548f6dadbd..56c3b949569 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala @@ -19,17 +19,19 @@ package org.apache.spark.sql.comet.execution.arrow -import java.io.{DataInputStream, DataOutputStream} -import java.nio.ByteBuffer +import java.io.ByteArrayInputStream import java.nio.channels.Channels -import org.apache.arrow.vector.ipc.ArrowStreamReader -import org.apache.spark.SparkEnv -import org.apache.spark.io.CompressionCodec -import org.apache.spark.sql.columnar.CachedBatch -import org.apache.spark.util.io.{ChunkedByteBuffer, ChunkedByteBufferOutputStream} +import scala.jdk.CollectionConverters._ -import org.apache.comet.CometArrowAllocator +import org.apache.arrow.flatbuf.{MessageHeader, RecordBatch => FlatBufRecordBatch} +import org.apache.arrow.vector.TypeLayout +import org.apache.arrow.vector.ipc.ReadChannel +import org.apache.arrow.vector.ipc.message.MessageSerializer +import org.apache.arrow.vector.types.pojo.Field +import org.apache.spark.sql.columnar.CachedBatch +import org.apache.spark.sql.comet.util.Utils +import org.apache.spark.sql.types.StructType /** * Test-only access to the internals of `CometCachedBatch`. @@ -37,83 +39,210 @@ import org.apache.comet.CometArrowAllocator * A top-level `private` class in Scala is visible to its own package, so this shim needs no * reflection; it exists so tests outside `org.apache.spark.sql.comet.execution.arrow` can assert * on the cached payload's shape. + * + * The IPC buffer arithmetic below is deliberately re-derived here rather than reused from + * `CachedBatchIpc`. A helper that called into the code under test would inherit any bug in it and + * still agree with itself, so the assertions built on this would pass for the wrong reason. */ object CometCachedBatchHelper { - /** Number of independently decodable column streams in a cached batch. */ - def numColumnStreams(batch: CachedBatch): Int = - batch.asInstanceOf[CometCachedBatch].columns.length + /** The raw cached payload: one encapsulated Arrow IPC record batch message and its body. */ + private def payload(batch: CachedBatch): Array[Byte] = + batch.asInstanceOf[CometCachedBatch].bytes + + /** Stored size of the whole cached batch, in bytes. */ + def payloadSize(batch: CachedBatch): Long = payload(batch).length.toLong + + /** + * Whether the payload begins with a Schema message rather than going straight to the record + * batch. + * + * Comet stores no schema per cached batch -- the reader rebuilds it from the cached relation's + * attributes -- so this is false, and a regression to a self-describing stream would show up + * here rather than only as a footprint number. + */ + def hasSchemaMessage(batch: CachedBatch): Boolean = + readMetadata(payload(batch))._1.headerType() == MessageHeader.Schema + + /** + * The on-body (offset, length) of every Arrow buffer belonging to each top-level column, in + * column order. + */ + def columnBufferRanges(batch: CachedBatch, cacheSchema: StructType): Seq[Seq[(Long, Long)]] = { + val data = payload(batch) + val (_, recordBatch) = readMetadata(data) + val fields = arrowFields(cacheSchema) + val starts = fields.scanLeft(0)(_ + bufferCount(_)).toArray + + fields.indices.map { i => + (starts(i) until starts(i) + bufferCount(fields(i))).map { j => + val buffer = recordBatch.buffers(j) + (buffer.offset(), buffer.length()) + } + } + } - /** Serialized size of each column stream, in column order. */ - def columnStreamSizes(batch: CachedBatch): Seq[Long] = - batch.asInstanceOf[CometCachedBatch].columns.map(_.size).toSeq + /** Stored size of each top-level column: the sum of its buffers' on-body lengths. */ + def columnSizes(batch: CachedBatch, cacheSchema: StructType): Seq[Long] = + columnBufferRanges(batch, cacheSchema).map(_.map(_._2).sum) /** - * Replace one column's stream with bytes that cannot be decoded, in place. + * Whether any of a column's buffers is actually stored compressed. * - * Reading a column this has corrupted fails; reading any other column only succeeds if that - * column's stream was never touched. That is the difference between decoding what was projected - * and decoding everything and projecting afterwards, so it is what the projection tests assert - * on rather than timings. + * Arrow prefixes each compressed buffer with its uncompressed length, and falls back to storing + * a buffer verbatim (length prefix `-1`) when compressing it would not make it smaller. Small + * buffers routinely take that fallback, so [[corruptColumn]] only has something to corrupt when + * this is true; the projection tests assert it as a precondition rather than assuming it. */ - def corruptColumnStream(batch: CachedBatch, index: Int): Unit = { - val columns = batch.asInstanceOf[CometCachedBatch].columns - columns(index) = new ChunkedByteBuffer(Array(ByteBuffer.wrap(Array[Byte](1, 2, 3, 4)))) + def columnIsCompressed(batch: CachedBatch, cacheSchema: StructType, index: Int): Boolean = { + val data = payload(batch) + val start = bodyStart(data) + columnBufferRanges(batch, cacheSchema)(index).exists { case (offset, length) => + length > 8 && uncompressedLength(data, start + offset.toInt) > 0 + } } - /** Whether each column's stream stores that column dictionary encoded, in column order. */ - def columnsAreDictionaryEncoded(batch: CachedBatch): Seq[Boolean] = - batch.asInstanceOf[CometCachedBatch].columns.toSeq.map { buffer => - val in = new DataInputStream(codec.compressedInputStream(buffer.toInputStream())) - val reader = new ArrowStreamReader(Channels.newChannel(in), CometArrowAllocator) - try { - reader.getVectorSchemaRoot.getSchema.getFields.get(0).getDictionary != null - } finally { - reader.close() + /** + * Scramble one column's compressed bytes in place, leaving every other column byte-identical. + * + * Reading a column this has corrupted fails in the decompressor; reading any other column only + * succeeds if this column's buffers were never copied out of the payload. That is the + * difference between decoding what was projected and decoding everything and projecting + * afterwards, so it is what the projection tests assert on rather than timings. + * + * Each buffer's 8-byte uncompressed-length prefix is left intact and only the compressed bytes + * after it are overwritten, so a read of this column fails while decompressing rather than by + * trying to allocate a nonsense length. Requires the column to have a genuinely compressed + * buffer -- see [[columnIsCompressed]]. + */ + def corruptColumn(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { + val data = payload(batch) + val start = bodyStart(data) + var corrupted = false + + columnBufferRanges(batch, cacheSchema)(index).foreach { case (offset, length) => + val bufferStart = start + offset.toInt + if (length > 8 && uncompressedLength(data, bufferStart) > 0) { + var i = bufferStart + 8 + while (i < bufferStart + length.toInt) { + // A fixed pattern rather than random bytes, so a failure reproduces exactly. + data(i) = (0xa5 ^ i).toByte + i += 1 + } + corrupted = true } } + require( + corrupted, + s"column $index of the cached batch has no compressed buffer to corrupt; " + + "the test needs data that Arrow actually compresses") + } + /** - * Drop the last `dropBytes` of one column's decoded Arrow stream, in place. + * Truncate the tail of one column's compressed bytes, in place, padding with zeros so every + * other column keeps its offset. * - * [[corruptColumnStream]] replaces the stream outright, so a reader over it fails on the very - * first message, before it has allocated anything. This keeps the stream genuine up to the cut: - * the reader parses the schema and loads the column's dictionary, and only then runs out of - * input part way through the record batch that indexes into it. The cut is made on the decoded - * bytes rather than the compressed ones because a small column compresses to a single block, - * and truncating that fails the decompressor before Arrow reads anything at all. + * [[corruptColumn]] rewrites the whole compressed payload, which fails as soon as the + * decompressor looks at it. This keeps the leading bytes genuine, so a decoder gets a stream + * that starts out valid and then runs out, exercising a failure part way through a column + * rather than at its first byte. */ - def truncateColumnStream(batch: CachedBatch, index: Int, dropBytes: Int): Unit = { - val columns = batch.asInstanceOf[CometCachedBatch].columns - - val decodedStream = new DataInputStream( - codec.compressedInputStream(columns(index).toInputStream())) - val decoded = - try { - val buffer = new java.io.ByteArrayOutputStream() - val chunk = new Array[Byte](8192) - var read = decodedStream.read(chunk) - while (read >= 0) { - buffer.write(chunk, 0, read) - read = decodedStream.read(chunk) + def truncateColumn(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { + val data = payload(batch) + val start = bodyStart(data) + var truncated = false + + columnBufferRanges(batch, cacheSchema)(index).foreach { case (offset, length) => + val bufferStart = start + offset.toInt + if (length > 32 && uncompressedLength(data, bufferStart) > 0 && !truncated) { + var i = bufferStart + length.toInt - 16 + while (i < bufferStart + length.toInt) { + data(i) = 0 + i += 1 } - buffer.toByteArray - } finally { - decodedStream.close() + truncated = true } + } + + require( + truncated, + s"column $index of the cached batch has no compressed buffer long enough to truncate") + } + + /** + * Scramble only the last compressed buffer of one column, leaving its earlier buffers genuine. + * + * A string column stores offsets and data as separate compressed buffers, so this makes the + * decoder decompress one buffer of the column successfully and then fail on the next. That is a + * different failure point from [[corruptColumn]], which takes out a column's first buffer and + * so fails before anything of it has been decompressed. + */ + def corruptTrailingBuffer(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { + val data = payload(batch) + val start = bodyStart(data) + val compressed = columnBufferRanges(batch, cacheSchema)(index).filter { + case (offset, length) => + length > 8 && uncompressedLength(data, start + offset.toInt) > 0 + } require( - decoded.length > dropBytes, - s"column $index decodes to ${decoded.length} bytes, too few to drop $dropBytes") - - val cbbos = new ChunkedByteBufferOutputStream(1024 * 1024, ByteBuffer.allocate) - val out = new DataOutputStream(codec.compressedOutputStream(cbbos)) - try { - out.write(decoded, 0, decoded.length - dropBytes) - } finally { - out.close() + compressed.length > 1, + s"column $index has ${compressed.length} compressed buffers; this needs at least two so a " + + "decode can succeed on one and then fail on the next") + + val (offset, length) = compressed.last + val bufferStart = start + offset.toInt + var i = bufferStart + 8 + while (i < bufferStart + length.toInt) { + data(i) = (0xa5 ^ i).toByte + i += 1 } - columns(index) = cbbos.toChunkedByteBuffer } - private def codec: CompressionCodec = CompressionCodec.createCodec(SparkEnv.get.conf) + /** The Arrow fields the read path rebuilds for `cacheSchema`. */ + private def arrowFields(cacheSchema: StructType): Seq[Field] = + Utils + .toArrowSchema(cacheSchema, CometArrowStream.NATIVE_TIMEZONE) + .getFields + .asScala + .toSeq + + /** + * Buffers a field occupies in the record batch body, including every descendant, in the + * depth-first order the body lays them out. + */ + private def bufferCount(field: Field): Int = + TypeLayout.getTypeBufferCount(field.getType) + + field.getChildren.asScala.map(bufferCount).sum + + private def readMetadata(data: Array[Byte]) = { + val channel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) + val metadata = MessageSerializer.readMessage(channel) + require(metadata != null, "cached payload holds no IPC message") + ( + metadata.getMessage, + metadata.getMessage.header(new FlatBufRecordBatch()).asInstanceOf[FlatBufRecordBatch]) + } + + /** Offset of the record batch body within the payload; the body is its tail. */ + private def bodyStart(data: Array[Byte]): Int = { + val channel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) + val metadata = MessageSerializer.readMessage(channel) + require(metadata != null, "cached payload holds no IPC message") + data.length - metadata.getMessageBodyLength.toInt + } + + /** + * The uncompressed-length prefix Arrow writes ahead of a compressed buffer, little-endian. A + * value of -1 means the buffer was stored verbatim because compressing it did not pay. + */ + private def uncompressedLength(data: Array[Byte], bufferStart: Int): Long = { + var value = 0L + var i = 7 + while (i >= 0) { + value = (value << 8) | (data(bufferStart + i) & 0xffL) + i -= 1 + } + value + } } From f59c9dc6743486cef696e0d5e52af07038d6311c Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 29 Aug 2026 08:52:31 -0600 Subject: [PATCH 02/24] refactor: use Spark's interpreted ordering for bounds, hoist projection layout Cleanup pass over the cache format change. No behaviour change. Drop the `compareStrings` shim in favour of `TypeUtils.getInterpretedOrdering`. That method is public with the same signature on every supported Spark version, and on Spark 4 it resolves a `StringType` through `CollationFactory.fetchCollation(collationId).comparator` -- the comparison the shim was reaching for. So the collation awareness comes from Spark itself and the shim, its Spark 3.x stub and the hand-rolled per-type `compare` all go. The ordering is now resolved once per column per partition rather than being re-dispatched on the `DataType` twice per row. Build the projection's index layout once per partition instead of per batch. The node, buffer and variadic index arithmetic is a pure function of the cached schema and the selected columns, but it walks every field of the relation, so recomputing it per batch made the bookkeeping O(total columns) against O(selected columns) of useful work -- worst in the wide-relation, narrow-projection case the format exists for. `CachedBatchIpc.Projection` now holds that layout and the projected schema, and owns the whole decode; `ProjectedBatch` is left with ownership only. This also puts the projected schema next to the code that packs buffers in the same order, an invariant that previously spanned two files unstated. Smaller cleanups: use Arrow's `DataSizeRoundingUtil.roundUpTo8Multiple` rather than open-coding IPC body alignment; size the serialization buffer from the record batch's known body length instead of growing from 32 bytes; resolve decompressors once instead of per batch; share the dictionary lookup guard between `Utils.combineDictionaryProviders` and the cache writer; read the codec config through one helper carrying the driver-vs-executor rationale; and collapse the duplicated compressed-buffer predicate and scramble loop in the test helper. Corrects two `Utils` scaladocs that still described the per-column stream format this change replaced. Benchmark and codec figures in the docs re-measured against the current code. --- .../user-guide/latest/in-memory-cache.md | 14 +- .../arrow/ArrowCachedBatchSerializer.scala | 162 +++++----- .../execution/arrow/CachedBatchIpc.scala | 281 ++++++++++-------- .../apache/spark/sql/comet/util/Utils.scala | 44 ++- .../apache/comet/shims/CometTypeShim.scala | 12 - .../apache/comet/shims/CometTypeShim.scala | 11 - .../arrow/CometCachedBatchHelper.scala | 182 +++++------- 7 files changed, 353 insertions(+), 353 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 67f0b0e9755..efd949ff92c 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -70,8 +70,8 @@ six-column relation: | Codec | Materialize | Footprint | Read 1 of 6 | Read 6 of 6 | | ------ | ----------: | --------: | ----------: | ----------: | -| `zstd` | 347 ms | 2 MiB | 52 ms | 63 ms | -| `none` | 1743 ms | 13 MiB | 74 ms | 79 ms | +| `zstd` | 363 ms | 2 MiB | 56 ms | 62 ms | +| `none` | 1776 ms | 13 MiB | 78 ms | 81 ms | Arrow's other IPC codec, LZ4, is deliberately not offered. It is commons-compress's pure-Java implementation and is unrelated to the JNI-accelerated lz4-java behind `spark.io.compression.codec`; @@ -101,11 +101,11 @@ SPARK_GENERATE_BENCHMARK_FILES=1 \ | Query shape | Spark cache scan + convert | `CometInMemoryTableScan` | Relative | | ------------------------------ | -------------------------: | -----------------------: | -------: | -| Repeated scan (3 of 6 columns) | 156 ms | 118 ms | 1.3x | -| Selective filter | 44 ms | 39 ms | 1.1x | -| Row count only (0 of 6) | 32 ms | 28 ms | 1.1x | -| Narrow projection (1 of 6) | 50 ms | 39 ms | 1.3x | -| Full projection (6 of 6) | 316 ms | 135 ms | 2.3x | +| Repeated scan (3 of 6 columns) | 157 ms | 116 ms | 1.4x | +| Selective filter | 44 ms | 38 ms | 1.1x | +| Row count only (0 of 6) | 30 ms | 28 ms | 1.1x | +| Narrow projection (1 of 6) | 49 ms | 39 ms | 1.3x | +| Full projection (6 of 6) | 299 ms | 135 ms | 2.2x | Read what this compares carefully. Comet execution is on in both columns, so the aggregation runs on Comet either way and only the cache-scan boundary moves: on the left, Spark's diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 24eba56f737..343d4f68c85 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -22,12 +22,11 @@ package org.apache.spark.sql.comet.execution.arrow import scala.collection.JavaConverters._ import scala.util.control.NonFatal -import org.apache.arrow.vector.VectorSchemaRoot -import org.apache.arrow.vector.types.pojo.{Field, Schema} 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.catalyst.util.TypeUtils import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch, SimpleMetricsCachedBatchSerializer} import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.columnar.{DefaultCachedBatch, DefaultCachedBatchSerializer} @@ -38,7 +37,6 @@ import org.apache.spark.storage.StorageLevel import org.apache.spark.unsafe.types.UTF8String import org.apache.comet.{CometArrowAllocator, CometConf} -import org.apache.comet.shims.CometTypeShim import org.apache.comet.vector.NativeUtil /** @@ -73,18 +71,35 @@ private case class CometCachedBatch( * 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 with CometTypeShim { +class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { import ArrowCachedBatchSerializer.supportsSchema private val fallback = new DefaultCachedBatchSerializer() + /** + * How each cached column's bounds are compared, or null for a column that records none. + * + * Spark's own interpreted ordering for the type, which is the comparison its expressions use, + * so bounds recorded with it order the same way a predicate over the column does. That matters + * for collated strings, where it resolves to the collation's comparator rather than byte order, + * and it is why this needs no per-Spark-version shim: the collation awareness comes from Spark. + * + * Resolved once per partition rather than per row -- `getInterpretedOrdering` walks the type + * and, for a collated string, looks the collation up by id. + */ + private def boundsOrderings(attrs: Seq[Attribute]): Array[Ordering[Any]] = + attrs.map { attr => + if (tracksBounds(attr.dataType)) TypeUtils.getInterpretedOrdering(attr.dataType) else null + }.toArray + // 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]) = { + attrs: Seq[Attribute], + orderings: Array[Ordering[Any]]): (Array[Any], Array[Any], Array[Int]) = { val numCols = attrs.length val lower = new Array[Any](numCols) val upper = new Array[Any](numCols) @@ -95,16 +110,17 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with while (c < numCols) { val dt = attrs(c).dataType val col = batch.column(c) + val ordering = orderings(c) var r = 0 while (r < numRows) { if (col.isNullAt(r)) { nulls(c) += 1 - } else if (tracksBounds(dt)) { + } else if (ordering != null) { val value = readValue(col, dt, r) - if (lower(c) == null || compare(dt, value, lower(c)) < 0) { + if (lower(c) == null || ordering.compare(value, lower(c)) < 0) { lower(c) = value } - if (upper(c) == null || compare(dt, value, upper(c)) > 0) { + if (upper(c) == null || ordering.compare(value, upper(c)) > 0) { upper(c) = value } } @@ -149,11 +165,10 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with // 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. // - // Every StringType qualifies, collated ones included: bounds are recorded with the collation's - // own comparison, which is the same ordering the predicate Spark generates over that column - // uses. Matching the bare `StringType` object instead would exclude collated columns, since a - // collated StringType is not equal to the default one, and they would then get null bounds and - // no pruning at all. + // Every StringType qualifies, collated ones included. Matching the bare `StringType` object + // instead would exclude them, since a collated StringType is not equal to the default one, and + // they would then get null bounds and no pruning at all. See boundsOrderings for how a collated + // column's bounds are compared. private def tracksBounds(dt: DataType): Boolean = dt match { case BooleanType | ByteType | ShortType | IntegerType | LongType | FloatType | DoubleType | _: DecimalType | _: StringType | DateType | TimestampType | TimestampNTZType => @@ -176,30 +191,6 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with 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 st: StringType => - compareStrings(left.asInstanceOf[UTF8String], right.asInstanceOf[UTF8String], st) - 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. @@ -207,19 +198,31 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with // 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. + /** + * The configured write codec, read on the driver. + * + * Both write paths resolve this here rather than inside their `mapPartitions` closure: the + * closure ships to the executors, where `CometConf` would resolve against whatever `SQLConf` + * happens to be current on that thread rather than against this session's. + */ + private def codecSettings(conf: SQLConf): (String, Int) = + ( + CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.get(conf), + CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL.get(conf)) + private def encodeBatches( batches: Iterator[ColumnarBatch], attrs: Seq[Attribute], - codecName: String, - zstdLevel: Int): Iterator[CachedBatch] = { + codecSetting: (String, Int)): Iterator[CachedBatch] = { val arrowSchema = Utils.toArrowSchema(Utils.fromAttributes(attrs), CometArrowStream.NATIVE_TIMEZONE) - val codec = CachedBatchIpc.compressionCodec(codecName, zstdLevel) + val codec = CachedBatchIpc.compressionCodec(codecSetting._1, codecSetting._2) + val orderings = boundsOrderings(attrs) batches.map { batch => // Bounds and null counts are read from the input batch before it is serialized, and the row // is only assembled once the per-column sizes the message reports are known. - val (lower, upper, nulls) = gatherColumnStats(batch, attrs) + val (lower, upper, nulls) = gatherColumnStats(batch, attrs, orderings) val numRows = batch.numRows() val (bytes, columnSizes) = if (Utils.isArrowBacked(batch)) { @@ -313,13 +316,10 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with storageLevel: StorageLevel, conf: SQLConf): RDD[CachedBatch] = { - // Read on the driver: the closure ships to the executors, where CometConf would resolve - // against whatever SQLConf happens to be current on that thread rather than this session's. - val codecName = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.get(conf) - val zstdLevel = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL.get(conf) + val codec = codecSettings(conf) input.mapPartitions { batches => - encodeBatches(batches, schema, codecName, zstdLevel) + encodeBatches(batches, schema, codec) } } @@ -342,12 +342,15 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with val cacheSchema = Utils.fromAttributes(cacheAttributes) input.mapPartitions { it => - val arrowFields = + // Built once per partition: resolving the Arrow schema and the projection's buffer layout + // walks every field of the cached relation, which would otherwise be paid per batch. + val projection = new CachedBatchIpc.Projection( Utils .toArrowSchema(cacheSchema, CometArrowStream.NATIVE_TIMEZONE) .getFields .asScala - .toSeq + .toIndexedSeq, + indices) // A ProjectedBatch owns the vectors of the batch it produced, and releases them only when // that batch has been consumed. A consumer that stops early -- LIMIT, take(), or a @@ -374,7 +377,7 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with // 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 projected = new ProjectedBatch(cb, arrowFields, indices) + val projected = new ProjectedBatch(cb, projection) current = projected projected.batches } @@ -387,48 +390,28 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with } /** - * Loads the projected columns of one cached batch into Arrow vectors. - * - * The schema is rebuilt from the cached relation's attributes rather than read from the - * payload, which stores none. `CachedBatchIpc.readProjected` then materializes only the - * selected columns' buffers, so the rest are never copied out of the cached bytes or - * decompressed. + * Owns the Arrow vectors decoded for one cached batch. * - * The decoded vectors stay owned by this object: closing it releases the batch, which is why - * this yields a single-element iterator that closes on exhaustion. + * The decode itself belongs to `CachedBatchIpc.Projection`, which is where knowledge of the + * payload format lives; what is left here is ownership. The vectors stay owned by this object + * -- closing it releases them -- which is why this yields a single-element iterator that closes + * on exhaustion. */ - private class ProjectedBatch( - cached: CometCachedBatch, - arrowFields: Seq[Field], - indices: Array[Int]) { - - // Allocated before anything can throw, so that a failure below has a root to release. - private val root = VectorSchemaRoot.create( - new Schema(indices.map(arrowFields).toSeq.asJava), - CometArrowAllocator) + private class ProjectedBatch(cached: CometCachedBatch, projection: CachedBatchIpc.Projection) { + + // Decoding happens during construction, so `batches` below can hand out the root directly. + // `load` releases everything it allocated if it throws, so there is nothing to unwind here. + private val root = projection.load(cached.bytes, CometArrowAllocator) private var closed = false - // Loading happens during construction, so `batches` below can hand out the root directly. - try { - val recordBatch = - CachedBatchIpc.readProjected(cached.bytes, arrowFields, indices, CometArrowAllocator) - try { - CachedBatchIpc.loaderFor(root).load(recordBatch) - } finally { - recordBatch.close() - } - // A cached batch's columns all cover the same rows. Check rather than trust: a mismatch - // would otherwise build a batch whose columns disagree with the row count recorded beside - // them, which reads as corrupt data far from here. - if (root.getRowCount != cached.numRows) { - throw new IllegalStateException( - s"Cached batch decoded ${root.getRowCount} rows, expected ${cached.numRows}") - } - } catch { - case NonFatal(e) => - try root.close() - catch { case NonFatal(closeError) => e.addSuppressed(closeError) } - throw e + // A cached batch's columns all cover the same rows. Check rather than trust: a mismatch would + // otherwise build a batch whose columns disagree with the row count recorded beside them, + // which reads as corrupt data far from here. + if (root.getRowCount != cached.numRows) { + val decoded = root.getRowCount + close() + throw new IllegalStateException( + s"Cached batch decoded $decoded rows, expected ${cached.numRows}") } def close(): Unit = synchronized { @@ -471,8 +454,7 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with fallback.convertInternalRowToCachedBatch(input, schema, storageLevel, conf) } else { val batchSize = conf.columnBatchSize - val codecName = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.get(conf) - val zstdLevel = CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL.get(conf) + val codec = codecSettings(conf) input.mapPartitions { rows => val iter = CometArrowConverters.rowToArrowBatchIter( @@ -490,7 +472,7 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer with CometArrowStream.NATIVE_TIMEZONE, CometArrowAllocator) - encodeBatches(iter, schema, codecName, zstdLevel) + encodeBatches(iter, schema, codec) } } } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala index 55215ef548a..43a78ab858b 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala @@ -29,12 +29,13 @@ import scala.util.control.NonFatal import org.apache.arrow.compression.{CommonsCompressionFactory, ZstdCompressionCodec} import org.apache.arrow.flatbuf.{RecordBatch => FlatBufRecordBatch} import org.apache.arrow.memory.{ArrowBuf, BufferAllocator} -import org.apache.arrow.vector.{FieldVector, TypeLayout, ValueVector, VectorSchemaRoot, VectorUnloader} +import org.apache.arrow.vector.{FieldVector, TypeLayout, ValueVector, VectorLoader, VectorSchemaRoot, VectorUnloader} import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil, NoCompressionCodec} import org.apache.arrow.vector.dictionary.DictionaryEncoder import org.apache.arrow.vector.ipc.{ReadChannel, WriteChannel} import org.apache.arrow.vector.ipc.message.{ArrowBodyCompression, ArrowFieldNode, ArrowRecordBatch, MessageSerializer} -import org.apache.arrow.vector.types.pojo.{ArrowType, Field} +import org.apache.arrow.vector.types.pojo.{ArrowType, Field, Schema} +import org.apache.arrow.vector.util.DataSizeRoundingUtil import org.apache.spark.SparkException import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.vectorized.ColumnarBatch @@ -80,6 +81,24 @@ private[comet] object CachedBatchIpc { "Supported values: none, zstd") } + // Room for the encapsulated metadata message that precedes the body. The message is a small + // flatbuffer whose size grows with the field count, not the data, so this is a starting size for + // the output buffer rather than a bound -- it grows if a very wide schema needs more. + private val METADATA_SIZE_HINT = 8 * 1024 + + // Decompressors are stateless and shared. Resolving one per cached batch would allocate a codec + // per batch on every scan, and the enum lookup walks the CodecType values each time. + private val readCodecs: Map[CompressionUtil.CodecType, CompressionCodec] = + CompressionUtil.CodecType + .values() + .filter(_ != CompressionUtil.CodecType.NO_COMPRESSION) + .map(t => t -> CommonsCompressionFactory.INSTANCE.createCodec(t)) + .toMap + + /** The decompressor for a body-compression byte, or None when the batch is stored plain. */ + private def readCodec(compressionType: Byte): Option[CompressionCodec] = + readCodecs.get(CompressionUtil.CodecType.fromCompressionType(compressionType)) + /** * Serialize `batch` into one encapsulated IPC RecordBatch message. * @@ -125,7 +144,12 @@ private[comet] object CachedBatchIpc { // does the same, so both writers leave a batch they were handed in the same state. root.clear() - val out = new ByteArrayOutputStream() + // Sized up front from the body length the record batch already knows, plus room for the + // metadata message. An unsized ByteArrayOutputStream starts at 32 bytes and doubles, so a + // multi-MiB payload would be reallocated and recopied a dozen-odd times per batch. + val sizeHint = recordBatch.computeBodyLength() + METADATA_SIZE_HINT + val out = new ByteArrayOutputStream( + math.min(math.max(sizeHint, METADATA_SIZE_HINT), Int.MaxValue.toLong).toInt) val channel = new WriteChannel(Channels.newChannel(out)) MessageSerializer.serialize(channel, recordBatch) (out.toByteArray, columnSizes(fields, recordBatch)) @@ -141,111 +165,151 @@ private[comet] object CachedBatchIpc { } /** - * Read an encapsulated IPC RecordBatch message, materializing off-heap only the buffers of the - * requested top-level columns. + * Everything about reading one projection of this format that does not change between batches. * - * The body is a flat, depth-first sequence of buffers in schema order, so each top-level column - * owns a contiguous run of buffers whose length is [[fieldBufferCount]]; field nodes and - * variadic buffer counts run in the same order. The selected columns' bytes are copied into a - * single off-heap allocation, each buffer 8-byte aligned exactly as Arrow's IPC body lays them - * out, and the returned batch's buffers are windows into it -- one allocation, no per-buffer - * bookkeeping. + * The index arithmetic here is a pure function of the cached schema and the selected columns, + * both fixed for the life of a scan, but it walks every field of the whole relation rather than + * just the projected ones. Recomputing it per batch would make the bookkeeping O(total columns) + * while the useful work is O(selected columns) -- worst in exactly the wide-relation, + * narrow-projection case this format exists for. A scan builds one of these per partition. * - * Only the selected buffers are ever decompressed. A buffer's recorded (offset, length) covers - * its on-body bytes including the uncompressed-length prefix, so a copied window is exactly - * what the writer emitted; the columns that were not selected are never read, let alone - * inflated. The copied windows are then decompressed in one pass -- see [[decompressed]] for - * why that is not left to `VectorLoader` -- so what comes back is an uncompressed batch. - * - * The returned batch owns its buffers; the caller closes it. + * Holding the projected `Schema` here too is what keeps it consistent with the buffers: + * [[load]] packs field nodes and buffers by walking `selectedIndices` in order, and the schema + * is built from the same walk, so the two cannot drift apart. */ - def readProjected( - data: Array[Byte], - schemaFields: Seq[Field], - selectedIndices: Array[Int], - allocator: BufferAllocator): ArrowRecordBatch = { - val readChannel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) - // Reads the message metadata only. The body stays in `data` and is copied selectively below. - val metadata = MessageSerializer.readMessage(readChannel) - if (metadata == null) { - throw new SparkException("Unexpected end of input reading a Comet cached batch") - } - val batch = - metadata.getMessage.header(new FlatBufRecordBatch()).asInstanceOf[FlatBufRecordBatch] - // serialize writes exactly [encapsulated message][body] and nothing after it, so the body is - // the tail of `data`. - val bodyStart = data.length - metadata.getMessageBodyLength.toInt + final class Projection(arrowFields: Seq[Field], selectedIndices: Array[Int]) { - val compression = - if (batch.compression() == null) NoCompressionCodec.DEFAULT_BODY_COMPRESSION - else new ArrowBodyCompression(batch.compression().codec(), batch.compression().method()) + private val schema = new Schema(selectedIndices.map(arrowFields).toSeq.asJava) - val nodeStarts = schemaFields.scanLeft(0)(_ + fieldNodeCount(_)).toArray - val bufferStarts = schemaFields.scanLeft(0)(_ + fieldBufferCount(_)).toArray - val variadicStarts = schemaFields.scanLeft(0)(_ + fieldVariadicCount(_)).toArray - val hasVariadic = batch.variadicBufferCountsLength() > 0 + // A record batch body is a flat, depth-first sequence of buffers in schema order, so each + // top-level column owns a contiguous run of it; field nodes and variadic buffer counts run in + // the same order. + private val nodeIndices = selectedRange(arrowFields, selectedIndices, fieldNodeCount) + private val bufferIndices = selectedRange(arrowFields, selectedIndices, fieldBufferCount) + private val variadicIndices = selectedRange(arrowFields, selectedIndices, fieldVariadicCount) + + /** + * Decode the projected columns of one cached payload into a fresh root the caller owns. + * + * Only the selected buffers are ever materialized off-heap or decompressed. The message + * metadata records every buffer's offset and length within the body, so the selected columns' + * bytes are copied into a single allocation -- each 8-byte aligned exactly as Arrow's IPC + * body lays them out -- and the columns that were not selected are never read, let alone + * inflated. + * + * A buffer's recorded (offset, length) covers its on-body bytes including the + * uncompressed-length prefix, so a copied window is exactly what the writer emitted. The + * windows are then decompressed in one pass; see [[decompressed]] for why that is not left to + * `VectorLoader`. + */ + def load(data: Array[Byte], allocator: BufferAllocator): VectorSchemaRoot = { + val readChannel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) + // Reads the message metadata only. The body stays in `data` and is copied selectively. + val metadata = MessageSerializer.readMessage(readChannel) + if (metadata == null) { + throw new SparkException("Unexpected end of input reading a Comet cached batch") + } + val batch = + metadata.getMessage.header(new FlatBufRecordBatch()).asInstanceOf[FlatBufRecordBatch] + // serialize writes exactly [encapsulated message][body] and nothing after it, so the body is + // the tail of `data`. + val bodyStart = data.length - metadata.getMessageBodyLength.toInt - // The selected columns' field nodes, buffer indices and variadic counts, in output order. - val nodes = new java.util.ArrayList[ArrowFieldNode]() - val bufferIndices = mutable.ArrayBuffer.empty[Int] - val variadicCounts = new java.util.ArrayList[java.lang.Long]() - selectedIndices.foreach { i => - val field = schemaFields(i) - val nodeStart = nodeStarts(i) - (nodeStart until nodeStart + fieldNodeCount(field)).foreach { j => + val compression = + if (batch.compression() == null) NoCompressionCodec.DEFAULT_BODY_COMPRESSION + else new ArrowBodyCompression(batch.compression().codec(), batch.compression().method()) + + val nodes = new java.util.ArrayList[ArrowFieldNode](nodeIndices.length) + nodeIndices.foreach { j => val node = batch.nodes(j) nodes.add(new ArrowFieldNode(node.length(), node.nullCount())) } - val bufferStart = bufferStarts(i) - (bufferStart until bufferStart + fieldBufferCount(field)).foreach(bufferIndices += _) - if (hasVariadic) { - val variadicStart = variadicStarts(i) - (variadicStart until variadicStart + fieldVariadicCount(field)) - .foreach(j => variadicCounts.add(batch.variadicBufferCounts(j))) + val variadicCounts = new java.util.ArrayList[java.lang.Long](variadicIndices.length) + if (batch.variadicBufferCountsLength() > 0) { + variadicIndices.foreach(j => variadicCounts.add(batch.variadicBufferCounts(j))) } - } - val layout = bufferIndices.map { j => - val buffer = batch.buffers(j) - (buffer.offset(), buffer.length()) - } - val alignedSizes = layout.map { case (_, length) => ((length + 7) / 8) * 8 } - // allocator.buffer(0) is legal but yields a buffer no window can be sliced from, and an - // all-empty projection (every selected column a NullVector, say) would ask for exactly that. - val body = allocator.buffer(math.max(alignedSizes.sum, 1L)) - val compressedBatch = - try { - val buffers = new java.util.ArrayList[ArrowBuf]() - var position = 0L - layout.indices.foreach { k => - val (sourceOffset, length) = layout(k) - if (length > 0) { - body.setBytes(position, data, bodyStart + sourceOffset.toInt, length.toInt) + val offsets = new Array[Long](bufferIndices.length) + val lengths = new Array[Long](bufferIndices.length) + var total = 0L + var k = 0 + while (k < bufferIndices.length) { + val buffer = batch.buffers(bufferIndices(k)) + offsets(k) = buffer.offset() + lengths(k) = buffer.length() + total += DataSizeRoundingUtil.roundUpTo8Multiple(lengths(k)) + k += 1 + } + + // allocator.buffer(0) is legal but yields a buffer no window can be sliced from, and an + // all-empty projection (every selected column a NullVector, say) would ask for exactly that. + val body = allocator.buffer(math.max(total, 1L)) + val compressedBatch = + try { + val buffers = new java.util.ArrayList[ArrowBuf](bufferIndices.length) + var position = 0L + var i = 0 + while (i < bufferIndices.length) { + val length = lengths(i) + if (length > 0) { + body.setBytes(position, data, bodyStart + offsets(i).toInt, length.toInt) + } + val window = body.slice(position, length) + window.writerIndex(length) + buffers.add(window) + position += DataSizeRoundingUtil.roundUpTo8Multiple(length) + i += 1 } - val window = body.slice(position, length) - window.writerIndex(length) - buffers.add(window) - position += alignedSizes(k) + new ArrowRecordBatch( + batch.length().toInt, + nodes, + buffers, + compression, + variadicCounts, + false) + } catch { + case NonFatal(e) => + body.close() + throw e } - new ArrowRecordBatch( - batch.length().toInt, - nodes, - buffers, - compression, - variadicCounts, - false) + + // The constructor retained each window; slice() alone does not. Dropping `body`'s own + // reference leaves the batch as sole owner of the one allocation, so closing the batch is + // what frees it -- and closing `body` again would drive its reference count negative. + body.close() + val plainBatch = + try decompressed(compressedBatch, allocator) + finally compressedBatch.close() + + // The loader needs no compression factory: every buffer is decompressed by this point. + val root = VectorSchemaRoot.create(schema, allocator) + try { + new VectorLoader(root).load(plainBatch) + root } catch { case NonFatal(e) => - body.close() + try root.close() + catch { case NonFatal(closeError) => e.addSuppressed(closeError) } throw e + } finally { + plainBatch.close() } + } + } - // The constructor retained each window; slice() alone does not. Dropping `body`'s own - // reference leaves the batch as sole owner of the one allocation, so closing the batch below - // is what frees it -- and closing `body` again here would drive its reference count negative. - body.close() - try decompressed(compressedBatch, allocator) - finally compressedBatch.close() + /** + * The indices, within a record batch's flat depth-first sequence, that the selected columns + * own. + * + * `count` gives how many entries of the sequence a field occupies including its descendants, so + * a running total over every field turns a column index into its run within the sequence. + */ + private def selectedRange( + arrowFields: Seq[Field], + selectedIndices: Array[Int], + count: Field => Int): Array[Int] = { + val starts = arrowFields.scanLeft(0)(_ + count(_)).toArray + selectedIndices.flatMap(i => starts(i) until starts(i + 1)) } /** @@ -269,14 +333,10 @@ private[comet] object CachedBatchIpc { batch: ArrowRecordBatch, allocator: BufferAllocator): ArrowRecordBatch = { // getCodec is the raw IPC byte; the factory keys off the enum. Both sides of the comparison - // below have to be CodecType: NoCompressionCodec.COMPRESSION_TYPE is the byte -1, and Scala - // compares a CodecType against it by universal equality, which is quietly always unequal. - val codecType = - CompressionUtil.CodecType.fromCompressionType(batch.getBodyCompression.getCodec) - val compressed = codecType != CompressionUtil.CodecType.NO_COMPRESSION - val codec: CompressionCodec = - if (compressed) CommonsCompressionFactory.INSTANCE.createCodec(codecType) - else NoCompressionCodec.INSTANCE + // in readCodec have to be CodecType: NoCompressionCodec.COMPRESSION_TYPE is the byte -1, and + // Scala compares a CodecType against it by universal equality, which is quietly always + // unequal. + val codec = readCodec(batch.getBodyCompression.getCodec) val buffers = new java.util.ArrayList[ArrowBuf]() try { @@ -285,8 +345,10 @@ private[comet] object CachedBatchIpc { val plain = try { // An empty buffer carries no compressed length prefix to read. - if (compressed && buffer.writerIndex() > 0) codec.decompress(allocator, buffer) - else buffer + codec match { + case Some(c) if buffer.writerIndex() > 0 => c.decompress(allocator, buffer) + case _ => buffer + } } catch { case NonFatal(e) => buffer.getReferenceManager.release() @@ -315,15 +377,6 @@ private[comet] object CachedBatchIpc { } } - /** - * A `VectorLoader` for what [[readProjected]] returns. - * - * No compression factory: [[readProjected]] has already decompressed every buffer, so the - * loader only ever sees a batch marked uncompressed. - */ - def loaderFor(root: VectorSchemaRoot): org.apache.arrow.vector.VectorLoader = - new org.apache.arrow.vector.VectorLoader(root) - /** * The on-body compressed size of each top-level column. * @@ -355,16 +408,10 @@ private[comet] object CachedBatchIpc { try { val vectors = Utils.getBatchFieldVectorsWithProviders(batch).map { case (vector, providerOpt) => - val encoding = vector.getField.getDictionary - if (encoding == null) { + if (vector.getField.getDictionary == null) { vector } else { - val dictionary = providerOpt.map(_.lookup(encoding.getId)).orNull - if (dictionary == null) { - throw new SparkException( - s"Column ${vector.getField.getName} is dictionary encoded with ID " + - s"${encoding.getId}, but no dictionary with that ID was provided") - } + val dictionary = Utils.lookupDictionary(vector, providerOpt) val decoded = DictionaryEncoder.decode(vector, dictionary, allocator) hydrated += decoded decoded.asInstanceOf[FieldVector] diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala index f418c31a626..97e78f3eaff 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala @@ -446,12 +446,12 @@ object Utils extends CometTypeShim with Logging { /** * The dictionaries every dictionary-encoded column of `columns` refers to, as one provider. * - * Columns of a batch need not share a provider. Comet's cache decodes each column from its own - * Arrow stream, so a dictionary-backed column arrives carrying the provider its reader built, - * and a batch that reaches [[serializeBatches]] -- a native broadcast of a cache scan, say -- - * can hold several. Writing the whole batch emits one schema covering every column and resolves - * each column's dictionary ID against the single provider the writer was given, so handing it - * any one column's provider fails with "Could not find dictionary with ID n" for the others. + * Columns of a batch need not share a provider. A batch assembled from several upstream readers + * -- a shuffle reader's output, or a broadcast that coalesces many blocks -- carries a + * dictionary-backed column with whichever provider its own reader built, so one batch can hold + * several. Writing the whole batch emits one schema covering every column and resolves each + * column's dictionary ID against the single provider the writer was given, so handing it any + * one column's provider fails with "Could not find dictionary with ID n" for the others. */ private def combineDictionaryProviders( columns: Seq[(FieldVector, Option[DictionaryProvider])]): Option[DictionaryProvider] = { @@ -461,12 +461,7 @@ object Utils extends CometTypeShim with Logging { val encoding = vector.getField.getDictionary if (encoding != null) { val id = encoding.getId - val dictionary = providerOpt.map(_.lookup(id)).orNull - if (dictionary == null) { - throw new SparkException( - s"Column ${vector.getField.getName} is dictionary encoded with ID $id, but no " + - "dictionary with that ID was provided") - } + val dictionary = lookupDictionary(vector, providerOpt) dictionaries.get(id) match { // Every provider seen here descends from one upstream reader, which numbers the // dictionaries it hands out, so two columns sharing an ID share the dictionary itself. @@ -484,13 +479,32 @@ object Utils extends CometTypeShim with Logging { else Some(new MapDictionaryProvider(dictionaries.values.toSeq: _*)) } + /** + * The dictionary a dictionary-encoded column refers to, or a failure naming the column. + * + * Shared with the cache serializer, which decodes dictionary-encoded columns rather than + * folding their providers together, so that both report a missing dictionary the same way. + */ + def lookupDictionary( + vector: FieldVector, + providerOpt: Option[DictionaryProvider]): Dictionary = { + val id = vector.getField.getDictionary.getId + val dictionary = providerOpt.map(_.lookup(id)).orNull + if (dictionary == null) { + throw new SparkException( + s"Column ${vector.getField.getName} is dictionary encoded with ID $id, but no " + + "dictionary with that ID was provided") + } + dictionary + } + /** * Field vectors of `batch` paired with the dictionary provider each column was decoded with. * * [[getBatchFieldVectors]] folds these into one provider covering the whole batch, which is - * what a single stream over every column needs. Comet's cache decodes each column from its own - * stream and writes it back the same way, so it keeps the pairing instead: each column is - * written with the provider it was decoded with. + * what a single stream over every column needs. Comet's cache serializer keeps the pairing + * instead: its payload has no schema message to describe a dictionary encoding, so it decodes + * each such column against the provider that column arrived with. */ def getBatchFieldVectorsWithProviders( batch: ColumnarBatch): Seq[(FieldVector, Option[DictionaryProvider])] = { diff --git a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala index a6b34b74e7d..b71476c3dd1 100644 --- a/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-3.x/org/apache/comet/shims/CometTypeShim.scala @@ -22,7 +22,6 @@ package org.apache.comet.shims import scala.annotation.nowarn import org.apache.spark.sql.types.{DataType, StructType} -import org.apache.spark.unsafe.types.{ByteArray, UTF8String} trait CometTypeShim { @nowarn // Spark 4 feature; stubbed to false in Spark 3.x for compatibility. @@ -42,15 +41,4 @@ trait CometTypeShim { @nowarn // Spark 4.1 feature; TimeType doesn't exist in Spark 3.x. def isTimeType(dt: DataType): Boolean = false - - /** - * Compare two strings under the collation of `dt`, which must be a `StringType`. - * - * Spark 3.x has no collations, so every string comparison is byte order. Callers that record - * comparable bounds (Comet's cache statistics, for instance) use this so the ordering they - * store is the one Spark's own comparison would produce. - */ - @nowarn // Collation is a Spark 4 feature; on 3.x every StringType compares as bytes. - def compareStrings(left: UTF8String, right: UTF8String, dt: DataType): Int = - ByteArray.compareBinary(left.getBytes, right.getBytes) } diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala index 71e72dd2de4..f48955a7da5 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/CometTypeShim.scala @@ -21,7 +21,6 @@ package org.apache.comet.shims import org.apache.spark.sql.execution.datasources.VariantMetadata import org.apache.spark.sql.types.{ArrayType, DataType, MapType, StringType, StructType, VariantType} -import org.apache.spark.unsafe.types.UTF8String trait CometTypeShim { // A `StringType` carries collation metadata in Spark 4.0. Only non-default (non-UTF8_BINARY) @@ -65,14 +64,4 @@ trait CometTypeShim { dt.getClass.getSimpleName.startsWith("TimeType") def hasCollationSupport: Boolean = true - - /** - * Compare two strings under the collation of `dt`, which must be a `StringType`. - * - * `semanticCompare` is the comparison Spark's own expressions use for the type, so bounds - * recorded with it order the same way a predicate over the column does. For the default - * UTF8_BINARY collation it is byte order, which is what Spark 3.x always does. - */ - def compareStrings(left: UTF8String, right: UTF8String, dt: DataType): Int = - left.semanticCompare(right, dt.asInstanceOf[StringType].collationId) } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala index 56c3b949569..e1a089f20f8 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala @@ -27,7 +27,7 @@ import scala.jdk.CollectionConverters._ import org.apache.arrow.flatbuf.{MessageHeader, RecordBatch => FlatBufRecordBatch} import org.apache.arrow.vector.TypeLayout import org.apache.arrow.vector.ipc.ReadChannel -import org.apache.arrow.vector.ipc.message.MessageSerializer +import org.apache.arrow.vector.ipc.message.{MessageMetadataResult, MessageSerializer} import org.apache.arrow.vector.types.pojo.Field import org.apache.spark.sql.columnar.CachedBatch import org.apache.spark.sql.comet.util.Utils @@ -50,9 +50,6 @@ object CometCachedBatchHelper { private def payload(batch: CachedBatch): Array[Byte] = batch.asInstanceOf[CometCachedBatch].bytes - /** Stored size of the whole cached batch, in bytes. */ - def payloadSize(batch: CachedBatch): Long = payload(batch).length.toLong - /** * Whether the payload begins with a Schema message rather than going straight to the record * batch. @@ -62,25 +59,7 @@ object CometCachedBatchHelper { * here rather than only as a footprint number. */ def hasSchemaMessage(batch: CachedBatch): Boolean = - readMetadata(payload(batch))._1.headerType() == MessageHeader.Schema - - /** - * The on-body (offset, length) of every Arrow buffer belonging to each top-level column, in - * column order. - */ - def columnBufferRanges(batch: CachedBatch, cacheSchema: StructType): Seq[Seq[(Long, Long)]] = { - val data = payload(batch) - val (_, recordBatch) = readMetadata(data) - val fields = arrowFields(cacheSchema) - val starts = fields.scanLeft(0)(_ + bufferCount(_)).toArray - - fields.indices.map { i => - (starts(i) until starts(i) + bufferCount(fields(i))).map { j => - val buffer = recordBatch.buffers(j) - (buffer.offset(), buffer.length()) - } - } - } + readMessage(payload(batch)).getMessage.headerType() == MessageHeader.Schema /** Stored size of each top-level column: the sum of its buffers' on-body lengths. */ def columnSizes(batch: CachedBatch, cacheSchema: StructType): Seq[Long] = @@ -94,13 +73,8 @@ object CometCachedBatchHelper { * buffers routinely take that fallback, so [[corruptColumn]] only has something to corrupt when * this is true; the projection tests assert it as a precondition rather than assuming it. */ - def columnIsCompressed(batch: CachedBatch, cacheSchema: StructType, index: Int): Boolean = { - val data = payload(batch) - val start = bodyStart(data) - columnBufferRanges(batch, cacheSchema)(index).exists { case (offset, length) => - length > 8 && uncompressedLength(data, start + offset.toInt) > 0 - } - } + def columnIsCompressed(batch: CachedBatch, cacheSchema: StructType, index: Int): Boolean = + compressedRanges(batch, cacheSchema, index).nonEmpty /** * Scramble one column's compressed bytes in place, leaving every other column byte-identical. @@ -110,38 +84,38 @@ object CometCachedBatchHelper { * difference between decoding what was projected and decoding everything and projecting * afterwards, so it is what the projection tests assert on rather than timings. * - * Each buffer's 8-byte uncompressed-length prefix is left intact and only the compressed bytes - * after it are overwritten, so a read of this column fails while decompressing rather than by - * trying to allocate a nonsense length. Requires the column to have a genuinely compressed - * buffer -- see [[columnIsCompressed]]. + * Requires the column to have a genuinely compressed buffer -- see [[columnIsCompressed]]. */ def corruptColumn(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { - val data = payload(batch) - val start = bodyStart(data) - var corrupted = false - - columnBufferRanges(batch, cacheSchema)(index).foreach { case (offset, length) => - val bufferStart = start + offset.toInt - if (length > 8 && uncompressedLength(data, bufferStart) > 0) { - var i = bufferStart + 8 - while (i < bufferStart + length.toInt) { - // A fixed pattern rather than random bytes, so a failure reproduces exactly. - data(i) = (0xa5 ^ i).toByte - i += 1 - } - corrupted = true - } - } - + val ranges = compressedRanges(batch, cacheSchema, index) require( - corrupted, + ranges.nonEmpty, s"column $index of the cached batch has no compressed buffer to corrupt; " + "the test needs data that Arrow actually compresses") + ranges.foreach { case (start, length) => scramble(payload(batch), start, length) } } /** - * Truncate the tail of one column's compressed bytes, in place, padding with zeros so every - * other column keeps its offset. + * Scramble only the last compressed buffer of one column, leaving its earlier buffers genuine. + * + * A string column stores offsets and data as separate compressed buffers, so this makes the + * decoder decompress one buffer of the column successfully and then fail on the next. That is a + * different failure point from [[corruptColumn]], which takes out a column's first buffer and + * so fails before anything of it has been decompressed. + */ + def corruptTrailingBuffer(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { + val ranges = compressedRanges(batch, cacheSchema, index) + require( + ranges.length > 1, + s"column $index has ${ranges.length} compressed buffers; this needs at least two so a " + + "decode can succeed on one and then fail on the next") + val (start, length) = ranges.last + scramble(payload(batch), start, length) + } + + /** + * Zero the tail of one column's compressed bytes, in place, leaving every other column's bytes + * and offsets untouched. * * [[corruptColumn]] rewrites the whole compressed payload, which fails as soon as the * decompressor looks at it. This keeps the leading bytes genuine, so a decoder gets a stream @@ -150,55 +124,70 @@ object CometCachedBatchHelper { */ def truncateColumn(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { val data = payload(batch) - val start = bodyStart(data) - var truncated = false - - columnBufferRanges(batch, cacheSchema)(index).foreach { case (offset, length) => - val bufferStart = start + offset.toInt - if (length > 32 && uncompressedLength(data, bufferStart) > 0 && !truncated) { - var i = bufferStart + length.toInt - 16 - while (i < bufferStart + length.toInt) { - data(i) = 0 - i += 1 - } - truncated = true - } + val target = compressedRanges(batch, cacheSchema, index).find { case (_, length) => + length > 32 } - require( - truncated, + target.isDefined, s"column $index of the cached batch has no compressed buffer long enough to truncate") + val (start, length) = target.get + java.util.Arrays.fill(data, (start + length - 16).toInt, (start + length).toInt, 0.toByte) } /** - * Scramble only the last compressed buffer of one column, leaving its earlier buffers genuine. + * The absolute (start, length) of each of a column's buffers that Arrow actually compressed. * - * A string column stores offsets and data as separate compressed buffers, so this makes the - * decoder decompress one buffer of the column successfully and then fail on the next. That is a - * different failure point from [[corruptColumn]], which takes out a column's first buffer and - * so fails before anything of it has been decompressed. + * `start` is an index into the payload, not an offset within the body, so callers can write + * through it directly. A buffer shorter than its 8-byte uncompressed-length prefix, or one + * whose prefix reads `-1`, was stored verbatim and is excluded: overwriting it would change the + * values a read returns rather than making the read fail. */ - def corruptTrailingBuffer(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { + private def compressedRanges( + batch: CachedBatch, + cacheSchema: StructType, + index: Int): Seq[(Long, Long)] = { val data = payload(batch) - val start = bodyStart(data) - val compressed = columnBufferRanges(batch, cacheSchema)(index).filter { - case (offset, length) => - length > 8 && uncompressedLength(data, start + offset.toInt) > 0 + val bodyStart = data.length - readMessage(data).getMessageBodyLength + columnBufferRanges(batch, cacheSchema)(index).collect { + case (offset, length) if length > 8 && uncompressedLength(data, bodyStart + offset) > 0 => + (bodyStart + offset, length) } - require( - compressed.length > 1, - s"column $index has ${compressed.length} compressed buffers; this needs at least two so a " + - "decode can succeed on one and then fail on the next") + } - val (offset, length) = compressed.last - val bufferStart = start + offset.toInt - var i = bufferStart + 8 - while (i < bufferStart + length.toInt) { + /** Overwrite a compressed buffer's payload, leaving its uncompressed-length prefix intact. */ + private def scramble(data: Array[Byte], start: Long, length: Long): Unit = { + var i = (start + 8).toInt + val end = (start + length).toInt + while (i < end) { + // A fixed pattern rather than random bytes, so a failure reproduces exactly. data(i) = (0xa5 ^ i).toByte i += 1 } } + /** + * The on-body (offset, length) of every Arrow buffer belonging to each top-level column, in + * column order. + */ + private def columnBufferRanges( + batch: CachedBatch, + cacheSchema: StructType): Seq[Seq[(Long, Long)]] = { + val data = payload(batch) + val recordBatch = + readMessage(data).getMessage + .header(new FlatBufRecordBatch()) + .asInstanceOf[FlatBufRecordBatch] + val fields = arrowFields(cacheSchema) + val starts = fields.scanLeft(0)(_ + bufferCount(_)).toArray + + fields.indices.map { i => + (starts(i) until starts(i + 1)).map { j => + val buffer = recordBatch.buffers(j) + (buffer.offset(), buffer.length()) + } + } + } + /** The Arrow fields the read path rebuilds for `cacheSchema`. */ private def arrowFields(cacheSchema: StructType): Seq[Field] = Utils @@ -215,32 +204,23 @@ object CometCachedBatchHelper { TypeLayout.getTypeBufferCount(field.getType) + field.getChildren.asScala.map(bufferCount).sum - private def readMetadata(data: Array[Byte]) = { - val channel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) - val metadata = MessageSerializer.readMessage(channel) - require(metadata != null, "cached payload holds no IPC message") - ( - metadata.getMessage, - metadata.getMessage.header(new FlatBufRecordBatch()).asInstanceOf[FlatBufRecordBatch]) - } - - /** Offset of the record batch body within the payload; the body is its tail. */ - private def bodyStart(data: Array[Byte]): Int = { + /** The payload's leading IPC message, carrying both its header and its body length. */ + private def readMessage(data: Array[Byte]): MessageMetadataResult = { val channel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) val metadata = MessageSerializer.readMessage(channel) require(metadata != null, "cached payload holds no IPC message") - data.length - metadata.getMessageBodyLength.toInt + metadata } /** * The uncompressed-length prefix Arrow writes ahead of a compressed buffer, little-endian. A * value of -1 means the buffer was stored verbatim because compressing it did not pay. */ - private def uncompressedLength(data: Array[Byte], bufferStart: Int): Long = { + private def uncompressedLength(data: Array[Byte], bufferStart: Long): Long = { var value = 0L var i = 7 while (i >= 0) { - value = (value << 8) | (data(bufferStart + i) & 0xffL) + value = (value << 8) | (data(bufferStart.toInt + i) & 0xffL) i -= 1 } value From ccd469e47b66227564ceca27069540ae47dbfd79 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 29 Aug 2026 09:18:17 -0600 Subject: [PATCH 03/24] fix: relocate the arrow-compression service file when shading arrow-compression ships META-INF/services/org.apache.arrow.vector.compression.CompressionCodec$Factory. The shade plugin copies it verbatim without a ServicesResourceTransformer, so the jar declared a provider for Spark's own unshaded Arrow interface while naming a class that exists here only under the relocated package. Every ServiceLoader lookup Spark's Arrow made then failed with a ServiceConfigurationError, which took CompressionCodec.Factory's static initializer down with it and broke unrelated Arrow IPC reads, including mapInArrow. Add ServicesResourceTransformer so the service file name and its contents are both relocated. arrow-compression is the only bundled artifact that ships one. Also drop an unused NonFatal import that scalafix flagged. --- spark/pom.xml | 10 ++++++++++ .../execution/arrow/ArrowCachedBatchSerializer.scala | 1 - 2 files changed, 10 insertions(+), 1 deletion(-) diff --git a/spark/pom.xml b/spark/pom.xml index d30afc75914..5d8ef9d4c69 100644 --- a/spark/pom.xml +++ b/spark/pom.xml @@ -634,6 +634,16 @@ under the License. ${comet.shade.packageName}.guava.thirdparty + + + + diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 343d4f68c85..6679671a05a 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -20,7 +20,6 @@ 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 From e71d8025c58c443e1d185da9ec89c6e3b5c10524 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 29 Aug 2026 13:41:59 -0600 Subject: [PATCH 04/24] test: drop the cache leak test that depends on zstd corruption detection "releases its vectors when a column fails part way through" zeroed the last 16 bytes of a compressed buffer and required the read to fail. Whether that fails is a property of the zstd runtime, not of Comet: the cached payload is byte-identical across Spark versions, but Comet takes zstd-jni from Spark rather than from arrow-compression, and 1.5.5 (Spark 3.4, 3.5) decodes that frame while 1.5.7 (Spark 4.x) reports it corrupt. So the test passed on 4.x and failed on 3.4 and 3.5. The scenario it claimed to cover is also unreachable: CachedBatchIpc decompresses every selected buffer before VectorLoader runs, so no content corruption can fail part way through the load. The two remaining leak tests corrupt a frame from its header onwards, which every zstd release rejects, and already cover a failure at a column's first buffer and a failure after an earlier buffer of the same column decoded. Records the constraint on scramble so a future test does not reach for a tail-only corruption again, and drops the now unused truncateColumn helper and the dictionary fixture's payload argument. --- .../comet/exec/CometInMemoryCacheSuite.scala | 30 +++--------------- .../arrow/CometCachedBatchHelper.scala | 31 ++++++------------- 2 files changed, 14 insertions(+), 47 deletions(-) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 3464e5e3ddc..3e154094a90 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -1406,12 +1406,12 @@ class CometInMemoryCacheSuite extends CometTestBase { } /** - * Cache two low-cardinality string columns and hand the test the cached payload. + * Cache two low-cardinality string columns and hand the test the cached relation. * * The shuffle is what makes this worth its own fixture: its reader hands the cache writer * dictionary-encoded columns, which the writer has to decode before storing them. */ - private def withDictionaryCache(f: (InMemoryRelation, Array[CachedBatch]) => Unit): Unit = { + private def withDictionaryCache(f: InMemoryRelation => Unit): Unit = { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", @@ -1438,7 +1438,7 @@ class CometInMemoryCacheSuite extends CometTestBase { .cachedRepresentation try { - f(relation, relation.cacheBuilder.cachedColumnBuffers.collect()) + f(relation) } finally { spark.catalog.clearCache() } @@ -1473,7 +1473,7 @@ class CometInMemoryCacheSuite extends CometTestBase { // column either way, so a writer that stored the index vector as-is would hand the loader // integer indices to read as strings. Reading the values back correctly is what proves the // writer decoded them first; a row count alone would not. - withDictionaryCache { (relation, _) => + withDictionaryCache { relation => assert(relation.output.length == 2) val df = spark.sql("SELECT s1, s2 FROM dictionary_cache") @@ -1489,7 +1489,7 @@ class CometInMemoryCacheSuite extends CometTestBase { test("Comet in-memory cache broadcasts a batch read back from the cache") { // A broadcast of a cache scan re-serializes each decoded batch through serializeBatches, // which is a different writer from the one that produced the cached payload. - withDictionaryCache { (relation, _) => + withDictionaryCache { relation => assert(relation.output.length == 2) val df = spark.sql( @@ -1499,26 +1499,6 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache releases its vectors when a column fails part way through") { - // Distinct from the corrupted-column case above: there the compressed bytes are wrong from - // their first byte, so the decompressor rejects them outright. Here the bytes start out - // genuine and only the tail is destroyed, so the failure lands after the loader has already - // begun filling vectors. Nothing else can release them -- the holder's constructor never - // returns, so no caller holds it and the task-completion listener has not been told about it. - withDictionaryCache { (relation, batches) => - val cacheSchema = Utils.fromAttributes(relation.output) - batches.foreach(b => CometCachedBatchHelper.truncateColumn(b, cacheSchema, 0)) - - val before = CometArrowAllocator.getAllocatedMemory - interceptDecodeFailure { - decodedRowCount(relation, batches, Seq(relation.output.head)) - } - assert( - CometArrowAllocator.getAllocatedMemory == before, - "a read that fails while loading must release what it already allocated") - } - } - test("Comet in-memory cache scans of one cache canonicalize equal, so exchanges are reused") { // The wrapped Spark scan is a plan-typed field rather than a child, so canonicalization walks // past it and leaves in place the expression IDs of whichever occurrence of the relation diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala index e1a089f20f8..f8f9430e268 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala @@ -113,27 +113,6 @@ object CometCachedBatchHelper { scramble(payload(batch), start, length) } - /** - * Zero the tail of one column's compressed bytes, in place, leaving every other column's bytes - * and offsets untouched. - * - * [[corruptColumn]] rewrites the whole compressed payload, which fails as soon as the - * decompressor looks at it. This keeps the leading bytes genuine, so a decoder gets a stream - * that starts out valid and then runs out, exercising a failure part way through a column - * rather than at its first byte. - */ - def truncateColumn(batch: CachedBatch, cacheSchema: StructType, index: Int): Unit = { - val data = payload(batch) - val target = compressedRanges(batch, cacheSchema, index).find { case (_, length) => - length > 32 - } - require( - target.isDefined, - s"column $index of the cached batch has no compressed buffer long enough to truncate") - val (start, length) = target.get - java.util.Arrays.fill(data, (start + length - 16).toInt, (start + length).toInt, 0.toByte) - } - /** * The absolute (start, length) of each of a column's buffers that Arrow actually compressed. * @@ -154,7 +133,15 @@ object CometCachedBatchHelper { } } - /** Overwrite a compressed buffer's payload, leaving its uncompressed-length prefix intact. */ + /** + * Overwrite a compressed buffer's payload, leaving its uncompressed-length prefix intact. + * + * The whole payload is rewritten, frame header included, so every zstd release rejects it + * outright. Corrupting only a frame's tail is not enough: whether that is detected depends on + * the zstd-jni each Spark version ships -- Comet takes it from Spark rather than from + * arrow-compression, and 1.5.5 (Spark 3.4, 3.5) decodes a frame whose last bytes have been + * zeroed that 1.5.7 (Spark 4.x) reports as corrupt. + */ private def scramble(data: Array[Byte], start: Long, length: Long): Unit = { var i = (start + 8).toInt val end = (start + length).toInt From d8d4196f2e7402194ab903895a8bd3421da252e5 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 2 Sep 2026 14:33:40 -0600 Subject: [PATCH 05/24] feat: enable Comet's in-memory cache by default Flip spark.comet.exec.inMemoryCache.enabled to true so cached tables are stored and scanned in Comet's Arrow format without an opt-in. CometDriverPlugin.maybeSetCacheSerializer read the config out of SparkConf with a hardcoded false default, so flipping the ConfigEntry alone would have left the serializer uninstalled unless the user set the key explicitly. It now falls back to the entry's own default, matching how the plugin reads spark.comet.metrics.enabled. Stacked on #5543. --- docs/source/user-guide/latest/in-memory-cache.md | 7 ++++--- spark/src/main/scala/org/apache/comet/CometConf.scala | 2 +- spark/src/main/scala/org/apache/spark/Plugins.scala | 4 +++- 3 files changed, 8 insertions(+), 5 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index efd949ff92c..235262eb71d 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -24,10 +24,11 @@ format that Comet operators read directly. Without it, a cached table is stored format and every scan of it has to convert each batch before Comet can continue, which shows up in the plan as a `CometSparkColumnarToColumnar` above the cache scan. -This feature is **experimental and disabled by default**. +This feature is **experimental and enabled by default**. To turn it off, set the config before the +`SparkContext` is created: ```scala -spark.conf.set("spark.comet.exec.inMemoryCache.enabled", "true") +spark.conf.set("spark.comet.exec.inMemoryCache.enabled", "false") ``` ## What changes when it is enabled @@ -85,7 +86,7 @@ nowhere to record either that a column is dictionary encoded or the dictionary i | Config | Default | Description | | ------------------------------------------------------- | ------- | ---------------------------------------------------------------------------------------------------------------------------------------------- | -| `spark.comet.exec.inMemoryCache.enabled` | `false` | Whether to store and scan Spark's in-memory cache in Comet's format. Read at startup. | +| `spark.comet.exec.inMemoryCache.enabled` | `true` | Whether to store and scan Spark's in-memory cache in Comet's format. Read at startup. | | `spark.comet.exec.inMemoryCache.compression.codec` | `zstd` | Arrow IPC compression codec for cached data: `zstd` or `none`. Affects newly cached data only — a batch records the codec it was written with. | | `spark.comet.exec.inMemoryCache.compression.zstd.level` | `1` | Compression level when the codec is `zstd`. Ignored otherwise. | diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 4183ffef457..92bad31952f 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -276,7 +276,7 @@ object CometConf extends ShimCometConf { "SparkContext, otherwise caching fails as soon as a block is serialized, including " + "the disk half of the default MEMORY_AND_DISK storage level.") .booleanConf - .createWithDefault(false) + .createWithDefault(true) val COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC: ConfigEntry[String] = conf("spark.comet.exec.inMemoryCache.compression.codec") diff --git a/spark/src/main/scala/org/apache/spark/Plugins.scala b/spark/src/main/scala/org/apache/spark/Plugins.scala index eaeac316655..9736d523633 100644 --- a/spark/src/main/scala/org/apache/spark/Plugins.scala +++ b/spark/src/main/scala/org/apache/spark/Plugins.scala @@ -130,7 +130,9 @@ object CometDriverPlugin extends Logging { private[apache] def maybeSetCacheSerializer( conf: SparkConf, extraConfs: ju.HashMap[String, String]): Unit = { - if (conf.getBoolean(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key, false)) { + if (conf.getBoolean( + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key, + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.defaultValue.get)) { val serializerKey = StaticSQLConf.SPARK_CACHE_SERIALIZER.key val serializerValue = "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer" From 4010f78e7d007fc2ef93d8246855ff84b5b14062 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 5 Sep 2026 10:02:41 -0600 Subject: [PATCH 06/24] test: cover nested columns in the cached-batch projection tests and benchmark Addresses review feedback asking whether nested data should be tested and benchmarked. Nested columns were already round-tripped, but only under a full projection, which cannot see the part of the format that is nontrivial for them. A flat column always owns one field node and two or three buffers; a nested one owns a run as long as its subtree, and selecting every column covers the whole sequence however it is partitioned. So the buffer-span arithmetic was only exercised in the one shape where getting it wrong does not show. Adds two tests over a six-column relation whose middle four columns are a struct, an array, a map and a struct wrapping an array: - Each column takes its turn as the sole projection while the other five are corrupted, so a run computed short or long is caught by reaching into a corrupted neighbour. - Values are compared against the uncached query across single-column, paired and out-of-order projections. Row counts cannot catch a window that is misaligned but still decompresses, and out-of-order is the case a full projection cannot stand in for. The per-column statistics test now runs over the nested relation too, since a nested column's recorded size is the sum of its whole subtree. Both new tests fail if fieldNodeCount stops recursing into children. In the benchmark, adds the three projection widths over a relation of struct columns, and asserts the width each case claims. That assertion caught the existing "full projection (6 of 6 columns)" case reading three: count() over a non-nullable column is rewritten to count(1) by NullPropagation, which prunes the column out of the scan, and only k, s1 and s2 were nullable -- and those only incidentally, because Remainder can divide by zero. Every column of both relations is now nullable so count(c) genuinely reads c, and the documented numbers are regenerated. Array and map columns are left out of the benchmark deliberately: the baseline arm needs Spark's cache scan to bridge into Comet operators, and CometSparkToColumnarExec declines ArrayType and MapType, so for those the arm does not exist and the two cases stop measuring the same boundary. The docs say so rather than leaving it to be rediscovered. --- .../user-guide/latest/in-memory-cache.md | 36 +++- .../comet/exec/CometInMemoryCacheSuite.scala | 165 +++++++++++++--- .../CometInMemoryCacheBenchmark.scala | 180 +++++++++++++++--- 3 files changed, 325 insertions(+), 56 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index efd949ff92c..e0e2755c74d 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -91,21 +91,43 @@ nowhere to record either that a column is dictionary encoded or the dictionary i ## Performance -Measured with `CometInMemoryCacheBenchmark` on a 5M-row, six-column relation (Apple M3 Ultra, -JDK 17, Spark 4.1, release build). Regenerate with: +Measured with `CometInMemoryCacheBenchmark` (Apple M3 Max, JDK 17, Spark 4.1, release build). +Regenerate with: ```sh SPARK_GENERATE_BENCHMARK_FILES=1 \ make benchmark-org.apache.spark.sql.benchmark.CometInMemoryCacheBenchmark ``` +On a 5M-row relation of six flat columns: + | Query shape | Spark cache scan + convert | `CometInMemoryTableScan` | Relative | | ------------------------------ | -------------------------: | -----------------------: | -------: | -| Repeated scan (3 of 6 columns) | 157 ms | 116 ms | 1.4x | -| Selective filter | 44 ms | 38 ms | 1.1x | -| Row count only (0 of 6) | 30 ms | 28 ms | 1.1x | -| Narrow projection (1 of 6) | 49 ms | 39 ms | 1.3x | -| Full projection (6 of 6) | 299 ms | 135 ms | 2.2x | +| Repeated scan (3 of 6 columns) | 201 ms | 167 ms | 1.2x | +| Selective filter | 69 ms | 61 ms | 1.1x | +| Row count only (0 of 6) | 45 ms | 47 ms | 1.0x | +| Narrow projection (1 of 6) | 70 ms | 57 ms | 1.2x | +| Full projection (6 of 6) | 556 ms | 290 ms | 1.9x | + +And on a 1M-row relation of six columns whose middle three are structs, one of them nested two +levels deep: + +| Query shape | Spark cache scan + convert | `CometInMemoryTableScan` | Relative | +| -------------------------- | -------------------------: | -----------------------: | -------: | +| Row count only (0 of 6) | 39 ms | 35 ms | 1.1x | +| Narrow projection (1 of 6) | 109 ms | 61 ms | 1.8x | +| Full projection (6 of 6) | 282 ms | 126 ms | 2.2x | + +The two relations are not comparable to each other — different row counts, and a struct column +carries several values per row. Within the struct relation the gap is wider than the flat one at +every width, because the conversion the left column pays scales with the values per row rather than +with the columns. + +Array and map columns are deliberately absent from the benchmark, not from the format — the cache +stores and projects them, and `CometInMemoryCacheSuite` covers them. They cannot be measured _here_ +because the left column would not exist: it needs Spark's cache scan to bridge into Comet operators, +and `CometSparkToColumnarExec` declines `ArrayType` and `MapType`, so a query projecting one falls +back to Spark row execution above the scan and the two columns stop measuring the same boundary. Read what this compares carefully. Comet execution is on in both columns, so the aggregation runs on Comet either way and only the cache-scan boundary moves: on the left, Spark's diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 3e154094a90..81e54db6a1c 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -1019,11 +1019,49 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + // Enough rows that every column's buffers are big enough for Arrow to actually compress them. + // Arrow stores a buffer verbatim when compressing it would not make it smaller, and a boolean + // column of a few hundred rows is a few dozen bytes, which takes that fallback -- leaving the + // corruption the projection tests rely on with nothing to corrupt. + private val projectionCacheRows = 8000 + + private val flatProjectionColumns = Seq( + "id", + "id % 100 AS k", + "cast(id as double) / 3 AS d", + "concat('a_', cast(id as string)) AS s1", + "concat('b_', cast(id % 17 as string)) AS s2", + "cast(id % 2 = 0 as boolean) AS flag") + + // A flat column always owns one field node and two or three buffers; a nested one owns a run + // whose length is a property of its whole subtree. That arithmetic is what turns a column index + // into a window of the payload, so a run computed short or long by a single buffer misaligns + // every column after it -- which a full projection cannot see, because selecting everything + // covers the whole sequence however it is partitioned. These are the shapes that exercise it: a + // struct, an array, a map (which Arrow stores as a list of two-child structs, so one column + // spans four field nodes), a struct wrapping an array, and flat columns on both sides of them. + private val nestedProjectionColumns = Seq( + "id", + "named_struct('a', id, 'b', concat('sa_', cast(id as string))) AS sc", + "array(concat('e0_', cast(id as string)), concat('e1_', cast(id as string))) AS ar", + "map(concat('k_', cast(id % 97 as string)), id) AS mp", + "named_struct('nums', array(id, id + 1, id + 2)) AS deep", + "concat('t_', cast(id as string)) AS tail") + /** - * Cache a six-column relation and hand the collected batches to `f` along with the relation, so - * a test can doctor the payload before decoding it again through the serializer. + * Cache a six-column flat relation and hand the collected batches to `f` along with the + * relation, so a test can doctor the payload before decoding it again through the serializer. */ private def withProjectionCache( + f: (org.apache.spark.sql.execution.columnar.InMemoryRelation, Array[CachedBatch]) => Unit) + : Unit = withCachedProjection("projection_cache", flatProjectionColumns)(f) + + /** The same, over a relation of the same width whose middle four columns are nested. */ + private def withNestedProjectionCache( + f: (org.apache.spark.sql.execution.columnar.InMemoryRelation, Array[CachedBatch]) => Unit) + : Unit = withCachedProjection("nested_projection_cache", nestedProjectionColumns)(f) + + private def withCachedProjection(view: String, columns: Seq[String])( f: (org.apache.spark.sql.execution.columnar.InMemoryRelation, Array[CachedBatch]) => Unit) : Unit = { withSQLConf( @@ -1033,28 +1071,18 @@ class CometInMemoryCacheSuite extends CometTestBase { SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> "true") { spark.catalog.clearCache() - // Wide enough that every column's buffers are big enough for Arrow to actually compress - // them. Arrow stores a buffer verbatim when compressing it would not make it smaller, and a - // boolean column of a few hundred rows is a few dozen bytes, which takes that fallback -- - // leaving the corruption the projection tests rely on with nothing to corrupt. spark - .range(0, 8000, 1, 2) - .selectExpr( - "id", - "id % 100 AS k", - "cast(id as double) / 3 AS d", - "concat('a_', cast(id as string)) AS s1", - "concat('b_', cast(id % 17 as string)) AS s2", - "cast(id % 2 = 0 as boolean) AS flag") - .createOrReplaceTempView("projection_cache") - spark.catalog.cacheTable("projection_cache") - assert(spark.table("projection_cache").count() == 8000) + .range(0, projectionCacheRows, 1, 2) + .selectExpr(columns: _*) + .createOrReplaceTempView(view) + spark.catalog.cacheTable(view) + assert(spark.table(view).count() == projectionCacheRows) assert( - cachedBatchTypes("projection_cache").sameElements( + cachedBatchTypes(view).sameElements( Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) val relation = spark.sharedState.cacheManager - .lookupCachedData(spark.table("projection_cache")) + .lookupCachedData(spark.table(view)) .get .cachedRepresentation @@ -1189,7 +1217,7 @@ class CometInMemoryCacheSuite extends CometTestBase { } assert( - decodedRowCount(relation, batches, selected) == 8000, + decodedRowCount(relation, batches, selected) == projectionCacheRows, "reading one column must not decompress the other five") batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, selectedIdx)) @@ -1199,6 +1227,87 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + test("Comet in-memory cache decodes only the projected columns of a nested relation") { + // The flat case above pins one column and corrupts the rest. Here every column takes its turn, + // because a nested column's run of buffers is as long as its subtree rather than a fixed two or + // three: a run computed short or long shifts every column after it, so which column is selected + // decides whether the misalignment reaches into a corrupted neighbour. + nestedProjectionColumns.indices.foreach { selectedIdx => + withNestedProjectionCache { (relation, batches) => + val cacheSchema = Utils.fromAttributes(relation.output) + val selected = Seq(relation.output(selectedIdx)) + val name = relation.output(selectedIdx).name + + relation.output.indices.foreach { i => + assert( + batches.forall(b => CometCachedBatchHelper.columnIsCompressed(b, cacheSchema, i)), + s"column ${relation.output(i).name} is not stored compressed, so corrupting it " + + "would prove nothing") + } + + relation.output.indices.filter(_ != selectedIdx).foreach { i => + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) + } + + assert( + decodedRowCount(relation, batches, selected) == projectionCacheRows, + s"reading $name must not decompress the other ${relation.output.length - 1} columns") + + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, selectedIdx)) + interceptDecodeFailure { + decodedRowCount(relation, batches, selected) + } + } + } + } + + test("Comet in-memory cache reads correct nested values under a narrow projection") { + // The corruption test above proves the projected read leaves the other columns' bytes alone, + // but it asserts on row counts, and a row count comes from the record batch header rather than + // from any buffer. A window taken from the wrong place within the selected column's own subtree + // -- a child's buffers swapped, say -- still decodes to the right number of rows and the wrong + // values. Comparing values against the uncached query is what rules that out. + withNativeCache { + val query = + s"SELECT ${nestedProjectionColumns.mkString(", ")} FROM range($projectionCacheRows)" + spark.sql(query).createOrReplaceTempView("nested_value_cache") + spark.catalog.cacheTable("nested_value_cache") + spark.table("nested_value_cache").count() + + assert( + cachedBatchTypes("nested_value_cache").sameElements( + Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) + + val names = Seq("id", "sc", "ar", "mp", "deep", "tail") + + // Each nested column on its own, then paired with `id`, then two projections that ask for + // columns out of cache-schema order, then the whole relation. Spark selects cached columns in + // whatever order the query wants them, and a full projection cannot stand in for that: with + // every column selected in order, the projected schema and the node/buffer windows are both + // the whole sequence, so nothing distinguishes them from windows taken in a different order. + val projections = + names.filter(_ != "id").map(Seq(_)) ++ + names.filter(_ != "id").map(n => Seq("id", n)) ++ + Seq(Seq("tail", "mp", "id"), Seq("deep", "sc")) ++ + Seq(names) + + projections.foreach { cols => + val list = cols.mkString(", ") + // Ordering by the JSON form rather than by the columns themselves, since a projection that + // excludes `id` has no orderable key of its own and ORDER BY over a map is not allowed. + val ordered = s"SELECT to_json(struct($list)) AS j FROM %s ORDER BY j" + val expected = spark.sql(ordered.format(s"($query)")).collect() + assert(expected.length == projectionCacheRows) + + val df = spark.sql(ordered.format("nested_value_cache")) + assert( + df.queryExecution.executedPlan.toString().contains("CometInMemoryTableScan"), + s"projection ($list) should read the cache natively") + assert(df.collect() === expected, s"projection ($list) read the wrong values") + } + } + } + test("Comet in-memory cache decodes no columns for a row-count-only read") { // SELECT count(*) selects no columns. Every column's bytes are corrupted, so the read can // only succeed by touching none of them and answering from the row count the cached batch @@ -1209,15 +1318,19 @@ class CometInMemoryCacheSuite extends CometTestBase { batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) } - assert(decodedRowCount(relation, batches, Seq.empty) == 8000) + assert(decodedRowCount(relation, batches, Seq.empty) == projectionCacheRows) } } test("Comet in-memory cache records per-column sizes in its statistics") { // SimpleMetricsCachedBatch reserves a fifth field per column for its size. A column owns a // known run of buffers in the payload, so the real stored size is known and must be reported - // rather than left at zero. - withProjectionCache { (relation, batches) => + // rather than left at zero. Run over the nested relation as well: a nested column's size is + // the sum of its whole subtree, so this is also where a size attributed to the wrong column + // surfaces. + def checkSizes( + relation: org.apache.spark.sql.execution.columnar.InMemoryRelation, + batches: Array[CachedBatch]): Unit = { val cacheSchema = Utils.fromAttributes(relation.output) batches.foreach { batch => val sizes = CometCachedBatchHelper.columnSizes(batch, cacheSchema) @@ -1225,10 +1338,14 @@ class CometInMemoryCacheSuite extends CometTestBase { sizes.zipWithIndex.foreach { case (size, i) => assert( stats.getLong(i * 5 + 4) == size, - s"column $i should report the stored size of its own buffers in the statistics row") + s"column ${relation.output(i).name} should report the stored size of its own " + + "buffers in the statistics row") } } } + + withProjectionCache(checkSizes) + withNestedProjectionCache(checkSizes) } test("Comet in-memory cache scans no columns for a row-count-only query") { diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index c269dc220c6..10c81ac001e 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -22,14 +22,46 @@ package org.apache.spark.sql.benchmark import org.apache.spark.SparkConf import org.apache.spark.benchmark.Benchmark import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.comet.CometInMemoryTableScanExec +import org.apache.spark.sql.execution.columnar.InMemoryTableScanExec import org.apache.spark.sql.internal.SQLConf import org.apache.comet.{CometConf, CometSparkSessionExtensions} object CometInMemoryCacheBenchmark extends CometBenchmarkBase { private val numRows = 5 * 1000 * 1000 + + // A struct column holds several values per row, so caching the nested relation at the flat row + // count would multiply its footprint for no extra insight. A fifth of the rows keeps the two in + // the same order of magnitude; the arms are only ever compared against their own relation, never + // across the two. + private val nestedNumRows = 1000 * 1000 + private val cacheTable = "comet_cache_bench" private val sourceTable = "comet_cache_bench_src" + private val nestedCacheTable = "comet_cache_bench_nested" + private val nestedSourceTable = "comet_cache_bench_nested_src" + + /** + * A relation cached once and then read under several projections. + * + * `columns` is the select list that builds the cached relation, so the projection widths the + * case labels quote are counted against it. + */ + private case class CachedRelation( + table: String, + source: String, + columns: Seq[String], + rows: Int) + + private val flatRelation = + CachedRelation(cacheTable, sourceTable, Seq("id", "k", "v", "s1", "s2", "s3"), numRows) + + private val nestedRelation = CachedRelation( + nestedCacheTable, + nestedSourceTable, + Seq("id", "sc", "deep", "wide", "tail", "d"), + nestedNumRows) override def getSparkSession: SparkSession = { val conf = new SparkConf() @@ -61,58 +93,141 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { } override def runCometBenchmark(args: Array[String]): Unit = { - withTempTable(sourceTable, cacheTable) { + withTempTable(sourceTable, cacheTable, nestedSourceTable, nestedCacheTable) { + // Every column nullable, in both relations, so that `count(c)` genuinely reads c. Spark's + // NullPropagation rewrites a count over a non-nullable column to `count(1)`, which then + // prunes that column out of the scan -- and since the projection-width cases below measure + // nothing but which columns are read, a case labelled "6 of 6" would quietly be measuring + // three. verifyPlan now asserts the widths rather than trusting them. Note that a column is + // otherwise nullable or not for incidental reasons: `id % 1000` is nullable only because + // Remainder can divide by zero, while `id + 1` is not, so relying on that is what let the + // mislabelling through in the first place. spark .range(0, numRows, 1, 16) .selectExpr( - "id", - "id % 1000 AS k", - "id + 1 AS v", - "concat('str_a_', cast(id % 100000 as string)) AS s1", - "concat('str_b_', cast(id % 7919 as string)) AS s2", - "concat('str_c_', cast(id as string)) AS s3") + "if(id % 8 = 0, null, id) AS id", + "if(id % 8 = 1, null, id % 1000) AS k", + "if(id % 8 = 2, null, id + 1) AS v", + "if(id % 8 = 3, null, concat('str_a_', cast(id % 100000 as string))) AS s1", + "if(id % 8 = 4, null, concat('str_b_', cast(id % 7919 as string))) AS s2", + "if(id % 8 = 5, null, concat('str_c_', cast(id as string))) AS s3") .createOrReplaceTempView(sourceTable) + // Struct columns, not arrays or maps. The baseline arm needs Spark's cache scan to bridge into + // Comet operators, and CometSparkToColumnarExec declines ArrayType and MapType outright, so + // for a relation projecting one of those the arm simply does not exist -- the partial + // aggregate stays on Spark and the two cases stop being a scan-boundary comparison. Structs + // are what can be measured here, and they are the shape that matters for the format anyway: a + // struct is where one cached column owns several field nodes and a validity buffer per level. + // Array and map coverage lives in CometInMemoryCacheSuite instead. + // + // The structs themselves are non-nullable and carry nullable fields, rather than the other way + // round. Comet cannot evaluate `if(c, null, named_struct(...))` at all: the Spark type keeps + // saying the fields are non-nullable while the batch has nulls in them wherever the parent is + // null, and native execution rejects that with "Cannot cast nullable struct field to + // non-nullable field". That is a CometProject limitation hit while the source rows are built, + // nothing to do with the cache. Counting a nullable field reads the whole column regardless, + // since the cache scan selects whole top-level columns. + spark + .range(0, nestedNumRows, 1, 16) + .selectExpr( + "if(id % 8 = 0, null, id) AS id", + "named_struct(" + + "'a', if(id % 8 = 1, null, id), " + + "'b', concat('sa_', cast(id as string))) AS sc", + "named_struct('n', named_struct(" + + "'v', if(id % 8 = 2, null, id), " + + "'w', concat('sw_', cast(id as string)))) AS deep", + "named_struct(" + + "'p', if(id % 8 = 3, null, id % 1000), " + + "'q', if(id % 8 = 3, null, id + 1), " + + "'r', concat('sr_', cast(id % 7919 as string))) AS wide", + "if(id % 8 = 4, null, concat('t_', cast(id as string))) AS tail", + "if(id % 8 = 5, null, cast(id as double) / 3) AS d") + .createOrReplaceTempView(nestedSourceTable) + runCacheBenchmark( + flatRelation, "in-memory cache repeated scan", - s"SELECT sum(id), sum(k), sum(v) FROM $cacheTable") + s"SELECT sum(id), sum(k), sum(v) FROM $cacheTable", + scanned = 3) runCacheBenchmark( + flatRelation, "in-memory cache selective filter", s""" |SELECT sum(id), sum(k), sum(v) |FROM $cacheTable |WHERE id >= 4500000 AND id < 4750000 - """.stripMargin) + """.stripMargin, + scanned = 3) // A CometCachedBatch records where each column's buffers sit in its payload, so a scan // copies out and decompresses only what it projected and cost tracks the width of the // projection. These three cases span that range over one cached relation: no columns, one // column, and all six. runCacheBenchmark( + flatRelation, "in-memory cache row count only (0 of 6 columns)", - s"SELECT count(*) FROM $cacheTable") + s"SELECT count(*) FROM $cacheTable", + scanned = 0) runCacheBenchmark( + flatRelation, "in-memory cache narrow projection (1 of 6 columns)", - s"SELECT count(k) FROM $cacheTable") + s"SELECT count(k) FROM $cacheTable", + scanned = 1) runCacheBenchmark( + flatRelation, "in-memory cache full projection (6 of 6 columns)", - s"SELECT count(id), count(k), count(v), count(s1), count(s2), count(s3) FROM $cacheTable") + s"SELECT count(id), count(k), count(v), count(s1), count(s2), count(s3) FROM $cacheTable", + scanned = 6) + + // The same three widths over a relation whose columns are structs. A struct column's buffers + // are a run as long as its subtree rather than the two or three a flat column owns, so the + // per-column bookkeeping the projected read does is proportionally a smaller share of the work + // here -- which is what these cases measure against the flat ones above. + // + // The aggregates reach into a field rather than counting the struct whole, because a struct + // built this way is non-nullable and `count(c)` over a non-nullable column is rewritten to + // `count(1)`. Either way the cache scan selects whole top-level columns, so one field is + // enough to decode all of that column's buffers. + runCacheBenchmark( + nestedRelation, + "in-memory cache nested row count only (0 of 6 columns)", + s"SELECT count(*) FROM $nestedCacheTable", + scanned = 0) + + runCacheBenchmark( + nestedRelation, + "in-memory cache nested narrow projection (1 of 6 columns)", + s"SELECT count(deep.n.v) FROM $nestedCacheTable", + scanned = 1) + + runCacheBenchmark( + nestedRelation, + "in-memory cache nested full projection (6 of 6 columns)", + s"SELECT count(id), count(sc.a), count(deep.n.v), count(wide.p), count(tail), count(d) " + + s"FROM $nestedCacheTable", + scanned = 6) } } - private def runCacheBenchmark(name: String, query: String): Unit = { - withCachedTable { + private def runCacheBenchmark( + relation: CachedRelation, + name: String, + query: String, + scanned: Int): Unit = { + withCachedTable(relation) { withSQLConf(cacheConf(nativeCacheEnabled = false): _*) { - verifyPlan(query, nativeCacheEnabled = false) + verifyPlan(query, nativeCacheEnabled = false, scanned) } withSQLConf(cacheConf(nativeCacheEnabled = true): _*) { - verifyPlan(query, nativeCacheEnabled = true) + verifyPlan(query, nativeCacheEnabled = true, scanned) } - val benchmark = new Benchmark(name, numRows, output = output) + val benchmark = new Benchmark(name, relation.rows, output = output) benchmark.addCase("Spark cache scan + CometSparkColumnarToColumnar") { _ => withSQLConf(cacheConf(nativeCacheEnabled = false): _*) { @@ -130,7 +245,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { } } - private def withCachedTable(f: => Unit): Unit = { + private def withCachedTable(relation: CachedRelation)(f: => Unit): Unit = { spark.catalog.clearCache() // Materialize the cache once using Comet's cache serializer, then read it both ways. @@ -148,15 +263,15 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { // against; both cases read the same Comet-written CometCachedBatch. withSQLConf(cacheConf(nativeCacheEnabled = true): _*) { spark - .sql(s"SELECT id, k, v, s1, s2, s3 FROM $sourceTable") - .createOrReplaceTempView(cacheTable) - spark.catalog.cacheTable(cacheTable) - spark.table(cacheTable).count() + .sql(s"SELECT ${relation.columns.mkString(", ")} FROM ${relation.source}") + .createOrReplaceTempView(relation.table) + spark.catalog.cacheTable(relation.table) + spark.table(relation.table).count() } try f finally { - spark.catalog.uncacheTable(cacheTable) + spark.catalog.uncacheTable(relation.table) spark.catalog.clearCache() } } @@ -165,8 +280,13 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { // disabled reads it through Spark's cache scan and a CometSparkColumnarToColumnar bridge. The // bridge is what makes the disabled case a scan-boundary comparison rather than a Spark-vs-Comet // execution one, since a Spark-columnar-to-Arrow transition only exists to feed Comet operators. - private def verifyPlan(query: String, nativeCacheEnabled: Boolean): Unit = { - val plan = spark.sql(query).queryExecution.executedPlan.toString() + // + // The projection width is checked too, because a case that reads fewer columns than its label + // says is not slightly off, it is measuring a different query: an optimizer rule that rewrites + // the aggregate can prune a column out of the scan entirely. + private def verifyPlan(query: String, nativeCacheEnabled: Boolean, scanned: Int): Unit = { + val executed = spark.sql(query).queryExecution.executedPlan + val plan = executed.toString() if (nativeCacheEnabled) { assert(plan.contains("CometInMemoryTableScan"), s"Expected native cache scan:\n$plan") @@ -179,6 +299,16 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { plan.contains("CometSparkColumnarToColumnar"), s"Expected the fallback read to bridge into Comet operators:\n$plan") } + + val scanOutputs = executed.collect { + case s: CometInMemoryTableScanExec => s.scanOutput + case s: InMemoryTableScanExec => s.attributes + } + assert(scanOutputs.length == 1, s"Expected exactly one cache scan:\n$plan") + assert( + scanOutputs.head.length == scanned, + s"Expected the scan to read $scanned columns, got " + + s"${scanOutputs.head.map(_.name).mkString("[", ",", "]")}:\n$plan") } private def cacheConf(nativeCacheEnabled: Boolean): Seq[(String, String)] = { From 8c19267bde1a4bb609d7b6c7f1222102fb72715e Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 5 Sep 2026 10:19:27 -0600 Subject: [PATCH 07/24] fix: drop a redundant string interpolator flagged by scalafix RedundantSyntax --- .../spark/sql/benchmark/CometInMemoryCacheBenchmark.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index 10c81ac001e..79b7feaf2bb 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -208,7 +208,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { runCacheBenchmark( nestedRelation, "in-memory cache nested full projection (6 of 6 columns)", - s"SELECT count(id), count(sc.a), count(deep.n.v), count(wide.p), count(tail), count(d) " + + "SELECT count(id), count(sc.a), count(deep.n.v), count(wide.p), count(tail), count(d) " + s"FROM $nestedCacheTable", scanned = 6) } From 3e74e9dcde021a043a9ad6ccd0322ade2279151c Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 15 Sep 2026 10:57:10 -0600 Subject: [PATCH 08/24] review: check the cached layout, and address the rest of the review Reader-side: a cached payload carries no schema, so `Projection` derived every node and buffer window from `Utils.toArrowSchema(cacheAttributes)` with nothing checking the writer had produced that layout. `load` now compares `nodesLength()`/`buffersLength()` against the totals `selectedRange` already computes, before any unchecked `batch.buffers(j)`. Writer-side: `isArrowBacked` accepts a `FixedSizeBinaryVector` for a `BinaryType` column, which is two buffers where the reader rebuilds three, and it answers for the top-level vector only -- so a struct of large strings passes it and is stored with 64-bit offsets. `matchesReaderLayout` compares the batch's Arrow types against the reader's recursively, and a batch that disagrees takes the conversion path instead. A dictionary column's field carries the index type, so the dictionary's field is what is compared. Also: an unrecognized body-compression byte is rejected rather than read as plain bytes, `fieldVariadicCount` and the variadic plumbing are gone (the length check covers view vectors, which the counts would not have), `columnSizes` no longer re-walks each column's subtree, the write codec is a case class rather than a bare tuple, the per-partition `Projection` is lazy so a row-count-only read never builds it, `hydrateDictionaries` is `decodeDictionaries`, `Projection` takes an `IndexedSeq`, and the stale `readProjected` links and some over-long comments are fixed. Tests: the two projection tests become one parameterized over both relations, caching once and restoring the payload between columns instead of re-caching; the two leak tests become one with two corruption points. New tests cover the reader's layout check and the writer declining a fixed-size-binary batch. --- .../arrow/ArrowCachedBatchSerializer.scala | 56 +++-- .../execution/arrow/CachedBatchIpc.scala | 237 +++++++++++------- .../comet/exec/CometInMemoryCacheSuite.scala | 237 +++++++++++------- .../arrow/CometCachedBatchHelper.scala | 27 ++ 4 files changed, 357 insertions(+), 200 deletions(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 5777bb3e76a..86f069ef7e6 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -47,7 +47,7 @@ import org.apache.comet.vector.NativeUtil * and no end-of-stream marker, produced by `CachedBatchIpc.serialize`. Compression is applied per * Arrow buffer rather than over the payload as a whole, which is what lets a scan decompress only * the columns it projected: the message records every buffer's offset and length, so - * `CachedBatchIpc.readProjected` copies out just the selected columns' byte ranges. The cache + * `CachedBatchIpc.Projection.load` copies out just the selected columns' byte ranges. The cache * manager still owns storage and eviction; this class only changes the cached payload. */ private case class CometCachedBatch( @@ -57,6 +57,15 @@ private case class CometCachedBatch( bytes: Array[Byte]) extends SimpleMetricsCachedBatch +/** + * The write codec, resolved on the driver and shipped to the executors in the write closure. + * + * Both write paths resolve it there rather than inside their `mapPartitions` closure: on an + * executor `CometConf` would resolve against whatever `SQLConf` happens to be current on that + * thread rather than against this session's. + */ +private case class CacheCodecSettings(name: String, zstdLevel: Int) + /** * Cache serializer that stores Comet-compatible Arrow batches in Spark's in-memory cache. * @@ -363,32 +372,26 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { case _ => false } - // 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. - /** - * The configured write codec, read on the driver. - * - * Both write paths resolve this here rather than inside their `mapPartitions` closure: the - * closure ships to the executors, where `CometConf` would resolve against whatever `SQLConf` - * happens to be current on that thread rather than against this session's. - */ - private def codecSettings(conf: SQLConf): (String, Int) = - ( + private def codecSettings(conf: SQLConf): CacheCodecSettings = + CacheCodecSettings( CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.get(conf), CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_ZSTD_LEVEL.get(conf)) + // Serialize each batch to Arrow, gathering the Spark-compatible cache stats first. The stats are + // stored beside the Arrow bytes so Spark's cache filter can prune a CometCachedBatch without + // decoding it. + // + // A columnar input batch is not guaranteed to be Arrow-backed, nor to be laid out the way the + // reader will read it; see supportsColumnarInput and CachedBatchIpc.matchesReaderLayout. Batches + // that are not get copied into Arrow first. private def encodeBatches( batches: Iterator[ColumnarBatch], attrs: Seq[Attribute], - codecSetting: (String, Int)): Iterator[CachedBatch] = { + codecSetting: CacheCodecSettings): Iterator[CachedBatch] = { val arrowSchema = Utils.toArrowSchema(Utils.fromAttributes(attrs), CometArrowStream.NATIVE_TIMEZONE) - val codec = CachedBatchIpc.compressionCodec(codecSetting._1, codecSetting._2) + val readerFields = arrowSchema.getFields.asScala.toIndexedSeq + val codec = CachedBatchIpc.compressionCodec(codecSetting.name, codecSetting.zstdLevel) val orderings = boundsOrderings(attrs) batches.map { batch => @@ -397,7 +400,14 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { val (lower, upper, nulls) = gatherColumnStats(batch, attrs, orderings) val numRows = batch.numRows() - val (bytes, columnSizes) = if (Utils.isArrowBacked(batch)) { + // Written as it stands only if its vectors are ones the writer accepts and are already laid + // out the way the schema-less payload will be read back; see CachedBatchIpc's + // matchesReaderLayout. Anything else is converted, which is what makes the fast path safe + // rather than merely usual. + val writeDirectly = + Utils.isArrowBacked(batch) && CachedBatchIpc.matchesReaderLayout(batch, readerFields) + + val (bytes, columnSizes) = if (writeDirectly) { CachedBatchIpc.serialize(batch, codec, CometArrowAllocator) } else { val arrowBatch = @@ -515,8 +525,10 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { input.mapPartitions { it => // Built once per partition: resolving the Arrow schema and the projection's buffer layout - // walks every field of the cached relation, which would otherwise be paid per batch. - val projection = new CachedBatchIpc.Projection( + // walks every field of the cached relation, which would otherwise be paid per batch. Lazy + // because a row-count-only read selects nothing and never decodes, and that walk is the + // whole cost of such a scan over a wide relation. + lazy val projection = new CachedBatchIpc.Projection( Utils .toArrowSchema(cacheSchema, CometArrowStream.NATIVE_TIMEZONE) .getFields diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala index 43a78ab858b..e268e805fc4 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala @@ -40,6 +40,8 @@ import org.apache.spark.SparkException import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.vectorized.ColumnarBatch +import org.apache.comet.vector.CometVector + /** * The on-disk shape of a `CometCachedBatch` payload, and the two operations over it. * @@ -52,9 +54,9 @@ import org.apache.spark.sql.vectorized.ColumnarBatch * * Compression is applied by Arrow per buffer rather than by wrapping the whole payload in a Spark * `CompressionCodec`. That is what makes projection cheap: the message metadata records every - * buffer's offset and length within the body, so [[readProjected]] can copy out only the buffers - * of the columns a scan selected and let `VectorLoader` decompress just those. A whole-payload - * codec would have to inflate everything before any column could be read. + * buffer's offset and length within the body, so [[Projection.load]] can copy out only the + * buffers of the columns a scan selected and let `VectorLoader` decompress just those. A + * whole-payload codec would have to inflate everything before any column could be read. */ private[comet] object CachedBatchIpc { @@ -95,9 +97,89 @@ private[comet] object CachedBatchIpc { .map(t => t -> CommonsCompressionFactory.INSTANCE.createCodec(t)) .toMap - /** The decompressor for a body-compression byte, or None when the batch is stored plain. */ + /** + * The decompressor for a body-compression byte, or None when the batch is stored plain. + * + * A byte this build does not recognize is rejected rather than read as plain bytes. + * `CodecType.fromCompressionType` answers `NO_COMPRESSION` for anything outside its enum, so + * taking its word for it would turn a corrupt payload into garbage values instead of an error. + */ private def readCodec(compressionType: Byte): Option[CompressionCodec] = - readCodecs.get(CompressionUtil.CodecType.fromCompressionType(compressionType)) + if (compressionType == NoCompressionCodec.COMPRESSION_TYPE) { + None + } else { + val codecType = CompressionUtil.CodecType.fromCompressionType(compressionType) + if (codecType == CompressionUtil.CodecType.NO_COMPRESSION) { + throw new SparkException( + s"Comet cached batch records an unknown Arrow compression codec: $compressionType") + } + Some(readCodecs(codecType)) + } + + /** + * Whether `batch`'s vectors can be unloaded as they stand, or have to be converted first. + * + * The payload records no schema, so [[Projection]] rebuilds the fields from the cached + * relation's Spark attributes and reads the body against them. The direct write path unloads + * whatever vectors the cached plan produced, and one Spark type can arrive as more than one + * Arrow type: `BinaryType` is a `VarBinaryVector` from Comet's own scans but a + * `FixedSizeBinaryVector` from an accelerated `mapInArrow` or an Iceberg `fixed[N]` read, and + * those occupy three buffers and two. Writing one and reading the other shifts every buffer + * from that column on, which is wrong values rather than an error, so a batch that does not + * already carry the reader's types is converted instead. + * + * The same holds inside a nested column, which `Utils.isArrowBacked` does not look at: it + * answers for the top-level vector only, so a struct of large strings passes it while its child + * is stored with 64-bit offsets and read with 32-bit ones. + * + * Names, nullability and a timestamp's timezone are not compared. None of them changes how the + * reader interprets the body, and the writer's legitimately differ -- a Comet scan labels + * timestamps with the session's zone where the reader rebuilds them as UTC, which is a label + * only: Spark's representation is micros since the epoch either way. + */ + def matchesReaderLayout(batch: ColumnarBatch, readerFields: Seq[Field]): Boolean = + batch.numCols() == readerFields.length && + (0 until batch.numCols()).forall { i => + batch.column(i) match { + case v: CometVector => sameLayout(writtenField(v), readerFields(i)) + case _ => false + } + } + + /** + * The field a column reaches the body as. + * + * A dictionary-encoded vector's own field carries the index type, not the values', because + * [[decodeDictionaries]] replaces it with the decoded form before anything is unloaded. + * Resolved through the same `lookupDictionary` the write path uses, so a batch missing its + * dictionary fails here exactly as it would there. + */ + private def writtenField(column: CometVector): Field = { + val vector = column.getValueVector + if (vector.getField.getDictionary == null) { + vector.getField + } else { + Utils + .lookupDictionary(vector.asInstanceOf[FieldVector], Option(column.getDictionaryProvider)) + .getVector + .getField + } + } + + private def sameLayout(written: Field, read: Field): Boolean = + layoutType(written.getType) == layoutType(read.getType) && { + val writtenChildren = written.getChildren + val readChildren = read.getChildren + writtenChildren.size == readChildren.size && + (0 until writtenChildren.size).forall(i => + sameLayout(writtenChildren.get(i), readChildren.get(i))) + } + + private def layoutType(t: ArrowType): ArrowType = t match { + case ts: ArrowType.Timestamp if ts.getTimezone != null => + new ArrowType.Timestamp(ts.getUnit, "UTC") + case other => other + } /** * Serialize `batch` into one encapsulated IPC RecordBatch message. @@ -119,7 +201,7 @@ private[comet] object CachedBatchIpc { batch: ColumnarBatch, codec: CompressionCodec, allocator: BufferAllocator): (Array[Byte], Array[Long]) = { - val (vectors, hydrated) = hydrateDictionaries(batch, allocator) + val (vectors, decoded) = decodeDictionaries(batch, allocator) try { val root = new VectorSchemaRoot(vectors.asJava) // A batch of zero columns carries only a row count, which a VectorSchemaRoot cannot infer @@ -128,20 +210,15 @@ private[comet] object CachedBatchIpc { root.setRowCount(batch.numRows()) } - // alignBuffers=true matches the 8-byte buffer alignment readProjected reproduces when it + // alignBuffers=true matches the 8-byte buffer alignment Projection.load reproduces when it // repacks the selected buffers. val unloader = new VectorUnloader(root, true, codec, true) val recordBatch = unloader.getRecordBatch try { val fields = vectors.map(_.getField) - // Serializing consumes the batch, as it does in Utils.serializeBatches. The record batch - // holds its own buffers by now -- compressed copies, or retained references when the codec - // is none -- so releasing the vectors here does not touch it. getField still answers - // afterwards: clearing releases buffers, not the schema. - // - // Not load bearing for memory: the plan that produced the batch releases its vectors - // either way, and dropping this line leaks nothing. It is here because serializeBatches - // does the same, so both writers leave a batch they were handed in the same state. + // Leaves the batch in the state serializeBatches leaves one. The record batch holds its + // own buffers by now, so this does not touch it, and getField still answers afterwards: + // clearing releases buffers, not the schema. root.clear() // Sized up front from the body length the record batch already knows, plus room for the @@ -158,7 +235,7 @@ private[comet] object CachedBatchIpc { } } finally { // Only the vectors this method allocated. The rest belong to the input batch. - hydrated.foreach(v => + decoded.foreach(v => try v.close() catch { case NonFatal(_) => () }) } @@ -166,41 +243,37 @@ private[comet] object CachedBatchIpc { /** * Everything about reading one projection of this format that does not change between batches. + * A scan builds one of these per partition. * - * The index arithmetic here is a pure function of the cached schema and the selected columns, - * both fixed for the life of a scan, but it walks every field of the whole relation rather than - * just the projected ones. Recomputing it per batch would make the bookkeeping O(total columns) - * while the useful work is O(selected columns) -- worst in exactly the wide-relation, - * narrow-projection case this format exists for. A scan builds one of these per partition. + * The index arithmetic walks every field of the cached relation rather than just the projected + * ones, so recomputing it per batch would make the bookkeeping O(total columns) against + * O(selected columns) of useful work -- worst in exactly the wide-relation, narrow-projection + * case this format exists for. * * Holding the projected `Schema` here too is what keeps it consistent with the buffers: - * [[load]] packs field nodes and buffers by walking `selectedIndices` in order, and the schema - * is built from the same walk, so the two cannot drift apart. + * [[Projection.load]] packs field nodes and buffers by walking `selectedIndices` in order, and + * the schema is built from the same walk, so the two cannot drift apart. */ - final class Projection(arrowFields: Seq[Field], selectedIndices: Array[Int]) { + final class Projection(arrowFields: IndexedSeq[Field], selectedIndices: Array[Int]) { private val schema = new Schema(selectedIndices.map(arrowFields).toSeq.asJava) // A record batch body is a flat, depth-first sequence of buffers in schema order, so each - // top-level column owns a contiguous run of it; field nodes and variadic buffer counts run in - // the same order. - private val nodeIndices = selectedRange(arrowFields, selectedIndices, fieldNodeCount) - private val bufferIndices = selectedRange(arrowFields, selectedIndices, fieldBufferCount) - private val variadicIndices = selectedRange(arrowFields, selectedIndices, fieldVariadicCount) + // top-level column owns a contiguous run of it; field nodes run in the same order. The totals + // are what a payload is checked against in load. + private val (nodeIndices, totalNodes) = + selectedRange(arrowFields, selectedIndices, fieldNodeCount) + private val (bufferIndices, totalBuffers) = + selectedRange(arrowFields, selectedIndices, fieldBufferCount) /** * Decode the projected columns of one cached payload into a fresh root the caller owns. * - * Only the selected buffers are ever materialized off-heap or decompressed. The message - * metadata records every buffer's offset and length within the body, so the selected columns' - * bytes are copied into a single allocation -- each 8-byte aligned exactly as Arrow's IPC - * body lays them out -- and the columns that were not selected are never read, let alone - * inflated. - * * A buffer's recorded (offset, length) covers its on-body bytes including the - * uncompressed-length prefix, so a copied window is exactly what the writer emitted. The - * windows are then decompressed in one pass; see [[decompressed]] for why that is not left to - * `VectorLoader`. + * uncompressed-length prefix, so a window copied out of the payload is exactly what the + * writer emitted, 8-byte aligned as Arrow's IPC body lays it out. The columns that were not + * selected are never read, let alone inflated. The windows are then decompressed in one pass; + * see [[decompressed]] for why that is not left to `VectorLoader`. */ def load(data: Array[Byte], allocator: BufferAllocator): VectorSchemaRoot = { val readChannel = new ReadChannel(Channels.newChannel(new ByteArrayInputStream(data))) @@ -211,6 +284,19 @@ private[comet] object CachedBatchIpc { } val batch = metadata.getMessage.header(new FlatBufRecordBatch()).asInstanceOf[FlatBufRecordBatch] + + // The payload carries no schema, so nothing in it says the writer laid the body out the way + // these windows read it. batch.buffers(j) is an unchecked flatbuffer accessor, so a + // disagreement would otherwise surface as wrong values, or as an out-of-range read from + // inside the copy below, rather than as an error naming the cause. See matchesReaderLayout + // for how the write path avoids producing one. + if (batch.nodesLength() != totalNodes || batch.buffersLength() != totalBuffers) { + throw new SparkException( + s"Comet cached batch does not match the cached schema: the payload holds " + + s"${batch.nodesLength()} field nodes and ${batch.buffersLength()} buffers, but the " + + s"schema describes $totalNodes and $totalBuffers") + } + // serialize writes exactly [encapsulated message][body] and nothing after it, so the body is // the tail of `data`. val bodyStart = data.length - metadata.getMessageBodyLength.toInt @@ -224,10 +310,6 @@ private[comet] object CachedBatchIpc { val node = batch.nodes(j) nodes.add(new ArrowFieldNode(node.length(), node.nullCount())) } - val variadicCounts = new java.util.ArrayList[java.lang.Long](variadicIndices.length) - if (batch.variadicBufferCountsLength() > 0) { - variadicIndices.foreach(j => variadicCounts.add(batch.variadicBufferCounts(j))) - } val offsets = new Array[Long](bufferIndices.length) val lengths = new Array[Long](bufferIndices.length) @@ -260,13 +342,7 @@ private[comet] object CachedBatchIpc { position += DataSizeRoundingUtil.roundUpTo8Multiple(length) i += 1 } - new ArrowRecordBatch( - batch.length().toInt, - nodes, - buffers, - compression, - variadicCounts, - false) + new ArrowRecordBatch(batch.length().toInt, nodes, buffers, compression, false) } catch { case NonFatal(e) => body.close() @@ -299,17 +375,18 @@ private[comet] object CachedBatchIpc { /** * The indices, within a record batch's flat depth-first sequence, that the selected columns - * own. + * own, paired with the length of the whole sequence. * * `count` gives how many entries of the sequence a field occupies including its descendants, so - * a running total over every field turns a column index into its run within the sequence. + * a running total over every field turns a column index into its run within the sequence. The + * final total is what [[Projection.load]] checks a payload against. */ private def selectedRange( - arrowFields: Seq[Field], + arrowFields: IndexedSeq[Field], selectedIndices: Array[Int], - count: Field => Int): Array[Int] = { + count: Field => Int): (Array[Int], Int) = { val starts = arrowFields.scanLeft(0)(_ + count(_)).toArray - selectedIndices.flatMap(i => starts(i) until starts(i + 1)) + (selectedIndices.flatMap(i => starts(i) until starts(i + 1)), starts.last) } /** @@ -332,10 +409,8 @@ private[comet] object CachedBatchIpc { private def decompressed( batch: ArrowRecordBatch, allocator: BufferAllocator): ArrowRecordBatch = { - // getCodec is the raw IPC byte; the factory keys off the enum. Both sides of the comparison - // in readCodec have to be CodecType: NoCompressionCodec.COMPRESSION_TYPE is the byte -1, and - // Scala compares a CodecType against it by universal equality, which is quietly always - // unequal. + // getBodyCompression().getCodec() is the raw IPC byte, which readCodec turns into a codec or + // rejects. val codec = readCodec(batch.getBodyCompression.getCodec) val buffers = new java.util.ArrayList[ArrowBuf]() @@ -387,11 +462,19 @@ private[comet] object CachedBatchIpc { private def columnSizes(fields: Seq[Field], recordBatch: ArrowRecordBatch): Array[Long] = { val buffers = recordBatch.getBuffersLayout val starts = fields.scanLeft(0)(_ + fieldBufferCount(_)).toArray - fields.indices.map { i => - (starts(i) until starts(i) + fieldBufferCount(fields(i))) - .map(j => buffers.get(j).getSize) - .sum - }.toArray + val sizes = new Array[Long](fields.length) + var i = 0 + while (i < sizes.length) { + var size = 0L + var j = starts(i) + while (j < starts(i + 1)) { + size += buffers.get(j).getSize + j += 1 + } + sizes(i) = size + i += 1 + } + sizes } /** @@ -401,10 +484,10 @@ private[comet] object CachedBatchIpc { * exactly those. Columns that needed no decoding are returned as they are and stay owned by * `batch`. */ - private def hydrateDictionaries( + private def decodeDictionaries( batch: ColumnarBatch, allocator: BufferAllocator): (Seq[FieldVector], Seq[ValueVector]) = { - val hydrated = mutable.ArrayBuffer.empty[ValueVector] + val decoded = mutable.ArrayBuffer.empty[ValueVector] try { val vectors = Utils.getBatchFieldVectorsWithProviders(batch).map { case (vector, providerOpt) => @@ -412,15 +495,15 @@ private[comet] object CachedBatchIpc { vector } else { val dictionary = Utils.lookupDictionary(vector, providerOpt) - val decoded = DictionaryEncoder.decode(vector, dictionary, allocator) - hydrated += decoded - decoded.asInstanceOf[FieldVector] + val plain = DictionaryEncoder.decode(vector, dictionary, allocator) + decoded += plain + plain.asInstanceOf[FieldVector] } } - (vectors, hydrated.toSeq) + (vectors, decoded.toSeq) } catch { case NonFatal(e) => - hydrated.foreach(v => + decoded.foreach(v => try v.close() catch { case NonFatal(closeError) => e.addSuppressed(closeError) }) throw e @@ -439,18 +522,4 @@ private[comet] object CachedBatchIpc { /** Number of field nodes a field occupies: itself plus every descendant. */ private def fieldNodeCount(field: Field): Int = 1 + field.getChildren.asScala.map(fieldNodeCount).sum - - /** - * Number of variadic buffer counts a field contributes, one per view-type buffer, recursively. - * - * Only Utf8View and BinaryView carry one. Comet's cache never writes view vectors today, but - * the span arithmetic above has to stay correct if that changes. - */ - private def fieldVariadicCount(field: Field): Int = { - val own = field.getType match { - case _: ArrowType.Utf8View | _: ArrowType.BinaryView => 1 - case _ => 0 - } - own + field.getChildren.asScala.map(fieldVariadicCount).sum - } } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 573278092c4..d0e07659a36 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -21,11 +21,12 @@ package org.apache.comet.exec import java.{util => ju} +import org.apache.arrow.vector.{FixedSizeBinaryVector, VarBinaryVector} import org.apache.arrow.vector.types.pojo.ArrowType import org.apache.spark.CometDriverPlugin import org.apache.spark.SparkConf import org.apache.spark.sql.{CometTestBase, Row} -import org.apache.spark.sql.catalyst.expressions.{And, Attribute, Expression, GreaterThanOrEqual, LessThan, Literal} +import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, Expression, GreaterThanOrEqual, LessThan, Literal} import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} import org.apache.spark.sql.comet.execution.arrow.CometCachedBatchHelper @@ -38,11 +39,12 @@ import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, SortMergeJoi import org.apache.spark.sql.functions.max import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} import org.apache.spark.sql.types._ +import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} import org.apache.spark.storage.StorageLevel import org.apache.comet.{CometArrowAllocator, CometConf} import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus} -import org.apache.comet.vector.CometVector +import org.apache.comet.vector.{CometPlainVector, CometVector} class CometInMemoryCacheSuite extends CometTestBase { @@ -1377,15 +1379,21 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - /** Decode `batches` through the cache serializer, selecting `selected`, and total the rows. */ + /** + * Decode `batches` through the cache serializer, selecting `selected`, and total the rows. + * + * `cacheAttributes` defaults to the relation's own, and is overridable so a test can hand the + * reader a schema the writer did not use. + */ private def decodedRowCount( relation: org.apache.spark.sql.execution.columnar.InMemoryRelation, batches: Array[CachedBatch], - selected: Seq[Attribute]): Long = { + selected: Seq[Attribute], + cacheAttributes: Option[Seq[Attribute]] = None): Long = { relation.cacheBuilder.serializer .convertCachedBatchToColumnarBatch( spark.sparkContext.parallelize(batches.toSeq, 1), - relation.output, + cacheAttributes.getOrElse(relation.output), selected, spark.sessionState.conf) // ColumnarBatch is not serializable, so reduce to a count inside the closure. @@ -1405,10 +1413,8 @@ class CometInMemoryCacheSuite extends CometTestBase { */ private def interceptDecodeFailure(f: => Unit): Throwable = { val thrown = intercept[Exception](f) - val chain = - Iterator.iterate(thrown: Throwable)(_.getCause).takeWhile(_ != null).take(20).toSeq assert( - !chain.exists { t => + !causeChain(thrown).exists { t => t.getClass.getName.contains("IllegalReferenceCount") || Option(t.getMessage).exists(m => m.contains("RefCnt") || m.contains("refCnt")) }, @@ -1477,70 +1483,110 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache decodes only the projected columns") { - // Timings would be a weak assertion here, so this scrambles the compressed bytes of the - // columns the read must not touch, leaving every other byte of the payload identical. - // Reading still has to succeed, which it only can if those columns' buffers were never copied - // out of the payload and handed to the decompressor. The second half checks the corruption is - // detectable at all, so the first half cannot pass just because the bad bytes decode silently - // to nothing. - withProjectionCache { (relation, batches) => - val cacheSchema = Utils.fromAttributes(relation.output) - val selectedIdx = 1 - val selected = Seq(relation.output(selectedIdx)) + // Every column of both relations takes a turn as the sole projection. Timings would be a weak + // assertion here, so each turn scrambles the compressed bytes of the columns the read must not + // touch, leaving every other byte of the payload identical. Reading still has to succeed, which + // it only can if those columns' buffers were never copied out of the payload and handed to the + // decompressor. Each turn then corrupts the selected column too, so the assertion cannot pass + // just because the bad bytes decode silently to nothing. + // + // The nested relation is what exercises the span arithmetic. A flat column always owns one field + // node and two or three buffers, whereas a nested one owns a run as long as its whole subtree, + // so a run computed short or long by a buffer shifts every column after it -- and which column + // is selected decides whether that misalignment reaches into a corrupted neighbour. + Seq(("flat", withProjectionCache _), ("nested", withNestedProjectionCache _)).foreach { + case (shape, withCache) => + test(s"Comet in-memory cache decodes only the projected columns of a $shape relation") { + withCache { (relation, batches) => + val cacheSchema = Utils.fromAttributes(relation.output) + val pristine = CometCachedBatchHelper.snapshotPayloads(batches) + + relation.output.indices.foreach { i => + assert( + batches.forall(b => CometCachedBatchHelper.columnIsCompressed(b, cacheSchema, i)), + s"column ${relation.output(i).name} is not stored compressed, so corrupting it " + + "would prove nothing") + } - relation.output.indices.foreach { i => - assert( - batches.forall(b => CometCachedBatchHelper.columnIsCompressed(b, cacheSchema, i)), - s"column $i is not stored compressed, so corrupting it would prove nothing") - } + relation.output.indices.foreach { selectedIdx => + CometCachedBatchHelper.restorePayloads(batches, pristine) + val selected = Seq(relation.output(selectedIdx)) + val name = relation.output(selectedIdx).name - relation.output.indices.filter(_ != selectedIdx).foreach { i => - batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) + relation.output.indices.filter(_ != selectedIdx).foreach { i => + batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) + } + assert( + decodedRowCount(relation, batches, selected) == projectionCacheRows, + s"reading $name must not decompress the other ${relation.output.length - 1} columns") + + batches.foreach(b => + CometCachedBatchHelper.corruptColumn(b, cacheSchema, selectedIdx)) + interceptDecodeFailure { + decodedRowCount(relation, batches, selected) + } + } + } } + } - assert( - decodedRowCount(relation, batches, selected) == projectionCacheRows, - "reading one column must not decompress the other five") - - batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, selectedIdx)) - interceptDecodeFailure { - decodedRowCount(relation, batches, selected) + test("Comet in-memory cache rejects a payload that disagrees with the cached schema") { + // Nothing in the payload says which schema wrote it, and `batch.buffers(j)` is an unchecked + // flatbuffer accessor, so a reader working from a wider schema than the writer used would + // otherwise copy windows from wherever the arithmetic landed: wrong values, or an + // out-of-range read reported from inside the copy rather than as the layout problem it is. + withProjectionCache { (relation, batches) => + val extra = AttributeReference("extra", LongType)() + val thrown = intercept[Exception] { + decodedRowCount( + relation, + batches, + Seq(relation.output.head), + cacheAttributes = Some(relation.output :+ extra)) } + assert( + causeChain(thrown).exists(t => + Option(t.getMessage).exists(_.contains("does not match the cached schema"))), + s"a layout mismatch must be reported as itself: $thrown") } } - test("Comet in-memory cache decodes only the projected columns of a nested relation") { - // The flat case above pins one column and corrupts the rest. Here every column takes its turn, - // because a nested column's run of buffers is as long as its subtree rather than a fixed two or - // three: a run computed short or long shifts every column after it, so which column is selected - // decides whether the misalignment reaches into a corrupted neighbour. - nestedProjectionColumns.indices.foreach { selectedIdx => - withNestedProjectionCache { (relation, batches) => - val cacheSchema = Utils.fromAttributes(relation.output) - val selected = Seq(relation.output(selectedIdx)) - val name = relation.output(selectedIdx).name - - relation.output.indices.foreach { i => - assert( - batches.forall(b => CometCachedBatchHelper.columnIsCompressed(b, cacheSchema, i)), - s"column ${relation.output(i).name} is not stored compressed, so corrupting it " + - "would prove nothing") - } - - relation.output.indices.filter(_ != selectedIdx).foreach { i => - batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, i)) - } - - assert( - decodedRowCount(relation, batches, selected) == projectionCacheRows, - s"reading $name must not decompress the other ${relation.output.length - 1} columns") + test("Comet in-memory cache converts a batch whose vectors do not match the cached layout") { + // BinaryType is an Arrow Binary to the reader -- validity, offsets, data -- but a CometVector + // may wrap a FixedSizeBinaryVector for the same Spark type, which has no offsets buffer. An + // accelerated mapInArrow returning pa.binary(n) and an Iceberg fixed[N] read both produce one. + // Since the payload stores no schema, writing that and reading a Binary shifts every buffer + // from that column on, so the write path has to notice and convert instead. isArrowBacked + // cannot: it accepts both vectors, as the first assertion of each case records. + val cacheSchema = StructType(Seq(StructField("b", BinaryType))) + val rows = 4 + + val fixed = new FixedSizeBinaryVector("b", CometArrowAllocator, 3) + try { + fixed.allocateNew(rows) + (0 until rows).foreach(i => fixed.set(i, Array[Byte](i.toByte, 1, 2))) + fixed.setValueCount(rows) + val batch = new ColumnarBatch(Array[ColumnVector](new CometPlainVector(fixed)), rows) + assert(Utils.isArrowBacked(batch)) + assert( + !CometCachedBatchHelper.writesDirectly(batch, cacheSchema), + "a fixed-size binary vector does not have the layout the reader rebuilds for BinaryType") + } finally { + fixed.close() + } - batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, selectedIdx)) - interceptDecodeFailure { - decodedRowCount(relation, batches, selected) - } - } + val varBinary = new VarBinaryVector("b", CometArrowAllocator) + try { + varBinary.allocateNew(rows) + (0 until rows).foreach(i => varBinary.set(i, Array[Byte](i.toByte, 1, 2))) + varBinary.setValueCount(rows) + val batch = new ColumnarBatch(Array[ColumnVector](new CometPlainVector(varBinary)), rows) + assert(Utils.isArrowBacked(batch)) + assert( + CometCachedBatchHelper.writesDirectly(batch, cacheSchema), + "the vector Comet's own scans produce for BinaryType must still take the direct path") + } finally { + varBinary.close() } } @@ -1789,19 +1835,44 @@ class CometInMemoryCacheSuite extends CometTestBase { // else -- the holder is published to the task-completion listener only once its constructor // returns -- so a failure that does not release them leaks off-heap for the life of the // executor. + // + // Two corruption points, because they fail at different depths. Taking out a column's first + // buffer fails before anything of that column has been decompressed. Taking out only the last + // buffer of a string column -- whose offsets and data are separately compressed -- decompresses + // one buffer into a fresh allocation and then throws on the next, leaving that allocation + // reachable from nothing the failure path can see. The second is the one that catches a leak + // in `VectorLoader`; the first is the one that catches a cleanup path releasing the shared + // body twice. withProjectionCache { (relation, batches) => val cacheSchema = Utils.fromAttributes(relation.output) - // Corrupt the second selected column, so the first is copied out successfully first. - val selected = Seq(relation.output(0), relation.output(1)) - batches.foreach(b => CometCachedBatchHelper.corruptColumn(b, cacheSchema, 1)) + val pristine = CometCachedBatchHelper.snapshotPayloads(batches) + val stringIdx = 3 + assert(relation.output(stringIdx).dataType.typeName == "string") - val before = CometArrowAllocator.getAllocatedMemory - interceptDecodeFailure { - decodedRowCount(relation, batches, selected) + val cases = Seq( + ( + "a column's first buffer", + // Corrupt the second selected column, so the first is copied out successfully first. + (b: CachedBatch) => CometCachedBatchHelper.corruptColumn(b, cacheSchema, 1), + Seq(relation.output(0), relation.output(1))), + ( + "a string column's trailing buffer", + (b: CachedBatch) => + CometCachedBatchHelper.corruptTrailingBuffer(b, cacheSchema, stringIdx), + Seq(relation.output(stringIdx)))) + + cases.foreach { case (where, corrupt, selected) => + CometCachedBatchHelper.restorePayloads(batches, pristine) + batches.foreach(corrupt) + + val before = CometArrowAllocator.getAllocatedMemory + interceptDecodeFailure { + decodedRowCount(relation, batches, selected) + } + assert( + CometArrowAllocator.getAllocatedMemory == before, + s"everything allocated before a failure in $where must be released") } - assert( - CometArrowAllocator.getAllocatedMemory == before, - "everything allocated before the failure must be released") } } @@ -1845,28 +1916,6 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache releases its vectors when a column fails after a partial decode") { - // Tighter than the two cases above, and the one that actually catches a leak. A string column - // stores its offsets and its data as separate compressed buffers, so corrupting only the - // second makes the decoder decompress one buffer of the column into a fresh allocation and - // then throw on the next, with the first reachable from nothing the failure path can see. - withProjectionCache { (relation, batches) => - val cacheSchema = Utils.fromAttributes(relation.output) - val stringIdx = 3 - assert(relation.output(stringIdx).dataType.typeName == "string") - batches.foreach(b => - CometCachedBatchHelper.corruptTrailingBuffer(b, cacheSchema, stringIdx)) - - val before = CometArrowAllocator.getAllocatedMemory - interceptDecodeFailure { - decodedRowCount(relation, batches, Seq(relation.output(stringIdx))) - } - assert( - CometArrowAllocator.getAllocatedMemory == before, - "a buffer decoded before the failure must be released") - } - } - test("Comet in-memory cache decodes dictionary-encoded columns before storing them") { // The payload carries no schema, so it has nowhere to record that a column is dictionary // encoded, nor the dictionary itself. The reader rebuilds a plain Utf8 field for a string diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala index f8f9430e268..ad2ecd8cbf7 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala @@ -32,6 +32,7 @@ import org.apache.arrow.vector.types.pojo.Field import org.apache.spark.sql.columnar.CachedBatch import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.types.StructType +import org.apache.spark.sql.vectorized.ColumnarBatch /** * Test-only access to the internals of `CometCachedBatch`. @@ -50,6 +51,32 @@ object CometCachedBatchHelper { private def payload(batch: CachedBatch): Array[Byte] = batch.asInstanceOf[CometCachedBatch].bytes + /** + * A copy of every batch's payload, so a test can undo what [[corruptColumn]] scrambled. + * + * Materializing the cache is the expensive part of a corruption test, and a test that gives + * each column of a relation a turn needs the payload back as it was between turns. Restoring + * beats re-caching by the column count. + */ + def snapshotPayloads(batches: Array[CachedBatch]): Array[Array[Byte]] = + batches.map(payload(_).clone()) + + /** Put back what [[snapshotPayloads]] captured. Corruption never changes a payload's length. */ + def restorePayloads(batches: Array[CachedBatch], snapshot: Array[Array[Byte]]): Unit = + batches.zip(snapshot).foreach { case (batch, bytes) => + System.arraycopy(bytes, 0, payload(batch), 0, bytes.length) + } + + /** + * Whether the write path would unload `batch`'s vectors as they stand, or convert them first. + * + * A thin shim rather than a re-derivation: this is the decision under test, not arithmetic to + * check it against. + */ + def writesDirectly(batch: ColumnarBatch, cacheSchema: StructType): Boolean = + Utils + .isArrowBacked(batch) && CachedBatchIpc.matchesReaderLayout(batch, arrowFields(cacheSchema)) + /** * Whether the payload begins with a Schema message rather than going straight to the record * batch. From b978e54474a87d16ed6070198512b020f845affe Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 15 Sep 2026 16:36:08 -0600 Subject: [PATCH 09/24] review: own the write-side compression buffers, fix the activation example Compressing through VectorUnloader leaks on the failure path: appendNodes retains each input buffer and accumulates the compressed ones into a list local to getRecordBatch, so a buffer that fails to compress strands that retain and leaves every buffer compressed before it reachable from nothing. Closing the input batch afterwards undoes neither. Unload plain and compress in CachedBatchIpc.compressed instead, mirroring what decompressed already does on the read side, so every allocation stays reachable from an error path that owns it. The docs enabled the cache with spark.conf.set, which cannot work: the driver plugin picks spark.sql.cache.serializer while the SparkContext is initializing. Show it as a startup --conf. Also drops a redundant s interpolator that the scalafix lint rejected. --- .../user-guide/latest/in-memory-cache.md | 16 ++-- .../execution/arrow/CachedBatchIpc.scala | 71 ++++++++++++++-- .../comet/exec/CometInMemoryCacheSuite.scala | 82 ++++++++++++++++++- .../arrow/CometCachedBatchHelper.scala | 15 ++++ 4 files changed, 173 insertions(+), 11 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index e0e2755c74d..d08f771d16e 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -24,16 +24,22 @@ format that Comet operators read directly. Without it, a cached table is stored format and every scan of it has to convert each batch before Comet can continue, which shows up in the plan as a `CometSparkColumnarToColumnar` above the cache scan. -This feature is **experimental and disabled by default**. +This feature is **experimental and disabled by default**. Turn it on at startup, alongside the rest +of Comet's configuration: -```scala -spark.conf.set("spark.comet.exec.inMemoryCache.enabled", "true") +```shell +$SPARK_HOME/bin/spark-shell \ + ... \ + --conf spark.comet.exec.inMemoryCache.enabled=true ``` +It has to be set before the `SparkContext` starts. Comet's driver plugin chooses +`spark.sql.cache.serializer` once, while the context is initializing, so a session that started +with the default goes on using Spark's cache format however the config is set afterwards. + ## What changes when it is enabled -`spark.comet.exec.inMemoryCache.enabled` is read at startup, and its value then decides whether -Comet installs its cache serializer as `spark.sql.cache.serializer`. When it is installed: +With Comet's serializer installed as `spark.sql.cache.serializer`: - Cached data is stored as `CometCachedBatch` rather than Spark's `DefaultCachedBatch`. - Cached tables are scanned by `CometInMemoryTableScan`, which feeds Comet operators directly. diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala index e268e805fc4..ee5993560ba 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala @@ -210,10 +210,13 @@ private[comet] object CachedBatchIpc { root.setRowCount(batch.numRows()) } - // alignBuffers=true matches the 8-byte buffer alignment Projection.load reproduces when it - // repacks the selected buffers. - val unloader = new VectorUnloader(root, true, codec, true) - val recordBatch = unloader.getRecordBatch + // Unloaded plain and compressed afterwards rather than by handing the codec to the unloader; + // see compressed for why. + val unloader = new VectorUnloader(root, true, NoCompressionCodec.INSTANCE, true) + val plainBatch = unloader.getRecordBatch + val recordBatch = + try compressed(plainBatch, codec, allocator) + finally plainBatch.close() try { val fields = vectors.map(_.getField) // Leaves the batch in the state serializeBatches leaves one. The record batch holds its @@ -292,7 +295,7 @@ private[comet] object CachedBatchIpc { // for how the write path avoids producing one. if (batch.nodesLength() != totalNodes || batch.buffersLength() != totalBuffers) { throw new SparkException( - s"Comet cached batch does not match the cached schema: the payload holds " + + "Comet cached batch does not match the cached schema: the payload holds " + s"${batch.nodesLength()} field nodes and ${batch.buffersLength()} buffers, but the " + s"schema describes $totalNodes and $totalBuffers") } @@ -389,6 +392,64 @@ private[comet] object CachedBatchIpc { (selectedIndices.flatMap(i => starts(i) until starts(i + 1)), starts.last) } + /** + * The same record batch with every buffer compressed, as a new batch the caller owns. + * + * `VectorUnloader` would do this itself if handed the codec, but it leaks on the failure path: + * `appendNodes` retains each input buffer and accumulates the compressed ones into a list local + * to `getRecordBatch`, so a buffer that fails to compress -- zstd unable to allocate its + * workspace, say -- strands that retain and leaves every buffer compressed before it reachable + * from nothing. Closing the input batch afterwards undoes neither, so one failed cache + * materialization leaks a batch's worth of off-heap for the life of the executor. Compressing + * here keeps every allocation reachable from this method's own error path, as [[decompressed]] + * does on the read side. + * + * The retain before each `compress` is where the reference on the buffer that comes back is + * from. A codec that allocates consumes it and hands back a buffer of its own; + * `NoCompressionCodec` hands back the input itself, and the retain is then the reference + * `result` ends up owning. Releasing it again is what a throw owes. + */ + private def compressed( + batch: ArrowRecordBatch, + codec: CompressionCodec, + allocator: BufferAllocator): ArrowRecordBatch = { + val buffers = new java.util.ArrayList[ArrowBuf](batch.getBuffers.size) + try { + batch.getBuffers.asScala.foreach { buffer => + buffer.getReferenceManager.retain() + val packed = + try codec.compress(allocator, buffer) + catch { + case NonFatal(e) => + buffer.getReferenceManager.release() + throw e + } + buffers.add(packed) + } + + val result = new ArrowRecordBatch( + batch.getLength, + batch.getNodes, + buffers, + CompressionUtil.createBodyCompression(codec), + batch.getVariadicBufferCounts, + // alignBuffers=true matches the 8-byte buffer alignment Projection.load reproduces when it + // repacks the selected buffers. This is the layout that gets written, so the unloader's is + // not the one that matters. + true) + // The constructor retained each buffer, so drop the references held here. + buffers.asScala.foreach(_.close()) + result + } catch { + case NonFatal(e) => + buffers.asScala.foreach { buffer => + try buffer.close() + catch { case NonFatal(closeError) => e.addSuppressed(closeError) } + } + throw e + } + } + /** * The same record batch with every buffer decompressed, as a new batch the caller owns. * diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index d0e07659a36..5592e71fd5f 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -21,7 +21,10 @@ package org.apache.comet.exec import java.{util => ju} -import org.apache.arrow.vector.{FixedSizeBinaryVector, VarBinaryVector} +import org.apache.arrow.compression.ZstdCompressionCodec +import org.apache.arrow.memory.{ArrowBuf, BufferAllocator} +import org.apache.arrow.vector.{FixedSizeBinaryVector, IntVector, VarBinaryVector, VarCharVector} +import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil} import org.apache.arrow.vector.types.pojo.ArrowType import org.apache.spark.CometDriverPlugin import org.apache.spark.SparkConf @@ -1876,6 +1879,83 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + /** + * A zstd codec that compresses the first `succeedFor` buffers and then throws. + * + * Delegating to the real codec until it fails is the point: the buffers already compressed are + * genuine off-heap allocations reachable only from inside the writer, which is what a real + * failure -- zstd unable to allocate its workspace part way through a batch -- leaves behind. + */ + private class FailAfterCompressionCodec(succeedFor: Int) extends CompressionCodec { + private val delegate = new ZstdCompressionCodec(1) + var compressed: Int = 0 + + override def compress(allocator: BufferAllocator, buffer: ArrowBuf): ArrowBuf = { + if (compressed == succeedFor) { + throw new RuntimeException(FailAfterCompressionCodec.Message) + } + compressed += 1 + delegate.compress(allocator, buffer) + } + + override def decompress(allocator: BufferAllocator, buffer: ArrowBuf): ArrowBuf = + delegate.decompress(allocator, buffer) + + override def getCodecType: CompressionUtil.CodecType = delegate.getCodecType + } + + private object FailAfterCompressionCodec { + val Message: String = "injected compression failure" + } + + test("Comet in-memory cache releases its buffers when a column fails to compress") { + // Writing a batch allocates a buffer per compressed buffer before any payload exists, and + // nothing outside the writer can reach them while it is still assembling the record batch they + // belong to. This is why the codec is not handed to `VectorUnloader`: it accumulates them in a + // list local to `getRecordBatch`, which is off the stack by the time a caller sees the failure, + // so a single failed materialization would leak a batch's worth of off-heap for the life of the + // executor. + val rows = 256 + val ints = new IntVector("i", CometArrowAllocator) + val strings = new VarCharVector("s", CometArrowAllocator) + try { + ints.allocateNew(rows) + (0 until rows).foreach(i => ints.set(i, i)) + ints.setValueCount(rows) + + strings.allocateNew(rows) + (0 until rows).foreach(i => strings.setSafe(i, s"value_$i".getBytes("UTF-8"))) + strings.setValueCount(rows) + + val batch = new ColumnarBatch( + Array[ColumnVector](new CometPlainVector(ints), new CometPlainVector(strings)), + rows) + + // An int vector is validity and data, a varchar validity, offsets and data: five buffers in + // all. Succeeding for two puts the failure at the string column's first buffer, with the int + // column's two already compressed into allocations only the writer can reach. Failing at the + // very first buffer would pass with no cleanup at all. + val codec = new FailAfterCompressionCodec(succeedFor = 2) + val before = CometArrowAllocator.getAllocatedMemory + val thrown = intercept[Exception] { + CometCachedBatchHelper.serialize(batch, codec, CometArrowAllocator) + } + + assert( + causeChain(thrown).exists(t => + Option(t.getMessage).contains(FailAfterCompressionCodec.Message)), + s"a compression failure must surface as itself: $thrown") + assert(codec.compressed == 2, "the failure must come after some buffers were compressed") + assert( + CometArrowAllocator.getAllocatedMemory == before, + "everything allocated before a write failure must be released") + } finally { + // Never reaches the writer's own clear(), which only runs once the payload is built. + ints.close() + strings.close() + } + } + /** * Cache two low-cardinality string columns and hand the test the cached relation. * diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala index ad2ecd8cbf7..5e09aa78948 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala @@ -25,7 +25,9 @@ import java.nio.channels.Channels import scala.jdk.CollectionConverters._ import org.apache.arrow.flatbuf.{MessageHeader, RecordBatch => FlatBufRecordBatch} +import org.apache.arrow.memory.BufferAllocator import org.apache.arrow.vector.TypeLayout +import org.apache.arrow.vector.compression.CompressionCodec import org.apache.arrow.vector.ipc.ReadChannel import org.apache.arrow.vector.ipc.message.{MessageMetadataResult, MessageSerializer} import org.apache.arrow.vector.types.pojo.Field @@ -77,6 +79,19 @@ object CometCachedBatchHelper { Utils .isArrowBacked(batch) && CachedBatchIpc.matchesReaderLayout(batch, arrowFields(cacheSchema)) + /** + * Serialize one batch the way the cache writer does, with a caller-chosen codec. + * + * Another thin shim, for tests about the write path itself rather than about what a cached + * relation reads back. The codec is the parameter that matters: handing in one that throws is + * how a compression failure part way through a batch is reached deterministically. + */ + def serialize( + batch: ColumnarBatch, + codec: CompressionCodec, + allocator: BufferAllocator): Array[Byte] = + CachedBatchIpc.serialize(batch, codec, allocator)._1 + /** * Whether the payload begins with a Schema message rather than going straight to the record * batch. From e483d083aecc32f17ae6da9f1f824f7c23280292 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Mon, 21 Sep 2026 13:17:17 -0600 Subject: [PATCH 10/24] fix: write cached batches to the schema width, not the batch width ArrowWriter.writeColumns drove its loop from the input ColumnarBatch's width while indexing the writer's fields, which come from the schema the batch is written under. That assumed every producer hands over a batch exactly as wide as the schema. Iceberg's vectorized reader does not. BatchDeleteFilter.filterBatch reads with the delete filter's requiredSchema, which carries _pos after the projected columns when a data file has position deletes, and trims the extras back only when the file also has equality deletes. A merge-on-read UPDATE writes position deletes and no equality deletes, so the extra column survives into the batch, and caching such a relation failed with ArrayIndexOutOfBoundsException inside the write loop. Drive the loop from the writer's fields instead, which writes exactly the columns the schema describes: the extras are trailing, the same prefix Iceberg keeps when it does trim. A batch narrower than the schema is a genuine contract violation and is now refused with a message naming both widths. Closes #6087. --- .../comet/execution/arrow/ArrowWriters.scala | 14 ++++- .../arrow/CometArrowStreamSuite.scala | 56 +++++++++++++++++++ 2 files changed, 69 insertions(+), 1 deletion(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala index e2632f563e3..bb883f8d2ab 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala @@ -165,9 +165,21 @@ class ArrowWriter(val root: VectorSchemaRoot, fields: Array[ArrowFieldWriter]) { count = input.numElements() } + // Driven by the writer's fields rather than by the input's width, because a producer may hand + // over a batch wider than the schema it is written under. Iceberg's vectorized reader does: + // it reads with the schema its delete filter required, which carries `_pos` after the projected + // columns when a data file has position deletes, and trims the extras back only when the file + // also has equality deletes. Those extras are trailing -- `removeExtraColumns` keeps the leading + // `expectedSchema` prefix when it does trim -- so writing the first `fields.length` columns + // writes exactly the columns the schema describes. A batch with fewer columns than the schema + // has no such reading and is refused rather than written short. def writeColumns(input: ColumnarBatch, startRow: Int, numRows: Int): Unit = { + require( + input.numCols() >= fields.length, + s"Cannot write ${fields.length} columns from a batch of ${input.numCols()} " + + (if (input.numCols() == 1) "column" else "columns")) var columnIndex = 0 - while (columnIndex < input.numCols()) { + while (columnIndex < fields.length) { fields(columnIndex).writeColumnSlice(input.column(columnIndex), startRow, numRows) columnIndex += 1 } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala index cfdefedd6e6..3ad59e96ef2 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala @@ -872,4 +872,60 @@ class CometArrowStreamSuite extends AnyFunSuite with Matchers { allocator.close() } } + + test("columnar writes are driven by the schema, not by the input batch width") { + val allocator = new RootAllocator(Long.MaxValue) + val numRows = 3 + val schema = StructType(Seq(StructField("name", StringType))) + val arrowSchema = Utils.toArrowSchema(schema, "UTC") + // A connector may hand over a batch wider than the schema it is read under. Iceberg's + // vectorized reader is the case that found this: it reads with the schema its delete filter + // required, which carries `_pos` after the projected columns when a data file has position + // deletes, and only trims the extras back when the file also has equality deletes. The + // trailing columns are the extras, so the schema's fields line up with the leading ones. + val names = new OnHeapColumnVector(numRows, StringType) + val positions = new OnHeapColumnVector(numRows, LongType) + val input = new ColumnarBatch(Array[ColumnVector](names, positions), numRows) + try { + (0 until numRows).foreach { i => + names.putByteArray(i, s"n$i".getBytes(StandardCharsets.UTF_8)) + positions.putLong(i, i.toLong) + } + val batch = CometArrowConverters.columnarBatchToArrowBatch(input, arrowSchema, allocator) + try { + batch.numCols() shouldBe 1 + batch.numRows() shouldBe numRows + (0 until numRows).foreach { i => + batch.column(0).getUTF8String(i).toString shouldBe s"n$i" + } + } finally batch.close() + } finally { + input.close() + allocator.close() + } + } + + test("a batch narrower than the schema is refused rather than written short") { + val allocator = new RootAllocator(Long.MaxValue) + val numRows = 2 + val schema = + StructType(Seq(StructField("name", StringType), StructField("id", LongType))) + val arrowSchema = Utils.toArrowSchema(schema, "UTC") + val names = new OnHeapColumnVector(numRows, StringType) + val input = new ColumnarBatch(Array[ColumnVector](names), numRows) + try { + (0 until numRows).foreach { i => + names.putByteArray(i, s"n$i".getBytes(StandardCharsets.UTF_8)) + } + val failure = intercept[IllegalArgumentException] { + CometArrowConverters.columnarBatchToArrowBatch(input, arrowSchema, allocator) + } + failure.getMessage should include("1 column") + failure.getMessage should include("2") + allocator.getAllocatedMemory shouldBe 0L + } finally { + input.close() + allocator.close() + } + } } From a371d646d0a43013a6382e34b7e5279fbb0acd34 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 29 Sep 2026 11:02:23 -0600 Subject: [PATCH 11/24] fix: report a cached relation's decoded size to the planner, not its compressed size Comet's cache format reported each batch's sizeInBytes as its stored, compressed payload. Spark's planner reads the sum of those as the size of a materialized cached relation, for the broadcast threshold and the shuffled hash join build side among others, and both of Spark's own cache formats report decoded sizes there. With the default zstd codec a cached relation could look several times smaller than in Spark's format and be broadcast where Spark would shuffle it. Record each column's decoded Arrow size, measured before compression as Spark's ArrowCachedBatchSerializer does, and let CometCachedBatch inherit sizeInBytes from SimpleMetricsCachedBatch as the sum of those. --- .../user-guide/latest/in-memory-cache.md | 3 + .../arrow/ArrowCachedBatchSerializer.scala | 19 +++-- .../execution/arrow/CachedBatchIpc.scala | 30 +++---- .../comet/exec/CometInMemoryCacheSuite.scala | 83 ++++++++++++++----- .../CometInMemoryCacheBenchmark.scala | 10 ++- .../arrow/CometCachedBatchHelper.scala | 17 +++- 6 files changed, 115 insertions(+), 47 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 8e39f00fe5e..e208908a998 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -45,6 +45,9 @@ With Comet's serializer installed as `spark.sql.cache.serializer`: - Cached tables are scanned by `CometInMemoryTableScan`, which feeds Comet operators directly. - Per-batch column statistics are recorded in the layout Spark's `SimpleMetricsCachedBatchSerializer` expects, so Spark can prune whole cached batches on a predicate before any of them is decoded. +- The size Spark's planner sees for a cached relation is its decoded Arrow size, not the compressed + size it occupies in memory, as with Spark's own cache formats. Compression therefore does not + change how queries over a cached relation are planned, such as whether a join broadcasts it. Relations whose schema Comet's Arrow writer cannot store — interval types, most notably — are delegated in full to Spark's default cache format, per relation. Which format a relation uses does diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 47691556848..d43b09b2628 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -52,10 +52,14 @@ import org.apache.comet.vector.NativeUtil * and length, so `CachedBatchIpc.Projection.load` copies out just the selected columns' byte * ranges. The cache manager still owns storage and eviction; this class only changes the cached * payload. + * + * `sizeInBytes` is not the payload's size. It is inherited from `SimpleMetricsCachedBatch`, which + * sums the per-column sizes in `stats`, and those are decoded sizes. Spark's planner reads it as + * the size of a materialized cached relation, as it does for Spark's own cache formats, which + * also report decoded sizes. See `statsRow`. */ private case class CometCachedBatch( override val numRows: Int, - override val sizeInBytes: Long, override val stats: InternalRow, bytes: ChunkedByteBuffer) extends SimpleMetricsCachedBatch @@ -350,11 +354,13 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { values(base + 1) = upper(c) values(base + 2) = nulls(c) values(base + 3) = numRows - // The stored size of the column's own Arrow buffers, taken from the message's buffer - // layout, so it is exact rather than an estimate. Cache pruning uses - // bounds/null-count/row-count rather than this field, but Spark reserves it and reports it, - // so record the real value. The per-batch message framing is not attributed to any column, - // so these sum to slightly less than sizeInBytes. + // The column's decoded size: the plain length of its own Arrow buffers before compression, + // which is what Spark's own Arrow cache format records here. SimpleMetricsCachedBatch sums + // these into sizeInBytes, which Spark's planner reads as the size of a materialized cached + // relation, for the broadcast threshold and the shuffled hash join build side among others. + // The compressed payload can be several times smaller, and a relation reported at that size + // would be planned differently from the same relation in Spark's cache format, for example + // broadcast where Spark's cache would have it shuffled. values(base + 4) = columnSizes(c) c += 1 } @@ -423,7 +429,6 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { CometCachedBatch( numRows = numRows, - sizeInBytes = bytes.size, stats = statsRow(lower, upper, nulls, numRows, columnSizes), bytes = bytes) } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala index d92971c2bfd..e0a4ee7a2cd 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchIpc.scala @@ -186,9 +186,10 @@ private[comet] object CachedBatchIpc { * batch arrives at whatever size the plan above produced. Chunks are appended rather than grown * and recopied, so the write also never holds the payload twice. * - * Returns the message and the on-body compressed size of each top-level column, which the - * caller records in the statistics row. The sizes come from the message's own buffer layout, so - * they are the real stored sizes rather than an estimate. + * Returns the message and the decoded size of each top-level column, which the caller records + * in the statistics row. Each size is measured on the batch before compression, from the plain + * lengths of the column's own buffers, which is what `getBufferSize` reports for a vector and + * what Spark's own Arrow cache format records. * * Dictionary-encoded columns are decoded to their plain form first. A payload with no Schema * message cannot describe a dictionary encoding, and the schema the reader rebuilds from Spark @@ -215,16 +216,18 @@ private[comet] object CachedBatchIpc { // Unloaded plain and compressed afterwards rather than by handing the codec to the unloader; // see compressed for why. + val fields = vectors.map(_.getField) val unloader = new VectorUnloader(root, true, NoCompressionCodec.INSTANCE, true) val plainBatch = unloader.getRecordBatch - val recordBatch = - try compressed(plainBatch, codec, allocator) - finally plainBatch.close() + val (sizes, recordBatch) = + try { + (columnSizes(fields, plainBatch), compressed(plainBatch, codec, allocator)) + } finally { + plainBatch.close() + } try { - val fields = vectors.map(_.getField) // Leaves the batch in the state serializeBatches leaves one. The record batch holds its - // own buffers by now, so this does not touch it, and getField still answers afterwards: - // clearing releases buffers, not the schema. + // own buffers by now, so this does not touch it. root.clear() val out = new ChunkedByteBufferOutputStream(chunkSize, ByteBuffer.allocate) @@ -234,7 +237,7 @@ private[comet] object CachedBatchIpc { } finally { out.close() } - (out.toChunkedByteBuffer, columnSizes(fields, recordBatch)) + (out.toChunkedByteBuffer, sizes) } finally { recordBatch.close() } @@ -608,11 +611,10 @@ private[comet] object CachedBatchIpc { } /** - * The on-body compressed size of each top-level column. + * The size of each top-level column's buffers in `recordBatch`. * - * Each column owns the run of buffers its subtree occupies, so its stored size is the sum of - * those buffers' recorded lengths. With one payload per batch these are the only per-column - * sizes available -- there is no separate stream to measure -- and they are exact. + * Each column owns the run of buffers its subtree occupies, so its size is the sum of those + * buffers' recorded lengths. */ private def columnSizes(fields: Seq[Field], recordBatch: ArrowRecordBatch): Array[Long] = { val buffers = recordBatch.getBuffersLayout diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 5933f7fa9c2..a554dbcda96 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -22,6 +22,7 @@ package org.apache.comet.exec import java.{util => ju} import java.nio.charset.StandardCharsets +import scala.collection.mutable import scala.jdk.CollectionConverters._ import org.apache.arrow.compression.ZstdCompressionCodec @@ -191,10 +192,11 @@ class CometInMemoryCacheSuite extends CometTestBase { // https://github.com/apache/spark/blob/v4.1.2/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala#L2780-L2832 test("AQE SPARK-37742: use valid Comet cache statistics for join selection") { withAQECache { - // Comet reports compressed Arrow bytes, so use a threshold below the compressed - // large cache as well as below its logical estimate. The single-row side still fits. + // Spark's own threshold. The large cache has to stay above it once materialized: its 60k + // 20-byte keys are about 1.4 MB decoded but compress to a small fraction of that, so a cache + // that reported its compressed size would be broadcast by the third join. withSQLConf( - SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "1024", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "1048584", SQLConf.ADAPTIVE_OPTIMIZER_EXCLUDED_RULES.key -> "org.apache.spark.sql.execution.adaptive.AQEPropagateEmptyRelation") { withTempView("cache_large", "cache_other", "cache_small") { @@ -235,7 +237,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val stats = relation.computeStats() assert(stats.rowCount.contains(BigInt(60000))) assert(stats.sizeInBytes == batches.map(_.sizeInBytes).sum) - assert(stats.sizeInBytes > 1024L) + assert(stats.sizeInBytes > 1048584L) } } } @@ -1968,9 +1970,6 @@ class CometInMemoryCacheSuite extends CometTestBase { assert( CometCachedBatchHelper.chunkCount(batch) > 1, "a payload larger than the chunk size must be stored in more than one chunk") - assert( - CometCachedBatchHelper.payloadSize(batch) == batch.sizeInBytes, - "sizeInBytes must report the whole payload across its chunks") } projections.zip(expected).foreach { case (cols, rows) => @@ -1998,25 +1997,32 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - test("Comet in-memory cache records per-column sizes in its statistics") { - // SimpleMetricsCachedBatch reserves a fifth field per column for its size. A column owns a - // known run of buffers in the payload, so the real stored size is known and must be reported - // rather than left at zero. Run over the nested relation as well: a nested column's size is - // the sum of its whole subtree, so this is also where a size attributed to the wrong column - // surfaces. + test("Comet in-memory cache records per-column decoded sizes in its statistics") { + // SimpleMetricsCachedBatch reserves a fifth field per column for its size and sums those into + // the batch's sizeInBytes, which is what Spark's planner reads as the size of a materialized + // cached relation. Spark's own formats record a column's decoded size there, so each field is + // compared with its column decoded back out of the payload, not with what the column occupies + // compressed. Run over the nested relation as well: a nested column's size is the sum of its + // whole subtree, so this is also where a size attributed to the wrong column surfaces. def checkSizes( relation: org.apache.spark.sql.execution.columnar.InMemoryRelation, batches: Array[CachedBatch]): Unit = { val cacheSchema = Utils.fromAttributes(relation.output) - batches.foreach { batch => - val sizes = CometCachedBatchHelper.columnSizes(batch, cacheSchema) - val stats = batch.asInstanceOf[SimpleMetricsCachedBatch].stats - sizes.zipWithIndex.foreach { case (size, i) => - assert( - stats.getLong(i * 5 + 4) == size, - s"column ${relation.output(i).name} should report the stored size of its own " + - "buffers in the statistics row") + val allocator = CometArrowAllocator.newChildAllocator("decoded-sizes", 0, Long.MaxValue) + try { + batches.foreach { batch => + val sizes = CometCachedBatchHelper.decodedColumnSizes(batch, cacheSchema, allocator) + val stats = batch.asInstanceOf[SimpleMetricsCachedBatch].stats + sizes.zipWithIndex.foreach { case (size, i) => + assert( + stats.getLong(i * 5 + 4) == size, + s"column ${relation.output(i).name} should report its decoded size in the " + + "statistics row") + } + assert(batch.sizeInBytes == sizes.sum) } + } finally { + allocator.close() } } @@ -2024,6 +2030,41 @@ class CometInMemoryCacheSuite extends CometTestBase { withNestedProjectionCache(checkSizes _) } + test("Comet in-memory cache reports the same relation size under every codec") { + // Broadcast thresholds and the shuffled hash join build side are compared against this size. + // Reported compressed, a relation that zstd shrinks several times over would be broadcast + // where the same relation in Spark's cache format is shuffled, so it must not depend on the + // codec. + val sizes = mutable.ArrayBuffer.empty[(String, Long, Long)] + Seq("zstd", "none").foreach { codec => + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_COMPRESSION_CODEC.key -> codec) { + spark.catalog.clearCache() + spark + .range(0, 20000, 1, 2) + .selectExpr("id", "id % 7 AS k", "concat('v_', cast(id % 100 AS string)) AS s") + .createOrReplaceTempView("codec_size_cache") + spark.catalog.cacheTable("codec_size_cache") + assert(spark.table("codec_size_cache").count() == 20000) + val relation = spark.sharedState.cacheManager + .lookupCachedData(spark.table("codec_size_cache")) + .get + .cachedRepresentation + val batches = relation.cacheBuilder.cachedColumnBuffers.collect() + val payload = batches.map(CometCachedBatchHelper.payloadSize).sum + sizes += ((codec, relation.computeStats().sizeInBytes.toLong, payload)) + spark.catalog.clearCache() + } + } + + val (_, zstdSize, zstdPayload) = sizes(0) + val (_, plainSize, _) = sizes(1) + assert(zstdPayload * 2 < zstdSize, s"zstd should compress this relation: $sizes") + assert(zstdSize == plainSize, s"the relation size should not depend on the codec: $sizes") + } + test("Comet in-memory cache scans no columns for a row-count-only query") { // SELECT count(*) selects no columns, and the scan must keep it that way. Widening it -- to // the whole cache schema, or to a single placeholder column -- makes the emitted batches diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index 1b0d6eff78a..9b0f652a19e 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -26,7 +26,7 @@ import org.apache.spark.benchmark.Benchmark import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.comet.CometInMemoryTableScanExec -import org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer +import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, DefaultCachedBatchSerializer, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.vectorized.OnHeapColumnVector import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} @@ -452,10 +452,12 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { .optimizedPlan .collectFirst { case r: InMemoryRelation => r } .getOrElse(sys.error(s"$view is not cached")) - // computeStats rather than the builder's size accumulator, which Spark 4.2 replaced. Before - // the buffers load it falls back to the plan's estimate, so insist they have. + // The stored payloads rather than computeStats, which reports the relation's decoded size and + // so is the same for every codec. assert(relation.cacheBuilder.isCachedColumnBuffersLoaded, s"$view is not materialized") - relation.computeStats().sizeInBytes.toLong + relation.cacheBuilder.cachedColumnBuffers + .map(CometCachedBatchHelper.payloadSize) + .fold(0L)(_ + _) } private def runStatsBenchmark(): Unit = { diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala index 14975cdd21b..3b88ef96b06 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometCachedBatchHelper.scala @@ -127,7 +127,7 @@ object CometCachedBatchHelper { /** A payload [[serialize]] wrote, as the cached batch the writer would have stored it in. */ def cachedBatch(payload: ChunkedByteBuffer, numRows: Int): CachedBatch = - CometCachedBatch(numRows, payload.size, InternalRow.empty, payload) + CometCachedBatch(numRows, InternalRow.empty, payload) /** * Decode the `selected` columns of a cached batch the way a scan does, into a root the caller @@ -197,6 +197,21 @@ object CometCachedBatchHelper { def columnSizes(batch: CachedBatch, cacheSchema: StructType): Seq[Long] = columnBufferRanges(batch, cacheSchema).map(_.map(_._2).sum) + /** + * Decoded size of each top-level column: the column read back out of the payload the way a scan + * reads it, measured the way Spark's own Arrow cache format measures a column for its + * statistics. + */ + def decodedColumnSizes( + batch: CachedBatch, + cacheSchema: StructType, + allocator: BufferAllocator): Seq[Long] = + cacheSchema.indices.map { i => + val root = load(batch, cacheSchema, Array(i), allocator) + try root.getVector(0).getBufferSize.toLong + finally root.close() + } + /** * Whether any of a column's buffers is actually stored compressed. * From 689938082cd97d85a91e3cc49bb03dea3bcfba4b Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 29 Sep 2026 12:29:48 -0600 Subject: [PATCH 12/24] fix: keep Spark's cache scan for a relation whose cached plan records observed metrics Spark collects Dataset.observe metrics after a query with CollectMetricsExec.collect, which reaches the metrics recorded inside a cached plan only through an InMemoryTableScanExec over it. With Comet's native cache scan in its place those metrics were lost: QueryExecution.observedMetrics came back empty, Observation.get returned an empty map on Spark 3.5 and later, and on Spark 3.4 it never returned. Keep Spark's scan for such a relation, with a fallback reason. The data stays in Comet's format and is read through the existing fallback path. The check walks the cached plan the way CollectMetricsExec.collect does, through subqueries, adaptive plans and nested caches. --- .../user-guide/latest/in-memory-cache.md | 4 ++ .../apache/comet/rules/CometExecRule.scala | 14 ++++- .../comet/CometInMemoryTableScanExec.scala | 25 +++++++- .../comet/exec/CometInMemoryCacheSuite.scala | 59 ++++++++++++++++++- 4 files changed, 95 insertions(+), 7 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 8e39f00fe5e..4414ea49530 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -55,6 +55,10 @@ under one setting stays readable after the setting changes. Turning `spark.comet.exec.inMemoryCache.enabled` off at runtime only sends cached scans back to Spark's execution path; the cached data stays readable either way. +A relation whose cached plan records observed metrics, from `Dataset.observe`, is still stored in +Comet's format but is scanned by Spark's `InMemoryTableScanExec`, because Spark collects those +metrics only through that scan. + ## Storage format Each cached batch is stored as a single Arrow IPC record batch message and its body. diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 9b56631aee8..b24e4263be5 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -377,8 +377,12 @@ case class CometExecRule(session: SparkSession) val cometCacheFormat = usesCometCacheSerializer && ArrowCachedBatchSerializer.supportsSchema(scan.relation.output) val nativeCacheEnabled = CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.get(conf) + // Walks the cached plan, so it is only consulted once the native scan is otherwise + // possible. See CometInMemoryTableScanExec.recordsObservedMetrics. + val nativeScan = nativeCacheEnabled && cometCacheFormat && + !CometInMemoryTableScanExec.recordsObservedMetrics(scan.relation) - if (nativeCacheEnabled && cometCacheFormat) { + if (nativeScan) { convertToComet(scan, CometInMemoryTableScanExec).getOrElse(scan) } else { // The native cache scan is not available for this relation. Record why, then take the @@ -389,7 +393,7 @@ case class CometExecRule(session: SparkSession) scan, s"Comet in-memory cache requires ${classOf[ArrowCachedBatchSerializer].getName} " + s"but this relation was cached with ${serializer.getClass.getName}") - } else if (nativeCacheEnabled) { + } else if (nativeCacheEnabled && !cometCacheFormat) { val unsupported = scan.relation.output .filterNot(a => ArrowCachedBatchSerializer.supportsType(a.dataType)) .map(a => s"${a.name}: ${a.dataType.simpleString}") @@ -397,6 +401,12 @@ case class CometExecRule(session: SparkSession) scan, "Comet in-memory cache does not support the type of these cached columns, so the " + s"relation was cached in Spark's default format: ${unsupported.mkString(", ")}") + } else if (nativeCacheEnabled) { + withFallbackReason( + scan, + "Comet in-memory cache does not scan a relation whose cached plan records " + + "Dataset.observe metrics, because Spark collects those metrics only through " + + "InMemoryTableScanExec") } else if (usesCometCacheSerializer) { withFallbackReason( scan, diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometInMemoryTableScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometInMemoryTableScanExec.scala index 103f95f2e53..7943c7431b2 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometInMemoryTableScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometInMemoryTableScanExec.scala @@ -26,8 +26,9 @@ import org.apache.spark.sql.catalyst.expressions.Attribute import org.apache.spark.sql.catalyst.plans.logical.Statistics import org.apache.spark.sql.columnar.{CachedBatch, CachedBatchSerializer} import org.apache.spark.sql.comet.shims.ShimCometInMemoryTableScanExec -import org.apache.spark.sql.execution.SparkPlan -import org.apache.spark.sql.execution.columnar.{CachedRDDBuilder, InMemoryTableScanExec} +import org.apache.spark.sql.execution.{CollectMetricsExec, SparkPlan} +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.columnar.{CachedRDDBuilder, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} import org.apache.spark.sql.vectorized.ColumnarBatch @@ -164,4 +165,24 @@ object CometInMemoryTableScanExec extends CometOperatorSerde[InMemoryTableScanEx op.output)) } + /** + * Whether `relation`'s cached plan records observed metrics, from `Dataset.observe`. + * + * Spark collects those metrics once a query finishes, with `CollectMetricsExec.collect`, and + * that reaches the ones recorded inside a cached plan only through an `InMemoryTableScanExec` + * over it. It does not know this node, so replacing the scan of such a relation leaves its + * metrics empty, and on Spark 3.4 leaves `Observation.get` waiting for good. The walk mirrors + * `CollectMetricsExec.collect`, through subqueries, adaptive plans and nested caches. + */ + def recordsObservedMetrics(relation: InMemoryRelation): Boolean = + ObservedMetrics.recordedIn(relation.cachedPlan) + + private object ObservedMetrics extends AdaptiveSparkPlanHelper { + def recordedIn(plan: SparkPlan): Boolean = + collectWithSubqueries(plan) { + case _: CollectMetricsExec => true + case scan: InMemoryTableScanExec => recordedIn(scan.relation.cachedPlan) + }.contains(true) + } + } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 5933f7fa9c2..f4383ca8c3b 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -22,6 +22,9 @@ package org.apache.comet.exec import java.{util => ju} import java.nio.charset.StandardCharsets +import scala.concurrent.{Await, Future} +import scala.concurrent.ExecutionContext.Implicits.global +import scala.concurrent.duration.DurationInt import scala.jdk.CollectionConverters._ import org.apache.arrow.compression.ZstdCompressionCodec @@ -31,7 +34,7 @@ import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil, N import org.apache.arrow.vector.types.pojo.ArrowType import org.apache.spark.CometDriverPlugin import org.apache.spark.SparkConf -import org.apache.spark.sql.{CometTestBase, QueryTest, Row} +import org.apache.spark.sql.{CometTestBase, Observation, QueryTest, Row} import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, EqualTo, Expression, GreaterThan, GreaterThanOrEqual, LessThan, Literal} import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} @@ -42,13 +45,13 @@ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffl import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.exchange.{Exchange, ReusedExchangeExec, ShuffleExchangeLike} import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, SortMergeJoinExec} -import org.apache.spark.sql.functions.max +import org.apache.spark.sql.functions.{count, lit, max, min, sum} import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} import org.apache.spark.sql.types._ import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} import org.apache.spark.storage.StorageLevel -import org.apache.comet.{CometArrowAllocator, CometConf} +import org.apache.comet.{CometArrowAllocator, CometConf, ExtendedExplainInfo} import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus} import org.apache.comet.vector.{CometPlainVector, CometVector} @@ -2504,4 +2507,54 @@ class CometInMemoryCacheSuite extends CometTestBase { spark.catalog.clearCache() } } + + test("Comet in-memory cache keeps the observed metrics recorded in a cached plan") { + // Spark collects the metrics of an observe() inside a cached plan only through an + // InMemoryTableScanExec over it, so the scan of such a relation has to stay Spark's. Replaced, + // the metrics come back empty, and on Spark 3.4 Observation.get never returns. Nested the way + // SPARK-35695's test nests it, with a shuffle in the inner cached plan so that AQE plans it. + withAQECache { + val df = spark + .range(0, 100, 1, 2) + .repartition(4) + .observe("inner_event", count(lit(1)).as("rows"), max($"id").as("max_id")) + .persist() + .observe("outer_event", min($"id").as("min_id")) + .persist() + df.collect() + assert( + df.queryExecution.observedMetrics == + Map("inner_event" -> Row(100L, 99L), "outer_event" -> Row(0L))) + val plan = df.queryExecution.executedPlan + assert(collect(plan) { case s: CometInMemoryTableScanExec => s }.isEmpty) + assert( + new ExtendedExplainInfo() + .generateExtendedInfo(plan) + .contains("records Dataset.observe metrics")) + // Still stored in Comet's format: only the scan changes. + assert( + spark.sharedState.cacheManager + .lookupCachedData(df) + .get + .cachedRepresentation + .cacheBuilder + .cachedColumnBuffers + .map(_.getClass.getName) + .distinct() + .collect() + .sameElements(Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) + + val observation = Observation("cached_observation") + val observed = spark.range(10).observe(observation, sum($"id").as("total")).persist() + observed.collect() + // Bounded, so that a regression fails here rather than hanging the suite on Spark 3.4. + assert(Await.result(Future(observation.get), 1.minute) == Map("total" -> 45L)) + + val plain = spark.range(0, 100, 1, 2).persist() + plain.collect() + assert(collect(plain.queryExecution.executedPlan) { case s: CometInMemoryTableScanExec => + s + }.nonEmpty) + } + } } From 0ba6716ca7a7d74cc9437f2b1e2224bfaadde32d Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 30 Sep 2026 07:04:22 -0600 Subject: [PATCH 13/24] test: install Comet's cache serializer in the Spark SQL test sessions The Spark SQL diffs turn Comet on through SharedSparkSession and TestHive rather than by loading CometPlugin, and the plugin is what installs Comet's cache serializer. So even with spark.comet.exec.inMemoryCache.enabled on by default, every Spark SQL suite cached in Spark's own format and none of them exercised Comet's. Set spark.sql.cache.serializer in both places when Comet is enabled, as the plugin would, and adapt the tests that assume Spark's cache scan: - plan checks that only need a cache scan to be present also accept CometInMemoryTableScanExec, reading its originalPlan where the test uses the scan's relation; - PartitionBatchPruningSuite checks its answers as before, but reads the test-only accumulators only from Spark's scan, which Comet's lacks; - CacheTableInKryoSuite registers Comet's classes with CometKryoRegistrator, as Comet asks of any application that sets spark.kryo.registrationRequired; - tests of Spark internals that Comet's scan replaces are tagged IgnoreComet: exact cache size estimates, the union's columnar support, subquery reuse through the scan's predicates, and AQE coalescing under a union (#6454). Each diff was regenerated from a clone of its Spark tag with the existing diff applied, after checking that it round-tripped unchanged. --- dev/diffs/3.4.3.diff | 326 +++++++++++++++++++++++++++++-- dev/diffs/3.5.9.diff | 369 +++++++++++++++++++++++++++++++++-- dev/diffs/4.0.4.diff | 405 ++++++++++++++++++++++++++++++++++++-- dev/diffs/4.1.3.diff | 455 +++++++++++++++++++++++++++++++++++++++++-- dev/diffs/4.2.0.diff | 455 +++++++++++++++++++++++++++++++++++++++++-- 5 files changed, 1927 insertions(+), 83 deletions(-) diff --git a/dev/diffs/3.4.3.diff b/dev/diffs/3.4.3.diff index 8fc6451ef12..2080f0eda7b 100644 --- a/dev/diffs/3.4.3.diff +++ b/dev/diffs/3.4.3.diff @@ -237,10 +237,14 @@ index 0efe0877e9b..423d3b3d76d 100644 -- SELECT_HAVING -- https://github.com/postgres/postgres/blob/REL_12_BETA2/src/test/regress/sql/select_having.sql diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index cf40e944c09..bdd5be4f462 100644 +index cf40e944c09..a10f3d46a69 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -@@ -38,7 +38,7 @@ import org.apache.spark.sql.catalyst.util.DateTimeConstants +@@ -35,10 +35,11 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException + import org.apache.spark.sql.catalyst.expressions.SubqueryExpression + import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, Join, JoinStrategyHint, SHUFFLE_HASH} + import org.apache.spark.sql.catalyst.util.DateTimeConstants ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec import org.apache.spark.sql.execution.{ColumnarToRowExec, ExecSubqueryExpression, RDDScanExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.execution.columnar._ @@ -249,7 +253,27 @@ index cf40e944c09..bdd5be4f462 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -516,7 +516,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -113,6 +114,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => + getNumInMemoryTablesRecursively(relation.cachedPlan) + + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ case cometScan: CometInMemoryTableScanExec => ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + ++ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 + case p => + getNumInMemoryTablesInSubquery(p) + }.sum +@@ -393,7 +397,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + assert(isExpectStorageLevel(rddId, Disk)) + } + +- test("InMemoryRelation statistics") { ++ test("InMemoryRelation statistics", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + sql("CACHE TABLE testData") + spark.table("testData").queryExecution.withCachedData.collect { + case cached: InMemoryRelation => +@@ -516,7 +521,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils */ private def verifyNumExchanges(df: DataFrame, expected: Int): Unit = { assert( @@ -259,6 +283,16 @@ index cf40e944c09..bdd5be4f462 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { +@@ -1559,7 +1565,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + } + } + +- test("SPARK-36120: Support cache/uncache table with TimestampNTZ type") { ++ test("SPARK-36120: Support cache/uncache table with TimestampNTZ type", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + val tableName = "ntzCache" + withTable(tableName) { + sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala index 1cc09c3d7fc..9e1e883d450 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala @@ -315,6 +349,20 @@ index 56e9520fdab..917932336df 100644 spark.range(50).write.saveAsTable(s"$dbName.$table1Name") spark.range(100).write.saveAsTable(s"$dbName.$table2Name") +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +index 61724a39dfa..8aa517c1575 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +@@ -1376,7 +1376,8 @@ class DataFrameSetOperationsSuite extends QueryTest with SharedSparkSession { + Row(Row(Seq(Seq(Row(null, "ba"))))) :: Nil) + } + +- test("SPARK-37371: UnionExec should support columnar if all children support columnar") { ++ test("SPARK-37371: UnionExec should support columnar if all children support columnar", ++ IgnoreComet("Comet replaces the cache scans and the union with its own operators")) { + def checkIfColumnar( + plan: SparkPlan, + targetPlan: (SparkPlan) => Boolean, diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala index a9f69ab28a1..760ea0e9565 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala @@ -438,6 +486,66 @@ index 433b4741979..e13e69deb79 100644 case _ => false } case _ => false +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +index a657c6212aa..c90f10c8fa5 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +@@ -20,6 +20,8 @@ package org.apache.spark.sql + import org.scalatest.concurrent.TimeLimits + import org.scalatest.time.SpanSugar._ + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec ++import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.columnar.{InMemoryRelation, InMemoryTableScanExec} + import org.apache.spark.sql.functions._ +@@ -34,6 +36,10 @@ class DatasetCacheSuite extends QueryTest + with AdaptiveSparkPlanHelper { + import testImplicits._ + ++ // A scan of a cached relation, Spark's or Comet's. ++ private def isCacheScan(plan: SparkPlan): Boolean = ++ plan.isInstanceOf[InMemoryTableScanExec] || plan.isInstanceOf[CometInMemoryTableScanExec] ++ + /** + * Asserts that a cached [[Dataset]] will be built using the given number of other cached results. + */ +@@ -41,7 +47,7 @@ class DatasetCacheSuite extends QueryTest + val plan = df.queryExecution.withCachedData + assert(plan.isInstanceOf[InMemoryRelation]) + val internalPlan = plan.asInstanceOf[InMemoryRelation].cacheBuilder.cachedPlan +- assert(find(internalPlan)(_.isInstanceOf[InMemoryTableScanExec]).size ++ assert(find(internalPlan)(isCacheScan).size + == numOfCachesDependedUpon) + } + +@@ -251,7 +257,7 @@ class DatasetCacheSuite extends QueryTest + case i: InMemoryRelation => i.cacheBuilder.cachedPlan + } + assert(df2LimitInnerPlan.isDefined && +- !df2LimitInnerPlan.get.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ !df2LimitInnerPlan.get.exists(isCacheScan)) + } + + test("SPARK-27739 Save stats from optimized plan") { +@@ -283,7 +289,7 @@ class DatasetCacheSuite extends QueryTest + val unionDf = df1.union(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.union(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative") { +@@ -291,7 +297,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.union(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index daef11ae4d6..9f3cc9181f2 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -986,10 +1094,18 @@ index b5b34922694..a72403780c4 100644 protected val baseResourcePath = { // use the same way as `SQLQueryTestSuite` to get the resource path diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -index 525d97e4998..aded8906d75 100644 +index 525d97e4998..c15eaf98ddc 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -@@ -1508,7 +1508,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -36,6 +36,7 @@ import org.apache.spark.sql.catalyst.expressions.aggregate.{Complete, Partial} + import org.apache.spark.sql.catalyst.optimizer.{ConvertToLocalRelation, NestedColumnAliasingSuite} + import org.apache.spark.sql.catalyst.plans.logical.{LocalLimit, Project, RepartitionByExpression, Sort} + import org.apache.spark.sql.connector.catalog.CatalogManager.SESSION_CATALOG_NAME ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.{CommandResultExec, UnionExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.aggregate._ +@@ -1508,7 +1509,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark checkAnswer(sql("select -0.001"), Row(BigDecimal("-0.001"))) } @@ -999,7 +1115,7 @@ index 525d97e4998..aded8906d75 100644 AccumulatorSuite.verifyPeakExecutionMemorySet(sparkContext, "external sort") { sql("SELECT * FROM testData2 ORDER BY a ASC, b ASC").collect() } -@@ -1960,8 +1961,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -1960,8 +1962,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark countAcc.add(1) x }) @@ -1016,7 +1132,15 @@ index 525d97e4998..aded8906d75 100644 verifyCallCount( df.selectExpr("testUdf(a + 1) + testUdf(1 + a)", "testUdf(a + 1)"), Row(4, 2), 1) -@@ -3730,7 +3738,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -3269,6 +3278,7 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark + + val inMemoryTableScan = collect(queryDf.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c + } + assert(inMemoryTableScan.size == 2) + checkAnswer(queryDf, Row(0, 1) :: Row(1, 2) :: Row(2, 3) :: Row(3, 4) :: Row(4, 5) :: Nil) +@@ -3730,7 +3740,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark } } @@ -1056,6 +1180,32 @@ index 2dabcf01be7..8fcec0d1ce4 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +index 48ad10992c5..23bf476d4dc 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +@@ -208,6 +208,7 @@ class SparkSessionExtensionSuite extends SparkFunSuite with SQLHelper { + df.select("i").filter($"i" > 1).cache() + assert(df.filter($"i" > 1).select("i").queryExecution.executedPlan.find { + case _: org.apache.spark.sql.execution.columnar.InMemoryTableScanExec => true ++ case _: org.apache.spark.sql.comet.CometInMemoryTableScanExec => true + case _ => false + }.isDefined) + } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +index e6b74a328e5..d4aa93eaebb 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +@@ -811,7 +811,8 @@ class StatisticsCollectionSuite extends StatisticsCollectionTestBase with Shared + } + } + +- test("SPARK-33687: analyze all tables in a specific database") { ++ test("SPARK-33687: analyze all tables in a specific database", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + withTempDatabase { database => + spark.catalog.setCurrentDatabase(database) + withTempDir { dir => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala index 18123a4d6ec..0fe185baa33 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala @@ -1964,6 +2114,151 @@ index 593bd7bb4ba..b327d84d5cc 100644 val emptyDf = spark.range(1).where("false") val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df1")) val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df2")) +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +index d15fabd9403..9fc89bbd2a0 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +@@ -22,11 +22,12 @@ import java.sql.{Date, Timestamp} + import java.util.concurrent.atomic.AtomicInteger + + import org.apache.spark.rdd.RDD +-import org.apache.spark.sql.{DataFrame, QueryTest, Row} ++import org.apache.spark.sql.{DataFrame, IgnoreComet, QueryTest, Row} + import org.apache.spark.sql.catalyst.InternalRow + import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSet, In} + import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning + import org.apache.spark.sql.columnar.CachedBatch ++import org.apache.spark.sql.comet.{CometFilterExec, CometInMemoryTableScanExec} + import org.apache.spark.sql.execution.{FilterExec, InputAdapter, WholeStageCodegenExec} + import org.apache.spark.sql.functions._ + import org.apache.spark.sql.internal.SQLConf +@@ -507,11 +508,17 @@ class InMemoryColumnarQuerySuite extends QueryTest with SharedSparkSession { + val planBeforeFilter = df2.queryExecution.executedPlan.collect { + case f: FilterExec => f.child + case WholeStageCodegenExec(FilterExec(_, i: InputAdapter)) => i.child ++ case f: CometFilterExec => f.child + } +- assert(planBeforeFilter.head.isInstanceOf[InMemoryTableScanExec]) +- + val execPlan = planBeforeFilter.head +- assert(execPlan.executeCollectPublic().length == 0) ++ execPlan match { ++ // Comet's cache scan is columnar only, so count the rows of the batches it emits. ++ case c: CometInMemoryTableScanExec => ++ assert(c.executeColumnar().map(_.numRows().toLong).collect().sum == 0) ++ case _ => ++ assert(execPlan.isInstanceOf[InMemoryTableScanExec]) ++ assert(execPlan.executeCollectPublic().length == 0) ++ } + } + + test("SPARK-25727 - otherCopyArgs in InMemoryRelation does not include outputOrdering") { +@@ -520,7 +527,8 @@ class InMemoryColumnarQuerySuite extends QueryTest with SharedSparkSession { + assert(json.contains("outputOrdering")) + } + +- test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached") { ++ test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + Seq("orc", "").foreach { useV1SourceReaderList => + // This test case depends on the size of ORC in statistics. + withSQLConf( +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +index e032e0c2b27..89d7867c250 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +@@ -18,6 +18,7 @@ + package org.apache.spark.sql.execution.columnar + + import org.apache.spark.SparkFunSuite ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions.expr +@@ -38,11 +39,15 @@ class InMemoryRelationSuite extends SparkFunSuite + test("SPARK-47177: Cached SQL plan do not display final AQE plan in explain string") { + def findIMRInnerChild(p: SparkPlan): SparkPlan = { + val tableCache = find(p) { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + case _ => false + } + assert(tableCache.isDefined) +- tableCache.get.asInstanceOf[InMemoryTableScanExec].relation.innerChildren.head ++ val scan = tableCache.get match { ++ case c: CometInMemoryTableScanExec => c.originalPlan ++ case s => s.asInstanceOf[InMemoryTableScanExec] ++ } ++ scan.relation.innerChildren.head + } + + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true") { +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +index a22cb664744..3831e428785 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +@@ -17,6 +17,7 @@ + + package org.apache.spark.sql.execution.columnar + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.internal.SQLConf + import org.apache.spark.sql.test.SharedSparkSession + import org.apache.spark.sql.test.SQLTestData._ +@@ -180,11 +181,16 @@ class PartitionBatchPruningSuite extends SharedSparkSession { + val result = df.collect().map(_(0)).toArray + assert(result.length === 1) + +- val (readPartitions, readBatches) = df.queryExecution.executedPlan.collect { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- assert(readPartitions === 5) +- assert(readBatches === 10) ++ val scans = df.queryExecution.executedPlan.collect { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan has none of these test-only accumulators, so there is nothing to count. ++ if (scans.isEmpty) { ++ assert(df.queryExecution.executedPlan.collect { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ assert(scans.head.readPartitions.value === 5) ++ assert(scans.head.readBatches.value === 10) ++ } + } + + def checkBatchPruning( +@@ -201,14 +207,23 @@ class PartitionBatchPruningSuite extends SharedSparkSession { + df.collect().map(_(0)).toArray + } + +- val (readPartitions, readBatches) = df.queryExecution.executedPlan.collect { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- +- assert(readBatches === expectedReadBatches, s"Wrong number of read batches: $queryExecution") +- assert( +- readPartitions === expectedReadPartitions, +- s"Wrong number of read partitions: $queryExecution") ++ val scans = df.queryExecution.executedPlan.collect { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan prunes on the same statistics, but has none of these test-only ++ // accumulators to read, so only the answer above is checked for it. ++ if (scans.isEmpty) { ++ assert(df.queryExecution.executedPlan.collect { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ val readPartitions = scans.head.readPartitions.value ++ val readBatches = scans.head.readBatches.value ++ assert( ++ readBatches === expectedReadBatches, ++ s"Wrong number of read batches: $queryExecution") ++ assert( ++ readPartitions === expectedReadPartitions, ++ s"Wrong number of read partitions: $queryExecution") ++ } + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala index bd9c79e5b96..2ada8c28842 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala @@ -3053,10 +3348,10 @@ index dd55fcfe42c..d9a3f2df535 100644 spark.internalCreateDataFrame(withoutFilters.execute(), schema) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -index ed2e309fa07..54d417624ff 100644 +index ed2e309fa07..8e3aaa888c7 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -@@ -74,6 +74,19 @@ trait SharedSparkSessionBase +@@ -74,6 +74,24 @@ trait SharedSparkSessionBase // this rule may potentially block testing of other optimization rules such as // ConstantPropagation etc. .set(SQLConf.OPTIMIZER_EXCLUDED_RULES.key, ConvertToLocalRelation.ruleName) @@ -3072,6 +3367,11 @@ index ed2e309fa07..54d417624ff 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ // CometDriverPlugin installs Comet's cache serializer when ++ // spark.comet.exec.inMemoryCache.enabled is on, as it is by default. These sessions do not ++ // load the plugin, so install it here. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } conf.set( StaticSQLConf.WAREHOUSE_PATH, @@ -3157,10 +3457,10 @@ index a902cb3a69e..e652edd9f81 100644 test("SPARK-4963 DataFrame sample on mutable row return wrong result") { diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -index 07361cfdce9..af6dcfc2302 100644 +index 07361cfdce9..f9002ce0d98 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -@@ -55,25 +55,41 @@ object TestHive +@@ -55,25 +55,45 @@ object TestHive new SparkContext( System.getProperty("spark.sql.test.master", "local[1]"), "TestSQLContext", @@ -3212,6 +3512,10 @@ index 07361cfdce9..af6dcfc2302 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ ++ // As in SharedSparkSession: what CometDriverPlugin would install. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } + conf diff --git a/dev/diffs/3.5.9.diff b/dev/diffs/3.5.9.diff index 462554a1047..fbc99fa97ec 100644 --- a/dev/diffs/3.5.9.diff +++ b/dev/diffs/3.5.9.diff @@ -218,10 +218,14 @@ index 0efe0877e9b..423d3b3d76d 100644 -- SELECT_HAVING -- https://github.com/postgres/postgres/blob/REL_12_BETA2/src/test/regress/sql/select_having.sql diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index e5494726695..00937f025c2 100644 +index e5494726695..eaa3ea7a222 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -@@ -38,7 +38,7 @@ import org.apache.spark.sql.catalyst.util.DateTimeConstants +@@ -35,10 +35,11 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException + import org.apache.spark.sql.catalyst.expressions.SubqueryExpression + import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, Join, JoinStrategyHint, SHUFFLE_HASH} + import org.apache.spark.sql.catalyst.util.DateTimeConstants ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec import org.apache.spark.sql.execution.{ColumnarToRowExec, ExecSubqueryExpression, RDDScanExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEPropagateEmptyRelation} import org.apache.spark.sql.execution.columnar._ @@ -230,7 +234,27 @@ index e5494726695..00937f025c2 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -519,7 +519,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -113,6 +114,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => + getNumInMemoryTablesRecursively(relation.cachedPlan) + + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ case cometScan: CometInMemoryTableScanExec => ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + ++ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 + case p => + getNumInMemoryTablesInSubquery(p) + }.sum +@@ -393,7 +397,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + assert(isExpectStorageLevel(rddId, Disk)) + } + +- test("InMemoryRelation statistics") { ++ test("InMemoryRelation statistics", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + sql("CACHE TABLE testData") + spark.table("testData").queryExecution.withCachedData.collect { + case cached: InMemoryRelation => +@@ -519,7 +524,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils df.collect() } assert( @@ -240,6 +264,16 @@ index e5494726695..00937f025c2 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { +@@ -1574,7 +1580,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + } + } + +- test("SPARK-36120: Support cache/uncache table with TimestampNTZ type") { ++ test("SPARK-36120: Support cache/uncache table with TimestampNTZ type", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + val tableName = "ntzCache" + withTable(tableName) { + sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala index 6f3090d8908..4774aad5019 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala @@ -296,6 +330,20 @@ index 56e9520fdab..917932336df 100644 spark.range(50).write.saveAsTable(s"$dbName.$table1Name") spark.range(100).write.saveAsTable(s"$dbName.$table2Name") +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +index 747f43fa2a7..2e52a4acb79 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +@@ -1398,7 +1398,8 @@ class DataFrameSetOperationsSuite extends QueryTest + Row(Row(Seq(Seq(Row(null, "ba"))))) :: Nil) + } + +- test("SPARK-37371: UnionExec should support columnar if all children support columnar") { ++ test("SPARK-37371: UnionExec should support columnar if all children support columnar", ++ IgnoreComet("Comet replaces the cache scans and the union with its own operators")) { + def checkIfColumnar( + plan: SparkPlan, + targetPlan: (SparkPlan) => Boolean, diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala index 7ee18df3756..d09f70e5d99 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala @@ -419,6 +467,83 @@ index a1d5d579338..8825683ebcd 100644 case _ => false } case _ => false +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +index bda8c7f2608..10a70745d07 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +@@ -20,6 +20,8 @@ package org.apache.spark.sql + import org.scalatest.concurrent.TimeLimits + import org.scalatest.time.SpanSugar._ + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec ++import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.columnar.{InMemoryRelation, InMemoryTableScanExec} + import org.apache.spark.sql.functions._ +@@ -35,6 +37,10 @@ class DatasetCacheSuite extends QueryTest + with AdaptiveSparkPlanHelper { + import testImplicits._ + ++ // A scan of a cached relation, Spark's or Comet's. ++ private def isCacheScan(plan: SparkPlan): Boolean = ++ plan.isInstanceOf[InMemoryTableScanExec] || plan.isInstanceOf[CometInMemoryTableScanExec] ++ + /** + * Asserts that a cached [[Dataset]] will be built using the given number of other cached results. + */ +@@ -42,7 +48,7 @@ class DatasetCacheSuite extends QueryTest + val plan = df.queryExecution.withCachedData + assert(plan.isInstanceOf[InMemoryRelation]) + val internalPlan = plan.asInstanceOf[InMemoryRelation].cacheBuilder.cachedPlan +- assert(find(internalPlan)(_.isInstanceOf[InMemoryTableScanExec]).size ++ assert(find(internalPlan)(isCacheScan).size + == numOfCachesDependedUpon) + } + +@@ -252,7 +258,7 @@ class DatasetCacheSuite extends QueryTest + case i: InMemoryRelation => i.cacheBuilder.cachedPlan + } + assert(df2LimitInnerPlan.isDefined && +- !df2LimitInnerPlan.get.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ !df2LimitInnerPlan.get.exists(isCacheScan)) + } + + test("SPARK-27739 Save stats from optimized plan") { +@@ -285,14 +291,14 @@ class DatasetCacheSuite extends QueryTest + val unionDf = df1.union(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.union(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("positive: union by name") { + val unionDf = df1.unionByName(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.unionByName(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by position") { +@@ -300,7 +306,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.union(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by name") { +@@ -308,7 +314,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.unionByName(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index c4fb4fa943c..a04b23870a8 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -993,10 +1118,18 @@ index c26757c9cff..d55775f09d7 100644 protected val baseResourcePath = { // use the same way as `SQLQueryTestSuite` to get the resource path diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -index 3cf2bfd17ab..5bcf9478e9b 100644 +index 3cf2bfd17ab..11d3ca4d5a9 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -@@ -1521,7 +1521,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -38,6 +38,7 @@ import org.apache.spark.sql.catalyst.optimizer.{ConvertToLocalRelation, NestedCo + import org.apache.spark.sql.catalyst.parser.ParseException + import org.apache.spark.sql.catalyst.plans.logical.{LocalLimit, Project, RepartitionByExpression, Sort} + import org.apache.spark.sql.connector.catalog.CatalogManager.SESSION_CATALOG_NAME ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.{CommandResultExec, UnionExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.aggregate._ +@@ -1521,7 +1522,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark checkAnswer(sql("select -0.001"), Row(BigDecimal("-0.001"))) } @@ -1006,7 +1139,7 @@ index 3cf2bfd17ab..5bcf9478e9b 100644 AccumulatorSuite.verifyPeakExecutionMemorySet(sparkContext, "external sort") { sql("SELECT * FROM testData2 ORDER BY a ASC, b ASC").collect() } -@@ -1979,8 +1980,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -1979,8 +1981,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark countAcc.add(1) x }) @@ -1023,7 +1156,15 @@ index 3cf2bfd17ab..5bcf9478e9b 100644 verifyCallCount( df.selectExpr("testUdf(a + 1) + testUdf(1 + a)", "testUdf(a + 1)"), Row(4, 2), 1) -@@ -3750,7 +3758,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -3289,6 +3298,7 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark + + val inMemoryTableScan = collect(queryDf.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c + } + assert(inMemoryTableScan.size == 2) + checkAnswer(queryDf, Row(0, 1) :: Row(1, 2) :: Row(2, 3) :: Row(3, 4) :: Row(4, 5) :: Nil) +@@ -3750,7 +3760,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark } } @@ -1064,6 +1205,32 @@ index 71af1fd69c3..81a04c93c9c 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +index 8b4ac474f87..096146cd69a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +@@ -210,6 +210,7 @@ class SparkSessionExtensionSuite extends SparkFunSuite with SQLHelper with Adapt + df.select("i").filter($"i" > 1).cache() + assert(find(df.filter($"i" > 1).select("i").queryExecution.executedPlan) { + case _: org.apache.spark.sql.execution.columnar.InMemoryTableScanExec => true ++ case _: org.apache.spark.sql.comet.CometInMemoryTableScanExec => true + case _ => false + }.isDefined) + } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +index e827396009d..066d1a6bf04 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +@@ -811,7 +811,8 @@ class StatisticsCollectionSuite extends StatisticsCollectionTestBase with Shared + } + } + +- test("SPARK-33687: analyze all tables in a specific database") { ++ test("SPARK-33687: analyze all tables in a specific database", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + withTempDatabase { database => + spark.catalog.setCurrentDatabase(database) + withTempDir { dir => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SubquerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SubquerySuite.scala index 04702201f82..4d38d8d6e51 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SubquerySuite.scala @@ -1561,7 +1728,7 @@ index 5a413c77754..207b66e1d7b 100644 import testImplicits._ diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala -index 2f8e401e743..7849c685b19 100644 +index 2f8e401e743..1b3a6ff17b5 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala @@ -27,9 +27,11 @@ import org.scalatest.time.SpanSugar._ @@ -1967,7 +2134,27 @@ index 2f8e401e743..7849c685b19 100644 }.size == (if (firstAccess) 2 else 0)) assert(collect(initialExecutedPlan) { case i: InMemoryTableScanLike => i -@@ -2938,7 +2981,8 @@ class AdaptiveQueryExecSuite +@@ -2898,7 +2941,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage") { ++ test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage", ++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/6454")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1") { + val cached = Seq(1).toDF("c").cache() +@@ -2923,7 +2967,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-43376: Improve reuse subquery with table cache") { ++ test("SPARK-43376: Improve reuse subquery with table cache", ++ IgnoreComet("Comet's cache scan does not plan the subqueries in its pruning predicates")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true") { + withTable("t1", "t2") { + withCache("t1") { +@@ -2938,7 +2983,8 @@ class AdaptiveQueryExecSuite } } @@ -1977,7 +2164,7 @@ index 2f8e401e743..7849c685b19 100644 val emptyDf = spark.range(1).where("false") val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df1")) val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df2")) -@@ -2980,7 +3024,9 @@ class AdaptiveQueryExecSuite +@@ -2980,7 +3026,9 @@ class AdaptiveQueryExecSuite val plan = df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec] assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec]) @@ -1988,6 +2175,151 @@ index 2f8e401e743..7849c685b19 100644 plan.inputPlan.output.zip(plan.finalPhysicalPlan.output).foreach { case (o1, o2) => assert(o1.semanticEquals(o2), "Different output column order after AQE optimization") } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +index de04938f247..fa019925b6f 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +@@ -22,11 +22,12 @@ import java.sql.{Date, Timestamp} + import java.util.concurrent.atomic.AtomicInteger + + import org.apache.spark.rdd.RDD +-import org.apache.spark.sql.{DataFrame, QueryTest, Row} ++import org.apache.spark.sql.{DataFrame, IgnoreComet, QueryTest, Row} + import org.apache.spark.sql.catalyst.InternalRow + import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSet, In} + import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning + import org.apache.spark.sql.columnar.CachedBatch ++import org.apache.spark.sql.comet.{CometFilterExec, CometInMemoryTableScanExec} + import org.apache.spark.sql.execution.{FilterExec, InputAdapter, WholeStageCodegenExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions._ +@@ -509,11 +510,17 @@ class InMemoryColumnarQuerySuite extends QueryTest + val planBeforeFilter = collect(df2.queryExecution.executedPlan) { + case f: FilterExec => f.child + case WholeStageCodegenExec(FilterExec(_, i: InputAdapter)) => i.child ++ case f: CometFilterExec => f.child + } +- assert(planBeforeFilter.head.isInstanceOf[InMemoryTableScanExec]) +- + val execPlan = planBeforeFilter.head +- assert(execPlan.executeCollectPublic().length == 0) ++ execPlan match { ++ // Comet's cache scan is columnar only, so count the rows of the batches it emits. ++ case c: CometInMemoryTableScanExec => ++ assert(c.executeColumnar().map(_.numRows().toLong).collect().sum == 0) ++ case _ => ++ assert(execPlan.isInstanceOf[InMemoryTableScanExec]) ++ assert(execPlan.executeCollectPublic().length == 0) ++ } + } + + test("SPARK-25727 - otherCopyArgs in InMemoryRelation does not include outputOrdering") { +@@ -522,7 +529,8 @@ class InMemoryColumnarQuerySuite extends QueryTest + assert(json.contains("outputOrdering")) + } + +- test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached") { ++ test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + Seq("orc", "").foreach { useV1SourceReaderList => + // This test case depends on the size of ORC in statistics. + withSQLConf( +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +index 2c73622739a..5d0efeb263a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +@@ -18,6 +18,7 @@ + package org.apache.spark.sql.execution.columnar + + import org.apache.spark.SparkFunSuite ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions.expr +@@ -37,11 +38,15 @@ class InMemoryRelationSuite extends SparkFunSuite + test("SPARK-47177: Cached SQL plan do not display final AQE plan in explain string") { + def findIMRInnerChild(p: SparkPlan): SparkPlan = { + val tableCache = find(p) { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + case _ => false + } + assert(tableCache.isDefined) +- tableCache.get.asInstanceOf[InMemoryTableScanExec].relation.innerChildren.head ++ val scan = tableCache.get match { ++ case c: CometInMemoryTableScanExec => c.originalPlan ++ case s => s.asInstanceOf[InMemoryTableScanExec] ++ } ++ scan.relation.innerChildren.head + } + + val d1 = spark.range(1).withColumn("key", expr("id % 100")) +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +index 885286843a1..f0f805a6cd9 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +@@ -17,6 +17,7 @@ + + package org.apache.spark.sql.execution.columnar + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.internal.SQLConf + import org.apache.spark.sql.test.SharedSparkSession +@@ -181,11 +182,16 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + val result = df.collect().map(_(0)).toArray + assert(result.length === 1) + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- assert(readPartitions === 5) +- assert(readBatches === 10) ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan has none of these test-only accumulators, so there is nothing to count. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ assert(scans.head.readPartitions.value === 5) ++ assert(scans.head.readBatches.value === 10) ++ } + } + + def checkBatchPruning( +@@ -202,14 +208,23 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + df.collect().map(_(0)).toArray + } + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- +- assert(readBatches === expectedReadBatches, s"Wrong number of read batches: $queryExecution") +- assert( +- readPartitions === expectedReadPartitions, +- s"Wrong number of read partitions: $queryExecution") ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan prunes on the same statistics, but has none of these test-only ++ // accumulators to read, so only the answer above is checked for it. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ val readPartitions = scans.head.readPartitions.value ++ val readBatches = scans.head.readBatches.value ++ assert( ++ readBatches === expectedReadBatches, ++ s"Wrong number of read batches: $queryExecution") ++ assert( ++ readPartitions === expectedReadPartitions, ++ s"Wrong number of read partitions: $queryExecution") ++ } + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala index fd52d038ca6..154c800be67 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala @@ -3064,10 +3396,10 @@ index e937173a590..263934fbe7b 100644 spark.internalCreateDataFrame(withoutFilters.execute(), schema) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -index c23bf4204f7..07d215aad2b 100644 +index c23bf4204f7..d922f0b3d23 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -@@ -97,6 +97,19 @@ trait SharedSparkSessionBase +@@ -97,6 +97,24 @@ trait SharedSparkSessionBase // this rule may potentially block testing of other optimization rules such as // ConstantPropagation etc. .set(SQLConf.OPTIMIZER_EXCLUDED_RULES.key, ConvertToLocalRelation.ruleName) @@ -3083,6 +3415,11 @@ index c23bf4204f7..07d215aad2b 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ // CometDriverPlugin installs Comet's cache serializer when ++ // spark.comet.exec.inMemoryCache.enabled is on, as it is by default. These sessions do not ++ // load the plugin, so install it here. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } conf.set( StaticSQLConf.WAREHOUSE_PATH, @@ -3168,10 +3505,10 @@ index 6160c3e5f6c..bfc0c618a9b 100644 test("SPARK-4963 DataFrame sample on mutable row return wrong result") { diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -index 1d646f40b3e..c8192f52f98 100644 +index 1d646f40b3e..b8a043c6248 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -@@ -53,25 +53,41 @@ object TestHive +@@ -53,25 +53,45 @@ object TestHive new SparkContext( System.getProperty("spark.sql.test.master", "local[1]"), "TestSQLContext", @@ -3223,6 +3560,10 @@ index 1d646f40b3e..c8192f52f98 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ ++ // As in SharedSparkSession: what CometDriverPlugin would install. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } + conf diff --git a/dev/diffs/4.0.4.diff b/dev/diffs/4.0.4.diff index 38ec78543aa..30201a5155f 100644 --- a/dev/diffs/4.0.4.diff +++ b/dev/diffs/4.0.4.diff @@ -333,10 +333,14 @@ index 21a3ce1e122..f4762ab98f0 100644 -- In COMPENSATION views get invalidated if the type can't cast diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index 0f42502f1d9..e9ff802141f 100644 +index 0f42502f1d9..18e5636f4cb 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -@@ -39,7 +39,7 @@ import org.apache.spark.sql.catalyst.util.DateTimeConstants +@@ -36,10 +36,11 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException + import org.apache.spark.sql.catalyst.expressions.SubqueryExpression + import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, Join, JoinStrategyHint, SHUFFLE_HASH} + import org.apache.spark.sql.catalyst.util.DateTimeConstants ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec import org.apache.spark.sql.execution.{ColumnarToRowExec, ExecSubqueryExpression, RDDScanExec, SparkPlan, SparkPlanInfo} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEPropagateEmptyRelation} import org.apache.spark.sql.execution.columnar._ @@ -345,7 +349,27 @@ index 0f42502f1d9..e9ff802141f 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -520,7 +520,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -114,6 +115,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => + getNumInMemoryTablesRecursively(relation.cachedPlan) + + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ case cometScan: CometInMemoryTableScanExec => ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + ++ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 + case p => + getNumInMemoryTablesInSubquery(p) + }.sum +@@ -394,7 +398,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + assert(isExpectStorageLevel(rddId, Disk)) + } + +- test("InMemoryRelation statistics") { ++ test("InMemoryRelation statistics", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + sql("CACHE TABLE testData") + spark.table("testData").queryExecution.withCachedData.collect { + case cached: InMemoryRelation => +@@ -520,7 +525,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils df.collect() } assert( @@ -355,7 +379,17 @@ index 0f42502f1d9..e9ff802141f 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1659,9 +1660,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1581,7 +1587,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + } + } + +- test("SPARK-36120: Support cache/uncache table with TimestampNTZ type") { ++ test("SPARK-36120: Support cache/uncache table with TimestampNTZ type", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + val tableName = "ntzCache" + withTable(tableName) { + sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") +@@ -1659,9 +1666,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils _.nodeName.contains("TableCacheQueryStage")) val aqeNode = findNodeInSparkPlanInfo(inMemoryScanNode.get, _.nodeName.contains("AdaptiveSparkPlan")) @@ -377,6 +411,14 @@ index 0f42502f1d9..e9ff802141f 100644 } withTempView("t0", "t1", "t2") { +@@ -1750,6 +1766,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + val cached = spark.table("t") + val tableCache = collect(cached.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + if (expected == StorageLevel.NONE) { + assert(tableCache.isEmpty) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala index 9db406ff12f..b3d55394d25 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala @@ -433,6 +475,20 @@ index ed182322aec..1ae6afa686a 100644 spark.range(50).write.saveAsTable(s"$dbName.$table1Name") spark.range(100).write.saveAsTable(s"$dbName.$table2Name") +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +index 332be4c7bbc..02899683f81 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +@@ -1474,7 +1474,8 @@ class DataFrameSetOperationsSuite extends QueryTest + Row(Row(Seq(Seq(Row(null, "ba"))))) :: Nil) + } + +- test("SPARK-37371: UnionExec should support columnar if all children support columnar") { ++ test("SPARK-37371: UnionExec should support columnar if all children support columnar", ++ IgnoreComet("Comet replaces the cache scans and the union with its own operators")) { + def checkIfColumnar( + plan: SparkPlan, + targetPlan: (SparkPlan) => Boolean, diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala index d9ce3000a0c..f2d044ed6b8 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSuite.scala @@ -556,6 +612,97 @@ index 552e2b2e274..17a5ae20f0f 100644 case _ => false } case _ => false +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +index 9d8aaf8d90e..41afdbdafc1 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +@@ -20,6 +20,8 @@ package org.apache.spark.sql + import org.scalatest.concurrent.TimeLimits + import org.scalatest.time.SpanSugar._ + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec ++import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.columnar.{InMemoryRelation, InMemoryTableScanExec} + import org.apache.spark.sql.functions._ +@@ -36,6 +38,10 @@ class DatasetCacheSuite extends QueryTest + with AdaptiveSparkPlanHelper { + import testImplicits._ + ++ // A scan of a cached relation, Spark's or Comet's. ++ private def isCacheScan(plan: SparkPlan): Boolean = ++ plan.isInstanceOf[InMemoryTableScanExec] || plan.isInstanceOf[CometInMemoryTableScanExec] ++ + /** + * Asserts that a cached [[Dataset]] will be built using the given number of other cached results. + */ +@@ -43,7 +49,7 @@ class DatasetCacheSuite extends QueryTest + val plan = df.queryExecution.withCachedData + assert(plan.isInstanceOf[InMemoryRelation]) + val internalPlan = plan.asInstanceOf[InMemoryRelation].cacheBuilder.cachedPlan +- assert(find(internalPlan)(_.isInstanceOf[InMemoryTableScanExec]).size ++ assert(find(internalPlan)(isCacheScan).size + == numOfCachesDependedUpon) + } + +@@ -253,7 +259,7 @@ class DatasetCacheSuite extends QueryTest + case i: InMemoryRelation => i.cacheBuilder.cachedPlan + } + assert(df2LimitInnerPlan.isDefined && +- !df2LimitInnerPlan.get.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ !df2LimitInnerPlan.get.exists(isCacheScan)) + } + + test("SPARK-27739 Save stats from optimized plan") { +@@ -286,14 +292,14 @@ class DatasetCacheSuite extends QueryTest + val unionDf = df1.union(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.union(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("positive: union by name") { + val unionDf = df1.unionByName(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.unionByName(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by position") { +@@ -301,7 +307,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.union(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by name") { +@@ -309,7 +315,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.unionByName(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + } + } +@@ -321,11 +327,11 @@ class DatasetCacheSuite extends QueryTest + df1.cache() + // This is exactly the same as df1. + val df2 = spark.range(5).select(struct($"id".as("name", metadata))) +- assert(df2.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(df2.queryExecution.executedPlan.exists(isCacheScan)) + + val metadata2 = Metadata.fromJson("""{"k2": "v2"}""") + // Same with df1 except for the Alias metadata + val df3 = spark.range(5).select(struct($"id".as("name", metadata2))) +- assert(!df3.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!df3.queryExecution.executedPlan.exists(isCacheScan)) + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 81713c777bc..b5f92ed9742 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -1140,10 +1287,18 @@ index ad424b3a7cc..4ece0117a34 100644 protected val baseResourcePath = { // use the same way as `SQLQueryTestSuite` to get the resource path diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -index f294ff81021..02d72be8d29 100644 +index f294ff81021..17d07ab6103 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -@@ -1524,7 +1524,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -38,6 +38,7 @@ import org.apache.spark.sql.catalyst.optimizer.{ConvertToLocalRelation, NestedCo + import org.apache.spark.sql.catalyst.parser.ParseException + import org.apache.spark.sql.catalyst.plans.logical.{LocalLimit, Project, RepartitionByExpression, Sort} + import org.apache.spark.sql.connector.catalog.CatalogManager.SESSION_CATALOG_NAME ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.{CommandResultExec, UnionExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.aggregate._ +@@ -1524,7 +1525,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark checkAnswer(sql("select -0.001"), Row(BigDecimal("-0.001"))) } @@ -1153,7 +1308,7 @@ index f294ff81021..02d72be8d29 100644 AccumulatorSuite.verifyPeakExecutionMemorySet(sparkContext, "external sort") { sql("SELECT * FROM testData2 ORDER BY a ASC, b ASC").collect() } -@@ -1985,8 +1986,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -1985,8 +1987,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark countAcc.add(1) x }) @@ -1170,6 +1325,14 @@ index f294ff81021..02d72be8d29 100644 verifyCallCount( df.selectExpr("testUdf(a + 1) + testUdf(1 + a)", "testUdf(a + 1)"), Row(4, 2), 1) +@@ -3278,6 +3287,7 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark + + val inMemoryTableScan = collect(queryDf.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c + } + assert(inMemoryTableScan.size == 2) + checkAnswer(queryDf, Row(0, 1) :: Row(1, 2) :: Row(2, 3) :: Row(3, 4) :: Row(4, 5) :: Nil) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala index 575a4ae69d1..129d9f27232 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala @@ -1201,6 +1364,18 @@ index 575a4ae69d1..129d9f27232 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +index c1c041509c3..580d320394d 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +@@ -222,6 +222,7 @@ class SparkSessionExtensionSuite extends SparkFunSuite with SQLHelper with Adapt + df.select("i").filter($"i" > 1).cache() + assert(find(df.filter($"i" > 1).select("i").queryExecution.executedPlan) { + case _: org.apache.spark.sql.execution.columnar.InMemoryTableScanExec => true ++ case _: org.apache.spark.sql.comet.CometInMemoryTableScanExec => true + case _ => false + }.isDefined) + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala index 5ba69c8f9d9..ac1256afe88 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala @@ -1216,6 +1391,20 @@ index 5ba69c8f9d9..ac1256afe88 100644 sc = new SparkContext("local[2]", "test") val session = classic.SparkSession.builder().sparkContext(sc).getOrCreate() import session.implicits._ +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +index 5222d5ce266..67f663c2764 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +@@ -829,7 +829,8 @@ class StatisticsCollectionSuite extends StatisticsCollectionTestBase with Shared + } + } + +- test("SPARK-33687: analyze all tables in a specific database") { ++ test("SPARK-33687: analyze all tables in a specific database", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + withTempDatabase { database => + spark.catalog.setCurrentDatabase(database) + withTempDir { dir => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala index 0df7f806272..9cdfe8b8f46 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala @@ -1415,7 +1604,7 @@ index a40e34d94d0..abc1f035d15 100644 } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala -index 11e9547dfc5..ba340c4ebcf 100644 +index 11e9547dfc5..327e0b1bc5b 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala @@ -24,6 +24,8 @@ import org.apache.spark.sql.{AnalysisException, Row} @@ -1423,7 +1612,7 @@ index 11e9547dfc5..ba340c4ebcf 100644 import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.util.CollationFactory +import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometHashJoinExec, CometSortMergeJoinExec} -+import org.apache.spark.sql.comet.CometHashAggregateExec ++import org.apache.spark.sql.comet.{CometHashAggregateExec, CometInMemoryTableScanExec} import org.apache.spark.sql.connector.{DatasourceV2SQLBase, FakeV2ProviderWithCustomSchema} import org.apache.spark.sql.connector.catalog.{Identifier, InMemoryTable} import org.apache.spark.sql.connector.catalog.CatalogV2Implicits.CatalogHelper @@ -1482,6 +1671,15 @@ index 11e9547dfc5..ba340c4ebcf 100644 }.head.isInstanceOf[ArrayTransform]) } } +@@ -1897,7 +1909,7 @@ class CollationSuite extends DatasourceV2SQLBase with AdaptiveSparkPlanHelper { + // Checks in-memory fetching code path. + val all = sql("SELECT col FROM tbl") + assert(all.queryExecution.executedPlan.collectFirst { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + }.nonEmpty) + checkAnswer(all, Row("a")) + // Checks column stats code path. diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala index 3eeed2e4175..9f21d547c1c 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala @@ -2208,7 +2406,7 @@ index a3cfdc5a240..3793b6191bf 100644 }) checkAnswer(distinctWithId, Seq(Row(1, 0), Row(1, 0))) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala -index fb8fab6a80f..403eb411920 100644 +index fb8fab6a80f..8a8dcee18db 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala @@ -28,12 +28,14 @@ import org.apache.spark.SparkException @@ -2693,7 +2891,27 @@ index fb8fab6a80f..403eb411920 100644 }.isEmpty) assert(collect(initialExecutedPlan) { case i: InMemoryTableScanLike => i -@@ -3048,7 +3101,8 @@ class AdaptiveQueryExecSuite +@@ -3003,7 +3056,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage") { ++ test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage", ++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/6454")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1") { + val cached = Seq(1).toDF("c").cache() +@@ -3033,7 +3087,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-43376: Improve reuse subquery with table cache") { ++ test("SPARK-43376: Improve reuse subquery with table cache", ++ IgnoreComet("Comet's cache scan does not plan the subqueries in its pruning predicates")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true") { + withTable("t1", "t2") { + withCache("t1") { +@@ -3048,7 +3103,8 @@ class AdaptiveQueryExecSuite } } @@ -2703,7 +2921,7 @@ index fb8fab6a80f..403eb411920 100644 val emptyDf = spark.range(1).where("false") val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df1")) val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df2")) -@@ -3138,7 +3192,8 @@ class AdaptiveQueryExecSuite +@@ -3138,7 +3194,8 @@ class AdaptiveQueryExecSuite val plan = df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec] assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec]) @@ -2713,6 +2931,152 @@ index fb8fab6a80f..403eb411920 100644 plan.inputPlan.output.zip(plan.finalPhysicalPlan.output).foreach { case (o1, o2) => assert(o1.semanticEquals(o2), "Different output column order after AQE optimization") } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +index 4f07d3d1c03..c0e1c829bea 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +@@ -22,12 +22,13 @@ import java.sql.{Date, Timestamp} + import java.util.concurrent.atomic.AtomicInteger + + import org.apache.spark.rdd.RDD +-import org.apache.spark.sql.{QueryTest, Row} ++import org.apache.spark.sql.{IgnoreComet, QueryTest, Row} + import org.apache.spark.sql.catalyst.InternalRow + import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSet, In} + import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning + import org.apache.spark.sql.classic.DataFrame + import org.apache.spark.sql.columnar.CachedBatch ++import org.apache.spark.sql.comet.{CometFilterExec, CometInMemoryTableScanExec} + import org.apache.spark.sql.execution.{FilterExec, InputAdapter, WholeStageCodegenExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions._ +@@ -510,11 +511,17 @@ class InMemoryColumnarQuerySuite extends QueryTest + val planBeforeFilter = collect(df2.queryExecution.executedPlan) { + case f: FilterExec => f.child + case WholeStageCodegenExec(FilterExec(_, i: InputAdapter)) => i.child ++ case f: CometFilterExec => f.child + } +- assert(planBeforeFilter.head.isInstanceOf[InMemoryTableScanExec]) +- + val execPlan = planBeforeFilter.head +- assert(execPlan.executeCollectPublic().length == 0) ++ execPlan match { ++ // Comet's cache scan is columnar only, so count the rows of the batches it emits. ++ case c: CometInMemoryTableScanExec => ++ assert(c.executeColumnar().map(_.numRows().toLong).collect().sum == 0) ++ case _ => ++ assert(execPlan.isInstanceOf[InMemoryTableScanExec]) ++ assert(execPlan.executeCollectPublic().length == 0) ++ } + } + + test("SPARK-25727 - otherCopyArgs in InMemoryRelation does not include outputOrdering") { +@@ -523,7 +530,8 @@ class InMemoryColumnarQuerySuite extends QueryTest + assert(json.contains("outputOrdering")) + } + +- test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached") { ++ test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + Seq("orc", "").foreach { useV1SourceReaderList => + // This test case depends on the size of ORC in statistics. + withSQLConf( +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +index 2c73622739a..5d0efeb263a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +@@ -18,6 +18,7 @@ + package org.apache.spark.sql.execution.columnar + + import org.apache.spark.SparkFunSuite ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions.expr +@@ -37,11 +38,15 @@ class InMemoryRelationSuite extends SparkFunSuite + test("SPARK-47177: Cached SQL plan do not display final AQE plan in explain string") { + def findIMRInnerChild(p: SparkPlan): SparkPlan = { + val tableCache = find(p) { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + case _ => false + } + assert(tableCache.isDefined) +- tableCache.get.asInstanceOf[InMemoryTableScanExec].relation.innerChildren.head ++ val scan = tableCache.get match { ++ case c: CometInMemoryTableScanExec => c.originalPlan ++ case s => s.asInstanceOf[InMemoryTableScanExec] ++ } ++ scan.relation.innerChildren.head + } + + val d1 = spark.range(1).withColumn("key", expr("id % 100")) +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +index 88ff51d0ff4..7d00fca3b0a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +@@ -17,6 +17,7 @@ + + package org.apache.spark.sql.execution.columnar + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.internal.SQLConf + import org.apache.spark.sql.test.SharedSparkSession +@@ -181,11 +182,16 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + val result = df.collect().map(_(0)).toArray + assert(result.length === 1) + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- assert(readPartitions === 5) +- assert(readBatches === 10) ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan has none of these test-only accumulators, so there is nothing to count. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ assert(scans.head.readPartitions.value === 5) ++ assert(scans.head.readBatches.value === 10) ++ } + } + + def checkBatchPruning( +@@ -202,14 +208,23 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + df.collect().map(_(0)).toArray + } + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- +- assert(readBatches === expectedReadBatches, s"Wrong number of read batches: $queryExecution") +- assert( +- readPartitions === expectedReadPartitions, +- s"Wrong number of read partitions: $queryExecution") ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan prunes on the same statistics, but has none of these test-only ++ // accumulators to read, so only the answer above is checked for it. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ val readPartitions = scans.head.readPartitions.value ++ val readBatches = scans.head.readBatches.value ++ assert( ++ readBatches === expectedReadBatches, ++ s"Wrong number of read batches: $queryExecution") ++ assert( ++ readPartitions === expectedReadPartitions, ++ s"Wrong number of read partitions: $queryExecution") ++ } + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala index 269990d7d14..140ee4112b1 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala @@ -3856,10 +4220,10 @@ index f0f3f94b811..b7d18771314 100644 spark.internalCreateDataFrame(withoutFilters.execute(), schema) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -index 720b13b812e..e3ac2cebc6e 100644 +index 720b13b812e..388dad6b4f3 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -@@ -98,6 +98,20 @@ trait SharedSparkSessionBase +@@ -98,6 +98,25 @@ trait SharedSparkSessionBase // this rule may potentially block testing of other optimization rules such as // ConstantPropagation etc. .set(SQLConf.OPTIMIZER_EXCLUDED_RULES.key, ConvertToLocalRelation.ruleName) @@ -3876,6 +4240,11 @@ index 720b13b812e..e3ac2cebc6e 100644 + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") + ++ // CometDriverPlugin installs Comet's cache serializer when ++ // spark.comet.exec.inMemoryCache.enabled is on, as it is by default. These sessions do not ++ // load the plugin, so install it here. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } conf.set( StaticSQLConf.WAREHOUSE_PATH, @@ -3998,10 +4367,10 @@ index b67370f6eb9..746b3974b29 100644 override def beforeEach(): Unit = { super.beforeEach() diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -index a394d0b7393..7056c74759c 100644 +index a394d0b7393..9c626b05ab2 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -@@ -53,24 +53,34 @@ object TestHive +@@ -53,24 +53,38 @@ object TestHive new SparkContext( System.getProperty("spark.sql.test.master", "local[1]"), "TestSQLContext", @@ -4046,6 +4415,10 @@ index a394d0b7393..7056c74759c 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ ++ // As in SharedSparkSession: what CometDriverPlugin would install. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } + + conf diff --git a/dev/diffs/4.1.3.diff b/dev/diffs/4.1.3.diff index a8d9383ced1..371225e9fed 100644 --- a/dev/diffs/4.1.3.diff +++ b/dev/diffs/4.1.3.diff @@ -358,11 +358,41 @@ index 21a3ce1e122..f4762ab98f0 100644 SET spark.sql.ansi.enabled = false; -- In COMPENSATION views get invalidated if the type can't cast +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala +index 26d8f750f6e..c888f8e0844 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala +@@ -32,9 +32,16 @@ class CacheTableInKryoSuite extends QueryTest + with SharedSparkSession { + + override def sparkConf: SparkConf = { +- super.sparkConf ++ val conf = super.sparkConf + .set("spark.kryo.registrationRequired", "true") + .set("spark.serializer", "org.apache.spark.serializer.KryoSerializer") ++ // Comet's cache format needs its classes registered too, which is what Comet asks of any ++ // application that sets registrationRequired. ++ if (isCometEnabled) { ++ conf.set("spark.kryo.registrator", "org.apache.comet.CometKryoRegistrator") ++ } else { ++ conf ++ } + } + + test("SPARK-51777: sql.columnar.* classes registered in KryoSerializer") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index 0d807aeae4d..6d7744e771b 100644 +index 0d807aeae4d..354a0757d31 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -@@ -49,7 +49,7 @@ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEProp +@@ -37,6 +37,7 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException + import org.apache.spark.sql.catalyst.expressions.SubqueryExpression + import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, Join, JoinStrategyHint, SHUFFLE_HASH} + import org.apache.spark.sql.catalyst.util.DateTimeConstants ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.connector.catalog.BasicInMemoryTableCatalog + import org.apache.spark.sql.connector.catalog.CatalogPlugin + import org.apache.spark.sql.connector.catalog.CatalogV2Implicits.CatalogHelper +@@ -49,7 +50,7 @@ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEProp import org.apache.spark.sql.execution.columnar._ import org.apache.spark.sql.execution.command.CommandUtils import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation @@ -371,7 +401,27 @@ index 0d807aeae4d..6d7744e771b 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -534,7 +534,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -128,6 +129,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => + getNumInMemoryTablesRecursively(relation.cachedPlan) + + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ case cometScan: CometInMemoryTableScanExec => ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + ++ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 + case p => + getNumInMemoryTablesInSubquery(p) + }.sum +@@ -408,7 +412,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + assert(isExpectStorageLevel(rddId, Disk)) + } + +- test("InMemoryRelation statistics") { ++ test("InMemoryRelation statistics", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + sql("CACHE TABLE testData") + spark.table("testData").queryExecution.withCachedData.collect { + case cached: InMemoryRelation => +@@ -534,7 +539,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils df.collect() } assert( @@ -381,7 +431,17 @@ index 0d807aeae4d..6d7744e771b 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1673,9 +1674,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1595,7 +1601,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + } + } + +- test("SPARK-36120: Support cache/uncache table with TimestampNTZ type") { ++ test("SPARK-36120: Support cache/uncache table with TimestampNTZ type", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + val tableName = "ntzCache" + withTable(tableName) { + sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") +@@ -1673,9 +1680,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils _.nodeName.contains("TableCacheQueryStage")) val aqeNode = findNodeInSparkPlanInfo(inMemoryScanNode.get, _.nodeName.contains("AdaptiveSparkPlan")) @@ -403,6 +463,30 @@ index 0d807aeae4d..6d7744e771b 100644 } withTempView("t0", "t1", "t2") { +@@ -1764,6 +1780,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + val cached = spark.table("t") + val tableCache = collect(cached.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + if (expected == StorageLevel.NONE) { + assert(tableCache.isEmpty) +@@ -2630,6 +2647,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + + val inMemoryTableScan = collect(df.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + assert(inMemoryTableScan.size == 1) + checkAnswer(df, Row(5) :: Nil) +@@ -2657,6 +2675,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + + val subqueryInMemoryTableScan = collect(cteInSubquery.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + assert(subqueryInMemoryTableScan.size == 1) + checkAnswer(cteInSubquery, Row(1) :: Nil) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala index bfe15b33768..13aeb3f6610 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala @@ -460,7 +544,7 @@ index ed182322aec..1ae6afa686a 100644 spark.range(100).write.saveAsTable(s"$dbName.$table2Name") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala -index 93ff7becaec..87537a25b3b 100644 +index 93ff7becaec..27c366a9e4f 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala @@ -23,10 +23,11 @@ import java.util.Locale @@ -476,7 +560,17 @@ index 93ff7becaec..87537a25b3b 100644 import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.{ExamplePoint, ExamplePointUDT, SharedSparkSession, SQLTestData} -@@ -1519,11 +1520,12 @@ class DataFrameSetOperationsSuite extends QueryTest +@@ -1476,7 +1477,8 @@ class DataFrameSetOperationsSuite extends QueryTest + Row(Row(Seq(Seq(Row(null, "ba"))))) :: Nil) + } + +- test("SPARK-37371: UnionExec should support columnar if all children support columnar") { ++ test("SPARK-37371: UnionExec should support columnar if all children support columnar", ++ IgnoreComet("Comet replaces the cache scans and the union with its own operators")) { + def checkIfColumnar( + plan: SparkPlan, + targetPlan: (SparkPlan) => Boolean, +@@ -1519,11 +1521,12 @@ class DataFrameSetOperationsSuite extends QueryTest val union = df1.repartition($"a").union(df2.repartition($"a")) val unionExec = union.queryExecution.executedPlan.collect { case u: UnionExec => u @@ -490,7 +584,7 @@ index 93ff7becaec..87537a25b3b 100644 } assert(shuffle.size == 1) -@@ -1554,11 +1556,12 @@ class DataFrameSetOperationsSuite extends QueryTest +@@ -1554,11 +1557,12 @@ class DataFrameSetOperationsSuite extends QueryTest val union = df1.repartition($"a").union(df2.repartition($"d")) val unionExec = union.queryExecution.executedPlan.collect { case u: UnionExec => u @@ -504,7 +598,7 @@ index 93ff7becaec..87537a25b3b 100644 } assert(shuffle.size == 1) -@@ -1573,10 +1576,10 @@ class DataFrameSetOperationsSuite extends QueryTest +@@ -1573,10 +1577,10 @@ class DataFrameSetOperationsSuite extends QueryTest // Avoid unnecessary shuffle if union output partitioning is enabled val shuffledUnion = union.repartition($"a") val shuffleNumBefore = union.queryExecution.executedPlan.collect { @@ -517,7 +611,7 @@ index 93ff7becaec..87537a25b3b 100644 } if (enabled) { -@@ -1605,6 +1608,7 @@ class DataFrameSetOperationsSuite extends QueryTest +@@ -1605,6 +1609,7 @@ class DataFrameSetOperationsSuite extends QueryTest val union = df1.repartitionByRange($"a").union(df2.repartitionByRange($"d")) val unionExec = union.queryExecution.executedPlan.collect { case u: UnionExec => u @@ -658,6 +752,99 @@ index 4a070becfa6..61d515d127e 100644 val df = Seq((1, "1"), (2, "2"), (1, "3"), (2, "4")).toDF("key", "value") val window = Window.partitionBy($"key").orderBy($"value") +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +index 627811eaecf..7269aeab949 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +@@ -22,7 +22,8 @@ import java.time.LocalTime + import org.scalatest.concurrent.TimeLimits + import org.scalatest.time.SpanSugar._ + +-import org.apache.spark.sql.execution.ColumnarToRowExec ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec ++import org.apache.spark.sql.execution.{ColumnarToRowExec, SparkPlan} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.columnar.{InMemoryRelation, InMemoryTableScanExec} + import org.apache.spark.sql.functions._ +@@ -39,6 +40,10 @@ class DatasetCacheSuite extends QueryTest + with AdaptiveSparkPlanHelper { + import testImplicits._ + ++ // A scan of a cached relation, Spark's or Comet's. ++ private def isCacheScan(plan: SparkPlan): Boolean = ++ plan.isInstanceOf[InMemoryTableScanExec] || plan.isInstanceOf[CometInMemoryTableScanExec] ++ + /** + * Asserts that a cached [[Dataset]] will be built using the given number of other cached results. + */ +@@ -46,7 +51,7 @@ class DatasetCacheSuite extends QueryTest + val plan = df.queryExecution.withCachedData + assert(plan.isInstanceOf[InMemoryRelation]) + val internalPlan = plan.asInstanceOf[InMemoryRelation].cacheBuilder.cachedPlan +- assert(find(internalPlan)(_.isInstanceOf[InMemoryTableScanExec]).size ++ assert(find(internalPlan)(isCacheScan).size + == numOfCachesDependedUpon) + } + +@@ -256,7 +261,7 @@ class DatasetCacheSuite extends QueryTest + case i: InMemoryRelation => i.cacheBuilder.cachedPlan + } + assert(df2LimitInnerPlan.isDefined && +- !df2LimitInnerPlan.get.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ !df2LimitInnerPlan.get.exists(isCacheScan)) + } + + test("SPARK-27739 Save stats from optimized plan") { +@@ -289,14 +294,14 @@ class DatasetCacheSuite extends QueryTest + val unionDf = df1.union(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.union(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("positive: union by name") { + val unionDf = df1.unionByName(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.unionByName(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by position") { +@@ -304,7 +309,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.union(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by name") { +@@ -312,7 +317,7 @@ class DatasetCacheSuite extends QueryTest + unionDf.cache() + val finalDf = unionDf.unionByName(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + } + } +@@ -324,12 +329,12 @@ class DatasetCacheSuite extends QueryTest + df1.cache() + // This is exactly the same as df1. + val df2 = spark.range(5).select(struct($"id".as("name", metadata))) +- assert(df2.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(df2.queryExecution.executedPlan.exists(isCacheScan)) + + val metadata2 = Metadata.fromJson("""{"k2": "v2"}""") + // Same with df1 except for the Alias metadata + val df3 = spark.range(5).select(struct($"id".as("name", metadata2))) +- assert(!df3.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!df3.queryExecution.executedPlan.exists(isCacheScan)) + } + + test("SPARK-53418: Handle TimeType in ColumnAccessor") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 6df8d66ee7f..35e270c7241 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -1247,10 +1434,18 @@ index cb9d0909554..084d6515e8b 100644 } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -index 74cdee49e55..f7452c9abb7 100644 +index 74cdee49e55..6a544644c32 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -@@ -1521,7 +1521,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -36,6 +36,7 @@ import org.apache.spark.sql.catalyst.optimizer.{ConvertToLocalRelation, NestedCo + import org.apache.spark.sql.catalyst.parser.ParseException + import org.apache.spark.sql.catalyst.plans.logical.{LocalLimit, Project, RepartitionByExpression, Sort} + import org.apache.spark.sql.connector.catalog.CatalogManager.SESSION_CATALOG_NAME ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.{CommandResultExec, OneRowRelationExec, UnionExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.aggregate._ +@@ -1521,7 +1522,8 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark checkAnswer(sql("select -0.001"), Row(BigDecimal("-0.001"))) } @@ -1260,7 +1455,7 @@ index 74cdee49e55..f7452c9abb7 100644 AccumulatorSuite.verifyPeakExecutionMemorySet(sparkContext, "external sort") { sql("SELECT * FROM testData2 ORDER BY a ASC, b ASC").collect() } -@@ -1982,8 +1983,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark +@@ -1982,8 +1984,15 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark countAcc.add(1) x }) @@ -1277,6 +1472,14 @@ index 74cdee49e55..f7452c9abb7 100644 verifyCallCount( df.selectExpr("testUdf(a + 1) + testUdf(1 + a)", "testUdf(a + 1)"), Row(4, 2), 1) +@@ -3275,6 +3284,7 @@ class SQLQuerySuite extends QueryTest with SharedSparkSession with AdaptiveSpark + + val inMemoryTableScan = collect(queryDf.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c + } + assert(inMemoryTableScan.size == 2) + checkAnswer(queryDf, Row(0, 1) :: Row(1, 2) :: Row(2, 3) :: Row(3, 4) :: Row(4, 5) :: Nil) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala index 23f0144dcec..40d536bb23a 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala @@ -1321,6 +1524,18 @@ index 23f0144dcec..40d536bb23a 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +index 66826a9ca76..efeac3e4529 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +@@ -239,6 +239,7 @@ class SparkSessionExtensionSuite extends SparkFunSuite with SQLHelper with Adapt + df.select("i").filter($"i" > 1).cache() + assert(find(df.filter($"i" > 1).select("i").queryExecution.executedPlan) { + case _: org.apache.spark.sql.execution.columnar.InMemoryTableScanExec => true ++ case _: org.apache.spark.sql.comet.CometInMemoryTableScanExec => true + case _ => false + }.isDefined) + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala index d7b2511eac2..d5f5b940b94 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala @@ -1336,6 +1551,20 @@ index d7b2511eac2..d5f5b940b94 100644 sc = new SparkContext("local[2]", "test") val session = classic.SparkSession.builder().sparkContext(sc).getOrCreate() import session.implicits._ +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +index 5222d5ce266..67f663c2764 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +@@ -829,7 +829,8 @@ class StatisticsCollectionSuite extends StatisticsCollectionTestBase with Shared + } + } + +- test("SPARK-33687: analyze all tables in a specific database") { ++ test("SPARK-33687: analyze all tables in a specific database", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + withTempDatabase { database => + spark.catalog.setCurrentDatabase(database) + withTempDir { dir => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala index 7bfc8cf4fa6..4bd387801db 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala @@ -1535,7 +1764,7 @@ index 8a0e2c29653..d276a51cbc6 100644 } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala -index 8f7a68bcbe6..88dbe1793c9 100644 +index 8f7a68bcbe6..c09c5d74309 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala @@ -26,6 +26,8 @@ import org.apache.spark.sql.{AnalysisException, Row} @@ -1543,7 +1772,7 @@ index 8f7a68bcbe6..88dbe1793c9 100644 import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.util.CollationFactory +import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometHashJoinExec, CometSortMergeJoinExec} -+import org.apache.spark.sql.comet.CometHashAggregateExec ++import org.apache.spark.sql.comet.{CometHashAggregateExec, CometInMemoryTableScanExec} import org.apache.spark.sql.connector.{DatasourceV2SQLBase, FakeV2ProviderWithCustomSchema} import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Identifier, InMemoryTable} import org.apache.spark.sql.connector.catalog.CatalogV2Implicits.CatalogHelper @@ -1602,6 +1831,15 @@ index 8f7a68bcbe6..88dbe1793c9 100644 }.head.isInstanceOf[ArrayTransform]) } } +@@ -1948,7 +1960,7 @@ class CollationSuite extends DatasourceV2SQLBase with AdaptiveSparkPlanHelper { + // Checks in-memory fetching code path. + val all = sql("SELECT col FROM tbl") + assert(all.queryExecution.executedPlan.collectFirst { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + }.nonEmpty) + checkAnswer(all, Row("a")) + // Checks column stats code path. diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala index a09b7e0827c..ffc29f764bc 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala @@ -2328,7 +2566,7 @@ index a3cfdc5a240..3793b6191bf 100644 }) checkAnswer(distinctWithId, Seq(Row(1, 0), Row(1, 0))) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala -index 188a28ff1c0..8fdccf31749 100644 +index 188a28ff1c0..cfc84dcba01 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala @@ -27,12 +27,14 @@ import org.apache.spark.SparkException @@ -2813,7 +3051,27 @@ index 188a28ff1c0..8fdccf31749 100644 }.isEmpty) assert(collect(initialExecutedPlan) { case i: InMemoryTableScanLike => i -@@ -3229,7 +3282,8 @@ class AdaptiveQueryExecSuite +@@ -3184,7 +3237,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage") { ++ test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage", ++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/6454")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1") { + val cached = Seq(1).toDF("c").cache() +@@ -3214,7 +3268,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-43376: Improve reuse subquery with table cache") { ++ test("SPARK-43376: Improve reuse subquery with table cache", ++ IgnoreComet("Comet's cache scan does not plan the subqueries in its pruning predicates")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true") { + withTable("t1", "t2") { + withCache("t1") { +@@ -3229,7 +3284,8 @@ class AdaptiveQueryExecSuite } } @@ -2823,7 +3081,7 @@ index 188a28ff1c0..8fdccf31749 100644 val emptyDf = spark.range(1).where("false") val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df1")) val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df2")) -@@ -3319,7 +3373,8 @@ class AdaptiveQueryExecSuite +@@ -3319,7 +3375,8 @@ class AdaptiveQueryExecSuite val plan = df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec] assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec]) @@ -2862,6 +3120,152 @@ index 47b935a2880..3fdeab3113c 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +index 4f07d3d1c03..c0e1c829bea 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +@@ -22,12 +22,13 @@ import java.sql.{Date, Timestamp} + import java.util.concurrent.atomic.AtomicInteger + + import org.apache.spark.rdd.RDD +-import org.apache.spark.sql.{QueryTest, Row} ++import org.apache.spark.sql.{IgnoreComet, QueryTest, Row} + import org.apache.spark.sql.catalyst.InternalRow + import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSet, In} + import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning + import org.apache.spark.sql.classic.DataFrame + import org.apache.spark.sql.columnar.CachedBatch ++import org.apache.spark.sql.comet.{CometFilterExec, CometInMemoryTableScanExec} + import org.apache.spark.sql.execution.{FilterExec, InputAdapter, WholeStageCodegenExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions._ +@@ -510,11 +511,17 @@ class InMemoryColumnarQuerySuite extends QueryTest + val planBeforeFilter = collect(df2.queryExecution.executedPlan) { + case f: FilterExec => f.child + case WholeStageCodegenExec(FilterExec(_, i: InputAdapter)) => i.child ++ case f: CometFilterExec => f.child + } +- assert(planBeforeFilter.head.isInstanceOf[InMemoryTableScanExec]) +- + val execPlan = planBeforeFilter.head +- assert(execPlan.executeCollectPublic().length == 0) ++ execPlan match { ++ // Comet's cache scan is columnar only, so count the rows of the batches it emits. ++ case c: CometInMemoryTableScanExec => ++ assert(c.executeColumnar().map(_.numRows().toLong).collect().sum == 0) ++ case _ => ++ assert(execPlan.isInstanceOf[InMemoryTableScanExec]) ++ assert(execPlan.executeCollectPublic().length == 0) ++ } + } + + test("SPARK-25727 - otherCopyArgs in InMemoryRelation does not include outputOrdering") { +@@ -523,7 +530,8 @@ class InMemoryColumnarQuerySuite extends QueryTest + assert(json.contains("outputOrdering")) + } + +- test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached") { ++ test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + Seq("orc", "").foreach { useV1SourceReaderList => + // This test case depends on the size of ORC in statistics. + withSQLConf( +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +index 2c73622739a..5d0efeb263a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +@@ -18,6 +18,7 @@ + package org.apache.spark.sql.execution.columnar + + import org.apache.spark.SparkFunSuite ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions.expr +@@ -37,11 +38,15 @@ class InMemoryRelationSuite extends SparkFunSuite + test("SPARK-47177: Cached SQL plan do not display final AQE plan in explain string") { + def findIMRInnerChild(p: SparkPlan): SparkPlan = { + val tableCache = find(p) { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + case _ => false + } + assert(tableCache.isDefined) +- tableCache.get.asInstanceOf[InMemoryTableScanExec].relation.innerChildren.head ++ val scan = tableCache.get match { ++ case c: CometInMemoryTableScanExec => c.originalPlan ++ case s => s.asInstanceOf[InMemoryTableScanExec] ++ } ++ scan.relation.innerChildren.head + } + + val d1 = spark.range(1).withColumn("key", expr("id % 100")) +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +index 88ff51d0ff4..7d00fca3b0a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +@@ -17,6 +17,7 @@ + + package org.apache.spark.sql.execution.columnar + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.internal.SQLConf + import org.apache.spark.sql.test.SharedSparkSession +@@ -181,11 +182,16 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + val result = df.collect().map(_(0)).toArray + assert(result.length === 1) + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- assert(readPartitions === 5) +- assert(readBatches === 10) ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan has none of these test-only accumulators, so there is nothing to count. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ assert(scans.head.readPartitions.value === 5) ++ assert(scans.head.readBatches.value === 10) ++ } + } + + def checkBatchPruning( +@@ -202,14 +208,23 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + df.collect().map(_(0)).toArray + } + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- +- assert(readBatches === expectedReadBatches, s"Wrong number of read batches: $queryExecution") +- assert( +- readPartitions === expectedReadPartitions, +- s"Wrong number of read partitions: $queryExecution") ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan prunes on the same statistics, but has none of these test-only ++ // accumulators to read, so only the answer above is checked for it. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ val readPartitions = scans.head.readPartitions.value ++ val readBatches = scans.head.readBatches.value ++ assert( ++ readBatches === expectedReadBatches, ++ s"Wrong number of read batches: $queryExecution") ++ assert( ++ readPartitions === expectedReadPartitions, ++ s"Wrong number of read partitions: $queryExecution") ++ } + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala index 269990d7d14..140ee4112b1 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala @@ -4155,10 +4559,10 @@ index f0f3f94b811..b7d18771314 100644 spark.internalCreateDataFrame(withoutFilters.execute(), schema) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -index 720b13b812e..e3ac2cebc6e 100644 +index 720b13b812e..388dad6b4f3 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -@@ -98,6 +98,20 @@ trait SharedSparkSessionBase +@@ -98,6 +98,25 @@ trait SharedSparkSessionBase // this rule may potentially block testing of other optimization rules such as // ConstantPropagation etc. .set(SQLConf.OPTIMIZER_EXCLUDED_RULES.key, ConvertToLocalRelation.ruleName) @@ -4175,6 +4579,11 @@ index 720b13b812e..e3ac2cebc6e 100644 + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") + ++ // CometDriverPlugin installs Comet's cache serializer when ++ // spark.comet.exec.inMemoryCache.enabled is on, as it is by default. These sessions do not ++ // load the plugin, so install it here. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } conf.set( StaticSQLConf.WAREHOUSE_PATH, @@ -4297,10 +4706,10 @@ index b67370f6eb9..746b3974b29 100644 override def beforeEach(): Unit = { super.beforeEach() diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -index a394d0b7393..7056c74759c 100644 +index a394d0b7393..9c626b05ab2 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -@@ -53,24 +53,34 @@ object TestHive +@@ -53,24 +53,38 @@ object TestHive new SparkContext( System.getProperty("spark.sql.test.master", "local[1]"), "TestSQLContext", @@ -4345,6 +4754,10 @@ index a394d0b7393..7056c74759c 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ ++ // As in SharedSparkSession: what CometDriverPlugin would install. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } + + conf diff --git a/dev/diffs/4.2.0.diff b/dev/diffs/4.2.0.diff index 60ad822d6d2..dfb411ba504 100644 --- a/dev/diffs/4.2.0.diff +++ b/dev/diffs/4.2.0.diff @@ -376,11 +376,41 @@ index 21a3ce1e122..f4762ab98f0 100644 SET spark.sql.ansi.enabled = false; -- In COMPENSATION views get invalidated if the type can't cast +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala +index 72a2da16054..54146036f8a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/CacheTableInKryoSuite.scala +@@ -30,9 +30,16 @@ import org.apache.spark.storage.StorageLevel + class CacheTableInKryoSuite extends SharedSparkSession { + + override def sparkConf: SparkConf = { +- super.sparkConf ++ val conf = super.sparkConf + .set("spark.kryo.registrationRequired", "true") + .set("spark.serializer", "org.apache.spark.serializer.KryoSerializer") ++ // Comet's cache format needs its classes registered too, which is what Comet asks of any ++ // application that sets registrationRequired. ++ if (isCometEnabled) { ++ conf.set("spark.kryo.registrator", "org.apache.comet.CometKryoRegistrator") ++ } else { ++ conf ++ } + } + + test("SPARK-51777: sql.columnar.* classes registered in KryoSerializer") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index 085dbcd8046..3090d321b6c 100644 +index 085dbcd8046..280d020ea33 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -@@ -49,7 +49,7 @@ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEProp +@@ -37,6 +37,7 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException + import org.apache.spark.sql.catalyst.expressions.SubqueryExpression + import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, Join, JoinStrategyHint, SHUFFLE_HASH} + import org.apache.spark.sql.catalyst.util.DateTimeConstants ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.connector.catalog.BasicInMemoryTableCatalog + import org.apache.spark.sql.connector.catalog.CatalogPlugin + import org.apache.spark.sql.connector.catalog.CatalogV2Implicits.CatalogHelper +@@ -49,7 +50,7 @@ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanHelper, AQEProp import org.apache.spark.sql.execution.columnar._ import org.apache.spark.sql.execution.command.CommandUtils import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation @@ -389,7 +419,27 @@ index 085dbcd8046..3090d321b6c 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -564,7 +564,8 @@ class CachedTableSuite extends SharedSparkSession +@@ -127,6 +128,9 @@ class CachedTableSuite extends SharedSparkSession + case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => + getNumInMemoryTablesRecursively(relation.cachedPlan) + + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ case cometScan: CometInMemoryTableScanExec => ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + ++ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 + case p => + getNumInMemoryTablesInSubquery(p) + }.sum +@@ -407,7 +411,8 @@ class CachedTableSuite extends SharedSparkSession + assert(isExpectStorageLevel(rddId, Disk)) + } + +- test("InMemoryRelation statistics") { ++ test("InMemoryRelation statistics", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + sql("CACHE TABLE testData") + spark.table("testData").queryExecution.withCachedData.collect { + case cached: InMemoryRelation => +@@ -564,7 +569,8 @@ class CachedTableSuite extends SharedSparkSession df.collect() } assert( @@ -399,7 +449,17 @@ index 085dbcd8046..3090d321b6c 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1703,9 +1704,18 @@ class CachedTableSuite extends SharedSparkSession +@@ -1625,7 +1631,8 @@ class CachedTableSuite extends SharedSparkSession + } + } + +- test("SPARK-36120: Support cache/uncache table with TimestampNTZ type") { ++ test("SPARK-36120: Support cache/uncache table with TimestampNTZ type", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + val tableName = "ntzCache" + withTable(tableName) { + sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") +@@ -1703,9 +1710,18 @@ class CachedTableSuite extends SharedSparkSession _.nodeName.contains("TableCacheQueryStage")) val aqeNode = findNodeInSparkPlanInfo(inMemoryScanNode.get, _.nodeName.contains("AdaptiveSparkPlan")) @@ -421,6 +481,30 @@ index 085dbcd8046..3090d321b6c 100644 } withTempView("t0", "t1", "t2") { +@@ -1794,6 +1810,7 @@ class CachedTableSuite extends SharedSparkSession + val cached = spark.table("t") + val tableCache = collect(cached.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + if (expected == StorageLevel.NONE) { + assert(tableCache.isEmpty) +@@ -2660,6 +2677,7 @@ class CachedTableSuite extends SharedSparkSession + + val inMemoryTableScan = collect(df.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + assert(inMemoryTableScan.size == 1) + checkAnswer(df, Row(5) :: Nil) +@@ -2687,6 +2705,7 @@ class CachedTableSuite extends SharedSparkSession + + val subqueryInMemoryTableScan = collect(cteInSubquery.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c.originalPlan + } + assert(subqueryInMemoryTableScan.size == 1) + checkAnswer(cteInSubquery, Row(1) :: Nil) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala index 5b8154d2900..f01366b66bc 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala @@ -489,7 +573,7 @@ index 9733d51a91c..395a108abc8 100644 spark.range(100).write.saveAsTable(s"$dbName.$table2Name") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala -index d838ba4c234..cb0573d56d0 100644 +index d838ba4c234..4661d627a55 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameSetOperationsSuite.scala @@ -23,10 +23,11 @@ import java.util.Locale @@ -505,7 +589,17 @@ index d838ba4c234..cb0573d56d0 100644 import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.{ExamplePoint, ExamplePointUDT, SharedSparkSession, SQLTestData} -@@ -1518,11 +1519,12 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP +@@ -1475,7 +1476,8 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP + Row(Row(Seq(Seq(Row(null, "ba"))))) :: Nil) + } + +- test("SPARK-37371: UnionExec should support columnar if all children support columnar") { ++ test("SPARK-37371: UnionExec should support columnar if all children support columnar", ++ IgnoreComet("Comet replaces the cache scans and the union with its own operators")) { + def checkIfColumnar( + plan: SparkPlan, + targetPlan: (SparkPlan) => Boolean, +@@ -1518,11 +1520,12 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP val union = df1.repartition($"a").union(df2.repartition($"a")) val unionExec = union.queryExecution.executedPlan.collect { case u: UnionExec => u @@ -519,7 +613,7 @@ index d838ba4c234..cb0573d56d0 100644 } assert(shuffle.size == 1) -@@ -1553,11 +1555,12 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP +@@ -1553,11 +1556,12 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP val union = df1.repartition($"a").union(df2.repartition($"d")) val unionExec = union.queryExecution.executedPlan.collect { case u: UnionExec => u @@ -533,7 +627,7 @@ index d838ba4c234..cb0573d56d0 100644 } assert(shuffle.size == 1) -@@ -1572,10 +1575,10 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP +@@ -1572,10 +1576,10 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP // Avoid unnecessary shuffle if union output partitioning is enabled val shuffledUnion = union.repartition($"a") val shuffleNumBefore = union.queryExecution.executedPlan.collect { @@ -546,7 +640,7 @@ index d838ba4c234..cb0573d56d0 100644 } if (enabled) { -@@ -1604,6 +1607,7 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP +@@ -1604,6 +1608,7 @@ class DataFrameSetOperationsSuite extends SharedSparkSession with AdaptiveSparkP val union = df1.repartitionByRange($"a").union(df2.repartitionByRange($"d")) val unionExec = union.queryExecution.executedPlan.collect { case u: UnionExec => u @@ -687,6 +781,99 @@ index f79824de8ff..5432984960f 100644 val df = Seq((1, "1"), (2, "2"), (1, "3"), (2, "4")).toDF("key", "value") val window = Window.partitionBy($"key").orderBy($"value") +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +index 0d1b0e1d981..e12c11e7b59 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetCacheSuite.scala +@@ -22,7 +22,8 @@ import java.time.LocalTime + import org.scalatest.concurrent.TimeLimits + import org.scalatest.time.SpanSugar._ + +-import org.apache.spark.sql.execution.ColumnarToRowExec ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec ++import org.apache.spark.sql.execution.{ColumnarToRowExec, SparkPlan} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.columnar.{InMemoryRelation, InMemoryTableScanExec} + import org.apache.spark.sql.functions._ +@@ -38,6 +39,10 @@ class DatasetCacheSuite extends SharedSparkSession + with AdaptiveSparkPlanHelper { + import testImplicits._ + ++ // A scan of a cached relation, Spark's or Comet's. ++ private def isCacheScan(plan: SparkPlan): Boolean = ++ plan.isInstanceOf[InMemoryTableScanExec] || plan.isInstanceOf[CometInMemoryTableScanExec] ++ + /** + * Asserts that a cached [[Dataset]] will be built using the given number of other cached results. + */ +@@ -45,7 +50,7 @@ class DatasetCacheSuite extends SharedSparkSession + val plan = df.queryExecution.withCachedData + assert(plan.isInstanceOf[InMemoryRelation]) + val internalPlan = plan.asInstanceOf[InMemoryRelation].cacheBuilder.cachedPlan +- assert(find(internalPlan)(_.isInstanceOf[InMemoryTableScanExec]).size ++ assert(find(internalPlan)(isCacheScan).size + == numOfCachesDependedUpon) + } + +@@ -255,7 +260,7 @@ class DatasetCacheSuite extends SharedSparkSession + case i: InMemoryRelation => i.cacheBuilder.cachedPlan + } + assert(df2LimitInnerPlan.isDefined && +- !df2LimitInnerPlan.get.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ !df2LimitInnerPlan.get.exists(isCacheScan)) + } + + test("SPARK-27739 Save stats from optimized plan") { +@@ -288,14 +293,14 @@ class DatasetCacheSuite extends SharedSparkSession + val unionDf = df1.union(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.union(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("positive: union by name") { + val unionDf = df1.unionByName(df2).select($"i") + unionDf.cache() + val finalDf = unionDf.unionByName(df3.select($"i")) +- assert(finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by position") { +@@ -303,7 +308,7 @@ class DatasetCacheSuite extends SharedSparkSession + unionDf.cache() + val finalDf = unionDf.union(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + + withClue("negative: union by name") { +@@ -311,7 +316,7 @@ class DatasetCacheSuite extends SharedSparkSession + unionDf.cache() + val finalDf = unionDf.unionByName(df3) + // It's by design to break caching here. +- assert(!finalDf.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!finalDf.queryExecution.executedPlan.exists(isCacheScan)) + } + } + } +@@ -323,12 +328,12 @@ class DatasetCacheSuite extends SharedSparkSession + df1.cache() + // This is exactly the same as df1. + val df2 = spark.range(5).select(struct($"id".as("name", metadata))) +- assert(df2.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(df2.queryExecution.executedPlan.exists(isCacheScan)) + + val metadata2 = Metadata.fromJson("""{"k2": "v2"}""") + // Same with df1 except for the Alias metadata + val df3 = spark.range(5).select(struct($"id".as("name", metadata2))) +- assert(!df3.queryExecution.executedPlan.exists(_.isInstanceOf[InMemoryTableScanExec])) ++ assert(!df3.queryExecution.executedPlan.exists(isCacheScan)) + } + + test("SPARK-53418: Handle TimeType in ColumnAccessor") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 879569045b6..f3ff89067d2 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -1307,10 +1494,18 @@ index 291aa7cab72..7783c37683e 100644 super.test(testName, testTags: _*) { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -index da6f6aca2ad..c02b7c99490 100644 +index da6f6aca2ad..c62307049a2 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/SQLQuerySuite.scala -@@ -1529,7 +1529,8 @@ class SQLQuerySuite extends SharedSparkSession with AdaptiveSparkPlanHelper +@@ -40,6 +40,7 @@ import org.apache.spark.sql.catalyst.parser.ParseException + import org.apache.spark.sql.catalyst.plans.logical.{LocalLimit, Project, RepartitionByExpression, Sort} + import org.apache.spark.sql.connector.catalog.CatalogManager + import org.apache.spark.sql.connector.catalog.CatalogManager.SESSION_CATALOG_NAME ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.{CommandResultExec, OneRowRelationExec, UnionExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.execution.aggregate._ +@@ -1529,7 +1530,8 @@ class SQLQuerySuite extends SharedSparkSession with AdaptiveSparkPlanHelper checkAnswer(sql("select -0.001"), Row(BigDecimal("-0.001"))) } @@ -1320,7 +1515,7 @@ index da6f6aca2ad..c02b7c99490 100644 AccumulatorSuite.verifyPeakExecutionMemorySet(sparkContext, "external sort") { sql("SELECT * FROM testData2 ORDER BY a ASC, b ASC").collect() } -@@ -1990,8 +1991,15 @@ class SQLQuerySuite extends SharedSparkSession with AdaptiveSparkPlanHelper +@@ -1990,8 +1992,15 @@ class SQLQuerySuite extends SharedSparkSession with AdaptiveSparkPlanHelper countAcc.add(1) x }) @@ -1337,6 +1532,14 @@ index da6f6aca2ad..c02b7c99490 100644 verifyCallCount( df.selectExpr("testUdf(a + 1) + testUdf(1 + a)", "testUdf(a + 1)"), Row(4, 2), 1) +@@ -3283,6 +3292,7 @@ class SQLQuerySuite extends SharedSparkSession with AdaptiveSparkPlanHelper + + val inMemoryTableScan = collect(queryDf.queryExecution.executedPlan) { + case i: InMemoryTableScanExec => i ++ case c: CometInMemoryTableScanExec => c + } + assert(inMemoryTableScan.size == 2) + checkAnswer(queryDf, Row(0, 1) :: Row(1, 2) :: Row(2, 3) :: Row(3, 4) :: Row(4, 5) :: Nil) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala index 395cb67f441..33ac6ed19af 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SQLQueryTestSuite.scala @@ -1402,6 +1605,18 @@ index 395cb67f441..33ac6ed19af 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +index bfcf583a705..3228c78b603 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionExtensionSuite.scala +@@ -260,6 +260,7 @@ class SparkSessionExtensionSuite extends PlanTest with SQLHelper with AdaptiveSp + df.select("i").filter($"i" > 1).cache() + assert(find(df.filter($"i" > 1).select("i").queryExecution.executedPlan) { + case _: org.apache.spark.sql.execution.columnar.InMemoryTableScanExec => true ++ case _: org.apache.spark.sql.comet.CometInMemoryTableScanExec => true + case _ => false + }.isDefined) + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala index d7b2511eac2..d5f5b940b94 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/SparkSessionJobTaggingAndCancellationSuite.scala @@ -1417,6 +1632,20 @@ index d7b2511eac2..d5f5b940b94 100644 sc = new SparkContext("local[2]", "test") val session = classic.SparkSession.builder().sparkContext(sc).getOrCreate() import session.implicits._ +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +index 5222d5ce266..67f663c2764 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/StatisticsCollectionSuite.scala +@@ -829,7 +829,8 @@ class StatisticsCollectionSuite extends StatisticsCollectionTestBase with Shared + } + } + +- test("SPARK-33687: analyze all tables in a specific database") { ++ test("SPARK-33687: analyze all tables in a specific database", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + withTempDatabase { database => + spark.catalog.setCurrentDatabase(database) + withTempDir { dir => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala index 63589472854..f8c07a9b037 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/StringFunctionsSuite.scala @@ -1616,7 +1845,7 @@ index 2d26356890d..2c5994f5fbc 100644 } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala -index 37684c7fce3..f3574dec867 100644 +index 37684c7fce3..a0a2f72b6c9 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/collation/CollationSuite.scala @@ -26,6 +26,8 @@ import org.apache.spark.sql.{AnalysisException, Row} @@ -1624,7 +1853,7 @@ index 37684c7fce3..f3574dec867 100644 import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.util.CollationFactory +import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometHashJoinExec, CometSortMergeJoinExec} -+import org.apache.spark.sql.comet.CometHashAggregateExec ++import org.apache.spark.sql.comet.{CometHashAggregateExec, CometInMemoryTableScanExec} import org.apache.spark.sql.connector.{DatasourceV2SQLBase, FakeV2ProviderWithCustomSchema} import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Identifier, InMemoryTable} import org.apache.spark.sql.connector.catalog.CatalogV2Implicits.CatalogHelper @@ -1683,6 +1912,15 @@ index 37684c7fce3..f3574dec867 100644 }.head.isInstanceOf[ArrayTransform]) } } +@@ -2005,7 +2017,7 @@ class CollationSuite extends DatasourceV2SQLBase with AdaptiveSparkPlanHelper { + // Checks in-memory fetching code path. + val all = sql("SELECT col FROM tbl") + assert(all.queryExecution.executedPlan.collectFirst { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + }.nonEmpty) + checkAnswer(all, Row("a")) + // Checks column stats code path. diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala index 5ae23bc3338..5c2c3fff284 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/DataSourceV2Suite.scala @@ -2419,7 +2657,7 @@ index d70bd715879..074a9fa29d9 100644 }) checkAnswer(distinctWithId, Seq(Row(1, 0), Row(1, 0))) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala -index d6d19d21e65..751ad50a569 100644 +index d6d19d21e65..702e77758fa 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala @@ -27,13 +27,15 @@ import org.apache.spark.SparkException @@ -2895,7 +3133,27 @@ index d6d19d21e65..751ad50a569 100644 }.isEmpty) assert(collect(initialExecutedPlan) { case i: InMemoryTableScanLike => i -@@ -3278,7 +3330,8 @@ class AdaptiveQueryExecSuite +@@ -3233,7 +3285,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage") { ++ test("SPARK-42101: Coalesce shuffle partition with union even if exists TableCacheQueryStage", ++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/6454")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1") { + val cached = Seq(1).toDF("c").cache() +@@ -3263,7 +3316,8 @@ class AdaptiveQueryExecSuite + } + } + +- test("SPARK-43376: Improve reuse subquery with table cache") { ++ test("SPARK-43376: Improve reuse subquery with table cache", ++ IgnoreComet("Comet's cache scan does not plan the subqueries in its pruning predicates")) { + withSQLConf(SQLConf.CAN_CHANGE_CACHED_PLAN_OUTPUT_PARTITIONING.key -> "true") { + withTable("t1", "t2") { + withCache("t1") { +@@ -3278,7 +3332,8 @@ class AdaptiveQueryExecSuite } } @@ -2905,7 +3163,7 @@ index d6d19d21e65..751ad50a569 100644 val emptyDf = spark.range(1).where("false") val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df1")) val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name", lit("df2")) -@@ -3368,7 +3421,8 @@ class AdaptiveQueryExecSuite +@@ -3368,7 +3423,8 @@ class AdaptiveQueryExecSuite val plan = df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec] assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec]) @@ -2944,6 +3202,152 @@ index 88be4adb6a4..f8fe831744e 100644 } } } +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +index 57da12e8797..413b5cf31a4 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryColumnarQuerySuite.scala +@@ -22,12 +22,13 @@ import java.sql.{Date, Timestamp} + import java.util.concurrent.atomic.AtomicInteger + + import org.apache.spark.rdd.RDD +-import org.apache.spark.sql.Row ++import org.apache.spark.sql.{IgnoreComet, Row} + import org.apache.spark.sql.catalyst.InternalRow + import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSet, In} + import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning + import org.apache.spark.sql.classic.DataFrame + import org.apache.spark.sql.columnar.CachedBatch ++import org.apache.spark.sql.comet.{CometFilterExec, CometInMemoryTableScanExec} + import org.apache.spark.sql.execution.{FilterExec, InputAdapter, WholeStageCodegenExec} + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions._ +@@ -509,11 +510,17 @@ class InMemoryColumnarQuerySuite extends SharedSparkSession with AdaptiveSparkPl + val planBeforeFilter = collect(df2.queryExecution.executedPlan) { + case f: FilterExec => f.child + case WholeStageCodegenExec(FilterExec(_, i: InputAdapter)) => i.child ++ case f: CometFilterExec => f.child + } +- assert(planBeforeFilter.head.isInstanceOf[InMemoryTableScanExec]) +- + val execPlan = planBeforeFilter.head +- assert(execPlan.executeCollectPublic().length == 0) ++ execPlan match { ++ // Comet's cache scan is columnar only, so count the rows of the batches it emits. ++ case c: CometInMemoryTableScanExec => ++ assert(c.executeColumnar().map(_.numRows().toLong).collect().sum == 0) ++ case _ => ++ assert(execPlan.isInstanceOf[InMemoryTableScanExec]) ++ assert(execPlan.executeCollectPublic().length == 0) ++ } + } + + test("SPARK-25727 - otherCopyArgs in InMemoryRelation does not include outputOrdering") { +@@ -522,7 +529,8 @@ class InMemoryColumnarQuerySuite extends SharedSparkSession with AdaptiveSparkPl + assert(json.contains("outputOrdering")) + } + +- test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached") { ++ test("SPARK-22673: InMemoryRelation should utilize existing stats of the plan to be cached", ++ IgnoreComet("Comet's cache format reports Arrow buffer sizes")) { + Seq("orc", "").foreach { useV1SourceReaderList => + // This test case depends on the size of ORC in statistics. + withSQLConf( +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +index 2c73622739a..5d0efeb263a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/InMemoryRelationSuite.scala +@@ -18,6 +18,7 @@ + package org.apache.spark.sql.execution.columnar + + import org.apache.spark.SparkFunSuite ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.SparkPlan + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.functions.expr +@@ -37,11 +38,15 @@ class InMemoryRelationSuite extends SparkFunSuite + test("SPARK-47177: Cached SQL plan do not display final AQE plan in explain string") { + def findIMRInnerChild(p: SparkPlan): SparkPlan = { + val tableCache = find(p) { +- case _: InMemoryTableScanExec => true ++ case _: InMemoryTableScanExec | _: CometInMemoryTableScanExec => true + case _ => false + } + assert(tableCache.isDefined) +- tableCache.get.asInstanceOf[InMemoryTableScanExec].relation.innerChildren.head ++ val scan = tableCache.get match { ++ case c: CometInMemoryTableScanExec => c.originalPlan ++ case s => s.asInstanceOf[InMemoryTableScanExec] ++ } ++ scan.relation.innerChildren.head + } + + val d1 = spark.range(1).withColumn("key", expr("id % 100")) +diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +index 88ff51d0ff4..7d00fca3b0a 100644 +--- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala ++++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/PartitionBatchPruningSuite.scala +@@ -17,6 +17,7 @@ + + package org.apache.spark.sql.execution.columnar + ++import org.apache.spark.sql.comet.CometInMemoryTableScanExec + import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper + import org.apache.spark.sql.internal.SQLConf + import org.apache.spark.sql.test.SharedSparkSession +@@ -181,11 +182,16 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + val result = df.collect().map(_(0)).toArray + assert(result.length === 1) + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- assert(readPartitions === 5) +- assert(readBatches === 10) ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan has none of these test-only accumulators, so there is nothing to count. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ assert(scans.head.readPartitions.value === 5) ++ assert(scans.head.readBatches.value === 10) ++ } + } + + def checkBatchPruning( +@@ -202,14 +208,23 @@ class PartitionBatchPruningSuite extends SharedSparkSession with AdaptiveSparkPl + df.collect().map(_(0)).toArray + } + +- val (readPartitions, readBatches) = collect(df.queryExecution.executedPlan) { +- case in: InMemoryTableScanExec => (in.readPartitions.value, in.readBatches.value) +- }.head +- +- assert(readBatches === expectedReadBatches, s"Wrong number of read batches: $queryExecution") +- assert( +- readPartitions === expectedReadPartitions, +- s"Wrong number of read partitions: $queryExecution") ++ val scans = collect(df.queryExecution.executedPlan) { case in: InMemoryTableScanExec => in } ++ // Comet's cache scan prunes on the same statistics, but has none of these test-only ++ // accumulators to read, so only the answer above is checked for it. ++ if (scans.isEmpty) { ++ assert(collect(df.queryExecution.executedPlan) { ++ case c: CometInMemoryTableScanExec => c ++ }.nonEmpty) ++ } else { ++ val readPartitions = scans.head.readPartitions.value ++ val readBatches = scans.head.readBatches.value ++ assert( ++ readBatches === expectedReadBatches, ++ s"Wrong number of read batches: $queryExecution") ++ assert( ++ readPartitions === expectedReadPartitions, ++ s"Wrong number of read partitions: $queryExecution") ++ } + } + } + } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala index fd8d1308e99..0e1f80045a3 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala @@ -4302,10 +4706,10 @@ index e2c74533e7f..a12d55848ea 100644 val tblTargetName = "tbl_target" val tblSourceQualified = s"default.$tblSourceName" diff --git a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -index fb26d3311eb..13dc8b89a91 100644 +index fb26d3311eb..188b4d439a2 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/test/SharedSparkSession.scala -@@ -102,6 +102,20 @@ trait SharedSparkSessionBase +@@ -102,6 +102,25 @@ trait SharedSparkSessionBase // this rule may potentially block testing of other optimization rules such as // ConstantPropagation etc. .set(SQLConf.OPTIMIZER_EXCLUDED_RULES.key, ConvertToLocalRelation.ruleName) @@ -4322,6 +4726,11 @@ index fb26d3311eb..13dc8b89a91 100644 + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") + ++ // CometDriverPlugin installs Comet's cache serializer when ++ // spark.comet.exec.inMemoryCache.enabled is on, as it is by default. These sessions do not ++ // load the plugin, so install it here. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } conf.set( StaticSQLConf.WAREHOUSE_PATH, @@ -4400,10 +4809,10 @@ index 59022deaed7..f9aeacb5a9b 100644 override def beforeEach(): Unit = { super.beforeEach() diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -index 8e7ff526a95..ea1072a7195 100644 +index 8e7ff526a95..518e1bae289 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/hive/test/TestHive.scala -@@ -54,24 +54,34 @@ object TestHive +@@ -54,24 +54,38 @@ object TestHive new SparkContext( System.getProperty("spark.sql.test.master", "local[1]"), "TestSQLContext", @@ -4448,6 +4857,10 @@ index 8e7ff526a95..ea1072a7195 100644 + .set("spark.shuffle.manager", + "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager") + .set("spark.comet.shuffle.enabled", "true") ++ ++ // As in SharedSparkSession: what CometDriverPlugin would install. ++ conf.set("spark.sql.cache.serializer", ++ "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + } + + conf From 37560b81e80322a198da65bfc3dcc704e70de5ae Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 30 Sep 2026 09:27:02 -0600 Subject: [PATCH 14/24] test: fix three Spark SQL test adaptations for Comet's cache format - CachedBatchSerializerNoUnwrapSuite sets its own cache serializer, but Spark keeps the first one it loads for the rest of the JVM, so the suite ran with Comet's. Clear it before and after the suite, as CachedBatchSerializerSuite does (4.1.3, 4.2.0). - The SPARK-19993 helper counted the subqueries in a Comet cache scan's original plan, which subquery reuse never reaches, on top of the same subqueries in the filter above the scan. - SPARK-35332's AQE check reads the cached plan from the SQL UI's plan tree, which shows it only under Spark's own cache scan (#6463). Skip it under Comet (4.0.4 and later). --- dev/diffs/3.4.3.diff | 15 ++++++------ dev/diffs/3.5.9.diff | 15 ++++++------ dev/diffs/4.0.4.diff | 29 ++++++++++++++++------- dev/diffs/4.1.3.diff | 56 ++++++++++++++++++++++++++++++++++---------- dev/diffs/4.2.0.diff | 56 ++++++++++++++++++++++++++++++++++---------- 5 files changed, 122 insertions(+), 49 deletions(-) diff --git a/dev/diffs/3.4.3.diff b/dev/diffs/3.4.3.diff index 2080f0eda7b..537cef243cb 100644 --- a/dev/diffs/3.4.3.diff +++ b/dev/diffs/3.4.3.diff @@ -237,7 +237,7 @@ index 0efe0877e9b..423d3b3d76d 100644 -- SELECT_HAVING -- https://github.com/postgres/postgres/blob/REL_12_BETA2/src/test/regress/sql/select_having.sql diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index cf40e944c09..a10f3d46a69 100644 +index cf40e944c09..3dc5574f819 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala @@ -35,10 +35,11 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException @@ -253,17 +253,18 @@ index cf40e944c09..a10f3d46a69 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -113,6 +114,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -113,6 +114,10 @@ class CachedTableSuite extends QueryTest with SQLTestUtils case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => getNumInMemoryTablesRecursively(relation.cachedPlan) + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ // Comet's cache scan keeps the predicates pushed into it on its original plan, out of reach ++ // of subquery reuse. The filter above it evaluates the same subqueries, and counts them. + case cometScan: CometInMemoryTableScanExec => -+ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + -+ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + 1 case p => getNumInMemoryTablesInSubquery(p) }.sum -@@ -393,7 +397,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -393,7 +398,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils assert(isExpectStorageLevel(rddId, Disk)) } @@ -273,7 +274,7 @@ index cf40e944c09..a10f3d46a69 100644 sql("CACHE TABLE testData") spark.table("testData").queryExecution.withCachedData.collect { case cached: InMemoryRelation => -@@ -516,7 +521,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -516,7 +522,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils */ private def verifyNumExchanges(df: DataFrame, expected: Int): Unit = { assert( @@ -283,7 +284,7 @@ index cf40e944c09..a10f3d46a69 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1559,7 +1565,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1559,7 +1566,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils } } diff --git a/dev/diffs/3.5.9.diff b/dev/diffs/3.5.9.diff index fbc99fa97ec..653566b7010 100644 --- a/dev/diffs/3.5.9.diff +++ b/dev/diffs/3.5.9.diff @@ -218,7 +218,7 @@ index 0efe0877e9b..423d3b3d76d 100644 -- SELECT_HAVING -- https://github.com/postgres/postgres/blob/REL_12_BETA2/src/test/regress/sql/select_having.sql diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index e5494726695..eaa3ea7a222 100644 +index e5494726695..7a2a2d8b721 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala @@ -35,10 +35,11 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException @@ -234,17 +234,18 @@ index e5494726695..eaa3ea7a222 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -113,6 +114,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -113,6 +114,10 @@ class CachedTableSuite extends QueryTest with SQLTestUtils case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => getNumInMemoryTablesRecursively(relation.cachedPlan) + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ // Comet's cache scan keeps the predicates pushed into it on its original plan, out of reach ++ // of subquery reuse. The filter above it evaluates the same subqueries, and counts them. + case cometScan: CometInMemoryTableScanExec => -+ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + -+ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + 1 case p => getNumInMemoryTablesInSubquery(p) }.sum -@@ -393,7 +397,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -393,7 +398,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils assert(isExpectStorageLevel(rddId, Disk)) } @@ -254,7 +255,7 @@ index e5494726695..eaa3ea7a222 100644 sql("CACHE TABLE testData") spark.table("testData").queryExecution.withCachedData.collect { case cached: InMemoryRelation => -@@ -519,7 +524,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -519,7 +525,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils df.collect() } assert( @@ -264,7 +265,7 @@ index e5494726695..eaa3ea7a222 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1574,7 +1580,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1574,7 +1581,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils } } diff --git a/dev/diffs/4.0.4.diff b/dev/diffs/4.0.4.diff index 30201a5155f..37521ebcf98 100644 --- a/dev/diffs/4.0.4.diff +++ b/dev/diffs/4.0.4.diff @@ -333,7 +333,7 @@ index 21a3ce1e122..f4762ab98f0 100644 -- In COMPENSATION views get invalidated if the type can't cast diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index 0f42502f1d9..18e5636f4cb 100644 +index 0f42502f1d9..81990d4a97f 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala @@ -36,10 +36,11 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException @@ -349,17 +349,18 @@ index 0f42502f1d9..18e5636f4cb 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -114,6 +115,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -114,6 +115,10 @@ class CachedTableSuite extends QueryTest with SQLTestUtils case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => getNumInMemoryTablesRecursively(relation.cachedPlan) + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ // Comet's cache scan keeps the predicates pushed into it on its original plan, out of reach ++ // of subquery reuse. The filter above it evaluates the same subqueries, and counts them. + case cometScan: CometInMemoryTableScanExec => -+ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + -+ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + 1 case p => getNumInMemoryTablesInSubquery(p) }.sum -@@ -394,7 +398,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -394,7 +399,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils assert(isExpectStorageLevel(rddId, Disk)) } @@ -369,7 +370,7 @@ index 0f42502f1d9..18e5636f4cb 100644 sql("CACHE TABLE testData") spark.table("testData").queryExecution.withCachedData.collect { case cached: InMemoryRelation => -@@ -520,7 +525,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -520,7 +526,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils df.collect() } assert( @@ -379,7 +380,7 @@ index 0f42502f1d9..18e5636f4cb 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1581,7 +1587,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1581,7 +1588,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils } } @@ -389,7 +390,17 @@ index 0f42502f1d9..18e5636f4cb 100644 val tableName = "ntzCache" withTable(tableName) { sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") -@@ -1659,9 +1666,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1626,7 +1634,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + } + } + +- test("SPARK-35332: Make cache plan disable configs configurable - check AQE") { ++ test("SPARK-35332: Make cache plan disable configs configurable - check AQE", ++ IgnoreComet("Spark's SQL UI shows a cached plan only under Spark's own cache scan")) { + withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true") { +@@ -1659,9 +1668,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils _.nodeName.contains("TableCacheQueryStage")) val aqeNode = findNodeInSparkPlanInfo(inMemoryScanNode.get, _.nodeName.contains("AdaptiveSparkPlan")) @@ -411,7 +422,7 @@ index 0f42502f1d9..18e5636f4cb 100644 } withTempView("t0", "t1", "t2") { -@@ -1750,6 +1766,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1750,6 +1768,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils val cached = spark.table("t") val tableCache = collect(cached.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i diff --git a/dev/diffs/4.1.3.diff b/dev/diffs/4.1.3.diff index 371225e9fed..5de0a7f4f46 100644 --- a/dev/diffs/4.1.3.diff +++ b/dev/diffs/4.1.3.diff @@ -381,7 +381,7 @@ index 26d8f750f6e..c888f8e0844 100644 test("SPARK-51777: sql.columnar.* classes registered in KryoSerializer") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index 0d807aeae4d..354a0757d31 100644 +index 0d807aeae4d..ce07fba149d 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala @@ -37,6 +37,7 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException @@ -401,17 +401,18 @@ index 0d807aeae4d..354a0757d31 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -128,6 +129,9 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -128,6 +129,10 @@ class CachedTableSuite extends QueryTest with SQLTestUtils case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => getNumInMemoryTablesRecursively(relation.cachedPlan) + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ // Comet's cache scan keeps the predicates pushed into it on its original plan, out of reach ++ // of subquery reuse. The filter above it evaluates the same subqueries, and counts them. + case cometScan: CometInMemoryTableScanExec => -+ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + -+ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + 1 case p => getNumInMemoryTablesInSubquery(p) }.sum -@@ -408,7 +412,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -408,7 +413,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils assert(isExpectStorageLevel(rddId, Disk)) } @@ -421,7 +422,7 @@ index 0d807aeae4d..354a0757d31 100644 sql("CACHE TABLE testData") spark.table("testData").queryExecution.withCachedData.collect { case cached: InMemoryRelation => -@@ -534,7 +539,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -534,7 +540,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils df.collect() } assert( @@ -431,7 +432,7 @@ index 0d807aeae4d..354a0757d31 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1595,7 +1601,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1595,7 +1602,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils } } @@ -441,7 +442,17 @@ index 0d807aeae4d..354a0757d31 100644 val tableName = "ntzCache" withTable(tableName) { sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") -@@ -1673,9 +1680,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1640,7 +1648,8 @@ class CachedTableSuite extends QueryTest with SQLTestUtils + } + } + +- test("SPARK-35332: Make cache plan disable configs configurable - check AQE") { ++ test("SPARK-35332: Make cache plan disable configs configurable - check AQE", ++ IgnoreComet("Spark's SQL UI shows a cached plan only under Spark's own cache scan")) { + withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true") { +@@ -1673,9 +1682,18 @@ class CachedTableSuite extends QueryTest with SQLTestUtils _.nodeName.contains("TableCacheQueryStage")) val aqeNode = findNodeInSparkPlanInfo(inMemoryScanNode.get, _.nodeName.contains("AdaptiveSparkPlan")) @@ -463,7 +474,7 @@ index 0d807aeae4d..354a0757d31 100644 } withTempView("t0", "t1", "t2") { -@@ -1764,6 +1780,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -1764,6 +1782,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils val cached = spark.table("t") val tableCache = collect(cached.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i @@ -471,7 +482,7 @@ index 0d807aeae4d..354a0757d31 100644 } if (expected == StorageLevel.NONE) { assert(tableCache.isEmpty) -@@ -2630,6 +2647,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -2630,6 +2649,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils val inMemoryTableScan = collect(df.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i @@ -479,7 +490,7 @@ index 0d807aeae4d..354a0757d31 100644 } assert(inMemoryTableScan.size == 1) checkAnswer(df, Row(5) :: Nil) -@@ -2657,6 +2675,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils +@@ -2657,6 +2677,7 @@ class CachedTableSuite extends QueryTest with SQLTestUtils val subqueryInMemoryTableScan = collect(cteInSubquery.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i @@ -3092,10 +3103,29 @@ index 188a28ff1c0..cfc84dcba01 100644 assert(o1.semanticEquals(o2), "Different output column order after AQE optimization") } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala -index 47b935a2880..3fdeab3113c 100644 +index 47b935a2880..65ee66c1975 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala -@@ -230,9 +230,21 @@ class CachedBatchSerializerNoUnwrapSuite extends QueryTest +@@ -212,6 +212,18 @@ class CachedBatchSerializerNoUnwrapSuite extends QueryTest + classOf[DefaultCachedBatchSerializerNoUnwrap].getName) + } + ++ // Spark keeps the first cache serializer it loads for the rest of the JVM, which the suites ++ // before this one set to Comet's. Clear it on both sides, as CachedBatchSerializerSuite does. ++ protected override def beforeAll(): Unit = { ++ super.beforeAll() ++ clearSerializer() ++ } ++ ++ protected override def afterAll(): Unit = { ++ clearSerializer() ++ super.afterAll() ++ } ++ + test("Do not unwrap ColumnarToRowExec") { + withTempPath { workDir => + val workDirPath = workDir.getAbsolutePath +@@ -230,9 +242,21 @@ class CachedBatchSerializerNoUnwrapSuite extends QueryTest assert(cachedPlans.length == 2) cachedPlans.foreach { cachedPlan => diff --git a/dev/diffs/4.2.0.diff b/dev/diffs/4.2.0.diff index dfb411ba504..ea06e8cbd40 100644 --- a/dev/diffs/4.2.0.diff +++ b/dev/diffs/4.2.0.diff @@ -399,7 +399,7 @@ index 72a2da16054..54146036f8a 100644 test("SPARK-51777: sql.columnar.* classes registered in KryoSerializer") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala -index 085dbcd8046..280d020ea33 100644 +index 085dbcd8046..ac98c64ae92 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/CachedTableSuite.scala @@ -37,6 +37,7 @@ import org.apache.spark.sql.catalyst.analysis.TempTableAlreadyExistsException @@ -419,17 +419,18 @@ index 085dbcd8046..280d020ea33 100644 import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -@@ -127,6 +128,9 @@ class CachedTableSuite extends SharedSparkSession +@@ -127,6 +128,10 @@ class CachedTableSuite extends SharedSparkSession case inMemoryTable @ InMemoryTableScanExec(_, _, relation) => getNumInMemoryTablesRecursively(relation.cachedPlan) + getNumInMemoryTablesInSubquery(inMemoryTable) + 1 ++ // Comet's cache scan keeps the predicates pushed into it on its original plan, out of reach ++ // of subquery reuse. The filter above it evaluates the same subqueries, and counts them. + case cometScan: CometInMemoryTableScanExec => -+ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + -+ getNumInMemoryTablesInSubquery(cometScan.originalPlan) + 1 ++ getNumInMemoryTablesRecursively(cometScan.originalPlan.relation.cachedPlan) + 1 case p => getNumInMemoryTablesInSubquery(p) }.sum -@@ -407,7 +411,8 @@ class CachedTableSuite extends SharedSparkSession +@@ -407,7 +412,8 @@ class CachedTableSuite extends SharedSparkSession assert(isExpectStorageLevel(rddId, Disk)) } @@ -439,7 +440,7 @@ index 085dbcd8046..280d020ea33 100644 sql("CACHE TABLE testData") spark.table("testData").queryExecution.withCachedData.collect { case cached: InMemoryRelation => -@@ -564,7 +569,8 @@ class CachedTableSuite extends SharedSparkSession +@@ -564,7 +570,8 @@ class CachedTableSuite extends SharedSparkSession df.collect() } assert( @@ -449,7 +450,7 @@ index 085dbcd8046..280d020ea33 100644 } test("A cached table preserves the partitioning and ordering of its cached SparkPlan") { -@@ -1625,7 +1631,8 @@ class CachedTableSuite extends SharedSparkSession +@@ -1625,7 +1632,8 @@ class CachedTableSuite extends SharedSparkSession } } @@ -459,7 +460,17 @@ index 085dbcd8046..280d020ea33 100644 val tableName = "ntzCache" withTable(tableName) { sql(s"CACHE TABLE $tableName AS SELECT TIMESTAMP_NTZ'2021-01-01 00:00:00'") -@@ -1703,9 +1710,18 @@ class CachedTableSuite extends SharedSparkSession +@@ -1670,7 +1678,8 @@ class CachedTableSuite extends SharedSparkSession + } + } + +- test("SPARK-35332: Make cache plan disable configs configurable - check AQE") { ++ test("SPARK-35332: Make cache plan disable configs configurable - check AQE", ++ IgnoreComet("Spark's SQL UI shows a cached plan only under Spark's own cache scan")) { + withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true") { +@@ -1703,9 +1712,18 @@ class CachedTableSuite extends SharedSparkSession _.nodeName.contains("TableCacheQueryStage")) val aqeNode = findNodeInSparkPlanInfo(inMemoryScanNode.get, _.nodeName.contains("AdaptiveSparkPlan")) @@ -481,7 +492,7 @@ index 085dbcd8046..280d020ea33 100644 } withTempView("t0", "t1", "t2") { -@@ -1794,6 +1810,7 @@ class CachedTableSuite extends SharedSparkSession +@@ -1794,6 +1812,7 @@ class CachedTableSuite extends SharedSparkSession val cached = spark.table("t") val tableCache = collect(cached.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i @@ -489,7 +500,7 @@ index 085dbcd8046..280d020ea33 100644 } if (expected == StorageLevel.NONE) { assert(tableCache.isEmpty) -@@ -2660,6 +2677,7 @@ class CachedTableSuite extends SharedSparkSession +@@ -2660,6 +2679,7 @@ class CachedTableSuite extends SharedSparkSession val inMemoryTableScan = collect(df.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i @@ -497,7 +508,7 @@ index 085dbcd8046..280d020ea33 100644 } assert(inMemoryTableScan.size == 1) checkAnswer(df, Row(5) :: Nil) -@@ -2687,6 +2705,7 @@ class CachedTableSuite extends SharedSparkSession +@@ -2687,6 +2707,7 @@ class CachedTableSuite extends SharedSparkSession val subqueryInMemoryTableScan = collect(cteInSubquery.queryExecution.executedPlan) { case i: InMemoryTableScanExec => i @@ -3174,10 +3185,29 @@ index d6d19d21e65..702e77758fa 100644 assert(o1.semanticEquals(o2), "Different output column order after AQE optimization") } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala -index 88be4adb6a4..f8fe831744e 100644 +index 88be4adb6a4..23ec5374f8a 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/columnar/CachedBatchSerializerSuite.scala -@@ -228,9 +228,21 @@ class CachedBatchSerializerNoUnwrapSuite extends SharedSparkSession with Adaptiv +@@ -210,6 +210,18 @@ class CachedBatchSerializerNoUnwrapSuite extends SharedSparkSession with Adaptiv + classOf[DefaultCachedBatchSerializerNoUnwrap].getName) + } + ++ // Spark keeps the first cache serializer it loads for the rest of the JVM, which the suites ++ // before this one set to Comet's. Clear it on both sides, as CachedBatchSerializerSuite does. ++ protected override def beforeAll(): Unit = { ++ super.beforeAll() ++ clearSerializer() ++ } ++ ++ protected override def afterAll(): Unit = { ++ clearSerializer() ++ super.afterAll() ++ } ++ + test("Do not unwrap ColumnarToRowExec") { + withTempPath { workDir => + val workDirPath = workDir.getAbsolutePath +@@ -228,9 +240,21 @@ class CachedBatchSerializerNoUnwrapSuite extends SharedSparkSession with Adaptiv assert(cachedPlans.length == 2) cachedPlans.foreach { cachedPlan => From 29edf939f2078747b1ab300b5db22bd2c9376942 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 04:58:53 -0600 Subject: [PATCH 15/24] fix: keep Spark's cache format when Kryo would reject Comet's or Comet disables itself CometDriverPlugin installed ArrowCachedBatchSerializer whenever Comet, its native execution and the cache config were enabled at startup. Two more startup settings decide whether that format can work: - With spark.kryo.registrationRequired=true and no CometKryoRegistrator, Kryo rejects CometCachedBatch, so caching failed with "Class is not registered" as soon as Spark serialized a cached block, where Spark's own format works. The plugin only warned. It now keeps Spark's format, and the warning says so. - With Comet shuffle enabled but neither of Comet's shuffle managers configured, Comet disables itself, so every cache was stored in Comet's format with only Spark operators to read it. The plugin now keeps Spark's format there too. The plugin's boolean config reads also honor a deprecated alternative key, such as spark.comet.exec.shuffle.enabled, as a session does. --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../contributor-guide/plugin_overview.md | 2 + .../user-guide/latest/in-memory-cache.md | 18 +++- docs/source/user-guide/latest/installation.md | 3 +- .../scala/org/apache/comet/CometConf.scala | 12 ++- .../main/scala/org/apache/spark/Plugins.scala | 48 +++++++--- .../exec/CometInMemoryCacheKryoSuite.scala | 92 ++++++++++++++++++- .../comet/exec/CometInMemoryCacheSuite.scala | 66 +++++++++++-- 9 files changed, 207 insertions(+), 36 deletions(-) diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index d85eb14d2c4..0f28229ae8b 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -515,6 +515,7 @@ jobs: org.apache.comet.exec.CometInMemoryCacheSuite org.apache.comet.exec.CometInMemoryCachePruningSuite org.apache.comet.exec.CometInMemoryCacheKryoSuite + org.apache.comet.exec.CometInMemoryCacheKryoUnregisteredSuite org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometJoinSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 02d3077e4fc..7838bf14d5a 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -219,6 +219,7 @@ jobs: org.apache.comet.exec.CometInMemoryCacheSuite org.apache.comet.exec.CometInMemoryCachePruningSuite org.apache.comet.exec.CometInMemoryCacheKryoSuite + org.apache.comet.exec.CometInMemoryCacheKryoUnregisteredSuite org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometJoinSuite diff --git a/docs/source/contributor-guide/plugin_overview.md b/docs/source/contributor-guide/plugin_overview.md index 89586bf7b6d..072f59e01a2 100644 --- a/docs/source/contributor-guide/plugin_overview.md +++ b/docs/source/contributor-guide/plugin_overview.md @@ -48,6 +48,8 @@ and skips the remaining steps. Otherwise it: - Appends `CometSparkSessionExtensions` to `spark.sql.extensions`, unless it is already listed. - Sets `spark.sql.cache.serializer` to Comet's `ArrowCachedBatchSerializer` when `spark.comet.exec.inMemoryCache.enabled=true`, unless the application has chosen a different serializer. + It leaves Spark's serializer in place where Comet could not scan its format natively or Kryo would + reject it; see [In-Memory Cache](../user-guide/latest/in-memory-cache.md). - Registers `CometSource` with Spark's metrics system and adds `CometMetricsListener` to `spark.sql.queryExecutionListeners` when `spark.comet.metrics.enabled=true`. - Logs a warning for settings that are likely to cause problems, such as an unset `spark.executor.memoryOverhead`. diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 8de4759957f..52a5f62f2b8 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -38,7 +38,9 @@ It has to be set before the `SparkContext` starts. Comet's driver plugin chooses with the default goes on using Spark's cache format however the config is set afterwards. The plugin installs Comet's serializer only if `spark.comet.enabled` and `spark.comet.exec.enabled` are enabled at that point too, because an application that starts without native execution could -not scan Comet's format natively. +not scan Comet's format natively. It also keeps Spark's format when Comet shuffle is enabled but +`spark.shuffle.manager` is not one of Comet's shuffle managers, since Comet then disables itself, +and when Kryo would reject Comet's format; see [Kryo](#kryo). ## What changes when it is enabled @@ -179,10 +181,16 @@ spark.kryo.registrator=org.apache.comet.CometKryoRegistrator Comet cannot set `spark.kryo.registrator` for you the way it sets `spark.sql.cache.serializer`: `KryoSerializer` reads it when `SparkEnv` builds the serializer, which happens before any plugin -runs. Without it, caching fails with a "Class is not registered" error that does not name this -feature. Comet's driver plugin warns at startup when it sees Kryo, `registrationRequired`, and no -registrator. Native broadcast needs the same registrator even when the cache is disabled; see -[Kryo serialization](installation.md#kryo-serialization). +runs. Without it, Kryo would reject Comet's cached batch with a "Class is not registered" error +that does not name this feature, so Comet's driver plugin does not install Comet's serializer, +and caches stay in Spark's format. The plugin warns at startup when it sees Kryo, +`registrationRequired`, and no registrator. An application that sets `spark.sql.cache.serializer` +to Comet's serializer itself gets the error instead. Native broadcast needs the same registrator +even when the cache is disabled; see [Kryo serialization](installation.md#kryo-serialization). + +Spark registers its own cached batch with Kryo only from Spark 4.1, so on earlier versions caching +in either format under `registrationRequired` needs a registrator. `CometKryoRegistrator` registers +Spark's cached batch too. ## Limitations diff --git a/docs/source/user-guide/latest/installation.md b/docs/source/user-guide/latest/installation.md index b9d28b74b03..273ea664772 100644 --- a/docs/source/user-guide/latest/installation.md +++ b/docs/source/user-guide/latest/installation.md @@ -269,7 +269,8 @@ If the application uses Kryo (`spark.serializer=org.apache.spark.serializer.Kryo Without it, any query that uses Comet's native broadcast exchange, which is enabled by default, fails with Kryo's "Class is not registered" error, for example on the first broadcast hash join. -The [in-memory cache](in-memory-cache.md#kryo) needs the same registrator. Set it before the +Comet's [in-memory cache](in-memory-cache.md#kryo) format needs the same registrator, and without +it Comet's plugin keeps caches in Spark's format. Set it before the `SparkContext` is created: `KryoSerializer` reads it before Comet's plugin runs, so Comet cannot add it for you. `spark.kryo.registrator` accepts a comma-separated list, so an application with its own registrator can list both. Comet logs a warning at startup when Kryo requires registration diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 6cfe5bc92bd..ced5860d685 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -275,7 +275,8 @@ object CometConf extends ShimCometConf { .doc("Whether to enable Comet native execution for in-memory cached tables. Its value at " + "startup also decides whether CometDriverPlugin installs Comet's cache serializer, " + "which stores cached data in Arrow format. The plugin installs it only if " + - "spark.comet.enabled and spark.comet.exec.enabled are also enabled at startup. " + + "spark.comet.enabled and spark.comet.exec.enabled are also enabled at startup, and " + + "only with one of Comet's shuffle managers while Comet shuffle is enabled. " + "Because spark.sql.cache.serializer is a " + "static config, the cached format is fixed for the application, and disabling this " + "at runtime only sends cached scans back to Spark's execution path. Relations whose " + @@ -284,10 +285,11 @@ object CometConf extends ShimCometConf { "zstd compression, and a scan copies out only the buffers of the columns it projected, " + "so the unselected ones are never decompressed. Reads that feed Spark operators rather " + "than Comet ones still pay a row conversion the default format avoids, and can be " + - "slower than Spark's cache. With spark.kryo.registrationRequired=true, also set " + - "spark.kryo.registrator=org.apache.comet.CometKryoRegistrator before creating the " + - "SparkContext, otherwise caching fails as soon as a block is serialized, including " + - "the disk half of the default MEMORY_AND_DISK storage level.") + "slower than Spark's cache. With spark.kryo.registrationRequired=true, the plugin " + + "installs it only if spark.kryo.registrator includes " + + "org.apache.comet.CometKryoRegistrator, which has to be set before creating the " + + "SparkContext, because Kryo would otherwise reject a cached block as soon as it is " + + "serialized, including the disk half of the default MEMORY_AND_DISK storage level.") .booleanConf .createWithDefault(false) diff --git a/spark/src/main/scala/org/apache/spark/Plugins.scala b/spark/src/main/scala/org/apache/spark/Plugins.scala index c3246679d02..c420b905ae6 100644 --- a/spark/src/main/scala/org/apache/spark/Plugins.scala +++ b/spark/src/main/scala/org/apache/spark/Plugins.scala @@ -30,6 +30,7 @@ import org.apache.spark.api.plugin.{DriverPlugin, ExecutorPlugin, PluginContext, import org.apache.spark.internal.Logging import org.apache.spark.internal.config.{EVENT_LOG_ENABLED, EXECUTOR_MEMORY_OVERHEAD, EXECUTOR_MEMORY_OVERHEAD_FACTOR} import org.apache.spark.scheduler.{SparkListener, SparkListenerApplicationEnd, SparkListenerExecutorMetricsUpdate, SparkListenerExecutorRemoved} +import org.apache.spark.sql.comet.execution.shuffle.{CometCelebornShuffleManager, CometShuffleManager} import org.apache.spark.sql.internal.StaticSQLConf import org.apache.spark.util.{Clock, SystemClock} @@ -174,7 +175,12 @@ object CometDriverPlugin extends Logging { // Use Comet's cache serializer only when the native in-memory cache scan can run, which needs // Comet and its native execution as well as the cache config. spark.sql.cache.serializer is // static, so an application that starts with Comet or native execution off would otherwise - // store every cache in Comet's format, with only Spark operators to read it. + // store every cache in Comet's format, with only Spark operators to read it. So would one that + // leaves Comet shuffle enabled without Comet's shuffle manager, since Comet then disables + // itself. + // Nor is it used where Kryo would reject its cached batches, under + // spark.kryo.registrationRequired without CometKryoRegistrator: caching that works in Spark's + // format would then fail the first time Spark serialized a cached block. // If the application already set spark.sql.cache.serializer, leave that value // unchanged so Comet does not replace a user-selected cache format. private[apache] def maybeSetCacheSerializer( @@ -182,7 +188,9 @@ object CometDriverPlugin extends Logging { extraConfs: ju.HashMap[String, String]): Unit = { if (getBooleanConf(conf, CometConf.COMET_ENABLED) && getBooleanConf(conf, CometConf.COMET_EXEC_ENABLED) && - getBooleanConf(conf, CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED)) { + getBooleanConf(conf, CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED) && + (!getBooleanConf(conf, CometConf.COMET_SHUFFLE_ENABLED) || isCometShuffleManager(conf)) && + !isKryoRegistratorMissing(conf)) { val serializerKey = StaticSQLConf.SPARK_CACHE_SERIALIZER.key val serializerValue = "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer" @@ -208,6 +216,18 @@ object CometDriverPlugin extends Logging { // serializer, before any plugin runs, so it cannot be set from here. Say so while the // application is still starting up rather than leaving the user to attribute the failure later. private[apache] def warnIfKryoRegistratorMissing(conf: SparkConf): Unit = { + if (isKryoRegistratorMissing(conf)) { + logWarning( + "spark.kryo.registrationRequired=true but spark.kryo.registrator does not include " + + s"${CometKryoRegistrator.CLASS_NAME}. Comet's native broadcast will fail with " + + "Kryo's \"Class is not registered\" as soon as its payload is serialized, and " + + "Comet does not install its in-memory cache format, which would fail the same way. " + + s"Add spark.kryo.registrator=${CometKryoRegistrator.CLASS_NAME} before creating the " + + "SparkContext; it cannot be set later.") + } + } + + private def isKryoRegistratorMissing(conf: SparkConf): Boolean = { val usingKryo = conf.get("spark.serializer", "") == "org.apache.spark.serializer.KryoSerializer" val registrationRequired = conf.getBoolean("spark.kryo.registrationRequired", false) @@ -216,18 +236,15 @@ object CometDriverPlugin extends Logging { .split(',') .map(_.trim) .contains(CometKryoRegistrator.CLASS_NAME) - - if (usingKryo && registrationRequired && !registered) { - logWarning( - "spark.kryo.registrationRequired=true but spark.kryo.registrator does not include " + - s"${CometKryoRegistrator.CLASS_NAME}. Comet's native broadcast and its in-memory " + - "cache format will fail with Kryo's \"Class is not registered\" as soon as their " + - "payloads are serialized. Add " + - s"spark.kryo.registrator=${CometKryoRegistrator.CLASS_NAME} before creating the " + - "SparkContext; it cannot be set later.") - } + usingKryo && registrationRequired && !registered } + // Comet's shuffle managers have no short name, so spark.shuffle.manager names one only by its + // class name. + private def isCometShuffleManager(conf: SparkConf): Boolean = + Set(classOf[CometShuffleManager].getName, classOf[CometCelebornShuffleManager].getName) + .contains(conf.get("spark.shuffle.manager", "")) + // Comet's native allocations are made by the Rust global allocator and live in the native heap. // In off-heap mode the share that operators reserve is charged against a memory pool, but // everything else -- expression kernels and Arrow array builders, decompression buffers, Parquet @@ -298,8 +315,13 @@ object CometDriverPlugin extends Logging { } } + // Reads a deprecated alternative too, such as spark.comet.exec.shuffle.enabled, as a session + // would. private def getBooleanConf(conf: SparkConf, entry: ConfigEntry[Boolean]): Boolean = - conf.getBoolean(entry.key, entry.defaultValue.get) + (entry.key +: entry.alternatives) + .find(conf.contains) + .map(conf.getBoolean(_, entry.defaultValue.get)) + .getOrElse(entry.defaultValue.get) def registerCometMetrics(sc: SparkContext): Unit = { if (sc.getConf.getBoolean( diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala index b403442531c..ea7d2b38307 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala @@ -20,10 +20,15 @@ package org.apache.comet.exec import org.apache.spark.SparkConf +import org.apache.spark.serializer.KryoRegistrator import org.apache.spark.sql.{CometTestBase, Row} -import org.apache.spark.sql.execution.columnar.CometInMemoryRelationHelper -import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, DefaultCachedBatch} +import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} import org.apache.spark.storage.StorageLevel +import org.apache.spark.unsafe.types.UTF8String + +import com.esotericsoftware.kryo.Kryo import org.apache.comet.{CometConf, CometKryoRegistrator} @@ -194,3 +199,86 @@ class CometInMemoryCacheKryoSuite extends CometTestBase { } } } + +/** + * The application that `CometDriverPlugin` warns about: Kryo with registration required and no + * [[CometKryoRegistrator]]. The plugin then leaves `spark.sql.cache.serializer` alone rather than + * install Comet's serializer, whose cached batch Kryo would reject, so caching works as it does + * without Comet. + * + * Spark registers its own cached batch with Kryo only from 4.1, so on earlier versions an + * application whose caches work under registration has to register it itself. The registrator + * this suite installs does that, and nothing of Comet's. + */ +class CometInMemoryCacheKryoUnregisteredSuite extends CometTestBase { + + override protected def beforeAll(): Unit = { + CometInMemoryRelationHelper.clearSerializer() + super.beforeAll() + } + + override protected def afterAll(): Unit = { + try { + super.afterAll() + } finally { + CometInMemoryRelationHelper.clearSerializer() + } + } + + override protected def sparkConf: SparkConf = { + val conf = super.sparkConf + conf.set("spark.plugins", "org.apache.spark.CometPlugin") + conf.set(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key, "true") + conf.set("spark.serializer", "org.apache.spark.serializer.KryoSerializer") + conf.set("spark.kryo.registrationRequired", "true") + conf.set("spark.kryo.registrator", classOf[SparkCachedBatchKryoRegistrator].getName) + conf + } + + private def cachedBatchTypes(table: String): Array[String] = { + val cached = spark.sharedState.cacheManager.lookupCachedData(spark.table(table)).get + cached.cachedRepresentation.cacheBuilder.cachedColumnBuffers + .map(_.getClass.getName) + .distinct() + .collect() + } + + test("Comet plugin keeps Spark's cache format when Kryo would reject Comet's") { + assert(!spark.sparkContext.getConf.contains(StaticSQLConf.SPARK_CACHE_SERIALIZER.key)) + + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + spark.catalog.clearCache() + try { + spark + .range(0, 100, 1, 2) + .selectExpr("id", "cast(id as string) AS s") + .createOrReplaceTempView("kryo_unregistered") + + // DISK_ONLY serializes every block as it is put, so this reaches Kryo in local mode. + spark.catalog.cacheTable("kryo_unregistered", StorageLevel.DISK_ONLY) + assert(spark.table("kryo_unregistered").count() == 100) + assert( + cachedBatchTypes("kryo_unregistered").sameElements( + Array("org.apache.spark.sql.execution.columnar.DefaultCachedBatch"))) + + checkAnswer( + spark.sql("SELECT s FROM kryo_unregistered WHERE id > 97"), + Seq(Row("98"), Row("99"))) + } finally { + spark.catalog.clearCache() + } + } + } +} + +/** Registers what Spark's own cached batch needs, over a long and a string column, with Kryo. */ +class SparkCachedBatchKryoRegistrator extends KryoRegistrator { + override def registerClasses(kryo: Kryo): Unit = { + // The batch and its statistics row, whose bounds include UTF8String. + Seq( + classOf[DefaultCachedBatch], + classOf[GenericInternalRow], + classOf[Array[Any]], + classOf[UTF8String]).foreach(kryo.register) + } +} diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index ecce9067522..488e57b7a97 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -39,6 +39,7 @@ import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeRefer import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} +import org.apache.spark.sql.comet.execution.shuffle.CometCelebornShuffleManager import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.SortExec import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, QueryStageExec, ShuffleQueryStageExec} @@ -51,7 +52,7 @@ import org.apache.spark.sql.types._ import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} import org.apache.spark.storage.StorageLevel -import org.apache.comet.{CometArrowAllocator, CometConf, ExtendedExplainInfo} +import org.apache.comet.{CometArrowAllocator, CometConf, CometKryoRegistrator, ExtendedExplainInfo} import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus} import org.apache.comet.vector.{CometPlainVector, CometVector} @@ -1021,6 +1022,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val defaultConf = new SparkConf() .set(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key, "true") + .set("spark.shuffle.manager", shuffleManager) val defaultExtraConfs = new ju.HashMap[String, String]() // With no user serializer configured, the plugin should install Comet's @@ -1032,6 +1034,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val userConf = new SparkConf() .set(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key, "true") + .set("spark.shuffle.manager", shuffleManager) .set(serializerKey, userSerializer) val userExtraConfs = new ju.HashMap[String, String]() @@ -1043,20 +1046,25 @@ class CometInMemoryCacheSuite extends CometTestBase { assert(!userExtraConfs.containsKey(serializerKey)) } - test("Comet plugin installs its cache serializer only if Comet can scan the cache natively") { + /** Whether the Comet plugin installs its cache serializer for an application's `settings`. */ + private def installsCacheSerializer(settings: (String, String)*): Boolean = { val serializerKey = StaticSQLConf.SPARK_CACHE_SERIALIZER.key + val conf = new SparkConf().setAll(settings) + val extraConfs = new ju.HashMap[String, String]() + CometDriverPlugin.maybeSetCacheSerializer(conf, extraConfs) + assert(conf.contains(serializerKey) == extraConfs.containsKey(serializerKey)) + extraConfs.containsKey(serializerKey) + } - def installed(settings: (String, String)*): Boolean = { - val conf = new SparkConf().setAll(settings) - val extraConfs = new ju.HashMap[String, String]() - CometDriverPlugin.maybeSetCacheSerializer(conf, extraConfs) - assert(conf.contains(serializerKey) == extraConfs.containsKey(serializerKey)) - extraConfs.containsKey(serializerKey) - } - + test("Comet plugin installs its cache serializer only if Comet can scan the cache natively") { val cometOn = CometConf.COMET_ENABLED.key -> "true" val execOn = CometConf.COMET_EXEC_ENABLED.key -> "true" val cacheOn = CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true" + // Without Comet's shuffle manager Comet disables itself, which the next test covers. + val cometShuffle = "spark.shuffle.manager" -> shuffleManager + + def installed(settings: (String, String)*): Boolean = + installsCacheSerializer(cometShuffle +: settings: _*) assert(installed(cometOn, execOn, cacheOn)) // An application that starts with Comet or its native execution off can never plan @@ -1074,6 +1082,44 @@ class CometInMemoryCacheSuite extends CometTestBase { CometConf.COMET_EXEC_ENABLED.defaultValue.get)) } + test("Comet plugin keeps Spark's cache format where Comet disables itself or Kryo rejects it") { + val cometShuffle = "spark.shuffle.manager" -> shuffleManager + + def installed(settings: (String, String)*): Boolean = + installsCacheSerializer( + Seq( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true") ++ settings: _*) + + assert(installed(cometShuffle)) + // Comet shuffle is enabled by default, and Comet disables itself while it is unless the + // application runs one of Comet's shuffle managers. + assert(!installed()) + assert(!installed("spark.shuffle.manager" -> "sort")) + assert(installed("spark.shuffle.manager" -> classOf[CometCelebornShuffleManager].getName)) + // Without Comet shuffle, under the key or its deprecated name, the shuffle manager does not + // matter. + assert(installed(CometConf.COMET_SHUFFLE_ENABLED.key -> "false")) + assert(installed("spark.comet.exec.shuffle.enabled" -> "false")) + + // Kryo with registration required rejects Comet's cached batch unless the application lists + // CometKryoRegistrator, on its own or beside a registrator of its own. + val kryo = "spark.serializer" -> "org.apache.spark.serializer.KryoSerializer" + val registrationRequired = "spark.kryo.registrationRequired" -> "true" + val registrator = "spark.kryo.registrator" + assert(!installed(cometShuffle, kryo, registrationRequired)) + assert(!installed(cometShuffle, kryo, registrationRequired, registrator -> "com.example.R")) + assert( + installed( + cometShuffle, + kryo, + registrationRequired, + registrator -> s"com.example.R, ${CometKryoRegistrator.CLASS_NAME}")) + // Without registrationRequired, Kryo writes the class name of anything unregistered instead. + assert(installed(cometShuffle, kryo)) + } + test("Comet in-memory cache supports empty projection scan") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", From 567956c765ac518c61265403e7e1211516479993 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:01:24 -0600 Subject: [PATCH 16/24] docs: describe the in-memory cache default in the upgrade guide Add a 1.2.0 section to the upgrade guide for the new default of spark.comet.exec.inMemoryCache.enabled: what changes, how to keep Spark's cache format, and when Comet keeps it without being asked. The operator and compatibility pages still called the feature disabled by default, and the cache guide still said an application that started with the default kept Spark's format. --- .../latest/compatibility/operators.md | 4 +-- .../user-guide/latest/in-memory-cache.md | 4 +-- .../user-guide/latest/migration-guide.md | 29 +++++++++++++++++++ docs/source/user-guide/latest/operators.md | 2 +- 4 files changed, 34 insertions(+), 5 deletions(-) diff --git a/docs/source/user-guide/latest/compatibility/operators.md b/docs/source/user-guide/latest/compatibility/operators.md index 53b290726aa..0ea0ac5b3e1 100644 --- a/docs/source/user-guide/latest/compatibility/operators.md +++ b/docs/source/user-guide/latest/compatibility/operators.md @@ -35,8 +35,8 @@ readable empty output files and their schema metadata. ## In-Memory Cache Comet can store cached relations (`df.cache()`, `CACHE TABLE`) in Arrow format and scan them -natively. This is experimental and disabled by default; see [In-Memory Cache](../in-memory-cache.md) -for how to enable it. Comet does not replace a `spark.sql.cache.serializer` that the application +natively. This is experimental and enabled by default; see [In-Memory Cache](../in-memory-cache.md) +for how to turn it off. Comet does not replace a `spark.sql.cache.serializer` that the application has already set. Relations whose schema Comet's Arrow writer does not support are cached in Spark's default format, and their scans fall back to Spark. Reads that feed Spark operators rather than Comet operators can be slower than Spark's cache. diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 685a03e6fa0..128950ab8f0 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -34,8 +34,8 @@ $SPARK_HOME/bin/spark-shell \ ``` It has to be set before the `SparkContext` starts. Comet's driver plugin chooses -`spark.sql.cache.serializer` once, while the context is initializing, so a session that started -with the default goes on using Spark's cache format however the config is set afterwards. The +`spark.sql.cache.serializer` once, while the context is initializing, so an application keeps the +cache format it started with however the config is set afterwards. The plugin installs Comet's serializer only if `spark.comet.enabled` and `spark.comet.exec.enabled` are enabled at that point too, because an application that starts without native execution could not scan Comet's format natively. It also keeps Spark's format when Comet shuffle is enabled but diff --git a/docs/source/user-guide/latest/migration-guide.md b/docs/source/user-guide/latest/migration-guide.md index 342f8d21182..e634dafa729 100644 --- a/docs/source/user-guide/latest/migration-guide.md +++ b/docs/source/user-guide/latest/migration-guide.md @@ -57,6 +57,35 @@ Treat setting one of these keys as a temporary measure. If you find you cannot s legacy behavior, please open an issue describing your use case so it can be considered before the key is removed. +## Upgrading to Comet 1.2.0 + +Comet `1.2.0` makes no behavior changes that need a `spark.comet.legacy.*` key. The changes below +need none either, but check whether any of them applies to your deployment. + +### In-Memory Cache Enabled by Default + +`spark.comet.exec.inMemoryCache.enabled` now defaults to `true`. An application that loads +`CometPlugin` now stores what it caches with `CACHE TABLE`, `df.cache()` or `df.persist()` in +Comet's Arrow format instead of Spark's, and Comet scans it natively. The format does not change +query results, but it can change performance: Spark operators read Comet's format more slowly than +Spark's own, which matters when a session turns Comet or its native execution off after caching. +Comet records a fallback reason on such a scan. See +[In-Memory Cache](in-memory-cache.md#limitations). + +The format is chosen once, when the application starts. To keep Spark's format, set +`spark.comet.exec.inMemoryCache.enabled=false` then. Comet also keeps Spark's format without that +setting when the application: + +- starts with `spark.comet.enabled` or `spark.comet.exec.enabled` set to `false`. +- leaves Comet shuffle enabled without one of Comet's shuffle managers, so that Comet disables + itself. +- uses Kryo with `spark.kryo.registrationRequired=true` and does not list + `org.apache.comet.CometKryoRegistrator` in `spark.kryo.registrator`, because Kryo would reject + Comet's cached batches. Add the registrator to use Comet's format; see + [Kryo](in-memory-cache.md#kryo). + +An application that sets `spark.sql.cache.serializer` itself keeps the serializer it chose. + ## Upgrading to Comet 1.1.0 Comet `1.1.0` makes no behavior changes that need a `spark.comet.legacy.*` key. The changes below diff --git a/docs/source/user-guide/latest/operators.md b/docs/source/user-guide/latest/operators.md index a47055df7e0..db29377ff42 100644 --- a/docs/source/user-guide/latest/operators.md +++ b/docs/source/user-guide/latest/operators.md @@ -55,7 +55,7 @@ omitted from the tables below and may be reconsidered based on demand: | `BatchScanExec` | ✅ | Apache Iceberg Parquet scans run natively. Native CSV scans are experimental and disabled by default. DataSource V2 Parquet scans are not accelerated. See [Parquet Scan Compatibility](compatibility/scans.md) and the [Iceberg Guide](iceberg.md). | | `LocalTableScanExec` | ⚠️ | Disabled by default; there is no acceleration advantage and this operator is typically only used in test code. Can be opted into via config ([#4393](https://github.com/apache/datafusion-comet/pull/4393)). | | `EmptyRelationExec` | ✅ | Spark 4.0 and later. See [Empty Relations](compatibility/operators.md#empty-relations) for native-input support and writer fallback. | -| `InMemoryTableScanExec` | ⚠️ | Experimental, disabled by default. Set `spark.comet.exec.inMemoryCache.enabled=true` before the application starts so Comet installs its Arrow cache serializer. Relations with unsupported column types stay in Spark's cache format and fall back. See [In-Memory Cache](in-memory-cache.md). | +| `InMemoryTableScanExec` | ⚠️ | Experimental, enabled by default. Comet installs its Arrow cache serializer as the application starts, unless `spark.comet.exec.inMemoryCache.enabled` is false. Relations with unsupported column types stay in Spark's cache format and fall back. See [In-Memory Cache](in-memory-cache.md). | ## Projection and filtering From 6cf61b686ef545020dab9582824031b4824e1eca Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:01:24 -0600 Subject: [PATCH 17/24] test: benchmark the in-memory cache against Spark's format with AQE on The published comparisons either read the same Comet-written cache both ways, with spark.comet.sparkToColumnar.enabled turned on, or read both formats with Comet off. Neither is what an application gets from the feature's default. Add cases with Comet and AQE on and Comet's other settings at their defaults, reading a relation cached in Spark's format and in Comet's, with Comet operators above the cache scan and with a Spark operator above it. --- .../CometInMemoryCacheBenchmark.scala | 153 ++++++++++++++++-- 1 file changed, 140 insertions(+), 13 deletions(-) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index 9b0f652a19e..f21de0046cd 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -27,6 +27,8 @@ import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.comet.CometInMemoryTableScanExec import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanExec +import org.apache.spark.sql.execution.aggregate.HashAggregateExec import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, DefaultCachedBatchSerializer, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.vectorized.OnHeapColumnVector import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} @@ -236,6 +238,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { runCodecBenchmark(flatRelation) runSparkOperatorBenchmark(flatRelation) + runAdaptiveBenchmark(flatRelation) } } @@ -346,9 +349,9 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { * Reads that feed Spark operators rather than Comet ones, against Spark's own cache format. * * Comet is off in every case, so this measures Spark consuming the cached data: the shape where - * Comet's format has something to lose, and the reason the feature is off by default. Both - * formats are cached from the same relation, one copy at a time as in runCodecBenchmark, and - * each case checks which serializer cached the relation it reads. + * Comet's format has the most to lose. Both formats are cached from the same relation, one copy + * at a time as in runCodecBenchmark, and each case checks which serializer cached the relation + * it reads. */ private def runSparkOperatorBenchmark(relation: CachedRelation): Unit = { val view = s"${relation.table}_spark_operators" @@ -375,14 +378,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { cachedBy = serializer } - Seq( - ("row count only (0 of 6 columns)", s"SELECT count(*) FROM $view", 0), - ("narrow projection (1 of 6 columns)", s"SELECT count(k) FROM $view", 1), - ("3 of 6 columns", s"SELECT sum(id), sum(k), sum(v) FROM $view", 3), - ( - "full projection (6 of 6 columns)", - s"SELECT count(id), count(k), count(v), count(s1), count(s2), count(s3) FROM $view", - 6)).foreach { case (label, query, scanned) => + readShapes(view).foreach { case (label, query, scanned) => val benchmark = new Benchmark( s"in-memory cache read by Spark operators, $label", relation.rows.toLong, @@ -411,6 +407,127 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { } } + /** + * What the feature changes for a query that runs with Comet, against Spark's own cache format, + * with AQE on and Comet's other settings at their defaults. This is the comparison an + * application gets from turning the feature on or off. Both formats are cached from the same + * relation, one copy at a time as in runSparkOperatorBenchmark. + * + * Two shapes of plan read the cache. With Comet operators above the cache scan, Comet's format + * runs the whole query natively, while Spark's leaves the operators directly above its scan on + * Spark: Comet reads Spark's cache scan only through spark.comet.sparkToColumnar.enabled, which + * is off by default. With a Spark operator above the scan, Comet's format is read by the native + * scan and converted to rows for that operator, where Spark's is read by Spark's own scan. The + * Spark operator is the aggregate, with Comet's turned off, standing in for any operator Comet + * does not support. + */ + private def runAdaptiveBenchmark(relation: CachedRelation): Unit = { + val view = s"${relation.table}_adaptive" + val formats = Seq( + "Spark's cache format" -> classOf[DefaultCachedBatchSerializer].getName, + "Comet's cache format" -> classOf[ArrowCachedBatchSerializer].getName) + val operatorsAbove = Seq( + "Comet operators" -> Seq.empty[(String, String)], + "a Spark operator" -> Seq(CometConf.COMET_EXEC_AGGREGATE_ENABLED.key -> "false")) + + spark.catalog.clearCache() + withTempTable(view) { + spark + .sql(s"SELECT ${relation.columns.mkString(", ")} FROM ${relation.source}") + .createOrReplaceTempView(view) + + var cachedBy: String = null + def cacheBy(serializer: String): Unit = if (cachedBy != serializer) { + spark.catalog.uncacheTable(view) + cachedBy = null + withCacheSerializer(serializer) { + withSQLConf(adaptiveConf: _*) { + spark.catalog.cacheTable(view) + spark.table(view).count() + } + } + cachedBy = serializer + } + + for { + (operators, operatorConf) <- operatorsAbove + (label, query, scanned) <- readShapes(view) + } { + val benchmark = new Benchmark( + s"in-memory cache with AQE, $operators above the scan, $label", + relation.rows.toLong, + output = output) + formats.foreach { case (name, serializer) => + var verified = false + // Re-caching in this case's format is setup, so it is outside the timer, and it only + // happens on the case's first call, which is a warmup iteration. + benchmark.addTimerCase(name) { timer => + cacheBy(serializer) + withSQLConf(adaptiveConf ++ operatorConf: _*) { + if (!verified) { + verifyAdaptiveRead(query, scanned, serializer, operatorConf.nonEmpty) + verified = true + } + timer.startTiming() + spark.sql(query).noop() + timer.stopTiming() + } + } + } + benchmark.run() + } + + spark.catalog.uncacheTable(view) + } + } + + // The reads runSparkOperatorBenchmark and runAdaptiveBenchmark measure: no columns of the flat + // relation, one, three, and all six. + private def readShapes(view: String): Seq[(String, String, Int)] = Seq( + ("row count only (0 of 6 columns)", s"SELECT count(*) FROM $view", 0), + ("narrow projection (1 of 6 columns)", s"SELECT count(k) FROM $view", 1), + ("3 of 6 columns", s"SELECT sum(id), sum(k), sum(v) FROM $view", 3), + ( + "full projection (6 of 6 columns)", + s"SELECT count(id), count(k), count(v), count(s1), count(s2), count(s3) FROM $view", + 6)) + + // Pins what an adaptive case claims, in the plan AQE settles on, which it does only by running + // the query: one cache scan, native exactly when it reads Comet's format, reading the columns its + // label counts from a relation the named serializer cached. Spark aggregates run only where the + // case puts them: above Spark's scan, which nothing bridges into Comet, or wherever Comet's + // aggregate is turned off. + private def verifyAdaptiveRead( + query: String, + scanned: Int, + serializer: String, + sparkOperator: Boolean): Unit = { + val df = spark.sql(query) + df.collect() + val executed = df.queryExecution.executedPlan + val plan = executed.toString() + assert(executed.isInstanceOf[AdaptiveSparkPlanExec], s"Expected an adaptive plan:\n$plan") + + val nativeScans = collect(executed) { case s: CometInMemoryTableScanExec => s } + val sparkScans = collect(executed) { case s: InMemoryTableScanExec => s } + assert(nativeScans.length + sparkScans.length == 1, s"Expected exactly one cache scan:\n$plan") + val cometFormat = serializer == classOf[ArrowCachedBatchSerializer].getName + assert( + nativeScans.nonEmpty == cometFormat, + s"Expected a native scan exactly for Comet's format:\n$plan") + val (relation, columns) = nativeScans.headOption + .map(s => (s.originalPlan.relation, s.scanOutput.length)) + .getOrElse((sparkScans.head.relation, sparkScans.head.attributes.length)) + assert(columns == scanned, s"Expected the scan to read $scanned columns:\n$plan") + val actual = relation.cacheBuilder.serializer.getClass.getName + assert(actual == serializer, s"Expected a relation cached by $serializer, not $actual") + + val sparkAggregates = collect(executed) { case a: HashAggregateExec => a } + assert( + sparkAggregates.nonEmpty == (sparkOperator || !cometFormat), + s"Expected Spark aggregates only where this case puts them:\n$plan") + } + // spark.sql.cache.serializer is static, and InMemoryRelation memoizes the serializer it names // for the life of the JVM. It looks the name up in the active session's conf when a relation is // cached, though, so setting it there directly and clearing the memoized instance around one @@ -540,8 +657,8 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { // scan and convert" -- which is the overhead this feature exists to remove. // // Neither case is a baseline for Spark's own cache format: both read the same Comet-written - // CometCachedBatch. Spark's format is only measured by runSparkOperatorBenchmark, with Comet - // off, since that is the only comparison it answers. + // CometCachedBatch. Spark's format is measured by runSparkOperatorBenchmark, with Comet off, + // and by runAdaptiveBenchmark, with Comet on. withSQLConf(cacheConf(nativeCacheEnabled = true): _*) { spark .sql(s"SELECT ${relation.columns.mkString(", ")} FROM ${relation.source}") @@ -611,4 +728,14 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { CometConf.COMET_ENABLED.key -> "false", CometConf.COMET_EXEC_ENABLED.key -> "false", "spark.sql.inMemoryColumnarStorage.batchSize" -> "10000") + + // Comet and AQE on, and Comet's other settings at their defaults, unlike cacheConf. The batch + // size matches the other confs, and on-heap mode is what lets Comet run in this session. + private val adaptiveConf: Seq[(String, String)] = Seq( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", + "spark.comet.exec.onHeap.enabled" -> "true", + "spark.sql.inMemoryColumnarStorage.batchSize" -> "10000") } From d124b1f3ccc45627cda9ca1e646c9d30fbb8d123 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:03:30 -0600 Subject: [PATCH 18/24] feat: record why Spark scans a relation cached in Comet's format when Comet is off spark.sql.cache.serializer is static, so a relation cached in Comet's format keeps it for the life of the application. A session that turns Comet or its native execution off then reads it with Spark's InMemoryTableScanExec, which is slower than reading Spark's own format (#5485). CometExecRule records a fallback reason for the other ways Spark ends up scanning a cache, but it returned before looking at the plan in these two cases, so nothing explained them. Record a reason on each such scan, naming the cause and the startup setting that keeps caches in Spark's format. --- .../apache/comet/rules/CometExecRule.scala | 39 +++++++++++++-- .../comet/exec/CometInMemoryCacheSuite.scala | 48 +++++++++++++++++++ 2 files changed, 82 insertions(+), 5 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index c1aa237970a..aa6a07fa077 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -62,6 +62,15 @@ import org.apache.comet.shims.{CometTypeShim, ShimCometStreaming, ShimCometWindo object CometExecRule { + /** + * Whether `scan` reads a relation stored in Comet's cache format. Comet's serializer stores + * that format only for schemas it supports and delegates everything else to Spark's default + * cache format, which the native scan cannot read. + */ + private[rules] def readsCometCacheFormat(scan: InMemoryTableScanExec): Boolean = + scan.relation.cacheBuilder.serializer.isInstanceOf[ArrowCachedBatchSerializer] && + ArrowCachedBatchSerializer.supportsSchema(scan.relation.output) + private[rules] def removePlaceholders(plan: SparkPlan): SparkPlan = plan.transformUp { // revertUnsafePartialAggregates re-runs transform over already wrapped query stages, which // can produce CometSinkPlaceHolder(CometSinkPlaceHolder(stage)). Remove sinks bottom-up. @@ -381,10 +390,7 @@ case class CometExecRule(session: SparkSession) case scan: InMemoryTableScanExec => val serializer = scan.relation.cacheBuilder.serializer val usesCometCacheSerializer = serializer.isInstanceOf[ArrowCachedBatchSerializer] - // The serializer only stores Comet's Arrow format for schemas it supports and delegates - // everything else to Spark's default cache format, which the native scan cannot read. - val cometCacheFormat = usesCometCacheSerializer && - ArrowCachedBatchSerializer.supportsSchema(scan.relation.output) + val cometCacheFormat = CometExecRule.readsCometCacheFormat(scan) val nativeCacheEnabled = CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.get(conf) // Walks the cached plan, so it is lazy: only consulted once the native scan is otherwise // possible. See CometInMemoryTableScanExec.recordsObservedMetrics. @@ -756,6 +762,25 @@ case class CometExecRule(session: SparkSession) } } + /** + * A relation keeps the cache format it was stored in, since `spark.sql.cache.serializer` is + * static, so a plan that runs without Comet's native execution still reads relations cached in + * Comet's format. Spark's `InMemoryTableScanExec` reads that format more slowly than Spark's + * own (https://github.com/apache/datafusion-comet/issues/5485), and nothing else records a + * fallback reason in such a plan, so record one on each scan that does. + */ + private def explainSparkReadsOfCometCache(plan: SparkPlan, cause: String): Unit = + plan.foreach { + case scan: InMemoryTableScanExec if CometExecRule.readsCometCacheFormat(scan) => + withFallbackReason( + scan, + s"$cause, so Spark reads this relation from Comet's cache format, which is slower " + + "than reading Spark's own. Set " + + s"${CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key}=false when the application " + + "starts to cache in Spark's format instead.") + case _ => + } + override def apply(plan: SparkPlan): SparkPlan = { val newPlan = _apply(plan) if (showTransformations && !newPlan.fastEquals(plan)) { @@ -769,13 +794,17 @@ case class CometExecRule(session: SparkSession) private def _apply(plan: SparkPlan): SparkPlan = { // We shouldn't transform Spark query plan if Comet is not loaded. - if (!isCometLoaded(conf)) return plan + if (!isCometLoaded(conf)) { + explainSparkReadsOfCometCache(plan, "Comet is disabled") + return plan + } // Comet does not support structured streaming. Fall back to Spark for any plan that // belongs to a streaming query (detected via StreamSourceAwareSparkPlan.getStream). if (ShimCometStreaming.isStreamingPlan(plan)) return plan if (!CometConf.COMET_EXEC_ENABLED.get(conf)) { + explainSparkReadsOfCometCache(plan, s"${CometConf.COMET_EXEC_ENABLED.key} is false") // Comet exec is disabled, but for Spark shuffle, we still can use Comet columnar shuffle if (isCometShuffleEnabled(conf)) { applyCometShuffle(plan) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index ecce9067522..b2c3764fd6d 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -474,6 +474,54 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + test("Comet explains Spark's scan of a relation cached in Comet's format") { + // spark.sql.cache.serializer is static, so a relation cached in Comet's format stays in it + // after a session turns Comet or its native execution off, and from then on Spark's + // InMemoryTableScanExec reads it, which nothing else in the plan would record. A relation + // that Comet's serializer delegated to Spark's format gets no such reason. + withNativeCache { + spark + .sql("SELECT id, id % 7 AS k FROM range(100)") + .createOrReplaceTempView("comet_format_cache") + spark + .sql(s"SELECT id, ${unsupportedForArrowCache.head} FROM range(100)") + .createOrReplaceTempView("spark_format_cache") + spark.catalog.cacheTable("comet_format_cache") + spark.catalog.cacheTable("spark_format_cache") + assert( + cachedBatchTypes("comet_format_cache").sameElements( + Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) + assert( + cachedBatchTypes("spark_format_cache").sameElements( + Array("org.apache.spark.sql.execution.columnar.DefaultCachedBatch"))) + + def reasons(query: String): Seq[String] = { + val df = spark.sql(query) + df.collect() + new ExtendedExplainInfo().getFallbackReasons(df.queryExecution.executedPlan) + } + + for { + (key, cause) <- Seq( + CometConf.COMET_ENABLED.key -> "Comet is disabled", + CometConf.COMET_EXEC_ENABLED.key -> s"${CometConf.COMET_EXEC_ENABLED.key} is false") + aqe <- Seq("false", "true") + } { + withSQLConf(key -> "false", SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + val explained = reasons("SELECT k, count(*) FROM comet_format_cache GROUP BY k") + assert( + explained.exists( + _.startsWith(s"$cause, so Spark reads this relation from Comet's cache format")), + s"$key=false, AQE $aqe: $explained") + assert( + !reasons("SELECT count(id) FROM spark_format_cache").exists( + _.contains("Comet's cache format")), + s"$key=false, AQE $aqe") + } + } + } + } + test("Comet in-memory cache handles multi-partition cache") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", From 73326f86925c1c5ae092a664c56d6f46285edc2a Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:08:11 -0600 Subject: [PATCH 19/24] docs: say an application keeps the cache format it started with The sentence described a session started with the default, which keeps Spark's format only while the default is off. --- docs/source/user-guide/latest/in-memory-cache.md | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 52a5f62f2b8..63191a815ad 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -34,13 +34,13 @@ $SPARK_HOME/bin/spark-shell \ ``` It has to be set before the `SparkContext` starts. Comet's driver plugin chooses -`spark.sql.cache.serializer` once, while the context is initializing, so a session that started -with the default goes on using Spark's cache format however the config is set afterwards. The -plugin installs Comet's serializer only if `spark.comet.enabled` and `spark.comet.exec.enabled` -are enabled at that point too, because an application that starts without native execution could -not scan Comet's format natively. It also keeps Spark's format when Comet shuffle is enabled but -`spark.shuffle.manager` is not one of Comet's shuffle managers, since Comet then disables itself, -and when Kryo would reject Comet's format; see [Kryo](#kryo). +`spark.sql.cache.serializer` once, while the context is initializing, so an application keeps the +cache format it started with however the config is set afterwards. The plugin installs Comet's +serializer only if `spark.comet.enabled` and `spark.comet.exec.enabled` are enabled at that point +too, because an application that starts without native execution could not scan Comet's format +natively. It also keeps Spark's format when Comet shuffle is enabled but `spark.shuffle.manager` is +not one of Comet's shuffle managers, since Comet then disables itself, and when Kryo would reject +Comet's format; see [Kryo](#kryo). ## What changes when it is enabled From 972c47f468b1551358c37085fd4edd8d87c037fd Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:10:03 -0600 Subject: [PATCH 20/24] test: format the in-memory cache benchmark --- .../spark/sql/benchmark/CometInMemoryCacheBenchmark.scala | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index f21de0046cd..8c8bbcfbe17 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -510,7 +510,9 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { val nativeScans = collect(executed) { case s: CometInMemoryTableScanExec => s } val sparkScans = collect(executed) { case s: InMemoryTableScanExec => s } - assert(nativeScans.length + sparkScans.length == 1, s"Expected exactly one cache scan:\n$plan") + assert( + nativeScans.length + sparkScans.length == 1, + s"Expected exactly one cache scan:\n$plan") val cometFormat = serializer == classOf[ArrowCachedBatchSerializer].getName assert( nativeScans.nonEmpty == cometFormat, From 5528fc46006b5edd9a3598fdcdca0b33ecf06299 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:13:13 -0600 Subject: [PATCH 21/24] fix: bind the fallback reason result for the strict Scala warnings build -Ywarn-value-discard rejects the node that withFallbackReason returns, discarded as the last expression of a case. --- spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index aa6a07fa077..ce354aed339 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -772,7 +772,7 @@ case class CometExecRule(session: SparkSession) private def explainSparkReadsOfCometCache(plan: SparkPlan, cause: String): Unit = plan.foreach { case scan: InMemoryTableScanExec if CometExecRule.readsCometCacheFormat(scan) => - withFallbackReason( + val _ = withFallbackReason( scan, s"$cause, so Spark reads this relation from Comet's cache format, which is slower " + "than reading Spark's own. Set " + From b9ab5927c90976de3f0b6c6558586d50e3d00fe9 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 05:23:39 -0600 Subject: [PATCH 22/24] docs: compare the in-memory cache with Spark's format as an application gets it Add the benchmark's adaptive results to the cache guide: with Comet and AQE on and Comet's other settings at their defaults, Comet's format is as fast or faster than Spark's in every shape but a read of three long columns, which zstd decompression makes about 10% slower. That includes a Spark operator above the native scan. The guide said that, without the feature, a CometSparkColumnarToColumnar converts each batch of Spark's cache scan for Comet. That takes spark.comet.sparkToColumnar.enabled, which is off by default, so by default the operators above Spark's scan run on Spark. The limitation that Spark operators read Comet's format slowly applies to Spark's own cache scan, not to Spark operators above Comet's. --- .../user-guide/latest/in-memory-cache.md | 53 ++++++++++++++++--- .../user-guide/latest/migration-guide.md | 6 +-- 2 files changed, 48 insertions(+), 11 deletions(-) diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 87b094c275a..1b6b642e6b8 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -21,8 +21,9 @@ Comet can store Spark's in-memory cache (`CACHE TABLE`, `df.cache()`, `df.persist()`) in an Arrow format that Comet operators read directly. Without it, a cached table is stored in Spark's own -format and every scan of it has to convert each batch before Comet can continue, which shows up in -the plan as a `CometSparkColumnarToColumnar` above the cache scan. +format, which Comet operators cannot read. Under Comet's default settings the operators above the +cache scan then run on Spark. With `spark.comet.sparkToColumnar.enabled`, a +`CometSparkColumnarToColumnar` above the scan converts each batch for Comet operators instead. This feature is **experimental and enabled by default**. To turn it off, set the config at startup, alongside the rest of Comet's configuration: @@ -160,10 +161,43 @@ back to Spark row execution above the scan and the two columns stop measuring th Read what this compares carefully. Comet execution is on in both columns, so the aggregation runs on Comet either way and only the cache-scan boundary moves: on the left, Spark's `InMemoryTableScanExec` feeds those same Comet operators through a `CometSparkColumnarToColumnar` -bridge; on the right, `CometInMemoryTableScan` feeds them directly. Both columns read the same +bridge, which the benchmark turns on with `spark.comet.sparkToColumnar.enabled`; on the right, +`CometInMemoryTableScan` feeds them directly. Both columns read the same Comet-written `CometCachedBatch`. These numbers are therefore "keep the cached scan native" against "fall back to a Spark cache scan and convert", not Comet against Spark execution, and not a -comparison with Spark's own cache format. That comparison is under [Limitations](#limitations). +comparison with Spark's own cache format, which follows. + +### Against Spark's cache format + +What turning the feature on changes for a query that Comet runs is measured against Spark's own +cache format by the benchmark's adaptive cases. Comet and AQE are on, Comet's other settings are at +their defaults, and the same 5M-row relation is cached in each format. The defaults leave +`spark.comet.sparkToColumnar.enabled` off, so Comet operators cannot read Spark's cache scan, and +with Spark's format the operators directly above the scan run on Spark. Measured on an AMD Ryzen 9 +7950X3D (JDK 17, Spark 4.1, release build): + +| Query shape | Spark's cache format | Comet's cache format | Relative | +| -------------------------- | -------------------: | -------------------: | -------: | +| Row count only (0 of 6) | 29 ms | 24 ms | 1.2x | +| Narrow projection (1 of 6) | 52 ms | 34 ms | 1.5x | +| 3 of 6 columns | 102 ms | 112 ms | 0.9x | +| Full projection (6 of 6) | 299 ms | 224 ms | 1.3x | + +A Spark operator above the cache scan, standing in for any operator Comet does not support, is +measured the same way, with Comet's aggregate turned off. With Comet's format, the native scan feeds +that operator through a columnar-to-row transition: + +| Query shape | Spark's cache format | Comet's cache format | Relative | +| -------------------------- | -------------------: | -------------------: | -------: | +| Row count only (0 of 6) | 39 ms | 16 ms | 2.4x | +| Narrow projection (1 of 6) | 50 ms | 27 ms | 1.8x | +| 3 of 6 columns | 97 ms | 112 ms | 0.9x | +| Full projection (6 of 6) | 303 ms | 299 ms | 1.0x | + +Comet's format is as fast or faster in every shape but one: the read of three of the six columns, +all of them longs, is about 10% slower under either kind of operator. That cost is `zstd` +decompression. With the `none` codec, the same read is 2.7x faster than Spark's format with Comet +operators above the scan, and 1.6x faster with a Spark operator above it. ## Kryo @@ -194,9 +228,11 @@ Spark's cached batch too. ## Limitations -Reads that feed **Spark** operators rather than Comet ones are slower than Spark's own cache -format, and the narrower the read, the wider the gap. Measured by the same benchmark over the same -5M-row relation, with Comet off so that Spark operators consume the cached data: +Spark's own cache scan, `InMemoryTableScanExec`, reads Comet's format more slowly than Spark's, and +the narrower the read, the wider the gap. Spark's scan reads a cached relation when a session turns +Comet or its native execution off after caching, and when the relation's cached plan records +`Dataset.observe` metrics, and Comet records a fallback reason on the scan in either case. Measured +by the same benchmark over the same 5M-row relation, with Comet off: | Read shape | Spark's cache format | Comet's cache format | Slowdown | | ----------------------- | -------------------: | -------------------: | -------: | @@ -205,7 +241,8 @@ format, and the narrower the read, the wider the gap. Measured by the same bench | 3 of 6 columns | 98 ms | 331 ms | 3.4x | | 6 of 6 columns | 410 ms | 623 ms | 1.5x | -This is the main reason the feature is still described as experimental. The cause is not yet +A Spark operator above Comet's native cache scan does not pay this; see [Performance](#performance). +This gap is the main reason the feature is still described as experimental. The cause is not yet established; [#5485](https://github.com/apache/datafusion-comet/issues/5485) tracks it. Comet's serializer exists because Spark's own Arrow cache format diff --git a/docs/source/user-guide/latest/migration-guide.md b/docs/source/user-guide/latest/migration-guide.md index e634dafa729..fec4e92c2f5 100644 --- a/docs/source/user-guide/latest/migration-guide.md +++ b/docs/source/user-guide/latest/migration-guide.md @@ -67,9 +67,9 @@ need none either, but check whether any of them applies to your deployment. `spark.comet.exec.inMemoryCache.enabled` now defaults to `true`. An application that loads `CometPlugin` now stores what it caches with `CACHE TABLE`, `df.cache()` or `df.persist()` in Comet's Arrow format instead of Spark's, and Comet scans it natively. The format does not change -query results, but it can change performance: Spark operators read Comet's format more slowly than -Spark's own, which matters when a session turns Comet or its native execution off after caching. -Comet records a fallback reason on such a scan. See +query results, but it can change performance. Spark's own cache scan reads Comet's format more +slowly than Spark's, which matters when a session turns Comet or its native execution off after +caching, and Comet records a fallback reason on such a scan. See [In-Memory Cache](in-memory-cache.md#limitations). The format is chosen once, when the application starts. To keep Spark's format, set From 667ea36256fa8ef662f1b586a04fd9fcfa863ddd Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 08:43:36 -0600 Subject: [PATCH 23/24] fix: count Kryo registrations however the application made them The gate kept Spark's cache format whenever spark.kryo.registrator did not list CometKryoRegistrator, even where the application had registered Comet's cached batch another way, such as spark.kryo.classesToRegister. Comet's format works there, while Spark's does not before 4.1, which registers DefaultCachedBatch itself only from then on, so the gate broke caches that worked. Ask a Kryo instance built from the application's conf which of Comet's classes it has registered, and keep Spark's format only if Comet's cached batch is not among them. If Kryo cannot be built, fall back to looking for CometKryoRegistrator in spark.kryo.registrator. The startup warning uses the same check and names the missing classes. --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../user-guide/latest/in-memory-cache.md | 13 ++-- docs/source/user-guide/latest/installation.md | 12 ++-- .../scala/org/apache/comet/CometConf.scala | 9 +-- .../main/scala/org/apache/spark/Plugins.scala | 65 ++++++++++++------- .../arrow/ArrowCachedBatchSerializer.scala | 5 ++ .../exec/CometInMemoryCacheKryoSuite.scala | 63 ++++++++++++------ .../comet/exec/CometInMemoryCacheSuite.scala | 45 ++++++++++++- 9 files changed, 154 insertions(+), 60 deletions(-) diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 0f28229ae8b..f6b85808f2d 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -516,6 +516,7 @@ jobs: org.apache.comet.exec.CometInMemoryCachePruningSuite org.apache.comet.exec.CometInMemoryCacheKryoSuite org.apache.comet.exec.CometInMemoryCacheKryoUnregisteredSuite + org.apache.comet.exec.CometInMemoryCacheKryoClassesToRegisterSuite org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometJoinSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 7838bf14d5a..7d6b1ac9fc9 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -220,6 +220,7 @@ jobs: org.apache.comet.exec.CometInMemoryCachePruningSuite org.apache.comet.exec.CometInMemoryCacheKryoSuite org.apache.comet.exec.CometInMemoryCacheKryoUnregisteredSuite + org.apache.comet.exec.CometInMemoryCacheKryoClassesToRegisterSuite org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometJoinSuite diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index 63191a815ad..fb3ba9cc632 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -182,11 +182,14 @@ spark.kryo.registrator=org.apache.comet.CometKryoRegistrator Comet cannot set `spark.kryo.registrator` for you the way it sets `spark.sql.cache.serializer`: `KryoSerializer` reads it when `SparkEnv` builds the serializer, which happens before any plugin runs. Without it, Kryo would reject Comet's cached batch with a "Class is not registered" error -that does not name this feature, so Comet's driver plugin does not install Comet's serializer, -and caches stay in Spark's format. The plugin warns at startup when it sees Kryo, -`registrationRequired`, and no registrator. An application that sets `spark.sql.cache.serializer` -to Comet's serializer itself gets the error instead. Native broadcast needs the same registrator -even when the cache is disabled; see [Kryo serialization](installation.md#kryo-serialization). +that does not name this feature. So when Kryo requires registration and has not registered +Comet's cached batch, Comet's driver plugin does not install Comet's serializer, and caches stay in +Spark's format. Registrations made another way, through a registrator of the application's own or +`spark.kryo.classesToRegister`, count as well. The plugin warns at startup when Kryo requires +registration and has not registered every class `CometKryoRegistrator` registers. An application +that sets `spark.sql.cache.serializer` to Comet's serializer itself gets the error instead. Native +broadcast needs the same registrator even when the cache is disabled; see +[Kryo serialization](installation.md#kryo-serialization). Spark registers its own cached batch with Kryo only from Spark 4.1, so on earlier versions caching in either format under `registrationRequired` needs a registrator. `CometKryoRegistrator` registers diff --git a/docs/source/user-guide/latest/installation.md b/docs/source/user-guide/latest/installation.md index 273ea664772..c6ca172ab06 100644 --- a/docs/source/user-guide/latest/installation.md +++ b/docs/source/user-guide/latest/installation.md @@ -269,9 +269,9 @@ If the application uses Kryo (`spark.serializer=org.apache.spark.serializer.Kryo Without it, any query that uses Comet's native broadcast exchange, which is enabled by default, fails with Kryo's "Class is not registered" error, for example on the first broadcast hash join. -Comet's [in-memory cache](in-memory-cache.md#kryo) format needs the same registrator, and without -it Comet's plugin keeps caches in Spark's format. Set it before the -`SparkContext` is created: `KryoSerializer` reads it before Comet's plugin runs, so Comet cannot -add it for you. `spark.kryo.registrator` accepts a comma-separated list, so an application with -its own registrator can list both. Comet logs a warning at startup when Kryo requires registration -and this registrator is missing. +Comet's [in-memory cache](in-memory-cache.md#kryo) format needs the same registrations, and while +Kryo has not registered Comet's cached batch, Comet's plugin keeps caches in Spark's format. Set it +before the `SparkContext` is created: `KryoSerializer` reads it before Comet's plugin runs, so +Comet cannot add it for you. `spark.kryo.registrator` accepts a comma-separated list, so an +application with its own registrator can list both. Comet logs a warning at startup when Kryo +requires registration and has not registered the classes this registrator covers. diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index ced5860d685..985b18b9cb1 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -286,10 +286,11 @@ object CometConf extends ShimCometConf { "so the unselected ones are never decompressed. Reads that feed Spark operators rather " + "than Comet ones still pay a row conversion the default format avoids, and can be " + "slower than Spark's cache. With spark.kryo.registrationRequired=true, the plugin " + - "installs it only if spark.kryo.registrator includes " + - "org.apache.comet.CometKryoRegistrator, which has to be set before creating the " + - "SparkContext, because Kryo would otherwise reject a cached block as soon as it is " + - "serialized, including the disk half of the default MEMORY_AND_DISK storage level.") + "installs it only if Kryo has registered Comet's cached batch, as " + + "spark.kryo.registrator=org.apache.comet.CometKryoRegistrator does when set before " + + "creating the SparkContext, because Kryo would otherwise reject a cached block as soon " + + "as it is serialized, including the disk half of the default MEMORY_AND_DISK storage " + + "level.") .booleanConf .createWithDefault(false) diff --git a/spark/src/main/scala/org/apache/spark/Plugins.scala b/spark/src/main/scala/org/apache/spark/Plugins.scala index c420b905ae6..e79074f45d2 100644 --- a/spark/src/main/scala/org/apache/spark/Plugins.scala +++ b/spark/src/main/scala/org/apache/spark/Plugins.scala @@ -30,6 +30,8 @@ import org.apache.spark.api.plugin.{DriverPlugin, ExecutorPlugin, PluginContext, import org.apache.spark.internal.Logging import org.apache.spark.internal.config.{EVENT_LOG_ENABLED, EXECUTOR_MEMORY_OVERHEAD, EXECUTOR_MEMORY_OVERHEAD_FACTOR} import org.apache.spark.scheduler.{SparkListener, SparkListenerApplicationEnd, SparkListenerExecutorMetricsUpdate, SparkListenerExecutorRemoved} +import org.apache.spark.serializer.KryoSerializer +import org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer import org.apache.spark.sql.comet.execution.shuffle.{CometCelebornShuffleManager, CometShuffleManager} import org.apache.spark.sql.internal.StaticSQLConf import org.apache.spark.util.{Clock, SystemClock} @@ -111,7 +113,7 @@ class CometDriverPlugin private[spark] (clock: Clock) extends DriverPlugin with val extraConfs = new ju.HashMap[String, String]() CometDriverPlugin.maybeSetCacheSerializer(sc.conf, extraConfs) - CometDriverPlugin.warnIfKryoRegistratorMissing(sc.conf) + CometDriverPlugin.warnIfKryoRegistrationsMissing(sc.conf) // register CometSparkSessionExtensions if it isn't already registered CometDriverPlugin.registerCometSessionExtension(sc.conf) @@ -178,9 +180,10 @@ object CometDriverPlugin extends Logging { // store every cache in Comet's format, with only Spark operators to read it. So would one that // leaves Comet shuffle enabled without Comet's shuffle manager, since Comet then disables // itself. - // Nor is it used where Kryo would reject its cached batches, under - // spark.kryo.registrationRequired without CometKryoRegistrator: caching that works in Spark's - // format would then fail the first time Spark serialized a cached block. + // Nor is it used where Kryo requires registration and has not registered Comet's cached batch: + // caching would then fail the first time Spark serialized a cached block. Where Kryo has + // registered it, by whatever means, Comet's format is used, since Spark registers its own + // cached batch only from 4.1. // If the application already set spark.sql.cache.serializer, leave that value // unchanged so Comet does not replace a user-selected cache format. private[apache] def maybeSetCacheSerializer( @@ -190,7 +193,7 @@ object CometDriverPlugin extends Logging { getBooleanConf(conf, CometConf.COMET_EXEC_ENABLED) && getBooleanConf(conf, CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED) && (!getBooleanConf(conf, CometConf.COMET_SHUFFLE_ENABLED) || isCometShuffleManager(conf)) && - !isKryoRegistratorMissing(conf)) { + !unregisteredKryoClasses(conf).contains(ArrowCachedBatchSerializer.cachedBatchClass)) { val serializerKey = StaticSQLConf.SPARK_CACHE_SERIALIZER.key val serializerValue = "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer" @@ -215,28 +218,46 @@ object CometDriverPlugin extends Logging { // CometKryoRegistrator covers both, but spark.kryo.registrator is read when SparkEnv builds the // serializer, before any plugin runs, so it cannot be set from here. Say so while the // application is still starting up rather than leaving the user to attribute the failure later. - private[apache] def warnIfKryoRegistratorMissing(conf: SparkConf): Unit = { - if (isKryoRegistratorMissing(conf)) { - logWarning( - "spark.kryo.registrationRequired=true but spark.kryo.registrator does not include " + - s"${CometKryoRegistrator.CLASS_NAME}. Comet's native broadcast will fail with " + - "Kryo's \"Class is not registered\" as soon as its payload is serialized, and " + - "Comet does not install its in-memory cache format, which would fail the same way. " + - s"Add spark.kryo.registrator=${CometKryoRegistrator.CLASS_NAME} before creating the " + - "SparkContext; it cannot be set later.") + private[apache] def warnIfKryoRegistrationsMissing(conf: SparkConf): Unit = { + val unregistered = unregisteredKryoClasses(conf) + if (unregistered.nonEmpty) { + logWarning("spark.kryo.registrationRequired=true but Kryo has not registered " + + s"${unregistered.map(_.getName).mkString(", ")}, which " + + s"${CometKryoRegistrator.CLASS_NAME} registers. Comet's native broadcast and in-memory " + + "cache fail with Kryo's \"Class is not registered\" when they serialize one of them, " + + "and Comet keeps Spark's cache format while its own cached batch is unregistered. " + + s"Add spark.kryo.registrator=${CometKryoRegistrator.CLASS_NAME} before creating the " + + "SparkContext; it cannot be set later.") } } - private def isKryoRegistratorMissing(conf: SparkConf): Boolean = { + // The classes CometKryoRegistrator registers that Kryo, configured as the application + // configured it, would reject: none unless it requires registration. They can be registered + // through CometKryoRegistrator, a registrator of the application's own or + // spark.kryo.classesToRegister, so ask a Kryo instance built from the conf rather than read the + // confs. If one cannot be built, take them as registered only if spark.kryo.registrator lists + // CometKryoRegistrator. + private[apache] def unregisteredKryoClasses(conf: SparkConf): Seq[Class[_]] = { val usingKryo = conf.get("spark.serializer", "") == "org.apache.spark.serializer.KryoSerializer" - val registrationRequired = conf.getBoolean("spark.kryo.registrationRequired", false) - val registered = conf - .get("spark.kryo.registrator", "") - .split(',') - .map(_.trim) - .contains(CometKryoRegistrator.CLASS_NAME) - usingKryo && registrationRequired && !registered + if (!usingKryo || !conf.getBoolean("spark.kryo.registrationRequired", false)) { + Nil + } else { + // Qualified, because in this package org.apache.spark.Success, a TaskEndReason, hides an + // imported scala.util.Success on Scala 2.12. + Try(new KryoSerializer(conf).newKryo()) match { + case scala.util.Success(kryo) => + CometKryoRegistrator.classes.filter(kryo.getClassResolver.getRegistration(_) == null) + case scala.util.Failure(e) => + logDebug("Could not build Kryo to check Comet's registrations", e) + val listed = conf + .get("spark.kryo.registrator", "") + .split(',') + .map(_.trim) + .contains(CometKryoRegistrator.CLASS_NAME) + if (listed) Nil else CometKryoRegistrator.classes + } + } } // Comet's shuffle managers have no short name, so spark.shuffle.manager names one only by its diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 103df61e17f..853c731a94a 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -739,6 +739,11 @@ object ArrowCachedBatchSerializer { def supportsSchema(schema: Seq[Attribute]): Boolean = schema.forall(a => supportsType(a.dataType)) + /** + * The class of Comet's cached batch, which Kryo has to have registered to store this format. + */ + private[apache] val cachedBatchClass: Class[_] = classOf[CometCachedBatch] + /** * The classes a `CometCachedBatch` adds on top of [[org.apache.comet.CometKryoRegistrator]]'s * shared Arrow-bytes classes. diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala index ea7d2b38307..4518f8ab2ee 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala @@ -23,8 +23,9 @@ import org.apache.spark.SparkConf import org.apache.spark.serializer.KryoRegistrator import org.apache.spark.sql.{CometTestBase, Row} import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, DefaultCachedBatch} -import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.storage.StorageLevel import org.apache.spark.unsafe.types.UTF8String @@ -201,16 +202,20 @@ class CometInMemoryCacheKryoSuite extends CometTestBase { } /** - * The application that `CometDriverPlugin` warns about: Kryo with registration required and no - * [[CometKryoRegistrator]]. The plugin then leaves `spark.sql.cache.serializer` alone rather than - * install Comet's serializer, whose cached batch Kryo would reject, so caching works as it does - * without Comet. + * Comet's driver plugin under Kryo with registration required, in an application that does not + * list [[CometKryoRegistrator]]. The plugin installs Comet's cache serializer only if Kryo has + * Comet's cached batch registered, whatever registered it, so the format it picks has to survive + * a `DISK_ONLY` cache, which serializes every block as it is put. * * Spark registers its own cached batch with Kryo only from 4.1, so on earlier versions an - * application whose caches work under registration has to register it itself. The registrator - * this suite installs does that, and nothing of Comet's. + * application whose caches work under registration registers it itself, as the suite that expects + * Spark's format does. */ -class CometInMemoryCacheKryoUnregisteredSuite extends CometTestBase { +abstract class CometInMemoryCacheKryoRegistrationSuite(expectedBatch: String) + extends CometTestBase { + + /** The Kryo registrations the application makes instead of listing CometKryoRegistrator. */ + protected def registrations: Seq[(String, String)] override protected def beforeAll(): Unit = { CometInMemoryRelationHelper.clearSerializer() @@ -231,8 +236,7 @@ class CometInMemoryCacheKryoUnregisteredSuite extends CometTestBase { conf.set(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key, "true") conf.set("spark.serializer", "org.apache.spark.serializer.KryoSerializer") conf.set("spark.kryo.registrationRequired", "true") - conf.set("spark.kryo.registrator", classOf[SparkCachedBatchKryoRegistrator].getName) - conf + conf.setAll(registrations) } private def cachedBatchTypes(table: String): Array[String] = { @@ -243,26 +247,21 @@ class CometInMemoryCacheKryoUnregisteredSuite extends CometTestBase { .collect() } - test("Comet plugin keeps Spark's cache format when Kryo would reject Comet's") { - assert(!spark.sparkContext.getConf.contains(StaticSQLConf.SPARK_CACHE_SERIALIZER.key)) - + test("Comet plugin picks a cache format that Kryo can store") { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { spark.catalog.clearCache() try { spark .range(0, 100, 1, 2) .selectExpr("id", "cast(id as string) AS s") - .createOrReplaceTempView("kryo_unregistered") + .createOrReplaceTempView("kryo_registration") - // DISK_ONLY serializes every block as it is put, so this reaches Kryo in local mode. - spark.catalog.cacheTable("kryo_unregistered", StorageLevel.DISK_ONLY) - assert(spark.table("kryo_unregistered").count() == 100) - assert( - cachedBatchTypes("kryo_unregistered").sameElements( - Array("org.apache.spark.sql.execution.columnar.DefaultCachedBatch"))) + spark.catalog.cacheTable("kryo_registration", StorageLevel.DISK_ONLY) + assert(spark.table("kryo_registration").count() == 100) + assert(cachedBatchTypes("kryo_registration").sameElements(Array(expectedBatch))) checkAnswer( - spark.sql("SELECT s FROM kryo_unregistered WHERE id > 97"), + spark.sql("SELECT s FROM kryo_registration WHERE id > 97"), Seq(Row("98"), Row("99"))) } finally { spark.catalog.clearCache() @@ -271,6 +270,28 @@ class CometInMemoryCacheKryoUnregisteredSuite extends CometTestBase { } } +/** Registers Spark's cached batch and nothing of Comet's, so the plugin keeps Spark's format. */ +class CometInMemoryCacheKryoUnregisteredSuite + extends CometInMemoryCacheKryoRegistrationSuite(classOf[DefaultCachedBatch].getName) { + override protected def registrations: Seq[(String, String)] = + Seq("spark.kryo.registrator" -> classOf[SparkCachedBatchKryoRegistrator].getName) +} + +/** + * Registers Comet's classes through `spark.kryo.classesToRegister` rather than + * [[CometKryoRegistrator]], and not Spark's cached batch, so before Spark 4.1 only Comet's format + * can be stored, and the plugin has to install it. + */ +class CometInMemoryCacheKryoClassesToRegisterSuite + extends CometInMemoryCacheKryoRegistrationSuite( + ArrowCachedBatchSerializer.cachedBatchClass.getName) { + override protected def registrations: Seq[(String, String)] = Seq( + "spark.kryo.classesToRegister" -> CometKryoRegistrator.classes + .filterNot(_ == classOf[DefaultCachedBatch]) + .map(_.getName) + .mkString(",")) +} + /** Registers what Spark's own cached batch needs, over a long and a string column, with Kryo. */ class SparkCachedBatchKryoRegistrator extends KryoRegistrator { override def registerClasses(kryo: Kryo): Unit = { diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 488e57b7a97..e913a21a69b 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -1103,12 +1103,28 @@ class CometInMemoryCacheSuite extends CometTestBase { assert(installed(CometConf.COMET_SHUFFLE_ENABLED.key -> "false")) assert(installed("spark.comet.exec.shuffle.enabled" -> "false")) - // Kryo with registration required rejects Comet's cached batch unless the application lists - // CometKryoRegistrator, on its own or beside a registrator of its own. + // Kryo with registration required rejects Comet's cached batch unless something registered + // it: CometKryoRegistrator, on its own or beside a registrator of the application's, or the + // application's own registrations. val kryo = "spark.serializer" -> "org.apache.spark.serializer.KryoSerializer" val registrationRequired = "spark.kryo.registrationRequired" -> "true" val registrator = "spark.kryo.registrator" + val sparkOnly = classOf[SparkCachedBatchKryoRegistrator].getName assert(!installed(cometShuffle, kryo, registrationRequired)) + assert(!installed(cometShuffle, kryo, registrationRequired, registrator -> sparkOnly)) + assert( + installed( + cometShuffle, + kryo, + registrationRequired, + registrator -> s"$sparkOnly, ${CometKryoRegistrator.CLASS_NAME}")) + assert( + installed( + cometShuffle, + kryo, + registrationRequired, + "spark.kryo.classesToRegister" -> ArrowCachedBatchSerializer.cachedBatchClass.getName)) + // A registrator that cannot be loaded leaves only spark.kryo.registrator to go by. assert(!installed(cometShuffle, kryo, registrationRequired, registrator -> "com.example.R")) assert( installed( @@ -1120,6 +1136,31 @@ class CometInMemoryCacheSuite extends CometTestBase { assert(installed(cometShuffle, kryo)) } + test("Comet plugin finds the Kryo registrations Comet needs however they were made") { + def unregistered(settings: (String, String)*): Seq[Class[_]] = + CometDriverPlugin.unregisteredKryoClasses(new SparkConf().setAll(settings)) + + val kryo = "spark.serializer" -> "org.apache.spark.serializer.KryoSerializer" + val registrationRequired = "spark.kryo.registrationRequired" -> "true" + assert(unregistered().isEmpty) + assert(unregistered(kryo).isEmpty) + assert( + unregistered(kryo, registrationRequired).contains( + ArrowCachedBatchSerializer.cachedBatchClass)) + assert( + unregistered( + kryo, + registrationRequired, + "spark.kryo.registrator" -> CometKryoRegistrator.CLASS_NAME).isEmpty) + assert( + unregistered( + kryo, + registrationRequired, + "spark.kryo.classesToRegister" -> CometKryoRegistrator.classes + .map(_.getName) + .mkString(",")).isEmpty) + } + test("Comet in-memory cache supports empty projection scan") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", From 1330b7505e8635f308f36c6806cdeac718ed6ea6 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 08:43:54 -0600 Subject: [PATCH 24/24] docs: describe the Kryo condition by registration in the upgrade guide The plugin now keeps Spark's cache format when Kryo has not registered Comet's cached batch, however the application registers classes, rather than when spark.kryo.registrator does not list CometKryoRegistrator. --- docs/source/user-guide/latest/migration-guide.md | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/source/user-guide/latest/migration-guide.md b/docs/source/user-guide/latest/migration-guide.md index fec4e92c2f5..a28c600c8fb 100644 --- a/docs/source/user-guide/latest/migration-guide.md +++ b/docs/source/user-guide/latest/migration-guide.md @@ -79,9 +79,9 @@ setting when the application: - starts with `spark.comet.enabled` or `spark.comet.exec.enabled` set to `false`. - leaves Comet shuffle enabled without one of Comet's shuffle managers, so that Comet disables itself. -- uses Kryo with `spark.kryo.registrationRequired=true` and does not list - `org.apache.comet.CometKryoRegistrator` in `spark.kryo.registrator`, because Kryo would reject - Comet's cached batches. Add the registrator to use Comet's format; see +- uses Kryo with `spark.kryo.registrationRequired=true` and has not registered Comet's cached + batch, because Kryo would reject it. To use Comet's format, register Comet's classes with + `spark.kryo.registrator=org.apache.comet.CometKryoRegistrator`; see [Kryo](in-memory-cache.md#kryo). An application that sets `spark.sql.cache.serializer` itself keeps the serializer it chose.