Skip to main content

paimon_datafusion/
variant_functions.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 std::sync::Arc;
19
20use datafusion::arrow::array::{
21    Array, ArrayRef, BinaryArray, BinaryBuilder, BooleanBuilder, LargeStringArray, StringArray,
22    StringViewArray, StructArray,
23};
24use datafusion::arrow::buffer::{BooleanBuffer, NullBuffer};
25use datafusion::arrow::datatypes::{DataType as ArrowDataType, Field, FieldRef, Fields};
26use datafusion::common::{DataFusionError, Result as DFResult, ScalarValue};
27use datafusion::logical_expr::{
28    ColumnarValue, ReturnFieldArgs, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, Signature,
29    Volatility,
30};
31use datafusion::prelude::SessionContext;
32use paimon::variant::{GenericVariant, VariantDecimal, VariantKind, VariantRef};
33
34pub fn register_variant_functions(ctx: &SessionContext) {
35    ctx.register_udf(ScalarUDF::from(ParseJsonFunc::new(false)));
36    ctx.register_udf(ScalarUDF::from(ParseJsonFunc::new(true)));
37    ctx.register_udf(ScalarUDF::from(IsVariantNullFunc::new()));
38    ctx.register_udf(ScalarUDF::from(VariantGetFunc::new(false)));
39    ctx.register_udf(ScalarUDF::from(VariantGetFunc::new(true)));
40}
41
42#[derive(Debug, Clone, PartialEq, Eq, Hash)]
43struct ParseJsonFunc {
44    try_parse: bool,
45    signature: Signature,
46}
47
48impl ParseJsonFunc {
49    fn new(try_parse: bool) -> Self {
50        Self {
51            try_parse,
52            signature: Signature::string(1, Volatility::Immutable),
53        }
54    }
55}
56
57impl ScalarUDFImpl for ParseJsonFunc {
58    fn name(&self) -> &str {
59        if self.try_parse {
60            "try_parse_json"
61        } else {
62            "parse_json"
63        }
64    }
65
66    fn signature(&self) -> &Signature {
67        &self.signature
68    }
69
70    fn return_type(&self, _arg_types: &[ArrowDataType]) -> DFResult<ArrowDataType> {
71        Ok(variant_arrow_type())
72    }
73
74    fn return_field_from_args(&self, _args: ReturnFieldArgs) -> DFResult<FieldRef> {
75        Ok(Arc::new(Field::new(
76            self.name(),
77            variant_arrow_type(),
78            true,
79        )))
80    }
81
82    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> DFResult<ColumnarValue> {
83        if args.args.len() != 1 {
84            return plan_err(format!("{} expects 1 argument", self.name()));
85        }
86        let arrays = ColumnarValue::values_to_arrays(&args.args)?;
87        let input = arrays[0].as_ref();
88        let mut values = Vec::with_capacity(input.len());
89        for row in 0..input.len() {
90            let Some(json) = string_at(input, row)? else {
91                values.push(None);
92                continue;
93            };
94            match GenericVariant::parse_json(&json) {
95                Ok(variant) => values.push(Some(variant)),
96                Err(e) if self.try_parse => {
97                    let _ = e;
98                    values.push(None);
99                }
100                Err(e) => return Err(to_df_error(e)),
101            }
102        }
103        Ok(ColumnarValue::Array(variant_array(values)?))
104    }
105}
106
107#[derive(Debug, Clone, PartialEq, Eq, Hash)]
108struct IsVariantNullFunc {
109    signature: Signature,
110}
111
112impl IsVariantNullFunc {
113    fn new() -> Self {
114        Self {
115            signature: Signature::any(1, Volatility::Immutable),
116        }
117    }
118}
119
120impl ScalarUDFImpl for IsVariantNullFunc {
121    fn name(&self) -> &str {
122        "is_variant_null"
123    }
124
125    fn signature(&self) -> &Signature {
126        &self.signature
127    }
128
129    fn return_type(&self, _arg_types: &[ArrowDataType]) -> DFResult<ArrowDataType> {
130        Ok(ArrowDataType::Boolean)
131    }
132
133    fn return_field_from_args(&self, _args: ReturnFieldArgs) -> DFResult<FieldRef> {
134        Ok(Arc::new(Field::new(
135            self.name(),
136            ArrowDataType::Boolean,
137            false,
138        )))
139    }
140
141    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> DFResult<ColumnarValue> {
142        if args.args.len() != 1 {
143            return plan_err("is_variant_null expects 1 argument");
144        }
145        let arrays = ColumnarValue::values_to_arrays(&args.args)?;
146        let input = arrays[0].as_ref();
147        let mut builder = BooleanBuilder::new();
148        let Some((values, metadata)) = variant_children(input)? else {
149            for _ in 0..input.len() {
150                builder.append_value(false);
151            }
152            return Ok(ColumnarValue::Array(Arc::new(builder.finish())));
153        };
154
155        for row in 0..input.len() {
156            if input.is_null(row) {
157                builder.append_value(false);
158            } else {
159                let variant = VariantRef::new(values.value(row), metadata.value(row), 0)
160                    .map_err(to_df_error)?;
161                builder.append_value(variant.is_null().map_err(to_df_error)?);
162            }
163        }
164        Ok(ColumnarValue::Array(Arc::new(builder.finish())))
165    }
166}
167
168#[derive(Debug, Clone, PartialEq, Eq, Hash)]
169struct VariantGetFunc {
170    try_get: bool,
171    signature: Signature,
172}
173
174impl VariantGetFunc {
175    fn new(try_get: bool) -> Self {
176        Self {
177            try_get,
178            signature: Signature::variadic_any(Volatility::Immutable),
179        }
180    }
181}
182
183impl ScalarUDFImpl for VariantGetFunc {
184    fn name(&self) -> &str {
185        if self.try_get {
186            "try_variant_get"
187        } else {
188            "variant_get"
189        }
190    }
191
192    fn signature(&self) -> &Signature {
193        &self.signature
194    }
195
196    fn return_type(&self, _arg_types: &[ArrowDataType]) -> DFResult<ArrowDataType> {
197        internal_err("return_field_from_args should be used for variant_get")
198    }
199
200    fn return_field_from_args(&self, args: ReturnFieldArgs) -> DFResult<FieldRef> {
201        if args.arg_fields.len() != 2 && args.arg_fields.len() != 3 {
202            return plan_err(format!("{} expects 2 or 3 arguments", self.name()));
203        }
204        let output = match args.arg_fields.len() {
205            2 => variant_get_output_type(None)?,
206            3 => {
207                let Some(type_arg) = args.scalar_arguments.get(2).and_then(|v| *v) else {
208                    return plan_err("variant_get type argument must be a string literal");
209                };
210                variant_get_output_type(Some(type_arg))?
211            }
212            _ => unreachable!("argument count checked above"),
213        };
214        Ok(Arc::new(Field::new(
215            self.name(),
216            output.arrow_type().clone(),
217            true,
218        )))
219    }
220
221    fn invoke_with_args(&self, args: ScalarFunctionArgs) -> DFResult<ColumnarValue> {
222        if args.args.len() != 2 && args.args.len() != 3 {
223            return plan_err(format!("{} expects 2 or 3 arguments", self.name()));
224        }
225        let output = if args.return_type() == &variant_arrow_type() {
226            VariantGetOutput::Variant
227        } else {
228            VariantGetOutput::Scalar(args.return_type().clone())
229        };
230        let arrays = ColumnarValue::values_to_arrays(&args.args)?;
231        let variants = arrays[0].as_ref();
232        let paths = arrays[1].as_ref();
233        let Some((values, metadata)) = variant_children(variants)? else {
234            return Ok(ColumnarValue::Array(null_array(
235                output.arrow_type(),
236                variants.len(),
237            )));
238        };
239
240        match output {
241            VariantGetOutput::Variant => {
242                let mut result = Vec::with_capacity(variants.len());
243                for row in 0..variants.len() {
244                    result.push(self.variant_at_path(variants, values, metadata, paths, row)?);
245                }
246                Ok(ColumnarValue::Array(variant_array(result)?))
247            }
248            VariantGetOutput::Scalar(data_type) => {
249                let mut scalars = Vec::with_capacity(variants.len());
250                for row in 0..variants.len() {
251                    match self.variant_at_path_ref(variants, values, metadata, paths, row)? {
252                        Some(variant) => scalars.push(cast_variant_to_scalar(
253                            variant,
254                            &data_type,
255                            !self.try_get,
256                        )?),
257                        None => scalars.push(ScalarValue::try_from(&data_type)?),
258                    }
259                }
260                if scalars.is_empty() {
261                    return Ok(ColumnarValue::Array(null_array(&data_type, 0)));
262                }
263                Ok(ColumnarValue::Array(ScalarValue::iter_to_array(scalars)?))
264            }
265        }
266    }
267}
268
269impl VariantGetFunc {
270    fn variant_at_path(
271        &self,
272        variants: &dyn Array,
273        values: &BinaryArray,
274        metadata: &BinaryArray,
275        paths: &dyn Array,
276        row: usize,
277    ) -> DFResult<Option<GenericVariant>> {
278        self.variant_at_path_ref(variants, values, metadata, paths, row)?
279            .map(|variant| variant.to_owned_variant().map_err(to_df_error))
280            .transpose()
281    }
282
283    fn variant_at_path_ref<'a>(
284        &self,
285        variants: &dyn Array,
286        values: &'a BinaryArray,
287        metadata: &'a BinaryArray,
288        paths: &dyn Array,
289        row: usize,
290    ) -> DFResult<Option<VariantRef<'a>>> {
291        if variants.is_null(row) || paths.is_null(row) {
292            return Ok(None);
293        }
294        let path = string_at(paths, row)?;
295        let Some(path) = path else {
296            return Ok(None);
297        };
298        let variant =
299            VariantRef::new(values.value(row), metadata.value(row), 0).map_err(to_df_error)?;
300        match variant.get_path(&path) {
301            Ok(value) => Ok(value),
302            Err(e) if self.try_get => {
303                let _ = e;
304                Ok(None)
305            }
306            Err(e) => Err(to_df_error(e)),
307        }
308    }
309}
310
311#[derive(Clone, Debug)]
312enum VariantGetOutput {
313    Variant,
314    Scalar(ArrowDataType),
315}
316
317impl VariantGetOutput {
318    fn arrow_type(&self) -> &ArrowDataType {
319        match self {
320            Self::Variant => {
321                static VARIANT_TYPE: std::sync::LazyLock<ArrowDataType> =
322                    std::sync::LazyLock::new(variant_arrow_type);
323                &VARIANT_TYPE
324            }
325            Self::Scalar(data_type) => data_type,
326        }
327    }
328}
329
330fn variant_get_output_type(type_arg: Option<&ScalarValue>) -> DFResult<VariantGetOutput> {
331    let Some(type_arg) = type_arg else {
332        return Ok(VariantGetOutput::Variant);
333    };
334    let type_name = match type_arg {
335        ScalarValue::Utf8(Some(v))
336        | ScalarValue::LargeUtf8(Some(v))
337        | ScalarValue::Utf8View(Some(v)) => v,
338        ScalarValue::Utf8(None) | ScalarValue::LargeUtf8(None) | ScalarValue::Utf8View(None) => {
339            return plan_err("variant_get type argument must not be NULL");
340        }
341        _ => return plan_err("variant_get type argument must be a string literal"),
342    };
343    parse_variant_get_type(type_name)
344}
345
346fn parse_variant_get_type(type_name: &str) -> DFResult<VariantGetOutput> {
347    let normalized = type_name.trim().to_ascii_lowercase();
348    match normalized.as_str() {
349        "variant" => Ok(VariantGetOutput::Variant),
350        "boolean" | "bool" => Ok(VariantGetOutput::Scalar(ArrowDataType::Boolean)),
351        "byte" | "tinyint" => Ok(VariantGetOutput::Scalar(ArrowDataType::Int8)),
352        "short" | "smallint" => Ok(VariantGetOutput::Scalar(ArrowDataType::Int16)),
353        "int" | "integer" => Ok(VariantGetOutput::Scalar(ArrowDataType::Int32)),
354        "long" | "bigint" => Ok(VariantGetOutput::Scalar(ArrowDataType::Int64)),
355        "float" | "real" => Ok(VariantGetOutput::Scalar(ArrowDataType::Float32)),
356        "double" => Ok(VariantGetOutput::Scalar(ArrowDataType::Float64)),
357        "string" | "varchar" | "text" => Ok(VariantGetOutput::Scalar(ArrowDataType::Utf8)),
358        "decimal" => Ok(VariantGetOutput::Scalar(ArrowDataType::Decimal128(10, 0))),
359        _ if normalized.starts_with("decimal(") && normalized.ends_with(')') => {
360            let inner = &normalized["decimal(".len()..normalized.len() - 1];
361            let Some((precision, scale)) = inner.split_once(',') else {
362                return plan_err(format!("Invalid decimal type for variant_get: {type_name}"));
363            };
364            let precision = precision
365                .trim()
366                .parse::<u8>()
367                .map_err(|e| DataFusionError::Plan(format!("Invalid decimal precision: {e}")))?;
368            let scale = scale
369                .trim()
370                .parse::<i8>()
371                .map_err(|e| DataFusionError::Plan(format!("Invalid decimal scale: {e}")))?;
372            Ok(VariantGetOutput::Scalar(ArrowDataType::Decimal128(
373                precision, scale,
374            )))
375        }
376        _ => plan_err(format!("Unsupported variant_get type: {type_name}")),
377    }
378}
379
380fn cast_variant_to_scalar(
381    variant: VariantRef<'_>,
382    target: &ArrowDataType,
383    fail_on_error: bool,
384) -> DFResult<ScalarValue> {
385    if variant.is_null().map_err(to_df_error)? {
386        return ScalarValue::try_from(target);
387    }
388    let result = match target {
389        ArrowDataType::Boolean => cast_to_boolean(variant),
390        ArrowDataType::Int8 => cast_to_i64(variant).and_then(|v| {
391            i8::try_from(v)
392                .map(ScalarValue::from)
393                .map_err(|_| invalid_cast())
394        }),
395        ArrowDataType::Int16 => cast_to_i64(variant).and_then(|v| {
396            i16::try_from(v)
397                .map(ScalarValue::from)
398                .map_err(|_| invalid_cast())
399        }),
400        ArrowDataType::Int32 => cast_to_i64(variant).and_then(|v| {
401            i32::try_from(v)
402                .map(ScalarValue::from)
403                .map_err(|_| invalid_cast())
404        }),
405        ArrowDataType::Int64 => cast_to_i64(variant).map(ScalarValue::from),
406        ArrowDataType::Float32 => {
407            cast_to_f64(variant).map(|v| ScalarValue::Float32(Some(v as f32)))
408        }
409        ArrowDataType::Float64 => cast_to_f64(variant).map(ScalarValue::from),
410        ArrowDataType::Utf8 => cast_to_string(variant).map(ScalarValue::from),
411        ArrowDataType::Decimal128(precision, scale) => cast_to_decimal(variant, *precision, *scale),
412        _ => Err(invalid_cast()),
413    };
414
415    match result {
416        Ok(value) => Ok(value),
417        Err(e) if !fail_on_error => {
418            let _ = e;
419            ScalarValue::try_from(target)
420        }
421        Err(e) => Err(e),
422    }
423}
424
425fn cast_to_boolean(variant: VariantRef<'_>) -> DFResult<ScalarValue> {
426    match variant.kind().map_err(to_df_error)? {
427        VariantKind::Boolean => Ok(ScalarValue::Boolean(Some(
428            variant.get_boolean().map_err(to_df_error)?,
429        ))),
430        VariantKind::String => match variant
431            .get_string()
432            .map_err(to_df_error)?
433            .to_ascii_lowercase()
434            .as_str()
435        {
436            "true" => Ok(ScalarValue::Boolean(Some(true))),
437            "false" => Ok(ScalarValue::Boolean(Some(false))),
438            _ => Err(invalid_cast()),
439        },
440        _ => Err(invalid_cast()),
441    }
442}
443
444fn cast_to_i64(variant: VariantRef<'_>) -> DFResult<i64> {
445    match variant.kind().map_err(to_df_error)? {
446        VariantKind::Long
447        | VariantKind::Date
448        | VariantKind::Timestamp
449        | VariantKind::TimestampNtz => variant.get_long().map_err(to_df_error),
450        VariantKind::String => variant
451            .get_string()
452            .map_err(to_df_error)?
453            .parse::<i64>()
454            .map_err(|_| invalid_cast()),
455        VariantKind::Decimal => {
456            let decimal = variant.get_decimal().map_err(to_df_error)?;
457            rescale_decimal(decimal.unscaled, decimal.scale, 0)
458                .and_then(|v| i64::try_from(v).map_err(|_| invalid_cast()))
459        }
460        _ => Err(invalid_cast()),
461    }
462}
463
464fn cast_to_f64(variant: VariantRef<'_>) -> DFResult<f64> {
465    match variant.kind().map_err(to_df_error)? {
466        VariantKind::Long
467        | VariantKind::Date
468        | VariantKind::Timestamp
469        | VariantKind::TimestampNtz => Ok(variant.get_long().map_err(to_df_error)? as f64),
470        VariantKind::Double => variant.get_double().map_err(to_df_error),
471        VariantKind::Float => Ok(variant.get_float().map_err(to_df_error)? as f64),
472        VariantKind::Decimal => {
473            let decimal = variant.get_decimal().map_err(to_df_error)?;
474            Ok(decimal.unscaled as f64 / 10f64.powi(decimal.scale as i32))
475        }
476        VariantKind::String => variant
477            .get_string()
478            .map_err(to_df_error)?
479            .parse::<f64>()
480            .map_err(|_| invalid_cast()),
481        _ => Err(invalid_cast()),
482    }
483}
484
485fn cast_to_string(variant: VariantRef<'_>) -> DFResult<String> {
486    match variant.kind().map_err(to_df_error)? {
487        VariantKind::Object | VariantKind::Array => variant.to_json().map_err(to_df_error),
488        VariantKind::Boolean => Ok(variant.get_boolean().map_err(to_df_error)?.to_string()),
489        VariantKind::Long
490        | VariantKind::Date
491        | VariantKind::Timestamp
492        | VariantKind::TimestampNtz => Ok(variant.get_long().map_err(to_df_error)?.to_string()),
493        VariantKind::String => variant.get_string().map_err(to_df_error),
494        VariantKind::Double => Ok(variant.get_double().map_err(to_df_error)?.to_string()),
495        VariantKind::Decimal => Ok(variant
496            .get_decimal()
497            .map_err(to_df_error)?
498            .to_plain_string()),
499        VariantKind::Float => Ok(variant.get_float().map_err(to_df_error)?.to_string()),
500        _ => variant.to_json().map_err(to_df_error),
501    }
502}
503
504fn cast_to_decimal(variant: VariantRef<'_>, precision: u8, scale: i8) -> DFResult<ScalarValue> {
505    let unscaled = match variant.kind().map_err(to_df_error)? {
506        VariantKind::Long
507        | VariantKind::Date
508        | VariantKind::Timestamp
509        | VariantKind::TimestampNtz => {
510            rescale_decimal(variant.get_long().map_err(to_df_error)? as i128, 0, scale)?
511        }
512        VariantKind::Decimal => {
513            let decimal = variant.get_decimal().map_err(to_df_error)?;
514            rescale_decimal(decimal.unscaled, decimal.scale, scale)?
515        }
516        VariantKind::String => {
517            let parsed = parse_decimal_string(&variant.get_string().map_err(to_df_error)?)
518                .ok_or_else(invalid_cast)?;
519            rescale_decimal(parsed.unscaled, parsed.scale, scale)?
520        }
521        _ => return Err(invalid_cast()),
522    };
523    if decimal_precision(unscaled) > precision {
524        return Err(invalid_cast());
525    }
526    Ok(ScalarValue::Decimal128(Some(unscaled), precision, scale))
527}
528
529fn rescale_decimal(unscaled: i128, from_scale: i8, to_scale: i8) -> DFResult<i128> {
530    match to_scale.cmp(&from_scale) {
531        std::cmp::Ordering::Equal => Ok(unscaled),
532        std::cmp::Ordering::Greater => {
533            let factor = 10_i128
534                .checked_pow((to_scale - from_scale) as u32)
535                .ok_or_else(invalid_cast)?;
536            unscaled.checked_mul(factor).ok_or_else(invalid_cast)
537        }
538        std::cmp::Ordering::Less => {
539            let factor = 10_i128
540                .checked_pow((from_scale - to_scale) as u32)
541                .ok_or_else(invalid_cast)?;
542            if unscaled % factor == 0 {
543                Ok(unscaled / factor)
544            } else {
545                Err(invalid_cast())
546            }
547        }
548    }
549}
550
551fn parse_decimal_string(input: &str) -> Option<VariantDecimal> {
552    let input = input.trim();
553    if input.is_empty() || input.contains(['e', 'E']) {
554        return None;
555    }
556    let negative = input.starts_with('-');
557    let unsigned = input.strip_prefix('-').unwrap_or(input);
558    if unsigned.is_empty()
559        || unsigned.matches('.').count() > 1
560        || !unsigned.bytes().all(|ch| ch == b'.' || ch.is_ascii_digit())
561    {
562        return None;
563    }
564    let scale = unsigned
565        .split_once('.')
566        .map(|(_, fraction)| fraction.len())
567        .unwrap_or(0);
568    let digits: String = unsigned
569        .bytes()
570        .filter(|ch| *ch != b'.')
571        .map(char::from)
572        .collect();
573    let significant = digits.trim_start_matches('0');
574    let precision = if significant.is_empty() {
575        1
576    } else {
577        significant.len()
578    };
579    if precision > 38 || scale > 38 {
580        return None;
581    }
582    let mut unscaled = digits.parse::<i128>().ok()?;
583    if negative {
584        unscaled = -unscaled;
585    }
586    Some(VariantDecimal {
587        unscaled,
588        precision: precision as u8,
589        scale: scale as i8,
590    })
591}
592
593fn decimal_precision(unscaled: i128) -> u8 {
594    let mut value = unscaled.unsigned_abs();
595    if value == 0 {
596        return 1;
597    }
598    let mut precision = 0;
599    while value > 0 {
600        precision += 1;
601        value /= 10;
602    }
603    precision
604}
605
606fn string_at(array: &dyn Array, row: usize) -> DFResult<Option<String>> {
607    if array.is_null(row) {
608        return Ok(None);
609    }
610    match array.data_type() {
611        ArrowDataType::Utf8 => Ok(Some(
612            array
613                .as_any()
614                .downcast_ref::<StringArray>()
615                .ok_or_else(|| DataFusionError::Internal("Expected Utf8 array".to_string()))?
616                .value(row)
617                .to_string(),
618        )),
619        ArrowDataType::LargeUtf8 => Ok(Some(
620            array
621                .as_any()
622                .downcast_ref::<LargeStringArray>()
623                .ok_or_else(|| DataFusionError::Internal("Expected LargeUtf8 array".to_string()))?
624                .value(row)
625                .to_string(),
626        )),
627        ArrowDataType::Utf8View => Ok(Some(
628            array
629                .as_any()
630                .downcast_ref::<StringViewArray>()
631                .ok_or_else(|| DataFusionError::Internal("Expected Utf8View array".to_string()))?
632                .value(row)
633                .to_string(),
634        )),
635        other => plan_err(format!("Expected string array, got {other:?}")),
636    }
637}
638
639fn variant_children(array: &dyn Array) -> DFResult<Option<(&BinaryArray, &BinaryArray)>> {
640    let ArrowDataType::Struct(fields) = array.data_type() else {
641        return Ok(None);
642    };
643    if fields.len() != 2
644        || fields[0].name() != "value"
645        || fields[0].data_type() != &ArrowDataType::Binary
646        || fields[1].name() != "metadata"
647        || fields[1].data_type() != &ArrowDataType::Binary
648    {
649        return Ok(None);
650    }
651    let array = array
652        .as_any()
653        .downcast_ref::<StructArray>()
654        .ok_or_else(|| DataFusionError::Internal("Expected Variant StructArray".to_string()))?;
655    let values = array
656        .column(0)
657        .as_any()
658        .downcast_ref::<BinaryArray>()
659        .ok_or_else(|| {
660            DataFusionError::Internal("Expected Variant.value BinaryArray".to_string())
661        })?;
662    let metadata = array
663        .column(1)
664        .as_any()
665        .downcast_ref::<BinaryArray>()
666        .ok_or_else(|| {
667            DataFusionError::Internal("Expected Variant.metadata BinaryArray".to_string())
668        })?;
669    Ok(Some((values, metadata)))
670}
671
672fn variant_array(values: Vec<Option<GenericVariant>>) -> DFResult<ArrayRef> {
673    let len = values.len();
674    let mut value_builder = BinaryBuilder::new();
675    let mut metadata_builder = BinaryBuilder::new();
676    let mut validities = Vec::with_capacity(len);
677    for value in values {
678        match value {
679            Some(variant) => {
680                value_builder.append_value(variant.value());
681                metadata_builder.append_value(variant.metadata());
682                validities.push(true);
683            }
684            None => {
685                value_builder.append_value(&[] as &[u8]);
686                metadata_builder.append_value(&[] as &[u8]);
687                validities.push(false);
688            }
689        }
690    }
691    let nulls = if validities.iter().all(|valid| *valid) {
692        None
693    } else {
694        Some(NullBuffer::new(BooleanBuffer::from(validities)))
695    };
696    let array = StructArray::try_new(
697        variant_fields(),
698        vec![
699            Arc::new(value_builder.finish()),
700            Arc::new(metadata_builder.finish()),
701        ],
702        nulls,
703    )?;
704    Ok(Arc::new(array))
705}
706
707fn variant_arrow_type() -> ArrowDataType {
708    paimon::arrow::variant_arrow_type()
709}
710
711fn variant_fields() -> Fields {
712    match variant_arrow_type() {
713        ArrowDataType::Struct(fields) => fields,
714        _ => unreachable!("variant_arrow_type must be a struct"),
715    }
716}
717
718fn null_array(data_type: &ArrowDataType, len: usize) -> ArrayRef {
719    datafusion::arrow::array::new_null_array(data_type, len)
720}
721
722fn invalid_cast() -> DataFusionError {
723    DataFusionError::Execution("Invalid Variant cast".to_string())
724}
725
726fn to_df_error(error: paimon::Error) -> DataFusionError {
727    DataFusionError::External(Box::new(error))
728}
729
730fn plan_err<T>(message: impl Into<String>) -> DFResult<T> {
731    Err(DataFusionError::Plan(message.into()))
732}
733
734fn internal_err<T>(message: impl Into<String>) -> DFResult<T> {
735    Err(DataFusionError::Internal(message.into()))
736}
737
738#[cfg(test)]
739mod tests {
740    use super::*;
741    use datafusion::arrow::array::{BooleanArray, Int32Array, StringArray};
742
743    async fn collect_one(sql: &str) -> datafusion::arrow::record_batch::RecordBatch {
744        let ctx = SessionContext::new();
745        register_variant_functions(&ctx);
746        let batches = ctx.sql(sql).await.unwrap().collect().await.unwrap();
747        assert_eq!(batches.len(), 1);
748        batches.into_iter().next().unwrap()
749    }
750
751    #[tokio::test]
752    async fn parse_json_and_variant_get_scalars() {
753        let batch = collect_one(
754            r#"
755            SELECT
756              variant_get(parse_json('{"age":26,"city":"Beijing","nested":{"name":"Alice"},"arr":[1,2,3]}'), '$.age', 'int') AS age,
757              variant_get(parse_json('{"age":26,"city":"Beijing","nested":{"name":"Alice"},"arr":[1,2,3]}'), '$.city', 'string') AS city,
758              variant_get(parse_json('{"age":26,"city":"Beijing","nested":{"name":"Alice"},"arr":[1,2,3]}'), '$.nested.name', 'string') AS name,
759              variant_get(parse_json('{"age":26,"city":"Beijing","nested":{"name":"Alice"},"arr":[1,2,3]}'), '$.arr[1]', 'int') AS arr_value
760            "#,
761        )
762        .await;
763
764        assert_eq!(
765            batch
766                .column(0)
767                .as_any()
768                .downcast_ref::<Int32Array>()
769                .unwrap()
770                .value(0),
771            26
772        );
773        assert_eq!(
774            batch
775                .column(1)
776                .as_any()
777                .downcast_ref::<StringArray>()
778                .unwrap()
779                .value(0),
780            "Beijing"
781        );
782        assert_eq!(
783            batch
784                .column(2)
785                .as_any()
786                .downcast_ref::<StringArray>()
787                .unwrap()
788                .value(0),
789            "Alice"
790        );
791        assert_eq!(
792            batch
793                .column(3)
794                .as_any()
795                .downcast_ref::<Int32Array>()
796                .unwrap()
797                .value(0),
798            2
799        );
800    }
801
802    #[tokio::test]
803    async fn variant_null_is_distinct_from_sql_null() {
804        let batch = collect_one(
805            "SELECT is_variant_null(parse_json('null')) AS variant_null, is_variant_null(NULL) AS sql_null",
806        )
807        .await;
808        let variant_null = batch
809            .column(0)
810            .as_any()
811            .downcast_ref::<BooleanArray>()
812            .unwrap();
813        let sql_null = batch
814            .column(1)
815            .as_any()
816            .downcast_ref::<BooleanArray>()
817            .unwrap();
818        assert!(variant_null.value(0));
819        assert!(!sql_null.value(0));
820    }
821
822    #[tokio::test]
823    async fn try_functions_return_null_on_invalid_input() {
824        let batch = collect_one(
825            r#"
826            SELECT
827              try_parse_json('{bad json') AS bad_json,
828              try_variant_get(parse_json('{"age":"not an int"}'), '$.age', 'int') AS bad_cast,
829              variant_get(parse_json('{}'), '$.missing', 'int') AS missing_path
830            "#,
831        )
832        .await;
833        assert!(batch.column(0).is_null(0));
834        assert!(batch.column(1).is_null(0));
835        assert!(batch.column(2).is_null(0));
836    }
837
838    #[tokio::test]
839    async fn strict_functions_surface_errors() {
840        let ctx = SessionContext::new();
841        register_variant_functions(&ctx);
842        let err = ctx
843            .sql("SELECT parse_json('{bad json')")
844            .await
845            .unwrap()
846            .collect()
847            .await
848            .unwrap_err();
849        assert!(err.to_string().contains("Expected"));
850
851        let err = ctx
852            .sql("SELECT variant_get(parse_json('{\"age\":\"not an int\"}'), '$.age', 'int')")
853            .await
854            .unwrap()
855            .collect()
856            .await
857            .unwrap_err();
858        assert!(err.to_string().contains("Invalid Variant cast"));
859    }
860
861    #[tokio::test]
862    async fn variant_get_rejects_non_literal_type_argument() {
863        let ctx = SessionContext::new();
864        register_variant_functions(&ctx);
865        let sql = r#"
866            SELECT variant_get(parse_json('{"age":26}'), '$.age', type_name)
867            FROM (VALUES ('int')) AS t(type_name)
868        "#;
869        let err = match ctx.sql(sql).await {
870            Ok(df) => df.collect().await.unwrap_err(),
871            Err(err) => err,
872        };
873        assert!(err
874            .to_string()
875            .contains("variant_get type argument must be a string literal"));
876    }
877}