From 45018b51ecb357c1087d731c976253931a28b6e0 Mon Sep 17 00:00:00 2001 From: Dustin Smith Date: Tue, 29 Sep 2026 07:48:00 +0700 Subject: [PATCH] perf: keep the sort-merge join when the forced hash join's build side is too large With spark.comet.exec.forceShuffledHashJoin on, RewriteJoin turned every eligible SortMergeJoinExec into a shuffled hash join and read the build side's size estimate only to pick the build side. DataFusion's hash join has to hold that side in memory, so one large join failed the task while the rest of the query would have gained from the rewrite. The rule now rewrites only when the build side's estimate is under a limit and otherwise keeps the SortMergeJoinExec, which still runs natively, recording the reason as plan info. The limit is Spark's own rule for choosing a shuffled hash join, spark.sql.autoBroadcastJoinThreshold times the initial shuffle partition count, with Spark's 10 MB default threshold standing in when broadcasts are disabled. spark.comet.exec.forceShuffledHashJoin.maxBuildSize sets a fixed limit, and a non-positive value removes it, which is what the hash join benchmark configurations and the macOS benchmarking guide now pin so they keep measuring every join. --- benchmarks/pyspark/run_all_benchmarks.sh | 4 + benchmarks/tpc/engines/comet-hashjoin.toml | 2 + .../tpc/engines/comet-iceberg-hashjoin.toml | 2 + .../contributor-guide/benchmarking_macos.md | 1 + .../contributor-guide/memory_management.md | 2 +- .../user-guide/latest/tuning/operators.md | 8 + .../scala/org/apache/comet/CometConf.scala | 16 ++ .../apache/comet/rules/CometExecRule.scala | 2 +- .../org/apache/comet/rules/RewriteJoin.scala | 78 ++++++- .../apache/comet/exec/CometJoinSuite.scala | 202 +++++++++++++++++- 10 files changed, 302 insertions(+), 15 deletions(-) diff --git a/benchmarks/pyspark/run_all_benchmarks.sh b/benchmarks/pyspark/run_all_benchmarks.sh index 0eafffe640d..eeae04255b0 100755 --- a/benchmarks/pyspark/run_all_benchmarks.sh +++ b/benchmarks/pyspark/run_all_benchmarks.sh @@ -67,6 +67,7 @@ $SPARK_HOME/bin/spark-submit \ # Run Comet JVM shuffle echo "" echo ">>> Running COMET JVM shuffle benchmark..." +# maxBuildSize=-1 measures the hash join alone: no build-side size limit on the rewrite. $SPARK_HOME/bin/spark-submit \ --master "$SPARK_MASTER" \ --executor-memory "$EXECUTOR_MEMORY" \ @@ -86,6 +87,7 @@ $SPARK_HOME/bin/spark-submit \ --conf spark.comet.shuffle.mode=jvm \ --conf spark.comet.shuffle.mode=jvm \ --conf spark.comet.exec.replaceSortMergeJoin=true \ + --conf spark.comet.exec.forceShuffledHashJoin.maxBuildSize=-1 \ --conf spark.shuffle.manager=org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager \ --conf spark.sql.extensions=org.apache.comet.CometSparkSessionExtensions \ --conf spark.comet.cast.allowIncompatible=true \ @@ -96,6 +98,7 @@ $SPARK_HOME/bin/spark-submit \ # Run Comet Native shuffle echo "" echo ">>> Running COMET NATIVE shuffle benchmark..." +# maxBuildSize=-1 measures the hash join alone: no build-side size limit on the rewrite. $SPARK_HOME/bin/spark-submit \ --master "$SPARK_MASTER" \ --executor-memory "$EXECUTOR_MEMORY" \ @@ -114,6 +117,7 @@ $SPARK_HOME/bin/spark-submit \ --conf spark.comet.explain.fallback.enabled=true \ --conf spark.comet.shuffle.mode=native \ --conf spark.comet.exec.replaceSortMergeJoin=true \ + --conf spark.comet.exec.forceShuffledHashJoin.maxBuildSize=-1 \ --conf spark.shuffle.manager=org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager \ --conf spark.sql.extensions=org.apache.comet.CometSparkSessionExtensions \ --conf spark.comet.cast.allowIncompatible=true \ diff --git a/benchmarks/tpc/engines/comet-hashjoin.toml b/benchmarks/tpc/engines/comet-hashjoin.toml index 202dcad914b..50e184b049b 100644 --- a/benchmarks/tpc/engines/comet-hashjoin.toml +++ b/benchmarks/tpc/engines/comet-hashjoin.toml @@ -31,4 +31,6 @@ driver_class_path = ["$COMET_JAR"] "spark.plugins" = "org.apache.spark.CometPlugin" "spark.shuffle.manager" = "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager" "spark.comet.exec.replaceSortMergeJoin" = "true" +# Measure the hash join alone: no build-side size limit on the rewrite. +"spark.comet.exec.forceShuffledHashJoin.maxBuildSize" = "-1" "spark.comet.expression.Cast.allowIncompatible" = "true" diff --git a/benchmarks/tpc/engines/comet-iceberg-hashjoin.toml b/benchmarks/tpc/engines/comet-iceberg-hashjoin.toml index 421d3201315..4a563613108 100644 --- a/benchmarks/tpc/engines/comet-iceberg-hashjoin.toml +++ b/benchmarks/tpc/engines/comet-iceberg-hashjoin.toml @@ -34,6 +34,8 @@ driver_class_path = ["$COMET_JAR", "$ICEBERG_JAR"] "spark.plugins" = "org.apache.spark.CometPlugin" "spark.shuffle.manager" = "org.apache.spark.sql.comet.execution.shuffle.CometShuffleManager" "spark.comet.exec.replaceSortMergeJoin" = "true" +# Measure the hash join alone: no build-side size limit on the rewrite. +"spark.comet.exec.forceShuffledHashJoin.maxBuildSize" = "-1" "spark.comet.expression.Cast.allowIncompatible" = "true" "spark.comet.enabled" = "true" "spark.comet.exec.enabled" = "true" diff --git a/docs/source/contributor-guide/benchmarking_macos.md b/docs/source/contributor-guide/benchmarking_macos.md index 13ed2c26a27..4f3541811ce 100644 --- a/docs/source/contributor-guide/benchmarking_macos.md +++ b/docs/source/contributor-guide/benchmarking_macos.md @@ -154,6 +154,7 @@ $SPARK_HOME/bin/spark-submit \ --conf spark.comet.enabled=true \ --conf spark.comet.shuffle.enabled=true \ --conf spark.comet.exec.forceShuffledHashJoin=true \ + --conf spark.comet.exec.forceShuffledHashJoin.maxBuildSize=-1 \ $DF_BENCH/runners/datafusion-comet/tpcbench.py \ --benchmark tpch \ --data $BENCH_DATA/ \ diff --git a/docs/source/contributor-guide/memory_management.md b/docs/source/contributor-guide/memory_management.md index f768baf366e..243b686e1e6 100644 --- a/docs/source/contributor-guide/memory_management.md +++ b/docs/source/contributor-guide/memory_management.md @@ -567,4 +567,4 @@ A checklist for triaging an executor OOM kill: `batch_size * columns`, and wide or deeply nested schemas amplify it. 4. Check whether the operators involved can spill at all. `ShuffledHashJoin` cannot, so `spark.comet.exec.forceShuffledHashJoin=true` converts a spillable sort-merge join into one that - is not. + is not, for build sides under `spark.comet.exec.forceShuffledHashJoin.maxBuildSize`. diff --git a/docs/source/user-guide/latest/tuning/operators.md b/docs/source/user-guide/latest/tuning/operators.md index 9dffa080eef..59a600a8897 100644 --- a/docs/source/user-guide/latest/tuning/operators.md +++ b/docs/source/user-guide/latest/tuning/operators.md @@ -30,6 +30,14 @@ to configure Comet to convert `SortMergeJoin` to `ShuffledHashJoin`. Comet does to test with both for your specific workloads. To configure Comet to convert `SortMergeJoin` to `ShuffledHashJoin`, set `spark.comet.exec.forceShuffledHashJoin=true`. +The conversion only happens when the build side is under a size limit, and a join whose build side has no statistics +is left as a `SortMergeJoin`. The size is Spark's planning estimate of the build side, or under AQE the materialized +shuffle size of the build side. By default the limit is Spark's own rule for choosing a `ShuffledHashJoin`: +`spark.sql.autoBroadcastJoinThreshold` times the initial shuffle partition count +(`spark.sql.adaptive.coalescePartitions.initialPartitionNum` when AQE and partition coalescing are both on and it is +set, else `spark.sql.shuffle.partitions`). When broadcasts are disabled with a non-positive threshold, Spark's default +threshold of 10 MB is used instead. Set `spark.comet.exec.forceShuffledHashJoin.maxBuildSize` to a size in bytes to use a fixed +limit, or to a non-positive value to convert every eligible join regardless of size. ### Join Runtime Filters diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 64a3ee4e6ff..d7716bf303b 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -435,6 +435,22 @@ object CometConf extends ShimCometConf { .booleanConf .createWithDefault(false) + val COMET_FORCE_SHJ_MAX_BUILD_SIZE: OptionalConfigEntry[Long] = + conf(s"$COMET_EXEC_CONFIG_PREFIX.forceShuffledHashJoin.maxBuildSize") + .category(CATEGORY_EXEC) + .doc(s"The build side size below which `${COMET_FORCE_SHJ.key}` converts a " + + "SortMergeJoin to ShuffledHashJoin. The size is Spark's planning estimate of the build " + + "child, or under AQE the materialized shuffle size of the build side. A build side at " + + "or over this size, or one with no statistics, keeps the SortMergeJoin. When unset, " + + "Spark's own rule applies: `spark.sql.autoBroadcastJoinThreshold` times the initial " + + "shuffle partition count (`spark.sql.adaptive.coalescePartitions.initialPartitionNum` " + + "when AQE and partition coalescing are both on and it is set, else " + + "`spark.sql.shuffle.partitions`). When broadcasts are disabled with a non-positive " + + "threshold, Spark's default threshold of 10 MB is used instead. A " + + s"non-positive value removes the limit. $TUNING_GUIDE.") + .bytesConf(ByteUnit.BYTE) + .createOptional + val COMET_EXEC_JOIN_DYNAMIC_FILTER_ENABLED: ConfigEntry[Boolean] = conf(s"$COMET_EXEC_CONFIG_PREFIX.join.dynamicFilter.enabled") .category(CATEGORY_EXEC) 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 e97ab50dea2..5cdd39116bd 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -767,7 +767,7 @@ case class CometExecRule(session: SparkSession) val planWithJoinRewritten = if (CometConf.COMET_FORCE_SHJ.get()) { normalizedPlan.transformUp { case p => - RewriteJoin.rewrite(p) + RewriteJoin.rewrite(p, conf) } } else { normalizedPlan diff --git a/spark/src/main/scala/org/apache/comet/rules/RewriteJoin.scala b/spark/src/main/scala/org/apache/comet/rules/RewriteJoin.scala index 2864eea4ed8..2b6bdc4f539 100644 --- a/spark/src/main/scala/org/apache/comet/rules/RewriteJoin.scala +++ b/spark/src/main/scala/org/apache/comet/rules/RewriteJoin.scala @@ -24,8 +24,10 @@ import org.apache.spark.sql.catalyst.plans.LeftSemi import org.apache.spark.sql.catalyst.plans.logical.Join import org.apache.spark.sql.execution.{SortExec, SparkPlan} import org.apache.spark.sql.execution.joins.{ShuffledHashJoinExec, SortMergeJoinExec} +import org.apache.spark.sql.internal.SQLConf -import org.apache.comet.CometSparkSessionExtensions.withFallbackReason +import org.apache.comet.CometConf +import org.apache.comet.CometSparkSessionExtensions.{withFallbackReason, withInfo} /** * Adapted from equivalent rule in Apache Gluten. @@ -64,7 +66,53 @@ object RewriteJoin extends JoinSelectionHelper { case _ => plan } - def rewrite(plan: SparkPlan): SparkPlan = plan match { + /** + * The largest build side the rewrite converts, or None when the limit is switched off. Unset + * means Spark's own rule for choosing a shuffled hash join by size: the broadcast threshold + * times the initial shuffle partition count. A non-positive threshold only disables broadcast + * joins and says nothing about the hash table an executor can hold, so Spark's default + * threshold stands in for it. + */ + private def maxBuildSize(conf: SQLConf): Option[BigInt] = + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.get(conf) match { + case Some(limit) if limit <= 0 => None + case Some(limit) => Some(BigInt(limit)) + case None => + val threshold = conf.autoBroadcastJoinThreshold + val perPartition = + if (threshold > 0) threshold else SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.defaultValue.get + Some(BigInt(perPartition) * conf.numShufflePartitions) + } + + /** + * Why the build side is too large to rewrite, or None when it fits. Like Spark's + * canBuildLocalHashMapBySize, the size must be strictly under the limit. The hash join cannot + * spill, so a build side without statistics is refused as well. Spark reports an unknown size + * as Long.MaxValue, which fails the comparison on its own. + */ + private def buildSideOverLimit( + smj: SortMergeJoinExec, + buildSide: BuildSide, + conf: SQLConf): Option[String] = + maxBuildSize(conf).flatMap { limit => + smj.logicalLink match { + case Some(join: Join) => + val buildSize = buildSide match { + case BuildLeft => join.left.stats.sizeInBytes + case BuildRight => join.right.stats.sizeInBytes + } + if (buildSize >= limit) { + Some( + s"build side size estimate of $buildSize bytes is not under the limit of " + + s"$limit bytes") + } else { + None + } + case _ => Some("no statistics are available for the build side") + } + } + + def rewrite(plan: SparkPlan, conf: SQLConf): SparkPlan = plan match { case smj: SortMergeJoinExec => getSmjBuildSide(smj) match { case Some(BuildRight) if smj.joinType == LeftSemi => @@ -75,15 +123,23 @@ object RewriteJoin extends JoinSelectionHelper { s"BuildRight with ${smj.joinType} is not supported") plan case Some(buildSide) => - ShuffledHashJoinExec( - smj.leftKeys, - smj.rightKeys, - smj.joinType, - buildSide, - smj.condition, - removeSort(smj.left), - removeSort(smj.right), - smj.isSkewJoin) + buildSideOverLimit(smj, buildSide, conf) match { + case Some(reason) => + // The sort-merge join still runs natively, so this is information for extended + // explain rather than a fallback reason. + withInfo(smj, s"Cannot rewrite SortMergeJoin to HashJoin: $reason") + plan + case None => + ShuffledHashJoinExec( + smj.leftKeys, + smj.rightKeys, + smj.joinType, + buildSide, + smj.condition, + removeSort(smj.left), + removeSort(smj.right), + smj.isSkewJoin) + } case _ => plan } case _ => plan diff --git a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala index 2163ab05e26..e93e01e4cc1 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala @@ -31,15 +31,20 @@ import org.apache.spark.sql.catalyst.TableIdentifier import org.apache.spark.sql.catalyst.analysis.UnresolvedRelation import org.apache.spark.sql.catalyst.expressions.{And, AttributeReference, IsNotNull} import org.apache.spark.sql.catalyst.optimizer.{BuildLeft, BuildRight} -import org.apache.spark.sql.comet.{CometBroadcastExchangeExec, CometBroadcastHashJoinExec, CometBroadcastNestedLoopJoinExec, CometFilterExec, CometHashJoinExec, CometNativeScanExec, CometSortMergeJoinExec, CometUnionExec} +import org.apache.spark.sql.catalyst.plans.Inner +import org.apache.spark.sql.catalyst.plans.logical.Join +import org.apache.spark.sql.comet.{CometBroadcastExchangeExec, CometBroadcastHashJoinExec, CometBroadcastNestedLoopJoinExec, CometFilterExec, CometHashJoinExec, CometNativeScanExec, CometSortExec, CometSortMergeJoinExec, CometUnionExec} import org.apache.spark.sql.execution.{LocalTableScanExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec} import org.apache.spark.sql.execution.datasources.SchemaColumnConvertNotSupportedException import org.apache.spark.sql.execution.exchange.ReusedExchangeExec +import org.apache.spark.sql.execution.joins.{ShuffledHashJoinExec, SortMergeJoinExec} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{ArrayType, IntegerType, MetadataBuilder, StructField, StructType} -import org.apache.comet.CometConf +import org.apache.comet.{CometConf, CometExplainInfo, ExtendedExplainInfo} +import org.apache.comet.CometSparkSessionExtensions.hasFallbackReason +import org.apache.comet.rules.RewriteJoin class CometJoinSuite extends CometTestBase { @@ -889,6 +894,199 @@ class CometJoinSuite extends CometTestBase { } } + // The MERGE hint makes Spark plan a SortMergeJoin, which forceShuffledHashJoin then rewrites. + private val forcedSmjQuery = + "SELECT /*+ MERGE(tbl_a) */ * FROM tbl_a JOIN tbl_b ON tbl_a._2 = tbl_b._1" + + // tbl_b is the smaller side, so the rewrite builds on the right. + private def withForcedSmjTables(f: => Unit): Unit = { + withParquetTable((0 until 1000).map(i => (i, i % 5)), "tbl_a") { + withParquetTable((0 until 10).map(i => (i % 10, i + 2)), "tbl_b") { + f + } + } + } + + private def rewriteReason(reason: String): String = + s"Cannot rewrite SortMergeJoin to HashJoin: $reason" + + // Checks that the SortMergeJoin and its sorts survived and that the kept join, which still + // runs natively, is not reported as a fallback. Returns Spark's size estimate of the build + // side, read through the same logical link the rewrite used. + private def assertSortMergeJoinKept(plan: SparkPlan): BigInt = { + val joins = collect(plan) { case j: CometSortMergeJoinExec => j } + assert(joins.size == 1, plan) + assert(collect(plan) { case s: CometSortExec => s }.size == 2, plan) + assert(collect(plan) { case j: CometHashJoinExec => j }.isEmpty, plan) + val kept = joins.head + assert(!hasFallbackReason(kept) && !hasFallbackReason(kept.originalPlan), plan) + val reasons = new ExtendedExplainInfo().getFallbackReasons(plan) + assert(!reasons.exists(_.contains("SortMergeJoin")), reasons) + kept.originalPlan.logicalLink match { + case Some(join: Join) => join.right.stats.sizeInBytes + case other => fail(s"Expected a Join logical link, got $other") + } + } + + // The refusal is recorded as information on the kept join and rendered as COMET-INFO. + private def assertKeptWithInfo(plan: SparkPlan, expected: String): Unit = { + val join = collect(plan) { case j: CometSortMergeJoinExec => j }.head + assert(join.getTagValue(CometExplainInfo.EXTENSION_INFO).exists(_.contains(expected)), plan) + val verbose = new ExtendedExplainInfo().generateVerboseInfo(plan) + assert(verbose.contains("[COMET-INFO:") && verbose.contains(expected), verbose) + } + + private def assertRewrittenToHashJoin(plan: SparkPlan): Unit = { + val joins = collect(plan) { case j: CometHashJoinExec => j } + assert(joins.map(_.buildSide) == Seq(BuildRight), plan) + assert(collect(plan) { case s: CometSortExec => s }.isEmpty, plan) + assert(collect(plan) { case j: CometSortMergeJoinExec => j }.isEmpty, plan) + } + + private def assertKeptWithReason(plan: SparkPlan, limit: Long): Unit = { + val estimate = assertSortMergeJoinKept(plan) + assert(estimate >= limit, estimate) + assertKeptWithInfo( + plan, + rewriteReason( + s"build side size estimate of $estimate bytes is not under the limit of $limit bytes")) + } + + for (adaptive <- Seq(false, true)) { + test(s"forceShuffledHashJoin keeps SortMergeJoin over maxBuildSize, AQE=$adaptive") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_FORCE_SHJ.key -> "true", + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> "1") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertKeptWithReason(cometPlan, 1) + } + } + } + + test(s"forceShuffledHashJoin rewrites SortMergeJoin under maxBuildSize, AQE=$adaptive") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_FORCE_SHJ.key -> "true", + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> "1g") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertRewrittenToHashJoin(cometPlan) + } + } + } + } + + test("forceShuffledHashJoin matches Spark's strict size rule at the limit") { + // AQE off, so the estimate read from the first plan is the one the later runs compare. + withForcedSmjTables { + var estimate = BigInt(0) + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_FORCE_SHJ.key -> "true", + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> "1") { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + estimate = assertSortMergeJoinKept(cometPlan) + } + // A build side exactly at the limit is not under it, as in canBuildLocalHashMapBySize. + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_FORCE_SHJ.key -> "true", + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> estimate.toString) { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertKeptWithReason(cometPlan, estimate.toLong) + } + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_FORCE_SHJ.key -> "true", + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> (estimate + 1).toString) { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertRewrittenToHashJoin(cometPlan) + } + } + } + + test("forceShuffledHashJoin with a non-positive maxBuildSize rewrites regardless of size") { + withSQLConf( + CometConf.COMET_FORCE_SHJ.key -> "true", + CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> "0", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "1") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertRewrittenToHashJoin(cometPlan) + } + } + } + + test("forceShuffledHashJoin without maxBuildSize uses Spark's shuffled hash join size rule") { + // The limit is autoBroadcastJoinThreshold times the shuffle partition count. + withSQLConf( + CometConf.COMET_FORCE_SHJ.key -> "true", + SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "1") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertKeptWithReason(cometPlan, 2) + } + } + withSQLConf( + CometConf.COMET_FORCE_SHJ.key -> "true", + SQLConf.SHUFFLE_PARTITIONS.key -> "2", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "1g") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertRewrittenToHashJoin(cometPlan) + } + } + } + + test("forceShuffledHashJoin uses Spark's default threshold when broadcasts are disabled") { + // Disabling broadcasts says nothing about the hash table an executor can hold, so the limit + // falls back to Spark's default threshold times the partition count. AQE is off so the + // decision uses the planning estimate, which fileCompressionFactor scales past 10 MB without + // writing a 10 MB table. + val defaultThreshold = SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.defaultValue.get + assert(defaultThreshold == 10L * 1024 * 1024) + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_FORCE_SHJ.key -> "true", + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertRewrittenToHashJoin(cometPlan) + } + withSQLConf(SQLConf.FILE_COMPRESSION_FACTOR.key -> "1000000") { + withForcedSmjTables { + val (_, cometPlan) = checkSparkAnswerAndOperator(sql(forcedSmjQuery)) + assertKeptWithReason(cometPlan, defaultThreshold) + } + } + } + } + + test("forceShuffledHashJoin keeps a SortMergeJoin that has no statistics") { + val left = spark.range(1).toDF("a").queryExecution.sparkPlan + val right = spark.range(1).toDF("b").queryExecution.sparkPlan + def smj: SortMergeJoinExec = + SortMergeJoinExec(left.output, right.output, Inner, None, left, right) + assert(smj.logicalLink.isEmpty) + + val kept = smj + assert(RewriteJoin.rewrite(kept, spark.sessionState.conf) eq kept) + assert(!hasFallbackReason(kept)) + val info = kept.getTagValue(CometExplainInfo.EXTENSION_INFO) + assert( + info.exists(_.contains(rewriteReason("no statistics are available for the build side"))), + info) + + withSQLConf(CometConf.COMET_FORCE_SHJ_MAX_BUILD_SIZE.key -> "-1") { + val rewritten = RewriteJoin.rewrite(smj, spark.sessionState.conf) + assert(rewritten.isInstanceOf[ShuffledHashJoinExec], rewritten) + } + } + test("BroadcastHashJoin with LeftAnti and NOT IN subquery (null-aware)") { withSQLConf( SQLConf.PREFER_SORTMERGEJOIN.key -> "false",