1use 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}