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
Expand Up @@ -1214,7 +1214,8 @@ private ComparisonJoinKey comparisonJoinKey(
break;
case CUSTOM_FUNCTION_REFERENCE:
comparisonType =
ComparisonJoinKey.CustomComparison.of(comparison.getCustomFunctionReference());
ComparisonJoinKey.CustomComparison.of(
lookup.getScalarFunction(comparison.getCustomFunctionReference(), extensions));
break;
default:
throw new IllegalArgumentException(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -221,7 +221,9 @@ public io.substrait.proto.ComparisonJoinKey.ComparisonType visit(
public io.substrait.proto.ComparisonJoinKey.ComparisonType visit(
ComparisonJoinKey.CustomComparison customComparison) {
return io.substrait.proto.ComparisonJoinKey.ComparisonType.newBuilder()
.setCustomFunctionReference(customComparison.getCustomFunctionReference())
.setCustomFunctionReference(
extensionCollector.getFunctionReference(
customComparison.getDeclaration()))
.build();
}
});
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
package io.substrait.relation.physical;

import io.substrait.expression.FieldReference;
import io.substrait.extension.SimpleExtension;
import io.substrait.type.Type;
import java.util.Arrays;
import org.immutables.value.Value;

/**
Expand Down Expand Up @@ -35,6 +38,29 @@ public abstract class ComparisonJoinKey {
*/
public abstract ComparisonType getComparison();

/** Validates parameterized comparison returns, which require the actual key types to resolve. */
@Value.Check
protected void checkCustomComparisonReturnType() {
if (getComparison() instanceof CustomComparison) {
SimpleExtension.ScalarFunctionVariant declaration =
((CustomComparison) getComparison()).getDeclaration();
if (declaration.returnType() instanceof Type) {
return;
}
Type resolvedReturnType;
try {
resolvedReturnType =
declaration.resolveType(Arrays.asList(getLeft().getType(), getRight().getType()));
} catch (UnsupportedOperationException e) {
throw new IllegalArgumentException(
"Custom comparison function cannot be resolved with the join key types", e);
}
if (!(resolvedReturnType instanceof Type.Bool)) {
throw new IllegalArgumentException("Custom comparison function must return boolean");
}
}
}

/**
* Creates a builder for {@link ComparisonJoinKey}.
*
Expand Down Expand Up @@ -128,28 +154,40 @@ public <R, E extends Throwable> R accept(ComparisonTypeVisitor<R, E> visitor) th
}

/**
* A custom comparison behavior, given by a reference to a binary function with a boolean return
* type.
* A custom comparison behavior, given by a binary function with a boolean return type.
* Substrait-java resolves this function as a scalar function.
*/
@Value.Immutable
public abstract static class CustomComparison implements ComparisonType {
/**
* Returns the reference to the binary boolean-returning comparison function.
* Returns the {@link io.substrait.extension.SimpleExtension.ScalarFunctionVariant} declaring
* the binary boolean-returning comparison function. Its plan-local reference is assigned during
* protobuf conversion.
*
* @return the custom function reference
* @return the comparison function declaration
*/
public abstract int getCustomFunctionReference();
public abstract SimpleExtension.ScalarFunctionVariant getDeclaration();

/** Validates the comparator's arity and any concrete return type. */
@Value.Check
protected void checkDeclaration() {
if (!getDeclaration().getRange().within(2)) {
throw new IllegalArgumentException("Custom comparison function must accept two arguments");
}
if (getDeclaration().returnType() instanceof Type
&& !(getDeclaration().returnType() instanceof Type.Bool)) {
throw new IllegalArgumentException("Custom comparison function must return boolean");
}
}

/**
* Creates a {@link CustomComparison} referencing the given comparison function.
* Creates a {@link CustomComparison} using the given comparison function declaration.
*
* @param customFunctionReference the reference to the comparison function
* @param declaration the binary boolean-returning comparison function declaration
* @return a new custom comparison
*/
public static CustomComparison of(int customFunctionReference) {
return ImmutableComparisonJoinKey.CustomComparison.builder()
.customFunctionReference(customFunctionReference)
.build();
public static CustomComparison of(SimpleExtension.ScalarFunctionVariant declaration) {
Comment thread
nielspardon marked this conversation as resolved.
return ImmutableComparisonJoinKey.CustomComparison.builder().declaration(declaration).build();
}

@Override
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
package io.substrait.relation;

import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
import static org.junit.jupiter.api.Assertions.assertThrows;

import io.substrait.TestBase;
import io.substrait.expression.FieldReference;
import io.substrait.extension.DefaultExtensionCatalog;
import io.substrait.extension.SimpleExtension;
import io.substrait.relation.physical.ComparisonJoinKey;
import io.substrait.relation.physical.ComparisonJoinKey.CustomComparison;
import io.substrait.type.Type;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.ValueSource;

class ComparisonJoinKeyTest extends TestBase {

@Test
void rejectsNonBinaryComparator() {
SimpleExtension.ScalarFunctionVariant not =
scalar(DefaultExtensionCatalog.FUNCTIONS_BOOLEAN, "not:bool");
assertThrows(IllegalArgumentException.class, () -> CustomComparison.of(not));
}

@Test
void rejectsNonBooleanComparator() {
SimpleExtension.ScalarFunctionVariant add =
scalar(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "add:i32_i32");
assertThrows(IllegalArgumentException.class, () -> CustomComparison.of(add));
}

@ParameterizedTest
@ValueSource(strings = {"nullif:any_any", "coalesce:any"})
void acceptsBooleanReturnResolvedFromKeyTypes(String key) {
CustomComparison comparison =
CustomComparison.of(scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, key));
assertDoesNotThrow(() -> joinKey(comparison, N.BOOLEAN));
}

@ParameterizedTest
@ValueSource(strings = {"nullif:any_any", "coalesce:any"})
void rejectsNonBooleanReturnResolvedFromKeyTypes(String key) {
CustomComparison comparison =
CustomComparison.of(scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, key));
assertThrows(IllegalArgumentException.class, () -> joinKey(comparison, R.I32));
}

@Test
void rejectsKeyTypesThatCannotBindToComparatorParameters() {
CustomComparison decimalAdd =
CustomComparison.of(
scalar(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC_DECIMAL, "add:dec_dec"));
assertThrows(IllegalArgumentException.class, () -> joinKey(decimalAdd, R.I32));

CustomComparison nullIf =
CustomComparison.of(scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any"));
assertThrows(IllegalArgumentException.class, () -> joinKey(nullIf, N.BOOLEAN, R.I32));
}

private SimpleExtension.ScalarFunctionVariant scalar(String urn, String key) {
return extensions.getScalarFunction(SimpleExtension.FunctionAnchor.of(urn, key));
}

private ComparisonJoinKey joinKey(CustomComparison comparison, Type type) {
return joinKey(comparison, type, type);
}

private ComparisonJoinKey joinKey(CustomComparison comparison, Type leftType, Type rightType) {
return ComparisonJoinKey.builder()
.left(FieldReference.newRootStructReference(0, leftType))
.right(FieldReference.newRootStructReference(0, rightType))
.comparison(comparison)
.build();
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,196 @@
package io.substrait.type.proto;

import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertThrows;

import io.substrait.TestBase;
import io.substrait.expression.FieldReference;
import io.substrait.extension.DefaultExtensionCatalog;
import io.substrait.extension.ImmutableExtensionLookup;
import io.substrait.extension.ImmutableSimpleExtension;
import io.substrait.extension.SimpleExtension;
import io.substrait.plan.PlanProtoConverter;
import io.substrait.plan.ProtoPlanConverter;
import io.substrait.proto.ComparisonJoinKey;
import io.substrait.proto.ExecutionBehavior;
import io.substrait.proto.Expression;
import io.substrait.proto.FunctionArgument;
import io.substrait.proto.HashJoinRel;
import io.substrait.proto.MergeJoinRel;
import io.substrait.proto.Plan;
import io.substrait.proto.PlanRel;
import io.substrait.proto.Rel;
import io.substrait.proto.RelRoot;
import io.substrait.proto.SimpleExtensionDeclaration;
import io.substrait.proto.SimpleExtensionURN;
import java.util.Arrays;
import java.util.List;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.CsvSource;
import org.junit.jupiter.params.provider.ValueSource;

class CustomComparisonPlanRoundtripTest extends TestBase {

private static final int POST_FILTER_REFERENCE = 99;

private final SimpleExtension.ScalarFunctionVariant equal =
extensions.getScalarFunction(
SimpleExtension.FunctionAnchor.of(
DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "equal:any_any"));

@ParameterizedTest
@CsvSource({"false, 1, true", "true, 1, true", "false, 0, false", "true, -1, false"})
void preservesComparisonIdentity(boolean merge, int anchor, boolean withFilter) {
Plan original = plan(merge, anchor, equal, withFilter);
io.substrait.plan.Plan pojo = new ProtoPlanConverter().from(original);
Plan converted = new PlanProtoConverter().toProto(pojo);

assertComparisonDeclarations(converted, merge, equal, extensions);
assertEquals(withFilter ? 2 : 1, converted.getExtensionsCount());
if (withFilter) {
Rel rel = converted.getRelations(0).getRoot().getInput();
Expression filter =
merge ? rel.getMergeJoin().getPostJoinFilter() : rel.getHashJoin().getPostJoinFilter();
assertEquals(
"not_equal:any_any",
ImmutableExtensionLookup.builder()
.from(converted)
.build()
.getScalarFunction(filter.getScalarFunction().getFunctionReference(), extensions)
.key());
}
assertEquals(pojo, new ProtoPlanConverter().from(converted));
}

@ParameterizedTest
@ValueSource(booleans = {false, true})
void usesConfiguredExtensionCollection(boolean merge) {
SimpleExtension.ScalarFunctionVariant custom =
ImmutableSimpleExtension.ScalarFunctionVariant.builder()
.from(equal)
.urn("extension:example:comparisons")
.name("matches")
.build();
SimpleExtension.ExtensionCollection collection =
SimpleExtension.ExtensionCollection.builder().addScalarFunctions(custom).build();
Plan original = plan(merge, 42, custom, false);
io.substrait.plan.Plan pojo = new ProtoPlanConverter(collection).from(original);
Plan converted = new PlanProtoConverter(collection).toProto(pojo);

assertComparisonDeclarations(converted, merge, custom, collection);
assertEquals(1, converted.getExtensionsCount());
assertEquals(pojo, new ProtoPlanConverter(collection).from(converted));
}

@ParameterizedTest
@ValueSource(booleans = {false, true})
void rejectsUndeclaredComparisonReference(boolean merge) {
Plan original = plan(merge, 42, equal, false).toBuilder().clearExtensions().build();
assertThrows(IllegalArgumentException.class, () -> new ProtoPlanConverter().from(original));
}

private void assertComparisonDeclarations(
Plan plan,
boolean merge,
SimpleExtension.ScalarFunctionVariant expected,
SimpleExtension.ExtensionCollection collection) {
Rel rel = plan.getRelations(0).getRoot().getInput();
List<ComparisonJoinKey> keys =
merge ? rel.getMergeJoin().getKeysList() : rel.getHashJoin().getKeysList();
assertEquals(2, keys.size());
int reference = keys.get(0).getComparison().getCustomFunctionReference();
assertEquals(
expected,
ImmutableExtensionLookup.builder()
.from(plan)
.build()
.getScalarFunction(reference, collection));
}

private Plan plan(
boolean merge,
int comparisonAnchor,
SimpleExtension.ScalarFunctionVariant comparison,
boolean withFilter) {
Rel input =
relProtoConverter.toProto(
sb.namedScan(Arrays.asList("t"), Arrays.asList("x"), Arrays.asList(R.I32)));
ComparisonJoinKey key =
ComparisonJoinKey.newBuilder()
.setLeft(field(0).getSelection())
.setRight(field(0).getSelection())
.setComparison(
ComparisonJoinKey.ComparisonType.newBuilder()
.setCustomFunctionReference(comparisonAnchor))
.build();
Expression postFilter =
Expression.newBuilder()
.setScalarFunction(
Expression.ScalarFunction.newBuilder()
.setFunctionReference(POST_FILTER_REFERENCE)
.setOutputType(relProtoConverter.getTypeProtoConverter().toProto(R.BOOLEAN))
.addArguments(FunctionArgument.newBuilder().setValue(field(0)))
.addArguments(FunctionArgument.newBuilder().setValue(field(1))))
.build();
Rel.Builder relation = Rel.newBuilder();
if (merge) {
MergeJoinRel.Builder join =
MergeJoinRel.newBuilder()
.setLeft(input)
.setRight(input)
.setType(MergeJoinRel.JoinType.JOIN_TYPE_INNER)
.addKeys(key)
.addKeys(key);
if (withFilter) {
join.setPostJoinFilter(postFilter);
}
relation.setMergeJoin(join);
} else {
HashJoinRel.Builder join =
HashJoinRel.newBuilder()
.setLeft(input)
.setRight(input)
.setType(HashJoinRel.JoinType.JOIN_TYPE_INNER)
.addKeys(key)
.addKeys(key);
if (withFilter) {
join.setPostJoinFilter(postFilter);
}
relation.setHashJoin(join);
}
Plan.Builder plan =
Plan.newBuilder()
.setExecutionBehavior(
ExecutionBehavior.newBuilder()
.setVariableEvalMode(
ExecutionBehavior.VariableEvaluationMode.VARIABLE_EVALUATION_MODE_PER_PLAN))
.addExtensionUrns(
SimpleExtensionURN.newBuilder().setExtensionUrnAnchor(1).setUrn(comparison.urn()))
.addExtensions(function(comparisonAnchor, comparison.key()))
.addRelations(
PlanRel.newBuilder()
.setRoot(
RelRoot.newBuilder()
.setInput(relation)
.addNames("left_x")
.addNames("right_x")));
if (withFilter) {
plan.addExtensions(function(POST_FILTER_REFERENCE, "not_equal:any_any"));
}
return plan.build();
}

private Expression field(int index) {
return expressionProtoConverter.toProto(FieldReference.newRootStructReference(index, R.I32));
}

private static SimpleExtensionDeclaration function(int reference, String name) {
return SimpleExtensionDeclaration.newBuilder()
.setExtensionFunction(
SimpleExtensionDeclaration.ExtensionFunction.newBuilder()
.setExtensionUrnReference(1)
.setFunctionAnchor(reference)
.setName(name))
.build();
}
}
Loading
Loading