Skip to content

Commit 11e582b

Browse files
authored
Spark: Backport aggregate pushdown tests with NaN's (#16316)
1 parent e57247b commit 11e582b

3 files changed

Lines changed: 135 additions & 0 deletions

File tree

spark/v3.4/spark/src/test/java/org/apache/iceberg/spark/sql/TestAggregatePushDown.java

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -767,6 +767,51 @@ public void testNaN() {
767767
assertEquals("expected and actual should equal", expected, actual);
768768
}
769769

770+
@TestTemplate
771+
public void testNanWithLowerAndUpperBoundMetrics() {
772+
sql("CREATE TABLE %s (id int, data float) USING iceberg PARTITIONED BY (id)", tableName);
773+
sql(
774+
"INSERT INTO %s VALUES (1, float('nan')),"
775+
+ "(1, float('nan')), "
776+
+ "(1, 10.0), "
777+
+ "(2, 2), "
778+
+ "(2, float('nan')), "
779+
+ "(3, float('nan')), "
780+
+ "(3, 1)",
781+
tableName);
782+
783+
// Validate all files has upper bound, lower bound and nan count
784+
String countsQuery =
785+
"select readable_metrics.data.nan_value_count > 0, "
786+
+ "isnull(readable_metrics.data.lower_bound), "
787+
+ "isnull(readable_metrics.data.upper_bound) "
788+
+ "from %s.files";
789+
790+
Object[] expectedResult = new Object[] {true, false, false};
791+
assertThat(sql(countsQuery, tableName))
792+
.as("Data files should contain nan count, lower bound and upper bound.")
793+
.allMatch(row -> Arrays.equals(row, expectedResult));
794+
795+
// Check aggregates are not pushed down
796+
String select = "SELECT count(*), max(data), min(data), count(data) FROM %s";
797+
798+
List<Object[]> explain = sql("EXPLAIN " + select, tableName);
799+
String explainString = explain.get(0)[0].toString().toLowerCase(Locale.ROOT);
800+
boolean explainContainsPushDownAggregates =
801+
(explainString.contains("max(data)")
802+
|| explainString.contains("min(data)")
803+
|| explainString.contains("count(data)"));
804+
805+
assertThat(explainContainsPushDownAggregates)
806+
.as("explain should not contain the pushed down aggregates")
807+
.isFalse();
808+
809+
List<Object[]> actual = sql(select, tableName);
810+
List<Object[]> expected = Lists.newArrayList();
811+
expected.add(new Object[] {7L, Float.NaN, 1.0F, 7L});
812+
assertEquals("expected and actual should equal", expected, actual);
813+
}
814+
770815
@TestTemplate
771816
public void testInfinity() {
772817
sql(

spark/v3.5/spark/src/test/java/org/apache/iceberg/spark/sql/TestAggregatePushDown.java

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -767,6 +767,51 @@ public void testNaN() {
767767
assertEquals("expected and actual should equal", expected, actual);
768768
}
769769

770+
@TestTemplate
771+
public void testNanWithLowerAndUpperBoundMetrics() {
772+
sql("CREATE TABLE %s (id int, data float) USING iceberg PARTITIONED BY (id)", tableName);
773+
sql(
774+
"INSERT INTO %s VALUES (1, float('nan')),"
775+
+ "(1, float('nan')), "
776+
+ "(1, 10.0), "
777+
+ "(2, 2), "
778+
+ "(2, float('nan')), "
779+
+ "(3, float('nan')), "
780+
+ "(3, 1)",
781+
tableName);
782+
783+
// Validate all files has upper bound, lower bound and nan count
784+
String countsQuery =
785+
"select readable_metrics.data.nan_value_count > 0, "
786+
+ "isnull(readable_metrics.data.lower_bound), "
787+
+ "isnull(readable_metrics.data.upper_bound) "
788+
+ "from %s.files";
789+
790+
Object[] expectedResult = new Object[] {true, false, false};
791+
assertThat(sql(countsQuery, tableName))
792+
.as("Data files should contain nan count, lower bound and upper bound.")
793+
.allMatch(row -> Arrays.equals(row, expectedResult));
794+
795+
// Check aggregates are not pushed down
796+
String select = "SELECT count(*), max(data), min(data), count(data) FROM %s";
797+
798+
List<Object[]> explain = sql("EXPLAIN " + select, tableName);
799+
String explainString = explain.get(0)[0].toString().toLowerCase(Locale.ROOT);
800+
boolean explainContainsPushDownAggregates =
801+
(explainString.contains("max(data)")
802+
|| explainString.contains("min(data)")
803+
|| explainString.contains("count(data)"));
804+
805+
assertThat(explainContainsPushDownAggregates)
806+
.as("explain should not contain the pushed down aggregates")
807+
.isFalse();
808+
809+
List<Object[]> actual = sql(select, tableName);
810+
List<Object[]> expected = Lists.newArrayList();
811+
expected.add(new Object[] {7L, Float.NaN, 1.0F, 7L});
812+
assertEquals("expected and actual should equal", expected, actual);
813+
}
814+
770815
@TestTemplate
771816
public void testInfinity() {
772817
sql(

spark/v4.0/spark/src/test/java/org/apache/iceberg/spark/sql/TestAggregatePushDown.java

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -767,6 +767,51 @@ public void testNaN() {
767767
assertEquals("expected and actual should equal", expected, actual);
768768
}
769769

770+
@TestTemplate
771+
public void testNanWithLowerAndUpperBoundMetrics() {
772+
sql("CREATE TABLE %s (id int, data float) USING iceberg PARTITIONED BY (id)", tableName);
773+
sql(
774+
"INSERT INTO %s VALUES (1, float('nan')),"
775+
+ "(1, float('nan')), "
776+
+ "(1, 10.0), "
777+
+ "(2, 2), "
778+
+ "(2, float('nan')), "
779+
+ "(3, float('nan')), "
780+
+ "(3, 1)",
781+
tableName);
782+
783+
// Validate all files has upper bound, lower bound and nan count
784+
String countsQuery =
785+
"select readable_metrics.data.nan_value_count > 0, "
786+
+ "isnull(readable_metrics.data.lower_bound), "
787+
+ "isnull(readable_metrics.data.upper_bound) "
788+
+ "from %s.files";
789+
790+
Object[] expectedResult = new Object[] {true, false, false};
791+
assertThat(sql(countsQuery, tableName))
792+
.as("Data files should contain nan count, lower bound and upper bound.")
793+
.allMatch(row -> Arrays.equals(row, expectedResult));
794+
795+
// Check aggregates are not pushed down
796+
String select = "SELECT count(*), max(data), min(data), count(data) FROM %s";
797+
798+
List<Object[]> explain = sql("EXPLAIN " + select, tableName);
799+
String explainString = explain.get(0)[0].toString().toLowerCase(Locale.ROOT);
800+
boolean explainContainsPushDownAggregates =
801+
(explainString.contains("max(data)")
802+
|| explainString.contains("min(data)")
803+
|| explainString.contains("count(data)"));
804+
805+
assertThat(explainContainsPushDownAggregates)
806+
.as("explain should not contain the pushed down aggregates")
807+
.isFalse();
808+
809+
List<Object[]> actual = sql(select, tableName);
810+
List<Object[]> expected = Lists.newArrayList();
811+
expected.add(new Object[] {7L, Float.NaN, 1.0F, 7L});
812+
assertEquals("expected and actual should equal", expected, actual);
813+
}
814+
770815
@TestTemplate
771816
public void testInfinity() {
772817
sql(

0 commit comments

Comments
 (0)