Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions docs/api/sql/NearestNeighbourSearching.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)) =>
Expand All @@ -640,7 +649,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy {
rightShape,
spatialPredicate = SpatialPredicate.KNN,
isGeography = useSpheroidUnwrapped,
condition,
extraCondition,
Some(k)))

case _ => None
Expand Down Expand Up @@ -841,7 +850,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy {
spatialPredicate = null,
isGeography,
condition,
extractExtraKNNJoinCondition(condition)) :: Nil
extraCondition) :: Nil
}

private def planDistanceJoin(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand All @@ -995,7 +986,7 @@ class JoinQueryDetector(sparkSession: SparkSession) extends SparkStrategy {
spatialPredicate,
isGeography,
condition = null,
extraCondition = None) :: Nil
extraCondition = extraCondition) :: Nil
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 = {
Expand Down
Loading