use arrow::array::{
Array, ArrayRef, AsArray, Datum, Int64Array, Int64Builder, StringArrayType,
};
use arrow::buffer::NullBuffer;
use arrow::datatypes::{DataType, Int64Type};
use arrow::datatypes::{
DataType::Int64, DataType::LargeUtf8, DataType::Utf8, DataType::Utf8View,
};
use arrow::error::ArrowError;
use datafusion_common::{Result, ScalarValue, exec_err, internal_err};
use datafusion_expr::{
ColumnarValue, Documentation, ScalarFunctionArgs, ScalarUDFImpl, Signature,
TypeSignature::Exact, TypeSignature::Uniform, Volatility,
};
use datafusion_macros::user_doc;
use regex::Regex;
use std::collections::HashMap;
use std::collections::hash_map::Entry;
use std::sync::Arc;
use crate::regex::{compile_regex, start_to_byte_offset};
#[user_doc(
doc_section(label = "Regular Expression Functions"),
description = "Returns the position in a string where the specified occurrence of a POSIX regular expression is located.",
syntax_example = "regexp_instr(str, regexp[, start[, N[, flags[, subexpr]]]])",
sql_example = r#"```sql
> SELECT regexp_instr('ABCDEF', 'C(.)(..)');
+---------------------------------------------------------------+
| regexp_instr(Utf8("ABCDEF"),Utf8("C(.)(..)")) |
+---------------------------------------------------------------+
| 3 |
+---------------------------------------------------------------+
```"#,
standard_argument(name = "str", prefix = "String"),
standard_argument(name = "regexp", prefix = "Regular"),
argument(
name = "start",
description = "Optional start position (the first position is 1) to search for the regular expression. Can be a constant, column, or function. Defaults to 1"
),
argument(
name = "N",
description = "Optional The N-th occurrence of pattern to find. Defaults to 1 (first match). Can be a constant, column, or function."
),
argument(
name = "flags",
description = r#"Optional regular expression flags that control the behavior of the regular expression. Refer to the flags reference above for supported flags."#
),
argument(
name = "subexpr",
description = "Optional Specifies which capture group (subexpression) to return the position for. Defaults to 0, which returns the position of the entire match."
)
)]
#[derive(Debug, PartialEq, Eq, Hash)]
pub struct RegexpInstrFunc {
signature: Signature,
}
impl Default for RegexpInstrFunc {
fn default() -> Self {
Self::new()
}
}
impl RegexpInstrFunc {
pub fn new() -> Self {
Self {
signature: Signature::one_of(
vec![
Uniform(2, vec![Utf8View, LargeUtf8, Utf8]),
Exact(vec![Utf8View, Utf8View, Int64]),
Exact(vec![LargeUtf8, LargeUtf8, Int64]),
Exact(vec![Utf8, Utf8, Int64]),
Exact(vec![Utf8View, Utf8View, Int64, Int64]),
Exact(vec![LargeUtf8, LargeUtf8, Int64, Int64]),
Exact(vec![Utf8, Utf8, Int64, Int64]),
Exact(vec![Utf8View, Utf8View, Int64, Int64, Utf8View]),
Exact(vec![LargeUtf8, LargeUtf8, Int64, Int64, LargeUtf8]),
Exact(vec![Utf8, Utf8, Int64, Int64, Utf8]),
Exact(vec![Utf8View, Utf8View, Int64, Int64, Utf8View, Int64]),
Exact(vec![LargeUtf8, LargeUtf8, Int64, Int64, LargeUtf8, Int64]),
Exact(vec![Utf8, Utf8, Int64, Int64, Utf8, Int64]),
],
Volatility::Immutable,
),
}
}
}
impl ScalarUDFImpl for RegexpInstrFunc {
fn name(&self) -> &str {
"regexp_instr"
}
fn signature(&self) -> &Signature {
&self.signature
}
fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
Ok(Int64)
}
fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
let args = &args.args;
let len = args
.iter()
.fold(Option::<usize>::None, |acc, arg| match arg {
ColumnarValue::Scalar(_) => acc,
ColumnarValue::Array(a) => Some(a.len()),
});
let is_scalar = len.is_none();
let inferred_length = len.unwrap_or(1);
let args = args
.iter()
.map(|arg| arg.to_array(inferred_length))
.collect::<Result<Vec<_>>>()?;
let result = regexp_instr_func(&args);
if is_scalar {
let result = result.and_then(|arr| ScalarValue::try_from_array(&arr, 0));
result.map(ColumnarValue::Scalar)
} else {
result.map(ColumnarValue::Array)
}
}
fn documentation(&self) -> Option<&Documentation> {
self.doc()
}
}
pub fn regexp_instr_func(args: &[ArrayRef]) -> Result<ArrayRef> {
let args_len = args.len();
if !(2..=6).contains(&args_len) {
return exec_err!(
"regexp_instr was called with {args_len} arguments. It requires at least 2 and at most 6."
);
}
let values = &args[0];
match values.data_type() {
Utf8 | LargeUtf8 | Utf8View => (),
other => {
return internal_err!(
"Unsupported data type {other:?} for function regexp_instr"
);
}
}
regexp_instr(
values,
&args[1],
if args_len > 2 { Some(&args[2]) } else { None },
if args_len > 3 { Some(&args[3]) } else { None },
if args_len > 4 { Some(&args[4]) } else { None },
if args_len > 5 { Some(&args[5]) } else { None },
)
.map_err(|e| e.into())
}
fn regexp_instr(
values: &dyn Array,
regex_array: &dyn Datum,
start_array: Option<&dyn Datum>,
nth_array: Option<&dyn Datum>,
flags_array: Option<&dyn Datum>,
subexpr_array: Option<&dyn Datum>,
) -> Result<ArrayRef, ArrowError> {
let (regex_array, _) = regex_array.get();
let start_array = start_array.map(|start| {
let (start, _) = start.get();
start
});
let nth_array = nth_array.map(|nth| {
let (nth, _) = nth.get();
nth
});
let flags_array = flags_array.map(|flags| {
let (flags, _) = flags.get();
flags
});
let subexpr_array = subexpr_array.map(|subexpr| {
let (subexpr, _) = subexpr.get();
subexpr
});
match (values.data_type(), regex_array.data_type(), flags_array) {
(Utf8, Utf8, None) => regexp_instr_inner(
&values.as_string::<i32>(),
®ex_array.as_string::<i32>(),
start_array.map(|start| start.as_primitive::<Int64Type>()),
nth_array.map(|nth| nth.as_primitive::<Int64Type>()),
None,
subexpr_array.map(|subexpr| subexpr.as_primitive::<Int64Type>()),
),
(Utf8, Utf8, Some(flags_array)) if *flags_array.data_type() == Utf8 => regexp_instr_inner(
&values.as_string::<i32>(),
®ex_array.as_string::<i32>(),
start_array.map(|start| start.as_primitive::<Int64Type>()),
nth_array.map(|nth| nth.as_primitive::<Int64Type>()),
Some(&flags_array.as_string::<i32>()),
subexpr_array.map(|subexpr| subexpr.as_primitive::<Int64Type>()),
),
(LargeUtf8, LargeUtf8, None) => regexp_instr_inner(
&values.as_string::<i64>(),
®ex_array.as_string::<i64>(),
start_array.map(|start| start.as_primitive::<Int64Type>()),
nth_array.map(|nth| nth.as_primitive::<Int64Type>()),
None,
subexpr_array.map(|subexpr| subexpr.as_primitive::<Int64Type>()),
),
(LargeUtf8, LargeUtf8, Some(flags_array)) if *flags_array.data_type() == LargeUtf8 => regexp_instr_inner(
&values.as_string::<i64>(),
®ex_array.as_string::<i64>(),
start_array.map(|start| start.as_primitive::<Int64Type>()),
nth_array.map(|nth| nth.as_primitive::<Int64Type>()),
Some(&flags_array.as_string::<i64>()),
subexpr_array.map(|subexpr| subexpr.as_primitive::<Int64Type>()),
),
(Utf8View, Utf8View, None) => regexp_instr_inner(
&values.as_string_view(),
®ex_array.as_string_view(),
start_array.map(|start| start.as_primitive::<Int64Type>()),
nth_array.map(|nth| nth.as_primitive::<Int64Type>()),
None,
subexpr_array.map(|subexpr| subexpr.as_primitive::<Int64Type>()),
),
(Utf8View, Utf8View, Some(flags_array)) if *flags_array.data_type() == Utf8View => regexp_instr_inner(
&values.as_string_view(),
®ex_array.as_string_view(),
start_array.map(|start| start.as_primitive::<Int64Type>()),
nth_array.map(|nth| nth.as_primitive::<Int64Type>()),
Some(&flags_array.as_string_view()),
subexpr_array.map(|subexpr| subexpr.as_primitive::<Int64Type>()),
),
_ => Err(ArrowError::ComputeError(
"regexp_instr() 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(),
)),
}
}
fn regexp_instr_inner<'a, S>(
values: &S,
regex_array: &S,
start_array: Option<&Int64Array>,
nth_array: Option<&Int64Array>,
flags_array: Option<&S>,
subexp_array: Option<&Int64Array>,
) -> Result<ArrayRef, ArrowError>
where
S: StringArrayType<'a>,
{
let len = values.len();
let mut regex_cache = RegexCache::default();
let mut result = Int64Builder::with_capacity(len);
let nulls = NullBuffer::union_many([
values.nulls(),
regex_array.nulls(),
start_array.and_then(|array| array.nulls()),
nth_array.and_then(|array| array.nulls()),
flags_array.and_then(|array| array.nulls()),
subexp_array.and_then(|array| array.nulls()),
]);
for i in 0..len {
if nulls.as_ref().is_some_and(|nulls| nulls.is_null(i)) {
result.append_null();
continue;
}
let value = values.value(i);
let regex = regex_array.value(i);
let flags = flags_array.map(|array| array.value(i));
let pattern = regex_cache.get_or_compile(regex, flags)?;
let start = start_array.map_or(1, |array| array.value(i));
let nth = nth_array.map_or(1, |array| array.value(i));
let subexp = subexp_array.map_or(0, |array| array.value(i));
result.append_value(get_index(value, pattern, start, nth, subexp)?);
}
Ok(Arc::new(result.finish()))
}
#[derive(Default)]
struct RegexCache<'a> {
compiled: Vec<Regex>,
indices: HashMap<(&'a str, Option<&'a str>), usize>,
last: Option<((&'a str, Option<&'a str>), usize)>,
}
impl<'a> RegexCache<'a> {
fn get_or_compile(
&mut self,
regex: &'a str,
flags: Option<&'a str>,
) -> Result<&Regex, ArrowError> {
let key = (regex, flags);
let index = match self.last {
Some((last_key, index)) if last_key == key => index,
_ => {
let index = match self.indices.entry(key) {
Entry::Occupied(entry) => *entry.get(),
Entry::Vacant(entry) => {
self.compiled.push(compile_regex(regex, flags)?);
*entry.insert(self.compiled.len() - 1)
}
};
self.last = Some((key, index));
index
}
};
Ok(&self.compiled[index])
}
}
fn get_index(
value: &str,
pattern: &Regex,
start: i64,
n: i64,
subexpr: i64,
) -> Result<i64, ArrowError> {
if start < 1 {
return Err(ArrowError::ComputeError(
"regexp_instr() requires start to be 1-based".to_string(),
));
}
if n < 1 {
return Err(ArrowError::ComputeError(
"N must be 1 or greater".to_string(),
));
}
let Some(byte_start_offset) = start_to_byte_offset(value, start) else {
return Ok(0);
};
let search_slice = &value[byte_start_offset..];
let match_start = if subexpr > 0 {
pattern
.captures(search_slice)
.and_then(|captures| captures.get(subexpr as usize))
.map(|matched| matched.start())
} else {
pattern
.find_iter(search_slice)
.nth((n - 1) as usize)
.map(|matched| matched.start())
};
Ok(match_start.map_or(0, |offset| {
value[..byte_start_offset + offset].chars().count() as i64 + 1
}))
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::{GenericStringArray, StringViewArray};
use arrow::datatypes::Field;
use datafusion_common::config::ConfigOptions;
use itertools::izip;
#[test]
fn test_regexp_instr() {
test_case_sensitive_regexp_instr_nulls();
test_case_sensitive_regexp_instr_scalar();
test_case_sensitive_regexp_instr_scalar_start();
test_case_sensitive_regexp_instr_scalar_nth();
test_case_sensitive_regexp_instr_scalar_subexp();
test_case_sensitive_regexp_instr_array::<GenericStringArray<i32>>();
test_case_sensitive_regexp_instr_array::<GenericStringArray<i64>>();
test_case_sensitive_regexp_instr_array::<StringViewArray>();
test_case_sensitive_regexp_instr_array_start::<GenericStringArray<i32>>();
test_case_sensitive_regexp_instr_array_start::<GenericStringArray<i64>>();
test_case_sensitive_regexp_instr_array_start::<StringViewArray>();
test_case_sensitive_regexp_instr_array_nth::<GenericStringArray<i32>>();
test_case_sensitive_regexp_instr_array_nth::<GenericStringArray<i64>>();
test_case_sensitive_regexp_instr_array_nth::<StringViewArray>();
test_case_sensitive_regexp_instr_empty_pattern::<GenericStringArray<i32>>();
test_case_sensitive_regexp_instr_empty_pattern::<GenericStringArray<i64>>();
test_case_sensitive_regexp_instr_empty_pattern::<StringViewArray>();
test_case_sensitive_regexp_instr_zero_width_pattern::<GenericStringArray<i32>>();
test_case_sensitive_regexp_instr_zero_width_pattern::<GenericStringArray<i64>>();
test_case_sensitive_regexp_instr_zero_width_pattern::<StringViewArray>();
test_regexp_instr_null_scalar_args();
test_regexp_instr_null_array_rows::<GenericStringArray<i32>>();
test_regexp_instr_null_array_rows::<GenericStringArray<i64>>();
test_regexp_instr_null_array_rows::<StringViewArray>();
}
fn regexp_instr_with_scalar_values(args: &[ScalarValue]) -> Result<ColumnarValue> {
let args_values: Vec<ColumnarValue> = args
.iter()
.map(|sv| ColumnarValue::Scalar(sv.clone()))
.collect();
let arg_fields = args
.iter()
.enumerate()
.map(|(idx, a)| {
Arc::new(Field::new(format!("arg_{idx}"), a.data_type(), true))
})
.collect::<Vec<_>>();
RegexpInstrFunc::new().invoke_with_args(ScalarFunctionArgs {
args: args_values,
arg_fields,
number_rows: args.len(),
return_field: Arc::new(Field::new("f", Int64, true)),
config_options: Arc::new(ConfigOptions::default()),
})
}
fn test_case_sensitive_regexp_instr_nulls() {
let v = "";
let r = "";
let expected = 1;
let regex_sv = ScalarValue::Utf8(Some(r.to_string()));
let re = regexp_instr_with_scalar_values(&[v.to_string().into(), regex_sv]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, Some(expected), "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
for (value, regex) in [
(
ScalarValue::Utf8(None),
ScalarValue::Utf8(Some(String::new())),
),
(
ScalarValue::LargeUtf8(None),
ScalarValue::LargeUtf8(Some(String::new())),
),
(
ScalarValue::Utf8View(None),
ScalarValue::Utf8View(Some(String::new())),
),
] {
let re = regexp_instr_with_scalar_values(&[value, regex]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, None, "regexp_instr NULL scalar test failed");
}
_ => panic!("Unexpected result"),
}
}
}
fn test_case_sensitive_regexp_instr_scalar() {
let values = [
"hello world",
"abcdefg",
"xyz123xyz",
"no match here",
"abc",
"ДатаФусион数据融合📊🔥",
];
let regex = ["o", "d", "123", "z", "gg", "📊"];
let expected: Vec<i64> = vec![5, 4, 4, 0, 0, 15];
izip!(values.iter(), regex.iter())
.enumerate()
.for_each(|(pos, (&v, &r))| {
let v_sv = ScalarValue::Utf8(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8(Some(r.to_string()));
let expected = expected.get(pos).cloned();
let re = regexp_instr_with_scalar_values(&[v_sv, regex_sv]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
let regex_sv = ScalarValue::LargeUtf8(Some(r.to_string()));
let re = regexp_instr_with_scalar_values(&[v_sv, regex_sv]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8View(Some(r.to_string()));
let re = regexp_instr_with_scalar_values(&[v_sv, regex_sv]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
});
}
fn test_case_sensitive_regexp_instr_scalar_start() {
let values = ["abcabcabc", "abcabcabc", ""];
let regex = ["abc", "abc", "gg"];
let start = [4, 5, 5];
let expected: Vec<i64> = vec![4, 7, 0];
izip!(values.iter(), regex.iter(), start.iter())
.enumerate()
.for_each(|(pos, (&v, &r, &s))| {
let v_sv = ScalarValue::Utf8(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let expected = expected.get(pos).cloned();
let re =
regexp_instr_with_scalar_values(&[v_sv, regex_sv, start_sv.clone()]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
let regex_sv = ScalarValue::LargeUtf8(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let re =
regexp_instr_with_scalar_values(&[v_sv, regex_sv, start_sv.clone()]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8View(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let re =
regexp_instr_with_scalar_values(&[v_sv, regex_sv, start_sv.clone()]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
});
}
fn test_case_sensitive_regexp_instr_scalar_nth() {
let values = ["abcabcabc", "abcabcabc", "abcabcabc", "abcabcabc"];
let regex = ["abc", "abc", "abc", "abc"];
let start = [1, 1, 1, 1];
let nth = [1, 2, 3, 4];
let expected: Vec<i64> = vec![1, 4, 7, 0];
izip!(values.iter(), regex.iter(), start.iter(), nth.iter())
.enumerate()
.for_each(|(pos, (&v, &r, &s, &n))| {
let v_sv = ScalarValue::Utf8(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let nth_sv = ScalarValue::Int64(Some(n));
let expected = expected.get(pos).cloned();
let re = regexp_instr_with_scalar_values(&[
v_sv,
regex_sv,
start_sv.clone(),
nth_sv.clone(),
]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
let regex_sv = ScalarValue::LargeUtf8(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let nth_sv = ScalarValue::Int64(Some(n));
let re = regexp_instr_with_scalar_values(&[
v_sv,
regex_sv,
start_sv.clone(),
nth_sv.clone(),
]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8View(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let nth_sv = ScalarValue::Int64(Some(n));
let re = regexp_instr_with_scalar_values(&[
v_sv,
regex_sv,
start_sv.clone(),
nth_sv.clone(),
]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
});
}
fn test_case_sensitive_regexp_instr_scalar_subexp() {
let values = ["12 abc def ghi 34"];
let regex = ["(abc) (def) (ghi)"];
let start = [1];
let nth = [1];
let flags = ["i"];
let subexps = [2];
let expected: Vec<i64> = vec![8];
izip!(
values.iter(),
regex.iter(),
start.iter(),
nth.iter(),
flags.iter(),
subexps.iter()
)
.enumerate()
.for_each(|(pos, (&v, &r, &s, &n, &flag, &subexp))| {
let v_sv = ScalarValue::Utf8(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let nth_sv = ScalarValue::Int64(Some(n));
let flags_sv = ScalarValue::Utf8(Some(flag.to_string()));
let subexp_sv = ScalarValue::Int64(Some(subexp));
let expected = expected.get(pos).cloned();
let re = regexp_instr_with_scalar_values(&[
v_sv,
regex_sv,
start_sv.clone(),
nth_sv.clone(),
flags_sv,
subexp_sv.clone(),
]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::LargeUtf8(Some(v.to_string()));
let regex_sv = ScalarValue::LargeUtf8(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let nth_sv = ScalarValue::Int64(Some(n));
let flags_sv = ScalarValue::LargeUtf8(Some(flag.to_string()));
let subexp_sv = ScalarValue::Int64(Some(subexp));
let re = regexp_instr_with_scalar_values(&[
v_sv,
regex_sv,
start_sv.clone(),
nth_sv.clone(),
flags_sv,
subexp_sv.clone(),
]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
let v_sv = ScalarValue::Utf8View(Some(v.to_string()));
let regex_sv = ScalarValue::Utf8View(Some(r.to_string()));
let start_sv = ScalarValue::Int64(Some(s));
let nth_sv = ScalarValue::Int64(Some(n));
let flags_sv = ScalarValue::Utf8View(Some(flag.to_string()));
let subexp_sv = ScalarValue::Int64(Some(subexp));
let re = regexp_instr_with_scalar_values(&[
v_sv,
regex_sv,
start_sv.clone(),
nth_sv.clone(),
flags_sv,
subexp_sv.clone(),
]);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, expected, "regexp_instr scalar test failed");
}
_ => panic!("Unexpected result"),
}
});
}
fn test_regexp_instr_null_scalar_args() {
let cases: Vec<Vec<ScalarValue>> = vec![
vec![
ScalarValue::Utf8(Some("abc".to_string())),
ScalarValue::Utf8(Some("b".to_string())),
ScalarValue::Int64(None),
],
vec![
ScalarValue::Utf8(Some("abc".to_string())),
ScalarValue::Utf8(Some("b".to_string())),
ScalarValue::Int64(Some(1)),
ScalarValue::Int64(None),
],
vec![
ScalarValue::Utf8(Some("abc".to_string())),
ScalarValue::Utf8(Some("b".to_string())),
ScalarValue::Int64(Some(1)),
ScalarValue::Int64(Some(1)),
ScalarValue::Utf8(None),
],
vec![
ScalarValue::Utf8(Some("abc".to_string())),
ScalarValue::Utf8(Some("(b)".to_string())),
ScalarValue::Int64(Some(1)),
ScalarValue::Int64(Some(1)),
ScalarValue::Utf8(Some("i".to_string())),
ScalarValue::Int64(None),
],
];
for args in cases {
let re = regexp_instr_with_scalar_values(&args);
match re {
Ok(ColumnarValue::Scalar(ScalarValue::Int64(v))) => {
assert_eq!(v, None, "regexp_instr null scalar test failed");
}
_ => panic!("Unexpected result"),
}
}
}
fn test_regexp_instr_null_array_rows<A>()
where
A: From<Vec<Option<&'static str>>> + Array + 'static,
{
let values = A::from(vec![
None,
Some("abc"),
Some("abc"),
Some("abc"),
Some("abc"),
Some("abc"),
Some("abc"),
]);
let regex = A::from(vec![
Some("b"),
None,
Some("b"),
Some("b"),
Some("b"),
Some("(b)"),
Some("b"),
]);
let start = Int64Array::from(vec![
Some(1),
Some(1),
None,
Some(1),
Some(1),
Some(1),
Some(1),
]);
let nth = Int64Array::from(vec![
Some(1),
Some(1),
Some(1),
None,
Some(1),
Some(1),
Some(1),
]);
let flags = A::from(vec![
Some(""),
Some(""),
Some(""),
Some(""),
None,
Some("i"),
Some(""),
]);
let subexp = Int64Array::from(vec![
Some(0),
Some(0),
Some(0),
Some(0),
Some(0),
None,
Some(0),
]);
let expected =
Int64Array::from(vec![None, None, None, None, None, None, Some(2)]);
let re = regexp_instr_func(&[
Arc::new(values),
Arc::new(regex),
Arc::new(start),
Arc::new(nth),
Arc::new(flags),
Arc::new(subexp),
])
.unwrap();
assert_eq!(re.as_ref(), &expected);
}
fn test_case_sensitive_regexp_instr_array<A>()
where
A: From<Vec<&'static str>> + Array + 'static,
{
let values = A::from(vec![
"hello world",
"abcdefg",
"xyz123xyz",
"no match here",
"",
]);
let regex = A::from(vec!["o", "d", "123", "z", "gg"]);
let expected = Int64Array::from(vec![5, 4, 4, 0, 0]);
let re = regexp_instr_func(&[Arc::new(values), Arc::new(regex)]).unwrap();
assert_eq!(re.as_ref(), &expected);
}
fn test_case_sensitive_regexp_instr_array_start<A>()
where
A: From<Vec<&'static str>> + Array + 'static,
{
let values = A::from(vec!["abcabcabc", "abcabcabc", ""]);
let regex = A::from(vec!["abc", "abc", "gg"]);
let start = Int64Array::from(vec![4, 5, 5]);
let expected = Int64Array::from(vec![4, 7, 0]);
let re = regexp_instr_func(&[Arc::new(values), Arc::new(regex), Arc::new(start)])
.unwrap();
assert_eq!(re.as_ref(), &expected);
}
fn test_case_sensitive_regexp_instr_array_nth<A>()
where
A: From<Vec<&'static str>> + Array + 'static,
{
let values = A::from(vec!["abcabcabc", "abcabcabc", "abcabcabc", "abcabcabc"]);
let regex = A::from(vec!["abc", "abc", "abc", "abc"]);
let start = Int64Array::from(vec![1, 1, 1, 1]);
let nth = Int64Array::from(vec![1, 2, 3, 4]);
let expected = Int64Array::from(vec![1, 4, 7, 0]);
let re = regexp_instr_func(&[
Arc::new(values),
Arc::new(regex),
Arc::new(start),
Arc::new(nth),
])
.unwrap();
assert_eq!(re.as_ref(), &expected);
}
fn test_case_sensitive_regexp_instr_empty_pattern<A>()
where
A: From<Vec<&'static str>> + Array + 'static,
{
let values = A::from(vec!["abc", "", "abc", "abc", "😀"]);
let regex = A::from(vec!["", "", "", "", ""]);
let start = Int64Array::from(vec![1, 1, 4, 5, 1]);
let nth = Int64Array::from(vec![1, 1, 1, 1, 2]);
let expected = Int64Array::from(vec![1, 1, 4, 0, 2]);
let re = regexp_instr_func(&[
Arc::new(values),
Arc::new(regex),
Arc::new(start),
Arc::new(nth),
])
.unwrap();
assert_eq!(re.as_ref(), &expected);
}
fn test_case_sensitive_regexp_instr_zero_width_pattern<A>()
where
A: From<Vec<&'static str>> + Array + 'static,
{
let values = A::from(vec!["abc"]);
let regex = A::from(vec!["x*"]);
let start = Int64Array::from(vec![4]);
let expected = Int64Array::from(vec![4]);
let re = regexp_instr_func(&[Arc::new(values), Arc::new(regex), Arc::new(start)])
.unwrap();
assert_eq!(re.as_ref(), &expected);
}
}