diff --git a/linq4j/src/main/java/org/apache/calcite/linq4j/Linq4j.java b/linq4j/src/main/java/org/apache/calcite/linq4j/Linq4j.java index df9e39e35bd2..c59303d4fdc0 100644 --- a/linq4j/src/main/java/org/apache/calcite/linq4j/Linq4j.java +++ b/linq4j/src/main/java/org/apache/calcite/linq4j/Linq4j.java @@ -587,10 +587,14 @@ static class ListEnumerable extends CollectionEnumerable { @Override public Enumerable skip(int count) { final List list = toList(); - if (count >= list.size()) { + // Clamp to zero, as the BigDecimal overload does, so that a negative + // count skips nothing and matches EnumerableDefaults.skip rather than + // throwing from List.subList. + final int rows = Math.max(count, 0); + if (rows >= list.size()) { return Linq4j.emptyEnumerable(); } - return new ListEnumerable<>(list.subList(count, list.size())); + return new ListEnumerable<>(list.subList(rows, list.size())); } @Override public Enumerable skip(BigDecimal count) { @@ -605,10 +609,14 @@ static class ListEnumerable extends CollectionEnumerable { @Override public Enumerable take(int count) { final List list = toList(); - if (count >= list.size()) { + // Clamp to zero, as the BigDecimal overload does, so that a negative + // count yields an empty enumerable and matches EnumerableDefaults.take + // rather than throwing from List.subList. + final int rows = Math.max(count, 0); + if (rows >= list.size()) { return this; } - return new ListEnumerable<>(list.subList(0, count)); + return new ListEnumerable<>(list.subList(0, rows)); } @Override public Enumerable take(BigDecimal count) { diff --git a/linq4j/src/test/java/org/apache/calcite/linq4j/test/Linq4jTest.java b/linq4j/src/test/java/org/apache/calcite/linq4j/test/Linq4jTest.java index 62e7dd9f5eef..cd894f964b83 100644 --- a/linq4j/src/test/java/org/apache/calcite/linq4j/test/Linq4jTest.java +++ b/linq4j/src/test/java/org/apache/calcite/linq4j/test/Linq4jTest.java @@ -2238,4 +2238,20 @@ public String toString() { new Department("HR", 20, ImmutableList.of()), new Department("Marketing", 30, ImmutableList.of(emps[1])), }; + + @Test void testTakeListEnumerableNegativeSize() { + final List values = Arrays.asList(1, 2, 3); + + assertThat(EnumerableDefaults.take(Linq4j.asEnumerable(values), -1).toList(), + is(empty())); + assertThat(Linq4j.asEnumerable(values).take(-1).toList(), is(empty())); + } + + @Test void testSkipListEnumerableNegativeSize() { + final List values = Arrays.asList(1, 2, 3); + + assertThat(EnumerableDefaults.skip(Linq4j.asEnumerable(values), -1).toList(), + is(values)); + assertThat(Linq4j.asEnumerable(values).skip(-1).toList(), is(values)); + } }