diff --git a/docs/api/sql/NearestNeighbourSearching.md b/docs/api/sql/NearestNeighbourSearching.md index 26712a07a18..c081a2c2489 100644 --- a/docs/api/sql/NearestNeighbourSearching.md +++ b/docs/api/sql/NearestNeighbourSearching.md @@ -68,6 +68,10 @@ CACHE TABLE knnResult; SELECT * FROM knnResult WHERE condition; ``` +Additional `ON` predicates that remain in the kNN join condition are evaluated on the neighbors selected by `ST_KNN`. If a predicate rejects a selected pair, the join does not choose another neighbor to replace it, so it may return fewer than `k` neighbors for a query geometry. Spark may push predicates that reference only one input below the join, as described above. + +Use only one `ST_KNN` predicate per join condition. When combining it with another spatial predicate, place `ST_KNN` before predicates such as `ST_Intersects` in the `ON` condition; otherwise, query planning may reject the join. + ### Optimization Barrier Use the `barrier` function to prevent filter pushdown and control predicate evaluation order in complex spatial joins. This function creates an optimization barrier by evaluating boolean expressions at runtime. diff --git a/spark/common/src/main/scala/org/apache/spark/sql/sedona_sql/strategy/join/JoinQueryDetector.scala b/spark/common/src/main/scala/org/apache/spark/sql/sedona_sql/strategy/join/JoinQueryDetector.scala index 68aad366b8b..e228fca7ee8 100644 --- a/spark/common/src/main/scala/org/apache/spark/sql/sedona_sql/strategy/join/JoinQueryDetector.scala +++ b/spark/common/src/main/scala/org/apache/spark/sql/sedona_sql/strategy/join/JoinQueryDetector.scala @@ -273,6 +273,15 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { val joinConditionMatcher = OptimizableJoinCondition(left, right) val queryDetection: Option[JoinQueryDetection] = condition.flatMap { case joinConditionMatcher(predicate, extraCondition) => + // ST_KNN must be implemented by the join, never evaluated as a per-pair filter. + if (extraCondition.exists(_.exists(_.isInstanceOf[ST_KNN]))) { + val message = if (predicate.isInstanceOf[ST_KNN]) { + "Only one ST_KNN predicate is supported per join condition" + } else { + "Place ST_KNN before other spatial predicates in the join condition" + } + throw new UnsupportedOperationException(message) + } predicate match { // ST_Contains / ST_Intersects / ST_Within / ST_Equals are InferredExpression (not // ST_Predicate) so they can't sit inside getJoinDetection; they're also the only @@ -627,7 +636,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { rightShape, spatialPredicate = SpatialPredicate.KNN, isGeography = false, - condition, + extraCondition, Some(k))) case ST_KNN(Seq(leftShape, rightShape, k, useSpheroid)) => @@ -640,7 +649,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { rightShape, spatialPredicate = SpatialPredicate.KNN, isGeography = useSpheroidUnwrapped, - condition, + extraCondition, Some(k))) case _ => None @@ -841,7 +850,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { spatialPredicate = null, isGeography, condition, - extractExtraKNNJoinCondition(condition)) :: Nil + extraCondition) :: Nil } private def planDistanceJoin( @@ -903,24 +912,6 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { } } - private def extractExtraKNNJoinCondition(condition: Expression): Option[Expression] = { - condition match { - case and: And => - // Check both left and right sides for ST_KNN or ST_AKNN - if (and.left.isInstanceOf[ST_KNN]) { - Some(and.right) - } else if (and.right.isInstanceOf[ST_KNN]) { - Some(and.left) - } else { - None - } - case _: ST_KNN => - None - case _ => - Some(condition) - } - } - private def planBroadcastJoin( left: LogicalPlan, right: LogicalPlan, @@ -980,7 +971,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { spatialPredicate, isGeography, condition = null, - extraCondition = None) :: Nil + extraCondition = extraCondition) :: Nil } else { // broadcast is on object side return BroadcastObjectSideKNNJoinExec( @@ -995,7 +986,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy { spatialPredicate, isGeography, condition = null, - extraCondition = None) :: Nil + extraCondition = extraCondition) :: Nil } } } diff --git a/spark/common/src/test/scala/org/apache/sedona/sql/KnnJoinSuite.scala b/spark/common/src/test/scala/org/apache/sedona/sql/KnnJoinSuite.scala index 4ca17598d30..b21910a6dd4 100644 --- a/spark/common/src/test/scala/org/apache/sedona/sql/KnnJoinSuite.scala +++ b/spark/common/src/test/scala/org/apache/sedona/sql/KnnJoinSuite.scala @@ -22,7 +22,7 @@ import org.apache.spark.sql.catalyst.expressions.Literal import org.apache.spark.sql.functions.col import org.apache.spark.sql.sedona_sql.UDT.GeometryUDT import org.apache.spark.sql.sedona_sql.expressions.st_constructors.ST_GeomFromText -import org.apache.spark.sql.sedona_sql.strategy.join.KNNJoinExec +import org.apache.spark.sql.sedona_sql.strategy.join.{BroadcastObjectSideKNNJoinExec, BroadcastQuerySideKNNJoinExec, KNNJoinExec} import org.apache.spark.sql.types.{IntegerType, StringType, StructField, StructType} import org.apache.spark.sql.{Column, DataFrame, Row, SparkSession} import org.apache.spark.sql.functions.expr @@ -486,6 +486,134 @@ class KnnJoinSuite extends TestBaseScala with TableDrivenPropertyChecks { assert(knnJoined.count() > 0) } + + val residualJoinStrategies: Seq[(String, String, Class[_])] = Seq( + ("regular", "", classOf[KNNJoinExec]), + ("broadcast query", "/*+ BROADCAST(q) */", classOf[BroadcastQuerySideKNNJoinExec]), + ("broadcast object", "/*+ BROADCAST(o) */", classOf[BroadcastObjectSideKNNJoinExec])) + val knnPredicate = "ST_KNN(q.g, o.g, 1, false)" + val scorePredicate = "q.score > o.score" + val ceilingPredicate = "q.ceiling > o.floor" + val residualConditions = Seq( + ( + "single residual with default distance metric", + s"ST_KNN(q.g, o.g, 1) AND $scorePredicate", + Seq((1, 101), (4, 104))), + ("KNN first", s"($knnPredicate AND $scorePredicate) AND $ceilingPredicate", Seq((1, 101))), + ("KNN middle", s"$scorePredicate AND ($knnPredicate AND $ceilingPredicate)", Seq((1, 101))), + ("KNN last", s"($scorePredicate AND $ceilingPredicate) AND $knnPredicate", Seq((1, 101))), + ( + "spatial residual with KNN first", + s"$knnPredicate AND ST_DWithin(q.g, o.g, CAST(q.id AS DOUBLE))", + Seq((2, 102), (3, 103), (4, 104)))) + + for { + (strategy, hint, expectedPlan) <- residualJoinStrategies + queriesFirst <- Seq(true, false) + (conditionName, condition, expected) <- residualConditions + } { + it(s"KNN residuals: $strategy, queriesFirst=$queriesFirst, $conditionName") { + withKnnResidualInputs { + val relations = if (queriesFirst) { + "knn_pred_queries q JOIN knn_pred_objects o" + } else { + "knn_pred_objects o JOIN knn_pred_queries q" + } + val joined = sparkSession.sql(s"SELECT $hint q.id, o.id FROM $relations ON $condition") + val plan = joined.queryExecution.executedPlan + assert(plan.find(_.getClass == expectedPlan).isDefined, plan.toString) + val pairs = joined.collect().map(row => (row.getInt(0), row.getInt(1))).sorted.toSeq + pairs should be(expected) + } + } + } + + val additionalKnnPredicates = Seq( + ("different k", s"ST_KNN(q.g, o.g, 2, false) AND $knnPredicate"), + ("reciprocal", s"$knnPredicate AND ST_KNN(o.g, q.g, 1, false)"), + ("nested in OR", s"$knnPredicate AND (ST_KNN(q.g, o.g, 2, false) OR $scorePredicate)")) + for { + (strategy, hint, _) <- residualJoinStrategies + (conditionName, condition) <- additionalKnnPredicates + } { + it(s"KNN residuals reject multiple predicates during planning: $strategy, $conditionName") { + withKnnResidualInputs { + val exception = intercept[UnsupportedOperationException] { + sparkSession + .sql(s"SELECT $hint q.id, o.id FROM knn_pred_queries q JOIN knn_pred_objects o " + + s"ON $condition") + .queryExecution + .executedPlan + } + exception.getMessage should be( + "Only one ST_KNN predicate is supported per join condition") + } + } + } + + val spatialPredicatesBeforeKnn = + Seq("ST_Intersects(ST_Buffer(q.g, 2), o.g)", "ST_DWithin(q.g, o.g, 2)") + for { + (strategy, hint, _) <- residualJoinStrategies + spatialPredicate <- spatialPredicatesBeforeKnn + } { + it(s"KNN residuals reject KNN after $spatialPredicate during planning: $strategy") { + withKnnResidualInputs { + val exception = intercept[UnsupportedOperationException] { + sparkSession + .sql(s"SELECT $hint q.id, o.id FROM knn_pred_queries q JOIN knn_pred_objects o " + + s"ON $spatialPredicate AND $knnPredicate") + .queryExecution + .executedPlan + } + exception.getMessage should include( + "Place ST_KNN before other spatial predicates in the join condition") + } + } + } + } + + private def withKnnResidualInputs(body: => Unit): Unit = { + withConf( + Map( + "spark.sql.adaptive.enabled" -> "false", + "spark.sql.autoBroadcastJoinThreshold" -> "-1", + "spark.sedona.join.autoBroadcastJoinThreshold" -> "-1", + "spark.sedona.join.knn.includeTieBreakers" -> "false")) { + try { + sparkSession + .sql(""" + |SELECT id, ST_Point(x, x) AS g, score, ceiling + |FROM VALUES + | (1, 0.0D, 30, 100), + | (2, 10.0D, 10, 100), + | (3, 20.0D, CAST(NULL AS INT), 100), + | (4, 30.0D, 30, 0) + |AS q(id, x, score, ceiling) + |""".stripMargin) + .coalesce(1) + .createOrReplaceTempView("knn_pred_queries") + // Different column layouts expose incorrectly bound residual attributes. Query 2's + // nearest object fails its score predicate; the farther passing object must not replace it. + sparkSession + .sql(""" + |SELECT 'object' AS label, score, ST_Point(x, x) AS g, id, floor + |FROM VALUES + | (101, 1.0D, 20, 50), + | (102, 11.0D, 20, 50), + | (103, 21.0D, 20, 50), + | (104, 31.0D, 20, 50), + | (105, 12.0D, 0, 50) + |AS o(id, x, score, floor) + |""".stripMargin) + .coalesce(1) + .createOrReplaceTempView("knn_pred_objects") + body + } finally { + sparkSession.catalog.dropTempView("knn_pred_queries") + sparkSession.catalog.dropTempView("knn_pred_objects") + } + } } private def withOptimizationMode(mode: String)(body: => Unit): Unit = {