dwsmith1983 commented on code in PR #6785:
URL: https://github.com/apache/datafusion-comet/pull/6785#discussion_r4233338965


##########
spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala:
##########
@@ -1219,6 +1251,237 @@ class CometJoinSuite extends CometTestBase {
     }
   }
 
+  // Spark removes a sort above a sort-merge join whose output ordering 
satisfies it, so the
+  // forced hash join must be chosen before that happens or the sort is lost 
(#6770).
+  private def withSortLossConf(adaptive: Boolean)(f: => Unit): Unit = 
withSQLConf(
+    SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString,
+    SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+    SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+    SQLConf.SHUFFLE_PARTITIONS.key -> "2",
+    CometConf.COMET_FORCE_SHJ.key -> "true") {
+    withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+      withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+        f
+      }
+    }
+  }
+
+  // Checks that a hash join ran, that the rows match Spark's, and that every 
partition is sorted
+  // on the first column. A local sort leaves rows with equal keys in any 
order, so the rows are
+  // compared without their order.
+  private def checkPartitionsSortedOverHashJoin(df: => DataFrame): Unit = {
+    // withSQLConf returns Unit on Spark 3.x, so the expected rows are 
assigned inside it.
+    var expected = Seq.empty[String]
+    withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+      expected = df.collect().map(_.toString).sorted.toSeq
+    }
+    val cometDf = df
+    assert(cometDf.collect().map(_.toString).sorted.toSeq == expected)
+    val cometPlan = cometDf.queryExecution.executedPlan
+    assert(collect(cometPlan) { case j: CometHashJoinExec => j }.nonEmpty, 
cometPlan)
+    val sorted = cometDf.queryExecution.toRdd
+      .mapPartitions { it =>
+        val keys = it.map(_.getInt(0)).toArray
+        Iterator(keys.sameElements(keys.sorted))
+      }
+      .collect()
+    assert(sorted.forall(identity), sorted.mkString(", "))
+  }
+
+  private val sortLossJoin =
+    "SELECT big._1 AS k, big._2 AS v FROM big JOIN small ON big._1 = small._1"
+
+  for (adaptive <- Seq(false, true)) {
+    test(s"forceShuffledHashJoin keeps sortWithinPartitions on the join key, 
AQE=$adaptive") {
+      withSortLossConf(adaptive) {
+        
checkPartitionsSortedOverHashJoin(sql(sortLossJoin).sortWithinPartitions("k"))
+      }
+    }
+
+    test(s"forceShuffledHashJoin keeps SORT BY on the join key, 
AQE=$adaptive") {
+      withSortLossConf(adaptive) {
+        checkPartitionsSortedOverHashJoin(sql(s"$sortLossJoin SORT BY k"))

Review Comment:
   > Would one of them be enough, given that the second adds a run per AQE 
setting without a different plan?
   
   Yes. Both forms analyze to the same `Sort` with `global = false`, so I 
dropped the `SORT BY` test and kept the `sortWithinPartitions` one.
   



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

To unsubscribe, e-mail: [email protected]

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


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

Reply via email to