diff --git a/src/main/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSides.java b/src/main/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSides.java new file mode 100644 index 0000000..ac499c5 --- /dev/null +++ b/src/main/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSides.java @@ -0,0 +1,141 @@ +/***************************************************************************** + * + * This MobilityDB code is provided under The PostgreSQL License. + * Copyright (c) 2020-2026, Université libre de Bruxelles and MobilityDB + * contributors + * + * Permission to use, copy, modify, and distribute this software and its + * documentation for any purpose, without fee, and without a written + * agreement is hereby retained provided that the above copyright notice and + * this paragraph and the following two paragraphs appear in all copies. + * + * IN NO EVENT SHALL UNIVERSITE LIBRE DE BRUXELLES BE LIABLE TO ANY PARTY FOR + * DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES + * INCLUDING LOST PROFITS, ARISING OUT OF THE USE OF THIS SOFTWARE AND ITS + * DOCUMENTATION, EVEN IF UNIVERSITE LIBRE DE BRUXELLES HAS BEEN ADVISED OF + * THE POSSIBILITY OF SUCH DAMAGE. + * + * UNIVERSITE LIBRE DE BRUXELLES SPECIFICALLY DISCLAIMS ANY WARRANTIES, + * INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY + * AND FITNESS FOR A PARTICULAR PURPOSE. THE SOFTWARE PROVIDED HEREUNDER IS + * ON AN "AS IS" BASIS, AND UNIVERSITE LIBRE DE BRUXELLES HAS NO OBLIGATIONS + * TO PROVIDE MAINTENANCE, SUPPORT, UPDATES, ENHANCEMENTS, OR MODIFICATIONS. + * + *****************************************************************************/ + +package org.mobilitydb.spark.catalyst; + +import java.util.ArrayList; +import java.util.List; + +import org.apache.spark.sql.catalyst.expressions.And; +import org.apache.spark.sql.catalyst.expressions.EqualTo; +import org.apache.spark.sql.catalyst.expressions.Expression; +import org.apache.spark.sql.catalyst.plans.logical.Join; +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan; +import org.apache.spark.sql.catalyst.plans.logical.Repartition; +import org.apache.spark.sql.catalyst.rules.Rule; + +import scala.runtime.AbstractPartialFunction; + +/** + * Computes each side of a nested-loop join once, where a side calls a user-defined function. + * + * A join with no equality between its two sides is executed as a nested loop over pairs of + * partitions, and a side that is a chain of scans and projections is RECOMPUTED for every + * partition of the other side: with thirty-two partitions a side, each row is read, filtered and + * projected thirty-two times. Where that projection calls a user-defined function, as a clip of a + * trajectory does, the recomputation dominates the join. Repartitioning a side puts a boundary + * under it, so its rows are computed once and every pair reads them back. + * + * The rule leaves alone a join that has an equality between its sides, which is executed by hash + * and reads each side once already, and a side that carries no call, whose recomputation is a + * scan Spark is good at. A side already behind a boundary is left as it is, so the rule reaches a + * fixed point. + */ +public final class MaterializeNestedLoopSides extends Rule { + + private final int partitions; + + public MaterializeNestedLoopSides(int partitions) { + this.partitions = Math.max(1, partitions); + } + + @Override + public LogicalPlan apply(LogicalPlan plan) { + return plan.transformUp(new AbstractPartialFunction() { + @Override + public boolean isDefinedAt(LogicalPlan node) { + return node instanceof Join; + } + + @Override + public LogicalPlan apply(LogicalPlan node) { + Join join = (Join) node; + if (join.condition().isEmpty() || joinsByEquality(join)) { + return join; + } + LogicalPlan left = materialize(join.left()); + LogicalPlan right = materialize(join.right()); + return left == join.left() && right == join.right() ? join + : new Join(left, right, join.joinType(), join.condition(), join.hint()); + } + }); + } + + /** The side behind a boundary, or the side as it is */ + private LogicalPlan materialize(LogicalPlan side) { + if (side instanceof Repartition || !callsUdf(side)) { + return side; + } + return new Repartition(partitions, true, side); + } + + /** Whether an equality relates the two sides, which Spark executes by hash */ + private static boolean joinsByEquality(Join join) { + List conjuncts = new ArrayList<>(); + split(join.condition().get(), conjuncts); + for (Expression conjunct : conjuncts) { + if (!(conjunct instanceof EqualTo)) { + continue; + } + EqualTo equality = (EqualTo) conjunct; + boolean leftThenRight = equality.left().references().subsetOf(join.left().outputSet()) + && equality.right().references().subsetOf(join.right().outputSet()); + boolean rightThenLeft = equality.left().references().subsetOf(join.right().outputSet()) + && equality.right().references().subsetOf(join.left().outputSet()); + if (leftThenRight || rightThenLeft) { + return true; + } + } + return false; + } + + /** The conjuncts of an `AND` tree, left to right */ + private static void split(Expression condition, List into) { + if (condition instanceof And) { + And and = (And) condition; + split(and.left(), into); + split(and.right(), into); + } else { + into.add(condition); + } + } + + /** Whether any expression of the plan, or of a plan under it, calls a user-defined function */ + private static boolean callsUdf(LogicalPlan plan) { + scala.collection.Iterator expressions = plan.expressions().iterator(); + while (expressions.hasNext()) { + if (OrderConjunctsByCost.callsUdf(expressions.next())) { + return true; + } + } + scala.collection.Iterator children = plan.children().iterator(); + while (children.hasNext()) { + if (callsUdf(children.next())) { + return true; + } + } + return false; + } +} diff --git a/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java b/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java index ad3ac77..14fbb4e 100644 --- a/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java +++ b/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java @@ -62,6 +62,13 @@ public Rule apply(SparkSession session) { return new OrderConjunctsByCost(); } }); + extensions.injectOptimizerRule(new AbstractFunction1>() { + @Override + public Rule apply(SparkSession session) { + return new MaterializeNestedLoopSides( + session.sparkContext().defaultParallelism()); + } + }); return BoxedUnit.UNIT; } } diff --git a/src/main/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCost.java b/src/main/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCost.java index b62cc93..05e178f 100644 --- a/src/main/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCost.java +++ b/src/main/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCost.java @@ -124,7 +124,7 @@ private static void split(Expression condition, List into) { } /** Whether the expression calls a user-defined function anywhere inside it */ - private static boolean callsUdf(Expression expression) { + static boolean callsUdf(Expression expression) { if (expression instanceof ScalaUDF) { return true; } diff --git a/src/test/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSidesControlTest.java b/src/test/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSidesControlTest.java new file mode 100644 index 0000000..ec3889a --- /dev/null +++ b/src/test/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSidesControlTest.java @@ -0,0 +1,84 @@ +/***************************************************************************** + * + * This MobilityDB code is provided under The PostgreSQL License. + * Copyright (c) 2020-2026, Université libre de Bruxelles and MobilityDB + * contributors + * + * Permission to use, copy, modify, and distribute this software and its + * documentation for any purpose, without fee, and without a written + * agreement is hereby retained provided that the above copyright notice and + * this paragraph and the following two paragraphs appear in all copies. + * + * IN NO EVENT SHALL UNIVERSITE LIBRE DE BRUXELLES BE LIABLE TO ANY PARTY FOR + * DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES + * INCLUDING LOST PROFITS, ARISING OUT OF THE USE OF THIS SOFTWARE AND ITS + * DOCUMENTATION, EVEN IF UNIVERSITE LIBRE DE BRUXELLES HAS BEEN ADVISED OF + * THE POSSIBILITY OF SUCH DAMAGE. + * + * UNIVERSITE LIBRE DE BRUXELLES SPECIFICALLY DISCLAIMS ANY WARRANTIES, + * INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY + * AND FITNESS FOR A PARTICULAR PURPOSE. THE SOFTWARE PROVIDED HEREUNDER IS + * ON AN "AS IS" BASIS, AND UNIVERSITE LIBRE DE BRUXELLES HAS NO OBLIGATIONS + * TO PROVIDE MAINTENANCE, SUPPORT, UPDATES, ENHANCEMENTS, OR MODIFICATIONS. + * + *****************************************************************************/ + +package org.mobilitydb.spark.catalyst; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; + +import org.apache.spark.sql.Dataset; +import org.apache.spark.sql.Row; +import org.apache.spark.sql.SparkSession; +import org.apache.spark.sql.api.java.UDF2; +import org.apache.spark.sql.types.DataTypes; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; + +/** + * The control for {@link MaterializeNestedLoopSidesTest}: the same query in a session WITHOUT the + * extension carries no boundary under the join, which is what makes the other class's assertion a + * statement about the rule rather than about Spark's own optimizer. A session carries its + * extensions from the builder and a JVM holds one Spark context, so the control needs its own + * class: surefire forks per class and does not reuse a fork. + */ +public class MaterializeNestedLoopSidesControlTest { + + private static SparkSession spark; + + @BeforeAll + static void session() { + spark = SparkSession.builder().appName("materialize-sides-control").master("local[1]") + .config("spark.ui.enabled", "false") + .getOrCreate(); + spark.sparkContext().setLogLevel("WARN"); + spark.udf().register("near", (UDF2) (a, b) -> Math.abs(a - b) <= 1, + DataTypes.BooleanType); + spark.range(0, 8).createOrReplaceTempView("t"); + } + + @AfterAll + static void stop() { + if (spark != null) { + spark.stop(); + } + } + + @Test + void withoutTheRuleNoBoundaryAppears() { + // The same query the rule's own test runs, its condition reading each side's computed + // column so column pruning keeps the call on the side + Dataset query = spark.sql( + "SELECT count(*) AS n FROM (SELECT id, near(id, id) AS flag FROM t) a " + + "JOIN (SELECT id, near(id, id) AS flag FROM t) b " + + "ON a.id < b.id AND a.flag AND b.flag AND near(a.id, b.id)"); + String plan = query.queryExecution().optimizedPlan().toString(); + assertFalse(plan.contains("Repartition"), + "Spark puts a boundary there by itself, so the other test proves nothing: " + plan); + // b - a <= 1 and a < b hold for b = a + 1 alone, so seven pairs of the eight ids + assertEquals(7L, query.collectAsList().get(0).getLong(0), + "the query answers seven pairs either way"); + } +} diff --git a/src/test/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSidesTest.java b/src/test/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSidesTest.java new file mode 100644 index 0000000..3fe431d --- /dev/null +++ b/src/test/java/org/mobilitydb/spark/catalyst/MaterializeNestedLoopSidesTest.java @@ -0,0 +1,99 @@ +/***************************************************************************** + * + * This MobilityDB code is provided under The PostgreSQL License. + * Copyright (c) 2020-2026, Université libre de Bruxelles and MobilityDB + * contributors + * + * Permission to use, copy, modify, and distribute this software and its + * documentation for any purpose, without fee, and without a written + * agreement is hereby retained provided that the above copyright notice and + * this paragraph and the following two paragraphs appear in all copies. + * + * IN NO EVENT SHALL UNIVERSITE LIBRE DE BRUXELLES BE LIABLE TO ANY PARTY FOR + * DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES + * INCLUDING LOST PROFITS, ARISING OUT OF THE USE OF THIS SOFTWARE AND ITS + * DOCUMENTATION, EVEN IF UNIVERSITE LIBRE DE BRUXELLES HAS BEEN ADVISED OF + * THE POSSIBILITY OF SUCH DAMAGE. + * + * UNIVERSITE LIBRE DE BRUXELLES SPECIFICALLY DISCLAIMS ANY WARRANTIES, + * INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY + * AND FITNESS FOR A PARTICULAR PURPOSE. THE SOFTWARE PROVIDED HEREUNDER IS + * ON AN "AS IS" BASIS, AND UNIVERSITE LIBRE DE BRUXELLES HAS NO OBLIGATIONS + * TO PROVIDE MAINTENANCE, SUPPORT, UPDATES, ENHANCEMENTS, OR MODIFICATIONS. + * + *****************************************************************************/ + +package org.mobilitydb.spark.catalyst; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import org.apache.spark.sql.Dataset; +import org.apache.spark.sql.Row; +import org.apache.spark.sql.SparkSession; +import org.apache.spark.sql.api.java.UDF2; +import org.apache.spark.sql.types.DataTypes; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; + +/** + * The rule in a session: a join with no equality between its sides, whose condition calls a + * user-defined function, puts each side behind a boundary, and a join that has such an equality + * is left alone. Both queries answer what they answer without the rule. + */ +public class MaterializeNestedLoopSidesTest { + + private static SparkSession spark; + + @BeforeAll + static void session() { + spark = SparkSession.builder().appName("materialize-sides").master("local[1]") + .config("spark.ui.enabled", "false") + .config("spark.sql.extensions", MobilitySparkExtensions.class.getName()) + .getOrCreate(); + spark.sparkContext().setLogLevel("WARN"); + spark.udf().register("near", (UDF2) (a, b) -> Math.abs(a - b) <= 1, + DataTypes.BooleanType); + spark.range(0, 8).createOrReplaceTempView("t"); + spark.sql("SELECT id, id AS k FROM t").createOrReplaceTempView("u"); + } + + @AfterAll + static void stop() { + if (spark != null) { + spark.stop(); + } + } + + private static String plan(Dataset query) { + return query.queryExecution().optimizedPlan().toString(); + } + + @Test + void aNestedLoopSideCallingTheFunctionGoesBehindABoundary() { + // The condition reads each side's computed column, so column pruning keeps the call on + // the side: an unused projection is pruned and the side optimizes to its bare source + Dataset query = spark.sql( + "SELECT count(*) AS n FROM (SELECT id, near(id, id) AS flag FROM t) a " + + "JOIN (SELECT id, near(id, id) AS flag FROM t) b " + + "ON a.id < b.id AND a.flag AND b.flag AND near(a.id, b.id)"); + assertTrue(plan(query).contains("Repartition"), + "no boundary under the nested-loop join: " + plan(query)); + // b - a <= 1 and a < b hold for b = a + 1 alone, so seven pairs of the eight ids + assertEquals(7L, query.collectAsList().get(0).getLong(0), + "the boundary answers what the query answered"); + } + + @Test + void aJoinByEqualityIsLeftAlone() { + Dataset query = spark.sql( + "SELECT count(*) AS n FROM (SELECT id, near(id, id) AS flag FROM t) a " + + "JOIN (SELECT id AS id2, k FROM u) b ON a.id = b.k"); + assertFalse(plan(query).contains("Repartition"), + "a join by equality reads each side once already: " + plan(query)); + assertEquals(8L, query.collectAsList().get(0).getLong(0), + "the query answers one row per id"); + } +}