sunchao commented on code in PR #5042:
URL: https://github.com/apache/datafusion-comet/pull/5042#discussion_r4117477417


##########
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:
   [P2] Preserve independent string offset types. The new dispatch chooses both 
downcasts from the left argument's type, so a `LargeStringArray` compared with 
an ordinary `Utf8` literal now fails instead of returning a distance. This 
affects both argument orders and both arities, including the analogous 
threshold dispatch. Base supported these combinations through separate 
`LeftOffset` and `RightOffset` parameters. Could you restore that generic 
evaluator and retain mixed-offset regression coverage?
   
   Evidence: Compiled unchanged base/head modules side by side. Calling 
`spark_levenshtein` with `ScalarValue::Utf8(Some("kitten"))` and 
`LargeStringArray::from(vec!["sitting"])` returns `[Some(3)]` at base but an 
internal LargeUtf8-to-GenericStringArray<i32> cast error at head. Reversing 
arguments or adding threshold 3 also fails. Reproduced with and without 
overflow checks. Harness: `/tmp/comet-5042-6858-review/repro.rs`. Results: 
`repro_debug.log` and `repro_release.log` in that directory.



##########
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;

Review Comment:
   [P2] Keep threshold arithmetic overflow-safe. With the default debug build 
used by `make core`, `levenshtein("frog", "fog", 2147483647)` now panics at 
`threshold + 1` instead of returning 1. Threshold 2147483646 also panics later 
at `j_1 + threshold`. Both ASCII and Unicode branches regress from base's 
`usize`/saturating arithmetic. Spark explicitly tests `Integer.MAX_VALUE` 
thresholds. Could you restore safe arithmetic in both branches and add these 
boundary cases? Release builds disable overflow checks and did not reproduce 
this panic.
   
   Evidence: The source-verified public-API harness returns `[Some(1)]` at base 
and catches a panic at head for thresholds 2147483647 and 2147483646, using 
both `"frog"/"fog"` and `"café"/"cafe"`. Recompiled with `-O -C 
overflow-checks=no`, all four return 1. Spark 3.5.9–4.2.0 uses an explicit 
overflow guard for the band endpoint and tests maximum thresholds. Reproduction 
and compilation commands: `/tmp/comet-5042-6858-review/repro.rs` and 
`compile_commands.txt`.



-- 
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