Skip to main content

arrow_string/
regexp.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
18//! Defines kernel to extract substrings based on a regular
19//! expression of a \[Large\]StringArray
20
21use crate::like::StringArrayType;
22
23use arrow_array::builder::{
24    BooleanBufferBuilder, GenericStringBuilder, ListBuilder, StringViewBuilder,
25};
26use arrow_array::cast::AsArray;
27use arrow_array::*;
28use arrow_buffer::{BooleanBuffer, NullBuffer};
29use arrow_data::ArrayDataBuilder;
30use arrow_schema::{ArrowError, DataType, Field};
31use regex::Regex;
32
33use std::collections::HashMap;
34use std::sync::Arc;
35
36/// Return BooleanArray indicating which strings in an array match an array of
37/// regular expressions.
38///
39/// This is equivalent to the SQL `array ~ regex_array`, supporting
40/// [`StringArray`] / [`LargeStringArray`] / [`StringViewArray`].
41///
42/// If `regex_array` element has an empty value, the corresponding result value is always true.
43///
44/// `flags_array` are optional [`StringArray`] / [`LargeStringArray`] / [`StringViewArray`] flag,
45/// which allow special search modes, such as case-insensitive and multi-line mode.
46/// See the documentation [here](https://docs.rs/regex/1.5.4/regex/#grouping-and-flags)
47/// for more information.
48///
49/// # See Also
50/// * [`regexp_is_match_scalar`] for matching a single regular expression against an array of strings
51/// * [`regexp_match`] for extracting groups from a string array based on a regular expression
52///
53/// # Example
54/// ```
55/// # use arrow_array::{StringArray, BooleanArray};
56/// # use arrow_string::regexp::regexp_is_match;
57/// // First array is the array of strings to match
58/// let array = StringArray::from(vec!["Foo", "Bar", "FooBar", "Baz"]);
59/// // Second array is the array of regular expressions to match against
60/// let regex_array = StringArray::from(vec!["^Foo", "^Foo", "Bar$", "Baz"]);
61/// // Third array is the array of flags to use for each regular expression, if desired
62/// // (the type must be provided to satisfy type inference for the third parameter)
63/// let flags_array: Option<&StringArray> = None;
64/// // The result is a BooleanArray indicating when each string in `array`
65/// // matches the corresponding regular expression in `regex_array`
66/// let result = regexp_is_match(&array, &regex_array, flags_array).unwrap();
67/// assert_eq!(result, BooleanArray::from(vec![true, false, true, true]));
68/// ```
69pub fn regexp_is_match<'a, S1, S2, S3>(
70    array: &'a S1,
71    regex_array: &'a S2,
72    flags_array: Option<&'a S3>,
73) -> Result<BooleanArray, ArrowError>
74where
75    &'a S1: StringArrayType<'a>,
76    &'a S2: StringArrayType<'a>,
77    &'a S3: StringArrayType<'a>,
78{
79    if array.len() != regex_array.len() {
80        return Err(ArrowError::ComputeError(
81            "Cannot perform comparison operation on arrays of different length".to_string(),
82        ));
83    }
84
85    let nulls = NullBuffer::union(array.nulls(), regex_array.nulls());
86
87    let mut patterns: HashMap<String, Regex> = HashMap::new();
88    let mut result = BooleanBufferBuilder::new(array.len());
89
90    let complete_pattern = match flags_array {
91        Some(flags) => Box::new(
92            regex_array
93                .iter()
94                .zip(flags.iter())
95                .map(|(pattern, flags)| {
96                    pattern.map(|pattern| match flags {
97                        Some(flag) => format!("(?{flag}){pattern}"),
98                        None => pattern.to_string(),
99                    })
100                }),
101        ) as Box<dyn Iterator<Item = Option<String>>>,
102        None => Box::new(
103            regex_array
104                .iter()
105                .map(|pattern| pattern.map(|pattern| pattern.to_string())),
106        ),
107    };
108
109    array
110        .iter()
111        .zip(complete_pattern)
112        .map(|(value, pattern)| {
113            match (value, pattern) {
114                // Required for Postgres compatibility:
115                // SELECT 'foobarbequebaz' ~ ''); = true
116                (Some(_), Some(pattern)) if pattern == *"" => {
117                    result.append(true);
118                }
119                (Some(value), Some(pattern)) => {
120                    let existing_pattern = patterns.get(&pattern);
121                    let re = match existing_pattern {
122                        Some(re) => re,
123                        None => {
124                            let re = Regex::new(pattern.as_str()).map_err(|e| {
125                                ArrowError::ComputeError(format!(
126                                    "Regular expression did not compile: {e:?}"
127                                ))
128                            })?;
129                            patterns.entry(pattern).or_insert(re)
130                        }
131                    };
132                    result.append(re.is_match(value));
133                }
134                _ => result.append(false),
135            }
136            Ok(())
137        })
138        .collect::<Result<Vec<()>, ArrowError>>()?;
139
140    let data = unsafe {
141        ArrayDataBuilder::new(DataType::Boolean)
142            .len(array.len())
143            .buffers(vec![result.into()])
144            .nulls(nulls)
145            .build_unchecked()
146    };
147
148    Ok(BooleanArray::from(data))
149}
150
151/// Return BooleanArray indicating which strings in an array match a single regular expression.
152///
153/// This is equivalent to the SQL `array ~ regex_array`, supporting
154/// [`StringArray`] / [`LargeStringArray`] / [`StringViewArray`] and a scalar.
155///
156/// See the documentation on [`regexp_is_match`] for more details on arguments
157///
158/// # See Also
159/// * [`regexp_is_match`] for matching an array of regular expression against an array of strings
160/// * [`regexp_match`] for extracting groups from a string array based on a regular expression
161///
162/// # Example
163/// ```
164/// # use arrow_array::{StringArray, BooleanArray};
165/// # use arrow_string::regexp::regexp_is_match_scalar;
166/// // array of strings to match
167/// let array = StringArray::from(vec!["Foo", "Bar", "FooBar", "Baz"]);
168/// let regexp = "^Foo"; // regular expression to match against
169/// let flags: Option<&str> = None;  // flags can control the matching behavior
170/// // The result is a BooleanArray indicating when each string in `array`
171/// // matches the regular expression `regexp`
172/// let result = regexp_is_match_scalar(&array, regexp, None).unwrap();
173/// assert_eq!(result, BooleanArray::from(vec![true, false, true, false]));
174/// ```
175pub fn regexp_is_match_scalar<'a, S>(
176    array: &'a S,
177    regex: &str,
178    flag: Option<&str>,
179) -> Result<BooleanArray, ArrowError>
180where
181    &'a S: StringArrayType<'a>,
182{
183    let mut result = BooleanBufferBuilder::new(array.len());
184
185    let pattern = match flag {
186        Some(flag) => format!("(?{flag}){regex}"),
187        None => regex.to_string(),
188    };
189
190    if pattern.is_empty() {
191        result.append_n(array.len(), true);
192    } else {
193        let re = Regex::new(pattern.as_str()).map_err(|e| {
194            ArrowError::ComputeError(format!("Regular expression did not compile: {e:?}"))
195        })?;
196        for i in 0..array.len() {
197            let value = array.value(i);
198            result.append(re.is_match(value));
199        }
200    }
201
202    let values = BooleanBuffer::from(result);
203    let nulls = array
204        .nulls()
205        .map(|n| n.inner().sliced())
206        .and_then(|b| NullBuffer::from_unsliced_buffer(b, array.len()));
207    Ok(BooleanArray::new(values, nulls))
208}
209
210macro_rules! process_regexp_array_match {
211    ($array:expr, $regex_array:expr, $flags_array:expr, $list_builder:expr) => {
212        let mut patterns: HashMap<String, Regex> = HashMap::new();
213
214        let complete_pattern = match $flags_array {
215            Some(flags) => Box::new($regex_array.iter().zip(flags.iter()).map(
216                |(pattern, flags)| {
217                    pattern.map(|pattern| match flags {
218                        Some(value) => format!("(?{value}){pattern}"),
219                        None => pattern.to_string(),
220                    })
221                },
222            )) as Box<dyn Iterator<Item = Option<String>>>,
223            None => Box::new(
224                $regex_array
225                    .iter()
226                    .map(|pattern| pattern.map(|pattern| pattern.to_string())),
227            ),
228        };
229
230        $array
231            .iter()
232            .zip(complete_pattern)
233            .map(|(value, pattern)| {
234                match (value, pattern) {
235                    // Required for Postgres compatibility:
236                    // SELECT regexp_match('foobarbequebaz', ''); = {""}
237                    (Some(_), Some(pattern)) if pattern == *"" => {
238                        $list_builder.values().append_value("");
239                        $list_builder.append(true);
240                    }
241                    (Some(value), Some(pattern)) => {
242                        let existing_pattern = patterns.get(&pattern);
243                        let re = match existing_pattern {
244                            Some(re) => re,
245                            None => {
246                                let re = Regex::new(pattern.as_str()).map_err(|e| {
247                                    ArrowError::ComputeError(format!(
248                                        "Regular expression did not compile: {e:?}"
249                                    ))
250                                })?;
251                                patterns.entry(pattern).or_insert(re)
252                            }
253                        };
254                        match re.captures(value) {
255                            Some(caps) => {
256                                let mut iter = caps.iter();
257                                if caps.len() > 1 {
258                                    iter.next();
259                                }
260                                for m in iter.flatten() {
261                                    $list_builder.values().append_value(m.as_str());
262                                }
263
264                                $list_builder.append(true);
265                            }
266                            None => $list_builder.append(false),
267                        }
268                    }
269                    _ => $list_builder.append(false),
270                }
271                Ok(())
272            })
273            .collect::<Result<Vec<()>, ArrowError>>()?;
274    };
275}
276
277fn regexp_array_match<OffsetSize: OffsetSizeTrait>(
278    array: &GenericStringArray<OffsetSize>,
279    regex_array: &GenericStringArray<OffsetSize>,
280    flags_array: Option<&GenericStringArray<OffsetSize>>,
281) -> Result<ArrayRef, ArrowError> {
282    let builder: GenericStringBuilder<OffsetSize> = GenericStringBuilder::with_capacity(0, 0);
283    let mut list_builder = ListBuilder::new(builder);
284
285    process_regexp_array_match!(array, regex_array, flags_array, list_builder);
286
287    Ok(Arc::new(list_builder.finish()))
288}
289
290fn regexp_array_match_utf8view(
291    array: &StringViewArray,
292    regex_array: &StringViewArray,
293    flags_array: Option<&StringViewArray>,
294) -> Result<ArrayRef, ArrowError> {
295    let builder = StringViewBuilder::with_capacity(0);
296    let mut list_builder = ListBuilder::new(builder);
297
298    process_regexp_array_match!(array, regex_array, flags_array, list_builder);
299
300    Ok(Arc::new(list_builder.finish()))
301}
302
303fn get_scalar_pattern_flag<'a, OffsetSize: OffsetSizeTrait>(
304    regex_array: &'a dyn Array,
305    flag_array: Option<&'a dyn Array>,
306) -> (Option<&'a str>, Option<&'a str>) {
307    let regex = regex_array.as_string::<OffsetSize>();
308    let regex = regex.is_valid(0).then(|| regex.value(0));
309
310    if let Some(flag_array) = flag_array {
311        let flag = flag_array.as_string::<OffsetSize>();
312        (regex, flag.is_valid(0).then(|| flag.value(0)))
313    } else {
314        (regex, None)
315    }
316}
317
318fn get_scalar_pattern_flag_utf8view<'a>(
319    regex_array: &'a dyn Array,
320    flag_array: Option<&'a dyn Array>,
321) -> (Option<&'a str>, Option<&'a str>) {
322    let regex = regex_array.as_string_view();
323    let regex = regex.is_valid(0).then(|| regex.value(0));
324
325    if let Some(flag_array) = flag_array {
326        let flag = flag_array.as_string_view();
327        (regex, flag.is_valid(0).then(|| flag.value(0)))
328    } else {
329        (regex, None)
330    }
331}
332
333macro_rules! process_regexp_match {
334    ($array:expr, $regex:expr, $list_builder:expr) => {
335        $array
336            .iter()
337            .map(|value| {
338                match value {
339                    // Required for Postgres compatibility:
340                    // SELECT regexp_match('foobarbequebaz', ''); = {""}
341                    Some(_) if $regex.as_str().is_empty() => {
342                        $list_builder.values().append_value("");
343                        $list_builder.append(true);
344                    }
345                    Some(value) => match $regex.captures(value) {
346                        Some(caps) => {
347                            let mut iter = caps.iter();
348                            if caps.len() > 1 {
349                                iter.next();
350                            }
351                            for m in iter.flatten() {
352                                $list_builder.values().append_value(m.as_str());
353                            }
354                            $list_builder.append(true);
355                        }
356                        None => $list_builder.append(false),
357                    },
358                    None => $list_builder.append(false),
359                }
360                Ok(())
361            })
362            .collect::<Result<Vec<()>, ArrowError>>()?
363    };
364}
365
366fn regexp_scalar_match<OffsetSize: OffsetSizeTrait>(
367    array: &GenericStringArray<OffsetSize>,
368    regex: &Regex,
369) -> Result<ArrayRef, ArrowError> {
370    let builder: GenericStringBuilder<OffsetSize> = GenericStringBuilder::with_capacity(0, 0);
371    let mut list_builder = ListBuilder::new(builder);
372
373    process_regexp_match!(array, regex, list_builder);
374
375    Ok(Arc::new(list_builder.finish()))
376}
377
378fn regexp_scalar_match_utf8view(
379    array: &StringViewArray,
380    regex: &Regex,
381) -> Result<ArrayRef, ArrowError> {
382    let builder = StringViewBuilder::with_capacity(0);
383    let mut list_builder = ListBuilder::new(builder);
384
385    process_regexp_match!(array, regex, list_builder);
386
387    Ok(Arc::new(list_builder.finish()))
388}
389
390/// Extract all groups matched by a regular expression for a given String array.
391///
392/// Modelled after the Postgres [regexp_match].
393///
394/// Returns a ListArray of [`GenericStringArray`] with each element containing the leftmost-first
395/// match of the corresponding index in `regex_array` to string in `array`
396///
397/// If there is no match, the list element is NULL.
398///
399/// If a match is found, and the pattern contains no capturing parenthesized subexpressions,
400/// then the list element is a single-element [`GenericStringArray`] containing the substring
401/// matching the whole pattern.
402///
403/// If a match is found, and the pattern contains capturing parenthesized subexpressions, then the
404/// list element is a [`GenericStringArray`] whose n'th element is the substring matching
405/// the n'th capturing parenthesized subexpression of the pattern.
406///
407/// The flags parameter is an optional text string containing zero or more single-letter flags
408/// that change the function's behavior.
409///
410/// # See Also
411/// * [`regexp_is_match`] for matching (rather than extracting) a regular expression against an array of strings
412///
413/// [regexp_match]: https://www.postgresql.org/docs/current/functions-matching.html#FUNCTIONS-POSIX-REGEXP
414pub fn regexp_match(
415    array: &dyn Array,
416    regex_array: &dyn Datum,
417    flags_array: Option<&dyn Datum>,
418) -> Result<ArrayRef, ArrowError> {
419    let (rhs, is_rhs_scalar) = regex_array.get();
420
421    if array.data_type() != rhs.data_type() {
422        return Err(ArrowError::ComputeError(
423            "regexp_match() requires both array and pattern to be either Utf8, Utf8View or LargeUtf8"
424                .to_string(),
425        ));
426    }
427
428    let (flags, is_flags_scalar) = match flags_array {
429        Some(flags) => {
430            let (flags, is_flags_scalar) = flags.get();
431            (Some(flags), Some(is_flags_scalar))
432        }
433        None => (None, None),
434    };
435
436    if is_flags_scalar.is_some_and(|is_flags_scalar| is_rhs_scalar != is_flags_scalar) {
437        return Err(ArrowError::ComputeError(
438            "regexp_match() requires both pattern and flags to be either scalar or array"
439                .to_string(),
440        ));
441    }
442
443    if flags.is_some_and(|flags| rhs.data_type() != flags.data_type()) {
444        return Err(ArrowError::ComputeError(
445            "regexp_match() requires both pattern and flags to be either Utf8, Utf8View or LargeUtf8"
446                .to_string(),
447        ));
448    }
449
450    if is_rhs_scalar {
451        // Regex and flag is scalars
452        let (regex, flag) = match rhs.data_type() {
453            DataType::Utf8View => get_scalar_pattern_flag_utf8view(rhs, flags),
454            DataType::Utf8 => get_scalar_pattern_flag::<i32>(rhs, flags),
455            DataType::LargeUtf8 => get_scalar_pattern_flag::<i64>(rhs, flags),
456            _ => {
457                return Err(ArrowError::ComputeError(
458                    "regexp_match() requires pattern to be either Utf8, Utf8View or LargeUtf8"
459                        .to_string(),
460                ));
461            }
462        };
463
464        let Some(regex) = regex else {
465            return Ok(new_null_array(
466                &DataType::List(Arc::new(Field::new_list_field(
467                    array.data_type().clone(),
468                    true,
469                ))),
470                array.len(),
471            ));
472        };
473
474        let pattern = if let Some(flag) = flag {
475            format!("(?{flag}){regex}")
476        } else {
477            regex.to_string()
478        };
479
480        let re = Regex::new(pattern.as_str()).map_err(|e| {
481            ArrowError::ComputeError(format!("Regular expression did not compile: {e:?}"))
482        })?;
483
484        match array.data_type() {
485            DataType::Utf8View => regexp_scalar_match_utf8view(array.as_string_view(), &re),
486            DataType::Utf8 => regexp_scalar_match(array.as_string::<i32>(), &re),
487            DataType::LargeUtf8 => regexp_scalar_match(array.as_string::<i64>(), &re),
488            _ => Err(ArrowError::ComputeError(
489                "regexp_match() requires array to be either Utf8, Utf8View or LargeUtf8"
490                    .to_string(),
491            )),
492        }
493    } else {
494        match array.data_type() {
495            DataType::Utf8View => {
496                let regex_array = rhs.as_string_view();
497                let flags_array = flags.map(|flags| flags.as_string_view());
498                regexp_array_match_utf8view(array.as_string_view(), regex_array, flags_array)
499            }
500            DataType::Utf8 => {
501                let regex_array = rhs.as_string();
502                let flags_array = flags.map(|flags| flags.as_string());
503                regexp_array_match(array.as_string::<i32>(), regex_array, flags_array)
504            }
505            DataType::LargeUtf8 => {
506                let regex_array = rhs.as_string();
507                let flags_array = flags.map(|flags| flags.as_string());
508                regexp_array_match(array.as_string::<i64>(), regex_array, flags_array)
509            }
510            _ => Err(ArrowError::ComputeError(
511                "regexp_match() requires array to be either Utf8, Utf8View or LargeUtf8"
512                    .to_string(),
513            )),
514        }
515    }
516}
517
518#[cfg(test)]
519mod tests {
520    use super::*;
521
522    macro_rules! test_match_single_group {
523        ($test_name:ident, $values:expr, $patterns:expr, $arr_type:ty, $builder_type:ty, $expected:expr) => {
524            #[test]
525            fn $test_name() {
526                let array: $arr_type = <$arr_type>::from($values);
527                let pattern: $arr_type = <$arr_type>::from($patterns);
528
529                let actual = regexp_match(&array, &pattern, None).unwrap();
530
531                let elem_builder: $builder_type = <$builder_type>::new();
532                let mut expected_builder = ListBuilder::new(elem_builder);
533
534                for val in $expected {
535                    match val {
536                        Some(v) => {
537                            expected_builder.values().append_value(v);
538                            expected_builder.append(true);
539                        }
540                        None => expected_builder.append(false),
541                    }
542                }
543
544                let expected = expected_builder.finish();
545                let result = actual.as_any().downcast_ref::<ListArray>().unwrap();
546                assert_eq!(&expected, result);
547            }
548        };
549    }
550
551    test_match_single_group!(
552        match_single_group_string,
553        vec![
554            Some("abc-005-def"),
555            Some("X-7-5"),
556            Some("X545"),
557            None,
558            Some("foobarbequebaz"),
559            Some("foobarbequebaz"),
560        ],
561        vec![
562            r".*-(\d*)-.*",
563            r".*-(\d*)-.*",
564            r".*-(\d*)-.*",
565            r".*-(\d*)-.*",
566            r"(bar)(bequ1e)",
567            ""
568        ],
569        StringArray,
570        GenericStringBuilder<i32>,
571        [Some("005"), Some("7"), None, None, None, Some("")]
572    );
573    test_match_single_group!(
574        match_single_group_string_view,
575        vec![
576            Some("abc-005-def"),
577            Some("X-7-5"),
578            Some("X545"),
579            None,
580            Some("foobarbequebaz"),
581            Some("foobarbequebaz"),
582        ],
583        vec![
584            r".*-(\d*)-.*",
585            r".*-(\d*)-.*",
586            r".*-(\d*)-.*",
587            r".*-(\d*)-.*",
588            r"(bar)(bequ1e)",
589            ""
590        ],
591        StringViewArray,
592        StringViewBuilder,
593        [Some("005"), Some("7"), None, None, None, Some("")]
594    );
595
596    macro_rules! test_match_single_group_with_flags {
597        ($test_name:ident, $values:expr, $patterns:expr, $flags:expr, $array_type:ty, $builder_type:ty, $expected:expr) => {
598            #[test]
599            fn $test_name() {
600                let array: $array_type = <$array_type>::from($values);
601                let pattern: $array_type = <$array_type>::from($patterns);
602                let flags: $array_type = <$array_type>::from($flags);
603
604                let actual = regexp_match(&array, &pattern, Some(&flags)).unwrap();
605
606                let elem_builder: $builder_type = <$builder_type>::new();
607                let mut expected_builder = ListBuilder::new(elem_builder);
608
609                for val in $expected {
610                    match val {
611                        Some(v) => {
612                            expected_builder.values().append_value(v);
613                            expected_builder.append(true);
614                        }
615                        None => {
616                            expected_builder.append(false);
617                        }
618                    }
619                }
620
621                let expected = expected_builder.finish();
622                let result = actual.as_any().downcast_ref::<ListArray>().unwrap();
623                assert_eq!(&expected, result);
624            }
625        };
626    }
627
628    test_match_single_group_with_flags!(
629        match_single_group_with_flags_string,
630        vec![Some("abc-005-def"), Some("X-7-5"), Some("X545"), None],
631        vec![r"x.*-(\d*)-.*"; 4],
632        vec!["i"; 4],
633        StringArray,
634        GenericStringBuilder<i32>,
635        [None, Some("7"), None, None]
636    );
637    test_match_single_group_with_flags!(
638        match_single_group_with_flags_stringview,
639        vec![Some("abc-005-def"), Some("X-7-5"), Some("X545"), None],
640        vec![r"x.*-(\d*)-.*"; 4],
641        vec!["i"; 4],
642        StringViewArray,
643        StringViewBuilder,
644        [None, Some("7"), None, None]
645    );
646
647    macro_rules! test_match_scalar_pattern {
648        ($test_name:ident, $values:expr, $pattern:expr, $flag:expr, $array_type:ty, $builder_type:ty, $expected:expr) => {
649            #[test]
650            fn $test_name() {
651                let array: $array_type = <$array_type>::from($values);
652
653                let pattern_scalar = Scalar::new(<$array_type>::from(vec![$pattern; 1]));
654                let flag_scalar = Scalar::new(<$array_type>::from(vec![$flag; 1]));
655
656                let actual = regexp_match(&array, &pattern_scalar, Some(&flag_scalar)).unwrap();
657
658                let elem_builder: $builder_type = <$builder_type>::new();
659                let mut expected_builder = ListBuilder::new(elem_builder);
660
661                for val in $expected {
662                    match val {
663                        Some(v) => {
664                            expected_builder.values().append_value(v);
665                            expected_builder.append(true);
666                        }
667                        None => expected_builder.append(false),
668                    }
669                }
670
671                let expected = expected_builder.finish();
672                let result = actual.as_any().downcast_ref::<ListArray>().unwrap();
673                assert_eq!(&expected, result);
674            }
675        };
676    }
677
678    test_match_scalar_pattern!(
679        match_scalar_pattern_string_with_flags,
680        vec![
681            Some("abc-005-def"),
682            Some("x-7-5"),
683            Some("X-0-Y"),
684            Some("X545"),
685            None
686        ],
687        r"x.*-(\d*)-.*",
688        Some("i"),
689        StringArray,
690        GenericStringBuilder<i32>,
691        [None, Some("7"), Some("0"), None, None]
692    );
693    test_match_scalar_pattern!(
694        match_scalar_pattern_stringview_with_flags,
695        vec![
696            Some("abc-005-def"),
697            Some("x-7-5"),
698            Some("X-0-Y"),
699            Some("X545"),
700            None
701        ],
702        r"x.*-(\d*)-.*",
703        Some("i"),
704        StringViewArray,
705        StringViewBuilder,
706        [None, Some("7"), Some("0"), None, None]
707    );
708
709    test_match_scalar_pattern!(
710        match_scalar_pattern_string_no_flags,
711        vec![
712            Some("abc-005-def"),
713            Some("x-7-5"),
714            Some("X-0-Y"),
715            Some("X545"),
716            None
717        ],
718        r"x.*-(\d*)-.*",
719        None::<&str>,
720        StringArray,
721        GenericStringBuilder<i32>,
722        [None, Some("7"), None, None, None]
723    );
724    test_match_scalar_pattern!(
725        match_scalar_pattern_stringview_no_flags,
726        vec![
727            Some("abc-005-def"),
728            Some("x-7-5"),
729            Some("X-0-Y"),
730            Some("X545"),
731            None
732        ],
733        r"x.*-(\d*)-.*",
734        None::<&str>,
735        StringViewArray,
736        StringViewBuilder,
737        [None, Some("7"), None, None, None]
738    );
739
740    macro_rules! test_match_scalar_no_pattern {
741        ($test_name:ident, $values:expr, $array_type:ty, $pattern_type:expr, $builder_type:ty, $expected:expr) => {
742            #[test]
743            fn $test_name() {
744                let array: $array_type = <$array_type>::from($values);
745                let pattern = Scalar::new(new_null_array(&$pattern_type, 1));
746
747                let actual = regexp_match(&array, &pattern, None).unwrap();
748
749                let elem_builder: $builder_type = <$builder_type>::new();
750                let mut expected_builder = ListBuilder::new(elem_builder);
751
752                for val in $expected {
753                    match val {
754                        Some(v) => {
755                            expected_builder.values().append_value(v);
756                            expected_builder.append(true);
757                        }
758                        None => expected_builder.append(false),
759                    }
760                }
761
762                let expected = expected_builder.finish();
763                let result = actual.as_any().downcast_ref::<ListArray>().unwrap();
764                assert_eq!(&expected, result);
765            }
766        };
767    }
768
769    test_match_scalar_no_pattern!(
770        match_scalar_no_pattern_string,
771        vec![Some("abc-005-def"), Some("X-7-5"), Some("X545"), None],
772        StringArray,
773        DataType::Utf8,
774        GenericStringBuilder<i32>,
775        [None::<&str>, None, None, None]
776    );
777    test_match_scalar_no_pattern!(
778        match_scalar_no_pattern_stringview,
779        vec![Some("abc-005-def"), Some("X-7-5"), Some("X545"), None],
780        StringViewArray,
781        DataType::Utf8View,
782        StringViewBuilder,
783        [None::<&str>, None, None, None]
784    );
785
786    macro_rules! test_match_single_group_not_skip {
787        ($test_name:ident, $values:expr, $pattern:expr, $array_type:ty, $builder_type:ty, $expected:expr) => {
788            #[test]
789            fn $test_name() {
790                let array: $array_type = <$array_type>::from($values);
791                let pattern: $array_type = <$array_type>::from(vec![$pattern]);
792
793                let actual = regexp_match(&array, &pattern, None).unwrap();
794
795                let elem_builder: $builder_type = <$builder_type>::new();
796                let mut expected_builder = ListBuilder::new(elem_builder);
797
798                for val in $expected {
799                    match val {
800                        Some(v) => {
801                            expected_builder.values().append_value(v);
802                            expected_builder.append(true);
803                        }
804                        None => expected_builder.append(false),
805                    }
806                }
807
808                let expected = expected_builder.finish();
809                let result = actual.as_any().downcast_ref::<ListArray>().unwrap();
810                assert_eq!(&expected, result);
811            }
812        };
813    }
814
815    test_match_single_group_not_skip!(
816        match_single_group_not_skip_string,
817        vec![Some("foo"), Some("bar")],
818        r"foo",
819        StringArray,
820        GenericStringBuilder<i32>,
821        [Some("foo")]
822    );
823    test_match_single_group_not_skip!(
824        match_single_group_not_skip_stringview,
825        vec![Some("foo"), Some("bar")],
826        r"foo",
827        StringViewArray,
828        StringViewBuilder,
829        [Some("foo")]
830    );
831
832    macro_rules! test_flag_utf8 {
833        ($test_name:ident, $left:expr, $right:expr, $op:expr, $expected:expr) => {
834            #[test]
835            fn $test_name() {
836                let left = $left;
837                let right = $right;
838                let res = $op(&left, &right, None).unwrap();
839                let expected = $expected;
840                assert_eq!(expected.len(), res.len());
841                for i in 0..res.len() {
842                    let v = res.value(i);
843                    assert_eq!(v, expected[i]);
844                }
845            }
846        };
847        ($test_name:ident, $left:expr, $right:expr, $flag:expr, $op:expr, $expected:expr) => {
848            #[test]
849            fn $test_name() {
850                let left = $left;
851                let right = $right;
852                let flag = Some($flag);
853                let res = $op(&left, &right, flag.as_ref()).unwrap();
854                let expected = $expected;
855                assert_eq!(expected.len(), res.len());
856                for i in 0..res.len() {
857                    let v = res.value(i);
858                    assert_eq!(v, expected[i]);
859                }
860            }
861        };
862    }
863
864    macro_rules! test_flag_utf8_scalar {
865        ($test_name:ident, $left:expr, $right:expr, $op:expr, $expected:expr) => {
866            #[test]
867            fn $test_name() {
868                let left = $left;
869                let res = $op(&left, $right, None).unwrap();
870                let expected = $expected;
871                assert_eq!(expected.len(), res.len());
872                for i in 0..res.len() {
873                    let v = res.value(i);
874                    assert_eq!(
875                        v,
876                        expected[i],
877                        "unexpected result when comparing {} at position {} to {} ",
878                        left.value(i),
879                        i,
880                        $right
881                    );
882                }
883            }
884        };
885        ($test_name:ident, $left:expr, $right:expr, $flag:expr, $op:expr, $expected:expr) => {
886            #[test]
887            fn $test_name() {
888                let left = $left;
889                let flag = Some($flag);
890                let res = $op(&left, $right, flag).unwrap();
891                let expected = $expected;
892                assert_eq!(expected.len(), res.len());
893                for i in 0..res.len() {
894                    let v = res.value(i);
895                    assert_eq!(
896                        v,
897                        expected[i],
898                        "unexpected result when comparing {} at position {} to {} ",
899                        left.value(i),
900                        i,
901                        $right
902                    );
903                }
904            }
905        };
906    }
907
908    test_flag_utf8!(
909        test_array_regexp_is_match_utf8,
910        StringArray::from(vec!["arrow", "arrow", "arrow", "arrow", "arrow", "arrow"]),
911        StringArray::from(vec!["^ar", "^AR", "ow$", "OW$", "foo", ""]),
912        regexp_is_match::<StringArray, StringArray, StringArray>,
913        [true, false, true, false, false, true]
914    );
915    test_flag_utf8!(
916        test_array_regexp_is_match_utf8_insensitive,
917        StringArray::from(vec!["arrow", "arrow", "arrow", "arrow", "arrow", "arrow"]),
918        StringArray::from(vec!["^ar", "^AR", "ow$", "OW$", "foo", ""]),
919        StringArray::from(vec!["i"; 6]),
920        regexp_is_match,
921        [true, true, true, true, false, true]
922    );
923
924    test_flag_utf8_scalar!(
925        test_array_regexp_is_match_utf8_scalar,
926        StringArray::from(vec!["arrow", "ARROW", "parquet", "PARQUET"]),
927        "^ar",
928        regexp_is_match_scalar,
929        [true, false, false, false]
930    );
931    test_flag_utf8_scalar!(
932        test_array_regexp_is_match_utf8_scalar_empty,
933        StringArray::from(vec!["arrow", "ARROW", "parquet", "PARQUET"]),
934        "",
935        regexp_is_match_scalar,
936        [true, true, true, true]
937    );
938    test_flag_utf8_scalar!(
939        test_array_regexp_is_match_utf8_scalar_insensitive,
940        StringArray::from(vec!["arrow", "ARROW", "parquet", "PARQUET"]),
941        "^ar",
942        "i",
943        regexp_is_match_scalar,
944        [true, true, false, false]
945    );
946
947    test_flag_utf8!(
948        tes_array_regexp_is_match,
949        StringViewArray::from(vec!["arrow", "arrow", "arrow", "arrow", "arrow", "arrow"]),
950        StringViewArray::from(vec!["^ar", "^AR", "ow$", "OW$", "foo", ""]),
951        regexp_is_match::<StringViewArray, StringViewArray, StringViewArray>,
952        [true, false, true, false, false, true]
953    );
954    test_flag_utf8!(
955        test_array_regexp_is_match_2,
956        StringViewArray::from(vec!["arrow", "arrow", "arrow", "arrow", "arrow", "arrow"]),
957        StringArray::from(vec!["^ar", "^AR", "ow$", "OW$", "foo", ""]),
958        regexp_is_match::<StringViewArray, GenericStringArray<i32>, GenericStringArray<i32>>,
959        [true, false, true, false, false, true]
960    );
961    test_flag_utf8!(
962        test_array_regexp_is_match_insensitive,
963        StringViewArray::from(vec![
964            "Official Rust implementation of Apache Arrow",
965            "apache/arrow-rs",
966            "apache/arrow-rs",
967            "parquet",
968            "parquet",
969            "row",
970            "row",
971        ]),
972        StringViewArray::from(vec![
973            ".*rust implement.*",
974            "^ap",
975            "^AP",
976            "et$",
977            "ET$",
978            "foo",
979            ""
980        ]),
981        StringViewArray::from(vec!["i"; 7]),
982        regexp_is_match::<StringViewArray, StringViewArray, StringViewArray>,
983        [true, true, true, true, true, false, true]
984    );
985    test_flag_utf8!(
986        test_array_regexp_is_match_insensitive_2,
987        LargeStringArray::from(vec!["arrow", "arrow", "arrow", "arrow", "arrow", "arrow"]),
988        StringViewArray::from(vec!["^ar", "^AR", "ow$", "OW$", "foo", ""]),
989        StringArray::from(vec!["i"; 6]),
990        regexp_is_match::<GenericStringArray<i64>, StringViewArray, GenericStringArray<i32>>,
991        [true, true, true, true, false, true]
992    );
993
994    test_flag_utf8_scalar!(
995        test_array_regexp_is_match_scalar,
996        StringViewArray::from(vec![
997            "apache/arrow-rs",
998            "APACHE/ARROW-RS",
999            "parquet",
1000            "PARQUET",
1001        ]),
1002        "^ap",
1003        regexp_is_match_scalar::<StringViewArray>,
1004        [true, false, false, false]
1005    );
1006    test_flag_utf8_scalar!(
1007        test_array_regexp_is_match_scalar_empty,
1008        StringViewArray::from(vec![
1009            "apache/arrow-rs",
1010            "APACHE/ARROW-RS",
1011            "parquet",
1012            "PARQUET",
1013        ]),
1014        "",
1015        regexp_is_match_scalar::<StringViewArray>,
1016        [true, true, true, true]
1017    );
1018    test_flag_utf8_scalar!(
1019        test_array_regexp_is_match_scalar_insensitive,
1020        StringViewArray::from(vec![
1021            "apache/arrow-rs",
1022            "APACHE/ARROW-RS",
1023            "parquet",
1024            "PARQUET",
1025        ]),
1026        "^ap",
1027        "i",
1028        regexp_is_match_scalar::<StringViewArray>,
1029        [true, true, false, false]
1030    );
1031}