From d1b27d3ad212f86734e6dac437a742e7fcf9d2a6 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Mon, 7 Sep 2026 15:56:22 +0300 Subject: [PATCH 1/2] fix(isthmus): index extended expression fields against the combined schema Extended expressions share one `base_schema`, but references to fields from later tables used indices local to each table. After a three-column table, `B2` in a second table was emitted as field 1 instead of field 4. Build references in the same registration order as the combined schema, following spec v0.102.0. Closes #1200 --- .../isthmus/SqlExpressionToSubstrait.java | 4 +- .../SimpleExtendedExpressionsTest.java | 74 +++++++++++++++++++ 2 files changed, 77 insertions(+), 1 deletion(-) diff --git a/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java b/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java index 620d12927..35b48c30e 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java @@ -208,6 +208,8 @@ private Result registerCreateTablesForExtendedExpression(List tables) for (SubstraitTable t : tList) { rootSchema.add(t.getName(), t); for (RelDataTypeField field : t.getRowType(factory).getFieldList()) { + // Field references index the combined base schema in insertion order. + int fieldIndex = nameToTypeMap.size(); nameToTypeMap.merge( // to validate the sql expression tree field.getName(), field.getType(), @@ -217,7 +219,7 @@ private Result registerCreateTablesForExtendedExpression(List tables) }); nameToNodeMap.merge( // to convert sql expression into RexNode field.getName(), - new RexInputRef(field.getIndex(), field.getType()), + new RexInputRef(fieldIndex, field.getType()), (v1, v2) -> { throw new IllegalArgumentException( "There is no support for duplicate column names: " + field.getName()); diff --git a/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java b/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java index cc31ce2e9..f753a5243 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java @@ -6,9 +6,11 @@ import static org.junit.jupiter.api.Assertions.assertTrue; import io.substrait.isthmus.expression.RexExpressionConverter; +import io.substrait.proto.Expression; import io.substrait.proto.Expression.RexTypeCase; import io.substrait.proto.ExtendedExpression; import java.io.IOException; +import java.util.List; import java.util.stream.Stream; import org.apache.calcite.sql.parser.SqlParseException; import org.junit.jupiter.api.Test; @@ -20,6 +22,78 @@ class SimpleExtendedExpressionsTest extends ExtendedExpressionTestBase { private static final String MARKER = "provider hook reached"; + private static final String TABLE_A = "CREATE TABLE A (A1 BIGINT, A2 BIGINT, A3 BIGINT)"; + private static final String TABLE_B = "CREATE TABLE B (B1 BIGINT, B2 BIGINT)"; + private static final String TABLE_C = "CREATE TABLE C (C1 BIGINT)"; + + private static Stream columnSchemaProvider() { + return Stream.of( + Arguments.of(List.of(TABLE_A), List.of("A1", "A2", "A3")), + Arguments.of( + List.of(TABLE_A, TABLE_B, TABLE_C), List.of("A1", "A2", "A3", "B1", "B2", "C1")), + Arguments.of( + List.of(TABLE_A + ";" + TABLE_B + ";" + TABLE_C), + List.of("A1", "A2", "A3", "B1", "B2", "C1")), + Arguments.of( + List.of(TABLE_B, TABLE_C, TABLE_A), List.of("B1", "B2", "C1", "A1", "A2", "A3"))); + } + + @ParameterizedTest + @MethodSource("columnSchemaProvider") + void fieldReferencesIndexTheCombinedSchema(List tables, List columnNames) + throws SqlParseException { + // Reverse the expression order so a reference's index cannot accidentally be its position + // in the output expression list. All columns have the same type, so types cannot detect this. + String[] expressions = new String[columnNames.size()]; + for (int index = 0; index < expressions.length; index++) { + expressions[index] = columnNames.get(columnNames.size() - index - 1); + } + ExtendedExpression converted = new SqlExpressionToSubstrait().convert(expressions, tables); + + assertEquals(columnNames, converted.getBaseSchema().getNamesList()); + assertEquals(columnNames.size(), converted.getBaseSchema().getStruct().getTypesCount()); + assertEquals(expressions.length, converted.getReferredExprCount()); + for (int index = 0; index < expressions.length; index++) { + assertEquals( + columnNames.size() - index - 1, + selectedField(converted.getReferredExpr(index).getExpression()), + expressions[index]); + } + } + + @Test + void functionArgumentsIndexTheCombinedSchema() throws SqlParseException { + ExtendedExpression converted = + new SqlExpressionToSubstrait() + .convert(new String[] {"A1 = B1", "B2 + A3"}, List.of(TABLE_A, TABLE_B)); + + Expression.ScalarFunction filter = + converted.getReferredExpr(0).getExpression().getScalarFunction(); + assertEquals(0, selectedField(filter.getArguments(0).getValue())); + assertEquals(3, selectedField(filter.getArguments(1).getValue())); + Expression.ScalarFunction projection = + converted.getReferredExpr(1).getExpression().getScalarFunction(); + assertEquals(4, selectedField(projection.getArguments(0).getValue())); + assertEquals(2, selectedField(projection.getArguments(1).getValue())); + } + + @Test + void eachConversionBuildsItsOwnColumnIndices() throws SqlParseException { + SqlExpressionToSubstrait converter = new SqlExpressionToSubstrait(); + ExtendedExpression multipleTables = converter.convert("B2", List.of(TABLE_A, TABLE_B)); + ExtendedExpression singleTable = converter.convert("B2", List.of(TABLE_B)); + + assertEquals(4, selectedField(multipleTables.getReferredExpr(0).getExpression())); + assertEquals(1, selectedField(singleTable.getReferredExpr(0).getExpression())); + } + + private static int selectedField(Expression expression) { + assertEquals(RexTypeCase.SELECTION, expression.getRexTypeCase()); + assertTrue(expression.getSelection().hasRootReference()); + assertTrue(expression.getSelection().getDirectReference().hasStructField()); + return expression.getSelection().getDirectReference().getStructField().getField(); + } + private static Stream expressionTypeProvider() { return Stream.of( Arguments.of("2"), // I32LiteralExpression From 8c830d820c91eb607072dac9f03666c6adb7551e Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 23 Sep 2026 16:47:10 +0300 Subject: [PATCH 2/2] test(isthmus): check extended expression references through a round trip Document the combined base_schema order, the unqualified-name rule and the duplicate-name failure on convert. Give table B a VARCHAR column so a reference read back from base_schema shows the wrong index as the wrong type, and derive each expected index from the schema rather than from the input construction. --- .../isthmus/SqlExpressionToSubstrait.java | 15 ++++- .../SimpleExtendedExpressionsTest.java | 55 ++++++++++++++----- 2 files changed, 54 insertions(+), 16 deletions(-) diff --git a/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java b/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java index 35b48c30e..689adbc74 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java @@ -85,10 +85,16 @@ private static final class Result { /** * Converts a single SQL expression to a Substrait {@link io.substrait.proto.ExtendedExpression}. * + *

