From 98759cfcd6078e5c058eee36c926803bbc649da1 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 09:11:58 -0700 Subject: [PATCH 1/9] feat: indexed nearest-by join over Lance (Spark 4.2 SQL) Accelerate Spark 4.2's `APPROX NEAREST k BY DISTANCE ...` join (SPARK-56395) over Lance tables by intercepting the `NearestByJoin` operator and running a per-row native ANN/exact probe against the Lance vector index instead of the default O(|L|x|R|) cross-product + MaxMinByK lowering. New opt-in module `lance-spark-knn-4.2_2.13`: - Catalyst postHocResolutionRule (`IndexedNearestByJoinRule`) that recognizes `NearestByJoin` and rewrites it to a single no-shuffle logical node; must run in postHoc (not the optimizer) because Spark's `RewriteNearestByJoin` fires first in `FinishAnalysis`. - Planner strategy + exec (`LanceKnnJoinStrategy`/`LanceKnnJoinExec`) lowering to one `mapPartitions` -> per-left-row native probe -> bounded top-K heap -> late materialize by row address. No shuffle, no broadcast. - Probe core (`LanceProbe`/`LanceKnnJoinStage`/`TopKHeap`/`Metric`) supporting L2, cosine, and dot metrics with `nprobes`/`refineFactor` tuning. - Right-side WHERE pushdown as a Lance prefilter, including nested struct-field predicates rendered as dotted paths (`col.field`). Conservative: if any sub-expression is unsupported the whole rewrite is refused and Spark's brute-force lowering runs unchanged. Off by default (`spark.lance.knn.indexedNearestByJoin.enabled=false`); wired via `LanceKnnSparkSessionExtensions`. Part of the effort tracked in #541. Co-Authored-By: Claude Opus 4.8 --- lance-spark-knn-4.2_2.13/pom.xml | 112 ++++ .../catalyst/IndexedNearestByJoinRule.scala | 454 +++++++++++++++++ .../spark/knn/catalyst/LanceKnnJoinExec.scala | 78 +++ .../catalyst/LanceKnnJoinLogicalPlan.scala | 55 ++ .../knn/catalyst/LanceKnnJoinStrategy.scala | 36 ++ .../LanceKnnSparkSessionExtensions.scala | 55 ++ .../knn/internal/LanceKnnJoinStage.scala | 207 ++++++++ .../lance/spark/knn/internal/LanceProbe.scala | 356 +++++++++++++ .../org/lance/spark/knn/internal/Metric.scala | 68 +++ .../spark/knn/internal/ScoredRowRef.scala | 40 ++ .../lance/spark/knn/internal/TopKHeap.scala | 114 +++++ .../IndexedNearestByJoinRuleTest.scala | 462 +++++++++++++++++ .../IndexedNearestByJoinSqlTest.scala | 481 ++++++++++++++++++ .../internal/LanceProbeValidationTest.scala | 202 ++++++++ .../internal/LanceVectorIndexBuilder.scala | 105 ++++ .../spark/knn/internal/TopKHeapTest.scala | 92 ++++ .../knn/testutil/ClusteredEmbeddings.scala | 137 +++++ pom.xml | 1 + 18 files changed, 3055 insertions(+) create mode 100644 lance-spark-knn-4.2_2.13/pom.xml create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinLogicalPlan.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinStrategy.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/extensions/LanceKnnSparkSessionExtensions.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala create mode 100644 lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceVectorIndexBuilder.scala create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/testutil/ClusteredEmbeddings.scala diff --git a/lance-spark-knn-4.2_2.13/pom.xml b/lance-spark-knn-4.2_2.13/pom.xml new file mode 100644 index 000000000..3482611e5 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/pom.xml @@ -0,0 +1,112 @@ + + + 4.0.0 + + + org.lance + lance-spark-root + 0.7.1 + ../pom.xml + + + lance-spark-knn-4.2_2.13 + ${project.artifactId} + Indexed nearest-by join on Lance — Spark 4.2 SQL Catalyst integration (SPARK-56395) + jar + + + ${scala213.version} + ${scala213.compat.version} + ${arrow19.version} + ${java17.release} + + + + + org.apache.spark + spark-sql_${scala.compat.version} + ${spark42.version} + provided + + + org.apache.spark + spark-catalyst_${scala.compat.version} + ${spark42.version} + provided + + + + org.lance + lance-spark-base_${scala.compat.version} + ${project.version} + + + org.apache.arrow + arrow-memory-netty-buffer-patch + + + + + + org.lance + lance-spark-4.2_${scala.compat.version} + ${project.version} + test + + + + + + + + net.alchim31.maven + scala-maven-plugin + ${scala-maven-plugin.version} + + + scala-compile-first + process-resources + + compile + + + + scala-test-compile + process-test-resources + + testCompile + + + + + + -feature + -release + ${java.release} + + + + + org.apache.maven.plugins + maven-compiler-plugin + ${maven-compiler-plugin.version} + + ${java.release} + + + + + + + java21 + + 21 + + + + diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala new file mode 100644 index 000000000..e83f02e68 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala @@ -0,0 +1,454 @@ +/* + * Licensed 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.lance.spark.knn.catalyst + +import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, AttributeSet, EqualTo, Expression, GetStructField, GreaterThan, GreaterThanOrEqual, In, IsNotNull, IsNull, LessThan, LessThanOrEqual, Literal, Not, Or, VectorCosineSimilarity, VectorInnerProduct, VectorL2Distance} +import org.apache.spark.sql.catalyst.plans.{JoinType, LeftOuter, NearestByDirection, NearestByDistance, NearestBySimilarity} +import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, NearestByJoin, Project, SubqueryAlias} +import org.apache.spark.sql.catalyst.rules.Rule +import org.apache.spark.sql.connector.catalog.Table +import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation +import org.apache.spark.sql.types.{BooleanType, ByteType, DoubleType, FloatType, IntegerType, LongType, ShortType, StringType, StructField, StructType} +import org.apache.spark.unsafe.types.UTF8String +import org.lance.spark.knn.internal.{LanceKnnJoinStage, Metric} + +/** + * Catalyst rule that rewrites a Spark [[NearestByJoin]] (`approx = true`) over a Lance scan with + * a recognized vector-distance ranking expression into a single [[LanceKnnJoinLogicalPlan]], + * wrapped in a top-level [[Project]] that restores `NearestByJoin.output` exactly. The paired + * [[LanceKnnJoinStrategy]] then lowers that node to [[LanceKnnJoinExec]], which drives the + * no-shuffle `LanceKnnJoinStage.runPartition` per-partition probe. + * + * == Why this rule must be a `postHocResolutionRule`, not an optimizer rule == + * + * Spark's built-in [[org.apache.spark.sql.catalyst.optimizer.RewriteNearestByJoin]] rule runs in + * the optimizer's `FinishAnalysis` batch — the very first batch. `injectOptimizerRule` adds + * rules to `operatorOptimizationBatch`, which runs AFTER `FinishAnalysis`. By the time an + * injected optimizer rule fires, the `NearestByJoin` operator has already been replaced with the + * cross-product + `MaxMinByK` rewrite, and we have nothing to pattern-match. + * + * `injectPostHocResolutionRule` runs after analysis but before any optimizer batch — it is the + * only injection point that sees the unrewritten `NearestByJoin`. The same constraint applies to + * any future engine wanting to substitute a different physical strategy for `NearestByJoin`. + * + * == Pattern match == + * + * The rule fires on the conjunction of: + * - `NearestByJoin(_, right, joinType, approx = true, k, rankingExpression, direction)` + * - `right` resolves to a Lance DSv2 relation (immediate or under a `SubqueryAlias`) + * - `rankingExpression` is one of three recognized vector functions, AND its direction matches + * the direction declared on `NearestByJoin`: + * + * | Spark expression | direction | metric | + * |---------------------------------|------------------------|----------------| + * | `VectorL2Distance(L, R)` | `NearestByDistance` | `Metric.L2` | + * | `VectorCosineSimilarity(L, R)` | `NearestBySimilarity` | `Metric.Cosine`| + * | `VectorInnerProduct(L, R)` | `NearestBySimilarity` | `Metric.Dot` | + * + * Any other shape is left alone — Spark's default cross-product rewrite handles it. + * + * The two arguments of the ranking function must each resolve to an [[Attribute]] from one + * specific side of the join. Mixed-side compounds (e.g. `l2_distance(left.vec, left.vec)`) and + * derived expressions (e.g. `l2_distance(left.vec, slice(right.vec, ...))`) are out of scope and + * fall through to the cross-product rewrite. + * + * == Lance scan detection == + * + * Class-name match: `getClass.getName.contains("Lance")`. The probe / materialize path drives + * Lance's Java API directly, so the indexed-path executor is Lance-specific by construction — + * there's no general "any vector-capable backend" extension point here. URI extracted from the + * standard `path` / `datasetUri` option. The rule is opt-in via + * `spark.lance.knn.indexedNearestByJoin.enabled`, so a false positive can only fire when the user + * explicitly enabled the feature against a non-Lance backend; the runtime probe would surface the + * mismatch. + * + * == Prefilter pushdown == + * + * If the right side is a `Filter(cond, lance)` (a `WHERE` clause on the indexed table), the + * rule translates the predicate to a Lance SQL filter string and threads it through to the + * probe. Lance applies the filter BEFORE the index lookup (we always pass `prefilter = true`), + * so the top-K is computed over only the rows matching the filter — the only correct behavior + * for `right WHERE p APPROX NEAREST K`. + * + * Translation is conservative: it handles binary comparisons (=, !=, <, <=, >, >=), `IN`, + * `IS [NOT] NULL`, `AND`/`OR`/`NOT` over right-side columns (top-level or nested struct fields) + * vs. literals. Anything else + * (UDFs, subqueries, computed expressions) means the rule REFUSES the rewrite and returns the + * original `NearestByJoin`, falling through to Spark's brute-force cross-product. Refusal — not + * "push what we can, drop the rest" — because dropping a residual would silently change result + * semantics. The job becomes slow rather than wrong. + * + * Filter pushdown into the V2 relation does NOT happen at this point: this rule runs as a + * `postHocResolutionRule` (before the optimizer), so the right side is still the freshly + * analyzed `Filter` over `DataSourceV2Relation` — the V2 `SupportsPushDownFilters` step has not + * yet run. After we rewrite, the right side is absorbed into our plan, so V2 pushdown never + * gets a chance to drop the filter on the floor. + */ +object IndexedNearestByJoinRule extends Rule[LogicalPlan] { + + /** Configuration key that gates the rule. Off by default to keep the rule opt-in for now. */ + val EnabledConfKey: String = "spark.lance.knn.indexedNearestByJoin.enabled" + + /** + * IVF cluster count to visit per query. Higher = better recall, more compute. Default + * (None) leaves Lance's index-default (typically 1). + */ + val NprobesConfKey: String = "spark.lance.knn.nprobes" + + /** + * IVF-PQ refine factor — Lance fetches `K * refineFactor` PQ candidates and re-ranks them + * with exact distance using full vectors. Highest-leverage recall knob for IVF-PQ. Default + * (None) leaves Lance's index-default (= 1, no re-rank). + */ + val RefineFactorConfKey: String = "spark.lance.knn.refineFactor" + + override def apply(plan: LogicalPlan): LogicalPlan = { + if (!conf.getConfString(EnabledConfKey, "false").toBoolean) { + return plan + } + val nprobes = optInt(NprobesConfKey) + val refineFactor = optInt(RefineFactorConfKey) + plan.transformDown { + case j @ NearestByJoin(left, right, joinType, true, k, rankingExpr, direction) => + rewriteIfApplicable( + j, + left, + right, + joinType, + k, + rankingExpr, + direction, + nprobes, + refineFactor).getOrElse(j) + } + } + + private def optInt(key: String): Option[Int] = + Option(conf.getConfString(key, null)).map(_.toInt) + + /** + * Rewrite `NearestByJoin` into a single [[LanceKnnJoinLogicalPlan]] carrying the + * [[LanceKnnJoinStage.Conf]] the executor runs per partition — the same stage the DataFrame + * API path drives: + * + * {{{ + * Project(j.output, drop __score) + * +- LanceKnnJoinLogicalPlan output = left ++ right ++ __score + * +- left + * }}} + * + * We add a top-level `Project` because `NearestByJoin.output` is `left ++ right` (no + * score), but the join node emits `left ++ right ++ __score`. The Project slices the trailing + * score attribute — Catalyst's ColumnPruning won't interfere because + * `LanceKnnJoinLogicalPlan` overrides `references = child.outputSet`. + */ + private def rewriteIfApplicable( + j: NearestByJoin, + left: LogicalPlan, + right: LogicalPlan, + joinType: JoinType, + k: Int, + rankingExpr: Expression, + direction: NearestByDirection, + nprobes: Option[Int], + refineFactor: Option[Int]): Option[LogicalPlan] = { + for { + (metric, leftVecAttr, rightVecCol) <- recognizeRanking(rankingExpr, direction, left, right) + lance <- unwrapLanceScan(right) + } yield { + val leftVecIdx = left.output.indexWhere(_.exprId == leftVecAttr.exprId) + require(leftVecIdx >= 0, s"left vector attr not found in left.output: $leftVecAttr") + + val rightFields: Seq[StructField] = + lance.output.map(a => StructField(a.name, a.dataType, nullable = true)) + val rightProjection: Seq[String] = lance.output.map(_.name) + + val stageConf = LanceKnnJoinStage.Conf( + datasetUri = lance.uri, + version = lance.version, + vectorColumn = rightVecCol, + metric = metric, + k = k, + internalK = k, // no overfetch on the SQL path + nprobes = nprobes, + refineFactor = refineFactor, + ef = None, + prefilter = lance.prefilter, + leftVecIdx = leftVecIdx, + rightProjection = rightProjection, + rightFields = rightFields, + leftFieldCount = left.output.size, + outerJoin = joinType == LeftOuter, + smallerIsBetter = metric.smallerIsBetter) + + // The join node emits left ++ right ++ __score. The SQL output is j.output (= left ++ right, + // no score). Set finalOutput = j.output :+ scoreAttr so the node's output is stable; the + // top-level Project drops __score. + // + // `NearestByJoin.output` widens every left+right attribute to `nullable = true` — a contract + // the base Spark rewrite also honors. `finalSchema` feeds the `ExpressionEncoder` in + // `LanceKnnJoinExec.doExecute`; if we left left fields at raw `nullable = false` while the + // logical output declares them nullable, the encoder's binary layout would drift from what + // downstream consumers expect. Widen left here to keep the encoder consistent. + val leftSchemaStruct = StructType( + left.output.map(a => StructField(a.name, a.dataType, a.nullable))) + val scoreAttr = AttributeReference("__score", FloatType, nullable = true)() + val finalSchema = StructType( + leftSchemaStruct.fields.map(_.copy(nullable = true)) ++ + rightFields.map(f => f.copy(nullable = true)) :+ + StructField("__score", FloatType, nullable = true)) + val finalOutput: Seq[Attribute] = j.output :+ scoreAttr + + val node = LanceKnnJoinLogicalPlan( + child = left, + stageConf = stageConf, + leftSchema = leftSchemaStruct, + finalSchema = finalSchema, + finalOutput = finalOutput) + + // Top-level Project drops the __score attr so the plan's external output matches + // NearestByJoin.output exactly. + Project(j.output, node) + } + } + + /** Lance scan info extracted from a DSv2 relation, optionally with a translated prefilter. */ + final private case class LanceScanInfo( + uri: String, + version: Option[Long], + output: Seq[Attribute], + prefilter: Option[String]) + + private def unwrapLanceScan(plan: LogicalPlan): Option[LanceScanInfo] = plan match { + case SubqueryAlias(_, child) => unwrapLanceScan(child) + case v: org.apache.spark.sql.catalyst.plans.logical.View => + // SQL `createOrReplaceTempView` + `spark.sql(... FROM ...)` wraps the underlying + // DataSourceV2Relation in a `View`. Unwrap to find the actual relation underneath. + unwrapLanceScan(v.children.head) + case Filter(cond, child) => + // Right-side `WHERE` clause. Recurse first so we have the relation's output to validate + // attribute references against, then translate the predicate. If translation fails we + // bail entirely (return None, no rewrite) — pushing only PART of a `WHERE` would silently + // change query semantics. The user's filter must be pushed in full or not at all. + unwrapLanceScan(child).flatMap { info => + translateFilter(cond, AttributeSet(info.output)).map { sql => + val combined = info.prefilter match { + case Some(prev) => Some(s"($prev) AND ($sql)") + case None => Some(sql) + } + info.copy(prefilter = combined) + } + } + case Project(projectList, child) if isPassthroughProject(projectList, child) => + // `SELECT * FROM lance` analyzes to `Project(, lance)` — a pass-through + // that preserves attrs and exprIds. Unwrap it. Non-pass-through Projects (renames, drops, + // computed columns) would change the schema we rely on for `j.output` mapping, so we + // refuse those by falling through to the default `_ => None` case. + unwrapLanceScan(child) + case rel: DataSourceV2Relation if isLanceTable(rel.table) => + // The probe / materialize path drives Lance's Java API directly, so this rule is + // Lance-specific by construction — there's no plug-in point for a non-Lance backend + // here. We detect Lance via class-name match and pull the URI from the standard `path` + // / `datasetUri` option. If neither is present we fall through (returning None lets + // Spark's brute-force rewrite handle the query). + val opts = rel.options + val uri = Option(opts.get("path")).orElse(Option(opts.get("datasetUri"))) + uri.map { u => + LanceScanInfo( + uri = u, + version = Option(opts.get("version")).map(_.toLong), + output = rel.output, + prefilter = None) + } + case _ => None + } + + /** + * Translate a Spark `Filter` predicate into a Lance SQL filter string. Returns `None` if any + * sub-expression isn't supported — refusal, not partial pushdown. + * + * Supported shapes, where `attr` is a right-side top-level column OR a nested struct field + * (the latter rendered as a dotted path `col.field`, arbitrarily deep). Array/map element + * access (`col[i]`) is NOT supported and falls through to refusal: + * - `attr literal` and `literal attr` for `=`, `!=`, `<`, `<=`, `>`, `>=` + * - `attr IS NULL` / `attr IS NOT NULL` + * - `attr IN (lit, lit, ...)` (the IN list must be all foldable literals) + * - `AND` / `OR` over supported children + * - `NOT` over supported child + * + * Anything else — UDFs, joins, subqueries, expressions on both sides referencing the LEFT + * input, computed sub-expressions on the right (e.g. `year(ts) = 2025`) — returns `None`. + * Lance's SQL dialect is DataFusion-flavored; the constructs above all parse identically + * there, so we don't need to translate operator names beyond literal serialization. + */ + private[catalyst] def translateFilter( + expr: Expression, + rightAttrs: AttributeSet): Option[String] = expr match { + case And(l, r) => + for { + a <- translateFilter(l, rightAttrs) + b <- translateFilter(r, rightAttrs) + } yield s"($a) AND ($b)" + case Or(l, r) => + for { + a <- translateFilter(l, rightAttrs) + b <- translateFilter(r, rightAttrs) + } yield s"($a) OR ($b)" + case Not(EqualTo(l, r)) => + // Render `NOT (a = b)` as `(a != b)` so it's the natural Lance form. + binaryOp(l, r, rightAttrs, "!=") + case Not(child) => + translateFilter(child, rightAttrs).map(s => s"NOT ($s)") + case IsNull(c) => + asRightColumn(c, rightAttrs).map(name => s"$name IS NULL") + case IsNotNull(c) => + asRightColumn(c, rightAttrs).map(name => s"$name IS NOT NULL") + case EqualTo(l, r) => binaryOp(l, r, rightAttrs, "=") + case GreaterThan(l, r) => binaryOp(l, r, rightAttrs, ">") + case GreaterThanOrEqual(l, r) => binaryOp(l, r, rightAttrs, ">=") + case LessThan(l, r) => binaryOp(l, r, rightAttrs, "<") + case LessThanOrEqual(l, r) => binaryOp(l, r, rightAttrs, "<=") + case In(value, list) if list.nonEmpty => + for { + col <- asRightColumn(value, rightAttrs) + lits <- list.foldLeft(Option(Vector.empty[String])) { (accOpt, e) => + accOpt.flatMap(acc => asLiteral(e).map(acc :+ _)) + } + } yield s"$col IN (${lits.mkString(", ")})" + case _ => None + } + + private def binaryOp( + l: Expression, + r: Expression, + rightAttrs: AttributeSet, + op: String): Option[String] = { + // attr literal — the natural shape + val attrLit = for { + col <- asRightColumn(l, rightAttrs) + lit <- asLiteral(r) + } yield s"$col $op $lit" + // literal attr — flip when the parser/optimizer emitted args in this order. Renders + // as `lit op col`, which DataFusion also accepts. + attrLit.orElse { + for { + col <- asRightColumn(r, rightAttrs) + lit <- asLiteral(l) + } yield s"$lit $op $col" + } + } + + private def asRightColumn(e: Expression, rightAttrs: AttributeSet): Option[String] = e match { + case a: Attribute if rightAttrs.contains(a) => Some(a.name) + case g: GetStructField => + // Nested struct field: render as a dotted path `col.field` (recursing so `a.b.c` works). + // Lance's filter dialect treats a dotted identifier as a nested column path — its scan + // planner runs with enable_relations = false — so this maps 1:1. The recursion also gates + // on the ROOT resolving to a right-side attribute, so a left-side or foreign root refuses. + val fieldName = g.name.getOrElse(g.childSchema(g.ordinal).name) + asRightColumn(g.child, rightAttrs).map(base => s"$base.$fieldName") + case _ => None + } + + /** + * Render a Spark literal as a Lance SQL literal. Dispatch is by `dataType`, NOT by the boxed + * value class — Catalyst stores e.g. `Literal(0, DateType)` with the value as a plain `Int`, + * so a value-class match would silently let a date literal through as the integer "0", a + * recall-corrupting mistranslation. + * + * Supports nulls, booleans, numeric primitives, and strings (with `'`-escaped quoting). Bails + * on dates, timestamps, decimals, binary, arrays, structs — those have non-trivial cross- + * dialect renderings and we'd rather refuse pushdown than risk a wrong filter. + */ + private def asLiteral(e: Expression): Option[String] = e match { + case Literal(null, _) => Some("NULL") + case Literal(v, BooleanType) => Some(v.toString) + case Literal(v, ByteType) => Some(v.toString) + case Literal(v, ShortType) => Some(v.toString) + case Literal(v, IntegerType) => Some(v.toString) + case Literal(v, LongType) => Some(v.toString) + case Literal(v, FloatType) => Some(v.toString) + case Literal(v, DoubleType) => Some(v.toString) + case Literal(v: UTF8String, StringType) => Some(quoteString(v.toString)) + case Literal(v: String, StringType) => Some(quoteString(v)) + case _ => None + } + + private def quoteString(s: String): String = "'" + s.replace("'", "''") + "'" + + /** + * True iff the Project is the canonical `SELECT *` form: same number of outputs as the child, + * each entry a bare `AttributeReference` whose `exprId` matches the child's output in order. + * Any aliasing, reordering, dropping, or computed column — return false and refuse to + * unwrap, since those change the schema we'd surface as the join's right-side output. + */ + private def isPassthroughProject( + projectList: Seq[org.apache.spark.sql.catalyst.expressions.NamedExpression], + child: LogicalPlan): Boolean = { + val childOut = child.output + if (projectList.size != childOut.size) return false + projectList.zip(childOut).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + } + + private def isLanceTable(table: Table): Boolean = { + val cls = table.getClass.getName + // Loose by design — the rule is opt-in via spark.lance.knn.indexedNearestByJoin.enabled, so + // a false positive here would only fire when the user explicitly turned the feature on + // against a non-Lance backend, and the runtime probe would surface the mismatch. + cls.contains("Lance") || cls.contains("lance") + } + + /** + * Recognize `rankingExpr` as one of the supported vector-distance functions, AND verify the + * declared `direction` on `NearestByJoin` matches the function's natural ordering. + * + * Returns `(metric, leftVecAttr, rightVecColName)` on success. + */ + private def recognizeRanking( + rankingExpr: Expression, + direction: NearestByDirection, + left: LogicalPlan, + right: LogicalPlan): Option[(Metric, Attribute, String)] = { + val (metric, lhs, rhs) = rankingExpr match { + case VectorL2Distance(l, r) if direction == NearestByDistance => (Metric.L2, l, r) + case VectorCosineSimilarity(l, r) if direction == NearestBySimilarity => (Metric.Cosine, l, r) + case VectorInnerProduct(l, r) if direction == NearestBySimilarity => (Metric.Dot, l, r) + case _ => return None + } + // Each argument must be a bare attribute from one side of the join. + (asAttr(lhs), asAttr(rhs)) match { + case (Some(la), Some(ra)) => + val leftAttrIds = left.outputSet + val rightAttrIds = right.outputSet + if (leftAttrIds.contains(la) && rightAttrIds.contains(ra)) { + Some((metric, la, ra.name)) + } else if (leftAttrIds.contains(ra) && rightAttrIds.contains(la)) { + // Argument order swapped — still valid for symmetric metrics. All three of L2/Cosine/Dot + // are symmetric so we don't have to retain the original orientation. + Some((metric, ra, la.name)) + } else { + None + } + case _ => None + } + } + + private def asAttr(e: Expression): Option[Attribute] = e match { + case a: Attribute => Some(a) + case _ => None + } +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala new file mode 100644 index 000000000..ae62196d3 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala @@ -0,0 +1,78 @@ +/* + * Licensed 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.lance.spark.knn.catalyst + +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.Row +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet} +import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} +import org.apache.spark.sql.types.StructType +import org.lance.spark.knn.internal.LanceKnnJoinStage + +/** + * Physical operator for the indexed nearest-by join. The whole join is one no-shuffle + * `mapPartitions` over the left input — [[requiredChildDistribution]] is intentionally NOT + * overridden, so Catalyst inserts NO `Exchange` above this node. Each task: + * + * 1. decodes the child's `RDD[InternalRow]` (the left input) into typed `Row`s, + * 2. drives [[LanceKnnJoinStage.runPartition]], which opens R's index once, probes + trims + + * late-materializes per left row (see that object's doc for why a shuffle/merge pipeline only + * adds cost), and + * 3. re-encodes the assembled `left ++ right ++ __score` rows back to `RDD[InternalRow]`. + * + * The decode/encode uses `ExpressionEncoder`; `.copy()` on both sides because Spark reuses the + * `InternalRow` buffer across iterations of the upstream/downstream operators. + */ +case class LanceKnnJoinExec( + override val child: SparkPlan, + stageConf: LanceKnnJoinStage.Conf, + leftSchema: StructType, + finalSchema: StructType, + finalOutput: Seq[Attribute]) + extends UnaryExecNode { + + override def output: Seq[Attribute] = finalOutput + + override def nodeName: String = "LanceKnnJoin" + + // The right-side + score attrs in `output` are synthesised per row from the probe results; + // they do not appear in `child.output`. Declare them produced so Spark's `missingInput` check + // (and the `!` marker in tree-string output) doesn't flag this node. + override def producedAttributes: AttributeSet = AttributeSet(output) -- child.outputSet + + override protected def doExecute(): RDD[InternalRow] = { + val childRdd = child.execute() + val leftSchemaCaptured = leftSchema + val finalSchemaCaptured = finalSchema + val confCaptured = stageConf + + // Encoders are created on the driver and captured into the closure (ExpressionEncoder is + // serializable). The deserializer/serializer instances are NOT thread-safe, so build them + // per partition inside `mapPartitions`. + val leftEnc = ExpressionEncoder(leftSchemaCaptured).resolveAndBind() + val finalEnc = ExpressionEncoder(finalSchemaCaptured).resolveAndBind() + + childRdd.mapPartitions { iter => + val deser = leftEnc.createDeserializer() + val ser = finalEnc.createSerializer() + val leftRows: Iterator[Row] = iter.map(ir => deser(ir.copy())) + LanceKnnJoinStage.runPartition(leftRows, confCaptured).map(row => ser(row).copy()) + } + } + + override protected def withNewChildInternal(newChild: SparkPlan): LanceKnnJoinExec = + copy(child = newChild) +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinLogicalPlan.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinLogicalPlan.scala new file mode 100644 index 000000000..e67bc43e1 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinLogicalPlan.scala @@ -0,0 +1,55 @@ +/* + * Licensed 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.lance.spark.knn.catalyst + +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet} +import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode} +import org.apache.spark.sql.types.StructType +import org.lance.spark.knn.internal.LanceKnnJoinStage + +/** + * The single logical node the SQL rewrite ([[IndexedNearestByJoinRule]]) emits for an indexed + * `APPROX NEAREST K` join. Its one child is the LEFT input; the right (Lance) side is captured in + * `stageConf` (URI + version + probe parameters), NOT as a plan child — the join runs entirely + * inside one `mapPartitions` over the left rows, with NO shuffle. The matching + * [[LanceKnnJoinStrategy]] lowers this to [[LanceKnnJoinExec]], which drives the + * [[LanceKnnJoinStage.runPartition]] per-partition probe. + * + * `output` is `left ++ right ++ __score` — the right-side and score attributes are synthesised + * here from the probe results, so `producedAttributes = output -- child.outputSet` marks them as + * introduced by this node (Catalyst's `missingInput` check would otherwise flag them). + * + * The `references = child.outputSet` override is load-bearing: the matching exec decodes the WHOLE + * left row per partition to feed the probe, so no left column can be pruned. Without this override + * Catalyst's `ColumnPruning` sees a downstream consumer that references only a subset (or nothing, + * e.g. `count(*)`) and wraps the child in a narrowing `Project`, which would change the row shape + * the executor's left-side encoder expects. + */ +case class LanceKnnJoinLogicalPlan( + override val child: LogicalPlan, + stageConf: LanceKnnJoinStage.Conf, + leftSchema: StructType, + finalSchema: StructType, + finalOutput: Seq[Attribute]) + extends UnaryNode { + + override def output: Seq[Attribute] = finalOutput + + override def producedAttributes: AttributeSet = AttributeSet(output) -- child.outputSet + + override lazy val references: AttributeSet = child.outputSet + + override protected def withNewChildInternal(newChild: LogicalPlan): LanceKnnJoinLogicalPlan = + copy(child = newChild) +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinStrategy.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinStrategy.scala new file mode 100644 index 000000000..23eb85b1c --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinStrategy.scala @@ -0,0 +1,36 @@ +/* + * Licensed 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.lance.spark.knn.catalyst + +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan +import org.apache.spark.sql.execution.{SparkPlan, SparkStrategy} + +/** + * Lowers [[LanceKnnJoinLogicalPlan]] to [[LanceKnnJoinExec]]. Registered via + * `SparkSessionExtensions.injectPlannerStrategy` in + * [[org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions]]. `planLater(p.child)` defers + * planning of the left input to the rest of the planner. + */ +object LanceKnnJoinStrategy extends SparkStrategy { + override def apply(plan: LogicalPlan): Seq[SparkPlan] = plan match { + case p: LanceKnnJoinLogicalPlan => + LanceKnnJoinExec( + planLater(p.child), + p.stageConf, + p.leftSchema, + p.finalSchema, + p.finalOutput) :: Nil + case _ => Nil + } +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/extensions/LanceKnnSparkSessionExtensions.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/extensions/LanceKnnSparkSessionExtensions.scala new file mode 100644 index 000000000..bbaf2d088 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/extensions/LanceKnnSparkSessionExtensions.scala @@ -0,0 +1,55 @@ +/* + * Licensed 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.lance.spark.knn.extensions + +import org.apache.spark.sql.SparkSessionExtensions +import org.lance.spark.knn.catalyst.{IndexedNearestByJoinRule, LanceKnnJoinStrategy} + +/** + * Registers the Catalyst integration for the indexed nearest-by join (the SQL + * `APPROX NEAREST K BY DISTANCE ...` syntax added in Spark 4.2 by SPARK-56395). + * + * Wire this into a SparkSession with: + * + * {{{ + * SparkSession.builder() + * .config("spark.sql.extensions", + * "org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions") + * .config("spark.lance.knn.indexedNearestByJoin.enabled", "true") + * ... + * }}} + * + * The `enabled` flag gates the rule itself — see [[IndexedNearestByJoinRule.EnabledConfKey]]. Off + * by default to keep the integration opt-in. + * + * == Injection point: postHocResolutionRule, NOT optimizerRule == + * + * Spark's `RewriteNearestByJoin` runs in `FinishAnalysis`, which precedes the + * `operatorOptimizationBatch` that `injectOptimizerRule` adds rules to. By the time an injected + * optimizer rule fires, the `NearestByJoin` operator has already been replaced with the + * cross-product + `MaxMinByK` rewrite. `injectPostHocResolutionRule` runs after analysis but + * before any optimizer batch — this is the only injection point that sees the unrewritten + * `NearestByJoin`. See [[IndexedNearestByJoinRule]]'s class doc for the full rationale. + * + * Coexistence: this extension does not replace `LanceSparkSessionExtensions` from the connector + * modules; both can be wired together in a comma-separated `spark.sql.extensions` value. + */ +class LanceKnnSparkSessionExtensions extends (SparkSessionExtensions => Unit) { + override def apply(extensions: SparkSessionExtensions): Unit = { + extensions.injectPostHocResolutionRule(_ => IndexedNearestByJoinRule) + // Lowers the single `LanceKnnJoinLogicalPlan` the rule emits to `LanceKnnJoinExec` — the + // no-shuffle `LanceKnnJoinStage.runPartition` per-partition probe. + extensions.injectPlannerStrategy(_ => LanceKnnJoinStrategy) + } +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala new file mode 100644 index 000000000..41805e366 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -0,0 +1,207 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.apache.spark.sql.Row +import org.apache.spark.sql.types.StructField + +import scala.collection.mutable + +/** + * The whole indexed nearest-by join, done per Spark partition with NO shuffle. + * + * A single native `LanceProbe.probe(...)` call is already a complete distributed search: Lance + * probes the IVF index and scans the candidate fragments across its own threads, heap-merges in + * process, and returns the final top-K for one query. Handing that orchestration to Spark (a + * probe → shuffle → merge → materialize pipeline) only adds a shuffle round-trip, a redundant + * merge stage, a second materialize scan, and `M × N_frag × K` refs crossing the Rust→JVM + * boundary. So this stage keeps everything local: + * + * {{{ + * left.rdd.mapPartitions { rows => + * val probe = new LanceProbe(uri, fragmentIds = None, version) // whole-index, once per task + * rows.flatMap { leftRow => + * val hits = probe.probe(query(leftRow), internalK, ...) // native top-K search + * val topK = trimToK(hits) // overfetch → K + * val payloads = probe.materialize(topK.map(_.rowAddr)) // late point-fetch by _rowid + * topK.map(ref => assembleRow(leftRow, payloads(ref), ref.score)) + * } + * } + * }}} + * + * No `requiredChildDistribution`, no Exchange. Each task opens R's whole index (`fragmentIds = + * None`) — Lance does the cross-fragment merge internally — so per-executor resident memory grows + * with `|R|`. + * + * The SQL Catalyst node ([[org.lance.spark.knn.catalyst.LanceKnnJoinExec]]) drives this + * `runPartition`, so probe/trim/materialize semantics stay defined in exactly one place. + */ +object LanceKnnJoinStage { + + /** + * Everything a probe task needs, shipped from the driver. `internalK` is the overfetch count + * handed to Lance (`k × overfetch`); `k` is the final per-left-row cut applied after the native + * search. `leftVecIdx` is the position of the query-vector column in the left row. + */ + final case class Conf( + datasetUri: String, + version: Option[Long], + vectorColumn: String, + metric: Metric, + k: Int, + internalK: Int, + nprobes: Option[Int], + refineFactor: Option[Int], + ef: Option[Int], + prefilter: Option[String], + leftVecIdx: Int, + rightProjection: Seq[String], + rightFields: Seq[StructField], + leftFieldCount: Int, + outerJoin: Boolean, + smallerIsBetter: Boolean) + extends Serializable + + /** + * Run the join for one partition of left rows. Opens the probe once, probes + materializes per + * row, and returns assembled join rows. Materializing into an `ArrayBuffer` before returning the + * iterator is deliberate: Spark pulls from `mapPartitions` lazily, so a bare lazy iterator would + * let the consumer outlive the `try`/`finally` and read from a closed probe handle. + */ + def runPartition(leftRows: Iterator[Row], conf: Conf): Iterator[Row] = { + if (leftRows.isEmpty) return Iterator.empty + + val probe = new LanceProbe(conf.datasetUri, fragmentIds = None, version = conf.version) + val out = mutable.ArrayBuffer.empty[Row] + try { + leftRows.foreach { leftRow => + val q = extractVector(leftRow, conf.leftVecIdx) + if (q == null) { + // Null query vector: nothing to search. Emit a null-right row only for an outer join. + if (conf.outerJoin) { + out += assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null) + } + } else { + // Overfetch `internalK` candidates natively, then trim to the final `k`. Lance already + // returns them best-first, so when it hands back no more than `k` we keep them as-is and + // skip the heap entirely. + val refs = probe + .probe( + conf.vectorColumn, + q, + conf.internalK, + conf.metric, + conf.nprobes, + conf.refineFactor, + conf.ef, + conf.prefilter) + .toArray + val trimmed = + if (refs.length <= conf.k) refs + else { + val heap = new TopKHeap(conf.k, conf.smallerIsBetter) + heap.offerAll(refs) + heap.drain() + } + + if (trimmed.isEmpty) { + if (conf.outerJoin) { + out += assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null) + } + } else { + // Late materialization: point-fetch the surviving right rows by `_rowid`. Building the + // `rowAddr -> row` map collapses any duplicate rowAddr to one payload; the loop below + // still emits one output row per surviving ref. + val materialized: Map[Long, Map[String, Any]] = probe + .materialize(trimmed.iterator.map(_.rowAddr).toSeq, conf.rightProjection) + .map(m => extractRowAddr(m) -> m) + .toMap + trimmed.foreach { ref => + val rightMap = materialized.getOrElse(ref.rowAddr, null) + out += assembleRow( + leftRow, + conf.leftFieldCount, + conf.rightFields, + rightMap, + ref.score) + } + } + } + } + } finally probe.close() + out.iterator + } + + /** + * Pull a query vector out of a Spark `Row`'s ArrayType column. The Scala 2.13 `Seq` gotcha is + * real: `Row.get` on `ArrayType` returns `mutable.ArraySeq`, which `case s: Seq[_]` only matches + * against the root `scala.collection.Seq` trait (the default `Seq` alias is `immutable.Seq` on + * 2.13). + */ + private[knn] def extractVector(row: Row, idx: Int): Array[Float] = { + if (row.isNullAt(idx)) return null + row.get(idx) match { + case s: scala.collection.Seq[_] => + s.iterator.map { + case f: java.lang.Float => f.floatValue() + case f: Float => f + case d: java.lang.Double => d.doubleValue().toFloat + case d: Double => d.toFloat + case other => + throw new IllegalStateException( + s"Unsupported vector element type: ${other.getClass.getName}") + }.toArray + case arr: Array[Float] => arr + case arr: Array[java.lang.Float] => arr.map(_.floatValue()) + case other => + throw new IllegalStateException( + s"Unsupported vector column representation: ${other.getClass.getName}") + } + } + + /** Read the `_rowid` key out of a materialized row map (tolerating boxed / stringy longs). */ + private def extractRowAddr(m: Map[String, Any]): Long = + m.get(LanceProbe.RowIdColumn) match { + case Some(l: java.lang.Long) => l.longValue() + case Some(l: Long) => l + case Some(other) => other.toString.toLong + case None => + throw new IllegalStateException( + s"Materialized row missing ${LanceProbe.RowIdColumn}; " + + s"got keys: ${m.keys.mkString(", ")}") + } + + /** + * Assemble one output row: `left fields ++ right fields ++ score`. A null `rightValues` (outer + * join with no hit) fills the right side with nulls. + */ + private def assembleRow( + leftRow: Row, + leftFieldCount: Int, + rightFields: Seq[StructField], + rightValues: Map[String, Any], + score: Any): Row = { + val arr = new Array[Any](leftFieldCount + rightFields.size + 1) + var i = 0 + while (i < leftFieldCount) { arr(i) = leftRow.get(i); i += 1 } + var j = 0 + while (j < rightFields.size) { + arr(leftFieldCount + j) = + if (rightValues == null) null else rightValues.getOrElse(rightFields(j).name, null) + j += 1 + } + arr(leftFieldCount + rightFields.size) = score + Row.fromSeq(arr.toSeq) + } +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala new file mode 100644 index 000000000..0ebac1ab5 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -0,0 +1,356 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.apache.arrow.memory.BufferAllocator +import org.apache.arrow.vector.{BigIntVector, FieldVector, Float4Vector, Float8Vector, UInt8Vector, VectorSchemaRoot} +import org.apache.arrow.vector.ipc.ArrowReader +import org.lance.{Dataset, ReadOptions} +import org.lance.ipc.{LanceScanner, Query, ScanOptions} +import org.lance.spark.{LanceConstant, LanceRuntime} + +import java.util + +import scala.collection.JavaConverters._ +import scala.collection.mutable + +/** + * Per-task vector-index probe primitive. Opens a Lance dataset once and runs many `probe()` calls + * against a fixed set of fragments. Returns row references + scores only — no payload. Late + * materialization happens elsewhere (`LanceMaterialize`). + * + * This is the core primitive Phase 0 of the indexed nearest-by design depends on. Validating its + * cost profile is the first thing to do on a new Lance build: + * - dataset open should be one-time cost + * - per-probe cost should be index traversal + small overhead, not full fragment scan + * - returning top-K row addrs should match Lance's native nearest search recall + * + * Lifecycle: instantiate per task, call `probe(...)` repeatedly, close at end. + * + * @param datasetUri Lance dataset URI (passed straight to `Dataset.open`). + * @param fragmentIds Fragments this probe is restricted to. Pass `None` for whole-dataset search. + * @param version Optional Lance version to pin. Required when used inside a join, so all + * probe / materialize stages see the same snapshot. + * @param allocator Arrow allocator. Defaults to lance-spark's shared `LanceRuntime.allocator()`. + */ +final class LanceProbe( + datasetUri: String, + fragmentIds: Option[Seq[Int]], + version: Option[Long] = None, + allocator: BufferAllocator = LanceRuntime.allocator()) + extends AutoCloseable { + + // Open the dataset once. Lance's Java binding caches index metadata against the Dataset handle, + // so reusing it across probes keeps subsequent calls index-warm. + private val dataset: Dataset = openDataset() + + private val javaFragmentIds: Option[util.List[Integer]] = fragmentIds.map { ids => + val javaList = new util.ArrayList[Integer](ids.size) + ids.foreach(i => javaList.add(Integer.valueOf(i))) + javaList: util.List[Integer] + } + + private def openDataset(): Dataset = { + val readOpts = { + val b = new ReadOptions.Builder() + version.foreach(v => b.setVersion(v)) + b.build() + } + Dataset.open() + .uri(datasetUri) + .allocator(allocator) + .readOptions(readOpts) + .build() + } + + /** + * Run a single nearest-neighbor query. Returns up to `k` row references for the configured + * fragments, ordered best-first by `metric`. + * + * Implementation note: lance-spark mandates `prefilter = true` for fragmented vector queries + * (see `LanceFragmentScanner.create`). We mirror that here — Lance's index probe semantics + * require it when fragment scope is restricted. + * + * `vectorColumn` is a per-call argument (not a constructor field) because the same + * `LanceProbe` instance also serves the materialize stage via [[materialize]], which + * doesn't reference any vector column. Keeping it on the call sidesteps the smell of + * passing a placeholder string when constructing for materialize-only use. + * + * `prefilter` is a Lance SQL filter string (DataFusion-flavored). Lance applies it BEFORE the + * vector index lookup when `prefilter = true` (which we always set), so the top-K is computed + * over only the rows matching the filter — exactly what a `Filter(cond, lance) RIGHT JOIN ... + * APPROX NEAREST K` should do. Without prefilter pushdown, a per-fragment vector probe could + * return K rows that are all later filtered out post-join, masking truly-nearest-but-also- + * matching rows further down the index — a recall bug. The translator in + * `IndexedNearestByJoinRule` is responsible for producing only safely-translated SQL; here we + * just hand it through. + */ + def probe( + vectorColumn: String, + query: Array[Float], + k: Int, + metric: Metric, + nprobes: Option[Int] = None, + refineFactor: Option[Int] = None, + ef: Option[Int] = None, + prefilter: Option[String] = None): Seq[ScoredRowRef] = { + require(vectorColumn != null && vectorColumn.nonEmpty, "vectorColumn must be non-empty") + require(query != null && query.length > 0, "Query vector must be non-empty") + require(k > 0, "k must be positive") + + val q = { + val b = new Query.Builder() + .setColumn(vectorColumn) + .setKey(query) + .setK(k) + .setDistanceType(metric.lanceType) + nprobes.foreach(b.setNprobes(_)) + // refineFactor: IVF-PQ recall knob. Lance fetches `k * refineFactor` approximate + // candidates, then re-ranks them with exact distance and trims to k. Bigger factor = + // better recall, more compute. None leaves Lance's default (= 1, no re-rank). + refineFactor.foreach(b.setRefineFactor(_)) + // ef: HNSW search depth. Higher = better recall, more compute. None leaves Lance's + // index-default. Only meaningful for HNSW indexes; ignored for IVF-PQ. + ef.foreach(b.setEf(_)) + b.build() + } + + val opts = new ScanOptions.Builder() + .nearest(q) + // EXPERIMENT: drop prefilter(true). The single-machine reference path + // doesn't set it; this LanceProbe call does. Comparing wallclock with + // and without isolates whether the prefilter branch in + // vector_search_source forces a slower index plan than the postfilter + // (default) branch. Re-enable when fragmented probe + prefilter + // pushdown is needed (we know fragmentIds requires prefilter from the + // Lance-side error, but at probeParallelism=1 there are no fragments). + .withRowId(true) + // Project only what we need into the result. The vector column is implied by `nearest`; + // requesting an empty user column list keeps the Arrow batch narrow (just the rowid + + // distance metadata). Materialization fetches payload columns later. + .columns(java.util.Collections.emptyList[String]()) + + if (prefilter.nonEmpty || javaFragmentIds.nonEmpty) { + // Real prefilter or fragment scope is requested — keep prefilter(true) + // so Lance applies the filter / restricts to fragments correctly. + opts.prefilter(true) + } + + prefilter.filter(_.nonEmpty).foreach(opts.filter) + javaFragmentIds.foreach(opts.fragmentIds) + + val scanner: LanceScanner = LanceScanner.create(dataset, opts.build(), allocator) + try { + readScored(scanner.scanBatches()) + } finally { + scanner.close() + } + } + + /** + * Drain the Arrow stream from a nearest-search scan into `(rowId, score)` pairs. + * + * Expected schema: + * - `_rowid` : UInt8 / BigInt — Lance logical row identifier + * - `_distance` (or score column added by `nearest`) : Float4 / Float8 — ranking value + * + * We resolve columns by name to be encoding-version-agnostic; the underlying primitive type + * (UInt8 vs BigInt for the id, Float4 vs Float8 for score) varies across Arrow / Lance combos + * and we tolerate both. + */ + private def readScored(reader: ArrowReader): Seq[ScoredRowRef] = { + val out = mutable.ArrayBuffer.empty[ScoredRowRef] + try { + while (reader.loadNextBatch()) { + val root = reader.getVectorSchemaRoot + val addrVec: FieldVector = root.getVector(LanceProbe.RowIdColumn) + val scoreVec: FieldVector = LanceProbe.ScoreColumns.iterator + .map(name => Option(root.getVector(name)).orNull) + .find(_ != null) + .getOrElse(throw new IllegalStateException( + s"Lance nearest scan did not return a score column. Got: " + + root.getSchema.getFields.asScala.map(_.getName).mkString(", "))) + + val n = root.getRowCount + var i = 0 + while (i < n) { + val addr = addrVec match { + case v: UInt8Vector => v.get(i) + case v: BigIntVector => v.get(i) + case other => + throw new IllegalStateException( + s"Unexpected row-address vector type: ${other.getClass.getName}") + } + val score = scoreVec match { + case v: Float4Vector => v.get(i) + case v: Float8Vector => v.get(i).toFloat + case other => + throw new IllegalStateException( + s"Unexpected score vector type: ${other.getClass.getName}") + } + out += ScoredRowRef(addr, score) + i += 1 + } + } + } finally { + reader.close() + } + out.toSeq + } + + /** + * Materialize a set of right-side rows by their `_rowaddr`s. Used by the join's materialize + * stage to fetch full payloads after the probe + merge has decided which rows survive. + * + * The row addresses are pushed down as a `_rowaddr IN (...)` filter, which Lance executes via + * its row-address index — the natural point-fetch path. The result is unordered with respect + * to the input list; the caller re-aligns by `_rowaddr`. + * + * @param rowAddrs list of Lance `_rowid` values (parameter name retained for source + * compatibility with callers — semantically these are now row IDs). + * @param projection projected column list. `Seq.empty` means "all columns". + * @return a sequence of materialized rows, each represented as a `Map[String, Any]` for the + * projected columns plus an entry under `LanceProbe.RowIdColumn` so the caller can + * re-key. Returning a Map keeps this primitive Spark-agnostic; conversion to + * `InternalRow` happens in the API layer. + */ + def materialize( + rowAddrs: Seq[Long], + projection: Seq[String] = Seq.empty): Seq[Map[String, Any]] = { + if (rowAddrs.isEmpty) return Seq.empty + + val opts = new ScanOptions.Builder().withRowId(true) + if (projection.nonEmpty) { + opts.columns(projection.toList.asJava) + } + // `_rowid IN (a, b, c)` — Lance lowers this to its row-id lookup path. Same point-fetch + // semantics as `_rowaddr IN (...)` previously used here, but `_rowid` is the universal + // identifier (works on indexed + non-indexed scan paths alike). + // + // Each row ID is rendered as `arrow_cast('', 'UInt64')` for two + // compounding reasons: + // + // 1. Lance row IDs are 64-bit UNSIGNED; storing them as Java signed `long` means + // values >= 2^63 come back negative. `mkString(", ")` would render them as + // negative integer literals and Lance/DataFusion would reject (`Int64(-...) + // cannot convert to UInt64`). + // 2. Even after `Long.toUnsignedString` produces a positive 20-digit decimal, + // DataFusion's SQL parser tries `Int64` first, overflows, then falls back to + // `Float64`. `Float64` loses precision past 2^53 — the literal becomes a + // different number — and DataFusion then can't downcast `Float64` to `UInt64`. + // + // `arrow_cast(string, 'UInt64')` bypasses both: the string literal goes through + // `arrow_cast`'s own coercion, which is precision-preserving for UInt64. + // + // At 100K rows row IDs stay below 2^53 and both layers of the bug are invisible; at + // 1M+ rows they bite. Caught when the DataFrame benchmark hit 1M-row scale. + val rowIdLiterals = rowAddrs.iterator + .map(addr => s"arrow_cast('${java.lang.Long.toUnsignedString(addr)}', 'UInt64')") + .mkString(", ") + opts.filter(s"${LanceProbe.RowIdColumn} IN ($rowIdLiterals)") + javaFragmentIds.foreach(opts.fragmentIds) + + val scanner: LanceScanner = LanceScanner.create(dataset, opts.build(), allocator) + try { + readRows(scanner.scanBatches()) + } finally { + scanner.close() + } + } + + private def readRows(reader: ArrowReader): Seq[Map[String, Any]] = { + val out = mutable.ArrayBuffer.empty[Map[String, Any]] + try { + while (reader.loadNextBatch()) { + val root: VectorSchemaRoot = reader.getVectorSchemaRoot + val n = root.getRowCount + var i = 0 + while (i < n) { + val rowMap = mutable.LinkedHashMap.empty[String, Any] + val fields = root.getSchema.getFields.asScala + var f = 0 + while (f < fields.size) { + val name = fields(f).getName + val v = root.getVector(name) + rowMap(name) = if (v.isNull(i)) null else LanceProbe.toSparkValue(v.getObject(i)) + f += 1 + } + out += rowMap.toMap + i += 1 + } + } + } finally { + reader.close() + } + out.toSeq + } + + override def close(): Unit = dataset.close() +} + +object LanceProbe { + + /** + * Lance row-identity virtual column name. We use `_rowid` rather than `_rowaddr` because + * Lance's INDEXED nearest-search path materializes `_rowid` but not `_rowaddr`, while + * non-indexed scans materialize both. `_rowid` therefore works on every code path that + * calls `probe()` (with or without a vector index built on the column). Sourced from + * `LanceConstant` to keep the literal defined in exactly one place. + */ + val RowIdColumn: String = LanceConstant.ROW_ID + + /** + * Candidate names for the score column in a Lance nearest-search result. Lance's vector indexes + * have used `_distance` historically; tolerate `_score` too in case future versions rename it. + * The lookup is name-based so the consumer is agnostic to where Lance puts the column in its + * output schema. + */ + val ScoreColumns: Seq[String] = Seq("_distance", "_score") + + /** + * Convert an Arrow-returned cell value into something Spark's encoders accept when stuffed + * into a `Row`. Arrow's `FieldVector.getObject` returns Java types (boxed primitives, + * `JsonStringArrayList` for list cells, `Text` for utf8) which Spark's `RowEncoder` does not + * always understand directly — most painfully, a `java.util.ArrayList` can't satisfy a Spark + * `ArrayType` slot, which expects a `scala.collection.Seq`. + * + * Conversion rules, in order: + * - `java.util.List` → recursively-converted `Seq` + * - `java.util.Map` → recursively-converted Scala `Map` + * - `org.apache.arrow.vector.util.Text` → `String` + * - `Number` boxed primitives → returned as-is (Spark handles them) + * - everything else → returned as-is (caller's responsibility) + * + * Recursive on lists/maps to handle nested types (arrays of structs, etc.) without surprises + * for callers. + */ + def toSparkValue(value: Any): Any = value match { + case null => null + case list: java.util.List[_] => + val out = scala.collection.mutable.ArrayBuffer.empty[Any] + val it = list.iterator + while (it.hasNext) out += toSparkValue(it.next()) + out.toSeq + case map: java.util.Map[_, _] => + val out = scala.collection.mutable.LinkedHashMap.empty[Any, Any] + val it = map.entrySet().iterator + while (it.hasNext) { + val e = it.next() + out(toSparkValue(e.getKey)) = toSparkValue(e.getValue) + } + out.toMap + case t: org.apache.arrow.vector.util.Text => t.toString + case other => other + } +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala new file mode 100644 index 000000000..798bcedd4 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala @@ -0,0 +1,68 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.lance.index.DistanceType + +/** + * Vector distance / similarity metric. Mirrors `org.lance.index.DistanceType` but exposed as a + * Scala enumeration so callers don't have to import Lance internals. Each metric fixes the + * "best-first" direction used during merge: + * + * - L2: smaller score is better (distance) + * - Cosine / Dot: larger score is better (similarity) + */ +sealed trait Metric { + + /** The Lance distance type used when configuring a `Query`. */ + def lanceType: DistanceType + + /** True if smaller scores rank better (distance), false if larger (similarity). */ + def smallerIsBetter: Boolean +} + +object Metric { + + case object L2 extends Metric { + val lanceType: DistanceType = DistanceType.L2 + val smallerIsBetter: Boolean = true + } + + case object Cosine extends Metric { + val lanceType: DistanceType = DistanceType.Cosine + val smallerIsBetter: Boolean = false + } + + case object Dot extends Metric { + val lanceType: DistanceType = DistanceType.Dot + val smallerIsBetter: Boolean = false + } + + /** + * Parse a metric name. Accepts the same set of names Lance accepts plus a few synonyms commonly + * used in Spark vector functions: + * + * - "l2" | "euclidean" → L2 + * - "cosine" → Cosine + * - "dot" | "inner" | "ip" → Dot + */ + def fromName(name: String): Metric = name.trim.toLowerCase match { + case "l2" | "euclidean" => L2 + case "cosine" => Cosine + case "dot" | "inner" | "ip" => Dot + case other => + throw new IllegalArgumentException( + s"Unknown metric '$other'. Expected one of: l2, cosine, dot.") + } +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala new file mode 100644 index 000000000..ac1b194fc --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala @@ -0,0 +1,40 @@ +/* + * Licensed 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.lance.spark.knn.internal + +/** + * A reference to a single right-side row produced by a vector index probe, along with the ranking + * score. Carries no payload — payloads are fetched in the materialize stage by row address. Pairing + * a tiny ref with a score is the unit of work passing through the shuffle and is what keeps the + * shuffle volume to `O(|L| × tasks × K × ~24B)` instead of `O(|L| × tasks × K × payload_bytes)`. + * + * @param rowAddr Lance row address (`_rowaddr`): packed `(frag_id << 32) | row_in_frag`. Stable + * within a Lance dataset version. + * @param score Distance or similarity returned by Lance's vector search. Smaller-is-better for + * distance metrics (L2), larger-is-better for similarity metrics (cosine/dot). + * Direction is carried out-of-band in the operator config; this struct stays metric- + * agnostic. + */ +final case class ScoredRowRef(rowAddr: Long, score: Float) + +object ScoredRowRef { + + /** Order best-first for distance metrics (smallest score wins). */ + val distanceOrdering: Ordering[ScoredRowRef] = + Ordering.by[ScoredRowRef, Float](_.score) + + /** Order best-first for similarity metrics (largest score wins). */ + val similarityOrdering: Ordering[ScoredRowRef] = + Ordering.by[ScoredRowRef, Float](-_.score) +} diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala new file mode 100644 index 000000000..673179c8b --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala @@ -0,0 +1,114 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import scala.collection.mutable + +/** + * Bounded top-K heap with metric-aware ordering. Used by the probe stage for map-side combine + * across fragments owned by a single task — keeps the per-left-row state at exactly K entries no + * matter how many fragments contribute, and again on the reduce side to merge contributions from + * different tasks for the same `leftId`. + * + * Semantics: + * - `smallerIsBetter = true` (distance, e.g. L2): retain the K smallest-score entries. + * - `smallerIsBetter = false` (similarity, e.g. cosine): retain the K largest-score entries. + * + * Internally, the heap's *head* holds the worst surviving element so eviction is O(log K). Scala's + * `mutable.PriorityQueue` is a max-heap by the supplied `Ordering`, so the ordering is chosen to + * place "worst surviving" at the top: + * - distance → max-heap on `score` (largest score is worst) + * - similarity → max-heap on `-score` (smallest score is worst) + * + * Not thread-safe. Each left row in a probe stage gets its own heap. + */ +final class TopKHeap(k: Int, smallerIsBetter: Boolean) { + require(k > 0, "k must be positive") + + private val ord: Ordering[ScoredRowRef] = + if (smallerIsBetter) Ordering.by[ScoredRowRef, Float](_.score) + else Ordering.by[ScoredRowRef, Float](-_.score) + + private val heap = new mutable.PriorityQueue[ScoredRowRef]()(ord) + + /** + * Insert `ref` if it would survive the top-K cut. Either grows the heap up to K or evicts the + * current worst-surviving element if `ref` is strictly better than it. + */ + def offer(ref: ScoredRowRef): Unit = { + if (heap.size < k) { + heap.enqueue(ref) + } else { + val worst = heap.head + val isBetter = + if (smallerIsBetter) ref.score < worst.score + else ref.score > worst.score + if (isBetter) { + heap.dequeue() + heap.enqueue(ref) + } + } + } + + def offerAll(refs: TraversableOnce[ScoredRowRef]): Unit = refs.foreach(offer) + + /** + * Drain the heap into a best-first sorted Array. After this call the heap is empty. Best-first + * means index 0 is the top-ranked entry (smallest score for distance, largest for similarity). + */ + def drain(): Array[ScoredRowRef] = { + val out = new Array[ScoredRowRef](heap.size) + var i = heap.size - 1 + // PriorityQueue.dequeue returns the worst surviving element first; walking the array in + // reverse places best at index 0. + while (i >= 0) { + out(i) = heap.dequeue() + i -= 1 + } + out + } + + def size: Int = heap.size + def isEmpty: Boolean = heap.isEmpty +} + +object TopKHeap { + + /** + * Convenience: merge several already-sorted (best-first) ref arrays into one top-K array. Used + * by the merge stage as the `reduceByKey` combine function. + */ + def merge( + a: Array[ScoredRowRef], + b: Array[ScoredRowRef], + k: Int, + smallerIsBetter: Boolean): Array[ScoredRowRef] = { + if (a.isEmpty) return takeBest(b, k, smallerIsBetter) + if (b.isEmpty) return takeBest(a, k, smallerIsBetter) + val heap = new TopKHeap(k, smallerIsBetter) + heap.offerAll(a) + heap.offerAll(b) + heap.drain() + } + + private def takeBest( + arr: Array[ScoredRowRef], + k: Int, + smallerIsBetter: Boolean): Array[ScoredRowRef] = { + if (arr.length <= k) return arr + val heap = new TopKHeap(k, smallerIsBetter) + heap.offerAll(arr) + heap.drain() + } +} diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala new file mode 100644 index 000000000..f4a22bfe1 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala @@ -0,0 +1,462 @@ +/* + * Licensed 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.lance.spark.knn.catalyst + +import org.apache.spark.sql.{RowFactory, SparkSession} +import org.apache.spark.sql.catalyst.expressions.{Add, And, Attribute, AttributeSet, EqualTo, Expression, GetStructField, GreaterThan, In, IsNotNull, IsNull, LessThanOrEqual, Literal, Not, Or, VectorCosineSimilarity, VectorInnerProduct, VectorL2Distance} +import org.apache.spark.sql.catalyst.plans.Inner +import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, NearestByJoin, Project, SubqueryAlias} +import org.apache.spark.sql.catalyst.plans.{NearestByDistance, NearestBySimilarity} +import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation +import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String +import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} +import org.junit.jupiter.api.Assertions._ +import org.junit.jupiter.api.io.TempDir +import org.lance.spark.knn.internal.Metric + +import java.nio.file.Path + +import scala.collection.JavaConverters._ + +/** + * Unit tests for [[IndexedNearestByJoinRule]]. The rule's responsibility is purely Catalyst-side + * pattern-matching — we don't need a Lance backend to exercise it. Each test constructs a small + * resolved plan and runs the rule, asserting either a rewrite to + * `Project(..., LanceKnnJoinLogicalPlan(left, ...))` or a no-op fallthrough. + * + * Coverage: + * - Happy path: VectorL2Distance + NearestByDistance over a Lance DSv2 relation rewrites. + * - Direction mismatch (e.g. L2 distance with NearestBySimilarity) does NOT rewrite. + * - EXACT (`approx = false`) does NOT rewrite — Spark's brute-force keeps owning that path. + * - Non-Lance right side does NOT rewrite (duck-type check via class name). + * - Disabled by default — fires only when the gating config is set. + * - Prefilter pushdown: right-side `WHERE` translates to a Lance SQL filter string, or refuses + * the rewrite entirely when the predicate can't be pushed in full. + * + * The rule's runtime behavior beyond the rewrite (probe execution against real Lance) is covered + * by the oracle tests in lance-spark-knn_2.12 and the e2e test in this module. + */ +class IndexedNearestByJoinRuleTest { + + @TempDir var tempDir: Path = _ + private var spark: SparkSession = _ + + @BeforeEach def setup(): Unit = { + spark = SparkSession.builder() + .appName("indexed-nearest-by-join-rule-test") + .master("local[2]") + .config("spark.driver.bindAddress", "127.0.0.1") + .config("spark.driver.host", "127.0.0.1") + .getOrCreate() + } + + @AfterEach def teardown(): Unit = if (spark != null) spark.stop() + + /** L2 + NearestByDistance + Lance scan + enabled config → rewrite. */ + @Test def testL2RewritesToIndexedPlan(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val join = NearestByJoin( + left = left, + right = right, + joinType = Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + val plan = expectRewritten(rewritten) + assertEquals(Metric.L2, plan.metric) + assertEquals(5, plan.k) + assertEquals(rightVec.name, plan.rightVecCol) + assertEquals(leftVec.exprId, plan.leftVecAttr.exprId) + } + + /** Cosine similarity + NearestBySimilarity → rewrite. */ + @Test def testCosineRewrites(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "cosine") + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 3, + rankingExpression = VectorCosineSimilarity(leftVec, rightVec), + direction = NearestBySimilarity) + val rewritten = IndexedNearestByJoinRule(join) + assertEquals(Metric.Cosine, expectRewritten(rewritten).metric) + } + + /** Inner product + NearestBySimilarity → rewrite as Dot. */ + @Test def testDotRewrites(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "dot") + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 4, + rankingExpression = VectorInnerProduct(leftVec, rightVec), + direction = NearestBySimilarity) + val rewritten = IndexedNearestByJoinRule(join) + assertEquals(Metric.Dot, expectRewritten(rewritten).metric) + } + + /** L2 distance with NearestBySimilarity is inconsistent — rule should NOT fire. */ + @Test def testDirectionMismatchDoesNotRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestBySimilarity) + val rewritten = IndexedNearestByJoinRule(join) + assertSame(join, rewritten, "rule should not fire on direction/metric mismatch") + } + + /** EXACT mode (approx = false) is owned by Spark's brute-force rewrite. */ + @Test def testExactModeDoesNotRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val join = NearestByJoin( + left, + right, + Inner, + approx = false, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + assertSame(join, rewritten, "EXACT queries must not be intercepted") + } + + /** Disabled flag (default) → no rewrite even when otherwise applicable. */ + @Test def testDisabledByDefault(): Unit = { + spark.conf.unset(IndexedNearestByJoinRule.EnabledConfKey) + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + assertSame(join, rewritten, "rule must be opt-in") + } + + /** Non-Lance right side (regular DataFrame as Project, no DSv2 relation) → no rewrite. */ + @Test def testNonLanceRightDoesNotRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val left = trivialPlan("lid", "lvec") + val right = trivialPlan("rid", "rvec") + val leftVec = left.output.find(_.name == "lvec").get + val rightVec = right.output.find(_.name == "rvec").get + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + assertSame(join, rewritten, "non-Lance right must fall through") + } + + /** Right side wrapped in SubqueryAlias still rewrites — alias unwrapping happens in the rule. */ + @Test def testSubqueryAliasOnRightStillRewrites(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val aliased = SubqueryAlias("d", right) + val join = NearestByJoin( + left, + aliased, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + // Rule emits `Project(j.output, LanceKnnJoinLogicalPlan(left, ...))`. Asserting on the top + // Project wrapping the join node is enough for the "did the rule fire" check. + assertTrue( + rewritten.isInstanceOf[Project] && + rewritten.asInstanceOf[Project].child.isInstanceOf[LanceKnnJoinLogicalPlan], + s"expected Project(..., LanceKnnJoinLogicalPlan(...)), got: " + + s"${rewritten.getClass.getSimpleName}") + } + + // -- prefilter pushdown ------------------------------------------------------------------- + + /** + * Right side wrapped in `Filter(simple predicate)` rewrites AND the predicate lands on the + * indexed plan as a Lance SQL filter string. The filter must be pushed in full (not dropped) + * for the result to be semantically equivalent to the original plan. + */ + @Test def testFilterOverLancePushesAsPrefilter(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val category = right.output.find(_.name == "category").get + val bucket = right.output.find(_.name == "bucket").get + val cond = And( + EqualTo(category, Literal(UTF8String.fromString("A"), StringType)), + GreaterThan(bucket, Literal(5, IntegerType))) + val filtered = Filter(cond, right) + val join = NearestByJoin( + left, + filtered, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + val plan = expectRewritten(rewritten) + assertTrue(plan.prefilter.isDefined, "prefilter should be populated") + val sql = plan.prefilter.get + assertTrue(sql.contains("category"), s"prefilter missing column ref: $sql") + assertTrue(sql.contains("'A'"), s"prefilter missing string literal: $sql") + assertTrue(sql.contains("bucket"), s"prefilter missing column ref: $sql") + assertTrue(sql.contains("> 5"), s"prefilter missing numeric comparison: $sql") + assertTrue(sql.contains("AND"), s"prefilter missing conjunction: $sql") + } + + /** + * Predicate touches a left-side attribute — translator can't safely render that as a Lance + * SQL string (Lance only sees the right table's columns). Rule must REFUSE the rewrite, not + * drop the predicate. We verify the original `NearestByJoin` is returned unchanged. + */ + @Test def testPredicateReferencingLeftAttrRefusesRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val lid = left.output.find(_.name == "lid").get + val cond = EqualTo(lid, Literal(0, IntegerType)) + val filtered = Filter(cond, right) + val join = NearestByJoin( + left, + filtered, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + assertSame( + join, + rewritten, + "predicate touching left side must refuse pushdown — not partial-push") + } + + /** + * Predicate is a computed expression (e.g. `bucket + 1 = 6`), not a bare attr-vs-literal + * comparison. Translator returns None, rule refuses. + */ + @Test def testComputedPredicateRefusesRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val bucket = right.output.find(_.name == "bucket").get + val cond = EqualTo(Add(bucket, Literal(1, IntegerType)), Literal(6, IntegerType)) + val filtered = Filter(cond, right) + val join = NearestByJoin( + left, + filtered, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + assertSame(join, rewritten, "computed expression must refuse pushdown") + } + + /** Filter wrapped in SubqueryAlias still pushes — order of unwrap shouldn't matter. */ + @Test def testFilterUnderSubqueryAliasPushes(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val (left, leftVec, right, rightVec) = buildPlans(metricFunction = "l2") + val category = right.output.find(_.name == "category").get + val cond = EqualTo(category, Literal(UTF8String.fromString("X"), StringType)) + val plan = SubqueryAlias("d", Filter(cond, right)) + val join = NearestByJoin( + left, + plan, + Inner, + approx = true, + numResults = 3, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + val p = expectRewritten(rewritten) + assertTrue(p.prefilter.isDefined, s"prefilter should be set; got ${p.prefilter}") + } + + // -- predicate translator unit tests ----------------------------------------------------- + + /** + * Direct unit tests on `translateFilter` to lock in the supported shapes. Uses a synthetic + * AttributeSet so we don't need a logical plan. + */ + @Test def testTranslatorHandlesSupportedShapes(): Unit = { + val rid = makeAttr("rid", IntegerType) + val category = makeAttr("category", StringType) + val bucket = makeAttr("bucket", IntegerType) + val meta = makeAttr("meta", new StructType().add("category", StringType).add("bucket", IntegerType)) + val attrs = AttributeSet(Seq(rid, category, bucket, meta)) + + val cases: Seq[(Expression, String)] = Seq( + EqualTo(category, lit("A")) -> "category = 'A'", + Not(EqualTo(category, lit("A"))) -> "category != 'A'", + GreaterThan(bucket, lit(5)) -> "bucket > 5", + LessThanOrEqual(bucket, lit(5)) -> "bucket <= 5", + IsNull(category) -> "category IS NULL", + IsNotNull(category) -> "category IS NOT NULL", + In(bucket, Seq(lit(1), lit(2), lit(3))) -> "bucket IN (1, 2, 3)", + And(EqualTo(category, lit("A")), GreaterThan(bucket, lit(5))) -> + "(category = 'A') AND (bucket > 5)", + Or(EqualTo(category, lit("A")), EqualTo(category, lit("B"))) -> + "(category = 'A') OR (category = 'B')", + // String-literal escape — single quotes inside the value get doubled. + EqualTo(category, lit("O'Brien")) -> "category = 'O''Brien'", + // literal-on-left flip + EqualTo(lit(5), bucket) -> "5 = bucket", + // nested struct field access -> dotted path (distinct from the top-level `category` attr) + EqualTo(GetStructField(meta, 0, Some("category")), lit("A")) -> "meta.category = 'A'") + cases.foreach { case (expr, expected) => + val got = IndexedNearestByJoinRule.translateFilter(expr, attrs) + assertEquals(Some(expected), got, s"translation mismatch for: $expr") + } + } + + /** Translator must return None for unsupported expressions so the rule refuses pushdown. */ + @Test def testTranslatorRefusesUnsupportedShapes(): Unit = { + val rid = makeAttr("rid", IntegerType) + val ts = makeAttr("ts", DateType) // date literals not in our supported set + val foreignMeta = makeAttr("fmeta", new StructType().add("category", StringType)) + val attrs = AttributeSet(Seq(rid, ts)) + + val rejected: Seq[Expression] = Seq( + // Two attributes — no literal — translator can't render `attr op attr` safely (Lance can, + // but we don't promise it; refuse to keep the rule conservative). + EqualTo(rid, makeAttr("rid2", IntegerType)), + // Foreign attribute (not in `attrs`) — translator must reject. + EqualTo(makeAttr("foreign", IntegerType), lit(1)), + // Empty IN list. + In(rid, Seq.empty), + // Date literal — out of supported types. + EqualTo(ts, Literal(0, DateType)), + // Nested struct field over a FOREIGN root attr (not in `attrs`) — the recursion must gate + // on the root and refuse. (Array/map element access like `col[i]` refuses the same way, + // via the translator's catch-all.) + EqualTo(GetStructField(foreignMeta, 0, Some("category")), lit("A"))) + rejected.foreach { e => + assertEquals( + None, + IndexedNearestByJoinRule.translateFilter(e, attrs), + s"expected refusal for: $e") + } + } + + // -- helpers ------------------------------------------------------------------------------ + + /** + * Construct a left-side regular plan and a right-side that resembles a Lance DSv2 scan via the + * duck-type check. Avoids the need for a real Lance reader. + */ + private def buildPlans(metricFunction: String) + : (LogicalPlan, Attribute, LogicalPlan, Attribute) = { + val left = trivialPlan("lid", "lvec") + val rightLance = lanceLikeDsv2Relation() + val leftVec = left.output.find(_.name == "lvec").get + val rightVec = rightLance.output.find(_.name == "rvec").get + (left, leftVec, rightLance, rightVec) + } + + private def trivialPlan(idCol: String, vecCol: String): LogicalPlan = { + val schema = new StructType(Array( + StructField(idCol, IntegerType, nullable = false), + StructField(vecCol, ArrayType(FloatType, containsNull = false), nullable = false))) + val rows = (0 until 4).map(i => RowFactory.create(Integer.valueOf(i), Array.fill(8)(0.0f))) + spark.createDataFrame(rows.asJava, schema).queryExecution.analyzed + } + + /** + * Build a `DataSourceV2Relation` whose `table.getClass.getName.contains("Lance")` so the + * rule's duck-type check accepts it. We don't actually run any I/O. Includes a `category` + * (string) and `bucket` (int) column so prefilter-pushdown tests can build realistic + * filter predicates without needing to extend the schema separately. + */ + private def lanceLikeDsv2Relation(): LogicalPlan = { + val schema = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + StructField("category", StringType, nullable = true), + StructField("bucket", IntegerType, nullable = true), + StructField("rvec", ArrayType(FloatType, containsNull = false), nullable = false))) + val table = new FakeLanceTable(schema) + val opts = new java.util.HashMap[String, String]() + opts.put("path", tempDir.resolve("fake_lance").toString) + val cims = new org.apache.spark.sql.util.CaseInsensitiveStringMap(opts) + DataSourceV2Relation.create(table, None, None, cims) + } + + /** + * Extract an assertion-friendly summary of the rule's rewrite output. The rule produces + * `Project(j.output, LanceKnnJoinLogicalPlan(left, stageConf, ...))`; this helper pulls out the + * fields the test cases want to check straight off `stageConf`. + */ + private case class RewriteSummary( + metric: Metric, + k: Int, + rightVecCol: String, + leftVecAttr: Attribute, + prefilter: Option[String]) + + private def expectRewritten(plan: LogicalPlan): RewriteSummary = plan match { + case Project(_, node: LanceKnnJoinLogicalPlan) => + val conf = node.stageConf + RewriteSummary( + metric = conf.metric, + k = conf.k, + rightVecCol = conf.vectorColumn, + leftVecAttr = node.child.output(conf.leftVecIdx), + prefilter = conf.prefilter) + case other => + fail(s"expected Project(LanceKnnJoinLogicalPlan(...)), got: $other"); ??? + } + + private def makeAttr(name: String, dt: DataType): Attribute = + org.apache.spark.sql.catalyst.expressions.AttributeReference(name, dt, nullable = true)() + + private def lit(v: Int): Literal = Literal(v, IntegerType) + private def lit(s: String): Literal = Literal(UTF8String.fromString(s), StringType) +} + +/** + * Stub Table whose class name ends with "Lance" so the rule's duck-type check accepts it. No I/O + * — the rule only reads schema and options. Lives in the test source tree. + */ +class FakeLanceTable(_schema: StructType) extends org.apache.spark.sql.connector.catalog.Table { + override def name(): String = "fake_lance" + override def schema(): StructType = _schema + override def capabilities() + : java.util.Set[org.apache.spark.sql.connector.catalog.TableCapability] = + java.util.Collections.emptySet() +} diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala new file mode 100644 index 000000000..f19991d3c --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala @@ -0,0 +1,481 @@ +/* + * Licensed 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.lance.spark.knn.catalyst + +import org.apache.spark.sql.{DataFrame, Row, RowFactory, SparkSession} +import org.apache.spark.sql.types._ +import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} +import org.junit.jupiter.api.Assertions._ +import org.junit.jupiter.api.io.TempDir +import org.lance.spark.knn.internal.LanceVectorIndexBuilder +import org.lance.spark.knn.testutil.ClusteredEmbeddings + +import java.nio.file.Path +import java.util.Random + +import scala.collection.JavaConverters._ + +/** + * SQL end-to-end tests for the Catalyst integration. Every test drives the full path: + * + * ANTLR parser ─▶ Analyzer ─▶ IndexedNearestByJoinRule (our postHoc) ─▶ + * Optimizer ─▶ LanceKnnJoinStrategy ─▶ LanceKnnJoinExec ─▶ + * Lance native per-row probe + late materialize ─▶ Rows + * + * Requires Spark 4.2 (the release where `NearestByJoin` exists, added by SPARK-56395) AND the + * `lance-spark-4.2_2.13` connector built against the same Spark version. + * + * Two groups of tests, sharing the same SparkSession + Lance scaffolding: + * + * 1. EXACT path (no vector index) — Lance does an exact per-fragment scan, so results must + * equal the brute-force oracle exactly. Covers oracle equivalence, right-side `WHERE` + * prefilter pushdown, and the rule's opt-in gating (disabled → falls through to Spark's + * `RewriteNearestByJoin`, still correct). + * 2. APPROXIMATE path (IVF-PQ index) — Lance returns approximate top-K, so recall is < 1.0. + * Covers that the indexed path engages, recall stays in a sane range at default settings, + * and `refineFactor > 1` improves (or matches) recall. + * + * The rule's plan-side pattern-matching (metric/direction, alias/filter unwrapping, prefilter + * translation) is unit-tested separately in `IndexedNearestByJoinRuleTest`. + */ +class IndexedNearestByJoinSqlTest { + + @TempDir var tempDir: Path = _ + private var spark: SparkSession = _ + // Monotonic suffix so each SQL invocation gets fresh temp-view names (a single test builds + // more than one left/right pair). + private var viewSeq: Int = 0 + + // Exact-path scale — kept tiny for speed; no index, so size doesn't affect correctness. + private val ExactDim = 16 + private val ExactRight = 64 + private val ExactLeft = 8 + // Approximate-path scale — IVF-PQ needs more rows to be meaningful (1024 rows / 4 partitions + // ≈ 256 per cluster). Still small enough to run in a few seconds. + private val Dim = 32 + private val NumRight = 1024 + private val NumLeft = 32 + private val K = 10 + private val Seed = 0xCAFEL + + @BeforeEach def setup(): Unit = { + spark = SparkSession.builder() + .appName("indexed-nearest-by-join-sql") + .master("local[2]") + .config("spark.driver.bindAddress", "127.0.0.1") + .config("spark.driver.host", "127.0.0.1") + .config( + "spark.sql.extensions", + "org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions") + .config("spark.sql.crossJoin.enabled", "true") + .getOrCreate() + spark.sparkContext.setLogLevel("WARN") + // Enabled by default; the "rule disabled" test flips it off explicitly. + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + } + + @AfterEach def teardown(): Unit = if (spark != null) spark.stop() + + // -- exact path (no index): oracle equivalence, WHERE pushdown, rule-off fallthrough ------- + + /** + * Full SQL path with the rule enabled. The physical plan must contain the `LanceKnnJoin` exec + * AND the result must match the brute-force oracle on every left row (exact, no index). + */ + @Test def testSqlApproxNearestRoutesThroughIndexedPathAndMatchesOracle(): Unit = { + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed) + val (leftDf, leftIds, _) = buildLeftDf(leftVecs, ExactDim) + val rightVecs = generateUniform(ExactRight, ExactDim, Seed + 1) + val (rightUri, rightIds, _) = writeRightDf(rightVecs, ExactDim, idBase = 1000) + + val k = 5 + val (q, d) = registerViews(leftDf, rightUri) + val df = spark.sql( + s"""SELECT q.lid, d.rid + |FROM $q q INNER JOIN $d d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + + // Plan-shape: confirm the rule fired (logical node present, AQE-independent) AND the strategy + // lowered it to the `LanceKnnJoin` physical exec. + val joinLogicals = df.queryExecution.optimizedPlan.collect { + case p: LanceKnnJoinLogicalPlan => p + } + assertTrue( + joinLogicals.nonEmpty, + s"expected LanceKnnJoinLogicalPlan in optimized plan; got:\n${df.queryExecution.optimizedPlan}") + val tree = df.queryExecution.executedPlan.treeString + assertTrue(tree.contains("LanceKnnJoin"), s"expected LanceKnnJoin exec in tree:\n$tree") + + // Correctness: oracle equivalence. + val rows = df.collect() + assertEquals(ExactLeft * k, rows.length, "expected k results per left row") + val byLid = rows.groupBy(_.getAs[Int]("lid")) + leftIds.zip(leftVecs).foreach { case (lid, lvec) => + val oracle = oracleTopKIds(lvec, rightVecs.indices, rightIds, rightVecs, k) + val actual = byLid(lid).map(_.getAs[Int]("rid")).toSet + assertEquals(oracle, actual, s"top-K mismatch for lid=$lid (rule on, brute-force oracle)") + } + } + + /** + * Right-side `WHERE` clause must round-trip through the prefilter pushdown — Lance computes + * top-K only over rows matching the filter, so the result must equal the brute-force oracle + * computed AFTER applying the same filter. If the rule pushed the filter wrong (or dropped + * it), this test would diverge from the oracle. Two right-side rows share each `category`, so + * `WHERE category = 'A'` shrinks the candidate pool meaningfully without zeroing it out. + */ + @Test def testSqlWherePushdownMatchesFilteredOracle(): Unit = { + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed + 200) + val (leftDf, leftIds, _) = buildLeftDf(leftVecs, ExactDim) + val (rightVecs, rightIds, rightCategories, rightUri) = + writeRightWithCategories(ExactRight, ExactDim, Seed + 201) + + val k = 4 + val targetCat = "A" + val (q, d) = registerViews(leftDf, rightUri) + val df = spark.sql( + s"""SELECT q.lid, d.rid + |FROM $q q INNER JOIN (SELECT * FROM $d WHERE category = '$targetCat') d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + + val joinLogicals = df.queryExecution.optimizedPlan.collect { + case p: LanceKnnJoinLogicalPlan => p + } + assertTrue( + joinLogicals.nonEmpty, + s"expected LanceKnnJoinLogicalPlan; optimized plan was:\n${df.queryExecution.optimizedPlan}") + val prefilter = joinLogicals.head.stageConf.prefilter + assertTrue( + prefilter.exists(_.contains(s"'$targetCat'")), + s"expected prefilter to carry category='$targetCat'; got: $prefilter") + + // Oracle: brute-force top-K computed AFTER applying the same filter on the right side. + val filteredIdxs = rightCategories.indices.filter(rightCategories(_) == targetCat) + val rows = df.collect() + val byLid = rows.groupBy(_.getAs[Int]("lid")) + leftIds.zip(leftVecs).foreach { case (lid, lvec) => + val oracle = oracleTopKIds(lvec, filteredIdxs, rightIds, rightVecs, k) + assertTrue( + oracle.nonEmpty, + s"oracle is empty for lid=$lid — test setup didn't produce filterable rows") + val actual = byLid(lid).map(_.getAs[Int]("rid")).toSet + assertEquals( + oracle, + actual, + s"top-K mismatch under WHERE pushdown for lid=$lid (filtered brute-force oracle)") + } + } + + /** + * With the gating config disabled, the SAME SQL falls through to Spark's + * `RewriteNearestByJoin` (cross-product + `MaxMinByK`). The plan contains NO + * `LanceKnnJoinLogicalPlan` and (importantly) results still match the oracle — proving the + * rule's opt-in behavior: turning it off doesn't break correctness. + */ + @Test def testSqlFallsThroughToBruteForceWhenRuleDisabled(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "false") + + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed + 100) + val (leftDf, leftIds, _) = buildLeftDf(leftVecs, ExactDim) + val rightVecs = generateUniform(ExactRight, ExactDim, Seed + 101) + val (rightUri, rightIds, _) = writeRightDf(rightVecs, ExactDim, idBase = 1000) + + val k = 4 + val (q, d) = registerViews(leftDf, rightUri) + val df = spark.sql( + s"""SELECT q.lid, d.rid + |FROM $q q INNER JOIN $d d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + + val joinLogicals = df.queryExecution.optimizedPlan.collect { + case p: LanceKnnJoinLogicalPlan => p + } + assertTrue( + joinLogicals.isEmpty, + s"rule disabled — expected NO LanceKnnJoinLogicalPlan; got:\n${df.queryExecution.optimizedPlan}") + + val rows = df.collect() + assertEquals(ExactLeft * k, rows.length) + val byLid = rows.groupBy(_.getAs[Int]("lid")) + leftIds.zip(leftVecs).foreach { case (lid, lvec) => + val oracle = oracleTopKIds(lvec, rightVecs.indices, rightIds, rightVecs, k) + val actual = byLid(lid).map(_.getAs[Int]("rid")).toSet + assertEquals(oracle, actual, s"top-K mismatch for lid=$lid (rule off, brute-force fallback)") + } + } + + // -- approximate path (IVF-PQ): recall floors, refineFactor ------------------------------- + + /** + * Build IVF-PQ, run the SQL `APPROX NEAREST`, measure recall@10 against the brute-force oracle. + * With 1024 rows × 4 IVF partitions, each partition holds ~256 rows; a default-`nprobes` query + * hits ~1 partition, so recall should be substantially > 0 but below 1.0. A threshold this + * loose (0.3) would only fail on a real bug (index path not engaging), not IVF's inherent + * approximation. + */ + @Test def testIvfPqRecallReasonableAtDefaults(): Unit = { + val leftVecs = generateUniform(NumLeft, Dim, Seed) + val (leftDf, leftIds, _) = buildLeftDf(leftVecs, Dim) + val rightVecs = generateUniform(NumRight, Dim, Seed + 1) + val (rightUri, rightIds, _) = writeRightDf(rightVecs, Dim, idBase = 100000) + LanceVectorIndexBuilder.buildIvfPq( + datasetUri = rightUri, + vectorColumn = "rvec", + numPartitions = 4, + numSubVectors = 8, + numBits = 8) + assertEquals( + 1, + LanceVectorIndexBuilder.listIndexCount(rightUri), + "expected exactly one index after build") + + val rows = runKnnSql(leftDf, rightUri, K, refineFactor = None) + val recall = computeRecallAtK(rows, leftIds, leftVecs, rightIds, rightVecs, K) + println(s" IVF-PQ recall@$K (no refine, default nprobes): $recall") + assertTrue(recall > 0.3, s"recall@$K=$recall too low; index path probably not engaging") + } + + /** + * Production-realistic distribution: clustered Gaussian mixture, unit-sphere-normalized — the + * geometry of typical sentence-transformer / image-feature embeddings (uniform-random vectors + * are IVF's WORST case; k-means has no cluster structure to latch onto). Asserts clustered + * recall@K >= 0.5 at default IVF-PQ settings; prints both uniform and clustered so a reviewer + * can see the realistic case helps. + * + * We don't `assert(clustered >= uniform)`: Lance's k-means init is non-deterministic across JVM + * sessions, so on a tiny 1024-row dataset run-to-run noise routinely exceeds the structural + * advantage. A reliable comparison would need many seeds or much larger N — we chose to print + * both and assert only the realistic-floor invariant. + */ + @Test def testClusteredEmbeddingsRecallSurvives(): Unit = { + val (uniformDf, uniformIds, uniformVecs) = + buildLeftDf(generateUniform(NumLeft, Dim, Seed), Dim) + val (uniformUri, uniformRightIds, uniformRightVecs) = + writeRightDf(generateUniform(NumRight, Dim, Seed + 1), Dim, idBase = 100000) + LanceVectorIndexBuilder.buildIvfPq(uniformUri, "rvec", numPartitions = 4, numSubVectors = 8) + + val (clusteredDf, clusteredIds, clusteredVecs) = buildLeftDf( + ClusteredEmbeddings.generate(NumLeft, Dim, numClusters = 4, seed = Seed + 2), Dim) + val (clusteredUri, clusteredRightIds, clusteredRightVecs) = writeRightDf( + ClusteredEmbeddings.generate(NumRight, Dim, numClusters = 16, seed = Seed + 3), + Dim, + idBase = 100000) + LanceVectorIndexBuilder.buildIvfPq(clusteredUri, "rvec", numPartitions = 4, numSubVectors = 8) + + val uniformRecall = recallAgainst( + uniformDf, uniformUri, uniformIds, uniformVecs, uniformRightIds, uniformRightVecs) + val clusteredRecall = recallAgainst( + clusteredDf, clusteredUri, clusteredIds, clusteredVecs, clusteredRightIds, clusteredRightVecs) + println( + s" IVF-PQ recall@$K: uniform=$uniformRecall, clustered=$clusteredRecall " + + "(uniform = IVF worst case; clustered = production-shaped)") + + assertTrue( + clusteredRecall >= 0.5, + s"clustered-data recall@$K=$clusteredRecall is unexpectedly low; " + + "defaults should comfortably exceed 0.5 on production-shaped embeddings — " + + "if this fails, suspect a regression in Lance's index path or in our probe wiring") + } + + /** + * `refineFactor > 1` engages Lance's exact-distance re-rank: fetch `K * refineFactor` + * approximate candidates, re-rank, trim back to K. Strictly improves (or matches) recall vs. no + * refine. We assert `>=` rather than a strict `>` so the test isn't flaky on tiny datasets where + * both paths find the same K rows. The knob is set through `spark.lance.knn.refineFactor`. + */ + @Test def testRefineFactorImprovesRecall(): Unit = { + val leftVecs = generateUniform(NumLeft, Dim, Seed) + val (leftDf, leftIds, _) = buildLeftDf(leftVecs, Dim) + val rightVecs = generateUniform(NumRight, Dim, Seed + 1) + val (rightUri, rightIds, _) = writeRightDf(rightVecs, Dim, idBase = 100000) + LanceVectorIndexBuilder.buildIvfPq(rightUri, "rvec", numPartitions = 4) + + val baselineRows = runKnnSql(leftDf, rightUri, K, refineFactor = None) + val refinedRows = runKnnSql(leftDf, rightUri, K, refineFactor = Some(8)) + + val recallBaseline = computeRecallAtK(baselineRows, leftIds, leftVecs, rightIds, rightVecs, K) + val recallRefined = computeRecallAtK(refinedRows, leftIds, leftVecs, rightIds, rightVecs, K) + println(s" IVF-PQ recall@$K: no refine = $recallBaseline, refineFactor=8 = $recallRefined") + assertTrue( + recallRefined >= recallBaseline, + s"refineFactor should not hurt recall: baseline=$recallBaseline, refined=$recallRefined") + } + + // -- helpers ------------------------------------------------------------------------------ + + /** Register the left DataFrame and the right Lance dataset as fresh temp views. */ + private def registerViews(leftDf: DataFrame, rightUri: String): (String, String) = { + viewSeq += 1 + val q = s"queries_$viewSeq" + val d = s"docs_$viewSeq" + leftDf.createOrReplaceTempView(q) + spark.read.format("lance").load(rightUri).createOrReplaceTempView(d) + (q, d) + } + + /** + * Run a simple `INNER JOIN ... APPROX NEAREST k` (no WHERE) through the indexed rule and collect. + * `refineFactor` is threaded via the `spark.lance.knn.refineFactor` config the rule reads + * (unset => Lance default = no re-rank). + */ + private def runKnnSql( + leftDf: DataFrame, + rightUri: String, + k: Int, + refineFactor: Option[Int]): Array[Row] = { + val (q, d) = registerViews(leftDf, rightUri) + refineFactor match { + case Some(rf) => spark.conf.set(IndexedNearestByJoinRule.RefineFactorConfKey, rf.toString) + case None => spark.conf.unset(IndexedNearestByJoinRule.RefineFactorConfKey) + } + spark.sql( + s"""SELECT q.lid, d.rid + |FROM $q q INNER JOIN $d d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin).collect() + } + + private def leftSchema(dim: Int): StructType = new StructType(Array( + StructField("lid", IntegerType, nullable = false), + fixedSizeVec("lvec", dim))) + + private def rightSchema(dim: Int): StructType = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + fixedSizeVec("rvec", dim))) + + private def rightSchemaWithCategories(dim: Int): StructType = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + StructField("category", StringType, nullable = false), + fixedSizeVec("rvec", dim))) + + private def fixedSizeVec(name: String, dim: Int): StructField = + StructField( + name, + ArrayType(FloatType, containsNull = false), + nullable = false, + new MetadataBuilder().putLong("arrow.fixed-size-list.size", dim.toLong).build()) + + /** Build a left query DataFrame (lid, lvec) from pre-generated vectors. */ + private def buildLeftDf( + vectors: Array[Array[Float]], + dim: Int): (DataFrame, Array[Int], Array[Array[Float]]) = { + val ids = vectors.indices.toArray + val rows = ids.zip(vectors).map { case (id, v) => RowFactory.create(Integer.valueOf(id), v) } + (spark.createDataFrame(rows.toSeq.asJava, leftSchema(dim)), ids, vectors) + } + + /** Write a right Lance dataset (rid, rvec) from pre-generated vectors. */ + private def writeRightDf( + vectors: Array[Array[Float]], + dim: Int, + idBase: Int): (String, Array[Int], Array[Array[Float]]) = { + val ids = vectors.indices.map(_ + idBase).toArray + val rows = ids.zip(vectors).map { case (id, v) => RowFactory.create(Integer.valueOf(id), v) } + val df = spark.createDataFrame(rows.toSeq.asJava, rightSchema(dim)) + val out = tempDir.resolve(s"right_${System.nanoTime()}").toString + df.write.format("lance").save(out) + (out, ids, vectors) + } + + /** + * Write a right Lance dataset whose rows also carry a `category` from a small alphabet, so the + * WHERE-pushdown test has a non-trivial filter to apply. Returns the vectors, ids, per-row + * categories, and the dataset URI. + */ + private def writeRightWithCategories( + n: Int, + dim: Int, + seed: Long): (Array[Array[Float]], Array[Int], Array[String], String) = { + val vectors = generateUniform(n, dim, seed) + val ids = vectors.indices.map(_ + 2000).toArray + val alphabet = Array("A", "B", "C", "D") + val categories = vectors.indices.map(i => alphabet(i % alphabet.length)).toArray + val rows = ids.zip(vectors).zip(categories).map { case ((id, v), cat) => + RowFactory.create(Integer.valueOf(id), cat, v) + } + val df = spark.createDataFrame(rows.toSeq.asJava, rightSchemaWithCategories(dim)) + val out = tempDir.resolve(s"right_cat_${System.nanoTime()}").toString + df.write.format("lance").save(out) + (vectors, ids, categories, out) + } + + /** Run an indexed nearest join (SQL path) against the given right dataset and compute recall@K. */ + private def recallAgainst( + leftDf: DataFrame, + rightUri: String, + leftIds: Array[Int], + leftVecs: Array[Array[Float]], + rightIds: Array[Int], + rightVecs: Array[Array[Float]]): Double = { + val rows = runKnnSql(leftDf, rightUri, K, refineFactor = None) + computeRecallAtK(rows, leftIds, leftVecs, rightIds, rightVecs, K) + } + + /** Uniform-random vectors over the unit hypercube — the IVF-worst-case data distribution. */ + private def generateUniform(n: Int, dim: Int, seed: Long): Array[Array[Float]] = { + val rng = new Random(seed) + Array.fill(n)(randomVector(rng, dim)) + } + + /** + * Brute-force top-K right ids for a query vector, over the given candidate right indices (pass + * `rightVecs.indices` for the whole dataset, or a filtered subset for WHERE-pushdown oracles). + */ + private def oracleTopKIds( + lvec: Array[Float], + candidateIdxs: Seq[Int], + rightIds: Array[Int], + rightVecs: Array[Array[Float]], + k: Int): Set[Int] = + candidateIdxs + .map(i => (rightIds(i), l2(lvec, rightVecs(i)))) + .sortBy(_._2) + .take(k) + .map(_._1) + .toSet + + /** + * Mean recall@K across all left rows: |indexed top-K ∩ brute-force top-K| / K. 1.0 means the + * indexed path returned the same K rows as brute force; lower means the IVF cluster cut excluded + * some true neighbors. + */ + private def computeRecallAtK( + joinedRows: Array[Row], + leftIds: Array[Int], + leftVecs: Array[Array[Float]], + rightIds: Array[Int], + rightVecs: Array[Array[Float]], + k: Int): Double = { + val byLid = joinedRows.groupBy(_.getAs[Int]("lid")) + val perLidRecall = leftIds.zip(leftVecs).map { case (lid, lvec) => + val oracle = oracleTopKIds(lvec, rightVecs.indices, rightIds, rightVecs, k) + val actual = byLid.getOrElse(lid, Array.empty).map(_.getAs[Int]("rid")).toSet + (oracle intersect actual).size.toDouble / k + } + perLidRecall.sum / perLidRecall.length + } + + private def randomVector(rng: Random, dim: Int): Array[Float] = { + val v = new Array[Float](dim) + var i = 0 + while (i < dim) { v(i) = rng.nextFloat(); i += 1 } + v + } + + private def l2(a: Array[Float], b: Array[Float]): Float = { + var s = 0.0f + var i = 0 + while (i < a.length) { val d = a(i) - b(i); s += d * d; i += 1 } + s + } +} diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala new file mode 100644 index 000000000..0ae287138 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala @@ -0,0 +1,202 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.apache.spark.sql.{Row, RowFactory, SparkSession} +import org.apache.spark.sql.types._ +import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} +import org.junit.jupiter.api.Assertions._ +import org.junit.jupiter.api.io.TempDir + +import java.nio.file.Path +import java.util.Random + +import scala.collection.JavaConverters._ + +/** + * End-to-end validation of [[LanceProbe]] against a real Lance dataset written by Spark. These are + * the day-1 validation tasks the implementation plan calls out: + * + * - Per-probe call should succeed and return Lance's nearest neighbors. + * - Repeated probes against the same `LanceProbe` instance should reuse the open dataset + * handle; the second call should not re-pay the dataset open cost. + * - `fragmentIds` restriction should narrow the search to specified fragments only. + * - Without an explicit vector index the probe falls back to a brute-force per-fragment scan, + * which gives recall = 1.0 — making the no-index path the natural correctness oracle. + * + * These tests do NOT require an actual vector index; that is exercised in the indexed test + * suites which build IVF-PQ via Lance's index DDL. Validating the brute-force path first lets us + * isolate any LanceProbe bugs from index-quality issues. + */ +class LanceProbeValidationTest { + + @TempDir var tempDir: Path = _ + private var spark: SparkSession = _ + + // Small synthetic dataset: 64 vectors, dim 8. Enough to exercise the probe loop without making + // the test slow. + private val NumRows = 64 + private val VectorDim = 8 + private val Seed = 42L + + @BeforeEach def setup(): Unit = { + spark = SparkSession.builder() + .appName("lance-probe-validation") + .master("local[2]") + // Pin the driver to loopback so test JVMs in restricted networks (CI sandboxes, dev + // containers) can bind without scanning the host's interfaces. + .config("spark.driver.bindAddress", "127.0.0.1") + .config("spark.driver.host", "127.0.0.1") + .getOrCreate() + } + + @AfterEach def teardown(): Unit = { + if (spark != null) spark.stop() + } + + /** + * Smoke test: write a dataset, probe it, get K rows back. No correctness assertion beyond + * "result has the right shape" — the brute-force-equivalence test below covers semantics. + */ + @Test def testProbeReturnsKResults(): Unit = { + val datasetUri = writeSyntheticDataset() + val query = randomVector(new Random(7L), VectorDim) + + val probe = new LanceProbe(datasetUri, fragmentIds = None) + try { + val results = probe.probe(vectorColumn = "vec", query, k = 5, metric = Metric.L2) + assertEquals(5, results.size, "probe should return exactly k results") + // Distances must be monotonically non-decreasing for L2 (best-first). + val scores = results.map(_.score) + assertEquals(scores, scores.sorted, "L2 results should be sorted ascending by distance") + // Row addresses are stable u64s; we just sanity-check they aren't all zero. + assertTrue(results.exists(_.rowAddr != 0L), "row addresses should be populated") + } finally probe.close() + } + + /** + * Without a vector index, Lance does an exact per-fragment scan. That makes it a recall = 1.0 + * oracle: the probe result should equal the ground-truth top-K computed in plain Scala. + */ + @Test def testProbeMatchesBruteForceOracle(): Unit = { + val rng = new Random(Seed) + val (rows, vectors) = generateRows(rng, NumRows, VectorDim) + val datasetUri = writeRows(rows) + + val query = randomVector(new Random(123L), VectorDim) + val k = 10 + + val oracle: Seq[(Int, Float)] = vectors.zipWithIndex + .map { case (v, idx) => (idx, l2Distance(query, v)) } + .sortBy(_._2) + .take(k) + + val probe = new LanceProbe(datasetUri, fragmentIds = None) + val actual = + try probe.probe("vec", query, k, Metric.L2) + finally probe.close() + + assertEquals(k, actual.size) + // Compare scores within float tolerance. + val expectedScores = oracle.map(_._2) + val actualScores = actual.map(_.score) + expectedScores.zip(actualScores).foreach { case (expected, actualScore) => + assertEquals( + expected, + actualScore, + 1e-4f, + s"top-K distance mismatch: oracle=$expectedScores actual=$actualScores") + } + } + + /** + * Validate the dataset handle is reused across calls. The exact perf invariant ("second call + * faster than first by some factor") is too brittle for CI, so we only assert that repeated + * probes succeed and don't OOM — i.e., no JNI handle / Arrow buffer leak per call. + */ + @Test def testRepeatedProbesShareDatasetHandle(): Unit = { + val datasetUri = writeSyntheticDataset() + val probe = new LanceProbe(datasetUri, None) + try { + val rng = new Random(99L) + val k = 4 + var i = 0 + while (i < 50) { + val results = probe.probe("vec", randomVector(rng, VectorDim), k, Metric.L2) + assertEquals(k, results.size, s"iteration $i returned wrong size") + i += 1 + } + } finally probe.close() + } + + /** Empty fragment-id list ⇒ no rows match. Confirms the pushdown actually narrows search. */ + @Test def testEmptyFragmentRestrictionReturnsNothing(): Unit = { + val datasetUri = writeSyntheticDataset() + val probe = new LanceProbe(datasetUri, Some(Seq.empty)) + try { + val results = probe.probe("vec", randomVector(new Random(1L), VectorDim), 5, Metric.L2) + assertTrue(results.isEmpty, s"empty fragmentIds should yield no results, got ${results.size}") + } finally probe.close() + } + + // -- helpers ------------------------------------------------------------------------------ + + /** Write a fresh dataset and return its file:// URI. */ + private def writeSyntheticDataset(): String = { + val rng = new Random(Seed) + val (rows, _) = generateRows(rng, NumRows, VectorDim) + writeRows(rows) + } + + private def writeRows(rows: Seq[Row]): String = { + val schema = new StructType(Array( + StructField("id", IntegerType, nullable = false), + StructField( + "vec", + ArrayType(FloatType, containsNull = false), + nullable = false, + new MetadataBuilder().putLong("arrow.fixed-size-list.size", VectorDim.toLong).build()))) + val df = spark.createDataFrame(rows.asJava, schema) + + val outDir = tempDir.resolve(s"probe_test_${System.nanoTime()}").toString + df.write.format("lance").save(outDir) + outDir + } + + private def generateRows(rng: Random, n: Int, dim: Int): (Seq[Row], Seq[Array[Float]]) = { + val vectors = (0 until n).map(_ => randomVector(rng, dim)) + val rows = vectors.zipWithIndex.map { case (v, idx) => + RowFactory.create(Integer.valueOf(idx), v) + } + (rows, vectors) + } + + private def randomVector(rng: Random, dim: Int): Array[Float] = { + val v = new Array[Float](dim) + var i = 0 + while (i < dim) { v(i) = rng.nextFloat(); i += 1 } + v + } + + private def l2Distance(a: Array[Float], b: Array[Float]): Float = { + var s = 0.0f + var i = 0 + while (i < a.length) { + val d = a(i) - b(i) + s += d * d + i += 1 + } + s + } +} diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceVectorIndexBuilder.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceVectorIndexBuilder.scala new file mode 100644 index 000000000..2a8f319c7 --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceVectorIndexBuilder.scala @@ -0,0 +1,105 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.lance.{Dataset, ReadOptions} +import org.lance.index.{IndexOptions, IndexParams, IndexType} +import org.lance.index.vector.VectorIndexParams +import org.lance.spark.LanceRuntime + +import scala.collection.JavaConverters._ + +/** + * Test-only helper to build an IVF-PQ vector index on a Lance dataset via + * `Dataset.createIndex`. Exists so recall tests can construct the indexed scan path + * without writing the Lance Java boilerplate inline. + * + * Lives in `src/test/scala` because the production code path doesn't need to build + * indexes — users build them via Lance's Python / Rust / SQL DDL on their own datasets, + * and we just probe whatever's there. The helper exists for closed-loop recall validation. + */ +object LanceVectorIndexBuilder { + + /** + * Build an IVF-PQ index on `vectorColumn` of the dataset at `datasetUri`. Defaults are + * tuned for tiny test datasets — production users would size these much larger. + * + * @param numPartitions IVF cluster count. Should divide cleanly into the dataset row count. + * For a 4K-row dataset, 4-8 partitions is reasonable. + * @param numSubVectors PQ sub-vector count. Must divide vector dim evenly. + * @param numBits PQ bits per sub-vector. 8 is the standard. + * @param metric distance type. Must match the metric used at probe time. + * @param maxIters KMeans iteration cap during IVF training. 50 is enough for tests. + */ + def buildIvfPq( + datasetUri: String, + vectorColumn: String, + numPartitions: Int = 4, + numSubVectors: Int = 8, + numBits: Int = 8, + metric: Metric = Metric.L2, + maxIters: Int = 50): Unit = { + val dataset = openDataset(datasetUri) + try { + // Arg order in lance-core is (numPartitions, numBits, numSubVectors, distanceType, maxIters) + // — numBits precedes numSubVectors. Both default to 8 here so a swap is silent; pin the + // documented order explicitly. + val vectorParams = + VectorIndexParams.ivfPq(numPartitions, numBits, numSubVectors, metric.lanceType, maxIters) + val indexParams = IndexParams.builder().setVectorIndexParams(vectorParams).build() + val opts = IndexOptions + .builder(java.util.Collections.singletonList(vectorColumn), IndexType.VECTOR, indexParams) + .build() + dataset.createIndex(opts) + } finally dataset.close() + } + + /** + * Build an IVF_FLAT index — IVF clustering without PQ compression. Exact distances within + * visited clusters (no PQ noise), so recall depends purely on `nprobes` coverage. Higher + * memory/disk footprint than IVF-PQ (full vectors stored per cluster) but better recall on + * high-dim or random workloads where PQ compression drops too much information. + */ + def buildIvfFlat( + datasetUri: String, + vectorColumn: String, + numPartitions: Int = 4, + metric: Metric = Metric.L2): Unit = { + val dataset = openDataset(datasetUri) + try { + val vectorParams = VectorIndexParams.ivfFlat(numPartitions, metric.lanceType) + val indexParams = IndexParams.builder().setVectorIndexParams(vectorParams).build() + val opts = IndexOptions + .builder(java.util.Collections.singletonList(vectorColumn), IndexType.VECTOR, indexParams) + .build() + dataset.createIndex(opts) + } finally dataset.close() + } + + private def openDataset(uri: String): Dataset = { + Dataset + .open() + .uri(uri) + .allocator(LanceRuntime.allocator()) + .readOptions(new ReadOptions.Builder().build()) + .build() + } + + /** Number of indexes on the dataset (sanity check after building). */ + def listIndexCount(datasetUri: String): Int = { + val dataset = openDataset(datasetUri) + try dataset.listIndexes.asScala.size + finally dataset.close() + } +} diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala new file mode 100644 index 000000000..41bd1de7d --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala @@ -0,0 +1,92 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.junit.jupiter.api.Assertions._ +import org.junit.jupiter.api.Test + +/** + * Unit tests for [[TopKHeap]]. The heap's correctness is the foundation of the merge stage — + * any off-by-one or wrong-direction ordering would silently corrupt top-K results. We test both + * metric directions explicitly. + */ +class TopKHeapTest { + + private def ref(addr: Long, score: Float): ScoredRowRef = ScoredRowRef(addr, score) + + /** Distance metric: smaller score is better. Top-K must hold the K smallest. */ + @Test def testDistanceKeepsKSmallest(): Unit = { + val heap = new TopKHeap(k = 3, smallerIsBetter = true) + Seq(5.0f, 1.0f, 4.0f, 2.0f, 8.0f, 0.5f).zipWithIndex.foreach { case (s, i) => + heap.offer(ref(i.toLong, s)) + } + val out = heap.drain() + val scores = out.map(_.score).toSeq + assertEquals(Seq(0.5f, 1.0f, 2.0f), scores, "distance heap should retain three smallest") + } + + /** Similarity metric: larger score is better. Top-K must hold the K largest. */ + @Test def testSimilarityKeepsKLargest(): Unit = { + val heap = new TopKHeap(k = 3, smallerIsBetter = false) + Seq(5.0f, 1.0f, 4.0f, 2.0f, 8.0f, 0.5f).zipWithIndex.foreach { case (s, i) => + heap.offer(ref(i.toLong, s)) + } + val out = heap.drain() + val scores = out.map(_.score).toSeq + assertEquals(Seq(8.0f, 5.0f, 4.0f), scores, "similarity heap should retain three largest") + } + + /** Drain order is best-first regardless of insertion order. */ + @Test def testDrainOrderIsBestFirst(): Unit = { + val heap = new TopKHeap(k = 4, smallerIsBetter = true) + heap.offerAll(Seq(ref(1, 9f), ref(2, 1f), ref(3, 5f), ref(4, 3f), ref(5, 2f))) + val drained = heap.drain() + val scores = drained.map(_.score).toSeq + assertEquals(Seq(1f, 2f, 3f, 5f), scores) + assertTrue(heap.isEmpty, "drain should leave the heap empty") + } + + /** Heap with fewer than K elements drains them all in best-first order. */ + @Test def testFewerThanKReturnsAll(): Unit = { + val heap = new TopKHeap(k = 10, smallerIsBetter = true) + heap.offerAll(Seq(ref(1, 3f), ref(2, 1f), ref(3, 2f))) + assertEquals(Seq(1f, 2f, 3f), heap.drain().map(_.score).toSeq) + } + + /** A worse-than-current-worst candidate is rejected. */ + @Test def testRejectsWorseCandidate(): Unit = { + val heap = new TopKHeap(k = 2, smallerIsBetter = true) + heap.offer(ref(1, 1f)) + heap.offer(ref(2, 2f)) + heap.offer(ref(3, 5f)) // worse than existing 2 → rejected + val drained = heap.drain() + assertEquals(Seq(1f, 2f), drained.map(_.score).toSeq) + assertEquals(Seq(1L, 2L), drained.map(_.rowAddr).toSeq) + } + + /** `merge` combines two pre-sorted arrays preserving top-K. */ + @Test def testMergeCombinesTwoArrays(): Unit = { + val a = Array(ref(1, 1f), ref(2, 3f), ref(3, 5f)) + val b = Array(ref(4, 2f), ref(5, 4f), ref(6, 6f)) + val merged = TopKHeap.merge(a, b, k = 4, smallerIsBetter = true) + assertEquals(Seq(1f, 2f, 3f, 4f), merged.map(_.score).toSeq) + } + + /** Merging with one empty input is a noop modulo trim to K. */ + @Test def testMergeWithEmpty(): Unit = { + val a = Array(ref(1, 1f), ref(2, 2f), ref(3, 3f)) + val merged = TopKHeap.merge(a, Array.empty[ScoredRowRef], k = 2, smallerIsBetter = true) + assertEquals(Seq(1f, 2f), merged.map(_.score).toSeq) + } +} diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/testutil/ClusteredEmbeddings.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/testutil/ClusteredEmbeddings.scala new file mode 100644 index 000000000..47ccc3c6d --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/testutil/ClusteredEmbeddings.scala @@ -0,0 +1,137 @@ +/* + * Licensed 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.lance.spark.knn.testutil + +import java.util.Random + +/** + * Generate a clustered Gaussian-mixture embedding sample as a stand-in for real production + * embeddings (SIFT / sentence-transformer / image features). Real embeddings are not uniform + * over the unit hypercube — they cluster around a small number of topic centroids with each + * cluster occupying a relatively narrow region of the space. Uniform-random vectors are the + * worst case for IVF: there's no natural cluster structure for k-means to latch onto, so the + * IVF partitions cover the space arbitrarily and the per-cluster recall is essentially random. + * + * Method: + * 1. Pick `numClusters` cluster centers, each drawn uniformly from the unit hypercube. + * 2. For each row, pick a cluster (round-robin so each cluster gets equal mass) and sample + * a Gaussian centered on it with standard deviation `sigma * cluster_separation`. + * 3. L2-normalize so vectors live on the unit sphere — the natural geometry for cosine / + * inner-product retrieval, and what most production embedding models produce. + * + * The cluster-separation factor is the median pairwise distance between centers; scaling sigma + * by it keeps the cluster radius proportional to inter-cluster spacing regardless of `dim` or + * `numClusters`. With sigma ≈ 0.15 the clusters overlap a little but stay distinguishable — + * a reasonable proxy for production embedding distributions. + * + * The generator is deterministic given the seed so test runs are reproducible. + */ +object ClusteredEmbeddings { + + /** + * Build a clustered-Gaussian-mixture sample. + * + * @param n number of vectors to generate + * @param dim vector dimension + * @param numClusters number of cluster centers (small relative to `n` — typical 16-64) + * @param sigma per-cluster standard deviation, in units of inter-cluster distance. + * 0.05 = tight clusters (high recall floor); 0.5 = loose, near-uniform + * @param seed RNG seed for reproducibility + * @return an array of `n` float vectors of dimension `dim`, L2-normalized + */ + def generate( + n: Int, + dim: Int, + numClusters: Int, + sigma: Double = 0.15, + seed: Long = 0L): Array[Array[Float]] = { + require(n > 0 && dim > 0 && numClusters > 0, "n, dim, numClusters must all be positive") + require(numClusters <= n, "numClusters cannot exceed n") + val rng = new Random(seed) + + // Step 1: cluster centers, uniform on [0, 1]^dim. Stored as Doubles so the noise pass keeps + // numerical headroom — L2 normalization at the end folds back to Float precision. + val centers = Array.fill(numClusters)(Array.fill(dim)(rng.nextDouble())) + + // Step 2: median pairwise distance between centers, used to scale sigma. We don't want sigma + // expressed in absolute distance units — the right notion is "fraction of cluster spacing," + // which keeps clustering tightness behavior stable across (dim, numClusters) settings. + val sep = medianPairwiseDistance(centers) + val scaledSigma = sigma * sep + + // Step 3: sample each row from a Gaussian centered on a round-robin cluster. Round-robin + // (rather than uniformly random cluster choice) gives every cluster the same mass — a more + // controlled benchmark setup than letting some clusters get sparsely populated. + val out = new Array[Array[Float]](n) + var i = 0 + while (i < n) { + val center = centers(i % numClusters) + val v = new Array[Float](dim) + var d = 0 + while (d < dim) { + v(d) = (center(d) + rng.nextGaussian() * scaledSigma).toFloat + d += 1 + } + l2Normalize(v) + out(i) = v + i += 1 + } + out + } + + /** + * Median pairwise L2 distance between centers. We sample up to 1024 random center pairs + * rather than computing all `O(K^2)` of them — for `numClusters = 64` that's 2016 pairs, + * trivial; for larger K we'd otherwise pay cost the rest of the test doesn't need. + */ + private def medianPairwiseDistance(centers: Array[Array[Double]]): Double = { + val k = centers.length + if (k < 2) return 1.0 + val rng = new Random(0L) + val numPairs = math.min(1024, k * (k - 1) / 2) + val dists = new Array[Double](numPairs) + var p = 0 + while (p < numPairs) { + var i = rng.nextInt(k) + var j = rng.nextInt(k) + while (j == i) j = rng.nextInt(k) + dists(p) = euclidean(centers(i), centers(j)) + p += 1 + } + java.util.Arrays.sort(dists) + dists(dists.length / 2) + } + + private def euclidean(a: Array[Double], b: Array[Double]): Double = { + var s = 0.0 + var i = 0 + while (i < a.length) { + val d = a(i) - b(i) + s += d * d + i += 1 + } + math.sqrt(s) + } + + private def l2Normalize(v: Array[Float]): Unit = { + var s = 0.0 + var i = 0 + while (i < v.length) { s += v(i) * v(i); i += 1 } + val norm = math.sqrt(s).toFloat + if (norm > 0f) { + i = 0 + while (i < v.length) { v(i) = v(i) / norm; i += 1 } + } + } +} diff --git a/pom.xml b/pom.xml index 488425ef2..5e400b7cf 100644 --- a/pom.xml +++ b/pom.xml @@ -149,6 +149,7 @@ lance-spark-bundle-4.1_2.13 lance-spark-4.2_2.13 lance-spark-bundle-4.2_2.13 + lance-spark-knn-4.2_2.13 From 8c14203f6dcf7057e35830a5987e3715632c5ad0 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 10:16:51 -0700 Subject: [PATCH 2/9] Address lance-gatekeeper review on indexed APPROX NEAREST join Fixes the six findings raised on the PR: 1. pom parent version -> 0.8.0-beta.1 (deps already ${project.version}). 2. Carry the full pinned Lance read context through the join: the rule captures the base read options AND the DataSourceV2 relation options (where a DataFrame read hides branch/version/storage credentials) plus the namespace context; LanceKnnJoinExec merges and version-pins that context once on the driver so every executor probe opens one snapshot. 3. Materialize right-side payloads schema-aware via coerceToSpark (Map -> Row, recursive for array/map/struct); LanceProbe stays Spark-agnostic. 4. Quote/escape non-simple identifiers in the Lance prefilter string (bare for [A-Za-z_][A-Za-z0-9_]*, double-quoted otherwise, per dotted-path segment). 5. Stream join output through a lazy iterator and close the probe on task completion (eager drain-and-close fallback when there is no TaskContext), so the native handle can't leak or close early. 6. Add docs: a DQL operations page for the APPROX NEAREST join and a config.md subsection for the KNN extension and its tuning configs. Adds regression tests: relation context capture, identifier quoting, payload coercion, and iterator laziness. Verified locally: lance-spark-knn-4.2_2.13 builds and all 40 tests pass, including the native e2e suite (oracle equivalence, WHERE-prefilter pushdown, IVF-PQ recall, and rule-disabled fallthrough). Co-Authored-By: Claude Opus 4.8 --- docs/src/config.md | 21 ++ docs/src/operations/dql/.pages | 1 + .../operations/dql/nearest-neighbor-join.md | 145 +++++++++ lance-spark-knn-4.2_2.13/pom.xml | 2 +- .../catalyst/IndexedNearestByJoinRule.scala | 101 ++++--- .../spark/knn/catalyst/LanceKnnJoinExec.scala | 6 +- .../knn/internal/LanceKnnJoinStage.scala | 284 +++++++++++++----- .../lance/spark/knn/internal/LanceProbe.scala | 76 +++-- .../IndexedNearestByJoinRuleTest.scala | 121 ++++++-- .../IndexedNearestByJoinSqlTest.scala | 17 +- .../knn/internal/LanceKnnJoinStageTest.scala | 121 ++++++++ 11 files changed, 747 insertions(+), 148 deletions(-) create mode 100644 docs/src/operations/dql/nearest-neighbor-join.md create mode 100644 lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala diff --git a/docs/src/config.md b/docs/src/config.md index 84dcfd75b..45fc3bc2a 100644 --- a/docs/src/config.md +++ b/docs/src/config.md @@ -58,6 +58,27 @@ The following features require the Lance Spark SQL extension to be enabled: - [OPTIMIZE](operations/ddl/optimize.md) - Compact table fragments for improved query performance - [VACUUM](operations/ddl/vacuum.md) - Remove old versions and reclaim storage space +### Indexed Nearest-Neighbor Join Extension + +The `APPROX NEAREST` join (Spark 4.2 `APPROX NEAREST ... BY DISTANCE` / `BY SIMILARITY` syntax) can +be rewritten to probe the Lance vector index instead of running a brute-force cross product. This is +a **separate** extension, packaged in the `lance-spark-knn-4.2` module, and requires Spark 4.2 or +later. Enable it alongside (or instead of) the connector extension: + +``` +spark.sql.extensions = org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions +spark.lance.knn.indexedNearestByJoin.enabled = true +``` + +| Configuration | Type | Description | +|--------------------------------------------------|---------|---------------------------------------------------------------------------------| +| `spark.lance.knn.indexedNearestByJoin.enabled` | Boolean | Enable the indexed rewrite. Default `false` (falls through to Spark brute force). | +| `spark.lance.knn.nprobes` | Integer | IVF partitions to probe per query. Higher improves recall at more compute. | +| `spark.lance.knn.refineFactor` | Integer | IVF-PQ refine factor — fetch `k * refineFactor` candidates and re-rank exactly. | + +See [APPROX NEAREST Join](operations/dql/nearest-neighbor-join.md) for usage, supported ranking +functions, and `WHERE` pushdown. + ## Basic Setup Configure Spark with the `LanceNamespaceSparkCatalog` by setting the appropriate Spark catalog implementation diff --git a/docs/src/operations/dql/.pages b/docs/src/operations/dql/.pages index 60c6db968..516d2a956 100644 --- a/docs/src/operations/dql/.pages +++ b/docs/src/operations/dql/.pages @@ -3,5 +3,6 @@ nav: - select.md - fts.md - vector-search.md + - nearest-neighbor-join.md - search.md - hybrid-search.md diff --git a/docs/src/operations/dql/nearest-neighbor-join.md b/docs/src/operations/dql/nearest-neighbor-join.md new file mode 100644 index 000000000..9d5fe25c8 --- /dev/null +++ b/docs/src/operations/dql/nearest-neighbor-join.md @@ -0,0 +1,145 @@ +# APPROX NEAREST Join + +Join a query (left) table against a Lance (right) table so that each left row is matched with its +approximate _k_ nearest neighbors in the Lance table, using the Lance vector index instead of a +brute-force cross product. + +This uses the `APPROX NEAREST ... BY DISTANCE` / `BY SIMILARITY` join syntax added to Spark SQL in +Spark 4.2 ([SPARK-56395](https://issues.apache.org/jira/browse/SPARK-56395)). When the right side of +the join is a Lance table and the ranking expression is a recognized vector-distance function, the +Lance Spark KNN extension rewrites the join into a single no-shuffle operator that probes the Lance +vector index directly. + +!!! warning "Spark 4.2 Required" + The `APPROX NEAREST` join syntax is only available in Spark 4.2 or later. The indexed rewrite is + packaged in the `lance-spark-knn-4.2` module. + +!!! warning "KNN Extension Required" + The indexed rewrite requires the Lance Spark KNN SQL extension to be enabled. This is a separate + extension from the connector's `LanceSparkSessionExtensions` — see + [Enabling the Extension](#enabling-the-extension) below. Both can be enabled together in a + comma-separated `spark.sql.extensions` value. + +!!! note "Opt-in" + The rewrite is off by default. It fires only when + `spark.lance.knn.indexedNearestByJoin.enabled` is set to `true`. When it is off (or when the join + shape is not supported), the query falls through to Spark's built-in brute-force + `APPROX NEAREST` rewrite, so results are unchanged either way. + +## Enabling the Extension + +=== "Scala" + ```scala + val spark = SparkSession.builder() + .appName("lance-knn-example") + .config("spark.sql.extensions", + "org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions") + .config("spark.lance.knn.indexedNearestByJoin.enabled", "true") + .getOrCreate() + ``` + +=== "PySpark" + ```python + spark = SparkSession.builder \ + .appName("lance-knn-example") \ + .config("spark.sql.extensions", + "org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions") \ + .config("spark.lance.knn.indexedNearestByJoin.enabled", "true") \ + .getOrCreate() + ``` + +=== "Spark Submit" + ```shell + spark-submit \ + --conf spark.sql.extensions=org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions \ + --conf spark.lance.knn.indexedNearestByJoin.enabled=true \ + your-application.jar + ``` + +## Basic Usage + +The right side of the join is a Lance table (loaded through the Lance data source or a Lance +namespace catalog table). The ranking expression takes the left query vector and the right table's +vector column. + +=== "SQL" + ```sql + SELECT q.id, d.id, d.title + FROM queries q INNER JOIN documents d + APPROX NEAREST 10 BY DISTANCE vector_l2_distance(q.embedding, d.embedding); + ``` + +`documents` must resolve to a Lance table, for example a temp view over +`spark.read.format("lance").load(...)` or a Lance namespace catalog table such as +`lance.default.documents`. + +## Supported Ranking Functions + +The rewrite fires only when the ranking function and the `BY` direction are consistent: + +| Ranking function | Direction | Lance metric | +|-------------------------------------------|-------------------|--------------| +| `vector_l2_distance(left, right)` | `BY DISTANCE` | `l2` | +| `vector_cosine_similarity(left, right)` | `BY SIMILARITY` | `cosine` | +| `vector_inner_product(left, right)` | `BY SIMILARITY` | `dot` | + +Each argument must resolve to a single column — the left query vector and the right table's vector +column. Mixed-side or computed arguments (for example `vector_l2_distance(q.vec, slice(d.vec, ...))`) +are not rewritten and fall through to Spark's brute-force path. + +Only `APPROX` joins are rewritten. An exact (`EXACT`) nearest join is always handled by Spark's +brute-force rewrite. + +## WHERE Pushdown + +A `WHERE` clause on the right (Lance) side is translated into a Lance filter and applied by the +index **before** the top-_k_ search (a prefilter), so the neighbors are drawn only from rows matching +the filter: + +=== "SQL" + ```sql + SELECT q.id, d.id + FROM queries q + INNER JOIN (SELECT * FROM documents WHERE category = 'news' AND score > 5) d + APPROX NEAREST 10 BY DISTANCE vector_l2_distance(q.embedding, d.embedding); + ``` + +Translation is conservative. It supports comparisons (`=`, `!=`, `<`, `<=`, `>`, `>=`), `IN`, +`IS [NOT] NULL`, and `AND` / `OR` / `NOT` over the right table's columns (top-level or nested struct +fields) compared against literals. If any part of the predicate cannot be translated (UDFs, +subqueries, computed expressions, or a reference to a left-side column), the rewrite is refused +entirely and the query falls through to the brute-force path — the predicate is never partially +applied. + +## Tuning + +These options tune the Lance index search. Both are optional; when unset, Lance's index defaults +apply. + +| Configuration | Type | Description | +|------------------------------------------|---------|---------------------------------------------------------------------------------------------------| +| `spark.lance.knn.indexedNearestByJoin.enabled` | Boolean | Enable the indexed rewrite. Default `false`. | +| `spark.lance.knn.nprobes` | Integer | Number of IVF partitions to probe per query. Higher improves recall at more compute. | +| `spark.lance.knn.refineFactor` | Integer | IVF-PQ refine factor. Lance fetches `k * refineFactor` candidates and re-ranks them with exact distance. Highest-leverage recall knob for IVF-PQ. | + +## Snapshot Consistency + +The Lance table version is resolved and pinned once on the driver before the join runs. Every +partition probes that same snapshot, so a concurrent write to the Lance table does not change the +result mid-query. A `version` / `branch` supplied through the Lance read options (for example +`spark.read.format("lance").option("branch", "...")`) is honored. + +## Execution + +The rewrite produces a single physical operator that runs as one no-shuffle `mapPartitions` over the +left input. No `Exchange` is inserted above it. Each task opens the Lance vector index once, then for +every left row probes the index, keeps the top _k_, and late-materializes the surviving right rows by +row id. Because there is no shuffle or broadcast of the right table, the right table is opened per +task rather than moved across the network. + +## Validation + +The `lance-spark-knn-4.2` module covers this path with unit tests for the Catalyst rewrite rule and +end-to-end SQL tests that build a Lance IVF-PQ index, run `APPROX NEAREST` through the indexed +operator, and check recall against a brute-force oracle — including `WHERE` pushdown and the +rule-disabled fallthrough. diff --git a/lance-spark-knn-4.2_2.13/pom.xml b/lance-spark-knn-4.2_2.13/pom.xml index 3482611e5..dfdcdafad 100644 --- a/lance-spark-knn-4.2_2.13/pom.xml +++ b/lance-spark-knn-4.2_2.13/pom.xml @@ -5,7 +5,7 @@ org.lance lance-spark-root - 0.7.1 + 0.8.0-beta.1 ../pom.xml diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala index e83f02e68..7f686b91a 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala @@ -17,10 +17,10 @@ import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeRefer import org.apache.spark.sql.catalyst.plans.{JoinType, LeftOuter, NearestByDirection, NearestByDistance, NearestBySimilarity} import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, NearestByJoin, Project, SubqueryAlias} import org.apache.spark.sql.catalyst.rules.Rule -import org.apache.spark.sql.connector.catalog.Table import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation import org.apache.spark.sql.types.{BooleanType, ByteType, DoubleType, FloatType, IntegerType, LongType, ShortType, StringType, StructField, StructType} import org.apache.spark.unsafe.types.UTF8String +import org.lance.spark.{LanceDataset, LanceSparkReadOptions} import org.lance.spark.knn.internal.{LanceKnnJoinStage, Metric} /** @@ -65,13 +65,15 @@ import org.lance.spark.knn.internal.{LanceKnnJoinStage, Metric} * * == Lance scan detection == * - * Class-name match: `getClass.getName.contains("Lance")`. The probe / materialize path drives - * Lance's Java API directly, so the indexed-path executor is Lance-specific by construction — - * there's no general "any vector-capable backend" extension point here. URI extracted from the - * standard `path` / `datasetUri` option. The rule is opt-in via - * `spark.lance.knn.indexedNearestByJoin.enabled`, so a false positive can only fire when the user - * explicitly enabled the feature against a non-Lance backend; the runtime probe would surface the - * mismatch. + * The right relation's V2 table must be a connector [[org.lance.spark.LanceDataset]]. The probe / + * materialize path drives Lance's Java API directly, so the indexed-path executor is Lance-specific + * by construction — there's no general "any vector-capable backend" extension point here. From the + * `LanceDataset` the rule captures the FULL read context: the base `LanceSparkReadOptions`, the + * driver-side initial storage options, and the runtime namespace impl / properties. The relation's + * OWN options are captured too — that is where a DataFrame read carries branch / version / storage + * credentials, since `LanceDataSource` is a `SupportsCatalogOptions` whose identifier is the URI + * alone. [[org.lance.spark.knn.catalyst.LanceKnnJoinExec]] merges + pins that context once on the + * driver (as `LanceScanBuilder` does) and opens on executors through `Utils.openDatasetBuilder`. * * == Prefilter pushdown == * @@ -175,8 +177,11 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { val rightProjection: Seq[String] = lance.output.map(_.name) val stageConf = LanceKnnJoinStage.Conf( - datasetUri = lance.uri, - version = lance.version, + readOptions = lance.readOptions, + relationOptions = lance.relationOptions, + initialStorageOptions = lance.initialStorageOptions, + namespaceImpl = lance.namespaceImpl, + namespaceProperties = lance.namespaceProperties, vectorColumn = rightVecCol, metric = metric, k = k, @@ -223,10 +228,19 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { } } - /** Lance scan info extracted from a DSv2 relation, optionally with a translated prefilter. */ + /** + * Lance scan context extracted from a DSv2 relation, optionally with a translated prefilter. Carries + * the full read context so the executor can merge + pin + open exactly as the connector's scan path + * does — see [[LanceKnnJoinStage.resolveReadContext]]. `relationOptions` is where a DataFrame read + * carries branch / version / storage credentials (the `SupportsCatalogOptions` identifier is the URI + * alone), so it must be captured alongside the base `readOptions`. + */ final private case class LanceScanInfo( - uri: String, - version: Option[Long], + readOptions: LanceSparkReadOptions, + relationOptions: java.util.Map[String, String], + initialStorageOptions: java.util.Map[String, String], + namespaceImpl: String, + namespaceProperties: java.util.Map[String, String], output: Seq[Attribute], prefilter: Option[String]) @@ -256,21 +270,29 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { // computed columns) would change the schema we rely on for `j.output` mapping, so we // refuse those by falling through to the default `_ => None` case. unwrapLanceScan(child) - case rel: DataSourceV2Relation if isLanceTable(rel.table) => + case rel: DataSourceV2Relation if rel.table.isInstanceOf[LanceDataset] => // The probe / materialize path drives Lance's Java API directly, so this rule is - // Lance-specific by construction — there's no plug-in point for a non-Lance backend - // here. We detect Lance via class-name match and pull the URI from the standard `path` - // / `datasetUri` option. If neither is present we fall through (returning None lets - // Spark's brute-force rewrite handle the query). - val opts = rel.options - val uri = Option(opts.get("path")).orElse(Option(opts.get("datasetUri"))) - uri.map { u => + // Lance-specific by construction — there's no plug-in point for a non-Lance backend here. + // Capture the FULL read context from the connector's `LanceDataset` (base read options + + // runtime namespace) PLUS the relation's own options, which is where a DataFrame read carries + // branch / version / storage credentials (`LanceDataSource` is a `SupportsCatalogOptions` + // whose identifier is the URI alone). `LanceKnnJoinExec.doExecute` merges + pins these exactly + // as `LanceScanBuilder` does. Namespace / storage maps are copied into fresh serializable + // HashMaps so they ship cleanly to executors inside the stage Conf. + val ds = rel.table.asInstanceOf[LanceDataset] + Some( LanceScanInfo( - uri = u, - version = Option(opts.get("version")).map(_.toLong), + readOptions = ds.readOptions(), + relationOptions = new java.util.HashMap[String, String](rel.options.asCaseSensitiveMap()), + initialStorageOptions = Option(ds.getInitialStorageOptions()) + .map(new java.util.HashMap[String, String](_)) + .orNull, + namespaceImpl = ds.getNamespaceImpl(), + namespaceProperties = Option(ds.getNamespaceProperties()) + .map(new java.util.HashMap[String, String](_)) + .orNull, output = rel.output, - prefilter = None) - } + prefilter = None)) case _ => None } @@ -350,17 +372,30 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { } private def asRightColumn(e: Expression, rightAttrs: AttributeSet): Option[String] = e match { - case a: Attribute if rightAttrs.contains(a) => Some(a.name) + case a: Attribute if rightAttrs.contains(a) => Some(quoteIdentifier(a.name)) case g: GetStructField => // Nested struct field: render as a dotted path `col.field` (recursing so `a.b.c` works). // Lance's filter dialect treats a dotted identifier as a nested column path — its scan - // planner runs with enable_relations = false — so this maps 1:1. The recursion also gates - // on the ROOT resolving to a right-side attribute, so a left-side or foreign root refuses. + // planner runs with enable_relations = false — so this maps 1:1. Each path segment is quoted + // independently. The recursion also gates on the ROOT resolving to a right-side attribute, so + // a left-side or foreign root refuses. val fieldName = g.name.getOrElse(g.childSchema(g.ordinal).name) - asRightColumn(g.child, rightAttrs).map(base => s"$base.$fieldName") + asRightColumn(g.child, rightAttrs).map(base => s"$base.${quoteIdentifier(fieldName)}") case _ => None } + /** + * Render a column identifier for a Lance (DataFusion-flavored) filter. A plain identifier + * (`[A-Za-z_][A-Za-z0-9_]*`) is emitted bare so the common case stays readable; anything else — + * spaces, punctuation, a leading digit, an embedded quote — is double-quoted with embedded quotes + * doubled. Without this, a column named e.g. `my col` renders as a bare `my col` and produces a + * malformed / mis-parsed filter string. Applied per dotted-path segment, so `outer.inner field` + * becomes `outer."inner field"`. + */ + private def quoteIdentifier(name: String): String = + if (name.matches("[A-Za-z_][A-Za-z0-9_]*")) name + else "\"" + name.replace("\"", "\"\"") + "\"" + /** * Render a Spark literal as a Lance SQL literal. Dispatch is by `dataType`, NOT by the boxed * value class — Catalyst stores e.g. `Literal(0, DateType)` with the value as a plain `Int`, @@ -404,14 +439,6 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { } } - private def isLanceTable(table: Table): Boolean = { - val cls = table.getClass.getName - // Loose by design — the rule is opt-in via spark.lance.knn.indexedNearestByJoin.enabled, so - // a false positive here would only fire when the user explicitly turned the feature on - // against a non-Lance backend, and the runtime probe would surface the mismatch. - cls.contains("Lance") || cls.contains("lance") - } - /** * Recognize `rankingExpr` as one of the supported vector-distance functions, AND verify the * declared `direction` on `NearestByJoin` matches the function's natural ordering. diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala index ae62196d3..01e79003f 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/LanceKnnJoinExec.scala @@ -57,7 +57,11 @@ case class LanceKnnJoinExec( val childRdd = child.execute() val leftSchemaCaptured = leftSchema val finalSchemaCaptured = finalSchema - val confCaptured = stageConf + // Resolve + pin the Lance read context ONCE here on the driver (merge the relation options over + // the base table options, open the dataset, pin its version) before capturing it into the RDD + // closure — so every executor probe opens the SAME snapshot. This is the driver-side I/O the + // Catalyst rule deliberately avoids; it never runs in pure-plan unit tests, which don't execute. + val confCaptured = LanceKnnJoinStage.resolveReadContext(stageConf) // Encoders are created on the driver and captured into the closure (ExpressionEncoder is // serializable). The deserializer/serializer instances are NOT thread-safe, so build them diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala index 41805e366..dd85b80be 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -13,10 +13,12 @@ */ package org.lance.spark.knn.internal +import org.apache.spark.TaskContext import org.apache.spark.sql.Row -import org.apache.spark.sql.types.StructField - -import scala.collection.mutable +import org.apache.spark.sql.types.{ArrayType, DataType, MapType, StructField, StructType} +import org.lance.Dataset +import org.lance.spark.LanceSparkReadOptions +import org.lance.spark.utils.Utils /** * The whole indexed nearest-by join, done per Spark partition with NO shuffle. @@ -53,10 +55,20 @@ object LanceKnnJoinStage { * Everything a probe task needs, shipped from the driver. `internalK` is the overfetch count * handed to Lance (`k × overfetch`); `k` is the final per-left-row cut applied after the native * search. `leftVecIdx` is the position of the query-vector column in the left row. + * + * The read context — `readOptions`, `relationOptions`, `initialStorageOptions`, `namespaceImpl`, + * `namespaceProperties` — is the full Lance scan context captured from the connector's + * `LanceDataset` plus the DataSourceV2 relation options (where a DataFrame read carries branch / + * version / storage credentials). [[resolveReadContext]] merges + pins these once on the driver + * before the probe RDD launches; the resolved `readOptions` carry the pinned ref so every task + * probes one consistent snapshot. All five fields are serializable so they ship to executors. */ final case class Conf( - datasetUri: String, - version: Option[Long], + readOptions: LanceSparkReadOptions, + relationOptions: java.util.Map[String, String], + initialStorageOptions: java.util.Map[String, String], + namespaceImpl: String, + namespaceProperties: java.util.Map[String, String], vectorColumn: String, metric: Metric, k: Int, @@ -74,73 +86,166 @@ object LanceKnnJoinStage { extends Serializable /** - * Run the join for one partition of left rows. Opens the probe once, probes + materializes per - * row, and returns assembled join rows. Materializing into an `ArrayBuffer` before returning the - * iterator is deliberate: Spark pulls from `mapPartitions` lazily, so a bare lazy iterator would - * let the consumer outlive the `try`/`finally` and read from a closed probe handle. + * Resolve and pin the Lance read context ONCE, on the driver, before the probe RDD is launched — + * the same thing the connector's `LanceScanBuilder` does at scan-build time. Merges the + * DataSourceV2 relation options over the base table read options (branch / version / storage + * credentials), opens the dataset to read its current version, and pins that version so every + * executor probe sees one consistent snapshot even under concurrent writes. + * + * Returns a [[Conf]] whose `readOptions` carry the pinned ref; `relationOptions` is cleared since + * it has been folded in. Call this from the physical operator's `doExecute` (driver side) — never + * from the Catalyst rule, which must stay I/O-free so it can pattern-match against fake-URI + * relations in unit tests. + */ + def resolveReadContext(conf: Conf): Conf = { + val merged = mergeReadOptions(conf.readOptions, conf.relationOptions) + val builder = Utils.openDatasetBuilder(merged) + if (conf.initialStorageOptions != null) { + builder.initialStorageOptions(conf.initialStorageOptions) + } + builder.runtimeNamespace(conf.namespaceImpl, conf.namespaceProperties, merged.getTableId()) + val dataset: Dataset = builder.build() + val pinned = + try merged.withRef(Utils.pinOpenedRef(dataset, merged.getRef())) + finally dataset.close() + conf.copy(readOptions = pinned, relationOptions = new java.util.HashMap[String, String]()) + } + + /** + * Port of the connector's `LanceDataset.mergeScanOptions`: overlay the DataSourceV2 relation + * options on the base table's read options. Storage options merge with the relation winning, and + * any stale `version` / `branch` keys are stripped from the base first so the relation's ref is + * not shadowed by a leftover storage entry. An empty relation returns the base options untouched. + * + * The connector additionally guards against combining an incompatible pre-pinned table ref with a + * relation-specified branch / version; that guard only fires when the base options already carry + * a ref (a catalog table time-travelled at load), which does not happen on the DataFrame-relation + * path this rule matches, so it is intentionally omitted here. + */ + private def mergeReadOptions( + base: LanceSparkReadOptions, + relationOptions: java.util.Map[String, String]): LanceSparkReadOptions = { + if (relationOptions == null || relationOptions.isEmpty) { + return base + } + val merged = new java.util.HashMap[String, String](base.getStorageOptions) + merged.remove(LanceSparkReadOptions.CONFIG_VERSION) + merged.remove(LanceSparkReadOptions.CONFIG_BRANCH) + merged.putAll(relationOptions) + LanceSparkReadOptions + .builder() + .datasetUri(base.getDatasetUri) + .namespace(base.getNamespace) + .tableId(base.getTableId) + .catalogName(base.getCatalogName) + .indexCacheBackend(base.getIndexCacheBackend) + .metadataCacheBackend(base.getMetadataCacheBackend) + .ref(base.getRef) + .fromOptions(merged) + .build() + } + + /** + * Run the join for one partition of left rows. Opens the probe once, then streams the output: each + * left row expands to its (≤k) join rows on demand via [[lazyJoinIterator]], so the whole + * partition is never buffered. + * + * Because Spark pulls from `mapPartitions` lazily, the returned iterator can outlive this method — + * so the probe is closed on task completion (success OR failure) via the `TaskContext` listener, + * NOT a `try`/`finally` here (which would release the native handle before the consumer reads it). + * When there is no `TaskContext` (a direct call outside a Spark task, e.g. a JVM-only test), we + * fall back to draining eagerly and closing before returning so the handle cannot leak. */ def runPartition(leftRows: Iterator[Row], conf: Conf): Iterator[Row] = { if (leftRows.isEmpty) return Iterator.empty - val probe = new LanceProbe(conf.datasetUri, fragmentIds = None, version = conf.version) - val out = mutable.ArrayBuffer.empty[Row] - try { - leftRows.foreach { leftRow => - val q = extractVector(leftRow, conf.leftVecIdx) - if (q == null) { - // Null query vector: nothing to search. Emit a null-right row only for an outer join. - if (conf.outerJoin) { - out += assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null) - } - } else { - // Overfetch `internalK` candidates natively, then trim to the final `k`. Lance already - // returns them best-first, so when it hands back no more than `k` we keep them as-is and - // skip the heap entirely. - val refs = probe - .probe( - conf.vectorColumn, - q, - conf.internalK, - conf.metric, - conf.nprobes, - conf.refineFactor, - conf.ef, - conf.prefilter) - .toArray - val trimmed = - if (refs.length <= conf.k) refs - else { - val heap = new TopKHeap(conf.k, conf.smallerIsBetter) - heap.offerAll(refs) - heap.drain() - } + val probe = new LanceProbe( + conf.readOptions, + conf.initialStorageOptions, + conf.namespaceImpl, + conf.namespaceProperties, + fragmentIds = None) + + val output = lazyJoinIterator(leftRows, leftRow => processRow(leftRow, probe, conf)) + TaskContext.get() match { + case null => + try output.toList.iterator + finally probe.close() + case tc => + tc.addTaskCompletionListener[Unit](_ => probe.close()) + output + } + } + + /** + * Lazily compose per-partition output: each left row expands to its (≤k) join rows on demand. + * Extracted so a unit test can assert laziness — the left iterator is pulled element-by-element, + * not drained up front — with a stub expander and no Lance dataset. [[runPartition]] wires the + * real per-row probe / trim / materialize expansion through here. + */ + private[knn] def lazyJoinIterator( + leftRows: Iterator[Row], + expand: Row => Iterator[Row]): Iterator[Row] = + leftRows.flatMap(expand) - if (trimmed.isEmpty) { - if (conf.outerJoin) { - out += assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null) - } - } else { - // Late materialization: point-fetch the surviving right rows by `_rowid`. Building the - // `rowAddr -> row` map collapses any duplicate rowAddr to one payload; the loop below - // still emits one output row per surviving ref. - val materialized: Map[Long, Map[String, Any]] = probe - .materialize(trimmed.iterator.map(_.rowAddr).toSeq, conf.rightProjection) - .map(m => extractRowAddr(m) -> m) - .toMap - trimmed.foreach { ref => - val rightMap = materialized.getOrElse(ref.rowAddr, null) - out += assembleRow( - leftRow, - conf.leftFieldCount, - conf.rightFields, - rightMap, - ref.score) - } - } + /** + * Expand one left row into its join output rows: probe R's index, trim to `k`, late-materialize + * the surviving right rows by `_rowid`, and assemble `left ++ right ++ score`. Returns an empty + * iterator (inner join) or a single null-right row (outer join) when there is no query vector or + * no hit. + */ + private def processRow(leftRow: Row, probe: LanceProbe, conf: Conf): Iterator[Row] = { + val q = extractVector(leftRow, conf.leftVecIdx) + if (q == null) { + // Null query vector: nothing to search. Emit a null-right row only for an outer join. + if (conf.outerJoin) { + Iterator.single(assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null)) + } else { + Iterator.empty + } + } else { + // Overfetch `internalK` candidates natively, then trim to the final `k`. Lance already + // returns them best-first, so when it hands back no more than `k` we keep them as-is and + // skip the heap entirely. + val refs = probe + .probe( + conf.vectorColumn, + q, + conf.internalK, + conf.metric, + conf.nprobes, + conf.refineFactor, + conf.ef, + conf.prefilter) + .toArray + val trimmed = + if (refs.length <= conf.k) refs + else { + val heap = new TopKHeap(conf.k, conf.smallerIsBetter) + heap.offerAll(refs) + heap.drain() + } + + if (trimmed.isEmpty) { + if (conf.outerJoin) { + Iterator.single(assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null)) + } else { + Iterator.empty + } + } else { + // Late materialization: point-fetch the surviving right rows by `_rowid`. Building the + // `rowAddr -> row` map collapses any duplicate rowAddr to one payload; we still emit one + // output row per surviving ref. Bounded by `k`, so this stays per-row, not per-partition. + val materialized: Map[Long, Map[String, Any]] = probe + .materialize(trimmed.iterator.map(_.rowAddr).toSeq, conf.rightProjection) + .map(m => extractRowAddr(m) -> m) + .toMap + trimmed.iterator.map { ref => + val rightMap = materialized.getOrElse(ref.rowAddr, null) + assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, rightMap, ref.score) } } - } finally probe.close() - out.iterator + } } /** @@ -184,7 +289,8 @@ object LanceKnnJoinStage { /** * Assemble one output row: `left fields ++ right fields ++ score`. A null `rightValues` (outer - * join with no hit) fills the right side with nulls. + * join with no hit) fills the right side with nulls. Each right value is shaped to its target + * Spark type via [[coerceToSpark]] so the join's `ExpressionEncoder` accepts it. */ private def assembleRow( leftRow: Row, @@ -197,11 +303,55 @@ object LanceKnnJoinStage { while (i < leftFieldCount) { arr(i) = leftRow.get(i); i += 1 } var j = 0 while (j < rightFields.size) { + val field = rightFields(j) arr(leftFieldCount + j) = - if (rightValues == null) null else rightValues.getOrElse(rightFields(j).name, null) + if (rightValues == null) null + else coerceToSpark(rightValues.getOrElse(field.name, null), field.dataType) j += 1 } arr(leftFieldCount + rightFields.size) = score Row.fromSeq(arr.toSeq) } + + /** + * Shape a materialized right-side value to match its target Spark [[DataType]] so the assembled + * row satisfies the join's `ExpressionEncoder`. [[LanceProbe]] returns payloads Spark-agnostically + * — an Arrow struct cell arrives as a `Map[String, Any]` keyed by child-field name and a list cell + * as a `Seq` — but a Spark `StructType` slot expects a positional `Row`, not a `Map`. Without this + * coercion a nested-struct payload column is handed to the encoder as a `Map` and either fails + * encoding or materializes as garbage. + * + * Recurses so nested shapes all land correctly: + * - `StructType` → `Row` built in declared field order, each field coerced to its type + * - `ArrayType` → `Seq` with every element coerced to the element type + * - `MapType` → map with keys and values coerced + * - anything else (numeric / string / boolean primitives) → passed through unchanged + */ + private[knn] def coerceToSpark(value: Any, dataType: DataType): Any = { + if (value == null) return null + dataType match { + case s: StructType => + value match { + case m: scala.collection.Map[_, _] => + val byName = m.asInstanceOf[scala.collection.Map[String, Any]] + Row.fromSeq( + s.fields.map(f => coerceToSpark(byName.getOrElse(f.name, null), f.dataType)).toSeq) + case r: Row => r + case _ => value + } + case ArrayType(elementType, _) => + value match { + case seq: scala.collection.Seq[_] => seq.map(v => coerceToSpark(v, elementType)) + case arr: Array[_] => arr.toSeq.map(v => coerceToSpark(v, elementType)) + case _ => value + } + case MapType(keyType, valueType, _) => + value match { + case m: scala.collection.Map[_, _] => + m.map { case (k, v) => coerceToSpark(k, keyType) -> coerceToSpark(v, valueType) } + case _ => value + } + case _ => value + } + } } diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala index 0ebac1ab5..d86bec278 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -16,9 +16,10 @@ package org.lance.spark.knn.internal import org.apache.arrow.memory.BufferAllocator import org.apache.arrow.vector.{BigIntVector, FieldVector, Float4Vector, Float8Vector, UInt8Vector, VectorSchemaRoot} import org.apache.arrow.vector.ipc.ArrowReader -import org.lance.{Dataset, ReadOptions} +import org.lance.Dataset import org.lance.ipc.{LanceScanner, Query, ScanOptions} -import org.lance.spark.{LanceConstant, LanceRuntime} +import org.lance.spark.{LanceConstant, LanceRef, LanceRuntime, LanceSparkReadOptions} +import org.lance.spark.utils.Utils import java.util @@ -38,21 +39,48 @@ import scala.collection.mutable * * Lifecycle: instantiate per task, call `probe(...)` repeatedly, close at end. * - * @param datasetUri Lance dataset URI (passed straight to `Dataset.open`). - * @param fragmentIds Fragments this probe is restricted to. Pass `None` for whole-dataset search. - * @param version Optional Lance version to pin. Required when used inside a join, so all - * probe / materialize stages see the same snapshot. - * @param allocator Arrow allocator. Defaults to lance-spark's shared `LanceRuntime.allocator()`. + * @param readOptions Fully resolved Lance read context — dataset URI, storage + * credentials, catalog / cache backends, and the pinned branch / + * version ref. Inside a join this is resolved and pinned once on the + * driver (see [[LanceKnnJoinStage.resolveReadContext]]) so every task + * probes the same snapshot. + * @param initialStorageOptions Driver-side storage options from `namespace.describeTable()`, merged + * into the base storage options at open. May be `null`. + * @param namespaceImpl Namespace implementation type, for reconnecting the namespace on an + * executor where the live handle did not survive serialization. May be + * `null` for a plain URI dataset. + * @param namespaceProperties Namespace connection properties for that reconnection. May be `null`. + * @param fragmentIds Fragments this probe is restricted to. Pass `None` for whole-dataset + * search. + * @param allocator Arrow allocator. Defaults to lance-spark's shared + * `LanceRuntime.allocator()`. */ final class LanceProbe( - datasetUri: String, + readOptions: LanceSparkReadOptions, + initialStorageOptions: java.util.Map[String, String], + namespaceImpl: String, + namespaceProperties: java.util.Map[String, String], fragmentIds: Option[Seq[Int]], - version: Option[Long] = None, allocator: BufferAllocator = LanceRuntime.allocator()) extends AutoCloseable { + /** + * URI + optional pinned-version convenience constructor. Builds read options straight from a bare + * dataset URI and pins `version` on the main branch when present. For callers / tests that have + * only a plain URI and no namespace or storage context. + */ + def this(datasetUri: String, fragmentIds: Option[Seq[Int]], version: Option[Long]) = + this(LanceProbe.readOptionsFor(datasetUri, version), null, null, null, fragmentIds) + + def this(datasetUri: String, fragmentIds: Option[Seq[Int]]) = + this(datasetUri, fragmentIds, None) + // Open the dataset once. Lance's Java binding caches index metadata against the Dataset handle, - // so reusing it across probes keeps subsequent calls index-warm. + // so reusing it across probes keeps subsequent calls index-warm. Opening through + // `Utils.openDatasetBuilder` (rather than a bare `Dataset.open().uri(...)`) makes the probe honor + // the full resolved read context — storage credentials, catalog / cache backends, the pinned + // branch / version ref, and (on executors, where the live namespace handle is transient) the + // runtime-namespace reconnection — exactly as the connector's own scan path does. private val dataset: Dataset = openDataset() private val javaFragmentIds: Option[util.List[Integer]] = fragmentIds.map { ids => @@ -62,16 +90,12 @@ final class LanceProbe( } private def openDataset(): Dataset = { - val readOpts = { - val b = new ReadOptions.Builder() - version.foreach(v => b.setVersion(v)) - b.build() + val builder = Utils.openDatasetBuilder(readOptions) + if (initialStorageOptions != null) { + builder.initialStorageOptions(initialStorageOptions) } - Dataset.open() - .uri(datasetUri) - .allocator(allocator) - .readOptions(readOpts) - .build() + builder.runtimeNamespace(namespaceImpl, namespaceProperties, readOptions.getTableId()) + builder.build() } /** @@ -318,6 +342,20 @@ object LanceProbe { */ val ScoreColumns: Seq[String] = Seq("_distance", "_score") + /** + * Build read options from a bare dataset URI, pinning `version` on the main branch when present. + * Backs the URI convenience constructor. + */ + private[knn] def readOptionsFor( + datasetUri: String, + version: Option[Long]): LanceSparkReadOptions = { + val base = LanceSparkReadOptions.from(datasetUri) + version match { + case Some(v) => base.withRef(LanceRef.ofMain(v)) + case None => base + } + } + /** * Convert an Arrow-returned cell value into something Spark's encoders accept when stuffed * into a `Row`. Arrow's `FieldVector.getObject` returns Java types (boxed primitives, diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala index f4a22bfe1..7904a54fd 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala @@ -15,9 +15,9 @@ package org.lance.spark.knn.catalyst import org.apache.spark.sql.{RowFactory, SparkSession} import org.apache.spark.sql.catalyst.expressions.{Add, And, Attribute, AttributeSet, EqualTo, Expression, GetStructField, GreaterThan, In, IsNotNull, IsNull, LessThanOrEqual, Literal, Not, Or, VectorCosineSimilarity, VectorInnerProduct, VectorL2Distance} +import org.apache.spark.sql.catalyst.plans.{NearestByDistance, NearestBySimilarity} import org.apache.spark.sql.catalyst.plans.Inner import org.apache.spark.sql.catalyst.plans.logical.{Filter, LogicalPlan, NearestByJoin, Project, SubqueryAlias} -import org.apache.spark.sql.catalyst.plans.{NearestByDistance, NearestBySimilarity} import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation import org.apache.spark.sql.types._ import org.apache.spark.unsafe.types.UTF8String @@ -40,7 +40,7 @@ import scala.collection.JavaConverters._ * - Happy path: VectorL2Distance + NearestByDistance over a Lance DSv2 relation rewrites. * - Direction mismatch (e.g. L2 distance with NearestBySimilarity) does NOT rewrite. * - EXACT (`approx = false`) does NOT rewrite — Spark's brute-force keeps owning that path. - * - Non-Lance right side does NOT rewrite (duck-type check via class name). + * - Non-Lance right side does NOT rewrite (right relation's table is not a `LanceDataset`). * - Disabled by default — fires only when the gating config is set. * - Prefilter pushdown: right-side `WHERE` translates to a Lance SQL filter string, or refuses * the rewrite entirely when the predicate can't be pushed in full. @@ -319,7 +319,8 @@ class IndexedNearestByJoinRuleTest { val rid = makeAttr("rid", IntegerType) val category = makeAttr("category", StringType) val bucket = makeAttr("bucket", IntegerType) - val meta = makeAttr("meta", new StructType().add("category", StringType).add("bucket", IntegerType)) + val meta = + makeAttr("meta", new StructType().add("category", StringType).add("bucket", IntegerType)) val attrs = AttributeSet(Seq(rid, category, bucket, meta)) val cases: Seq[(Expression, String)] = Seq( @@ -375,11 +376,83 @@ class IndexedNearestByJoinRuleTest { } } + /** + * Identifiers that aren't plain `[A-Za-z_][A-Za-z0-9_]*` — spaces, punctuation — must be + * double-quoted (embedded quotes doubled) so the Lance filter string is well-formed rather than + * malformed. Plain identifiers stay bare (covered by the shapes test above). Applies per + * dotted-path segment for nested struct fields. + */ + @Test def testTranslatorQuotesUnsafeIdentifiers(): Unit = { + val spaced = makeAttr("weird col", StringType) + val quoteName = makeAttr("has\"quote", IntegerType) + val outer = makeAttr("outer", new StructType().add("inner field", StringType)) + val attrs = AttributeSet(Seq(spaced, quoteName, outer)) + + val cases: Seq[(Expression, String)] = Seq( + EqualTo(spaced, lit("A")) -> "\"weird col\" = 'A'", + GreaterThan(quoteName, lit(5)) -> "\"has\"\"quote\" > 5", + // nested field with a space -> only the unsafe segment gets quoted + EqualTo(GetStructField(outer, 0, Some("inner field")), lit("A")) -> + "outer.\"inner field\" = 'A'") + cases.foreach { case (expr, expected) => + assertEquals( + Some(expected), + IndexedNearestByJoinRule.translateFilter(expr, attrs), + s"identifier quoting mismatch for: $expr") + } + } + + /** + * A DataFrame read carries branch / version / storage credentials in the DSv2 RELATION options, + * not the base read options (`LanceDataSource` is a `SupportsCatalogOptions` whose identifier is + * the URI alone). The rule must capture `rel.options` into the stage `Conf` so the executor can + * merge + pin them; capturing only the base read options would silently read `main` HEAD without + * credentials — the exact bug this regression guards against. + */ + @Test def testCapturesBranchAndStorageOptionsFromRelation(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val left = trivialPlan("lid", "lvec") + val schema = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + StructField("rvec", ArrayType(FloatType, containsNull = false), nullable = false))) + val uri = tempDir.resolve("branch_lance").toString + val table = new FakeLanceTable(schema, uri) + val opts = new java.util.HashMap[String, String]() + opts.put("path", uri) + opts.put("branch", "frozen") + opts.put("storage.account_key", "secret") + val cims = new org.apache.spark.sql.util.CaseInsensitiveStringMap(opts) + val right = DataSourceV2Relation.create(table, None, None, cims) + val leftVec = left.output.find(_.name == "lvec").get + val rightVec = right.output.find(_.name == "rvec").get + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val conf = IndexedNearestByJoinRule(join) match { + case Project(_, node: LanceKnnJoinLogicalPlan) => node.stageConf + case other => fail(s"expected rewrite, got: $other"); ??? + } + assertEquals(uri, conf.readOptions.getDatasetUri, "base read options should carry the URI") + assertEquals( + "frozen", + conf.relationOptions.get("branch"), + "relation branch must be captured for merge + pin") + assertEquals( + "secret", + conf.relationOptions.get("storage.account_key"), + "relation storage credential must be captured") + } + // -- helpers ------------------------------------------------------------------------------ /** - * Construct a left-side regular plan and a right-side that resembles a Lance DSv2 scan via the - * duck-type check. Avoids the need for a real Lance reader. + * Construct a left-side regular plan and a right-side Lance DSv2 scan (a `FakeLanceTable`, which + * IS a `LanceDataset`). Avoids the need for a real Lance reader. */ private def buildPlans(metricFunction: String) : (LogicalPlan, Attribute, LogicalPlan, Attribute) = { @@ -399,10 +472,11 @@ class IndexedNearestByJoinRuleTest { } /** - * Build a `DataSourceV2Relation` whose `table.getClass.getName.contains("Lance")` so the - * rule's duck-type check accepts it. We don't actually run any I/O. Includes a `category` - * (string) and `bucket` (int) column so prefilter-pushdown tests can build realistic - * filter predicates without needing to extend the schema separately. + * Build a `DataSourceV2Relation` backed by a `FakeLanceTable` (a real connector `LanceDataset` + * subclass) so the rule's `instanceof LanceDataset` check accepts it. We don't run any I/O — the + * `LanceDataset` constructor only stores its options + schema. Includes a `category` (string) and + * `bucket` (int) column so prefilter-pushdown tests can build realistic filter predicates without + * needing to extend the schema separately. */ private def lanceLikeDsv2Relation(): LogicalPlan = { val schema = new StructType(Array( @@ -410,9 +484,10 @@ class IndexedNearestByJoinRuleTest { StructField("category", StringType, nullable = true), StructField("bucket", IntegerType, nullable = true), StructField("rvec", ArrayType(FloatType, containsNull = false), nullable = false))) - val table = new FakeLanceTable(schema) + val uri = tempDir.resolve("fake_lance").toString + val table = new FakeLanceTable(schema, uri) val opts = new java.util.HashMap[String, String]() - opts.put("path", tempDir.resolve("fake_lance").toString) + opts.put("path", uri) val cims = new org.apache.spark.sql.util.CaseInsensitiveStringMap(opts) DataSourceV2Relation.create(table, None, None, cims) } @@ -450,13 +525,19 @@ class IndexedNearestByJoinRuleTest { } /** - * Stub Table whose class name ends with "Lance" so the rule's duck-type check accepts it. No I/O - * — the rule only reads schema and options. Lives in the test source tree. + * Stub table that IS a connector `LanceDataset` (the rule requires `instanceof LanceDataset`). The + * `LanceDataset` constructor does no I/O — it only stores its read options + schema — so building + * one over a fake URI is safe and keeps these tests backend-free. `readOptions()` returns options + * carrying the fake URI; `getInitialStorageOptions()`/namespace getters return the empty / null + * values the constructor stores, which is exactly the read context the rule captures. Lives in the + * test source tree. */ -class FakeLanceTable(_schema: StructType) extends org.apache.spark.sql.connector.catalog.Table { - override def name(): String = "fake_lance" - override def schema(): StructType = _schema - override def capabilities() - : java.util.Set[org.apache.spark.sql.connector.catalog.TableCapability] = - java.util.Collections.emptySet() -} +class FakeLanceTable(_schema: StructType, uri: String) + extends org.lance.spark.LanceDataset( + org.lance.spark.LanceSparkReadOptions.from(uri), + _schema, + java.util.Collections.emptyMap[String, String](), + null, + null, + false, + null) diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala index f19991d3c..7ae4f4605 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala @@ -266,7 +266,8 @@ class IndexedNearestByJoinSqlTest { LanceVectorIndexBuilder.buildIvfPq(uniformUri, "rvec", numPartitions = 4, numSubVectors = 8) val (clusteredDf, clusteredIds, clusteredVecs) = buildLeftDf( - ClusteredEmbeddings.generate(NumLeft, Dim, numClusters = 4, seed = Seed + 2), Dim) + ClusteredEmbeddings.generate(NumLeft, Dim, numClusters = 4, seed = Seed + 2), + Dim) val (clusteredUri, clusteredRightIds, clusteredRightVecs) = writeRightDf( ClusteredEmbeddings.generate(NumRight, Dim, numClusters = 16, seed = Seed + 3), Dim, @@ -274,9 +275,19 @@ class IndexedNearestByJoinSqlTest { LanceVectorIndexBuilder.buildIvfPq(clusteredUri, "rvec", numPartitions = 4, numSubVectors = 8) val uniformRecall = recallAgainst( - uniformDf, uniformUri, uniformIds, uniformVecs, uniformRightIds, uniformRightVecs) + uniformDf, + uniformUri, + uniformIds, + uniformVecs, + uniformRightIds, + uniformRightVecs) val clusteredRecall = recallAgainst( - clusteredDf, clusteredUri, clusteredIds, clusteredVecs, clusteredRightIds, clusteredRightVecs) + clusteredDf, + clusteredUri, + clusteredIds, + clusteredVecs, + clusteredRightIds, + clusteredRightVecs) println( s" IVF-PQ recall@$K: uniform=$uniformRecall, clustered=$clusteredRecall " + "(uniform = IVF worst case; clustered = production-shaped)") diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala new file mode 100644 index 000000000..a7dfe96ae --- /dev/null +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala @@ -0,0 +1,121 @@ +/* + * Licensed 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.lance.spark.knn.internal + +import org.apache.spark.sql.Row +import org.apache.spark.sql.types.{ArrayType, IntegerType, MapType, StringType, StructType} +import org.junit.jupiter.api.Assertions._ +import org.junit.jupiter.api.Test + +import java.util.concurrent.atomic.AtomicInteger + +/** + * Backend-free unit tests for the two pieces of [[LanceKnnJoinStage]] that don't need a Lance + * dataset: the lazy per-partition output composition ([[LanceKnnJoinStage.lazyJoinIterator]]) and + * the Spark-type coercion of materialized right-side payloads + * ([[LanceKnnJoinStage.coerceToSpark]]). Both are `private[knn]`, so this test lives in the + * `org.lance.spark.knn.internal` package to reach them. The full probe → trim → materialize + * pipeline is covered by the e2e test in this module against a real Lance dataset. + */ +class LanceKnnJoinStageTest { + + // -- fix #5: streaming output must be lazy ------------------------------------------------ + + /** + * `lazyJoinIterator` must pull left rows ON DEMAND, never drain them up front — otherwise the + * whole partition would materialize in memory before a single output row is produced (the exact + * regression the streaming rewrite fixed). We feed a large source with a side-effect pull counter, + * take only the first 3 outputs (one per left row here), and assert only 3 left rows were pulled. + */ + @Test def testLazyJoinIteratorPullsLeftRowsOnDemand(): Unit = { + val pulled = new AtomicInteger(0) + val left: Iterator[Row] = Iterator.range(0, 1000000).map { i => + pulled.incrementAndGet() + Row(i) + } + val out = LanceKnnJoinStage.lazyJoinIterator(left, r => Iterator.single(r)) + val first3 = out.take(3).toList + assertEquals(3, first3.size, "should surface exactly the requested rows") + assertTrue( + pulled.get() <= 3, + s"lazyJoinIterator must pull left rows on demand, not drain them; pulled ${pulled.get()}") + } + + /** + * Fan-out (each left row expands to several join rows) stays lazy too: taking 3 outputs at 2 rows + * per left row must pull only the first 2 left rows, not the whole source. + */ + @Test def testLazyJoinIteratorFansOutLazily(): Unit = { + val pulled = new AtomicInteger(0) + val left: Iterator[Row] = Iterator.range(0, 1000000).map { i => + pulled.incrementAndGet() + Row(i) + } + val out = LanceKnnJoinStage.lazyJoinIterator( + left, + r => { + val i = r.getInt(0) + Iterator(Row(i, 0), Row(i, 1)) + }) + val first3 = out.take(3).toList + assertEquals(3, first3.size) + assertTrue(pulled.get() <= 2, s"expected <= 2 left pulls for 3 outputs, got ${pulled.get()}") + } + + // -- fix #3: schema-aware materialization ------------------------------------------------- + + /** A `Map` payload for a `StructType` slot becomes a positional `Row` in declared field order. */ + @Test def testCoerceStructFromMapToRowInFieldOrder(): Unit = { + val dt = new StructType().add("a", IntegerType).add("b", StringType) + // Keyed by name and deliberately out of declared order — coercion must reorder by field. + val value = Map("b" -> "x", "a" -> Integer.valueOf(1)) + val row = LanceKnnJoinStage.coerceToSpark(value, dt).asInstanceOf[Row] + assertEquals(1, row.getInt(0)) + assertEquals("x", row.getString(1)) + } + + /** A field absent from the payload map materializes as null, not a missing slot. */ + @Test def testCoerceStructFillsMissingFieldWithNull(): Unit = { + val dt = new StructType().add("a", IntegerType).add("b", StringType) + val value = Map[String, Any]("a" -> Integer.valueOf(1)) + val row = LanceKnnJoinStage.coerceToSpark(value, dt).asInstanceOf[Row] + assertEquals(1, row.getInt(0)) + assertTrue(row.isNullAt(1), "missing struct field should be null") + } + + /** `ArrayType` recurses: an array of struct payloads becomes a `Seq[Row]`. */ + @Test def testCoerceArrayOfStructsRecurses(): Unit = { + val dt = ArrayType(new StructType().add("a", IntegerType)) + val value = Seq(Map("a" -> Integer.valueOf(7))) + val out = LanceKnnJoinStage.coerceToSpark(value, dt).asInstanceOf[Seq[_]] + assertEquals(1, out.size) + assertEquals(7, out.head.asInstanceOf[Row].getInt(0)) + } + + /** `MapType` recurses on its values: a struct-valued map entry becomes a `Row`. */ + @Test def testCoerceMapOfStructsRecurses(): Unit = { + val dt = MapType(StringType, new StructType().add("a", IntegerType)) + val value = Map("k" -> Map("a" -> Integer.valueOf(9))) + val out = + LanceKnnJoinStage.coerceToSpark(value, dt).asInstanceOf[scala.collection.Map[String, Any]] + assertEquals(9, out("k").asInstanceOf[Row].getInt(0)) + } + + /** Primitives pass straight through; null stays null. */ + @Test def testCoercePrimitivesAndNullPassThrough(): Unit = { + assertEquals("hello", LanceKnnJoinStage.coerceToSpark("hello", StringType)) + assertEquals(42, LanceKnnJoinStage.coerceToSpark(Integer.valueOf(42), IntegerType)) + assertNull(LanceKnnJoinStage.coerceToSpark(null, IntegerType).asInstanceOf[AnyRef]) + } +} From eef95288f6a66e58f80aea0cea6d0dc7886ca94c Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 10:21:54 -0700 Subject: [PATCH 3/9] Add end-to-end version-pinning test for the KNN join read context MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Covers finding #2 at the runtime level (the earlier test only asserted capture into the stage Conf): a `version` read option supplied on the DataFrame read (`spark.read.format("lance").option("version", "1")`) lands only on the DataSourceV2 relation options, so it must be captured, merged, and pinned on the driver for the probe to scan that snapshot. The test writes v1, appends v2 rows that duplicate each left query vector (distance 0), then asserts: pinned to v1 the appended rows are invisible and every left row still gets k v1 hits; on latest each left row's own duplicate is a nearest hit — so the pin is what excluded them (the assertion is discriminating, not vacuous). Verified: IndexedNearestByJoinSqlTest 7/7 pass. Co-Authored-By: Claude Opus 4.8 --- .../IndexedNearestByJoinSqlTest.scala | 90 +++++++++++++++++++ 1 file changed, 90 insertions(+) diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala index 7ae4f4605..a5622ea8b 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala @@ -215,6 +215,69 @@ class IndexedNearestByJoinSqlTest { } } + /** + * Read-context end-to-end: a `version` supplied through the Lance read option + * (`spark.read.format("lance").option("version", "1")`) must be captured from the DataSourceV2 + * relation options, merged, and pinned on the driver — so every probe scans that snapshot, not + * HEAD. `LanceDataSource` is a `SupportsCatalogOptions` whose identifier is the URI alone, so this + * option lives ONLY on the relation options; a rule that captured just the base read options would + * silently read HEAD. + * + * We write v1, then append v2 rows that are EXACT duplicates of each left query vector (distance + * 0 — they would be the #1 nearest neighbor if visible), then assert: + * - pinned to v1: no appended row (`rid >= DupIdBase`) may appear, and every left row still gets + * `k` v1 hits; + * - latest (no version option): each left row's own duplicate (`rid == DupIdBase + lid`) DOES + * appear — proving the pin is what excluded them, so the pinned assertion is discriminating, + * not vacuously satisfied. + */ + @Test def testSqlVersionReadOptionPinsSnapshot(): Unit = { + val DupIdBase = 5000 + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed + 300) + val (leftDf, leftIds, _) = buildLeftDf(leftVecs, ExactDim) + + // v1: baseline right rows (ids 1000..1063, all < DupIdBase). + val v1Vecs = generateUniform(ExactRight, ExactDim, Seed + 301) + val (rightUri, _, _) = writeRightDf(v1Vecs, ExactDim, idBase = 1000) + + // v2: append one exact duplicate of each left query vector (id DupIdBase + lid). + val dupRows = leftIds.zip(leftVecs).map { case (lid, v) => + RowFactory.create(Integer.valueOf(DupIdBase + lid), v) + } + spark + .createDataFrame(dupRows.toSeq.asJava, rightSchema(ExactDim)) + .write + .format("lance") + .mode("append") + .save(rightUri) + + val k = 5 + + // Pinned to version 1: the appended duplicates must be invisible. + val pinnedRows = runKnnSqlVersioned(leftDf, rightUri, k, version = Some(1L)) + assertEquals( + ExactLeft * k, + pinnedRows.length, + "expected k v1 hits per left row when pinned to v1") + val pinnedByLid = pinnedRows.groupBy(_.getAs[Int]("lid")) + leftIds.foreach { lid => + val rids = pinnedByLid(lid).map(_.getAs[Int]("rid")) + assertTrue( + rids.forall(_ < DupIdBase), + s"pinned-to-v1 result for lid=$lid leaked an appended (v2) row: ${rids.mkString(",")}") + } + + // Latest (no version option): each left row's own duplicate is a nearest hit (distance 0). + val latestRows = runKnnSqlVersioned(leftDf, rightUri, k, version = None) + val latestByLid = latestRows.groupBy(_.getAs[Int]("lid")) + leftIds.foreach { lid => + val rids = latestByLid(lid).map(_.getAs[Int]("rid")).toSet + assertTrue( + rids.contains(DupIdBase + lid), + s"latest result for lid=$lid should contain its duplicate ${DupIdBase + lid}; got $rids") + } + } + // -- approximate path (IVF-PQ): recall floors, refineFactor ------------------------------- /** @@ -356,6 +419,33 @@ class IndexedNearestByJoinSqlTest { |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin).collect() } + /** + * Like [[runKnnSql]] but registers the right Lance view through an optional `version` READ OPTION, + * exercising the DataSourceV2-relation-options path that carries branch / version / storage + * credentials. Also asserts the indexed rule still fires when a version option is present. + */ + private def runKnnSqlVersioned( + leftDf: DataFrame, + rightUri: String, + k: Int, + version: Option[Long]): Array[Row] = { + viewSeq += 1 + val q = s"queries_$viewSeq" + val d = s"docs_$viewSeq" + leftDf.createOrReplaceTempView(q) + val reader = spark.read.format("lance") + version.foreach(v => reader.option("version", v.toString)) + reader.load(rightUri).createOrReplaceTempView(d) + val df = spark.sql( + s"""SELECT q.lid, d.rid + |FROM $q q INNER JOIN $d d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + assertTrue( + df.queryExecution.optimizedPlan.collect { case p: LanceKnnJoinLogicalPlan => p }.nonEmpty, + s"expected indexed rewrite even with a version option; plan:\n${df.queryExecution.optimizedPlan}") + df.collect() + } + private def leftSchema(dim: Int): StructType = new StructType(Array( StructField("lid", IntegerType, nullable = false), fixedSizeVec("lvec", dim))) From d31fd248439321ea73c766b78a0cdb356499cd23 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 12:07:35 -0700 Subject: [PATCH 4/9] Address second-round lance-gatekeeper review on indexed APPROX NEAREST join MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes the six findings from the re-review, each with a regression test matching the reviewer's reproducer: 1. LanceProbe.openDataset no longer calls runtimeNamespace unconditionally (which force-loaded the namespace impl class). It now applies the exact worker-open namespace policy of LanceFragmentScanner.create: touch the namespace only when a namespace impl is configured AND executor credential refresh is enabled, then rebuild or clear it accordingly. 2. LanceKnnJoinStage.mergeReadOptions now preserves the connector's incompatible-ref guard: when the base carries a pinned ref and the relation also pins version/branch, a same-named branch keeps the table ref while an incompatible combination is rejected, rather than letting the relation silently win and read a different snapshot. 3. coerceToSpark rebuilds a real Spark map from Arrow's map representation (a LIST of {key,value} entry structs surfaces as a Seq, not a Scala Map), so a MapType payload column encodes correctly instead of failing/garbling. 4. quoteIdentifier always-quotes column identifiers so a name colliding with a SQL keyword/literal (e.g. `true`) can't collapse into a tautology. The quote character is a BACKTICK: Lance's filter dialect treats a double-quoted token as a string literal, so `"category" = 'A'` prefilters to zero rows — verified against a real Lance dataset via the WHERE-pushdown oracle test. 5. The Catalyst rewrite now declines when a standard analysis guard would fire (spark.sql.crossJoin.enabled=false or a streaming child), leaving Spark's own path to reject with the expected error instead of silently rewriting past it. 6. Docs gain an Installation section with the module Maven coordinate (org.lance:lance-spark-knn-4.2_2.13) and a --packages spark-submit example. All 47 tests in lance-spark-knn-4.2_2.13 pass; spotless clean. Co-Authored-By: Claude Opus 4.8 --- .../operations/dql/nearest-neighbor-join.md | 45 +++++++++++ .../catalyst/IndexedNearestByJoinRule.scala | 43 ++++++++--- .../knn/internal/LanceKnnJoinStage.scala | 67 +++++++++++++++-- .../lance/spark/knn/internal/LanceProbe.scala | 20 ++++- .../IndexedNearestByJoinRuleTest.scala | 74 +++++++++++++------ .../IndexedNearestByJoinSqlTest.scala | 36 ++++++++- .../knn/internal/LanceKnnJoinStageTest.scala | 55 ++++++++++++++ .../internal/LanceProbeValidationTest.scala | 36 ++++++++- 8 files changed, 333 insertions(+), 43 deletions(-) diff --git a/docs/src/operations/dql/nearest-neighbor-join.md b/docs/src/operations/dql/nearest-neighbor-join.md index 9d5fe25c8..9ef4304f8 100644 --- a/docs/src/operations/dql/nearest-neighbor-join.md +++ b/docs/src/operations/dql/nearest-neighbor-join.md @@ -26,6 +26,51 @@ vector index directly. shape is not supported), the query falls through to Spark's built-in brute-force `APPROX NEAREST` rewrite, so results are unchanged either way. +## Installation + +The indexed rewrite ships in its own module, **separate** from the base connector artifacts — it is +built only for Spark 4.2 / Scala 2.13 (the Spark release where `APPROX NEAREST` exists). Add it +alongside the connector. + +| Artifact | Coordinate | +|-----------------------------------|-------------------------------------------------| +| KNN SQL extension (Spark 4.2) | `org.lance:lance-spark-knn-4.2_2.13:` | +| Lance connector (Spark 4.2) | `org.lance:lance-spark-bundle-4.2_2.13:` | + +Use the same `` as the connector release. + +=== "Maven" + ```xml + + org.lance + lance-spark-knn-4.2_2.13 + VERSION + + ``` + +=== "Gradle" + ```gradle + dependencies { + implementation 'org.lance:lance-spark-knn-4.2_2.13:VERSION' + } + ``` + +=== "sbt" + ```scala + libraryDependencies += "org.lance" % "lance-spark-knn-4.2_2.13" % "VERSION" + ``` + +To supply it to a running cluster, add the coordinate to `--packages` (comma-separated, together +with the Lance connector bundle): + +```shell +spark-submit \ + --packages org.lance:lance-spark-bundle-4.2_2.13:VERSION,org.lance:lance-spark-knn-4.2_2.13:VERSION \ + --conf spark.sql.extensions=org.lance.spark.knn.extensions.LanceKnnSparkSessionExtensions \ + --conf spark.lance.knn.indexedNearestByJoin.enabled=true \ + your-application.jar +``` + ## Enabling the Extension === "Scala" diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala index 7f686b91a..38aa3c55d 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala @@ -122,7 +122,8 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { val nprobes = optInt(NprobesConfKey) val refineFactor = optInt(RefineFactorConfKey) plan.transformDown { - case j @ NearestByJoin(left, right, joinType, true, k, rankingExpr, direction) => + case j @ NearestByJoin(left, right, joinType, true, k, rankingExpr, direction) + if preservesAnalysisGuards(left, right) => rewriteIfApplicable( j, left, @@ -136,6 +137,24 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { } } + /** + * This rule runs as a `postHocResolutionRule`, i.e. BEFORE the optimizer's `FinishAnalysis` batch + * (which lowers `NearestByJoin` into a Cartesian product + `MaxMinByK`) and BEFORE the checks that + * batch relies on. Replacing a `NearestByJoin` that one of those checks would REJECT would let an + * illegal query succeed silently. So we decline the rewrite — leaving the original `NearestByJoin` + * in place for Spark's own path to reject with the exact error the user expects — whenever a + * standard analysis guard would fire: + * + * - `spark.sql.crossJoin.enabled = false`: a `NearestByJoin` carries no equi-condition, so its + * default lowering is a Cartesian product, which `CheckCartesianProducts` rejects with + * `CROSS_JOIN_NOT_ENABLED`. Reading `conf.crossJoinEnabled` here uses the same value (and + * default) that check uses. + * - a streaming child: `NearestByJoin` over a streaming input is unsupported + * (`STREAMING_NOT_SUPPORTED`); leave it for Spark's unsupported-operation check to reject. + */ + private def preservesAnalysisGuards(left: LogicalPlan, right: LogicalPlan): Boolean = + conf.crossJoinEnabled && !left.isStreaming && !right.isStreaming + private def optInt(key: String): Option[Int] = Option(conf.getConfString(key, null)).map(_.toInt) @@ -385,16 +404,22 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { } /** - * Render a column identifier for a Lance (DataFusion-flavored) filter. A plain identifier - * (`[A-Za-z_][A-Za-z0-9_]*`) is emitted bare so the common case stays readable; anything else — - * spaces, punctuation, a leading digit, an embedded quote — is double-quoted with embedded quotes - * doubled. Without this, a column named e.g. `my col` renders as a bare `my col` and produces a - * malformed / mis-parsed filter string. Applied per dotted-path segment, so `outer.inner field` - * becomes `outer."inner field"`. + * Render a column identifier for a Lance filter by ALWAYS back-quoting it (any embedded backtick + * doubled). A bare identifier is not merely a readability question — it is ambiguous with SQL + * keywords and literals: a column literally named `true`, `null`, or `select` emitted bare parses + * as the keyword/literal, not a column reference (`true = true` is a tautology, not `col = true`), + * silently corrupting the prefilter. A "quote only the non-word identifiers" exception cannot + * enumerate every such reserved word, so we quote unconditionally — a delimited identifier is + * unambiguous for every name. + * + * The quote character is a BACKTICK, not a double-quote. Lance's filter dialect is MySQL-flavored: + * a double-quoted token is a STRING LITERAL, so `"category" = 'A'` parses as the constant + * comparison `'category' = 'A'` (always false) and silently prefilters to zero rows — verified + * against a real Lance dataset. Backticks delimit an identifier. Applied per dotted-path segment, + * so a nested field `outer.inner` becomes `` `outer`.`inner` ``. */ private def quoteIdentifier(name: String): String = - if (name.matches("[A-Za-z_][A-Za-z0-9_]*")) name - else "\"" + name.replace("\"", "\"\"") + "\"" + "`" + name.replace("`", "``") + "`" /** * Render a Spark literal as a Lance SQL literal. Dispatch is by `dataType`, NOT by the boxed diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala index dd85b80be..cfb4edd60 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -17,7 +17,7 @@ import org.apache.spark.TaskContext import org.apache.spark.sql.Row import org.apache.spark.sql.types.{ArrayType, DataType, MapType, StructField, StructType} import org.lance.Dataset -import org.lance.spark.LanceSparkReadOptions +import org.lance.spark.{LanceRef, LanceSparkReadOptions} import org.lance.spark.utils.Utils /** @@ -117,22 +117,29 @@ object LanceKnnJoinStage { * any stale `version` / `branch` keys are stripped from the base first so the relation's ref is * not shadowed by a leftover storage entry. An empty relation returns the base options untouched. * - * The connector additionally guards against combining an incompatible pre-pinned table ref with a - * relation-specified branch / version; that guard only fires when the base options already carry - * a ref (a catalog table time-travelled at load), which does not happen on the DataFrame-relation - * path this rule matches, so it is intentionally omitted here. + * Preserves the connector's incompatible-ref guard: if the base options already carry a pinned + * `ref` (e.g. a catalog table time-travelled at load) and the relation ALSO sets `version` / + * `branch`, a same-named branch keeps the table ref while an incompatible combination is rejected + * with an `IllegalArgumentException` rather than silently letting the relation value win. On the + * plain DataFrame-relation path this rule usually matches, the base carries no ref and the guard + * is inert — but keeping it makes the merge behave identically to `LanceDataset.mergeScanOptions` + * for a pinned base, so a caller cannot smuggle a conflicting snapshot past the merge. */ - private def mergeReadOptions( + private[knn] def mergeReadOptions( base: LanceSparkReadOptions, relationOptions: java.util.Map[String, String]): LanceSparkReadOptions = { if (relationOptions == null || relationOptions.isEmpty) { return base } + val tableRef = base.getRef + val scanSetsBranchOrVersion = + relationOptions.containsKey(LanceSparkReadOptions.CONFIG_VERSION) || + relationOptions.containsKey(LanceSparkReadOptions.CONFIG_BRANCH) val merged = new java.util.HashMap[String, String](base.getStorageOptions) merged.remove(LanceSparkReadOptions.CONFIG_VERSION) merged.remove(LanceSparkReadOptions.CONFIG_BRANCH) merged.putAll(relationOptions) - LanceSparkReadOptions + val scanOptions = LanceSparkReadOptions .builder() .datasetUri(base.getDatasetUri) .namespace(base.getNamespace) @@ -140,11 +147,29 @@ object LanceKnnJoinStage { .catalogName(base.getCatalogName) .indexCacheBackend(base.getIndexCacheBackend) .metadataCacheBackend(base.getMetadataCacheBackend) - .ref(base.getRef) + .ref(tableRef) .fromOptions(merged) .build() + if (tableRef != null && scanSetsBranchOrVersion) { + val scanRef = scanOptions.getRef + if (sameNamedBranch(tableRef, scanRef)) { + // Same named branch, different version → keep the table's pinned ref (snapshot isolation). + return scanOptions.withRef(tableRef) + } + require( + tableRef == scanRef, + s"Cannot combine $tableRef with $scanRef") + } + scanOptions } + /** Same-named-branch check, mirroring the connector's `LanceDataset.sameNamedBranch`. */ + private def sameNamedBranch(tableRef: LanceRef, scanRef: LanceRef): Boolean = + tableRef.isBranch && + scanRef != null && + scanRef.isBranch && + tableRef.getBranchName.equals(scanRef.getBranchName) + /** * Run the join for one partition of left rows. Opens the probe once, then streams the output: each * left row expands to its (≤k) join rows on demand via [[lazyJoinIterator]], so the whole @@ -349,9 +374,35 @@ object LanceKnnJoinStage { value match { case m: scala.collection.Map[_, _] => m.map { case (k, v) => coerceToSpark(k, keyType) -> coerceToSpark(v, valueType) } + case entries: scala.collection.Seq[_] => + // The shape a real Arrow map cell actually arrives in. Arrow represents a `MapType` + // cell as a LIST of `{key, value}` entry structs, so `LanceProbe.toSparkValue` turns it + // into a `Seq(Map("key" -> …, "value" -> …), …)` — NOT a Scala map. Left uncoerced the + // encoder would see a sequence where a `MapType` slot is expected and either fail or + // materialize garbage. Rebuild a real map from the entries. (The `Map` case above stays + // for direct-map payloads and JVM-only tests.) + entries.iterator.map { entry => + val (k, v) = mapEntryKeyValue(entry) + coerceToSpark(k, keyType) -> coerceToSpark(v, valueType) + }.toMap case _ => value } case _ => value } } + + /** + * Pull `(key, value)` out of a single Arrow map entry. `LanceProbe.toSparkValue` renders each + * entry as a `Map("key" -> …, "value" -> …)` (Arrow's `MapVector` names the entry-struct children + * `key` / `value`); a positional `Row(key, value)` is tolerated defensively. + */ + private def mapEntryKeyValue(entry: Any): (Any, Any) = entry match { + case m: scala.collection.Map[_, _] => + val byName = m.asInstanceOf[scala.collection.Map[String, Any]] + (byName.getOrElse("key", null), byName.getOrElse("value", null)) + case r: Row if r.length >= 2 => (r.get(0), r.get(1)) + case other => + throw new IllegalStateException( + s"Unexpected Arrow map-entry representation: ${other.getClass.getName}") + } } diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala index d86bec278..41c7b2f11 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -90,11 +90,29 @@ final class LanceProbe( } private def openDataset(): Dataset = { + // Apply the connector's exact worker-open namespace policy (`LanceFragmentScanner.create`): + // touch the namespace ONLY when a namespace impl is configured AND executor credential refresh + // is enabled. When refresh is on, either rebuild the namespace client (impls that must run on + // workers) or clear it (impls that must not, so the open falls back to the URI + initial storage + // options). When refresh is OFF we leave the read options' namespace exactly as shipped — this + // is the whole point of `executor_credential_refresh=false`, and forcing a rebuild here would + // turn that policy into a namespace class-load / RPC that can fail before the dataset even opens. + // + // Notably this does NOT call `builder.runtimeNamespace(namespaceImpl, ...)` unconditionally + // (which always loads the namespace impl class) — matching the fragment scanner, which reaches + // the namespace only through the guarded `setNamespace` path above. + if (namespaceImpl != null && readOptions.isExecutorCredentialRefresh()) { + if (LanceRuntime.useNamespaceOnWorkers(namespaceImpl)) { + readOptions.setNamespace( + LanceRuntime.getOrCreateNamespace(namespaceImpl, namespaceProperties)) + } else { + readOptions.setNamespace(null) + } + } val builder = Utils.openDatasetBuilder(readOptions) if (initialStorageOptions != null) { builder.initialStorageOptions(initialStorageOptions) } - builder.runtimeNamespace(namespaceImpl, namespaceProperties, readOptions.getTableId()) builder.build() } diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala index 7904a54fd..ae09fde07 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala @@ -59,6 +59,11 @@ class IndexedNearestByJoinRuleTest { .master("local[2]") .config("spark.driver.bindAddress", "127.0.0.1") .config("spark.driver.host", "127.0.0.1") + // A `NearestByJoin` lowers to a Cartesian product, so the rule (like a real query) only fires + // when cross joins are permitted — otherwise `preservesAnalysisGuards` declines the rewrite so + // Spark can reject the query itself. Enable it here so these rewrite-shape tests exercise the + // rewrite path; the crossJoin-guard behavior is covered end-to-end in the SQL test. + .config("spark.sql.crossJoin.enabled", "true") .getOrCreate() } @@ -323,30 +328,52 @@ class IndexedNearestByJoinRuleTest { makeAttr("meta", new StructType().add("category", StringType).add("bucket", IntegerType)) val attrs = AttributeSet(Seq(rid, category, bucket, meta)) + // Every column identifier is back-quoted (see `quoteIdentifier`): a delimited identifier is + // unambiguous with SQL keywords/literals, so a column named e.g. `true` can't collapse into a + // tautology. Backticks (not double-quotes) delimit an identifier in Lance's filter dialect — + // a double-quoted token is a string literal there. Literals are unquoted; string values keep + // single-quote escaping. val cases: Seq[(Expression, String)] = Seq( - EqualTo(category, lit("A")) -> "category = 'A'", - Not(EqualTo(category, lit("A"))) -> "category != 'A'", - GreaterThan(bucket, lit(5)) -> "bucket > 5", - LessThanOrEqual(bucket, lit(5)) -> "bucket <= 5", - IsNull(category) -> "category IS NULL", - IsNotNull(category) -> "category IS NOT NULL", - In(bucket, Seq(lit(1), lit(2), lit(3))) -> "bucket IN (1, 2, 3)", + EqualTo(category, lit("A")) -> "`category` = 'A'", + Not(EqualTo(category, lit("A"))) -> "`category` != 'A'", + GreaterThan(bucket, lit(5)) -> "`bucket` > 5", + LessThanOrEqual(bucket, lit(5)) -> "`bucket` <= 5", + IsNull(category) -> "`category` IS NULL", + IsNotNull(category) -> "`category` IS NOT NULL", + In(bucket, Seq(lit(1), lit(2), lit(3))) -> "`bucket` IN (1, 2, 3)", And(EqualTo(category, lit("A")), GreaterThan(bucket, lit(5))) -> - "(category = 'A') AND (bucket > 5)", + "(`category` = 'A') AND (`bucket` > 5)", Or(EqualTo(category, lit("A")), EqualTo(category, lit("B"))) -> - "(category = 'A') OR (category = 'B')", + "(`category` = 'A') OR (`category` = 'B')", // String-literal escape — single quotes inside the value get doubled. - EqualTo(category, lit("O'Brien")) -> "category = 'O''Brien'", + EqualTo(category, lit("O'Brien")) -> "`category` = 'O''Brien'", // literal-on-left flip - EqualTo(lit(5), bucket) -> "5 = bucket", - // nested struct field access -> dotted path (distinct from the top-level `category` attr) - EqualTo(GetStructField(meta, 0, Some("category")), lit("A")) -> "meta.category = 'A'") + EqualTo(lit(5), bucket) -> "5 = `bucket`", + // nested struct field access -> dotted path, each segment quoted independently + EqualTo(GetStructField(meta, 0, Some("category")), lit("A")) -> + "`meta`.`category` = 'A'") cases.foreach { case (expr, expected) => val got = IndexedNearestByJoinRule.translateFilter(expr, attrs) assertEquals(Some(expected), got, s"translation mismatch for: $expr") } } + /** + * A column whose name collides with a SQL keyword/literal (`true`, `null`, `select`, …) must be + * emitted as a delimited identifier, NOT bare — otherwise `col = true` on a Boolean column named + * `true` would render as the tautology `true = true`, matching every row. This is the exact + * failure a "quote only non-word identifiers" exception cannot cover, so `quoteIdentifier` quotes + * unconditionally. + */ + @Test def testTranslatorQuotesSqlLiteralIdentifier(): Unit = { + val boolCol = makeAttr("true", BooleanType) + val attrs = AttributeSet(Seq(boolCol)) + assertEquals( + Some("`true` = true"), + IndexedNearestByJoinRule.translateFilter(EqualTo(boolCol, Literal(true)), attrs), + "a column named `true` must be quoted, not collapsed into a tautology") + } + /** Translator must return None for unsupported expressions so the rule refuses pushdown. */ @Test def testTranslatorRefusesUnsupportedShapes(): Unit = { val rid = makeAttr("rid", IntegerType) @@ -377,23 +404,24 @@ class IndexedNearestByJoinRuleTest { } /** - * Identifiers that aren't plain `[A-Za-z_][A-Za-z0-9_]*` — spaces, punctuation — must be - * double-quoted (embedded quotes doubled) so the Lance filter string is well-formed rather than - * malformed. Plain identifiers stay bare (covered by the shapes test above). Applies per - * dotted-path segment for nested struct fields. + * Identifiers with spaces, punctuation, or an embedded backtick must be back-quoted (an embedded + * backtick doubled) so the Lance filter string is well-formed rather than malformed. Every segment + * is quoted — including a nested struct field's own segments, independently — so `outer.inner + * field` becomes `` `outer`.`inner field` ``. */ @Test def testTranslatorQuotesUnsafeIdentifiers(): Unit = { val spaced = makeAttr("weird col", StringType) - val quoteName = makeAttr("has\"quote", IntegerType) + val tickName = makeAttr("has`tick", IntegerType) val outer = makeAttr("outer", new StructType().add("inner field", StringType)) - val attrs = AttributeSet(Seq(spaced, quoteName, outer)) + val attrs = AttributeSet(Seq(spaced, tickName, outer)) val cases: Seq[(Expression, String)] = Seq( - EqualTo(spaced, lit("A")) -> "\"weird col\" = 'A'", - GreaterThan(quoteName, lit(5)) -> "\"has\"\"quote\" > 5", - // nested field with a space -> only the unsafe segment gets quoted + EqualTo(spaced, lit("A")) -> "`weird col` = 'A'", + // embedded backtick in the identifier is doubled inside the delimiters + GreaterThan(tickName, lit(5)) -> "`has``tick` > 5", + // nested field with a space -> every segment quoted independently EqualTo(GetStructField(outer, 0, Some("inner field")), lit("A")) -> - "outer.\"inner field\" = 'A'") + "`outer`.`inner field` = 'A'") cases.foreach { case (expr, expected) => assertEquals( Some(expected), diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala index a5622ea8b..3186cbb38 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala @@ -13,7 +13,7 @@ */ package org.lance.spark.knn.catalyst -import org.apache.spark.sql.{DataFrame, Row, RowFactory, SparkSession} +import org.apache.spark.sql.{AnalysisException, DataFrame, Row, RowFactory, SparkSession} import org.apache.spark.sql.types._ import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} import org.junit.jupiter.api.Assertions._ @@ -278,6 +278,40 @@ class IndexedNearestByJoinSqlTest { } } + /** + * The rewrite runs as a `postHocResolutionRule`, BEFORE Spark's `FinishAnalysis` / + * `CheckCartesianProducts`. It must NOT let a query that Spark would reject slip through: with + * `spark.sql.crossJoin.enabled = false`, an `APPROX NEAREST` join (which lowers to a Cartesian + * product) must still fail with the standard `CROSS_JOIN_NOT_ENABLED` analysis error, exactly as + * it would without the extension. The rule declines the rewrite in this case, leaving the + * `NearestByJoin` for Spark's own path to reject. + */ + @Test def testIndexedRewritePreservesCrossJoinAnalysisGuard(): Unit = { + spark.conf.set("spark.sql.crossJoin.enabled", "false") + try { + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed + 400) + val (leftDf, _, _) = buildLeftDf(leftVecs, ExactDim) + val rightVecs = generateUniform(ExactRight, ExactDim, Seed + 401) + val (rightUri, _, _) = writeRightDf(rightVecs, ExactDim, idBase = 1000) + val (q, d) = registerViews(leftDf, rightUri) + val ex = assertThrows( + classOf[AnalysisException], + () => { + val df = spark.sql( + s"""SELECT q.lid, d.rid + |FROM $q q INNER JOIN $d d + |APPROX NEAREST 5 BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + // Force optimization — the Cartesian-product check runs in the optimizer, not analysis. + df.queryExecution.optimizedPlan + }) + assertTrue( + String.valueOf(ex.getMessage).toUpperCase.contains("CROSS"), + s"expected a cross-join analysis error; got: ${ex.getMessage}") + } finally { + spark.conf.set("spark.sql.crossJoin.enabled", "true") + } + } + // -- approximate path (IVF-PQ): recall floors, refineFactor ------------------------------- /** diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala index a7dfe96ae..de209a44a 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala @@ -17,6 +17,7 @@ import org.apache.spark.sql.Row import org.apache.spark.sql.types.{ArrayType, IntegerType, MapType, StringType, StructType} import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.Test +import org.lance.spark.{LanceRef, LanceSparkReadOptions} import java.util.concurrent.atomic.AtomicInteger @@ -112,6 +113,60 @@ class LanceKnnJoinStageTest { assertEquals(9, out("k").asInstanceOf[Row].getInt(0)) } + /** + * The shape a REAL Arrow map cell actually arrives in. Arrow encodes a `MapType` cell as a LIST + * of `{key, value}` entry structs, so [[LanceProbe.toSparkValue]] hands `coerceToSpark` a + * `Seq(Map("key" -> …, "value" -> …), …)` — a sequence, not a Scala map. Coercion must rebuild a + * real Spark map keyed/valued by type; leaving it a sequence is the exact bug the reviewer flagged + * (a `MapType` slot filled with a `Seq` fails the encoder or materializes garbage). + */ + @Test def testCoerceMapFromArrowEntryList(): Unit = { + val dt = MapType(StringType, IntegerType) + val arrowShape = Seq( + Map[String, Any]("key" -> "a", "value" -> Integer.valueOf(1)), + Map[String, Any]("key" -> "b", "value" -> Integer.valueOf(2))) + val out = + LanceKnnJoinStage.coerceToSpark(arrowShape, dt).asInstanceOf[scala.collection.Map[ + String, + Any]] + assertEquals(2, out.size, "both entries should survive") + assertEquals(1, out("a")) + assertEquals(2, out("b")) + } + + // -- fix #2: merge must reject a conflicting pinned ref ----------------------------------- + + /** + * When the base read options already carry a pinned `ref` and the relation options ALSO pin a + * different `version` / `branch`, the merge must REJECT the combination (as + * `LanceDataset.mergeScanOptions` does) rather than silently letting the relation value win — that + * would read a different snapshot than the caller pinned. Here base = `main@1`, relation = + * `version=2` (`main@2`), which is neither the same ref nor the same named branch, so the merge + * throws `IllegalArgumentException`. + */ + @Test def testMergeRejectsConflictingPinnedRef(): Unit = { + val base = LanceSparkReadOptions.from("/tmp/knn_merge_guard").withRef(LanceRef.ofMain(1L)) + val relation = new java.util.HashMap[String, String]() + relation.put(LanceSparkReadOptions.CONFIG_VERSION, "2") + assertThrows( + classOf[IllegalArgumentException], + () => LanceKnnJoinStage.mergeReadOptions(base, relation)) + } + + /** + * The same-named-branch case is allowed and keeps the table's pinned ref: base pinned to branch + * `dev` (version 1), relation re-specifies `branch=dev` — the merge must NOT throw and must retain + * the base ref for snapshot isolation. + */ + @Test def testMergeKeepsPinnedRefForSameNamedBranch(): Unit = { + val base = + LanceSparkReadOptions.from("/tmp/knn_merge_branch").withRef(LanceRef.ofBranch("dev", 1L)) + val relation = new java.util.HashMap[String, String]() + relation.put(LanceSparkReadOptions.CONFIG_BRANCH, "dev") + val merged = LanceKnnJoinStage.mergeReadOptions(base, relation) + assertEquals(base.getRef, merged.getRef, "same-named branch must keep the base's pinned ref") + } + /** Primitives pass straight through; null stays null. */ @Test def testCoercePrimitivesAndNullPassThrough(): Unit = { assertEquals("hello", LanceKnnJoinStage.coerceToSpark("hello", StringType)) diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala index 0ae287138..e24d8c09f 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala @@ -18,9 +18,10 @@ import org.apache.spark.sql.types._ import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.io.TempDir +import org.lance.spark.LanceSparkReadOptions import java.nio.file.Path -import java.util.Random +import java.util.{Collections, Random} import scala.collection.JavaConverters._ @@ -150,6 +151,39 @@ class LanceProbeValidationTest { } finally probe.close() } + /** + * When `executor_credential_refresh = false`, the probe must NOT rebuild (or even load) the + * runtime namespace on the worker — exactly the policy `LanceFragmentScanner.create` applies. + * Regression against the earlier `openDataset` that called `builder.runtimeNamespace(impl, ...)` + * unconditionally, which forced the namespace impl class to load regardless of the refresh flag. + * + * We pass a namespace impl that does not exist on the classpath and point the probe at a + * non-existent dataset URI. The open must fail because the DATASET is missing — not because it + * tried to load the bogus namespace class. Asserting the failure message does not mention the + * namespace class proves the namespace path was skipped. + */ + @Test def testExecutorCredentialRefreshFalseSkipsNamespaceRebuild(): Unit = { + val missingUri = tempDir.resolve("does_not_exist").toString + val readOptions = LanceSparkReadOptions + .builder() + .datasetUri(missingUri) + .executorCredentialRefresh(false) + .build() + val ex = assertThrows( + classOf[RuntimeException], + () => + new LanceProbe( + readOptions, + null, + "example.namespace.MustNotBeLoaded", + Collections.emptyMap[String, String](), + None)) + assertFalse( + String.valueOf(ex.getMessage).contains("MustNotBeLoaded"), + s"namespace impl must not be loaded when executor credential refresh is off; got: " + + ex.getMessage) + } + // -- helpers ------------------------------------------------------------------------------ /** Write a fresh dataset and return its file:// URI. */ From 491f463bc545c2c757145e2d8b149327799da674 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 12:48:40 -0700 Subject: [PATCH 5/9] knn-4.2: gate rewrite on fixed-size-vector metadata + schema-aware payload materialization MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Address two gatekeeper findings on PR #796: Finding A — require the connector's canonical fixed-size-vector metadata on the right attribute before rewriting. A plain variable-length List and a searchable Lance fixed-size-list vector both map to Spark ArrayType(FloatType); only the `arrow.fixed-size-list.size` field metadata distinguishes them. The rule now gates recognizeRanking on VectorUtils.isVectorField (the connector's canonical marker) and leaves the NearestByJoin unchanged otherwise, so a variable list falls through to Spark's brute-force path instead of handing Lance a column it cannot search. Regression: variableListVectorFallsBackInsteadOfFailing. Finding B — materialize projected right-side fields through the canonical connector Arrow-to-Spark adapter (LanceArrowColumnVector -> ColumnarBatch row -> CatalystTypeConverters external values) instead of raw Arrow getObject, so non-numeric payload columns (DateType, TimestampType) round-trip as the external java.sql.Date / java.sql.Timestamp the join's ExpressionEncoder expects. `_rowid` (not in the Spark schema) keeps the getObject fallback. Regression: dateAndTimestampPayloadMatchesOrdinarySparkTypes (payload parity vs an ordinary Spark Lance read). All 49 knn-4.2 module tests pass. Co-Authored-By: Claude Opus 4.8 --- .../catalyst/IndexedNearestByJoinRule.scala | 26 +++++- .../knn/internal/LanceKnnJoinStage.scala | 5 +- .../lance/spark/knn/internal/LanceProbe.scala | 87 ++++++++++++++++--- .../IndexedNearestByJoinRuleTest.scala | 58 ++++++++++++- .../IndexedNearestByJoinSqlTest.scala | 81 +++++++++++++++++ 5 files changed, 237 insertions(+), 20 deletions(-) diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala index 38aa3c55d..64edad2c7 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala @@ -22,6 +22,7 @@ import org.apache.spark.sql.types.{BooleanType, ByteType, DoubleType, FloatType, import org.apache.spark.unsafe.types.UTF8String import org.lance.spark.{LanceDataset, LanceSparkReadOptions} import org.lance.spark.knn.internal.{LanceKnnJoinStage, Metric} +import org.lance.spark.utils.VectorUtils /** * Catalyst rule that rewrites a Spark [[NearestByJoin]] (`approx = true`) over a Lance scan with @@ -481,14 +482,16 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { case VectorInnerProduct(l, r) if direction == NearestBySimilarity => (Metric.Dot, l, r) case _ => return None } - // Each argument must be a bare attribute from one side of the join. + // Each argument must be a bare attribute from one side of the join. The RIGHT (Lance) argument + // must additionally be a real fixed-size vector column — see [[isFixedSizeVector]] — otherwise + // there is nothing for the index to probe and we must fall through to the brute-force path. (asAttr(lhs), asAttr(rhs)) match { case (Some(la), Some(ra)) => val leftAttrIds = left.outputSet val rightAttrIds = right.outputSet - if (leftAttrIds.contains(la) && rightAttrIds.contains(ra)) { + if (leftAttrIds.contains(la) && rightAttrIds.contains(ra) && isFixedSizeVector(ra)) { Some((metric, la, ra.name)) - } else if (leftAttrIds.contains(ra) && rightAttrIds.contains(la)) { + } else if (leftAttrIds.contains(ra) && rightAttrIds.contains(la) && isFixedSizeVector(la)) { // Argument order swapped — still valid for symmetric metrics. All three of L2/Cosine/Dot // are symmetric so we don't have to retain the original orientation. Some((metric, ra, la.name)) @@ -503,4 +506,21 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { case a: Attribute => Some(a) case _ => None } + + /** + * True iff `attr` is the connector's canonical fixed-size vector shape: an `ArrayType` of float / + * double carrying the `arrow.fixed-size-list.size` metadata the connector stamps on a Lance + * `FixedSizeList` column. A plain variable-length `List` analyzes to the SAME Spark + * `ArrayType(FloatType)` but WITHOUT that metadata and is NOT index-searchable — probing it would + * fail at run time. Gating the rewrite on this marker (reusing the connector's own + * [[org.lance.spark.utils.VectorUtils.isVectorField]] so the definition lives in exactly one + * place) lets a non-vector right column fall through to Spark's brute-force path instead of + * producing a broken operator. A right attribute that is not an `AttributeReference` (so carries + * no field metadata) is conservatively treated as non-vector. + */ + private def isFixedSizeVector(attr: Attribute): Boolean = attr match { + case ref: AttributeReference => + VectorUtils.isVectorField(StructField(ref.name, ref.dataType, ref.nullable, ref.metadata)) + case _ => false + } } diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala index cfb4edd60..68df87bd1 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -262,7 +262,10 @@ object LanceKnnJoinStage { // `rowAddr -> row` map collapses any duplicate rowAddr to one payload; we still emit one // output row per surviving ref. Bounded by `k`, so this stays per-row, not per-partition. val materialized: Map[Long, Map[String, Any]] = probe - .materialize(trimmed.iterator.map(_.rowAddr).toSeq, conf.rightProjection) + .materialize( + trimmed.iterator.map(_.rowAddr).toSeq, + conf.rightProjection, + conf.rightFields) .map(m => extractRowAddr(m) -> m) .toMap trimmed.iterator.map { ref => diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala index 41c7b2f11..63c77acc1 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -16,10 +16,14 @@ package org.lance.spark.knn.internal import org.apache.arrow.memory.BufferAllocator import org.apache.arrow.vector.{BigIntVector, FieldVector, Float4Vector, Float8Vector, UInt8Vector, VectorSchemaRoot} import org.apache.arrow.vector.ipc.ArrowReader +import org.apache.spark.sql.catalyst.CatalystTypeConverters +import org.apache.spark.sql.types.{DataType, StructField} +import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} import org.lance.Dataset import org.lance.ipc.{LanceScanner, Query, ScanOptions} import org.lance.spark.{LanceConstant, LanceRef, LanceRuntime, LanceSparkReadOptions} import org.lance.spark.utils.Utils +import org.lance.spark.vectorized.LanceArrowColumnVector import java.util @@ -262,14 +266,20 @@ final class LanceProbe( * @param rowAddrs list of Lance `_rowid` values (parameter name retained for source * compatibility with callers — semantically these are now row IDs). * @param projection projected column list. `Seq.empty` means "all columns". - * @return a sequence of materialized rows, each represented as a `Map[String, Any]` for the - * projected columns plus an entry under `LanceProbe.RowIdColumn` so the caller can - * re-key. Returning a Map keeps this primitive Spark-agnostic; conversion to - * `InternalRow` happens in the API layer. + * @param projectionFields Spark target fields (name + dataType) for the projected columns. When + * non-empty, each projected payload cell is converted to the Spark EXTERNAL + * value its declared type expects (see [[readRows]]); when empty, cells fall + * back to a generic Arrow-object conversion. + * @return a sequence of materialized rows, each a `Map[String, Any]` keyed by column name. With + * `projectionFields` supplied, each projected value is already the Spark external + * representation of its target type — the shape the join's `ExpressionEncoder` accepts — + * plus an entry under `LanceProbe.RowIdColumn` (a plain long) so the caller can re-key. + * Building those values into a `Row` / `InternalRow` happens in the join stage. */ def materialize( rowAddrs: Seq[Long], - projection: Seq[String] = Seq.empty): Seq[Map[String, Any]] = { + projection: Seq[String] = Seq.empty, + projectionFields: Seq[StructField] = Seq.empty): Seq[Map[String, Any]] = { if (rowAddrs.isEmpty) return Seq.empty val opts = new ScanOptions.Builder().withRowId(true) @@ -305,28 +315,77 @@ final class LanceProbe( val scanner: LanceScanner = LanceScanner.create(dataset, opts.build(), allocator) try { - readRows(scanner.scanBatches()) + readRows(scanner.scanBatches(), projectionFields) } finally { scanner.close() } } - private def readRows(reader: ArrowReader): Seq[Map[String, Any]] = { + /** + * Drain the materialize scan into row maps, converting each projected payload cell to the Spark + * EXTERNAL value its target type expects — the shape the join's `ExpressionEncoder` (external + * `Row` → `InternalRow`) accepts. + * + * Projected data columns are materialized through the connector's canonical Arrow→Spark adapter + * [[org.lance.spark.vectorized.LanceArrowColumnVector]] — the same vector-and-schema-aware path + * the connector's own reader uses — so every payload type Lance can store round-trips to the + * value Spark produces for that column on an ordinary read: DateType, the various timestamp / time + * units, unsigned integers, decimals, fixed/large binary and varchar, and nested structs / lists / + * maps. A raw Arrow `getObject` would instead hand back e.g. a `LocalDate` for a DateType column, + * which the encoder then rejects. The adapter yields Spark INTERNAL values (days for a date, + * micros for a timestamp, …), so each is run back through [[CatalystTypeConverters]] to the + * external representation the `Row` encoder wants. + * + * Columns absent from `projectionFields` — notably the `_rowid` virtual column, which carries no + * Spark type here — keep the generic Arrow `getObject` fallback; the caller reads `_rowid` only as + * a plain long. + */ + private def readRows( + reader: ArrowReader, + projectionFields: Seq[StructField]): Seq[Map[String, Any]] = { + val schemaByName: Map[String, StructField] = + projectionFields.iterator.map(f => f.name -> f).toMap val out = mutable.ArrayBuffer.empty[Map[String, Any]] try { while (reader.loadNextBatch()) { val root: VectorSchemaRoot = reader.getVectorSchemaRoot val n = root.getRowCount + val fields = root.getSchema.getFields.asScala.toIndexedSeq + + // Columns with a Spark target type are converted through the canonical connector adapter; + // the rest (e.g. `_rowid`) keep the raw Arrow-object fallback. + val mapped = fields.filter(af => schemaByName.contains(af.getName)) + val unmapped = fields.filterNot(af => schemaByName.contains(af.getName)) + val mappedNames: Array[String] = mapped.iterator.map(_.getName).toArray + val mappedTypes: Array[DataType] = + mapped.iterator.map(af => schemaByName(af.getName).dataType).toArray + val mappedConverters: Array[Any => Any] = mapped.iterator + .map(af => + CatalystTypeConverters.createToScalaConverter(schemaByName(af.getName).dataType)) + .toArray + val mappedVectors: Array[ColumnVector] = mapped.iterator + .map(af => + new LanceArrowColumnVector(root.getVector(af.getName), false, schemaByName(af.getName)) + .asInstanceOf[ColumnVector]) + .toArray + // The batch is a thin view over the (reader-owned) Arrow vectors; closeVectorOnClose=false + // above means neither the batch nor its column vectors free the underlying buffers — the + // `reader.close()` in the finally does, once the values below have been copied out. + val batch = new ColumnarBatch(mappedVectors, n) + var i = 0 while (i < n) { val rowMap = mutable.LinkedHashMap.empty[String, Any] - val fields = root.getSchema.getFields.asScala - var f = 0 - while (f < fields.size) { - val name = fields(f).getName - val v = root.getVector(name) - rowMap(name) = if (v.isNull(i)) null else LanceProbe.toSparkValue(v.getObject(i)) - f += 1 + val internalRow = batch.getRow(i) + var j = 0 + while (j < mappedNames.length) { + val internal = internalRow.get(j, mappedTypes(j)) + rowMap(mappedNames(j)) = if (internal == null) null else mappedConverters(j)(internal) + j += 1 + } + unmapped.foreach { af => + val v = root.getVector(af.getName) + rowMap(af.getName) = if (v.isNull(i)) null else LanceProbe.toSparkValue(v.getObject(i)) } out += rowMap.toMap i += 1 diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala index ae09fde07..2ef95e26c 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala @@ -188,6 +188,45 @@ class IndexedNearestByJoinRuleTest { assertSame(join, rewritten, "non-Lance right must fall through") } + /** + * The right `rvec` column is a plain variable-length `List` — `ArrayType(FloatType)` with + * NO `arrow.fixed-size-list.size` metadata — not a searchable fixed-size vector. Both shapes map + * to the same Spark `ArrayType(FloatType)`, so only the connector's canonical fixed-size-list + * metadata distinguishes them. Lance can only index/probe a fixed-size-list vector, so the rule + * must require that marker on the right attribute and otherwise leave the `NearestByJoin` + * unchanged for Spark's brute-force path to own — rewriting a variable list would hand Lance a + * column it cannot search. + */ + @Test def variableListVectorFallsBackInsteadOfFailing(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val left = trivialPlan("lid", "lvec") + // rvec deliberately carries NO fixed-size-list metadata → a variable-length list. + val schema = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + StructField("rvec", ArrayType(FloatType, containsNull = false), nullable = false))) + val uri = tempDir.resolve("variable_list_lance").toString + val table = new FakeLanceTable(schema, uri) + val opts = new java.util.HashMap[String, String]() + opts.put("path", uri) + val cims = new org.apache.spark.sql.util.CaseInsensitiveStringMap(opts) + val right = DataSourceV2Relation.create(table, None, None, cims) + val leftVec = left.output.find(_.name == "lvec").get + val rightVec = right.output.find(_.name == "rvec").get + val join = NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + val rewritten = IndexedNearestByJoinRule(join) + assertSame( + join, + rewritten, + "variable-length List lacks fixed-size-vector metadata — rule must fall through") + } + /** Right side wrapped in SubqueryAlias still rewrites — alias unwrapping happens in the rule. */ @Test def testSubqueryAliasOnRightStillRewrites(): Unit = { spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") @@ -442,7 +481,7 @@ class IndexedNearestByJoinRuleTest { val left = trivialPlan("lid", "lvec") val schema = new StructType(Array( StructField("rid", IntegerType, nullable = false), - StructField("rvec", ArrayType(FloatType, containsNull = false), nullable = false))) + fixedSizeVectorField("rvec", 8))) val uri = tempDir.resolve("branch_lance").toString val table = new FakeLanceTable(schema, uri) val opts = new java.util.HashMap[String, String]() @@ -499,6 +538,21 @@ class IndexedNearestByJoinRuleTest { spark.createDataFrame(rows.asJava, schema).queryExecution.analyzed } + /** + * A `StructField` shaped exactly like the connector emits for a searchable Lance fixed-size-list + * vector column: `ArrayType(FloatType)` carrying the canonical `arrow.fixed-size-list.size` + * metadata key. `IndexedNearestByJoinRule` gates the rewrite on this marker (via + * `VectorUtils.isVectorField`), so the rule's rewrite-path scaffolds must stamp it — a plain + * `ArrayType(FloatType)` without it is a variable-length list Lance cannot search, and must fall + * through (see `variableListVectorFallsBackInsteadOfFailing`). + */ + private def fixedSizeVectorField(name: String, dim: Int): StructField = + StructField( + name, + ArrayType(FloatType, containsNull = false), + nullable = false, + new MetadataBuilder().putLong("arrow.fixed-size-list.size", dim.toLong).build()) + /** * Build a `DataSourceV2Relation` backed by a `FakeLanceTable` (a real connector `LanceDataset` * subclass) so the rule's `instanceof LanceDataset` check accepts it. We don't run any I/O — the @@ -511,7 +565,7 @@ class IndexedNearestByJoinRuleTest { StructField("rid", IntegerType, nullable = false), StructField("category", StringType, nullable = true), StructField("bucket", IntegerType, nullable = true), - StructField("rvec", ArrayType(FloatType, containsNull = false), nullable = false))) + fixedSizeVectorField("rvec", 8))) val uri = tempDir.resolve("fake_lance").toString val table = new FakeLanceTable(schema, uri) val opts = new java.util.HashMap[String, String]() diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala index 3186cbb38..613d3510e 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala @@ -278,6 +278,57 @@ class IndexedNearestByJoinSqlTest { } } + /** + * Payload parity for non-numeric projected columns. A projected `DateType` / `TimestampType` + * column must materialize through the indexed join EXACTLY as an ordinary Spark Lance read + * produces it. The join late-materializes right-side payloads by `_rowid` from native Arrow, so + * those cells must be shaped through the canonical connector Arrow-to-Spark adapter + * (vector-and-schema-aware) — NOT handed back as raw Arrow (`java.time.LocalDate` for a date, + * a boxed `Integer` day / `Long` micro). If the adapter is bypassed, the assembled row's external + * `java.sql.Date` / `java.sql.Timestamp` would diverge from the ordinary-read oracle (or fail to + * encode through the join's `ExpressionEncoder`). We build the oracle by reading the same columns + * back the ordinary way and assert every joined row's `(dt, ts)` equals the oracle for its `rid`. + */ + @Test def dateAndTimestampPayloadMatchesOrdinarySparkTypes(): Unit = { + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed + 500) + val (leftDf, _, _) = buildLeftDf(leftVecs, ExactDim) + val (_, _, rightUri) = writeRightWithDateTime(ExactRight, ExactDim, Seed + 501) + + // Oracle: ordinary Spark Lance read of the same columns -> external java.sql.Date/Timestamp. + val oracle: Map[Int, (java.sql.Date, java.sql.Timestamp)] = + spark.read.format("lance").load(rightUri).select("rid", "dt", "ts").collect().map { r => + r.getAs[Int]("rid") -> ((r.getAs[java.sql.Date]("dt"), r.getAs[java.sql.Timestamp]("ts"))) + }.toMap + + val k = 5 + val (q, d) = registerViews(leftDf, rightUri) + val df = spark.sql( + s"""SELECT q.lid, d.rid, d.dt, d.ts + |FROM $q q INNER JOIN $d d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + + val joinLogicals = df.queryExecution.optimizedPlan.collect { + case p: LanceKnnJoinLogicalPlan => p + } + assertTrue( + joinLogicals.nonEmpty, + s"expected LanceKnnJoinLogicalPlan; optimized plan was:\n${df.queryExecution.optimizedPlan}") + + val rows = df.collect() + assertEquals(ExactLeft * k, rows.length, "expected k results per left row") + rows.foreach { r => + val rid = r.getAs[Int]("rid") + val (expectedDt, expectedTs) = oracle(rid) + val actualDt = r.getAs[java.sql.Date]("dt") + val actualTs = r.getAs[java.sql.Timestamp]("ts") + // Guard against a silent all-null column masquerading as parity, then assert equality. + assertNotNull(actualDt, s"date payload unexpectedly null for rid=$rid") + assertNotNull(actualTs, s"timestamp payload unexpectedly null for rid=$rid") + assertEquals(expectedDt, actualDt, s"date payload mismatch for rid=$rid") + assertEquals(expectedTs, actualTs, s"timestamp payload mismatch for rid=$rid") + } + } + /** * The rewrite runs as a `postHocResolutionRule`, BEFORE Spark's `FinishAnalysis` / * `CheckCartesianProducts`. It must NOT let a query that Spark would reject slip through: with @@ -544,6 +595,36 @@ class IndexedNearestByJoinSqlTest { (vectors, ids, categories, out) } + /** + * Write a right Lance dataset whose rows carry a `DateType` and a `TimestampType` column + * alongside the vector, so the payload-parity test has non-numeric projected columns to + * round-trip through the indexed join's late materialization. Each row's date/timestamp is + * derived deterministically from its index. Returns the vectors, ids, and the dataset URI. + */ + private def writeRightWithDateTime( + n: Int, + dim: Int, + seed: Long): (Array[Array[Float]], Array[Int], String) = { + val vectors = generateUniform(n, dim, seed) + val ids = vectors.indices.map(_ + 3000).toArray + val rows = ids.zip(vectors).map { case (id, v) => + val i = id - 3000 + val dt = java.sql.Date.valueOf(java.time.LocalDate.of(2020, 1, 1).plusDays(i.toLong)) + val ts = java.sql.Timestamp.valueOf( + java.time.LocalDateTime.of(2020, 1, 1, 0, 0, 0).plusHours(i.toLong).plusSeconds(i.toLong)) + RowFactory.create(Integer.valueOf(id), dt, ts, v) + } + val schema = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + StructField("dt", DateType, nullable = false), + StructField("ts", TimestampType, nullable = false), + fixedSizeVec("rvec", dim))) + val df = spark.createDataFrame(rows.toSeq.asJava, schema) + val out = tempDir.resolve(s"right_dt_${System.nanoTime()}").toString + df.write.format("lance").save(out) + (vectors, ids, out) + } + /** Run an indexed nearest join (SQL path) against the given right dataset and compute recall@K. */ private def recallAgainst( leftDf: DataFrame, From 07f72151921d4a10835b02c01bbee539d149fa4f Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 13:19:16 -0700 Subject: [PATCH 6/9] Decline indexed KNN rewrite for blob-carrying Lance relations MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Blob columns are late-materialized: the connector's canonical reader threads the dataset URI, column name, and row addresses into BlobStructAccessor.setBlobReferenceContext so a blob descriptor resolves to its payload. The no-shuffle probe path fetches only _rowid and wraps vectors without that context, so a non-null legacy (v1) blob would resolve to an empty payload. Conservatively decline the rewrite when the Lance relation output carries any blob column (v1 lance-encoding:blob, or v2 lance.blob.v2 extension), detected via the connector's own BlobUtils.isBlobReadColumn. The query then falls through to Spark's brute-force cross-product, whose canonical blob-aware scan returns the true payload — so payload parity is owned by the fallback. Adds rule-level v1/v2 decline regression tests, each with a positive control (same schema minus the blob column rewrites) proving the blob column — not a missing fixed-size vector — is the discriminating cause. Module suite: 51 tests, 0 failures. Co-Authored-By: Claude Opus 4.8 --- .../catalyst/IndexedNearestByJoinRule.scala | 28 ++++- .../IndexedNearestByJoinRuleTest.scala | 116 ++++++++++++++++++ 2 files changed, 143 insertions(+), 1 deletion(-) diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala index 64edad2c7..67a18cac0 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala @@ -22,7 +22,7 @@ import org.apache.spark.sql.types.{BooleanType, ByteType, DoubleType, FloatType, import org.apache.spark.unsafe.types.UTF8String import org.lance.spark.{LanceDataset, LanceSparkReadOptions} import org.lance.spark.knn.internal.{LanceKnnJoinStage, Metric} -import org.lance.spark.utils.VectorUtils +import org.lance.spark.utils.{BlobUtils, VectorUtils} /** * Catalyst rule that rewrites a Spark [[NearestByJoin]] (`approx = true`) over a Lance scan with @@ -188,6 +188,7 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { for { (metric, leftVecAttr, rightVecCol) <- recognizeRanking(rankingExpr, direction, left, right) lance <- unwrapLanceScan(right) + if !hasBlobColumn(lance.output) } yield { val leftVecIdx = left.output.indexWhere(_.exprId == leftVecAttr.exprId) require(leftVecIdx >= 0, s"left vector attr not found in left.output: $leftVecAttr") @@ -523,4 +524,29 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { VectorUtils.isVectorField(StructField(ref.name, ref.dataType, ref.nullable, ref.metadata)) case _ => false } + + /** + * True iff any column in the Lance relation output is a blob column — legacy v1 + * (`lance-encoding:blob` metadata on a `BinaryType`) or v2 (`ARROW:extension:name = lance.blob.v2` + * metadata on the descriptor struct). Reuses the connector's own + * [[org.lance.spark.utils.BlobUtils.isBlobReadColumn]] so the marker definition lives in one place; + * both markers survive into the relation's `AttributeReference.metadata`. + * + * When present, the rule DECLINES the rewrite and falls through to Spark's brute-force cross- + * product. Blob columns are late-materialized: the connector's canonical reader + * ([[org.lance.spark.internal.LanceFragmentColumnarBatchScanner]]) threads the dataset URI, column + * name, and row addresses into `BlobStructAccessor.setBlobReferenceContext` so the descriptor can + * be resolved to its payload. The no-shuffle probe path fetches only `_rowid` and wraps vectors + * without that context, so a non-null legacy blob would resolve to an empty payload. Declining is + * the conservative, always-correct choice: Spark's own scan (with full blob context) returns the + * true payload. A non-`AttributeReference` attribute carries no field metadata and is treated as + * non-blob. + */ + private def hasBlobColumn(output: Seq[Attribute]): Boolean = + output.exists { + case ref: AttributeReference => + BlobUtils.isBlobReadColumn( + StructField(ref.name, ref.dataType, ref.nullable, ref.metadata)) + case _ => false + } } diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala index 2ef95e26c..4ad7ce066 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala @@ -25,6 +25,7 @@ import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.io.TempDir import org.lance.spark.knn.internal.Metric +import org.lance.spark.utils.BlobUtils import java.nio.file.Path @@ -41,6 +42,9 @@ import scala.collection.JavaConverters._ * - Direction mismatch (e.g. L2 distance with NearestBySimilarity) does NOT rewrite. * - EXACT (`approx = false`) does NOT rewrite — Spark's brute-force keeps owning that path. * - Non-Lance right side does NOT rewrite (right relation's table is not a `LanceDataset`). + * - A variable-length `List` right vector (no fixed-size-list metadata) does NOT rewrite. + * - A blob column (v1 or v2) on the right relation does NOT rewrite — even alongside a searchable + * vector — so blob payloads are materialized by Spark's canonical (blob-aware) fallback reader. * - Disabled by default — fires only when the gating config is set. * - Prefilter pushdown: right-side `WHERE` translates to a Lance SQL filter string, or refuses * the rewrite entirely when the predicate can't be pushed in full. @@ -227,6 +231,64 @@ class IndexedNearestByJoinRuleTest { "variable-length List lacks fixed-size-vector metadata — rule must fall through") } + /** + * A legacy (v1) blob column on the Lance relation forces the rule to DECLINE, even though the + * relation ALSO carries a searchable fixed-size vector. Blob columns are late-materialized: the + * connector's canonical reader threads dataset-URI / column-name / row-address context into + * `BlobStructAccessor.setBlobReferenceContext` so a descriptor resolves to its payload. The + * no-shuffle probe path fetches only `_rowid` and wraps vectors WITHOUT that context, so a + * non-null legacy blob would resolve to an empty payload. Declining hands the query to Spark's + * brute-force cross-product, whose canonical (blob-aware) scan returns the true payload — so the + * payload parity is owned by the fallback, and this test locks in that we take it. + * + * The positive control (same schema MINUS the blob column) rewrites, proving the fixed-size-vector + * gate is satisfied and the blob column is the sole discriminating cause of the decline. + */ + @Test def blobV1ColumnDeclinesRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val left = trivialPlan("lid", "lvec") + val baseFields = Array( + StructField("rid", IntegerType, nullable = false), + fixedSizeVectorField("rvec", 8)) + // Control: a searchable vector-only schema WITHOUT any blob column must rewrite. + val control = lanceRelationWithSchema(new StructType(baseFields), "blob_v1_control") + assertTrue( + IndexedNearestByJoinRule(l2Join(left, control)).isInstanceOf[Project], + "control: vector-only schema must rewrite (proves the vector gate passes)") + // Same schema + a legacy v1 blob column → decline. + val withBlob = + lanceRelationWithSchema(new StructType(baseFields :+ blobV1Field("payload")), "blob_v1") + val join = l2Join(left, withBlob) + assertSame( + join, + IndexedNearestByJoinRule(join), + "legacy v1 blob column must force fallback to Spark's canonical (blob-aware) reader") + } + + /** + * As [[blobV1ColumnDeclinesRewrite]] but for a blob v2 column — carried as the descriptor struct + * with the `ARROW:extension:name = lance.blob.v2` marker. The same late-materialization reasoning + * applies, so the rule must decline and leave payload materialization to Spark's canonical reader. + */ + @Test def blobV2ColumnDeclinesRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val left = trivialPlan("lid", "lvec") + val baseFields = Array( + StructField("rid", IntegerType, nullable = false), + fixedSizeVectorField("rvec", 8)) + val control = lanceRelationWithSchema(new StructType(baseFields), "blob_v2_control") + assertTrue( + IndexedNearestByJoinRule(l2Join(left, control)).isInstanceOf[Project], + "control: vector-only schema must rewrite (proves the vector gate passes)") + val withBlob = + lanceRelationWithSchema(new StructType(baseFields :+ blobV2Field("payload")), "blob_v2") + val join = l2Join(left, withBlob) + assertSame( + join, + IndexedNearestByJoinRule(join), + "blob v2 descriptor column must force fallback to Spark's canonical (blob-aware) reader") + } + /** Right side wrapped in SubqueryAlias still rewrites — alias unwrapping happens in the rule. */ @Test def testSubqueryAliasOnRightStillRewrites(): Unit = { spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") @@ -553,6 +615,60 @@ class IndexedNearestByJoinRuleTest { nullable = false, new MetadataBuilder().putLong("arrow.fixed-size-list.size", dim.toLong).build()) + /** + * A `StructField` shaped like the connector emits for a LEGACY (v1) blob column: a `BinaryType` + * carrying the `lance-encoding:blob = true` metadata that `BlobUtils.isBlobReadColumn` keys on. + * The marker survives into the relation's `AttributeReference.metadata`, which is what the rule's + * `hasBlobColumn` gate inspects. + */ + private def blobV1Field(name: String): StructField = + StructField( + name, + BinaryType, + nullable = true, + new MetadataBuilder() + .putString(BlobUtils.LANCE_ENCODING_BLOB_KEY, BlobUtils.LANCE_ENCODING_BLOB_VALUE) + .build()) + + /** + * A `StructField` shaped like the connector emits for a blob v2 column at read time: the + * `BLOB_DESCRIPTOR_STRUCT` data type carrying the `ARROW:extension:name = lance.blob.v2` extension + * metadata `BlobUtils.isBlobV2SparkField` keys on. (Detection is metadata-based, so the descriptor + * struct is used only to faithfully mirror the real relation output.) + */ + private def blobV2Field(name: String): StructField = + StructField( + name, + BlobUtils.BLOB_DESCRIPTOR_STRUCT, + nullable = true, + new MetadataBuilder() + .putString(BlobUtils.ARROW_EXTENSION_NAME_KEY, BlobUtils.ARROW_EXTENSION_BLOB_V2) + .build()) + + /** Build a Lance-backed DSv2 relation over the given schema (no I/O — `FakeLanceTable` is inert). */ + private def lanceRelationWithSchema(schema: StructType, dirName: String): LogicalPlan = { + val uri = tempDir.resolve(dirName).toString + val table = new FakeLanceTable(schema, uri) + val opts = new java.util.HashMap[String, String]() + opts.put("path", uri) + val cims = new org.apache.spark.sql.util.CaseInsensitiveStringMap(opts) + DataSourceV2Relation.create(table, None, None, cims) + } + + /** Build an `approx` L2 `NearestByJoin` over `left.lvec` and `right.rvec` (numResults = 5). */ + private def l2Join(left: LogicalPlan, right: LogicalPlan): NearestByJoin = { + val leftVec = left.output.find(_.name == "lvec").get + val rightVec = right.output.find(_.name == "rvec").get + NearestByJoin( + left, + right, + Inner, + approx = true, + numResults = 5, + rankingExpression = VectorL2Distance(leftVec, rightVec), + direction = NearestByDistance) + } + /** * Build a `DataSourceV2Relation` backed by a `FakeLanceTable` (a real connector `LanceDataset` * subclass) so the rule's `instanceof LanceDataset` check accepts it. We don't run any I/O — the From b42d049225badff783253b6415245223495d94c8 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 13:47:24 -0700 Subject: [PATCH 7/9] knn: fold probe+materialize into one scan on the no-overfetch path MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit On the SQL path the rule sets internalK == k, so every probed row is kept and the separate late-materialization scan just re-fetches the exact rows the search already found — a second Lance scan plus a k-element `_rowid IN (arrow_cast(...))` filter parse per query, for nothing. Add LanceProbe.probeRows, which runs the nearest search AND projects the payload columns in a single scan, returning MaterializedHit(rowAddr, score, row). LanceKnnJoinStage.processRow uses it when internalK <= k and keeps the existing probe -> trim -> materialize path for a future over-fetching caller (internalK > k), where deferring the payload fetch to survivors is the win. Factor the shared nearest-query builder and the _rowid / score vector readers so probe and probeRows search identically. Add a probeRows-vs- (probe+materialize) parity test asserting identical (rowAddr, score) hits and identical projected payload. Document the fixed-size-vector and blob right-side fallback triggers. 52 tests pass. Co-Authored-By: Claude Opus 4.8 --- .../operations/dql/nearest-neighbor-join.md | 13 + .../knn/internal/LanceKnnJoinStage.scala | 45 +++- .../lance/spark/knn/internal/LanceProbe.scala | 225 +++++++++++++++--- .../spark/knn/internal/ScoredRowRef.scala | 15 ++ .../internal/LanceProbeValidationTest.scala | 54 +++++ 5 files changed, 306 insertions(+), 46 deletions(-) diff --git a/docs/src/operations/dql/nearest-neighbor-join.md b/docs/src/operations/dql/nearest-neighbor-join.md index 9ef4304f8..6cd7f68fd 100644 --- a/docs/src/operations/dql/nearest-neighbor-join.md +++ b/docs/src/operations/dql/nearest-neighbor-join.md @@ -135,6 +135,19 @@ are not rewritten and fall through to Spark's brute-force path. Only `APPROX` joins are rewritten. An exact (`EXACT`) nearest join is always handled by Spark's brute-force rewrite. +## Right-Side Table Requirements + +Beyond the ranking function, the rewrite inspects the right (Lance) table's columns and declines — +falling through to Spark's brute-force `APPROX NEAREST` path, with identical results — when: + +- **The vector column is not a fixed-size vector.** The right ranking column must be a fixed-size + list (an `ARRAY` written with the `arrow.fixed-size-list.size` schema hint, the shape Lance + builds a vector index over). A variable-length `ARRAY` column is not a probeable vector and + is not rewritten. +- **The table has a blob column.** If any column in the scanned Lance relation is a blob + (`lance-encoding:blob` v1, or the `lance.blob.v2` extension type), the rewrite declines so that + Spark's canonical Lance reader materializes the true blob payload on the fallback path. + ## WHERE Pushdown A `WHERE` clause on the right (Lance) side is translated into a Lance filter and applied by the diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala index 68df87bd1..b72d06030 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -34,10 +34,13 @@ import org.lance.spark.utils.Utils * left.rdd.mapPartitions { rows => * val probe = new LanceProbe(uri, fragmentIds = None, version) // whole-index, once per task * rows.flatMap { leftRow => - * val hits = probe.probe(query(leftRow), internalK, ...) // native top-K search - * val topK = trimToK(hits) // overfetch → K - * val payloads = probe.materialize(topK.map(_.rowAddr)) // late point-fetch by _rowid - * topK.map(ref => assembleRow(leftRow, payloads(ref), ref.score)) + * // No over-fetch (SQL path, internalK == k): search + project payload in ONE scan. + * val hits = probe.probeRows(query(leftRow), k, projection, ...) + * hits.map(hit => assembleRow(leftRow, hit.row, hit.score)) + * // Over-fetch path (internalK > k): search cheap refs → trim → materialize survivors. + * // val topK = trimToK(probe.probe(query, internalK, ...)) + * // val payloads = probe.materialize(topK.map(_.rowAddr)) + * // topK.map(ref => assembleRow(leftRow, payloads(ref), ref.score)) * } * } * }}} @@ -228,10 +231,38 @@ object LanceKnnJoinStage { } else { Iterator.empty } + } else if (conf.internalK <= conf.k) { + // No JVM-side over-fetch (the SQL path: internalK == k). Every probed row is kept, so probe + // and materialize in ONE native scan: Lance searches the index AND projects the payload + // columns in a single pass. A split probe → trim → materialize would re-scan the exact rows + // the search already found, for nothing. Lance returns them best-first, so no trim is needed. + val hits = probe.probeRows( + conf.vectorColumn, + q, + conf.internalK, + conf.metric, + conf.rightProjection, + conf.rightFields, + conf.nprobes, + conf.refineFactor, + conf.ef, + conf.prefilter) + if (hits.isEmpty) { + if (conf.outerJoin) { + Iterator.single(assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, null, null)) + } else { + Iterator.empty + } + } else { + hits.iterator.map { hit => + assembleRow(leftRow, conf.leftFieldCount, conf.rightFields, hit.row, hit.score) + } + } } else { - // Overfetch `internalK` candidates natively, then trim to the final `k`. Lance already - // returns them best-first, so when it hands back no more than `k` we keep them as-is and - // skip the heap entirely. + // Over-fetch path (internalK > k): fetch `internalK` cheap refs natively, trim to the final + // `k`, then late-materialize ONLY the survivors by `_rowid`. Not exercised by the SQL rule + // (which sets internalK == k) but retained for a future over-fetching caller, where deferring + // the payload fetch to the survivors is the win. Lance already returns refs best-first. val refs = probe .probe( conf.vectorColumn, diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala index 63c77acc1..98e200f9a 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -31,9 +31,14 @@ import scala.collection.JavaConverters._ import scala.collection.mutable /** - * Per-task vector-index probe primitive. Opens a Lance dataset once and runs many `probe()` calls - * against a fixed set of fragments. Returns row references + scores only — no payload. Late - * materialization happens elsewhere (`LanceMaterialize`). + * Per-task vector-index probe primitive. Opens a Lance dataset once and serves many queries against + * a fixed set of fragments. Two query shapes: + * - [[probe]] returns row references + scores only (no payload); the payload is fetched later via + * [[materialize]]. This split late-materialization is the right shape when the caller OVER-FETCHES + * candidates and trims before materializing — only the survivors are paid for. + * - [[probeRows]] folds the nearest search AND the payload projection into a SINGLE scan. This is + * the right shape when there is NO JVM-side over-fetch (every probed row is kept), so a separate + * materialize scan would just re-fetch the exact rows the search already found. * * This is the core primitive Phase 0 of the indexed nearest-by design depends on. Validating its * cost profile is the first thing to do on a new Lance build: @@ -155,22 +160,7 @@ final class LanceProbe( require(query != null && query.length > 0, "Query vector must be non-empty") require(k > 0, "k must be positive") - val q = { - val b = new Query.Builder() - .setColumn(vectorColumn) - .setKey(query) - .setK(k) - .setDistanceType(metric.lanceType) - nprobes.foreach(b.setNprobes(_)) - // refineFactor: IVF-PQ recall knob. Lance fetches `k * refineFactor` approximate - // candidates, then re-ranks them with exact distance and trims to k. Bigger factor = - // better recall, more compute. None leaves Lance's default (= 1, no re-rank). - refineFactor.foreach(b.setRefineFactor(_)) - // ef: HNSW search depth. Higher = better recall, more compute. None leaves Lance's - // index-default. Only meaningful for HNSW indexes; ignored for IVF-PQ. - ef.foreach(b.setEf(_)) - b.build() - } + val q = buildNearestQuery(vectorColumn, query, k, metric, nprobes, refineFactor, ef) val opts = new ScanOptions.Builder() .nearest(q) @@ -204,6 +194,92 @@ final class LanceProbe( } } + /** + * Nearest search AND payload projection in a SINGLE scan. Runs the top-`k` search projecting the + * requested payload `projection` columns directly into the result batch, so each hit comes back + * with its row-id, ranking score, AND materialized payload — no second point-fetch scan. + * + * This is the no-overfetch fast path. When the caller does not over-fetch at the JVM level (the + * SQL path: `internalK == k`), every probed row is kept, so the split [[probe]] → trim → + * [[materialize]] would just re-scan the exact rows the search already found — a second Lance scan + * plus a `k`-element `_rowid IN (arrow_cast …)` filter parse, per query, for nothing. Folding the + * two removes both. Rows come back best-first (Lance's native ordering), same as [[probe]]. + * + * Use [[probe]] + [[materialize]] instead when the caller OVER-fetches candidates and trims before + * materializing — there, deferring the payload fetch to only the survivors is the win. `projection` + * / `projectionFields` follow [[materialize]]'s contract: an empty `projection` means "all columns", + * and each projected cell is converted to the Spark EXTERNAL value its declared type expects (see + * [[readRows]]). + */ + def probeRows( + vectorColumn: String, + query: Array[Float], + k: Int, + metric: Metric, + projection: Seq[String], + projectionFields: Seq[StructField], + nprobes: Option[Int] = None, + refineFactor: Option[Int] = None, + ef: Option[Int] = None, + prefilter: Option[String] = None): Seq[MaterializedHit] = { + require(vectorColumn != null && vectorColumn.nonEmpty, "vectorColumn must be non-empty") + require(query != null && query.length > 0, "Query vector must be non-empty") + require(k > 0, "k must be positive") + + val q = buildNearestQuery(vectorColumn, query, k, metric, nprobes, refineFactor, ef) + + val opts = new ScanOptions.Builder() + .nearest(q) + .withRowId(true) + // Project the payload columns into the nearest scan itself. `_distance` is added by `nearest` + // regardless (Lance includes it even when explicit columns omit it), so the drain still finds a + // score column. An empty projection leaves columns unset → all columns, matching `materialize`. + if (projection.nonEmpty) { + opts.columns(projection.toList.asJava) + } + + if (prefilter.nonEmpty || javaFragmentIds.nonEmpty) { + opts.prefilter(true) + } + prefilter.filter(_.nonEmpty).foreach(opts.filter) + javaFragmentIds.foreach(opts.fragmentIds) + + val scanner: LanceScanner = LanceScanner.create(dataset, opts.build(), allocator) + try { + readScoredRows(scanner.scanBatches(), projectionFields) + } finally { + scanner.close() + } + } + + /** + * Build the Lance nearest-neighbor query. Shared by [[probe]] (refs only) and [[probeRows]] + * (refs + folded payload) so both search the index identically. + */ + private def buildNearestQuery( + vectorColumn: String, + query: Array[Float], + k: Int, + metric: Metric, + nprobes: Option[Int], + refineFactor: Option[Int], + ef: Option[Int]): Query = { + val b = new Query.Builder() + .setColumn(vectorColumn) + .setKey(query) + .setK(k) + .setDistanceType(metric.lanceType) + nprobes.foreach(b.setNprobes(_)) + // refineFactor: IVF-PQ recall knob. Lance fetches `k * refineFactor` approximate + // candidates, then re-ranks them with exact distance and trims to k. Bigger factor = + // better recall, more compute. None leaves Lance's default (= 1, no re-rank). + refineFactor.foreach(b.setRefineFactor(_)) + // ef: HNSW search depth. Higher = better recall, more compute. None leaves Lance's + // index-default. Only meaningful for HNSW indexes; ignored for IVF-PQ. + ef.foreach(b.setEf(_)) + b.build() + } + /** * Drain the Arrow stream from a nearest-search scan into `(rowId, score)` pairs. * @@ -221,31 +297,75 @@ final class LanceProbe( while (reader.loadNextBatch()) { val root = reader.getVectorSchemaRoot val addrVec: FieldVector = root.getVector(LanceProbe.RowIdColumn) - val scoreVec: FieldVector = LanceProbe.ScoreColumns.iterator - .map(name => Option(root.getVector(name)).orNull) - .find(_ != null) - .getOrElse(throw new IllegalStateException( - s"Lance nearest scan did not return a score column. Got: " + - root.getSchema.getFields.asScala.map(_.getName).mkString(", "))) + val scoreVec: FieldVector = resolveScoreVector(root) val n = root.getRowCount var i = 0 while (i < n) { - val addr = addrVec match { - case v: UInt8Vector => v.get(i) - case v: BigIntVector => v.get(i) - case other => - throw new IllegalStateException( - s"Unexpected row-address vector type: ${other.getClass.getName}") - } - val score = scoreVec match { - case v: Float4Vector => v.get(i) - case v: Float8Vector => v.get(i).toFloat - case other => - throw new IllegalStateException( - s"Unexpected score vector type: ${other.getClass.getName}") + out += ScoredRowRef(rowAddrAt(addrVec, i), scoreAt(scoreVec, i)) + i += 1 + } + } + } finally { + reader.close() + } + out.toSeq + } + + /** + * Drain a nearest-search scan that ALSO projected payload columns into `(rowId, score, payload)` + * hits — the folded counterpart of [[readScored]] + [[readRows]]. Row-id and score are read out of + * the `_rowid` / `_distance` vectors directly; the payload columns (those with a Spark target type + * in `projectionFields`) go through the canonical [[org.lance.spark.vectorized.LanceArrowColumnVector]] + * adapter and back through [[CatalystTypeConverters]] to the external value the encoder expects — + * exactly as [[readRows]] does. `_rowid` / `_distance` are handled here, so they are deliberately + * excluded from the payload map (unlike [[readRows]], which surfaces `_rowid` for re-keying). + */ + private def readScoredRows( + reader: ArrowReader, + projectionFields: Seq[StructField]): Seq[MaterializedHit] = { + val schemaByName: Map[String, StructField] = + projectionFields.iterator.map(f => f.name -> f).toMap + val out = mutable.ArrayBuffer.empty[MaterializedHit] + try { + while (reader.loadNextBatch()) { + val root: VectorSchemaRoot = reader.getVectorSchemaRoot + val n = root.getRowCount + val fields = root.getSchema.getFields.asScala.toIndexedSeq + + val addrVec: FieldVector = root.getVector(LanceProbe.RowIdColumn) + val scoreVec: FieldVector = resolveScoreVector(root) + + // Payload columns only (those with a Spark target type). `_rowid` / `_distance` are read out + // separately, so they are deliberately excluded from the payload map here. + val mapped = fields.filter(af => schemaByName.contains(af.getName)) + val mappedNames: Array[String] = mapped.iterator.map(_.getName).toArray + val mappedTypes: Array[DataType] = + mapped.iterator.map(af => schemaByName(af.getName).dataType).toArray + val mappedConverters: Array[Any => Any] = mapped.iterator + .map(af => + CatalystTypeConverters.createToScalaConverter(schemaByName(af.getName).dataType)) + .toArray + val mappedVectors: Array[ColumnVector] = mapped.iterator + .map(af => + new LanceArrowColumnVector(root.getVector(af.getName), false, schemaByName(af.getName)) + .asInstanceOf[ColumnVector]) + .toArray + // Thin view over the reader-owned Arrow vectors (closeVectorOnClose=false): `reader.close()` + // in the finally frees the buffers once the values below are copied out. + val batch = new ColumnarBatch(mappedVectors, n) + + var i = 0 + while (i < n) { + val rowMap = mutable.LinkedHashMap.empty[String, Any] + val internalRow = batch.getRow(i) + var j = 0 + while (j < mappedNames.length) { + val internal = internalRow.get(j, mappedTypes(j)) + rowMap(mappedNames(j)) = if (internal == null) null else mappedConverters(j)(internal) + j += 1 } - out += ScoredRowRef(addr, score) + out += MaterializedHit(rowAddrAt(addrVec, i), scoreAt(scoreVec, i), rowMap.toMap) i += 1 } } @@ -255,6 +375,33 @@ final class LanceProbe( out.toSeq } + /** Locate the nearest-search score column by name, tolerant of `_distance` vs `_score`. */ + private def resolveScoreVector(root: VectorSchemaRoot): FieldVector = + LanceProbe.ScoreColumns.iterator + .map(name => Option(root.getVector(name)).orNull) + .find(_ != null) + .getOrElse(throw new IllegalStateException( + "Lance nearest scan did not return a score column. Got: " + + root.getSchema.getFields.asScala.map(_.getName).mkString(", "))) + + /** Read a Lance row id out of an Arrow `_rowid` vector (UInt8 or BigInt, encoding-dependent). */ + private def rowAddrAt(addrVec: FieldVector, i: Int): Long = addrVec match { + case v: UInt8Vector => v.get(i) + case v: BigIntVector => v.get(i) + case other => + throw new IllegalStateException( + s"Unexpected row-address vector type: ${other.getClass.getName}") + } + + /** Read a ranking score out of an Arrow score vector (Float4 or Float8, encoding-dependent). */ + private def scoreAt(scoreVec: FieldVector, i: Int): Float = scoreVec match { + case v: Float4Vector => v.get(i) + case v: Float8Vector => v.get(i).toFloat + case other => + throw new IllegalStateException( + s"Unexpected score vector type: ${other.getClass.getName}") + } + /** * Materialize a set of right-side rows by their `_rowaddr`s. Used by the join's materialize * stage to fetch full payloads after the probe + merge has decided which rows survive. diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala index ac1b194fc..bacd984c9 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/ScoredRowRef.scala @@ -38,3 +38,18 @@ object ScoredRowRef { val similarityOrdering: Ordering[ScoredRowRef] = Ordering.by[ScoredRowRef, Float](-_.score) } + +/** + * A single probe hit WITH its materialized payload — the unit produced by the no-overfetch fast path + * ([[LanceProbe.probeRows]]), which folds the nearest search and the payload projection into one + * scan. Unlike [[ScoredRowRef]] (a payload-free ref bound for a later materialize point-fetch), this + * already carries the right row. + * + * @param rowAddr Lance row id of the hit (kept for de-duplication / re-keying by the join stage). + * @param score Distance or similarity from Lance's vector search; direction is carried out-of-band + * in the operator config, so this stays metric-agnostic (see [[ScoredRowRef.score]]). + * @param row Materialized right-side payload keyed by column name, each value already the Spark + * EXTERNAL representation its target type expects (the join's `ExpressionEncoder` + * accepts it directly). + */ +final case class MaterializedHit(rowAddr: Long, score: Float, row: Map[String, Any]) diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala index e24d8c09f..f7a2f7516 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala @@ -121,6 +121,53 @@ class LanceProbeValidationTest { } } + /** + * The folded fast path ([[LanceProbe.probeRows]]) must be observationally identical to the split + * probe + materialize path: the SAME (rowAddr, score) hits in the SAME order, and the SAME + * materialized payload per row. This is the invariant the SQL join relies on when it single-scans + * (internalK == k) instead of probing then late-materializing. Runs on the brute-force path so the + * search itself is deterministic. + */ + @Test def testProbeRowsMatchesProbeThenMaterialize(): Unit = { + val datasetUri = writeSyntheticDataset() + val query = randomVector(new Random(321L), VectorDim) + val k = 8 + val projection = Seq("id", "vec") + val projectionFields = Seq( + StructField("id", IntegerType, nullable = false), + StructField("vec", ArrayType(FloatType, containsNull = false), nullable = false)) + + val probe = new LanceProbe(datasetUri, fragmentIds = None) + try { + // Split path: search for refs, then late-materialize the payload by _rowid. + val refs = probe.probe("vec", query, k, Metric.L2) + val expectedPayload: Map[Long, Map[String, Any]] = probe + .materialize(refs.map(_.rowAddr), projection, projectionFields) + .map(m => rowAddrOf(m) -> m) + .toMap + + // Folded path: search AND project the payload in one scan. + val hits = probe.probeRows("vec", query, k, Metric.L2, projection, projectionFields) + + assertEquals(k, hits.size, "probeRows should return exactly k hits") + // Same search: identical (rowAddr, score) sequence, in order. + assertEquals( + refs.map(r => (r.rowAddr, r.score)), + hits.map(h => (h.rowAddr, h.score)), + "probeRows hits must match probe refs (rowAddr + score), in order") + // Same payload: each folded hit equals the split materialize's row for that rowAddr, on the + // projected columns. + hits.foreach { h => + val expected = expectedPayload(h.rowAddr) + assertEquals(expected("id"), h.row("id"), s"id mismatch for rowAddr=${h.rowAddr}") + assertEquals( + expected("vec").asInstanceOf[Seq[_]].toList, + h.row("vec").asInstanceOf[Seq[_]].toList, + s"vec mismatch for rowAddr=${h.rowAddr}") + } + } finally probe.close() + } + /** * Validate the dataset handle is reused across calls. The exact perf invariant ("second call * faster than first by some factor") is too brittle for CI, so we only assert that repeated @@ -216,6 +263,13 @@ class LanceProbeValidationTest { (rows, vectors) } + /** Read the `_rowid` key out of a materialized row map (a boxed / stringy long). */ + private def rowAddrOf(m: Map[String, Any]): Long = m(LanceProbe.RowIdColumn) match { + case l: java.lang.Long => l.longValue() + case l: Long => l + case other => other.toString.toLong + } + private def randomVector(rng: Random, dim: Int): Array[Float] = { val v = new Array[Float](dim) var i = 0 From 261e7e3329ec2916030f3bed2904d1c7c49b5785 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 14:33:01 -0700 Subject: [PATCH 8/9] fix(knn): correct metric direction, NaN admission, payload fold, reserved-name routing Carry the probe-core review fixes into the umbrella revision and close the fold's headline finding at the stage layer: - Metric: cosine/dot are DISTANCES in Lance (1 - similarity), so smaller is better for all three metrics; flip smallerIsBetter and fix the doc. - TopKHeap: admit by the metric Ordering (ord.lt), not a raw float compare, so a NaN score sorts as worst-and-evictable instead of silently winning. - LanceProbe.probeRows: preserve projected payload columns that lack a Spark type via the generic Arrow fallback; guard + fusesCleanly reject a projection that collides with the injected _rowid / _distance / _score columns. - LanceKnnJoinStage: route reserved-name projections away from the fold to the split probe -> materialize path (materialize injects only _rowid), via the new unit-testable foldsInOneScan predicate. Tests: +3 fold-routing regressions in LanceKnnJoinStageTest, +metric-direction and reserved-name regressions carried from probe-core. Full module: 59 pass. Co-Authored-By: Claude Opus 4.8 --- .bumpversion.toml | 5 + .../knn/internal/LanceKnnJoinStage.scala | 27 +++-- .../lance/spark/knn/internal/LanceProbe.scala | 57 +++++++++-- .../org/lance/spark/knn/internal/Metric.scala | 17 +++- .../lance/spark/knn/internal/TopKHeap.scala | 10 +- .../knn/internal/LanceKnnJoinStageTest.scala | 34 +++++++ .../internal/LanceProbeValidationTest.scala | 99 +++++++++++++++++++ .../spark/knn/internal/TopKHeapTest.scala | 16 +++ 8 files changed, 243 insertions(+), 22 deletions(-) diff --git a/.bumpversion.toml b/.bumpversion.toml index 82e0d1485..2043101ea 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -85,6 +85,11 @@ filename = "lance-spark-4.2_2.13/pom.xml" search = "{current_version}" replace = "{new_version}" +[[tool.bumpversion.files]] +filename = "lance-spark-knn-4.2_2.13/pom.xml" +search = "{current_version}" +replace = "{new_version}" + # Bundle module pom.xml files - parent version [[tool.bumpversion.files]] filename = "lance-spark-bundle-3.4_2.12/pom.xml" diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala index b72d06030..9f655d3c6 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -216,6 +216,18 @@ object LanceKnnJoinStage { expand: Row => Iterator[Row]): Iterator[Row] = leftRows.flatMap(expand) + /** + * Whether [[processRow]] may take the folded one-scan fast path. True only when there is no + * JVM-side over-fetch (`internalK <= k`, so every probed row is kept and no trim is needed) AND + * the payload projection does not name a column the nearest scan injects itself — `_rowid`, + * `_distance`, `_score` (see [[LanceProbe.fusesCleanly]]). A colliding projection must fall + * through to the split probe → materialize path, whose scan injects only `_rowid` and therefore + * reads any `_distance` / `_score` payload column straight from the dataset without collision. + * Extracted so the routing decision is unit-testable without a Lance dataset. + */ + private[knn] def foldsInOneScan(internalK: Int, k: Int, projection: Seq[String]): Boolean = + internalK <= k && LanceProbe.fusesCleanly(projection) + /** * Expand one left row into its join output rows: probe R's index, trim to `k`, late-materialize * the surviving right rows by `_rowid`, and assemble `left ++ right ++ score`. Returns an empty @@ -231,8 +243,9 @@ object LanceKnnJoinStage { } else { Iterator.empty } - } else if (conf.internalK <= conf.k) { - // No JVM-side over-fetch (the SQL path: internalK == k). Every probed row is kept, so probe + } else if (foldsInOneScan(conf.internalK, conf.k, conf.rightProjection)) { + // No JVM-side over-fetch (the SQL path: internalK == k) AND the payload projection does not + // collide with an injected column (see [[foldsInOneScan]]). Every probed row is kept, so probe // and materialize in ONE native scan: Lance searches the index AND projects the payload // columns in a single pass. A split probe → trim → materialize would re-scan the exact rows // the search already found, for nothing. Lance returns them best-first, so no trim is needed. @@ -259,10 +272,12 @@ object LanceKnnJoinStage { } } } else { - // Over-fetch path (internalK > k): fetch `internalK` cheap refs natively, trim to the final - // `k`, then late-materialize ONLY the survivors by `_rowid`. Not exercised by the SQL rule - // (which sets internalK == k) but retained for a future over-fetching caller, where deferring - // the payload fetch to the survivors is the win. Lance already returns refs best-first. + // Split probe → trim → materialize path. Taken when the caller over-fetches (internalK > k) + // OR when the payload projection collides with an injected column name and the fold above + // had to decline. Fetch `internalK` cheap refs natively, trim to the final `k`, then + // late-materialize ONLY the survivors by `_rowid`. With internalK == k (the SQL rule) nothing + // is trimmed; the split still buys collision-safety, since materialize injects only `_rowid`. + // Lance already returns refs best-first. val refs = probe .probe( conf.vectorColumn, diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala index 98e200f9a..6f7466147 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -225,6 +225,17 @@ final class LanceProbe( require(vectorColumn != null && vectorColumn.nonEmpty, "vectorColumn must be non-empty") require(query != null && query.length > 0, "Query vector must be non-empty") require(k > 0, "k must be positive") + // The fused scan injects `_rowid` (via withRowId) and the score column (via nearest). A payload + // column that shares one of those names would collide with the injected column inside the SAME + // scan — Lance fails the scan with "merge incompatible fields". Such a schema must go through the + // split probe + materialize path (whose materialize scan injects no score column), so reject it + // here rather than let the collision surface as an opaque Arrow error. The join stage checks + // [[LanceProbe.fusesCleanly]] and routes accordingly. + require( + LanceProbe.fusesCleanly(projection), + "probeRows cannot project a column whose name collides with Lance search metadata " + + s"(${LanceProbe.ReservedProjectionColumns.toSeq.sorted.mkString(", ")}); the nearest scan " + + "injects those columns itself. Route such schemas through the split probe + materialize path.") val q = buildNearestQuery(vectorColumn, query, k, metric, nprobes, refineFactor, ef) @@ -315,11 +326,13 @@ final class LanceProbe( /** * Drain a nearest-search scan that ALSO projected payload columns into `(rowId, score, payload)` * hits — the folded counterpart of [[readScored]] + [[readRows]]. Row-id and score are read out of - * the `_rowid` / `_distance` vectors directly; the payload columns (those with a Spark target type - * in `projectionFields`) go through the canonical [[org.lance.spark.vectorized.LanceArrowColumnVector]] - * adapter and back through [[CatalystTypeConverters]] to the external value the encoder expects — - * exactly as [[readRows]] does. `_rowid` / `_distance` are handled here, so they are deliberately - * excluded from the payload map (unlike [[readRows]], which surfaces `_rowid` for re-keying). + * the `_rowid` / score vectors directly and excluded from the payload map. Every OTHER column the + * scan returned is payload and follows [[readRows]]' contract exactly: a column with a Spark target + * type in `projectionFields` goes through the canonical + * [[org.lance.spark.vectorized.LanceArrowColumnVector]] adapter and back through + * [[CatalystTypeConverters]] to the external value the encoder expects; a projected column WITHOUT + * a supplied type falls back to the generic Arrow conversion ([[LanceProbe.toSparkValue]]) rather + * than being silently dropped. */ private def readScoredRows( reader: ArrowReader, @@ -336,9 +349,15 @@ final class LanceProbe( val addrVec: FieldVector = root.getVector(LanceProbe.RowIdColumn) val scoreVec: FieldVector = resolveScoreVector(root) - // Payload columns only (those with a Spark target type). `_rowid` / `_distance` are read out - // separately, so they are deliberately excluded from the payload map here. - val mapped = fields.filter(af => schemaByName.contains(af.getName)) + // `_rowid` and the injected score column are read out-of-band (above) and excluded from the + // payload. Everything else the scan returned IS payload: columns with a Spark target type go + // through the canonical adapter, the rest fall back to the generic Arrow conversion — the + // same split `readRows` / `materialize` apply, so a projected column with no supplied type is + // surfaced (via `toSparkValue`) instead of being silently dropped. + val reserved = Set(LanceProbe.RowIdColumn, scoreVec.getField.getName) + val payloadFields = fields.filterNot(af => reserved.contains(af.getName)) + val mapped = payloadFields.filter(af => schemaByName.contains(af.getName)) + val unmapped = payloadFields.filterNot(af => schemaByName.contains(af.getName)) val mappedNames: Array[String] = mapped.iterator.map(_.getName).toArray val mappedTypes: Array[DataType] = mapped.iterator.map(af => schemaByName(af.getName).dataType).toArray @@ -365,6 +384,10 @@ final class LanceProbe( rowMap(mappedNames(j)) = if (internal == null) null else mappedConverters(j)(internal) j += 1 } + unmapped.foreach { af => + val v = root.getVector(af.getName) + rowMap(af.getName) = if (v.isNull(i)) null else LanceProbe.toSparkValue(v.getObject(i)) + } out += MaterializedHit(rowAddrAt(addrVec, i), scoreAt(scoreVec, i), rowMap.toMap) i += 1 } @@ -566,6 +589,24 @@ object LanceProbe { */ val ScoreColumns: Seq[String] = Seq("_distance", "_score") + /** + * Column names Lance's nearest scan injects itself: `_rowid` (from `withRowId`) and the score + * column (from `nearest`). A right-side payload column with one of these names collides with the + * injected column inside the fused [[LanceProbe.probeRows]] scan, which Lance rejects with a + * "merge incompatible fields" error. A schema that projects one of them must be served through the + * split [[LanceProbe.probe]] + [[LanceProbe.materialize]] path instead — the materialize scan does + * not inject a score column, so there is no collision. See [[fusesCleanly]]. + */ + val ReservedProjectionColumns: Set[String] = ScoreColumns.toSet + RowIdColumn + + /** + * True if `projection` can be served by the fused [[LanceProbe.probeRows]] scan — i.e. it names no + * column the nearest scan injects (see [[ReservedProjectionColumns]]). The join stage calls this to + * decide between the fold and the split probe + materialize path. + */ + def fusesCleanly(projection: Seq[String]): Boolean = + !projection.exists(ReservedProjectionColumns.contains) + /** * Build read options from a bare dataset URI, pinning `version` on the main branch when present. * Backs the URI convenience constructor. diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala index 798bcedd4..2a5ed1726 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/Metric.scala @@ -18,10 +18,15 @@ import org.lance.index.DistanceType /** * Vector distance / similarity metric. Mirrors `org.lance.index.DistanceType` but exposed as a * Scala enumeration so callers don't have to import Lance internals. Each metric fixes the - * "best-first" direction used during merge: + * "best-first" direction used during merge. * - * - L2: smaller score is better (distance) - * - Cosine / Dot: larger score is better (similarity) + * Lance's vector search returns a DISTANCE for every metric — including the similarity-flavored + * ones, which it reports as `1 - cosine_similarity` and `1 - dot_product`. So smaller is better for + * ALL three metrics; there is no larger-is-better case: + * + * - L2: smaller score is better (squared L2 distance) + * - Cosine: smaller score is better (`1 - cosine_similarity`) + * - Dot: smaller score is better (`1 - dot_product`) */ sealed trait Metric { @@ -41,12 +46,14 @@ object Metric { case object Cosine extends Metric { val lanceType: DistanceType = DistanceType.Cosine - val smallerIsBetter: Boolean = false + // Lance returns `1 - cosine_similarity`, a distance: smaller is better. + val smallerIsBetter: Boolean = true } case object Dot extends Metric { val lanceType: DistanceType = DistanceType.Dot - val smallerIsBetter: Boolean = false + // Lance returns `1 - dot_product`, a distance: smaller is better. + val smallerIsBetter: Boolean = true } /** diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala index 673179c8b..7c562517f 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/TopKHeap.scala @@ -51,9 +51,13 @@ final class TopKHeap(k: Int, smallerIsBetter: Boolean) { heap.enqueue(ref) } else { val worst = heap.head - val isBetter = - if (smallerIsBetter) ref.score < worst.score - else ref.score > worst.score + // Admission must use the SAME total ordering as the heap, not a raw float `<` / `>`. Float + // comparisons involving NaN are always false, so a NaN worst-survivor would never test as + // "beatable" and could never be evicted — a single NaN score would then pin a slot forever. + // `ord` is `java.lang.Float.compare`-based (NaN sorts as the largest Float, i.e. the worst + // for a distance), and `ord.lt(ref, worst)` means exactly "ref ranks strictly better than the + // current worst" in both directions — so a finite score correctly displaces a NaN worst. + val isBetter = ord.lt(ref, worst) if (isBetter) { heap.dequeue() heap.enqueue(ref) diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala index de209a44a..8aae70968 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala @@ -167,6 +167,40 @@ class LanceKnnJoinStageTest { assertEquals(base.getRef, merged.getRef, "same-named branch must keep the base's pinned ref") } + // -- fold vs split routing: reserved-name projections must decline the fold ---------------- + + /** + * The folded one-scan path is sound only when there is no over-fetch (internalK == k) AND the + * payload projection does not collide with a column the nearest scan injects. A plain projection + * at internalK == k folds. + */ + @Test def testFoldsInOneScanForCleanProjectionAtNoOverfetch(): Unit = { + assertTrue( + LanceKnnJoinStage.foldsInOneScan(internalK = 5, k = 5, projection = Seq("id", "vec")), + "clean projection with internalK == k should fold in one scan") + } + + /** + * A projection naming ANY reserved column (`_rowid`, `_distance`, `_score`) must decline the fold + * even at internalK == k, so the stage routes it to the split probe → materialize path whose scan + * injects only `_rowid`. This is the gatekeeper's headline finding: the fold must preserve Spark's + * fallback for a payload column whose name collides with the injected search metadata. + */ + @Test def testDoesNotFoldForReservedProjection(): Unit = { + LanceProbe.ReservedProjectionColumns.foreach { reserved => + assertFalse( + LanceKnnJoinStage.foldsInOneScan(internalK = 5, k = 5, projection = Seq("id", reserved)), + s"projection naming reserved column '$reserved' must NOT fold; route to the split path") + } + } + + /** Over-fetch (internalK > k) always takes the split path, regardless of projection. */ + @Test def testDoesNotFoldWhenOverfetching(): Unit = { + assertFalse( + LanceKnnJoinStage.foldsInOneScan(internalK = 20, k = 5, projection = Seq("id")), + "over-fetch (internalK > k) must take the split probe → trim → materialize path") + } + /** Primitives pass straight through; null stays null. */ @Test def testCoercePrimitivesAndNullPassThrough(): Unit = { assertEquals("hello", LanceKnnJoinStage.coerceToSpark("hello", StringType)) diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala index f7a2f7516..575caac04 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala @@ -121,6 +121,105 @@ class LanceProbeValidationTest { } } + /** + * Cosine and Dot must rank best-first the same direction L2 does. Lance returns a DISTANCE for + * every metric (`1 - cosine_similarity`, `1 - dot_product` for the similarity-flavored ones), so + * smaller is better for all three — [[Metric.smallerIsBetter]] must be `true` for each. This + * reproduces the gatekeeper's failure directly: run a real Lance query, then merge the results + * through the size-1 [[TopKHeap]] the join stage uses, keyed by the metric's own direction flag. + * The nearest ref (Lance returns best-first, so `refs.head`) must survive; a wrong flag (treating a + * Lance distance as larger-is-better) would retain the FARTHEST ref instead. + */ + @Test def testMetricFlagsKeepNearestThroughHeap(): Unit = { + val datasetUri = writeSyntheticDataset() + val query = randomVector(new Random(555L), VectorDim) + + val probe = new LanceProbe(datasetUri, fragmentIds = None) + try { + Seq[Metric](Metric.L2, Metric.Cosine, Metric.Dot).foreach { metric => + val refs = probe.probe("vec", query, k = 5, metric) + assertEquals(5, refs.size, s"$metric probe should return k results") + // Lance returns best-first, so refs.head is the true nearest for this metric. + val nearest = refs.head + val heap = new TopKHeap(k = 1, metric.smallerIsBetter) + heap.offerAll(refs) + val survivor = heap.drain() + assertEquals(1, survivor.length, s"$metric: size-1 heap should retain one ref") + assertEquals( + nearest.rowAddr, + survivor.head.rowAddr, + s"$metric: size-1 heap must keep the nearest ref (rowAddr=${nearest.rowAddr}), " + + s"got ${survivor.head.rowAddr} — wrong smallerIsBetter direction?") + assertEquals(nearest.score, survivor.head.score, 1e-6f, s"$metric: kept score mismatch") + } + } finally probe.close() + } + + /** + * A projected payload column WITHOUT a supplied Spark type must be preserved through the generic + * Arrow conversion — the same fallback [[LanceProbe.materialize]] / `readRows` apply — not silently + * dropped. Regression: `projection = Seq("id")` with empty `projectionFields` must return payload + * keys `Set("id")` (the injected `_rowid` / score columns stay out of the payload). + */ + @Test def testProbeRowsPreservesUnmappedProjectedFields(): Unit = { + val datasetUri = writeSyntheticDataset() + val query = randomVector(new Random(321L), VectorDim) + + val probe = new LanceProbe(datasetUri, fragmentIds = None) + try { + val hits = probe.probeRows( + "vec", + query, + k = 5, + Metric.L2, + projection = Seq("id"), + projectionFields = Seq.empty) + assertEquals(5, hits.size, "probeRows should return k hits") + hits.foreach { h => + assertEquals( + Set("id"), + h.row.keySet, + s"unmapped projected field must be preserved (id only, no _rowid/_distance); " + + s"got ${h.row.keySet}") + assertNotNull(h.row("id"), "unmapped id payload must be populated") + } + } finally probe.close() + } + + /** + * The fused [[LanceProbe.probeRows]] scan injects `_rowid` and the score column itself, so a + * payload projection naming one of those reserved columns would collide inside the single scan. + * `probeRows` must reject such a projection (so the join stage can route it through the split + * probe + materialize path) rather than surfacing an opaque Arrow "merge incompatible fields" + * error, and [[LanceProbe.fusesCleanly]] must report it as not fusible. + */ + @Test def testProbeRowsRejectsReservedProjection(): Unit = { + val datasetUri = writeSyntheticDataset() + val query = randomVector(new Random(1L), VectorDim) + + val probe = new LanceProbe(datasetUri, fragmentIds = None) + try { + LanceProbe.ReservedProjectionColumns.foreach { reserved => + assertFalse( + LanceProbe.fusesCleanly(Seq("id", reserved)), + s"projection naming reserved column '$reserved' must not be fusible") + val ex = assertThrows( + classOf[IllegalArgumentException], + () => + probe.probeRows( + "vec", + query, + k = 5, + Metric.L2, + projection = Seq("id", reserved), + projectionFields = Seq.empty)) + assertTrue( + String.valueOf(ex.getMessage).contains(reserved), + s"guard message should name reserved column '$reserved'; got: ${ex.getMessage}") + } + } finally probe.close() + } + /** * The folded fast path ([[LanceProbe.probeRows]]) must be observationally identical to the split * probe + materialize path: the SAME (rowAddr, score) hits in the SAME order, and the SAME diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala index 41bd1de7d..33481b4f8 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/TopKHeapTest.scala @@ -75,6 +75,22 @@ class TopKHeapTest { assertEquals(Seq(1L, 2L), drained.map(_.rowAddr).toSeq) } + /** + * A NaN score must not pin a heap slot forever. Admission derives from the heap's total ordering + * (`java.lang.Float.compare`, which sorts NaN as the largest Float), not a raw float `<` — with + * raw `<`, `1.0f < NaN` is false, so a NaN worst-survivor would test as un-beatable and never be + * evicted. Regression for exactly that: offer NaN then a finite score into a size-1 distance heap; + * the finite score must win. + */ + @Test def testNaNScoreIsEvictable(): Unit = { + val heap = new TopKHeap(k = 1, smallerIsBetter = true) + heap.offer(ref(1, Float.NaN)) + heap.offer(ref(2, 1.0f)) + val drained = heap.drain() + assertEquals(Seq(2L), drained.map(_.rowAddr).toSeq, "finite score must evict the NaN worst") + assertEquals(1.0f, drained.head.score, "size-1 distance heap must retain the finite minimum") + } + /** `merge` combines two pre-sorted arrays preserving top-K. */ @Test def testMergeCombinesTwoArrays(): Unit = { val a = Array(ref(1, 1f), ref(2, 3f), ref(3, 5f)) From 124f5cf9eb4e4ab435deceaa9820ef648081e111 Mon Sep 17 00:00:00 2001 From: Eunjin Song Date: Wed, 26 Aug 2026 18:34:38 -0700 Subject: [PATCH 9/9] fix(knn): decline indexed rewrite on schema with reserved column names MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Make the nearest eligibility contract schema-level, not projection-level, end to end (rule + stage + probe). Lance's nearest scan always injects _rowid and the _distance/_score metadata. If the right table's own schema already has a column by one of those names, the injected metadata shadows it and NO indexed route recovers the physical column: an all-columns fused scan reads the user's _distance out-of-band as the ranking score and silently drops it from the payload — data loss, not an error. Fold-vs-split routing does not help; both scans inject _rowid. - IndexedNearestByJoinRule: add hasReservedColumn, a decline guard modeled on the existing hasBlobColumn — when the right schema owns a reserved-named column the rule leaves the NearestByJoin in place so Spark's brute-force nearest-by (whose ordinary scan returns the true payload) handles it. - LanceProbe: schemaSupportsNearest / reservedSchemaColumns eligibility primitive + requireNearestCompatibleSchema() backstop called at the top of probe()/probeRows(), superseding the projection-only fusesCleanly guard (which missed the empty/all-columns form). - LanceKnnJoinStage.foldsInOneScan: drop the removed fusesCleanly term — reserved-column tables are declined upstream, so the fold decision is now purely about over-fetch (internalK <= k). Tests: - rule: reservedColumnDeclinesRewrite — a right schema owning each reserved name declines, with a vector-only control that rewrites. - SQL e2e (real Lance): sqlUserDistanceColumnDeclinesAndReturnsStoredPayload — a table with a physical _distance column declines and Spark's fallback returns the true stored value, not the search score. - probe: real _distance-column dataset regression — probe()/all-columns probeRows() fail fast naming _distance instead of dropping it; pure schemaSupportsNearest / reservedSchemaColumns contract test. - stage routing tests reworked to the over-fetch-only decision. Co-Authored-By: Claude Opus 4.8 --- .../catalyst/IndexedNearestByJoinRule.scala | 19 ++- .../knn/internal/LanceKnnJoinStage.scala | 39 +++--- .../lance/spark/knn/internal/LanceProbe.scala | 83 +++++++++--- .../IndexedNearestByJoinRuleTest.scala | 41 +++++- .../IndexedNearestByJoinSqlTest.scala | 77 +++++++++++ .../knn/internal/LanceKnnJoinStageTest.scala | 35 ++--- .../internal/LanceProbeValidationTest.scala | 120 ++++++++++++++---- 7 files changed, 324 insertions(+), 90 deletions(-) diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala index 67a18cac0..cc751e07d 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRule.scala @@ -21,7 +21,7 @@ import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation import org.apache.spark.sql.types.{BooleanType, ByteType, DoubleType, FloatType, IntegerType, LongType, ShortType, StringType, StructField, StructType} import org.apache.spark.unsafe.types.UTF8String import org.lance.spark.{LanceDataset, LanceSparkReadOptions} -import org.lance.spark.knn.internal.{LanceKnnJoinStage, Metric} +import org.lance.spark.knn.internal.{LanceKnnJoinStage, LanceProbe, Metric} import org.lance.spark.utils.{BlobUtils, VectorUtils} /** @@ -189,6 +189,7 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { (metric, leftVecAttr, rightVecCol) <- recognizeRanking(rankingExpr, direction, left, right) lance <- unwrapLanceScan(right) if !hasBlobColumn(lance.output) + if !hasReservedColumn(lance.output) } yield { val leftVecIdx = left.output.indexWhere(_.exprId == leftVecAttr.exprId) require(leftVecIdx >= 0, s"left vector attr not found in left.output: $leftVecAttr") @@ -549,4 +550,20 @@ object IndexedNearestByJoinRule extends Rule[LogicalPlan] { StructField(ref.name, ref.dataType, ref.nullable, ref.metadata)) case _ => false } + + /** + * True iff the Lance relation output has a column whose name collides with the metadata a nearest + * scan injects — `_rowid`, `_distance`, `_score` (see [[LanceProbe.schemaSupportsNearest]]). + * + * When present, the rule DECLINES the rewrite and falls through to Spark's brute-force cross- + * product. Every indexed route runs a `nearest` scan that injects those columns, so the injected + * metadata would shadow the physical column: an all-columns fused scan reads the user's `_distance` + * out-of-band as the ranking score and silently drops it from the payload, and an explicit + * projection of it collides outright. No fold-vs-split routing recovers it — the eligibility is + * schema-level. Declining is the conservative, always-correct choice: Spark's own nearest-by + * rewrite returns the true payload including that column. [[LanceProbe]] enforces the same contract + * defensively at probe time. + */ + private def hasReservedColumn(output: Seq[Attribute]): Boolean = + !LanceProbe.schemaSupportsNearest(output.map(_.name)) } diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala index 9f655d3c6..ddee2c525 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceKnnJoinStage.scala @@ -217,16 +217,19 @@ object LanceKnnJoinStage { leftRows.flatMap(expand) /** - * Whether [[processRow]] may take the folded one-scan fast path. True only when there is no - * JVM-side over-fetch (`internalK <= k`, so every probed row is kept and no trim is needed) AND - * the payload projection does not name a column the nearest scan injects itself — `_rowid`, - * `_distance`, `_score` (see [[LanceProbe.fusesCleanly]]). A colliding projection must fall - * through to the split probe → materialize path, whose scan injects only `_rowid` and therefore - * reads any `_distance` / `_score` payload column straight from the dataset without collision. - * Extracted so the routing decision is unit-testable without a Lance dataset. + * Whether [[processRow]] may take the folded one-scan fast path. True when there is no JVM-side + * over-fetch (`internalK <= k`, so every probed row is kept and no trim is needed), so probe and + * materialize collapse into a single native scan; false when the caller over-fetches candidates + * (`internalK > k`) and must trim before paying to materialize only the survivors. + * + * This decision is ONLY about over-fetch. Reserved-column collisions (a right schema owning + * `_rowid` / `_distance` / `_score`) are handled upstream: the Catalyst rule declines the indexed + * rewrite for such a table (see [[LanceProbe.schemaSupportsNearest]]) and [[LanceProbe]] enforces + * the same contract defensively — so by the time a row reaches here the projection is always + * fusible. Extracted so the routing decision is unit-testable without a Lance dataset. */ - private[knn] def foldsInOneScan(internalK: Int, k: Int, projection: Seq[String]): Boolean = - internalK <= k && LanceProbe.fusesCleanly(projection) + private[knn] def foldsInOneScan(internalK: Int, k: Int): Boolean = + internalK <= k /** * Expand one left row into its join output rows: probe R's index, trim to `k`, late-materialize @@ -243,10 +246,10 @@ object LanceKnnJoinStage { } else { Iterator.empty } - } else if (foldsInOneScan(conf.internalK, conf.k, conf.rightProjection)) { - // No JVM-side over-fetch (the SQL path: internalK == k) AND the payload projection does not - // collide with an injected column (see [[foldsInOneScan]]). Every probed row is kept, so probe - // and materialize in ONE native scan: Lance searches the index AND projects the payload + } else if (foldsInOneScan(conf.internalK, conf.k)) { + // No JVM-side over-fetch (the SQL path: internalK == k); see [[foldsInOneScan]]. Every probed + // row is kept, so probe and materialize in ONE native scan: Lance searches the index AND + // projects the payload // columns in a single pass. A split probe → trim → materialize would re-scan the exact rows // the search already found, for nothing. Lance returns them best-first, so no trim is needed. val hits = probe.probeRows( @@ -272,12 +275,10 @@ object LanceKnnJoinStage { } } } else { - // Split probe → trim → materialize path. Taken when the caller over-fetches (internalK > k) - // OR when the payload projection collides with an injected column name and the fold above - // had to decline. Fetch `internalK` cheap refs natively, trim to the final `k`, then - // late-materialize ONLY the survivors by `_rowid`. With internalK == k (the SQL rule) nothing - // is trimmed; the split still buys collision-safety, since materialize injects only `_rowid`. - // Lance already returns refs best-first. + // Split probe → trim → materialize path. Taken when the caller over-fetches (internalK > k): + // fetch `internalK` cheap refs natively, trim to the final `k` with the top-K heap, then + // late-materialize ONLY the survivors by `_rowid` — so the payload fetch is paid for just the + // rows that make the cut. Lance already returns refs best-first. val refs = probe .probe( conf.vectorColumn, diff --git a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala index 6f7466147..5dd557f5e 100644 --- a/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala +++ b/lance-spark-knn-4.2_2.13/src/main/scala/org/lance/spark/knn/internal/LanceProbe.scala @@ -125,6 +125,33 @@ final class LanceProbe( builder.build() } + // The dataset's own top-level column names, read once from the open handle. Used by the schema + // eligibility backstop below; cheap (metadata only) and stable for the probe's lifetime. + private lazy val datasetColumnNames: Seq[String] = + dataset.getSchema.getFields.asScala.map(_.getName).toSeq + + /** + * Reject a dataset whose own schema collides with a column the nearest scan injects + * (`_rowid` / `_distance` / `_score`). Every probe route ([[probe]] and [[probeRows]], fused or + * split) runs a `nearest` scan that injects those columns, so such a table cannot be served by ANY + * indexed route: the injected metadata shadows the physical column and it is read out-of-band as + * the ranking score (silently dropped from an all-columns payload) or collides outright when + * projected. There is no fold-vs-split routing that recovers it. The Catalyst rule is expected to + * DECLINE the indexed rewrite for such a table via [[LanceProbe.schemaSupportsNearest]] and fall + * back to the engine's default nearest-by execution; this is the defensive backstop for any caller + * that reached the probe anyway. + */ + private def requireNearestCompatibleSchema(): Unit = { + val collisions = LanceProbe.reservedSchemaColumns(datasetColumnNames) + require( + collisions.isEmpty, + s"Lance dataset schema has column(s) ${collisions.toSeq.sorted.mkString(", ")} whose name(s) " + + "collide with the metadata a nearest scan injects (_rowid, _distance, _score). No indexed " + + "probe can serve this table — the injected metadata shadows the physical column. Decline the " + + "indexed rewrite (see LanceProbe.schemaSupportsNearest) and fall back to default nearest-by " + + "execution.") + } + /** * Run a single nearest-neighbor query. Returns up to `k` row references for the configured * fragments, ordered best-first by `metric`. @@ -159,6 +186,7 @@ final class LanceProbe( require(vectorColumn != null && vectorColumn.nonEmpty, "vectorColumn must be non-empty") require(query != null && query.length > 0, "Query vector must be non-empty") require(k > 0, "k must be positive") + requireNearestCompatibleSchema() val q = buildNearestQuery(vectorColumn, query, k, metric, nprobes, refineFactor, ef) @@ -225,17 +253,15 @@ final class LanceProbe( require(vectorColumn != null && vectorColumn.nonEmpty, "vectorColumn must be non-empty") require(query != null && query.length > 0, "Query vector must be non-empty") require(k > 0, "k must be positive") - // The fused scan injects `_rowid` (via withRowId) and the score column (via nearest). A payload - // column that shares one of those names would collide with the injected column inside the SAME - // scan — Lance fails the scan with "merge incompatible fields". Such a schema must go through the - // split probe + materialize path (whose materialize scan injects no score column), so reject it - // here rather than let the collision surface as an opaque Arrow error. The join stage checks - // [[LanceProbe.fusesCleanly]] and routes accordingly. - require( - LanceProbe.fusesCleanly(projection), - "probeRows cannot project a column whose name collides with Lance search metadata " + - s"(${LanceProbe.ReservedProjectionColumns.toSeq.sorted.mkString(", ")}); the nearest scan " + - "injects those columns itself. Route such schemas through the split probe + materialize path.") + // The nearest scan injects `_rowid` (via withRowId) and the score columns (via nearest). If the + // dataset's own schema has a column by one of those names the injected metadata shadows it, and + // NO projection shape recovers the physical column: an empty (all-columns) projection reads it + // out-of-band as the ranking score and silently drops it from the payload, while an explicit + // projection of it collides inside the scan. Routing to the split path does not help — its scan + // also injects `_rowid`. So the eligibility is schema-level, not projection-level: reject the + // whole table. The Catalyst rule declines such a table up front via + // [[LanceProbe.schemaSupportsNearest]]; this is the defensive backstop. + requireNearestCompatibleSchema() val q = buildNearestQuery(vectorColumn, query, k, metric, nprobes, refineFactor, ef) @@ -591,21 +617,36 @@ object LanceProbe { /** * Column names Lance's nearest scan injects itself: `_rowid` (from `withRowId`) and the score - * column (from `nearest`). A right-side payload column with one of these names collides with the - * injected column inside the fused [[LanceProbe.probeRows]] scan, which Lance rejects with a - * "merge incompatible fields" error. A schema that projects one of them must be served through the - * split [[LanceProbe.probe]] + [[LanceProbe.materialize]] path instead — the materialize scan does - * not inject a score column, so there is no collision. See [[fusesCleanly]]. + * columns (from `nearest`). These are metadata a nearest scan always produces; a right-side + * table column sharing one of these names collides with the injected column. See + * [[reservedSchemaColumns]] / [[schemaSupportsNearest]]. */ val ReservedProjectionColumns: Set[String] = ScoreColumns.toSet + RowIdColumn /** - * True if `projection` can be served by the fused [[LanceProbe.probeRows]] scan — i.e. it names no - * column the nearest scan injects (see [[ReservedProjectionColumns]]). The join stage calls this to - * decide between the fold and the split probe + materialize path. + * The subset of `schemaColumnNames` that collide with a column the nearest scan injects + * (`_rowid` / `_distance` / `_score`). Empty means the table is nearest-compatible. + */ + def reservedSchemaColumns(schemaColumnNames: Iterable[String]): Set[String] = + schemaColumnNames.iterator.filter(ReservedProjectionColumns.contains).toSet + + /** + * Whether an indexed nearest scan can run against a right-side table with these column names. + * + * Lance's nearest scan ALWAYS injects `_rowid` and the `_distance` / `_score` metadata. If the + * table's own schema already has a column by one of those names, the injected metadata shadows + * it and no projection shape recovers the physical column: the fused scan reads it out-of-band as + * the ranking score and silently drops it from the payload (observed with an empty / all-columns + * projection), and an explicit projection of it collides outright. There is no safe fold-vs-split + * routing around this — every indexed route runs the same nearest scan. The indexed rewrite must + * therefore DECLINE such a table and fall back to the engine's default nearest-by execution. + * + * This is the shared eligibility contract the Catalyst rule consults before intercepting; the + * probe enforces it defensively too (see the schema guard in [[LanceProbe.probe]] / + * [[LanceProbe.probeRows]]). */ - def fusesCleanly(projection: Seq[String]): Boolean = - !projection.exists(ReservedProjectionColumns.contains) + def schemaSupportsNearest(schemaColumnNames: Iterable[String]): Boolean = + reservedSchemaColumns(schemaColumnNames).isEmpty /** * Build read options from a bare dataset URI, pinning `version` on the main branch when present. diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala index 4ad7ce066..bbc6eb1d7 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinRuleTest.scala @@ -24,7 +24,7 @@ import org.apache.spark.unsafe.types.UTF8String import org.junit.jupiter.api.{AfterEach, BeforeEach, Test} import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.io.TempDir -import org.lance.spark.knn.internal.Metric +import org.lance.spark.knn.internal.{LanceProbe, Metric} import org.lance.spark.utils.BlobUtils import java.nio.file.Path @@ -289,6 +289,45 @@ class IndexedNearestByJoinRuleTest { "blob v2 descriptor column must force fallback to Spark's canonical (blob-aware) reader") } + /** + * A right-side schema owning a column whose name collides with the metadata a nearest scan injects + * (`_rowid` / `_distance` / `_score`) forces the rule to DECLINE, even though the relation ALSO + * carries a searchable fixed-size vector. Every indexed route runs a `nearest` scan that injects + * those columns, so the injected metadata shadows the physical column — an all-columns fused scan + * reads the user's `_distance` out-of-band as the ranking score and silently drops it from the + * payload. No fold-vs-split routing recovers it (both scans inject `_rowid`), so the eligibility is + * schema-level: decline and hand the query to Spark's brute-force cross-product, whose canonical + * scan returns the true payload including that column. `LanceProbe` enforces the same contract + * defensively at probe time (see `LanceProbeValidationTest`). + * + * The positive control (same schema MINUS the reserved column) rewrites, proving the fixed-size- + * vector gate is satisfied and the reserved column is the sole discriminating cause of the decline. + * Exercised for each reserved name. + */ + @Test def reservedColumnDeclinesRewrite(): Unit = { + spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") + val left = trivialPlan("lid", "lvec") + val baseFields = Array( + StructField("rid", IntegerType, nullable = false), + fixedSizeVectorField("rvec", 8)) + // Control: a searchable vector-only schema WITHOUT any reserved column must rewrite. + val control = lanceRelationWithSchema(new StructType(baseFields), "reserved_control") + assertTrue( + IndexedNearestByJoinRule(l2Join(left, control)).isInstanceOf[Project], + "control: vector-only schema must rewrite (proves the vector gate passes)") + // Same schema + a column named like injected search metadata → decline, one name at a time. + LanceProbe.ReservedProjectionColumns.foreach { reserved => + val withReserved = lanceRelationWithSchema( + new StructType(baseFields :+ StructField(reserved, FloatType, nullable = false)), + s"reserved_${reserved.stripPrefix("_")}") + val join = l2Join(left, withReserved) + assertSame( + join, + IndexedNearestByJoinRule(join), + s"reserved column '$reserved' must force fallback to Spark's brute-force nearest-by") + } + } + /** Right side wrapped in SubqueryAlias still rewrites — alias unwrapping happens in the rule. */ @Test def testSubqueryAliasOnRightStillRewrites(): Unit = { spark.conf.set(IndexedNearestByJoinRule.EnabledConfKey, "true") diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala index 613d3510e..c28088ca5 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/catalyst/IndexedNearestByJoinSqlTest.scala @@ -329,6 +329,58 @@ class IndexedNearestByJoinSqlTest { } } + /** + * Real-Lance end-to-end: a right table whose schema OWNS a `_distance` column (a name the nearest + * scan injects) must DECLINE the indexed rewrite and fall through to Spark's brute-force nearest-by, + * whose ordinary scan returns the true stored `_distance` value. This is the gatekeeper's headline + * finding executed end-to-end: the indexed path would read that physical column out-of-band as the + * ranking score and silently drop it, so the only correct behavior is to decline and let Spark's + * canonical scan own the payload. + * + * Asserts: (a) the optimized plan has NO `LanceKnnJoinLogicalPlan` — the rule declined; (b) the + * query returns `k` hits per left row; (c) every joined row's `_distance` equals the value an + * ordinary Spark Lance read produces for that `rid` (proving the stored column is preserved, not + * clobbered by the search score). + */ + @Test def sqlUserDistanceColumnDeclinesAndReturnsStoredPayload(): Unit = { + val leftVecs = generateUniform(ExactLeft, ExactDim, Seed + 700) + val (leftDf, _, _) = buildLeftDf(leftVecs, ExactDim) + val (_, _, rightUri) = writeRightWithDistanceColumn(ExactRight, ExactDim, Seed + 701) + + // Oracle: ordinary Spark Lance read of the physical `_distance` column, per rid. + val oracle: Map[Int, Float] = + spark.read.format("lance").load(rightUri).select("rid", "_distance").collect().map { r => + r.getAs[Int]("rid") -> r.getAs[Float]("_distance") + }.toMap + + val k = 5 + val (q, d) = registerViews(leftDf, rightUri) + val df = spark.sql( + s"""SELECT q.lid, d.rid, d.`_distance` + |FROM $q q INNER JOIN $d d + |APPROX NEAREST $k BY DISTANCE vector_l2_distance(q.lvec, d.rvec)""".stripMargin) + + // (a) The rule must decline — a `_distance` column is present, so no indexed rewrite. + assertTrue( + df.queryExecution.optimizedPlan.collect { case p: LanceKnnJoinLogicalPlan => p }.isEmpty, + s"a user `_distance` column must force fallback (no LanceKnnJoinLogicalPlan); optimized plan:" + + s"\n${df.queryExecution.optimizedPlan}") + + // (b) + (c) The fallback runs and returns the true stored `_distance`, not the ranking score. + val rows = df.collect() + assertEquals(ExactLeft * k, rows.length, "expected k results per left row") + rows.foreach { r => + val rid = r.getAs[Int]("rid") + val actual = r.getAs[Float]("_distance") + assertEquals( + oracle(rid), + actual, + 1e-6f, + s"stored `_distance` payload must survive the fallback for rid=$rid " + + "(indexed path would have clobbered it with the ranking score)") + } + } + /** * The rewrite runs as a `postHocResolutionRule`, BEFORE Spark's `FinishAnalysis` / * `CheckCartesianProducts`. It must NOT let a query that Spark would reject slip through: with @@ -625,6 +677,31 @@ class IndexedNearestByJoinSqlTest { (vectors, ids, out) } + /** + * Write a right Lance dataset whose schema OWNS a `_distance` column — a name Lance's nearest scan + * injects. Each row's `_distance` is a deterministic payload value (NOT a ranking score). Used to + * prove the indexed rule DECLINES such a table and Spark's fallback returns the true stored value. + */ + private def writeRightWithDistanceColumn( + n: Int, + dim: Int, + seed: Long): (Array[Array[Float]], Array[Int], String) = { + val vectors = generateUniform(n, dim, seed) + val ids = vectors.indices.map(_ + 4000).toArray + val rows = ids.zip(vectors).map { case (id, v) => + // Stored payload distinct from any plausible ranking score, so a clobber would be visible. + RowFactory.create(Integer.valueOf(id), v, java.lang.Float.valueOf(id.toFloat * 0.25f)) + } + val schema = new StructType(Array( + StructField("rid", IntegerType, nullable = false), + fixedSizeVec("rvec", dim), + StructField("_distance", FloatType, nullable = false))) + val df = spark.createDataFrame(rows.toSeq.asJava, schema) + val out = tempDir.resolve(s"right_reserved_${System.nanoTime()}").toString + df.write.format("lance").save(out) + (vectors, ids, out) + } + /** Run an indexed nearest join (SQL path) against the given right dataset and compute recall@K. */ private def recallAgainst( leftDf: DataFrame, diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala index 8aae70968..7bf157bd3 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceKnnJoinStageTest.scala @@ -167,37 +167,26 @@ class LanceKnnJoinStageTest { assertEquals(base.getRef, merged.getRef, "same-named branch must keep the base's pinned ref") } - // -- fold vs split routing: reserved-name projections must decline the fold ---------------- + // -- fold vs split routing: purely an over-fetch decision ---------------------------------- + // Reserved-name collisions (a right schema owning `_rowid` / `_distance` / `_score`) are NOT this + // helper's concern: such a table is declined upstream by the Catalyst rule and by LanceProbe's + // schema backstop (see LanceProbe.schemaSupportsNearest / LanceProbeValidationTest), so by the + // time routing runs the projection is always fusible. /** - * The folded one-scan path is sound only when there is no over-fetch (internalK == k) AND the - * payload projection does not collide with a column the nearest scan injects. A plain projection - * at internalK == k folds. + * No over-fetch (internalK == k) folds probe + materialize into one scan: every probed row is + * kept, so a split path would just re-scan the exact rows the search already found. */ - @Test def testFoldsInOneScanForCleanProjectionAtNoOverfetch(): Unit = { + @Test def testFoldsInOneScanAtNoOverfetch(): Unit = { assertTrue( - LanceKnnJoinStage.foldsInOneScan(internalK = 5, k = 5, projection = Seq("id", "vec")), - "clean projection with internalK == k should fold in one scan") + LanceKnnJoinStage.foldsInOneScan(internalK = 5, k = 5), + "internalK == k should fold in one scan") } - /** - * A projection naming ANY reserved column (`_rowid`, `_distance`, `_score`) must decline the fold - * even at internalK == k, so the stage routes it to the split probe → materialize path whose scan - * injects only `_rowid`. This is the gatekeeper's headline finding: the fold must preserve Spark's - * fallback for a payload column whose name collides with the injected search metadata. - */ - @Test def testDoesNotFoldForReservedProjection(): Unit = { - LanceProbe.ReservedProjectionColumns.foreach { reserved => - assertFalse( - LanceKnnJoinStage.foldsInOneScan(internalK = 5, k = 5, projection = Seq("id", reserved)), - s"projection naming reserved column '$reserved' must NOT fold; route to the split path") - } - } - - /** Over-fetch (internalK > k) always takes the split path, regardless of projection. */ + /** Over-fetch (internalK > k) takes the split probe → trim → materialize path. */ @Test def testDoesNotFoldWhenOverfetching(): Unit = { assertFalse( - LanceKnnJoinStage.foldsInOneScan(internalK = 20, k = 5, projection = Seq("id")), + LanceKnnJoinStage.foldsInOneScan(internalK = 20, k = 5), "over-fetch (internalK > k) must take the split probe → trim → materialize path") } diff --git a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala index 575caac04..05bfe4580 100644 --- a/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala +++ b/lance-spark-knn-4.2_2.13/src/test/scala/org/lance/spark/knn/internal/LanceProbeValidationTest.scala @@ -187,39 +187,84 @@ class LanceProbeValidationTest { } /** - * The fused [[LanceProbe.probeRows]] scan injects `_rowid` and the score column itself, so a - * payload projection naming one of those reserved columns would collide inside the single scan. - * `probeRows` must reject such a projection (so the join stage can route it through the split - * probe + materialize path) rather than surfacing an opaque Arrow "merge incompatible fields" - * error, and [[LanceProbe.fusesCleanly]] must report it as not fusible. + * The eligibility contract is SCHEMA-level, not projection-level. Lance's nearest scan always + * injects `_rowid` and the `_distance` / `_score` metadata; if the dataset's OWN schema has a + * column by one of those names, the injected metadata shadows it and no probe route recovers the + * physical column. Empirically, a physical `_distance` column writes fine and probes silently drop + * it: an all-columns fused scan reads it out-of-band as the ranking score, so its payload comes + * back without `_distance`. That silent data loss is why the indexed rewrite must DECLINE such a + * table (see [[LanceProbe.schemaSupportsNearest]]) rather than pick a materialization shape. + * + * This regression writes a real dataset WITH a `_distance` column and asserts (a) the schema is + * reported non-nearest-compatible, and (b) BOTH probe entry points — `probe` and the all-columns + * `probeRows(projection = empty)` — fail fast with a clear error naming the offending column, + * instead of returning results that silently omit it. */ - @Test def testProbeRowsRejectsReservedProjection(): Unit = { - val datasetUri = writeSyntheticDataset() + @Test def testReservedSchemaColumnIsDeclinedNotSilentlyDropped(): Unit = { + val datasetUri = writeDatasetWithReservedColumn() val query = randomVector(new Random(1L), VectorDim) + // (a) schema eligibility primitive reports the collision. + assertFalse( + LanceProbe.schemaSupportsNearest(Seq("id", "vec", "_distance")), + "a schema owning a reserved column (_distance) must not be nearest-compatible") + val probe = new LanceProbe(datasetUri, fragmentIds = None) try { - LanceProbe.ReservedProjectionColumns.foreach { reserved => - assertFalse( - LanceProbe.fusesCleanly(Seq("id", reserved)), - s"projection naming reserved column '$reserved' must not be fusible") - val ex = assertThrows( - classOf[IllegalArgumentException], - () => - probe.probeRows( - "vec", - query, - k = 5, - Metric.L2, - projection = Seq("id", reserved), - projectionFields = Seq.empty)) - assertTrue( - String.valueOf(ex.getMessage).contains(reserved), - s"guard message should name reserved column '$reserved'; got: ${ex.getMessage}") - } + // (b) probe() declines with a clear error naming the offending column. + val probeEx = assertThrows( + classOf[IllegalArgumentException], + () => probe.probe("vec", query, k = 5, Metric.L2)) + assertTrue( + String.valueOf(probeEx.getMessage).contains("_distance"), + s"probe guard message should name the offending column '_distance'; got: ${probeEx.getMessage}") + + // (b) all-columns probeRows — the exact path that used to SILENTLY drop the physical + // `_distance` — declines with the same clear error rather than returning a lossy payload. + val probeRowsEx = assertThrows( + classOf[IllegalArgumentException], + () => + probe.probeRows( + "vec", + query, + k = 5, + Metric.L2, + projection = Seq.empty, + projectionFields = Seq.empty)) + assertTrue( + String.valueOf(probeRowsEx.getMessage).contains("_distance"), + s"probeRows guard message should name the offending column '_distance'; " + + s"got: ${probeRowsEx.getMessage}") } finally probe.close() } + /** + * Pure eligibility primitive: [[LanceProbe.schemaSupportsNearest]] is false iff the schema names + * any column the nearest scan injects, and [[LanceProbe.reservedSchemaColumns]] returns exactly + * that colliding set. No backend needed — this is the contract the Catalyst rule consults. + */ + @Test def testSchemaSupportsNearestContract(): Unit = { + assertTrue( + LanceProbe.schemaSupportsNearest(Seq("id", "vec", "payload")), + "a schema with no reserved column names must be nearest-compatible") + assertTrue( + LanceProbe.schemaSupportsNearest(Seq.empty), + "an empty schema must be nearest-compatible") + LanceProbe.ReservedProjectionColumns.foreach { reserved => + assertFalse( + LanceProbe.schemaSupportsNearest(Seq("id", reserved, "vec")), + s"a schema owning reserved column '$reserved' must not be nearest-compatible") + assertEquals( + Set(reserved), + LanceProbe.reservedSchemaColumns(Seq("id", reserved, "vec")), + s"reservedSchemaColumns must report exactly the colliding column '$reserved'") + } + assertEquals( + LanceProbe.ReservedProjectionColumns, + LanceProbe.reservedSchemaColumns(Seq("id") ++ LanceProbe.ReservedProjectionColumns.toSeq), + "reservedSchemaColumns must report every reserved column present") + } + /** * The folded fast path ([[LanceProbe.probeRows]]) must be observationally identical to the split * probe + materialize path: the SAME (rowAddr, score) hits in the SAME order, and the SAME @@ -354,6 +399,31 @@ class LanceProbeValidationTest { outDir } + /** + * Write a dataset whose OWN schema carries a `_distance` column — a name Lance's nearest scan + * injects. Used to prove the schema eligibility backstop declines such a table. The `_distance` + * values are arbitrary payload; the point is that the physical column exists in the stored schema. + */ + private def writeDatasetWithReservedColumn(): String = { + val rng = new Random(Seed) + val (baseRows, _) = generateRows(rng, NumRows, VectorDim) + val rows = baseRows.zipWithIndex.map { case (r, idx) => + RowFactory.create(r.get(0), r.get(1), java.lang.Float.valueOf(idx.toFloat)) + } + val schema = new StructType(Array( + StructField("id", IntegerType, nullable = false), + StructField( + "vec", + ArrayType(FloatType, containsNull = false), + nullable = false, + new MetadataBuilder().putLong("arrow.fixed-size-list.size", VectorDim.toLong).build()), + StructField("_distance", FloatType, nullable = false))) + val df = spark.createDataFrame(rows.asJava, schema) + val outDir = tempDir.resolve(s"reserved_col_test_${System.nanoTime()}").toString + df.write.format("lance").save(outDir) + outDir + } + private def generateRows(rng: Random, n: Int, dim: Int): (Seq[Row], Seq[Array[Float]]) = { val vectors = (0 until n).map(_ => randomVector(rng, dim)) val rows = vectors.zipWithIndex.map { case (v, idx) =>