diff --git a/core/src/main/java/io/substrait/relation/ProtoRelConverter.java b/core/src/main/java/io/substrait/relation/ProtoRelConverter.java index 5f43ace2b..40fc1e3bd 100644 --- a/core/src/main/java/io/substrait/relation/ProtoRelConverter.java +++ b/core/src/main/java/io/substrait/relation/ProtoRelConverter.java @@ -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( diff --git a/core/src/main/java/io/substrait/relation/RelProtoConverter.java b/core/src/main/java/io/substrait/relation/RelProtoConverter.java index 885cd87c9..e15907b9b 100644 --- a/core/src/main/java/io/substrait/relation/RelProtoConverter.java +++ b/core/src/main/java/io/substrait/relation/RelProtoConverter.java @@ -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(); } }); diff --git a/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java b/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java index 9d3b18e4a..779dac02b 100644 --- a/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java +++ b/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java @@ -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; /** @@ -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}. * @@ -128,28 +154,40 @@ public R accept(ComparisonTypeVisitor 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) { + return ImmutableComparisonJoinKey.CustomComparison.builder().declaration(declaration).build(); } @Override diff --git a/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java b/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java new file mode 100644 index 000000000..01a58a372 --- /dev/null +++ b/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java @@ -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(); + } +} diff --git a/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java new file mode 100644 index 000000000..8be5177b3 --- /dev/null +++ b/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java @@ -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 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(); + } +} diff --git a/core/src/test/java/io/substrait/type/proto/HashMergeJoinKeysTest.java b/core/src/test/java/io/substrait/type/proto/HashMergeJoinKeysTest.java index 2b8e3f572..61b3b57d6 100644 --- a/core/src/test/java/io/substrait/type/proto/HashMergeJoinKeysTest.java +++ b/core/src/test/java/io/substrait/type/proto/HashMergeJoinKeysTest.java @@ -3,6 +3,8 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import io.substrait.TestBase; +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.SimpleExtension; import io.substrait.relation.Rel; import io.substrait.relation.physical.ComparisonJoinKey; import io.substrait.relation.physical.ComparisonJoinKey.SimpleComparisonType; @@ -81,7 +83,11 @@ void fullFidelityRoundTrip() { ComparisonJoinKey.builder() .left(sb.fieldReference(leftTable, 2)) .right(sb.fieldReference(rightTable, 1)) - .comparison(ComparisonJoinKey.CustomComparison.of(42)) + .comparison( + ComparisonJoinKey.CustomComparison.of( + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "equal:any_any")))) .build()); Rel hash =