kazantsev-maksim commented on code in PR #5042:
URL: https://github.com/apache/datafusion-comet/pull/5042#discussion_r4130058613


##########
native/spark-expr/src/string_funcs/levenshtein.rs:
##########
@@ -87,399 +179,405 @@ fn levenshtein_distance_with_threshold(s: &str, t: &str, 
threshold: i32) -> i32
         return -1;
     }
 
+    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 as i32 - n as i32).abs() > threshold {
+            return -1;
+        }
+        if m == 0 {
+            return if n as i32 <= threshold { n as i32 } else { -1 };
+        }
+        if n == 0 {
+            return if m as i32 <= threshold { m as i32 } else { -1 };
+        }
+
+        let (s_bytes, t_bytes, m, n) = if m > n {
+            (t_bytes, s_bytes, n, m)
+        } else {
+            (s_bytes, t_bytes, m, n)
+        };
+
+        if (n as i32 - m as i32) > threshold {
+            return -1;
+        }
+
+        let out_of_band = threshold + 1;
+
+        return with_scratch_buffers(m + 1, out_of_band, |prev, curr| {
+            for (i, val) in prev.iter_mut().enumerate() {
+                *val = if i as i32 <= threshold {
+                    i as i32
+                } else {
+                    out_of_band
+                };
+            }
+
+            for (j, &t_byte) in t_bytes.iter().enumerate().take(n) {
+                let j_1 = (j + 1) as i32;
+                curr[0] = if j_1 <= threshold { j_1 } else { out_of_band };
+
+                let min_i = (j_1 - threshold).max(1) as usize;
+                let max_i = ((j_1 + threshold) as usize).min(m);
+
+                if min_i > 1 {
+                    curr[min_i - 1] = out_of_band;
+                }
+
+                assert!(prev.len() > m && curr.len() > m);
+                assert!(s_bytes.len() >= m);
+
+                for i in min_i..=max_i {
+                    let cost = if s_bytes[i - 1] == t_byte { 0 } else { 1 };
+                    curr[i] = (prev[i] + 1).min(curr[i - 1] + 1).min(prev[i - 
1] + cost);
+                }
+
+                if max_i < m {
+                    curr[max_i + 1] = out_of_band;
+                }
+
+                std::mem::swap(prev, curr);
+            }
+
+            let result = prev[m];
+            if result <= threshold {
+                result
+            } else {
+                -1
+            }
+        });
+    }
+
     let s_chars: Vec<char> = s.chars().collect();
     let t_chars: Vec<char> = t.chars().collect();
-    let (shorter, longer) = if s_chars.len() <= t_chars.len() {
-        (s_chars, t_chars)
-    } else {
-        (t_chars, s_chars)
-    };
-    let m = shorter.len();
-    let n = longer.len();
-    let threshold = threshold as usize;
+    let m = s_chars.len();
+    let n = t_chars.len();
 
-    if n - m > threshold {
+    if (m as i32 - n as i32).abs() > threshold {
         return -1;
     }
     if m == 0 {
-        return if n <= threshold { n as i32 } else { -1 };
+        return if n as i32 <= threshold { n as i32 } else { -1 };
     }
+    if n == 0 {
+        return if m as i32 <= threshold { m as i32 } else { -1 };
+    }
+
+    let (s_chars, t_chars, m, n) = if m > n {
+        (t_chars, s_chars, n, m)
+    } else {
+        (s_chars, t_chars, m, n)
+    };
 
-    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;
+    if (n as i32 - m as i32) > threshold {
+        return -1;
     }
 
-    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;
-        }
+    let out_of_band = threshold + 1;
 
