diff --git a/sql/core/src/test/scala/org/apache/spark/sql/jdbc/v2/JDBCV2JoinPushdownIntegrationSuiteBase.scala b/sql/core/src/test/scala/org/apache/spark/sql/jdbc/v2/JDBCV2JoinPushdownIntegrationSuiteBase.scala index 4c46f5587a3f1..3b36a3e73de21 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/jdbc/v2/JDBCV2JoinPushdownIntegrationSuiteBase.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/jdbc/v2/JDBCV2JoinPushdownIntegrationSuiteBase.scala @@ -229,8 +229,6 @@ trait JDBCV2JoinPushdownIntegrationSuiteBase protected val supportsColumnPruning: Boolean = true - protected val supportsJoinPushdown: Boolean = true - // Condition-less joins are not supported in join pushdown test("Test that 2-way join without condition should not have join pushed down") { val sqlQuery = @@ -497,6 +495,55 @@ trait JDBCV2JoinPushdownIntegrationSuiteBase } } + test("Test aggregate with group by on top of join") { + val sqlQuery = + s""" + |SELECT t1.id, t1.address, min(t2.salary), count(1) + |FROM $catalogAndNamespace.$casedJoinTableName1 t1 + |JOIN $catalogAndNamespace.$casedJoinTableName2 t2 ON t1.id = t2.id + |WHERE t1.amount > 1000 + |GROUP BY t1.id, t1.address + |""".stripMargin + + val rowsNoPushdown = withSQLConf(SQLConf.DATA_SOURCE_V2_JOIN_PUSHDOWN.key -> "false") { + sql(sqlQuery).collect().toSeq + } + + assert(rowsNoPushdown.nonEmpty) + + withSQLConf(SQLConf.DATA_SOURCE_V2_JOIN_PUSHDOWN.key -> "true") { + val df = sql(sqlQuery) + checkJoinPushed(df) + checkAggregateRemoved(df, supportsAggregatePushdown) + checkAnswer(df, rowsNoPushdown) + } + } + + test("Test multi-way join with function in join condition") { + val sqlQuery = + s""" + |SELECT a.id, c.address, d.address, a.amount + |FROM $catalogAndNamespace.$casedJoinTableName1 a + |JOIN $catalogAndNamespace.$casedJoinTableName1 c + | ON a.address = c.address + |JOIN $catalogAndNamespace.$casedJoinTableName1 d + | ON LOWER(a.address) = LOWER(d.address) + |WHERE a.amount >= 1000 + |""".stripMargin + + val rowsNoPushdown = withSQLConf(SQLConf.DATA_SOURCE_V2_JOIN_PUSHDOWN.key -> "false") { + sql(sqlQuery).collect().toSeq + } + + assert(rowsNoPushdown.nonEmpty) + + withSQLConf(SQLConf.DATA_SOURCE_V2_JOIN_PUSHDOWN.key -> "true") { + val df = sql(sqlQuery) + checkJoinPushed(df) + checkAnswer(df, rowsNoPushdown) + } + } + test("Test sort limit on top of join is pushed down") { val sqlQuery = s""" |SELECT min(a.id + b.id), a.id, b.id