1use 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 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
161fn 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 ®ex_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 ®ex_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 ®ex_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 ®ex_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 ®ex_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 ®ex_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 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 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 let start_scalar =
284 is_start_scalar.then(|| start_array.map_or(1, |array| array.value(0)));
285 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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}