This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git
The following commit(s) were added to refs/heads/master by this push:
new 1f3040a879 [spark] Repartition limited lateral vector search input
(#8800)
1f3040a879 is described below
commit 1f3040a8792d65899e8e619ef12b6519a0088600
Author: shyjsarah <[email protected]>
AuthorDate: Sun Jul 26 15:58:32 2026 +0800
[spark] Repartition limited lateral vector search input (#8800)
---
docs/generated/spark_connector_configuration.html | 6 +
.../spark/sql/paimon/shims/MinorVersionShim.scala | 4 +
.../spark/sql/paimon/shims/MinorVersionShim.scala | 7 +
.../apache/spark/sql/paimon/shims/Spark4Shim.scala | 9 +
.../apache/paimon/spark/SparkConnectorOptions.java | 8 +
.../RepartitionLateralVectorSearchInput.scala | 91 ++++++
.../paimon/spark/execution/PaimonStrategy.scala | 31 +-
.../extensions/PaimonSparkSessionExtensions.scala | 3 +-
.../apache/spark/sql/paimon/shims/SparkShim.scala | 3 +
.../spark/sql/TableValuedFunctionsTest.scala | 364 ++++++++++++++++++++-
.../spark/sql/paimon/shims/MinorVersionShim.scala | 7 +
.../apache/spark/sql/paimon/shims/Spark3Shim.scala | 6 +
.../apache/spark/sql/paimon/shims/Spark4Shim.scala | 9 +
13 files changed, 544 insertions(+), 4 deletions(-)
diff --git a/docs/generated/spark_connector_configuration.html
b/docs/generated/spark_connector_configuration.html
index 6937d85757..cd95fd5fd4 100644
--- a/docs/generated/spark_connector_configuration.html
+++ b/docs/generated/spark_connector_configuration.html
@@ -86,6 +86,12 @@ under the License.
<td>Boolean</td>
<td>Whether to adjust the target split size based on pruned
(projected) columns. If enabled, split size estimation uses only the columns
actually being read.</td>
</tr>
+ <tr>
+ <td><h5>vector-search.lateral-join.parallelism</h5></td>
+ <td style="word-wrap: break-word;">16</td>
+ <td>Integer</td>
+ <td>Parallelism used to repartition a single-partition LIMIT input
before executing a lateral vector search.</td>
+ </tr>
<tr>
<td><h5>write.data-evolution.update-conflict-retry.max-attempts</h5></td>
<td style="word-wrap: break-word;">20</td>
diff --git
a/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
b/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
index dffd53beb9..aed66477bc 100644
---
a/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
+++
b/paimon-spark/paimon-spark-3.2/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
@@ -20,6 +20,7 @@ package org.apache.spark.sql.paimon.shims
import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression}
import org.apache.spark.sql.catalyst.plans.logical.{CTERelationRef,
LogicalPlan, MergeAction, MergeIntoTable}
+import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution,
Distribution}
object MinorVersionShim {
@@ -29,6 +30,9 @@ object MinorVersionShim {
output: Seq[Attribute],
isStreaming: Boolean): CTERelationRef = CTERelationRef(cteId, resolved,
output)
+ def createClusteredDistribution(expressions: Seq[Expression], numPartitions:
Int): Distribution =
+ ClusteredDistribution(expressions, Some(numPartitions))
+
def createMergeIntoTable(
targetTable: LogicalPlan,
sourceTable: LogicalPlan,
diff --git
a/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
b/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
index ae0d96e8f8..41e5ed4c12 100644
---
a/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
+++
b/paimon-spark/paimon-spark-3.3/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
@@ -20,6 +20,7 @@ package org.apache.spark.sql.paimon.shims
import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression}
import org.apache.spark.sql.catalyst.plans.logical.{CTERelationRef,
LogicalPlan, MergeAction, MergeIntoTable}
+import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution,
Distribution}
object MinorVersionShim {
@@ -29,6 +30,12 @@ object MinorVersionShim {
output: Seq[Attribute],
isStreaming: Boolean): CTERelationRef = CTERelationRef(cteId, resolved,
output)
+ def createClusteredDistribution(expressions: Seq[Expression], numPartitions:
Int): Distribution =
+ ClusteredDistribution(
+ expressions,
+ requireAllClusterKeys = false,
+ requiredNumPartitions = Some(numPartitions))
+
def createMergeIntoTable(
targetTable: LogicalPlan,
sourceTable: LogicalPlan,
diff --git
a/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
b/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
index e3046f564f..854dda9fa2 100644
---
a/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
+++
b/paimon-spark/paimon-spark-4.0/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
@@ -39,6 +39,7 @@ import
org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression
import org.apache.spark.sql.catalyst.parser.ParserInterface
import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Assignment,
ColumnDefinition, CTERelationRef, InsertAction, LogicalPlan, MergeAction,
MergeIntoTable, MergeRows, SubqueryAlias, TableSpec, UnresolvedWith,
UpdateAction}
import org.apache.spark.sql.catalyst.plans.logical.MergeRows.Keep
+import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution,
Distribution}
import org.apache.spark.sql.catalyst.rules.Rule
import org.apache.spark.sql.catalyst.util.{ArrayData, GeneratedColumn,
IdentityColumn, ResolveDefaultColumns}
import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Column,
Identifier, StagingTableCatalog, Table, TableCatalog}
@@ -243,6 +244,14 @@ class Spark4Shim extends SparkShim {
CTERelationRef(cteId, resolved, output.toSeq, isStreaming)
}
+ override def createClusteredDistribution(
+ expressions: Seq[Expression],
+ numPartitions: Int): Distribution =
+ ClusteredDistribution(
+ expressions,
+ requireAllClusterKeys = false,
+ requiredNumPartitions = Some(numPartitions))
+
override def supportsHashAggregate(
aggregateBufferAttributes: Seq[Attribute],
groupingExpression: Seq[Expression]): Boolean = {
diff --git
a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java
b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java
index 6e0bbf9c58..2f315b8df0 100644
---
a/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java
+++
b/paimon-spark/paimon-spark-common/src/main/java/org/apache/paimon/spark/SparkConnectorOptions.java
@@ -39,6 +39,14 @@ public class SparkConnectorOptions {
.withDescription(
"If true, map Paimon TIMESTAMP to Spark TIMESTAMP
instead of TIMESTAMP_NTZ.");
+ public static final ConfigOption<Integer>
VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM =
+ key("vector-search.lateral-join.parallelism")
+ .intType()
+ .defaultValue(16)
+ .withDescription(
+ "Parallelism used to repartition a
single-partition LIMIT input before "
+ + "executing a lateral vector search.");
+
public static final ConfigOption<Boolean> MERGE_SCHEMA =
key("write.merge-schema")
.booleanType()
diff --git
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala
new file mode 100644
index 0000000000..c44d47c16f
--- /dev/null
+++
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/catalyst/optimizer/RepartitionLateralVectorSearchInput.scala
@@ -0,0 +1,91 @@
+/*
+ * 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.paimon.spark.catalyst.optimizer
+
+import org.apache.paimon.spark.SparkConnectorOptions
+import org.apache.paimon.spark.catalyst.plans.logical.LateralVectorSearch
+import org.apache.paimon.spark.util.OptionUtils
+
+import org.apache.spark.sql.catalyst.plans.logical.{BROADCAST, CTERelationRef,
GlobalLimit, HintInfo, Join, LogicalPlan, Repartition, RepartitionOperation,
ResolvedHint, UnaryNode, WithCTE}
+import org.apache.spark.sql.catalyst.rules.Rule
+
+/** Restores parallelism lost by a global limit before executing a lateral
vector search. */
+object RepartitionLateralVectorSearchInput extends Rule[LogicalPlan] {
+
+ override def apply(plan: LogicalPlan): LogicalPlan = {
+ val cteDefinitions = plan
+ .collect { case withCTE: WithCTE => withCTE.cteDefs }
+ .flatten
+ .map(definition => definition.id -> definition.child)
+ .toMap
+
+ plan.transformUp {
+ case lateralVectorSearch: LateralVectorSearch
+ if hasUnrepartitionedGlobalLimit(lateralVectorSearch.left,
cteDefinitions, Set.empty) =>
+ lateralVectorSearch.copy(
+ left = Repartition(parallelism, shuffle = true,
lateralVectorSearch.left))
+ }
+ }
+
+ private[spark] def parallelism: Int = {
+ val value =
+ OptionUtils
+
.getOptionString(SparkConnectorOptions.VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM)
+ .toInt
+ require(
+ value > 0,
+
s"spark.paimon.${SparkConnectorOptions.VECTOR_SEARCH_LATERAL_JOIN_PARALLELISM.key()}
" +
+ s"must be positive, but got $value")
+ value
+ }
+
+ private def hasUnrepartitionedGlobalLimit(
+ plan: LogicalPlan,
+ cteDefinitions: Map[Long, LogicalPlan],
+ visitedCTEs: Set[Long]): Boolean = plan match {
+ case repartition: RepartitionOperation if repartition.shuffle => false
+ case repartition: RepartitionOperation =>
+ hasUnrepartitionedGlobalLimit(repartition.child, cteDefinitions,
visitedCTEs)
+ case _: GlobalLimit => true
+ case reference: CTERelationRef if !visitedCTEs.contains(reference.cteId) =>
+ cteDefinitions
+ .get(reference.cteId)
+ .exists(hasUnrepartitionedGlobalLimit(_, cteDefinitions, visitedCTEs +
reference.cteId))
+ case join: Join
+ if hasBroadcastHint(join.hint.rightHint) ||
hasResolvedBroadcastHint(join.right) =>
+ hasUnrepartitionedGlobalLimit(join.left, cteDefinitions, visitedCTEs)
+ case join: Join
+ if hasBroadcastHint(join.hint.leftHint) ||
hasResolvedBroadcastHint(join.left) =>
+ hasUnrepartitionedGlobalLimit(join.right, cteDefinitions, visitedCTEs)
+ case unary: UnaryNode =>
+ hasUnrepartitionedGlobalLimit(unary.child, cteDefinitions, visitedCTEs)
+ case _ => false
+ }
+
+ private def hasBroadcastHint(hint: Option[HintInfo]): Boolean = {
+ hint.flatMap(_.strategy).contains(BROADCAST)
+ }
+
+ private def hasResolvedBroadcastHint(plan: LogicalPlan): Boolean = {
+ plan.exists {
+ case hint: ResolvedHint => hasBroadcastHint(Some(hint.hints))
+ case _ => false
+ }
+ }
+}
diff --git
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala
index 2686ef7d76..77c50acf70 100644
---
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala
+++
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/execution/PaimonStrategy.scala
@@ -27,6 +27,7 @@ import org.apache.paimon.predicate.{Predicate,
PredicateBuilder}
import org.apache.paimon.spark.{PaimonRecordReaderIterator, SparkCatalog,
SparkGenericCatalog, SparkTable, SparkUtils}
import org.apache.paimon.spark.catalog.{SparkBaseCatalog, SupportView}
import org.apache.paimon.spark.catalyst.analysis.ResolvedPaimonView
+import
org.apache.paimon.spark.catalyst.optimizer.RepartitionLateralVectorSearchInput
import
org.apache.paimon.spark.catalyst.plans.logical.{CopyIntoLocationCommand,
CopyIntoLocationSource, CopyIntoTableCommand, CreateOrReplaceTagCommand,
CreatePaimonView, DeleteTagCommand, DropPaimonView, LateralVectorSearch,
PaimonCallCommand, PaimonDropPartitions, PaimonTableValuedFunctions,
RenameTagCommand, ResolvedIdentifier, ShowPaimonViews, ShowTagsCommand,
TruncatePaimonTableWithFilter}
import org.apache.paimon.spark.data.SparkInternalRow
import org.apache.paimon.spark.format.PaimonFormatTable
@@ -43,11 +44,15 @@ import org.apache.spark.sql.SparkSession
import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.catalyst.analysis.{ResolvedNamespace,
ResolvedTable}
import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet,
Expression, GenericInternalRow, JoinedRow, PredicateHelper, UnsafeProjection}
+import org.apache.spark.sql.catalyst.optimizer.BuildRight
import org.apache.spark.sql.catalyst.plans.logical.{AddPartitions,
CreateTableAsSelect, DescribeRelation, DropPartitions, LogicalPlan,
RepairTable, ReplaceTable, ReplaceTableAsSelect, ShowCreateTable}
+import org.apache.spark.sql.catalyst.plans.physical.{Distribution,
UnspecifiedDistribution}
import org.apache.spark.sql.catalyst.util.ArrayData
import org.apache.spark.sql.connector.catalog.{Identifier,
PaimonLookupCatalog, TableCatalog}
-import org.apache.spark.sql.execution.{PaimonDescribeTableExec, SparkPlan,
SparkStrategy}
+import org.apache.spark.sql.execution.{GlobalLimitExec,
PaimonDescribeTableExec, SparkPlan, SparkStrategy, UnaryExecNode}
import org.apache.spark.sql.execution.datasources.v2.{DataSourceV2Implicits,
DataSourceV2Relation}
+import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike
+import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec,
BroadcastNestedLoopJoinExec}
import org.apache.spark.sql.execution.shim.{PaimonCreateTableAsSelectStrategy,
PaimonReplaceTableAsSelectStrategy, PaimonReplaceTableStrategy}
import org.apache.spark.sql.paimon.shims.SparkShimLoader
@@ -309,6 +314,30 @@ case class LateralVectorSearchExec(
override def output: Seq[Attribute] = child.output ++ projectOutput
+ // Statistics-based broadcast selection is only known after physical
planning. Request a
+ // distribution here so EnsureRequirements can restore the streamed LIMIT
side's parallelism.
+ override def requiredChildDistribution: Seq[Distribution] = {
+ if (hasUnrepartitionedGlobalLimit(child)) {
+ Seq(
+ SparkShimLoader.shim.createClusteredDistribution(
+ child.output,
+ RepartitionLateralVectorSearchInput.parallelism))
+ } else {
+ Seq(UnspecifiedDistribution)
+ }
+ }
+
+ private def hasUnrepartitionedGlobalLimit(plan: SparkPlan): Boolean = plan
match {
+ case _: ShuffleExchangeLike => false
+ case _: GlobalLimitExec => true
+ case join: BroadcastHashJoinExec =>
+ hasUnrepartitionedGlobalLimit(if (join.buildSide == BuildRight)
join.left else join.right)
+ case join: BroadcastNestedLoopJoinExec =>
+ hasUnrepartitionedGlobalLimit(if (join.buildSide == BuildRight)
join.left else join.right)
+ case unary: UnaryExecNode => hasUnrepartitionedGlobalLimit(unary.child)
+ case _ => false
+ }
+
@transient override lazy val producedAttributes: AttributeSet = {
AttributeSet(vectorSearchOutput ++ output.filterNot(attr =>
inputSet.contains(attr)))
}
diff --git
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala
index 388889bbea..69b3b5f28b 100644
---
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala
+++
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/paimon/spark/extensions/PaimonSparkSessionExtensions.scala
@@ -19,7 +19,7 @@
package org.apache.paimon.spark.extensions
import org.apache.paimon.spark.catalyst.analysis.{PaimonAnalysis,
PaimonDeleteTable, PaimonFunctionResolver, PaimonIncompatibleResolutionRules,
PaimonMergeInto, PaimonPostHocResolutionRules, PaimonProcedureResolver,
PaimonUpdateTable, PaimonViewResolver, ReplacePaimonFunctions,
RewriteUpsertTable}
-import
org.apache.paimon.spark.catalyst.optimizer.{MergePaimonScalarSubqueries,
OptimizeMetadataOnlyDeleteFromPaimonTable, PushDownLateralVectorSearchFilter}
+import
org.apache.paimon.spark.catalyst.optimizer.{MergePaimonScalarSubqueries,
OptimizeMetadataOnlyDeleteFromPaimonTable, PushDownLateralVectorSearchFilter,
RepartitionLateralVectorSearchInput}
import
org.apache.paimon.spark.catalyst.plans.logical.PaimonTableValuedFunctions
import org.apache.paimon.spark.commands.BucketExpression
import org.apache.paimon.spark.execution.{OldCompatibleStrategy,
PaimonStrategy}
@@ -104,6 +104,7 @@ class PaimonSparkSessionExtensions extends
(SparkSessionExtensions => Unit) {
// TODO: Enable MAP selected-key pushdown after core reader supports
// __PAIMON_MAP_SELECTED_KEYS read type.
extensions.injectOptimizerRule(_ => MergePaimonScalarSubqueries)
+ extensions.injectOptimizerRule(_ => RepartitionLateralVectorSearchInput)
extensions.injectOptimizerRule(_ => PushDownLateralVectorSearchFilter)
// planner extensions
diff --git
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala
index 85c923325d..b65079d067 100644
---
a/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala
+++
b/paimon-spark/paimon-spark-common/src/main/scala/org/apache/spark/sql/paimon/shims/SparkShim.scala
@@ -33,6 +33,7 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute,
Expression}
import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression
import org.apache.spark.sql.catalyst.parser.ParserInterface
import org.apache.spark.sql.catalyst.plans.logical.{Assignment,
CTERelationRef, InsertAction, LogicalPlan, MergeAction, MergeIntoTable,
SubqueryAlias, TableSpec, UnresolvedWith, UpdateAction}
+import org.apache.spark.sql.catalyst.plans.physical.Distribution
import org.apache.spark.sql.catalyst.rules.Rule
import org.apache.spark.sql.catalyst.util.ArrayData
import org.apache.spark.sql.connector.catalog.{Column, Identifier,
StagingTableCatalog, Table, TableCatalog}
@@ -159,6 +160,8 @@ trait SparkShim {
output: Seq[Attribute],
isStreaming: Boolean): CTERelationRef
+ def createClusteredDistribution(expressions: Seq[Expression], numPartitions:
Int): Distribution
+
def supportsHashAggregate(
aggregateBufferAttributes: Seq[Attribute],
groupingExpression: Seq[Expression]): Boolean
diff --git
a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala
b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala
index a44711d5b9..e46a00358d 100644
---
a/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala
+++
b/paimon-spark/paimon-spark-ut/src/test/scala/org/apache/paimon/spark/sql/TableValuedFunctionsTest.scala
@@ -22,15 +22,19 @@ import org.apache.paimon.data.{BinaryString, GenericRow,
Timestamp}
import org.apache.paimon.manifest.ManifestCommittable
import org.apache.paimon.spark.PaimonHiveTestBase
import org.apache.paimon.spark.catalyst.plans.logical.{LateralVectorSearch,
PaimonTableValuedFunctions}
+import org.apache.paimon.spark.execution.LateralVectorSearchExec
import org.apache.paimon.utils.DateTimeUtils
import org.apache.spark.sql.{DataFrame, Row}
-import org.apache.spark.sql.catalyst.plans.logical.Filter
+import org.apache.spark.sql.catalyst.optimizer.{BuildLeft, BuildRight}
+import org.apache.spark.sql.catalyst.plans.logical.{Filter, Repartition}
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
+import org.apache.spark.sql.execution.joins.BroadcastHashJoinExec
import java.time.LocalDateTime
import java.util.Collections
-class TableValuedFunctionsTest extends PaimonHiveTestBase {
+class TableValuedFunctionsTest extends PaimonHiveTestBase with
AdaptiveSparkPlanHelper {
test("parse positive limit rejects overflowing long") {
val longValue: Long = 4294967297L
@@ -42,6 +46,348 @@ class TableValuedFunctionsTest extends PaimonHiveTestBase {
assert(error.getMessage.contains("Limit must be no greater than"))
}
+ test("lateral vector search repartitions global limit input") {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val optimizedPlan = spark
+ .sql("""
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM (
+ | SELECT gid, embs
+ | FROM vector_search_source
+ | WHERE dt = '20260629'
+ | LIMIT 1000
+ |) AS q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val lateralVectorSearch = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .getOrElse(fail(optimizedPlan.toString))
+ val repartitions = lateralVectorSearch.left.collect {
+ case repartition: Repartition => repartition
+ }
+
+ assert(repartitions.size == 1, optimizedPlan.toString)
+ assert(repartitions.head.shuffle, optimizedPlan.toString)
+ assert(repartitions.head.numPartitions == 16, optimizedPlan.toString)
+ }
+ }
+
+ test("lateral vector search uses configured repartition parallelism") {
+ val parallelismKey = "spark.paimon.vector-search.lateral-join.parallelism"
+ spark.conf.set(parallelismKey, "4")
+ try {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val optimizedPlan = spark
+ .sql("""
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM (
+ | SELECT gid, embs
+ | FROM vector_search_source
+ | LIMIT 1000
+ |) AS q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs,
3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val repartition = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .flatMap(_.left.collectFirst { case repartition: Repartition =>
repartition })
+ .getOrElse(fail(optimizedPlan.toString))
+
+ assert(repartition.numPartitions == 4, optimizedPlan.toString)
+ }
+ } finally {
+ spark.conf.unset(parallelismKey)
+ }
+ }
+
+ test("lateral vector search repartitions above a limited repartition") {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val optimizedPlan = spark
+ .sql("""
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM (
+ | SELECT /*+ REPARTITION(4) */ gid, embs
+ | FROM vector_search_source
+ | LIMIT 1000
+ |) AS q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val lateralVectorSearch = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .getOrElse(fail(optimizedPlan.toString))
+ val repartitions = lateralVectorSearch.left.collect {
+ case repartition: Repartition => repartition
+ }
+
+ assert(lateralVectorSearch.left.isInstanceOf[Repartition],
optimizedPlan.toString)
+ assert(repartitions.map(_.numPartitions) == Seq(16, 4),
optimizedPlan.toString)
+ }
+ }
+
+ test("lateral vector search preserves repartition above limit") {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val optimizedPlan = spark
+ .sql("""
+ |WITH q_limit AS (
+ | SELECT gid, embs
+ | FROM vector_search_source
+ | LIMIT 1000
+ |),
+ |q AS (
+ | SELECT /*+ REPARTITION(4) */ gid, embs
+ | FROM q_limit
+ |)
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val lateralVectorSearch = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .getOrElse(fail(optimizedPlan.toString))
+ val repartitions = lateralVectorSearch.left.collect {
+ case repartition: Repartition => repartition
+ }
+
+ assert(repartitions.map(_.numPartitions) == Seq(4),
optimizedPlan.toString)
+ }
+ }
+
+ test("lateral vector search repartitions CTE limited input") {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val optimizedPlan = spark
+ .sql("""
+ |WITH q_limit AS (
+ | SELECT gid, embs
+ | FROM vector_search_source
+ | LIMIT 1000
+ |),
+ |q AS (
+ | SELECT gid, embs
+ | FROM q_limit
+ |)
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val lateralVectorSearch = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .getOrElse(fail(optimizedPlan.toString))
+
+ assert(lateralVectorSearch.left.isInstanceOf[Repartition],
optimizedPlan.toString)
+ val repartition = lateralVectorSearch.left.asInstanceOf[Repartition]
+ assert(repartition.shuffle, optimizedPlan.toString)
+ assert(repartition.numPartitions == 16, optimizedPlan.toString)
+ }
+ }
+
+ test("lateral vector search repartitions broadcast join streamed limited
input") {
+ withTable("vector_search_source", "vector_search_dimension") {
+ createVectorSearchSource()
+ spark.sql("CREATE TABLE vector_search_dimension (gid BIGINT) USING
paimon")
+
+ val optimizedPlan = spark
+ .sql("""
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM (
+ | SELECT /*+ BROADCAST(d) */ s.gid, s.embs
+ | FROM (
+ | SELECT gid, embs
+ | FROM vector_search_source
+ | LIMIT 1000
+ | ) s
+ | JOIN vector_search_dimension d
+ | ON s.gid = d.gid
+ |) q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val lateralVectorSearch = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .getOrElse(fail(optimizedPlan.toString))
+
+ assert(lateralVectorSearch.left.isInstanceOf[Repartition],
optimizedPlan.toString)
+ val repartition = lateralVectorSearch.left.asInstanceOf[Repartition]
+ assert(repartition.shuffle, optimizedPlan.toString)
+ assert(repartition.numPartitions == 16, optimizedPlan.toString)
+ }
+ }
+
+ test("lateral vector search repartitions automatically broadcast join
streamed limited input") {
+ Seq(false, true).foreach {
+ aqeEnabled =>
+ withSparkSQLConf(
+ "spark.sql.adaptive.enabled" -> aqeEnabled.toString,
+ "spark.sql.autoBroadcastJoinThreshold" -> "1024",
+ "spark.sql.adaptive.autoBroadcastJoinThreshold" -> "1024",
+ "spark.paimon.vector-search.lateral-join.parallelism" -> "4"
+ ) {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val result = spark.sql("""
+ |SELECT q.gid AS query_gid, r.gid AS
result_gid
+ |FROM (
+ | SELECT s.gid, s.embs
+ | FROM (
+ | SELECT id AS gid, array(1.0F, 2.0F,
3.0F) AS embs
+ | FROM range(0, 10000, 1, 8)
+ | LIMIT 1000
+ | ) s
+ | JOIN VALUES (0L) AS d(gid)
+ | ON s.gid = d.gid
+ |) q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search(
+ | 'vector_search_source', 'embs',
q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ val executedPlan = result.queryExecution.executedPlan
+ val broadcastJoin = collect(executedPlan) {
+ case join: BroadcastHashJoinExec => join
+ }.headOption.getOrElse(fail(executedPlan.toString))
+ val lateralVectorSearch = collect(executedPlan) {
+ case exec: LateralVectorSearchExec => exec
+ }.headOption.getOrElse(fail(executedPlan.toString))
+
+ withClue(s"AQE enabled: $aqeEnabled\n$executedPlan") {
+ assert(broadcastJoin.buildSide == BuildRight)
+
assert(lateralVectorSearch.child.outputPartitioning.numPartitions == 4)
+ }
+ }
+ }
+ }
+ }
+
+ test("lateral vector search preserves automatically broadcast join streamed
parallelism") {
+ Seq(false, true).foreach {
+ aqeEnabled =>
+ withSparkSQLConf(
+ "spark.sql.adaptive.enabled" -> aqeEnabled.toString,
+ "spark.sql.autoBroadcastJoinThreshold" -> "1024",
+ "spark.sql.adaptive.autoBroadcastJoinThreshold" -> "1024",
+ "spark.paimon.vector-search.lateral-join.parallelism" -> "4"
+ ) {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val result = spark.sql("""
+ |SELECT q.gid AS query_gid, r.gid AS
result_gid
+ |FROM (
+ | SELECT d.gid, d.embs
+ | FROM (
+ | SELECT id AS gid
+ | FROM range(0, 10, 1, 1)
+ | LIMIT 10
+ | ) s
+ | JOIN (
+ | SELECT id AS gid, array(1.0F, 2.0F,
3.0F) AS embs
+ | FROM range(0, 10000, 1, 8)
+ | ) d
+ | ON s.gid = d.gid
+ |) q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search(
+ | 'vector_search_source', 'embs',
q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ val executedPlan = result.queryExecution.executedPlan
+ val broadcastJoin = collect(executedPlan) {
+ case join: BroadcastHashJoinExec => join
+ }.headOption.getOrElse(fail(executedPlan.toString))
+ val lateralVectorSearch = collect(executedPlan) {
+ case exec: LateralVectorSearchExec => exec
+ }.headOption.getOrElse(fail(executedPlan.toString))
+
+ withClue(s"AQE enabled: $aqeEnabled\n$executedPlan") {
+ assert(broadcastJoin.buildSide == BuildLeft)
+
assert(lateralVectorSearch.child.outputPartitioning.numPartitions == 8)
+ }
+ }
+ }
+ }
+ }
+
+ test("lateral vector search repartitions coalesced limited input") {
+ withTable("vector_search_source") {
+ createVectorSearchSource()
+
+ val optimizedPlan = spark
+ .sql("""
+ |WITH q_limit AS (
+ | SELECT gid, embs
+ | FROM vector_search_source
+ | LIMIT 1000
+ |)
+ |SELECT q.gid AS query_gid, r.gid AS result_gid
+ |FROM (
+ | SELECT /*+ COALESCE(16) */ gid, embs
+ | FROM q_limit
+ |) q,
+ |LATERAL (
+ | SELECT gid
+ | FROM vector_search('vector_search_source', 'embs', q.embs, 3)
+ |) AS r
+ |""".stripMargin)
+ .queryExecution
+ .optimizedPlan
+
+ val lateralVectorSearch = optimizedPlan
+ .collectFirst { case lvs: LateralVectorSearch => lvs }
+ .getOrElse(fail(optimizedPlan.toString))
+
+ assert(lateralVectorSearch.left.isInstanceOf[Repartition],
optimizedPlan.toString)
+ val repartition = lateralVectorSearch.left.asInstanceOf[Repartition]
+ assert(repartition.shuffle, optimizedPlan.toString)
+ assert(repartition.numPartitions == 16, optimizedPlan.toString)
+ }
+ }
+
test("lateral vector search preserves subquery alias qualifiers") {
withTable("vector_search_source", "vector_search_result") {
spark.sql("""
@@ -542,6 +888,20 @@ class TableValuedFunctionsTest extends PaimonHiveTestBase {
.table(tableIdent)
}
+ private def createVectorSearchSource(): Unit = {
+ spark.sql("""
+ |CREATE TABLE vector_search_source (gid BIGINT, embs
ARRAY<FLOAT>, dt STRING)
+ |USING paimon
+ |TBLPROPERTIES (
+ | 'vector.file.format' = 'lance',
+ | 'vector-field' = 'embs',
+ | 'field.embs.vector-dim' = '3',
+ | 'row-tracking.enabled' = 'true',
+ | 'data-evolution.enabled' = 'true')
+ |PARTITIONED BY (dt)
+ |""".stripMargin)
+ }
+
private def utcMills(timestamp: String) =
Timestamp.fromLocalDateTime(LocalDateTime.parse(timestamp)).getMillisecond
}
diff --git
a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
index 8877446d76..0784530909 100644
---
a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
+++
b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/MinorVersionShim.scala
@@ -20,6 +20,7 @@ package org.apache.spark.sql.paimon.shims
import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression}
import org.apache.spark.sql.catalyst.plans.logical.{CTERelationRef,
LogicalPlan, MergeAction, MergeIntoTable}
+import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution,
Distribution}
object MinorVersionShim {
@@ -29,6 +30,12 @@ object MinorVersionShim {
output: Seq[Attribute],
isStreaming: Boolean): CTERelationRef = CTERelationRef(cteId, resolved,
output, isStreaming)
+ def createClusteredDistribution(expressions: Seq[Expression], numPartitions:
Int): Distribution =
+ ClusteredDistribution(
+ expressions,
+ requireAllClusterKeys = false,
+ requiredNumPartitions = Some(numPartitions))
+
def createMergeIntoTable(
targetTable: LogicalPlan,
sourceTable: LogicalPlan,
diff --git
a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala
b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala
index 9bde530b28..f14c212590 100644
---
a/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala
+++
b/paimon-spark/paimon-spark3-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark3Shim.scala
@@ -38,6 +38,7 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute,
AttributeReference,
import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression
import org.apache.spark.sql.catalyst.parser.ParserInterface
import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Assignment,
CTERelationRef, InsertAction, LogicalPlan, MergeAction, MergeIntoTable,
SubqueryAlias, TableSpec, UnresolvedWith, UpdateAction}
+import org.apache.spark.sql.catalyst.plans.physical.Distribution
// NOTE: `MergeRows` / `MergeRows.Keep` were introduced in Spark 3.4. We
access them only via
// reflection inside the `mergeRowsKeep*` method bodies so that loading
`Spark3Shim` does not fail
// on Spark 3.2 / 3.3 runtimes that still ship `paimon-spark3-common` (the
module targets 3.5.8 at
@@ -228,6 +229,11 @@ class Spark3Shim extends SparkShim {
isStreaming: Boolean): CTERelationRef =
MinorVersionShim.createCTERelationRef(cteId, resolved, output, isStreaming)
+ override def createClusteredDistribution(
+ expressions: Seq[Expression],
+ numPartitions: Int): Distribution =
+ MinorVersionShim.createClusteredDistribution(expressions, numPartitions)
+
override def supportsHashAggregate(
aggregateBufferAttributes: Seq[Attribute],
groupingExpression: Seq[Expression]): Boolean =
diff --git
a/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
b/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
index 7a7cdc70f5..91664858ac 100644
---
a/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
+++
b/paimon-spark/paimon-spark4-common/src/main/scala/org/apache/spark/sql/paimon/shims/Spark4Shim.scala
@@ -39,6 +39,7 @@ import
org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression
import org.apache.spark.sql.catalyst.parser.ParserInterface
import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, Assignment,
ColumnDefinition, CTERelationRef, InsertAction, LogicalPlan, MergeAction,
MergeIntoTable, MergeRows, SubqueryAlias, TableSpec, UnresolvedWith,
UpdateAction}
import org.apache.spark.sql.catalyst.plans.logical.MergeRows.{Copy, Insert,
Keep, Update}
+import org.apache.spark.sql.catalyst.plans.physical.{ClusteredDistribution,
Distribution}
import org.apache.spark.sql.catalyst.rules.Rule
import org.apache.spark.sql.catalyst.util.{ArrayData, GeneratedColumn,
IdentityColumn, ResolveDefaultColumns}
import org.apache.spark.sql.connector.catalog.{CatalogV2Util, Column,
Identifier, StagingTableCatalog, Table, TableCatalog}
@@ -227,6 +228,14 @@ class Spark4Shim extends SparkShim {
CTERelationRef(cteId, resolved, output.toSeq, isStreaming)
}
+ override def createClusteredDistribution(
+ expressions: Seq[Expression],
+ numPartitions: Int): Distribution =
+ ClusteredDistribution(
+ expressions,
+ requireAllClusterKeys = false,
+ requiredNumPartitions = Some(numPartitions))
+
override def supportsHashAggregate(
aggregateBufferAttributes: Seq[Attribute],
groupingExpression: Seq[Expression]): Boolean = {