diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index 6faa9fec4e..58eeeadf67 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/benches/contains.rs b/native/spark-expr/benches/contains.rs new file mode 100644 index 0000000000..29f885f143 --- /dev/null +++ b/native/spark-expr/benches/contains.rs @@ -0,0 +1,117 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +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 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 + )) + } + }) + .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); diff --git a/native/spark-expr/src/string_funcs/contains.rs b/native/spark-expr/src/string_funcs/contains.rs index 537227efdf..1c21fc82f2 100644 --- a/native/spark-expr/src/string_funcs/contains.rs +++ b/native/spark-expr/src/string_funcs/contains.rs @@ -83,15 +83,14 @@ 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 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_array(haystack_scalar, needle_array)?; + Ok(ColumnarValue::Array(result)) } // Both scalars - compute single result @@ -102,9 +101,24 @@ 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( +fn contains_array_scalar( haystack_array: &ArrayRef, needle_scalar: &ScalarValue, ) -> Result { @@ -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_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); @@ -134,6 +138,21 @@ fn contains_with_arrow_scalar( Ok(Arc::new(result)) } +fn contains_scalar_array( + 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_scalar_value(haystack_scalar, "haystack")?; + let haystack_scalar_array = StringArray::new_scalar(haystack_str); + + 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_scalar_value(haystack_scalar, "haystack")?; + let needle_str = get_string_scalar_value(needle_scalar, "needle")?; Ok(ScalarValue::Boolean(Some( haystack_str.contains(needle_str), @@ -188,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" @@ -218,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 @@ -231,11 +229,47 @@ 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)); + } }