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


##########
spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala:
##########
@@ -1219,6 +1219,298 @@ class CometJoinSuite extends CometTestBase {
     }
   }
 
+  // A hash join reports its streamed side's ordering, so Spark drops a sort 
above it when the
+  // streamed side already arrives partitioned and sorted. Here that side is a 
cached table that
+  // reaches the native join through Comet's JVM to native conversion.
+  private def withSortedCache(data: DataFrame, order: Column*)(f: DataFrame => 
Unit): Unit = {
+    withSQLConf(
+      CometConf.COMET_CONVERT_FROM_IN_MEMORY_CACHE_ENABLED.key -> "true",
+      SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+      SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+      SQLConf.SHUFFLE_PARTITIONS.key -> "2") {
+      val cached = data.repartition(2, data("_1")).sortWithinPartitions(order: 
_*).cache()
+      try {
+        cached.count()
+        f(cached)
+      } finally {
+        cached.unpersist()
+      }
+    }
+  }
+
+  /**
+   * Checks the query's rows against Spark and that every partition is sorted 
by `order`, which is
+   * bound to the query's output. Returns the plan's single native hash join 
and its build side.
+   */
+  private def checkJoinKeepsOrder(
+      query: => DataFrame,
+      order: SortOrder*): (CometBinaryExec, BuildSide) = {
+    // 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 = query.collect().map(_.toString).sorted.toSeq
+    }
+    val df = query
+    val rows = df.collect()
+    assert(rows.map(_.toString).sorted.toSeq == expected)
+    val plan = df.queryExecution.executedPlan
+    val joins = collect(plan) {
+      case j: CometHashJoinExec => (j, j.buildSide)
+      case j: CometBroadcastHashJoinExec => (j, j.buildSide)
+    }
+    assert(joins.length == 1, plan)
+    // Read before the plan runs again below, so the count covers the one run 
above.
+    assert(joins.head._1.metrics("output_rows").value == rows.length, plan)
+    assert(joins.head._1.metrics("build_input_rows").value > 0, plan)
+    // Spark dropped the sort above the join, so the order comes from the join 
alone.
+    val sortsAboveJoin = collect(plan) {
+      case s: SortExec => s
+      case s: CometSortExec => s
+    }.filter(s => find(s)(_ eq joins.head._1).isDefined)
+    assert(sortsAboveJoin.isEmpty, plan)
+    val ordering = new InterpretedOrdering(order)
+    // Propagates the session's confs, such as the batch size, to the tasks as 
collect() does.
+    val parts = SQLExecution.withSQLConfPropagated(spark) {
+      df.queryExecution.toRdd.map(_.copy()).glom().collect()
+    }
+    parts.zipWithIndex.foreach { case (part, i) =>
+      val outOfOrder = part.zip(part.drop(1)).indexWhere { case (a, b) => 
ordering.gt(a, b) }
+      assert(outOfOrder < 0, s"partition $i is out of order at row 
$outOfOrder\n$plan")
+    }
+    joins.head
+  }
+
+  private def keyOrder(
+      dataType: DataType,
+      direction: SortDirection,
+      ordinal: Int = 0): SortOrder =
+    SortOrder(BoundReference(ordinal, dataType, nullable = true), direction)
+
+  // (hint, join type, extra condition, AQE, native batch size)
+  private val streamedOrderCases =
+    (for (hint <- Seq("shuffle_hash", "broadcast"); joinType <- 
Seq("left_outer", "right_outer"))
+      yield (hint, joinType, false, false, 8192)) ++ Seq(
+      ("shuffle_hash", "left_outer", true, false, 8192),
+      ("shuffle_hash", "right_outer", false, true, 8192),
+      ("shuffle_hash", "left_outer", false, false, 100))
+
+  for ((hint, joinType, withCondition, adaptive, batchSize) <- 
streamedOrderCases) {
+    test(
+      s"$hint $joinType join keeps a cached streamed side's order, " +
+        s"condition=$withCondition, AQE=$adaptive, batchSize=$batchSize") {
+      withSQLConf(CometConf.COMET_BATCH_SIZE.key -> batchSize.toString) {
+        withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+          withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+            withSortedCache(spark.table("big"), $"_1") { streamed =>
+              val build = spark.table("small").hint(hint)
+              val keys = streamed("_1") === build("_1")
+              val on = if (withCondition) keys && build("_2") > 3 else keys
+              def query: DataFrame = {
+                val joined =
+                  if (joinType == "left_outer") streamed.join(build, on, 
joinType)
+                  else build.join(streamed, on, joinType)
+                joined
+                  .select(streamed("_1").as("k"), streamed("_2").as("v"), 
build("_2").as("w"))
+                  .sortWithinPartitions("k")
+              }
+              withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> 
adaptive.toString) {
+                val (join, buildSide) =
+                  checkJoinKeepsOrder(query, keyOrder(IntegerType, Ascending))
+                assert(buildSide == (if (joinType == "left_outer") BuildRight 
else BuildLeft))
+                hint match {
+                  case "broadcast" => 
assert(join.isInstanceOf[CometBroadcastHashJoinExec])
+                  case _ => assert(join.isInstanceOf[CometHashJoinExec])
+                }
+                if (batchSize < 1000) {
+                  // Several batches per partition come back to the JVM from 
the native join.
+                  val partitions = 2
+                  val plan = query.queryExecution.executedPlan
+                  val batches = collectFirst(plan) { case c: 
ColumnarToRowTransition =>
+                    SQLExecution.withSQLConfPropagated(spark) {
+                      c.child.executeColumnar().count()
+                    }
+                  }
+                  assert(batches.exists(_ > partitions), plan)
+                }
+              }
+            }
+          }
+        }
+      }
+    }
+  }
+
+  test("shuffle_hash join keeps a descending, nulls last streamed order with 
null keys") {
+    val data = (0 until 10000).map(i => (if (i % 7 == 0) None else Some(i % 
100), i))
+    withParquetTable(data, "big") {
+      withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+        withSortedCache(spark.table("big"), $"_1".desc_nulls_last) { streamed 
=>
+          val build = spark.table("small").hint("shuffle_hash")
+          checkJoinKeepsOrder(
+            streamed
+              .join(build, streamed("_1") === build("_1"), "left_outer")
+              .select(streamed("_1").as("k"), streamed("_2").as("v"), 
build("_2").as("w"))
+              .sortWithinPartitions($"k".desc_nulls_last),
+            keyOrder(IntegerType, Descending))
+        }
+      }
+    }
+  }
+
+  test("shuffle_hash join keeps a descending streamed order on a double with 
NaN and -0.0") {
+    def d(i: Int): Option[Double] = i % 9 match {
+      case 0 => None
+      case 1 => Some(Double.NaN)
+      case 2 => Some(-0.0)
+      case 3 => Some(0.0)
+      case _ => Some(i % 50 - 25.5)
+    }
+    withParquetTable((0 until 10000).map(i => (i % 100, i, d(i))), "big") {
+      withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+        withSortedCache(spark.table("big"), $"_3".desc_nulls_last) { streamed 
=>
+          val build = spark.table("small").hint("shuffle_hash")
+          checkJoinKeepsOrder(
+            streamed
+              .join(build, streamed("_1") === build("_1"), "left_outer")
+              .select(streamed("_3").as("d"), streamed("_2").as("v"), 
build("_2").as("w"))
+              .sortWithinPartitions($"d".desc_nulls_last),
+            keyOrder(DoubleType, Descending))
+        }
+      }
+    }
+  }
+
+  test("shuffle_hash join keeps a two-key streamed order") {
+    withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+      withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+        withSortedCache(spark.table("big"), $"_1", $"_2".desc) { streamed =>
+          val build = spark.table("small").hint("shuffle_hash")
+          checkJoinKeepsOrder(
+            streamed
+              .join(build, streamed("_1") === build("_1"), "left_outer")
+              .select(streamed("_1").as("k"), streamed("_2").as("v"), 
build("_2").as("w"))
+              .sortWithinPartitions($"k", $"v".desc),
+            keyOrder(IntegerType, Ascending),
+            keyOrder(IntegerType, Descending, ordinal = 1))
+        }
+      }
+    }
+  }
+
+  // The streamed side is sorted inside the native plan, and the join's output 
keeps that order in
+  // every partition. EliminateSorts would drop that sort as one below a join.
+  test("shuffle_hash join keeps a natively sorted streamed side's order") {
+    withSQLConf(
+      SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> EliminateSorts.ruleName,
+      SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+      SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+      SQLConf.SHUFFLE_PARTITIONS.key -> "2") {
+      withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+        withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+          val streamed = spark.table("big").repartition(2, 
$"_1").sortWithinPartitions("_1")
+          val build = spark.table("small").hint("shuffle_hash")
+          val (join, buildSide) = checkJoinKeepsOrder(
+            streamed
+              .join(build, streamed("_1") === build("_1"), "left_outer")
+              .select(streamed("_1").as("k"), streamed("_2").as("v"), 
build("_2").as("w"))
+              .sortWithinPartitions("k"),
+            keyOrder(IntegerType, Ascending))
+          assert(join.isInstanceOf[CometHashJoinExec])
+          assert(buildSide == BuildRight)
+        }
+      }
+    }
+  }
+
+  // Spark plans NOT IN as a null-aware anti join that streams the left side. 
Comet runs it with
+  // the left side as DataFusion's build side, which gives up that side's 
batch order. Sparse keys
+  // keep DataFusion on its hash map build, which stores the batches last to 
first.
+  test("null-aware anti join keeps a cached streamed side's order across 
batches") {
+    withSQLConf(CometConf.COMET_BATCH_SIZE.key -> "100") {
+      withParquetTable((0 until 10000).map(i => (i % 100 * 1000, i)), "big") {
+        withParquetTable((0 until 10).map(i => (i * 10000, i)), "small") {
+          withSortedCache(spark.table("big"), $"_1") { streamed =>
+            withSQLConf(SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> 
"10485760") {
+              withTempView("streamed") {
+                streamed.createOrReplaceTempView("streamed")
+                def query: DataFrame =
+                  sql("SELECT _1 AS k, _2 AS v FROM streamed WHERE _1 NOT IN " 
+
+                    "(SELECT _1 FROM small)").sortWithinPartitions("k")
+                val (join, _) = checkJoinKeepsOrder(query, 
keyOrder(IntegerType, Ascending))
+                join match {
+                  case j: CometBroadcastHashJoinExec => 
assert(j.isNullAwareAntiJoin, j)
+                  case j => fail(s"expected a null-aware broadcast hash join: 
$j")
+                }
+                // The build side spans several batches per partition. Read on 
a fresh plan that
+                // runs once, so the count covers that run alone.
+                val partitions = 2
+                val df = query
+                assert(df.collect().nonEmpty)
+                val plan = df.queryExecution.executedPlan
+                val buildBatches = collectFirst(plan) { case j: 
CometBroadcastHashJoinExec =>
+                  j.metrics("build_input_batches").value
+                }
+                assert(buildBatches.exists(_ > partitions), plan)
+              }
+            }
+          }
+        }
+      }
+    }
+  }
+
+  test("only outer hash joins that can move unmatched streamed rows send their 
ordering") {
+    withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+      withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+        withSortedCache(spark.table("big"), $"_1") { streamed =>
+          val build = spark.table("small").hint("shuffle_hash")
+          def query(joinType: String): DataFrame = streamed
+            .join(build, streamed("_1") === build("_1"), joinType)
+            .select(streamed("_1").as("k"), streamed("_2").as("v"), 
build("_2").as("w"))
+            .sortWithinPartitions("k")
+          val (inner, _) = checkJoinKeepsOrder(query("inner"), 
keyOrder(IntegerType, Ascending))
+          assert(inner.nativeOp.getHashJoin.getOutputOrderingCount == 0)
+          val (leftOuter, buildSide) =
+            checkJoinKeepsOrder(query("left_outer"), keyOrder(IntegerType, 
Ascending))
+          assert(buildSide == BuildRight)
+          assert(leftOuter.nativeOp.getHashJoin.getOutputOrderingCount > 0)

Review Comment:
   > Could the `getOutputOrderingCount > 0` assertion move into the matrix loop 
instead?
   
   Moved. Every matrix case now asserts that the join sends its ordering, and 
the separate test keeps only the inner check.
   



-- 
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