This is an automated email from the ASF dual-hosted git repository.
lincoln pushed a commit to branch release-2.0
in repository https://gitbox.apache.org/repos/asf/flink.git
The following commit(s) were added to refs/heads/release-2.0 by this push:
new c3e0acab15d [FLINK-30687][table] Fix wrong result of agg with fiter
which references first input column
c3e0acab15d is described below
commit c3e0acab15dd6610404de93775671fc108915224
Author: lincoln lee <[email protected]>
AuthorDate: Mon Mar 17 19:57:10 2025 +0800
[FLINK-30687][table] Fix wrong result of agg with fiter which references
first input column
This closes #26302.
---
.../codegen/agg/AggsHandlerCodeGenerator.scala | 2 +-
.../table/planner/codegen/agg/AggTestBase.scala | 9 +++++++++
.../runtime/batch/sql/agg/AggregateITCaseBase.scala | 5 +++++
.../runtime/stream/sql/AggregateITCase.scala | 21 ++++++++++++++++++++-
4 files changed, 35 insertions(+), 2 deletions(-)
diff --git
a/flink-table/flink-table-planner/src/main/scala/org/apache/flink/table/planner/codegen/agg/AggsHandlerCodeGenerator.scala
b/flink-table/flink-table-planner/src/main/scala/org/apache/flink/table/planner/codegen/agg/AggsHandlerCodeGenerator.scala
index c9c7ee5d500..deec91f894a 100644
---
a/flink-table/flink-table-planner/src/main/scala/org/apache/flink/table/planner/codegen/agg/AggsHandlerCodeGenerator.scala
+++
b/flink-table/flink-table-planner/src/main/scala/org/apache/flink/table/planner/codegen/agg/AggsHandlerCodeGenerator.scala
@@ -330,7 +330,7 @@ class AggsHandlerCodeGenerator(
aggIndex: Int,
aggName: String): Option[Expression] = {
- if (filterArg > 0) {
+ if (filterArg >= 0) {
val filterType = inputFieldTypes(filterArg)
if (!filterType.isInstanceOf[BooleanType]) {
throw new TableException(
diff --git
a/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/codegen/agg/AggTestBase.scala
b/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/codegen/agg/AggTestBase.scala
index b6c760a4cca..f178deb3d2f 100644
---
a/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/codegen/agg/AggTestBase.scala
+++
b/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/codegen/agg/AggTestBase.scala
@@ -66,6 +66,7 @@ abstract class AggTestBase(isBatchMode: Boolean) {
val aggInfo1: AggregateInfo = {
val aggInfo = mock(classOf[AggregateInfo])
val call = mock(classOf[AggregateCall])
+ updateFilter(call, -1)
when(aggInfo.agg).thenReturn(call)
when(call.getName).thenReturn("avg1")
when(call.hasFilter).thenReturn(false)
@@ -81,6 +82,7 @@ abstract class AggTestBase(isBatchMode: Boolean) {
val aggInfo2: AggregateInfo = {
val aggInfo = mock(classOf[AggregateInfo])
val call = mock(classOf[AggregateCall])
+ updateFilter(call, -1)
when(aggInfo.agg).thenReturn(call)
when(call.getName).thenReturn("avg2")
when(call.hasFilter).thenReturn(false)
@@ -97,6 +99,7 @@ abstract class AggTestBase(isBatchMode: Boolean) {
val aggInfo3: AggregateInfo = {
val aggInfo = mock(classOf[AggregateInfo])
val call = mock(classOf[AggregateCall])
+ updateFilter(call, -1)
when(aggInfo.agg).thenReturn(call)
when(call.getName).thenReturn("avg3")
when(call.hasFilter).thenReturn(false)
@@ -117,4 +120,10 @@ abstract class AggTestBase(isBatchMode: Boolean) {
val classLoader: ClassLoader = Thread.currentThread().getContextClassLoader
val context: ExecutionContext = mock(classOf[ExecutionContext])
when(context.getRuntimeContext).thenReturn(mock(classOf[RuntimeContext]))
+
+ private def updateFilter(call: AggregateCall, v: Int): Unit = {
+ val field = call.getClass.getField("filterArg")
+ field.setAccessible(true)
+ field.set(call, v)
+ }
}
diff --git
a/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/batch/sql/agg/AggregateITCaseBase.scala
b/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/batch/sql/agg/AggregateITCaseBase.scala
index 2e9cf7de432..206d25f86e9 100644
---
a/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/batch/sql/agg/AggregateITCaseBase.scala
+++
b/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/batch/sql/agg/AggregateITCaseBase.scala
@@ -1224,6 +1224,11 @@ abstract class AggregateITCaseBase(testName: String)
extends BatchTestBase {
checkResult(sql, Seq(row("11, 11"), row("12, 12"), row("null, null")))
}
+ @Test
+ def testAggFilterReferenceFirstColumn(): Unit = {
+ checkResult("select count(*) filter (where a < 10) from Table3",
Seq(row(9)))
+ }
+
// TODO support csv
// @Test
// def testMultiGroupBys(): Unit = {
diff --git
a/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/stream/sql/AggregateITCase.scala
b/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/stream/sql/AggregateITCase.scala
index b0d1cc09317..2ae87808ce9 100644
---
a/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/stream/sql/AggregateITCase.scala
+++
b/flink-table/flink-table-planner/src/test/scala/org/apache/flink/table/planner/runtime/stream/sql/AggregateITCase.scala
@@ -22,7 +22,6 @@ import org.apache.flink.streaming.api.datastream.DataStream
import org.apache.flink.table.api._
import org.apache.flink.table.api.bridge.scala._
import org.apache.flink.table.api.config.ExecutionConfigOptions
-import org.apache.flink.table.api.internal.TableEnvironmentInternal
import org.apache.flink.table.connector.ChangelogMode
import org.apache.flink.table.legacy.api.Types
import org.apache.flink.table.planner.factories.TestValuesTableFactory
@@ -1659,6 +1658,26 @@ class AggregateITCase(
assertThat(sink.getRetractResults.sorted).isEqualTo(expected.sorted)
}
+ @TestTemplate
+ def testAggFilterReferenceFirstColumn(): Unit = {
+ val t = failingDataSource(TestData.tupleData3).toTable(tEnv).as("a", "b",
"c")
+ tEnv.createTemporaryView("MyTable", t)
+
+ val sqlQuery =
+ s"""
+ |SELECT
+ | COUNT(*) filter (where a < 10)
+ |FROM MyTable
+ """.stripMargin
+
+ val sink = new TestingRetractSink
+ val result = tEnv.sqlQuery(sqlQuery).toRetractStream[Row]
+ result.addSink(sink).setParallelism(1)
+ env.execute()
+ val expected = List("9")
+ assertThat(sink.getRetractResults.sorted).isEqualTo(expected.sorted)
+ }
+
@TestTemplate
def testPruneUselessAggCall(): Unit = {
val data = new mutable.MutableList[(Int, Long, String)]