Skip to main content

datafusion_functions/string/
concat_ws.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 arrow::array::Array;
19use arrow::datatypes::DataType;
20
21use crate::binaries::{
22    ConcatBinaryBuilder, ConcatBinaryViewBuilder, ConcatLargeBinaryBuilder,
23};
24use crate::string::concat;
25use crate::string::concat::{coerce_arg_types, deduce_return_type, simplify_concat};
26use crate::string::concat_ws;
27use crate::strings::{
28    ColumnarValueRef, ConcatBuilder, ConcatLargeStringBuilder, ConcatStringBuilder,
29    ConcatStringViewBuilder,
30};
31use datafusion_common::{Result, ScalarValue, exec_err, internal_err, plan_err};
32use datafusion_expr::expr::ScalarFunction;
33use datafusion_expr::simplify::{ExprSimplifyResult, SimplifyContext};
34use datafusion_expr::{ColumnarValue, Documentation, Expr, Volatility, lit};
35use datafusion_expr::{ScalarFunctionArgs, ScalarUDFImpl, Signature};
36use datafusion_macros::user_doc;
37
38#[user_doc(
39    doc_section(label = "String Functions"),
40    description = "Concatenates multiple strings together with a specified separator.",
41    syntax_example = "concat_ws(separator, str[, ..., str_n])",
42    sql_example = r#"```sql
43> select concat_ws('_', 'data', 'fusion');
44+--------------------------------------------------+
45| concat_ws(Utf8("_"),Utf8("data"),Utf8("fusion")) |
46+--------------------------------------------------+
47| data_fusion                                      |
48+--------------------------------------------------+
49```"#,
50    argument(
51        name = "separator",
52        description = "Separator to insert between concatenated strings."
53    ),
54    argument(
55        name = "str",
56        description = "String expression to operate on. Can be a constant, column, or function, and any combination of operators."
57    ),
58    argument(
59        name = "str_n",
60        description = "Subsequent string expressions to concatenate."
61    ),
62    related_udf(name = "concat")
63)]
64#[derive(Debug, PartialEq, Eq, Hash)]
65pub struct ConcatWsFunc {
66    signature: Signature,
67}
68
69impl Default for ConcatWsFunc {
70    fn default() -> Self {
71        ConcatWsFunc::new()
72    }
73}
74
75impl ConcatWsFunc {
76    pub fn new() -> Self {
77        Self {
78            // Use `Signature::UserDefined` to allow different argument types.
79            // `Variadic` requires every argument to be coerced to the same string type,
80            // so the UDF cannot distinguish between binary and string inputs.
81            signature: Signature::user_defined(Volatility::Immutable),
82        }
83    }
84}
85
86impl ScalarUDFImpl for ConcatWsFunc {
87    fn name(&self) -> &str {
88        "concat_ws"
89    }
90
91    fn signature(&self) -> &Signature {
92        &self.signature
93    }
94
95    /// Coerce all arguments to the widest type within the binary / string family
96    fn coerce_types(&self, arg_types: &[DataType]) -> Result<Vec<DataType>> {
97        if arg_types.len() < 2 {
98            plan_err!(
99                "concat_ws expects at least 2 arguments, got {}",
100                arg_types.len()
101            )
102        } else {
103            coerce_arg_types(arg_types)
104        }
105    }
106
107    /// Match the return type to the input types. Delegates to `concat` implementation.
108    fn return_type(&self, arg_types: &[DataType]) -> Result<DataType> {
109        Ok(deduce_return_type(arg_types))
110    }
111
112    /// Concatenates all but the first argument, with separators. The first
113    /// argument is used as the separator string, and should not be NULL. Other
114    /// NULL arguments are ignored.
115    /// concat_ws(',', 'abcde', 2, NULL, 22) = 'abcde,2,22'
116    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result<ColumnarValue> {
117        let return_datatype = args.return_type().clone();
118        let ScalarFunctionArgs { args, .. } = args;
119
120        if args.len() < 2 {
121            return exec_err!(
122                "concat_ws was called with {} arguments. It requires at least 2.",
123                args.len()
124            );
125        }
126
127        let arg_types: Vec<DataType> = args.iter().map(|c| c.data_type()).collect();
128
129        let with_binary = arg_types.iter().any(|dt| dt.is_binary());
130
131        let array_len = args.iter().find_map(|x| match x {
132            ColumnarValue::Array(array) => Some(array.len()),
133            _ => None,
134        });
135
136        // Scalar
137        if array_len.is_none() {
138            let ColumnarValue::Scalar(scalar) = &args[0] else {
139                unreachable!()
140            };
141
142            return if with_binary {
143                // Binary scalar path
144                let sep_bytes: &[u8] = match scalar {
145                    ScalarValue::Binary(Some(v))
146                    | ScalarValue::LargeBinary(Some(v))
147                    | ScalarValue::BinaryView(Some(v)) => v.as_slice(),
148                    ScalarValue::FixedSizeBinary(_, Some(v)) => v.as_slice(),
149                    scalar if scalar.is_null() => {
150                        return Ok(null_scalar(&return_datatype));
151                    }
152                    other => {
153                        return internal_err!("Expected binary separator, got {other:?}");
154                    }
155                };
156
157                let mut values: Vec<&[u8]> = Vec::with_capacity(args.len() - 1);
158                for arg in &args[1..] {
159                    let ColumnarValue::Scalar(s) = arg else {
160                        unreachable!()
161                    };
162                    match s {
163                        ScalarValue::Binary(Some(v))
164                        | ScalarValue::LargeBinary(Some(v))
165                        | ScalarValue::BinaryView(Some(v)) => values.push(v.as_slice()),
166                        ScalarValue::FixedSizeBinary(_, Some(v)) => {
167                            values.push(v.as_slice())
168                        }
169                        // skip null
170                        scalar if scalar.is_null() => {}
171                        other => {
172                            return internal_err!("Expected binary value, got {other:?}");
173                        }
174                    }
175                }
176                let result = values.join(sep_bytes);
177
178                match return_datatype {
179                    DataType::Binary => {
180                        Ok(ColumnarValue::Scalar(ScalarValue::Binary(Some(result))))
181                    }
182                    DataType::LargeBinary => Ok(ColumnarValue::Scalar(
183                        ScalarValue::LargeBinary(Some(result)),
184                    )),
185                    DataType::BinaryView => {
186                        Ok(ColumnarValue::Scalar(ScalarValue::BinaryView(Some(result))))
187                    }
188                    other => {
189                        plan_err!("concat_ws does not support return type {other}")
190                    }
191                }
192            } else {
193                // String scalar path
194                let sep = match scalar.try_as_str() {
195                    Some(Some(s)) => s,
196                    Some(None) => {
197                        return Ok(null_scalar(&return_datatype));
198                    }
199                    None => {
200                        return internal_err!("Expected string literal, got {scalar:?}");
201                    }
202                };
203
204                let mut values = Vec::with_capacity(args.len() - 1);
205                for arg in &args[1..] {
206                    let ColumnarValue::Scalar(scalar) = arg else {
207                        unreachable!()
208                    };
209
210                    match scalar.try_as_str() {
211                        Some(Some(v)) => values.push(v),
212                        Some(None) => {} // null literal string
213                        None => {
214                            return internal_err!(
215                                "Expected string literal, got {scalar:?}"
216                            );
217                        }
218                    }
219                }
220                let result = values.join(sep);
221
222                match return_datatype {
223                    DataType::Utf8View => {
224                        Ok(ColumnarValue::Scalar(ScalarValue::Utf8View(Some(result))))
225                    }
226                    DataType::LargeUtf8 => {
227                        Ok(ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some(result))))
228                    }
229                    DataType::Utf8 => {
230                        Ok(ColumnarValue::Scalar(ScalarValue::Utf8(Some(result))))
231                    }
232                    other => {
233                        plan_err!("concat_ws does not support return type {other}")
234                    }
235                }
236            };
237        }
238
239        // Array
240        let len = array_len.unwrap();
241        let mut data_size = 0;
242
243        let sep_column = &args[0];
244
245        // A null scalar separator makes the entire result null for all rows.
246        if matches!(sep_column, ColumnarValue::Scalar(s) if s.is_null()) {
247            return Ok(null_scalar(&return_datatype));
248        }
249
250        let sep: ColumnarValueRef = ColumnarValueRef::from_columnar_value(sep_column, &mut data_size, len, args.len() - 2, true)?
251            .map(Ok)
252            .unwrap_or_else(|| plan_err!(
253                "Input {sep_column} which is not a supported datatype for concat_ws separator"
254            ))?;
255
256        let mut columns = Vec::with_capacity(args.len() - 1);
257        for arg in &args[1..] {
258            if let Some(column) =
259                ColumnarValueRef::from_columnar_value(arg, &mut data_size, len, 1, false)?
260            {
261                columns.push(column);
262            }
263        }
264
265        match return_datatype {
266            DataType::Utf8 => build_concat_ws(
267                ConcatStringBuilder::with_capacity(len, data_size),
268                &sep,
269                &columns,
270                len,
271            ),
272            DataType::LargeUtf8 => build_concat_ws(
273                ConcatLargeStringBuilder::with_capacity(len, data_size),
274                &sep,
275                &columns,
276                len,
277            ),
278            DataType::Utf8View => build_concat_ws(
279                ConcatStringViewBuilder::with_capacity(len, data_size),
280                &sep,
281                &columns,
282                len,
283            ),
284            DataType::Binary => build_concat_ws(
285                ConcatBinaryBuilder::with_capacity(len, data_size),
286                &sep,
287                &columns,
288                len,
289            ),
290            DataType::LargeBinary => build_concat_ws(
291                ConcatLargeBinaryBuilder::with_capacity(len, data_size),
292                &sep,
293                &columns,
294                len,
295            ),
296            DataType::BinaryView => build_concat_ws(
297                ConcatBinaryViewBuilder::with_capacity(len, data_size),
298                &sep,
299                &columns,
300                len,
301            ),
302            other => plan_err!("concat_ws does not support return type {other}"),
303        }
304    }
305
306    /// Simply the `concat_ws` function by
307    /// 1. folding to `null` if the delimiter is null
308    /// 2. filtering out `null` arguments
309    /// 3. using `concat` to replace `concat_ws` if the delimiter is an empty string
310    /// 4. concatenating contiguous literals if the delimiter is a literal.
311    fn simplify(
312        &self,
313        args: Vec<Expr>,
314        _info: &SimplifyContext,
315    ) -> Result<ExprSimplifyResult> {
316        match &args[..] {
317            [delimiter, vals @ ..] => simplify_concat_ws(delimiter, vals),
318            _ => Ok(ExprSimplifyResult::Original(args)),
319        }
320    }
321
322    fn documentation(&self) -> Option<&Documentation> {
323        self.doc()
324    }
325}
326
327/// Build a `concat_ws` output array using a generic [`ConcatBuilder`].
328/// Write non-null column values per row, inserting the separator between them
329fn build_concat_ws<B: ConcatBuilder>(
330    mut builder: B,
331    sep: &ColumnarValueRef,
332    columns: &[ColumnarValueRef],
333    len: usize,
334) -> Result<ColumnarValue> {
335    for i in 0..len {
336        if !sep.is_valid(i) {
337            builder.append_offset()?;
338            continue;
339        }
340        let mut first = true;
341        for column in columns {
342            if column.is_valid(i) {
343                if !first {
344                    builder.write::<false>(sep, i)?;
345                }
346                builder.write::<false>(column, i)?;
347                first = false;
348            }
349        }
350        builder.append_offset()?;
351    }
352    let array = builder.finish(sep.nulls())?;
353    Ok(ColumnarValue::Array(array))
354}
355
356fn null_scalar(dt: &DataType) -> ColumnarValue {
357    ColumnarValue::Scalar(
358        ScalarValue::try_new_null(dt).unwrap_or(ScalarValue::Utf8(None)),
359    )
360}
361
362fn simplify_concat_ws(delimiter: &Expr, args: &[Expr]) -> Result<ExprSimplifyResult> {
363    // Preserve the delimiter's string type for any new literals produced
364    // during simplification.
365    let delimiter_type = match delimiter {
366        Expr::Literal(v, _) => v.data_type(),
367        _ => DataType::Utf8,
368    };
369
370    // Shortcut for binary delimiters
371    if delimiter_type.is_binary() {
372        let mut args = args
373            .iter()
374            .filter(|x| !is_null(x))
375            .cloned()
376            .collect::<Vec<_>>();
377        args.insert(0, delimiter.clone());
378        return Ok(ExprSimplifyResult::Original(args));
379    }
380
381    let typed_lit = |s: String| -> Expr {
382        match delimiter_type {
383            DataType::LargeUtf8 => lit(ScalarValue::LargeUtf8(Some(s))),
384            DataType::Utf8View => lit(ScalarValue::Utf8View(Some(s))),
385            _ => lit(s),
386        }
387    };
388
389    match delimiter {
390        Expr::Literal(
391            ScalarValue::Utf8(delimiter)
392            | ScalarValue::LargeUtf8(delimiter)
393            | ScalarValue::Utf8View(delimiter),
394            _,
395        ) => {
396            match delimiter {
397                // When the delimiter is the empty string, replace `concat_ws`
398                // with `concat`
399                Some(delimiter) if delimiter.is_empty() => {
400                    match simplify_concat(args.to_vec())? {
401                        ExprSimplifyResult::Original(_) => {
402                            Ok(ExprSimplifyResult::Simplified(Expr::ScalarFunction(
403                                ScalarFunction {
404                                    func: concat(),
405                                    args: args.to_vec(),
406                                },
407                            )))
408                        }
409                        expr => Ok(expr),
410                    }
411                }
412                Some(delimiter) => {
413                    let mut new_args = Vec::with_capacity(args.len());
414                    new_args.push(typed_lit(delimiter.to_string()));
415                    let mut contiguous_scalar = None;
416                    for arg in args {
417                        match arg {
418                            // filter out null args
419                            Expr::Literal(
420                                ScalarValue::Utf8(None)
421                                | ScalarValue::LargeUtf8(None)
422                                | ScalarValue::Utf8View(None),
423                                _,
424                            ) => {}
425                            Expr::Literal(
426                                ScalarValue::Utf8(Some(v))
427                                | ScalarValue::LargeUtf8(Some(v))
428                                | ScalarValue::Utf8View(Some(v)),
429                                _,
430                            ) => match contiguous_scalar {
431                                None => contiguous_scalar = Some(v.to_string()),
432                                Some(mut pre) => {
433                                    pre += delimiter;
434                                    pre += v;
435                                    contiguous_scalar = Some(pre)
436                                }
437                            },
438                            Expr::Literal(s, _) => {
439                                return internal_err!(
440                                    "The scalar {s} should be casted to string type during the type coercion."
441                                );
442                            }
443                            // If the arg is not a literal, we should first push the current `contiguous_scalar`
444                            // to the `new_args` and reset it to None.
445                            // Then pushing this arg to the `new_args`.
446                            arg => {
447                                if let Some(val) = contiguous_scalar {
448                                    new_args.push(typed_lit(val));
449                                }
450                                new_args.push(arg.clone());
451                                contiguous_scalar = None;
452                            }
453                        }
454                    }
455                    if let Some(val) = contiguous_scalar {
456                        new_args.push(typed_lit(val));
457                    }
458
459                    Ok(ExprSimplifyResult::Simplified(Expr::ScalarFunction(
460                        ScalarFunction {
461                            func: concat_ws(),
462                            args: new_args,
463                        },
464                    )))
465                }
466                // If the delimiter is null, then the value of the whole expression is null.
467                None => {
468                    let null_scalar = match delimiter_type {
469                        DataType::LargeUtf8 => ScalarValue::LargeUtf8(None),
470                        DataType::Utf8View => ScalarValue::Utf8View(None),
471                        _ => ScalarValue::Utf8(None),
472                    };
473                    Ok(ExprSimplifyResult::Simplified(Expr::Literal(
474                        null_scalar,
475                        None,
476                    )))
477                }
478            }
479        }
480        Expr::Literal(d, _) => internal_err!(
481            "The scalar {d} should be casted to string type during the type coercion."
482        ),
483        _ => {
484            let mut args = args
485                .iter()
486                .filter(|&x| !is_null(x))
487                .cloned()
488                .collect::<Vec<Expr>>();
489            args.insert(0, delimiter.clone());
490            Ok(ExprSimplifyResult::Original(args))
491        }
492    }
493}
494
495fn is_null(expr: &Expr) -> bool {
496    match expr {
497        Expr::Literal(v, _) => v.is_null(),
498        _ => false,
499    }
500}
501
502#[cfg(test)]
503mod tests {
504    use std::sync::Arc;
505
506    use crate::string::concat_ws::ConcatWsFunc;
507    use arrow::array::{
508        Array, ArrayRef, BinaryArray, LargeBinaryArray, LargeStringArray, StringArray,
509        StringViewArray,
510    };
511    use arrow::datatypes::DataType::{Binary, LargeBinary, LargeUtf8, Utf8, Utf8View};
512    use arrow::datatypes::Field;
513    use datafusion_common::Result;
514    use datafusion_common::ScalarValue;
515    use datafusion_common::config::ConfigOptions;
516    use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, ScalarUDFImpl};
517
518    use crate::utils::test::test_function;
519
520    #[test]
521    fn test_functions() -> Result<()> {
522        test_function!(
523            ConcatWsFunc::new(),
524            vec![
525                ColumnarValue::Scalar(ScalarValue::from("|")),
526                ColumnarValue::Scalar(ScalarValue::from("aa")),
527                ColumnarValue::Scalar(ScalarValue::from("bb")),
528                ColumnarValue::Scalar(ScalarValue::from("cc")),
529            ],
530            Ok(Some("aa|bb|cc")),
531            &str,
532            Utf8,
533            StringArray
534        );
535        test_function!(
536            ConcatWsFunc::new(),
537            vec![
538                ColumnarValue::Scalar(ScalarValue::from("|")),
539                ColumnarValue::Scalar(ScalarValue::Utf8(None)),
540            ],
541            Ok(Some("")),
542            &str,
543            Utf8,
544            StringArray
545        );
546        test_function!(
547            ConcatWsFunc::new(),
548            vec![
549                ColumnarValue::Scalar(ScalarValue::Utf8(None)),
550                ColumnarValue::Scalar(ScalarValue::from("aa")),
551                ColumnarValue::Scalar(ScalarValue::from("bb")),
552                ColumnarValue::Scalar(ScalarValue::from("cc")),
553            ],
554            Ok(None),
555            &str,
556            Utf8,
557            StringArray
558        );
559        test_function!(
560            ConcatWsFunc::new(),
561            vec![
562                ColumnarValue::Scalar(ScalarValue::from("|")),
563                ColumnarValue::Scalar(ScalarValue::from("aa")),
564                ColumnarValue::Scalar(ScalarValue::Utf8(None)),
565                ColumnarValue::Scalar(ScalarValue::from("cc")),
566            ],
567            Ok(Some("aa|cc")),
568            &str,
569            Utf8,
570            StringArray
571        );
572
573        Ok(())
574    }
575
576    #[test]
577    fn concat_ws() -> Result<()> {
578        // sep is scalar
579        let c0 = ColumnarValue::Scalar(ScalarValue::Utf8(Some(",".to_string())));
580        let c1 =
581            ColumnarValue::Array(Arc::new(StringArray::from(vec!["foo", "bar", "baz"])));
582        let c2 = ColumnarValue::Array(Arc::new(StringArray::from(vec![
583            Some("x"),
584            None,
585            Some("z"),
586        ])));
587
588        let arg_fields = vec![
589            Field::new("a", Utf8, true).into(),
590            Field::new("a", Utf8, true).into(),
591            Field::new("a", Utf8, true).into(),
592        ];
593        let args = ScalarFunctionArgs {
594            args: vec![c0, c1, c2],
595            arg_fields,
596            number_rows: 3,
597            return_field: Field::new("f", Utf8, true).into(),
598            config_options: Arc::new(ConfigOptions::default()),
599        };
600
601        let result = ConcatWsFunc::new().invoke_with_args(args)?;
602        let expected =
603            Arc::new(StringArray::from(vec!["foo,x", "bar", "baz,z"])) as ArrayRef;
604        match &result {
605            ColumnarValue::Array(array) => {
606                assert_eq!(&expected, array);
607            }
608            _ => panic!(),
609        }
610
611        // sep is nullable array
612        let c0 = ColumnarValue::Array(Arc::new(StringArray::from(vec![
613            Some(","),
614            None,
615            Some("+"),
616        ])));
617        let c1 =
618            ColumnarValue::Array(Arc::new(StringArray::from(vec!["foo", "bar", "baz"])));
619        let c2 = ColumnarValue::Array(Arc::new(StringArray::from(vec![
620            Some("x"),
621            Some("y"),
622            Some("z"),
623        ])));
624
625        let arg_fields = vec![
626            Field::new("a", Utf8, true).into(),
627            Field::new("a", Utf8, true).into(),
628            Field::new("a", Utf8, true).into(),
629        ];
630        let args = ScalarFunctionArgs {
631            args: vec![c0, c1, c2],
632            arg_fields,
633            number_rows: 3,
634            return_field: Field::new("f", Utf8, true).into(),
635            config_options: Arc::new(ConfigOptions::default()),
636        };
637
638        let result = ConcatWsFunc::new().invoke_with_args(args)?;
639        let expected =
640            Arc::new(StringArray::from(vec![Some("foo,x"), None, Some("baz+z")]))
641                as ArrayRef;
642        match &result {
643            ColumnarValue::Array(array) => {
644                assert_eq!(&expected, array);
645            }
646            _ => panic!(),
647        }
648
649        Ok(())
650    }
651
652    #[test]
653    fn concat_ws_utf8view_scalar_separator() -> Result<()> {
654        let c0 = ColumnarValue::Scalar(ScalarValue::Utf8View(Some(",".to_string())));
655        let c1 =
656            ColumnarValue::Array(Arc::new(StringArray::from(vec!["foo", "bar", "baz"])));
657        let c2 = ColumnarValue::Array(Arc::new(StringArray::from(vec![
658            Some("x"),
659            None,
660            Some("z"),
661        ])));
662
663        let arg_fields = vec![
664            Field::new("a", Utf8View, true).into(),
665            Field::new("a", Utf8, true).into(),
666            Field::new("a", Utf8, true).into(),
667        ];
668        let args = ScalarFunctionArgs {
669            args: vec![c0, c1, c2],
670            arg_fields,
671            number_rows: 3,
672            return_field: Field::new("f", Utf8View, true).into(),
673            config_options: Arc::new(ConfigOptions::default()),
674        };
675
676        let result = ConcatWsFunc::new().invoke_with_args(args)?;
677        let expected =
678            Arc::new(StringViewArray::from(vec!["foo,x", "bar", "baz,z"])) as ArrayRef;
679        match &result {
680            ColumnarValue::Array(array) => {
681                assert_eq!(&expected, array);
682            }
683            _ => panic!("Expected array result"),
684        }
685
686        Ok(())
687    }
688
689    #[test]
690    fn concat_ws_largeutf8_scalar_separator() -> Result<()> {
691        let c0 = ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some(",".to_string())));
692        let c1 =
693            ColumnarValue::Array(Arc::new(StringArray::from(vec!["foo", "bar", "baz"])));
694        let c2 = ColumnarValue::Array(Arc::new(StringArray::from(vec![
695            Some("x"),
696            None,
697            Some("z"),
698        ])));
699
700        let arg_fields = vec![
701            Field::new("a", LargeUtf8, true).into(),
702            Field::new("a", Utf8, true).into(),
703            Field::new("a", Utf8, true).into(),
704        ];
705        let args = ScalarFunctionArgs {
706            args: vec![c0, c1, c2],
707            arg_fields,
708            number_rows: 3,
709            return_field: Field::new("f", LargeUtf8, true).into(),
710            config_options: Arc::new(ConfigOptions::default()),
711        };
712
713        let result = ConcatWsFunc::new().invoke_with_args(args)?;
714        let expected =
715            Arc::new(LargeStringArray::from(vec!["foo,x", "bar", "baz,z"])) as ArrayRef;
716        match &result {
717            ColumnarValue::Array(array) => {
718                assert_eq!(&expected, array);
719            }
720            _ => panic!("Expected array result"),
721        }
722
723        Ok(())
724    }
725
726    #[test]
727    fn concat_ws_utf8view_nullable_separator() -> Result<()> {
728        let c0 = ColumnarValue::Array(Arc::new(StringViewArray::from(vec![
729            Some(","),
730            None,
731            Some("+"),
732        ])));
733        let c1 = ColumnarValue::Array(Arc::new(StringViewArray::from(vec![
734            "foo", "bar", "baz",
735        ])));
736        let c2 = ColumnarValue::Array(Arc::new(StringViewArray::from(vec![
737            Some("x"),
738            Some("y"),
739            Some("z"),
740        ])));
741
742        let arg_fields = vec![
743            Field::new("a", Utf8View, true).into(),
744            Field::new("a", Utf8View, true).into(),
745            Field::new("a", Utf8View, true).into(),
746        ];
747        let args = ScalarFunctionArgs {
748            args: vec![c0, c1, c2],
749            arg_fields,
750            number_rows: 3,
751            return_field: Field::new("f", Utf8View, true).into(),
752            config_options: Arc::new(ConfigOptions::default()),
753        };
754
755        let result = ConcatWsFunc::new().invoke_with_args(args)?;
756        let expected = Arc::new(StringViewArray::from(vec![
757            Some("foo,x"),
758            None,
759            Some("baz+z"),
760        ])) as ArrayRef;
761        match &result {
762            ColumnarValue::Array(array) => {
763                assert_eq!(&expected, array);
764            }
765            _ => panic!("Expected array result"),
766        }
767
768        Ok(())
769    }
770
771    #[test]
772    fn concat_ws_largeutf8_arrays() -> Result<()> {
773        let c0 = ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some(",".to_string())));
774        let c1 = ColumnarValue::Array(Arc::new(LargeStringArray::from(vec![
775            "foo", "bar", "baz",
776        ])));
777        let c2 = ColumnarValue::Array(Arc::new(LargeStringArray::from(vec![
778            Some("x"),
779            None,
780            Some("z"),
781        ])));
782
783        let arg_fields = vec![
784            Field::new("a", LargeUtf8, true).into(),
785            Field::new("a", LargeUtf8, true).into(),
786            Field::new("a", LargeUtf8, true).into(),
787        ];
788        let args = ScalarFunctionArgs {
789            args: vec![c0, c1, c2],
790            arg_fields,
791            number_rows: 3,
792            return_field: Field::new("f", LargeUtf8, true).into(),
793            config_options: Arc::new(ConfigOptions::default()),
794        };
795
796        let result = ConcatWsFunc::new().invoke_with_args(args)?;
797        let expected =
798            Arc::new(LargeStringArray::from(vec!["foo,x", "bar", "baz,z"])) as ArrayRef;
799        match &result {
800            ColumnarValue::Array(array) => {
801                assert_eq!(&expected, array);
802            }
803            _ => panic!("Expected array result"),
804        }
805
806        Ok(())
807    }
808
809    #[test]
810    fn concat_ws_utf8view_null_separator() -> Result<()> {
811        // All-scalar path: null Utf8View separator should return Utf8View(None)
812        let c0 = ColumnarValue::Scalar(ScalarValue::Utf8View(None));
813        let c1 = ColumnarValue::Scalar(ScalarValue::Utf8View(Some("aa".to_string())));
814        let c2 = ColumnarValue::Scalar(ScalarValue::Utf8View(Some("bb".to_string())));
815
816        let arg_fields = vec![
817            Field::new("a", Utf8View, true).into(),
818            Field::new("a", Utf8View, true).into(),
819            Field::new("a", Utf8View, true).into(),
820        ];
821        let args = ScalarFunctionArgs {
822            args: vec![c0, c1, c2],
823            arg_fields,
824            number_rows: 1,
825            return_field: Field::new("f", Utf8View, true).into(),
826            config_options: Arc::new(ConfigOptions::default()),
827        };
828
829        let result = ConcatWsFunc::new().invoke_with_args(args)?;
830        match result {
831            ColumnarValue::Scalar(ScalarValue::Utf8View(None)) => {}
832            other => panic!("Expected Utf8View(None), got {other:?}"),
833        }
834
835        // Array path: null Utf8View scalar separator with array args
836        let c0 = ColumnarValue::Scalar(ScalarValue::Utf8View(None));
837        let c1 =
838            ColumnarValue::Array(Arc::new(StringViewArray::from(vec!["foo", "bar"])));
839
840        let arg_fields = vec![
841            Field::new("a", Utf8View, true).into(),
842            Field::new("a", Utf8View, true).into(),
843        ];
844        let args = ScalarFunctionArgs {
845            args: vec![c0, c1],
846            arg_fields,
847            number_rows: 2,
848            return_field: Field::new("f", Utf8View, true).into(),
849            config_options: Arc::new(ConfigOptions::default()),
850        };
851
852        let result = ConcatWsFunc::new().invoke_with_args(args)?;
853        match result {
854            ColumnarValue::Scalar(ScalarValue::Utf8View(None)) => {}
855            other => panic!("Expected Utf8View(None), got {other:?}"),
856        }
857
858        Ok(())
859    }
860
861    #[test]
862    fn concat_ws_largeutf8_null_separator() -> Result<()> {
863        // All-scalar path: null LargeUtf8 separator should return LargeUtf8(None)
864        let c0 = ColumnarValue::Scalar(ScalarValue::LargeUtf8(None));
865        let c1 = ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some("aa".to_string())));
866        let c2 = ColumnarValue::Scalar(ScalarValue::LargeUtf8(Some("bb".to_string())));
867
868        let arg_fields = vec![
869            Field::new("a", LargeUtf8, true).into(),
870            Field::new("a", LargeUtf8, true).into(),
871            Field::new("a", LargeUtf8, true).into(),
872        ];
873        let args = ScalarFunctionArgs {
874            args: vec![c0, c1, c2],
875            arg_fields,
876            number_rows: 1,
877            return_field: Field::new("f", LargeUtf8, true).into(),
878            config_options: Arc::new(ConfigOptions::default()),
879        };
880
881        let result = ConcatWsFunc::new().invoke_with_args(args)?;
882        match result {
883            ColumnarValue::Scalar(ScalarValue::LargeUtf8(None)) => {}
884            other => panic!("Expected LargeUtf8(None), got {other:?}"),
885        }
886
887        // Array path: null LargeUtf8 scalar separator with array args
888        let c0 = ColumnarValue::Scalar(ScalarValue::LargeUtf8(None));
889        let c1 =
890            ColumnarValue::Array(Arc::new(LargeStringArray::from(vec!["foo", "bar"])));
891
892        let arg_fields = vec![
893            Field::new("a", LargeUtf8, true).into(),
894            Field::new("a", LargeUtf8, true).into(),
895        ];
896        let args = ScalarFunctionArgs {
897            args: vec![c0, c1],
898            arg_fields,
899            number_rows: 2,
900            return_field: Field::new("f", LargeUtf8, true).into(),
901            config_options: Arc::new(ConfigOptions::default()),
902        };
903
904        let result = ConcatWsFunc::new().invoke_with_args(args)?;
905        match result {
906            ColumnarValue::Scalar(ScalarValue::LargeUtf8(None)) => {}
907            other => panic!("Expected LargeUtf8(None), got {other:?}"),
908        }
909
910        Ok(())
911    }
912
913    #[test]
914    fn concat_ws_binary_scalars() -> Result<()> {
915        let c0 = ColumnarValue::Scalar(ScalarValue::Binary(Some(b"|".to_vec())));
916        let c1 = ColumnarValue::Scalar(ScalarValue::Binary(Some(b"aa".to_vec())));
917        let c2 = ColumnarValue::Scalar(ScalarValue::Binary(None));
918        let c3 = ColumnarValue::Scalar(ScalarValue::Binary(Some(b"cc".to_vec())));
919
920        let arg_fields = vec![
921            Field::new("a", Binary, true).into(),
922            Field::new("a", Binary, true).into(),
923            Field::new("a", Binary, true).into(),
924            Field::new("a", Binary, true).into(),
925        ];
926        let args = ScalarFunctionArgs {
927            args: vec![c0, c1, c2, c3],
928            arg_fields,
929            number_rows: 1,
930            return_field: Field::new("f", Binary, true).into(),
931            config_options: Arc::new(ConfigOptions::default()),
932        };
933        let result = ConcatWsFunc::new().invoke_with_args(args)?;
934        match result {
935            ColumnarValue::Scalar(ScalarValue::Binary(Some(v))) => {
936                assert_eq!(v, b"aa|cc");
937            }
938            other => panic!("Expected Binary scalar, got {other:?}"),
939        }
940
941        Ok(())
942    }
943
944    #[test]
945    fn concat_ws_binary_arrays() -> Result<()> {
946        for c1_large_binary in [false, true] {
947            let c0 = ColumnarValue::Scalar(ScalarValue::Binary(Some(b",".to_vec())));
948            let c1 = if c1_large_binary {
949                ColumnarValue::Array(Arc::new(LargeBinaryArray::from_vec(vec![
950                    b"foo".as_ref(),
951                    b"bar",
952                    b"baz",
953                ])))
954            } else {
955                ColumnarValue::Array(Arc::new(BinaryArray::from_vec(vec![
956                    b"foo".as_ref(),
957                    b"bar",
958                    b"baz",
959                ])))
960            };
961            let c2 =
962                ColumnarValue::Array(Arc::new(LargeBinaryArray::from_opt_vec(vec![
963                    Some(b"x".as_ref()),
964                    None,
965                    Some(b"z"),
966                ])));
967
968            let arg_fields = vec![
969                Field::new("a", Binary, true).into(),
970                Field::new("a", Binary, true).into(),
971                Field::new("a", LargeBinary, true).into(),
972            ];
973            let args = ScalarFunctionArgs {
974                args: vec![c0, c1, c2],
975                arg_fields,
976                number_rows: 3,
977                return_field: Field::new("f", LargeBinary, true).into(),
978                config_options: Arc::new(ConfigOptions::default()),
979            };
980
981            let result = ConcatWsFunc::new().invoke_with_args(args)?;
982            let expected = Arc::new(LargeBinaryArray::from_opt_vec(vec![
983                Some(b"foo,x".as_ref()),
984                Some(b"bar"),
985                Some(b"baz,z"),
986            ])) as ArrayRef;
987            match &result {
988                ColumnarValue::Array(array) => assert_eq!(&expected, array),
989                _ => panic!("Expected array result"),
990            }
991        }
992
993        Ok(())
994    }
995}