Skip to main content

datafusion_functions/regex/
regexpcount.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use crate::regex::{compile_and_cache_regex, compile_regex, start_to_byte_offset};
19use arrow::array::{Array, ArrayRef, AsArray, Datum, Int64Array, StringArrayType};
20use arrow::datatypes::{DataType, Int64Type};
21use arrow::datatypes::{
22    DataType::Int64, DataType::LargeUtf8, DataType::Utf8, DataType::Utf8View,
23};
24use arrow::error::ArrowError;
25use datafusion_common::{Result, ScalarValue, exec_err, internal_err};
26use datafusion_expr::{
27    ColumnarValue, Documentation, ScalarFunctionArgs, ScalarUDFImpl, Signature,
28    TypeSignature::Exact, TypeSignature::Uniform, Volatility,
29};
30use datafusion_macros::user_doc;
31use itertools::izip;
32use regex::Regex;
33use std::collections::HashMap;
34use std::sync::Arc;
35
36#[user_doc(
37    doc_section(label = "Regular Expression Functions"),
38    description = "Returns the number of matches that a [regular expression](https://docs.rs/regex/latest/regex/#syntax) has in a string.",
39    syntax_example = "regexp_count(str, regexp[, start[, flags]])",
40    sql_example = r#"```sql
41> select regexp_count('abcAbAbc', 'abc', 2, 'i');
42+---------------------------------------------------------------+
43| regexp_count(Utf8("abcAbAbc"),Utf8("abc"),Int64(2),Utf8("i")) |
44+---------------------------------------------------------------+
45| 1                                                             |
46+---------------------------------------------------------------+
47```"#,
48    standard_argument(name = "str", prefix = "String"),
49    standard_argument(name = "regexp", prefix = "Regular"),
50    argument(
51        name = "start",
52        description = "Optional start position (the first position is 1) to search for the regular expression. Can be a constant, column, or function."
53    ),
54    argument(
55        name = "flags",
56        description = r#"Optional regular expression flags that control the behavior of the regular expression. Refer to the flags reference above for supported flags."#
57    )
58)]
59#[derive(Debug, PartialEq, Eq, Hash)]
60pub struct RegexpCountFunc {
61    signature: Signature,
62}
63
64impl Default for RegexpCountFunc {
65    fn default() -> Self {
66        Self::new()
67    }
68}
69
70impl RegexpCountFunc {
71    pub fn new() -> Self {
72        Self {
73            signature: Signature::one_of(
74                vec![
75                    Uniform(2, vec![Utf8View, LargeUtf8, Utf8]),
76                    Exact(vec![Utf8View, Utf8View, Int64]),
77                    Exact(vec![LargeUtf8, LargeUtf8, Int64]),
78                    Exact(vec![Utf8, Utf8, Int64]),
79                    Exact(vec![Utf8View, Utf8View, Int64, Utf8View]),
80                    Exact(vec![LargeUtf8, LargeUtf8, Int64, LargeUtf8]),
81                    Exact(vec![Utf8, Utf8, Int64, Utf8]),
82                ],
83                Volatility::Immutable,
84            ),
85        }
86    }
87}
88
89impl ScalarUDFImpl for RegexpCountFunc {
90    fn name(&self) -> &str {
91        "regexp_count"
92    }
93
94    fn signature(&self) -> &Signature {
95        &self.signature
96    }
97
98    fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
99        Ok(Int64)
100    }
101
102    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
103        let args = &args.args;
104
105        let len = args
106            .iter()
107            .fold(Option::<usize>::None, |acc, arg| match arg {
108                ColumnarValue::Scalar(_) => acc,
109                ColumnarValue::Array(a) => Some(a.len()),
110            });
111
112        let is_scalar = len.is_none();
113        let inferred_length = len.unwrap_or(1);
114        let args = args
115            .iter()
116            .map(|arg| arg.to_array(inferred_length))
117            .collect::<Result<Vec<_>>>()?;
118
119        let result = regexp_count_func(&args);
120        if is_scalar {
121            // If all inputs are scalar, keeps output as scalar
122            let result = result.and_then(|arr| ScalarValue::try_from_array(&arr, 0));
123            result.map(ColumnarValue::Scalar)
124        } else {
125            result.map(ColumnarValue::Array)
126        }
127    }
128
129    fn documentation(&self) -> Option<&Documentation> {
130        self.doc()
131    }
132}
133
134pub fn regexp_count_func(args: &[ArrayRef]) -> Result<ArrayRef> {
135    let args_len = args.len();
136    if !(2..=4).contains(&args_len) {
137        return exec_err!(
138            "regexp_count was called with {args_len} arguments. It requires at least 2 and at most 4."
139        );
140    }
141
142    let values = &args[0];
143    match values.data_type() {
144        Utf8 | LargeUtf8 | Utf8View => (),
145        other => {
146            return internal_err!(
147                "Unsupported data type {other:?} for function regexp_count"
148            );
149        }
150    }
151
152    regexp_count(
153        values,
154        &args[1],
155        if args_len > 2 { Some(&args[2]) } else { None },
156        if args_len > 3 { Some(&args[3]) } else { None },
157    )
158    .map_err(|e| e.into())
159}
160
161/// `arrow-rs` style implementation of `regexp_count` function.
162/// This function `regexp_count` is responsible for counting the occurrences of a regular expression pattern
163/// within a string array. It supports optional start positions and flags for case insensitivity.
164///
165/// The function accepts a variable number of arguments:
166/// - `values`: The array of strings to search within.
167/// - `regex_array`: The array of regular expression patterns to search for.
168/// - `start_array` (optional): The array of start positions for the search.
169/// - `flags_array` (optional): The array of flags to modify the search behavior (e.g., case insensitivity).
170///
171/// The function handles different combinations of scalar and array inputs for the regex patterns, start positions,
172/// and flags. It uses a cache to store compiled regular expressions for efficiency.
173///
174/// # Errors
175/// Returns an error if the input arrays have mismatched lengths or if the regular expression fails to compile.
176fn regexp_count(
177    values: &dyn Array,
178    regex_array: &dyn Datum,
179    start_array: Option<&dyn Datum>,
180    flags_array: Option<&dyn Datum>,
181) -> Result<ArrayRef, ArrowError> {
182    let (regex_array, is_regex_scalar) = regex_array.get();
183    let (start_array, is_start_scalar) = start_array.map_or((None, true), |start| {
184        let (start, is_start_scalar) = start.get();
185        (Some(start), is_start_scalar)
186    });
187    let (flags_array, is_flags_scalar) = flags_array.map_or((None, true), |flags| {
188        let (flags, is_flags_scalar) = flags.get();
189        (Some(flags), is_flags_scalar)
190    });
191
192    match (values.data_type(), regex_array.data_type(), flags_array) {
193        (Utf8, Utf8, None) => regexp_count_inner(
194            &values.as_string::<i32>(),
195            &regex_array.as_string::<i32>(),
196            is_regex_scalar,
197            start_array.map(|start| start.as_primitive::<Int64Type>()),
198            is_start_scalar,
199            None,
200            is_flags_scalar,
201        ),
202        (Utf8, Utf8, Some(flags_array)) if *flags_array.data_type() == Utf8 => regexp_count_inner(
203            &values.as_string::<i32>(),
204            &regex_array.as_string::<i32>(),
205            is_regex_scalar,
206            start_array.map(|start| start.as_primitive::<Int64Type>()),
207            is_start_scalar,
208            Some(&flags_array.as_string::<i32>()),
209            is_flags_scalar,
210        ),
211        (LargeUtf8, LargeUtf8, None) => regexp_count_inner(
212            &values.as_string::<i64>(),
213            &regex_array.as_string::<i64>(),
214            is_regex_scalar,
215            start_array.map(|start| start.as_primitive::<Int64Type>()),
216            is_start_scalar,
217            None,
218            is_flags_scalar,
219        ),
220        (LargeUtf8, LargeUtf8, Some(flags_array)) if *flags_array.data_type() == LargeUtf8 => regexp_count_inner(
221            &values.as_string::<i64>(),
222            &regex_array.as_string::<i64>(),
223            is_regex_scalar,
224            start_array.map(|start| start.as_primitive::<Int64Type>()),
225            is_start_scalar,
226            Some(&flags_array.as_string::<i64>()),
227            is_flags_scalar,
228        ),
229        (Utf8View, Utf8View, None) => regexp_count_inner(
230            &values.as_string_view(),
231            &regex_array.as_string_view(),
232            is_regex_scalar,
233            start_array.map(|start| start.as_primitive::<Int64Type>()),
234            is_start_scalar,
235            None,
236            is_flags_scalar,
237        ),
238        (Utf8View, Utf8View, Some(flags_array)) if *flags_array.data_type() == Utf8View => regexp_count_inner(
239            &values.as_string_view(),
240            &regex_array.as_string_view(),
241            is_regex_scalar,
242            start_array.map(|start| start.as_primitive::<Int64Type>()),
243            is_start_scalar,
244            Some(&flags_array.as_string_view()),
245            is_flags_scalar,
246        ),
247        _ => Err(ArrowError::ComputeError(
248            "regexp_count() expected the input arrays to be of type Utf8, LargeUtf8, or Utf8View and the data types of the values, regex_array, and flags_array to match".to_string(),
249        )),
250    }
251}
252
253fn regexp_count_inner<'a, S>(
254    values: &S,
255    regex_array: &S,
256    is_regex_scalar: bool,
257    start_array: Option<&Int64Array>,
258    is_start_scalar: bool,
259    flags_array: Option<&S>,
260    is_flags_scalar: bool,
261) -> Result<ArrayRef, ArrowError>
262where
263    S: StringArrayType<'a>,
264{
265    // Treat single-element arrays as scalars, broadcast to every row. An
266    // absent optional argument behaves like a scalar set to its default.
267    let is_regex_scalar = is_regex_scalar || regex_array.len() == 1;
268    let is_start_scalar =
269        start_array.is_none_or(|array| is_start_scalar || array.len() == 1);
270    let is_flags_scalar =
271        flags_array.is_none_or(|array| is_flags_scalar || array.len() == 1);
272
273    // A NULL in any scalar argument produces a NULL result for every row
274    if (is_regex_scalar && regex_array.is_null(0))
275        || (is_start_scalar && start_array.is_some_and(|array| array.is_null(0)))
276        || (is_flags_scalar && flags_array.is_some_and(|array| array.is_null(0)))
277    {
278        return Ok(Arc::new(Int64Array::new_null(values.len())));
279    }
280
281    let regex_scalar = is_regex_scalar.then(|| regex_array.value(0));
282    // An absent `start` defaults to 1
283    let start_scalar =
284        is_start_scalar.then(|| start_array.map_or(1, |array| array.value(0)));
285    // A `flags_scalar` of None means no flags were supplied
286    let flags_scalar = if is_flags_scalar {
287        flags_array.map(|array| array.value(0))
288    } else {
289        None
290    };
291
292    let mut regex_cache = HashMap::new();
293
294    match (regex_scalar, is_start_scalar, is_flags_scalar) {
295        (Some(regex), true, true) => {
296            let pattern = compile_regex(regex, flags_scalar)?;
297
298            Ok(Arc::new(
299                values
300                    .iter()
301                    .map(|value| count_matches(value, &pattern, start_scalar))
302                    .collect::<Result<Int64Array, ArrowError>>()?,
303            ))
304        }
305        (Some(regex), true, false) => {
306            let flags_array = flags_array.unwrap();
307            if values.len() != flags_array.len() {
308                return Err(ArrowError::ComputeError(format!(
309                    "flags_array must be the same length as values array; got {} and {}",
310                    flags_array.len(),
311                    values.len(),
312                )));
313            }
314
315            Ok(Arc::new(
316                values
317                    .iter()
318                    .zip(flags_array.iter())
319                    .map(|(value, flags)| {
320                        let Some(flags) = flags else {
321                            return Ok(None);
322                        };
323
324                        let pattern = compile_and_cache_regex(
325                            regex,
326                            Some(flags),
327                            &mut regex_cache,
328                        )?;
329                        count_matches(value, pattern, start_scalar)
330                    })
331                    .collect::<Result<Int64Array, ArrowError>>()?,
332            ))
333        }
334        (Some(regex), false, true) => {
335            let pattern = compile_regex(regex, flags_scalar)?;
336
337            let start_array = start_array.unwrap();
338
339            Ok(Arc::new(
340                values
341                    .iter()
342                    .zip(start_array.iter())
343                    .map(|(value, start)| count_matches(value, &pattern, start))
344                    .collect::<Result<Int64Array, ArrowError>>()?,
345            ))
346        }
347        (Some(regex), false, false) => {
348            let flags_array = flags_array.unwrap();
349            if values.len() != flags_array.len() {
350                return Err(ArrowError::ComputeError(format!(
351                    "flags_array must be the same length as values array; got {} and {}",
352                    flags_array.len(),
353                    values.len(),
354                )));
355            }
356
357            Ok(Arc::new(
358                izip!(
359                    values.iter(),
360                    start_array.unwrap().iter(),
361                    flags_array.iter()
362                )
363                .map(|(value, start, flags)| {
364                    let Some(flags) = flags else {
365                        return Ok(None);
366                    };
367
368                    let pattern =
369                        compile_and_cache_regex(regex, Some(flags), &mut regex_cache)?;
370
371                    count_matches(value, pattern, start)
372                })
373                .collect::<Result<Int64Array, ArrowError>>()?,
374            ))
375        }
376        (None, true, true) => {
377            if values.len() != regex_array.len() {
378                return Err(ArrowError::ComputeError(format!(
379                    "regex_array must be the same length as values array; got {} and {}",
380                    regex_array.len(),
381                    values.len(),
382                )));
383            }
384
385            Ok(Arc::new(
386                values
387                    .iter()
388                    .zip(regex_array.iter())
389                    .map(|(value, regex)| {
390                        let Some(regex) = regex else {
391                            return Ok(None);
392                        };
393
394                        let pattern = compile_and_cache_regex(
395                            regex,
396                            flags_scalar,
397                            &mut regex_cache,
398                        )?;
399                        count_matches(value, pattern, start_scalar)
400                    })
401                    .collect::<Result<Int64Array, ArrowError>>()?,
402            ))
403        }
404        (None, true, false) => {
405            if values.len() != regex_array.len() {
406                return Err(ArrowError::ComputeError(format!(
407                    "regex_array must be the same length as values array; got {} and {}",
408                    regex_array.len(),
409                    values.len(),
410                )));
411            }
412
413            let flags_array = flags_array.unwrap();
414            if values.len() != flags_array.len() {
415                return Err(ArrowError::ComputeError(format!(
416                    "flags_array must be the same length as values array; got {} and {}",
417                    flags_array.len(),
418                    values.len(),
419                )));
420            }
421
422            Ok(Arc::new(
423                izip!(values.iter(), regex_array.iter(), flags_array.iter())
424                    .map(|(value, regex, flags)| {
425                        let (Some(regex), Some(flags)) = (regex, flags) else {
426                            return Ok(None);
427                        };
428
429                        let pattern = compile_and_cache_regex(
430                            regex,
431                            Some(flags),
432                            &mut regex_cache,
433                        )?;
434
435                        count_matches(value, pattern, start_scalar)
436                    })
437                    .collect::<Result<Int64Array, ArrowError>>()?,
438            ))
439        }
440        (None, false, true) => {
441            if values.len() != regex_array.len() {
442                return Err(ArrowError::ComputeError(format!(
443                    "regex_array must be the same length as values array; got {} and {}",
444                    regex_array.len(),
445                    values.len(),
446                )));
447            }
448
449            let start_array = start_array.unwrap();
450            if values.len() != start_array.len() {
451                return Err(ArrowError::ComputeError(format!(
452                    "start_array must be the same length as values array; got {} and {}",
453                    start_array.len(),
454                    values.len(),
455                )));
456            }
457
458            Ok(Arc::new(
459                izip!(values.iter(), regex_array.iter(), start_array.iter())
460                    .map(|(value, regex, start)| {
461                        let Some(regex) = regex else {
462                            return Ok(None);
463                        };
464
465                        let pattern = compile_and_cache_regex(
466                            regex,
467                            flags_scalar,
468                            &mut regex_cache,
469                        )?;
470                        count_matches(value, pattern, start)
471                    })
472                    .collect::<Result<Int64Array, ArrowError>>()?,
473            ))
474        }
475        (None, false, false) => {
476            if values.len() != regex_array.len() {
477                return Err(ArrowError::ComputeError(format!(
478                    "regex_array must be the same length as values array; got {} and {}",
479                    regex_array.len(),
480                    values.len(),
481                )));
482            }
483
484            let start_array = start_array.unwrap();
485            if values.len() != start_array.len() {
486                return Err(ArrowError::ComputeError(format!(
487                    "start_array must be the same length as values array; got {} and {}",
488                    start_array.len(),
489                    values.len(),
490                )));
491            }
492
493            let flags_array = flags_array.unwrap();
494            if values.len() != flags_array.len() {
495                return Err(ArrowError::ComputeError(format!(
496                    "flags_array must be the same length as values array; got {} and {}",
497                    flags_array.len(),
498                    values.len(),
499                )));
500            }
501
502            Ok(Arc::new(
503                izip!(
504                    values.iter(),
505                    regex_array.iter(),
506                    start_array.iter(),
507                    flags_array.iter()
508                )
509                .map(|(value, regex, start, flags)| {
510                    let (Some(regex), Some(flags)) = (regex, flags) else {
511                        return Ok(None);
512                    };
513
514                    let pattern =
515                        compile_and_cache_regex(regex, Some(flags), &mut regex_cache)?;
516                    count_matches(value, pattern, start)
517                })
518                .collect::<Result<Int64Array, ArrowError>>()?,
519            ))
520        }
521    }
522}
523
524fn count_matches(
525    value: Option<&str>,
526    pattern: &Regex,
527    start: Option<i64>,
528) -> Result<Option<i64>, ArrowError> {
529    // A NULL value or start position produces a NULL result.
530    let (Some(value), Some(start)) = (value, start) else {
531        return Ok(None);
532    };
533
534    if start < 1 {
535        return Err(ArrowError::ComputeError(
536            "regexp_count() requires start to be 1 based".to_string(),
537        ));
538    }
539
540    let Some(byte_offset) = start_to_byte_offset(value, start) else {
541        return Ok(Some(0));
542    };
543    let count = pattern.find_iter(&value[byte_offset..]).count();
544    Ok(Some(count as i64))
545}
546
547#[cfg(test)]
548mod tests {
549    use super::*;
550    use arrow::array::{GenericStringArray, StringViewArray};
551    use arrow::datatypes::Field;
552    use datafusion_common::config::ConfigOptions;
553
554    #[test]
555    fn test_regexp_count() {
556        test_case_sensitive_regexp_count_scalar();
557        test_case_sensitive_regexp_count_empty_pattern_scalar();
558        test_case_sensitive_regexp_count_scalar_start();
559        test_case_insensitive_regexp_count_scalar_flags();
560        test_case_sensitive_regexp_count_start_scalar_complex();
561
562        test_case_sensitive_regexp_count_array::<GenericStringArray<i32>>();
563        test_case_sensitive_regexp_count_array::<GenericStringArray<i64>>();
564        test_case_sensitive_regexp_count_array::<StringViewArray>();
565
566        test_case_sensitive_regexp_count_array_start::<GenericStringArray<i32>>();
567        test_case_sensitive_regexp_count_array_start::<GenericStringArray<i64>>();
568        test_case_sensitive_regexp_count_array_start::<StringViewArray>();
569
570        test_case_insensitive_regexp_count_array_flags::<GenericStringArray<i32>>();
571        test_case_insensitive_regexp_count_array_flags::<GenericStringArray<i64>>();
572        test_case_insensitive_regexp_count_array_flags::<StringViewArray>();
573
574        test_case_sensitive_regexp_count_array_complex::<GenericStringArray<i32>>();
575        test_case_sensitive_regexp_count_array_complex::<GenericStringArray<i64>>();
576        test_case_sensitive_regexp_count_array_complex::<StringViewArray>();
577
578        test_case_regexp_count_cache_check::<GenericStringArray<i32>>();
579
580        test_regexp_count_null_scalars();
581
582        test_regexp_count_null_array_rows::<GenericStringArray<i32>>();
583        test_regexp_count_null_array_rows::<GenericStringArray<i64>>();
584        test_regexp_count_null_array_rows::<StringViewArray>();
585
586        test_regexp_count_null_start_array::<GenericStringArray<i32>>();
587        test_regexp_count_null_start_array::<GenericStringArray<i64>>();
588        test_regexp_count_null_start_array::<StringViewArray>();
589
590        test_regexp_count_null_flags_array::<GenericStringArray<i32>>();
591        test_regexp_count_null_flags_array::<GenericStringArray<i64>>();
592        test_regexp_count_null_flags_array::<StringViewArray>();
593
594        test_regexp_count_null_scalar_regex_array_values::<GenericStringArray<i32>>();
595        test_regexp_count_null_scalar_regex_array_values::<GenericStringArray<i64>>();
596        test_regexp_count_null_scalar_regex_array_values::<StringViewArray>();
597    }
598
599    fn regexp_count_with_scalar_values(args: &[ScalarValue]) -> Result<ColumnarValue> {
600        let args_values = args
601            .iter()
602            .map(|sv| ColumnarValue::Scalar(sv.clone()))
603            .collect();
604
605        let arg_fields = args
606            .iter()
607            .enumerate()
608            .map(|(idx, a)| Field::new(format!("arg_{idx}"), a.data_type(), true).into())
609            .collect::<Vec<_>>();
610
611        RegexpCountFunc::new().invoke_with_args(ScalarFunctionArgs {
612            args: args_values,
613            arg_fields,
614            number_rows: args.len(),
615            return_field: Field::new("f", Int64, true).into(),
616            config_options: Arc::new(ConfigOptions::default()),
617        })
618    }
619
620    fn test_case_sensitive_regexp_count_scalar() {
621        let values = ["", "aabca", "abcabc", "abcAbcab", "abcabcabc"];
622        let regex = "abc";
623        let expected: Vec<i64> = vec![0, 1, 2, 1, 3];
624
625        values.iter().enumerate().for_each(|(pos, &v)| {
626            // utf8
627            let v_sv = ScalarValue::Utf8(Some(v.to_string()));
628            let regex_sv = ScalarValue::Utf8(Some(regex.to_string()));
629            let expected = expected.get(pos).cloned();
630            let re = regexp_count_with_scalar_values(&[v_sv, regex_sv]);
631            match re {
632                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
633                    assert_eq!(v, expected, "regexp_count scalar test failed");
634                }
635                _ => panic!("Unexpected result"),
636            }
637
638            // largeutf8
639            let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
640            let regex_sv = ScalarValue::LargeUtf8(Some(regex.to_string()));
641            let re = regexp_count_with_scalar_values(&[v_sv, regex_sv]);
642            match re {
643                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
644                    assert_eq!(v, expected, "regexp_count scalar test failed");
645                }
646                _ => panic!("Unexpected result"),
647            }
648
649            // utf8view
650            let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
651            let regex_sv = ScalarValue::Utf8View(Some(regex.to_string()));
652            let re = regexp_count_with_scalar_values(&[v_sv, regex_sv]);
653            match re {
654                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
655                    assert_eq!(v, expected, "regexp_count scalar test failed");
656                }
657                _ => panic!("Unexpected result"),
658            }
659        });
660    }
661
662    fn test_case_sensitive_regexp_count_empty_pattern_scalar() {
663        let values = ["", "abc", "abc"];
664        let start_positions = [1, 1, 2];
665        let expected: Vec<i64> = vec![1, 4, 3];
666
667        values
668            .iter()
669            .zip(start_positions.iter())
670            .enumerate()
671            .for_each(|(pos, (&value, &start))| {
672                let expected = expected.get(pos).cloned();
673                let start_sv = ScalarValue::Int64(Some(start));
674
675                let re = regexp_count_with_scalar_values(&[
676                    ScalarValue::Utf8(Some(value.to_string())),
677                    ScalarValue::Utf8(Some("".to_string())),
678                    start_sv.clone(),
679                ]);
680                match re {
681                    Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
682                        assert_eq!(v, expected, "regexp_count scalar test failed");
683                    }
684                    _ => panic!("Unexpected result"),
685                }
686
687                let re = regexp_count_with_scalar_values(&[
688                    ScalarValue::LargeUtf8(Some(value.to_string())),
689                    ScalarValue::LargeUtf8(Some("".to_string())),
690                    start_sv.clone(),
691                ]);
692                match re {
693                    Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
694                        assert_eq!(v, expected, "regexp_count scalar test failed");
695                    }
696                    _ => panic!("Unexpected result"),
697                }
698
699                let re = regexp_count_with_scalar_values(&[
700                    ScalarValue::Utf8View(Some(value.to_string())),
701                    ScalarValue::Utf8View(Some("".to_string())),
702                    start_sv,
703                ]);
704                match re {
705                    Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
706                        assert_eq!(v, expected, "regexp_count scalar test failed");
707                    }
708                    _ => panic!("Unexpected result"),
709                }
710            });
711    }
712
713    fn test_case_sensitive_regexp_count_scalar_start() {
714        let values = ["", "aabca", "abcabc", "abcAbcab", "abcabcabc"];
715        let regex = "abc";
716        let start = 2;
717        let expected: Vec<i64> = vec![0, 1, 1, 0, 2];
718
719        values.iter().enumerate().for_each(|(pos, &v)| {
720            // utf8
721            let v_sv = ScalarValue::Utf8(Some(v.to_string()));
722            let regex_sv = ScalarValue::Utf8(Some(regex.to_string()));
723            let start_sv = ScalarValue::Int64(Some(start));
724            let expected = expected.get(pos).cloned();
725            let re = regexp_count_with_scalar_values(&[v_sv, regex_sv, start_sv.clone()]);
726            match re {
727                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
728                    assert_eq!(v, expected, "regexp_count scalar test failed");
729                }
730                _ => panic!("Unexpected result"),
731            }
732
733            // largeutf8
734            let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
735            let regex_sv = ScalarValue::LargeUtf8(Some(regex.to_string()));
736            let re = regexp_count_with_scalar_values(&[v_sv, regex_sv, start_sv.clone()]);
737            match re {
738                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
739                    assert_eq!(v, expected, "regexp_count scalar test failed");
740                }
741                _ => panic!("Unexpected result"),
742            }
743
744            // utf8view
745            let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
746            let regex_sv = ScalarValue::Utf8View(Some(regex.to_string()));
747            let re = regexp_count_with_scalar_values(&[v_sv, regex_sv, start_sv.clone()]);
748            match re {
749                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
750                    assert_eq!(v, expected, "regexp_count scalar test failed");
751                }
752                _ => panic!("Unexpected result"),
753            }
754        });
755    }
756
757    fn test_case_insensitive_regexp_count_scalar_flags() {
758        let values = ["", "aabca", "abcabc", "abcAbcab", "abcabcabc"];
759        let regex = "abc";
760        let start = 1;
761        let flags = "i";
762        let expected: Vec<i64> = vec![0, 1, 2, 2, 3];
763
764        values.iter().enumerate().for_each(|(pos, &v)| {
765            // utf8
766            let v_sv = ScalarValue::Utf8(Some(v.to_string()));
767            let regex_sv = ScalarValue::Utf8(Some(regex.to_string()));
768            let start_sv = ScalarValue::Int64(Some(start));
769            let flags_sv = ScalarValue::Utf8(Some(flags.to_string()));
770            let expected = expected.get(pos).cloned();
771
772            let re = regexp_count_with_scalar_values(&[
773                v_sv,
774                regex_sv,
775                start_sv.clone(),
776                flags_sv.clone(),
777            ]);
778            match re {
779                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
780                    assert_eq!(v, expected, "regexp_count scalar test failed");
781                }
782                _ => panic!("Unexpected result"),
783            }
784
785            // largeutf8
786            let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
787            let regex_sv = ScalarValue::LargeUtf8(Some(regex.to_string()));
788            let flags_sv = ScalarValue::LargeUtf8(Some(flags.to_string()));
789
790            let re = regexp_count_with_scalar_values(&[
791                v_sv,
792                regex_sv,
793                start_sv.clone(),
794                flags_sv.clone(),
795            ]);
796            match re {
797                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
798                    assert_eq!(v, expected, "regexp_count scalar test failed");
799                }
800                _ => panic!("Unexpected result"),
801            }
802
803            // utf8view
804            let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
805            let regex_sv = ScalarValue::Utf8View(Some(regex.to_string()));
806            let flags_sv = ScalarValue::Utf8View(Some(flags.to_string()));
807
808            let re = regexp_count_with_scalar_values(&[
809                v_sv,
810                regex_sv,
811                start_sv.clone(),
812                flags_sv.clone(),
813            ]);
814            match re {
815                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
816                    assert_eq!(v, expected, "regexp_count scalar test failed");
817                }
818                _ => panic!("Unexpected result"),
819            }
820        });
821    }
822
823    fn test_case_sensitive_regexp_count_array<A>()
824    where
825        A: From<Vec<&'static str>> + Array + 'static,
826    {
827        let values = A::from(vec!["", "aabca", "abcabc", "abcAbcab", "abcabcAbc"]);
828        let regex = A::from(vec!["", "abc", "a", "bc", "ab"]);
829
830        let expected = Int64Array::from(vec![1, 1, 2, 2, 2]);
831
832        let re = regexp_count_func(&[Arc::new(values), Arc::new(regex)]).unwrap();
833        assert_eq!(re.as_ref(), &expected);
834    }
835
836    fn test_case_sensitive_regexp_count_array_start<A>()
837    where
838        A: From<Vec<&'static str>> + Array + 'static,
839    {
840        let values = A::from(vec!["", "aAbca", "abcabc", "abcAbcab", "abcabcAbc"]);
841        let regex = A::from(vec!["", "abc", "a", "bc", "ab"]);
842        let start = Int64Array::from(vec![1, 2, 3, 4, 5]);
843
844        let expected = Int64Array::from(vec![1, 0, 1, 1, 0]);
845
846        let re = regexp_count_func(&[Arc::new(values), Arc::new(regex), Arc::new(start)])
847            .unwrap();
848        assert_eq!(re.as_ref(), &expected);
849    }
850
851    fn test_case_insensitive_regexp_count_array_flags<A>()
852    where
853        A: From<Vec<&'static str>> + Array + 'static,
854    {
855        let values = A::from(vec!["", "aAbca", "abcabc", "abcAbcab", "abcabcAbc"]);
856        let regex = A::from(vec!["", "abc", "a", "bc", "ab"]);
857        let start = Int64Array::from(vec![1]);
858        let flags = A::from(vec!["", "i", "", "", "i"]);
859
860        let expected = Int64Array::from(vec![1, 1, 2, 2, 3]);
861
862        let re = regexp_count_func(&[
863            Arc::new(values),
864            Arc::new(regex),
865            Arc::new(start),
866            Arc::new(flags),
867        ])
868        .unwrap();
869        assert_eq!(re.as_ref(), &expected);
870    }
871
872    fn test_case_sensitive_regexp_count_start_scalar_complex() {
873        let values = ["", "aabca", "abcabc", "abcAbcab", "abcabcabc"];
874        let regex = ["", "abc", "a", "bc", "ab"];
875        let start = 5;
876        let flags = ["", "i", "", "", "i"];
877        let expected: Vec<i64> = vec![0, 0, 0, 1, 1];
878
879        values.iter().enumerate().for_each(|(pos, &v)| {
880            // utf8
881            let v_sv = ScalarValue::Utf8(Some(v.to_string()));
882            let regex_sv = ScalarValue::Utf8(regex.get(pos).map(|s| (*s).to_string()));
883            let start_sv = ScalarValue::Int64(Some(start));
884            let flags_sv = ScalarValue::Utf8(flags.get(pos).map(|f| (*f).to_string()));
885            let expected = expected.get(pos).cloned();
886            let re = regexp_count_with_scalar_values(&[
887                v_sv,
888                regex_sv,
889                start_sv.clone(),
890                flags_sv.clone(),
891            ]);
892            match re {
893                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
894                    assert_eq!(v, expected, "regexp_count scalar test failed");
895                }
896                _ => panic!("Unexpected result"),
897            }
898
899            // largeutf8
900            let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
901            let regex_sv =
902                ScalarValue::LargeUtf8(regex.get(pos).map(|s| (*s).to_string()));
903            let flags_sv =
904                ScalarValue::LargeUtf8(flags.get(pos).map(|f| (*f).to_string()));
905            let re = regexp_count_with_scalar_values(&[
906                v_sv,
907                regex_sv,
908                start_sv.clone(),
909                flags_sv.clone(),
910            ]);
911            match re {
912                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
913                    assert_eq!(v, expected, "regexp_count scalar test failed");
914                }
915                _ => panic!("Unexpected result"),
916            }
917
918            // utf8view
919            let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
920            let regex_sv =
921                ScalarValue::Utf8View(regex.get(pos).map(|s| (*s).to_string()));
922            let flags_sv =
923                ScalarValue::Utf8View(flags.get(pos).map(|f| (*f).to_string()));
924            let re = regexp_count_with_scalar_values(&[
925                v_sv,
926                regex_sv,
927                start_sv.clone(),
928                flags_sv.clone(),
929            ]);
930            match re {
931                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
932                    assert_eq!(v, expected, "regexp_count scalar test failed");
933                }
934                _ => panic!("Unexpected result"),
935            }
936        });
937    }
938
939    fn test_case_sensitive_regexp_count_array_complex<A>()
940    where
941        A: From<Vec<&'static str>> + Array + 'static,
942    {
943        let values = A::from(vec!["", "aAbca", "abcabc", "abcAbcab", "abcabcAbc"]);
944        let regex = A::from(vec!["", "abc", "a", "bc", "ab"]);
945        let start = Int64Array::from(vec![1, 2, 3, 4, 5]);
946        let flags = A::from(vec!["", "i", "", "", "i"]);
947
948        let expected = Int64Array::from(vec![1, 1, 1, 1, 1]);
949
950        let re = regexp_count_func(&[
951            Arc::new(values),
952            Arc::new(regex),
953            Arc::new(start),
954            Arc::new(flags),
955        ])
956        .unwrap();
957        assert_eq!(re.as_ref(), &expected);
958    }
959
960    fn test_regexp_count_null_scalars() {
961        // A NULL in any scalar argument produces a NULL result.
962        let cases: Vec<Vec<ScalarValue>> = vec![
963            vec![ScalarValue::Utf8(None), ScalarValue::Utf8(None)],
964            vec![
965                ScalarValue::Utf8(None),
966                ScalarValue::Utf8(Some("abc".to_string())),
967                ScalarValue::Int64(Some(1)),
968                ScalarValue::Utf8(Some("i".to_string())),
969            ],
970            vec![
971                ScalarValue::Utf8(Some("abc".to_string())),
972                ScalarValue::Utf8(None),
973                ScalarValue::Int64(Some(1)),
974                ScalarValue::Utf8(Some("i".to_string())),
975            ],
976            vec![
977                ScalarValue::Utf8(Some("abc".to_string())),
978                ScalarValue::Utf8(Some("abc".to_string())),
979                ScalarValue::Int64(None),
980                ScalarValue::Utf8(Some("i".to_string())),
981            ],
982            vec![
983                ScalarValue::Utf8(Some("abc".to_string())),
984                ScalarValue::Utf8(Some("abc".to_string())),
985                ScalarValue::Int64(Some(1)),
986                ScalarValue::Utf8(None),
987            ],
988        ];
989
990        for args in cases {
991            let re = regexp_count_with_scalar_values(&args);
992            match re {
993                Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
994                    assert_eq!(v, None, "regexp_count null scalar test failed");
995                }
996                _ => panic!("Unexpected result"),
997            }
998        }
999    }
1000
1001    fn test_regexp_count_null_array_rows<A>()
1002    where
1003        A: From<Vec<Option<&'static str>>> + Array + 'static,
1004    {
1005        let values = A::from(vec![
1006            None,
1007            Some("abc"),
1008            Some("abc"),
1009            Some("abc"),
1010            Some("abc"),
1011        ]);
1012        let regex = A::from(vec![
1013            Some("abc"),
1014            None,
1015            Some("abc"),
1016            Some("abc"),
1017            Some("abc"),
1018        ]);
1019        let start = Int64Array::from(vec![Some(1), Some(1), None, Some(1), Some(1)]);
1020        let flags = A::from(vec![Some("i"), Some("i"), Some("i"), None, Some("i")]);
1021
1022        let expected = Int64Array::from(vec![None, None, None, None, Some(1)]);
1023
1024        let re = regexp_count_func(&[
1025            Arc::new(values),
1026            Arc::new(regex),
1027            Arc::new(start),
1028            Arc::new(flags),
1029        ])
1030        .unwrap();
1031        assert_eq!(re.as_ref(), &expected);
1032    }
1033
1034    fn test_regexp_count_null_start_array<A>()
1035    where
1036        A: From<Vec<&'static str>> + Array + 'static,
1037    {
1038        let values = A::from(vec!["abc", "abcb"]);
1039        let regex = A::from(vec!["b"]);
1040        let start = Int64Array::from(vec![Some(1), None]);
1041
1042        let expected = Int64Array::from(vec![Some(1), None]);
1043
1044        let re = regexp_count_func(&[Arc::new(values), Arc::new(regex), Arc::new(start)])
1045            .unwrap();
1046        assert_eq!(re.as_ref(), &expected);
1047    }
1048
1049    fn test_regexp_count_null_flags_array<A>()
1050    where
1051        A: From<Vec<&'static str>> + From<Vec<Option<&'static str>>> + Array + 'static,
1052    {
1053        let values: A = vec!["aB", "aB"].into();
1054        let regex: A = vec!["b"].into();
1055        let start = Int64Array::from(vec![1]);
1056        let flags: A = vec![None, Some("i")].into();
1057
1058        let expected = Int64Array::from(vec![None, Some(1)]);
1059
1060        let re = regexp_count_func(&[
1061            Arc::new(values),
1062            Arc::new(regex),
1063            Arc::new(start),
1064            Arc::new(flags),
1065        ])
1066        .unwrap();
1067        assert_eq!(re.as_ref(), &expected);
1068    }
1069
1070    fn test_regexp_count_null_scalar_regex_array_values<A>()
1071    where
1072        A: From<Vec<&'static str>> + From<Vec<Option<&'static str>>> + Array + 'static,
1073    {
1074        let values: A = vec!["abc", "abcabc"].into();
1075        let regex: A = vec![Option::<&str>::None].into();
1076
1077        let expected = Int64Array::from(vec![None::<i64>, None]);
1078
1079        let re = regexp_count_func(&[Arc::new(values), Arc::new(regex)]).unwrap();
1080        assert_eq!(re.as_ref(), &expected);
1081    }
1082
1083    fn test_case_regexp_count_cache_check<A>()
1084    where
1085        A: From<Vec<&'static str>> + Array + 'static,
1086    {
1087        let values = A::from(vec!["aaa", "Aaa", "aaa"]);
1088        let regex = A::from(vec!["aaa", "aaa", "aaa"]);
1089        let start = Int64Array::from(vec![1, 1, 1]);
1090        let flags = A::from(vec!["", "i", ""]);
1091
1092        let expected = Int64Array::from(vec![1, 1, 1]);
1093
1094        let re = regexp_count_func(&[
1095            Arc::new(values),
1096            Arc::new(regex),
1097            Arc::new(start),
1098            Arc::new(flags),
1099        ])
1100        .unwrap();
1101        assert_eq!(re.as_ref(), &expected);
1102    }
1103}