From 768b3e90f261c7aea58bdb98dc698b90deeeae34 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 14 Dec 2025 16:24:01 +0400 Subject: [PATCH 1/4] impl map_from_entries --- native/core/src/execution/jni_api.rs | 2 + .../apache/comet/serde/QueryPlanSerde.scala | 3 +- .../scala/org/apache/comet/serde/maps.scala | 29 +++++++++++- .../comet/CometMapExpressionSuite.scala | 45 +++++++++++++++++++ 4 files changed, 77 insertions(+), 2 deletions(-) diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index a24d993059..4f53cea3e6 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -46,6 +46,7 @@ use datafusion_spark::function::datetime::date_add::SparkDateAdd; use datafusion_spark::function::datetime::date_sub::SparkDateSub; use datafusion_spark::function::hash::sha1::SparkSha1; use datafusion_spark::function::hash::sha2::SparkSha2; +use datafusion_spark::function::map::map_from_entries::MapFromEntries; use datafusion_spark::function::math::expm1::SparkExpm1; use datafusion_spark::function::string::char::CharFunc; use datafusion_spark::function::string::concat::SparkConcat; @@ -337,6 +338,7 @@ fn register_datafusion_spark_function(session_ctx: &SessionContext) { session_ctx.register_udf(ScalarUDF::new_from_impl(SparkSha1::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkConcat::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkBitwiseNot::default())); + session_ctx.register_udf(ScalarUDF::new_from_impl(MapFromEntries::default())); } /// Prepares arrow arrays for output. diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index 54df2f1688..a99cf3824b 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -125,7 +125,8 @@ object QueryPlanSerde extends Logging with CometExprShim { classOf[MapKeys] -> CometMapKeys, classOf[MapEntries] -> CometMapEntries, classOf[MapValues] -> CometMapValues, - classOf[MapFromArrays] -> CometMapFromArrays) + classOf[MapFromArrays] -> CometMapFromArrays, + classOf[MapFromEntries] -> CometMapFromEntries) private val structExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( classOf[CreateNamedStruct] -> CometCreateNamedStruct, diff --git a/spark/src/main/scala/org/apache/comet/serde/maps.scala b/spark/src/main/scala/org/apache/comet/serde/maps.scala index 2e217f6af0..498aa3594c 100644 --- a/spark/src/main/scala/org/apache/comet/serde/maps.scala +++ b/spark/src/main/scala/org/apache/comet/serde/maps.scala @@ -19,9 +19,12 @@ package org.apache.comet.serde +import scala.annotation.tailrec + import org.apache.spark.sql.catalyst.expressions._ -import org.apache.spark.sql.types.{ArrayType, MapType} +import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, MapType, StructType} +import org.apache.comet.serde.CometArrayReverse.containsBinary import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithInfo, scalarFunctionExprToProto, scalarFunctionExprToProtoWithReturnType} object CometMapKeys extends CometExpressionSerde[MapKeys] { @@ -89,3 +92,27 @@ object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { optExprWithInfo(mapFromArraysExpr, expr, expr.children: _*) } } + +object CometMapFromEntries extends CometScalarFunction[MapFromEntries]("map_from_entries") { + val keyUnsupportedReason = "Using BinaryType as Map keys is not allowed in map_from_entries" + val valueUnsupportedReason = "Using BinaryType as Map values is not allowed in map_from_entries" + + private def containsBinary(dataType: DataType): Boolean = { + dataType match { + case BinaryType => true + case StructType(fields) => fields.exists(field => containsBinary(field.dataType)) + case ArrayType(elementType, _) => containsBinary(elementType) + case _ => false + } + } + + override def getSupportLevel(expr: MapFromEntries): SupportLevel = { + if (containsBinary(expr.dataType.keyType)) { + return Incompatible(Some(keyUnsupportedReason)) + } + if (containsBinary(expr.dataType.valueType)) { + return Incompatible(Some(valueUnsupportedReason)) + } + Compatible(None) + } +} diff --git a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala index 88c13391a6..01b9744ed6 100644 --- a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala @@ -25,7 +25,9 @@ import org.apache.hadoop.fs.Path import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.BinaryType +import org.apache.comet.serde.CometMapFromEntries import org.apache.comet.testing.{DataGenOptions, ParquetGenerator, SchemaGenOptions} class CometMapExpressionSuite extends CometTestBase { @@ -125,4 +127,47 @@ class CometMapExpressionSuite extends CometTestBase { } } + test("map_from_entries") { + withTempDir { dir => + val path = new Path(dir.toURI.toString, "test.parquet") + val filename = path.toString + val random = new Random(42) + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + val schemaGenOptions = + SchemaGenOptions( + generateArray = true, + generateStruct = true, + primitiveTypes = SchemaGenOptions.defaultPrimitiveTypes.filterNot(_ == BinaryType)) + val dataGenOptions = DataGenOptions(allowNull = false, generateNegativeZero = false) + ParquetGenerator.makeParquetFile( + random, + spark, + filename, + 100, + schemaGenOptions, + dataGenOptions) + } + val df = spark.read.parquet(filename) + df.createOrReplaceTempView("t1") + for (field <- df.schema.fieldNames) { + checkSparkAnswerAndOperator( + spark.sql(s"SELECT map_from_entries(array(struct($field as a, $field as b))) FROM t1")) + } + } + } + + test("map_from_entries - fallback for binary type") { + val table = "t2" + withTable(table) { + sql( + s"create table $table using parquet as select cast(array() as array) as c1 from range(10)") + checkSparkAnswerAndFallbackReason( + sql(s"select map_from_entries(array(struct(c1, 0))) from $table"), + CometMapFromEntries.keyUnsupportedReason) + checkSparkAnswerAndFallbackReason( + sql(s"select map_from_entries(array(struct(0, c1))) from $table"), + CometMapFromEntries.valueUnsupportedReason) + } + } + } From c68c3428676b5d991e7ba9e13464bf2ce1ec84e8 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Tue, 16 Dec 2025 16:10:43 +0400 Subject: [PATCH 2/4] Revert "impl map_from_entries" This reverts commit 768b3e90f261c7aea58bdb98dc698b90deeeae34. --- native/core/src/execution/jni_api.rs | 2 - .../apache/comet/serde/QueryPlanSerde.scala | 3 +- .../scala/org/apache/comet/serde/maps.scala | 29 +----------- .../comet/CometMapExpressionSuite.scala | 45 ------------------- 4 files changed, 2 insertions(+), 77 deletions(-) diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index 4f53cea3e6..a24d993059 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -46,7 +46,6 @@ use datafusion_spark::function::datetime::date_add::SparkDateAdd; use datafusion_spark::function::datetime::date_sub::SparkDateSub; use datafusion_spark::function::hash::sha1::SparkSha1; use datafusion_spark::function::hash::sha2::SparkSha2; -use datafusion_spark::function::map::map_from_entries::MapFromEntries; use datafusion_spark::function::math::expm1::SparkExpm1; use datafusion_spark::function::string::char::CharFunc; use datafusion_spark::function::string::concat::SparkConcat; @@ -338,7 +337,6 @@ fn register_datafusion_spark_function(session_ctx: &SessionContext) { session_ctx.register_udf(ScalarUDF::new_from_impl(SparkSha1::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkConcat::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkBitwiseNot::default())); - session_ctx.register_udf(ScalarUDF::new_from_impl(MapFromEntries::default())); } /// Prepares arrow arrays for output. diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index a99cf3824b..54df2f1688 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -125,8 +125,7 @@ object QueryPlanSerde extends Logging with CometExprShim { classOf[MapKeys] -> CometMapKeys, classOf[MapEntries] -> CometMapEntries, classOf[MapValues] -> CometMapValues, - classOf[MapFromArrays] -> CometMapFromArrays, - classOf[MapFromEntries] -> CometMapFromEntries) + classOf[MapFromArrays] -> CometMapFromArrays) private val structExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( classOf[CreateNamedStruct] -> CometCreateNamedStruct, diff --git a/spark/src/main/scala/org/apache/comet/serde/maps.scala b/spark/src/main/scala/org/apache/comet/serde/maps.scala index 498aa3594c..2e217f6af0 100644 --- a/spark/src/main/scala/org/apache/comet/serde/maps.scala +++ b/spark/src/main/scala/org/apache/comet/serde/maps.scala @@ -19,12 +19,9 @@ package org.apache.comet.serde -import scala.annotation.tailrec - import org.apache.spark.sql.catalyst.expressions._ -import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, MapType, StructType} +import org.apache.spark.sql.types.{ArrayType, MapType} -import org.apache.comet.serde.CometArrayReverse.containsBinary import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithInfo, scalarFunctionExprToProto, scalarFunctionExprToProtoWithReturnType} object CometMapKeys extends CometExpressionSerde[MapKeys] { @@ -92,27 +89,3 @@ object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { optExprWithInfo(mapFromArraysExpr, expr, expr.children: _*) } } - -object CometMapFromEntries extends CometScalarFunction[MapFromEntries]("map_from_entries") { - val keyUnsupportedReason = "Using BinaryType as Map keys is not allowed in map_from_entries" - val valueUnsupportedReason = "Using BinaryType as Map values is not allowed in map_from_entries" - - private def containsBinary(dataType: DataType): Boolean = { - dataType match { - case BinaryType => true - case StructType(fields) => fields.exists(field => containsBinary(field.dataType)) - case ArrayType(elementType, _) => containsBinary(elementType) - case _ => false - } - } - - override def getSupportLevel(expr: MapFromEntries): SupportLevel = { - if (containsBinary(expr.dataType.keyType)) { - return Incompatible(Some(keyUnsupportedReason)) - } - if (containsBinary(expr.dataType.valueType)) { - return Incompatible(Some(valueUnsupportedReason)) - } - Compatible(None) - } -} diff --git a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala index 01b9744ed6..88c13391a6 100644 --- a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala @@ -25,9 +25,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.BinaryType -import org.apache.comet.serde.CometMapFromEntries import org.apache.comet.testing.{DataGenOptions, ParquetGenerator, SchemaGenOptions} class CometMapExpressionSuite extends CometTestBase { @@ -127,47 +125,4 @@ class CometMapExpressionSuite extends CometTestBase { } } - test("map_from_entries") { - withTempDir { dir => - val path = new Path(dir.toURI.toString, "test.parquet") - val filename = path.toString - val random = new Random(42) - withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - val schemaGenOptions = - SchemaGenOptions( - generateArray = true, - generateStruct = true, - primitiveTypes = SchemaGenOptions.defaultPrimitiveTypes.filterNot(_ == BinaryType)) - val dataGenOptions = DataGenOptions(allowNull = false, generateNegativeZero = false) - ParquetGenerator.makeParquetFile( - random, - spark, - filename, - 100, - schemaGenOptions, - dataGenOptions) - } - val df = spark.read.parquet(filename) - df.createOrReplaceTempView("t1") - for (field <- df.schema.fieldNames) { - checkSparkAnswerAndOperator( - spark.sql(s"SELECT map_from_entries(array(struct($field as a, $field as b))) FROM t1")) - } - } - } - - test("map_from_entries - fallback for binary type") { - val table = "t2" - withTable(table) { - sql( - s"create table $table using parquet as select cast(array() as array) as c1 from range(10)") - checkSparkAnswerAndFallbackReason( - sql(s"select map_from_entries(array(struct(c1, 0))) from $table"), - CometMapFromEntries.keyUnsupportedReason) - checkSparkAnswerAndFallbackReason( - sql(s"select map_from_entries(array(struct(0, c1))) from $table"), - CometMapFromEntries.valueUnsupportedReason) - } - } - } From 2f7c3087ccfdb65ab9c105185f18e9a32c2e160b Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 26 Jul 2026 17:53:33 +0400 Subject: [PATCH 3/4] perf --- native/spark-expr/Cargo.toml | 4 + native/spark-expr/benches/levenshtein.rs | 73 +++++ .../src/string_funcs/levenshtein.rs | 271 ++++++++++++------ 3 files changed, 267 insertions(+), 81 deletions(-) create mode 100644 native/spark-expr/benches/levenshtein.rs diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index c05ae89793..93c3ea4433 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -198,3 +198,7 @@ harness = false [[bench]] name = "unscaled_value" harness = false + +[[bench]] +name = "levenshtein" +harness = false diff --git a/native/spark-expr/benches/levenshtein.rs b/native/spark-expr/benches/levenshtein.rs new file mode 100644 index 0000000000..ff15d83112 --- /dev/null +++ b/native/spark-expr/benches/levenshtein.rs @@ -0,0 +1,73 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use arrow::array::{ArrayRef, Int32Array, StringArray}; +use criterion::{criterion_group, criterion_main, Criterion}; +use datafusion::physical_plan::ColumnarValue; +use datafusion_comet_spark_expr::spark_levenshtein; +use std::hint::black_box; +use std::sync::Arc; + +fn create_string_arrays(rows: usize) -> (ArrayRef, ArrayRef) { + let left_strings: Vec = (0..rows) + .map(|i| format!("apache_datafusion_comet_{}", i % 100)) + .collect(); + let right_strings: Vec = (0..rows) + .map(|i| format!("apache_comet_expr_{}", (i + 5) % 100)) + .collect(); + + let left_array = StringArray::from( + left_strings + .iter() + .map(|s| s.as_str()) + .collect::>(), + ); + let right_array = StringArray::from( + right_strings + .iter() + .map(|s| s.as_str()) + .collect::>(), + ); + + (Arc::new(left_array) as ArrayRef, Arc::new(right_array) as ArrayRef) +} + +fn criterion_benchmark(c: &mut Criterion) { + let rows = 8192; + let (left, right) = create_string_arrays(rows); + + c.bench_function("spark_levenshtein: 2 arguments (no threshold)", |b| { + let args = vec![ + ColumnarValue::Array(Arc::clone(&left)), + ColumnarValue::Array(Arc::clone(&right)), + ]; + b.iter(|| black_box(spark_levenshtein(black_box(&args)).unwrap())) + }); + + let threshold = Int32Array::from(vec![10; rows]); + c.bench_function("spark_levenshtein: 3 arguments (with threshold)", |b| { + let args = vec![ + ColumnarValue::Array(Arc::clone(&left)), + ColumnarValue::Array(Arc::clone(&right)), + ColumnarValue::Array(Arc::new(threshold.clone()) as ArrayRef), + ]; + b.iter(|| black_box(spark_levenshtein(black_box(&args)).unwrap())) + }); +} + +criterion_group!(benches, criterion_benchmark); +criterion_main!(benches); diff --git a/native/spark-expr/src/string_funcs/levenshtein.rs b/native/spark-expr/src/string_funcs/levenshtein.rs index 19a81f92ff..a2eb45df47 100644 --- a/native/spark-expr/src/string_funcs/levenshtein.rs +++ b/native/spark-expr/src/string_funcs/levenshtein.rs @@ -26,16 +26,66 @@ use datafusion::common::{cast::as_generic_string_array, DataFusionError, Result} use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; +// Thread-local scratch buffers to avoid heap allocations in the row processing loop +thread_local! { + static LEVENSHTEIN_SCRATCH: std::cell::RefCell<(Vec, Vec)> = + std::cell::RefCell::new((Vec::with_capacity(64), Vec::with_capacity(64))); +} + /// Computes the Levenshtein edit distance between two UTF-8 strings. -/// -/// This uses the standard dynamic programming algorithm with O(min(m,n)) space. fn levenshtein_distance(s: &str, t: &str) -> i32 { + // Fast path for ASCII strings: operate directly on raw bytes without vector allocations + if s.is_ascii() && t.is_ascii() { + let s_bytes = s.as_bytes(); + let t_bytes = t.as_bytes(); + let m = s_bytes.len(); + let n = t_bytes.len(); + + if m == 0 { + return n as i32; + } + if n == 0 { + return m as i32; + } + + let (s_bytes, t_bytes, m, n) = if m > n { + (t_bytes, s_bytes, n, m) + } else { + (s_bytes, t_bytes, m, n) + }; + + return LEVENSHTEIN_SCRATCH.with(|scratch| { + let mut borrow = scratch.borrow_mut(); + let (prev, curr) = &mut *borrow; + + prev.resize(m + 1, 0); + curr.resize(m + 1, 0); + + for (i, val) in prev.iter_mut().enumerate() { + *val = i as i32; + } + + for j in 1..=n { + curr[0] = j as i32; + for i in 1..=m { + let cost = if s_bytes[i - 1] == t_bytes[j - 1] { 0 } else { 1 }; + curr[i] = (prev[i] + 1) + .min(curr[i - 1] + 1) + .min(prev[i - 1] + cost); + } + std::mem::swap(prev, curr); + } + + prev[m] + }); + } + + // General Unicode path for non-ASCII strings let s_chars: Vec = s.chars().collect(); let t_chars: Vec = t.chars().collect(); let m = s_chars.len(); let n = t_chars.len(); - // Optimization: if one string is empty, distance is the length of the other if m == 0 { return n as i32; } @@ -43,50 +93,109 @@ fn levenshtein_distance(s: &str, t: &str) -> i32 { return m as i32; } - // Use the shorter string for the "column" to minimize space usage let (s_chars, t_chars, m, n) = if m > n { (t_chars, s_chars, n, m) } else { (s_chars, t_chars, m, n) }; - // Previous and current row of distances - let mut prev = vec![0i32; m + 1]; - let mut curr = vec![0i32; m + 1]; + LEVENSHTEIN_SCRATCH.with(|scratch| { + let mut borrow = scratch.borrow_mut(); + let (prev, curr) = &mut *borrow; - // Initialize base case: distance from empty string - for (i, val) in prev.iter_mut().enumerate() { - *val = i as i32; - } + prev.resize(m + 1, 0); + curr.resize(m + 1, 0); - for j in 1..=n { - curr[0] = j as i32; - for i in 1..=m { - let cost = if s_chars[i - 1] == t_chars[j - 1] { - 0 - } else { - 1 - }; - curr[i] = (prev[i] + 1) // deletion - .min(curr[i - 1] + 1) // insertion - .min(prev[i - 1] + cost); // substitution + for (i, val) in prev.iter_mut().enumerate() { + *val = i as i32; + } + + for j in 1..=n { + curr[0] = j as i32; + for i in 1..=m { + let cost = if s_chars[i - 1] == t_chars[j - 1] { 0 } else { 1 }; + curr[i] = (prev[i] + 1) + .min(curr[i - 1] + 1) + .min(prev[i - 1] + cost); + } + std::mem::swap(prev, curr); } - std::mem::swap(&mut prev, &mut curr); - } - prev[m] + prev[m] + }) } /// Computes the Levenshtein distance up to `threshold` using a diagonal band. -/// -/// Spark's three-argument form uses the threshold to avoid evaluating cells that cannot -/// contribute to a result within the requested distance. This keeps the complexity at -/// O(threshold * max(m, n)) when the threshold is small rather than always using O(m * n). fn levenshtein_distance_with_threshold(s: &str, t: &str, threshold: i32) -> i32 { if threshold < 0 { return -1; } + // Fast path for ASCII strings + if s.is_ascii() && t.is_ascii() { + let s_bytes = s.as_bytes(); + let t_bytes = t.as_bytes(); + let (shorter, longer) = if s_bytes.len() <= t_bytes.len() { + (s_bytes, t_bytes) + } else { + (t_bytes, s_bytes) + }; + let m = shorter.len(); + let n = longer.len(); + let threshold = threshold as usize; + + if n - m > threshold { + return -1; + } + if m == 0 { + return if n <= threshold { n as i32 } else { -1 }; + } + + let out_of_band = n.saturating_add(1) as i32; + + return LEVENSHTEIN_SCRATCH.with(|scratch| { + let mut borrow = scratch.borrow_mut(); + let (prev, curr) = &mut *borrow; + + prev.resize(m + 1, out_of_band); + curr.resize(m + 1, out_of_band); + + for (i, value) in prev.iter_mut().enumerate().take(m.min(threshold) + 1) { + *value = i as i32; + } + + for j in 1..=n { + let start = 1.max(j.saturating_sub(threshold)); + let end = m.min(j.saturating_add(threshold)); + if start > end { + return -1; + } + + curr[0] = if j <= threshold { j as i32 } else { out_of_band }; + curr[start - 1] = if start == 1 { curr[0] } else { out_of_band }; + + for i in start..=end { + let cost = if shorter[i - 1] == longer[j - 1] { 0 } else { 1 }; + curr[i] = prev[i] + .saturating_add(1) + .min(curr[i - 1].saturating_add(1)) + .min(prev[i - 1].saturating_add(cost)); + } + if end < m { + curr[end + 1] = out_of_band; + } + std::mem::swap(prev, curr); + } + + if prev[m] <= threshold as i32 { + prev[m] + } else { + -1 + } + }); + } + + // General Unicode path let s_chars: Vec = s.chars().collect(); let t_chars: Vec = t.chars().collect(); let (shorter, longer) = if s_chars.len() <= t_chars.len() { @@ -105,40 +214,48 @@ fn levenshtein_distance_with_threshold(s: &str, t: &str, threshold: i32) -> i32 return if n <= threshold { n as i32 } else { -1 }; } - let out_of_band = n.saturating_add(1); - let mut prev = vec![out_of_band; m + 1]; - let mut curr = vec![out_of_band; m + 1]; - for (i, value) in prev.iter_mut().enumerate().take(m.min(threshold) + 1) { - *value = i; - } + let out_of_band = n.saturating_add(1) as i32; - for j in 1..=n { - let start = 1.max(j.saturating_sub(threshold)); - let end = m.min(j.saturating_add(threshold)); - if start > end { - return -1; - } + LEVENSHTEIN_SCRATCH.with(|scratch| { + let mut borrow = scratch.borrow_mut(); + let (prev, curr) = &mut *borrow; - curr[0] = if j <= threshold { j } else { out_of_band }; - curr[start - 1] = if start == 1 { curr[0] } else { out_of_band }; - for i in start..=end { - let cost = usize::from(shorter[i - 1] != longer[j - 1]); - curr[i] = prev[i] - .saturating_add(1) - .min(curr[i - 1].saturating_add(1)) - .min(prev[i - 1].saturating_add(cost)); + prev.resize(m + 1, out_of_band); + curr.resize(m + 1, out_of_band); + + for (i, value) in prev.iter_mut().enumerate().take(m.min(threshold) + 1) { + *value = i as i32; } - if end < m { - curr[end + 1] = out_of_band; + + for j in 1..=n { + let start = 1.max(j.saturating_sub(threshold)); + let end = m.min(j.saturating_add(threshold)); + if start > end { + return -1; + } + + curr[0] = if j <= threshold { j as i32 } else { out_of_band }; + curr[start - 1] = if start == 1 { curr[0] } else { out_of_band }; + + for i in start..=end { + let cost = if shorter[i - 1] == longer[j - 1] { 0 } else { 1 }; + curr[i] = prev[i] + .saturating_add(1) + .min(curr[i - 1].saturating_add(1)) + .min(prev[i - 1].saturating_add(cost)); + } + if end < m { + curr[end + 1] = out_of_band; + } + std::mem::swap(prev, curr); } - std::mem::swap(&mut prev, &mut curr); - } - if prev[m] <= threshold { - prev[m] as i32 - } else { - -1 - } + if prev[m] <= threshold as i32 { + prev[m] + } else { + -1 + } + }) } fn evaluate_levenshtein( @@ -179,40 +296,34 @@ fn evaluate_string_arrays( threshold: Option<&Int32Array>, ) -> Result { match (left.data_type(), right.data_type()) { - (DataType::Utf8, DataType::Utf8) => Ok(evaluate_levenshtein( - as_generic_string_array::(left.as_ref())?, - as_generic_string_array::(right.as_ref())?, + (DataType::Utf8, DataType::Utf8) => Ok(evaluate_levenshtein::( + as_generic_string_array::(left)?, + as_generic_string_array::(right)?, threshold, )), - (DataType::Utf8, DataType::LargeUtf8) => Ok(evaluate_levenshtein( - as_generic_string_array::(left.as_ref())?, - as_generic_string_array::(right.as_ref())?, + (DataType::Utf8, DataType::LargeUtf8) => Ok(evaluate_levenshtein::( + as_generic_string_array::(left)?, + as_generic_string_array::(right)?, threshold, )), - (DataType::LargeUtf8, DataType::Utf8) => Ok(evaluate_levenshtein( - as_generic_string_array::(left.as_ref())?, - as_generic_string_array::(right.as_ref())?, + (DataType::LargeUtf8, DataType::Utf8) => Ok(evaluate_levenshtein::( + as_generic_string_array::(left)?, + as_generic_string_array::(right)?, threshold, )), - (DataType::LargeUtf8, DataType::LargeUtf8) => Ok(evaluate_levenshtein( - as_generic_string_array::(left.as_ref())?, - as_generic_string_array::(right.as_ref())?, + (DataType::LargeUtf8, DataType::LargeUtf8) => Ok(evaluate_levenshtein::( + as_generic_string_array::(left)?, + as_generic_string_array::(right)?, threshold, )), - (left_type, right_type) => Err(DataFusionError::Execution(format!( - "levenshtein expects Utf8 or LargeUtf8 arguments, got {left_type:?} and {right_type:?}" + other => Err(DataFusionError::Internal(format!( + "levenshtein not supported for types {:?} and {:?}", + other.0, other.1 ))), } } /// Spark-compatible levenshtein scalar function. -/// -/// Accepts two or three arguments: -/// - `levenshtein(str1, str2)` → edit distance -/// - `levenshtein(str1, str2, threshold)` → edit distance if <= threshold, else -1 -/// -/// The threshold argument can be either a scalar or a column (array). -/// NULL inputs produce NULL outputs. NULL threshold produces NULL output for that row. pub fn spark_levenshtein(args: &[ColumnarValue]) -> Result { if args.len() < 2 || args.len() > 3 { return Err(DataFusionError::Internal(format!( @@ -221,7 +332,6 @@ pub fn spark_levenshtein(args: &[ColumnarValue]) -> Result { ))); } - // Determine array length from any array argument let len = args .iter() .find_map(|arg| match arg { @@ -238,7 +348,6 @@ pub fn spark_levenshtein(args: &[ColumnarValue]) -> Result { )); } - // Handle the optional threshold argument (scalar or array) let threshold_array = if args.len() == 3 { let threshold_array = args[2].clone().into_array(len)?; if threshold_array.len() != len { From 6d7985d255d0654cb70d426e248b5c1af8c9cf81 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 26 Jul 2026 19:09:10 +0400 Subject: [PATCH 4/4] perf: optimize levenshtein function --- native/spark-expr/benches/levenshtein.rs | 5 ++- .../src/string_funcs/levenshtein.rs | 44 ++++++++++++++----- 2 files changed, 36 insertions(+), 13 deletions(-) diff --git a/native/spark-expr/benches/levenshtein.rs b/native/spark-expr/benches/levenshtein.rs index ff15d83112..236e33453c 100644 --- a/native/spark-expr/benches/levenshtein.rs +++ b/native/spark-expr/benches/levenshtein.rs @@ -43,7 +43,10 @@ fn create_string_arrays(rows: usize) -> (ArrayRef, ArrayRef) { .collect::>(), ); - (Arc::new(left_array) as ArrayRef, Arc::new(right_array) as ArrayRef) + ( + Arc::new(left_array) as ArrayRef, + Arc::new(right_array) as ArrayRef, + ) } fn criterion_benchmark(c: &mut Criterion) { diff --git a/native/spark-expr/src/string_funcs/levenshtein.rs b/native/spark-expr/src/string_funcs/levenshtein.rs index a2eb45df47..7914a32961 100644 --- a/native/spark-expr/src/string_funcs/levenshtein.rs +++ b/native/spark-expr/src/string_funcs/levenshtein.rs @@ -68,10 +68,12 @@ fn levenshtein_distance(s: &str, t: &str) -> i32 { for j in 1..=n { curr[0] = j as i32; for i in 1..=m { - let cost = if s_bytes[i - 1] == t_bytes[j - 1] { 0 } else { 1 }; - curr[i] = (prev[i] + 1) - .min(curr[i - 1] + 1) - .min(prev[i - 1] + cost); + let cost = if s_bytes[i - 1] == t_bytes[j - 1] { + 0 + } else { + 1 + }; + curr[i] = (prev[i] + 1).min(curr[i - 1] + 1).min(prev[i - 1] + cost); } std::mem::swap(prev, curr); } @@ -113,10 +115,12 @@ fn levenshtein_distance(s: &str, t: &str) -> i32 { for j in 1..=n { curr[0] = j as i32; for i in 1..=m { - let cost = if s_chars[i - 1] == t_chars[j - 1] { 0 } else { 1 }; - curr[i] = (prev[i] + 1) - .min(curr[i - 1] + 1) - .min(prev[i - 1] + cost); + let cost = if s_chars[i - 1] == t_chars[j - 1] { + 0 + } else { + 1 + }; + curr[i] = (prev[i] + 1).min(curr[i - 1] + 1).min(prev[i - 1] + cost); } std::mem::swap(prev, curr); } @@ -171,11 +175,19 @@ fn levenshtein_distance_with_threshold(s: &str, t: &str, threshold: i32) -> i32 return -1; } - curr[0] = if j <= threshold { j as i32 } else { out_of_band }; + curr[0] = if j <= threshold { + j as i32 + } else { + out_of_band + }; curr[start - 1] = if start == 1 { curr[0] } else { out_of_band }; for i in start..=end { - let cost = if shorter[i - 1] == longer[j - 1] { 0 } else { 1 }; + let cost = if shorter[i - 1] == longer[j - 1] { + 0 + } else { + 1 + }; curr[i] = prev[i] .saturating_add(1) .min(curr[i - 1].saturating_add(1)) @@ -234,11 +246,19 @@ fn levenshtein_distance_with_threshold(s: &str, t: &str, threshold: i32) -> i32 return -1; } - curr[0] = if j <= threshold { j as i32 } else { out_of_band }; + curr[0] = if j <= threshold { + j as i32 + } else { + out_of_band + }; curr[start - 1] = if start == 1 { curr[0] } else { out_of_band }; for i in start..=end { - let cost = if shorter[i - 1] == longer[j - 1] { 0 } else { 1 }; + let cost = if shorter[i - 1] == longer[j - 1] { + 0 + } else { + 1 + }; curr[i] = prev[i] .saturating_add(1) .min(curr[i - 1].saturating_add(1))