diff --git a/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java b/isthmus/src/main/java/io/substrait/isthmus/SqlExpressionToSubstrait.java
index 620d12927..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,6 +220,9 @@ private Result registerCreateTablesForExtendedExpression(List tables)
for (SubstraitTable t : tList) {
rootSchema.add(t.getName(), t);
for (RelDataTypeField field : t.getRowType(factory).getFieldList()) {
+ // 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(),
field.getType(),
@@ -217,7 +232,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..f86fb2173 100644
--- a/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java
+++ b/isthmus/src/test/java/io/substrait/isthmus/SimpleExtendedExpressionsTest.java
@@ -5,10 +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;
@@ -20,6 +25,100 @@ 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 VARCHAR(5))";
+ 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.
+ 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.indexOf(expressions[index]),
+ selectedField(converted.getReferredExpr(index).getExpression(), expressions[index]),
+ expressions[index]);
+ }
+ }
+
+ @Test
+ void functionArgumentsIndexTheCombinedSchema() throws SqlParseException {
+ ExtendedExpression converted =
+ new SqlExpressionToSubstrait()
+ .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(), "A1"));
+ assertEquals(3, selectedField(filter.getArguments(1).getValue(), "B1"));
+ Expression.ScalarFunction projection =
+ converted.getReferredExpr(1).getExpression().getScalarFunction();
+ assertEquals(3, selectedField(projection.getArguments(0).getValue(), "B1"));
+ assertEquals(2, selectedField(projection.getArguments(1).getValue(), "A3"));
+ }
+
+ @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(), "B2"));
+ assertEquals(1, selectedField(singleTable.getReferredExpr(0).getExpression(), "B2"));
+ }
+
+ @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();
+ }
+
private static Stream expressionTypeProvider() {
return Stream.of(
Arguments.of("2"), // I32LiteralExpression