Skip to main content

datafusion_functions/string/
concat.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
18use 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            // Use `Signature::UserDefined` to allow different argument types.
71            // `Variadic` requires every argument to be coerced to the same string type,
72            // so the UDF cannot distinguish between binary and string inputs.
73            signature: Signature::user_defined(Volatility::Immutable),
74        }
75    }
76}
77
78// Supports string + string concatenation, binary + binary concatenation,
79// and mixed string + binary concatenation (binary is coerced to the widest
80// string type).
81impl ScalarUDFImpl for ConcatFunc {
82    fn name(&self) -> &str {
83        "concat"
84    }
85
86    fn signature(&self) -> &Signature {
87        &self.signature
88    }
89
90    /// Coerce all arguments to the widest type within the binary / string family
91    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    /// mixed inputs, prefer Utf8View; prefer LargeUtf8 over Utf8 to avoid
100    /// potential overflow on LargeUtf8 input.
101    /// For binaries, use the similar hierarchy
102    fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
103        Ok(deduce_return_type(arg_types))
104    }
105
106    /// Concatenates the text representations of all the arguments. NULL arguments are ignored.
107    /// concat('abcde', 2, NULL, 22) = 'abcde222'
108    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        // Scalar
118        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                    // null binary scalar: skip (consistent with null string behaviour)
132                } else {
133                    // String case
134                    match scalar.try_as_str() {
135                        Some(Some(v)) => values.push(v.as_bytes()),
136                        Some(None) => {} // null literal
137                        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                // Serves LargeBinary and FixedSizeBinary inputs
175                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        // Array
188        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            // Serves LargeBinary and FixedSizeBinary inputs
222            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    /// Simplify the `concat` function by
237    /// 1. filtering out all `null` literals
238    /// 2. concatenating contiguous literal arguments
239    ///
240    /// For example:
241    /// `concat(col(a), 'hello ', 'world', col(b), null)`
242    /// will be optimized to
243    /// `concat(col(a), 'hello world', col(b))`
244    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        // Serves LargeBinary and FixedSizeBinary inputs
263        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
275/// Coerce all arguments to the widest type within the binary / string family
276pub(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        // Mixed string+binary: coerce everything to the widest string type
281        // This behaviour is seen for Spark, DuckDB
282        Ok(vec![widest_string_type(arg_types); arg_types.len()])
283    } else if has_binary {
284        // Pure binary+binary concatenation: coerce to the widest binary type
285        Ok(vec![widest_binary_type(arg_types); arg_types.len()])
286    } else {
287        // Pure string+string concatenation: coerce to the widest string type
288        Ok(vec![widest_string_type(arg_types); arg_types.len()])
289    }
290}
291
292/// Build a `concats` output array using a generic [`ConcatBuilder`].
293fn 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    // Skip simplification when binary literals are present, because it
311    // handles only strings
312    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            // filter out `null` args
342            // All literals have been converted to Utf8 or LargeUtf8 in type_coercion.
343            // Concatenate it with the `contiguous_scalar`.
344            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            // If the arg is not a literal, we should first push the current `contiguous_scalar`
360            // to the `new_args` (if it is not empty) and reset it to empty string.
361            // Then pushing this arg to the `new_args`.
362            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        // Skip one Binary(None)
551        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        // Skip all Binary(None), producing an empty array
563        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}