1use std::sync::Arc;
19
20use crate::strings::append_view;
21use crate::utils::make_scalar_function;
22use arrow::array::{
23 Array, ArrayRef, AsArray, GenericStringArray, Int64Array, OffsetSizeTrait,
24 StringArrayType, StringViewArray,
25};
26use arrow::buffer::{NullBuffer, ScalarBuffer};
27use arrow::datatypes::DataType;
28use datafusion_common::cast::as_int64_array;
29use datafusion_common::types::{
30 NativeType, logical_int32, logical_int64, logical_string,
31};
32use datafusion_common::{Result, exec_err};
33use datafusion_expr::{
34 Coercion, ColumnarValue, Documentation, ScalarFunctionArgs, ScalarUDFImpl, Signature,
35 TypeSignature, TypeSignatureClass, Volatility,
36};
37use datafusion_macros::user_doc;
38
39#[user_doc(
40 doc_section(label = "String Functions"),
41 description = "Extracts a substring of a specified number of characters from a specific starting position in a string.",
42 syntax_example = "substr(str, start_pos[, length])",
43 alternative_syntax = "substring(str from start_pos for length)",
44 sql_example = r#"```sql
45> select substr('datafusion', 5, 3);
46+----------------------------------------------+
47| substr(Utf8("datafusion"),Int64(5),Int64(3)) |
48+----------------------------------------------+
49| fus |
50+----------------------------------------------+
51```"#,
52 standard_argument(name = "str", prefix = "String"),
53 argument(
54 name = "start_pos",
55 description = "Character position to start the substring at. The first character in the string has a position of 1. If the start position is less than 1, it is treated as if it is before the start of the string and the (absolute) number of characters before position 1 is subtracted from `length` (if given). For example, `substr('abc', -3, 6)` returns `'ab'`."
56 ),
57 argument(
58 name = "length",
59 description = "Number of characters to extract. If not specified, returns the rest of the string after the start position."
60 )
61)]
62#[derive(Debug, PartialEq, Eq, Hash)]
63pub struct SubstrFunc {
64 signature: Signature,
65 aliases: Vec<String>,
66}
67
68impl Default for SubstrFunc {
69 fn default() -> Self {
70 Self::new()
71 }
72}
73
74impl SubstrFunc {
75 pub fn new() -> Self {
76 let string = Coercion::new_exact(TypeSignatureClass::Native(logical_string()));
77 let int64 = Coercion::new_implicit(
78 TypeSignatureClass::Native(logical_int64()),
79 vec![TypeSignatureClass::Native(logical_int32())],
80 NativeType::Int64,
81 );
82 Self {
83 signature: Signature::one_of(
84 vec![
85 TypeSignature::Coercible(vec![string.clone(), int64.clone()]),
86 TypeSignature::Coercible(vec![
87 string.clone(),
88 int64.clone(),
89 int64.clone(),
90 ]),
91 ],
92 Volatility::Immutable,
93 )
94 .with_parameter_names(vec![
95 "str".to_string(),
96 "start_pos".to_string(),
97 "length".to_string(),
98 ])
99 .expect("valid parameter names"),
100 aliases: vec![String::from("substring")],
101 }
102 }
103}
104
105impl ScalarUDFImpl for SubstrFunc {
106 fn name(&self) -> &str {
107 "substr"
108 }
109
110 fn signature(&self) -> &Signature {
111 &self.signature
112 }
113
114 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
115 Ok(arg_types[0].clone())
116 }
117
118 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
119 make_scalar_function(substr, vec![])(&args.args)
120 }
121
122 fn aliases(&self) -> &[String] {
123 &self.aliases
124 }
125
126 fn documentation(&self) -> Option<&Documentation> {
127 self.doc()
128 }
129}
130
131fn substr(args: &[ArrayRef]) -> Result<ArrayRef> {
133 match args[0].data_type() {
134 DataType::Utf8 => {
135 let string_array = args[0].as_string::<i32>();
136 generic_string_substr(string_array, &args[1..])
137 }
138 DataType::LargeUtf8 => {
139 let string_array = args[0].as_string::<i64>();
140 generic_string_substr(string_array, &args[1..])
141 }
142 DataType::Utf8View => {
143 let string_array = args[0].as_string_view();
144 string_view_substr(string_array, &args[1..])
145 }
146 other => exec_err!(
147 "Unsupported data type {other:?} for function substr,\
148 expected Utf8View, Utf8 or LargeUtf8."
149 ),
150 }
151}
152
153pub fn get_true_start_end(
170 input: &str,
171 start: i64,
172 count: Option<i64>,
173 is_input_ascii_only: bool,
174) -> Result<(usize, usize)> {
175 if let Some(count) = count
176 && count < 0
177 {
178 return exec_err!("negative count not allowed: {count}");
179 }
180
181 let Some(start) = start.checked_sub(1) else {
183 return exec_err!("start position overflow: {start}");
184 };
185
186 let end = match count {
187 Some(count) => start.saturating_add(count),
188 None => input.len() as i64,
189 };
190
191 let start = start.clamp(0, input.len() as i64) as usize;
192 let end = end.clamp(0, input.len() as i64) as usize;
193
194 if is_input_ascii_only {
196 return Ok((start, end));
197 }
198
199 let mut byte_start = input.len();
204 let mut byte_end = input.len();
205
206 for (char_idx, (byte_idx, _)) in input.char_indices().enumerate() {
207 if char_idx == start {
208 byte_start = byte_idx;
209 if count.is_none() {
211 break;
212 }
213 }
214 if char_idx == end {
215 byte_end = byte_idx;
216 break;
217 }
218 }
219
220 Ok((byte_start, byte_end))
221}
222
223pub fn enable_ascii_fast_path<'a, V: StringArrayType<'a>>(
234 string_array: &V,
235 start: &Int64Array,
236 count: Option<&Int64Array>,
237) -> bool {
238 let is_short_prefix = match count {
239 Some(count) => {
240 let short_prefix_threshold = 32.0;
241 let n_sample = 10;
242
243 let total_prefix_len = start
246 .iter()
247 .zip(count.iter())
248 .take(n_sample)
249 .map(|(start, count)| {
250 let start = start.unwrap_or(0);
251 let count = count.unwrap_or(0);
252 start.saturating_add(count)
254 })
255 .fold(0i64, |acc, val| acc.saturating_add(val));
256
257 (total_prefix_len as f64 / n_sample as f64) <= short_prefix_threshold
258 }
259 None => false,
260 };
261
262 if is_short_prefix {
263 false
265 } else {
266 string_array.is_ascii()
267 }
268}
269
270fn string_view_substr(
271 string_view_array: &StringViewArray,
272 args: &[ArrayRef],
273) -> Result<ArrayRef> {
274 let start_array = as_int64_array(&args[0])?;
275 let count_array_opt = args.get(1).map(|a| as_int64_array(a)).transpose()?;
276
277 let is_ascii =
278 enable_ascii_fast_path(&string_view_array, start_array, count_array_opt);
279
280 let nulls = NullBuffer::union_many([
282 string_view_array.nulls(),
283 start_array.nulls(),
284 count_array_opt.and_then(|a| a.nulls()),
285 ]);
286
287 let mut views_buf = Vec::with_capacity(string_view_array.len());
288
289 for (i, raw_view) in string_view_array.views().iter().enumerate() {
290 if nulls.as_ref().is_some_and(|n| n.is_null(i)) {
291 views_buf.push(0);
292 continue;
293 }
294
295 let string = string_view_array.value(i);
296 let start = start_array.value(i);
297 let count = count_array_opt.map(|a| a.value(i));
298
299 let (byte_start, byte_end) = get_true_start_end(string, start, count, is_ascii)?;
300 let substr = &string[byte_start..byte_end];
301
302 append_view(&mut views_buf, raw_view, substr, byte_start as u32);
303 }
304
305 let views_buf = ScalarBuffer::from(views_buf);
306
307 unsafe {
312 let array = StringViewArray::new_unchecked(
313 views_buf,
314 string_view_array.data_buffers().to_vec(),
315 nulls,
316 );
317 Ok(Arc::new(array) as ArrayRef)
318 }
319}
320
321fn generic_string_substr<T: OffsetSizeTrait>(
322 string_array: &GenericStringArray<T>,
323 args: &[ArrayRef],
324) -> Result<ArrayRef> {
325 let start_array = as_int64_array(&args[0])?;
326 let count_array_opt = args.get(1).map(|a| as_int64_array(a)).transpose()?;
327
328 let is_ascii = enable_ascii_fast_path(&string_array, start_array, count_array_opt);
329 let nulls = NullBuffer::union_many([
330 string_array.nulls(),
331 start_array.nulls(),
332 count_array_opt.and_then(|a| a.nulls()),
333 ]);
334
335 let result = (0..string_array.len())
336 .map(|i| {
337 if nulls.as_ref().is_some_and(|n| n.is_null(i)) {
338 return Ok(None);
339 }
340
341 let string = string_array.value(i);
342 let start = start_array.value(i);
343 let count = count_array_opt.map(|a| a.value(i));
344
345 let (byte_start, byte_end) =
346 get_true_start_end(string, start, count, is_ascii)?;
347 Ok(Some(&string[byte_start..byte_end]))
348 })
349 .collect::<Result<GenericStringArray<T>>>()?;
350
351 Ok(Arc::new(result) as ArrayRef)
352}
353
354#[cfg(test)]
355mod tests {
356 use std::sync::Arc;
357
358 use arrow::array::{
359 Array, ArrayRef, AsArray, Int64Array, LargeStringArray, StringArray,
360 StringViewArray,
361 };
362 use arrow::datatypes::DataType::{LargeUtf8, Utf8, Utf8View};
363
364 use datafusion_common::{Result, ScalarValue, exec_err};
365 use datafusion_expr::{ColumnarValue, ScalarUDFImpl};
366
367 use crate::unicode::substr::SubstrFunc;
368 use crate::utils::test::test_function;
369
370 #[test]
371 fn test_functions() -> Result<()> {
372 test_function!(
373 SubstrFunc::new(),
374 vec![
375 ColumnarValue::Scalar(ScalarValue::Utf8View(None)),
376 ColumnarValue::Scalar(ScalarValue::from(1i64)),
377 ],
378 Ok(None),
379 &str,
380 Utf8View,
381 StringViewArray
382 );
383 test_function!(
384 SubstrFunc::new(),
385 vec![
386 ColumnarValue::Scalar(ScalarValue::Utf8View(Some(String::from(
387 "alphabet"
388 )))),
389 ColumnarValue::Scalar(ScalarValue::from(0i64)),
390 ],
391 Ok(Some("alphabet")),
392 &str,
393 Utf8View,
394 StringViewArray
395 );
396 test_function!(
397 SubstrFunc::new(),
398 vec![
399 ColumnarValue::Scalar(ScalarValue::Utf8View(Some(String::from(
400 "this és longer than 12B"
401 )))),
402 ColumnarValue::Scalar(ScalarValue::from(5i64)),
403 ColumnarValue::Scalar(ScalarValue::from(2i64)),
404 ],
405 Ok(Some(" é")),
406 &str,
407 Utf8View,
408 StringViewArray
409 );
410 test_function!(
411 SubstrFunc::new(),
412 vec![
413 ColumnarValue::Scalar(ScalarValue::Utf8View(Some(String::from(
414 "this is longer than 12B"
415 )))),
416 ColumnarValue::Scalar(ScalarValue::from(5i64)),
417 ],
418 Ok(Some(" is longer than 12B")),
419 &str,
420 Utf8View,
421 StringViewArray
422 );
423 test_function!(
424 SubstrFunc::new(),
425 vec![
426 ColumnarValue::Scalar(ScalarValue::Utf8View(Some(String::from(
427 "joséésoj"
428 )))),
429 ColumnarValue::Scalar(ScalarValue::from(5i64)),
430 ],
431 Ok(Some("ésoj")),
432 &str,
433 Utf8View,
434 StringViewArray
435 );
436 test_function!(
437 SubstrFunc::new(),
438 vec![
439 ColumnarValue::Scalar(ScalarValue::Utf8View(Some(String::from(
440 "alphabet"
441 )))),
442 ColumnarValue::Scalar(ScalarValue::from(3i64)),
443 ColumnarValue::Scalar(ScalarValue::from(2i64)),
444 ],
445 Ok(Some("ph")),
446 &str,
447 Utf8View,
448 StringViewArray
449 );
450 test_function!(
451 SubstrFunc::new(),
452 vec![
453 ColumnarValue::Scalar(ScalarValue::Utf8View(Some(String::from(
454 "alphabet"
455 )))),
456 ColumnarValue::Scalar(ScalarValue::from(3i64)),
457 ColumnarValue::Scalar(ScalarValue::from(20i64)),
458 ],
459 Ok(Some("phabet")),
460 &str,
461 Utf8View,
462 StringViewArray
463 );
464 test_function!(
465 SubstrFunc::new(),
466 vec![
467 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
468 ColumnarValue::Scalar(ScalarValue::from(0i64)),
469 ],
470 Ok(Some("alphabet")),
471 &str,
472 Utf8,
473 StringArray
474 );
475 test_function!(
476 SubstrFunc::new(),
477 vec![
478 ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some(
479 "alphabet".to_string()
480 ))),
481 ColumnarValue::Scalar(ScalarValue::from(0i64)),
482 ],
483 Ok(Some("alphabet")),
484 &str,
485 LargeUtf8,
486 LargeStringArray
487 );
488 test_function!(
489 SubstrFunc::new(),
490 vec![
491 ColumnarValue::Scalar(ScalarValue::from("joséésoj")),
492 ColumnarValue::Scalar(ScalarValue::from(5i64)),
493 ],
494 Ok(Some("ésoj")),
495 &str,
496 Utf8,
497 StringArray
498 );
499 test_function!(
500 SubstrFunc::new(),
501 vec![
502 ColumnarValue::Scalar(ScalarValue::from("joséésoj")),
503 ColumnarValue::Scalar(ScalarValue::from(-5i64)),
504 ],
505 Ok(Some("joséésoj")),
506 &str,
507 Utf8,
508 StringArray
509 );
510 test_function!(
511 SubstrFunc::new(),
512 vec![
513 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
514 ColumnarValue::Scalar(ScalarValue::from(1i64)),
515 ],
516 Ok(Some("alphabet")),
517 &str,
518 Utf8,
519 StringArray
520 );
521 test_function!(
522 SubstrFunc::new(),
523 vec![
524 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
525 ColumnarValue::Scalar(ScalarValue::from(2i64)),
526 ],
527 Ok(Some("lphabet")),
528 &str,
529 Utf8,
530 StringArray
531 );
532 test_function!(
533 SubstrFunc::new(),
534 vec![
535 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
536 ColumnarValue::Scalar(ScalarValue::from(3i64)),
537 ],
538 Ok(Some("phabet")),
539 &str,
540 Utf8,
541 StringArray
542 );
543 test_function!(
544 SubstrFunc::new(),
545 vec![
546 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
547 ColumnarValue::Scalar(ScalarValue::from(-3i64)),
548 ],
549 Ok(Some("alphabet")),
550 &str,
551 Utf8,
552 StringArray
553 );
554 test_function!(
555 SubstrFunc::new(),
556 vec![
557 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
558 ColumnarValue::Scalar(ScalarValue::from(30i64)),
559 ],
560 Ok(Some("")),
561 &str,
562 Utf8,
563 StringArray
564 );
565 test_function!(
566 SubstrFunc::new(),
567 vec![
568 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
569 ColumnarValue::Scalar(ScalarValue::Int64(None)),
570 ],
571 Ok(None),
572 &str,
573 Utf8,
574 StringArray
575 );
576 test_function!(
577 SubstrFunc::new(),
578 vec![
579 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
580 ColumnarValue::Scalar(ScalarValue::from(3i64)),
581 ColumnarValue::Scalar(ScalarValue::from(2i64)),
582 ],
583 Ok(Some("ph")),
584 &str,
585 Utf8,
586 StringArray
587 );
588 test_function!(
589 SubstrFunc::new(),
590 vec![
591 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
592 ColumnarValue::Scalar(ScalarValue::from(3i64)),
593 ColumnarValue::Scalar(ScalarValue::from(20i64)),
594 ],
595 Ok(Some("phabet")),
596 &str,
597 Utf8,
598 StringArray
599 );
600 test_function!(
601 SubstrFunc::new(),
602 vec![
603 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
604 ColumnarValue::Scalar(ScalarValue::from(0i64)),
605 ColumnarValue::Scalar(ScalarValue::from(5i64)),
606 ],
607 Ok(Some("alph")),
608 &str,
609 Utf8,
610 StringArray
611 );
612 test_function!(
614 SubstrFunc::new(),
615 vec![
616 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
617 ColumnarValue::Scalar(ScalarValue::from(-5i64)),
618 ColumnarValue::Scalar(ScalarValue::from(10i64)),
619 ],
620 Ok(Some("alph")),
621 &str,
622 Utf8,
623 StringArray
624 );
625 test_function!(
627 SubstrFunc::new(),
628 vec![
629 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
630 ColumnarValue::Scalar(ScalarValue::from(-5i64)),
631 ColumnarValue::Scalar(ScalarValue::from(4i64)),
632 ],
633 Ok(Some("")),
634 &str,
635 Utf8,
636 StringArray
637 );
638 test_function!(
640 SubstrFunc::new(),
641 vec![
642 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
643 ColumnarValue::Scalar(ScalarValue::from(-5i64)),
644 ColumnarValue::Scalar(ScalarValue::from(5i64)),
645 ],
646 Ok(Some("")),
647 &str,
648 Utf8,
649 StringArray
650 );
651 test_function!(
652 SubstrFunc::new(),
653 vec![
654 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
655 ColumnarValue::Scalar(ScalarValue::Int64(None)),
656 ColumnarValue::Scalar(ScalarValue::from(20i64)),
657 ],
658 Ok(None),
659 &str,
660 Utf8,
661 StringArray
662 );
663 test_function!(
664 SubstrFunc::new(),
665 vec![
666 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
667 ColumnarValue::Scalar(ScalarValue::from(3i64)),
668 ColumnarValue::Scalar(ScalarValue::Int64(None)),
669 ],
670 Ok(None),
671 &str,
672 Utf8,
673 StringArray
674 );
675 test_function!(
676 SubstrFunc::new(),
677 vec![
678 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
679 ColumnarValue::Scalar(ScalarValue::from(1i64)),
680 ColumnarValue::Scalar(ScalarValue::from(-1i64)),
681 ],
682 exec_err!("negative count not allowed: -1"),
683 &str,
684 Utf8,
685 StringArray
686 );
687 test_function!(
688 SubstrFunc::new(),
689 vec![
690 ColumnarValue::Scalar(ScalarValue::from("joséésoj")),
691 ColumnarValue::Scalar(ScalarValue::from(5i64)),
692 ColumnarValue::Scalar(ScalarValue::from(2i64)),
693 ],
694 Ok(Some("és")),
695 &str,
696 Utf8,
697 StringArray
698 );
699 #[cfg(not(feature = "unicode_expressions"))]
700 test_function!(
701 SubstrFunc::new(),
702 &[
703 ColumnarValue::Scalar(ScalarValue::from("alphabet")),
704 ColumnarValue::Scalar(ScalarValue::from(0i64)),
705 ],
706 internal_err!(
707 "function substr requires compilation with feature flag: unicode_expressions."
708 ),
709 &str,
710 Utf8,
711 StringArray
712 );
713 test_function!(
714 SubstrFunc::new(),
715 vec![
716 ColumnarValue::Scalar(ScalarValue::from("abc")),
717 ColumnarValue::Scalar(ScalarValue::from(i64::MIN)),
718 ],
719 exec_err!("start position overflow: -9223372036854775808"),
720 &str,
721 Utf8,
722 StringArray
723 );
724 test_function!(
725 SubstrFunc::new(),
726 vec![
727 ColumnarValue::Scalar(ScalarValue::from("overflow")),
728 ColumnarValue::Scalar(ScalarValue::from(i64::MIN)),
729 ColumnarValue::Scalar(ScalarValue::from(1i64)),
730 ],
731 exec_err!("start position overflow: -9223372036854775808"),
732 &str,
733 Utf8,
734 StringArray
735 );
736 test_function!(
737 SubstrFunc::new(),
738 vec![
739 ColumnarValue::Scalar(ScalarValue::from("large count")),
740 ColumnarValue::Scalar(ScalarValue::from(2i64)),
741 ColumnarValue::Scalar(ScalarValue::from(i64::MAX)),
742 ],
743 Ok(Some("arge count")),
744 &str,
745 Utf8,
746 StringArray
747 );
748
749 Ok(())
750 }
751
752 #[test]
753 fn test_sliced_string_array_array_args() -> Result<()> {
754 let string_array = Arc::new(StringArray::from(vec![
755 "skipped_prefix_value",
756 "alphabet_long_string",
757 "joséésojanother_long",
758 ])) as ArrayRef;
759 let string_array = string_array.slice(1, 2);
760 let start_array = Arc::new(Int64Array::from(vec![3, 5])) as ArrayRef;
761 let count_array = Arc::new(Int64Array::from(vec![15, 14])) as ArrayRef;
762
763 let result = super::substr(&[string_array, start_array, count_array])?;
764 let result = result.as_string::<i32>();
765
766 assert_eq!(result.value(0), "phabet_long_str");
767 assert_eq!(result.value(1), "ésojanother_lo");
768
769 Ok(())
770 }
771}