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

github-bot pushed a commit to branch site
in repository https://gitbox.apache.org/repos/asf/calcite.git


The following commit(s) were added to refs/heads/site by this push:
     new 5447e1ddfc [CALCITE-6704] Limit result size of RelMdUniqueKeys handler
5447e1ddfc is described below

commit 5447e1ddfc2e8c056f9296c21dfeeda501da9695
Author: Stamatis Zampetakis <[email protected]>
AuthorDate: Fri Dec 6 12:07:46 2024 +0100

    [CALCITE-6704] Limit result size of RelMdUniqueKeys handler
    
    For certain query patterns RelMdUniqueKeys handler generates an
    exponentially large number of unique keys that results into crashes
    and OOM errors. The limit guards against the combinatorial explosion
    that may appear for such use-cases and provides the users of a way to
    tune further the upper bound if needed.
    
    Close apache/calcite#4089
---
 .../calcite/rel/metadata/RelMdUniqueKeys.java      | 126 ++++++---
 .../org/apache/calcite/test/RelMetadataTest.java   | 288 +++++++++++++++++++++
 site/_docs/history.md                              |  11 +
 .../apache/calcite/test/RelMetadataFixture.java    |  21 +-
 4 files changed, 395 insertions(+), 51 deletions(-)

diff --git 
a/core/src/main/java/org/apache/calcite/rel/metadata/RelMdUniqueKeys.java 
b/core/src/main/java/org/apache/calcite/rel/metadata/RelMdUniqueKeys.java
index 31707997ea..8a21497371 100644
--- a/core/src/main/java/org/apache/calcite/rel/metadata/RelMdUniqueKeys.java
+++ b/core/src/main/java/org/apache/calcite/rel/metadata/RelMdUniqueKeys.java
@@ -68,16 +68,41 @@ import static java.util.Objects.requireNonNull;
 /**
  * RelMdUniqueKeys supplies a default implementation of
  * {@link RelMetadataQuery#getUniqueKeys} for the standard logical algebra.
+ * The number of returned keys for each relational expression is bounded by a 
limit.
+ * The limit is used to restrict the exponential logic that can appear for 
certain query patterns
+ * and lead to CPU/memory exhaustion and crashes.
  */
 public class RelMdUniqueKeys
     implements MetadataHandler<BuiltInMetadata.UniqueKeys> {
   public static final RelMetadataProvider SOURCE =
       ReflectiveRelMetadataProvider.reflectiveSource(
           new RelMdUniqueKeys(), BuiltInMetadata.UniqueKeys.Handler.class);
-
+  /**
+   * A limit about the number of unique keys returned by the handler.
+   * The limit must be in the range [0, Integer.MAX_VALUE].
+   */
+  private final int limit;
   //~ Constructors -----------------------------------------------------------
 
-  private RelMdUniqueKeys() {}
+  /**
+   * Creates a metadata handler for unique keys with the default limit.
+   */
+  public RelMdUniqueKeys() {
+    this(1000);
+  }
+
+  /**
+   * Creates a metadata handler for unique keys with the specified limit.
+   *
+   * @param limit a non-negative integer that bounds the number of unique keys 
returned for each
+   *              relational expression.
+   */
+  public RelMdUniqueKeys(int limit) {
+    if (limit < 0) {
+      throw new IllegalArgumentException("Limit cannot be negative");
+    }
+    this.limit = limit;
+  }
 
   //~ Methods ----------------------------------------------------------------
 
@@ -105,6 +130,9 @@ public class RelMdUniqueKeys
 
   public @Nullable Set<ImmutableBitSet> getUniqueKeys(Sort rel, 
RelMetadataQuery mq,
       boolean ignoreNulls) {
+    if (limit == 0) {
+      return ImmutableSet.of();
+    }
     Double maxRowCount = mq.getMaxRowCount(rel);
     if (maxRowCount != null && maxRowCount <= 1.0d) {
       return ImmutableSet.of(ImmutableBitSet.of());
@@ -134,7 +162,7 @@ public class RelMdUniqueKeys
         Util.transform(program.getProjectList(), program::expandLocalRef));
   }
 
-  private static Set<ImmutableBitSet> getProjectUniqueKeys(SingleRel rel, 
RelMetadataQuery mq,
+  private Set<ImmutableBitSet> getProjectUniqueKeys(SingleRel rel, 
RelMetadataQuery mq,
       boolean ignoreNulls, List<RexNode> projExprs) {
     // LogicalProject maps a set of rows to a different set;
     // Without knowledge of the mapping function(whether it
@@ -175,30 +203,28 @@ public class RelMdUniqueKeys
     Map<Integer, ImmutableBitSet> mapInToOutPos =
         Maps.transformValues(inToOutPosBuilder.build().asMap(), 
ImmutableBitSet::of);
 
-    ImmutableSet.Builder<ImmutableBitSet> resultBuilder = 
ImmutableSet.builder();
+    Set<ImmutableBitSet> resultBuilder = new HashSet<>();
     // Now add to the projUniqueKeySet the child keys that are fully
     // projected.
+    outerLoop:
     for (ImmutableBitSet colMask : childUniqueKeySet) {
       if (!inColumnsUsed.contains(colMask)) {
         // colMask contains a column that is not projected as RexInput => the 
key is not unique
         continue;
       }
       // colMask is mapped to output project, however, the column can be 
mapped more than once:
-      // select id, id, id, unique2, unique2
-      // the resulting unique keys would be {{0},{3}}, {{0},{4}}, 
{{0},{1},{4}}, ...
+      // select key1, key1, val1, val2, key2 from ...
+      // the resulting unique keys would be {{0},{4}}, {{1},{4}}
 
-      Iterable<List<ImmutableBitSet>> product =
-          Linq4j.product(
-              Util.transform(colMask, in ->
-                  Util.filter(
-                      requireNonNull(mapInToOutPos.get(in),
-                          () -> "no entry for column " + in
-                              + " in mapInToOutPos: " + 
mapInToOutPos).powerSet(),
-                      bs -> !bs.isEmpty())));
-
-      resultBuilder.addAll(Util.transform(product, ImmutableBitSet::union));
+      Iterable<List<Integer>> product = Linq4j.product(Util.transform(colMask, 
mapInToOutPos::get));
+      for (List<Integer> passKey : product) {
+        if (resultBuilder.size() == limit) {
+          break outerLoop;
+        }
+        resultBuilder.add(ImmutableBitSet.of(passKey));
+      }
     }
-    return resultBuilder.build();
+    return resultBuilder;
   }
 
   public @Nullable Set<ImmutableBitSet> getUniqueKeys(Join rel, 
RelMetadataQuery mq,
@@ -283,17 +309,7 @@ public class RelMdUniqueKeys
           .forEach(retSet::add);
     }
 
-    // Remove sets that are supersets of other sets
-    final Set<ImmutableBitSet> reducedSet = new HashSet<>();
-    for (ImmutableBitSet bigger : retSet) {
-      if (retSet.stream()
-          .filter(smaller -> !bigger.equals(smaller))
-          .noneMatch(bigger::contains)) {
-        reducedSet.add(bigger);
-      }
-    }
-
-    return reducedSet;
+    return filterSupersets(retSet, limit);
   }
 
   /**
@@ -344,15 +360,23 @@ public class RelMdUniqueKeys
 
       // If an input's unique column(s) value is returned (passed through) by 
an aggregation
       // function, then the result of the function(s) is also unique.
-      final ImmutableSet.Builder<ImmutableBitSet> keysBuilder = 
ImmutableSet.builder();
+      Set<ImmutableBitSet> keysBuilder = new HashSet<>();
       if (inputUniqueKeys != null) {
+        outerLoop:
         for (ImmutableBitSet inputKey : inputUniqueKeys) {
-          keysBuilder.addAll(getPassedThroughCols(inputKey, rel));
+          Iterable<List<Integer>> product =
+              Linq4j.product(Util.transform(inputKey, i -> 
getPassedThroughCols(i, rel)));
+          for (List<Integer> passKey : product) {
+            if (keysBuilder.size() == limit) {
+              break outerLoop;
+            }
+            keysBuilder.add(ImmutableBitSet.of(passKey));
+          }
         }
       }
 
-      return filterSupersets(Sets.union(preciseUniqueKeys, 
keysBuilder.build()));
-    } else if (ignoreNulls) {
+      return filterSupersets(Sets.union(preciseUniqueKeys, keysBuilder), 
limit);
+    } else if (ignoreNulls && limit > 0) {
       // group by keys form a unique key
       return ImmutableSet.of(rel.getGroupSet());
     } else {
@@ -367,7 +391,7 @@ public class RelMdUniqueKeys
    * other keys.  Given {@code {0},{1},{1,2}}, returns {@code {0},{1}}.
    */
   private static Set<ImmutableBitSet> filterSupersets(
-      Set<ImmutableBitSet> uniqueKeys) {
+      Set<ImmutableBitSet> uniqueKeys, int limit) {
     Set<ImmutableBitSet> minimalKeys = new HashSet<>();
     outer:
     for (ImmutableBitSet candidateKey : uniqueKeys) {
@@ -377,6 +401,9 @@ public class RelMdUniqueKeys
           continue outer;
         }
       }
+      if (minimalKeys.size() == limit) {
+        break outer;
+      }
       minimalKeys.add(candidateKey);
     }
     return minimalKeys;
