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
Original file line number Diff line number Diff line change
@@ -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`:
*
* <pre>
* SparkSession.builder().config("spark.sql.extensions",
* "org.mobilitydb.spark.catalyst.MobilitySparkExtensions")
* </pre>
*
* 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<SparkSessionExtensions, BoxedUnit> {

@Override
public BoxedUnit apply(SparkSessionExtensions extensions) {
extensions.injectOptimizerRule(new AbstractFunction1<SparkSession, Rule<LogicalPlan>>() {
@Override
public Rule<LogicalPlan> apply(SparkSession session) {
return new OrderConjunctsByCost();
}
});
return BoxedUnit.UNIT;
}
}
139 changes: 139 additions & 0 deletions src/main/java/org/mobilitydb/spark/catalyst/OrderConjunctsByCost.java
Original file line number Diff line number Diff line change
@@ -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<LogicalPlan> {

@Override
public LogicalPlan apply(LogicalPlan plan) {
return plan.transformUp(new AbstractPartialFunction<LogicalPlan, LogicalPlan>() {
@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<Expression> conjuncts = new ArrayList<>();
split(condition, conjuncts);
if (conjuncts.size() < 2) {
return condition;
}
List<Expression> plain = new ArrayList<>();
List<Expression> 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<Expression> 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<Expression> children = expression.children().iterator();
while (children.hasNext()) {
if (callsUdf(children.next())) {
return true;
}
}
return false;
}
}
Original file line number Diff line number Diff line change
@@ -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<Long, Long, Boolean>) (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<Row> 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<Row> 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");
}
}
Original file line number Diff line number Diff line change
@@ -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<Long, Long, Boolean>) (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<Row> 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<Row> 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<Row> 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");
}
}
Loading