-        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));
-        }
-        if end < m {
-            curr[end + 1] = out_of_band;
+    with_scratch_buffers(m + 1, out_of_band, |prev, curr| {
+        for (i, val) in prev.iter_mut().enumerate() {
+            *val = if i as i32 <= threshold {
+                i as i32
+            } else {
+                out_of_band
+            };
         }
-        std::mem::swap(&mut prev, &mut curr);
-    }
 
-    if prev[m] <= threshold {
-        prev[m] as i32
-    } else {
-        -1
-    }
-}
+        for (j, &t_char) in t_chars.iter().enumerate().take(n) {
+            let j_1 = (j + 1) as i32;
+            curr[0] = if j_1 <= threshold { j_1 } else { out_of_band };
 
-fn evaluate_levenshtein<LeftOffset, RightOffset>(
-    left: &GenericStringArray<LeftOffset>,
-    right: &GenericStringArray<RightOffset>,
-    threshold: Option<&Int32Array>,
-) -> Int32Array
-where
-    LeftOffset: OffsetSizeTrait,
-    RightOffset: OffsetSizeTrait,
-{
-    left.iter()
-        .zip(right.iter())
-        .enumerate()
-        .map(|(i, (left_value, right_value))| {
-            if threshold.is_some_and(|values| values.is_null(i)) {
-                return None;
+            let min_i = (j_1 - threshold).max(1) as usize;
+            let max_i = ((j_1 + threshold) as usize).min(m);
+
+            if min_i > 1 {
+                curr[min_i - 1] = out_of_band;
             }
 
-            match (left_value, right_value) {
-                (Some(left_value), Some(right_value)) => Some(match threshold {
-                    Some(values) => levenshtein_distance_with_threshold(
-                        left_value,
-                        right_value,
-                        values.value(i),
-                    ),
-                    None => levenshtein_distance(left_value, right_value),
-                }),
-                _ => None,
+            assert!(prev.len() > m && curr.len() > m);
+            assert!(s_chars.len() >= m);
+
+            for i in min_i..=max_i {
+                let cost = if s_chars[i - 1] == t_char { 0 } else { 1 };
+                curr[i] = (prev[i] + 1).min(curr[i - 1] + 1).min(prev[i - 1] + 
cost);
             }
-        })
-        .collect()
+
+            if max_i < m {
+                curr[max_i + 1] = out_of_band;
+            }
+
+            std::mem::swap(prev, curr);
+        }
+
+        let result = prev[m];
+        if result <= threshold {
+            result
+        } else {
+            -1
+        }
+    })
 }
 
-fn evaluate_string_arrays(
-    left: &ArrayRef,
-    right: &ArrayRef,
-    threshold: Option<&Int32Array>,
-) -> Result<Int32Array> {
-    match (left.data_type(), right.data_type()) {
-        (DataType::Utf8, DataType::Utf8) => Ok(evaluate_levenshtein(
-            as_generic_string_array::<i32>(left.as_ref())?,
-            as_generic_string_array::<i32>(right.as_ref())?,
-            threshold,
-        )),
-        (DataType::Utf8, DataType::LargeUtf8) => Ok(evaluate_levenshtein(
-            as_generic_string_array::<i32>(left.as_ref())?,
-            as_generic_string_array::<i64>(right.as_ref())?,
-            threshold,
-        )),
-        (DataType::LargeUtf8, DataType::Utf8) => Ok(evaluate_levenshtein(
-            as_generic_string_array::<i64>(left.as_ref())?,
-            as_generic_string_array::<i32>(right.as_ref())?,
-            threshold,
-        )),
-        (DataType::LargeUtf8, DataType::LargeUtf8) => Ok(evaluate_levenshtein(
-            as_generic_string_array::<i64>(left.as_ref())?,
-            as_generic_string_array::<i64>(right.as_ref())?,
-            threshold,
-        )),
-        (left_type, right_type) => Err(DataFusionError::Execution(format!(
-            "levenshtein expects Utf8 or LargeUtf8 arguments, got 
{left_type:?} and {right_type:?}"
-        ))),
+fn levenshtein<O: OffsetSizeTrait>(
+    left: &GenericStringArray<O>,
+    right: &GenericStringArray<O>,
+) -> Result<ArrayRef> {
+    let mut builder = Int32Array::builder(left.len());
+    for i in 0..left.len() {
+        if left.is_null(i) || right.is_null(i) {
+            builder.append_null();
+        } else {
+            builder.append_value(levenshtein_distance(left.value(i), 
right.value(i)));
+        }
     }
+    Ok(Arc::new(builder.finish()) as ArrayRef)
 }
 
-/// 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<ColumnarValue> {
-    if args.len() < 2 || args.len() > 3 {
-        return Err(DataFusionError::Internal(format!(
-            "levenshtein requires 2 or 3 arguments, got {}",
-            args.len()
-        )));
+fn levenshtein_with_threshold<O: OffsetSizeTrait>(
+    left: &GenericStringArray<O>,
+    right: &GenericStringArray<O>,
+    threshold: &Int32Array,
+) -> Result<ArrayRef> {
+    let mut builder = Int32Array::builder(left.len());
+    for i in 0..left.len() {
+        if left.is_null(i) || right.is_null(i) || threshold.is_null(i) {
+            builder.append_null();
+        } else {
+            builder.append_value(levenshtein_distance_with_threshold(
+                left.value(i),
+                right.value(i),
+                threshold.value(i),
+            ));
+        }
     }
+    Ok(Arc::new(builder.finish()) as ArrayRef)
+}
 
-    // Determine array length from any array argument
-    let len = args
-        .iter()
-        .find_map(|arg| match arg {
-            ColumnarValue::Array(a) => Some(a.len()),
-            _ => None,
-        })
-        .unwrap_or(1);
-
-    let left = args[0].clone().into_array(len)?;
-    let right = args[1].clone().into_array(len)?;
-    if left.len() != len || right.len() != len {
-        return Err(DataFusionError::Internal(
-            "levenshtein arguments must have the same length".to_string(),
-        ));
-    }
+/// Computes the Levenshtein distance between two strings, matching Spark 
semantics.
+pub fn spark_levenshtein(args: &[ColumnarValue]) -> Result<ColumnarValue> {
+    match args.len() {
+        2 => {
+            if let (ColumnarValue::Scalar(s1), ColumnarValue::Scalar(s2)) = 
(&args[0], &args[1]) {
+                let res = match (s1, s2) {
+                    (ScalarValue::Utf8(Some(v1)), ScalarValue::Utf8(Some(v2)))
+                    | (ScalarValue::LargeUtf8(Some(v1)), 
ScalarValue::LargeUtf8(Some(v2)))
+                    | (ScalarValue::Utf8(Some(v1)), 
ScalarValue::LargeUtf8(Some(v2)))
+                    | (ScalarValue::LargeUtf8(Some(v1)), 
ScalarValue::Utf8(Some(v2))) => {
+                        Some(levenshtein_distance(v1, v2))
+                    }
+                    (ScalarValue::Utf8(None), _)
+                    | (_, ScalarValue::Utf8(None))
+                    | (ScalarValue::LargeUtf8(None), _)
+                    | (_, ScalarValue::LargeUtf8(None)) => None,
+                    _ => {
+                        return Err(DataFusionError::Internal(
+                            "Expected string scalar for 
levenshtein".to_string(),
+                        ))
+                    }
+                };
+                return Ok(ColumnarValue::Scalar(ScalarValue::Int32(res)));
+            }
 
-    // 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 {
-            return Err(DataFusionError::Internal(
-                "levenshtein threshold must have the same length as string 
arguments".to_string(),
-            ));
+            let num_rows = match (&args[0], &args[1]) {
+                (ColumnarValue::Array(a), _) | (_, ColumnarValue::Array(a)) => 
a.len(),
+                _ => unreachable!(),
+            };
+
+            let left = args[0].clone().into_array(num_rows)?;
+            let right = args[1].clone().into_array(num_rows)?;
+
+            let result = match left.data_type() {
+                DataType::Utf8 => {
+                    let left = as_generic_string_array::<i32>(&left)?;
+                    let right = as_generic_string_array::<i32>(&right)?;

Review Comment:
   Thanks for catching this regression!
   
   I have resolved the `[P2]` mixed string offset types finding:
   
   1. **Restored dual generic offsets:** Updated `levenshtein` and 
`levenshtein_with_threshold` to accept `left: &GenericStringArray<L>` and 
`right: &GenericStringArray<R>` parameterized over independent `L: 
OffsetSizeTrait` and `R: OffsetSizeTrait`.
   2. **Matrix dispatch in `spark_levenshtein`:** Replaced the single-type 
match with `match (left.data_type(), right.data_type())`, covering all 4 
combinations (`(Utf8, Utf8)`, `(Utf8, LargeUtf8)`, `(LargeUtf8, Utf8)`, and 
`(LargeUtf8, LargeUtf8)`) for both the 2-argument and 3-argument (thresholded) 
forms.
   3. **Regression tests:** Added `test_spark_levenshtein_mixed_offset_types` 
covering `Utf8` literals with `LargeUtf8` arrays, `LargeUtf8` arrays with 
`Utf8` arrays in both argument orders, and mixed-offset execution with 
threshold.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to