@@ -431,7 +458,7 @@ public class RelMdUniqueKeys
 
   public Set<ImmutableBitSet> getUniqueKeys(Union rel, RelMetadataQuery mq,
       boolean ignoreNulls) {
-    if (!rel.all) {
+    if (!rel.all && limit > 0) {
       return ImmutableSet.of(
           ImmutableBitSet.range(rel.getRowType().getFieldCount()));
     }
@@ -443,19 +470,24 @@ public class RelMdUniqueKeys
    */
   public Set<ImmutableBitSet> getUniqueKeys(Intersect rel,
       RelMetadataQuery mq, boolean ignoreNulls) {
-    ImmutableSet.Builder<ImmutableBitSet> keys = new ImmutableSet.Builder<>();
+    Set<ImmutableBitSet> keys = new HashSet<>();
+    outerLoop:
     for (RelNode input : rel.getInputs()) {
       Set<ImmutableBitSet> uniqueKeys = mq.getUniqueKeys(input, ignoreNulls);
       if (uniqueKeys != null) {
-        keys.addAll(uniqueKeys);
+        for (ImmutableBitSet inKey : uniqueKeys) {
+          if (keys.size() == limit) {
+            break outerLoop;
+          }
+          keys.add(inKey);
+        }
       }
     }
-    ImmutableSet<ImmutableBitSet> uniqueKeys = keys.build();
-    if (!uniqueKeys.isEmpty()) {
-      return uniqueKeys;
+    if (!keys.isEmpty()) {
+      return keys;
     }
 
-    if (!rel.all) {
+    if (!rel.all && limit > 0) {
       return ImmutableSet.of(
           ImmutableBitSet.range(rel.getRowType().getFieldCount()));
     }
@@ -472,7 +504,7 @@ public class RelMdUniqueKeys
       return uniqueKeys;
     }
 
-    if (!rel.all) {
+    if (!rel.all && limit > 0) {
       return ImmutableSet.of(
           ImmutableBitSet.range(rel.getRowType().getFieldCount()));
     }
@@ -491,10 +523,15 @@ public class RelMdUniqueKeys
     if (keys == null) {
       return null;
     }
+    Set<ImmutableBitSet> result = new HashSet<>(Math.min(keys.size(), limit));
     for (ImmutableBitSet key : keys) {
+      if (result.size() == limit) {
+        break;
+      }
       assert rel.getTable().isKey(key);
+      result.add(key);
     }
-    return ImmutableSet.copyOf(keys);
+    return result;
   }
 
   public @Nullable Set<ImmutableBitSet> getUniqueKeys(Values rel, 
RelMetadataQuery mq,
@@ -516,10 +553,15 @@ public class RelMdUniqueKeys
     }
 
     ImmutableSet.Builder<ImmutableBitSet> keySetBuilder = 
ImmutableSet.builder();
+    int keySetSize = 0;
     for (int i = 0; i < ranges.size(); i++) {
       final Set<RexLiteral> range = ranges.get(i);
+      if (keySetSize == limit) {
+        break;
+      }
       if (range.size() == tuples.size()) {
         keySetBuilder.add(ImmutableBitSet.of(i));
+        keySetSize++;
       }
     }
     return keySetBuilder.build();
diff --git a/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java 
b/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java
index d654612851..ed754f8e00 100644
--- a/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java
+++ b/core/src/test/java/org/apache/calcite/test/RelMetadataTest.java
@@ -45,6 +45,7 @@ import org.apache.calcite.rel.core.AggregateCall;
 import org.apache.calcite.rel.core.Correlate;
 import org.apache.calcite.rel.core.Exchange;
 import org.apache.calcite.rel.core.Filter;
+import org.apache.calcite.rel.core.Intersect;
 import org.apache.calcite.rel.core.Join;
 import org.apache.calcite.rel.core.JoinRelType;
 import org.apache.calcite.rel.core.Minus;
@@ -77,7 +78,11 @@ import 
org.apache.calcite.rel.metadata.ReflectiveRelMetadataProvider;
 import org.apache.calcite.rel.metadata.RelColumnOrigin;
 import org.apache.calcite.rel.metadata.RelMdCollation;
 import org.apache.calcite.rel.metadata.RelMdColumnUniqueness;
+import org.apache.calcite.rel.metadata.RelMdExplainVisibility;
+import org.apache.calcite.rel.metadata.RelMdMaxRowCount;
 import org.apache.calcite.rel.metadata.RelMdPopulationSize;
+import org.apache.calcite.rel.metadata.RelMdPredicates;
+import org.apache.calcite.rel.metadata.RelMdUniqueKeys;
 import org.apache.calcite.rel.metadata.RelMdUtil;
 import org.apache.calcite.rel.metadata.RelMetadataProvider;
 import org.apache.calcite.rel.metadata.RelMetadataQuery;
@@ -132,6 +137,8 @@ import java.util.List;
 import java.util.Map;
 import java.util.Set;
 import java.util.concurrent.locks.ReentrantLock;
+import java.util.stream.Collectors;
+import java.util.stream.IntStream;
 
 import static com.google.common.collect.ImmutableList.toImmutableList;
 
@@ -191,6 +198,13 @@ public class RelMetadataTest {
    * time. */
   private static final ReentrantLock LOCK = new ReentrantLock();
 
+  private static final SqlTestFactory.CatalogReaderFactory COMPOSITE_FACTORY =
+      (typeFactory, caseSensitive) -> {
+        CompositeKeysCatalogReader catalogReader =
+            new CompositeKeysCatalogReader(typeFactory, false);
+        catalogReader.init();
+        return catalogReader;
+      };
   //~ Methods ----------------------------------------------------------------
 
   /** Creates a fixture. */
@@ -1621,6 +1635,242 @@ public class RelMetadataTest {
         .assertThatUniqueKeysAre(bitSetOf(0, 1));
   }
 
+  @Test void testUniqueKeysWithLimitOnSortOneRow() {
+    sql("select ename, empno from emp order by ename limit 1")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .assertThatRel(is(instanceOf(Sort.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf());
+  }
+
+  @Test void testUniqueKeysWithLimitOnFilter() {
+    sql("select * from s.passenger t1 where t1.age > 35")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Filter.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(1));
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnProjectOverInputWithCompositeKeyAndRepeatedColumns() {
+    String cols = IntStream.range(0, 32).mapToObj(i -> "k" + 
i).collect(Collectors.joining(","));
+    sql("select " + cols + ", " + cols + " from s.composite_keys_32_table")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .assertThatRel(is(instanceOf(Project.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(ImmutableBitSet.range(0, 32),
+            ImmutableBitSet.range(0, 31).set(63));
+  }
+
+  @Test void testUniqueKeysWithLimitOnCrossJoin() {
+    sql("select *\n"
+        + "from s.passenger t1\n"
+        + "cross join s.passenger t2\n")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Join.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(1, 5), bitSetOf(1, 6));
+  }
+
+  @Test void testUniqueKeysWithLimitOnInnerJoinAndConditionOnKeys() {
+    sql("select *\n"
+        + "from s.passenger t1\n"
+        + "inner join s.passenger t2\n"
+        + "   on t1.passport=t2.passport")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Join.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(1), bitSetOf(6));
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnInnerJoinAndConditionOnLeftKeyRightNotKey() {
+    sql("select *\n"
+        + "from s.passenger t1\n"
+        + "inner join s.passenger t2\n"
+        + "   on t1.nid=t2.age")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Join.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(5), bitSetOf(6));
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnInnerJoinAndConditionOnLeftNotKeyRightKey() {
+    sql("select *\n"
+        + "from s.passenger t1\n"
+        + "inner join s.passenger t2\n"
+        + "   on t1.age=t2.nid")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Join.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(1));
+  }
+
+  @Test void testUniqueKeysWithLimitOnInnerJoinAndConditionOnNonKeys() {
+    sql("select *\n"
+        + "from s.passenger t1\n"
+        + "inner join s.passenger t2\n"
+        + "   on t1.fname=t2.fname")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Join.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(1, 5), bitSetOf(1, 6));
+  }
+
+  @Test void testUniqueKeysWithLimitOnSimpleAggregateOverInputWithSimpleKeys() 
{
+    sql("select passport, nid, ssn from s.passenger group by passport, nid, 
ssn")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .assertThatRel(is(instanceOf(Aggregate.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(1));
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnSimpleAggregateOverInputWithSimpleKeysAndPassthroughAggs()
 {
+    sql("select passport, nid, ssn, min(passport), max(passport), min(nid), 
max(nid)\n"
+        + "from s.passenger group by passport, nid, ssn\n")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .assertThatRel(is(instanceOf(Aggregate.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(1));
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnSimpleAggregateOverInputWithCompositeKeyAndPassthroughAggs()
 {
+    StringBuilder cols = new StringBuilder();
+    StringBuilder minCols = new StringBuilder();
+    StringBuilder maxCols = new StringBuilder();
+    for (int i = 0; i < 32; i++) {
+      if (i > 0) {
+        cols.append(',');
+        minCols.append(',');
+        maxCols.append(',');
+      }
+      cols.append("k").append(i);
+      minCols.append("min(k").append(i).append(")");
+      maxCols.append("max(k").append(i).append(")");
+    }
+    sql("select " + cols + ", " + minCols + ", " + maxCols
+        + " from s.composite_keys_32_table group by " + cols)
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Aggregate.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(
+            ImmutableBitSet.range(0, 32),
+            ImmutableBitSet.range(0, 31).set(63));
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnSimpleAggregateOverInputWithKeysNotInGroupBy() {
+    sql("select ename, job from emp group by ename, job")
+        .assertThatRel(is(instanceOf(Aggregate.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0, 1))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre();
+  }
+
+  @Test void 
testUniqueKeysWithLimitOnSimpleAggregateOverInputWithUnknownKeys() {
+    sql("select col1 from s.unknown_keys_table group by col1")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .assertThatRel(is(instanceOf(Aggregate.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre();
+  }
+
+  @Test void testUniqueKeysWithConfOnAggregateWithGroupingSets() {
+    sql("select ename, job from emp group by grouping sets ((ename), (ename, 
job))")
+        .assertThatRel(is(instanceOf(Aggregate.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(true, bitSetOf(0, 1))
+        .assertThatUniqueKeysAre(false);
+  }
+
+  @Test void testUniqueKeysWithLimitOnUnion() {
+    sql("select ename, job, mgr from emp union select ename, job, mgr from 
emp")
+        .assertThatRel(is(instanceOf(Union.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0, 1, 2))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre();
+  }
+
+  @Test void testUniqueKeysWithLimitOnUnionAll() {
+    sql("select ename, job, mgr from emp union all select ename, job, mgr from 
emp")
+        .assertThatRel(is(instanceOf(Union.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre();
+  }
+
+  @Test void testUniqueKeysWithLimitOnIntersect() {
+    sql("select empno, deptno from emp intersect select 100, deptno from dept")
+        .assertThatRel(is(instanceOf(Intersect.class)))
+        .withMetadataConfig(uniqueKeyConfig(1))
+        .assertThatUniqueKeysAre(bitSetOf(0));
+  }
+
+  @Test void testUniqueKeysWithLimitOnIntersectWhereInputKeysAreEmpty() {
+    sql("select ename, job, mgr from emp intersect select ename, job, mgr from 
emp")
+        .assertThatRel(is(instanceOf(Intersect.class)))
+        .withMetadataConfig(uniqueKeyConfig(0))
+        .assertThatUniqueKeysAre()
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0, 1, 2));
+  }
+
+
+  @Test void testUniqueKeysWithLimitOnIntersectAllWhereInputsKeysAreEmpty() {
+    sql("select ename, job, mgr from emp intersect all select ename, job, mgr 
from emp")
+        .assertThatRel(is(instanceOf(Intersect.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre();
+  }
+
+  @Test void testUniqueKeysWithLimitOnExceptWhereLeftInputHasKeys() {
+    sql("select * from s.passenger except select 1111, 2222, 3333, 'Rob', 40")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .assertThatRel(is(instanceOf(Minus.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(1));
+  }
+
+  @Test void testUniqueKeysWithLimitOnScan() {
+    sql("select * from s.passenger")
+        .withCatalogReaderFactory(COMPOSITE_FACTORY)
+        .withRelTransform(r -> r.getInput(0))
+        .assertThatRel(is(instanceOf(TableScan.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(1));
+  }
+
+  @Test void testUniqueKeysWithLimitOnValues() {
+    sql("select * from (values\n"
+        + "('X133345', 'Zimmer', 'Bob', '13-10-2022'),\n"
+        + "('Y223455', 'Zimmer', 'Alice', '22-11-2024'))\n")
+        .withRelTransform(project -> project.getInput(0))
+        .assertThatRel(is(instanceOf(Values.class)))
+        .withMetadataConfig(uniqueKeyConfig(2))
+        .assertThatUniqueKeysAre(bitSetOf(0), bitSetOf(2));
+  }
+
   private static ImmutableBitSet bitSetOf(int... bits) {
     return ImmutableBitSet.of(bits);
   }
@@ -3990,6 +4240,21 @@ public class RelMetadataTest {
                 + "true. join=" + join);
   }
 
+  private static RelMetadataFixture.MetadataConfig uniqueKeyConfig(int limit) {
+    ImmutableList.Builder<RelMetadataProvider> providers = 
ImmutableList.builder();
+    providers.add(
+        ReflectiveRelMetadataProvider.reflectiveSource(new 
RelMdUniqueKeys(limit),
+            BuiltInMetadata.UniqueKeys.Handler.class));
+    // The RelMdUniqueKeys handler relies on the following providers
+    providers.add(RelMdColumnUniqueness.SOURCE);
+    providers.add(RelMdPredicates.SOURCE);
+    providers.add(RelMdMaxRowCount.SOURCE);
+    // The visibility provider is needed for printing plans in tests
+    providers.add(RelMdExplainVisibility.SOURCE);
+    return new RelMetadataFixture.MetadataConfig("UQ", 
JaninoRelMetadataProvider::of,
+        () -> new ChainedRelMetadataProvider(providers.build()) {
+        }, false);
+  }
   //~ Inner classes and interfaces -------------------------------------------
 
   /** Custom metadata interface. */
@@ -4117,6 +4382,29 @@ public class RelMetadataTest {
       addDistinctRowcountHandler(t1);
       addUniqueKeyHandler(t1);
       registerTable(t1);
+      MockTable t2 = MockTable.create(this, tSchema, 
"composite_keys_32_table", false, 22.0, null);
+      for (int i = 0; i < 32; i++) {
+        t2.addColumn("k" + i, typeFactory.createSqlType(SqlTypeName.INTEGER));
+      }
+      t2.addKey(ImmutableBitSet.range(0, 32));
+      registerTable(t2);
+      MockTable t3 = MockTable.create(this, tSchema, "passenger", false, 10.0, 
null);
+      t3.addColumn("passport", typeFactory.createSqlType(SqlTypeName.INTEGER), 
true);
+      t3.addColumn("nid", typeFactory.createSqlType(SqlTypeName.INTEGER), 
true);
+      t3.addColumn("ssn", typeFactory.createSqlType(SqlTypeName.INTEGER), 
true);
+      t3.addColumn("fname", typeFactory.createSqlType(SqlTypeName.VARCHAR));
+      t3.addColumn("age", typeFactory.createSqlType(SqlTypeName.INTEGER));
+      registerTable(t3);
+      MockTable t4 = MockTable.create(this, tSchema, "unknown_keys_table", 
false, 15.0, null);
+      t4.addColumn("col1", typeFactory.createSqlType(SqlTypeName.INTEGER));
+      t4.addColumn("col2", typeFactory.createSqlType(SqlTypeName.INTEGER));
+      t4.addWrap(new BuiltInMetadata.UniqueKeys.Handler() {
+        @Override public @Nullable Set<ImmutableBitSet> getUniqueKeys(RelNode 
r,
+            RelMetadataQuery mq, boolean ignoreNulls) {
+          return null;
+        }
+      });
+      registerTable(t4);
       return this;
     }
 
diff --git a/site/_docs/history.md b/site/_docs/history.md
index 67a9d44adc..feaa06069d 100644
--- a/site/_docs/history.md
+++ b/site/_docs/history.md
@@ -57,6 +57,17 @@ this behavior.  The BIG_QUERY and SQL_SERVER_2008 
conformance have
 been changed to use checked arithmetic, matching the specification of
 these dialects.
 
+* [<a 
href="https://issues.apache.org/jira/browse/CALCITE-6704";>CALCITE-6704</a>]
+Limit result size of `RelMdUniqueKeys` handler. Certain query patterns can lead
+to an exponentially large number of unique keys that can cause crashes and OOM
+errors. To prevent this kind of issues the `RelMdUniqueKeys` handler is now 
using
+a limit to restrict the number of keys for each relational expression. The 
limit
+is set to `1000` by default. The value is reasonably large to ensure that
+most common use-cases will not be affected and at the same time bounds 
exponentially
+large results set to a manageable value. Users that need a bigger/smaller limit
+should create a new instance of `RelMdUniqueKeys` and register it using the
+metadata provider of their choice.
+
 #### New features
 {: #new-features-1-39-0}
 
diff --git 
a/testkit/src/main/java/org/apache/calcite/test/RelMetadataFixture.java 
b/testkit/src/main/java/org/apache/calcite/test/RelMetadataFixture.java
index 3dff044762..f51104bc4d 100644
--- a/testkit/src/main/java/org/apache/calcite/test/RelMetadataFixture.java
+++ b/testkit/src/main/java/org/apache/calcite/test/RelMetadataFixture.java
@@ -343,19 +343,24 @@ public class RelMetadataFixture {
     });
   }
 
+  @SuppressWarnings({"UnusedReturnValue"})
+  public RelMetadataFixture assertThatUniqueKeysAre(ImmutableBitSet... 
expectedUniqueKeys) {
+    return assertThatUniqueKeysAre(false, expectedUniqueKeys);
+  }
+
   /** Checks result of getting unique keys for SQL. */
   @SuppressWarnings({"UnusedReturnValue"})
-  public RelMetadataFixture assertThatUniqueKeysAre(
+  public RelMetadataFixture assertThatUniqueKeysAre(boolean ignoreNulls,
       ImmutableBitSet... expectedUniqueKeys) {
     RelNode rel = toRel();
     final RelMetadataQuery mq = rel.getCluster().getMetadataQuery();
-    Set<ImmutableBitSet> result = mq.getUniqueKeys(rel);
+    Set<ImmutableBitSet> result = mq.getUniqueKeys(rel, ignoreNulls);
     assertThat(result, notNullValue());
     assertThat("unique keys, sql: " + relSupplier
             + ", rel: " + RelOptUtil.toString(rel),
         ImmutableSortedSet.copyOf(result),
         is(ImmutableSortedSet.copyOf(expectedUniqueKeys)));
-    checkUniqueConsistent(rel);
+    checkUniqueConsistent(rel, ignoreNulls);
     return this;
   }
 
@@ -364,14 +369,12 @@ public class RelMetadataFixture {
    * and {@link RelMetadataQuery#areColumnsUnique(RelNode, ImmutableBitSet)}
    * return consistent results.
    */
-  private static void checkUniqueConsistent(RelNode rel) {
+  private static void checkUniqueConsistent(RelNode rel, boolean ignoreNulls) {
     final RelMetadataQuery mq = rel.getCluster().getMetadataQuery();
-    final Set<ImmutableBitSet> uniqueKeys = mq.getUniqueKeys(rel);
+    final Set<ImmutableBitSet> uniqueKeys = mq.getUniqueKeys(rel, ignoreNulls);
     assertThat(uniqueKeys, notNullValue());
-    final ImmutableBitSet allCols =
-        ImmutableBitSet.range(0, rel.getRowType().getFieldCount());
-    for (ImmutableBitSet key : allCols.powerSet()) {
-      Boolean result2 = mq.areColumnsUnique(rel, key);
+    for (ImmutableBitSet key : uniqueKeys) {
+      Boolean result2 = mq.areColumnsUnique(rel, key, ignoreNulls);
       assertThat("areColumnsUnique. key: " + key
           + ", uniqueKeys: " + uniqueKeys
           + ", rel: " + RelOptUtil.toString(rel),

Reply via email to