This is an automated email from the ASF dual-hosted git repository.

mihaibudiu pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/calcite.git


The following commit(s) were added to refs/heads/main by this push:
     new a2a973ee07 [CALCITE-7688] Support scalar subqueries in table function 
arguments
a2a973ee07 is described below

commit a2a973ee074065553f082827b654cf43693ae060
Author: Kirill Tkalenko <[email protected]>
AuthorDate: Fri Aug 7 18:23:14 2026 +0300

    [CALCITE-7688] Support scalar subqueries in table function arguments
---
 .../org/apache/calcite/rel/rules/CoreRules.java    |  10 ++
 .../calcite/rel/rules/SubQueryRemoveRule.java      |  87 +++++++++++++
 .../main/java/org/apache/calcite/rex/RexUtil.java  |  21 ++-
 .../java/org/apache/calcite/tools/Programs.java    |   2 +
 .../java/org/apache/calcite/rex/RexUtilTest.java   |  46 +++++++
 .../org/apache/calcite/test/TableFunctionTest.java | 144 +++++++++++++++++++++
 .../main/java/org/apache/calcite/util/Smalls.java  |  65 ++++++++++
 7 files changed, 374 insertions(+), 1 deletion(-)

diff --git a/core/src/main/java/org/apache/calcite/rel/rules/CoreRules.java 
b/core/src/main/java/org/apache/calcite/rel/rules/CoreRules.java
index 104e34bfae..11708c5fac 100644
--- a/core/src/main/java/org/apache/calcite/rel/rules/CoreRules.java
+++ b/core/src/main/java/org/apache/calcite/rel/rules/CoreRules.java
@@ -518,6 +518,16 @@ private CoreRules() {}
   public static final SubQueryRemoveRule JOIN_SUB_QUERY_TO_CORRELATE =
       SubQueryRemoveRule.Config.JOIN.toRule();
 
+  /** Rule that converts scalar sub-queries from table function arguments into
+   * {@link Correlate} instances.
+   *
+   * @see #PROJECT_SUB_QUERY_TO_CORRELATE
+   * @see #FILTER_SUB_QUERY_TO_CORRELATE
+   * @see #JOIN_SUB_QUERY_TO_CORRELATE */
+  public static final SubQueryRemoveRule
+      TABLE_FUNCTION_SCAN_SCALAR_QUERY_TO_CORRELATE =
+      SubQueryRemoveRule.Config.TABLE_FUNCTION_SCAN_SCALAR_QUERY.toRule();
+
   /** Rule that converts sub-queries from filter expressions into
    * {@link Correlate} instances. It will rewrite SOME/EXISTS/IN to a LEFT 
MARK type Correlate. */
   public static final SubQueryRemoveRule FILTER_SUB_QUERY_TO_MARK_CORRELATE =
diff --git 
a/core/src/main/java/org/apache/calcite/rel/rules/SubQueryRemoveRule.java 
b/core/src/main/java/org/apache/calcite/rel/rules/SubQueryRemoveRule.java
index f7a6ce552f..c14e654a01 100644
--- a/core/src/main/java/org/apache/calcite/rel/rules/SubQueryRemoveRule.java
+++ b/core/src/main/java/org/apache/calcite/rel/rules/SubQueryRemoveRule.java
@@ -28,6 +28,8 @@
 import org.apache.calcite.rel.core.Join;
 import org.apache.calcite.rel.core.JoinRelType;
 import org.apache.calcite.rel.core.Project;
+import org.apache.calcite.rel.core.TableFunctionScan;
+import org.apache.calcite.rel.logical.LogicalCorrelate;
 import org.apache.calcite.rel.metadata.RelMdUtil;
 import org.apache.calcite.rel.metadata.RelMetadataQuery;
 import org.apache.calcite.rex.LogicVisitor;
@@ -77,6 +79,7 @@
  * @see CoreRules#FILTER_SUB_QUERY_TO_CORRELATE
  * @see CoreRules#PROJECT_SUB_QUERY_TO_CORRELATE
  * @see CoreRules#JOIN_SUB_QUERY_TO_CORRELATE