The expressions' {@code base_schema} holds the columns of every table, in statement order + * and then in column order, and a field reference indexes that combined list. Expressions refer + * to columns by their bare, unqualified names, so each column name must be unique across all + * tables. + * * @param sqlExpression a SQL expression * @param createStatements table creation statements defining fields referenced by the expression * @return the Substrait extended expression proto * @throws SqlParseException if parsing or validation fails + * @throws IllegalArgumentException if a column name appears more than once across the tables */ public io.substrait.proto.ExtendedExpression convert( String sqlExpression, List createStatements) throws SqlParseException { @@ -98,10 +104,16 @@ public io.substrait.proto.ExtendedExpression convert( /** * Converts multiple SQL expressions to a Substrait {@link io.substrait.proto.ExtendedExpression}. * + *

The expressions' {@code base_schema} holds the columns of every table, in statement order + * and then in column order, and a field reference indexes that combined list. Expressions refer + * to columns by their bare, unqualified names, so each column name must be unique across all + * tables. + * * @param sqlExpressions array of SQL expressions * @param createStatements table creation statements defining fields referenced by the expressions * @return the Substrait extended expression proto * @throws SqlParseException if parsing or validation fails + * @throws IllegalArgumentException if a column name appears more than once across the tables */ public io.substrait.proto.ExtendedExpression convert( String[] sqlExpressions, List createStatements) throws SqlParseException { @@ -208,7 +220,8 @@ private Result registerCreateTablesForExtendedExpression(List tables) for (SubstraitTable t : tList) { rootSchema.add(t.getName(), t); for (RelDataTypeField field : t.getRowType(factory).getFieldList()) { - // Field references index the combined base schema in insertion order. + // Index into base_schema.struct.types, whose order is this map's insertion order. + // Read the size before merging, so it is the position this field will occupy. int fieldIndex = nameToTypeMap.size(); nameToTypeMap.merge( // to validate the sql expression tree field.getName(), diff --git a/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java b/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java index f753a5243..f86fb2173 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java @@ -5,12 +5,15 @@ import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import io.substrait.extendedexpression.ProtoExtendedExpressionConverter; import io.substrait.isthmus.expression.RexExpressionConverter; import io.substrait.proto.Expression; import io.substrait.proto.Expression.RexTypeCase; import io.substrait.proto.ExtendedExpression; +import io.substrait.type.TypeCreator; import java.io.IOException; import java.util.List; +import java.util.stream.Collectors; import java.util.stream.Stream; import org.apache.calcite.sql.parser.SqlParseException; import org.junit.jupiter.api.Test; @@ -23,7 +26,7 @@ class SimpleExtendedExpressionsTest extends ExtendedExpressionTestBase { private static final String MARKER = "provider hook reached"; private static final String TABLE_A = "CREATE TABLE A (A1 BIGINT, A2 BIGINT, A3 BIGINT)"; - private static final String TABLE_B = "CREATE TABLE B (B1 BIGINT, B2 BIGINT)"; + private static final String TABLE_B = "CREATE TABLE B (B1 BIGINT, B2 VARCHAR(5))"; private static final String TABLE_C = "CREATE TABLE C (C1 BIGINT)"; private static Stream columnSchemaProvider() { @@ -43,7 +46,7 @@ private static Stream columnSchemaProvider() { void fieldReferencesIndexTheCombinedSchema(List tables, List columnNames) throws SqlParseException { // Reverse the expression order so a reference's index cannot accidentally be its position - // in the output expression list. All columns have the same type, so types cannot detect this. + // in the output expression list. String[] expressions = new String[columnNames.size()]; for (int index = 0; index < expressions.length; index++) { expressions[index] = columnNames.get(columnNames.size() - index - 1); @@ -55,8 +58,8 @@ void fieldReferencesIndexTheCombinedSchema(List tables, List col assertEquals(expressions.length, converted.getReferredExprCount()); for (int index = 0; index < expressions.length; index++) { assertEquals( - columnNames.size() - index - 1, - selectedField(converted.getReferredExpr(index).getExpression()), + columnNames.indexOf(expressions[index]), + selectedField(converted.getReferredExpr(index).getExpression(), expressions[index]), expressions[index]); } } @@ -65,16 +68,16 @@ void fieldReferencesIndexTheCombinedSchema(List tables, List col void functionArgumentsIndexTheCombinedSchema() throws SqlParseException { ExtendedExpression converted = new SqlExpressionToSubstrait() - .convert(new String[] {"A1 = B1", "B2 + A3"}, List.of(TABLE_A, TABLE_B)); + .convert(new String[] {"A1 = B1", "B1 + A3"}, List.of(TABLE_A, TABLE_B)); Expression.ScalarFunction filter = converted.getReferredExpr(0).getExpression().getScalarFunction(); - assertEquals(0, selectedField(filter.getArguments(0).getValue())); - assertEquals(3, selectedField(filter.getArguments(1).getValue())); + assertEquals(0, selectedField(filter.getArguments(0).getValue(), "A1")); + assertEquals(3, selectedField(filter.getArguments(1).getValue(), "B1")); Expression.ScalarFunction projection = converted.getReferredExpr(1).getExpression().getScalarFunction(); - assertEquals(4, selectedField(projection.getArguments(0).getValue())); - assertEquals(2, selectedField(projection.getArguments(1).getValue())); + assertEquals(3, selectedField(projection.getArguments(0).getValue(), "B1")); + assertEquals(2, selectedField(projection.getArguments(1).getValue(), "A3")); } @Test @@ -83,14 +86,36 @@ void eachConversionBuildsItsOwnColumnIndices() throws SqlParseException { ExtendedExpression multipleTables = converter.convert("B2", List.of(TABLE_A, TABLE_B)); ExtendedExpression singleTable = converter.convert("B2", List.of(TABLE_B)); - assertEquals(4, selectedField(multipleTables.getReferredExpr(0).getExpression())); - assertEquals(1, selectedField(singleTable.getReferredExpr(0).getExpression())); + assertEquals(4, selectedField(multipleTables.getReferredExpr(0).getExpression(), "B2")); + assertEquals(1, selectedField(singleTable.getReferredExpr(0).getExpression(), "B2")); } - private static int selectedField(Expression expression) { - assertEquals(RexTypeCase.SELECTION, expression.getRexTypeCase()); - assertTrue(expression.getSelection().hasRootReference()); - assertTrue(expression.getSelection().getDirectReference().hasStructField()); + @Test + void referencesReadTheirColumnTypeBackFromTheCombinedSchema() throws SqlParseException { + // Reading the proto back re-derives each reference's type from base_schema, so a reference + // indexed against its own table would come back as A2's I64 rather than B2's VARCHAR(5). + ExtendedExpression converted = + new SqlExpressionToSubstrait() + .convert(new String[] {"B2", "A3"}, List.of(TABLE_A, TABLE_B)); + io.substrait.extendedexpression.ExtendedExpression roundTripped = + new ProtoExtendedExpressionConverter().from(converted); + + assertEquals( + List.of(TypeCreator.NULLABLE.varChar(5), TypeCreator.NULLABLE.I64), + roundTripped.getReferredExpressions().stream() + .map( + reference -> + ((io.substrait.extendedexpression.ExtendedExpression.ExpressionReference) + reference) + .getExpression() + .getType()) + .collect(Collectors.toList())); + } + + private static int selectedField(Expression expression, String column) { + assertEquals(RexTypeCase.SELECTION, expression.getRexTypeCase(), column); + assertTrue(expression.getSelection().hasRootReference(), column); + assertTrue(expression.getSelection().getDirectReference().hasStructField(), column); return expression.getSelection().getDirectReference().getStructField().getField(); }