From 439f84d570ac293f8eb17c3121a7fd1b040e3ad2 Mon Sep 17 00:00:00 2001 From: LinSimon-901101 Date: Sun, 27 Sep 2026 03:33:40 +0800 Subject: [PATCH 1/3] test: compare in-memory caches with independent answers --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../exec/CometInMemoryCacheKryoSuite.scala | 16 +- .../exec/CometInMemoryCachePruningSuite.scala | 253 ++++++++++++++++++ .../comet/exec/CometInMemoryCacheSuite.scala | 246 ++++++++++------- 5 files changed, 414 insertions(+), 103 deletions(-) create mode 100644 spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index ea8a23d204c..2400afca360 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -527,6 +527,7 @@ jobs: org.apache.comet.exec.CometExecSuite org.apache.comet.exec.CometEmptyRelationExecSuite org.apache.comet.exec.CometInMemoryCacheSuite + org.apache.comet.exec.CometInMemoryCachePruningSuite org.apache.comet.exec.CometInMemoryCacheKryoSuite org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 3a351d1e6f9..a9f0e58dfa9 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -175,6 +175,7 @@ jobs: org.apache.comet.exec.CometExecSuite org.apache.comet.exec.CometEmptyRelationExecSuite org.apache.comet.exec.CometInMemoryCacheSuite + org.apache.comet.exec.CometInMemoryCachePruningSuite org.apache.comet.exec.CometInMemoryCacheKryoSuite org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite 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 13d0623a355..f6d3a51936a 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala @@ -20,7 +20,7 @@ package org.apache.comet.exec import org.apache.spark.SparkConf -import org.apache.spark.sql.CometTestBase +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.storage.StorageLevel @@ -113,6 +113,12 @@ class CometInMemoryCacheKryoSuite extends CometTestBase { .selectExpr(statsColumns: _*) .createOrReplaceTempView("kryo_cache") + val query = "SELECT * FROM kryo_cache WHERE c_dec_short >= 100 AND c_string > '1'" + // Disabling Comet after caching would still read the same serialized payload. + val expected = withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark.sql(query).collect().toSeq + } + spark.catalog.cacheTable("kryo_cache", level) assert(spark.table("kryo_cache").count() == 200) @@ -124,9 +130,7 @@ class CometInMemoryCacheKryoSuite extends CometTestBase { // Read the payload back rather than only the row count, so a Kryo round trip that // silently mangles the Arrow bytes fails too. The predicate also exercises the // statistics row, which is what carries UTF8String and Decimal through Kryo. - checkSparkAnswer( - spark.sql("SELECT c_long, c_string, c_dec_long, c_ts FROM kryo_cache " + - "WHERE c_dec_short >= 100 AND c_string > '1'")) + checkAnswer(spark.sql(query), expected) } finally { spark.catalog.clearCache() } @@ -180,7 +184,9 @@ class CometInMemoryCacheKryoSuite extends CometTestBase { cachedBatchTypes("kryo_cache_fallback").sameElements( Array("org.apache.spark.sql.execution.columnar.DefaultCachedBatch"))) - checkSparkAnswer(spark.sql("SELECT id FROM kryo_cache_fallback WHERE id > 90")) + checkAnswer( + spark.sql("SELECT id FROM kryo_cache_fallback WHERE id > 90"), + (91L until 100L).map(Row(_))) } finally { spark.catalog.clearCache() } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala new file mode 100644 index 00000000000..b73faba21fc --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala @@ -0,0 +1,253 @@ +/* + * 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.comet.exec + +import java.sql.Timestamp +import java.time.Instant + +import org.apache.spark.SparkConf +import org.apache.spark.sql.{CometTestBase, DataFrame, Row} +import org.apache.spark.sql.comet.{CometInMemoryTableScanExec, CometNativeScanExec} +import org.apache.spark.sql.execution.FileSourceScanExec +import org.apache.spark.sql.execution.columnar.CometInMemoryRelationHelper +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types._ + +import org.apache.comet.CometConf + +class CometInMemoryCachePruningSuite 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 = super.sparkConf + .set("spark.plugins", "org.apache.spark.CometPlugin") + .set( + "spark.sql.cache.serializer", + "org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer") + + private val schema = StructType( + Seq( + StructField("id", IntegerType, nullable = false), + StructField("d", DoubleType), + StructField("f", FloatType), + StructField("n", IntegerType), + StructField("s", StringType), + StructField("dec", DecimalType(20, 3)), + StructField("ts", TimestampType), + StructField("b", BooleanType))) + + private def fixture(): DataFrame = { + def repeated(d: Double): Seq[Double] = Seq.fill(4)(d) + // Every four rows form one batch in all three writers. Keep NaN-only, mixed finite/NaN, + // signed-zero-only, infinity and all-null batches separate so incorrect bounds lose rows. + val values = Seq( + repeated(Double.NegativeInfinity), + repeated(-100.0), + repeated(-2.0), + repeated(-0.0), + repeated(0.0), + repeated(0.25), + repeated(1.0), + repeated(2.0), + repeated(100.0), + repeated(Double.PositiveInfinity), + repeated(Double.NaN), + Seq(1.0, Double.NaN, 3.0, 2.0), + repeated(0.0), // all-null batch + Seq(-0.0, 0.0, -0.0, 0.0), + Seq(-3.0, -2.0, -1.0, 0.0), + Seq(Double.PositiveInfinity, Double.NaN, Double.PositiveInfinity, Double.NaN)) + val strings = Seq( + "", + "a", + "ab", + "b", + "\u007f", + "\u0080", + "\ue000", + "\ud800\udc00", + "é", + "中", + "prefix-a", + "prefix-z", + null, + "z", + "e\u0301", + "😀") + val rows = values.zipWithIndex.flatMap { case (batch, group) => + batch.zipWithIndex.map { case (d, offset) => + val isNull = group == 12 || (group == 11 && offset == 2) + Row( + group * 4 + offset, + if (isNull) null else Double.box(d), + if (isNull) null else Float.box(d.toFloat), + if (isNull) null else Int.box(group), + strings(group), + if (group == 12) null else new java.math.BigDecimal(s"${group - 8}.125"), + if (group == 12) null + else + Timestamp.from( + Instant + .parse("1960-01-01T00:00:00Z") + .plusSeconds(group * 86400L) + .plusNanos(offset * 1000L)), + if (group == 12) null else Boolean.box(group % 2 == 0)) + } + } + spark.createDataFrame(spark.sparkContext.parallelize(rows, 1), schema) + } + + private val predicates = Seq( + "d = CAST('NaN' AS DOUBLE)", + "f = CAST('NaN' AS FLOAT)", + "d > CAST('Infinity' AS DOUBLE)", + "f < CAST('NaN' AS FLOAT)", + "d = 0.0D", + "f = CAST('-0.0' AS FLOAT)", + "d >= CAST('-0.0' AS DOUBLE) AND d <= 0.0D", + "f >= CAST(0.0 AS FLOAT) AND f <= CAST('-0.0' AS FLOAT)", + "d = CAST('-Infinity' AS DOUBLE)", + "f >= CAST('Infinity' AS FLOAT)", + "d > -2.0D AND d < 2.0D", + "d IS NULL", + "d IS NOT NULL", + "n IS NULL", + "s = '中'", + "s >= '\ue000'", + "s < '\u0080'", + "s LIKE 'prefix%'", + "dec >= -1.125 AND dec < 2.125", + "ts < TIMESTAMP '1960-01-05 00:00:00'", + "b = true", + "id IN (1, 9, 49)", + "d = -100.0D OR s = 'prefix-z'") + + Seq("native Arrow", "Spark columnar", "row").foreach { writer => + test(s"cache pruning matches uncached Spark with $writer input") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "UTC", + SQLConf.IN_MEMORY_PARTITION_PRUNING.key -> "true", + SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> "true", + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "true", + SQLConf.PARQUET_VECTORIZED_READER_BATCH_SIZE.key -> "4", + SQLConf.COLUMN_BATCH_SIZE.key -> "4", + CometConf.COMET_BATCH_SIZE.key -> "4", + CometConf.COMET_SHUFFLE_JVM_BATCH_SIZE.key -> "4", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + CometConf.COMET_SPARK_TO_ARROW_ENABLED.key -> "false", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> (writer == "native Arrow").toString) { + spark.catalog.clearCache() + val source = fixture() + // Collect every expected value and predicate result before registering any cache. A + // second query against the cached table, even with Comet disabled, is not an oracle: + // Spark still decodes the same Comet payload and applies the same cached statistics. + // Use the original in-memory data: Parquet row-group pruning can itself mishandle + // signed zero, which would make a file-based oracle hide a cache pruning regression. + val oracle = withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + source + .selectExpr((Seq("*") ++ predicates.zipWithIndex.map { case (p, i) => + s"($p) AS predicate_$i" + }): _*) + .collect() + .toSeq + } + val expectedRows = oracle.map(row => Row.fromSeq(row.toSeq.take(schema.length))) + + withTempPath { path => + val input = if (writer == "row") { + source + } else { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + source.write.option("parquet.enable.dictionary", "false").parquet(path.toString) + } + spark.read.parquet(path.toString) + } + input.createOrReplaceTempView("pruning_cache") + val cached = spark.table("pruning_cache").cache() + try { + val relation = + spark.sharedState.cacheManager.lookupCachedData(cached).get.cachedRepresentation + val plan = relation.cacheBuilder.cachedPlan + withClue(s"$writer cache writer:\n$plan\n") { + writer match { + case "native Arrow" => + assert(plan.supportsColumnar) + assert(plan.collect { case s: CometNativeScanExec => s }.nonEmpty) + case "Spark columnar" => + assert(plan.supportsColumnar) + assert(plan.collect { case s: FileSourceScanExec => s }.nonEmpty) + assert(plan.collect { case s: CometNativeScanExec => s }.isEmpty) + case "row" => assert(!plan.supportsColumnar) + } + } + checkCometAnswer(cached, expectedRows) + val batches = relation.cacheBuilder.cachedColumnBuffers.collect() + assert(batches.length == 16, "the fixture must produce many distinct small batches") + assert(batches.forall(_.numRows == 4)) + assert( + batches.forall(_.getClass.getName == + "org.apache.spark.sql.comet.execution.arrow.CometCachedBatch")) + + predicates.zipWithIndex.foreach { case (predicate, i) => + val expected = oracle + .filter { row => + !row.isNullAt(schema.length + i) && row.getBoolean(schema.length + i) + } + .map(_.getInt(0)) + .sorted + assert(expected.nonEmpty && expected.length < expectedRows.length) + val query = cached.where(predicate).select("id") + val actual = query.collect().map(_.getInt(0)).sorted.toSeq + val scans = query.queryExecution.executedPlan.collect { + case scan: CometInMemoryTableScanExec => scan + } + withClue(s"$writer input, predicate: $predicate\n") { + assert(actual == expected) + assert(scans.length == 1) + assert(scans.head.originalPlan.predicates.nonEmpty) + val scannedRows = scans.head.metrics("numOutputRows").value + assert(scannedRows >= expected.length) + assert( + scannedRows < expectedRows.length, + s"pruning must skip batches, but decoded $scannedRows rows") + } + } + } finally { + cached.unpersist(blocking = true) + spark.catalog.clearCache() + spark.catalog.dropTempView("pruning_cache") + } + } + } + } + } +} 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 b12ef02c5f7..be307930652 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -32,7 +32,7 @@ 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.catalyst.expressions.{And, Attribute, AttributeReference, Expression, GreaterThanOrEqual, LessThan, Literal} +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} import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} @@ -102,6 +102,18 @@ class CometInMemoryCacheSuite extends CometTestBase { .collect() } + // Disabling Comet does not bypass an existing cache: both readers would still consume the + // same serialized values. Materialize every Spark reference before registering the cache. + private def uncachedSparkAnswer(query: String): Array[Row] = { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + val df = spark.sql(query) + assert( + df.queryExecution.withCachedData.collect { case r: InMemoryRelation => r }.isEmpty, + "the reference answer must not read a cached relation") + df.collect() + } + } + // The tests below are ported from Spark 4.1.2's AdaptiveQueryExecSuite; see each source link. private def withAQECache(f: => Unit): Unit = { withSQLConf( @@ -300,7 +312,7 @@ class CometInMemoryCacheSuite extends CometTestBase { Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) val df = spark.sql("SELECT key, count(*) FROM abc GROUP BY key") - checkSparkAnswer(df) + checkAnswer(df, (0L until 1000L).map(i => Row(i, 1L))) val plan = df.queryExecution.executedPlan.toString() assert(plan.contains("CometInMemoryTableScan")) @@ -341,7 +353,7 @@ class CometInMemoryCacheSuite extends CometTestBase { "spark.comet.sparkToColumnar.enabled" -> "true") { val df = spark.sql("SELECT key, count(*) FROM comet_cache_disabled GROUP BY key") - checkSparkAnswer(df) + checkAnswer(df, (0L until 1000L).map(i => Row(i, 1L))) val plan = df.queryExecution.executedPlan.toString() assert(!plan.contains("CometInMemoryTableScan")) @@ -379,6 +391,19 @@ class CometInMemoryCacheSuite extends CometTestBase { .sql(s"SELECT id AS key, $column FROM range(1000)") .createOrReplaceTempView("default_cached_batch") + val columnarQuery = """ + SELECT key, payload + FROM default_cached_batch + WHERE key >= 10 AND key < 20 + """ + val rowQuery = """ + SELECT payload + FROM default_cached_batch + WHERE key >= 10 AND key < 20 + """ + val expectedColumnar = uncachedSparkAnswer(columnarQuery) + val expectedRows = uncachedSparkAnswer(rowQuery) + spark.catalog.cacheTable("default_cached_batch") spark.table("default_cached_batch").count() @@ -388,24 +413,16 @@ class CometInMemoryCacheSuite extends CometTestBase { s"$column was cached in Comet's format") // Columnar read path, delegated to Spark's serializer. - val columnarDf = spark.sql(""" - SELECT key, payload - FROM default_cached_batch - WHERE key >= 10 AND key < 20 - """) - checkSparkAnswer(columnarDf) + val columnarDf = spark.sql(columnarQuery) + checkAnswer(columnarDf, expectedColumnar.toSeq) assert( !columnarDf.queryExecution.executedPlan.toString().contains("CometInMemoryTableScan")) // Row read path: disabling the vectorized cache reader makes Spark use // convertCachedBatchToInternalRow. withSQLConf(SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> "false") { - val rowDf = spark.sql(""" - SELECT payload - FROM default_cached_batch - WHERE key >= 10 AND key < 20 - """) - checkSparkAnswer(rowDf) + val rowDf = spark.sql(rowQuery) + checkAnswer(rowDf, expectedRows.toSeq) assert(!rowDf.queryExecution.executedPlan.toString().contains("CometInMemoryTableScan")) } @@ -438,7 +455,7 @@ class CometInMemoryCacheSuite extends CometTestBase { FROM multi_partition_cache GROUP BY id % 100 """) - checkSparkAnswer(grouped) + checkAnswer(grouped, (0L until 100L).map(i => Row(i, 10L))) val groupedPlan = grouped.queryExecution.executedPlan.toString() assert(groupedPlan.contains("CometInMemoryTableScan")) @@ -463,7 +480,7 @@ class CometInMemoryCacheSuite extends CometTestBase { empty.count() val emptyDf = spark.sql("SELECT * FROM empty_cache") - checkSparkAnswer(emptyDf) + checkAnswer(emptyDf, Seq.empty[Row]) val emptyPlan = emptyDf.queryExecution.executedPlan.toString() assert(emptyPlan.contains("CometInMemoryTableScan")) @@ -497,7 +514,7 @@ class CometInMemoryCacheSuite extends CometTestBase { Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) val df = spark.sql("SELECT key FROM project_cache") - checkSparkAnswer(df) + checkAnswer(df, (0L until 1000L).map(Row(_))) val plan = df.queryExecution.executedPlan.toString() assert(plan.contains("CometInMemoryTableScan")) @@ -535,7 +552,7 @@ class CometInMemoryCacheSuite extends CometTestBase { Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) val df = spark.sql("SELECT group, count(*) FROM shuffle_cache GROUP BY group") - checkSparkAnswer(df) + checkAnswer(df, (0L until 100L).map(i => Row(i, 10L))) val plan = df.queryExecution.executedPlan.toString() assert(plan.contains("CometInMemoryTableScan")) @@ -622,7 +639,7 @@ class CometInMemoryCacheSuite extends CometTestBase { spark.catalog.dropTempView("typed_stats_input") } - withSparkColumnarCache("typed_stats_columnar")(path => df.write.parquet(path)) { + withSparkColumnarCache("typed_stats_columnar")(path => df.write.parquet(path)) { _ => val relation = spark.sharedState.cacheManager .lookupCachedData(spark.table("typed_stats_columnar")) .get @@ -689,7 +706,7 @@ class CometInMemoryCacheSuite extends CometTestBase { df.unpersist(blocking = true) spark.catalog.dropTempView("extreme_stats_input") } - withSparkColumnarCache("extreme_stats_columnar")(path => df.write.parquet(path)) { + withSparkColumnarCache("extreme_stats_columnar")(path => df.write.parquet(path)) { _ => checkStats("extreme_stats_columnar") } } @@ -759,7 +776,7 @@ class CometInMemoryCacheSuite extends CometTestBase { FROM prune_cache WHERE key >= 900 AND key < 905 """) - checkSparkAnswer(df) + checkAnswer(df, (900L until 905L).map(i => Row(i, i % 7))) val plan = df.queryExecution.executedPlan.toString() assert(plan.contains("CometInMemoryTableScan")) @@ -798,12 +815,8 @@ class CometInMemoryCacheSuite extends CometTestBase { val df = spark.sql("SELECT key, value FROM prune_conf_cache WHERE key >= 900 AND key < 905") - checkSparkAnswer(df) - - // checkSparkAnswer takes its argument by name and executes its own copies of the query, - // so this df's plan instance has not run and its metrics are all still zero. Force this - // exact plan before reading them, or the comparison below passes vacuously with 0 == 0. - df.collect() + // Run this exact plan once so its metrics describe the checked result. + QueryTest.checkAnswer(df, (900L until 905L).map(i => Row(i, i % 7)), checkToRDD = false) val scans = df.queryExecution.executedPlan.collect { case s: org.apache.spark.sql.comet.CometInMemoryTableScanExec => s @@ -867,7 +880,7 @@ class CometInMemoryCacheSuite extends CometTestBase { s"expected all ${info.numPartitions} partitions cached, got ${info.numCachedPartitions}") val df = spark.sql("SELECT key, value FROM disk_cache WHERE key >= 900 AND key < 905") - checkSparkAnswer(df) + checkAnswer(df, (900L until 905L).map(i => Row(i, i % 7))) val plan = df.queryExecution.executedPlan.toString() assert(plan.contains("CometInMemoryTableScan")) @@ -900,6 +913,10 @@ class CometInMemoryCacheSuite extends CometTestBase { (2, java.sql.Timestamp.valueOf("1970-01-01 00:00:00")), (3, null)) rows.toDF("id", "ts").createOrReplaceTempView("ts_cache") + val valuesQuery = "SELECT id, ts FROM ts_cache ORDER BY id" + val stringsQuery = "SELECT id, CAST(ts AS STRING) AS s FROM ts_cache ORDER BY id" + val expectedValues = uncachedSparkAnswer(valuesQuery) + val expectedStrings = uncachedSparkAnswer(stringsQuery) spark.catalog.cacheTable("ts_cache") assert(spark.table("ts_cache").count() == 3) @@ -945,9 +962,8 @@ class CometInMemoryCacheSuite extends CometTestBase { s"got ${labels.mkString("[", ",", "]")}") // The label change must not move any values. - checkSparkAnswer(spark.sql("SELECT id, ts FROM ts_cache ORDER BY id")) - checkSparkAnswer( - spark.sql("SELECT id, CAST(ts AS STRING) AS s FROM ts_cache ORDER BY id")) + checkAnswer(spark.sql(valuesQuery), expectedValues.toSeq) + checkAnswer(spark.sql(stringsQuery), expectedStrings.toSeq) spark.catalog.clearCache() } @@ -1003,7 +1019,7 @@ class CometInMemoryCacheSuite extends CometTestBase { spark.table("count_cache").count() val df = spark.sql("SELECT count(*) FROM count_cache") - checkSparkAnswer(df) + checkAnswer(df, Seq(Row(1000L))) val plan = df.queryExecution.executedPlan.toString() assert(plan.contains("CometInMemoryTableScan")) @@ -1052,7 +1068,7 @@ class CometInMemoryCacheSuite extends CometTestBase { // Expected values come from the uncached query so a wrong-but-consistent cached answer // cannot make this pass. - val expected = spark.sql(query).orderBy("l").collect() + val expected = uncachedSparkAnswer(s"$query ORDER BY l") spark.sql(query).createOrReplaceTempView("all_types_cache") spark.catalog.cacheTable("all_types_cache") @@ -1125,7 +1141,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val predicate = "startswith(s, 'i\u0307')" withNativeCache { // Collected before the relation is cached, so the cache cannot stand in for it. - val expected = spark.sql(s"SELECT id FROM ($query) WHERE $predicate ORDER BY id").collect() + val expected = uncachedSparkAnswer(s"SELECT id FROM ($query) WHERE $predicate ORDER BY id") assert(expected.length == 5) spark.sql(query).createOrReplaceTempView("collated_prefix_cache") @@ -1264,45 +1280,61 @@ class CometInMemoryCacheSuite extends CometTestBase { spark.catalog.clearCache() - spark - .sql(""" - SELECT * - FROM VALUES - (0, CAST('NaN' AS DOUBLE), CAST('NaN' AS FLOAT)), - (1, 1.0D, CAST(1.0 AS FLOAT)), - (2, -0.0D, CAST(-0.0 AS FLOAT)), - (3, 0.0D, CAST(0.0 AS FLOAT)) - AS t(id, d, f) - """) - .createOrReplaceTempView("nan_prune_cache") - - spark.catalog.cacheTable("nan_prune_cache") - spark.table("nan_prune_cache").count() - - val doubleDf = spark.sql(""" - SELECT id - FROM nan_prune_cache - WHERE isnan(d) - """) - checkSparkAnswer(doubleDf) + // A single row-input partition gives two deterministic two-row cached batches. + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .sql(""" + SELECT * + FROM VALUES + (0, CAST('NaN' AS DOUBLE), CAST('NaN' AS FLOAT)), + (1, 1.0D, CAST(1.0 AS FLOAT)), + (2, CAST('-0.0' AS DOUBLE), CAST('-0.0' AS FLOAT)), + (3, 0.0D, CAST(0.0 AS FLOAT)) + AS t(id, d, f) + """) + .coalesce(1) + .createOrReplaceTempView("nan_prune_cache") - val floatDf = spark.sql(""" - SELECT id - FROM nan_prune_cache - WHERE isnan(f) - """) - checkSparkAnswer(floatDf) + spark.catalog.cacheTable("nan_prune_cache") + spark.table("nan_prune_cache").count() + } - val zeroDf = spark.sql(""" - SELECT id - FROM nan_prune_cache - WHERE d = 0.0D OR f = CAST(0.0 AS FLOAT) - """) - checkSparkAnswer(zeroDf) + val relation = spark.sharedState.cacheManager + .lookupCachedData(spark.table("nan_prune_cache")) + .get + .cachedRepresentation + val batches = relation.cacheBuilder.cachedColumnBuffers + assert(batches.count() == 2, "the NaN/finite and signed-zero rows need separate batches") - val plan = doubleDf.queryExecution.executedPlan.toString() - assert(plan.contains("CometInMemoryTableScan")) - assert(!plan.contains("CometSparkColumnarToColumnar")) + Seq( + ("d", Literal(Double.NaN), Literal(0.0d), Literal(1.0d)), + ("f", Literal(Float.NaN), Literal(0.0f), Literal(1.0f))).foreach { + case (column, nan, zero, one) => + val attr = relation.output.find(_.name == column).get + val nanSql = s"CAST('NaN' AS ${nan.dataType.sql})" + // These comparisons become statistics filters. isnan alone would leave all batches + // eligible and could not catch an incorrectly recorded NaN upper bound. + val comparisons = Seq( + (s"$column = $nanSql", EqualTo(attr, nan), Seq(Row(0)), 1L), + (s"$column > 1", GreaterThan(attr, one), Seq(Row(0)), 1L), + (s"$column < $nanSql", LessThan(attr, nan), Seq(Row(1), Row(2), Row(3)), 2L), + (s"$column = 0", EqualTo(attr, zero), Seq(Row(2), Row(3)), 1L), + (s"$column > $nanSql", GreaterThan(attr, nan), Seq.empty[Row], 0L), + (s"$column < 0", LessThan(attr, zero), Seq.empty[Row], 0L)) + comparisons.foreach { case (predicate, expression, expected, expectedBatches) => + withClue(s"predicate: $predicate: ") { + val filter = relation.cacheBuilder.serializer + .buildFilter(Seq(expression), relation.output) + assert(batches.mapPartitionsWithIndex(filter).count() == expectedBatches) + + val df = spark.sql(s"SELECT id FROM nan_prune_cache WHERE $predicate") + checkAnswer(df, expected) + val plan = df.queryExecution.executedPlan.toString() + assert(plan.contains("CometInMemoryTableScan")) + assert(!plan.contains("CometSparkColumnarToColumnar")) + } + } + } spark.catalog.clearCache() } @@ -1314,10 +1346,10 @@ class CometInMemoryCacheSuite extends CometTestBase { * CometVector. Spark's InMemoryRelation strips the ColumnarToRow above that scan because * supportsColumnarInput is true for the schema, so the serializer receives non-Arrow columnar * batches. Asserts the relation really was stored in Comet's format before handing control to - * `f`. + * `f`, along with Spark's result collected before caching. */ private def withSparkColumnarCache(view: String, extraConfs: (String, String)*)( - write: String => Unit)(f: => Unit): Unit = { + write: String => Unit)(f: Seq[Row] => Unit): Unit = { withTempPath { path => write(path.toString) @@ -1329,6 +1361,7 @@ class CometInMemoryCacheSuite extends CometTestBase { SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "true") ++ extraConfs: _*) { spark.read.parquet(path.toString).createOrReplaceTempView(view) + val expected = uncachedSparkAnswer(s"SELECT * FROM $view") spark.catalog.cacheTable(view) spark.table(view).count() @@ -1336,7 +1369,7 @@ class CometInMemoryCacheSuite extends CometTestBase { cachedBatchTypes(view).sameElements( Array("org.apache.spark.sql.comet.execution.arrow.CometCachedBatch"))) - f + f(expected.toSeq) } } } @@ -1359,14 +1392,25 @@ class CometInMemoryCacheSuite extends CometTestBase { "timestamp_micros(id * 1000000) as ts") .write .parquet(path) - } { + } { expected => assert(spark.table("spark_columnar_cache").count() == 1000) - checkSparkAnswer( - spark.sql("SELECT * FROM spark_columnar_cache WHERE key >= 10 AND key < 20 ORDER BY key")) - checkSparkAnswer( - spark.sql("SELECT sum(key), sum(d), sum(dec), count(s), count(n), max(dt), max(ts) " + - "FROM spark_columnar_cache")) + checkAnswer( + spark.sql("SELECT * FROM spark_columnar_cache WHERE key >= 10 AND key < 20 ORDER BY key"), + expected.filter(row => row.getLong(0) >= 10 && row.getLong(0) < 20).sortBy(_.getLong(0))) + checkAnswer( + spark.sql( + "SELECT sum(key), sum(d), sum(dec), count(s), count(n), max(dt), max(ts) " + + "FROM spark_columnar_cache"), + Seq( + Row( + 499500L, + 499500.0d, + BigDecimal(499500), + 1000L, + 0L, + java.sql.Date.valueOf(java.time.LocalDate.of(2020, 1, 1).plusDays(999)), + java.sql.Timestamp.from(java.time.Instant.ofEpochSecond(999))))) } } @@ -1386,10 +1430,10 @@ class CometInMemoryCacheSuite extends CometTestBase { "cast(cast(id as string) as binary) as b") .write .parquet(path) - } { + } { expected => assert(spark.table("spark_columnar_complex").count() == 200) - checkSparkAnswer(spark.sql("SELECT key, a, st, m, b FROM spark_columnar_complex")) + checkAnswer(spark.sql("SELECT key, a, st, m, b FROM spark_columnar_complex"), expected) } } @@ -1563,8 +1607,8 @@ class CometInMemoryCacheSuite extends CometTestBase { // stores the wrong values: with and without Comet, both sides read the one cached payload. val full = s"SELECT * FROM $view ORDER BY id" val projected = s"SELECT s FROM $view WHERE id >= 3990 ORDER BY s" - val expectedFull = spark.sql(full).collect() - val expectedProjected = spark.sql(projected).collect() + val expectedFull = uncachedSparkAnswer(full) + val expectedProjected = uncachedSparkAnswer(projected) spark.catalog.cacheTable(view) assert( @@ -1576,8 +1620,10 @@ class CometInMemoryCacheSuite extends CometTestBase { // nothing -- the three shapes the read path distinguishes. assert(spark.sql(full).collect() === expectedFull, s"codec $codec read the wrong values") assert(spark.sql(projected).collect() === expectedProjected) - checkSparkAnswer(spark.sql(s"SELECT * FROM $view")) - checkSparkAnswer(spark.sql(s"SELECT s FROM $view WHERE id >= 3990")) + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + checkAnswer(spark.sql(full), expectedFull.toSeq) + checkAnswer(spark.sql(projected), expectedProjected.toSeq) + } 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) @@ -1630,8 +1676,8 @@ class CometInMemoryCacheSuite extends CometTestBase { // another column, it is a NullVector inside an ordinary payload instead. val query = "SELECT id, NULL AS n FROM range(0, 1000, 1, 2)" // Collected before anything is cached, so the cache cannot stand in for the reference. - val expectedNulls = spark.sql(s"SELECT n FROM ($query)").collect() - val expectedPairs = spark.sql(s"SELECT id, n FROM ($query) ORDER BY id").collect() + val expectedNulls = uncachedSparkAnswer(s"SELECT n FROM ($query)") + val expectedPairs = uncachedSparkAnswer(s"SELECT id, n FROM ($query) ORDER BY id") assert(expectedNulls.length == 1000) Seq("none", "zstd").foreach { codec => @@ -1877,7 +1923,7 @@ class CometInMemoryCacheSuite extends CometTestBase { // Collected before the relation is cached. Once it is, Spark answers this same query from the // cache too, so a reference taken afterwards would compare the cache with itself. val expected = - projections.map(cols => spark.sql(orderedByJson(cols, s"($query)")).collect()) + projections.map(cols => uncachedSparkAnswer(orderedByJson(cols, s"($query)"))) expected.foreach(rows => assert(rows.length == projectionCacheRows)) spark.sql(query).createOrReplaceTempView("nested_value_cache") @@ -1911,7 +1957,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val projections = Seq(Seq("id", "sc", "ar", "mp", "deep", "tail"), Seq("tail", "mp", "id"), Seq("deep", "sc")) // Collected before the fixture caches the relation, for the reason given in the test above. - val expected = projections.map(cols => spark.sql(orderedByJson(cols, s"($source)")).collect()) + val expected = projections.map(cols => uncachedSparkAnswer(orderedByJson(cols, s"($source)"))) expected.foreach(rows => assert(rows.length == projectionCacheRows)) withNestedProjectionCache(Some(tinyChunkSize)) { (relation, batches) => @@ -2010,7 +2056,7 @@ class CometInMemoryCacheSuite extends CometTestBase { scan.get.scanOutput.isEmpty, s"expected no scanned columns, got ${scan.get.scanOutput.map(_.name).mkString(",")}") - checkSparkAnswer(df) + checkAnswer(df, Seq(Row(500L))) spark.catalog.clearCache() } } @@ -2031,15 +2077,19 @@ class CometInMemoryCacheSuite extends CometTestBase { // 3 left rows joined to 2 right rows, summing only the right side: 3 * (0 + 1) == 3. // Leaking the left id column into the scan output made this read 10 + 11 + 12 twice. - checkSparkAnswer(spark.sql(""" + checkAnswer( + spark.sql(""" |SELECT /*+ BROADCAST(r) */ sum(r.id) |FROM cached_left l JOIN range(2) r ON true - """.stripMargin)) + """.stripMargin), + Seq(Row(3L))) - checkSparkAnswer(spark.sql(""" + checkAnswer( + spark.sql(""" |SELECT /*+ BROADCAST(r) */ r.id |FROM cached_left l JOIN range(2) r ON true - """.stripMargin)) + """.stripMargin), + Seq.fill(3)(Seq(Row(0L), Row(1L))).flatten) spark.catalog.clearCache() } @@ -2376,7 +2426,7 @@ class CometInMemoryCacheSuite extends CometTestBase { assert(relation.output.length == 2) val df = spark.sql("SELECT s1, s2 FROM dictionary_cache") - checkSparkAnswer(df) + checkAnswer(df, (0 until 2000).map(i => Row(s"a_${i % 3}", s"b_${i % 4}"))) val distinct = spark.sql("SELECT DISTINCT s1, s2 FROM dictionary_cache ORDER BY s1, s2").collect() @@ -2393,7 +2443,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val df = spark.sql( "SELECT /*+ BROADCAST(c) */ c.s1, c.s2 FROM range(1) r JOIN dictionary_cache c ON true") - checkSparkAnswer(df) + checkAnswer(df, (0 until 2000).map(i => Row(s"a_${i % 3}", s"b_${i % 4}"))) assert(df.count() == 2000) } } @@ -2420,7 +2470,7 @@ class CometInMemoryCacheSuite extends CometTestBase { val df = spark.sql( "SELECT k, count(*) AS c FROM reuse_cache GROUP BY k " + "UNION ALL SELECT k, count(*) AS c FROM reuse_cache GROUP BY k") - checkSparkAnswer(df) + checkAnswer(df, Seq.fill(2)((0L until 10L).map(k => Row(k, 40L))).flatten) val plan = df.queryExecution.executedPlan val exchanges = plan.collect { case e: Exchange => e } From 3695fdf14d2be5effdff2e9701994992f28ec15b Mon Sep 17 00:00:00 2001 From: LinSimon-901101 Date: Sun, 27 Sep 2026 03:39:47 +0800 Subject: [PATCH 2/3] test: support Spark 3.x and assert precise cache pruning --- .../exec/CometInMemoryCacheKryoSuite.scala | 5 +++-- .../exec/CometInMemoryCachePruningSuite.scala | 21 +++++++++++-------- .../comet/exec/CometInMemoryCacheSuite.scala | 4 +++- 3 files changed, 18 insertions(+), 12 deletions(-) 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 f6d3a51936a..b403442531c 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheKryoSuite.scala @@ -115,8 +115,9 @@ class CometInMemoryCacheKryoSuite extends CometTestBase { val query = "SELECT * FROM kryo_cache WHERE c_dec_short >= 100 AND c_string > '1'" // Disabling Comet after caching would still read the same serialized payload. - val expected = withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - spark.sql(query).collect().toSeq + var expected = Seq.empty[Row] + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + expected = spark.sql(query).collect().toSeq } spark.catalog.cacheTable("kryo_cache", level) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala index b73faba21fc..21ef0dd0c37 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCachePruningSuite.scala @@ -146,7 +146,7 @@ class CometInMemoryCachePruningSuite extends CometTestBase { "s LIKE 'prefix%'", "dec >= -1.125 AND dec < 2.125", "ts < TIMESTAMP '1960-01-05 00:00:00'", - "b = true", + "b <=> true", "id IN (1, 9, 49)", "d = -100.0D OR s = 'prefix-z'") @@ -172,8 +172,9 @@ class CometInMemoryCachePruningSuite extends CometTestBase { // Spark still decodes the same Comet payload and applies the same cached statistics. // Use the original in-memory data: Parquet row-group pruning can itself mishandle // signed zero, which would make a file-based oracle hide a cache pruning regression. - val oracle = withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - source + var oracle = Seq.empty[Row] + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + oracle = source .selectExpr((Seq("*") ++ predicates.zipWithIndex.map { case (p, i) => s"($p) AS predicate_$i" }): _*) @@ -218,12 +219,13 @@ class CometInMemoryCachePruningSuite extends CometTestBase { "org.apache.spark.sql.comet.execution.arrow.CometCachedBatch")) predicates.zipWithIndex.foreach { case (predicate, i) => + def matches(row: Row): Boolean = + !row.isNullAt(schema.length + i) && row.getBoolean(schema.length + i) val expected = oracle - .filter { row => - !row.isNullAt(schema.length + i) && row.getBoolean(schema.length + i) - } + .filter(matches) .map(_.getInt(0)) .sorted + val expectedScannedRows = oracle.grouped(4).count(_.exists(matches)) * 4 assert(expected.nonEmpty && expected.length < expectedRows.length) val query = cached.where(predicate).select("id") val actual = query.collect().map(_.getInt(0)).sorted.toSeq @@ -235,10 +237,11 @@ class CometInMemoryCachePruningSuite extends CometTestBase { assert(scans.length == 1) assert(scans.head.originalPlan.predicates.nonEmpty) val scannedRows = scans.head.metrics("numOutputRows").value - assert(scannedRows >= expected.length) + // Counting eligible fixture batches also rejects a filter that only drops the + // all-null batch, without applying the predicate's actual bounds. assert( - scannedRows < expectedRows.length, - s"pruning must skip batches, but decoded $scannedRows rows") + scannedRows == expectedScannedRows, + s"expected $expectedScannedRows rows from eligible batches, decoded $scannedRows") } } } finally { 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 be307930652..b60f2f19359 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -105,13 +105,15 @@ class CometInMemoryCacheSuite extends CometTestBase { // Disabling Comet does not bypass an existing cache: both readers would still consume the // same serialized values. Materialize every Spark reference before registering the cache. private def uncachedSparkAnswer(query: String): Array[Row] = { + var expected = Array.empty[Row] withSQLConf(CometConf.COMET_ENABLED.key -> "false") { val df = spark.sql(query) assert( df.queryExecution.withCachedData.collect { case r: InMemoryRelation => r }.isEmpty, "the reference answer must not read a cached relation") - df.collect() + expected = df.collect() } + expected } // The tests below are ported from Spark 4.1.2's AdaptiveQueryExecSuite; see each source link. From 4924c36d3e4655909349bb11238905aed1957179 Mon Sep 17 00:00:00 2001 From: LinSimon-901101 Date: Sun, 27 Sep 2026 03:47:22 +0800 Subject: [PATCH 3/3] test: preserve empty-cache comparisons on Spark 3.4 --- .../scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 b60f2f19359..db76ddbeb35 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -482,7 +482,7 @@ class CometInMemoryCacheSuite extends CometTestBase { empty.count() val emptyDf = spark.sql("SELECT * FROM empty_cache") - checkAnswer(emptyDf, Seq.empty[Row]) + checkCometAnswer(emptyDf, Seq.empty[Row]) val emptyPlan = emptyDf.queryExecution.executedPlan.toString() assert(emptyPlan.contains("CometInMemoryTableScan"))