1use crate::binaries::{
19 ConcatBinaryBuilder, ConcatBinaryViewBuilder, ConcatLargeBinaryBuilder,
20};
21use crate::string::concat;
22use crate::strings::{
23 ColumnarValueRef, ConcatBuilder, ConcatLargeStringBuilder, ConcatStringBuilder,
24 ConcatStringViewBuilder, widest_binary_type, widest_string_type,
25};
26use arrow::array::Array;
27use arrow::datatypes::DataType;
28use datafusion_common::{
29 Result, ScalarValue, exec_datafusion_err, internal_err, plan_err,
30};
31use datafusion_expr::expr::ScalarFunction;
32use datafusion_expr::simplify::{ExprSimplifyResult, SimplifyContext};
33use datafusion_expr::{ColumnarValue, Documentation, Expr, Volatility, lit};
34use datafusion_expr::{ScalarFunctionArgs, ScalarUDFImpl, Signature};
35use datafusion_macros::user_doc;
36
37#[user_doc(
38 doc_section(label = "String Functions"),
39 description = "Concatenates multiple strings together.",
40 syntax_example = "concat(str[, ..., str_n])",
41 sql_example = r#"```sql
42> select concat('data', 'f', 'us', 'ion');
43+-------------------------------------------------------+
44| concat(Utf8("data"),Utf8("f"),Utf8("us"),Utf8("ion")) |
45+-------------------------------------------------------+
46| datafusion |
47+-------------------------------------------------------+
48```"#,
49 standard_argument(name = "str", prefix = "String"),
50 argument(
51 name = "str_n",
52 description = "Subsequent string expressions to concatenate."
53 ),
54 related_udf(name = "concat_ws")
55)]
56#[derive(Debug, PartialEq, Eq, Hash)]
57pub struct ConcatFunc {
58 signature: Signature,
59}
60
61impl Default for ConcatFunc {
62 fn default() -> Self {
63 ConcatFunc::new()
64 }
65}
66
67impl ConcatFunc {
68 pub fn new() -> Self {
69 Self {
70 signature: Signature::user_defined(Volatility::Immutable),
74 }
75 }
76}
77
78impl ScalarUDFImpl for ConcatFunc {
82 fn name(&self) -> &str {
83 "concat"
84 }
85
86 fn signature(&self) -> &Signature {
87 &self.signature
88 }
89
90 fn coerce_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
92 if arg_types.is_empty() {
93 plan_err!("concat does not support zero arguments")
94 } else {
95 coerce_arg_types(arg_types)
96 }
97 }
98
99 fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
103 Ok(deduce_return_type(arg_types))
104 }
105
106 fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
109 let return_datatype = args.return_type().clone();
110 let ScalarFunctionArgs { args, .. } = args;
111
112 let array_len = args.iter().find_map(|x| match x {
113 ColumnarValue::Array(array) => Some(array.len()),
114 _ => None,
115 });
116
117 if array_len.is_none() {
119 let mut values: Vec<&[u8]> = Vec::with_capacity(args.len());
120 for arg in &args {
121 let ColumnarValue::Scalar(scalar) = arg else {
122 return internal_err!("concat expected scalar value, got {arg:?}");
123 };
124 if let ScalarValue::Binary(Some(value)) = scalar {
125 values.push(value);
126 } else if let ScalarValue::LargeBinary(Some(value)) = scalar {
127 values.push(value);
128 } else if let ScalarValue::BinaryView(Some(value)) = scalar {
129 values.push(value);
130 } else if scalar.is_null() {
131 } else {
133 match scalar.try_as_str() {
135 Some(Some(v)) => values.push(v.as_bytes()),
136 Some(None) => {} None => plan_err!(
138 "Concat function does not support scalar type {}",
139 scalar
140 )?,
141 }
142 }
143 }
144 let concat_bytes = values.concat();
145
146 return match return_datatype {
147 DataType::Utf8View => {
148 let result = std::str::from_utf8(&concat_bytes)
149 .map_err(|_| {
150 exec_datafusion_err!("invalid UTF-8 in binary literal")
151 })?
152 .to_string();
153 Ok(ColumnarValue::Scalar(ScalarValue::Utf8View(Some(result))))
154 }
155 DataType::Utf8 => {
156 let result = std::str::from_utf8(&concat_bytes)
157 .map_err(|_| {
158 exec_datafusion_err!("invalid UTF-8 in binary literal")
159 })?
160 .to_string();
161 Ok(ColumnarValue::Scalar(ScalarValue::Utf8(Some(result))))
162 }
163 DataType::LargeUtf8 => {
164 let result = std::str::from_utf8(&concat_bytes)
165 .map_err(|_| {
166 exec_datafusion_err!("invalid UTF-8 in binary literal")
167 })?
168 .to_string();
169 Ok(ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some(result))))
170 }
171 DataType::Binary => Ok(ColumnarValue::Scalar(ScalarValue::Binary(Some(
172 concat_bytes,
173 )))),
174 DataType::LargeBinary => Ok(ColumnarValue::Scalar(
176 ScalarValue::LargeBinary(Some(concat_bytes)),
177 )),
178 DataType::BinaryView => Ok(ColumnarValue::Scalar(
179 ScalarValue::BinaryView(Some(concat_bytes)),
180 )),
181 other => {
182 plan_err!("Concat function does not support datatype of {other}")
183 }
184 };
185 }
186
187 let len = array_len.unwrap();
189 let mut data_size = 0;
190 let mut columns = Vec::with_capacity(args.len());
191
192 for arg in &args {
193 if let Some(column) =
194 ColumnarValueRef::from_columnar_value(arg, &mut data_size, len, 1, false)?
195 {
196 columns.push(column);
197 }
198 }
199
200 match return_datatype {
201 DataType::Utf8 => build_concat(
202 ConcatStringBuilder::with_capacity(len, data_size),
203 &columns,
204 len,
205 ),
206 DataType::Utf8View => build_concat(
207 ConcatStringViewBuilder::with_capacity(len, data_size),
208 &columns,
209 len,
210 ),
211 DataType::LargeUtf8 => build_concat(
212 ConcatLargeStringBuilder::with_capacity(len, data_size),
213 &columns,
214 len,
215 ),
216 DataType::Binary => build_concat(
217 ConcatBinaryBuilder::with_capacity(len, data_size),
218 &columns,
219 len,
220 ),
221 DataType::LargeBinary => build_concat(
223 ConcatLargeBinaryBuilder::with_capacity(len, data_size),
224 &columns,
225 len,
226 ),
227 DataType::BinaryView => build_concat(
228 ConcatBinaryViewBuilder::with_capacity(len, data_size),
229 &columns,
230 len,
231 ),
232 _ => unreachable!("concat"),
233 }
234 }
235
236 fn simplify(
245 &self,
246 args: Vec<Expr>,
247 _info: &SimplifyContext,
248 ) -> Result<ExprSimplifyResult> {
249 simplify_concat(args)
250 }
251
252 fn documentation(&self) -> Option<&Documentation> {
253 self.doc()
254 }
255}
256
257pub(crate) fn deduce_return_type(arg_types: &[DataType]) -> DataType {
258 use DataType::*;
259 if arg_types.contains(&BinaryView) {
260 BinaryView
261 } else if arg_types.contains(&LargeBinary) {
262 LargeBinary
264 } else if arg_types.contains(&Binary) {
265 Binary
266 } else if arg_types.contains(&Utf8View) {
267 Utf8View
268 } else if arg_types.contains(&LargeUtf8) {
269 LargeUtf8
270 } else {
271 Utf8
272 }
273}
274
275pub(crate) fn coerce_arg_types(arg_types: &[DataType]) -> Result<Vec<DataType>> {
277 let has_binary = arg_types.iter().any(|dt| dt.is_binary());
278 let has_string = arg_types.iter().any(|dt| dt.is_string());
279 if has_binary && has_string {
280 Ok(vec![widest_string_type(arg_types); arg_types.len()])
283 } else if has_binary {
284 Ok(vec![widest_binary_type(arg_types); arg_types.len()])
286 } else {
287 Ok(vec![widest_string_type(arg_types); arg_types.len()])
289 }
290}
291
292fn build_concat<B: ConcatBuilder>(
294 mut builder: B,
295 columns: &[ColumnarValueRef],
296 len: usize,
297) -> Result<ColumnarValue> {
298 for i in 0..len {
299 for column in columns {
300 builder.write::<true>(column, i)?;
301 }
302 builder.append_offset()?;
303 }
304
305 let array = builder.finish(None)?;
306 Ok(ColumnarValue::Array(array))
307}
308
309pub(crate) fn simplify_concat(args: Vec<Expr>) -> Result<ExprSimplifyResult> {
310 for arg in &args {
313 match arg {
314 Expr::Literal(dt, _) if dt.data_type().is_binary() => {
315 return Ok(ExprSimplifyResult::Original(args));
316 }
317 _ => {}
318 }
319 }
320
321 let mut new_args = Vec::with_capacity(args.len());
322 let mut contiguous_scalar = "".to_string();
323
324 let return_type = {
325 let data_types: Vec<_> = args
326 .iter()
327 .filter_map(|expr| match expr {
328 Expr::Literal(l, _) => Some(l.data_type()),
329 _ => None,
330 })
331 .collect();
332 ConcatFunc::new().return_type(&data_types)
333 }?;
334
335 for arg in args.clone() {
336 match arg {
337 Expr::Literal(ScalarValue::Utf8(None), _) => {}
338 Expr::Literal(ScalarValue::LargeUtf8(None), _) => {}
339 Expr::Literal(ScalarValue::Utf8View(None), _) => {}
340
341 Expr::Literal(ScalarValue::Utf8(Some(v)), _) => {
345 contiguous_scalar += &v;
346 }
347 Expr::Literal(ScalarValue::LargeUtf8(Some(v)), _) => {
348 contiguous_scalar += &v;
349 }
350 Expr::Literal(ScalarValue::Utf8View(Some(v)), _) => {
351 contiguous_scalar += &v;
352 }
353
354 Expr::Literal(x, _) => {
355 return internal_err!(
356 "The scalar {x} should be casted to string type during the type coercion."
357 );
358 }
359 arg => {
363 if !contiguous_scalar.is_empty() {
364 match return_type {
365 DataType::Utf8 => new_args.push(lit(contiguous_scalar)),
366 DataType::LargeUtf8 => new_args
367 .push(lit(ScalarValue::LargeUtf8(Some(contiguous_scalar)))),
368 DataType::Utf8View => new_args
369 .push(lit(ScalarValue::Utf8View(Some(contiguous_scalar)))),
370 _ => unreachable!(),
371 }
372 contiguous_scalar = "".to_string();
373 }
374 new_args.push(arg);
375 }
376 }
377 }
378
379 if !contiguous_scalar.is_empty() {
380 match return_type {
381 DataType::Utf8 => new_args.push(lit(contiguous_scalar)),
382 DataType::LargeUtf8 => {
383 new_args.push(lit(ScalarValue::LargeUtf8(Some(contiguous_scalar))))
384 }
385 DataType::Utf8View => {
386 new_args.push(lit(ScalarValue::Utf8View(Some(contiguous_scalar))))
387 }
388 _ => unreachable!(),
389 }
390 }
391
392 if !args.eq(&new_args) {
393 Ok(ExprSimplifyResult::Simplified(Expr::ScalarFunction(
394 ScalarFunction {
395 func: concat(),
396 args: new_args,
397 },
398 )))
399 } else {
400 Ok(ExprSimplifyResult::Original(args))
401 }
402}
403
404#[cfg(test)]
405mod tests {
406 use super::*;
407 use crate::utils::test::test_function;
408 use DataType::*;
409 use arrow::array::{
410 ArrayRef, BinaryArray, BinaryViewArray, LargeBinaryArray, StringArray,
411 };
412 use arrow::array::{LargeStringArray, StringViewArray};
413 use arrow::datatypes::Field;
414 use datafusion_common::config::ConfigOptions;
415 use std::sync::Arc;
416
417 #[test]
418 fn test_functions() -> Result<()> {
419 test_function!(
420 ConcatFunc::new(),
421 vec![
422 ColumnarValue::Scalar(ScalarValue::from("aa")),
423 ColumnarValue::Scalar(ScalarValue::from("bb")),
424 ColumnarValue::Scalar(ScalarValue::from("cc")),
425 ],
426 Ok(Some("aabbcc")),
427 &str,
428 Utf8,
429 StringArray
430 );
431 test_function!(
432 ConcatFunc::new(),
433 vec![
434 ColumnarValue::Scalar(ScalarValue::from("aa")),
435 ColumnarValue::Scalar(ScalarValue::Utf8(None)),
436 ColumnarValue::Scalar(ScalarValue::from("cc")),
437 ],
438 Ok(Some("aacc")),
439 &str,
440 Utf8,
441 StringArray
442 );
443 test_function!(
444 ConcatFunc::new(),
445 vec![ColumnarValue::Scalar(ScalarValue::Utf8(None))],
446 Ok(Some("")),
447 &str,
448 Utf8,
449 StringArray
450 );
451 test_function!(
452 ConcatFunc::new(),
453 vec![
454 ColumnarValue::Scalar(ScalarValue::from("aa")),
455 ColumnarValue::Scalar(ScalarValue::Utf8View(None)),
456 ColumnarValue::Scalar(ScalarValue::LargeUtf8(None)),
457 ColumnarValue::Scalar(ScalarValue::from("cc")),
458 ],
459 Ok(Some("aacc")),
460 &str,
461 Utf8View,
462 StringViewArray
463 );
464 test_function!(
465 ConcatFunc::new(),
466 vec![
467 ColumnarValue::Scalar(ScalarValue::from("aa")),
468 ColumnarValue::Scalar(ScalarValue::LargeUtf8(None)),
469 ColumnarValue::Scalar(ScalarValue::from("cc")),
470 ],
471 Ok(Some("aacc")),
472 &str,
473 LargeUtf8,
474 LargeStringArray
475 );
476 test_function!(
477 ConcatFunc::new(),
478 vec![
479 ColumnarValue::Scalar(ScalarValue::Utf8View(Some("aa".to_string()))),
480 ColumnarValue::Scalar(ScalarValue::Utf8(Some("cc".to_string()))),
481 ],
482 Ok(Some("aacc")),
483 &str,
484 Utf8View,
485 StringViewArray
486 );
487 Ok(())
488 }
489
490 #[test]
491 fn test_scalar_binary() -> Result<()> {
492 test_function!(
493 ConcatFunc::new(),
494 vec![
495 ColumnarValue::Scalar(ScalarValue::Binary(Some(
496 "Café".as_bytes().into()
497 ))),
498 ColumnarValue::Scalar(ScalarValue::Binary(Some("cc".as_bytes().into()))),
499 ],
500 Ok(Some("Cafécc".as_bytes())),
501 &[u8],
502 Binary,
503 BinaryArray
504 );
505 test_function!(
506 ConcatFunc::new(),
507 vec![
508 ColumnarValue::Scalar(ScalarValue::Binary(Some(
509 "Café".as_bytes().into()
510 ))),
511 ColumnarValue::Scalar(ScalarValue::LargeBinary(Some(
512 "cc".as_bytes().into()
513 ))),
514 ],
515 Ok(Some("Cafécc".as_bytes())),
516 &[u8],
517 LargeBinary,
518 LargeBinaryArray
519 );
520 test_function!(
521 ConcatFunc::new(),
522 vec![
523 ColumnarValue::Scalar(ScalarValue::Binary(Some(
524 "Café".as_bytes().into()
525 ))),
526 ColumnarValue::Scalar(ScalarValue::BinaryView(Some(
527 "cc".as_bytes().into()
528 ))),
529 ],
530 Ok(Some("Cafécc".as_bytes())),
531 &[u8],
532 BinaryView,
533 BinaryViewArray
534 );
535 test_function!(
536 ConcatFunc::new(),
537 vec![
538 ColumnarValue::Scalar(ScalarValue::BinaryView(Some(
539 "Café".as_bytes().into()
540 ))),
541 ColumnarValue::Scalar(ScalarValue::BinaryView(Some(
542 "cc".as_bytes().into()
543 ))),
544 ],
545 Ok(Some("Cafécc".as_bytes())),
546 &[u8],
547 BinaryView,
548 BinaryViewArray
549 );
550 test_function!(
552 ConcatFunc::new(),
553 vec![
554 ColumnarValue::Scalar(ScalarValue::Binary(None)),
555 ColumnarValue::Scalar(ScalarValue::Binary(Some(b"hello".to_vec()))),
556 ],
557 Ok(Some(b"hello".as_ref())),
558 &[u8],
559 Binary,
560 BinaryArray
561 );
562 test_function!(
564 ConcatFunc::new(),
565 vec![ColumnarValue::Scalar(ScalarValue::Binary(None))],
566 Ok(Some(b"".as_ref())),
567 &[u8],
568 Binary,
569 BinaryArray
570 );
571 Ok(())
572 }
573
574 #[test]
575 fn test_array_string() -> Result<()> {
576 let c0 =
577 ColumnarValue::Array(Arc::new(StringArray::from(vec!["foo", "bar", "baz"])));
578 let c1 = ColumnarValue::Scalar(ScalarValue::Utf8(Some(",".to_string())));
579 let c2 = ColumnarValue::Array(Arc::new(StringArray::from(vec![
580 Some("x"),
581 None,
582 Some("z"),
583 ])));
584 let c3 = ColumnarValue::Scalar(ScalarValue::Utf8View(Some(",".to_string())));
585 let c4 = ColumnarValue::Array(Arc::new(StringViewArray::from(vec![
586 Some("a"),
587 None,
588 Some("b"),
589 ])));
590 let arg_fields = vec![
591 Field::new("a", Utf8, true),
592 Field::new("a", Utf8, true),
593 Field::new("a", Utf8, true),
594 Field::new("a", Utf8View, true),
595 Field::new("a", Utf8View, true),
596 ]
597 .into_iter()
598 .map(Arc::new)
599 .collect::<Vec<_>>();
600
601 let args = ScalarFunctionArgs {
602 args: vec![c0, c1, c2, c3, c4],
603 arg_fields,
604 number_rows: 3,
605 return_field: Field::new("f", Utf8View, true).into(),
606 config_options: Arc::new(ConfigOptions::default()),
607 };
608
609 let result = ConcatFunc::new().invoke_with_args(args)?;
610 let expected =
611 Arc::new(StringViewArray::from(vec!["foo,x,a", "bar,,", "baz,z,b"]))
612 as ArrayRef;
613 match &result {
614 ColumnarValue::Array(array) => {
615 assert_eq!(&expected, array);
616 }
617 _ => panic!(),
618 }
619 Ok(())
620 }
621
622 #[test]
623 fn test_array_binary() -> Result<()> {
624 let c0 = ColumnarValue::Array(Arc::new(BinaryArray::from_vec(vec![
625 b"foo", b"bar", b"baz",
626 ])));
627 let c1 = ColumnarValue::Scalar(ScalarValue::LargeBinary(Some(b",".to_vec())));
628 let c2 = ColumnarValue::Array(Arc::new(BinaryArray::from_opt_vec(vec![
629 Some(b"x"),
630 None,
631 Some(b"z"),
632 ])));
633 let c3 = ColumnarValue::Scalar(ScalarValue::BinaryView(Some(b",".to_vec())));
634 let c4 = ColumnarValue::Array(Arc::new(BinaryViewArray::from_iter(vec![
635 Some(b"a"),
636 None,
637 Some(b"b"),
638 ])));
639 let arg_fields = vec![
640 Field::new("a", Binary, true),
641 Field::new("a", LargeBinary, true),
642 Field::new("a", Binary, true),
643 Field::new("a", BinaryView, true),
644 Field::new("a", BinaryView, true),
645 ]
646 .into_iter()
647 .map(Arc::new)
648 .collect::<Vec<_>>();
649
650 let args = ScalarFunctionArgs {
651 args: vec![c0, c1, c2, c3, c4],
652 arg_fields,
653 number_rows: 3,
654 return_field: Field::new("f", BinaryView, true).into(),
655 config_options: Arc::new(ConfigOptions::default()),
656 };
657
658 let result = ConcatFunc::new().invoke_with_args(args)?;
659 let expected = Arc::new(BinaryViewArray::from_iter(vec![
660 Some(b"foo,x,a".to_vec()),
661 Some(b"bar,,".to_vec()),
662 Some(b"baz,z,b".to_vec()),
663 ])) as ArrayRef;
664 match &result {
665 ColumnarValue::Array(array) => {
666 assert_eq!(&expected, array);
667 }
668 _ => panic!(),
669 }
670 Ok(())
671 }
672}