Skip to content

Commit 91bb733

Browse files
authored
fix(isthmus): support variadic concat conversion (#1001)
1 parent 3875b89 commit 91bb733

2 files changed

Lines changed: 33 additions & 0 deletions

File tree

isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -521,6 +521,13 @@ public RexNode visit(Expression.ScalarFunctionInvocation expr, Context context)
521521
.collect(Collectors.toList());
522522

523523
RelDataType returnType = typeConverter.toCalcite(typeFactory, expr.outputType());
524+
if (operator == SqlStdOperatorTable.CONCAT && args.size() > 2) {
525+
return args.stream()
526+
.skip(1)
527+
.reduce(
528+
args.get(0),
529+
(left, right) -> rexBuilder.makeCall(returnType, operator, List.of(left, right)));
530+
}
524531
return rexBuilder.makeCall(returnType, operator, args);
525532
}
526533

isthmus/src/test/java/io/substrait/isthmus/FunctionConversionTest.java

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
import org.apache.calcite.rex.RexCall;
2121
import org.apache.calcite.rex.RexNode;
2222
import org.apache.calcite.sql.SqlKind;
23+
import org.apache.calcite.sql.fun.SqlStdOperatorTable;
2324
import org.junit.jupiter.api.Test;
2425

2526
/**
@@ -212,6 +213,31 @@ void concatStringLiteralAndChar() throws Exception {
212213
assertProtoPlanRoundrip("select 'brand_'||P_BRAND from PART");
213214
}
214215

216+
@Test
217+
void variadicConcat() {
218+
ScalarFunctionInvocation concat =
219+
sb.scalarFn(
220+
DefaultExtensionCatalog.FUNCTIONS_STRING,
221+
"concat:str",
222+
TypeCreator.REQUIRED.STRING,
223+
Expression.StrLiteral.builder().value("a").build(),
224+
Expression.StrLiteral.builder().value("b").build(),
225+
Expression.StrLiteral.builder().value("c").build());
226+
227+
RexCall outer =
228+
assertInstanceOf(
229+
RexCall.class, concat.accept(expressionRexConverter, Context.newContext()));
230+
assertEquals(SqlStdOperatorTable.CONCAT, outer.getOperator());
231+
assertEquals(2, outer.getOperands().size());
232+
233+
RexCall inner = assertInstanceOf(RexCall.class, outer.getOperands().get(0));
234+
assertEquals(SqlStdOperatorTable.CONCAT, inner.getOperator());
235+
assertEquals(2, inner.getOperands().size());
236+
assertEquals("'a':VARCHAR", inner.getOperands().get(0).toString());
237+
assertEquals("'b':VARCHAR", inner.getOperands().get(1).toString());
238+
assertEquals("'c':VARCHAR", outer.getOperands().get(1).toString());
239+
}
240+
215241
@Test
216242
void strptimeTime() {
217243
Expression.StrLiteral inputString = Expression.StrLiteral.builder().value("12:34:56").build();

0 commit comments

Comments
 (0)