diff --git a/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java b/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java new file mode 100644 index 00000000..ad3ac775 --- /dev/null +++ b/src/main/java/org/mobilitydb/spark/catalyst/MobilitySparkExtensions.java @@ -0,0 +1,67 @@ +/***************************************************************************** + * + * 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 org.apache.spark.sql.SparkSession; +import org.apache.spark.sql.SparkSessionExtensions; +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan; +import org.apache.spark.sql.catalyst.rules.Rule; + +import scala.Function1; +import scala.runtime.AbstractFunction1; +import scala.runtime.BoxedUnit; + +/** + * The Catalyst rules MobilitySpark adds to a Spark session, named in `spark.sql.extensions`: + * + *
+ *   SparkSession.builder().config("spark.sql.extensions",
+ *       "org.mobilitydb.spark.catalyst.MobilitySparkExtensions")
+ * 
+ * + * Spark applies the named class to the session's extensions, as a Scala function of one argument, + * which this class provides through scala.runtime.AbstractFunction1. + * + * It injects {@link OrderConjunctsByCost}, which moves a MobilitySpark function behind the + * comparisons beside it in one conjunction. Spark holds no cost for a user-defined function, so + * its optimizer never sinks one: in a proximity join over trajectories the distance is evaluated + * on every pair the join enumerates, while the box and time comparisons that reject almost all of + * them run afterwards. + */ +public final class MobilitySparkExtensions + extends AbstractFunction1 { + + @Override + public BoxedUnit apply(SparkSessionExtensions extensions) { + extensions.injectOptimizerRule(new AbstractFunction1>() { + @Override + public Rule apply(SparkSession session) { + return new OrderConjunctsByCost(); + } + }); + 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 new file mode 100644 index 00000000..b62cc93f --- /dev/null +++ b/src/main/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCost.java @@ -0,0 +1,139 @@ +/***************************************************************************** + * + * 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.Expression; +import org.apache.spark.sql.catalyst.expressions.ScalaUDF; +import org.apache.spark.sql.catalyst.plans.logical.Filter; +import org.apache.spark.sql.catalyst.plans.logical.Join; +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan; +import org.apache.spark.sql.catalyst.rules.Rule; + +import scala.Option; +import scala.runtime.AbstractPartialFunction; + +/** + * Moves every conjunct that calls a user-defined function behind the conjuncts that do not, in a + * join condition and in a filter. + * + * `AND` evaluates its operands in order and stops at the first false one, so a conjunct placed + * later is evaluated on fewer rows: the reordering can only reduce the number of evaluations, and + * it never introduces one, which is why it preserves both the answer and the exceptions a + * user-defined function can raise. The relative order inside each of the two groups is kept, so a + * condition already in this shape is returned unchanged. + * + * Spark attaches no cost to a user-defined function, so its optimizer leaves the order the query + * produced. A proximity join over trajectories shows what that costs: the plan evaluates the + * distance between two trajectories on every pair the join enumerates, while the box and time + * comparisons in the same conjunction, which reject almost all of those pairs, run after it. + */ +public final class OrderConjunctsByCost extends Rule { + + @Override + public LogicalPlan apply(LogicalPlan plan) { + return plan.transformUp(new AbstractPartialFunction() { + @Override + public boolean isDefinedAt(LogicalPlan node) { + return node instanceof Join || node instanceof Filter; + } + + @Override + public LogicalPlan apply(LogicalPlan node) { + if (node instanceof Filter) { + Filter filter = (Filter) node; + Expression ordered = order(filter.condition()); + return ordered == filter.condition() ? filter + : new Filter(ordered, filter.child()); + } + Join join = (Join) node; + if (join.condition().isEmpty()) { + return join; + } + Expression condition = join.condition().get(); + Expression ordered = order(condition); + return ordered == condition ? join + : new Join(join.left(), join.right(), join.joinType(), + Option.apply(ordered), join.hint()); + } + }); + } + + /** The same conjunction with the calls to a user-defined function last, or it unchanged */ + private static Expression order(Expression condition) { + List conjuncts = new ArrayList<>(); + split(condition, conjuncts); + if (conjuncts.size() < 2) { + return condition; + } + List plain = new ArrayList<>(); + List calls = new ArrayList<>(); + for (Expression conjunct : conjuncts) { + // A conjunct that is not deterministic keeps its place: its position is observable + if (!conjunct.deterministic()) { + return condition; + } + (callsUdf(conjunct) ? calls : plain).add(conjunct); + } + if (calls.isEmpty() || plain.isEmpty()) { + return condition; + } + plain.addAll(calls); + Expression ordered = plain.get(0); + for (int i = 1; i < plain.size(); i++) { + ordered = new And(ordered, plain.get(i)); + } + return ordered; + } + + /** 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 the expression calls a user-defined function anywhere inside it */ + private static boolean callsUdf(Expression expression) { + if (expression instanceof ScalaUDF) { + return true; + } + scala.collection.Iterator children = expression.children().iterator(); + while (children.hasNext()) { + if (callsUdf(children.next())) { + return true; + } + } + return false; + } +} diff --git a/src/test/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCostControlTest.java b/src/test/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCostControlTest.java new file mode 100644 index 00000000..694c6898 --- /dev/null +++ b/src/test/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCostControlTest.java @@ -0,0 +1,91 @@ +/***************************************************************************** + * + * 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.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 control for {@link OrderConjunctsByCostTest}: the same query in a session WITHOUT the + * extension keeps the call to the user-defined function ahead of the comparison, 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 OrderConjunctsByCostControlTest { + + private static SparkSession spark; + + @BeforeAll + static void session() { + spark = SparkSession.builder().appName("order-conjuncts-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(); + } + } + + /** The condition of a join, as the optimized plan prints it */ + private static String joinCondition(Dataset query) { + for (String line : query.queryExecution().optimizedPlan().toString().split("\n")) { + if (line.contains("Join ")) { + return line; + } + } + throw new AssertionError("no join in the optimized plan"); + } + + @Test + void withoutTheRuleTheCallStaysFirst() { + Dataset query = spark.sql( + "SELECT count(*) AS n FROM t a JOIN t b ON near(a.id, b.id) AND a.id < b.id"); + String condition = joinCondition(query); + assertTrue(condition.indexOf("near(") < condition.indexOf(" < "), + "Spark orders the conjunction by itself, so the other test proves nothing: " + + condition); + // 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/OrderConjunctsByCostTest.java b/src/test/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCostTest.java new file mode 100644 index 00000000..778bb29c --- /dev/null +++ b/src/test/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCostTest.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.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 condition naming a user-defined function first is optimized into + * one naming it last, and the query answers what it answered before. + */ +public class OrderConjunctsByCostTest { + + private static SparkSession spark; + + @BeforeAll + static void session() { + spark = SparkSession.builder().appName("order-conjuncts").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"); + } + + @AfterAll + static void stop() { + if (spark != null) { + spark.stop(); + } + } + + /** The condition of a join, as the optimized plan prints it */ + private static String joinCondition(Dataset query) { + for (String line : query.queryExecution().optimizedPlan().toString().split("\n")) { + if (line.contains("Join ")) { + return line; + } + } + throw new AssertionError("no join in the optimized plan"); + } + + @Test + void theCallGoesBehindTheComparisons() { + Dataset query = spark.sql( + "SELECT count(*) AS n FROM t a JOIN t b ON near(a.id, b.id) AND a.id < b.id"); + String condition = joinCondition(query); + assertTrue(condition.indexOf("near(") > condition.indexOf(" < "), + "the call to near stays ahead of the comparison: " + condition); + // 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 reordering answers what the query answered"); + } + + @Test + void aConditionAlreadyOrderedIsLeftAlone() { + Dataset query = spark.sql( + "SELECT count(*) AS n FROM t a JOIN t b ON a.id < b.id AND near(a.id, b.id)"); + String condition = joinCondition(query); + assertTrue(condition.indexOf("near(") > condition.indexOf(" < "), + "the call to near moved ahead of the comparison: " + condition); + assertEquals(7L, query.collectAsList().get(0).getLong(0), + "the query answers what it answered"); + } +}