+ * @see CoreRules#TABLE_FUNCTION_SCAN_SCALAR_QUERY_TO_CORRELATE
  */
 @Value.Enclosing
 public class SubQueryRemoveRule
@@ -1017,6 +1020,78 @@ private static void matchFilter(SubQueryRemoveRule rule,
     call.transformTo(builder.build());
   }
 
+  /**
+   * Rewrites one scalar sub-query in a table-function call into a
+   * {@link Correlate}.
+   *
+   * <p>For example, converts:
+   *
+   * <pre>{@code
+   * LogicalTableFunctionScan(invocation=[F(SCALAR_QUERY(SUB_QUERY_REL), 20)])
+   * }</pre>
+   *
+   * <p>into:
+   *
+   * <pre>{@code
+   * LogicalProject(TABLE_FUNCTION_FIELDS)
+   *   LogicalCorrelate(correlation=[$cor0], joinType=[inner],
+   *       requiredColumns=[{0}])
+   *     LogicalAggregate(group=[{}], scalarValue=[SINGLE_VALUE($0)])
+   *       SUB_QUERY_REL
+   *     LogicalTableFunctionScan(invocation=[F($cor0.scalarValue, 20)])
+   * }</pre>
+   *
+   * <p>The aggregate implements scalar-query cardinality: zero rows produce
+   * {@code null}, and more than one row raises an error. The correlate makes
+   * the resulting value available to the table-function invocation, and the
+   * project removes the helper scalar field from the output. The rule rewrites
+   * one scalar sub-query per invocation; subsequent invocations rewrite any
+   * remaining scalar sub-queries.
+   */
+  private static void matchTableFunctionScan(SubQueryRemoveRule rule,
+      RelOptRuleCall call) {
+    final TableFunctionScan scan = call.rel(0);
+    final RexSubQuery e =
+        requireNonNull(
+            RexUtil.SubQueryFinder.find(scan.getCall(), SqlKind.SCALAR_QUERY));
+
+    final RelBuilder builder = call.builder();
+    builder.push(e.rel);
+    builder.aggregate(builder.groupKey(),
+        builder.aggregateCall(SqlStdOperatorTable.SINGLE_VALUE,
+            builder.field(0)));
+    final RelNode scalarValue = builder.build();
+
+    final CorrelationId correlationId =
+        scan.getCluster().createCorrel();
+    final RexCorrelVariable correlationVariable =
+        (RexCorrelVariable) scan.getCluster().getRexBuilder()
+            .makeCorrel(scalarValue.getRowType(), correlationId);
+    final RexNode target =
+        scan.getCluster().getRexBuilder()
+            .makeFieldAccess(correlationVariable, 0);
+    final RexNode newCall =
+        scan.getCall().accept(new ReplaceSubQueryShuttle(e, target));
+    final TableFunctionScan newScan =
+        (TableFunctionScan) scan.copy(scan.getTraitSet(), scan.getInputs(),
+            newCall, scan.getElementType(), scan.getRowType(),
+            scan.getColumnMappings())
+            .withHints(scan.getHints());
+
+    final RelNode correlate =
+        LogicalCorrelate.create(scalarValue, newScan, ImmutableList.of(),
+            correlationId, ImmutableBitSet.of(0), JoinRelType.INNER);
+    builder.push(correlate);
+    final int scalarFieldCount =
+        scalarValue.getRowType().getFieldCount();
+    builder.project(
+        IntStream.range(0, scan.getRowType().getFieldCount())
+            .mapToObj(i -> builder.field(scalarFieldCount + i))
+            .collect(Collectors.toList()),
+        scan.getRowType().getFieldNames());
+    call.transformTo(builder.build());
+  }
+
   private static void matchJoin(SubQueryRemoveRule rule, RelOptRuleCall call) {
     final Join join = call.rel(0);
     final RelBuilder builder = call.builder();
@@ -1372,6 +1447,18 @@ public interface Config extends RelRule.Config {
                 .anyInputs())
         .withDescription("SubQueryRemoveRule:Join");
 
+    Config TABLE_FUNCTION_SCAN_SCALAR_QUERY =
+        ImmutableSubQueryRemoveRule.Config.builder()
+            .withMatchHandler(SubQueryRemoveRule::matchTableFunctionScan)
+            .build()
+            .withOperandSupplier(b ->
+                b.operand(TableFunctionScan.class)
+                    .predicate(scan ->
+                        RexUtil.SubQueryFinder.find(scan.getCall(),
+                            SqlKind.SCALAR_QUERY) != null)
+                    .anyInputs())
+            
.withDescription("SubQueryRemoveRule:TableFunctionScanScalarQuery");
+
     Config PROJECT_ENABLE_MARK_JOIN = 
ImmutableSubQueryRemoveRule.Config.builder()
         .withMatchHandler(SubQueryRemoveRule::matchProjectEnableMarkJoin)
         .build()
diff --git a/core/src/main/java/org/apache/calcite/rex/RexUtil.java 
b/core/src/main/java/org/apache/calcite/rex/RexUtil.java
index b7c30b7608..2faa2633b0 100644
--- a/core/src/main/java/org/apache/calcite/rex/RexUtil.java
+++ b/core/src/main/java/org/apache/calcite/rex/RexUtil.java
@@ -3365,6 +3365,7 @@ public static List<RexSubQuery> collect(Project project) {
    * applied to an expression that contains a {@link RexSubQuery}. */
   public static class SubQueryFinder extends RexVisitorImpl<Void> {
     public static final SubQueryFinder INSTANCE = new SubQueryFinder();
+    private final @Nullable SqlKind kind;
 
     @SuppressWarnings("Guava")
     @Deprecated // to be removed before 2.0
@@ -3382,7 +3383,12 @@ public static class SubQueryFinder extends 
RexVisitorImpl<Void> {
         SubQueryFinder::containsSubQuery;
 
     private SubQueryFinder() {
+      this(null);
+    }
+
+    private SubQueryFinder(@Nullable SqlKind kind) {
       super(true);
+      this.kind = kind;
     }
 
     /** Returns whether a {@link Project} contains a sub-query. */
@@ -3418,7 +3424,10 @@ public static boolean containsSubQuery(Join join) {
     }
 
     @Override public Void visitSubQuery(RexSubQuery subQuery) {
-      throw new Util.FoundOne(subQuery);
+      if (kind == null || subQuery.getKind() == kind) {
+        throw new Util.FoundOne(subQuery);
+      }
+      return super.visitSubQuery(subQuery);
     }
 
     public static @Nullable RexSubQuery find(Iterable<RexNode> nodes) {
@@ -3440,6 +3449,16 @@ public static boolean containsSubQuery(Join join) {
         return (RexSubQuery) e.getNode();
       }
     }
+
+    /** Returns the first sub-query of the given kind, or {@code null}. */
+    public static @Nullable RexSubQuery find(RexNode node, SqlKind kind) {
+      try {
+        node.accept(new SubQueryFinder(kind));
+        return null;
+      } catch (Util.FoundOne e) {
+        return (RexSubQuery) e.getNode();
+      }
+    }
   }
 
   /** Deep expressions simplifier.
diff --git a/core/src/main/java/org/apache/calcite/tools/Programs.java 
b/core/src/main/java/org/apache/calcite/tools/Programs.java
index 83f6834d28..1307e70100 100644
--- a/core/src/main/java/org/apache/calcite/tools/Programs.java
+++ b/core/src/main/java/org/apache/calcite/tools/Programs.java
@@ -260,6 +260,7 @@ public static Program subQuery(RelMetadataProvider 
metadataProvider) {
         ImmutableList.of(CoreRules.FILTER_SUB_QUERY_TO_CORRELATE,
             CoreRules.PROJECT_SUB_QUERY_TO_CORRELATE,
             CoreRules.JOIN_SUB_QUERY_TO_CORRELATE,
+            CoreRules.TABLE_FUNCTION_SCAN_SCALAR_QUERY_TO_CORRELATE,
             CoreRules.PROJECT_OVER_SUM_TO_SUM0_RULE));
     final Program oldProgram = of(builder.build(), true, metadataProvider);
 
@@ -268,6 +269,7 @@ public static Program subQuery(RelMetadataProvider 
metadataProvider) {
         ImmutableList.of(CoreRules.FILTER_SUB_QUERY_TO_MARK_CORRELATE,
             CoreRules.PROJECT_SUB_QUERY_TO_MARK_CORRELATE,
             CoreRules.JOIN_SUB_QUERY_TO_CORRELATE,
+            CoreRules.TABLE_FUNCTION_SCAN_SCALAR_QUERY_TO_CORRELATE,
             CoreRules.PROJECT_OVER_SUM_TO_SUM0_RULE));
     final Program newProgram = of(newBuilder.build(), true, metadataProvider);
 
diff --git a/core/src/test/java/org/apache/calcite/rex/RexUtilTest.java 
b/core/src/test/java/org/apache/calcite/rex/RexUtilTest.java
new file mode 100644
index 0000000000..229886f6d0
--- /dev/null
+++ b/core/src/test/java/org/apache/calcite/rex/RexUtilTest.java
@@ -0,0 +1,46 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to you under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.calcite.rex;
+
+import org.apache.calcite.rel.RelNode;
+import org.apache.calcite.sql.SqlKind;
+import org.apache.calcite.sql.fun.SqlStdOperatorTable;
+import org.apache.calcite.tools.Frameworks;
+import org.apache.calcite.tools.RelBuilder;
+
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertNull;
+import static org.junit.jupiter.api.Assertions.assertSame;
+
+/** Tests for {@link RexUtil}. */
+class RexUtilTest {
+  @Test void testSubQueryFinderByKind() {
+    final RelBuilder builder =
+        RelBuilder.create(Frameworks.newConfigBuilder().build());
+    final RelNode rel = builder.values(new String[] {"i"}, 1).build();
+    final RexSubQuery arrayQuery = RexSubQuery.array(rel);
+    final RexSubQuery scalarQuery = RexSubQuery.scalar(rel);
+    final RexNode expression = rel.getCluster().getRexBuilder()
+        .makeCall(SqlStdOperatorTable.ROW, arrayQuery, scalarQuery);
+
+    assertSame(arrayQuery, RexUtil.SubQueryFinder.find(expression));
+    assertSame(scalarQuery,
+        RexUtil.SubQueryFinder.find(expression, SqlKind.SCALAR_QUERY));
+    assertNull(RexUtil.SubQueryFinder.find(expression, SqlKind.EXISTS));
+  }
+}
diff --git a/core/src/test/java/org/apache/calcite/test/TableFunctionTest.java 
b/core/src/test/java/org/apache/calcite/test/TableFunctionTest.java
index e34845c1e6..71aed9c3b9 100644
--- a/core/src/test/java/org/apache/calcite/test/TableFunctionTest.java
+++ b/core/src/test/java/org/apache/calcite/test/TableFunctionTest.java
@@ -59,6 +59,9 @@ private CalciteAssert.AssertThat with() {
     final String m = Smalls.MULTIPLICATION_TABLE_METHOD.getName();
     final String m2 = Smalls.FIBONACCI_TABLE_METHOD.getName();
     final String m3 = Smalls.FIBONACCI_LIMIT_TABLE_METHOD.getName();
+    final String m4 = Smalls.SCALAR_QUERY_ARGUMENTS_TABLE_METHOD.getName();
+    final String m5 =
+        Smalls.SCALAR_QUERY_ARGUMENTS_TABLE_WITHOUT_COLUMN_METHOD.getName();
     return CalciteAssert.model("{\n"
         + "  version: '1.0',\n"
         + "   schemas: [\n"
@@ -77,6 +80,14 @@ private CalciteAssert.AssertThat with() {
         + "           name: 'fibonacci2',\n"
         + "           className: '" + c + "',\n"
         + "           methodName: '" + m3 + "'\n"
+        + "         }, {\n"
+        + "           name: 'scalar_query_arguments',\n"
+        + "           className: '" + c + "',\n"
+        + "           methodName: '" + m4 + "'\n"
+        + "         }, {\n"
+        + "           name: 'scalar_query_arguments_without_column',\n"
+        + "           className: '" + c + "',\n"
+        + "           methodName: '" + m5 + "'\n"
         + "         }\n"
         + "       ]\n"
         + "     }\n"
@@ -400,6 +411,139 @@ private Connection getConnectionWithMultiplyFunction() 
throws SQLException {
             "row_name=row 2; c1=103; c2=106");
   }
 
+  @Test void testTableFunctionWithScalarQueryLiteralAndColumnArguments() {
+    final String sql = "select d.n as outer_value,\n"
+        + "  f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value,\n"
+        + "  f.\"column_value\" as column_value\n"
+        + "from (values (100), (200)) as d(n)\n"
+        + "cross join lateral table(\"s\".\"scalar_query_arguments\"(\n"
+        + "  (select 10), 20, d.n)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "OUTER_VALUE=100; SCALAR_VALUE=10; LITERAL_VALUE=20; 
COLUMN_VALUE=100",
+            "OUTER_VALUE=200; SCALAR_VALUE=10; LITERAL_VALUE=20; 
COLUMN_VALUE=200");
+  }
+
+  @Test void testTableFunctionWithScalarQueriesAndColumnArguments() {
+    final String sql = "select d.n as outer_value,\n"
+        + "  f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value,\n"
+        + "  f.\"column_value\" as column_value\n"
+        + "from (values (100), (200)) as d(n)\n"
+        + "cross join lateral table(\"s\".\"scalar_query_arguments\"(\n"
+        + "  (select 10), (select 20), d.n)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "OUTER_VALUE=100; SCALAR_VALUE=10; LITERAL_VALUE=20; 
COLUMN_VALUE=100",
+            "OUTER_VALUE=200; SCALAR_VALUE=10; LITERAL_VALUE=20; 
COLUMN_VALUE=200");
+  }
+
+  @Test void testTableFunctionWithScalarQueryAndLiteralArguments() {
+    final String sql = "select f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value\n"
+        + "from table(\"s\".\"scalar_query_arguments_without_column\"(\n"
+        + "  (select 10), 20)) as f";
+    with().query(sql)
+        .returnsUnordered("SCALAR_VALUE=10; LITERAL_VALUE=20");
+  }
+
+  @Test void testTableFunctionWithEmptyScalarQuery() {
+    final String sql = "select f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value\n"
+        + "from table(\"s\".\"scalar_query_arguments_without_column\"(\n"
+        + "  (select v from (values (10)) as q(v) where v < 0), 20)) as f";
+    with().query(sql)
+        .returnsUnordered("SCALAR_VALUE=null; LITERAL_VALUE=20");
+  }
+
+  @Test void testTableFunctionWithMultiRowScalarQuery() {
+    final String sql = "select *\n"
+        + "from table(\"s\".\"scalar_query_arguments_without_column\"(\n"
+        + "  (select v from (values (10), (20)) as q(v)), 20))";
+    with().query(sql)
+        .throws_("more than one value in agg SINGLE_VALUE");
+  }
+
+  @Test void testTableFunctionWithCorrelatedScalarQueryAndLiteralArguments() {
+    final String sql = "select d.n as outer_value,\n"
+        + "  f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value\n"
+        + "from (values (100), (200)) as d(n)\n"
+        + "cross join lateral 
table(\"s\".\"scalar_query_arguments_without_column\"(\n"
+        + "  (select d.n + 1), 20)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "OUTER_VALUE=100; SCALAR_VALUE=101; LITERAL_VALUE=20",
+            "OUTER_VALUE=200; SCALAR_VALUE=201; LITERAL_VALUE=20");
+  }
+
+  @Test void
+      testTableFunctionWithCorrelatedScalarQueryLiteralAndColumnArguments() {
+    final String sql = "select d.n as outer_value,\n"
+        + "  f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value,\n"
+        + "  f.\"column_value\" as column_value\n"
+        + "from (values (100), (200)) as d(n)\n"
+        + "cross join lateral table(\"s\".\"scalar_query_arguments\"(\n"
+        + "  (select d.n + 1), 20, d.n)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "OUTER_VALUE=100; SCALAR_VALUE=101; LITERAL_VALUE=20; 
COLUMN_VALUE=100",
+            "OUTER_VALUE=200; SCALAR_VALUE=201; LITERAL_VALUE=20; 
COLUMN_VALUE=200");
+  }
+
+  @Test void 
testTableFunctionWithScalarQueryExpressionLiteralAndColumnArguments() {
+    final String sql = "select d.n as outer_value,\n"
+        + "  f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value,\n"
+        + "  f.\"column_value\" as column_value\n"
+        + "from (values (100), (200)) as d(n)\n"
+        + "cross join lateral table(\"s\".\"scalar_query_arguments\"(\n"
+        + "  (select 4) + (select 6), 20, d.n)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "OUTER_VALUE=100; SCALAR_VALUE=10; LITERAL_VALUE=20; 
COLUMN_VALUE=100",
+            "OUTER_VALUE=200; SCALAR_VALUE=10; LITERAL_VALUE=20; 
COLUMN_VALUE=200");
+  }
+
+  @Test void testTableFunctionWithScalarQueryExpressionAndLiteralArguments() {
+    final String sql = "select f.\"scalar_value\" as scalar_value,\n"
+        + "  f.\"literal_value\" as literal_value\n"
+        + "from table(\"s\".\"scalar_query_arguments_without_column\"(\n"
+        + "  (select 4) + (select 6), 20)) as f";
+    with().query(sql)
+        .returnsUnordered("SCALAR_VALUE=10; LITERAL_VALUE=20");
+  }
+
+  @Test void testTableFunctionWithRowScalarQueryLiteralAndColumnArguments() {
+    final String sql = "select d.n as outer_value,\n"
+        + "  f.\"row_value_0\" as row_value_0,\n"
+        + "  f.\"row_value_1\" as row_value_1,\n"
+        + "  f.\"literal_value\" as literal_value,\n"
+        + "  f.\"column_value\" as column_value\n"
+        + "from (values (100), (200)) as d(n)\n"
+        + "cross join lateral table(\"s\".\"scalar_query_arguments\"(\n"
+        + "  (select row(1, 2)), 20, d.n)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "OUTER_VALUE=100; ROW_VALUE_0=1; ROW_VALUE_1=2; "
+                + "LITERAL_VALUE=20; COLUMN_VALUE=100",
+            "OUTER_VALUE=200; ROW_VALUE_0=1; ROW_VALUE_1=2; "
+                + "LITERAL_VALUE=20; COLUMN_VALUE=200");
+  }
+
+  @Test void testTableFunctionWithRowScalarQueryAndLiteralArguments() {
+    final String sql = "select f.\"row_value_0\" as row_value_0,\n"
+        + "  f.\"row_value_1\" as row_value_1,\n"
+        + "  f.\"literal_value\" as literal_value\n"
+        + "from table(\"s\".\"scalar_query_arguments_without_column\"(\n"
+        + "  (select row(1, 2)), 20)) as f";
+    with().query(sql)
+        .returnsUnordered(
+            "ROW_VALUE_0=1; ROW_VALUE_1=2; LITERAL_VALUE=20");
+  }
+
   /** Tests a query with a table function in the FROM clause,
    * attempting to reference a column from the table function in the WHERE
    * clause but getting the case wrong.
diff --git a/testkit/src/main/java/org/apache/calcite/util/Smalls.java 
b/testkit/src/main/java/org/apache/calcite/util/Smalls.java
index 6f33f00669..fe4def986a 100644
--- a/testkit/src/main/java/org/apache/calcite/util/Smalls.java
+++ b/testkit/src/main/java/org/apache/calcite/util/Smalls.java
@@ -111,6 +111,12 @@ public class Smalls {
   public static final Method MULTIPLICATION_TABLE_METHOD =
       Types.lookupMethod(Smalls.class, "multiplicationTable", int.class,
         int.class, Integer.class);
+  public static final Method SCALAR_QUERY_ARGUMENTS_TABLE_METHOD =
+      Types.lookupMethod(Smalls.class, "scalarQueryArgumentsTable",
+          Object.class, Integer.class, Integer.class);
+  public static final Method 
SCALAR_QUERY_ARGUMENTS_TABLE_WITHOUT_COLUMN_METHOD =
+      Types.lookupMethod(Smalls.class,
+          "scalarQueryArgumentsTableWithoutColumn", Object.class, int.class);
   public static final Method FIBONACCI_TABLE_METHOD =
       Types.lookupMethod(Smalls.class, "fibonacciTable");
   public static final Method FIBONACCI_LIMIT_100_TABLE_METHOD =
@@ -283,6 +289,65 @@ public static QueryableTable multiplicationTable(final int 
ncol,
     };
   }
 
+  /** A one-row table containing the arguments passed to the function. */
+  public static QueryableTable scalarQueryArgumentsTable(
+      final @Nullable Object scalarValue,
+      final @Nullable Integer literalValue,
+      final @Nullable Integer columnValue) {
+    final @Nullable Integer normalizedScalarValue;
+    final @Nullable Integer rowValue0;
+    final @Nullable Integer rowValue1;
+    if (scalarValue == null) {
+      normalizedScalarValue = null;
+      rowValue0 = null;
+      rowValue1 = null;
+    } else if (scalarValue instanceof Number) {
+      normalizedScalarValue = ((Number) scalarValue).intValue();
+      rowValue0 = null;
+      rowValue1 = null;
+    } else if (scalarValue instanceof Object[]) {
+      final Object[] row = (Object[]) scalarValue;
+      if (row.length != 2
+          || !(row[0] instanceof Number)
+          || !(row[1] instanceof Number)) {
+        throw new IllegalArgumentException("expected ROW with two numbers");
+      }
+      normalizedScalarValue = null;
+      rowValue0 = ((Number) row[0]).intValue();
+      rowValue1 = ((Number) row[1]).intValue();
+    } else {
+      throw new IllegalArgumentException("expected a number or ROW");
+    }
+    return new AbstractQueryableTable(Object[].class) {
+      @Override public RelDataType getRowType(RelDataTypeFactory typeFactory) {
+        return typeFactory.builder()
+            .add("scalar_value", typeFactory.createJavaType(Integer.class))
+            .add("row_value_0", typeFactory.createJavaType(Integer.class))
+            .add("row_value_1", typeFactory.createJavaType(Integer.class))
+            .add("literal_value", typeFactory.createJavaType(Integer.class))
+            .add("column_value", typeFactory.createJavaType(Integer.class))
+            .build();
+      }
+
+      @Override public Queryable<Object[]> asQueryable(
+          QueryProvider queryProvider, SchemaPlus schema, String tableName) {
+        return Linq4j.asEnumerable(
+            Collections.singletonList(
+                new Object[] {
+                    normalizedScalarValue, rowValue0, rowValue1,
+                    literalValue, columnValue
+                }))
+            .asQueryable();
+      }
+    };
+  }
+
+  /** A two-argument version of {@link #scalarQueryArgumentsTable}. */
+  public static QueryableTable scalarQueryArgumentsTableWithoutColumn(
+      final @Nullable Object scalarValue, final int literalValue) {
+    return scalarQueryArgumentsTable(scalarValue, literalValue, null);
+  }
+
   /** A function that generates the Fibonacci sequence.
    *
    * <p>Interesting because it has one column and no arguments,

Reply via email to