diff --git a/.github/workflows/license.yml b/.github/workflows/license.yml
index 9166bc30434..fb74c19d172 100644
--- a/.github/workflows/license.yml
+++ b/.github/workflows/license.yml
@@ -44,7 +44,7 @@ jobs:
- run: |
build/mvn org.apache.rat:apache-rat-plugin:check \
-Ptpcds -Pkubernetes-it \
- -Pspark-3.3 -Pspark-3.4 -Pspark-3.5 -Pspark-4.0 -Pspark-4.1
+ -Pspark-3.3 -Pspark-3.4 -Pspark-3.5 -Pspark-4.0 -Pspark-4.1 -Pspark-4.2
- name: Upload rat report
if: failure()
uses: actions/upload-artifact@v7
diff --git a/.github/workflows/master.yml b/.github/workflows/master.yml
index bc461e26786..e3c51530cc2 100644
--- a/.github/workflows/master.yml
+++ b/.github/workflows/master.yml
@@ -74,6 +74,12 @@ jobs:
spark-archive: '-Pscala-2.13'
exclude-tags: ''
comment: 'normal'
+ - java: 21
+ python: '3.11'
+ spark: '4.2'
+ spark-archive: '-Pscala-2.13'
+ exclude-tags: '-Dmaven.plugin.scalatest.exclude.tags=org.scalatest.tags.Slow,org.apache.kyuubi.tags.DeltaTest,org.apache.kyuubi.tags.IcebergTest,org.apache.kyuubi.tags.PaimonTest,org.apache.kyuubi.tags.HudiTest'
+ comment: 'normal'
- java: 8
python: '3.9'
spark: '3.5'
@@ -98,6 +104,12 @@ jobs:
spark-archive: '-Pscala-2.13 -Dspark.archive.mirror=https://www.apache.org/dyn/closer.lua/spark/spark-4.1.2 -Dspark.archive.name=spark-4.1.2-bin-hadoop3.tgz'
exclude-tags: '-Dmaven.plugin.scalatest.exclude.tags=org.scalatest.tags.Slow,org.apache.kyuubi.tags.DeltaTest,org.apache.kyuubi.tags.IcebergTest,org.apache.kyuubi.tags.PaimonTest,org.apache.kyuubi.tags.HudiTest,org.apache.kyuubi.tags.SparkLocalClusterTest'
comment: 'verify-on-spark-4.1-binary'
+ - java: 17
+ python: '3.11'
+ spark: '3.5'
+ spark-archive: '-Pscala-2.13 -Dspark.archive.mirror=https://www.apache.org/dyn/closer.lua/spark/spark-4.2.0 -Dspark.archive.name=spark-4.2.0-bin-hadoop3.tgz'
+ exclude-tags: '-Dmaven.plugin.scalatest.exclude.tags=org.scalatest.tags.Slow,org.apache.kyuubi.tags.DeltaTest,org.apache.kyuubi.tags.IcebergTest,org.apache.kyuubi.tags.PaimonTest,org.apache.kyuubi.tags.HudiTest,org.apache.kyuubi.tags.SparkLocalClusterTest'
+ comment: 'verify-on-spark-4.2-binary'
env:
SPARK_LOCAL_IP: localhost
steps:
diff --git a/.github/workflows/style.yml b/.github/workflows/style.yml
index 2d1bf3d7a89..d60367b7296 100644
--- a/.github/workflows/style.yml
+++ b/.github/workflows/style.yml
@@ -71,6 +71,7 @@ jobs:
build/mvn clean install -pl extensions/spark/kyuubi-extension-spark-3-5,extensions/spark/kyuubi-spark-connector-hive -am -Pspark-3.5
build/mvn clean install -pl extensions/spark/kyuubi-extension-spark-4-0 -am -Pspark-4.0 -Pscala-2.13
build/mvn clean install -pl extensions/spark/kyuubi-extension-spark-4-1 -am -Pspark-4.1 -Pscala-2.13
+ build/mvn clean install -pl extensions/spark/kyuubi-extension-spark-4-2 -am -Pspark-4.2 -Pscala-2.13
- name: Scalastyle with maven
id: scalastyle-check
diff --git a/build/release/release.sh b/build/release/release.sh
index 31f96f458d7..09b5a18b9e1 100755
--- a/build/release/release.sh
+++ b/build/release/release.sh
@@ -130,6 +130,11 @@ upload_nexus_staging() {
-s "${KYUUBI_DIR}/build/release/asf-settings.xml" \
-pl extensions/spark/kyuubi-extension-spark-4-1 -am
+ # Spark Extension Plugin for Spark 4.2 and Scala 2.13
+ ${KYUUBI_DIR}/build/mvn clean deploy -DskipTests -Papache-release,flink-provided,spark-provided,hive-provided,spark-4.2,scala-2.13 \
+ -s "${KYUUBI_DIR}/build/release/asf-settings.xml" \
+ -pl extensions/spark/kyuubi-extension-spark-4-2 -am
+
# Spark Hive/TPC-DS/TPC-H Connector built with default Spark version (3.5) and Scala 2.13
${KYUUBI_DIR}/build/mvn clean deploy -DskipTests -Papache-release,flink-provided,spark-provided,hive-provided,spark-3.5,scala-2.13 \
-s "${KYUUBI_DIR}/build/release/asf-settings.xml" \
diff --git a/dev/kyuubi-codecov/pom.xml b/dev/kyuubi-codecov/pom.xml
index 7e912f36fe3..2f65cf73d47 100644
--- a/dev/kyuubi-codecov/pom.xml
+++ b/dev/kyuubi-codecov/pom.xml
@@ -280,6 +280,31 @@
+
+ spark-4.2
+
+
+ org.apache.kyuubi
+ kyuubi-extension-spark-4-2_${scala.binary.version}
+ ${project.version}
+
+
+ org.apache.kyuubi
+ kyuubi-spark-connector-hive_${scala.binary.version}
+ ${project.version}
+
+
+ org.apache.kyuubi
+ kyuubi-spark-authz_${scala.binary.version}
+ ${project.version}
+
+
+ org.apache.kyuubi
+ kyuubi-spark-lineage_${scala.binary.version}
+ ${project.version}
+
+
+
codecov
diff --git a/dev/reformat b/dev/reformat
index a2deea54a3a..28ce515f5f1 100755
--- a/dev/reformat
+++ b/dev/reformat
@@ -20,7 +20,7 @@ set -x
KYUUBI_HOME="$(cd "`dirname "$0"`/.."; pwd)"
-PROFILES="-Pflink-provided,hive-provided,spark-provided,spark-4.1,spark-4.0,spark-3.5,spark-3.4,spark-3.3,tpcds,kubernetes-it"
+PROFILES="-Pflink-provided,hive-provided,spark-provided,spark-4.2,spark-4.1,spark-4.0,spark-3.5,spark-3.4,spark-3.3,tpcds,kubernetes-it"
# python style checks rely on `black` in path
if ! command -v black &> /dev/null
diff --git a/docs/extensions/engines/spark/rules.md b/docs/extensions/engines/spark/rules.md
index 3999336c9d7..4b8c52ebe59 100644
--- a/docs/extensions/engines/spark/rules.md
+++ b/docs/extensions/engines/spark/rules.md
@@ -64,6 +64,9 @@ And don't worry, Kyuubi will support the new Apache Spark version in the future.
| kyuubi-extension-spark-3-3 | 3.3.x | 1.6.0-incubating | N/A | 1.6.0-incubating | spark-3.3 |
| kyuubi-extension-spark-3-4 | 3.4.x | 1.8.0 | N/A | 1.8.0 | spark-3.4 |
| kyuubi-extension-spark-3-5 | 3.5.x | 1.8.0 | N/A | 1.9.0 | spark-3.5 |
+| kyuubi-extension-spark-4-0 | 4.0.x | 1.10.0 | N/A | 1.10.0 | spark-4.0 |
+| kyuubi-extension-spark-4-1 | 4.1.x | 1.11.0 | N/A | 1.11.0 | spark-4.1 |
+| kyuubi-extension-spark-4-2 | 4.2.x | 1.12.0 | N/A | 1.12.0 | spark-4.2 |
1. Check the matrix that if you are using the supported Spark version, and find the corresponding Kyuubi Spark SQL Extension jar
2. Get the Kyuubi Spark SQL Extension jar
diff --git a/docs/quick_start/quick_start.rst b/docs/quick_start/quick_start.rst
index d09659c95c9..3e71beaecaa 100644
--- a/docs/quick_start/quick_start.rst
+++ b/docs/quick_start/quick_start.rst
@@ -43,7 +43,7 @@ pre-installed and the ``JAVA_HOME`` is correctly set to each component.
**Kyuubi** Gateway \ |release| \ - Kyuubi Server
Engine lib - Kyuubi Engine
Beeline - Kyuubi Beeline
- **Spark** Engine 3.3 to 3.5, 4.0 to 4.1 A Spark distribution
+ **Spark** Engine 3.3 to 3.5, 4.0 to 4.2 A Spark distribution
**Flink** Engine 1.17 to 1.20 A Flink distribution
**Trino** Engine N/A A Trino cluster allows to access via trino-client v411
**Doris** Engine N/A A Doris cluster
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/pom.xml b/extensions/spark/kyuubi-extension-spark-4-2/pom.xml
new file mode 100644
index 00000000000..a0c45b67c81
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/pom.xml
@@ -0,0 +1,222 @@
+
+
+
+ 4.0.0
+
+ org.apache.kyuubi
+ kyuubi-parent
+ 1.12.0-SNAPSHOT
+ ../../../pom.xml
+
+
+ kyuubi-extension-spark-4-2_${scala.binary.version}
+ jar
+ Kyuubi Dev Spark Extensions (for Spark 4.2)
+ https://kyuubi.apache.org/
+
+
+
+ org.scala-lang
+ scala-library
+ provided
+
+
+
+ org.apache.spark
+ spark-sql_${scala.binary.version}
+ provided
+
+
+
+ org.apache.spark
+ spark-hive_${scala.binary.version}
+ provided
+
+
+
+ org.apache.hadoop
+ hadoop-client-api
+ provided
+
+
+
+ org.apache.kyuubi
+ kyuubi-download
+ ${project.version}
+ pom
+ test
+
+
+
+ org.apache.kyuubi
+ kyuubi-util-scala_${scala.binary.version}
+ ${project.version}
+ test-jar
+ test
+
+
+
+ org.apache.spark
+ spark-core_${scala.binary.version}
+ test-jar
+ test
+
+
+
+ org.apache.spark
+ spark-catalyst_${scala.binary.version}
+ test-jar
+ test
+
+
+
+ org.scalatestplus
+ scalacheck-1-17_${scala.binary.version}
+ test
+
+
+
+ org.apache.spark
+ spark-sql_${scala.binary.version}
+ ${spark.version}
+ test-jar
+ test
+
+
+
+ org.apache.hadoop
+ hadoop-client-runtime
+ test
+
+
+
+ javax.servlet
+ javax.servlet-api
+ test
+
+
+
+ jakarta.servlet
+ jakarta.servlet-api
+ test
+
+
+
+ org.apache.logging.log4j
+ log4j-slf4j-impl
+ test
+
+
+
+
+
+
+
+ org.codehaus.mojo
+ build-helper-maven-plugin
+
+
+ regex-property
+
+ regex-property
+
+
+ spark.home
+ ${project.basedir}/../../../externals/kyuubi-download/target/${spark.archive.name}
+ (.+)\.tgz
+ $1
+
+
+
+
+
+ org.antlr
+ antlr4-maven-plugin
+
+ true
+ ${project.basedir}/src/main/antlr4
+
+
+
+
+ org.apache.maven.plugins
+ maven-shade-plugin
+
+ false
+
+
+ org.apache.kyuubi:*
+
+
+
+
+
+
+ shade
+
+ package
+
+
+
+
+ org.scalatest
+ scalatest-maven-plugin
+
+
+
+ ${scala.binary.version}
+
+
+
+
+ target/scala-${scala.binary.version}/classes
+ target/scala-${scala.binary.version}/test-classes
+
+
+
+
+ spark-home-from-archive
+
+
+ !env.SPARK_HOME
+
+
+
+
+
+ org.scalatest
+ scalatest-maven-plugin
+
+
+ ${spark.home}
+
+
+
+
+
+
+
+
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/antlr4/org/apache/kyuubi/sql/KyuubiSparkSQL.g4 b/extensions/spark/kyuubi-extension-spark-4-2/src/main/antlr4/org/apache/kyuubi/sql/KyuubiSparkSQL.g4
new file mode 100644
index 00000000000..e52b7f5cfeb
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/antlr4/org/apache/kyuubi/sql/KyuubiSparkSQL.g4
@@ -0,0 +1,191 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+grammar KyuubiSparkSQL;
+
+@members {
+ /**
+ * Verify whether current token is a valid decimal token (which contains dot).
+ * Returns true if the character that follows the token is not a digit or letter or underscore.
+ *
+ * For example:
+ * For char stream "2.3", "2." is not a valid decimal token, because it is followed by digit '3'.
+ * For char stream "2.3_", "2.3" is not a valid decimal token, because it is followed by '_'.
+ * For char stream "2.3W", "2.3" is not a valid decimal token, because it is followed by 'W'.
+ * For char stream "12.0D 34.E2+0.12 " 12.0D is a valid decimal token because it is followed
+ * by a space. 34.E2 is a valid decimal token because it is followed by symbol '+'
+ * which is not a digit or letter or underscore.
+ */
+ public boolean isValidDecimal() {
+ int nextChar = _input.LA(1);
+ if (nextChar >= 'A' && nextChar <= 'Z' || nextChar >= '0' && nextChar <= '9' ||
+ nextChar == '_') {
+ return false;
+ } else {
+ return true;
+ }
+ }
+ }
+
+tokens {
+ DELIMITER
+}
+
+singleStatement
+ : statement EOF
+ ;
+
+statement
+ : OPTIMIZE multipartIdentifier whereClause? zorderClause #optimizeZorder
+ | .*? #passThrough
+ ;
+
+whereClause
+ : WHERE partitionPredicate = predicateToken
+ ;
+
+zorderClause
+ : ZORDER BY order+=multipartIdentifier (',' order+=multipartIdentifier)*
+ ;
+
+// We don't have an expression rule in our grammar here, so we just grab the tokens and defer
+// parsing them to later.
+predicateToken
+ : .+?
+ ;
+
+multipartIdentifier
+ : parts+=identifier ('.' parts+=identifier)*
+ ;
+
+identifier
+ : strictIdentifier
+ ;
+
+strictIdentifier
+ : IDENTIFIER #unquotedIdentifier
+ | quotedIdentifier #quotedIdentifierAlternative
+ | nonReserved #unquotedIdentifier
+ ;
+
+quotedIdentifier
+ : BACKQUOTED_IDENTIFIER
+ ;
+
+nonReserved
+ : AND
+ | BY
+ | FALSE
+ | DATE
+ | INTERVAL
+ | OPTIMIZE
+ | OR
+ | TABLE
+ | TIMESTAMP
+ | TRUE
+ | WHERE
+ | ZORDER
+ ;
+
+AND: 'AND';
+BY: 'BY';
+FALSE: 'FALSE';
+DATE: 'DATE';
+INTERVAL: 'INTERVAL';
+OPTIMIZE: 'OPTIMIZE';
+OR: 'OR';
+TABLE: 'TABLE';
+TIMESTAMP: 'TIMESTAMP';
+TRUE: 'TRUE';
+WHERE: 'WHERE';
+ZORDER: 'ZORDER';
+
+MINUS: '-';
+
+BIGINT_LITERAL
+ : DIGIT+ 'L'
+ ;
+
+SMALLINT_LITERAL
+ : DIGIT+ 'S'
+ ;
+
+TINYINT_LITERAL
+ : DIGIT+ 'Y'
+ ;
+
+INTEGER_VALUE
+ : DIGIT+
+ ;
+
+DECIMAL_VALUE
+ : DIGIT+ EXPONENT
+ | DECIMAL_DIGITS EXPONENT? {isValidDecimal()}?
+ ;
+
+DOUBLE_LITERAL
+ : DIGIT+ EXPONENT? 'D'
+ | DECIMAL_DIGITS EXPONENT? 'D' {isValidDecimal()}?
+ ;
+
+BIGDECIMAL_LITERAL
+ : DIGIT+ EXPONENT? 'BD'
+ | DECIMAL_DIGITS EXPONENT? 'BD' {isValidDecimal()}?
+ ;
+
+BACKQUOTED_IDENTIFIER
+ : '`' ( ~'`' | '``' )* '`'
+ ;
+
+IDENTIFIER
+ : (LETTER | DIGIT | '_')+
+ ;
+
+fragment DECIMAL_DIGITS
+ : DIGIT+ '.' DIGIT*
+ | '.' DIGIT+
+ ;
+
+fragment EXPONENT
+ : 'E' [+-]? DIGIT+
+ ;
+
+fragment DIGIT
+ : [0-9]
+ ;
+
+fragment LETTER
+ : [A-Z]
+ ;
+
+SIMPLE_COMMENT
+ : '--' ~[\r\n]* '\r'? '\n'? -> channel(HIDDEN)
+ ;
+
+BRACKETED_COMMENT
+ : '/*' .*? '*/' -> channel(HIDDEN)
+ ;
+
+WS : [ \r\n\t]+ -> channel(HIDDEN)
+ ;
+
+// Catch-all for anything we can't recognize.
+// We use this to be able to ignore and recover all the text
+// when splitting statements with DelimiterLexer
+UNRECOGNIZED
+ : .
+ ;
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/DropIgnoreNonexistent.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/DropIgnoreNonexistent.scala
new file mode 100644
index 00000000000..26b4b5b9422
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/DropIgnoreNonexistent.scala
@@ -0,0 +1,56 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.kyuubi.sql
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.analysis.{UnresolvedFunctionName, UnresolvedRelation}
+import org.apache.spark.sql.catalyst.plans.logical._
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.command.{AlterTableDropPartitionCommand, DropFunctionCommand, DropTableCommand}
+
+import org.apache.kyuubi.sql.KyuubiSQLConf._
+
+case class DropIgnoreNonexistent(session: SparkSession) extends Rule[LogicalPlan] {
+
+ override def apply(plan: LogicalPlan): LogicalPlan = {
+ if (conf.getConf(DROP_IGNORE_NONEXISTENT)) {
+ plan match {
+ case i @ AlterTableDropPartitionCommand(_, _, false, _, _) =>
+ i.copy(ifExists = true)
+ case i @ DropTableCommand(_, false, _, _) =>
+ i.copy(ifExists = true)
+ case i @ DropTable(_, false, _) =>
+ i.copy(ifExists = true)
+ case i @ DropNamespace(_, false, _) =>
+ i.copy(ifExists = true)
+ case i @ DropFunctionCommand(_, false, _) =>
+ i.copy(ifExists = true)
+ case i @ DropView(_, false) =>
+ i.copy(ifExists = true)
+ // refer: org.apache.spark.sql.catalyst.analysis.ResolveCommandsWithIfExists
+ case UncacheTable(u: UnresolvedRelation, false, _) =>
+ NoopCommand("UNCACHE TABLE", u.multipartIdentifier)
+ case DropFunction(u: UnresolvedFunctionName, false) =>
+ NoopCommand("DROP FUNCTION", u.multipartIdentifier)
+ case _ => plan
+ }
+ } else {
+ plan
+ }
+ }
+
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/DynamicShufflePartitions.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/DynamicShufflePartitions.scala
new file mode 100644
index 00000000000..471e43d863d
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/DynamicShufflePartitions.scala
@@ -0,0 +1,98 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.kyuubi.sql
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning, RoundRobinPartitioning}
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.{FileSourceScanExec, SparkPlan}
+import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec
+import org.apache.spark.sql.execution.exchange.{REPARTITION_BY_NUM, ShuffleExchangeExec, ValidateRequirements}
+import org.apache.spark.sql.hive.HiveSparkPlanHelper.HiveTableScanExec
+import org.apache.spark.sql.internal.SQLConf._
+
+import org.apache.kyuubi.sql.KyuubiSQLConf.{DYNAMIC_SHUFFLE_PARTITIONS, DYNAMIC_SHUFFLE_PARTITIONS_MAX_NUM}
+
+/**
+ * Dynamically adjust the number of shuffle partitions according to the input data size
+ */
+case class DynamicShufflePartitions(spark: SparkSession) extends Rule[SparkPlan] {
+
+ override def apply(plan: SparkPlan): SparkPlan = {
+ if (!conf.getConf(DYNAMIC_SHUFFLE_PARTITIONS) || !conf.getConf(ADAPTIVE_EXECUTION_ENABLED)) {
+ plan
+ } else {
+ val maxDynamicShufflePartitions = conf.getConf(DYNAMIC_SHUFFLE_PARTITIONS_MAX_NUM)
+
+ def collectScanSizes(plan: SparkPlan): Seq[Long] = plan match {
+ case FileSourceScanExec(relation, _, _, _, _, _, _, _, _, _) =>
+ Seq(relation.location.sizeInBytes)
+ case t: HiveTableScanExec =>
+ t.relation.prunedPartitions match {
+ case Some(partitions) => Seq(partitions.flatMap(_.stats).map(_.sizeInBytes.toLong).sum)
+ case None => Seq(t.relation.computeStats().sizeInBytes.toLong)
+ .filter(_ != conf.defaultSizeInBytes)
+ }
+ case stage: ShuffleQueryStageExec if stage.isMaterialized && stage.mapStats.isDefined =>
+ Seq(stage.mapStats.get.bytesByPartitionId.sum)
+ case p =>
+ p.children.flatMap(collectScanSizes)
+ }
+
+ val scanSizes = collectScanSizes(plan)
+ if (scanSizes.isEmpty) {
+ return plan
+ }
+
+ val targetSize = conf.getConf(ADVISORY_PARTITION_SIZE_IN_BYTES)
+ val targetShufflePartitions = Math.min(
+ Math.max(scanSizes.sum / targetSize + 1, conf.numShufflePartitions).toInt,
+ maxDynamicShufflePartitions)
+
+ val newPlan = plan transformUp {
+ case exchange @ ShuffleExchangeExec(outputPartitioning, _, shuffleOrigin, _)
+ if shuffleOrigin != REPARTITION_BY_NUM =>
+ val newOutPartitioning = outputPartitioning match {
+ case RoundRobinPartitioning(numPartitions)
+ if targetShufflePartitions != numPartitions =>
+ Some(RoundRobinPartitioning(targetShufflePartitions))
+ case HashPartitioning(expressions, numPartitions)
+ if targetShufflePartitions != numPartitions =>
+ Some(HashPartitioning(expressions, targetShufflePartitions))
+ case RangePartitioning(ordering, numPartitions)
+ if targetShufflePartitions != numPartitions =>
+ Some(RangePartitioning(ordering, targetShufflePartitions))
+ case _ => None
+ }
+ if (newOutPartitioning.isDefined) {
+ exchange.copy(outputPartitioning = newOutPartitioning.get)
+ } else {
+ exchange
+ }
+ }
+
+ if (ValidateRequirements.validate(newPlan)) {
+ newPlan
+ } else {
+ logInfo("DynamicShufflePartitions rule generated an invalid plan. " +
+ "Falling back to the original plan.")
+ plan
+ }
+ }
+ }
+
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/InferRebalanceAndSortOrders.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/InferRebalanceAndSortOrders.scala
new file mode 100644
index 00000000000..ba80eda0782
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/InferRebalanceAndSortOrders.scala
@@ -0,0 +1,134 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import scala.annotation.tailrec
+
+import org.apache.spark.sql.catalyst.expressions.{Alias, Attribute, AttributeSet, BoundReference, Expression, ExtractValue, NamedExpression, OuterReference, UnaryExpression}
+import org.apache.spark.sql.catalyst.planning.ExtractEquiJoinKeys
+import org.apache.spark.sql.catalyst.plans.{FullOuter, Inner, LeftAnti, LeftOuter, LeftSemi, RightOuter}
+import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Filter, Generate, LogicalPlan, Project, Sort, SubqueryAlias, View, Window}
+
+/**
+ * Infer the columns for Rebalance and Sort to improve the compression ratio.
+ *
+ * For example
+ * {{{
+ * INSERT INTO TABLE t PARTITION(p='a')
+ * SELECT * FROM t1 JOIN t2 on t1.c1 = t2.c1
+ * }}}
+ * the inferred columns are: t1.c1
+ */
+object InferRebalanceAndSortOrders {
+
+ type PartitioningAndOrdering = (Seq[Expression], Seq[Expression])
+
+ private def getAliasMap(named: Seq[NamedExpression]): Map[Expression, Attribute] = {
+ @tailrec
+ def throughUnary(e: Expression): Expression = e match {
+ case u: UnaryExpression if u.deterministic =>
+ throughUnary(u.child)
+ case _ => e
+ }
+
+ named.flatMap {
+ case a @ Alias(child, _) =>
+ Some((throughUnary(child).canonicalized, a.toAttribute))
+ case _ => None
+ }.toMap
+ }
+
+ def isCheap(e: Expression): Boolean = e match {
+ case _: Attribute | _: OuterReference | _: BoundReference => true
+ case _ if e.foldable => true
+ case _: Alias | _: ExtractValue => e.children.forall(isCheap)
+ case _ => false
+ }
+
+ def infer(
+ plan: LogicalPlan,
+ onlyInferWithCheapColumns: Boolean): Option[PartitioningAndOrdering] = {
+ def candidateKeys(
+ input: LogicalPlan,
+ output: AttributeSet = AttributeSet.empty): Option[PartitioningAndOrdering] = {
+ input match {
+ case ExtractEquiJoinKeys(joinType, leftKeys, rightKeys, _, _, _, _, _) =>
+ joinType match {
+ case LeftSemi | LeftAnti | LeftOuter => Some((leftKeys, leftKeys))
+ case RightOuter => Some((rightKeys, rightKeys))
+ case Inner | FullOuter =>
+ if (output.isEmpty) {
+ Some((leftKeys ++ rightKeys, leftKeys ++ rightKeys))
+ } else {
+ assert(leftKeys.length == rightKeys.length)
+ val keys = leftKeys.zip(rightKeys).flatMap { case (left, right) =>
+ if (left.references.subsetOf(output)) {
+ Some(left)
+ } else if (right.references.subsetOf(output)) {
+ Some(right)
+ } else {
+ None
+ }
+ }
+ Some((keys, keys))
+ }
+ case _ => None
+ }
+ case agg: Aggregate =>
+ val aliasMap = getAliasMap(agg.aggregateExpressions)
+ Some((
+ agg.groupingExpressions.map(p => aliasMap.getOrElse(p.canonicalized, p)),
+ agg.groupingExpressions.map(o => aliasMap.getOrElse(o.canonicalized, o))))
+ case s: Sort => Some((s.order.map(_.child), s.order.map(_.child)))
+ case p: Project =>
+ val aliasMap = getAliasMap(p.projectList)
+ candidateKeys(p.child, p.references).map { case (partitioning, ordering) =>
+ (
+ partitioning.map(p => aliasMap.getOrElse(p.canonicalized, p)),
+ ordering.map(o => aliasMap.getOrElse(o.canonicalized, o)))
+ }
+ case f: Filter => candidateKeys(f.child, output)
+ case s: SubqueryAlias => candidateKeys(s.child, output)
+ case v: View => candidateKeys(v.child, output)
+ case g: Generate => candidateKeys(g.child, AttributeSet(g.requiredChildOutput))
+ case w: Window =>
+ val aliasMap = getAliasMap(w.windowExpressions)
+ Some((
+ w.partitionSpec.map(p => aliasMap.getOrElse(p.canonicalized, p)),
+ w.orderSpec.map(_.child).map(o => aliasMap.getOrElse(o.canonicalized, o))))
+
+ case _ => None
+ }
+ }
+
+ val partitioningAndSort = candidateKeys(plan).map { case (partitioning, ordering) =>
+ (
+ partitioning.filter(_.references.subsetOf(plan.outputSet)),
+ ordering.filter(_.references.subsetOf(plan.outputSet)))
+ }
+ val allCheap = partitioningAndSort.exists {
+ case (partitionings, sorts) =>
+ partitionings.forall(isCheap) && sorts.forall(isCheap)
+ }
+ if (!onlyInferWithCheapColumns || allCheap) {
+ partitioningAndSort
+ } else {
+ None
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/InsertShuffleNodeBeforeJoin.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/InsertShuffleNodeBeforeJoin.scala
new file mode 100644
index 00000000000..92626f02745
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/InsertShuffleNodeBeforeJoin.scala
@@ -0,0 +1,93 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import org.apache.spark.sql.catalyst.plans.physical.Distribution
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.{SortExec, SparkPlan}
+import org.apache.spark.sql.execution.adaptive.QueryStageExec
+import org.apache.spark.sql.execution.aggregate.BaseAggregateExec
+import org.apache.spark.sql.execution.exchange.{Exchange, ShuffleExchangeExec}
+import org.apache.spark.sql.execution.joins.{ShuffledHashJoinExec, SortMergeJoinExec}
+import org.apache.spark.sql.internal.SQLConf
+
+import org.apache.kyuubi.sql.KyuubiSQLConf._
+
+/**
+ * Insert shuffle node before join if it doesn't exist to make `OptimizeSkewedJoin` works.
+ */
+object InsertShuffleNodeBeforeJoin extends Rule[SparkPlan] {
+
+ override def apply(plan: SparkPlan): SparkPlan = {
+ // this rule has no meaning without AQE
+ if (!conf.getConf(FORCE_SHUFFLE_BEFORE_JOIN) ||
+ !conf.getConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED)) {
+ return plan
+ }
+
+ val newPlan = insertShuffleBeforeJoin(plan)
+ if (plan.fastEquals(newPlan)) {
+ plan
+ } else {
+ // make sure the output partitioning and ordering will not be broken.
+ KyuubiEnsureRequirements.apply(newPlan)
+ }
+ }
+
+ // SPARK-33832 (Spark 3.3) moves the rule OptimizeSkewedJoin from queryStageOptimizerRules
+ // to queryStagePreparationRules, injecting shuffle after OptimizeSkewedJoin may produce
+ // invalid query plan.
+ private def insertShuffleBeforeJoin(plan: SparkPlan): SparkPlan = plan transformUp {
+ case smj @ SortMergeJoinExec(_, _, _, _, l, r, isSkewJoin) if !isSkewJoin =>
+ smj.withNewChildren(checkAndInsertShuffle(smj.requiredChildDistribution.head, l) ::
+ checkAndInsertShuffle(smj.requiredChildDistribution(1), r) :: Nil)
+
+ case shj: ShuffledHashJoinExec if !shj.isSkewJoin =>
+ if (!shj.left.isInstanceOf[Exchange] && !shj.right.isInstanceOf[Exchange]) {
+ shj.withNewChildren(withShuffleExec(shj.requiredChildDistribution.head, shj.left) ::
+ withShuffleExec(shj.requiredChildDistribution(1), shj.right) :: Nil)
+ } else if (!shj.left.isInstanceOf[Exchange]) {
+ shj.withNewChildren(
+ withShuffleExec(shj.requiredChildDistribution.head, shj.left) :: shj.right :: Nil)
+ } else if (!shj.right.isInstanceOf[Exchange]) {
+ shj.withNewChildren(
+ shj.left :: withShuffleExec(shj.requiredChildDistribution(1), shj.right) :: Nil)
+ } else {
+ shj
+ }
+ }
+
+ private def checkAndInsertShuffle(
+ distribution: Distribution,
+ child: SparkPlan): SparkPlan = child match {
+ case SortExec(_, _, _: Exchange, _) =>
+ child
+ case SortExec(_, _, _: QueryStageExec, _) =>
+ child
+ case sort @ SortExec(_, _, agg: BaseAggregateExec, _) =>
+ sort.withNewChildren(withShuffleExec(distribution, agg) :: Nil)
+ case _ =>
+ withShuffleExec(distribution, child)
+ }
+
+ private def withShuffleExec(distribution: Distribution, child: SparkPlan): SparkPlan = {
+ val numPartitions = distribution.requiredNumPartitions
+ .getOrElse(conf.numShufflePartitions)
+ ShuffleExchangeExec(distribution.createPartitioning(numPartitions), child)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiEnsureRequirements.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiEnsureRequirements.scala
new file mode 100644
index 00000000000..a780d6e2fb6
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiEnsureRequirements.scala
@@ -0,0 +1,787 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import scala.collection.mutable
+import scala.collection.mutable.ArrayBuffer
+
+import org.apache.spark.SparkException
+import org.apache.spark.internal.LogKeys
+import org.apache.spark.sql.catalyst.expressions._
+import org.apache.spark.sql.catalyst.plans._
+import org.apache.spark.sql.catalyst.plans.JoinType
+import org.apache.spark.sql.catalyst.plans.physical._
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.catalyst.util.InternalRowComparableWrapper
+import org.apache.spark.sql.connector.catalog.functions.Reducer
+import org.apache.spark.sql.execution._
+import org.apache.spark.sql.execution.datasources.v2.GroupPartitionsExec
+import org.apache.spark.sql.execution.exchange._
+import org.apache.spark.sql.execution.joins.{ShuffledHashJoinExec, SortMergeJoinExec}
+import org.apache.spark.sql.internal.SQLConf
+
+/**
+ * Copy from Apache Spark `EnsureRequirements`
+ * 1. remove reorder join predicates
+ * 2. remove shuffle pruning
+ */
+object KyuubiEnsureRequirements extends Rule[SparkPlan] {
+
+ private def ensureDistributionAndOrdering(
+ parent: Option[SparkPlan],
+ originalChildren: Seq[SparkPlan],
+ requiredChildDistributions: Seq[Distribution],
+ requiredChildOrderings: Seq[Seq[SortOrder]],
+ shuffleOrigin: ShuffleOrigin): Seq[SparkPlan] = {
+ assert(requiredChildDistributions.length == originalChildren.length)
+ assert(requiredChildOrderings.length == originalChildren.length)
+ // Ensure that the operator's children satisfy their output distribution requirements.
+ var children = originalChildren.zip(requiredChildDistributions).map {
+ case (child, distribution) =>
+ // Split child's partitioning into categories
+ val (other, grouped, nonGrouped) = splitKeyedPartitionings(child.outputPartitioning)
+
+ // If non-KeyedPartitioning already satisfies, no changes needed
+ if (other.exists(_.satisfies(distribution))) {
+ child
+ } else {
+ // Check KeyedPartitioning satisfaction conditions
+ val groupedSatisfies = grouped.find(_.satisfies(distribution))
+ val nonGroupedSatisfiesAsIs = nonGrouped.exists(_.nonGroupedSatisfies(distribution))
+ val nonGroupedSatisfiesWhenGrouped = nonGrouped.find(_.groupedSatisfies(distribution))
+
+ // Check if any KeyedPartitioning satisfies the distribution
+ if (groupedSatisfies.isDefined || nonGroupedSatisfiesAsIs
+ || nonGroupedSatisfiesWhenGrouped.isDefined) {
+ distribution match {
+ case o: OrderedDistribution =>
+ // OrderedDistribution requires grouped KeyedPartitioning with sorted keys
+ // according to the distribution's ordering.
+ // Find any KeyedPartitioning that satisfies via groupedSatisfies.
+ val satisfyingKeyedPartitioning =
+ groupedSatisfies.orElse(nonGroupedSatisfiesWhenGrouped).get
+ val attrs = satisfyingKeyedPartitioning.expressions.flatMap(_.collectLeaves())
+ .map(_.asInstanceOf[Attribute])
+ val keyRowOrdering = RowOrdering.create(o.ordering, attrs)
+ val keyOrdering = keyRowOrdering.on((t: InternalRowComparableWrapper) => t.row)
+ if (satisfyingKeyedPartitioning.partitionKeys.sliding(2).forall {
+ case Seq(k1, k2) => keyOrdering.lteq(k1, k2)
+ }) {
+ child
+ } else {
+ // Use distributePartitions to spread splits across expected partitions
+ val sortedGroupedKeys = satisfyingKeyedPartitioning.partitionKeys
+ .groupBy(identity).view.mapValues(_.size)
+ .toSeq.sortBy(_._1)(keyOrdering)
+ GroupPartitionsExec(
+ child,
+ expectedPartitionKeys = Some(sortedGroupedKeys),
+ distributePartitions = true)
+ }
+
+ case _ if groupedSatisfies.isDefined =>
+ // Grouped KeyedPartitioning already satisfies
+ child
+
+ case _ if nonGroupedSatisfiesAsIs =>
+ // Non-grouped KeyedPartitioning satisfies without grouping
+ child
+
+ case _ =>
+ // Non-grouped KeyedPartitioning satisfies only after grouping
+ GroupPartitionsExec(child)
+ }
+ } else {
+ // No partitioning satisfies - need broadcast or shuffle
+ val numPartitions = distribution.requiredNumPartitions
+ .getOrElse(conf.numShufflePartitions)
+ distribution match {
+ case BroadcastDistribution(mode) =>
+ BroadcastExchangeExec(mode, child)
+ case _: StatefulOpClusteredDistribution =>
+ ShuffleExchangeExec(
+ distribution.createPartitioning(numPartitions),
+ child,
+ REQUIRED_BY_STATEFUL_OPERATOR)
+ case _ =>
+ ShuffleExchangeExec(
+ distribution.createPartitioning(numPartitions),
+ child,
+ shuffleOrigin)
+ }
+ }
+ }
+ }
+
+ // Get the indexes of children which have specified distribution requirements and need to be
+ // co-partitioned.
+ val childrenIndexes = requiredChildDistributions.zipWithIndex.filter {
+ case (_: ClusteredDistribution, _) => true
+ case _ => false
+ }.map(_._2)
+
+ // Special case: if all sides of the join are single partition and it's physical size less than
+ // or equal spark.sql.maxSinglePartitionBytes.
+ val preferSinglePartition = childrenIndexes.forall { i =>
+ children(i).outputPartitioning == SinglePartition &&
+ children(i).logicalLink
+ .forall(_.stats.sizeInBytes <= conf.getConf(SQLConf.MAX_SINGLE_PARTITION_BYTES))
+ }
+
+ // If there are more than one children, we'll need to check partitioning & distribution of them
+ // and see if extra shuffles are necessary.
+ if (childrenIndexes.length > 1 && !preferSinglePartition) {
+ val specs = childrenIndexes.map(i => {
+ val requiredDist = requiredChildDistributions(i)
+ assert(
+ requiredDist.isInstanceOf[ClusteredDistribution],
+ s"Expected ClusteredDistribution but found ${requiredDist.getClass.getSimpleName}")
+ i -> children(i).outputPartitioning.createShuffleSpec(
+ requiredDist.asInstanceOf[ClusteredDistribution])
+ }).toMap
+
+ // Find out the shuffle spec that gives better parallelism. Currently this is done by
+ // picking the spec with the largest number of partitions.
+ //
+ // NOTE: this is not optimal for the case when there are more than 2 children. Consider:
+ // (10, 10, 11)
+ // where the number represent the number of partitions for each child, it's better to pick 10
+ // here since we only need to shuffle one side - we'd need to shuffle two sides if we pick 11.
+ //
+ // However this should be sufficient for now since in Spark nodes with multiple children
+ // always have exactly 2 children.
+
+ // Whether we should consider `spark.sql.shuffle.partitions` and ensure enough parallelism
+ // during shuffle. To achieve a good trade-off between parallelism and shuffle cost, we only
+ // consider the minimum parallelism iff ALL children need to be re-shuffled.
+ //
+ // A child needs to be re-shuffled iff either one of below is true:
+ // 1. It can't create partitioning by itself, i.e., `canCreatePartitioning` returns false
+ // (as for the case of `RangePartitioning`), therefore it needs to be re-shuffled
+ // according to other shuffle spec.
+ // 2. It already has `ShuffleExchangeLike`, so we can re-use existing shuffle without
+ // introducing extra shuffle.
+ //
+ // On the other hand, in scenarios such as:
+ // HashPartitioning(5) <-> HashPartitioning(6)
+ // while `spark.sql.shuffle.partitions` is 10, we'll only re-shuffle the left side and make it
+ // HashPartitioning(6).
+ val shouldConsiderMinParallelism = specs.forall(p =>
+ !p._2.canCreatePartitioning || children(p._1).isInstanceOf[ShuffleExchangeLike])
+ // Choose all the specs that can be used to shuffle other children
+ val candidateSpecs = specs.filter { case (index, spec) =>
+ spec.canCreatePartitioning &&
+ (!shouldConsiderMinParallelism ||
+ children(index).outputPartitioning.numPartitions >= conf.defaultNumShufflePartitions)
+ }
+ val bestSpecOpt = if (candidateSpecs.isEmpty) {
+ None
+ } else {
+ // When choosing specs, we should consider those children with no `ShuffleExchangeLike` node
+ // first. For instance, if we have:
+ // A: (No_Exchange, 100) <---> B: (Exchange, 120)
+ // it's better to pick A and change B to (Exchange, 100) instead of picking B and insert a
+ // new shuffle for A.
+ val candidateSpecsWithoutShuffle = candidateSpecs.filter { case (k, _) =>
+ !children(k).isInstanceOf[ShuffleExchangeLike]
+ }
+ val finalCandidateSpecs = if (candidateSpecsWithoutShuffle.nonEmpty) {
+ candidateSpecsWithoutShuffle
+ } else {
+ candidateSpecs
+ }
+ // Pick the spec with the best parallelism
+ Some(finalCandidateSpecs.values.maxBy(_.numPartitions))
+ }
+
+ // Check if the following conditions are satisfied:
+ // 1. There are exactly two children (e.g., join). Note that Spark doesn't support
+ // multi-way join at the moment, so this check should be sufficient.
+ // 2. All children are of the compatible key group partitioning or
+ // compatible shuffle partition id pass through partitioning
+ // If both are true, skip shuffle.
+ val areChildrenCompatible = parent.isDefined &&
+ children.length == 2 && childrenIndexes.length == 2 && {
+ val left = children.head
+ val right = children(1)
+
+ // key group compatibility check
+ val newChildren = checkKeyGroupCompatible(
+ parent.get,
+ left,
+ right,
+ requiredChildDistributions)
+ if (newChildren.isDefined) {
+ children = newChildren.get
+ true
+ } else {
+ // If key group check fails, check ShufflePartitionIdPassThrough compatibility
+ checkShufflePartitionIdPassThroughCompatible(
+ left,
+ right,
+ requiredChildDistributions)
+ }
+ }
+
+ children = children.zip(requiredChildDistributions).zipWithIndex.map {
+ case ((child, _), idx)
+ if areChildrenCompatible ||
+ !childrenIndexes.contains(idx) =>
+ child
+ case ((child, dist), idx) =>
+ if (bestSpecOpt.isDefined && bestSpecOpt.get.isCompatibleWith(specs(idx))) {
+ bestSpecOpt match {
+ // If `areChildrenCompatible` is false, we can still perform SPJ
+ // by shuffling the other side based on join keys (see the else case below).
+ // Hence we need to ensure that after this call, the outputPartitioning of the
+ // partitioned side's BatchScanExec is grouped by join keys to match,
+ // and we do that by pushing down the join keys
+ case Some(KeyedShuffleSpec(_, _, Some(joinKeyPositions))) =>
+ withJoinKeyPositions(child, joinKeyPositions)
+ case _ => child
+ }
+ } else {
+ val newPartitioning = bestSpecOpt.map { bestSpec =>
+ // Use the best spec to create a new partitioning to re-shuffle this child
+ val clustering = dist.asInstanceOf[ClusteredDistribution].clustering
+ bestSpec.createPartitioning(clustering)
+ }.getOrElse {
+ // No best spec available, so we create default partitioning from the required
+ // distribution
+ val numPartitions = dist.requiredNumPartitions
+ .getOrElse(conf.numShufflePartitions)
+ dist.createPartitioning(numPartitions)
+ }
+
+ child match {
+ case ShuffleExchangeExec(_, c, so, ps) =>
+ ShuffleExchangeExec(newPartitioning, c, so, ps)
+ case gpe: GroupPartitionsExec => ShuffleExchangeExec(newPartitioning, gpe.child)
+ case _ => ShuffleExchangeExec(newPartitioning, child)
+ }
+ }
+ }
+ }
+
+ // Now that we've performed any necessary shuffles, add sorts to guarantee output orderings:
+ children = children.zip(requiredChildOrderings).map { case (child, requiredOrdering) =>
+ // If child.outputOrdering already satisfies the requiredOrdering, we do not need to sort.
+ if (SortOrder.orderingSatisfies(child.outputOrdering, requiredOrdering)) {
+ child
+ } else {
+ // Before adding a SortExec, check whether a GroupPartitionsExec anywhere in the child
+ // subtree can self-satisfy via sorted merge. tryEnableSortedMerge generates all alternative
+ // plans where one or more GPEs have sorted merge enabled; we take the first one whose
+ // outputOrdering satisfies the requirement.
+ tryEnableSortedMerge(child)
+ .find(newChild => SortOrder.orderingSatisfies(newChild.outputOrdering, requiredOrdering))
+ .getOrElse(SortExec(requiredOrdering, global = false, child = child))
+ }
+ }
+
+ children
+ }
+
+ private def hasKeyedPartitioning(p: Partitioning): Boolean = p match {
+ case e: Expression => e.exists(_.isInstanceOf[KeyedPartitioning])
+ case _ => false
+ }
+
+ // Generates all alternative plans in which one or more GroupPartitionsExec nodes in the subtree
+ // have sorted-merge enabled (every possible combination). Returns a LazyList so the caller can
+ // stop evaluating once a satisfying alternative is found.
+ //
+ // Pruning: traversal stops at SortExec (which reorders data, making sorted merge below it
+ // pointless) and at any node whose outputPartitioning no longer carries a KeyedPartitioning.
+ // This is a good heuristic, though not strictly equivalent to "ordering no longer propagates":
+ // partition-key expressions are constant within each coalesced partition and therefore usually
+ // prefix outputOrdering. When a node prunes the KeyedPartitioning (e.g. a Project that drops
+ // partition keys), it also prunes that ordering prefix. Since Spark has no notion of constant
+ // expressions in SortOrder, dropping a prefix invalidates the rest of the ordering too -- so in
+ // practice the two are always pruned together.
+ //
+ // At each GPE the rule emits [original, sorted-merge-enabled] alternatives (or just [original]
+ // when sorted merge cannot be enabled). multiTransformDownWithPruning then builds the Cartesian
+ // product across all GPEs in the subtree, giving every combination.
+ private[sql] def tryEnableSortedMerge(plan: SparkPlan): LazyList[SparkPlan] =
+ plan.multiTransformDownWithPruning(p =>
+ !p.isInstanceOf[SortExec] &&
+ hasKeyedPartitioning(p.asInstanceOf[SparkPlan].outputPartitioning)) {
+ case gpe: GroupPartitionsExec =>
+ // Include the original so that peer GPEs are still independently considered.
+ gpe +: gpe.tryEnableSortedMerge().toSeq
+ }
+
+ /**
+ * Checks whether two children, `left` and `right`, of a join operator have compatible
+ * `KeyedPartitioning`, and can benefit from storage-partitioned join.
+ *
+ * Returns the updated new children if the check is successful, otherwise `None`.
+ */
+ private def checkKeyGroupCompatible(
+ parent: SparkPlan,
+ left: SparkPlan,
+ right: SparkPlan,
+ requiredChildDistribution: Seq[Distribution]): Option[Seq[SparkPlan]] = {
+ parent match {
+ case smj: SortMergeJoinExec =>
+ checkKeyGroupCompatible(left, right, smj.joinType, requiredChildDistribution)
+ case sj: ShuffledHashJoinExec =>
+ checkKeyGroupCompatible(left, right, sj.joinType, requiredChildDistribution)
+ case _ =>
+ None
+ }
+ }
+
+ private def checkKeyGroupCompatible(
+ left: SparkPlan,
+ right: SparkPlan,
+ joinType: JoinType,
+ requiredChildDistribution: Seq[Distribution]): Option[Seq[SparkPlan]] = {
+ assert(requiredChildDistribution.length == 2)
+
+ var newLeft = left
+ var newRight = right
+
+ val specs = Seq(left, right).zip(requiredChildDistribution).map { case (p, d) =>
+ if (!d.isInstanceOf[ClusteredDistribution]) return None
+ val cd = d.asInstanceOf[ClusteredDistribution]
+ val specOpt = createKeyedShuffleSpec(p.outputPartitioning, cd)
+ if (specOpt.isEmpty) return None
+ specOpt.get
+ }
+
+ val leftSpec = specs.head
+ val rightSpec = specs(1)
+ val leftPartitioning = leftSpec.partitioning
+ val rightPartitioning = rightSpec.partitioning
+
+ // We don't need to alter the existing or add new `GroupPartitionsExec` when the child
+ // partitionings are not modified (projected) in specs and left and right side partitionings are
+ // compatible with each other.
+ // Left and right `outputPartitioning` is a `PartitioningCollection` or a `KeyedPartitioning`
+ // otherwise `createKeyedShuffleSpec()` would have returned `None`.
+ var isCompatible =
+ left.outputPartitioning.asInstanceOf[Expression].exists(_ == leftPartitioning) &&
+ right.outputPartitioning.asInstanceOf[Expression].exists(_ == rightPartitioning) &&
+ leftSpec.isCompatibleWith(rightSpec)
+ if ((!isCompatible || conf.v2BucketingPartiallyClusteredDistributionEnabled) &&
+ (conf.v2BucketingPushPartValuesEnabled ||
+ conf.v2BucketingAllowJoinKeysSubsetOfPartitionKeys)) {
+ logInfo("Pushing common partition values for storage-partitioned join")
+ isCompatible = leftSpec.areKeysCompatible(rightSpec)
+
+ // Partition expressions are compatible. Regardless of whether partition values
+ // match from both sides of children, we can calculate a superset of partition values and
+ // push-down to respective data sources so they can adjust their output partitioning by
+ // filling missing partition keys with empty partitions. As result, we can still avoid
+ // shuffle.
+ //
+ // For instance, if two sides of a join have partition expressions
+ // `day(a)` and `day(b)` respectively
+ // (the join query could be `SELECT ... FROM t1 JOIN t2 on t1.a = t2.b`), but
+ // with different partition values:
+ // `day(a)`: [0, 1]
+ // `day(b)`: [1, 2, 3]
+ // Following the case 2 above, we don't have to shuffle both sides, but instead can
+ // just push the common set of partition values: `[0, 1, 2, 3]` down to the two data
+ // sources.
+ if (isCompatible) {
+ val leftPartKeys = leftPartitioning.partitionKeys
+ val rightPartKeys = rightPartitioning.partitionKeys
+
+ val numLeftPartKeys = MDC(LogKeys.NUM_LEFT_PARTITION_VALUES, leftPartKeys.size)
+ val numRightPartKeys = MDC(LogKeys.NUM_RIGHT_PARTITION_VALUES, rightPartKeys.size)
+ logInfo(
+ log"""
+ |Left side # of partitions: $numLeftPartKeys
+ |Right side # of partitions: $numRightPartKeys
+ |""".stripMargin)
+
+ // in case of compatible but not identical partition expressions, we apply 'reduce'
+ // transforms to group one side's partitions as well as the common partition values
+ val leftReducers = leftSpec.reducers(rightSpec)
+ val rightReducers = rightSpec.reducers(leftSpec)
+ val (leftReducedDataTypes, leftReducedKeys) = leftReducers.fold(
+ (
+ leftPartitioning.expressionDataTypes,
+ leftPartitioning.partitionKeys))(leftPartitioning.reduceKeys)
+ val (rightReducedDataTypes, rightReducedKeys) = rightReducers.fold(
+ (
+ rightPartitioning.expressionDataTypes,
+ rightPartitioning.partitionKeys))(rightPartitioning.reduceKeys)
+ val reducedDataTypes = if (leftReducedDataTypes == rightReducedDataTypes) {
+ leftReducedDataTypes
+ } else {
+ // QueryExecutionErrors.storagePartitionJoinIncompatibleReducedTypesError is
+ // private[sql] in Spark and thus inaccessible from kyuubi; fall back to a generic
+ // internal error carrying the same diagnostic information.
+ throw SparkException.internalError(
+ s"Storage partition join has incompatible reduced key data types: " +
+ s"left=$leftReducedDataTypes with reducers=$leftReducers, " +
+ s"right=$rightReducedDataTypes with reducers=$rightReducers")
+ }
+
+ val reducedKeyRowOrdering = RowOrdering.createNaturalAscendingOrdering(reducedDataTypes)
+ val reducedKeyOrdering =
+ reducedKeyRowOrdering.on((t: InternalRowComparableWrapper) => t.row)
+
+ // merge values on both sides
+ var mergedPartitionKeys =
+ mergeAndDedupPartitions(leftReducedKeys, rightReducedKeys, joinType, reducedKeyOrdering)
+ .map((_, 1))
+
+ logInfo(log"After merging, there are " +
+ log"${MDC(LogKeys.NUM_PARTITIONS, mergedPartitionKeys.size)} partitions")
+
+ var replicateLeftSide = false
+ var replicateRightSide = false
+ var applyPartialClustering = false
+
+ // This means we allow partitions that are not clustered on their values,
+ // that is, multiple partitions with the same partition value. In the
+ // following, we calculate how many partitions that each distinct partition
+ // value has, and pushdown the information to scans, so they can adjust their
+ // final input partitions respectively.
+ if (conf.v2BucketingPartiallyClusteredDistributionEnabled) {
+ logInfo("Calculating partially clustered distribution for " +
+ "storage-partitioned join")
+
+ // Similar to `OptimizeSkewedJoin`, we need to check join type and decide
+ // whether partially clustered distribution can be applied. For instance, the
+ // optimization cannot be applied to a left outer join, where the left hand
+ // side is chosen as the side to replicate partitions according to stats.
+ // Otherwise, query result could be incorrect.
+ val canReplicateLeft = canReplicateLeftSide(joinType)
+ val canReplicateRight = canReplicateRightSide(joinType)
+
+ if (!canReplicateLeft && !canReplicateRight) {
+ logInfo(log"Skipping partially clustered distribution as it cannot be applied for " +
+ log"join type '${MDC(LogKeys.JOIN_TYPE, joinType)}'")
+ } else {
+ val unwrappedLeft = unwrapGroupPartitions(left)
+ val unwrappedRight = unwrapGroupPartitions(right)
+
+ val leftLink = unwrappedLeft.logicalLink
+ val rightLink = unwrappedRight.logicalLink
+
+ replicateLeftSide =
+ if (leftLink.isDefined && rightLink.isDefined &&
+ leftLink.get.stats.sizeInBytes > 1 &&
+ rightLink.get.stats.sizeInBytes > 1) {
+ val leftLinkStatsSizeInBytes =
+ MDC(LogKeys.LEFT_LOGICAL_PLAN_STATS_SIZE_IN_BYTES, leftLink.get.stats.sizeInBytes)
+ val rightLinkStatsSizeInBytes = MDC(
+ LogKeys.RIGHT_LOGICAL_PLAN_STATS_SIZE_IN_BYTES,
+ rightLink.get.stats.sizeInBytes)
+ logInfo(
+ log"""
+ |Using plan statistics to determine which side of join to fully
+ |cluster partition values:
+ |Left side size (in bytes): $leftLinkStatsSizeInBytes
+ |Right side size (in bytes): $rightLinkStatsSizeInBytes
+ |""".stripMargin)
+ leftLink.get.stats.sizeInBytes < rightLink.get.stats.sizeInBytes
+ } else {
+ // As a simple heuristic, we pick the side with fewer number of partitions
+ // to apply the grouping & replication of partitions
+ logInfo("Using number of partitions to determine which side of join " +
+ "to fully cluster partition values")
+ leftPartKeys.size < rightPartKeys.size
+ }
+
+ replicateRightSide = !replicateLeftSide
+
+ // Similar to skewed join, we need to check the join type to see whether replication
+ // of partitions can be applied. For instance, replication should not be allowed for
+ // the left-hand side of a right outer join.
+ if (replicateLeftSide && !canReplicateLeft) {
+ logInfo(log"Left-hand side is picked but cannot be applied to join type " +
+ log"'${MDC(LogKeys.JOIN_TYPE, joinType)}'. Skipping partially clustered " +
+ log"distribution.")
+ replicateLeftSide = false
+ } else if (replicateRightSide && !canReplicateRight) {
+ logInfo(log"Right-hand side is picked but cannot be applied to join type " +
+ log"'${MDC(LogKeys.JOIN_TYPE, joinType)}'. Skipping partially clustered " +
+ log"distribution.")
+ replicateRightSide = false
+ } else {
+ // In partially clustered distribution, we should use un-grouped partition values
+ val (partiallyClusteredChild, partiallyClusteredSpec) = if (replicateLeftSide) {
+ (unwrappedRight, rightSpec)
+ } else {
+ (unwrappedLeft, leftSpec)
+ }
+ // Original `KeyedPartitioning` can be obtained from the child directly if the child
+ // satisfied the distribution requirement; or from the child's child if it didn't as
+ // the child must be a `GroupPartitionsExec` inserted by `EnsureRequirement`
+ // to satisfy the distribution requirement.
+ val originalPartitioning =
+ partiallyClusteredChild.outputPartitioning.asInstanceOf[Expression]
+ // `outputPartitioning` is either a `PartitioningCollection` or a `KeyedPartitioning`
+ // otherwise `createKeyedShuffleSpec()` would have returned `None`.
+ val originalKeyedPartitioning =
+ originalPartitioning.collectFirst { case k: KeyedPartitioning => k }.get
+ val projectedOriginalPartitionKeys = partiallyClusteredSpec.joinKeyPositions
+ .fold(originalKeyedPartitioning.partitionKeys)(
+ originalKeyedPartitioning.projectKeys(_)._2)
+
+ val numExpectedPartitions =
+ projectedOriginalPartitionKeys.groupBy(identity).view.mapValues(_.size)
+
+ mergedPartitionKeys = mergedPartitionKeys.map { case (key, numParts) =>
+ (key, numExpectedPartitions.getOrElse(key, numParts))
+ }
+
+ logInfo(log"After applying partially clustered distribution, there are " +
+ log"${MDC(LogKeys.NUM_PARTITIONS, mergedPartitionKeys.map(_._2).sum)} partitions.")
+ applyPartialClustering = true
+ }
+ }
+ }
+
+ // Now we need to push-down the common partition information to the `GroupPartitionsExec`s.
+ newLeft = applyGroupPartitions(
+ left,
+ leftSpec.joinKeyPositions,
+ mergedPartitionKeys,
+ leftReducers,
+ distributePartitions = applyPartialClustering && !replicateLeftSide)
+ newRight = applyGroupPartitions(
+ right,
+ rightSpec.joinKeyPositions,
+ mergedPartitionKeys,
+ rightReducers,
+ distributePartitions = applyPartialClustering && !replicateRightSide)
+ }
+ }
+
+ if (isCompatible) Some(Seq(newLeft, newRight)) else None
+ }
+
+ private def checkShufflePartitionIdPassThroughCompatible(
+ left: SparkPlan,
+ right: SparkPlan,
+ requiredChildDistribution: Seq[Distribution]): Boolean = {
+ (left.outputPartitioning, right.outputPartitioning) match {
+ case (p1: ShufflePartitionIdPassThrough, p2: ShufflePartitionIdPassThrough) =>
+ assert(requiredChildDistribution.length == 2)
+ val leftSpec = p1.createShuffleSpec(
+ requiredChildDistribution.head.asInstanceOf[ClusteredDistribution])
+ val rightSpec = p2.createShuffleSpec(
+ requiredChildDistribution(1).asInstanceOf[ClusteredDistribution])
+ leftSpec.isCompatibleWith(rightSpec)
+ case _ =>
+ false
+ }
+ }
+
+ // Similar to `OptimizeSkewedJoin.canSplitRightSide`
+ private def canReplicateLeftSide(joinType: JoinType): Boolean = {
+ joinType == Inner || joinType == Cross || joinType == RightOuter
+ }
+
+ // Similar to `OptimizeSkewedJoin.canSplitLeftSide`
+ private def canReplicateRightSide(joinType: JoinType): Boolean = {
+ joinType == Inner || joinType == Cross || joinType == LeftSemi ||
+ joinType == LeftAnti || joinType == LeftOuter
+ }
+
+ /**
+ * Unwraps a GroupPartitionsExec to get the underlying child plan.
+ */
+ private def unwrapGroupPartitions(plan: SparkPlan): SparkPlan = plan match {
+ case g: GroupPartitionsExec => g.child
+ case other => other
+ }
+
+ /**
+ * Applies or updates `GroupPartitionsExec` with the given parameters.
+ *
+ * `GroupPartitionsExec` can be either the given plan node (child of the join inserted by
+ * `EnsureRequirement`) if the original child didn't satisfy the distribution requirement; or we
+ * can create a new one specifically for this join.
+ */
+ private def applyGroupPartitions(
+ plan: SparkPlan,
+ joinKeyPositions: Option[Seq[Int]],
+ mergedPartitionKeys: Seq[(InternalRowComparableWrapper, Int)],
+ reducers: Option[Seq[Option[Reducer[_, _]]]],
+ distributePartitions: Boolean): SparkPlan = {
+ plan match {
+ case g: GroupPartitionsExec =>
+ val newGroupPartitions = g.copy(
+ joinKeyPositions = joinKeyPositions,
+ expectedPartitionKeys = Some(mergedPartitionKeys),
+ reducers = reducers,
+ distributePartitions = distributePartitions)
+ newGroupPartitions.copyTagsFrom(g)
+ newGroupPartitions
+ case _ =>
+ GroupPartitionsExec(
+ plan,
+ joinKeyPositions,
+ Some(mergedPartitionKeys),
+ reducers,
+ distributePartitions)
+ }
+ }
+
+ /**
+ * Applies join key positions to a plan by wrapping or updating GroupPartitionsExec.
+ */
+ private def withJoinKeyPositions(plan: SparkPlan, positions: Seq[Int]): SparkPlan = {
+ plan match {
+ case g: GroupPartitionsExec =>
+ val newGroupPartitions = g.copy(joinKeyPositions = Some(positions))
+ newGroupPartitions.copyTagsFrom(g)
+ newGroupPartitions
+ case _ => GroupPartitionsExec(plan, joinKeyPositions = Some(positions))
+ }
+ }
+
+ /**
+ * Tries to create a [[KeyedShuffleSpec]] from the input partitioning and distribution, if the
+ * partitioning is a [[KeyedPartitioning]] (either directly or indirectly), and satisfies the
+ * given distribution.
+ */
+ private def createKeyedShuffleSpec(
+ partitioning: Partitioning,
+ distribution: ClusteredDistribution): Option[KeyedShuffleSpec] = {
+ def tryCreate(partitioning: KeyedPartitioning): Option[KeyedShuffleSpec] = {
+ val attributes = partitioning.expressions.flatMap(_.collectLeaves())
+ val clustering = distribution.clustering
+
+ val satisfies = if (SQLConf.get.getConf(SQLConf.REQUIRE_ALL_CLUSTER_KEYS_FOR_CO_PARTITION)) {
+ attributes.length == clustering.length && attributes.zip(clustering).forall {
+ case (l, r) => l.semanticEquals(r)
+ }
+ } else {
+ partitioning.satisfies(distribution)
+ }
+
+ if (satisfies) {
+ Some(partitioning.createShuffleSpec(distribution).asInstanceOf[KeyedShuffleSpec])
+ } else {
+ None
+ }
+ }
+
+ partitioning match {
+ case p: KeyedPartitioning => tryCreate(p)
+ case PartitioningCollection(partitionings) =>
+ partitionings.collectFirst(Function.unlift(createKeyedShuffleSpec(_, distribution)))
+ case _ => None
+ }
+ }
+
+ /**
+ * Merge, dedup and sort partitions keys for SPJ and optionally enable partition filtering.
+ * Both sides must have matching partition expressions.
+ * @param leftPartitionKeys left side partition keys
+ * @param rightPartitionKeys right side partition keys
+ * @param joinType join type for optional partition filtering
+ * @param keyOrdering ordering to sort partition keys
+ * @return merged and sorted partition values
+ */
+ def mergeAndDedupPartitions(
+ leftPartitionKeys: Seq[InternalRowComparableWrapper],
+ rightPartitionKeys: Seq[InternalRowComparableWrapper],
+ joinType: JoinType,
+ keyOrdering: Ordering[InternalRowComparableWrapper]): Seq[InternalRowComparableWrapper] = {
+ val merged = if (SQLConf.get.getConf(SQLConf.V2_BUCKETING_PARTITION_FILTER_ENABLED)) {
+ joinType match {
+ case Inner =>
+ mergeAndDedupPartitionKeys(leftPartitionKeys, rightPartitionKeys, intersect = true)
+ case LeftOuter => leftPartitionKeys.distinct
+ case RightOuter => rightPartitionKeys.distinct
+ case _ => mergeAndDedupPartitionKeys(leftPartitionKeys, rightPartitionKeys)
+ }
+ } else {
+ mergeAndDedupPartitionKeys(leftPartitionKeys, rightPartitionKeys)
+ }
+
+ // SPARK-41471: We keep to order of partitions to make sure the order of
+ // partitions is deterministic in different case.
+ merged.sorted(keyOrdering)
+ }
+
+ private def mergeAndDedupPartitionKeys(
+ leftPartitionKeys: Seq[InternalRowComparableWrapper],
+ rightPartitionKeys: Seq[InternalRowComparableWrapper],
+ intersect: Boolean = false) = {
+ val leftKeySet = mutable.HashSet.from(leftPartitionKeys)
+ val rightKeySet = mutable.HashSet.from(rightPartitionKeys)
+ val result = if (intersect) {
+ leftKeySet.intersect(rightKeySet)
+ } else {
+ leftKeySet.union(rightKeySet)
+ }
+ result.toSeq
+ }
+
+ /**
+ * Splits a partitioning into three categories:
+ * 1. Non-KeyedPartitioning (HashPartitioning, RangePartitioning, etc.)
+ * 2. Grouped KeyedPartitioning (isGrouped = true)
+ * 3. Non-grouped KeyedPartitioning (isGrouped = false)
+ *
+ * @param partitioning The partitioning to split
+ * @return A tuple of (other, grouped, nonGrouped) where:
+ * - other: Option containing non-KeyedPartitioning(s)
+ * - grouped: Seq of grouped KeyedPartitionings
+ * - nonGrouped: Seq of non-grouped KeyedPartitionings
+ */
+ private def splitKeyedPartitionings(partitioning: Partitioning) = {
+ val otherPartitionings = ArrayBuffer.empty[Partitioning]
+ val groupedKeyedPartitionings = ArrayBuffer.empty[KeyedPartitioning]
+ val nonGroupedKeyedPartitionings = ArrayBuffer.empty[KeyedPartitioning]
+
+ def split(p: Partitioning): Unit = p match {
+ case c: PartitioningCollection => c.partitionings.foreach(split)
+ case k: KeyedPartitioning =>
+ if (k.isGrouped) {
+ groupedKeyedPartitionings += k
+ } else {
+ nonGroupedKeyedPartitionings += k
+ }
+ case o => otherPartitionings += o
+ }
+
+ split(partitioning)
+
+ val other = otherPartitionings.length match {
+ case 0 => None
+ case 1 => Some(otherPartitionings.head)
+ case _ => Some(PartitioningCollection(otherPartitionings.toSeq))
+ }
+
+ (other, groupedKeyedPartitionings.toSeq, nonGroupedKeyedPartitionings.toSeq)
+ }
+
+ def apply(plan: SparkPlan): SparkPlan = plan.transformUp {
+ case operator: SparkPlan =>
+ val newChildren = ensureDistributionAndOrdering(
+ Some(operator),
+ operator.children,
+ operator.requiredChildDistribution,
+ operator.requiredChildOrdering,
+ ENSURE_REQUIREMENTS)
+ operator.withNewChildren(newChildren)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiQueryStagePreparation.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiQueryStagePreparation.scala
new file mode 100644
index 00000000000..a7fcbecd422
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiQueryStagePreparation.scala
@@ -0,0 +1,194 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.SparkPlan
+import org.apache.spark.sql.execution.adaptive.QueryStageExec
+import org.apache.spark.sql.execution.command.{ResetCommand, SetCommand}
+import org.apache.spark.sql.execution.exchange.{BroadcastExchangeLike, ReusedExchangeExec, ShuffleExchangeLike}
+import org.apache.spark.sql.internal.SQLConf
+
+import org.apache.kyuubi.sql.KyuubiSQLConf._
+
+/**
+ * This rule split stage into two parts:
+ * 1. previous stage
+ * 2. final stage
+ * For final stage, we can inject extra config. It's useful if we use repartition to optimize
+ * small files that needs bigger shuffle partition size than previous.
+ *
+ * Let's say we have a query with 3 stages, then the logical machine like:
+ *
+ * Set/Reset Command -> cleanup previousStage config if user set the spark config.
+ * Query -> AQE -> stage1 -> preparation (use previousStage to overwrite spark config)
+ * -> AQE -> stage2 -> preparation (use spark config)
+ * -> AQE -> stage3 -> preparation (use finalStage config to overwrite spark config,
+ * store spark config to previousStage.)
+ *
+ * An example of the new finalStage config:
+ * `spark.sql.adaptive.advisoryPartitionSizeInBytes` ->
+ * `spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes`
+ */
+case class FinalStageConfigIsolation(session: SparkSession) extends Rule[SparkPlan] {
+ import FinalStageConfigIsolation._
+
+ override def apply(plan: SparkPlan): SparkPlan = {
+ // this rule has no meaning without AQE
+ if (!conf.getConf(FINAL_STAGE_CONFIG_ISOLATION) ||
+ !conf.getConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED)) {
+ return plan
+ }
+
+ if (isFinalStage(plan)) {
+ // We can not get the whole plan at query preparation phase to detect if current plan is
+ // for writing, so we depend on a tag which is been injected at post resolution phase.
+ // Note: we should still do clean up previous config for non-final stage to avoid such case:
+ // the first statement is write, but the second statement is query.
+ if (conf.getConf(FINAL_STAGE_CONFIG_ISOLATION_WRITE_ONLY) &&
+ !WriteUtils.isWrite(session, plan)) {
+ return plan
+ }
+
+ // set config for final stage
+ session.conf.getAll.filter(_._1.startsWith(FINAL_STAGE_CONFIG_PREFIX)).foreach {
+ case (k, v) =>
+ val sparkConfigKey = s"spark.sql.${k.substring(FINAL_STAGE_CONFIG_PREFIX.length)}"
+ val previousStageConfigKey =
+ s"$PREVIOUS_STAGE_CONFIG_PREFIX${k.substring(FINAL_STAGE_CONFIG_PREFIX.length)}"
+ // store the previous config only if we have not stored, to avoid some query only
+ // have one stage that will overwrite real config.
+ if (!session.sessionState.conf.contains(previousStageConfigKey)) {
+ val originalValue =
+ if (session.conf.getOption(sparkConfigKey).isDefined) {
+ session.sessionState.conf.getConfString(sparkConfigKey)
+ } else {
+ // the default value of config is None, so we need to use a internal tag
+ INTERNAL_UNSET_CONFIG_TAG
+ }
+ logInfo(s"Store config: $sparkConfigKey to previousStage, " +
+ s"original value: $originalValue ")
+ session.sessionState.conf.setConfString(previousStageConfigKey, originalValue)
+ }
+ logInfo(s"For final stage: set $sparkConfigKey = $v.")
+ session.conf.set(sparkConfigKey, v)
+ }
+ } else {
+ // reset config for previous stage
+ session.conf.getAll.filter(_._1.startsWith(PREVIOUS_STAGE_CONFIG_PREFIX)).foreach {
+ case (k, v) =>
+ val sparkConfigKey = s"spark.sql.${k.substring(PREVIOUS_STAGE_CONFIG_PREFIX.length)}"
+ logInfo(s"For previous stage: set $sparkConfigKey = $v.")
+ if (v == INTERNAL_UNSET_CONFIG_TAG) {
+ session.conf.unset(sparkConfigKey)
+ } else {
+ session.conf.set(sparkConfigKey, v)
+ }
+ // unset config so that we do not need to reset configs for every previous stage
+ session.conf.unset(k)
+ }
+ }
+
+ plan
+ }
+
+ /**
+ * Currently formula depend on AQE in Spark 3.1.1, not sure it can work in future.
+ */
+ private def isFinalStage(plan: SparkPlan): Boolean = {
+ var shuffleNum = 0
+ var broadcastNum = 0
+ var reusedNum = 0
+ var queryStageNum = 0
+
+ def collectNumber(p: SparkPlan): SparkPlan = {
+ p transform {
+ case shuffle: ShuffleExchangeLike =>
+ shuffleNum += 1
+ shuffle
+
+ case broadcast: BroadcastExchangeLike =>
+ broadcastNum += 1
+ broadcast
+
+ case reusedExchangeExec: ReusedExchangeExec =>
+ reusedNum += 1
+ reusedExchangeExec
+
+ // query stage is leaf node so we need to transform it manually
+ // compatible with Spark 3.5:
+ // SPARK-42101: table cache is a independent query stage, so do not need include it.
+ case queryStage: QueryStageExec if queryStage.nodeName != "TableCacheQueryStage" =>
+ queryStageNum += 1
+ collectNumber(queryStage.plan)
+ queryStage
+ }
+ }
+ collectNumber(plan)
+
+ if (shuffleNum == 0) {
+ // we don not care about broadcast stage here since it won't change partition number.
+ true
+ } else if (shuffleNum + broadcastNum + reusedNum == queryStageNum) {
+ true
+ } else {
+ false
+ }
+ }
+}
+object FinalStageConfigIsolation {
+ final val SQL_PREFIX = "spark.sql."
+ final val FINAL_STAGE_CONFIG_PREFIX = "spark.sql.finalStage."
+ final val PREVIOUS_STAGE_CONFIG_PREFIX = "spark.sql.previousStage."
+ final val INTERNAL_UNSET_CONFIG_TAG = "__INTERNAL_UNSET_CONFIG_TAG__"
+
+ def getPreviousStageConfigKey(configKey: String): Option[String] = {
+ if (configKey.startsWith(SQL_PREFIX)) {
+ Some(s"$PREVIOUS_STAGE_CONFIG_PREFIX${configKey.substring(SQL_PREFIX.length)}")
+ } else {
+ None
+ }
+ }
+}
+
+case class FinalStageConfigIsolationCleanRule(session: SparkSession) extends Rule[LogicalPlan] {
+ import FinalStageConfigIsolation._
+
+ override def apply(plan: LogicalPlan): LogicalPlan = plan match {
+ case set @ SetCommand(Some((k, Some(_)))) if k.startsWith(SQL_PREFIX) =>
+ checkAndUnsetPreviousStageConfig(k)
+ set
+
+ case reset @ ResetCommand(Some(k)) if k.startsWith(SQL_PREFIX) =>
+ checkAndUnsetPreviousStageConfig(k)
+ reset
+
+ case other => other
+ }
+
+ private def checkAndUnsetPreviousStageConfig(configKey: String): Unit = {
+ getPreviousStageConfigKey(configKey).foreach { previousStageConfigKey =>
+ if (session.sessionState.conf.contains(previousStageConfigKey)) {
+ logInfo(s"For previous stage: unset $previousStageConfigKey")
+ session.conf.unset(previousStageConfigKey)
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSQLConf.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSQLConf.scala
new file mode 100644
index 00000000000..5e8f9f844b2
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSQLConf.scala
@@ -0,0 +1,332 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import org.apache.spark.network.util.{ByteUnit, JavaUtils}
+import org.apache.spark.sql.internal.SQLConf
+import org.apache.spark.sql.internal.SQLConf._
+
+object KyuubiSQLConf {
+
+ final val FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY =
+ "spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes"
+
+ val INSERT_REPARTITION_BEFORE_WRITE =
+ buildConf("spark.sql.optimizer.insertRepartitionBeforeWrite.enabled")
+ .doc("Add repartition node at the top of query plan. An approach of merging small files.")
+ .version("1.2.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ val FORCE_SHUFFLE_BEFORE_JOIN =
+ buildConf("spark.sql.optimizer.forceShuffleBeforeJoin.enabled")
+ .doc("Ensure shuffle node exists before shuffled join (shj and smj) to make AQE " +
+ "`OptimizeSkewedJoin` works (complex scenario join, multi table join).")
+ .version("1.2.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val FINAL_STAGE_CONFIG_ISOLATION =
+ buildConf("spark.sql.optimizer.finalStageConfigIsolation.enabled")
+ .doc("If true, the final stage support use different config with previous stage. " +
+ "The prefix of final stage config key should be `spark.sql.finalStage.`." +
+ "For example, the raw spark config: `spark.sql.adaptive.advisoryPartitionSizeInBytes`, " +
+ "then the final stage config should be: " +
+ "`spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes`.")
+ .version("1.2.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val INSERT_ZORDER_BEFORE_WRITING =
+ buildConf("spark.sql.optimizer.insertZorderBeforeWriting.enabled")
+ .doc("When true, we will follow target table properties to insert zorder or not. " +
+ "The key properties are: 1) kyuubi.zorder.enabled; if this property is true, we will " +
+ "insert zorder before writing data. 2) kyuubi.zorder.cols; string split by comma, we " +
+ "will zorder by these cols.")
+ .version("1.4.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ val ZORDER_GLOBAL_SORT_ENABLED =
+ buildConf("spark.sql.optimizer.zorderGlobalSort.enabled")
+ .doc("When true, we do a global sort using zorder. Note that, it can cause data skew " +
+ "issue if the zorder columns have less cardinality. When false, we only do local sort " +
+ "using zorder.")
+ .version("1.4.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ val REBALANCE_BEFORE_ZORDER =
+ buildConf("spark.sql.optimizer.rebalanceBeforeZorder.enabled")
+ .doc("When true, we do a rebalance before zorder in case data skew. " +
+ "Note that, if the insertion is dynamic partition we will use the partition " +
+ "columns to rebalance.")
+ .version("1.6.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val REBALANCE_ZORDER_COLUMNS_ENABLED =
+ buildConf("spark.sql.optimizer.rebalanceZorderColumns.enabled")
+ .doc(s"When true and ${REBALANCE_BEFORE_ZORDER.key} is true, we do rebalance before " +
+ s"Z-Order. If it's dynamic partition insert, the rebalance expression will include " +
+ s"both partition columns and Z-Order columns.")
+ .version("1.6.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val TWO_PHASE_REBALANCE_BEFORE_ZORDER =
+ buildConf("spark.sql.optimizer.twoPhaseRebalanceBeforeZorder.enabled")
+ .doc(s"When true and ${REBALANCE_BEFORE_ZORDER.key} is true, we do two phase rebalance " +
+ s"before Z-Order for the dynamic partition write. The first phase rebalance using " +
+ s"dynamic partition column; The second phase rebalance using dynamic partition column + " +
+ s"Z-Order columns.")
+ .version("1.6.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val ZORDER_USING_ORIGINAL_ORDERING_ENABLED =
+ buildConf("spark.sql.optimizer.zorderUsingOriginalOrdering.enabled")
+ .doc(s"When true and ${REBALANCE_BEFORE_ZORDER.key} is true, we do sort by " +
+ s"the original ordering i.e. lexicographical order.")
+ .version("1.6.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val WATCHDOG_MAX_PARTITIONS =
+ buildConf("spark.sql.watchdog.maxPartitions")
+ .doc("Set the max partition number when spark scans a data source. " +
+ "Enable maxPartitions Strategy by specifying this configuration. " +
+ "Add maxPartitions Strategy to avoid scan excessive partitions " +
+ "on partitioned table, it's optional that works with defined")
+ .version("1.4.0")
+ .intConf
+ .createOptional
+
+ val WATCHDOG_MAX_FILE_SIZE =
+ buildConf("spark.sql.watchdog.maxFileSize")
+ .doc("Set the maximum size in bytes of files when spark scans a data source. " +
+ "Enable maxFileSize Strategy by specifying this configuration. " +
+ "Add maxFileSize Strategy to avoid scan excessive size of files," +
+ " it's optional that works with defined")
+ .version("1.8.0")
+ .bytesConf(ByteUnit.BYTE)
+ .createOptional
+
+ val DROP_IGNORE_NONEXISTENT =
+ buildConf("spark.sql.optimizer.dropIgnoreNonExistent")
+ .doc("Do not report an error if DROP DATABASE/TABLE/VIEW/FUNCTION/PARTITION specifies " +
+ "a non-existent database/table/view/function/partition")
+ .version("1.5.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val INFER_REBALANCE_AND_SORT_ORDERS =
+ buildConf("spark.sql.optimizer.inferRebalanceAndSortOrders.enabled")
+ .doc("When true, infer columns for rebalance and sort orders from original query, " +
+ "e.g. the join keys from join. It can avoid compression ratio regression.")
+ .version("1.7.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val SKIP_INFER_SORT_ORDERS =
+ buildConf("spark.sql.optimizer.inferRebalanceAndSortOrders.skipSort")
+ .doc(s"When true and `${INFER_REBALANCE_AND_SORT_ORDERS.key}` is true, only infer the " +
+ s"rebalance partition columns and skip inferring the sort orders. Skipping the sort " +
+ s"avoids a local sort before writing when only the file layout from rebalance is wanted.")
+ .version("1.12.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val INFER_REBALANCE_AND_SORT_ORDERS_WITH_CHEAP_COLUMNS =
+ buildConf("spark.sql.optimizer.inferRebalanceAndSortOrders.cheapColumnsOnly")
+ .doc(s"When true and `${INFER_REBALANCE_AND_SORT_ORDERS.key}` is true, only infer the " +
+ s"rebalance and sort columns when all inferred columns are cheap expressions, i.e. " +
+ s"attributes, foldable values, or field extractions over cheap expressions. This avoids " +
+ s"evaluating expensive expressions during the inferred shuffle and sort.")
+ .version("1.12.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ val INFER_REBALANCE_AND_SORT_ORDERS_MAX_COLUMNS =
+ buildConf("spark.sql.optimizer.inferRebalanceAndSortOrdersMaxColumns")
+ .doc("The max columns of inferred columns.")
+ .version("1.7.0")
+ .intConf
+ .checkValue(_ > 0, "must be positive number")
+ .createWithDefault(3)
+
+ val INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE =
+ buildConf("spark.sql.optimizer.insertRepartitionBeforeWriteIfNoShuffle.enabled")
+ .doc("When true, add repartition even if the original plan does not have shuffle.")
+ .version("1.7.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val FINAL_STAGE_CONFIG_ISOLATION_WRITE_ONLY =
+ buildConf("spark.sql.optimizer.finalStageConfigIsolationWriteOnly.enabled")
+ .doc("When true, only enable final stage isolation for writing.")
+ .version("1.7.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ val FINAL_WRITE_STAGE_EAGERLY_KILL_EXECUTORS_ENABLED =
+ buildConf("spark.sql.finalWriteStage.eagerlyKillExecutors.enabled")
+ .doc("When true, eagerly kill redundant executors before running final write stage.")
+ .version("1.8.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val FINAL_WRITE_STAGE_EAGERLY_KILL_EXECUTORS_KILL_ALL =
+ buildConf("spark.sql.finalWriteStage.eagerlyKillExecutors.killAll")
+ .doc("When true, eagerly kill all executors before running final write stage. " +
+ "Mainly for test.")
+ .version("1.8.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val FINAL_WRITE_STAGE_SKIP_KILLING_EXECUTORS_FOR_TABLE_CACHE =
+ buildConf("spark.sql.finalWriteStage.skipKillingExecutorsForTableCache")
+ .doc("When true, skip killing executors if the plan has table caches.")
+ .version("1.8.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ val FINAL_WRITE_STAGE_PARTITION_FACTOR =
+ buildConf("spark.sql.finalWriteStage.retainExecutorsFactor")
+ .doc("If the target executors * factor < active executors, and " +
+ "target executors * factor > min executors, then kill redundant executors.")
+ .version("1.8.0")
+ .doubleConf
+ .checkValue(_ >= 1, "must be bigger than or equal to 1")
+ .createWithDefault(1.2)
+
+ val FINAL_WRITE_STAGE_RESOURCE_ISOLATION_ENABLED =
+ buildConf("spark.sql.finalWriteStage.resourceIsolation.enabled")
+ .doc(
+ "When true, make final write stage resource isolation using custom RDD resource profile.")
+ .version("1.8.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val FINAL_WRITE_STAGE_EXECUTOR_CORES =
+ buildConf("spark.sql.finalWriteStage.executorCores")
+ .doc("Specify the executor core request for final write stage. " +
+ "It would be passed to the RDD resource profile.")
+ .version("1.8.0")
+ .intConf
+ .createOptional
+
+ val FINAL_WRITE_STAGE_EXECUTOR_MEMORY =
+ buildConf("spark.sql.finalWriteStage.executorMemory")
+ .doc("Specify the executor on heap memory request for final write stage. " +
+ "It would be passed to the RDD resource profile.")
+ .version("1.8.0")
+ .stringConf
+ .createOptional
+
+ val FINAL_WRITE_STAGE_EXECUTOR_MEMORY_OVERHEAD =
+ buildConf("spark.sql.finalWriteStage.executorMemoryOverhead")
+ .doc("Specify the executor memory overhead request for final write stage. " +
+ "It would be passed to the RDD resource profile.")
+ .version("1.8.0")
+ .stringConf
+ .createOptional
+
+ val FINAL_WRITE_STAGE_EXECUTOR_OFF_HEAP_MEMORY =
+ buildConf("spark.sql.finalWriteStage.executorOffHeapMemory")
+ .doc("Specify the executor off heap memory request for final write stage. " +
+ "It would be passed to the RDD resource profile.")
+ .version("1.8.0")
+ .stringConf
+ .createOptional
+
+ val REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE =
+ buildConf("spark.sql.adaptive.rebalancePartitionsAdvisoryPartitionSizeInBytes")
+ .doc("The advisory partition size for RebalancePartitions operator. When set, it will be " +
+ "passed to `optAdvisoryPartitionSize` of the RebalancePartitions node. " +
+ "It takes precedence over `spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes`.")
+ .version("1.12.0")
+ .bytesConf(ByteUnit.BYTE)
+ .createOptional
+
+ val REMOVE_REBALANCE_SHUFFLE_ENABLED =
+ buildConf("spark.sql.adaptive.removeRebalanceShuffle.enabled")
+ .doc("When true, the rebalance shuffle injected before writing can be removed at AQE time " +
+ "if the materialized upstream data size makes the shuffle not worthwhile. Only takes " +
+ "effect on a `RebalancePartitions` shuffle that has no partition expressions and whose " +
+ "advisory partition size is larger than " +
+ s"`${SQLConf.ADVISORY_PARTITION_SIZE_IN_BYTES.key}`.")
+ .version("1.12.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val REMOVE_REBALANCE_SHUFFLE_SMALL_PARTITION_SIZE =
+ buildConf("spark.sql.adaptive.removeRebalanceShuffle.smallPartitionSize")
+ .doc("The advisory partition size below which a partition is considered small. In the " +
+ "large-data scenario, the rebalance shuffle is removed when the rebalance input has no " +
+ s"data-reducing operator and its stage size is larger than " +
+ s"`${SQLConf.SHUFFLE_PARTITIONS.key}` * this value.")
+ .version("1.12.0")
+ .bytesConf(ByteUnit.BYTE)
+ .createWithDefault(128L * 1024 * 1024)
+
+ val REMOVE_REBALANCE_SHUFFLE_TOLERABLE_SMALL_FILE_NUM =
+ buildConf("spark.sql.adaptive.removeRebalanceShuffle.tolerableSmallFileNum")
+ .doc("The tolerable number of small output files for the small-data scenario. When the " +
+ "rebalance input has no data-expanding operator and its stage size is smaller than " +
+ s"`${SQLConf.ADVISORY_PARTITION_SIZE_IN_BYTES.key}` * this value, the rebalance shuffle " +
+ "is removed.")
+ .version("1.12.0")
+ .intConf
+ .createWithDefault(3)
+
+ val DYNAMIC_SHUFFLE_PARTITIONS =
+ buildConf("spark.sql.optimizer.dynamicShufflePartitions")
+ .doc("If true, adjust the number of shuffle partitions dynamically based on the job" +
+ " input size. The new number of partitions is the maximum input size" +
+ " divided by `spark.sql.adaptive.advisoryPartitionSizeInBytes`.")
+ .version("1.9.0")
+ .booleanConf
+ .createWithDefault(false)
+
+ val DYNAMIC_SHUFFLE_PARTITIONS_MAX_NUM =
+ buildConf("spark.sql.optimizer.dynamicShufflePartitions.maxNum")
+ .doc("The maximum partition number of DynamicShufflePartitions.")
+ .version("1.9.0")
+ .intConf
+ .createWithDefault(2000)
+
+ val SCRIPT_TRANSFORMATION_ENABLED =
+ buildConf("spark.sql.execution.scriptTransformation.enabled")
+ .doc("When false, script transformation is not allowed.")
+ .version("1.9.0")
+ .booleanConf
+ .createWithDefault(true)
+
+ def getAdvisoryPartitionSize(conf: SQLConf): Option[Long] = {
+ conf.getConf(REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE).orElse {
+ if (conf.contains(FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY)) {
+ Some(JavaUtils.byteStringAs(
+ conf.getConfString(FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY),
+ ByteUnit.BYTE))
+ } else {
+ None
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSQLExtensionException.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSQLExtensionException.scala
new file mode 100644
index 00000000000..88c5a988fd9
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSQLExtensionException.scala
@@ -0,0 +1,28 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import java.sql.SQLException
+
+class KyuubiSQLExtensionException(reason: String, cause: Throwable)
+ extends SQLException(reason, cause) {
+
+ def this(reason: String) = {
+ this(reason, null)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLAstBuilder.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLAstBuilder.scala
new file mode 100644
index 00000000000..79a7d297f50
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLAstBuilder.scala
@@ -0,0 +1,190 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import scala.collection.mutable.ListBuffer
+import scala.jdk.CollectionConverters._
+
+import org.antlr.v4.runtime.ParserRuleContext
+import org.antlr.v4.runtime.misc.Interval
+import org.antlr.v4.runtime.tree.ParseTree
+import org.apache.spark.sql.catalyst.SQLConfHelper
+import org.apache.spark.sql.catalyst.analysis.{UnresolvedAttribute, UnresolvedRelation, UnresolvedStar}
+import org.apache.spark.sql.catalyst.expressions._
+import org.apache.spark.sql.catalyst.parser.ParserUtils.withOrigin
+import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, Project, Sort}
+
+import org.apache.kyuubi.sql.KyuubiSparkSQLParser._
+import org.apache.kyuubi.sql.zorder.{OptimizeZorderStatement, Zorder}
+
+class KyuubiSparkSQLAstBuilder extends KyuubiSparkSQLBaseVisitor[AnyRef] with SQLConfHelper {
+
+ def buildOptimizeStatement(
+ unparsedPredicateOptimize: UnparsedPredicateOptimize,
+ parseExpression: String => Expression): LogicalPlan = {
+
+ val UnparsedPredicateOptimize(tableIdent, tablePredicate, orderExpr) =
+ unparsedPredicateOptimize
+
+ val predicate = tablePredicate.map(parseExpression)
+ verifyPartitionPredicates(predicate)
+ val table = UnresolvedRelation(tableIdent)
+ val tableWithFilter = predicate match {
+ case Some(expr) => Filter(expr, table)
+ case None => table
+ }
+ val query =
+ Sort(
+ SortOrder(orderExpr, Ascending, NullsLast, Seq.empty) :: Nil,
+ conf.getConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED),
+ Project(Seq(UnresolvedStar(None)), tableWithFilter))
+ OptimizeZorderStatement(tableIdent, query)
+ }
+
+ private def verifyPartitionPredicates(predicates: Option[Expression]): Unit = {
+ predicates.foreach {
+ case p if !isLikelySelective(p) =>
+ throw new KyuubiSQLExtensionException(s"unsupported partition predicates: ${p.sql}")
+ case _ =>
+ }
+ }
+
+ /**
+ * Forked from Apache Spark's org.apache.spark.sql.catalyst.expressions.PredicateHelper
+ * The `PredicateHelper.isLikelySelective()` is available since Spark-3.3, forked for Spark
+ * that is lower than 3.3.
+ *
+ * Returns whether an expression is likely to be selective
+ */
+ private def isLikelySelective(e: Expression): Boolean = e match {
+ case Not(expr) => isLikelySelective(expr)
+ case And(l, r) => isLikelySelective(l) || isLikelySelective(r)
+ case Or(l, r) => isLikelySelective(l) && isLikelySelective(r)
+ case _: StringRegexExpression => true
+ case _: BinaryComparison => true
+ case _: In | _: InSet => true
+ case _: StringPredicate => true
+ case BinaryPredicate(_) => true
+ case _: MultiLikeBase => true
+ case _ => false
+ }
+
+ private object BinaryPredicate {
+ def unapply(expr: Expression): Option[Expression] = expr match {
+ case _: Contains => Option(expr)
+ case _: StartsWith => Option(expr)
+ case _: EndsWith => Option(expr)
+ case _ => None
+ }
+ }
+
+ /**
+ * Create an expression from the given context. This method just passes the context on to the
+ * visitor and only takes care of typing (We assume that the visitor returns an Expression here).
+ */
+ protected def expression(ctx: ParserRuleContext): Expression = typedVisit(ctx)
+
+ protected def multiPart(ctx: ParserRuleContext): Seq[String] = typedVisit(ctx)
+
+ override def visitSingleStatement(ctx: SingleStatementContext): LogicalPlan = {
+ visit(ctx.statement()).asInstanceOf[LogicalPlan]
+ }
+
+ override def visitOptimizeZorder(
+ ctx: OptimizeZorderContext): UnparsedPredicateOptimize = withOrigin(ctx) {
+ val tableIdent = multiPart(ctx.multipartIdentifier())
+
+ val predicate = Option(ctx.whereClause())
+ .map(_.partitionPredicate)
+ .map(extractRawText(_))
+
+ val zorderCols = ctx.zorderClause().order.asScala
+ .map(visitMultipartIdentifier)
+ .map(UnresolvedAttribute(_))
+ .toSeq
+
+ val orderExpr =
+ if (zorderCols.length == 1) {
+ zorderCols.head
+ } else {
+ Zorder(zorderCols)
+ }
+ UnparsedPredicateOptimize(tableIdent, predicate, orderExpr)
+ }
+
+ override def visitPassThrough(ctx: PassThroughContext): LogicalPlan = null
+
+ override def visitMultipartIdentifier(ctx: MultipartIdentifierContext): Seq[String] =
+ withOrigin(ctx) {
+ ctx.parts.asScala.map(typedVisit[String]).toSeq
+ }
+
+ override def visitIdentifier(ctx: IdentifierContext): String = {
+ withOrigin(ctx) {
+ ctx.strictIdentifier() match {
+ case quotedContext: QuotedIdentifierAlternativeContext =>
+ typedVisit[String](quotedContext)
+ case _ => ctx.getText
+ }
+ }
+ }
+
+ override def visitQuotedIdentifier(ctx: QuotedIdentifierContext): String = {
+ withOrigin(ctx) {
+ ctx.BACKQUOTED_IDENTIFIER().getText.stripPrefix("`").stripSuffix("`").replace("``", "`")
+ }
+ }
+
+ override def visitZorderClause(ctx: ZorderClauseContext): Seq[UnresolvedAttribute] =
+ withOrigin(ctx) {
+ val res = ListBuffer[UnresolvedAttribute]()
+ ctx.multipartIdentifier().forEach { identifier =>
+ res += UnresolvedAttribute(identifier.parts.asScala.map(typedVisit[String]).toSeq)
+ }
+ res.toSeq
+ }
+
+ private def typedVisit[T](ctx: ParseTree): T = {
+ ctx.accept(this).asInstanceOf[T]
+ }
+
+ private def extractRawText(exprContext: ParserRuleContext): String = {
+ // Extract the raw expression which will be parsed later
+ exprContext.getStart.getInputStream.getText(new Interval(
+ exprContext.getStart.getStartIndex,
+ exprContext.getStop.getStopIndex))
+ }
+}
+
+/**
+ * a logical plan contains an unparsed expression that will be parsed by spark.
+ */
+trait UnparsedExpressionLogicalPlan extends LogicalPlan {
+ override def output: Seq[Attribute] = throw new UnsupportedOperationException()
+
+ override def children: Seq[LogicalPlan] = throw new UnsupportedOperationException()
+
+ override def withNewChildrenInternal(
+ newChildren: IndexedSeq[LogicalPlan]): LogicalPlan =
+ throw new UnsupportedOperationException()
+}
+
+case class UnparsedPredicateOptimize(
+ tableIdent: Seq[String],
+ tablePredicate: Option[String],
+ orderExpr: Expression) extends UnparsedExpressionLogicalPlan {}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLExtension.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLExtension.scala
new file mode 100644
index 00000000000..4de056d594b
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLExtension.scala
@@ -0,0 +1,63 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import org.apache.spark.sql.{FinalStageResourceManager, InjectCustomResourceProfile, RemoveRebalanceShuffle, SparkSessionExtensions}
+
+import org.apache.kyuubi.sql.watchdog.{KyuubiUnsupportedOperationsCheck, MaxScanStrategy}
+import org.apache.kyuubi.sql.zorder.{InsertZorderBeforeWritingDatasource, InsertZorderBeforeWritingHive, ResolveZorder}
+
+// scalastyle:off line.size.limit
+/**
+ * Depend on Spark SQL Extension framework, we can use this extension follow steps
+ * 1. move this jar into $SPARK_HOME/jars
+ * 2. add config into `spark-defaults.conf`: `spark.sql.extensions=org.apache.kyuubi.sql.KyuubiSparkSQLExtension`
+ */
+// scalastyle:on line.size.limit
+class KyuubiSparkSQLExtension extends (SparkSessionExtensions => Unit) {
+ override def apply(extensions: SparkSessionExtensions): Unit = {
+ // inject zorder parser and related rules
+ extensions.injectParser { case (_, parser) => new SparkKyuubiSparkSQLParser(parser) }
+ extensions.injectResolutionRule(ResolveZorder)
+
+ // Note that:
+ // InsertZorderBeforeWriting* should be applied before RebalanceBeforeWriting*
+ // because we can only apply one of them (i.e. GlobalSort or Rebalance)
+ extensions.injectPostHocResolutionRule(InsertZorderBeforeWritingDatasource)
+ extensions.injectPostHocResolutionRule(InsertZorderBeforeWritingHive)
+ extensions.injectPostHocResolutionRule(FinalStageConfigIsolationCleanRule)
+ extensions.injectPostHocResolutionRule(RebalanceBeforeWritingDatasource)
+ extensions.injectPostHocResolutionRule(RebalanceBeforeWritingHive)
+ extensions.injectPostHocResolutionRule(DropIgnoreNonexistent)
+
+ // watchdog extension
+ extensions.injectCheckRule(_ => KyuubiUnsupportedOperationsCheck)
+ extensions.injectPlannerStrategy(MaxScanStrategy)
+
+ extensions.injectQueryStagePrepRule(_ => InsertShuffleNodeBeforeJoin)
+ extensions.injectQueryStagePrepRule(DynamicShufflePartitions)
+ extensions.injectQueryStagePrepRule(FinalStageConfigIsolation(_))
+ extensions.injectQueryStagePrepRule(FinalStageResourceManager(_))
+ extensions.injectQueryStagePrepRule(InjectCustomResourceProfile)
+
+ // Remove the rebalance shuffle at AQE time when the materialized upstream data size makes it
+ // not worthwhile. It runs as an AQE runtime optimizer rule so that it can rewrite the logical
+ // plan based on the materialized shuffle statistics.
+ extensions.injectRuntimeOptimizerRule(RemoveRebalanceShuffle(_))
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLParser.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLParser.scala
new file mode 100644
index 00000000000..17f55e7ba18
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/KyuubiSparkSQLParser.scala
@@ -0,0 +1,148 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import scala.jdk.CollectionConverters._
+
+import org.antlr.v4.runtime._
+import org.antlr.v4.runtime.atn.PredictionMode
+import org.antlr.v4.runtime.misc.{Interval, ParseCancellationException}
+import org.apache.spark.SparkThrowable
+import org.apache.spark.sql.catalyst.{FunctionIdentifier, SQLConfHelper, TableIdentifier}
+import org.apache.spark.sql.catalyst.expressions.Expression
+import org.apache.spark.sql.catalyst.parser.{ParseErrorListener, ParseException, ParserInterface, PostProcessor}
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
+import org.apache.spark.sql.catalyst.trees.WithOrigin
+import org.apache.spark.sql.types.{DataType, StructType}
+
+abstract class KyuubiSparkSQLParserBase extends ParserInterface with SQLConfHelper {
+ def delegate: ParserInterface
+ def astBuilder: KyuubiSparkSQLAstBuilder
+
+ override def parsePlan(sqlText: String): LogicalPlan = parse(sqlText) { parser =>
+ astBuilder.visit(parser.singleStatement()) match {
+ case optimize: UnparsedPredicateOptimize =>
+ astBuilder.buildOptimizeStatement(optimize, delegate.parseExpression)
+ case plan: LogicalPlan => plan
+ case _ => delegate.parsePlan(sqlText)
+ }
+ }
+
+ protected def parse[T](command: String)(toResult: KyuubiSparkSQLParser => T): T = {
+ val lexer = new KyuubiSparkSQLLexer(
+ new UpperCaseCharStream(CharStreams.fromString(command)))
+ lexer.removeErrorListeners()
+ lexer.addErrorListener(ParseErrorListener)
+
+ val tokenStream = new CommonTokenStream(lexer)
+ val parser = new KyuubiSparkSQLParser(tokenStream)
+ parser.addParseListener(PostProcessor)
+ parser.removeErrorListeners()
+ parser.addErrorListener(ParseErrorListener)
+
+ try {
+ try {
+ // first, try parsing with potentially faster SLL mode
+ parser.getInterpreter.setPredictionMode(PredictionMode.SLL)
+ toResult(parser)
+ } catch {
+ case _: ParseCancellationException =>
+ // if we fail, parse with LL mode
+ tokenStream.seek(0) // rewind input stream
+ parser.reset()
+
+ // Try Again.
+ parser.getInterpreter.setPredictionMode(PredictionMode.LL)
+ toResult(parser)
+ }
+ } catch {
+ case e: ParseException if e.command.isDefined =>
+ throw e
+ case e: ParseException =>
+ throw e.withCommand(command)
+ case e: SparkThrowable with WithOrigin =>
+ throw new ParseException(
+ command = Option(command),
+ start = e.origin,
+ errorClass = e.getCondition,
+ messageParameters = e.getMessageParameters.asScala.toMap,
+ queryContext = e.getQueryContext)
+ }
+ }
+
+ override def parseExpression(sqlText: String): Expression = {
+ delegate.parseExpression(sqlText)
+ }
+
+ override def parseTableIdentifier(sqlText: String): TableIdentifier = {
+ delegate.parseTableIdentifier(sqlText)
+ }
+
+ override def parseFunctionIdentifier(sqlText: String): FunctionIdentifier = {
+ delegate.parseFunctionIdentifier(sqlText)
+ }
+
+ override def parseMultipartIdentifier(sqlText: String): Seq[String] = {
+ delegate.parseMultipartIdentifier(sqlText)
+ }
+
+ override def parseTableSchema(sqlText: String): StructType = {
+ delegate.parseTableSchema(sqlText)
+ }
+
+ override def parseDataType(sqlText: String): DataType = {
+ delegate.parseDataType(sqlText)
+ }
+
+ // SPARK-37266 (3.3.0)
+ override def parseQuery(sqlText: String): LogicalPlan = {
+ delegate.parseQuery(sqlText)
+ }
+
+ // SPARK-51439 (4.0.0)
+ override def parseRoutineParam(sqlText: String): StructType = {
+ delegate.parseRoutineParam(sqlText)
+ }
+}
+
+class SparkKyuubiSparkSQLParser(
+ override val delegate: ParserInterface)
+ extends KyuubiSparkSQLParserBase {
+ def astBuilder: KyuubiSparkSQLAstBuilder = new KyuubiSparkSQLAstBuilder
+}
+
+/* Copied from Apache Spark's to avoid dependency on Spark Internals */
+class UpperCaseCharStream(wrapped: CodePointCharStream) extends CharStream {
+ override def consume(): Unit = wrapped.consume()
+ override def getSourceName(): String = wrapped.getSourceName
+ override def index(): Int = wrapped.index
+ override def mark(): Int = wrapped.mark
+ override def release(marker: Int): Unit = wrapped.release(marker)
+ override def seek(where: Int): Unit = wrapped.seek(where)
+ override def size(): Int = wrapped.size
+
+ override def getText(interval: Interval): String = wrapped.getText(interval)
+
+ // scalastyle:off
+ override def LA(i: Int): Int = {
+ val la = wrapped.LA(i)
+ if (la == 0 || la == IntStream.EOF) la
+ else Character.toUpperCase(la)
+ }
+ // scalastyle:on
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/RebalanceBeforeWriting.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/RebalanceBeforeWriting.scala
new file mode 100644
index 00000000000..5a8c43d969b
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/RebalanceBeforeWriting.scala
@@ -0,0 +1,165 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import scala.annotation.tailrec
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.expressions.{Ascending, Attribute, SortOrder}
+import org.apache.spark.sql.catalyst.plans.logical._
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.datasources.InsertIntoHadoopFsRelationCommand
+import org.apache.spark.sql.hive.execution.InsertIntoHiveTable
+import org.apache.spark.sql.internal.StaticSQLConf
+
+trait RebalanceBeforeWritingBase extends Rule[LogicalPlan] {
+
+ private def hasBenefit(plan: LogicalPlan): Boolean = {
+ def probablyHasShuffle: Boolean = plan.exists {
+ case _: Join => true
+ case _: Aggregate => true
+ case _: Distinct => true
+ case _: Deduplicate => true
+ case _: Window => true
+ case s: Sort if s.global => true
+ case _: RepartitionOperation => true
+ case _: RebalancePartitions => true
+ case _: GlobalLimit => true
+ case _ => false
+ }
+
+ conf.getConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE) || probablyHasShuffle
+ }
+
+ def canInsertRebalance(plan: LogicalPlan): Boolean = {
+ @tailrec
+ def canInsert(p: LogicalPlan): Boolean = p match {
+ case Project(_, child) => canInsert(child)
+ case SubqueryAlias(_, child) => canInsert(child)
+ case Limit(_, _) => false
+ case _: Sort => false
+ case _: RepartitionOperation => false
+ case _: RebalancePartitions => false
+ case _ => true
+ }
+
+ // 1. make sure AQE is enabled, otherwise it is no meaning to add a shuffle
+ // 2. make sure it does not break the semantics of original plan
+ // 3. try to avoid adding a shuffle if it has potential performance regression
+ conf.adaptiveExecutionEnabled && canInsert(plan) && hasBenefit(plan)
+ }
+
+ def buildRebalance(
+ dynamicPartitionColumns: Seq[Attribute],
+ query: LogicalPlan): LogicalPlan = {
+ val advisoryPartitionSize = KyuubiSQLConf.getAdvisoryPartitionSize(conf)
+ if (!conf.getConf(KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS) ||
+ dynamicPartitionColumns.nonEmpty) {
+ RebalancePartitions(
+ dynamicPartitionColumns,
+ query,
+ optAdvisoryPartitionSize = advisoryPartitionSize)
+ } else {
+ val maxColumns = conf.getConf(KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS_MAX_COLUMNS)
+ val onlyInferWithCheapColumns = conf.getConf(
+ KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS_WITH_CHEAP_COLUMNS)
+ val inferred = InferRebalanceAndSortOrders.infer(query, onlyInferWithCheapColumns)
+ if (inferred.isDefined) {
+ val (partitioning, ordering) = inferred.get
+ val rebalance = RebalancePartitions(
+ partitioning.take(maxColumns),
+ query,
+ optAdvisoryPartitionSize = advisoryPartitionSize)
+ val skipInferOrder = conf.getConf(KyuubiSQLConf.SKIP_INFER_SORT_ORDERS)
+ if (!skipInferOrder && ordering.nonEmpty) {
+ val sortOrders = ordering.take(maxColumns).map(o => SortOrder(o, Ascending))
+ Sort(sortOrders, false, rebalance)
+ } else {
+ rebalance
+ }
+ } else {
+ RebalancePartitions(
+ dynamicPartitionColumns,
+ query,
+ optAdvisoryPartitionSize = advisoryPartitionSize)
+ }
+ }
+ }
+}
+
+/**
+ * For datasource table, there two commands can write data to table
+ * 1. InsertIntoHadoopFsRelationCommand
+ * 2. CreateDataSourceTableAsSelectCommand
+ * This rule add a repartition node between write and query
+ */
+case class RebalanceBeforeWritingDatasource(session: SparkSession)
+ extends RebalanceBeforeWritingBase {
+
+ override def apply(plan: LogicalPlan): LogicalPlan = {
+ if (conf.getConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE)) {
+ addRebalance(plan)
+ } else {
+ plan
+ }
+ }
+
+ private def addRebalance(plan: LogicalPlan): LogicalPlan = plan match {
+ case i @ InsertIntoHadoopFsRelationCommand(_, sp, _, pc, bucket, _, _, query, _, _, _, _)
+ if query.resolved && bucket.isEmpty && canInsertRebalance(query) =>
+ val dynamicPartitionColumns = pc.filterNot(attr => sp.contains(attr.name))
+ i.copy(query = buildRebalance(dynamicPartitionColumns, query))
+
+ case u @ Union(children, _, _) =>
+ u.copy(children = children.map(addRebalance))
+
+ case _ => plan
+ }
+}
+
+/**
+ * For Hive table, there two commands can write data to table
+ * 1. InsertIntoHiveTable
+ * 2. CreateHiveTableAsSelectCommand
+ * This rule add a repartition node between write and query
+ */
+case class RebalanceBeforeWritingHive(session: SparkSession)
+ extends RebalanceBeforeWritingBase {
+
+ override def apply(plan: LogicalPlan): LogicalPlan = {
+ if (conf.getConf(StaticSQLConf.CATALOG_IMPLEMENTATION) == "hive" &&
+ conf.getConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE)) {
+ addRebalance(plan)
+ } else {
+ plan
+ }
+ }
+
+ private def addRebalance(plan: LogicalPlan): LogicalPlan = plan match {
+ case i @ InsertIntoHiveTable(table, partition, query, _, _, _, _, _, _, _, _)
+ if query.resolved && table.bucketSpec.isEmpty && canInsertRebalance(query) =>
+ val dynamicPartitionColumns = partition.filter(_._2.isEmpty).keys
+ .flatMap(name => query.output.find(_.name == name)).toSeq
+ i.copy(query = buildRebalance(dynamicPartitionColumns, query))
+
+ case u @ Union(children, _, _) =>
+ u.copy(children = children.map(addRebalance))
+
+ case _ => plan
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/RemoveRebalanceShuffle.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/RemoveRebalanceShuffle.scala
new file mode 100644
index 00000000000..8298e909c05
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/RemoveRebalanceShuffle.scala
@@ -0,0 +1,214 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.catalyst.plans.{Inner, LeftExistence}
+import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Expand, Filter, Generate, GlobalLimit, Join, LocalLimit, LogicalPlan, Offset, Project, RebalancePartitions, Sample, Sort, UnaryNode, Union, WindowGroupLimit}
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.adaptive.{LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec}
+import org.apache.spark.sql.execution.datasources.WriteFiles
+import org.apache.spark.sql.internal.SQLConf
+
+import org.apache.kyuubi.sql.KyuubiSQLConf._
+
+/**
+ * An AQE runtime-optimizer rule (injected via `injectRuntimeOptimizerRule`) that removes the
+ * `RebalancePartitions` operator injected before writing when the materialized rebalance input
+ * makes the rebalance shuffle not worthwhile. Because it runs in the AQE optimizer on the logical
+ * plan, the removal is a plain logical rewrite (replace `RebalancePartitions` with its child) and
+ * the physical planner re-plans without the rebalance shuffle.
+ *
+ * This rule only applies when all partition columns of the write have static values assigned
+ * ([[WriteFiles.staticPartitions]].size == [[WriteFiles.partitionColumns]].size), i.e., static
+ * partition inserts or non-partitioned tables. For dynamic partition inserts, where partition
+ * values come from the SELECT clause, the rebalance helps cluster data by partition values and
+ * is left untouched.
+ *
+ * Starting from the [[WriteFiles]] node, the rule walks down through at most [[Project]] and
+ * local [[Sort]] operators to reach the `RebalancePartitions`. Any other operator in between
+ * means the rebalance does not directly feed the write, so it is left untouched. A
+ * `RebalancePartitions` is only removed when:
+ * 1. [[KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED]] is on;
+ * 2. its `optAdvisoryPartitionSize` is present and larger than the session
+ * `spark.sql.adaptive.advisoryPartitionSizeInBytes` (otherwise the smaller advisory size
+ * suggests the user intentionally wants to balance data into smaller partitions).
+ *
+ * The `partitionExpressions` of the `RebalancePartitions` may be either empty (from
+ * `InsertAdaptor`, yielding physical `shuffleOrigin` `REBALANCE_PARTITIONS_BY_NONE`) or
+ * non-empty (inferred by `InferRebalanceAndSortOrders`, yielding
+ * `REBALANCE_PARTITIONS_BY_COL`). For the non-empty case, the local [[Sort]] that was
+ * injected alongside the rebalance serves no purpose after removal and is cleaned up by
+ * Spark's `RemoveRedundantSorts` rule.
+ *
+ * The rebalance input is split into independent "query stage groups" by
+ * [[collectQueryStageGroupSize]]: only sub-plans whose leaves are all materialized
+ * [[LogicalQueryStage]] exchange stages contribute a group of stage sizes; a `Union` produces one
+ * group per branch. The size used per stage is its runtime `sizeInBytes`. Every group must be
+ * removable for the shuffle to be removed. A group is removable when either:
+ * - Large data: the rebalance input has no data-reducing operator and the group's representative
+ * size (`max(maxStageSize, groupSum / 2)`) is larger than
+ * `numShufflePartitions * [[KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_SMALL_PARTITION_SIZE]]`; or
+ * - Small data: the rebalance input has no data-expanding operator and the representative size is
+ * smaller than `advisoryPartitionSize *
+ * [[KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_TOLERABLE_SMALL_FILE_NUM]]`.
+ */
+case class RemoveRebalanceShuffle(session: SparkSession) extends Rule[LogicalPlan] {
+
+ override def apply(plan: LogicalPlan): LogicalPlan = {
+ if (!conf.getConf(REMOVE_REBALANCE_SHUFFLE_ENABLED)) {
+ return plan
+ }
+
+ plan.transformDown {
+ case write: WriteFiles if write.staticPartitions.size == write.partitionColumns.size =>
+ // only remove the rebalance when all partition columns have static values assigned,
+ // i.e., static partition inserts or non-partitioned tables. For dynamic partition
+ // inserts the rebalance helps cluster data by partition values to reduce file count.
+ write.withNewChildren(Seq(removeRebalance(write.child)))
+ }
+ }
+
+ /**
+ * Walk down from the write through at most [[Project]] and local [[Sort]] operators to find the
+ * `RebalancePartitions` that feeds the write, and remove it (replace with its child) if eligible.
+ * Any other operator in between means the rebalance does not directly feed the write, so it is
+ * left untouched.
+ */
+ private def removeRebalance(plan: LogicalPlan): LogicalPlan = plan match {
+ case p: Project => p.withNewChildren(Seq(removeRebalance(p.child)))
+ case s: Sort if !s.global => s.withNewChildren(Seq(removeRebalance(s.child)))
+ case RebalancePartitions(_, child, _, optAdvisoryPartitionSize)
+ if optAdvisoryPartitionSize.isDefined && shouldRemove(
+ child,
+ optAdvisoryPartitionSize.get) =>
+ // We have already skipped remove for dynamic partition insert table before.
+ // Here if the partitionExpressions is non-empty, it should be inferred by
+ // `InferRebalanceAndSortOrders`, and that local sort should be removed by Spark rule
+ // `RemoveRedundantSorts`.
+ child
+ case other => other
+ }
+
+ /**
+ * Collect the runtime sizes of the materialized query stages below the rebalance, grouped so
+ * that stages that must be considered together stay in one group. A sub-plan contributes a group
+ * only when all of its leaves are [[LogicalQueryStage]] exchange stages and every such stage is
+ * materialized (a broadcast stage needs `isMaterialized`; a shuffle stage additionally needs
+ * `mapStats`); otherwise the sub-plan is skipped. A [[Union]] yields one group per child so the
+ * branches are evaluated independently. Returns an empty sequence when no group is found.
+ */
+ private def collectQueryStageGroupSize(plan: LogicalPlan): Seq[Seq[Long]] = plan match {
+ case u: UnaryNode => collectQueryStageGroupSize(u.child)
+ case u: Union => u.children.flatMap(collectQueryStageGroupSize)
+ case p if p.collectLeaves().forall(_.isInstanceOf[LogicalQueryStage]) =>
+ val stages = p.collect {
+ case stage: LogicalQueryStage => stage
+ }
+ val allValid = stages.map(_.physicalPlan).forall {
+ case s: ShuffleQueryStageExec => s.isMaterialized && s.mapStats.isDefined
+ case s: QueryStageExec => s.isMaterialized
+ case _ => false
+ }
+ if (allValid) {
+ Seq(stages.map(_.physicalPlan.asInstanceOf[QueryStageExec])
+ .map(_.getRuntimeStatistics.sizeInBytes.min(Long.MaxValue).toLong))
+ } else {
+ Seq.empty
+ }
+ case _ => Seq.empty
+ }
+
+ /**
+ * Decide whether the rebalance feeding the write can be removed. The rebalance advisory size must
+ * be larger than the session advisory size, every query stage group below the rebalance must be
+ * materialized, and every group must be removable under either the large-data (no data-reducing
+ * operator) or the small-data (no data-expanding operator) condition.
+ */
+ private def shouldRemove(child: LogicalPlan, rebalanceAdvisorySize: Long): Boolean = {
+ val sessionAdvisorySize = conf.getConf(SQLConf.ADVISORY_PARTITION_SIZE_IN_BYTES)
+ // Maybe intentionally by user to balance data
+ if (rebalanceAdvisorySize <= sessionAdvisorySize) {
+ logInfo(s"Keep rebalance shuffle since the rebalance advisory partition size " +
+ s"($rebalanceAdvisorySize bytes) is not larger than the session advisory partition size " +
+ s"($sessionAdvisorySize bytes).")
+ return false
+ }
+
+ val groupedStageSize = collectQueryStageGroupSize(child)
+ if (groupedStageSize.isEmpty || groupedStageSize.exists(_.isEmpty)) {
+ logInfo("Keep rebalance shuffle since no materialized query stage group is found below " +
+ "the rebalance.")
+ return false
+ }
+
+ val shufflePartitions = conf.numShufflePartitions
+ val smallPartitionSizeThreshold = conf.getConf(REMOVE_REBALANCE_SHUFFLE_SMALL_PARTITION_SIZE)
+ val hasReducingOperators = hasReducingOperator(child)
+ val hasExpandingOperators = hasExpandingOperator(child)
+ val tolerableSmallFileNum = conf.getConf(REMOVE_REBALANCE_SHUFFLE_TOLERABLE_SMALL_FILE_NUM)
+ val coalesceEnabled = conf.coalesceShufflePartitionsEnabled
+
+ def shouldRemoveOneGroup(stageSizes: Seq[Long]): Boolean = {
+ val maxStageSize = stageSizes.max.max(stageSizes.sum / 2)
+ (!hasReducingOperators && maxStageSize > shufflePartitions * smallPartitionSizeThreshold) ||
+ (coalesceEnabled && !hasExpandingOperators &&
+ maxStageSize < sessionAdvisorySize * tolerableSmallFileNum)
+ }
+
+ val remove = groupedStageSize.forall(shouldRemoveOneGroup)
+ logInfo(s"${if (remove) "Remove" else "Keep"} rebalance shuffle. " +
+ s"rebalanceAdvisorySize=$rebalanceAdvisorySize, sessionAdvisorySize=$sessionAdvisorySize, " +
+ s"shufflePartitions=$shufflePartitions, smallPartitionSize=$smallPartitionSizeThreshold, " +
+ s"tolerableSmallFileNum=$tolerableSmallFileNum," +
+ s"hasReducingOperators=$hasReducingOperators," +
+ s"hasExpandingOperators=$hasExpandingOperators, " +
+ s"groupedStageSize=${groupedStageSize.map(_.mkString("[", ",", "]")).mkString(", ")}.")
+ remove
+ }
+
+ /** Whether the rebalance input contains an operator that may reduce the data size. */
+ private def hasReducingOperator(plan: LogicalPlan): Boolean = {
+ plan.exists {
+ case _: Aggregate => true
+ case _: Filter => true
+ case _: Sample => true
+ case _: Offset => true
+ case _: GlobalLimit => true
+ case _: LocalLimit => true
+ case _: WindowGroupLimit => true
+ case Join(_, _, joinType, _, _) =>
+ joinType match {
+ case LeftExistence(_) => true
+ case Inner => true
+ case _ => false
+ }
+ case _ => false
+ }
+ }
+
+ /** Whether the rebalance input contains an operator that may expand the data size. */
+ private def hasExpandingOperator(plan: LogicalPlan): Boolean = {
+ plan.exists {
+ case _: Generate => true
+ case _: Expand => true
+ case Join(_, _, LeftExistence(_), _, _) => false
+ case _: Join => true
+ case _ => false
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/WriteUtils.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/WriteUtils.scala
new file mode 100644
index 00000000000..89dd8319480
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/WriteUtils.scala
@@ -0,0 +1,34 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.execution.{SparkPlan, UnionExec}
+import org.apache.spark.sql.execution.command.DataWritingCommandExec
+import org.apache.spark.sql.execution.datasources.v2.V2TableWriteExec
+
+object WriteUtils {
+ def isWrite(session: SparkSession, plan: SparkPlan): Boolean = {
+ plan match {
+ case _: DataWritingCommandExec => true
+ case _: V2TableWriteExec => true
+ case u: UnionExec if u.children.nonEmpty => u.children.forall(isWrite(session, _))
+ case _ => false
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/KyuubiUnsupportedOperationsCheck.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/KyuubiUnsupportedOperationsCheck.scala
new file mode 100644
index 00000000000..2b4d3940ada
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/KyuubiUnsupportedOperationsCheck.scala
@@ -0,0 +1,35 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.watchdog
+
+import org.apache.spark.sql.catalyst.SQLConfHelper
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, ScriptTransformation}
+
+import org.apache.kyuubi.sql.{KyuubiSQLConf, KyuubiSQLExtensionException}
+
+object KyuubiUnsupportedOperationsCheck extends (LogicalPlan => Unit) with SQLConfHelper {
+ override def apply(plan: LogicalPlan): Unit =
+ conf.getConf(KyuubiSQLConf.SCRIPT_TRANSFORMATION_ENABLED) match {
+ case false => plan foreach {
+ case _: ScriptTransformation =>
+ throw new KyuubiSQLExtensionException("Script transformation is not allowed")
+ case _ =>
+ }
+ case true =>
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/KyuubiWatchDogException.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/KyuubiWatchDogException.scala
new file mode 100644
index 00000000000..e44309192a9
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/KyuubiWatchDogException.scala
@@ -0,0 +1,30 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.watchdog
+
+import org.apache.kyuubi.sql.KyuubiSQLExtensionException
+
+final class MaxPartitionExceedException(
+ private val reason: String = "",
+ private val cause: Throwable = None.orNull)
+ extends KyuubiSQLExtensionException(reason, cause)
+
+final class MaxFileSizeExceedException(
+ private val reason: String = "",
+ private val cause: Throwable = None.orNull)
+ extends KyuubiSQLExtensionException(reason, cause)
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/MaxScanStrategy.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/MaxScanStrategy.scala
new file mode 100644
index 00000000000..4807101186b
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/watchdog/MaxScanStrategy.scala
@@ -0,0 +1,342 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.watchdog
+
+import org.apache.hadoop.fs.Path
+import org.apache.spark.sql.{PruneFileSourcePartitionHelper, SparkSession}
+import org.apache.spark.sql.catalyst.SQLConfHelper
+import org.apache.spark.sql.catalyst.catalog.{CatalogTable, HiveTableRelation}
+import org.apache.spark.sql.catalyst.planning.ScanOperation
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
+import org.apache.spark.sql.connector.read.SupportsReportStatistics
+import org.apache.spark.sql.execution.{SparkPlan, SparkStrategy}
+import org.apache.spark.sql.execution.datasources.{CatalogFileIndex, HadoopFsRelation, InMemoryFileIndex, LogicalRelation}
+import org.apache.spark.sql.execution.datasources.v2.DataSourceV2ScanRelation
+import org.apache.spark.sql.types.StructType
+
+import org.apache.kyuubi.sql.KyuubiSQLConf
+
+/**
+ * Add MaxScanStrategy to avoid scan excessive partitions or files
+ * 1. Check if scan exceed maxPartition of partitioned table
+ * 2. Check if scan exceed maxFileSize (calculated by hive table and partition statistics)
+ * This Strategy Add Planner Strategy after LogicalOptimizer
+ *
+ * @param session
+ */
+case class MaxScanStrategy(session: SparkSession)
+ extends SparkStrategy with SQLConfHelper with PruneFileSourcePartitionHelper {
+ override def apply(plan: LogicalPlan): Seq[SparkPlan] = {
+ val maxScanPartitionsOpt = conf.getConf(KyuubiSQLConf.WATCHDOG_MAX_PARTITIONS)
+ val maxFileSizeOpt = conf.getConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE)
+ if (maxScanPartitionsOpt.isDefined || maxFileSizeOpt.isDefined) {
+ checkScan(plan, maxScanPartitionsOpt, maxFileSizeOpt)
+ }
+ Nil
+ }
+
+ private def checkScan(
+ plan: LogicalPlan,
+ maxScanPartitionsOpt: Option[Int],
+ maxFileSizeOpt: Option[Long]): Unit = {
+ plan match {
+ case ScanOperation(_, _, _, relation: HiveTableRelation) =>
+ if (relation.isPartitioned) {
+ relation.prunedPartitions match {
+ case Some(prunedPartitions) =>
+ if (maxScanPartitionsOpt.exists(_ < prunedPartitions.size)) {
+ throw new MaxPartitionExceedException(
+ s"""
+ |SQL job scan hive partition: ${prunedPartitions.size}
+ |exceed restrict of hive scan maxPartition ${maxScanPartitionsOpt.get}
+ |You should optimize your SQL logical according partition structure
+ |or shorten query scope such as p_date, detail as below:
+ |Table: ${relation.tableMeta.qualifiedName}
+ |Owner: ${relation.tableMeta.owner}
+ |Partition Structure: ${relation.partitionCols.map(_.name).mkString(", ")}
+ |""".stripMargin)
+ }
+ lazy val scanFileSize = prunedPartitions.flatMap(_.stats).map(_.sizeInBytes).sum
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw partTableMaxFileExceedError(
+ scanFileSize,
+ maxFileSizeOpt.get,
+ Some(relation.tableMeta),
+ prunedPartitions.flatMap(_.storage.locationUri).map(_.toString),
+ relation.partitionCols.map(_.name))
+ }
+ case _ =>
+ lazy val scanPartitions: Int = session
+ .sessionState.catalog.externalCatalog.listPartitionNames(
+ relation.tableMeta.database,
+ relation.tableMeta.identifier.table).size
+ if (maxScanPartitionsOpt.exists(_ < scanPartitions)) {
+ throw new MaxPartitionExceedException(
+ s"""
+ |Your SQL job scan a whole huge table without any partition filter,
+ |You should optimize your SQL logical according partition structure
+ |or shorten query scope such as p_date, detail as below:
+ |Table: ${relation.tableMeta.qualifiedName}
+ |Owner: ${relation.tableMeta.owner}
+ |Partition Structure: ${relation.partitionCols.map(_.name).mkString(", ")}
+ |""".stripMargin)
+ }
+
+ lazy val scanFileSize: BigInt =
+ relation.tableMeta.stats.map(_.sizeInBytes).getOrElse {
+ session
+ .sessionState.catalog.externalCatalog.listPartitions(
+ relation.tableMeta.database,
+ relation.tableMeta.identifier.table).flatMap(_.stats).map(_.sizeInBytes).sum
+ }
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw new MaxFileSizeExceedException(
+ s"""
+ |Your SQL job scan a whole huge table without any partition filter,
+ |You should optimize your SQL logical according partition structure
+ |or shorten query scope such as p_date, detail as below:
+ |Table: ${relation.tableMeta.qualifiedName}
+ |Owner: ${relation.tableMeta.owner}
+ |Partition Structure: ${relation.partitionCols.map(_.name).mkString(", ")}
+ |""".stripMargin)
+ }
+ }
+ } else {
+ lazy val scanFileSize = relation.tableMeta.stats.map(_.sizeInBytes).sum
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw nonPartTableMaxFileExceedError(
+ scanFileSize,
+ maxFileSizeOpt.get,
+ Some(relation.tableMeta))
+ }
+ }
+ case ScanOperation(
+ _,
+ _,
+ filters,
+ relation @ LogicalRelation(
+ fsRelation @ HadoopFsRelation(
+ fileIndex: InMemoryFileIndex,
+ partitionSchema,
+ _,
+ _,
+ _,
+ _),
+ _,
+ _,
+ _,
+ _)) =>
+ if (fsRelation.partitionSchema.nonEmpty) {
+ val (partitionKeyFilters, dataFilter) =
+ getPartitionKeyFiltersAndDataFilters(
+ SparkSession.active,
+ relation,
+ partitionSchema,
+ filters,
+ relation.output)
+ val prunedPartitions = fileIndex.listFiles(
+ partitionKeyFilters.toSeq,
+ dataFilter)
+ if (maxScanPartitionsOpt.exists(_ < prunedPartitions.size)) {
+ throw maxPartitionExceedError(
+ prunedPartitions.size,
+ maxScanPartitionsOpt.get,
+ relation.catalogTable,
+ fileIndex.rootPaths,
+ fsRelation.partitionSchema)
+ }
+ lazy val scanFileSize = prunedPartitions.flatMap(_.files).map(_.getLen).sum
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw partTableMaxFileExceedError(
+ scanFileSize,
+ maxFileSizeOpt.get,
+ relation.catalogTable,
+ fileIndex.rootPaths.map(_.toString),
+ fsRelation.partitionSchema.map(_.name))
+ }
+ } else {
+ lazy val scanFileSize = fileIndex.sizeInBytes
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw nonPartTableMaxFileExceedError(
+ scanFileSize,
+ maxFileSizeOpt.get,
+ relation.catalogTable)
+ }
+ }
+ case ScanOperation(
+ _,
+ _,
+ filters,
+ logicalRelation @ LogicalRelation(
+ fsRelation @ HadoopFsRelation(
+ catalogFileIndex: CatalogFileIndex,
+ partitionSchema,
+ _,
+ _,
+ _,
+ _),
+ _,
+ _,
+ _,
+ _)) =>
+ if (fsRelation.partitionSchema.nonEmpty) {
+ val (partitionKeyFilters, _) =
+ getPartitionKeyFiltersAndDataFilters(
+ SparkSession.active,
+ logicalRelation,
+ partitionSchema,
+ filters,
+ logicalRelation.output)
+
+ val fileIndex = catalogFileIndex.filterPartitions(
+ partitionKeyFilters.toSeq)
+
+ lazy val prunedPartitionSize = fileIndex.partitionSpec().partitions.size
+ if (maxScanPartitionsOpt.exists(_ < prunedPartitionSize)) {
+ throw maxPartitionExceedError(
+ prunedPartitionSize,
+ maxScanPartitionsOpt.get,
+ logicalRelation.catalogTable,
+ catalogFileIndex.rootPaths,
+ fsRelation.partitionSchema)
+ }
+
+ lazy val scanFileSize = fileIndex
+ .listFiles(Nil, Nil).flatMap(_.files).map(_.getLen).sum
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw partTableMaxFileExceedError(
+ scanFileSize,
+ maxFileSizeOpt.get,
+ logicalRelation.catalogTable,
+ catalogFileIndex.rootPaths.map(_.toString),
+ fsRelation.partitionSchema.map(_.name))
+ }
+ } else {
+ lazy val scanFileSize = catalogFileIndex.sizeInBytes
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw nonPartTableMaxFileExceedError(
+ scanFileSize,
+ maxFileSizeOpt.get,
+ logicalRelation.catalogTable)
+ }
+ }
+ case ScanOperation(
+ _,
+ _,
+ _,
+ relation @ DataSourceV2ScanRelation(_, _: SupportsReportStatistics, _, _, _, _)) =>
+ val table = relation.relation.table
+ if (table.partitioning().nonEmpty) {
+ val partitionColumnNames = table.partitioning().map(_.describe())
+ val stats = relation.computeStats()
+ lazy val scanFileSize = stats.sizeInBytes
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw new MaxFileSizeExceedException(
+ s"""
+ |SQL job scan file size in bytes: $scanFileSize
+ |exceed restrict of table scan maxFileSize ${maxFileSizeOpt.get}
+ |You should optimize your SQL logical according partition structure
+ |or shorten query scope such as p_date, detail as below:
+ |Table: ${table.name()}
+ |Partition Structure: ${partitionColumnNames.mkString(",")}
+ |""".stripMargin)
+ }
+ } else {
+ val stats = relation.computeStats()
+ lazy val scanFileSize = stats.sizeInBytes
+ if (maxFileSizeOpt.exists(_ < scanFileSize)) {
+ throw new MaxFileSizeExceedException(
+ s"""
+ |SQL job scan file size in bytes: $scanFileSize
+ |exceed restrict of table scan maxFileSize ${maxFileSizeOpt.get}
+ |detail as below:
+ |Table: ${table.name()}
+ |""".stripMargin)
+ }
+ }
+ case _ =>
+ }
+ }
+
+ def maxPartitionExceedError(
+ prunedPartitionSize: Int,
+ maxPartitionSize: Int,
+ tableMeta: Option[CatalogTable],
+ rootPaths: Seq[Path],
+ partitionSchema: StructType): Throwable = {
+ val truncatedPaths =
+ if (rootPaths.length > 5) {
+ rootPaths.slice(0, 5).mkString(",") + """... """ + (rootPaths.length - 5) + " more paths"
+ } else {
+ rootPaths.mkString(",")
+ }
+
+ new MaxPartitionExceedException(
+ s"""
+ |SQL job scan data source partition: $prunedPartitionSize
+ |exceed restrict of data source scan maxPartition $maxPartitionSize
+ |You should optimize your SQL logical according partition structure
+ |or shorten query scope such as p_date, detail as below:
+ |Table: ${tableMeta.map(_.qualifiedName).getOrElse("")}
+ |Owner: ${tableMeta.map(_.owner).getOrElse("")}
+ |RootPaths: $truncatedPaths
+ |Partition Structure: ${partitionSchema.map(_.name).mkString(", ")}
+ |""".stripMargin)
+ }
+
+ private def partTableMaxFileExceedError(
+ scanFileSize: Number,
+ maxFileSize: Long,
+ tableMeta: Option[CatalogTable],
+ rootPaths: Seq[String],
+ partitions: Seq[String]): Throwable = {
+ val truncatedPaths =
+ if (rootPaths.length > 5) {
+ rootPaths.slice(0, 5).mkString(",") + """... """ + (rootPaths.length - 5) + " more paths"
+ } else {
+ rootPaths.mkString(",")
+ }
+
+ new MaxFileSizeExceedException(
+ s"""
+ |SQL job scan file size in bytes: $scanFileSize
+ |exceed restrict of table scan maxFileSize $maxFileSize
+ |You should optimize your SQL logical according partition structure
+ |or shorten query scope such as p_date, detail as below:
+ |Table: ${tableMeta.map(_.qualifiedName).getOrElse("")}
+ |Owner: ${tableMeta.map(_.owner).getOrElse("")}
+ |RootPaths: $truncatedPaths
+ |Partition Structure: ${partitions.mkString(", ")}
+ |""".stripMargin)
+ }
+
+ private def nonPartTableMaxFileExceedError(
+ scanFileSize: Number,
+ maxFileSize: Long,
+ tableMeta: Option[CatalogTable]): Throwable = {
+ new MaxFileSizeExceedException(
+ s"""
+ |SQL job scan file size in bytes: $scanFileSize
+ |exceed restrict of table scan maxFileSize $maxFileSize
+ |detail as below:
+ |Table: ${tableMeta.map(_.qualifiedName).getOrElse("")}
+ |Owner: ${tableMeta.map(_.owner).getOrElse("")}
+ |Location: ${tableMeta.map(_.location).getOrElse("")}
+ |""".stripMargin)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/InsertZorderBeforeWriting.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/InsertZorderBeforeWriting.scala
new file mode 100644
index 00000000000..2b19bef716e
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/InsertZorderBeforeWriting.scala
@@ -0,0 +1,179 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.zorder
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.catalog.CatalogTable
+import org.apache.spark.sql.catalyst.expressions.{Ascending, Attribute, Expression, NullsLast, SortOrder}
+import org.apache.spark.sql.catalyst.plans.logical._
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.datasources.InsertIntoHadoopFsRelationCommand
+import org.apache.spark.sql.hive.execution.InsertIntoHiveTable
+
+import org.apache.kyuubi.sql.{KyuubiSQLConf, KyuubiSQLExtensionException}
+
+trait InsertZorderBeforeWritingBase extends Rule[LogicalPlan] {
+ private val KYUUBI_ZORDER_ENABLED = "kyuubi.zorder.enabled"
+ private val KYUUBI_ZORDER_COLS = "kyuubi.zorder.cols"
+
+ def session: SparkSession
+
+ def applyInternal(plan: LogicalPlan): LogicalPlan
+
+ final override def apply(plan: LogicalPlan): LogicalPlan = {
+ if (conf.getConf(KyuubiSQLConf.INSERT_ZORDER_BEFORE_WRITING)) {
+ applyInternal(plan)
+ } else {
+ plan
+ }
+ }
+
+ def isZorderEnabled(props: Map[String, String]): Boolean = {
+ props.contains(KYUUBI_ZORDER_ENABLED) &&
+ "true".equalsIgnoreCase(props(KYUUBI_ZORDER_ENABLED)) &&
+ props.contains(KYUUBI_ZORDER_COLS)
+ }
+
+ def getZorderColumns(props: Map[String, String]): Seq[String] = {
+ val cols = props.get(KYUUBI_ZORDER_COLS)
+ assert(cols.isDefined)
+ cols.get.split(",").map(_.trim)
+ }
+
+ def canInsertZorder(query: LogicalPlan): Boolean = query match {
+ case Project(_, child) => canInsertZorder(child)
+ case _: RepartitionOperation | _: RebalancePartitions
+ if !conf.getConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED) => true
+ // TODO: actually, we can force zorder even if existed some shuffle
+ case _: Sort => false
+ case _: RepartitionOperation => false
+ case _: RebalancePartitions => false
+ case _ => true
+ }
+
+ def insertZorder(
+ catalogTable: CatalogTable,
+ plan: LogicalPlan,
+ dynamicPartitionColumns: Seq[Attribute]): LogicalPlan = {
+ if (!canInsertZorder(plan)) {
+ return plan
+ }
+ val cols = getZorderColumns(catalogTable.properties)
+ val resolver = session.sessionState.conf.resolver
+ val output = plan.output
+ val bound = cols.flatMap(col => output.find(attr => resolver(attr.name, col)))
+ if (bound.size < cols.size) {
+ logWarning(s"target table does not contain all zorder cols: ${cols.mkString(",")}, " +
+ s"please check your table properties $KYUUBI_ZORDER_COLS.")
+ plan
+ } else {
+ if (conf.getConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED) &&
+ conf.getConf(KyuubiSQLConf.REBALANCE_BEFORE_ZORDER)) {
+ throw new KyuubiSQLExtensionException(s"${KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED.key} " +
+ s"and ${KyuubiSQLConf.REBALANCE_BEFORE_ZORDER.key} can not be enabled together.")
+ }
+ if (conf.getConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED) &&
+ dynamicPartitionColumns.nonEmpty) {
+ logWarning(s"Dynamic partition insertion with global sort may produce small files.")
+ }
+
+ val zorderExpr =
+ if (bound.length == 1) {
+ bound
+ } else if (conf.getConf(KyuubiSQLConf.ZORDER_USING_ORIGINAL_ORDERING_ENABLED)) {
+ bound.asInstanceOf[Seq[Expression]]
+ } else {
+ Zorder(bound) :: Nil
+ }
+ val (global, orderExprs, child) =
+ if (conf.getConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED)) {
+ (true, zorderExpr, plan)
+ } else if (conf.getConf(KyuubiSQLConf.REBALANCE_BEFORE_ZORDER)) {
+ val rebalanceExpr =
+ if (dynamicPartitionColumns.isEmpty) {
+ // static partition insert
+ bound
+ } else if (conf.getConf(KyuubiSQLConf.REBALANCE_ZORDER_COLUMNS_ENABLED)) {
+ // improve data compression ratio
+ dynamicPartitionColumns.asInstanceOf[Seq[Expression]] ++ bound
+ } else {
+ dynamicPartitionColumns.asInstanceOf[Seq[Expression]]
+ }
+ // for dynamic partition insert, Spark always sort the partition columns,
+ // so here we sort partition columns + zorder.
+ val rebalance =
+ if (dynamicPartitionColumns.nonEmpty &&
+ conf.getConf(KyuubiSQLConf.TWO_PHASE_REBALANCE_BEFORE_ZORDER)) {
+ // improve compression ratio
+ RebalancePartitions(
+ rebalanceExpr,
+ RebalancePartitions(dynamicPartitionColumns, plan))
+ } else {
+ RebalancePartitions(rebalanceExpr, plan)
+ }
+ (false, dynamicPartitionColumns.asInstanceOf[Seq[Expression]] ++ zorderExpr, rebalance)
+ } else {
+ (false, zorderExpr, plan)
+ }
+ val order = orderExprs.map { expr =>
+ SortOrder(expr, Ascending, NullsLast, Seq.empty)
+ }
+ Sort(order, global, child)
+ }
+ }
+}
+
+case class InsertZorderBeforeWritingDatasource(session: SparkSession)
+ extends InsertZorderBeforeWritingBase {
+
+ override def applyInternal(plan: LogicalPlan): LogicalPlan = plan match {
+ case insert: InsertIntoHadoopFsRelationCommand
+ if insert.query.resolved && insert.bucketSpec.isEmpty && insert.catalogTable.isDefined &&
+ isZorderEnabled(insert.catalogTable.get.properties) =>
+ val dynamicPartition =
+ insert.partitionColumns.filterNot(attr => insert.staticPartitions.contains(attr.name))
+ val newQuery = insertZorder(insert.catalogTable.get, insert.query, dynamicPartition)
+ if (newQuery.eq(insert.query)) {
+ insert
+ } else {
+ insert.copy(query = newQuery)
+ }
+
+ case _ => plan
+ }
+}
+
+case class InsertZorderBeforeWritingHive(session: SparkSession)
+ extends InsertZorderBeforeWritingBase {
+
+ override def applyInternal(plan: LogicalPlan): LogicalPlan = plan match {
+ case insert: InsertIntoHiveTable
+ if insert.query.resolved && insert.table.bucketSpec.isEmpty &&
+ isZorderEnabled(insert.table.properties) =>
+ val dynamicPartition = insert.partition.filter(_._2.isEmpty).keys
+ .flatMap(name => insert.query.output.find(_.name == name)).toSeq
+ val newQuery = insertZorder(insert.table, insert.query, dynamicPartition)
+ if (newQuery.eq(insert.query)) {
+ insert
+ } else {
+ insert.copy(query = newQuery)
+ }
+
+ case _ => plan
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/OptimizeZorderCommand.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/OptimizeZorderCommand.scala
new file mode 100644
index 00000000000..65a714f37fa
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/OptimizeZorderCommand.scala
@@ -0,0 +1,70 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.zorder
+
+import org.apache.spark.sql.Row
+import org.apache.spark.sql.catalyst.catalog.CatalogTable
+import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
+import org.apache.spark.sql.classic.SparkSession
+import org.apache.spark.sql.execution.SparkPlan
+import org.apache.spark.sql.execution.command.DataWritingCommand
+import org.apache.spark.sql.hive.execution.InsertIntoHiveTable
+
+import org.apache.kyuubi.sql.KyuubiSQLExtensionException
+
+/**
+ * A runnable command for zorder, we delegate to real command to execute
+ */
+case class OptimizeZorderCommand(
+ catalogTable: CatalogTable,
+ override val query: LogicalPlan) extends DataWritingCommand {
+
+ override def outputColumnNames: Seq[String] = query.output.map(_.name)
+
+ private def isHiveTable: Boolean = {
+ catalogTable.provider.isEmpty ||
+ (catalogTable.provider.isDefined && "hive".equalsIgnoreCase(catalogTable.provider.get))
+ }
+
+ private def getWritingCommand(session: SparkSession): DataWritingCommand = {
+ // TODO: Support convert hive relation to datasource relation, can see
+ // [[org.apache.spark.sql.hive.RelationConversions]]
+ InsertIntoHiveTable(
+ catalogTable,
+ catalogTable.partitionColumnNames.map(p => (p, None)).toMap,
+ query,
+ overwrite = true,
+ ifPartitionNotExists = false,
+ outputColumnNames)
+ }
+
+ override def run(session: SparkSession, child: SparkPlan): Seq[Row] = {
+ // TODO: Support datasource relation
+ // TODO: Support read and insert overwrite the same table for some table format
+ if (!isHiveTable) {
+ throw new KyuubiSQLExtensionException("only support hive table")
+ }
+
+ val command = getWritingCommand(session)
+ command.run(session, child)
+ DataWritingCommand.propogateMetrics(session.sparkContext, command, metrics)
+ Seq.empty
+ }
+
+ override def withNewChildInternal(newChild: LogicalPlan): LogicalPlan = copy(query = newChild)
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/OptimizeZorderStatement.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/OptimizeZorderStatement.scala
new file mode 100644
index 00000000000..61a9bda35d1
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/OptimizeZorderStatement.scala
@@ -0,0 +1,33 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.zorder
+
+import org.apache.spark.sql.catalyst.expressions.Attribute
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode}
+
+/**
+ * A zorder statement that contains we parsed from SQL.
+ * We should convert this plan to certain command at Analyzer.
+ */
+case class OptimizeZorderStatement(
+ tableIdentifier: Seq[String],
+ query: LogicalPlan) extends UnaryNode {
+ override def child: LogicalPlan = query
+ override def output: Seq[Attribute] = child.output
+ protected def withNewChildInternal(newChild: LogicalPlan): LogicalPlan = copy(query = newChild)
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/ResolveZorder.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/ResolveZorder.scala
new file mode 100644
index 00000000000..0df83685b72
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/ResolveZorder.scala
@@ -0,0 +1,64 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.zorder
+
+import org.apache.spark.sql.SparkSession
+import org.apache.spark.sql.catalyst.TableIdentifier
+import org.apache.spark.sql.catalyst.catalog.HiveTableRelation
+import org.apache.spark.sql.catalyst.expressions.AttributeSet
+import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, SubqueryAlias}
+import org.apache.spark.sql.catalyst.rules.Rule
+
+import org.apache.kyuubi.sql.KyuubiSQLExtensionException
+
+/**
+ * Resolve `OptimizeZorderStatement` to `OptimizeZorderCommand`
+ */
+case class ResolveZorder(session: SparkSession) extends Rule[LogicalPlan] {
+
+ protected def checkQueryAllowed(query: LogicalPlan): Unit = query foreach {
+ case Filter(condition, SubqueryAlias(_, tableRelation: HiveTableRelation)) =>
+ if (tableRelation.partitionCols.isEmpty) {
+ throw new KyuubiSQLExtensionException("Filters are only supported for partitioned table")
+ }
+
+ val partitionKeyIds = AttributeSet(tableRelation.partitionCols)
+ if (condition.references.isEmpty || !condition.references.subsetOf(partitionKeyIds)) {
+ throw new KyuubiSQLExtensionException("Only partition column filters are allowed")
+ }
+
+ case _ =>
+ }
+
+ protected def getTableIdentifier(tableIdent: Seq[String]): TableIdentifier = tableIdent match {
+ case Seq(tbl) => TableIdentifier.apply(tbl)
+ case Seq(db, tbl) => TableIdentifier.apply(tbl, Some(db))
+ case _ => throw new KyuubiSQLExtensionException(
+ "only support session catalog table, please use db.table instead")
+ }
+
+ override def apply(plan: LogicalPlan): LogicalPlan = plan match {
+ case statement: OptimizeZorderStatement if statement.query.resolved =>
+ checkQueryAllowed(statement.query)
+ val tableIdentifier = getTableIdentifier(statement.tableIdentifier)
+ val catalogTable = session.sessionState.catalog.getTableMetadata(tableIdentifier)
+ OptimizeZorderCommand(catalogTable, statement.query)
+
+ case _ => plan
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/Zorder.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/Zorder.scala
new file mode 100644
index 00000000000..aa4da07ff16
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/Zorder.scala
@@ -0,0 +1,93 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.zorder
+
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.analysis.TypeCheckResult
+import org.apache.spark.sql.catalyst.expressions.Expression
+import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode, FalseLiteral}
+import org.apache.spark.sql.catalyst.expressions.codegen.Block._
+import org.apache.spark.sql.types.{BinaryType, DataType}
+
+import org.apache.kyuubi.sql.KyuubiSQLExtensionException
+
+case class Zorder(children: Seq[Expression]) extends Expression {
+ override def foldable: Boolean = children.forall(_.foldable)
+ override def nullable: Boolean = false
+ override def dataType: DataType = BinaryType
+ override def prettyName: String = "zorder"
+
+ override def checkInputDataTypes(): TypeCheckResult = {
+ try {
+ defaultNullValues
+ TypeCheckResult.TypeCheckSuccess
+ } catch {
+ case e: KyuubiSQLExtensionException =>
+ TypeCheckResult.TypeCheckFailure(e.getMessage)
+ }
+ }
+
+ @transient
+ private[this] lazy val defaultNullValues: Array[Any] =
+ children.map(_.dataType)
+ .map(ZorderBytesUtils.defaultValue)
+ .toArray
+
+ override def eval(input: InternalRow): Any = {
+ val childrenValues = children.zipWithIndex.map {
+ case (child: Expression, index) =>
+ val v = child.eval(input)
+ if (v == null) {
+ defaultNullValues(index)
+ } else {
+ v
+ }
+ }
+ ZorderBytesUtils.interleaveBits(childrenValues.toArray)
+ }
+
+ override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = {
+ val evals = children.map(_.genCode(ctx))
+ val defaultValues = ctx.addReferenceObj("defaultValues", defaultNullValues)
+ val values = ctx.freshName("values")
+ val util = ZorderBytesUtils.getClass.getName.stripSuffix("$")
+ val inputs = evals.zipWithIndex.map {
+ case (eval, index) =>
+ s"""
+ |${eval.code}
+ |if (${eval.isNull}) {
+ | $values[$index] = $defaultValues[$index];
+ |} else {
+ | $values[$index] = ${eval.value};
+ |}
+ |""".stripMargin
+ }
+ ev.copy(
+ code =
+ code"""
+ |byte[] ${ev.value} = null;
+ |Object[] $values = new Object[${evals.length}];
+ |${inputs.mkString("\n")}
+ |${ev.value} = $util.interleaveBits($values);
+ |""".stripMargin,
+ isNull = FalseLiteral)
+ }
+
+ override def withNewChildrenInternal(newChildren: IndexedSeq[Expression]): Expression =
+ copy(children = newChildren)
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/ZorderBytesUtils.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/ZorderBytesUtils.scala
new file mode 100644
index 00000000000..d249f1dc32f
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/kyuubi/sql/zorder/ZorderBytesUtils.scala
@@ -0,0 +1,517 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.kyuubi.sql.zorder
+
+import java.lang.{Double => jDouble, Float => jFloat}
+
+import org.apache.spark.sql.types._
+import org.apache.spark.unsafe.types.UTF8String
+
+import org.apache.kyuubi.sql.KyuubiSQLExtensionException
+
+object ZorderBytesUtils {
+ final private val BIT_8_MASK = 1 << 7
+ final private val BIT_16_MASK = 1 << 15
+ final private val BIT_32_MASK = 1 << 31
+ final private val BIT_64_MASK = 1L << 63
+
+ def interleaveBits(inputs: Array[Any]): Array[Byte] = {
+ inputs.length match {
+ // it's a more fast approach, use O(8 * 8)
+ // can see http://graphics.stanford.edu/~seander/bithacks.html#InterleaveTableObvious
+ case 1 => longToByte(toLong(inputs(0)))
+ case 2 => interleave2Longs(toLong(inputs(0)), toLong(inputs(1)))
+ case 3 => interleave3Longs(toLong(inputs(0)), toLong(inputs(1)), toLong(inputs(2)))
+ case 4 =>
+ interleave4Longs(toLong(inputs(0)), toLong(inputs(1)), toLong(inputs(2)), toLong(inputs(3)))
+ case 5 => interleave5Longs(
+ toLong(inputs(0)),
+ toLong(inputs(1)),
+ toLong(inputs(2)),
+ toLong(inputs(3)),
+ toLong(inputs(4)))
+ case 6 => interleave6Longs(
+ toLong(inputs(0)),
+ toLong(inputs(1)),
+ toLong(inputs(2)),
+ toLong(inputs(3)),
+ toLong(inputs(4)),
+ toLong(inputs(5)))
+ case 7 => interleave7Longs(
+ toLong(inputs(0)),
+ toLong(inputs(1)),
+ toLong(inputs(2)),
+ toLong(inputs(3)),
+ toLong(inputs(4)),
+ toLong(inputs(5)),
+ toLong(inputs(6)))
+ case 8 => interleave8Longs(
+ toLong(inputs(0)),
+ toLong(inputs(1)),
+ toLong(inputs(2)),
+ toLong(inputs(3)),
+ toLong(inputs(4)),
+ toLong(inputs(5)),
+ toLong(inputs(6)),
+ toLong(inputs(7)))
+
+ case _ =>
+ // it's the default approach, use O(64 * n), n is the length of inputs
+ interleaveBitsDefault(inputs.map(toByteArray))
+ }
+ }
+
+ private def interleave2Longs(l1: Long, l2: Long): Array[Byte] = {
+ // output 8 * 16 bits
+ val result = new Array[Byte](16)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toShort
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toShort
+
+ var z = 0
+ var j = 0
+ while (j < 8) {
+ val x_masked = tmp1 & (1 << j)
+ val y_masked = tmp2 & (1 << j)
+ z |= (x_masked << j)
+ z |= (y_masked << (j + 1))
+ j = j + 1
+ }
+ result((7 - i) * 2 + 1) = (z & 0xFF).toByte
+ result((7 - i) * 2) = ((z >> 8) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ private def interleave3Longs(l1: Long, l2: Long, l3: Long): Array[Byte] = {
+ // output 8 * 24 bits
+ val result = new Array[Byte](24)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toInt
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toInt
+ val tmp3 = ((l3 >> (i * 8)) & 0xFF).toInt
+
+ var z = 0
+ var j = 0
+ while (j < 8) {
+ val r1_mask = tmp1 & (1 << j)
+ val r2_mask = tmp2 & (1 << j)
+ val r3_mask = tmp3 & (1 << j)
+ z |= (r1_mask << (2 * j)) | (r2_mask << (2 * j + 1)) | (r3_mask << (2 * j + 2))
+ j = j + 1
+ }
+ result((7 - i) * 3 + 2) = (z & 0xFF).toByte
+ result((7 - i) * 3 + 1) = ((z >> 8) & 0xFF).toByte
+ result((7 - i) * 3) = ((z >> 16) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ private def interleave4Longs(l1: Long, l2: Long, l3: Long, l4: Long): Array[Byte] = {
+ // output 8 * 32 bits
+ val result = new Array[Byte](32)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toInt
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toInt
+ val tmp3 = ((l3 >> (i * 8)) & 0xFF).toInt
+ val tmp4 = ((l4 >> (i * 8)) & 0xFF).toInt
+
+ var z = 0
+ var j = 0
+ while (j < 8) {
+ val r1_mask = tmp1 & (1 << j)
+ val r2_mask = tmp2 & (1 << j)
+ val r3_mask = tmp3 & (1 << j)
+ val r4_mask = tmp4 & (1 << j)
+ z |= (r1_mask << (3 * j)) | (r2_mask << (3 * j + 1)) | (r3_mask << (3 * j + 2)) |
+ (r4_mask << (3 * j + 3))
+ j = j + 1
+ }
+ result((7 - i) * 4 + 3) = (z & 0xFF).toByte
+ result((7 - i) * 4 + 2) = ((z >> 8) & 0xFF).toByte
+ result((7 - i) * 4 + 1) = ((z >> 16) & 0xFF).toByte
+ result((7 - i) * 4) = ((z >> 24) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ private def interleave5Longs(
+ l1: Long,
+ l2: Long,
+ l3: Long,
+ l4: Long,
+ l5: Long): Array[Byte] = {
+ // output 8 * 40 bits
+ val result = new Array[Byte](40)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toLong
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toLong
+ val tmp3 = ((l3 >> (i * 8)) & 0xFF).toLong
+ val tmp4 = ((l4 >> (i * 8)) & 0xFF).toLong
+ val tmp5 = ((l5 >> (i * 8)) & 0xFF).toLong
+
+ var z = 0L
+ var j = 0
+ while (j < 8) {
+ val r1_mask = tmp1 & (1 << j)
+ val r2_mask = tmp2 & (1 << j)
+ val r3_mask = tmp3 & (1 << j)
+ val r4_mask = tmp4 & (1 << j)
+ val r5_mask = tmp5 & (1 << j)
+ z |= (r1_mask << (4 * j)) | (r2_mask << (4 * j + 1)) | (r3_mask << (4 * j + 2)) |
+ (r4_mask << (4 * j + 3)) | (r5_mask << (4 * j + 4))
+ j = j + 1
+ }
+ result((7 - i) * 5 + 4) = (z & 0xFF).toByte
+ result((7 - i) * 5 + 3) = ((z >> 8) & 0xFF).toByte
+ result((7 - i) * 5 + 2) = ((z >> 16) & 0xFF).toByte
+ result((7 - i) * 5 + 1) = ((z >> 24) & 0xFF).toByte
+ result((7 - i) * 5) = ((z >> 32) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ private def interleave6Longs(
+ l1: Long,
+ l2: Long,
+ l3: Long,
+ l4: Long,
+ l5: Long,
+ l6: Long): Array[Byte] = {
+ // output 8 * 48 bits
+ val result = new Array[Byte](48)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toLong
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toLong
+ val tmp3 = ((l3 >> (i * 8)) & 0xFF).toLong
+ val tmp4 = ((l4 >> (i * 8)) & 0xFF).toLong
+ val tmp5 = ((l5 >> (i * 8)) & 0xFF).toLong
+ val tmp6 = ((l6 >> (i * 8)) & 0xFF).toLong
+
+ var z = 0L
+ var j = 0
+ while (j < 8) {
+ val r1_mask = tmp1 & (1 << j)
+ val r2_mask = tmp2 & (1 << j)
+ val r3_mask = tmp3 & (1 << j)
+ val r4_mask = tmp4 & (1 << j)
+ val r5_mask = tmp5 & (1 << j)
+ val r6_mask = tmp6 & (1 << j)
+ z |= (r1_mask << (5 * j)) | (r2_mask << (5 * j + 1)) | (r3_mask << (5 * j + 2)) |
+ (r4_mask << (5 * j + 3)) | (r5_mask << (5 * j + 4)) | (r6_mask << (5 * j + 5))
+ j = j + 1
+ }
+ result((7 - i) * 6 + 5) = (z & 0xFF).toByte
+ result((7 - i) * 6 + 4) = ((z >> 8) & 0xFF).toByte
+ result((7 - i) * 6 + 3) = ((z >> 16) & 0xFF).toByte
+ result((7 - i) * 6 + 2) = ((z >> 24) & 0xFF).toByte
+ result((7 - i) * 6 + 1) = ((z >> 32) & 0xFF).toByte
+ result((7 - i) * 6) = ((z >> 40) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ private def interleave7Longs(
+ l1: Long,
+ l2: Long,
+ l3: Long,
+ l4: Long,
+ l5: Long,
+ l6: Long,
+ l7: Long): Array[Byte] = {
+ // output 8 * 56 bits
+ val result = new Array[Byte](56)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toLong
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toLong
+ val tmp3 = ((l3 >> (i * 8)) & 0xFF).toLong
+ val tmp4 = ((l4 >> (i * 8)) & 0xFF).toLong
+ val tmp5 = ((l5 >> (i * 8)) & 0xFF).toLong
+ val tmp6 = ((l6 >> (i * 8)) & 0xFF).toLong
+ val tmp7 = ((l7 >> (i * 8)) & 0xFF).toLong
+
+ var z = 0L
+ var j = 0
+ while (j < 8) {
+ val r1_mask = tmp1 & (1 << j)
+ val r2_mask = tmp2 & (1 << j)
+ val r3_mask = tmp3 & (1 << j)
+ val r4_mask = tmp4 & (1 << j)
+ val r5_mask = tmp5 & (1 << j)
+ val r6_mask = tmp6 & (1 << j)
+ val r7_mask = tmp7 & (1 << j)
+ z |= (r1_mask << (6 * j)) | (r2_mask << (6 * j + 1)) | (r3_mask << (6 * j + 2)) |
+ (r4_mask << (6 * j + 3)) | (r5_mask << (6 * j + 4)) | (r6_mask << (6 * j + 5)) |
+ (r7_mask << (6 * j + 6))
+ j = j + 1
+ }
+ result((7 - i) * 7 + 6) = (z & 0xFF).toByte
+ result((7 - i) * 7 + 5) = ((z >> 8) & 0xFF).toByte
+ result((7 - i) * 7 + 4) = ((z >> 16) & 0xFF).toByte
+ result((7 - i) * 7 + 3) = ((z >> 24) & 0xFF).toByte
+ result((7 - i) * 7 + 2) = ((z >> 32) & 0xFF).toByte
+ result((7 - i) * 7 + 1) = ((z >> 40) & 0xFF).toByte
+ result((7 - i) * 7) = ((z >> 48) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ private def interleave8Longs(
+ l1: Long,
+ l2: Long,
+ l3: Long,
+ l4: Long,
+ l5: Long,
+ l6: Long,
+ l7: Long,
+ l8: Long): Array[Byte] = {
+ // output 8 * 64 bits
+ val result = new Array[Byte](64)
+ var i = 0
+ while (i < 8) {
+ val tmp1 = ((l1 >> (i * 8)) & 0xFF).toLong
+ val tmp2 = ((l2 >> (i * 8)) & 0xFF).toLong
+ val tmp3 = ((l3 >> (i * 8)) & 0xFF).toLong
+ val tmp4 = ((l4 >> (i * 8)) & 0xFF).toLong
+ val tmp5 = ((l5 >> (i * 8)) & 0xFF).toLong
+ val tmp6 = ((l6 >> (i * 8)) & 0xFF).toLong
+ val tmp7 = ((l7 >> (i * 8)) & 0xFF).toLong
+ val tmp8 = ((l8 >> (i * 8)) & 0xFF).toLong
+
+ var z = 0L
+ var j = 0
+ while (j < 8) {
+ val r1_mask = tmp1 & (1 << j)
+ val r2_mask = tmp2 & (1 << j)
+ val r3_mask = tmp3 & (1 << j)
+ val r4_mask = tmp4 & (1 << j)
+ val r5_mask = tmp5 & (1 << j)
+ val r6_mask = tmp6 & (1 << j)
+ val r7_mask = tmp7 & (1 << j)
+ val r8_mask = tmp8 & (1 << j)
+ z |= (r1_mask << (7 * j)) | (r2_mask << (7 * j + 1)) | (r3_mask << (7 * j + 2)) |
+ (r4_mask << (7 * j + 3)) | (r5_mask << (7 * j + 4)) | (r6_mask << (7 * j + 5)) |
+ (r7_mask << (7 * j + 6)) | (r8_mask << (7 * j + 7))
+ j = j + 1
+ }
+ result((7 - i) * 8 + 7) = (z & 0xFF).toByte
+ result((7 - i) * 8 + 6) = ((z >> 8) & 0xFF).toByte
+ result((7 - i) * 8 + 5) = ((z >> 16) & 0xFF).toByte
+ result((7 - i) * 8 + 4) = ((z >> 24) & 0xFF).toByte
+ result((7 - i) * 8 + 3) = ((z >> 32) & 0xFF).toByte
+ result((7 - i) * 8 + 2) = ((z >> 40) & 0xFF).toByte
+ result((7 - i) * 8 + 1) = ((z >> 48) & 0xFF).toByte
+ result((7 - i) * 8) = ((z >> 56) & 0xFF).toByte
+ i = i + 1
+ }
+ result
+ }
+
+ def interleaveBitsDefault(arrays: Array[Array[Byte]]): Array[Byte] = {
+ var totalLength = 0
+ var maxLength = 0
+ arrays.foreach { array =>
+ totalLength += array.length
+ maxLength = maxLength.max(array.length * 8)
+ }
+ val result = new Array[Byte](totalLength)
+ var resultBit = 0
+
+ var bit = 0
+ while (bit < maxLength) {
+ val bytePos = bit / 8
+ val bitPos = bit % 8
+
+ for (arr <- arrays) {
+ val len = arr.length
+ if (bytePos < len) {
+ val resultBytePos = totalLength - 1 - resultBit / 8
+ val resultBitPos = resultBit % 8
+ result(resultBytePos) =
+ updatePos(result(resultBytePos), resultBitPos, arr(len - 1 - bytePos), bitPos)
+ resultBit += 1
+ }
+ }
+ bit += 1
+ }
+ result
+ }
+
+ def updatePos(a: Byte, apos: Int, b: Byte, bpos: Int): Byte = {
+ var temp = (b & (1 << bpos)).toByte
+ if (apos > bpos) {
+ temp = (temp << (apos - bpos)).toByte
+ } else if (apos < bpos) {
+ temp = (temp >> (bpos - apos)).toByte
+ }
+ val atemp = (a & (1 << apos)).toByte
+ if (atemp == temp) {
+ return a
+ }
+ (a ^ (1 << apos)).toByte
+ }
+
+ def toLong(a: Any): Long = {
+ a match {
+ case b: Boolean => (if (b) 1 else 0).toLong ^ BIT_64_MASK
+ case b: Byte => b.toLong ^ BIT_64_MASK
+ case s: Short => s.toLong ^ BIT_64_MASK
+ case i: Int => i.toLong ^ BIT_64_MASK
+ case l: Long => l ^ BIT_64_MASK
+ case f: Float => java.lang.Float.floatToRawIntBits(f).toLong ^ BIT_64_MASK
+ case d: Double => java.lang.Double.doubleToRawLongBits(d) ^ BIT_64_MASK
+ case str: UTF8String => str.getPrefix
+ case dec: Decimal => dec.toLong ^ BIT_64_MASK
+ case other: Any =>
+ throw new KyuubiSQLExtensionException("Unsupported z-order type: " + other.getClass)
+ }
+ }
+
+ def toByteArray(a: Any): Array[Byte] = {
+ a match {
+ case bo: Boolean =>
+ booleanToByte(bo)
+ case b: Byte =>
+ byteToByte(b)
+ case s: Short =>
+ shortToByte(s)
+ case i: Int =>
+ intToByte(i)
+ case l: Long =>
+ longToByte(l)
+ case f: Float =>
+ floatToByte(f)
+ case d: Double =>
+ doubleToByte(d)
+ case str: UTF8String =>
+ // truncate or padding str to 8 byte
+ paddingTo8Byte(str.getBytes)
+ case dec: Decimal =>
+ longToByte(dec.toLong)
+ case other: Any =>
+ throw new KyuubiSQLExtensionException("Unsupported z-order type: " + other.getClass)
+ }
+ }
+
+ def booleanToByte(a: Boolean): Array[Byte] = {
+ if (a) {
+ byteToByte(1.toByte)
+ } else {
+ byteToByte(0.toByte)
+ }
+ }
+
+ def byteToByte(a: Byte): Array[Byte] = {
+ val tmp = (a ^ BIT_8_MASK).toByte
+ Array(tmp)
+ }
+
+ def shortToByte(a: Short): Array[Byte] = {
+ val tmp = a ^ BIT_16_MASK
+ Array(((tmp >> 8) & 0xFF).toByte, (tmp & 0xFF).toByte)
+ }
+
+ def intToByte(a: Int): Array[Byte] = {
+ val result = new Array[Byte](4)
+ var i = 0
+ val tmp = a ^ BIT_32_MASK
+ while (i <= 3) {
+ val offset = i * 8
+ result(3 - i) = ((tmp >> offset) & 0xFF).toByte
+ i += 1
+ }
+ result
+ }
+
+ def longToByte(a: Long): Array[Byte] = {
+ val result = new Array[Byte](8)
+ var i = 0
+ val tmp = a ^ BIT_64_MASK
+ while (i <= 7) {
+ val offset = i * 8
+ result(7 - i) = ((tmp >> offset) & 0xFF).toByte
+ i += 1
+ }
+ result
+ }
+
+ def floatToByte(a: Float): Array[Byte] = {
+ val fi = jFloat.floatToRawIntBits(a)
+ intToByte(fi)
+ }
+
+ def doubleToByte(a: Double): Array[Byte] = {
+ val dl = jDouble.doubleToRawLongBits(a)
+ longToByte(dl)
+ }
+
+ def paddingTo8Byte(a: Array[Byte]): Array[Byte] = {
+ val len = a.length
+ if (len == 8) {
+ a
+ } else if (len > 8) {
+ val result = new Array[Byte](8)
+ System.arraycopy(a, 0, result, 0, 8)
+ result
+ } else {
+ val result = new Array[Byte](8)
+ System.arraycopy(a, 0, result, 8 - len, len)
+ result
+ }
+ }
+
+ def defaultByteArrayValue(dataType: DataType): Array[Byte] = toByteArray {
+ defaultValue(dataType)
+ }
+
+ def defaultValue(dataType: DataType): Any = {
+ dataType match {
+ case BooleanType =>
+ true
+ case ByteType =>
+ Byte.MaxValue
+ case ShortType =>
+ Short.MaxValue
+ case IntegerType | DateType =>
+ Int.MaxValue
+ case LongType | TimestampType | _: DecimalType =>
+ Long.MaxValue
+ case FloatType =>
+ Float.MaxValue
+ case DoubleType =>
+ Double.MaxValue
+ case StringType =>
+ // we pad string to 8 bytes so it's equal to long
+ UTF8String.fromBytes(longToByte(Long.MaxValue))
+ case other: Any =>
+ throw new KyuubiSQLExtensionException(s"Unsupported z-order type: ${other.catalogString}")
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/FinalStageResourceManager.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/FinalStageResourceManager.scala
new file mode 100644
index 00000000000..81873476cc4
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/FinalStageResourceManager.scala
@@ -0,0 +1,289 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import scala.annotation.tailrec
+import scala.collection.mutable
+import scala.collection.mutable.ArrayBuffer
+
+import org.apache.spark.{ExecutorAllocationClient, MapOutputTrackerMaster, SparkContext, SparkEnv}
+import org.apache.spark.internal.Logging
+import org.apache.spark.resource.ResourceProfile
+import org.apache.spark.scheduler.cluster.CoarseGrainedSchedulerBackend
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.{FilterExec, ProjectExec, SortExec, SparkPlan}
+import org.apache.spark.sql.execution.adaptive._
+import org.apache.spark.sql.execution.columnar.InMemoryTableScanExec
+import org.apache.spark.sql.execution.command.DataWritingCommandExec
+import org.apache.spark.sql.execution.datasources.WriteFilesExec
+import org.apache.spark.sql.execution.datasources.v2.V2TableWriteExec
+import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, ShuffleExchangeExec}
+
+import org.apache.kyuubi.sql.{KyuubiSQLConf, WriteUtils}
+
+/**
+ * This rule assumes the final write stage has less cores requirement than previous, otherwise
+ * this rule would take no effect.
+ *
+ * It provide a feature:
+ * 1. Kill redundant executors before running final write stage
+ */
+case class FinalStageResourceManager(session: SparkSession)
+ extends Rule[SparkPlan] with FinalRebalanceStageHelper {
+ override def apply(plan: SparkPlan): SparkPlan = {
+ if (!conf.getConf(KyuubiSQLConf.FINAL_WRITE_STAGE_EAGERLY_KILL_EXECUTORS_ENABLED)) {
+ return plan
+ }
+
+ if (!WriteUtils.isWrite(session, plan)) {
+ return plan
+ }
+
+ val sc = session.sparkContext
+ val dra = sc.getConf.getBoolean("spark.dynamicAllocation.enabled", false)
+ val coresPerExecutor = sc.getConf.getInt("spark.executor.cores", 1)
+ val minExecutors = sc.getConf.getInt("spark.dynamicAllocation.minExecutors", 0)
+ val maxExecutors = sc.getConf.getInt("spark.dynamicAllocation.maxExecutors", Int.MaxValue)
+ val factor = conf.getConf(KyuubiSQLConf.FINAL_WRITE_STAGE_PARTITION_FACTOR)
+ val hasImprovementRoom = maxExecutors - 1 > minExecutors * factor
+ // Fast fail if:
+ // 1. DRA off
+ // 2. only work with yarn and k8s
+ // 3. maxExecutors is not bigger than minExecutors * factor
+ if (!dra || !sc.schedulerBackend.isInstanceOf[CoarseGrainedSchedulerBackend] ||
+ !hasImprovementRoom) {
+ return plan
+ }
+
+ val stageOpt = findFinalRebalanceStage(plan)
+ if (stageOpt.isEmpty) {
+ return plan
+ }
+
+ // It's not safe to kill executors if this plan contains table cache.
+ // If the executor loses then the rdd would re-compute those partition.
+ if (hasTableCache(plan) &&
+ conf.getConf(KyuubiSQLConf.FINAL_WRITE_STAGE_SKIP_KILLING_EXECUTORS_FOR_TABLE_CACHE)) {
+ return plan
+ }
+
+ // TODO: move this to query stage optimizer when updating Spark to 3.5.x
+ // Since we are in `prepareQueryStage`, the AQE shuffle read has not been applied.
+ // So we need to apply it by self.
+ val shuffleRead = queryStageOptimizerRules.foldLeft(stageOpt.get.asInstanceOf[SparkPlan]) {
+ case (latest, rule) => rule.apply(latest)
+ }
+ val (targetCores, stage) = shuffleRead match {
+ case AQEShuffleReadExec(stage: ShuffleQueryStageExec, partitionSpecs) =>
+ (partitionSpecs.length, stage)
+ case stage: ShuffleQueryStageExec =>
+ // we can still kill executors if no AQE shuffle read, e.g., `.repartition(2)`
+ (stage.shuffle.numPartitions, stage)
+ case _ =>
+ // it should never happen in current Spark, but to be safe do nothing if happens
+ logWarning("BUG, Please report to Apache Kyuubi community")
+ return plan
+ }
+ // The condition whether inject custom resource profile:
+ // - target executors < active executors
+ // - active executors - target executors > min executors
+ val numActiveExecutors = sc.getExecutorIds().length
+ val targetExecutors = (math.ceil(targetCores.toFloat / coresPerExecutor) * factor).toInt
+ .max(1)
+ val hasBenefits = targetExecutors < numActiveExecutors &&
+ (numActiveExecutors - targetExecutors) > minExecutors
+ logInfo(s"The snapshot of current executors view, " +
+ s"active executors: $numActiveExecutors, min executor: $minExecutors, " +
+ s"target executors: $targetExecutors, has benefits: $hasBenefits")
+ if (hasBenefits) {
+ val shuffleId = stage.plan.asInstanceOf[ShuffleExchangeExec].shuffleDependency.shuffleId
+ val numReduce = stage.plan.asInstanceOf[ShuffleExchangeExec].numPartitions
+ // Now, there is only a final rebalance stage waiting to execute and all tasks of previous
+ // stage are finished. Kill redundant existed executors eagerly so the tasks of final
+ // stage can be centralized scheduled.
+ killExecutors(sc, targetExecutors, shuffleId, numReduce)
+ }
+
+ plan
+ }
+
+ /**
+ * The priority of kill executors follow:
+ * 1. kill executor who is younger than other (The older the JIT works better)
+ * 2. kill executor who produces less shuffle data first
+ */
+ private def findExecutorToKill(
+ sc: SparkContext,
+ targetExecutors: Int,
+ shuffleId: Int,
+ numReduce: Int): Seq[String] = {
+ val tracker = SparkEnv.get.mapOutputTracker.asInstanceOf[MapOutputTrackerMaster]
+ val shuffleStatusOpt = tracker.shuffleStatuses.get(shuffleId)
+ if (shuffleStatusOpt.isEmpty) {
+ return Seq.empty
+ }
+ val shuffleStatus = shuffleStatusOpt.get
+ val executorToBlockSize = new mutable.HashMap[String, Long]
+ shuffleStatus.withMapStatuses { mapStatus =>
+ mapStatus.foreach { status =>
+ var i = 0
+ var sum = 0L
+ while (i < numReduce) {
+ sum += status.getSizeForBlock(i)
+ i += 1
+ }
+ executorToBlockSize.getOrElseUpdate(status.location.executorId, sum)
+ }
+ }
+
+ val backend = sc.schedulerBackend.asInstanceOf[CoarseGrainedSchedulerBackend]
+ val executorsWithRegistrationTs = backend.getExecutorsWithRegistrationTs()
+ val existedExecutors = executorsWithRegistrationTs.keys.toSet
+ val expectedNumExecutorToKill = existedExecutors.size - targetExecutors
+ if (expectedNumExecutorToKill < 1) {
+ return Seq.empty
+ }
+
+ val executorIdsToKill = new ArrayBuffer[String]()
+ // We first kill executor who does not hold shuffle block. It would happen because
+ // the last stage is running fast and finished in a short time. The existed executors are
+ // from previous stages that have not been killed by DRA, so we can not find it by tracking
+ // shuffle status.
+ // We should evict executors by their alive time first and retain all of executors which
+ // have better locality for shuffle block.
+ executorsWithRegistrationTs.toSeq.sortBy(_._2).foreach { case (id, _) =>
+ if (executorIdsToKill.length < expectedNumExecutorToKill &&
+ !executorToBlockSize.contains(id)) {
+ executorIdsToKill.append(id)
+ }
+ }
+
+ // Evict the rest executors according to the shuffle block size
+ executorToBlockSize.toSeq.sortBy(_._2).foreach { case (id, _) =>
+ if (executorIdsToKill.length < expectedNumExecutorToKill && existedExecutors.contains(id)) {
+ executorIdsToKill.append(id)
+ }
+ }
+
+ executorIdsToKill.toSeq
+ }
+
+ private def killExecutors(
+ sc: SparkContext,
+ targetExecutors: Int,
+ shuffleId: Int,
+ numReduce: Int): Unit = {
+ val executorAllocationClient = sc.schedulerBackend.asInstanceOf[ExecutorAllocationClient]
+
+ val executorsToKill =
+ if (conf.getConf(KyuubiSQLConf.FINAL_WRITE_STAGE_EAGERLY_KILL_EXECUTORS_KILL_ALL)) {
+ executorAllocationClient.getExecutorIds()
+ } else {
+ findExecutorToKill(sc, targetExecutors, shuffleId, numReduce)
+ }
+ logInfo(s"Request to kill executors, total count ${executorsToKill.size}, " +
+ s"[${executorsToKill.mkString(", ")}].")
+ if (executorsToKill.isEmpty) {
+ return
+ }
+
+ // Note, `SparkContext#killExecutors` does not allow with DRA enabled,
+ // see `https://github.com/apache/spark/pull/20604`.
+ // It may cause the status in `ExecutorAllocationManager` inconsistent with
+ // `CoarseGrainedSchedulerBackend` for a while. But it should be synchronous finally.
+ //
+ // We should adjust target num executors, otherwise `YarnAllocator` might re-request original
+ // target executors if DRA has not updated target executors yet.
+ // Note, DRA would re-adjust executors if there are more tasks to be executed, so we are safe.
+ //
+ // * We kill executor
+ // * YarnAllocator re-request target executors
+ // * DRA can not release executors since they are new added
+ // ----------------------------------------------------------------> timeline
+ executorAllocationClient.killExecutors(
+ executorIds = executorsToKill,
+ adjustTargetNumExecutors = true,
+ countFailures = false,
+ force = false)
+
+ FinalStageResourceManager.getAdjustedTargetExecutors(sc)
+ .filter(_ < targetExecutors).foreach { adjustedExecutors =>
+ val delta = targetExecutors - adjustedExecutors
+ logInfo(s"Target executors after kill ($adjustedExecutors) is lower than required " +
+ s"($targetExecutors). Requesting $delta additional executor(s).")
+ executorAllocationClient.requestExecutors(delta)
+ }
+ }
+
+ @transient private val queryStageOptimizerRules: Seq[Rule[SparkPlan]] = Seq(
+ OptimizeSkewInRebalancePartitions,
+ CoalesceShufflePartitions(session),
+ OptimizeShuffleWithLocalRead)
+}
+
+object FinalStageResourceManager extends Logging {
+
+ private[sql] def getAdjustedTargetExecutors(sc: SparkContext): Option[Int] = {
+ sc.schedulerBackend match {
+ case schedulerBackend: CoarseGrainedSchedulerBackend =>
+ try {
+ val field = classOf[CoarseGrainedSchedulerBackend]
+ .getDeclaredField("requestedTotalExecutorsPerResourceProfile")
+ field.setAccessible(true)
+ schedulerBackend.synchronized {
+ val requestedTotalExecutorsPerResourceProfile =
+ field.get(schedulerBackend).asInstanceOf[mutable.HashMap[ResourceProfile, Int]]
+ val defaultRp = sc.resourceProfileManager.defaultResourceProfile
+ requestedTotalExecutorsPerResourceProfile.get(defaultRp)
+ }
+ } catch {
+ case e: Exception =>
+ logWarning("Failed to get requestedTotalExecutors of Default ResourceProfile", e)
+ None
+ }
+ case _ => None
+ }
+ }
+}
+
+trait FinalRebalanceStageHelper extends AdaptiveSparkPlanHelper {
+ @tailrec
+ final protected def findFinalRebalanceStage(plan: SparkPlan): Option[ShuffleQueryStageExec] = {
+ plan match {
+ case write: DataWritingCommandExec => findFinalRebalanceStage(write.child)
+ case write: V2TableWriteExec => findFinalRebalanceStage(write.child)
+ case write: WriteFilesExec => findFinalRebalanceStage(write.child)
+ case p: ProjectExec => findFinalRebalanceStage(p.child)
+ case f: FilterExec => findFinalRebalanceStage(f.child)
+ case s: SortExec if !s.global => findFinalRebalanceStage(s.child)
+ case stage: ShuffleQueryStageExec
+ if stage.isMaterialized && stage.mapStats.isDefined &&
+ stage.plan.isInstanceOf[ShuffleExchangeExec] &&
+ stage.plan.asInstanceOf[ShuffleExchangeExec].shuffleOrigin != ENSURE_REQUIREMENTS =>
+ Some(stage)
+ case _ => None
+ }
+ }
+
+ final protected def hasTableCache(plan: SparkPlan): Boolean = {
+ find(plan) {
+ case _: InMemoryTableScanExec => true
+ case _ => false
+ }.isDefined
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/InjectCustomResourceProfile.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/InjectCustomResourceProfile.scala
new file mode 100644
index 00000000000..64421d6bfab
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/InjectCustomResourceProfile.scala
@@ -0,0 +1,60 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.{CustomResourceProfileExec, SparkPlan}
+import org.apache.spark.sql.execution.adaptive._
+
+import org.apache.kyuubi.sql.{KyuubiSQLConf, WriteUtils}
+
+/**
+ * Inject custom resource profile for final write stage, so we can specify custom
+ * executor resource configs.
+ */
+case class InjectCustomResourceProfile(session: SparkSession)
+ extends Rule[SparkPlan] with FinalRebalanceStageHelper {
+ override def apply(plan: SparkPlan): SparkPlan = {
+ if (!conf.getConf(KyuubiSQLConf.FINAL_WRITE_STAGE_RESOURCE_ISOLATION_ENABLED)) {
+ return plan
+ }
+
+ if (!WriteUtils.isWrite(session, plan)) {
+ return plan
+ }
+
+ val stage = findFinalRebalanceStage(plan)
+ if (stage.isEmpty) {
+ return plan
+ }
+
+ // TODO: Ideally, We can call `CoarseGrainedSchedulerBackend.requestTotalExecutors` eagerly
+ // to reduce the task submit pending time, but it may lose task locality.
+ //
+ // By default, it would request executors when catch stage submit event.
+ injectCustomResourceProfile(plan, stage.get.id)
+ }
+
+ private def injectCustomResourceProfile(plan: SparkPlan, id: Int): SparkPlan = {
+ plan match {
+ case stage: ShuffleQueryStageExec if stage.id == id =>
+ CustomResourceProfileExec(stage)
+ case _ => plan.mapChildren(child => injectCustomResourceProfile(child, id))
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/PruneFileSourcePartitionHelper.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/PruneFileSourcePartitionHelper.scala
new file mode 100644
index 00000000000..ce496eb474c
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/PruneFileSourcePartitionHelper.scala
@@ -0,0 +1,46 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.catalyst.expressions.{AttributeReference, AttributeSet, Expression, ExpressionSet, PredicateHelper, SubqueryExpression}
+import org.apache.spark.sql.catalyst.plans.logical.LeafNode
+import org.apache.spark.sql.execution.datasources.DataSourceStrategy
+import org.apache.spark.sql.types.StructType
+
+trait PruneFileSourcePartitionHelper extends PredicateHelper {
+
+ def getPartitionKeyFiltersAndDataFilters(
+ sparkSession: SparkSession,
+ relation: LeafNode,
+ partitionSchema: StructType,
+ filters: Seq[Expression],
+ output: Seq[AttributeReference]): (ExpressionSet, Seq[Expression]) = {
+ val normalizedFilters = DataSourceStrategy.normalizeExprs(
+ filters.filter(f => f.deterministic && !SubqueryExpression.hasSubquery(f)),
+ output)
+ val partitionColumns =
+ relation.resolve(partitionSchema, sparkSession.sessionState.analyzer.resolver)
+ val partitionSet = AttributeSet(partitionColumns)
+ val (partitionFilters, dataFilters) = normalizedFilters.partition(f =>
+ f.references.subsetOf(partitionSet))
+ val extraPartitionFilter =
+ dataFilters.flatMap(extractPredicatesWithinOutputSet(_, partitionSet))
+
+ (ExpressionSet(partitionFilters ++ extraPartitionFilter), dataFilters)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/execution/CustomResourceProfileExec.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/execution/CustomResourceProfileExec.scala
new file mode 100644
index 00000000000..043d6496b22
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/execution/CustomResourceProfileExec.scala
@@ -0,0 +1,117 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.execution
+
+import org.apache.spark.network.util.{ByteUnit, JavaUtils}
+import org.apache.spark.rdd.RDD
+import org.apache.spark.resource.{ExecutorResourceRequests, ResourceProfileBuilder}
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, SortOrder}
+import org.apache.spark.sql.catalyst.plans.physical.Partitioning
+import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics}
+import org.apache.spark.sql.vectorized.ColumnarBatch
+import org.apache.spark.util.Utils
+
+import org.apache.kyuubi.sql.KyuubiSQLConf._
+
+/**
+ * This node wraps the final executed plan and inject custom resource profile to the RDD.
+ * It assumes that, the produced RDD would create the `ResultStage` in `DAGScheduler`,
+ * so it makes resource isolation between previous and final stage.
+ *
+ * Note that, Spark does not support config `minExecutors` for each resource profile.
+ * Which means, it would retain `minExecutors` for each resource profile.
+ * So, suggest set `spark.dynamicAllocation.minExecutors` to 0 if enable this feature.
+ */
+case class CustomResourceProfileExec(child: SparkPlan) extends UnaryExecNode {
+ override def output: Seq[Attribute] = child.output
+ override def outputPartitioning: Partitioning = child.outputPartitioning
+ override def outputOrdering: Seq[SortOrder] = child.outputOrdering
+ override def supportsColumnar: Boolean = child.supportsColumnar
+ override def supportsRowBased: Boolean = child.supportsRowBased
+ override protected def doCanonicalize(): SparkPlan = child.canonicalized
+
+ private val executorCores = conf.getConf(FINAL_WRITE_STAGE_EXECUTOR_CORES).getOrElse(
+ sparkContext.getConf.getInt("spark.executor.cores", 1))
+ private val executorMemory = conf.getConf(FINAL_WRITE_STAGE_EXECUTOR_MEMORY).getOrElse(
+ sparkContext.getConf.get("spark.executor.memory", "2G"))
+ private val executorMemoryOverhead =
+ conf.getConf(FINAL_WRITE_STAGE_EXECUTOR_MEMORY_OVERHEAD)
+ .getOrElse(sparkContext.getConf.get("spark.executor.memoryOverhead", "1G"))
+ private val executorOffHeapMemory =
+ if (sparkContext.getConf.getBoolean("spark.memory.offHeap.enabled", false)) {
+ conf.getConf(FINAL_WRITE_STAGE_EXECUTOR_OFF_HEAP_MEMORY)
+ } else {
+ None
+ }
+
+ override lazy val metrics: Map[String, SQLMetric] = {
+ val base = Map(
+ "executorCores" -> SQLMetrics.createMetric(sparkContext, "executor cores"),
+ "executorMemory" -> SQLMetrics.createMetric(sparkContext, "executor memory (MiB)"),
+ "executorMemoryOverhead" -> SQLMetrics.createMetric(
+ sparkContext,
+ "executor memory overhead (MiB)"))
+ val addition = executorOffHeapMemory.map(_ =>
+ "executorOffHeapMemory" ->
+ SQLMetrics.createMetric(sparkContext, "executor off heap memory (MiB)")).toMap
+ base ++ addition
+ }
+
+ private def wrapResourceProfile[T](rdd: RDD[T]): RDD[T] = {
+ if (Utils.isTesting) {
+ // do nothing for local testing
+ return rdd
+ }
+
+ metrics("executorCores") += executorCores
+ metrics("executorMemory") += JavaUtils.byteStringAs(executorMemory, ByteUnit.MiB)
+ metrics("executorMemoryOverhead") += JavaUtils.byteStringAs(
+ executorMemoryOverhead,
+ ByteUnit.MiB)
+ executorOffHeapMemory.foreach(m =>
+ metrics("executorOffHeapMemory") += JavaUtils.byteStringAs(m, ByteUnit.MiB))
+
+ val executionId = sparkContext.getLocalProperty(SQLExecution.EXECUTION_ID_KEY)
+ SQLMetrics.postDriverMetricUpdates(sparkContext, executionId, metrics.values.toSeq)
+
+ val resourceProfileBuilder = new ResourceProfileBuilder()
+ val executorResourceRequests = new ExecutorResourceRequests()
+ executorResourceRequests.cores(executorCores)
+ executorResourceRequests.memory(executorMemory)
+ executorResourceRequests.memoryOverhead(executorMemoryOverhead)
+ executorOffHeapMemory.foreach(executorResourceRequests.offHeapMemory)
+ resourceProfileBuilder.require(executorResourceRequests)
+ rdd.withResources(resourceProfileBuilder.build())
+ rdd
+ }
+
+ override protected def doExecute(): RDD[InternalRow] = {
+ val rdd = child.execute()
+ wrapResourceProfile(rdd)
+ }
+
+ override protected def doExecuteColumnar(): RDD[ColumnarBatch] = {
+ val rdd = child.executeColumnar()
+ wrapResourceProfile(rdd)
+ }
+
+ override protected def withNewChildInternal(newChild: SparkPlan): SparkPlan = {
+ this.copy(child = newChild)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/hive/HiveSparkPlanHelper.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/hive/HiveSparkPlanHelper.scala
new file mode 100644
index 00000000000..aa9a0459616
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/main/scala/org/apache/spark/sql/hive/HiveSparkPlanHelper.scala
@@ -0,0 +1,21 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.spark.sql.hive
+
+object HiveSparkPlanHelper {
+ type HiveTableScanExec = org.apache.spark.sql.hive.execution.HiveTableScanExec
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/resources/log4j2-test.xml b/extensions/spark/kyuubi-extension-spark-4-2/src/test/resources/log4j2-test.xml
new file mode 100644
index 00000000000..3110216c17c
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/resources/log4j2-test.xml
@@ -0,0 +1,43 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/DropIgnoreNonexistentSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/DropIgnoreNonexistentSuite.scala
new file mode 100644
index 00000000000..37ac61f8e52
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/DropIgnoreNonexistentSuite.scala
@@ -0,0 +1,62 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.spark.sql
+
+import org.apache.spark.sql.catalyst.plans.logical.{DropNamespace, DropTable}
+import org.apache.spark.sql.execution.command._
+
+import org.apache.kyuubi.sql.KyuubiSQLConf
+
+class DropIgnoreNonexistentSuite extends KyuubiSparkSQLExtensionTest {
+
+ test("drop ignore nonexistent") {
+ withSQLConf(KyuubiSQLConf.DROP_IGNORE_NONEXISTENT.key -> "true") {
+ // drop nonexistent database
+ val df1 = sql("DROP DATABASE nonexistent_database")
+ assert(df1.queryExecution.analyzed.asInstanceOf[DropNamespace].ifExists == true)
+
+ // drop nonexistent table
+ val df2 = sql("DROP TABLE nonexistent_table")
+ assert(df2.queryExecution.analyzed.asInstanceOf[DropTable].ifExists == true)
+
+ // drop nonexistent view
+ val df3 = sql("DROP VIEW nonexistent_view")
+ assert(df3.queryExecution.analyzed.asInstanceOf[DropTableCommand].isView == true &&
+ df3.queryExecution.analyzed.asInstanceOf[DropTableCommand].ifExists == true)
+
+ // drop nonexistent function
+ // SPARK-54866 (4.2.0) eagerly resolves DROP FUNCTION to the v1 DropFunctionCommand
+ // during analysis (deferring the existence check to execution), so the post-hoc rule
+ // no longer sees an unresolved DropFunction to turn into a NoopCommand. Setting
+ // ifExists=true preserves the "drop nonexistent does not fail" contract.
+ val df4 = sql("DROP FUNCTION nonexistent_function")
+ assert(df4.queryExecution.analyzed.asInstanceOf[DropFunctionCommand].ifExists == true)
+
+ // drop nonexistent temporary function
+ val df5 = sql("DROP TEMPORARY FUNCTION nonexistent_temp_function")
+ assert(df5.queryExecution.analyzed.asInstanceOf[DropFunctionCommand].ifExists == true)
+
+ // drop nonexistent PARTITION
+ withTable("test") {
+ sql("CREATE TABLE IF NOT EXISTS test(i int) PARTITIONED BY (p int)")
+ val df5 = sql("ALTER TABLE test DROP PARTITION (p = 1)")
+ assert(df5.queryExecution.analyzed
+ .asInstanceOf[AlterTableDropPartitionCommand].ifExists == true)
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/DynamicShufflePartitionsSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/DynamicShufflePartitionsSuite.scala
new file mode 100644
index 00000000000..6d252671d35
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/DynamicShufflePartitionsSuite.scala
@@ -0,0 +1,155 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.spark.sql
+
+import org.apache.spark.sql.catalyst.TableIdentifier
+import org.apache.spark.sql.execution.{CommandResultExec, SparkPlan}
+import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, ShuffleQueryStageExec}
+import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, ShuffleExchangeExec}
+import org.apache.spark.sql.hive.HiveUtils.CONVERT_METASTORE_PARQUET
+import org.apache.spark.sql.internal.SQLConf._
+
+import org.apache.kyuubi.sql.KyuubiSQLConf.{DYNAMIC_SHUFFLE_PARTITIONS, DYNAMIC_SHUFFLE_PARTITIONS_MAX_NUM}
+
+class DynamicShufflePartitionsSuite extends KyuubiSparkSQLExtensionTest {
+
+ override protected def beforeAll(): Unit = {
+ super.beforeAll()
+ setupData()
+ }
+
+ test("test dynamic shuffle partitions") {
+ def collectExchanges(plan: SparkPlan): Seq[ShuffleExchangeExec] = {
+ plan match {
+ case p: CommandResultExec => collectExchanges(p.commandPhysicalPlan)
+ case p: AdaptiveSparkPlanExec => collectExchanges(p.finalPhysicalPlan)
+ case p: ShuffleQueryStageExec => collectExchanges(p.plan)
+ case p: ShuffleExchangeExec => p +: collectExchanges(p.child)
+ case p => p.children.flatMap(collectExchanges)
+ }
+ }
+
+ // datasource scan
+ withTable("table1", "table2", "table3") {
+ sql("create table table1 stored as parquet as select c1, c2 from t1")
+ sql("create table table2 stored as parquet as select c1, c2 from t2")
+ sql("create table table3 (c1 int, c2 string) stored as parquet")
+ sql("ANALYZE TABLE table1 COMPUTE STATISTICS")
+ sql("ANALYZE TABLE table2 COMPUTE STATISTICS")
+
+ val initialPartitionNum: Int = 2
+ val advisoryPartitionSizeInBytes: Long = 500
+
+ val t1Size = spark.sessionState.catalog.getTableMetadata(TableIdentifier("table1"))
+ .stats.get.sizeInBytes.toLong
+ val t2Size = spark.sessionState.catalog.getTableMetadata(TableIdentifier("table2"))
+ .stats.get.sizeInBytes.toLong
+ val scanSize = t1Size + t2Size
+ val expectedJoinPartitionNum = Math.ceil(scanSize.toDouble / advisoryPartitionSizeInBytes)
+
+ Seq(false, true).foreach { dynamicShufflePartitions =>
+ val maxDynamicShufflePartitions = if (dynamicShufflePartitions) {
+ Seq(8, 2000)
+ } else {
+ Seq(2000)
+ }
+ maxDynamicShufflePartitions.foreach { maxDynamicShufflePartitionNum =>
+ withSQLConf(
+ DYNAMIC_SHUFFLE_PARTITIONS.key -> dynamicShufflePartitions.toString,
+ DYNAMIC_SHUFFLE_PARTITIONS_MAX_NUM.key -> maxDynamicShufflePartitionNum.toString,
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> initialPartitionNum.toString,
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> advisoryPartitionSizeInBytes.toString) {
+ val df = sql("insert overwrite table3 " +
+ " select a.c1 as c1, b.c2 as c2 from table1 a join table2 b on a.c1 = b.c1")
+
+ val exchanges = collectExchanges(df.queryExecution.executedPlan)
+ val (joinExchanges, rebalanceExchanges) = exchanges
+ .partition(_.shuffleOrigin == ENSURE_REQUIREMENTS)
+ assert(joinExchanges.size == 2)
+ if (dynamicShufflePartitions) {
+ joinExchanges.foreach { e =>
+ val expected = Math.min(expectedJoinPartitionNum, maxDynamicShufflePartitionNum)
+ assert(e.outputPartitioning.numPartitions == expected)
+ }
+ } else {
+ joinExchanges.foreach { e =>
+ assert(e.outputPartitioning.numPartitions == initialPartitionNum)
+ }
+ }
+
+ assert(rebalanceExchanges.size == 1)
+ if (dynamicShufflePartitions) {
+ if (maxDynamicShufflePartitionNum == 8) {
+ // shuffle query size: 1424 451 (the size may change with spark version updates
+ // or shuffle configuration updates)
+ val expected = Math.min(4, maxDynamicShufflePartitionNum)
+ assert(rebalanceExchanges.head.outputPartitioning.numPartitions == expected)
+ } else {
+ // shuffle query size: 2057 664 (the size may change with spark version updates
+ // or shuffle configuration updates)
+ val expected = Math.min(6, maxDynamicShufflePartitionNum)
+ assert(rebalanceExchanges.head.outputPartitioning.numPartitions == expected)
+ }
+ } else {
+ assert(
+ rebalanceExchanges.head.outputPartitioning.numPartitions == initialPartitionNum)
+ }
+ }
+
+ // hive table scan
+ withSQLConf(
+ DYNAMIC_SHUFFLE_PARTITIONS.key -> dynamicShufflePartitions.toString,
+ DYNAMIC_SHUFFLE_PARTITIONS_MAX_NUM.key -> maxDynamicShufflePartitionNum.toString,
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> initialPartitionNum.toString,
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> advisoryPartitionSizeInBytes.toString,
+ CONVERT_METASTORE_PARQUET.key -> "false") {
+ val df = sql("insert overwrite table3 " +
+ " select a.c1 as c1, b.c2 as c2 from table1 a join table2 b on a.c1 = b.c1")
+
+ val exchanges = collectExchanges(df.queryExecution.executedPlan)
+ val (joinExchanges, rebalanceExchanges) = exchanges
+ .partition(_.shuffleOrigin == ENSURE_REQUIREMENTS)
+ assert(joinExchanges.size == 2)
+ if (dynamicShufflePartitions) {
+ joinExchanges.foreach { e =>
+ val expected = Math.min(expectedJoinPartitionNum, maxDynamicShufflePartitionNum)
+ assert(e.outputPartitioning.numPartitions == expected)
+ }
+ } else {
+ joinExchanges.foreach { e =>
+ assert(e.outputPartitioning.numPartitions == initialPartitionNum)
+ }
+ }
+ // shuffle query size: 5154 720 (the size may change with spark version updates
+ // or shuffle configuration updates)
+ assert(rebalanceExchanges.size == 1)
+ if (dynamicShufflePartitions) {
+ val expected = Math.min(12, maxDynamicShufflePartitionNum)
+ assert(rebalanceExchanges.head.outputPartitioning.numPartitions == expected)
+ } else {
+ assert(rebalanceExchanges.head.outputPartitioning.numPartitions ==
+ initialPartitionNum)
+ }
+ }
+ }
+ }
+ }
+ }
+
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/FinalStageConfigIsolationSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/FinalStageConfigIsolationSuite.scala
new file mode 100644
index 00000000000..96c8ae6e8b0
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/FinalStageConfigIsolationSuite.scala
@@ -0,0 +1,203 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.execution.adaptive.{AQEShuffleReadExec, QueryStageExec}
+import org.apache.spark.sql.internal.SQLConf
+
+import org.apache.kyuubi.sql.{FinalStageConfigIsolation, KyuubiSQLConf}
+
+class FinalStageConfigIsolationSuite extends KyuubiSparkSQLExtensionTest {
+ override protected def beforeAll(): Unit = {
+ super.beforeAll()
+ setupData()
+ }
+
+ test("final stage config set reset check") {
+ withSQLConf(
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION.key -> "true",
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION_WRITE_ONLY.key -> "false",
+ "spark.sql.finalStage.adaptive.coalescePartitions.minPartitionNum" -> "1",
+ "spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes" -> "100") {
+ // use loop to double check final stage config doesn't affect the sql query each other
+ (1 to 3).foreach { _ =>
+ sql("SELECT COUNT(*) FROM VALUES(1) as t(c)").collect()
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.previousStage.adaptive.coalescePartitions.minPartitionNum") ===
+ FinalStageConfigIsolation.INTERNAL_UNSET_CONFIG_TAG)
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.adaptive.coalescePartitions.minPartitionNum") ===
+ "1")
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.finalStage.adaptive.coalescePartitions.minPartitionNum") ===
+ "1")
+
+ // 64MB
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.previousStage.adaptive.advisoryPartitionSizeInBytes") ===
+ "67108864b")
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.adaptive.advisoryPartitionSizeInBytes") ===
+ "100")
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes") ===
+ "100")
+ }
+
+ sql("SET spark.sql.adaptive.advisoryPartitionSizeInBytes=1")
+ assert(spark.sessionState.conf.getConfString(
+ "spark.sql.adaptive.advisoryPartitionSizeInBytes") ===
+ "1")
+ assert(!spark.sessionState.conf.contains(
+ "spark.sql.previousStage.adaptive.advisoryPartitionSizeInBytes"))
+
+ sql("SET a=1")
+ assert(spark.sessionState.conf.getConfString("a") === "1")
+
+ sql("RESET spark.sql.adaptive.coalescePartitions.minPartitionNum")
+ assert(!spark.sessionState.conf.contains(
+ "spark.sql.adaptive.coalescePartitions.minPartitionNum"))
+ assert(!spark.sessionState.conf.contains(
+ "spark.sql.previousStage.adaptive.coalescePartitions.minPartitionNum"))
+
+ sql("RESET a")
+ assert(!spark.sessionState.conf.contains("a"))
+ }
+ }
+
+ test("final stage config isolation") {
+ def checkPartitionNum(
+ sqlString: String,
+ previousPartitionNum: Int,
+ finalPartitionNum: Int): Unit = {
+ val df = sql(sqlString)
+ df.collect()
+ val shuffleReaders = collect(df.queryExecution.executedPlan) {
+ case customShuffleReader: AQEShuffleReadExec => customShuffleReader
+ }
+ assert(shuffleReaders.nonEmpty)
+ // reorder stage by stage id to ensure we get the right stage
+ val sortedShuffleReaders = shuffleReaders.sortWith {
+ case (s1, s2) =>
+ s1.child.asInstanceOf[QueryStageExec].id < s2.child.asInstanceOf[QueryStageExec].id
+ }
+ if (sortedShuffleReaders.length > 1) {
+ assert(sortedShuffleReaders.head.partitionSpecs.length === previousPartitionNum)
+ }
+ assert(sortedShuffleReaders.last.partitionSpecs.length === finalPartitionNum)
+ assert(df.rdd.partitions.length === finalPartitionNum)
+ }
+
+ withSQLConf(
+ SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1",
+ SQLConf.SHUFFLE_PARTITIONS.key -> "3",
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION.key -> "true",
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION_WRITE_ONLY.key -> "false",
+ "spark.sql.adaptive.advisoryPartitionSizeInBytes" -> "1",
+ "spark.sql.adaptive.coalescePartitions.minPartitionSize" -> "1",
+ "spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes" -> "10000000") {
+
+ // use loop to double check final stage config doesn't affect the sql query each other
+ (1 to 3).foreach { _ =>
+ checkPartitionNum(
+ "SELECT c1, count(*) FROM t1 GROUP BY c1",
+ 1,
+ 1)
+
+ checkPartitionNum(
+ "SELECT c2, count(*) FROM (SELECT c1, count(*) as c2 FROM t1 GROUP BY c1) GROUP BY c2",
+ 3,
+ 1)
+
+ checkPartitionNum(
+ "SELECT t1.c1, count(*) FROM t1 JOIN t2 ON t1.c2 = t2.c2 GROUP BY t1.c1",
+ 3,
+ 1)
+
+ checkPartitionNum(
+ """
+ | SELECT /*+ REPARTITION */
+ | t1.c1, count(*) FROM t1
+ | JOIN t2 ON t1.c2 = t2.c2
+ | JOIN t3 ON t1.c1 = t3.c1
+ | GROUP BY t1.c1
+ |""".stripMargin,
+ 3,
+ 1)
+
+ // one shuffle reader
+ checkPartitionNum(
+ """
+ | SELECT /*+ BROADCAST(t1) */
+ | t1.c1, t2.c2 FROM t1
+ | JOIN t2 ON t1.c2 = t2.c2
+ | DISTRIBUTE BY c1
+ |""".stripMargin,
+ 1,
+ 1)
+
+ // test ReusedExchange
+ checkPartitionNum(
+ """
+ |SELECT /*+ REPARTITION */ t0.c2 FROM (
+ |SELECT t1.c1, (count(*) + c1) as c2 FROM t1 GROUP BY t1.c1
+ |) t0 JOIN (
+ |SELECT t1.c1, (count(*) + c1) as c2 FROM t1 GROUP BY t1.c1
+ |) t1 ON t0.c2 = t1.c2
+ |""".stripMargin,
+ 3,
+ 1)
+
+ // one shuffle reader
+ checkPartitionNum(
+ """
+ |SELECT t0.c1 FROM (
+ |SELECT t1.c1 FROM t1 GROUP BY t1.c1
+ |) t0 JOIN (
+ |SELECT t1.c1 FROM t1 GROUP BY t1.c1
+ |) t1 ON t0.c1 = t1.c1
+ |""".stripMargin,
+ 1,
+ 1)
+ }
+ }
+ }
+
+ test("final stage config isolation write only") {
+ withSQLConf(
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION.key -> "true",
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION_WRITE_ONLY.key -> "true",
+ "spark.sql.finalStage.adaptive.advisoryPartitionSizeInBytes" -> "7") {
+ sql("set spark.sql.adaptive.advisoryPartitionSizeInBytes=5")
+ sql("SELECT * FROM t1").count()
+ assert(spark.conf.getOption("spark.sql.adaptive.advisoryPartitionSizeInBytes")
+ .contains("5"))
+
+ withTable("tmp") {
+ sql("CREATE TABLE t1 USING PARQUET SELECT /*+ repartition */ 1 AS c1, 'a' AS c2")
+ assert(spark.conf.getOption("spark.sql.adaptive.advisoryPartitionSizeInBytes")
+ .contains("7"))
+ }
+
+ sql("SELECT * FROM t1").count()
+ assert(spark.conf.getOption("spark.sql.adaptive.advisoryPartitionSizeInBytes")
+ .contains("5"))
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/FinalStageResourceManagerSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/FinalStageResourceManagerSuite.scala
new file mode 100644
index 00000000000..4b9991ef6f2
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/FinalStageResourceManagerSuite.scala
@@ -0,0 +1,62 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.SparkConf
+import org.scalatest.time.{Minutes, Span}
+
+import org.apache.kyuubi.sql.KyuubiSQLConf
+import org.apache.kyuubi.tags.SparkLocalClusterTest
+
+@SparkLocalClusterTest
+class FinalStageResourceManagerSuite extends KyuubiSparkSQLExtensionTest {
+
+ override def sparkConf(): SparkConf = {
+ // It is difficult to run spark in local-cluster mode when spark.testing is set.
+ sys.props.remove("spark.testing")
+
+ super.sparkConf().set("spark.master", "local-cluster[3, 1, 1024]")
+ .set("spark.dynamicAllocation.enabled", "true")
+ .set("spark.dynamicAllocation.initialExecutors", "3")
+ .set("spark.dynamicAllocation.minExecutors", "1")
+ .set("spark.dynamicAllocation.shuffleTracking.enabled", "true")
+ .set(KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION.key, "true")
+ .set(KyuubiSQLConf.FINAL_WRITE_STAGE_EAGERLY_KILL_EXECUTORS_ENABLED.key, "true")
+ }
+
+ test("[KYUUBI #5136][Bug] Final Stage hangs forever") {
+ // Prerequisite to reproduce the bug:
+ // 1. Dynamic allocation is enabled.
+ // 2. Dynamic allocation min executors is 1.
+ // 3. target executors < active executors.
+ // 4. No active executor is left after FinalStageResourceManager killed executors.
+ // This is possible because FinalStageResourceManager retained executors may already be
+ // requested to be killed but not died yet.
+ // 5. Final Stage required executors is 1.
+ withSQLConf(
+ (KyuubiSQLConf.FINAL_WRITE_STAGE_EAGERLY_KILL_EXECUTORS_KILL_ALL.key, "true")) {
+ withTable("final_stage") {
+ eventually(timeout(Span(10, Minutes))) {
+ sql(
+ "CREATE TABLE final_stage AS SELECT id, count(*) as num FROM (SELECT 0 id) GROUP BY id")
+ }
+ assert(FinalStageResourceManager.getAdjustedTargetExecutors(spark.sparkContext).get == 1)
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/InjectResourceProfileSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/InjectResourceProfileSuite.scala
new file mode 100644
index 00000000000..b0767b18708
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/InjectResourceProfileSuite.scala
@@ -0,0 +1,79 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.scheduler.{SparkListener, SparkListenerEvent}
+import org.apache.spark.sql.execution.ui.SparkListenerSQLAdaptiveExecutionUpdate
+
+import org.apache.kyuubi.sql.KyuubiSQLConf
+
+class InjectResourceProfileSuite extends KyuubiSparkSQLExtensionTest {
+ private def checkCustomResourceProfile(sqlString: String, exists: Boolean): Unit = {
+ @volatile var lastEvent: SparkListenerSQLAdaptiveExecutionUpdate = null
+ val listener = new SparkListener {
+ override def onOtherEvent(event: SparkListenerEvent): Unit = {
+ event match {
+ case e: SparkListenerSQLAdaptiveExecutionUpdate => lastEvent = e
+ case _ =>
+ }
+ }
+ }
+
+ spark.sparkContext.addSparkListener(listener)
+ try {
+ sql(sqlString).collect()
+ spark.sparkContext.listenerBus.waitUntilEmpty()
+ assert(lastEvent != null)
+ var current = lastEvent.sparkPlanInfo
+ var shouldStop = false
+ while (!shouldStop) {
+ if (current.nodeName != "CustomResourceProfile") {
+ if (current.children.isEmpty) {
+ assert(!exists)
+ shouldStop = true
+ } else {
+ current = current.children.head
+ }
+ } else {
+ assert(exists)
+ shouldStop = true
+ }
+ }
+ } finally {
+ spark.sparkContext.removeSparkListener(listener)
+ }
+ }
+
+ test("Inject resource profile") {
+ withTable("t") {
+ withSQLConf(
+ "spark.sql.adaptive.forceApply" -> "true",
+ KyuubiSQLConf.FINAL_STAGE_CONFIG_ISOLATION.key -> "true",
+ KyuubiSQLConf.FINAL_WRITE_STAGE_RESOURCE_ISOLATION_ENABLED.key -> "true") {
+
+ sql("CREATE TABLE t (c1 int, c2 string) USING PARQUET")
+
+ checkCustomResourceProfile("INSERT INTO TABLE t VALUES(1, 'a')", false)
+ checkCustomResourceProfile("SELECT 1", false)
+ checkCustomResourceProfile(
+ "INSERT INTO TABLE t SELECT /*+ rebalance */ * FROM VALUES(1, 'a')",
+ true)
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/InsertShuffleNodeBeforeJoinSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/InsertShuffleNodeBeforeJoinSuite.scala
new file mode 100644
index 00000000000..b6b07a5222f
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/InsertShuffleNodeBeforeJoinSuite.scala
@@ -0,0 +1,98 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.SparkConf
+import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, ShuffleExchangeLike}
+import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf}
+
+import org.apache.kyuubi.sql.KyuubiSQLConf
+
+class InsertShuffleNodeBeforeJoinSuite extends KyuubiSparkSQLExtensionTest {
+ override protected def beforeAll(): Unit = {
+ super.beforeAll()
+ setupData()
+ }
+
+ override def sparkConf(): SparkConf = {
+ super.sparkConf()
+ .set(
+ StaticSQLConf.SPARK_SESSION_EXTENSIONS.key,
+ "org.apache.kyuubi.sql.KyuubiSparkSQLExtension")
+ }
+
+ test("force shuffle before join") {
+ def checkShuffleNodeNum(sqlString: String, num: Int): Unit = {
+ var expectedResult: Seq[Row] = Seq.empty
+ withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
+ expectedResult = sql(sqlString).collect()
+ }
+ val df = sql(sqlString)
+ checkAnswer(df, expectedResult)
+ assert(
+ collect(df.queryExecution.executedPlan) {
+ case shuffle: ShuffleExchangeLike if shuffle.shuffleOrigin == ENSURE_REQUIREMENTS =>
+ shuffle
+ }.size == num)
+ }
+
+ withSQLConf(
+ SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ KyuubiSQLConf.FORCE_SHUFFLE_BEFORE_JOIN.key -> "true") {
+ Seq("SHUFFLE_HASH", "MERGE").foreach { joinHint =>
+ // positive case
+ checkShuffleNodeNum(
+ s"""
+ |SELECT /*+ $joinHint(t2, t3) */ t1.c1, t1.c2, t2.c1, t3.c1 from t1
+ | JOIN t2 ON t1.c1 = t2.c1
+ | JOIN t3 ON t1.c1 = t3.c1
+ | """.stripMargin,
+ 4)
+
+ // negative case
+ checkShuffleNodeNum(
+ s"""
+ |SELECT /*+ $joinHint(t2, t3) */ t1.c1, t1.c2, t2.c1, t3.c1 from t1
+ | JOIN t2 ON t1.c1 = t2.c1
+ | JOIN t3 ON t1.c2 = t3.c2
+ | """.stripMargin,
+ 4)
+ }
+
+ checkShuffleNodeNum(
+ """
+ |SELECT t1.c1, t2.c1, t3.c2 from t1
+ | JOIN t2 ON t1.c1 = t2.c1
+ | JOIN (
+ | SELECT c2, count(*) FROM t1 GROUP BY c2
+ | ) t3 ON t1.c1 = t3.c2
+ | """.stripMargin,
+ 5)
+
+ checkShuffleNodeNum(
+ """
+ |SELECT t1.c1, t2.c1, t3.c1 from t1
+ | JOIN t2 ON t1.c1 = t2.c1
+ | JOIN (
+ | SELECT c1, count(*) FROM t1 GROUP BY c1
+ | ) t3 ON t1.c1 = t3.c1
+ | """.stripMargin,
+ 5)
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/KyuubiSparkSQLExtensionTest.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/KyuubiSparkSQLExtensionTest.scala
new file mode 100644
index 00000000000..d1c439d2231
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/KyuubiSparkSQLExtensionTest.scala
@@ -0,0 +1,122 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.spark.sql
+
+import org.apache.hadoop.hive.conf.HiveConf.ConfVars
+import org.apache.spark.SparkConf
+import org.apache.spark.sql.classic.SparkSession
+import org.apache.spark.sql.execution.QueryExecution
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
+import org.apache.spark.sql.execution.command.{DataWritingCommand, DataWritingCommandExec}
+import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf}
+import org.apache.spark.sql.test.SQLTestData.TestData
+import org.apache.spark.sql.test.SQLTestUtils
+import org.apache.spark.sql.util.QueryExecutionListener
+import org.apache.spark.util.Utils
+
+trait KyuubiSparkSQLExtensionTest extends QueryTest
+ with SQLTestUtils
+ with AdaptiveSparkPlanHelper {
+ sys.props.put("spark.testing", "1")
+
+ private var _spark: Option[SparkSession] = None
+ protected def spark: SparkSession = _spark.getOrElse {
+ throw new RuntimeException("test spark session don't initial before using it.")
+ }
+
+ override protected def beforeAll(): Unit = {
+ if (_spark.isEmpty) {
+ _spark = Option(SparkSession.builder()
+ .master("local[1]")
+ .config(sparkConf())
+ .enableHiveSupport()
+ .getOrCreate())
+ }
+ super.beforeAll()
+ }
+
+ override protected def afterAll(): Unit = {
+ super.afterAll()
+ cleanupData()
+ _spark.foreach(_.stop)
+ }
+
+ protected def setupData(): Unit = {
+ val self = spark
+ import self.implicits._
+ spark.sparkContext.parallelize(
+ (1 to 100).map(i => TestData(i, i.toString)),
+ 10)
+ .toDF("c1", "c2").createOrReplaceTempView("t1")
+ spark.sparkContext.parallelize(
+ (1 to 10).map(i => TestData(i, i.toString)),
+ 5)
+ .toDF("c1", "c2").createOrReplaceTempView("t2")
+ spark.sparkContext.parallelize(
+ (1 to 50).map(i => TestData(i, i.toString)),
+ 2)
+ .toDF("c1", "c2").createOrReplaceTempView("t3")
+ }
+
+ private def cleanupData(): Unit = {
+ spark.sql("DROP VIEW IF EXISTS t1")
+ spark.sql("DROP VIEW IF EXISTS t2")
+ spark.sql("DROP VIEW IF EXISTS t3")
+ }
+
+ def sparkConf(): SparkConf = {
+ val basePath = Utils.createTempDir() + "/" + getClass.getCanonicalName
+ val metastorePath = basePath + "/metastore_db"
+ val warehousePath = basePath + "/warehouse"
+ new SparkConf()
+ .set(
+ StaticSQLConf.SPARK_SESSION_EXTENSIONS.key,
+ "org.apache.kyuubi.sql.KyuubiSparkSQLExtension")
+ .set(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key, "true")
+ .set("spark.hadoop.hive.exec.dynamic.partition.mode", "nonstrict")
+ .set("spark.hadoop.hive.metastore.client.capability.check", "false")
+ .set(
+ ConfVars.METASTORECONNECTURLKEY.varname,
+ s"jdbc:derby:;databaseName=$metastorePath;create=true")
+ .set(StaticSQLConf.WAREHOUSE_PATH, warehousePath)
+ .set("spark.ui.enabled", "false")
+ }
+
+ def withListener(sqlString: String)(callback: DataWritingCommand => Unit): Unit = {
+ withListener(sql(sqlString))(callback)
+ }
+
+ def withListener(df: => DataFrame)(callback: DataWritingCommand => Unit): Unit = {
+ val listener = new QueryExecutionListener {
+ override def onFailure(f: String, qe: QueryExecution, e: Exception): Unit = {}
+
+ override def onSuccess(funcName: String, qe: QueryExecution, duration: Long): Unit = {
+ qe.executedPlan match {
+ case write: DataWritingCommandExec => callback(write.cmd)
+ case _ =>
+ }
+ }
+ }
+ spark.listenerManager.register(listener)
+ try {
+ df.collect()
+ sparkContext.listenerBus.waitUntilEmpty()
+ } finally {
+ spark.listenerManager.unregister(listener)
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/RebalanceBeforeWritingSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/RebalanceBeforeWritingSuite.scala
new file mode 100644
index 00000000000..8a413b298f3
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/RebalanceBeforeWritingSuite.scala
@@ -0,0 +1,487 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.catalyst.expressions.Attribute
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, RebalancePartitions, Sort}
+import org.apache.spark.sql.execution.command.DataWritingCommand
+import org.apache.spark.sql.execution.datasources.InsertIntoHadoopFsRelationCommand
+import org.apache.spark.sql.hive.HiveUtils
+import org.apache.spark.sql.hive.execution.InsertIntoHiveTable
+
+import org.apache.kyuubi.sql.{InferRebalanceAndSortOrders, KyuubiSQLConf}
+
+class RebalanceBeforeWritingSuite extends KyuubiSparkSQLExtensionTest {
+
+ test("check rebalance exists") {
+ def check(
+ df: => DataFrame,
+ expectedRebalanceNumEnabled: Int = 1,
+ expectedRebalanceNumDisabled: Int = 0): Unit = {
+ withSQLConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true") {
+ withListener(df) { write =>
+ assert(write.collect {
+ case r: RebalancePartitions => r
+ }.size == expectedRebalanceNumEnabled)
+ }
+ }
+ withSQLConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "false") {
+ withListener(df) { write =>
+ assert(write.collect {
+ case r: RebalancePartitions => r
+ }.size == expectedRebalanceNumDisabled)
+ }
+ }
+ }
+
+ // It's better to set config explicitly in case of we change the default value.
+ withSQLConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true") {
+ Seq("USING PARQUET", "").foreach { storage =>
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage PARTITIONED BY (c2 string)")
+ check(sql("INSERT INTO TABLE tmp1 PARTITION(c2='a') " +
+ "SELECT * FROM VALUES(1),(2) AS t(c1)"))
+ }
+
+ withTable("tmp1", "tmp2") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage PARTITIONED BY (c2 string)")
+ sql(s"CREATE TABLE tmp2 (c1 int) $storage PARTITIONED BY (c2 string)")
+ check(
+ sql(
+ """FROM VALUES(1),(2)
+ |INSERT INTO TABLE tmp1 PARTITION(c2='a') SELECT *
+ |INSERT INTO TABLE tmp2 PARTITION(c2='a') SELECT *
+ |""".stripMargin),
+ 2)
+ }
+
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage")
+ check(sql("INSERT INTO TABLE tmp1 SELECT * FROM VALUES(1),(2),(3) AS t(c1)"))
+ }
+
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage")
+ check(
+ sql("INSERT INTO TABLE tmp1 SELECT /*+ REBALANCE */ * FROM VALUES(1),(2),(3) AS t(c1)"),
+ 1,
+ 1)
+ }
+
+ withTable("tmp1", "tmp2") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage")
+ sql(s"CREATE TABLE tmp2 (c1 int) $storage")
+ check(
+ sql(
+ """FROM VALUES(1),(2),(3)
+ |INSERT INTO TABLE tmp1 SELECT *
+ |INSERT INTO TABLE tmp2 SELECT *
+ |""".stripMargin),
+ 2)
+ }
+
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 $storage AS SELECT * FROM VALUES(1),(2),(3) AS t(c1)")
+ }
+
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 $storage PARTITIONED BY(c2) AS " +
+ s"SELECT * FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2)")
+ }
+ }
+ }
+ }
+
+ test("check rebalance does not exists") {
+ def check(df: DataFrame): Unit = {
+ withListener(df) { write =>
+ assert(write.collect {
+ case r: RebalancePartitions => r
+ }.isEmpty)
+ }
+ }
+
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true") {
+ // test no write command
+ check(sql("SELECT * FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2)"))
+ check(sql("SELECT count(*) FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2)"))
+
+ // test not supported plan
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) PARTITIONED BY (c2 string)")
+ check(sql("INSERT INTO TABLE tmp1 PARTITION(c2) " +
+ "SELECT /*+ repartition(10) */ * FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2)"))
+ check(sql("INSERT INTO TABLE tmp1 PARTITION(c2) " +
+ "SELECT * FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2) ORDER BY c1"))
+ check(sql("INSERT INTO TABLE tmp1 PARTITION(c2) " +
+ "SELECT * FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2) LIMIT 10"))
+ }
+ }
+
+ withSQLConf(KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "false") {
+ Seq("USING PARQUET", "").foreach { storage =>
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage PARTITIONED BY (c2 string)")
+ check(sql("INSERT INTO TABLE tmp1 PARTITION(c2) " +
+ "SELECT * FROM VALUES(1, 'a'),(2, 'b') AS t(c1, c2)"))
+ }
+
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage")
+ check(sql("INSERT INTO TABLE tmp1 SELECT * FROM VALUES(1),(2),(3) AS t(c1)"))
+ }
+ }
+ }
+ }
+
+ test("test dynamic partition write") {
+ def checkRepartitionExpression(sqlString: String): Unit = {
+ withListener(sqlString) { write =>
+ assert(write.isInstanceOf[InsertIntoHiveTable])
+ assert(write.collect {
+ case r: RebalancePartitions if r.partitionExpressions.size == 1 =>
+ assert(r.partitionExpressions.head.asInstanceOf[Attribute].name === "c2")
+ r
+ }.size == 1)
+ }
+ }
+
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true") {
+ Seq("USING PARQUET", "").foreach { storage =>
+ withTable("tmp1") {
+ sql(s"CREATE TABLE tmp1 (c1 int) $storage PARTITIONED BY (c2 string)")
+ checkRepartitionExpression("INSERT INTO TABLE tmp1 SELECT 1 as c1, 'a' as c2 ")
+ }
+
+ withTable("tmp1") {
+ checkRepartitionExpression(
+ "CREATE TABLE tmp1 PARTITIONED BY(C2) SELECT 1 as c1, 'a' as c2")
+ }
+ }
+ }
+ }
+
+ test("OptimizedCreateHiveTableAsSelectCommand") {
+ withSQLConf(
+ HiveUtils.CONVERT_METASTORE_PARQUET.key -> "true",
+ HiveUtils.CONVERT_METASTORE_CTAS.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true") {
+ withTable("t") {
+ withListener("CREATE TABLE t STORED AS parquet AS SELECT 1 as a") { write =>
+ assert(write.isInstanceOf[InsertIntoHadoopFsRelationCommand])
+ assert(write.collect {
+ case _: RebalancePartitions => true
+ }.size == 1)
+ }
+ }
+ }
+ }
+
+ test("Infer rebalance and sorder orders") {
+ def checkShuffleAndSort(dataWritingCommand: LogicalPlan, sSize: Int, rSize: Int): Unit = {
+ assert(dataWritingCommand.isInstanceOf[DataWritingCommand])
+ val plan = dataWritingCommand.asInstanceOf[DataWritingCommand].query
+ assert(plan.collect {
+ case s: Sort => s
+ }.size == sSize)
+ assert(plan.collect {
+ case r: RebalancePartitions if r.partitionExpressions.size == rSize => r
+ }.nonEmpty || rSize == 0)
+ }
+
+ withView("v") {
+ withTable("t", "t2", "input1", "input2") {
+ withSQLConf(KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS.key -> "true") {
+ sql(s"CREATE TABLE t (c1 int, c2 long) USING PARQUET PARTITIONED BY (p string)")
+ sql(s"CREATE TABLE t2 (c1 int, c2 long, c3 long) USING PARQUET PARTITIONED BY (p string)")
+ sql(s"CREATE TABLE input1 USING PARQUET AS SELECT * FROM VALUES(1,2),(1,3)")
+ sql(s"CREATE TABLE input2 USING PARQUET AS SELECT * FROM VALUES(1,3),(1,3)")
+ sql(s"CREATE VIEW v as SELECT col1, count(*) as col2 FROM input1 GROUP BY col1")
+
+ val df0 = sql(
+ s"""
+ |INSERT INTO TABLE t PARTITION(p='a')
+ |SELECT /*+ broadcast(input2) */ input1.col1, input2.col1
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 = input2.col1
+ |""".stripMargin)
+ checkShuffleAndSort(df0.queryExecution.analyzed, 1, 1)
+
+ val df1 = sql(
+ s"""
+ |INSERT INTO TABLE t PARTITION(p='a')
+ |SELECT /*+ broadcast(input2) */ input1.col1, input1.col2
+ |FROM input1
+ |LEFT JOIN input2
+ |ON input1.col1 = input2.col1 and input1.col2 = input2.col2
+ |""".stripMargin)
+ checkShuffleAndSort(df1.queryExecution.analyzed, 1, 2)
+
+ val df2 = sql(
+ s"""
+ |INSERT INTO TABLE t PARTITION(p='a')
+ |SELECT col1 as c1, count(*) as c2
+ |FROM input1
+ |GROUP BY col1
+ |HAVING count(*) > 0
+ |""".stripMargin)
+ checkShuffleAndSort(df2.queryExecution.analyzed, 1, 1)
+
+ // dynamic partition
+ val df3 = sql(
+ s"""
+ |INSERT INTO TABLE t PARTITION(p)
+ |SELECT /*+ broadcast(input2) */ input1.col1, input1.col2, input1.col2
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 = input2.col1
+ |""".stripMargin)
+ checkShuffleAndSort(df3.queryExecution.analyzed, 0, 1)
+
+ // non-deterministic
+ val df4 = sql(
+ s"""
+ |INSERT INTO TABLE t PARTITION(p='a')
+ |SELECT col1 + rand(), count(*) as c2
+ |FROM input1
+ |GROUP BY col1
+ |""".stripMargin)
+ checkShuffleAndSort(df4.queryExecution.analyzed, 0, 0)
+
+ // view
+ val df5 = sql(
+ s"""
+ |INSERT INTO TABLE t PARTITION(p='a')
+ |SELECT * FROM v
+ |""".stripMargin)
+ checkShuffleAndSort(df5.queryExecution.analyzed, 1, 1)
+
+ // generate
+ val df6 = sql(
+ s"""
+ |INSERT INTO TABLE t2 PARTITION(p='a')
+ |SELECT /*+ broadcast(input2) */ input1.col1, input2.col1, cast(cc.action1 as bigint)
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 = input2.col1
+ | lateral view explode(ARRAY(input1.col1, input1.col2)) cc as action1
+ |""".stripMargin)
+ checkShuffleAndSort(df6.queryExecution.analyzed, 1, 1)
+
+ // window
+ val df7 = sql(
+ s"""
+ |INSERT INTO TABLE t2 PARTITION(p='a')
+ |SELECT /*+ broadcast(input2) */ input1.col1, input2.col2,
+ | RANK() OVER (PARTITION BY input2.col2 ORDER BY input1.col1) AS rank
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 = input2.col1
+ |""".stripMargin)
+ checkShuffleAndSort(df7.queryExecution.analyzed, 1, 1)
+ }
+ }
+ }
+ }
+
+ test("Skip inferring sort orders") {
+ def checkShuffleAndSort(dataWritingCommand: LogicalPlan, sSize: Int, rSize: Int): Unit = {
+ assert(dataWritingCommand.isInstanceOf[DataWritingCommand])
+ val plan = dataWritingCommand.asInstanceOf[DataWritingCommand].query
+ assert(plan.collect { case s: Sort => s }.size == sSize)
+ assert(plan.collect {
+ case r: RebalancePartitions if r.partitionExpressions.size == rSize => r
+ }.nonEmpty || rSize == 0)
+ }
+
+ withTable("t", "input1", "input2") {
+ sql(s"CREATE TABLE t (c1 int, c2 long) USING PARQUET PARTITIONED BY (p string)")
+ sql(s"CREATE TABLE input1 USING PARQUET AS SELECT * FROM VALUES(1,2),(1,3)")
+ sql(s"CREATE TABLE input2 USING PARQUET AS SELECT * FROM VALUES(1,3),(1,3)")
+
+ val insert =
+ s"""
+ |INSERT INTO TABLE t PARTITION(p='a')
+ |SELECT /*+ broadcast(input2) */ input1.col1, input2.col1
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 = input2.col1
+ |""".stripMargin
+
+ // when skip is disabled (default), both rebalance and sort orders are inferred
+ withSQLConf(
+ KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS.key -> "true",
+ KyuubiSQLConf.SKIP_INFER_SORT_ORDERS.key -> "false") {
+ checkShuffleAndSort(sql(insert).queryExecution.analyzed, 1, 1)
+ }
+
+ // when skip is enabled, only the rebalance is inferred and the sort is skipped
+ withSQLConf(
+ KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS.key -> "true",
+ KyuubiSQLConf.SKIP_INFER_SORT_ORDERS.key -> "true") {
+ checkShuffleAndSort(sql(insert).queryExecution.analyzed, 0, 1)
+ }
+ }
+ }
+
+ test("Infer rebalance and sort orders only with cheap columns") {
+ withTable("input1", "input2") {
+ sql(s"CREATE TABLE input1 USING PARQUET AS SELECT * FROM VALUES(1,2),(1,3)")
+ sql(s"CREATE TABLE input2 USING PARQUET AS SELECT * FROM VALUES(1,3),(1,3)")
+
+ // the join keys `input1.col1 + 1` / `input2.col1 + 1` are not cheap expressions
+ val expensivePlan = sql(
+ s"""
+ |SELECT input1.col1, input2.col1
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 + 1 = input2.col1 + 1
+ |""".stripMargin).queryExecution.analyzed
+
+ // when only cheap columns are allowed, the expensive keys are not inferred
+ assert(InferRebalanceAndSortOrders.infer(expensivePlan, true).isEmpty)
+ // when the cheap-column restriction is disabled, the expensive keys are inferred
+ InferRebalanceAndSortOrders.infer(expensivePlan, false) match {
+ case Some((partitioning, ordering)) =>
+ assert(partitioning.nonEmpty)
+ assert(ordering.nonEmpty)
+ case None => fail("expected inferred columns when cheap-column restriction is disabled")
+ }
+
+ // the join keys `input1.col1` / `input2.col1` are cheap attributes
+ val cheapPlan = sql(
+ s"""
+ |SELECT input1.col1, input2.col1
+ |FROM input1
+ |JOIN input2
+ |ON input1.col1 = input2.col1
+ |""".stripMargin).queryExecution.analyzed
+
+ // cheap columns are inferred regardless of the restriction
+ Seq(true, false).foreach { onlyCheap =>
+ InferRebalanceAndSortOrders.infer(cheapPlan, onlyCheap) match {
+ case Some((partitioning, ordering)) =>
+ assert(partitioning.nonEmpty)
+ assert(ordering.nonEmpty)
+ case None => fail(s"expected inferred columns when onlyCheap=$onlyCheap")
+ }
+ }
+ }
+ }
+
+ test("advisory partition size - finalStage config takes precedence") {
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true",
+ KyuubiSQLConf.FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY -> "100MB") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET")
+ val df = sql("INSERT INTO TABLE tmp1 SELECT * FROM VALUES(1),(2),(3) AS t(c1)")
+ withListener(df) { write =>
+ val rebalances = write.collect { case r: RebalancePartitions => r }
+ assert(rebalances.size === 1)
+ assert(rebalances.head.optAdvisoryPartitionSize.isDefined)
+ assert(rebalances.head.optAdvisoryPartitionSize.get === 104857600L) // 100MB
+ }
+ }
+ }
+ }
+
+ test("advisory partition size - fallback to kyuubi config") {
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "200MB") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET")
+ val df = sql("INSERT INTO TABLE tmp1 SELECT * FROM VALUES(1),(2),(3) AS t(c1)")
+ withListener(df) { write =>
+ val rebalances = write.collect { case r: RebalancePartitions => r }
+ assert(rebalances.size === 1)
+ assert(rebalances.head.optAdvisoryPartitionSize.isDefined)
+ assert(rebalances.head.optAdvisoryPartitionSize.get === 209715200L) // 200MB
+ }
+ }
+ }
+ }
+
+ test("advisory partition size - kyuubi config takes precedence over finalStage") {
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "200MB",
+ KyuubiSQLConf.FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY -> "300MB") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET")
+ val df = sql("INSERT INTO TABLE tmp1 SELECT * FROM VALUES(1),(2),(3) AS t(c1)")
+ withListener(df) { write =>
+ val rebalances = write.collect { case r: RebalancePartitions => r }
+ assert(rebalances.size === 1)
+ assert(rebalances.head.optAdvisoryPartitionSize.isDefined)
+ // kyuubi config takes precedence over finalStage
+ assert(rebalances.head.optAdvisoryPartitionSize.get === 209715200L) // 200MB
+ }
+ }
+ }
+ }
+
+ test("advisory partition size - none when neither config is set") {
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET")
+ val df = sql("INSERT INTO TABLE tmp1 SELECT * FROM VALUES(1),(2),(3) AS t(c1)")
+ withListener(df) { write =>
+ val rebalances = write.collect { case r: RebalancePartitions => r }
+ assert(rebalances.size === 1)
+ assert(rebalances.head.optAdvisoryPartitionSize.isEmpty)
+ }
+ }
+ }
+ }
+
+ test("getAdvisoryPartitionSize - raw config resolution") {
+ val conf = spark.sessionState.conf
+
+ // Neither config set
+ assert(KyuubiSQLConf.getAdvisoryPartitionSize(conf).isEmpty)
+
+ // Only finalStage config set
+ withSQLConf(KyuubiSQLConf.FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY -> "50m") {
+ assert(KyuubiSQLConf.getAdvisoryPartitionSize(conf).get === 52428800L) // 50MB
+ }
+
+ // Only kyuubi config set
+ withSQLConf(KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "1gb") {
+ assert(KyuubiSQLConf.getAdvisoryPartitionSize(conf).get === 1073741824L) // 1GB
+ }
+
+ // Both set: kyuubi config takes precedence
+ withSQLConf(
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "1gb",
+ KyuubiSQLConf.FINAL_STAGE_ADVISORY_PARTITION_SIZE_KEY -> "500MB") {
+ assert(KyuubiSQLConf.getAdvisoryPartitionSize(conf).get === 1073741824L) // 1GB
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/RemoveRebalanceShuffleSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/RemoveRebalanceShuffleSuite.scala
new file mode 100644
index 00000000000..8dbc51af61c
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/RemoveRebalanceShuffleSuite.scala
@@ -0,0 +1,384 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.sql.execution.{CommandResultExec, SortExec, SparkPlan}
+import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, ShuffleQueryStageExec}
+import org.apache.spark.sql.execution.exchange.{REBALANCE_PARTITIONS_BY_COL, REBALANCE_PARTITIONS_BY_NONE, ShuffleExchangeExec}
+import org.apache.spark.sql.execution.joins.SortMergeJoinExec
+import org.apache.spark.sql.internal.SQLConf._
+
+import org.apache.kyuubi.sql.KyuubiSQLConf
+
+class RemoveRebalanceShuffleSuite extends KyuubiSparkSQLExtensionTest {
+
+ override protected def beforeAll(): Unit = {
+ super.beforeAll()
+ setupData()
+ }
+
+ private def collectExchanges(plan: SparkPlan): Seq[ShuffleExchangeExec] = plan match {
+ case p: CommandResultExec => collectExchanges(p.commandPhysicalPlan)
+ case p: AdaptiveSparkPlanExec => collectExchanges(p.finalPhysicalPlan)
+ case p: ShuffleQueryStageExec => collectExchanges(p.plan)
+ case p: ShuffleExchangeExec => p +: collectExchanges(p.child)
+ case p => p.children.flatMap(collectExchanges)
+ }
+
+ /**
+ * Collects all rebalance shuffle exchanges from the physical plan. Both
+ * [[REBALANCE_PARTITIONS_BY_NONE]] (empty partition expressions, e.g. inserted by
+ * `InsertAdaptor` for non-partitioned tables) and [[REBALANCE_PARTITIONS_BY_COL]]
+ * (non-empty partition expressions, e.g. inferred by `InferRebalanceAndSortOrders`
+ * for partitioned tables) are counted.
+ */
+ private def rebalanceExchanges(plan: SparkPlan): Seq[ShuffleExchangeExec] = {
+ collectExchanges(plan).filter(e =>
+ e.shuffleOrigin == REBALANCE_PARTITIONS_BY_NONE ||
+ e.shuffleOrigin == REBALANCE_PARTITIONS_BY_COL)
+ }
+
+ /**
+ * Collects local [[SortExec]] nodes from the plan that are NOT part of a
+ * [[SortMergeJoinExec]] (whose child Sort nodes are join-required sorts, not the
+ * local Sort injected by `InferRebalanceAndSortOrders`). Used to verify that
+ * [[org.apache.spark.sql.execution.RemoveRedundantSorts]] cleaned up the local Sort
+ * after the rebalance was removed.
+ */
+ private def collectSorts(plan: SparkPlan): Seq[SortExec] = plan match {
+ case p: CommandResultExec => collectSorts(p.commandPhysicalPlan)
+ case p: AdaptiveSparkPlanExec => collectSorts(p.finalPhysicalPlan)
+ case p: ShuffleQueryStageExec => collectSorts(p.plan)
+ // Don't recurse into SortMergeJoinExec, its child Sort nodes are join-required sorts,
+ // not the local Sort injected by InferRebalanceAndSortOrders.
+ case _: SortMergeJoinExec => Seq.empty
+ case p: SortExec => p +: p.children.flatMap(collectSorts).toSeq
+ case p => p.children.flatMap(collectSorts)
+ }
+
+ private def assertRebalanceRemoved(df: DataFrame, removed: Boolean): Unit = {
+ df.collect()
+ val rebalances = rebalanceExchanges(df.queryExecution.executedPlan)
+ val expected = if (removed) 0 else 1
+ assert(
+ rebalances.size == expected,
+ s"Expected $expected rebalance shuffle, got ${rebalances.size}:\n" +
+ s"${df.queryExecution.executedPlan}")
+ }
+
+ test("remove rebalance shuffle - small-data scenario with non-expanding semi join") {
+ // A left-semi join is reducing (not expanding), so the small-data condition applies: the
+ // materialized stage size is far below `advisoryPartitionSize * tolerableSmallFileNum`, so the
+ // rebalance shuffle is removed.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a LEFT SEMI JOIN t2 b ON a.c1 = b.c1"),
+ removed = true)
+ }
+ }
+ }
+
+ test("large-data scenario with non-reducing outer join") {
+ // A left-outer join is expanding (not reducing), so the large-data condition applies. With a
+ // tiny `smallPartitionSize`, the stage size exceeds
+ // `numShufflePartitions * smallPartitionSize`, so the rebalance shuffle is removed.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ SHUFFLE_PARTITIONS.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+
+ withSQLConf(KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_SMALL_PARTITION_SIZE.key -> "1b") {
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a LEFT OUTER JOIN t2 b ON a.c1 = b.c1"),
+ removed = true)
+ }
+
+ withSQLConf(KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_SMALL_PARTITION_SIZE.key -> "1tb") {
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a LEFT OUTER JOIN t2 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - advisory size not larger than session advisory size") {
+ // The rebalance advisory size equals the session advisory size, which suggests the user
+ // intentionally wants smaller partitions, so the rebalance shuffle is kept.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "64MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a LEFT SEMI JOIN t2 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - inner join is both reducing and expanding") {
+ // An inner join is classified as both reducing and expanding, so neither the large-data nor
+ // the small-data condition can be satisfied and the rebalance shuffle is kept.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a JOIN t2 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - reducing semi join fails the small-data condition") {
+ // A left-semi join is reducing, so the large-data condition never applies. Setting
+ // `tolerableSmallFileNum` to 0 also disables the small-data condition, so the shuffle is kept.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_TOLERABLE_SMALL_FILE_NUM.key -> "0",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a LEFT SEMI JOIN t2 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - switch off") {
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "false",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT a.c1 FROM t1 a LEFT SEMI JOIN t2 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+
+ test("remove rebalance shuffle - union of non-expanding branches") {
+ // A union-all of two left-semi joins: every branch is reducing (not expanding) and each forms
+ // its own materialized query stage group. All groups satisfy the small-data condition, so the
+ // rebalance shuffle is removed.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql(
+ "INSERT INTO TABLE tmp1 " +
+ "SELECT a.c1 FROM t1 a LEFT SEMI JOIN t2 b ON a.c1 = b.c1 " +
+ "UNION ALL " +
+ "SELECT a.c1 FROM t1 a LEFT SEMI JOIN t3 b ON a.c1 = b.c1"),
+ removed = true)
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - union with an expanding branch") {
+ // A union-all of a left-semi join and an inner join. The inner join makes the whole rebalance
+ // input both reducing and expanding, so neither the large-data nor the small-data condition can
+ // be satisfied for any group and the rebalance shuffle is kept.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql(
+ "INSERT INTO TABLE tmp1 " +
+ "SELECT a.c1 FROM t1 a LEFT SEMI JOIN t2 b ON a.c1 = b.c1 " +
+ "UNION ALL " +
+ "SELECT a.c1 FROM t1 a JOIN t3 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+
+ test("remove rebalance shuffle - cached table produces TableCacheQueryStageExec") {
+ // A cached table produces a TableCacheQueryStageExec (not an ExchangeQueryStageExec).
+ // The non-expanding cached table scan satisfies the small-data condition, so the
+ // rebalance shuffle should be removed.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CACHE TABLE t2")
+ try {
+ sql("CREATE TABLE tmp1 USING PARQUET as select * from range(10)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 SELECT c1 FROM t2"),
+ removed = true)
+ } finally {
+ sql("UNCACHE TABLE t2")
+ }
+ }
+ }
+ }
+
+ test("remove rebalance shuffle - static partition insert with non-expanding query") {
+ // staticPartitions.size == partitionColumns.size triggers the rule for static partition
+ // inserts. A simple SELECT has no reducing or expanding operators, so the small-data
+ // condition applies and the rebalance shuffle is removed.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET PARTITIONED BY (c2 string)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 PARTITION(c2='a') SELECT c1 FROM t2"),
+ removed = true)
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - static partition insert with inner join") {
+ // An inner join is both reducing and expanding, so neither the large-data nor the small-data
+ // condition can be satisfied and the rebalance shuffle is kept.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET PARTITIONED BY (c2 string)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 PARTITION(c2='a') " +
+ "SELECT a.c1 FROM t1 a JOIN t2 b ON a.c1 = b.c1"),
+ removed = false)
+ }
+ }
+ }
+
+ test("keep rebalance shuffle - dynamic partition insert skips rule") {
+ // Dynamic partition inserts (staticPartitions.size < partitionColumns.size) do not trigger
+ // the rule because the rebalance helps cluster data by partition values, so the rebalance
+ // shuffle is kept regardless of data size or operators.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp1") {
+ sql("CREATE TABLE tmp1 (c1 int) USING PARQUET PARTITIONED BY (c2 string)")
+ assertRebalanceRemoved(
+ sql("INSERT INTO TABLE tmp1 PARTITION(c2) SELECT c1, c2 FROM t2"),
+ removed = false)
+ }
+ }
+ }
+
+ test("remove both rebalance and local sort inferred by InferRebalanceAndSortOrders") {
+ // When InferRebalanceAndSortOrders infers non-empty partitionExpressions from a join query
+ // and injects both a RebalancePartitions and a wrapping local Sort, the
+ // RemoveRebalanceShuffle rule removes the RebalancePartitions (the relaxed partitionExpressions
+ // match allows non-empty expressions). The local Sort then serves no purpose and is
+ // removed by Spark's RemoveRedundantSorts rule.
+ // Uses a left-semi join: it is reducing (not expanding), so the small-data condition applies.
+ withSQLConf(
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE.key -> "true",
+ KyuubiSQLConf.INSERT_REPARTITION_BEFORE_WRITE_IF_NO_SHUFFLE.key -> "true",
+ KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS.key -> "true",
+ KyuubiSQLConf.INFER_REBALANCE_AND_SORT_ORDERS_MAX_COLUMNS.key -> "2",
+ KyuubiSQLConf.REMOVE_REBALANCE_SHUFFLE_ENABLED.key -> "true",
+ ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "64MB",
+ KyuubiSQLConf.REBALANCE_PARTITIONS_ADVISORY_PARTITION_SIZE.key -> "128MB",
+ COALESCE_PARTITIONS_INITIAL_PARTITION_NUM.key -> "2",
+ AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
+ withTable("tmp_pt", "input1", "input2") {
+ sql("CREATE TABLE tmp_pt (c1 int, c2 long) USING PARQUET PARTITIONED BY (p string)")
+ sql("CREATE TABLE input1 USING PARQUET AS SELECT * FROM VALUES(1,2),(1,3)")
+ sql("CREATE TABLE input2 USING PARQUET AS SELECT * FROM VALUES(1,3),(1,3)")
+
+ val df = sql(
+ """INSERT INTO TABLE tmp_pt PARTITION(p='a')
+ |SELECT input1.col1, input1.col2
+ |FROM input1 LEFT SEMI JOIN input2 ON input1.col1 = input2.col1
+ |""".stripMargin)
+
+ // Verify rebalance is removed
+ assertRebalanceRemoved(df, removed = true)
+
+ // Verify the local Sort injected by InferRebalanceAndSortOrders is also removed
+ df.collect()
+ val sorts = collectSorts(df.queryExecution.executedPlan)
+ assert(
+ sorts.isEmpty,
+ s"Expected no local Sort in the final plan after rebalance removal, " +
+ s"but found ${sorts.size}:\n${df.queryExecution.executedPlan}")
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ReportStatisticsAndPartitionAwareDataSource.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ReportStatisticsAndPartitionAwareDataSource.scala
new file mode 100644
index 00000000000..136ced538e1
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ReportStatisticsAndPartitionAwareDataSource.scala
@@ -0,0 +1,64 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import java.util.OptionalLong
+
+import org.apache.spark.sql.connector.{RangeInputPartition, SimpleBatchTable, SimpleScanBuilder, SimpleWritableDataSource}
+import org.apache.spark.sql.connector.catalog.Table
+import org.apache.spark.sql.connector.expressions.{Expressions, FieldReference, Transform}
+import org.apache.spark.sql.connector.read.{InputPartition, ScanBuilder, Statistics, SupportsReportPartitioning, SupportsReportStatistics}
+import org.apache.spark.sql.connector.read.partitioning.{KeyGroupedPartitioning, Partitioning}
+import org.apache.spark.sql.util.CaseInsensitiveStringMap
+
+class ReportStatisticsAndPartitionAwareDataSource extends SimpleWritableDataSource {
+
+ class MyScanBuilder(
+ val partitionKeys: Seq[String]) extends SimpleScanBuilder
+ with SupportsReportStatistics with SupportsReportPartitioning {
+
+ override def estimateStatistics(): Statistics = {
+ new Statistics {
+ override def sizeInBytes(): OptionalLong = OptionalLong.of(80)
+
+ override def numRows(): OptionalLong = OptionalLong.of(10)
+
+ }
+ }
+
+ override def planInputPartitions(): Array[InputPartition] = {
+ Array(RangeInputPartition(0, 5), RangeInputPartition(5, 10))
+ }
+
+ override def outputPartitioning(): Partitioning = {
+ new KeyGroupedPartitioning(partitionKeys.map(FieldReference(_)).toArray, 10)
+ }
+ }
+
+ override def getTable(options: CaseInsensitiveStringMap): Table = {
+ new SimpleBatchTable {
+ override def newScanBuilder(options: CaseInsensitiveStringMap): ScanBuilder = {
+ new MyScanBuilder(Seq("i"))
+ }
+
+ override def partitioning(): Array[Transform] = {
+ Array(Expressions.identity("i"))
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ReportStatisticsDataSource.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ReportStatisticsDataSource.scala
new file mode 100644
index 00000000000..2035d352562
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ReportStatisticsDataSource.scala
@@ -0,0 +1,53 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import java.util.OptionalLong
+
+import org.apache.spark.sql.connector._
+import org.apache.spark.sql.connector.catalog.Table
+import org.apache.spark.sql.connector.read._
+import org.apache.spark.sql.util.CaseInsensitiveStringMap
+
+class ReportStatisticsDataSource extends SimpleWritableDataSource {
+
+ class MyScanBuilder extends SimpleScanBuilder
+ with SupportsReportStatistics {
+
+ override def estimateStatistics(): Statistics = {
+ new Statistics {
+ override def sizeInBytes(): OptionalLong = OptionalLong.of(80)
+
+ override def numRows(): OptionalLong = OptionalLong.of(10)
+ }
+ }
+
+ override def planInputPartitions(): Array[InputPartition] = {
+ Array(RangeInputPartition(0, 5), RangeInputPartition(5, 10))
+ }
+
+ }
+
+ override def getTable(options: CaseInsensitiveStringMap): Table = {
+ new SimpleBatchTable {
+ override def newScanBuilder(options: CaseInsensitiveStringMap): ScanBuilder = {
+ new MyScanBuilder
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/WatchDogSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/WatchDogSuite.scala
new file mode 100644
index 00000000000..038ee3fe32b
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/WatchDogSuite.scala
@@ -0,0 +1,254 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import java.io.File
+
+import scala.jdk.CollectionConverters._
+
+import org.apache.commons.io.FileUtils
+import org.apache.spark.sql.execution.datasources.v2.DataSourceV2ScanRelation
+
+import org.apache.kyuubi.sql.{KyuubiSQLConf, KyuubiSQLExtensionException}
+import org.apache.kyuubi.sql.watchdog.{MaxFileSizeExceedException, MaxPartitionExceedException}
+
+class WatchDogSuite extends KyuubiSparkSQLExtensionTest {
+ override protected def beforeAll(): Unit = {
+ super.beforeAll()
+ setupData()
+ }
+
+ case class LimitAndExpected(limit: Int, expected: Int)
+
+ val limitAndExpecteds = List(LimitAndExpected(1, 1), LimitAndExpected(11, 10))
+
+ private def checkMaxPartition: Unit = {
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_PARTITIONS.key -> "100") {
+ checkAnswer(sql("SELECT count(distinct(p)) FROM test"), Row(10) :: Nil)
+ }
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_PARTITIONS.key -> "5") {
+ sql("SELECT * FROM test where p=1").queryExecution.sparkPlan
+
+ sql(s"SELECT * FROM test WHERE p in (${Range(0, 5).toList.mkString(",")})")
+ .queryExecution.sparkPlan
+
+ intercept[MaxPartitionExceedException](
+ sql("SELECT * FROM test where p != 1").queryExecution.sparkPlan)
+
+ intercept[MaxPartitionExceedException](
+ sql("SELECT * FROM test").queryExecution.sparkPlan)
+
+ intercept[MaxPartitionExceedException](sql(
+ s"SELECT * FROM test WHERE p in (${Range(0, 6).toList.mkString(",")})")
+ .queryExecution.sparkPlan)
+ }
+ }
+
+ test("watchdog with scan maxPartitions -- hive") {
+ Seq("textfile", "parquet").foreach { format =>
+ withTable("test", "temp") {
+ sql(
+ s"""
+ |CREATE TABLE test(i int)
+ |PARTITIONED BY (p int)
+ |STORED AS $format""".stripMargin)
+ spark.range(0, 10, 1).selectExpr("id as col")
+ .createOrReplaceTempView("temp")
+
+ for (part <- Range(0, 10)) {
+ sql(
+ s"""
+ |INSERT OVERWRITE TABLE test PARTITION (p='$part')
+ |select col from temp""".stripMargin)
+ }
+ checkMaxPartition
+ }
+ }
+ }
+
+ test("watchdog with scan maxPartitions -- data source") {
+ withTempDir { dir =>
+ withTempView("test") {
+ spark.range(10).selectExpr("id", "id as p")
+ .write
+ .partitionBy("p")
+ .mode("overwrite")
+ .save(dir.getCanonicalPath)
+ spark.read.load(dir.getCanonicalPath).createOrReplaceTempView("test")
+ checkMaxPartition
+ }
+ }
+ }
+
+ private def checkMaxFileSize(tableSize: Long, nonPartTableSize: Long): Unit = {
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> tableSize.toString) {
+ checkAnswer(sql("SELECT count(distinct(p)) FROM test"), Row(10) :: Nil)
+ }
+
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> (tableSize / 2).toString) {
+ sql("SELECT * FROM test where p=1").queryExecution.sparkPlan
+
+ sql(s"SELECT * FROM test WHERE p in (${Range(0, 3).toList.mkString(",")})")
+ .queryExecution.sparkPlan
+
+ intercept[MaxFileSizeExceedException](
+ sql("SELECT * FROM test where p != 1").queryExecution.sparkPlan)
+
+ intercept[MaxFileSizeExceedException](
+ sql("SELECT * FROM test").queryExecution.sparkPlan)
+
+ intercept[MaxFileSizeExceedException](sql(
+ s"SELECT * FROM test WHERE p in (${Range(0, 6).toList.mkString(",")})")
+ .queryExecution.sparkPlan)
+ }
+
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> nonPartTableSize.toString) {
+ checkAnswer(sql("SELECT count(*) FROM test_non_part"), Row(10000) :: Nil)
+ }
+
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> (nonPartTableSize - 1).toString) {
+ intercept[MaxFileSizeExceedException](
+ sql("SELECT * FROM test_non_part").queryExecution.sparkPlan)
+ }
+ }
+
+ test("watchdog with scan maxFileSize -- hive") {
+ Seq(false).foreach { convertMetastoreParquet =>
+ withTable("test", "test_non_part", "temp") {
+ spark.range(10000).selectExpr("id as col")
+ .createOrReplaceTempView("temp")
+
+ // partitioned table
+ sql(
+ s"""
+ |CREATE TABLE test(i int)
+ |PARTITIONED BY (p int)
+ |STORED AS parquet""".stripMargin)
+ for (part <- Range(0, 10)) {
+ sql(
+ s"""
+ |INSERT OVERWRITE TABLE test PARTITION (p='$part')
+ |select col from temp""".stripMargin)
+ }
+
+ val tablePath = new File(spark.sessionState.catalog.externalCatalog
+ .getTable("default", "test").location)
+ val tableSize = FileUtils.listFiles(tablePath, Array("parquet"), true).asScala
+ .map(_.length()).sum
+ assert(tableSize > 0)
+
+ // non-partitioned table
+ sql(
+ s"""
+ |CREATE TABLE test_non_part(i int)
+ |STORED AS parquet""".stripMargin)
+ sql(
+ s"""
+ |INSERT OVERWRITE TABLE test_non_part
+ |select col from temp""".stripMargin)
+ sql("ANALYZE TABLE test_non_part COMPUTE STATISTICS")
+
+ val nonPartTablePath = new File(spark.sessionState.catalog.externalCatalog
+ .getTable("default", "test_non_part").location)
+ val nonPartTableSize = FileUtils.listFiles(nonPartTablePath, Array("parquet"), true).asScala
+ .map(_.length()).sum
+ assert(nonPartTableSize > 0)
+
+ // check
+ withSQLConf("spark.sql.hive.convertMetastoreParquet" -> convertMetastoreParquet.toString) {
+ checkMaxFileSize(tableSize, nonPartTableSize)
+ }
+ }
+ }
+ }
+
+ test("watchdog with scan maxFileSize -- data source") {
+ withTempDir { dir =>
+ withTempView("test", "test_non_part") {
+ // partitioned table
+ val tablePath = new File(dir, "test")
+ spark.range(10).selectExpr("id", "id as p")
+ .write
+ .partitionBy("p")
+ .mode("overwrite")
+ .parquet(tablePath.getCanonicalPath)
+ spark.read.load(tablePath.getCanonicalPath).createOrReplaceTempView("test")
+
+ val tableSize = FileUtils.listFiles(tablePath, Array("parquet"), true).asScala
+ .map(_.length()).sum
+ assert(tableSize > 0)
+
+ // non-partitioned table
+ val nonPartTablePath = new File(dir, "test_non_part")
+ spark.range(10000).selectExpr("id", "id as p")
+ .write
+ .mode("overwrite")
+ .parquet(nonPartTablePath.getCanonicalPath)
+ spark.read.load(nonPartTablePath.getCanonicalPath).createOrReplaceTempView("test_non_part")
+
+ val nonPartTableSize = FileUtils.listFiles(nonPartTablePath, Array("parquet"), true).asScala
+ .map(_.length()).sum
+ assert(tableSize > 0)
+
+ // check
+ checkMaxFileSize(tableSize, nonPartTableSize)
+ }
+ }
+ }
+
+ test("disable script transformation") {
+ withSQLConf(KyuubiSQLConf.SCRIPT_TRANSFORMATION_ENABLED.key -> "false") {
+ val e = intercept[KyuubiSQLExtensionException] {
+ sql("SELECT TRANSFORM('') USING 'ls /'")
+ }
+ assert(e.getMessage == "Script transformation is not allowed")
+ }
+ }
+
+ test("watchdog with scan maxFileSize -- data source v2") {
+ val df = spark.read.format(classOf[ReportStatisticsAndPartitionAwareDataSource].getName).load()
+ df.createOrReplaceTempView("test")
+ val logical = df.queryExecution.optimizedPlan.collect {
+ case d: DataSourceV2ScanRelation => d
+ }.head
+ val tableSize = logical.computeStats().sizeInBytes.toLong
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> tableSize.toString) {
+ sql("SELECT * FROM test").queryExecution.sparkPlan
+ }
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> (tableSize / 2).toString) {
+ intercept[MaxFileSizeExceedException](
+ sql("SELECT * FROM test").queryExecution.sparkPlan)
+ }
+
+ val nonPartDf = spark.read.format(classOf[ReportStatisticsDataSource].getName).load()
+ nonPartDf.createOrReplaceTempView("test_non_part")
+ val nonPartLogical = nonPartDf.queryExecution.optimizedPlan.collect {
+ case d: DataSourceV2ScanRelation => d
+ }.head
+ val nonPartTableSize = nonPartLogical.computeStats().sizeInBytes.toLong
+
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> nonPartTableSize.toString) {
+ sql("SELECT * FROM test_non_part").queryExecution.sparkPlan
+ }
+
+ withSQLConf(KyuubiSQLConf.WATCHDOG_MAX_FILE_SIZE.key -> (nonPartTableSize / 2).toString) {
+ intercept[MaxFileSizeExceedException](
+ sql("SELECT * FROM test_non_part").queryExecution.sparkPlan)
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ZorderCoreBenchmark.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ZorderCoreBenchmark.scala
new file mode 100644
index 00000000000..e9f4fc45ea9
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ZorderCoreBenchmark.scala
@@ -0,0 +1,117 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.SparkConf
+import org.apache.spark.benchmark.Benchmark
+import org.apache.spark.sql.benchmark.KyuubiBenchmarkBase
+import org.apache.spark.sql.internal.StaticSQLConf
+
+import org.apache.kyuubi.sql.zorder.ZorderBytesUtils
+
+/**
+ * Benchmark to measure performance with zorder core.
+ *
+ * {{{
+ * RUN_BENCHMARK=1 ./build/mvn clean test \
+ * -pl extensions/spark/kyuubi-extension-spark-4-2 -am \
+ * -Pspark-4.2,kyuubi-extension-spark-4-2,scala-2.13 \
+ * -Dtest=none -DwildcardSuites=org.apache.spark.sql.ZorderCoreBenchmark
+ * }}}
+ */
+class ZorderCoreBenchmark extends KyuubiSparkSQLExtensionTest with KyuubiBenchmarkBase {
+ private val runBenchmark = sys.env.contains("RUN_BENCHMARK")
+ private val numRows = 1 * 1000 * 1000
+
+ private def randomInt(numColumns: Int): Seq[Array[Any]] = {
+ (1 to numRows).map { l =>
+ val arr = new Array[Any](numColumns)
+ (0 until numColumns).foreach(col => arr(col) = l)
+ arr
+ }
+ }
+
+ private def randomLong(numColumns: Int): Seq[Array[Any]] = {
+ (1 to numRows).map { l =>
+ val arr = new Array[Any](numColumns)
+ (0 until numColumns).foreach(col => arr(col) = l.toLong)
+ arr
+ }
+ }
+
+ private def interleaveMultiByteArrayBenchmark(): Unit = {
+ val benchmark =
+ new Benchmark(s"$numRows rows zorder core benchmark", numRows, output = output)
+ benchmark.addCase("2 int columns benchmark", 3) { _ =>
+ randomInt(2).foreach(ZorderBytesUtils.interleaveBits)
+ }
+
+ benchmark.addCase("3 int columns benchmark", 3) { _ =>
+ randomInt(3).foreach(ZorderBytesUtils.interleaveBits)
+ }
+
+ benchmark.addCase("4 int columns benchmark", 3) { _ =>
+ randomInt(4).foreach(ZorderBytesUtils.interleaveBits)
+ }
+
+ benchmark.addCase("2 long columns benchmark", 3) { _ =>
+ randomLong(2).foreach(ZorderBytesUtils.interleaveBits)
+ }
+
+ benchmark.addCase("3 long columns benchmark", 3) { _ =>
+ randomLong(3).foreach(ZorderBytesUtils.interleaveBits)
+ }
+
+ benchmark.addCase("4 long columns benchmark", 3) { _ =>
+ randomLong(4).foreach(ZorderBytesUtils.interleaveBits)
+ }
+
+ benchmark.run()
+ }
+
+ private def paddingTo8ByteBenchmark() {
+ val iterations = 10 * 1000 * 1000
+
+ val b2 = Array('a'.toByte, 'b'.toByte)
+ val benchmark =
+ new Benchmark(s"$iterations iterations paddingTo8Byte benchmark", iterations, output = output)
+ benchmark.addCase("2 length benchmark", 3) { _ =>
+ (1 to iterations).foreach(_ => ZorderBytesUtils.paddingTo8Byte(b2))
+ }
+
+ val b16 = Array.tabulate(16) { i => i.toByte }
+ benchmark.addCase("16 length benchmark", 3) { _ =>
+ (1 to iterations).foreach(_ => ZorderBytesUtils.paddingTo8Byte(b16))
+ }
+
+ benchmark.run()
+ }
+
+ test("zorder core benchmark") {
+ assume(runBenchmark)
+
+ withHeader {
+ interleaveMultiByteArrayBenchmark()
+ paddingTo8ByteBenchmark()
+ }
+ }
+
+ override def sparkConf(): SparkConf = {
+ super.sparkConf().remove(StaticSQLConf.SPARK_SESSION_EXTENSIONS.key)
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ZorderSuite.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ZorderSuite.scala
new file mode 100644
index 00000000000..53be2153dec
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/ZorderSuite.scala
@@ -0,0 +1,918 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql
+
+import org.apache.spark.SparkConf
+import org.apache.spark.sql.catalyst.{InternalRow, TableIdentifier}
+import org.apache.spark.sql.catalyst.analysis.{UnresolvedAttribute, UnresolvedFunction, UnresolvedRelation, UnresolvedStar}
+import org.apache.spark.sql.catalyst.expressions.{Alias, Ascending, AttributeReference, EqualTo, Expression, ExpressionEvalHelper, Literal, NullsLast, SortOrder}
+import org.apache.spark.sql.catalyst.parser.{ParseException, ParserInterface}
+import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, OneRowRelation, Project, RebalancePartitions, Sort}
+import org.apache.spark.sql.classic.ColumnConversions.toRichColumn
+import org.apache.spark.sql.classic.Dataset
+import org.apache.spark.sql.execution.datasources.InsertIntoHadoopFsRelationCommand
+import org.apache.spark.sql.functions._
+import org.apache.spark.sql.hive.execution.InsertIntoHiveTable
+import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf}
+import org.apache.spark.sql.types._
+
+import org.apache.kyuubi.sql.{KyuubiSQLConf, KyuubiSQLExtensionException, SparkKyuubiSparkSQLParser}
+import org.apache.kyuubi.sql.zorder.{OptimizeZorderCommand, OptimizeZorderStatement, Zorder, ZorderBytesUtils}
+
+trait ZorderSuite extends KyuubiSparkSQLExtensionTest with ExpressionEvalHelper {
+ override def sparkConf(): SparkConf = {
+ super.sparkConf()
+ .set(
+ StaticSQLConf.SPARK_SESSION_EXTENSIONS.key,
+ "org.apache.kyuubi.sql.KyuubiSparkSQLExtension")
+ }
+
+ test("optimize unpartitioned table") {
+ withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "1") {
+ withTable("up") {
+ sql(s"DROP TABLE IF EXISTS up")
+
+ val target = Seq(
+ Seq(0, 0),
+ Seq(1, 0),
+ Seq(0, 1),
+ Seq(1, 1),
+ Seq(2, 0),
+ Seq(3, 0),
+ Seq(2, 1),
+ Seq(3, 1),
+ Seq(0, 2),
+ Seq(1, 2),
+ Seq(0, 3),
+ Seq(1, 3),
+ Seq(2, 2),
+ Seq(3, 2),
+ Seq(2, 3),
+ Seq(3, 3))
+ sql(s"CREATE TABLE up (c1 INT, c2 INT, c3 INT) STORED AS PARQUET")
+ sql(s"INSERT INTO TABLE up VALUES" +
+ "(0,0,2),(0,1,2),(0,2,1),(0,3,3)," +
+ "(1,0,4),(1,1,2),(1,2,1),(1,3,3)," +
+ "(2,0,2),(2,1,1),(2,2,5),(2,3,5)," +
+ "(3,0,3),(3,1,4),(3,2,9),(3,3,0)")
+
+ val e = intercept[KyuubiSQLExtensionException] {
+ sql("OPTIMIZE up WHERE c1 > 1 ZORDER BY c1, c2")
+ }
+ assert(e.getMessage == "Filters are only supported for partitioned table")
+
+ sql("OPTIMIZE up ZORDER BY c1, c2")
+ val res = sql("SELECT c1, c2 FROM up").collect()
+
+ assert(res.length == 16)
+
+ for (i <- target.indices) {
+ val t = target(i)
+ val r = res(i)
+ assert(t(0) == r.getInt(0))
+ assert(t(1) == r.getInt(1))
+ }
+ }
+ }
+ }
+
+ test("optimize partitioned table") {
+ withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "1") {
+ withTable("p") {
+ sql("DROP TABLE IF EXISTS p")
+
+ val target = Seq(
+ Seq(0, 0),
+ Seq(1, 0),
+ Seq(0, 1),
+ Seq(1, 1),
+ Seq(2, 0),
+ Seq(3, 0),
+ Seq(2, 1),
+ Seq(3, 1),
+ Seq(0, 2),
+ Seq(1, 2),
+ Seq(0, 3),
+ Seq(1, 3),
+ Seq(2, 2),
+ Seq(3, 2),
+ Seq(2, 3),
+ Seq(3, 3))
+
+ sql(s"CREATE TABLE p (c1 INT, c2 INT, c3 INT) PARTITIONED BY (id INT) STORED AS PARQUET")
+ sql(s"ALTER TABLE p ADD PARTITION (id = 1)")
+ sql(s"ALTER TABLE p ADD PARTITION (id = 2)")
+ sql(s"INSERT INTO TABLE p PARTITION (id = 1) VALUES" +
+ "(0,0,2),(0,1,2),(0,2,1),(0,3,3)," +
+ "(1,0,4),(1,1,2),(1,2,1),(1,3,3)," +
+ "(2,0,2),(2,1,1),(2,2,5),(2,3,5)," +
+ "(3,0,3),(3,1,4),(3,2,9),(3,3,0)")
+ sql(s"INSERT INTO TABLE p PARTITION (id = 2) VALUES" +
+ "(0,0,2),(0,1,2),(0,2,1),(0,3,3)," +
+ "(1,0,4),(1,1,2),(1,2,1),(1,3,3)," +
+ "(2,0,2),(2,1,1),(2,2,5),(2,3,5)," +
+ "(3,0,3),(3,1,4),(3,2,9),(3,3,0)")
+
+ sql(s"OPTIMIZE p ZORDER BY c1, c2")
+
+ val res1 = sql(s"SELECT c1, c2 FROM p WHERE id = 1").collect()
+ val res2 = sql(s"SELECT c1, c2 FROM p WHERE id = 2").collect()
+
+ assert(res1.length == 16)
+ assert(res2.length == 16)
+
+ for (i <- target.indices) {
+ val t = target(i)
+ val r1 = res1(i)
+ assert(t(0) == r1.getInt(0))
+ assert(t(1) == r1.getInt(1))
+
+ val r2 = res2(i)
+ assert(t(0) == r2.getInt(0))
+ assert(t(1) == r2.getInt(1))
+ }
+ }
+ }
+ }
+
+ test("optimize partitioned table with filters") {
+ withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "1") {
+ withTable("p") {
+ sql("DROP TABLE IF EXISTS p")
+
+ val target1 = Seq(
+ Seq(0, 0),
+ Seq(1, 0),
+ Seq(0, 1),
+ Seq(1, 1),
+ Seq(2, 0),
+ Seq(3, 0),
+ Seq(2, 1),
+ Seq(3, 1),
+ Seq(0, 2),
+ Seq(1, 2),
+ Seq(0, 3),
+ Seq(1, 3),
+ Seq(2, 2),
+ Seq(3, 2),
+ Seq(2, 3),
+ Seq(3, 3))
+ val target2 = Seq(
+ Seq(0, 0),
+ Seq(0, 1),
+ Seq(0, 2),
+ Seq(0, 3),
+ Seq(1, 0),
+ Seq(1, 1),
+ Seq(1, 2),
+ Seq(1, 3),
+ Seq(2, 0),
+ Seq(2, 1),
+ Seq(2, 2),
+ Seq(2, 3),
+ Seq(3, 0),
+ Seq(3, 1),
+ Seq(3, 2),
+ Seq(3, 3))
+ sql(s"CREATE TABLE p (c1 INT, c2 INT, c3 INT) PARTITIONED BY (id INT) STORED AS PARQUET")
+ sql(s"ALTER TABLE p ADD PARTITION (id = 1)")
+ sql(s"ALTER TABLE p ADD PARTITION (id = 2)")
+ sql(s"INSERT INTO TABLE p PARTITION (id = 1) VALUES" +
+ "(0,0,2),(0,1,2),(0,2,1),(0,3,3)," +
+ "(1,0,4),(1,1,2),(1,2,1),(1,3,3)," +
+ "(2,0,2),(2,1,1),(2,2,5),(2,3,5)," +
+ "(3,0,3),(3,1,4),(3,2,9),(3,3,0)")
+ sql(s"INSERT INTO TABLE p PARTITION (id = 2) VALUES" +
+ "(0,0,2),(0,1,2),(0,2,1),(0,3,3)," +
+ "(1,0,4),(1,1,2),(1,2,1),(1,3,3)," +
+ "(2,0,2),(2,1,1),(2,2,5),(2,3,5)," +
+ "(3,0,3),(3,1,4),(3,2,9),(3,3,0)")
+
+ val e = intercept[KyuubiSQLExtensionException](
+ sql(s"OPTIMIZE p WHERE id = 1 AND c1 > 1 ZORDER BY c1, c2"))
+ assert(e.getMessage == "Only partition column filters are allowed")
+
+ sql(s"OPTIMIZE p WHERE id = 1 ZORDER BY c1, c2")
+
+ val res1 = sql(s"SELECT c1, c2 FROM p WHERE id = 1").collect()
+ val res2 = sql(s"SELECT c1, c2 FROM p WHERE id = 2").collect()
+
+ assert(res1.length == 16)
+ assert(res2.length == 16)
+
+ for (i <- target1.indices) {
+ val t1 = target1(i)
+ val r1 = res1(i)
+ assert(t1(0) == r1.getInt(0))
+ assert(t1(1) == r1.getInt(1))
+
+ val t2 = target2(i)
+ val r2 = res2(i)
+ assert(t2(0) == r2.getInt(0))
+ assert(t2(1) == r2.getInt(1))
+ }
+ }
+ }
+ }
+
+ test("optimize zorder with datasource table") {
+ // TODO remove this if we support datasource table
+ withTable("t") {
+ sql("CREATE TABLE t (c1 int, c2 int) USING PARQUET")
+ val msg = intercept[KyuubiSQLExtensionException] {
+ sql("OPTIMIZE t ZORDER BY c1, c2")
+ }.getMessage
+ assert(msg.contains("only support hive table"))
+ }
+ }
+
+ private def checkZorderTable(
+ enabled: Boolean,
+ cols: String,
+ planHasRepartition: Boolean,
+ resHasSort: Boolean): Unit = {
+ def checkSort(plan: LogicalPlan): Unit = {
+ assert(plan.isInstanceOf[Sort] === resHasSort)
+ plan match {
+ case sort: Sort =>
+ val colArr = cols.split(",")
+ val refs =
+ if (colArr.length == 1) {
+ sort.order.head
+ .child.asInstanceOf[AttributeReference] :: Nil
+ } else {
+ sort.order.head
+ .child.asInstanceOf[Zorder].children.map(_.references.head)
+ }
+ assert(refs.size === colArr.size)
+ refs.zip(colArr).foreach { case (ref, col) =>
+ assert(ref.name === col.trim)
+ }
+ case _ =>
+ }
+ }
+
+ val repartition =
+ if (planHasRepartition) {
+ "/*+ repartition */"
+ } else {
+ ""
+ }
+ withSQLConf("spark.sql.shuffle.partitions" -> "1") {
+ // hive
+ withSQLConf("spark.sql.hive.convertMetastoreParquet" -> "false") {
+ withTable("zorder_t1", "zorder_t2_true", "zorder_t2_false") {
+ sql(
+ s"""
+ |CREATE TABLE zorder_t1 (c1 int, c2 string, c3 long, c4 double) STORED AS PARQUET
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled' = '$enabled',
+ | 'kyuubi.zorder.cols' = '$cols')
+ |""".stripMargin)
+ val df1 = sql(
+ s"""
+ |INSERT INTO TABLE zorder_t1
+ |SELECT $repartition * FROM VALUES(1,'a',2,4D),(2,'b',3,6D)
+ |""".stripMargin)
+ assert(df1.queryExecution.analyzed.isInstanceOf[InsertIntoHiveTable])
+ checkSort(df1.queryExecution.analyzed.children.head)
+
+ Seq("true", "false").foreach { optimized =>
+ withSQLConf(
+ "spark.sql.hive.convertMetastoreCtas" -> optimized,
+ "spark.sql.hive.convertMetastoreParquet" -> optimized) {
+
+ withListener(
+ s"""
+ |CREATE TABLE zorder_t2_$optimized STORED AS PARQUET
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled' = '$enabled',
+ | 'kyuubi.zorder.cols' = '$cols')
+ |
+ |SELECT $repartition * FROM
+ |VALUES(1,'a',2,4D),(2,'b',3,6D) AS t(c1 ,c2 , c3, c4)
+ |""".stripMargin) { write =>
+ if (optimized.toBoolean) {
+ assert(write.isInstanceOf[InsertIntoHadoopFsRelationCommand])
+ } else {
+ assert(write.isInstanceOf[InsertIntoHiveTable])
+ }
+ checkSort(write.query)
+ }
+ }
+ }
+ }
+ }
+
+ // datasource
+ withTable("zorder_t3", "zorder_t4") {
+ sql(
+ s"""
+ |CREATE TABLE zorder_t3 (c1 int, c2 string, c3 long, c4 double) USING PARQUET
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled' = '$enabled',
+ | 'kyuubi.zorder.cols' = '$cols')
+ |""".stripMargin)
+ val df1 = sql(
+ s"""
+ |INSERT INTO TABLE zorder_t3
+ |SELECT $repartition * FROM VALUES(1,'a',2,4D),(2,'b',3,6D)
+ |""".stripMargin)
+ assert(df1.queryExecution.analyzed.isInstanceOf[InsertIntoHadoopFsRelationCommand])
+ checkSort(df1.queryExecution.analyzed.children.head)
+
+ withListener(
+ s"""
+ |CREATE TABLE zorder_t4 USING PARQUET
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled' = '$enabled',
+ | 'kyuubi.zorder.cols' = '$cols')
+ |
+ |SELECT $repartition * FROM
+ |VALUES(1,'a',2,4D),(2,'b',3,6D) AS t(c1 ,c2 , c3, c4)
+ |""".stripMargin) { write =>
+ assert(write.isInstanceOf[InsertIntoHadoopFsRelationCommand])
+ checkSort(write.query)
+ }
+ }
+ }
+ }
+
+ test("Support insert zorder by table properties") {
+ withSQLConf(KyuubiSQLConf.INSERT_ZORDER_BEFORE_WRITING.key -> "false") {
+ checkZorderTable(true, "c1", false, false)
+ checkZorderTable(false, "c1", false, false)
+ }
+ withSQLConf(KyuubiSQLConf.INSERT_ZORDER_BEFORE_WRITING.key -> "true") {
+ checkZorderTable(true, "", false, false)
+ checkZorderTable(true, "c5", false, false)
+ checkZorderTable(true, "c1,c5", false, false)
+ checkZorderTable(false, "c3", false, false)
+ checkZorderTable(true, "c3", true, false)
+ checkZorderTable(true, "c3", false, true)
+ checkZorderTable(true, "c2,c4", false, true)
+ checkZorderTable(true, "c4, c2, c1, c3", false, true)
+ }
+ }
+
+ test("zorder: check unsupported data type") {
+ def checkZorderPlan(zorder: Expression): Unit = {
+ val msg = intercept[AnalysisException] {
+ val plan = Project(Seq(Alias(zorder, "c")()), OneRowRelation())
+ spark.sessionState.analyzer.checkAnalysis(plan)
+ }.getMessage
+ assert(msg.contains("Unsupported z-order type: void"))
+ }
+
+ checkZorderPlan(Zorder(Seq(Literal(null, NullType))))
+ checkZorderPlan(Zorder(Seq(Literal(1, IntegerType), Literal(null, NullType))))
+ }
+
+ test("zorder: check supported data type") {
+ val children = Seq(
+ Literal.create(false, BooleanType),
+ Literal.create(null, BooleanType),
+ Literal.create(1.toByte, ByteType),
+ Literal.create(null, ByteType),
+ Literal.create(1.toShort, ShortType),
+ Literal.create(null, ShortType),
+ Literal.create(1, IntegerType),
+ Literal.create(null, IntegerType),
+ Literal.create(1L, LongType),
+ Literal.create(null, LongType),
+ Literal.create(1f, FloatType),
+ Literal.create(null, FloatType),
+ Literal.create(1d, DoubleType),
+ Literal.create(null, DoubleType),
+ Literal.create("1", StringType),
+ Literal.create(null, StringType),
+ Literal.create(1L, TimestampType),
+ Literal.create(null, TimestampType),
+ Literal.create(1, DateType),
+ Literal.create(null, DateType),
+ Literal.create(BigDecimal(1, 1), DecimalType(1, 1)),
+ Literal.create(null, DecimalType(1, 1)))
+ val zorder = Zorder(children)
+ val plan = Project(Seq(Alias(zorder, "c")()), OneRowRelation())
+ spark.sessionState.analyzer.checkAnalysis(plan)
+ assert(zorder.foldable)
+
+ // val resultGen = org.apache.commons.codec.binary.Hex.encodeHex(
+ // zorder.eval(InternalRow.fromSeq(children)).asInstanceOf[Array[Byte]], false)
+ // resultGen.grouped(2).zipWithIndex.foreach { case (char, i) =>
+ // print("0x" + char(0) + char(1) + ", ")
+ // if ((i + 1) % 10 == 0) {
+ // println()
+ // }
+ // }
+
+ val expected = Array(
+ 0xFB, 0xEA, 0xAA, 0xBA, 0xAE, 0xAB, 0xAA, 0xEA, 0xBA, 0xAE, 0xAB, 0xAA, 0xEA, 0xBA, 0xA6,
+ 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA,
+ 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xBA, 0xBB, 0xAA, 0xAA, 0xAA,
+ 0xBA, 0xAA, 0xBA, 0xAA, 0xBA, 0xAA, 0xBA, 0xAA, 0xBA, 0xAA, 0xBA, 0xAA, 0x9A, 0xAA, 0xAA,
+ 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xEA,
+ 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA,
+ 0xAA, 0xAA, 0xBE, 0xAA, 0xAA, 0x8A, 0xBA, 0xAA, 0x2A, 0xEA, 0xA8, 0xAA, 0xAA, 0xA2, 0xAA,
+ 0xAA, 0x8A, 0xAA, 0xAA, 0x2F, 0xEB, 0xFE)
+ .map(_.toByte)
+ checkEvaluation(zorder, expected, InternalRow.fromSeq(children))
+ }
+
+ private def checkSort(input: DataFrame, expected: Seq[Row], dataType: Array[DataType]): Unit = {
+ withTempDir { dir =>
+ input.repartition(3).write.mode("overwrite").format("parquet").save(dir.getCanonicalPath)
+ val df = spark.read.format("parquet")
+ .load(dir.getCanonicalPath)
+ .repartition(1)
+ assert(df.schema.fields.map(_.dataType).sameElements(dataType))
+ val exprs = Seq("c1", "c2").map(col).map(_.expr)
+ val sortOrder = SortOrder(Zorder(exprs), Ascending, NullsLast, Seq.empty)
+ val zorderSort = Sort(Seq(sortOrder), true, df.logicalPlan)
+ val result = Dataset.ofRows(spark, zorderSort)
+ checkAnswer(result, expected)
+ }
+ }
+
+ test("sort with zorder -- boolean column") {
+ val schema = StructType(StructField("c1", BooleanType) :: StructField("c2", BooleanType) :: Nil)
+ val nonNullDF = spark.createDataFrame(
+ spark.sparkContext.parallelize(
+ Seq(Row(false, false), Row(false, true), Row(true, false), Row(true, true))),
+ schema)
+ val expected =
+ Row(false, false) :: Row(true, false) :: Row(false, true) :: Row(true, true) :: Nil
+ checkSort(nonNullDF, expected, Array(BooleanType, BooleanType))
+ val df = spark.createDataFrame(
+ spark.sparkContext.parallelize(
+ Seq(Row(false, false), Row(false, null), Row(null, false), Row(null, null))),
+ schema)
+ val expected2 =
+ Row(false, false) :: Row(null, false) :: Row(false, null) :: Row(null, null) :: Nil
+ checkSort(df, expected2, Array(BooleanType, BooleanType))
+ }
+
+ test("sort with zorder -- int column") {
+ // TODO: add more datatype unit test
+ val session = spark
+ import session.implicits._
+ // generate 4 * 4 matrix
+ val len = 3
+ val input = spark.range(len + 1).selectExpr("cast(id as int) as c1")
+ .select($"c1", explode(sequence(lit(0), lit(len))) as "c2")
+ val expected =
+ Row(0, 0) :: Row(1, 0) :: Row(0, 1) :: Row(1, 1) ::
+ Row(2, 0) :: Row(3, 0) :: Row(2, 1) :: Row(3, 1) ::
+ Row(0, 2) :: Row(1, 2) :: Row(0, 3) :: Row(1, 3) ::
+ Row(2, 2) :: Row(3, 2) :: Row(2, 3) :: Row(3, 3) :: Nil
+ checkSort(input, expected, Array(IntegerType, IntegerType))
+
+ // contains null value case.
+ val nullDF = spark.range(1).selectExpr("cast(null as int) as c1")
+ val input2 = spark.range(len).selectExpr("cast(id as int) as c1")
+ .union(nullDF)
+ .select(
+ $"c1",
+ explode(concat(sequence(lit(0), lit(len - 1)), array(lit(null)))) as "c2")
+ val expected2 = Row(0, 0) :: Row(1, 0) :: Row(0, 1) :: Row(1, 1) ::
+ Row(2, 0) :: Row(2, 1) :: Row(0, 2) :: Row(1, 2) ::
+ Row(2, 2) :: Row(null, 0) :: Row(null, 1) :: Row(null, 2) ::
+ Row(0, null) :: Row(1, null) :: Row(2, null) :: Row(null, null) :: Nil
+ checkSort(input2, expected2, Array(IntegerType, IntegerType))
+ }
+
+ test("sort with zorder -- string column") {
+ val schema = StructType(StructField("c1", StringType) :: StructField("c2", StringType) :: Nil)
+ val rdd = spark.sparkContext.parallelize(Seq(
+ Row("a", "a"),
+ Row("a", "b"),
+ Row("a", "c"),
+ Row("a", "d"),
+ Row("b", "a"),
+ Row("b", "b"),
+ Row("b", "c"),
+ Row("b", "d"),
+ Row("c", "a"),
+ Row("c", "b"),
+ Row("c", "c"),
+ Row("c", "d"),
+ Row("d", "a"),
+ Row("d", "b"),
+ Row("d", "c"),
+ Row("d", "d")))
+ val input = spark.createDataFrame(rdd, schema)
+ val expected = Row("a", "a") :: Row("b", "a") :: Row("c", "a") :: Row("a", "b") ::
+ Row("a", "c") :: Row("b", "b") :: Row("c", "b") :: Row("b", "c") ::
+ Row("c", "c") :: Row("d", "a") :: Row("d", "b") :: Row("d", "c") ::
+ Row("a", "d") :: Row("b", "d") :: Row("c", "d") :: Row("d", "d") :: Nil
+ checkSort(input, expected, Array(StringType, StringType))
+
+ val rdd2 = spark.sparkContext.parallelize(Seq(
+ Row(null, "a"),
+ Row("a", "b"),
+ Row("a", "c"),
+ Row("a", null),
+ Row("b", "a"),
+ Row(null, "b"),
+ Row("b", null),
+ Row("b", "d"),
+ Row("c", "a"),
+ Row("c", null),
+ Row(null, "c"),
+ Row("c", "d"),
+ Row("d", null),
+ Row("d", "b"),
+ Row("d", "c"),
+ Row(null, "d"),
+ Row(null, null)))
+ val input2 = spark.createDataFrame(rdd2, schema)
+ val expected2 = Row("b", "a") :: Row("c", "a") :: Row("a", "b") :: Row("a", "c") ::
+ Row("d", "b") :: Row("d", "c") :: Row("b", "d") :: Row("c", "d") ::
+ Row(null, "a") :: Row(null, "b") :: Row(null, "c") :: Row(null, "d") ::
+ Row("a", null) :: Row("b", null) :: Row("c", null) :: Row("d", null) ::
+ Row(null, null) :: Nil
+ checkSort(input2, expected2, Array(StringType, StringType))
+ }
+
+ test("test special value of short int long type") {
+ val df1 = spark.createDataFrame(Seq(
+ (-1, -1L),
+ (Int.MinValue, Int.MinValue.toLong),
+ (1, 1L),
+ (Int.MaxValue - 1, Int.MaxValue.toLong),
+ (Int.MaxValue - 1, Int.MaxValue.toLong - 1),
+ (Int.MaxValue, Int.MaxValue.toLong + 1),
+ (Int.MaxValue, Int.MaxValue.toLong))).toDF("c1", "c2")
+ val expected1 =
+ Row(Int.MinValue, Int.MinValue.toLong) ::
+ Row(-1, -1L) ::
+ Row(1, 1L) ::
+ Row(Int.MaxValue - 1, Int.MaxValue.toLong - 1) ::
+ Row(Int.MaxValue - 1, Int.MaxValue.toLong) ::
+ Row(Int.MaxValue, Int.MaxValue.toLong) ::
+ Row(Int.MaxValue, Int.MaxValue.toLong + 1) :: Nil
+ checkSort(df1, expected1, Array(IntegerType, LongType))
+
+ val df2 = spark.createDataFrame(Seq(
+ (-1, -1.toShort),
+ (Short.MinValue.toInt, Short.MinValue),
+ (1, 1.toShort),
+ (Short.MaxValue.toInt, (Short.MaxValue - 1).toShort),
+ (Short.MaxValue.toInt + 1, (Short.MaxValue - 1).toShort),
+ (Short.MaxValue.toInt, Short.MaxValue),
+ (Short.MaxValue.toInt + 1, Short.MaxValue))).toDF("c1", "c2")
+ val expected2 =
+ Row(Short.MinValue.toInt, Short.MinValue) ::
+ Row(-1, -1.toShort) ::
+ Row(1, 1.toShort) ::
+ Row(Short.MaxValue.toInt, Short.MaxValue - 1) ::
+ Row(Short.MaxValue.toInt, Short.MaxValue) ::
+ Row(Short.MaxValue.toInt + 1, Short.MaxValue - 1) ::
+ Row(Short.MaxValue.toInt + 1, Short.MaxValue) :: Nil
+ checkSort(df2, expected2, Array(IntegerType, ShortType))
+
+ val df3 = spark.createDataFrame(Seq(
+ (-1L, -1.toShort),
+ (Short.MinValue.toLong, Short.MinValue),
+ (1L, 1.toShort),
+ (Short.MaxValue.toLong, (Short.MaxValue - 1).toShort),
+ (Short.MaxValue.toLong + 1, (Short.MaxValue - 1).toShort),
+ (Short.MaxValue.toLong, Short.MaxValue),
+ (Short.MaxValue.toLong + 1, Short.MaxValue))).toDF("c1", "c2")
+ val expected3 =
+ Row(Short.MinValue.toLong, Short.MinValue) ::
+ Row(-1L, -1.toShort) ::
+ Row(1L, 1.toShort) ::
+ Row(Short.MaxValue.toLong, Short.MaxValue - 1) ::
+ Row(Short.MaxValue.toLong, Short.MaxValue) ::
+ Row(Short.MaxValue.toLong + 1, Short.MaxValue - 1) ::
+ Row(Short.MaxValue.toLong + 1, Short.MaxValue) :: Nil
+ checkSort(df3, expected3, Array(LongType, ShortType))
+ }
+
+ test("skip zorder if only requires one column") {
+ withTable("t") {
+ withSQLConf("spark.sql.hive.convertMetastoreParquet" -> "false") {
+ sql("CREATE TABLE t (c1 int, c2 string) STORED AS PARQUET")
+ val order1 = sql("OPTIMIZE t ZORDER BY c1").queryExecution.analyzed
+ .asInstanceOf[OptimizeZorderCommand].query.asInstanceOf[Sort].order.head.child
+ assert(!order1.isInstanceOf[Zorder])
+ assert(order1.isInstanceOf[AttributeReference])
+ }
+ }
+ }
+
+ test("Add config to control if zorder using global sort") {
+ withTable("t") {
+ withSQLConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED.key -> "false") {
+ sql(
+ """
+ |CREATE TABLE t (c1 int, c2 string) STORED AS PARQUET
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled'= 'true',
+ | 'kyuubi.zorder.cols'= 'c1,c2')
+ |""".stripMargin)
+ val p1 = sql("OPTIMIZE t ZORDER BY c1, c2").queryExecution.analyzed
+ assert(p1.collect {
+ case shuffle: Sort if !shuffle.global => shuffle
+ }.size == 1)
+
+ val p2 = sql("INSERT INTO TABLE t SELECT * FROM VALUES(1,'a')").queryExecution.analyzed
+ assert(p2.collect {
+ case shuffle: Sort if !shuffle.global => shuffle
+ }.size == 1)
+ }
+ }
+ }
+
+ test("Allow insert zorder after repartition if zorder using local sort") {
+ withTable("t") {
+ sql(
+ """
+ |CREATE TABLE t (c1 int, c2 string) STORED AS PARQUET
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled'= 'true',
+ | 'kyuubi.zorder.cols'= 'c1,c2')
+ |""".stripMargin)
+ withSQLConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED.key -> "false") {
+ val p1 = sql("INSERT INTO TABLE t SELECT /*+ REPARTITION(1) */* FROM VALUES(1,'a')")
+ .queryExecution.analyzed
+ assert(p1.collect {
+ case sort: Sort if !sort.global => sort
+ }.size == 1)
+ }
+ withSQLConf(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED.key -> "true") {
+ val p2 = sql("INSERT INTO TABLE t SELECT /*+ REPARTITION(1) */* FROM VALUES(1,'a')")
+ .queryExecution.analyzed
+ assert(p2.collect {
+ case sort: Sort if !sort.global => sort
+ }.size == 0)
+ }
+ }
+ }
+
+ test("fast approach test") {
+ Seq[Seq[Any]](
+ Seq(1L, 2L),
+ Seq(1L, 2L, 3L),
+ Seq(1L, 2L, 3L, 4L),
+ Seq(1L, 2L, 3L, 4L, 5L),
+ Seq(1L, 2L, 3L, 4L, 5L, 6L),
+ Seq(1L, 2L, 3L, 4L, 5L, 6L, 7L),
+ Seq(1L, 2L, 3L, 4L, 5L, 6L, 7L, 8L))
+ .foreach { inputs =>
+ assert(java.util.Arrays.equals(
+ ZorderBytesUtils.interleaveBits(inputs.toArray),
+ ZorderBytesUtils.interleaveBitsDefault(inputs.map(ZorderBytesUtils.toByteArray).toArray)))
+ }
+ }
+
+ test("OPTIMIZE command is parsed as expected") {
+ val parser = createParser
+ val globalSort = spark.conf.get(KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED)
+
+ assert(parser.parsePlan("OPTIMIZE p zorder by c1") ===
+ OptimizeZorderStatement(
+ Seq("p"),
+ Sort(
+ SortOrder(UnresolvedAttribute("c1"), Ascending, NullsLast, Seq.empty) :: Nil,
+ globalSort,
+ Project(Seq(UnresolvedStar(None)), UnresolvedRelation(TableIdentifier("p"))))))
+
+ assert(parser.parsePlan("OPTIMIZE p zorder by c1, c2") ===
+ OptimizeZorderStatement(
+ Seq("p"),
+ Sort(
+ SortOrder(
+ Zorder(Seq(UnresolvedAttribute("c1"), UnresolvedAttribute("c2"))),
+ Ascending,
+ NullsLast,
+ Seq.empty) :: Nil,
+ globalSort,
+ Project(Seq(UnresolvedStar(None)), UnresolvedRelation(TableIdentifier("p"))))))
+
+ assert(parser.parsePlan("OPTIMIZE p where id = 1 zorder by c1") ===
+ OptimizeZorderStatement(
+ Seq("p"),
+ Sort(
+ SortOrder(UnresolvedAttribute("c1"), Ascending, NullsLast, Seq.empty) :: Nil,
+ globalSort,
+ Project(
+ Seq(UnresolvedStar(None)),
+ Filter(
+ EqualTo(UnresolvedAttribute("id"), Literal(1)),
+ UnresolvedRelation(TableIdentifier("p")))))))
+
+ assert(parser.parsePlan("OPTIMIZE p where id = 1 zorder by c1, c2") ===
+ OptimizeZorderStatement(
+ Seq("p"),
+ Sort(
+ SortOrder(
+ Zorder(Seq(UnresolvedAttribute("c1"), UnresolvedAttribute("c2"))),
+ Ascending,
+ NullsLast,
+ Seq.empty) :: Nil,
+ globalSort,
+ Project(
+ Seq(UnresolvedStar(None)),
+ Filter(
+ EqualTo(UnresolvedAttribute("id"), Literal(1)),
+ UnresolvedRelation(TableIdentifier("p")))))))
+
+ assert(parser.parsePlan("OPTIMIZE p where id = current_date() zorder by c1") ===
+ OptimizeZorderStatement(
+ Seq("p"),
+ Sort(
+ SortOrder(UnresolvedAttribute("c1"), Ascending, NullsLast, Seq.empty) :: Nil,
+ globalSort,
+ Project(
+ Seq(UnresolvedStar(None)),
+ Filter(
+ EqualTo(
+ UnresolvedAttribute("id"),
+ UnresolvedFunction("current_date", Seq.empty, false)),
+ UnresolvedRelation(TableIdentifier("p")))))))
+
+ // TODO: add following case support
+ intercept[ParseException] {
+ parser.parsePlan("OPTIMIZE p zorder by (c1)")
+ }
+
+ intercept[ParseException] {
+ parser.parsePlan("OPTIMIZE p zorder by (c1, c2)")
+ }
+ }
+
+ test("OPTIMIZE partition predicates constraint") {
+ withTable("p") {
+ sql("CREATE TABLE p (c1 INT, c2 INT) STORED AS PARQUET PARTITIONED BY (event_date DATE)")
+ val e1 = intercept[KyuubiSQLExtensionException] {
+ sql("OPTIMIZE p WHERE event_date = current_date as c ZORDER BY c1, c2")
+ }
+ assert(e1.getMessage.contains("unsupported partition predicates"))
+
+ val e2 = intercept[KyuubiSQLExtensionException] {
+ sql("OPTIMIZE p WHERE c1 = 1 ZORDER BY c1, c2")
+ }
+ assert(e2.getMessage == "Only partition column filters are allowed")
+ }
+ }
+
+ test("optimize sort by backquoted column name") {
+ withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "1") {
+ withTable("up") {
+ sql(s"DROP TABLE IF EXISTS up")
+ val target = Seq(
+ Seq(0, 0),
+ Seq(1, 0),
+ Seq(0, 1),
+ Seq(1, 1),
+ Seq(2, 0),
+ Seq(3, 0),
+ Seq(2, 1),
+ Seq(3, 1),
+ Seq(0, 2),
+ Seq(1, 2),
+ Seq(0, 3),
+ Seq(1, 3),
+ Seq(2, 2),
+ Seq(3, 2),
+ Seq(2, 3),
+ Seq(3, 3))
+ sql(s"CREATE TABLE up (c1 INT, `@c2` INT, c3 INT) STORED AS PARQUET")
+ sql(s"INSERT INTO TABLE up VALUES" +
+ "(0,0,2),(0,1,2),(0,2,1),(0,3,3)," +
+ "(1,0,4),(1,1,2),(1,2,1),(1,3,3)," +
+ "(2,0,2),(2,1,1),(2,2,5),(2,3,5)," +
+ "(3,0,3),(3,1,4),(3,2,9),(3,3,0)")
+
+ sql("OPTIMIZE up ZORDER BY c1, `@c2`")
+ val res = sql("SELECT c1, `@c2` FROM up").collect()
+
+ assert(res.length == 16)
+
+ for (i <- target.indices) {
+ val t = target(i)
+ val r = res(i)
+ assert(t(0) == r.getInt(0))
+ assert(t(1) == r.getInt(1))
+ }
+ }
+ }
+ }
+
+ test("Add rebalance before zorder") {
+ Seq("true" -> false, "false" -> true).foreach { case (useOriginalOrdering, zorder) =>
+ withSQLConf(
+ KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED.key -> "false",
+ KyuubiSQLConf.REBALANCE_BEFORE_ZORDER.key -> "true",
+ KyuubiSQLConf.REBALANCE_ZORDER_COLUMNS_ENABLED.key -> "true",
+ KyuubiSQLConf.ZORDER_USING_ORIGINAL_ORDERING_ENABLED.key -> useOriginalOrdering) {
+ withTable("t") {
+ sql(
+ """
+ |CREATE TABLE t (c1 int, c2 string) STORED AS PARQUET
+ |PARTITIONED BY (d string)
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled'= 'true',
+ | 'kyuubi.zorder.cols'= 'c1,C2')
+ |""".stripMargin)
+ val p = sql("INSERT INTO TABLE t PARTITION(d='a') SELECT * FROM VALUES(1,'a')")
+ .queryExecution.analyzed
+ assert(p.collect {
+ case sort: Sort
+ if !sort.global &&
+ ((sort.order.exists(_.child.isInstanceOf[Zorder]) && zorder) ||
+ (!sort.order.exists(_.child.isInstanceOf[Zorder]) && !zorder)) => sort
+ }.size == 1)
+ assert(p.collect {
+ case rebalance: RebalancePartitions
+ if rebalance.references.map(_.name).exists(_.equals("c1")) => rebalance
+ }.size == 1)
+
+ val p2 = sql("INSERT INTO TABLE t PARTITION(d) SELECT * FROM VALUES(1,'a','b')")
+ .queryExecution.analyzed
+ assert(p2.collect {
+ case sort: Sort
+ if (!sort.global && Seq("c1", "c2", "d").forall(x =>
+ sort.references.map(_.name).exists(_.equals(x)))) &&
+ ((sort.order.exists(_.child.isInstanceOf[Zorder]) && zorder) ||
+ (!sort.order.exists(_.child.isInstanceOf[Zorder]) && !zorder)) => sort
+ }.size == 1)
+ assert(p2.collect {
+ case rebalance: RebalancePartitions
+ if Seq("c1", "c2", "d").forall(x =>
+ rebalance.references.map(_.name).exists(_.equals(x))) => rebalance
+ }.size == 1)
+ }
+ }
+ }
+ }
+
+ test("Two phase rebalance before Z-Order") {
+ withSQLConf(
+ SQLConf.OPTIMIZER_EXCLUDED_RULES.key ->
+ "org.apache.spark.sql.catalyst.optimizer.CollapseRepartition",
+ KyuubiSQLConf.ZORDER_GLOBAL_SORT_ENABLED.key -> "false",
+ KyuubiSQLConf.REBALANCE_BEFORE_ZORDER.key -> "true",
+ KyuubiSQLConf.TWO_PHASE_REBALANCE_BEFORE_ZORDER.key -> "true",
+ KyuubiSQLConf.REBALANCE_ZORDER_COLUMNS_ENABLED.key -> "true") {
+ withTable("t") {
+ sql(
+ """
+ |CREATE TABLE t (c1 int) STORED AS PARQUET
+ |PARTITIONED BY (d string)
+ |TBLPROPERTIES (
+ | 'kyuubi.zorder.enabled'= 'true',
+ | 'kyuubi.zorder.cols'= 'c1')
+ |""".stripMargin)
+ val p = sql("INSERT INTO TABLE t PARTITION(d) SELECT * FROM VALUES(1,'a')")
+ val rebalance = p.queryExecution.optimizedPlan.innerChildren
+ .flatMap(_.collect { case r: RebalancePartitions => r })
+ assert(rebalance.size == 2)
+ assert(rebalance.head.partitionExpressions.flatMap(_.references.map(_.name))
+ .contains("d"))
+ assert(rebalance.head.partitionExpressions.flatMap(_.references.map(_.name))
+ .contains("c1"))
+
+ assert(rebalance(1).partitionExpressions.flatMap(_.references.map(_.name))
+ .contains("d"))
+ assert(!rebalance(1).partitionExpressions.flatMap(_.references.map(_.name))
+ .contains("c1"))
+ }
+ }
+ }
+
+ def createParser: ParserInterface =
+ new SparkKyuubiSparkSQLParser(spark.sessionState.sqlParser)
+}
+
+class ZorderWithCodegenEnabledSuite extends ZorderSuite {
+ override def sparkConf(): SparkConf = {
+ super.sparkConf()
+ .set(SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key, "true")
+ }
+}
+
+class ZorderWithCodegenDisabledSuite extends ZorderSuite {
+ override def sparkConf(): SparkConf = {
+ super.sparkConf()
+ .set(SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key, "false")
+ .set(SQLConf.CODEGEN_FACTORY_MODE.key, "NO_CODEGEN")
+ }
+}
diff --git a/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/benchmark/KyuubiBenchmarkBase.scala b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/benchmark/KyuubiBenchmarkBase.scala
new file mode 100644
index 00000000000..d209f100b6f
--- /dev/null
+++ b/extensions/spark/kyuubi-extension-spark-4-2/src/test/scala/org/apache/spark/sql/benchmark/KyuubiBenchmarkBase.scala
@@ -0,0 +1,71 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.benchmark
+
+import java.io.{File, FileOutputStream, OutputStream}
+
+import scala.jdk.CollectionConverters._
+
+import com.google.common.reflect.ClassPath
+import org.scalatest.Assertions._
+
+trait KyuubiBenchmarkBase {
+ var output: Option[OutputStream] = None
+
+ private val prefix = {
+ val benchmarkClasses = ClassPath.from(Thread.currentThread.getContextClassLoader)
+ .getTopLevelClassesRecursive("org.apache.spark.sql").asScala.toArray
+ assert(benchmarkClasses.nonEmpty)
+ val benchmark = benchmarkClasses.find(_.load().getName.endsWith("Benchmark"))
+ val targetDirOrProjDir =
+ new File(benchmark.get.load().getProtectionDomain.getCodeSource.getLocation.toURI)
+ .getParentFile.getParentFile
+ if (targetDirOrProjDir.getName == "target") {
+ targetDirOrProjDir.getParentFile.getCanonicalPath + "/"
+ } else {
+ targetDirOrProjDir.getCanonicalPath + "/"
+ }
+ }
+
+ def withHeader(func: => Unit): Unit = {
+ val version = System.getProperty("java.version").split("\\D+")(0).toInt
+ val jdkString = if (version > 8) s"-jdk$version" else ""
+ val resultFileName =
+ s"${this.getClass.getSimpleName.replace("$", "")}$jdkString-results.txt"
+ val dir = new File(s"${prefix}benchmarks/")
+ if (!dir.exists()) {
+ // scalastyle:off println
+ println(s"Creating ${dir.getAbsolutePath} for benchmark results.")
+ // scalastyle:on println
+ dir.mkdirs()
+ }
+ val file = new File(dir, resultFileName)
+ if (!file.exists()) {
+ file.createNewFile()
+ }
+ output = Some(new FileOutputStream(file))
+
+ func
+
+ output.foreach { o =>
+ if (o != null) {
+ o.close()
+ }
+ }
+ }
+}
diff --git a/extensions/spark/kyuubi-spark-authz/README.md b/extensions/spark/kyuubi-spark-authz/README.md
index bf23bfd43fe..55291122de3 100644
--- a/extensions/spark/kyuubi-spark-authz/README.md
+++ b/extensions/spark/kyuubi-spark-authz/README.md
@@ -33,6 +33,7 @@ build/mvn clean package -DskipTests -pl :kyuubi-spark-authz_2.12 -am -Dspark.ver
`-Dspark.version=`
+- [x] 4.2.x
- [x] 4.1.x
- [x] 4.0.x
- [x] 3.5.x (default)
diff --git a/extensions/spark/kyuubi-spark-authz/src/main/resources/function_command_spec.json b/extensions/spark/kyuubi-spark-authz/src/main/resources/function_command_spec.json
index 14dad8e2a3f..b0da1a95199 100644
--- a/extensions/spark/kyuubi-spark-authz/src/main/resources/function_command_spec.json
+++ b/extensions/spark/kyuubi-spark-authz/src/main/resources/function_command_spec.json
@@ -97,6 +97,13 @@
}, {
"classname" : "org.apache.spark.sql.execution.command.RefreshFunctionCommand",
"functionDescs" : [ {
+ "fieldName" : "identifier",
+ "fieldExtractor" : "FunctionIdentifierFunctionExtractor",
+ "databaseDesc" : null,
+ "functionTypeDesc" : null,
+ "isInput" : false,
+ "comment" : ""
+ }, {
"fieldName" : "functionName",
"fieldExtractor" : "StringFunctionExtractor",
"databaseDesc" : {
diff --git a/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/gen/FunctionCommands.scala b/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/gen/FunctionCommands.scala
index d5c849dd6bf..9eef2d3317b 100644
--- a/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/gen/FunctionCommands.scala
+++ b/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/gen/FunctionCommands.scala
@@ -80,7 +80,11 @@ object FunctionCommands extends CommandSpecs[FunctionCommandSpec] {
"functionName",
classOf[StringFunctionExtractor],
Some(databaseDesc))
- FunctionCommandSpec(cmd, Seq(functionDesc), RELOADFUNCTION)
+ // SPARK-54866 (4.2.0) consolidated databaseName/functionName into a single identifier field
+ val functionIdentifierDesc = FunctionDesc(
+ "identifier",
+ classOf[FunctionIdentifierFunctionExtractor])
+ FunctionCommandSpec(cmd, Seq(functionIdentifierDesc, functionDesc), RELOADFUNCTION)
}
override def specs: Seq[FunctionCommandSpec] = Seq(
diff --git a/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/datamasking/DataMaskingForIcebergSuite.scala b/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/datamasking/DataMaskingForIcebergSuite.scala
index 3a3cf6011ea..6554e183528 100644
--- a/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/datamasking/DataMaskingForIcebergSuite.scala
+++ b/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/datamasking/DataMaskingForIcebergSuite.scala
@@ -21,7 +21,9 @@ import org.apache.spark.SparkConf
import org.scalatest.Outcome
import org.apache.kyuubi.Utils
+import org.apache.kyuubi.tags.IcebergTest
+@IcebergTest
class DataMaskingForIcebergSuite extends DataMaskingTestBase {
override protected def extraSparkConf: SparkConf = {
super.extraSparkConf
diff --git a/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/rowfiltering/RowFilteringForIcebergSuite.scala b/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/rowfiltering/RowFilteringForIcebergSuite.scala
index 57a9e29b665..002923ec3e8 100644
--- a/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/rowfiltering/RowFilteringForIcebergSuite.scala
+++ b/extensions/spark/kyuubi-spark-authz/src/test/scala/org/apache/kyuubi/plugin/spark/authz/ranger/rowfiltering/RowFilteringForIcebergSuite.scala
@@ -21,7 +21,9 @@ import org.apache.spark.SparkConf
import org.scalatest.Outcome
import org.apache.kyuubi.Utils
+import org.apache.kyuubi.tags.IcebergTest
+@IcebergTest
class RowFilteringForIcebergSuite extends RowFilteringTestBase {
override protected val extraSparkConf: SparkConf = {
new SparkConf()
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/main/scala/org/apache/kyuubi/spark/connector/hive/HiveTableCatalog.scala b/extensions/spark/kyuubi-spark-connector-hive/src/main/scala/org/apache/kyuubi/spark/connector/hive/HiveTableCatalog.scala
index 80406c4e040..2cb6c79f3e5 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/main/scala/org/apache/kyuubi/spark/connector/hive/HiveTableCatalog.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/main/scala/org/apache/kyuubi/spark/connector/hive/HiveTableCatalog.scala
@@ -203,7 +203,58 @@ class HiveTableCatalog(sparkSession: SparkSession)
ignoredProperties: Map[String, String] = Map.empty,
viewOriginalText: Option[String] = None): CatalogTable = {
// scalastyle:on
- Try { // SPARK-50675 (4.0.0)
+ Try { // SPARK-52729 (4.2.0)
+ DynConstructors.builder()
+ .impl(
+ classOf[CatalogTable],
+ classOf[TableIdentifier],
+ classOf[CatalogTableType],
+ classOf[CatalogStorageFormat],
+ classOf[StructType],
+ classOf[Option[String]],
+ classOf[Seq[String]],
+ classOf[Option[BucketSpec]],
+ classOf[String],
+ classOf[Long],
+ classOf[Long],
+ classOf[String],
+ classOf[Map[String, String]],
+ classOf[Option[CatalogStatistics]],
+ classOf[Option[String]],
+ classOf[Option[String]],
+ classOf[Option[String]],
+ classOf[Seq[String]],
+ classOf[Boolean],
+ classOf[Boolean],
+ classOf[Map[String, String]],
+ classOf[Option[String]],
+ classOf[Option[Seq[String]]])
+ .buildChecked()
+ .invokeChecked[CatalogTable](
+ null,
+ identifier,
+ tableType,
+ storage,
+ schema,
+ provider,
+ partitionColumnNames,
+ bucketSpec,
+ owner,
+ createTime,
+ lastAccessTime,
+ createVersion,
+ properties,
+ stats,
+ viewText,
+ comment,
+ collation,
+ unsupportedFeatures,
+ tracksPartitionsInCatalog,
+ schemaPreservesCase,
+ ignoredProperties,
+ viewOriginalText,
+ None)
+ }.recover { case _: Exception => // SPARK-50675 (4.0.0)
DynConstructors.builder()
.impl(
classOf[CatalogTable],
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/DynamicPartitionPruningSuite.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/DynamicPartitionPruningSuite.scala
index 4ddbe9827d9..602c8cc1fc5 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/DynamicPartitionPruningSuite.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/DynamicPartitionPruningSuite.scala
@@ -62,7 +62,7 @@ class DynamicPartitionPruningSuite extends KyuubiHiveTest {
val fact = s"hive.default.dpp_fact_$suffix"
val dim = s"hive.default.dpp_dim_$suffix"
- withTable(fact, dim) {
+ dropTableAfter(fact, dim) {
spark.sql(
s"""
| CREATE TABLE $fact (id INT, v STRING) PARTITIONED BY (dt STRING)
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/HiveQuerySuite.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/HiveQuerySuite.scala
index 5cfde6bd550..33c69e7f912 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/HiveQuerySuite.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/HiveQuerySuite.scala
@@ -123,7 +123,7 @@ class HiveQuerySuite extends KyuubiHiveTest {
"spark.sql.shuffle.partitions" -> "1")) { spark =>
val table = "hive.default.test_part_table"
val tempTable = "hive.default.test_part_table_tmp"
- withTable(table, tempTable) {
+ dropTableAfter(table, tempTable) {
spark.sql(
s"""
| CREATE TABLE $table (
@@ -311,7 +311,7 @@ class HiveQuerySuite extends KyuubiHiveTest {
test("ORC filter pushdown") {
val table = "hive.default.orc_filter_pushdown"
- withTable(table) {
+ dropTableAfter(table) {
spark.sql(
s"""
| CREATE TABLE $table (
@@ -404,7 +404,7 @@ class HiveQuerySuite extends KyuubiHiveTest {
test("PARQUET filter pushdown") {
val table = "hive.default.parquet_filter_pushdown"
- withTable(table) {
+ dropTableAfter(table) {
spark.sql(
s"""
| CREATE TABLE $table (
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/KyuubiHiveTest.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/KyuubiHiveTest.scala
index fb5bcd62184..d5e137cbe4d 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/KyuubiHiveTest.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/KyuubiHiveTest.scala
@@ -81,7 +81,7 @@ abstract class KyuubiHiveTest extends QueryTest with Logging {
/**
* Drops table `tableName` after calling `f`.
*/
- protected def withTable(tableNames: String*)(f: => Unit): Unit = {
+ protected def dropTableAfter(tableNames: String*)(f: => Unit): Unit = {
Utils.tryWithSafeFinally(f) {
tableNames.foreach { name =>
spark.sql(s"DROP TABLE IF EXISTS $name")
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateNamespaceSuite.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateNamespaceSuite.scala
index d6b90cc0419..872c9f378e6 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateNamespaceSuite.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateNamespaceSuite.scala
@@ -45,7 +45,7 @@ trait CreateNamespaceSuiteBase extends DDLCommandTestUtils {
test("basic test") {
val ns = s"$catalogName.$namespace"
- withNamespace(ns) {
+ dropNamespaceAfter(ns) {
sql(s"CREATE NAMESPACE $ns")
assert(getCatalog(catalogName).asNamespaceCatalog.namespaceExists(namespaceArray))
}
@@ -53,7 +53,7 @@ trait CreateNamespaceSuiteBase extends DDLCommandTestUtils {
test("namespace with location") {
val ns = s"$catalogName.$namespace"
- withNamespace(ns) {
+ dropNamespaceAfter(ns) {
withTempDir { tmpDir =>
// The generated temp path is not qualified.
val path = tmpDir.getCanonicalPath
@@ -69,7 +69,7 @@ trait CreateNamespaceSuiteBase extends DDLCommandTestUtils {
sql(s"CREATE NAMESPACE $ns LOCATION '$uri'")
// Make sure the location is qualified.
- val expected = makeQualifiedPath(tmpDir.toString)
+ val expected = qualifyPath(tmpDir.toString)
assert("file" === expected.getScheme)
assert(new Path(getNamespaceLocation(catalogName, namespaceArray)).toUri === expected)
}
@@ -78,7 +78,7 @@ trait CreateNamespaceSuiteBase extends DDLCommandTestUtils {
test("Namespace already exists") {
val ns = s"$catalogName.$namespace"
- withNamespace(ns) {
+ dropNamespaceAfter(ns) {
sql(s"CREATE NAMESPACE $ns")
val e = intercept[NamespaceAlreadyExistsException] {
@@ -105,7 +105,7 @@ trait CreateNamespaceSuiteBase extends DDLCommandTestUtils {
}
withSQLConf((SQLConf.LEGACY_PROPERTY_NON_RESERVED.key, "true")) {
NAMESPACE_RESERVED_PROPERTIES.filterNot(_ == PROP_COMMENT).foreach { key =>
- withNamespace(ns) {
+ dropNamespaceAfter(ns) {
sql(s"CREATE NAMESPACE $ns WITH DBPROPERTIES('$key'='foo')")
assert(
sql(s"DESC NAMESPACE EXTENDED $ns")
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateTableSuite.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateTableSuite.scala
index d26ec420980..5262e289493 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateTableSuite.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/CreateTableSuite.scala
@@ -34,7 +34,7 @@ class CreateTableSuite extends DDLCommandTestUtils {
.catalog(catalogName).asInstanceOf[HiveTableCatalog]
val table = "hive.default.employee"
Seq("parquet", "orc").foreach { provider =>
- withTable(table) {
+ dropTableAfter(table) {
sql(
s"""
| CREATE TABLE IF NOT EXISTS
@@ -57,7 +57,7 @@ class CreateTableSuite extends DDLCommandTestUtils {
.catalog(catalogName).asInstanceOf[HiveTableCatalog]
val table = "hive.default.employee"
Seq("parquet", "orc").foreach { provider =>
- withTable(table) {
+ dropTableAfter(table) {
sql(
s"""
| CREATE TABLE IF NOT EXISTS
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/DDLCommandTestUtils.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/DDLCommandTestUtils.scala
index 2a56025540f..de32fa91ece 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/DDLCommandTestUtils.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/DDLCommandTestUtils.scala
@@ -48,7 +48,7 @@ trait DDLCommandTestUtils extends KyuubiHiveTest {
// So they share a metadata derby db.
protected val builtinNamespace: Seq[String] = Seq("default")
- protected def withNamespace(namespaces: String*)(f: => Unit): Unit = {
+ protected def dropNamespaceAfter(namespaces: String*)(f: => Unit): Unit = {
try {
f
} finally {
@@ -72,7 +72,7 @@ trait DDLCommandTestUtils extends KyuubiHiveTest {
}
}
- protected def makeQualifiedPath(path: String): URI = {
+ protected def qualifyPath(path: String): URI = {
val hadoopPath = new Path(path)
val fs = hadoopPath.getFileSystem(spark.sessionState.newHadoopConf())
fs.makeQualified(hadoopPath).toUri
@@ -83,10 +83,10 @@ trait DDLCommandTestUtils extends KyuubiHiveTest {
tableName: String,
cat: String = catalogName)(f: String => Unit): Unit = {
val nsCat = s"$cat.$ns"
- withNamespace(nsCat) {
+ dropNamespaceAfter(nsCat) {
sql(s"CREATE NAMESPACE $nsCat")
val t = s"$nsCat.$tableName"
- withTable(t) {
+ dropTableAfter(t) {
f(t)
}
}
@@ -95,7 +95,7 @@ trait DDLCommandTestUtils extends KyuubiHiveTest {
/**
* Restores the current catalog/database after calling `f`.
*/
- protected def withCurrentCatalogAndNamespace(f: => Unit): Unit = {
+ protected def resetCatalogAndNamespace(f: => Unit): Unit = {
val curCatalog = sql("select current_catalog()").head().getString(0)
val curDatabase = sql("select current_database()").head().getString(0)
try {
diff --git a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/ShowTablesSuite.scala b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/ShowTablesSuite.scala
index 445ca9fa7a5..abfdbf2202a 100644
--- a/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/ShowTablesSuite.scala
+++ b/extensions/spark/kyuubi-spark-connector-hive/src/test/scala/org/apache/kyuubi/spark/connector/hive/command/ShowTablesSuite.scala
@@ -39,10 +39,10 @@ trait ShowTablesSuiteBase extends DDLCommandTestUtils {
}
test("show tables with a pattern") {
- withNamespace(s"$catalogName.ns1", s"$catalogName.ns2") {
+ dropNamespaceAfter(s"$catalogName.ns1", s"$catalogName.ns2") {
sql(s"CREATE NAMESPACE $catalogName.ns1")
sql(s"CREATE NAMESPACE $catalogName.ns2")
- withTable(
+ dropTableAfter(
s"$catalogName.ns1.table",
s"$catalogName.ns1.table_name_1a",
s"$catalogName.ns1.table_name_2b",
@@ -79,7 +79,7 @@ trait ShowTablesSuiteBase extends DDLCommandTestUtils {
}
test("change current catalog and namespace with USE statements") {
- withCurrentCatalogAndNamespace {
+ resetCatalogAndNamespace {
withNamespaceAndTable("ns", "table") { t =>
sql(s"CREATE TABLE $t (name STRING, id INT) $defaultUsing")
diff --git a/extensions/spark/kyuubi-spark-lineage/README.md b/extensions/spark/kyuubi-spark-lineage/README.md
index 12d27664315..955ad06c352 100644
--- a/extensions/spark/kyuubi-spark-lineage/README.md
+++ b/extensions/spark/kyuubi-spark-lineage/README.md
@@ -33,6 +33,7 @@ build/mvn clean package -DskipTests -pl :kyuubi-spark-lineage_2.12 -am -Dspark.v
`-Dspark.version=`
+- [x] 4.2.x
- [x] 4.1.x
- [x] 4.0.x
- [x] 3.5.x (default)
diff --git a/extensions/spark/kyuubi-spark-lineage/src/test/scala/org/apache/kyuubi/plugin/lineage/helper/SparkSQLLineageParserHelperSuite.scala b/extensions/spark/kyuubi-spark-lineage/src/test/scala/org/apache/kyuubi/plugin/lineage/helper/SparkSQLLineageParserHelperSuite.scala
index 86be8bb0d8f..324ae41fa0a 100644
--- a/extensions/spark/kyuubi-spark-lineage/src/test/scala/org/apache/kyuubi/plugin/lineage/helper/SparkSQLLineageParserHelperSuite.scala
+++ b/extensions/spark/kyuubi-spark-lineage/src/test/scala/org/apache/kyuubi/plugin/lineage/helper/SparkSQLLineageParserHelperSuite.scala
@@ -1502,9 +1502,15 @@ abstract class SparkSQLLineageParserHelperSuite extends KyuubiFunSuite
}
private def extractLineage(sql: String): Lineage = {
+ // Avoid going through QueryExecution.analyzed: SPARK-55855 (4.2.0) made `.analyzed`
+ // lazily begin a transaction on transactional catalogs (QueryExecution.lazyTransactionOpt).
+ // Since this helper only inspects the analyzed plan and never executes the query, that
+ // transaction would be left in the Active state and the next call would trip
+ // InMemoryRowLevelOperationTableCatalog.beginTransaction's "no nested active transaction"
+ // assertion. Analyzing directly produces the same plan.
val parsed = spark.sessionState.sqlParser.parsePlan(sql)
- val qe = spark.sessionState.executePlan(parsed)
- val analyzed = qe.analyzed
+ val analyzed = spark.sessionState.analyzer.execute(parsed)
+ spark.sessionState.analyzer.checkAnalysis(analyzed)
SparkSQLLineageParseHelper(spark).transformToLineage(0, analyzed).get
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentCatalog.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentCatalog.scala
index 1d85d3d5adc..2da6ec230c6 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentCatalog.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentCatalog.scala
@@ -20,6 +20,7 @@ package org.apache.kyuubi.engine.spark.operation
import org.apache.spark.sql.Row
import org.apache.spark.sql.types.StructType
+import org.apache.kyuubi.engine.spark.util.SparkCatalogUtils
import org.apache.kyuubi.operation.IterableFetchIterator
import org.apache.kyuubi.operation.log.OperationLog
import org.apache.kyuubi.operation.meta.ResultSetSchemaConstant.TABLE_CAT
@@ -38,7 +39,7 @@ class GetCurrentCatalog(session: Session) extends SparkOperation(session) {
override protected def runInternal(): Unit = {
try {
- val currentCatalogName = spark.sessionState.catalogManager.currentCatalog.name()
+ val currentCatalogName = SparkCatalogUtils.currentCatalog(spark).name()
iter = new IterableFetchIterator(Seq(Row(currentCatalogName)))
} catch onError()
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentDatabase.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentDatabase.scala
index 2478fb6a49a..95e811c7a43 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentDatabase.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/GetCurrentDatabase.scala
@@ -20,7 +20,7 @@ package org.apache.kyuubi.engine.spark.operation
import org.apache.spark.sql.Row
import org.apache.spark.sql.types.StructType
-import org.apache.kyuubi.engine.spark.util.SparkCatalogUtils.quoteIfNeeded
+import org.apache.kyuubi.engine.spark.util.SparkCatalogUtils.{currentNamespace, quoteIfNeeded}
import org.apache.kyuubi.operation.IterableFetchIterator
import org.apache.kyuubi.operation.log.OperationLog
import org.apache.kyuubi.operation.meta.ResultSetSchemaConstant.TABLE_SCHEM
@@ -40,7 +40,7 @@ class GetCurrentDatabase(session: Session) extends SparkOperation(session) {
override protected def runInternal(): Unit = {
try {
val currentDatabaseName =
- spark.sessionState.catalogManager.currentNamespace.map(quoteIfNeeded).mkString(".")
+ currentNamespace(spark).map(quoteIfNeeded).mkString(".")
iter = new IterableFetchIterator(Seq(Row(currentDatabaseName)))
} catch onError()
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SetCurrentDatabase.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SetCurrentDatabase.scala
index 170be4dcb49..79c63bbf4a2 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SetCurrentDatabase.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SetCurrentDatabase.scala
@@ -17,6 +17,7 @@
package org.apache.kyuubi.engine.spark.operation
+import org.apache.kyuubi.engine.spark.util.SparkCatalogUtils
import org.apache.kyuubi.operation.log.OperationLog
import org.apache.kyuubi.session.Session
@@ -29,7 +30,7 @@ class SetCurrentDatabase(session: Session, database: String)
override protected def runInternal(): Unit = {
try {
- spark.sessionState.catalogManager.setCurrentNamespace(Array(database))
+ SparkCatalogUtils.setCurrentNamespace(spark, Array(database))
setHasResultSet(false)
} catch onError()
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SparkOperation.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SparkOperation.scala
index 6d7af170212..f9d3042822c 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SparkOperation.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/operation/SparkOperation.scala
@@ -23,6 +23,7 @@ import java.time.ZoneId
import org.apache.spark.kyuubi.{SparkProgressMonitor, SQLOperationListener}
import org.apache.spark.kyuubi.SparkUtilsHelper.redact
import org.apache.spark.sql.{DataFrame, Row, SparkSession}
+import org.apache.spark.sql.catalyst.CurrentUserContext.CURRENT_USER
import org.apache.spark.sql.execution.SparkSQLExecutionHelper
import org.apache.spark.sql.types.{BinaryType, StructField, StructType}
import org.apache.spark.ui.SparkUIUtils.formatDuration
@@ -162,6 +163,10 @@ abstract class SparkOperation(session: Session)
spark.sparkContext.setJobGroup(statementId, redactedStatement, forceCancel)
spark.sparkContext.setLocalProperty(KYUUBI_SESSION_USER_KEY, session.user)
spark.sparkContext.setLocalProperty(KYUUBI_STATEMENT_ID_KEY, statementId)
+ // Set the session (proxy) user as the current user so that Spark built-in functions
+ // current_user() (SPARK-21957, 3.2.0), session_user() (SPARK-44860, 4.0.0) and user
+ // return it instead of the engine user.
+ CURRENT_USER.set(session.user)
schedulerPool match {
case Some(pool) =>
spark.sparkContext.setLocalProperty(SPARK_SCHEDULER_POOL_KEY, pool)
@@ -176,6 +181,7 @@ abstract class SparkOperation(session: Session)
spark.sparkContext.setLocalProperty(SPARK_SCHEDULER_POOL_KEY, null)
spark.sparkContext.setLocalProperty(KYUUBI_SESSION_USER_KEY, null)
spark.sparkContext.setLocalProperty(KYUUBI_STATEMENT_ID_KEY, null)
+ CURRENT_USER.remove()
spark.sparkContext.clearJobGroup()
if (isSessionUserSignEnabled) {
clearSessionUserSign()
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/session/SparkSessionImpl.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/session/SparkSessionImpl.scala
index 829fba21866..0a0cb8f6290 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/session/SparkSessionImpl.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/session/SparkSessionImpl.scala
@@ -87,7 +87,7 @@ class SparkSessionImpl(
useCatalogAndDatabaseConf.get("use:database").foreach { database =>
try {
- spark.sessionState.catalogManager.setCurrentNamespace(Array(database))
+ SparkCatalogUtils.setCurrentNamespace(spark, Array(database))
} catch {
case e
if database == "default" &&
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/udf/KDFRegistry.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/udf/KDFRegistry.scala
index a2d50d1515b..bd8d5a22661 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/udf/KDFRegistry.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/udf/KDFRegistry.scala
@@ -26,6 +26,7 @@ import org.apache.spark.sql.functions.udf
import org.apache.kyuubi.{KYUUBI_VERSION, Utils}
import org.apache.kyuubi.config.KyuubiReservedKeys.{KYUUBI_ENGINE_URL, KYUUBI_SESSION_USER_KEY}
+import org.apache.kyuubi.engine.spark.KyuubiSparkUtil.SPARK_ENGINE_RUNTIME_VERSION
object KDFRegistry {
@@ -96,7 +97,13 @@ object KDFRegistry {
def registerAll(spark: SparkSession): Unit = {
for (func <- registeredFunctions) {
- spark.udf.register(func.name, func.udf)
+ // SPARK-44860 (4.0.0) added a built-in `session_user` (alias for `CurrentUser`).
+ // Respect the built-in instead of registering the Kyuubi UDF of the same name.
+ if (func.name == "session_user" && SPARK_ENGINE_RUNTIME_VERSION >= "4.0") {
+ // skip
+ } else {
+ spark.udf.register(func.name, func.udf)
+ }
}
}
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/util/SparkCatalogUtils.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/util/SparkCatalogUtils.scala
index b9a5028acdc..1d3b84edcdf 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/util/SparkCatalogUtils.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/kyuubi/engine/spark/util/SparkCatalogUtils.scala
@@ -44,16 +44,42 @@ object SparkCatalogUtils extends Logging {
// Catalog //
// ///////////////////////////////////////////////////////////////////////////////////////////////
+ // SPARK-46050 (4.2.0) changed CatalogManager from a class to an interface, breaking
+ // binary compatibility when compiled against Spark 3.5 and run against Spark 4.2.
+ // Access catalogManager reflectively to avoid IncompatibleClassChangeError.
+ def catalogManager(spark: SparkSession): AnyRef =
+ invokeAs[AnyRef](spark.sessionState, "catalogManager")
+
+ def setCurrentNamespace(spark: SparkSession, namespace: Array[String]): Unit = {
+ val mgr = catalogManager(spark)
+ val method = mgr.getClass.getMethod("setCurrentNamespace", classOf[Array[String]])
+ try {
+ method.invoke(mgr, namespace)
+ } catch {
+ case e: java.lang.reflect.InvocationTargetException => throw e.getCause
+ }
+ }
+
+ def currentCatalog(spark: SparkSession): CatalogPlugin =
+ invokeAs[CatalogPlugin](catalogManager(spark), "currentCatalog")
+
+ def currentNamespace(spark: SparkSession): Array[String] =
+ invokeAs[Array[String]](catalogManager(spark), "currentNamespace")
+
+ def catalog(spark: SparkSession, name: String): CatalogPlugin =
+ invokeAs[CatalogPlugin](catalogManager(spark), "catalog", (classOf[String], name))
+
+ def isCatalogRegistered(spark: SparkSession, name: String): Boolean =
+ invokeAs[Boolean](catalogManager(spark), "isCatalogRegistered", (classOf[String], name))
+
/**
* Get all register catalogs in Spark's `CatalogManager`
*/
def getCatalogs(spark: SparkSession): Seq[Row] = {
-
- // A [[CatalogManager]] is session unique
- val catalogMgr = spark.sessionState.catalogManager
+ val catalogMgr = catalogManager(spark)
// get the custom v2 session catalog or default spark_catalog
val sessionCatalog = invokeAs[AnyRef](catalogMgr, "v2SessionCatalog")
- val defaultCatalog = catalogMgr.currentCatalog
+ val defaultCatalog = invokeAs[CatalogPlugin](catalogMgr, "currentCatalog")
val defaults = Seq(sessionCatalog, defaultCatalog).distinct.map(invokeAs[String](_, "name"))
val catalogs = getField[scala.collection.Map[String, _]](catalogMgr, "catalogs")
@@ -61,18 +87,17 @@ object SparkCatalogUtils extends Logging {
}
def getCatalog(spark: SparkSession, catalogName: String): CatalogPlugin = {
- val catalogManager = spark.sessionState.catalogManager
if (StringUtils.isBlank(catalogName)) {
- catalogManager.currentCatalog
+ currentCatalog(spark)
} else {
- catalogManager.catalog(catalogName)
+ catalog(spark, catalogName)
}
}
def setCurrentCatalog(spark: SparkSession, catalog: String): Unit = {
- // SPARK-36841(3.3.0) Ensure setCurrentCatalog method catalog must exist
- if (spark.sessionState.catalogManager.isCatalogRegistered(catalog)) {
- spark.sessionState.catalogManager.setCurrentCatalog(catalog)
+ // SPARK-36841 (3.3.0) Ensure setCurrentCatalog method catalog must exist
+ if (isCatalogRegistered(spark, catalog)) {
+ invokeAs[Unit](catalogManager(spark), "setCurrentCatalog", (classOf[String], catalog))
} else {
throw new IllegalArgumentException(s"Cannot find catalog plugin class for catalog '$catalog'")
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/EngineTab.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/EngineTab.scala
index 0afcb5fec5f..0d9f4f2d4d2 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/EngineTab.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/EngineTab.scala
@@ -95,9 +95,14 @@ case class EngineTab(
sparkUI.foreach { ui =>
try {
val sparkServletContextHandlerClz = DynClasses.builder()
- // for official Spark releases and distributions built via Maven
+ // SPARK-47086 (4.2.0) upgraded Jetty to 12 (ee10) and relocated it under
+ // org.sparkproject.jetty.ee10 for official releases/distributions built via Maven
+ .impl("org.sparkproject.jetty.ee10.servlet.ServletContextHandler")
+ // for official Spark 3.x ~ 4.1 releases and distributions built via Maven
.impl("org.sparkproject.jetty.servlet.ServletContextHandler")
- // for distributions built via SBT
+ // for dev Spark 4.2+ distributions built via SBT
+ .impl("org.eclipse.jetty.ee10.servlet.ServletContextHandler")
+ // for dev Spark 3.x ~ 4.1 distributions built via SBT
.impl("org.eclipse.jetty.servlet.ServletContextHandler")
.buildChecked()
val attachHandlerMethod = DynMethods.builder("attachHandler")
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/HttpServletRequestLike.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/HttpServletRequestLike.scala
index 7b66cb87b59..817278edc48 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/HttpServletRequestLike.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/HttpServletRequestLike.scala
@@ -24,6 +24,17 @@ trait HttpServletRequestLike {
def getParameter(name: String): String
def getParameterMap: util.Map[String, Array[String]]
+
+ /**
+ * The underlying raw servlet request, which is required when invoking Spark
+ * private UI helpers (e.g. UIUtils.headerSparkPage) that take a concrete
+ * javax/jakarta HttpServletRequest. Keeping the raw request here, instead of
+ * having the adapter re-implement the whole servlet interface, avoids
+ * depending on servlet API methods that vary across versions (for example,
+ * jakarta.servlet-api 6.0 removed the deprecated isRequestedSessionIdFromUrl
+ * and getRealPath methods that 5.0 still declares abstract).
+ */
+ def underlying: AnyRef
}
object HttpServletRequestLike {
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JakartaHttpServletRequest.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JakartaHttpServletRequest.scala
index 9219f294a13..17c96894d53 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JakartaHttpServletRequest.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JakartaHttpServletRequest.scala
@@ -17,155 +17,21 @@
package org.apache.spark.ui
-import java.io.BufferedReader
-import java.security.Principal
import java.util
-import java.util.Locale
-import jakarta.servlet._
import jakarta.servlet.http._
+/**
+ * Adapts a [[jakarta.servlet.http.HttpServletRequest]] to the engine UI code via
+ * [[HttpServletRequestLike]]. Only the request parameters used by the engine UI
+ * are forwarded; Spark private UI helpers that need a real servlet request are
+ * given [[HttpServletRequestLike.underlying]] instead.
+ */
class JakartaHttpServletRequest(req: HttpServletRequest)
- extends HttpServletRequest with HttpServletRequestLike {
- override def getAuthType: String = req.getAuthType
-
- override def getCookies: Array[Cookie] = req.getCookies
-
- override def getDateHeader(s: String): Long = req.getDateHeader(s)
-
- override def getHeader(s: String): String = req.getHeader(s)
-
- override def getHeaders(s: String): util.Enumeration[String] = req.getHeaders(s)
-
- override def getHeaderNames: util.Enumeration[String] = req.getHeaderNames
-
- override def getIntHeader(s: String): Int = req.getIntHeader(s)
-
- override def getMethod: String = req.getMethod
-
- override def getPathInfo: String = req.getPathInfo
-
- override def getPathTranslated: String = req.getPathTranslated
-
- override def getContextPath: String = req.getContextPath
-
- override def getQueryString: String = req.getQueryString
-
- override def getRemoteUser: String = req.getRemoteUser
-
- override def isUserInRole(s: String): Boolean = req.isUserInRole(s)
-
- override def getUserPrincipal: Principal = req.getUserPrincipal
-
- override def getRequestedSessionId: String = req.getRequestedSessionId
-
- override def getRequestURI: String = req.getRequestURI
-
- override def getRequestURL: StringBuffer = req.getRequestURL
-
- override def getServletPath: String = req.getServletPath
-
- override def getSession(b: Boolean): HttpSession = req.getSession(b)
-
- override def getSession: HttpSession = req.getSession
-
- override def changeSessionId(): String = req.changeSessionId()
-
- override def isRequestedSessionIdValid: Boolean = req.isRequestedSessionIdValid
-
- override def isRequestedSessionIdFromCookie: Boolean = req.isRequestedSessionIdFromCookie
-
- override def isRequestedSessionIdFromURL: Boolean = req.isRequestedSessionIdFromURL
-
- override def isRequestedSessionIdFromUrl: Boolean = req.isRequestedSessionIdFromUrl
-
- override def authenticate(httpServletResponse: HttpServletResponse): Boolean =
- req.authenticate(httpServletResponse)
-
- override def login(s: String, s1: String): Unit = req.login(s, s1)
-
- override def logout(): Unit = req.logout()
-
- override def getParts: util.Collection[Part] = req.getParts
-
- override def getPart(s: String): Part = req.getPart(s)
-
- override def upgrade[T <: HttpUpgradeHandler](aClass: Class[T]): T = req.upgrade(aClass)
-
- override def getAttribute(s: String): AnyRef = req.getAttribute(s)
-
- override def getAttributeNames: util.Enumeration[String] = req.getAttributeNames
-
- override def getCharacterEncoding: String = req.getCharacterEncoding
-
- override def setCharacterEncoding(s: String): Unit = req.setCharacterEncoding(s)
-
- override def getContentLength: Int = req.getContentLength
-
- override def getContentLengthLong: Long = req.getContentLengthLong
-
- override def getContentType: String = req.getContentType
-
- override def getInputStream: ServletInputStream = req.getInputStream
-
- override def getParameter(s: String): String = req.getParameter(s)
-
- override def getParameterNames: util.Enumeration[String] = req.getParameterNames
-
- override def getParameterValues(s: String): Array[String] = req.getParameterValues(s)
+ extends HttpServletRequestLike {
+ override def getParameter(name: String): String = req.getParameter(name)
override def getParameterMap: util.Map[String, Array[String]] = req.getParameterMap
- override def getProtocol: String = req.getProtocol
-
- override def getScheme: String = req.getScheme
-
- override def getServerName: String = req.getServerName
-
- override def getServerPort: Int = req.getServerPort
-
- override def getReader: BufferedReader = req.getReader
-
- override def getRemoteAddr: String = req.getRemoteAddr
-
- override def getRemoteHost: String = req.getRemoteHost
-
- override def setAttribute(s: String, o: Any): Unit = req.setAttribute(s, o)
-
- override def removeAttribute(s: String): Unit = req.removeAttribute(s)
-
- override def getLocale: Locale = req.getLocale
-
- override def getLocales: util.Enumeration[Locale] = req.getLocales
-
- override def isSecure: Boolean = req.isSecure
-
- override def getRequestDispatcher(s: String): RequestDispatcher = req.getRequestDispatcher(s)
-
- override def getRealPath(s: String): String = req.getRealPath(s)
-
- override def getRemotePort: Int = req.getRemotePort
-
- override def getLocalName: String = req.getLocalName
-
- override def getLocalAddr: String = req.getLocalAddr
-
- override def getLocalPort: Int = req.getLocalPort
-
- override def getServletContext: ServletContext = req.getServletContext
-
- override def startAsync(): AsyncContext = req.startAsync()
-
- override def startAsync(
- servletRequest: ServletRequest,
- servletResponse: ServletResponse): AsyncContext =
- req.startAsync(servletRequest, servletResponse)
-
- override def isAsyncStarted: Boolean = req.isAsyncStarted
-
- override def isAsyncSupported: Boolean = req.isAsyncSupported
-
- override def getAsyncContext: AsyncContext = req.getAsyncContext
-
- override def getDispatcherType: DispatcherType = req.getDispatcherType
+ override def underlying: AnyRef = req
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JavaxHttpServletRequest.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JavaxHttpServletRequest.scala
index d77948001bb..6d71f8d82fb 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JavaxHttpServletRequest.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/JavaxHttpServletRequest.scala
@@ -17,154 +17,20 @@
package org.apache.spark.ui
-import java.io.BufferedReader
-import java.security.Principal
import java.util
-import java.util.Locale
-import javax.servlet._
import javax.servlet.http._
+/**
+ * Adapts a javax.servlet.http.HttpServletRequest to the engine UI code via
+ * [[HttpServletRequestLike]], used by Spark 3.x which still depends on the
+ * javax servlet API. See [[JakartaHttpServletRequest]] for the rationale of
+ * only forwarding the parameters the engine UI uses.
+ */
class JavaxHttpServletRequest(req: HttpServletRequest)
- extends HttpServletRequest with HttpServletRequestLike {
- override def getAuthType: String = req.getAuthType
-
- override def getCookies: Array[Cookie] = req.getCookies
-
- override def getDateHeader(s: String): Long = req.getDateHeader(s)
-
- override def getHeader(s: String): String = req.getHeader(s)
-
- override def getHeaders(s: String): util.Enumeration[String] = req.getHeaders(s)
-
- override def getHeaderNames: util.Enumeration[String] = req.getHeaderNames
-
- override def getIntHeader(s: String): Int = req.getIntHeader(s)
-
- override def getMethod: String = req.getMethod
-
- override def getPathInfo: String = req.getPathInfo
-
- override def getPathTranslated: String = req.getPathTranslated
-
- override def getContextPath: String = req.getContextPath
-
- override def getQueryString: String = req.getQueryString
-
- override def getRemoteUser: String = req.getRemoteUser
-
- override def isUserInRole(s: String): Boolean = req.isUserInRole(s)
-
- override def getUserPrincipal: Principal = req.getUserPrincipal
-
- override def getRequestedSessionId: String = req.getRequestedSessionId
-
- override def getRequestURI: String = req.getRequestURI
-
- override def getRequestURL: StringBuffer = req.getRequestURL
-
- override def getServletPath: String = req.getServletPath
-
- override def getSession(b: Boolean): HttpSession = req.getSession(b)
-
- override def getSession: HttpSession = req.getSession
-
- override def changeSessionId(): String = req.changeSessionId()
-
- override def isRequestedSessionIdValid: Boolean = req.isRequestedSessionIdValid
-
- override def isRequestedSessionIdFromCookie: Boolean = req.isRequestedSessionIdFromCookie
-
- override def isRequestedSessionIdFromURL: Boolean = req.isRequestedSessionIdFromURL
-
- override def isRequestedSessionIdFromUrl: Boolean = req.isRequestedSessionIdFromUrl
-
- override def authenticate(httpServletResponse: HttpServletResponse): Boolean =
- req.authenticate(httpServletResponse)
-
- override def login(s: String, s1: String): Unit = req.login(s, s1)
-
- override def logout(): Unit = req.logout()
-
- override def getParts: util.Collection[Part] = req.getParts
-
- override def getPart(s: String): Part = req.getPart(s)
-
- override def upgrade[T <: HttpUpgradeHandler](aClass: Class[T]): T = req.upgrade(aClass)
-
- override def getAttribute(s: String): AnyRef = req.getAttribute(s)
-
- override def getAttributeNames: util.Enumeration[String] = req.getAttributeNames
-
- override def getCharacterEncoding: String = req.getCharacterEncoding
-
- override def setCharacterEncoding(s: String): Unit = req.setCharacterEncoding(s)
-
- override def getContentLength: Int = req.getContentLength
-
- override def getContentLengthLong: Long = req.getContentLengthLong
-
- override def getContentType: String = req.getContentType
-
- override def getInputStream: ServletInputStream = req.getInputStream
-
- override def getParameter(s: String): String = req.getParameter(s)
-
- override def getParameterNames: util.Enumeration[String] = req.getParameterNames
-
- override def getParameterValues(s: String): Array[String] = req.getParameterValues(s)
+ extends HttpServletRequestLike {
+ override def getParameter(name: String): String = req.getParameter(name)
override def getParameterMap: util.Map[String, Array[String]] = req.getParameterMap
- override def getProtocol: String = req.getProtocol
-
- override def getScheme: String = req.getScheme
-
- override def getServerName: String = req.getServerName
-
- override def getServerPort: Int = req.getServerPort
-
- override def getReader: BufferedReader = req.getReader
-
- override def getRemoteAddr: String = req.getRemoteAddr
-
- override def getRemoteHost: String = req.getRemoteHost
-
- override def setAttribute(s: String, o: Any): Unit = req.setAttribute(s, o)
-
- override def removeAttribute(s: String): Unit = req.removeAttribute(s)
-
- override def getLocale: Locale = req.getLocale
-
- override def getLocales: util.Enumeration[Locale] = req.getLocales
-
- override def isSecure: Boolean = req.isSecure
-
- override def getRequestDispatcher(s: String): RequestDispatcher = req.getRequestDispatcher(s)
-
- override def getRealPath(s: String): String = req.getRealPath(s)
-
- override def getRemotePort: Int = req.getRemotePort
-
- override def getLocalName: String = req.getLocalName
-
- override def getLocalAddr: String = req.getLocalAddr
-
- override def getLocalPort: Int = req.getLocalPort
-
- override def getServletContext: ServletContext = req.getServletContext
-
- override def startAsync(): AsyncContext = req.startAsync()
-
- override def startAsync(
- servletRequest: ServletRequest,
- servletResponse: ServletResponse): AsyncContext =
- req.startAsync(servletRequest, servletResponse)
-
- override def isAsyncStarted: Boolean = req.isAsyncStarted
-
- override def isAsyncSupported: Boolean = req.isAsyncSupported
-
- override def getAsyncContext: AsyncContext = req.getAsyncContext
-
- override def getDispatcherType: DispatcherType = req.getDispatcherType
+ override def underlying: AnyRef = req
}
diff --git a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/SparkUIUtils.scala b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/SparkUIUtils.scala
index 4862372a19e..fec10606ce0 100644
--- a/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/SparkUIUtils.scala
+++ b/externals/kyuubi-spark-sql-engine/src/main/scala/org/apache/spark/ui/SparkUIUtils.scala
@@ -41,9 +41,21 @@ object SparkUIUtils {
activeTab: SparkUITab,
helpText: Option[String] = None,
showVisualization: JBoolean = false,
- useDataTables: JBoolean = false): Seq[Node] = {
+ useDataTables: JBoolean = false,
+ useTimeline: JBoolean = false): Seq[Node] = {
val headerSparkPageMethod = if (SPARK_ENGINE_RUNTIME_VERSION >= "4.0") {
DynMethods.builder("headerSparkPage")
+ // SPARK-56354 (4.2.0) added the useTimeline flag as an 8th parameter
+ .impl(
+ UIUtils.getClass,
+ classOf[jakarta.servlet.http.HttpServletRequest],
+ classOf[String],
+ classOf[() => Seq[Node]],
+ classOf[SparkUITab],
+ classOf[Option[String]],
+ classOf[Boolean],
+ classOf[Boolean],
+ classOf[Boolean])
.impl(
UIUtils.getClass,
classOf[jakarta.servlet.http.HttpServletRequest],
@@ -68,13 +80,14 @@ object SparkUIUtils {
.buildChecked(UIUtils)
}
headerSparkPageMethod.invoke[Seq[Node]](
- request,
+ request.underlying,
title,
() => content,
activeTab,
helpText,
showVisualization,
- useDataTables)
+ useDataTables,
+ useTimeline)
}
def prependBaseUri(
@@ -98,6 +111,6 @@ object SparkUIUtils {
classOf[String])
.buildChecked(UIUtils)
}
- prependBaseUriMethod.invoke[String](request, basePath, resource)
+ prependBaseUriMethod.invoke[String](request.underlying, basePath, resource)
}
}
diff --git a/kyuubi-server/src/test/scala/org/apache/kyuubi/operation/KyuubiOperationPerConnectionSuite.scala b/kyuubi-server/src/test/scala/org/apache/kyuubi/operation/KyuubiOperationPerConnectionSuite.scala
index 12f8520b6d7..64d3a25695a 100644
--- a/kyuubi-server/src/test/scala/org/apache/kyuubi/operation/KyuubiOperationPerConnectionSuite.scala
+++ b/kyuubi-server/src/test/scala/org/apache/kyuubi/operation/KyuubiOperationPerConnectionSuite.scala
@@ -83,7 +83,8 @@ class KyuubiOperationPerConnectionSuite extends WithKyuubiServer with HiveJDBCTe
assert(errMsg.contains("Caused by: java.net.SocketException: Connection reset") ||
errMsg.contains(s"Socket for ${SessionHandle(handle)} is closed") ||
errMsg.contains("Socket is closed by peer") ||
- errMsg.contains("SparkContext was shut down"))
+ errMsg.contains("SparkContext was shut down") ||
+ errMsg.contains("DAGScheduler.cleanupQueryJobs"))
}
}
diff --git a/pom.xml b/pom.xml
index 3876d22306c..0346467ffd3 100644
--- a/pom.xml
+++ b/pom.xml
@@ -2110,6 +2110,34 @@
+
+ spark-4.2
+
+ extensions/spark/kyuubi-extension-spark-4-2
+ extensions/spark/kyuubi-spark-connector-hive
+
+
+ 17
+ 17
+ 4.2.0
+
+ 4.1
+
+ 6.0.0
+ 4.13.1
+ 4.3.1
+ delta-spark_${spark.binary.version}_${scala.binary.version}
+ 1.2.0
+
+ 1.4.2
+ paimon-spark-4.0_${scala.binary.version}
+ org.scalatest.tags.Slow,org.apache.kyuubi.tags.DeltaTest,org.apache.kyuubi.tags.IcebergTest,org.apache.kyuubi.tags.PaimonTest,org.apache.kyuubi.tags.HudiTest
+ spark-${spark.version}-bin-hadoop3.tgz
+
+
+
spark-master