From cdf51872f16095b3a489deb4831ce3be1be98cf2 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Thu, 3 Sep 2026 19:12:17 +0000 Subject: [PATCH 1/3] fix(core)!: preserve custom comparison function identity Custom hash and merge join comparisons keep a raw function anchor while Plan conversion regenerates all declarations. A comparison referencing equal:any_any at anchor 1 can therefore resolve to not_equal:any_any after a round trip; a comparison-only function loses its declaration entirely. Store the resolved scalar function declaration in CustomComparison, resolve it against the input plan's lookup, and register it with the output collector. This preserves identity across anchor reassignment and works with custom extension collections. BREAKING CHANGE: CustomComparison.of(int), getCustomFunctionReference(), and the generated customFunctionReference(int) builder method are replaced by of(ScalarFunctionVariant), getDeclaration(), and declaration(ScalarFunctionVariant). Supply the comparator declaration from your extension collection instead of a plan-local integer anchor. --- .../substrait/relation/ProtoRelConverter.java | 3 +- .../substrait/relation/RelProtoConverter.java | 4 +- .../relation/physical/ComparisonJoinKey.java | 24 +- .../CustomComparisonPlanRoundtripTest.java | 213 ++++++++++++++++++ .../type/proto/HashMergeJoinKeysTest.java | 8 +- 5 files changed, 236 insertions(+), 16 deletions(-) create mode 100644 core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java 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..370af6efb 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,7 @@ package io.substrait.relation.physical; import io.substrait.expression.FieldReference; +import io.substrait.extension.SimpleExtension; import org.immutables.value.Value; /** @@ -127,29 +128,26 @@ 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 scalar function with a boolean return type. */ @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(); /** - * 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/type/proto/CustomComparisonPlanRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java new file mode 100644 index 000000000..2e03620f7 --- /dev/null +++ b/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java @@ -0,0 +1,213 @@ +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 io.substrait.proto.Version; +import java.util.Arrays; +import java.util.List; +import java.util.stream.Stream; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; +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")); + + private static Stream joinCases() { + return Stream.of(false, true) + .flatMap( + merge -> + // Both zero and the unsigned value represented by -1 are valid wire anchors. + Stream.of(0, 1, 42, -1) + .flatMap( + anchor -> + Stream.of(false, true) + .map(withFilter -> Arguments.of(merge, anchor, withFilter)))); + } + + @ParameterizedTest + @MethodSource("joinCases") + 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(reference, keys.get(1).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() + .setVersion(Version.newBuilder().setMinorNumber(102)) + .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 = From c822af34dfb7f03d6719e856cd32357e719f246b Mon Sep 17 00:00:00 2001 From: bvolpato Date: Sat, 3 Oct 2026 17:54:51 -0400 Subject: [PATCH 2/3] fix(core): validate custom comparison signatures Check comparator arity and boolean return types while retaining valid parameterized and variadic declarations. Clarify scalar resolution and simplify comparison round-trip cases. --- .../relation/physical/ComparisonJoinKey.java | 33 +++++++++- .../relation/ComparisonJoinKeyTest.java | 60 +++++++++++++++++++ .../CustomComparisonPlanRoundtripTest.java | 21 +------ 3 files changed, 94 insertions(+), 20 deletions(-) create mode 100644 core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java 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 370af6efb..b82b068eb 100644 --- a/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java +++ b/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java @@ -2,6 +2,8 @@ import io.substrait.expression.FieldReference; import io.substrait.extension.SimpleExtension; +import io.substrait.type.Type; +import java.util.Arrays; import org.immutables.value.Value; /** @@ -36,6 +38,20 @@ 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) + && !(declaration.resolveType(Arrays.asList(getLeft().getType(), getRight().getType())) + instanceof Type.Bool)) { + throw new IllegalArgumentException("Custom comparison function must return boolean"); + } + } + } + /** * Creates a builder for {@link ComparisonJoinKey}. * @@ -128,7 +144,10 @@ public R accept(ComparisonTypeVisitor visitor) th } } - /** A custom comparison behavior, given by a binary scalar 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 { /** @@ -140,6 +159,18 @@ public abstract static class CustomComparison implements ComparisonType { */ 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} using the given comparison function declaration. * 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..836b5fa98 --- /dev/null +++ b/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java @@ -0,0 +1,60 @@ +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)); + } + + private SimpleExtension.ScalarFunctionVariant scalar(String urn, String key) { + return extensions.getScalarFunction(SimpleExtension.FunctionAnchor.of(urn, key)); + } + + private ComparisonJoinKey joinKey(CustomComparison comparison, Type type) { + return ComparisonJoinKey.builder() + .left(FieldReference.newRootStructReference(0, type)) + .right(FieldReference.newRootStructReference(0, type)) + .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 index 2e03620f7..8be5177b3 100644 --- a/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java +++ b/core/src/test/java/io/substrait/type/proto/CustomComparisonPlanRoundtripTest.java @@ -23,13 +23,10 @@ import io.substrait.proto.RelRoot; import io.substrait.proto.SimpleExtensionDeclaration; import io.substrait.proto.SimpleExtensionURN; -import io.substrait.proto.Version; import java.util.Arrays; import java.util.List; -import java.util.stream.Stream; import org.junit.jupiter.params.ParameterizedTest; -import org.junit.jupiter.params.provider.Arguments; -import org.junit.jupiter.params.provider.MethodSource; +import org.junit.jupiter.params.provider.CsvSource; import org.junit.jupiter.params.provider.ValueSource; class CustomComparisonPlanRoundtripTest extends TestBase { @@ -41,20 +38,8 @@ class CustomComparisonPlanRoundtripTest extends TestBase { SimpleExtension.FunctionAnchor.of( DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "equal:any_any")); - private static Stream joinCases() { - return Stream.of(false, true) - .flatMap( - merge -> - // Both zero and the unsigned value represented by -1 are valid wire anchors. - Stream.of(0, 1, 42, -1) - .flatMap( - anchor -> - Stream.of(false, true) - .map(withFilter -> Arguments.of(merge, anchor, withFilter)))); - } - @ParameterizedTest - @MethodSource("joinCases") + @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); @@ -114,7 +99,6 @@ private void assertComparisonDeclarations( merge ? rel.getMergeJoin().getKeysList() : rel.getHashJoin().getKeysList(); assertEquals(2, keys.size()); int reference = keys.get(0).getComparison().getCustomFunctionReference(); - assertEquals(reference, keys.get(1).getComparison().getCustomFunctionReference()); assertEquals( expected, ImmutableExtensionLookup.builder() @@ -176,7 +160,6 @@ private Plan plan( } Plan.Builder plan = Plan.newBuilder() - .setVersion(Version.newBuilder().setMinorNumber(102)) .setExecutionBehavior( ExecutionBehavior.newBuilder() .setVariableEvalMode( From 6b506dcfa5efa27a132cd073a22ab340d2532db6 Mon Sep 17 00:00:00 2001 From: bvolpato Date: Wed, 7 Oct 2026 02:01:40 -0400 Subject: [PATCH 3/3] fix(core): report unbindable comparison keys as invalid arguments --- .../relation/physical/ComparisonJoinKey.java | 15 +++++++++++--- .../relation/ComparisonJoinKeyTest.java | 20 +++++++++++++++++-- 2 files changed, 30 insertions(+), 5 deletions(-) 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 b82b068eb..779dac02b 100644 --- a/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java +++ b/core/src/main/java/io/substrait/relation/physical/ComparisonJoinKey.java @@ -44,9 +44,18 @@ protected void checkCustomComparisonReturnType() { if (getComparison() instanceof CustomComparison) { SimpleExtension.ScalarFunctionVariant declaration = ((CustomComparison) getComparison()).getDeclaration(); - if (!(declaration.returnType() instanceof Type) - && !(declaration.resolveType(Arrays.asList(getLeft().getType(), getRight().getType())) - instanceof Type.Bool)) { + 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"); } } diff --git a/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java b/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java index 836b5fa98..01a58a372 100644 --- a/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java +++ b/core/src/test/java/io/substrait/relation/ComparisonJoinKeyTest.java @@ -46,14 +46,30 @@ void rejectsNonBooleanReturnResolvedFromKeyTypes(String 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, type)) - .right(FieldReference.newRootStructReference(0, type)) + .left(FieldReference.newRootStructReference(0, leftType)) + .right(FieldReference.newRootStructReference(0, rightType)) .comparison(comparison) .build(); }