From 768b3e90f261c7aea58bdb98dc698b90deeeae34 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 14 Dec 2025 16:24:01 +0400 Subject: [PATCH 01/12] 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 a24d9930597..4f53cea3e68 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 54df2f1688d..a99cf3824bf 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 2e217f6af0b..498aa3594cf 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 88c13391a67..01b9744ed6f 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 02/12] 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 4f53cea3e68..a24d9930597 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 a99cf3824bf..54df2f1688d 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 498aa3594cf..2e217f6af0b 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 01b9744ed6f..88c13391a67 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 57d469c8b5d2e8b6bf68cb30f98587c0d5492775 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 9 Aug 2026 21:20:24 +0400 Subject: [PATCH 03/12] work --- .../spark-expr/src/string_funcs/contains.rs | 72 +++++++++---------- 1 file changed, 35 insertions(+), 37 deletions(-) diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index 537227efdfc..398177dcd9a 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -89,9 +89,8 @@ fn spark_contains(haystack: &ColumnarValue, needle: &ColumnarValue) -> Result { - let haystack_array = haystack_scalar.to_array_of_size(needle_array.len())?; - let result = arrow_contains(&haystack_array, needle_array)?; - Ok(ColumnarValue::Array(Arc::new(result))) + let result = contains_scalar_with_arrow(haystack_scalar, needle_array)?; + Ok(ColumnarValue::Array(result)) } // Both scalars - compute single result @@ -102,6 +101,21 @@ fn spark_contains(haystack: &ColumnarValue, needle: &ColumnarValue) -> Result(scalar: &'a ScalarValue, arg_name: &str) -> Result<&'a str> { + match scalar { + ScalarValue::Utf8(Some(s)) + | ScalarValue::LargeUtf8(Some(s)) + | ScalarValue::Utf8View(Some(s)) => Ok(s.as_str()), + _ => exec_err!( + "contains function requires string type for {}, got {:?}", + arg_name, + scalar.data_type() + ), + } +} + /// Optimized contains for array haystack with scalar needle. /// Uses Arrow's native scalar handling for better performance. fn contains_with_arrow_scalar( @@ -114,17 +128,7 @@ fn contains_with_arrow_scalar( } // Extract the needle string - let needle_str = match needle_scalar { - ScalarValue::Utf8(Some(s)) - | ScalarValue::LargeUtf8(Some(s)) - | ScalarValue::Utf8View(Some(s)) => s.clone(), - _ => { - return exec_err!( - "contains function requires string type for needle, got {:?}", - needle_scalar.data_type() - ) - } - }; + let needle_str = get_string(needle_scalar, "needle")?; // Create scalar array for needle - tells Arrow to use optimized paths let needle_scalar_array = StringArray::new_scalar(needle_str); @@ -134,6 +138,21 @@ fn contains_with_arrow_scalar( Ok(Arc::new(result)) } +fn contains_scalar_with_arrow( + haystack_scalar: &ScalarValue, + needle_array: &ArrayRef, +) -> Result { + if haystack_scalar.is_null() { + return Ok(Arc::new(BooleanArray::new_null(needle_array.len()))); + } + + let haystack_str = get_string(haystack_scalar, "haystack")?; + let haystack_scalar_array = StringArray::new_scalar(haystack_str.to_string()); + + let result = arrow_contains(&haystack_scalar_array, needle_array)?; + Ok(Arc::new(result)) +} + /// Contains for two scalar values. fn contains_scalar_scalar( haystack_scalar: &ScalarValue, @@ -144,29 +163,8 @@ fn contains_scalar_scalar( return Ok(ScalarValue::Boolean(None)); } - let haystack_str = match haystack_scalar { - ScalarValue::Utf8(Some(s)) - | ScalarValue::LargeUtf8(Some(s)) - | ScalarValue::Utf8View(Some(s)) => s.as_str(), - _ => { - return exec_err!( - "contains function requires string type for haystack, got {:?}", - haystack_scalar.data_type() - ) - } - }; - - let needle_str = match needle_scalar { - ScalarValue::Utf8(Some(s)) - | ScalarValue::LargeUtf8(Some(s)) - | ScalarValue::Utf8View(Some(s)) => s.as_str(), - _ => { - return exec_err!( - "contains function requires string type for needle, got {:?}", - needle_scalar.data_type() - ) - } - }; + let haystack_str = get_string(haystack_scalar, "haystack")?; + let needle_str = get_string(needle_scalar, "needle")?; Ok(ScalarValue::Boolean(Some( haystack_str.contains(needle_str), From 3303a04cfd3282eb3e0e876dbe9b0d3ac16b15be Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 9 Aug 2026 22:11:03 +0400 Subject: [PATCH 04/12] work --- native/spark-expr/Cargo.toml | 4 ++ .../spark-expr/src/string_funcs/contains.rs | 61 +++++++++++++++---- 2 files changed, 53 insertions(+), 12 deletions(-) diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index 6faa9fec4ec..58eeeadf679 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -222,4 +222,8 @@ harness = false [[bench]] name = "cast_int_to_decimal" +harness = false + +[[bench]] +name = "contains" harness = false \ No newline at end of file diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index 398177dcd9a..5b7184f222c 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -83,13 +83,13 @@ fn spark_contains(haystack: &ColumnarValue, needle: &ColumnarValue) -> Result { - let result = contains_with_arrow_scalar(haystack_array, needle_scalar)?; + let result = contains_array_scalar(haystack_array, needle_scalar)?; Ok(ColumnarValue::Array(result)) } // Scalar haystack, array needle - less common (ColumnarValue::Scalar(haystack_scalar), ColumnarValue::Array(needle_array)) => { - let result = contains_scalar_with_arrow(haystack_scalar, needle_array)?; + let result = contains_scalar_array(haystack_scalar, needle_array)?; Ok(ColumnarValue::Array(result)) } @@ -103,7 +103,7 @@ fn spark_contains(haystack: &ColumnarValue, needle: &ColumnarValue) -> Result(scalar: &'a ScalarValue, arg_name: &str) -> Result<&'a str> { +fn get_string_scalar_value<'a>(scalar: &'a ScalarValue, arg_name: &str) -> Result<&'a str> { match scalar { ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) @@ -118,7 +118,7 @@ fn get_string<'a>(scalar: &'a ScalarValue, arg_name: &str) -> Result<&'a str> { /// Optimized contains for array haystack with scalar needle. /// Uses Arrow's native scalar handling for better performance. -fn contains_with_arrow_scalar( +fn contains_array_scalar( haystack_array: &ArrayRef, needle_scalar: &ScalarValue, ) -> Result { @@ -128,7 +128,7 @@ fn contains_with_arrow_scalar( } // Extract the needle string - let needle_str = get_string(needle_scalar, "needle")?; + let needle_str = get_string_scalar_value(needle_scalar, "needle")?; // Create scalar array for needle - tells Arrow to use optimized paths let needle_scalar_array = StringArray::new_scalar(needle_str); @@ -138,7 +138,7 @@ fn contains_with_arrow_scalar( Ok(Arc::new(result)) } -fn contains_scalar_with_arrow( +fn contains_scalar_array( haystack_scalar: &ScalarValue, needle_array: &ArrayRef, ) -> Result { @@ -146,7 +146,7 @@ fn contains_scalar_with_arrow( return Ok(Arc::new(BooleanArray::new_null(needle_array.len()))); } - let haystack_str = get_string(haystack_scalar, "haystack")?; + let haystack_str = get_string_scalar_value(haystack_scalar, "haystack")?; let haystack_scalar_array = StringArray::new_scalar(haystack_str.to_string()); let result = arrow_contains(&haystack_scalar_array, needle_array)?; @@ -163,8 +163,8 @@ fn contains_scalar_scalar( return Ok(ScalarValue::Boolean(None)); } - let haystack_str = get_string(haystack_scalar, "haystack")?; - let needle_str = get_string(needle_scalar, "needle")?; + let haystack_str = get_string_scalar_value(haystack_scalar, "haystack")?; + let needle_str = get_string_scalar_value(needle_scalar, "needle")?; Ok(ScalarValue::Boolean(Some( haystack_str.contains(needle_str), @@ -186,7 +186,7 @@ mod tests { ])) as ArrayRef; let needle = ScalarValue::Utf8(Some("world".to_string())); - let result = contains_with_arrow_scalar(&haystack, &needle).unwrap(); + let result = contains_array_scalar(&haystack, &needle).unwrap(); let bool_array = result.as_any().downcast_ref::().unwrap(); assert!(bool_array.value(0)); // "hello world" contains "world" @@ -216,7 +216,7 @@ mod tests { ])) as ArrayRef; let needle = ScalarValue::Utf8(None); - let result = contains_with_arrow_scalar(&haystack, &needle).unwrap(); + let result = contains_array_scalar(&haystack, &needle).unwrap(); let bool_array = result.as_any().downcast_ref::().unwrap(); // Null needle should produce null results @@ -229,11 +229,48 @@ mod tests { let haystack = Arc::new(StringArray::from(vec![Some("hello world"), Some("")])) as ArrayRef; let needle = ScalarValue::Utf8(Some("".to_string())); - let result = contains_with_arrow_scalar(&haystack, &needle).unwrap(); + let result = contains_array_scalar(&haystack, &needle).unwrap(); let bool_array = result.as_any().downcast_ref::().unwrap(); // Empty string is contained in any string assert!(bool_array.value(0)); assert!(bool_array.value(1)); } + + #[test] + fn test_contains_scalar_array_null_haystack() { + let haystack = ScalarValue::Utf8(None); + let needle = Arc::new(StringArray::from(vec![ + Some("hello world"), + Some("foo bar"), + ])) as ArrayRef; + + let result = contains_scalar_array(&haystack, &needle).unwrap(); + let bool_array = result.as_any().downcast_ref::().unwrap(); + + // Null haystack should produce null results for all array elements + assert!(bool_array.is_null(0)); + assert!(bool_array.is_null(1)); + } + + #[test] + fn test_spark_contains_dispatcher_scalar_array() { + let haystack = ColumnarValue::Scalar(ScalarValue::Utf8(Some("abc".to_string()))); + let needle = ColumnarValue::Array(Arc::new(StringArray::from(vec![ + Some("a"), + Some("bc"), + Some("d"), + ])) as ArrayRef); + + let result = spark_contains(&haystack, &needle).unwrap(); + let array = match result { + ColumnarValue::Array(arr) => arr, + _ => panic!("Expected ColumnarValue::Array"), + }; + let bool_array = array.as_any().downcast_ref::().unwrap(); + + assert!(bool_array.value(0)); + assert!(bool_array.value(1)); + assert!(!bool_array.value(2)); + } } From c1ec9a3d2da5610cf362d8aa41d8aa891bd2c548 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 9 Aug 2026 22:14:48 +0400 Subject: [PATCH 05/12] work --- native/spark-expr/benches/contains.rs | 122 ++++++++++++++++++++++++++ 1 file changed, 122 insertions(+) create mode 100644 native/spark-expr/benches/contains.rs diff --git a/native/spark-expr/benches/contains.rs b/native/spark-expr/benches/contains.rs new file mode 100644 index 00000000000..3eb89d93e69 --- /dev/null +++ b/native/spark-expr/benches/contains.rs @@ -0,0 +1,122 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use criterion::{criterion_group, criterion_main, Criterion}; +use arrow::array::{ArrayRef, StringArray}; +use arrow::datatypes::{DataType, Field}; +use datafusion::common::ScalarValue; +use datafusion::config::ConfigOptions; +use datafusion::logical_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl}; +use std::sync::Arc; +use datafusion_comet_spark_expr::SparkContains; + +fn generate_string_array(size: usize) -> ArrayRef { + let data: Vec> = (0..size) + .map(|i| { + if i % 10 == 0 { + None + } else { + Some(format!("hello string data sample number {} with some text", i)) + } + }) + .collect(); + Arc::new(StringArray::from(data)) +} + +fn bench_contains(c: &mut Criterion) { + let rows = 8192; + let udf = SparkContains::new(); + + let mut group = c.benchmark_group("string_funcs/contains"); + + let haystack_array = generate_string_array(rows); + let needle_scalar = ColumnarValue::Scalar(ScalarValue::Utf8(Some("sample".to_string()))); + let needle_array = generate_string_array(rows); + + // Общие метаданные для ScalarFunctionArgs + let arg_fields = vec![ + Arc::new(Field::new("haystack", DataType::Utf8, true)), + Arc::new(Field::new("needle", DataType::Utf8, true)), + ]; + let return_field = Arc::new(Field::new("result", DataType::Boolean, true)); + let config_options = Arc::new(ConfigOptions::new()); + + // 1. Array haystack vs Scalar needle (optimized path) + group.bench_function( + &format!("array_vs_scalar_size_{}", rows), + |b| { + b.iter(|| { + let args = ScalarFunctionArgs { + args: vec![ + ColumnarValue::Array(haystack_array.clone()), + needle_scalar.clone(), + ], + arg_fields: arg_fields.clone(), + number_rows: rows, + return_field: return_field.clone(), + config_options: config_options.clone(), + }; + std::hint::black_box(udf.invoke_with_args(args).unwrap()); + }); + }, + ); + + // 2. Array haystack vs Array needle + group.bench_function( + &format!("array_vs_array_size_{}", rows), + |b| { + b.iter(|| { + let args = ScalarFunctionArgs { + args: vec![ + ColumnarValue::Array(haystack_array.clone()), + ColumnarValue::Array(needle_array.clone()), + ], + arg_fields: arg_fields.clone(), + number_rows: rows, + return_field: return_field.clone(), + config_options: config_options.clone(), + }; + std::hint::black_box(udf.invoke_with_args(args).unwrap()); + }); + }, + ); + + let haystack_scalar_val = ColumnarValue::Scalar(ScalarValue::Utf8(Some("sample".to_string()))); + group.bench_function( + &format!("scalar_vs_array_size_{}", rows), + |b| { + b.iter(|| { + let args = ScalarFunctionArgs { + args: vec![ + haystack_scalar_val.clone(), + ColumnarValue::Array(needle_array.clone()), + ], + arg_fields: arg_fields.clone(), + number_rows: rows, + return_field: return_field.clone(), + config_options: config_options.clone(), + }; + std::hint::black_box(udf.invoke_with_args(args).unwrap()); + }); + }, + ); + + group.finish(); +} + +criterion_group!(benches, bench_contains); +criterion_main!(benches); From beda70bf0a31a319f9c416b70a7ed55256b28045 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 9 Aug 2026 22:20:31 +0400 Subject: [PATCH 06/12] fmt --- native/spark-expr/benches/contains.rs | 109 +++++++++--------- .../spark-expr/src/string_funcs/contains.rs | 9 +- 2 files changed, 56 insertions(+), 62 deletions(-) diff --git a/native/spark-expr/benches/contains.rs b/native/spark-expr/benches/contains.rs index 3eb89d93e69..29f885f143f 100644 --- a/native/spark-expr/benches/contains.rs +++ b/native/spark-expr/benches/contains.rs @@ -15,22 +15,26 @@ // specific language governing permissions and limitations // under the License. -use criterion::{criterion_group, criterion_main, Criterion}; use arrow::array::{ArrayRef, StringArray}; use arrow::datatypes::{DataType, Field}; +use criterion::{criterion_group, criterion_main, Criterion}; use datafusion::common::ScalarValue; use datafusion::config::ConfigOptions; use datafusion::logical_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl}; -use std::sync::Arc; use datafusion_comet_spark_expr::SparkContains; +use std::sync::Arc; + fn generate_string_array(size: usize) -> ArrayRef { let data: Vec> = (0..size) .map(|i| { if i % 10 == 0 { None } else { - Some(format!("hello string data sample number {} with some text", i)) + Some(format!( + "hello string data sample number {} with some text", + i + )) } }) .collect(); @@ -56,64 +60,55 @@ fn bench_contains(c: &mut Criterion) { let config_options = Arc::new(ConfigOptions::new()); // 1. Array haystack vs Scalar needle (optimized path) - group.bench_function( - &format!("array_vs_scalar_size_{}", rows), - |b| { - b.iter(|| { - let args = ScalarFunctionArgs { - args: vec![ - ColumnarValue::Array(haystack_array.clone()), - needle_scalar.clone(), - ], - arg_fields: arg_fields.clone(), - number_rows: rows, - return_field: return_field.clone(), - config_options: config_options.clone(), - }; - std::hint::black_box(udf.invoke_with_args(args).unwrap()); - }); - }, - ); + group.bench_function(&format!("array_vs_scalar_size_{}", rows), |b| { + b.iter(|| { + let args = ScalarFunctionArgs { + args: vec![ + ColumnarValue::Array(haystack_array.clone()), + needle_scalar.clone(), + ], + arg_fields: arg_fields.clone(), + number_rows: rows, + return_field: return_field.clone(), + config_options: config_options.clone(), + }; + std::hint::black_box(udf.invoke_with_args(args).unwrap()); + }); + }); // 2. Array haystack vs Array needle - group.bench_function( - &format!("array_vs_array_size_{}", rows), - |b| { - b.iter(|| { - let args = ScalarFunctionArgs { - args: vec![ - ColumnarValue::Array(haystack_array.clone()), - ColumnarValue::Array(needle_array.clone()), - ], - arg_fields: arg_fields.clone(), - number_rows: rows, - return_field: return_field.clone(), - config_options: config_options.clone(), - }; - std::hint::black_box(udf.invoke_with_args(args).unwrap()); - }); - }, - ); + group.bench_function(&format!("array_vs_array_size_{}", rows), |b| { + b.iter(|| { + let args = ScalarFunctionArgs { + args: vec![ + ColumnarValue::Array(haystack_array.clone()), + ColumnarValue::Array(needle_array.clone()), + ], + arg_fields: arg_fields.clone(), + number_rows: rows, + return_field: return_field.clone(), + config_options: config_options.clone(), + }; + std::hint::black_box(udf.invoke_with_args(args).unwrap()); + }); + }); let haystack_scalar_val = ColumnarValue::Scalar(ScalarValue::Utf8(Some("sample".to_string()))); - group.bench_function( - &format!("scalar_vs_array_size_{}", rows), - |b| { - b.iter(|| { - let args = ScalarFunctionArgs { - args: vec![ - haystack_scalar_val.clone(), - ColumnarValue::Array(needle_array.clone()), - ], - arg_fields: arg_fields.clone(), - number_rows: rows, - return_field: return_field.clone(), - config_options: config_options.clone(), - }; - std::hint::black_box(udf.invoke_with_args(args).unwrap()); - }); - }, - ); + group.bench_function(&format!("scalar_vs_array_size_{}", rows), |b| { + b.iter(|| { + let args = ScalarFunctionArgs { + args: vec![ + haystack_scalar_val.clone(), + ColumnarValue::Array(needle_array.clone()), + ], + arg_fields: arg_fields.clone(), + number_rows: rows, + return_field: return_field.clone(), + config_options: config_options.clone(), + }; + std::hint::black_box(udf.invoke_with_args(args).unwrap()); + }); + }); group.finish(); } diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index 5b7184f222c..6b805e30fdd 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -256,11 +256,10 @@ mod tests { #[test] fn test_spark_contains_dispatcher_scalar_array() { let haystack = ColumnarValue::Scalar(ScalarValue::Utf8(Some("abc".to_string()))); - let needle = ColumnarValue::Array(Arc::new(StringArray::from(vec![ - Some("a"), - Some("bc"), - Some("d"), - ])) as ArrayRef); + let needle = + ColumnarValue::Array( + Arc::new(StringArray::from(vec![Some("a"), Some("bc"), Some("d")])) as ArrayRef, + ); let result = spark_contains(&haystack, &needle).unwrap(); let array = match result { From ffc9e5fd80ce6d68bb375aa2f16258aaa9152c52 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Mon, 10 Aug 2026 22:32:10 +0400 Subject: [PATCH 07/12] clippy --- native/spark-expr/src/string_funcs/contains.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index 6b805e30fdd..1c21fc82f23 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -147,7 +147,7 @@ fn contains_scalar_array( } let haystack_str = get_string_scalar_value(haystack_scalar, "haystack")?; - let haystack_scalar_array = StringArray::new_scalar(haystack_str.to_string()); + let haystack_scalar_array = StringArray::new_scalar(haystack_str); let result = arrow_contains(&haystack_scalar_array, needle_array)?; Ok(Arc::new(result)) From b46a4f992f17d5c4af38f05c5ff2cc37fbbee5b5 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Tue, 11 Aug 2026 20:36:08 +0400 Subject: [PATCH 08/12] clippy --- native/spark-expr/benches/contains.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/native/spark-expr/benches/contains.rs b/native/spark-expr/benches/contains.rs index 29f885f143f..f90551de472 100644 --- a/native/spark-expr/benches/contains.rs +++ b/native/spark-expr/benches/contains.rs @@ -60,7 +60,7 @@ fn bench_contains(c: &mut Criterion) { let config_options = Arc::new(ConfigOptions::new()); // 1. Array haystack vs Scalar needle (optimized path) - group.bench_function(&format!("array_vs_scalar_size_{}", rows), |b| { + group.bench_function(format!("array_vs_scalar_size_{}", rows), |b| { b.iter(|| { let args = ScalarFunctionArgs { args: vec![ @@ -77,7 +77,7 @@ fn bench_contains(c: &mut Criterion) { }); // 2. Array haystack vs Array needle - group.bench_function(&format!("array_vs_array_size_{}", rows), |b| { + group.bench_function(format!("array_vs_array_size_{}", rows), |b| { b.iter(|| { let args = ScalarFunctionArgs { args: vec![ @@ -94,7 +94,7 @@ fn bench_contains(c: &mut Criterion) { }); let haystack_scalar_val = ColumnarValue::Scalar(ScalarValue::Utf8(Some("sample".to_string()))); - group.bench_function(&format!("scalar_vs_array_size_{}", rows), |b| { + group.bench_function(format!("scalar_vs_array_size_{}", rows), |b| { b.iter(|| { let args = ScalarFunctionArgs { args: vec![ From 2f8aef04c6f2c1bc815c97eb06f6dc8af44f8513 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Thu, 20 Aug 2026 12:13:42 +0400 Subject: [PATCH 09/12] address comments --- .../spark-expr/src/string_funcs/contains.rs | 48 +++++++++++++++++-- 1 file changed, 43 insertions(+), 5 deletions(-) diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index 1c21fc82f23..f75222a6f67 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -20,7 +20,7 @@ //! Optimized for scalar pattern case by passing scalar directly to arrow_contains //! instead of expanding to arrays like DataFusion's built-in contains. -use arrow::array::{Array, ArrayRef, BooleanArray, StringArray}; +use arrow::array::{Array, ArrayRef, BooleanArray, Scalar, StringArray}; use arrow::compute::kernels::comparison::contains as arrow_contains; use arrow::datatypes::DataType; use datafusion::common::{exec_err, Result, ScalarValue}; @@ -146,9 +146,7 @@ fn contains_scalar_array( return Ok(Arc::new(BooleanArray::new_null(needle_array.len()))); } - let haystack_str = get_string_scalar_value(haystack_scalar, "haystack")?; - let haystack_scalar_array = StringArray::new_scalar(haystack_str); - + let haystack_scalar_array = Scalar::new(haystack_scalar.to_array()?); let result = arrow_contains(&haystack_scalar_array, needle_array)?; Ok(Arc::new(result)) } @@ -174,7 +172,7 @@ fn contains_scalar_scalar( #[cfg(test)] mod tests { use super::*; - use arrow::array::StringArray; + use arrow::array::{LargeStringArray, StringArray, StringViewArray}; #[test] fn test_contains_array_scalar() { @@ -272,4 +270,44 @@ mod tests { assert!(bool_array.value(1)); assert!(!bool_array.value(2)); } + + #[test] + fn test_contains_scalar_large_utf8() { + let haystack = ScalarValue::LargeUtf8(Some("abc".to_string())); + let needle = Arc::new(LargeStringArray::from(vec![ + Some("a"), + Some("bc"), + None, + Some(""), + Some("d"), + ])) as ArrayRef; + + let res = contains_scalar_array(&haystack, &needle).unwrap(); + let res = res.as_any().downcast_ref::().unwrap(); + + let expected = + BooleanArray::from(vec![Some(true), Some(true), None, Some(true), Some(false)]); + + assert_eq!(res, &expected); + } + + #[test] + fn test_contains_scalar_utf8_view() { + let haystack = ScalarValue::Utf8View(Some("abc".to_string())); + let needle = Arc::new(StringViewArray::from(vec![ + Some("a"), + Some("bc"), + None, + Some(""), + Some("d"), + ])) as ArrayRef; + + let res = contains_scalar_array(&haystack, &needle).unwrap(); + let res = res.as_any().downcast_ref::().unwrap(); + + let expected = + BooleanArray::from(vec![Some(true), Some(true), None, Some(true), Some(false)]); + + assert_eq!(res, &expected); + } } From 79efafd8a7f46330f986941c374155099d0dcf39 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sat, 29 Aug 2026 21:26:59 +0400 Subject: [PATCH 10/12] address comments --- .../spark-expr/src/string_funcs/contains.rs | 53 ++++++++++++++++--- 1 file changed, 46 insertions(+), 7 deletions(-) diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index f75222a6f67..ff330a683ad 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -20,7 +20,7 @@ //! Optimized for scalar pattern case by passing scalar directly to arrow_contains //! instead of expanding to arrays like DataFusion's built-in contains. -use arrow::array::{Array, ArrayRef, BooleanArray, Scalar, StringArray}; +use arrow::array::{Array, ArrayRef, BooleanArray, Scalar}; use arrow::compute::kernels::comparison::contains as arrow_contains; use arrow::datatypes::DataType; use datafusion::common::{exec_err, Result, ScalarValue}; @@ -127,13 +127,9 @@ fn contains_array_scalar( return Ok(Arc::new(BooleanArray::new_null(haystack_array.len()))); } - // Extract the needle string - let needle_str = get_string_scalar_value(needle_scalar, "needle")?; - - // Create scalar array for needle - tells Arrow to use optimized paths - let needle_scalar_array = StringArray::new_scalar(needle_str); + let _ = get_string_scalar_value(needle_scalar, "needle")?; - // Use Arrow's contains which detects scalar case and uses optimized paths + let needle_scalar_array = Scalar::new(needle_scalar.to_array()?); let result = arrow_contains(haystack_array, &needle_scalar_array)?; Ok(Arc::new(result)) } @@ -146,6 +142,8 @@ fn contains_scalar_array( return Ok(Arc::new(BooleanArray::new_null(needle_array.len()))); } + let _ = get_string_scalar_value(haystack_scalar, "haystack")?; + let haystack_scalar_array = Scalar::new(haystack_scalar.to_array()?); let result = arrow_contains(&haystack_scalar_array, needle_array)?; Ok(Arc::new(result)) @@ -310,4 +308,45 @@ mod tests { assert_eq!(res, &expected); } + + #[test] + fn test_contains_scalar_array_all_cases() { + let haystack = ScalarValue::Utf8(Some("hello world".to_string())); + let needle = Arc::new(StringArray::from(vec![ + Some("hello"), + Some("world"), + Some("foo"), + None, + ])) as ArrayRef; + + let res = contains_scalar_array(&haystack, &needle).unwrap(); + let bool_arr = res.as_any().downcast_ref::().unwrap(); + + assert_eq!( + bool_arr, + &BooleanArray::from(vec![Some(true), Some(true), Some(false), None]) + ); + } + + #[test] + fn test_contains_scalar_array_empty_needle() { + let haystack = ScalarValue::Utf8(Some("hello world".to_string())); + let needle = Arc::new(StringArray::from(Vec::>::new())) as ArrayRef; + + let res = contains_scalar_array(&haystack, &needle).unwrap(); + assert_eq!(res.len(), 0); + } + + #[test] + fn test_contains_scalar_array_invalid_type_error() { + let haystack = ScalarValue::Int32(Some(123)); + let needle = Arc::new(StringArray::from(vec![Some("1")])) as ArrayRef; + + let err = contains_scalar_array(&haystack, &needle).unwrap_err(); + assert!( + err.to_string() + .contains("contains function requires string type for haystack, got Int32"), + "Actual error: {err}" + ); + } } From 37d1ddda797ef803e072edec97559bc218b022cd Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 13 Sep 2026 11:13:03 +0400 Subject: [PATCH 11/12] address comments --- .../spark-expr/src/string_funcs/contains.rs | 101 +++++++++++++----- 1 file changed, 76 insertions(+), 25 deletions(-) diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index ff330a683ad..545d1475600 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -1,29 +1,25 @@ // Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file +// or more contributor license agreements. See the NOTICE file // distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file +// 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 +// 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 +// KIND, either express or implied. See the License for the // specific language governing permissions and limitations // under the License. -//! Optimized `contains` string function for Spark compatibility. -//! -//! Optimized for scalar pattern case by passing scalar directly to arrow_contains -//! instead of expanding to arrays like DataFusion's built-in contains. - use arrow::array::{Array, ArrayRef, BooleanArray, Scalar}; +use arrow::compute::kernels::cast::cast; use arrow::compute::kernels::comparison::contains as arrow_contains; use arrow::datatypes::DataType; -use datafusion::common::{exec_err, Result, ScalarValue}; +use datafusion::common::{exec_err, DataFusionError, Result, ScalarValue}; use datafusion::logical_expr::{ ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl, Signature, Volatility, }; @@ -101,13 +97,15 @@ fn spark_contains(haystack: &ColumnarValue, needle: &ColumnarValue) -> Result(scalar: &'a ScalarValue, arg_name: &str) -> Result<&'a str> { match scalar { ScalarValue::Utf8(Some(s)) | ScalarValue::LargeUtf8(Some(s)) | ScalarValue::Utf8View(Some(s)) => Ok(s.as_str()), + ScalarValue::Dictionary(_, inner) => get_string_scalar_value(inner, arg_name), _ => exec_err!( "contains function requires string type for {}, got {:?}", arg_name, @@ -116,6 +114,23 @@ fn get_string_scalar_value<'a>(scalar: &'a ScalarValue, arg_name: &str) -> Resul } } +/// Materialize a scalar into a length-1 array whose type matches `target_type`, +/// so Arrow's CONTAINS kernel accepts the (scalar, array) pair. +/// Cost is O(1): the cast touches a single element. +fn scalar_to_aligned_array( + scalar: &ScalarValue, + target_type: &DataType, + arg_name: &str, +) -> Result { + let _ = get_string_scalar_value(scalar, arg_name)?; + let array = scalar.to_array()?; + if array.data_type() == target_type { + Ok(array) + } else { + cast(&array, target_type).map_err(DataFusionError::from) + } +} + /// Optimized contains for array haystack with scalar needle. /// Uses Arrow's native scalar handling for better performance. fn contains_array_scalar( @@ -126,26 +141,24 @@ fn contains_array_scalar( if needle_scalar.is_null() { return Ok(Arc::new(BooleanArray::new_null(haystack_array.len()))); } - - let _ = get_string_scalar_value(needle_scalar, "needle")?; - - let needle_scalar_array = Scalar::new(needle_scalar.to_array()?); - let result = arrow_contains(haystack_array, &needle_scalar_array)?; + let needle_array = + scalar_to_aligned_array(needle_scalar, haystack_array.data_type(), "needle")?; + let result = arrow_contains(haystack_array, &Scalar::new(needle_array))?; Ok(Arc::new(result)) } +/// Contains for scalar haystack with array needle - less common path. fn contains_scalar_array( haystack_scalar: &ScalarValue, needle_array: &ArrayRef, ) -> Result { + // Handle null haystack if haystack_scalar.is_null() { return Ok(Arc::new(BooleanArray::new_null(needle_array.len()))); } - - let _ = get_string_scalar_value(haystack_scalar, "haystack")?; - - let haystack_scalar_array = Scalar::new(haystack_scalar.to_array()?); - let result = arrow_contains(&haystack_scalar_array, needle_array)?; + let haystack_array = + scalar_to_aligned_array(haystack_scalar, needle_array.data_type(), "haystack")?; + let result = arrow_contains(&Scalar::new(haystack_array), needle_array)?; Ok(Arc::new(result)) } @@ -154,7 +167,6 @@ fn contains_scalar_scalar( haystack_scalar: &ScalarValue, needle_scalar: &ScalarValue, ) -> Result { - // Handle nulls if haystack_scalar.is_null() || needle_scalar.is_null() { return Ok(ScalarValue::Boolean(None)); } @@ -170,7 +182,8 @@ fn contains_scalar_scalar( #[cfg(test)] mod tests { use super::*; - use arrow::array::{LargeStringArray, StringArray, StringViewArray}; + use arrow::array::{DictionaryArray, LargeStringArray, StringArray, StringViewArray}; + use arrow::datatypes::Int32Type; #[test] fn test_contains_array_scalar() { @@ -309,6 +322,44 @@ mod tests { assert_eq!(res, &expected); } + #[test] + fn test_contains_scalar_dictionary() { + // Regression: a non-null dictionary-string scalar previously worked before + // the optimization, then started failing at `get_string_scalar_value`. + let haystack = ScalarValue::Dictionary( + Box::new(DataType::Int32), + Box::new(ScalarValue::Utf8(Some("abc".to_string()))), + ); + let needle = Arc::new(DictionaryArray::::from_iter(vec![ + Some("a"), + Some("bc"), + None, + Some(""), + Some("d"), + ])) as ArrayRef; + + let res = contains_scalar_array(&haystack, &needle).unwrap(); + let res = res.as_any().downcast_ref::().unwrap(); + + let expected = + BooleanArray::from(vec![Some(true), Some(true), None, Some(true), Some(false)]); + + assert_eq!(res, &expected); + } + + #[test] + fn test_contains_array_scalar_large_utf8_haystack() { + // Symmetric case: scalar needle must be aligned to the array's type, + // so a Utf8 needle works against a LargeUtf8 haystack. + let haystack = Arc::new(LargeStringArray::from(vec![Some("abc"), Some("xyz")])) as ArrayRef; + let needle = ScalarValue::Utf8(Some("bc".to_string())); + + let res = contains_array_scalar(&haystack, &needle).unwrap(); + let res = res.as_any().downcast_ref::().unwrap(); + + assert_eq!(res, &BooleanArray::from(vec![Some(true), Some(false)])); + } + #[test] fn test_contains_scalar_array_all_cases() { let haystack = ScalarValue::Utf8(Some("hello world".to_string())); @@ -345,8 +396,8 @@ mod tests { let err = contains_scalar_array(&haystack, &needle).unwrap_err(); assert!( err.to_string() - .contains("contains function requires string type for haystack, got Int32"), - "Actual error: {err}" + .contains("contains function requires string type for haystack"), + "unexpected error: {err}" ); } } From 486b4fc99dee856b251d9b87174ec2ff4b316970 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 13 Sep 2026 11:58:43 +0400 Subject: [PATCH 12/12] address comments --- native/spark-expr/benches/contains.rs | 119 ++++++++++++++++++++++---- 1 file changed, 102 insertions(+), 17 deletions(-) diff --git a/native/spark-expr/benches/contains.rs b/native/spark-expr/benches/contains.rs index 457ce79db35..22fa8eec23a 100644 --- a/native/spark-expr/benches/contains.rs +++ b/native/spark-expr/benches/contains.rs @@ -28,9 +28,36 @@ use std::sync::Arc; mod common; use common::{string_array, NULL_RATIOS, ROW_COUNTS}; +/// Scalar used as the haystack in the scalar/array shape. The matching needle +/// values below are chosen so the scalar/array shape performs real work. +const HAYSTACK_SCALAR: &str = "datafusion-comet"; + +/// Scalar used as the needle in the array/scalar shape. +const NEEDLE_SCALAR: &str = "comet"; + +fn build_args( + haystack: ColumnarValue, + needle: ColumnarValue, + number_rows: usize, +) -> ScalarFunctionArgs { + ScalarFunctionArgs { + args: vec![haystack, needle], + arg_fields: vec![], + number_rows, + return_field: Arc::new(Field::new("result", DataType::Boolean, true)), + config_options: Arc::new(ConfigOptions::default()), + } +} + fn criterion_benchmark(c: &mut Criterion) { let udf = SparkContains::new(); - let mut group = c.benchmark_group("spark_contains"); + + // ------------------------------------------------------------------ + // Shape 1: array haystack vs scalar needle (`contains_array_scalar`). + // This path already used a scalar representation on `main`; included as a + // regression control since this PR touches it incidentally. + // ------------------------------------------------------------------ + let mut group = c.benchmark_group("spark_contains/array_scalar"); for rows in ROW_COUNTS { for (null_ratio, tag) in NULL_RATIOS { let haystack = string_array(rows, null_ratio, |_| "datafusion-comet".to_string()); @@ -40,22 +67,80 @@ fn criterion_benchmark(c: &mut Criterion) { |b, haystack| { b.iter(|| { black_box( - udf.invoke_with_args(ScalarFunctionArgs { - args: vec![ - ColumnarValue::Array(Arc::clone(haystack)), - ColumnarValue::Scalar(ScalarValue::Utf8(Some( - "comet".to_string(), - ))), - ], - arg_fields: vec![], - number_rows: haystack.len(), - return_field: Arc::new(Field::new( - "result", - DataType::Boolean, - true, - )), - config_options: Arc::new(ConfigOptions::default()), - }) + udf.invoke_with_args(build_args( + ColumnarValue::Array(Arc::clone(haystack)), + ColumnarValue::Scalar(ScalarValue::Utf8(Some( + NEEDLE_SCALAR.to_string(), + ))), + haystack.len(), + )) + .unwrap(), + ) + }) + }, + ); + } + } + group.finish(); + + // ------------------------------------------------------------------ + // Shape 2: scalar haystack vs array needle (`contains_scalar_array`). + // This is the path the PR actually optimizes (it replaced + // `to_array_of_size(N)` with an O(1) broadcast), so it must be measured. + // The needle array is varied per row so the kernel does non-trivial work. + // ------------------------------------------------------------------ + let mut group = c.benchmark_group("spark_contains/scalar_array"); + for rows in ROW_COUNTS { + for (null_ratio, tag) in NULL_RATIOS { + let needle = string_array(rows, null_ratio, |i| { + if i % 2 == 0 { + "comet".to_string() + } else { + format!("comet-{i}") + } + }); + group.bench_with_input( + BenchmarkId::from_parameter(format!("{rows}/{tag}")), + &needle, + |b, needle| { + b.iter(|| { + black_box( + udf.invoke_with_args(build_args( + ColumnarValue::Scalar(ScalarValue::Utf8(Some( + HAYSTACK_SCALAR.to_string(), + ))), + ColumnarValue::Array(Arc::clone(needle)), + needle.len(), + )) + .unwrap(), + ) + }) + }, + ); + } + } + group.finish(); + + // ------------------------------------------------------------------ + // Shape 3: array haystack vs array needle (`arrow_contains` directly). + // Regression control for the straight-through path the PR does not touch. + // ------------------------------------------------------------------ + let mut group = c.benchmark_group("spark_contains/array_array"); + for rows in ROW_COUNTS { + for (null_ratio, tag) in NULL_RATIOS { + let haystack = string_array(rows, null_ratio, |_| "datafusion-comet".to_string()); + let needle = string_array(rows, null_ratio, |_| "comet".to_string()); + group.bench_with_input( + BenchmarkId::from_parameter(format!("{rows}/{tag}")), + &(haystack, needle), + |b, (haystack, needle)| { + b.iter(|| { + black_box( + udf.invoke_with_args(build_args( + ColumnarValue::Array(Arc::clone(haystack)), + ColumnarValue::Array(Arc::clone(needle)), + haystack.len(), + )) .unwrap(), ) })