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,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<LogicalPlan> {

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<LogicalPlan, LogicalPlan>() {
@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<Expression> 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<Expression> 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<Expression> expressions = plan.expressions().iterator();
while (expressions.hasNext()) {
if (OrderConjunctsByCost.callsUdf(expressions.next())) {
return true;
}
}
scala.collection.Iterator<LogicalPlan> children = plan.children().iterator();
while (children.hasNext()) {
if (callsUdf(children.next())) {
return true;
}
}
return false;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,13 @@ public Rule<LogicalPlan> apply(SparkSession session) {
return new OrderConjunctsByCost();
}
});
extensions.injectOptimizerRule(new AbstractFunction1<SparkSession, Rule<LogicalPlan>>() {
@Override
public Rule<LogicalPlan> apply(SparkSession session) {
return new MaterializeNestedLoopSides(
session.sparkContext().defaultParallelism());
}
});
return BoxedUnit.UNIT;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -124,7 +124,7 @@ private static void split(Expression condition, List<Expression> 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;
}
Expand Down
Original file line number Diff line number Diff line change
@@ -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<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();
}
}

@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<Row> 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");
}
}
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.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<Long, Long, Boolean>) (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<Row> 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<Row> 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<Row> 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");
}
}
Loading