1use std::collections::HashMap;
16use std::pin::Pin;
17use std::task::{Context, Poll};
18
19use datafusion_common::cast_column;
20use datafusion_common::format::DEFAULT_CAST_OPTIONS;
21use datatypes::arrow::array::{ArrayRef, new_null_array};
22use datatypes::arrow::datatypes::{DataType, Field, FieldRef, Schema, SchemaRef};
23use datatypes::arrow::record_batch::RecordBatch;
24use datatypes::extension::json::{JsonMetadata, is_json2_extension_type};
25use datatypes::json::{JsonSettings, TypeHintMismatchPolicy};
26use datatypes::vectors::json::array::JsonArray;
27use datatypes::vectors::json::json2_physical_data_type;
28use futures::Stream;
29use futures::stream::BoxStream;
30use serde_json::from_str;
31use snafu::{ResultExt, ensure};
32
33use crate::error::{
34 CastColumnSnafu, DataTypeMismatchSnafu, NewRecordBatchSnafu, Result, UnexpectedSnafu,
35};
36use crate::sst::parquet::Json2TargetLayout;
37
38pub(crate) type ProjectedRecordBatchStream = BoxStream<'static, Result<RecordBatch>>;
39
40#[derive(Debug)]
42pub(crate) enum AlignMode {
43 AlignToSchema,
45 Rewrite {
47 columns: HashMap<String, Json2TargetLayout>,
53 },
54}
55
56#[derive(Debug)]
58enum ResolvedAlignMode {
59 AlignToSchema,
60 Rewrite {
61 columns: HashMap<String, RewriteSettings>,
62 },
63}
64
65#[derive(derive_more::Debug)]
75pub struct JsonSchemaAligner<S> {
76 #[debug(skip)]
77 inner: S,
78 output_schema: SchemaRef,
80 projected_root_presence: Vec<bool>,
83 expected_input_col_num: usize,
85 all_roots_present: bool,
88 mode: ResolvedAlignMode,
90 is_schema_matched: Option<bool>,
92}
93
94impl<S> JsonSchemaAligner<S>
95where
96 S: Stream<Item = Result<RecordBatch>>,
97{
98 pub(crate) fn new(
101 inner: S,
102 projected_root_presence: Vec<bool>,
103 output_schema: SchemaRef,
104 mode: AlignMode,
105 ) -> Result<JsonSchemaAligner<S>> {
106 ensure!(
107 projected_root_presence.len() == output_schema.fields().len(),
108 UnexpectedSnafu {
109 reason: format!(
110 "JsonSchemaAligner projected root presence len {} does not match output schema columns {}",
111 projected_root_presence.len(),
112 output_schema.fields().len()
113 ),
114 }
115 );
116
117 let mode = resolve_align_mode(mode, output_schema.as_ref())?;
118
119 let expected_input_col_num = projected_root_presence
120 .iter()
121 .filter(|matched| **matched)
122 .count();
123 let all_roots_present = projected_root_presence.iter().all(|&m| m);
124 Ok(JsonSchemaAligner {
125 inner,
126 output_schema,
127 projected_root_presence,
128 expected_input_col_num,
129 all_roots_present,
130 mode,
131 is_schema_matched: None,
132 })
133 }
134}
135
136impl<S> Stream for JsonSchemaAligner<S>
137where
138 S: Stream<Item = Result<RecordBatch>> + Unpin,
139{
140 type Item = Result<RecordBatch>;
141
142 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
143 let this = self.get_mut();
144
145 match Pin::new(&mut this.inner).poll_next(cx) {
146 Poll::Ready(Some(Ok(rb))) => {
147 let is_schema_matched = matches!(this.mode, ResolvedAlignMode::AlignToSchema)
148 && this.all_roots_present
149 && *this
150 .is_schema_matched
151 .get_or_insert_with(|| rb.schema() == this.output_schema);
152
153 if is_schema_matched {
154 Poll::Ready(Some(Ok(rb)))
155 } else {
156 Poll::Ready(Some(align_projected_batch(
157 rb,
158 &this.output_schema,
159 &this.projected_root_presence,
160 this.expected_input_col_num,
161 &this.mode,
162 )))
163 }
164 }
165 Poll::Ready(Some(Err(err))) => Poll::Ready(Some(Err(err))),
166 Poll::Ready(None) => Poll::Ready(None),
167 Poll::Pending => Poll::Pending,
168 }
169 }
170}
171
172fn align_projected_batch(
173 rb: RecordBatch,
174 output_schema: &SchemaRef,
175 projected_root_presence: &[bool],
176 expected_input_col_num: usize,
177 mode: &ResolvedAlignMode,
178) -> Result<RecordBatch> {
179 ensure!(
180 rb.columns().len() == expected_input_col_num,
181 UnexpectedSnafu {
182 reason: format!(
183 "JsonSchemaAligner expected {} input columns but got {}",
184 expected_input_col_num,
185 rb.columns().len()
186 ),
187 }
188 );
189
190 let mut cols = Vec::with_capacity(projected_root_presence.len());
191 let mut idx = 0;
192 let input_schema = rb.schema_ref();
193
194 for (field, present) in output_schema.fields().iter().zip(projected_root_presence) {
195 if !present {
196 cols.push(new_null_array(field.data_type(), rb.num_rows()));
197 continue;
198 }
199
200 let array = match mode {
201 ResolvedAlignMode::AlignToSchema => {
202 align_array(rb.column(idx), input_schema.field(idx), field)?
203 }
204 ResolvedAlignMode::Rewrite { columns } => match columns.get(field.name()) {
205 Some(settings) => rewrite_array(rb.column(idx), input_schema.field(idx), settings)?,
206 None => rb.column(idx).clone(),
207 },
208 };
209 cols.push(array);
210 idx += 1;
211 }
212
213 RecordBatch::try_new(output_schema.clone(), cols).context(NewRecordBatchSnafu)
214}
215
216fn align_array(
217 source_array: &ArrayRef,
218 source_field: &Field,
219 target_field: &FieldRef,
220) -> Result<ArrayRef> {
221 if source_array.data_type() == target_field.data_type() {
222 return Ok(source_array.clone());
223 }
224
225 if is_json2_extension_type(target_field) {
226 if is_json2_extension_type(source_field) {
227 return JsonArray::from(source_array)
228 .project_to_v2(source_field, target_field.data_type())
229 .context(DataTypeMismatchSnafu);
230 }
231 return JsonArray::from(source_array)
232 .project_to(target_field.data_type())
233 .context(DataTypeMismatchSnafu);
234 }
235
236 if !matches!(target_field.data_type(), DataType::Struct(_)) {
237 return Ok(source_array.clone());
238 }
239
240 cast_column(
241 source_array,
242 target_field.data_type(),
243 &DEFAULT_CAST_OPTIONS,
244 )
245 .context(CastColumnSnafu)
246}
247
248fn rewrite_array(
249 source_array: &ArrayRef,
250 source_field: &Field,
251 settings: &RewriteSettings,
252) -> Result<ArrayRef> {
253 JsonArray::from(source_array)
254 .rewrite_to_v2_with_type_hint_mismatch_policy(
255 source_field,
256 &settings.logical_settings,
257 &settings.target_layout,
258 TypeHintMismatchPolicy::CoerceOrNull,
259 )
260 .context(DataTypeMismatchSnafu)
261}
262
263#[derive(Debug)]
268struct RewriteSettings {
269 logical_settings: JsonSettings,
272 target_layout: JsonSettings,
274}
275
276impl TryFrom<&Json2TargetLayout> for RewriteSettings {
277 type Error = crate::error::Error;
278
279 fn try_from(layout: &Json2TargetLayout) -> Result<Self> {
280 let metadata = from_str::<JsonMetadata>(&layout.extension_metadata).map_err(|e| {
281 UnexpectedSnafu {
282 reason: format!("invalid JSON2 extension metadata: {e}"),
283 }
284 .build()
285 })?;
286 Ok(Self {
287 logical_settings: metadata.into_json_settings(),
288 target_layout: layout.target_layout.clone(),
289 })
290 }
291}
292
293fn resolve_align_mode(mode: AlignMode, output_schema: &Schema) -> Result<ResolvedAlignMode> {
295 let AlignMode::Rewrite { columns } = mode else {
296 return Ok(ResolvedAlignMode::AlignToSchema);
297 };
298
299 let mut rewrite_columns = HashMap::with_capacity(columns.len());
300 for (name, layout) in columns {
301 let settings = RewriteSettings::try_from(&layout)?;
302 let field = output_schema.field_with_name(&name).map_err(|_| {
303 UnexpectedSnafu {
304 reason: format!("JSON2 rewrite column '{name}' is missing from output schema"),
305 }
306 .build()
307 })?;
308 ensure!(
309 is_json2_extension_type(field)
310 && field.data_type() == &json2_physical_data_type(&settings.target_layout),
311 UnexpectedSnafu {
312 reason: format!(
313 "JSON2 rewrite layout for column '{name}' does not match output field"
314 ),
315 }
316 );
317 rewrite_columns.insert(name, settings);
318 }
319
320 Ok(ResolvedAlignMode::Rewrite {
321 columns: rewrite_columns,
322 })
323}
324
325#[cfg(test)]
326mod tests {
327 use std::collections::HashMap;
328 use std::sync::Arc;
329
330 use datatypes::arrow::array::{
331 Array, ArrayRef, BinaryArray, Int64Array, StringArray, StringViewArray, StructArray,
332 };
333 use datatypes::arrow::datatypes::{DataType, Field, Fields, Schema};
334 use datatypes::extension::json::{Json2ExtensionType, JsonMetadata};
335 use datatypes::json::JsonTypeHint;
336 use datatypes::prelude::ConcreteDataType;
337 use datatypes::types::parse_string_to_jsonb;
338 use futures::{StreamExt, stream};
339
340 use super::*;
341
342 #[test]
343 fn test_aligner_resolves_json2_rewrite_settings()
344 -> std::result::Result<(), Box<dyn std::error::Error>> {
345 let logical_settings = JsonSettings::default();
346 let target_layout = JsonSettings::try_new(vec![], Some(0))?;
347 let rewrite_targets = HashMap::from([(
348 "j".to_string(),
349 Json2TargetLayout {
350 extension_metadata: serde_json::to_string(&JsonMetadata::new(
351 logical_settings.clone(),
352 ))?,
353 target_layout: target_layout.clone(),
354 },
355 )]);
356 let aligner = JsonSchemaAligner::new(
357 stream::empty::<Result<RecordBatch>>(),
358 vec![false],
359 schema([
360 Field::new("j", json2_physical_data_type(&target_layout), true)
361 .with_extension_type(Json2ExtensionType::default()),
362 ]),
363 AlignMode::Rewrite {
364 columns: rewrite_targets,
365 },
366 )?;
367 let ResolvedAlignMode::Rewrite { columns } = &aligner.mode else {
368 panic!("expected rewrite mode");
369 };
370 let settings = &columns["j"];
371 assert_eq!(logical_settings, settings.logical_settings);
372 assert_eq!(target_layout, settings.target_layout);
373 Ok(())
374 }
375
376 #[tokio::test]
377 async fn test_aligner_with_all_projected_roots_match() {
378 let output_schema = schema([
379 Field::new("a", DataType::Int64, true),
380 Field::new("b", DataType::Utf8, true),
381 ]);
382 let input = RecordBatch::try_new(
383 output_schema.clone(),
384 vec![int_array([1, 2, 3]), string_array(["x", "y", "z"])],
385 )
386 .unwrap();
387 let stream = stream::iter([Ok(input.clone())]);
388
389 let mut aligner = JsonSchemaAligner::new(
390 stream,
391 vec![true, true],
392 output_schema.clone(),
393 AlignMode::AlignToSchema,
394 )
395 .unwrap();
396 let output = aligner.next().await.unwrap().unwrap();
397
398 assert_eq!(input, output);
399 assert!(aligner.next().await.is_none());
400 }
401
402 #[tokio::test]
403 async fn test_aligner_with_fills_null_root_columns() {
404 let input_schema = schema([Field::new("a", DataType::Int64, true)]);
405 let output_schema = schema([
406 Field::new("a", DataType::Int64, true),
407 Field::new("missing", DataType::Utf8, true),
408 Field::new("c", DataType::Int64, true),
409 ]);
410 let input = RecordBatch::try_new(input_schema, vec![int_array([10, 20])]).unwrap();
411 let stream = stream::iter([Ok(input)]);
412
413 let mut aligner = JsonSchemaAligner::new(
414 stream,
415 vec![true, false, false],
416 output_schema.clone(),
417 AlignMode::AlignToSchema,
418 )
419 .unwrap();
420 let output = aligner.next().await.unwrap().unwrap();
421
422 assert_eq!(output_schema, output.schema());
423 assert_eq!(3, output.num_columns());
424 assert_eq!(
425 &[Some(10), Some(20)],
426 output
427 .column(0)
428 .as_any()
429 .downcast_ref::<Int64Array>()
430 .unwrap()
431 .iter()
432 .collect::<Vec<_>>()
433 .as_slice()
434 );
435 assert_eq!(DataType::Utf8, *output.column(1).data_type());
436 assert_eq!(output.num_rows(), output.column(1).null_count());
437 assert_eq!(DataType::Int64, *output.column(2).data_type());
438 assert_eq!(output.num_rows(), output.column(2).null_count());
439 }
440
441 #[tokio::test]
442 async fn test_aligner_with_fills_missing_struct_root_column() {
443 let input_schema = schema([Field::new("a", DataType::Int64, true)]);
444 let struct_type = DataType::Struct(Fields::from(vec![
445 Field::new("x", DataType::Int64, true),
446 Field::new("y", DataType::Utf8, true),
447 ]));
448 let output_schema = schema([
449 Field::new("a", DataType::Int64, true),
450 Field::new("missing_struct", struct_type.clone(), true),
451 ]);
452 let input = RecordBatch::try_new(input_schema, vec![int_array([10, 20])]).unwrap();
453 let stream = stream::iter([Ok(input)]);
454
455 let mut aligner = JsonSchemaAligner::new(
456 stream,
457 vec![true, false],
458 output_schema.clone(),
459 AlignMode::AlignToSchema,
460 )
461 .unwrap();
462 let output = aligner.next().await.unwrap().unwrap();
463
464 assert_eq!(output_schema, output.schema());
465 assert_eq!(2, output.num_columns());
466 assert_eq!(struct_type, output.column(1).data_type().clone());
467 assert_eq!(output.num_rows(), output.column(1).null_count());
468 }
469
470 #[tokio::test]
471 async fn test_aligner_reject_projection_len_mismatch() {
472 let output_schema = schema([Field::new("a", DataType::Int64, true)]);
473 let stream = stream::iter([]);
474
475 let err = match JsonSchemaAligner::new(
476 stream,
477 vec![true, false],
478 output_schema,
479 AlignMode::AlignToSchema,
480 ) {
481 Ok(_) => panic!("JsonSchemaAligner should reject projection length mismatch"),
482 Err(err) => err,
483 };
484
485 assert!(
486 err.to_string()
487 .contains("projected root presence len 2 does not match output schema columns 1")
488 );
489 }
490
491 #[tokio::test]
492 async fn test_aligner_reject_with_input_column_mismatch() {
493 let input_schema = schema([Field::new("a", DataType::Int64, true)]);
494 let output_schema = schema([
495 Field::new("a", DataType::Int64, true),
496 Field::new("b", DataType::Int64, true),
497 Field::new("missing", DataType::Int64, true),
498 ]);
499 let input = RecordBatch::try_new(input_schema, vec![int_array([1, 2])]).unwrap();
500 let stream = stream::iter([Ok(input)]);
501
502 let mut aligner = JsonSchemaAligner::new(
503 stream,
504 vec![true, true, false],
505 output_schema,
506 AlignMode::AlignToSchema,
507 )
508 .unwrap();
509 let err = aligner.next().await.unwrap().unwrap_err();
510
511 assert!(
512 err.to_string()
513 .contains("expected 2 input columns but got 1")
514 );
515 }
516
517 #[tokio::test]
518 async fn test_json_schema_aligner_aligns_struct_field() {
519 let output_schema = schema([Field::new(
520 "nested",
521 DataType::Struct(Fields::from(vec![
522 Field::new("x", DataType::Int64, true),
523 Field::new("y", DataType::Utf8, true),
524 ])),
525 true,
526 )]);
527 let input = RecordBatch::try_new(
528 schema([Field::new(
529 "nested",
530 DataType::Struct(Fields::from(vec![Field::new("x", DataType::Int64, true)])),
531 true,
532 )]),
533 vec![Arc::new(StructArray::from(vec![(
534 Arc::new(Field::new("x", DataType::Int64, true)),
535 int_array([1, 2]),
536 )]))],
537 )
538 .unwrap();
539
540 let mut aligner = JsonSchemaAligner::new(
541 stream::iter([Ok(input)]),
542 vec![true],
543 output_schema.clone(),
544 AlignMode::AlignToSchema,
545 )
546 .unwrap();
547 let output = aligner.next().await.unwrap().unwrap();
548
549 assert_eq!(output_schema, output.schema());
550 let nested = output
551 .column(0)
552 .as_any()
553 .downcast_ref::<StructArray>()
554 .unwrap();
555 assert_eq!(2, nested.columns().len());
556 assert_eq!(2, nested.column(1).null_count());
557 }
558
559 #[tokio::test]
560 async fn test_json_schema_aligner_decodes_variant_to_struct() {
561 let source_values = [
562 Some(parse_string_to_jsonb("1").unwrap()),
563 Some(parse_string_to_jsonb(r#"{"b":2}"#).unwrap()),
564 None,
565 ];
566 let source = Arc::new(BinaryArray::from_iter(
567 source_values.iter().map(|value| value.as_deref()),
568 )) as ArrayRef;
569 let input_fields = Fields::from(vec![Arc::new(Field::new("a", DataType::Binary, true))]);
570 let input = RecordBatch::try_new(
571 schema([Field::new(
572 "j",
573 DataType::Struct(input_fields.clone()),
574 true,
575 )]),
576 vec![Arc::new(StructArray::new(input_fields, vec![source], None))],
577 )
578 .unwrap();
579
580 let output_schema = schema([Field::new(
581 "j",
582 DataType::Struct(Fields::from(vec![Arc::new(Field::new(
583 "a",
584 DataType::Struct(Fields::from(vec![
585 Arc::new(Field::new("b", DataType::UInt64, true)),
586 Arc::new(Field::new("c", DataType::Utf8View, true)),
587 ])),
588 true,
589 ))])),
590 true,
591 )
592 .with_extension_type(Json2ExtensionType::default())]);
593 let mut aligner = JsonSchemaAligner::new(
594 stream::iter([Ok(input)]),
595 vec![true],
596 output_schema.clone(),
597 AlignMode::AlignToSchema,
598 )
599 .unwrap();
600 let output = aligner.next().await.unwrap().unwrap();
601
602 assert_eq!(output_schema, output.schema());
603 let j = output
604 .column(0)
605 .as_any()
606 .downcast_ref::<StructArray>()
607 .unwrap();
608 let a = j.column(0).as_any().downcast_ref::<StructArray>().unwrap();
609 assert_eq!(
610 &[None, Some(2), None],
611 a.column(0)
612 .as_any()
613 .downcast_ref::<datatypes::arrow::array::UInt64Array>()
614 .unwrap()
615 .iter()
616 .collect::<Vec<_>>()
617 .as_slice()
618 );
619 assert_eq!(
620 &[None, None, None],
621 a.column(1)
622 .as_any()
623 .downcast_ref::<StringViewArray>()
624 .unwrap()
625 .iter()
626 .collect::<Vec<_>>()
627 .as_slice()
628 );
629 }
630
631 #[tokio::test]
632 async fn test_json_schema_aligner_preserves_struct_siblings() {
633 let source_values = [
634 Some(parse_string_to_jsonb(r#"{"x":1}"#).unwrap()),
635 Some(parse_string_to_jsonb(r#"{"x":2}"#).unwrap()),
636 ];
637 let source = Arc::new(BinaryArray::from_iter(
638 source_values.iter().map(|value| value.as_deref()),
639 )) as ArrayRef;
640 let c = Arc::new(Int64Array::from_iter_values([10, 20])) as ArrayRef;
641
642 let a_fields = Fields::from(vec![
643 Arc::new(Field::new("b", DataType::Binary, true)),
644 Arc::new(Field::new("c", DataType::Int64, true)),
645 ]);
646 let input_fields = Fields::from(vec![Arc::new(Field::new(
647 "a",
648 DataType::Struct(a_fields.clone()),
649 true,
650 ))]);
651 let input = RecordBatch::try_new(
652 schema([Field::new(
653 "j",
654 DataType::Struct(input_fields.clone()),
655 true,
656 )]),
657 vec![Arc::new(StructArray::new(
658 input_fields,
659 vec![Arc::new(StructArray::new(a_fields, vec![source, c], None))],
660 None,
661 ))],
662 )
663 .unwrap();
664
665 let output_schema = schema([Field::new(
666 "j",
667 DataType::Struct(Fields::from(vec![Arc::new(Field::new(
668 "a",
669 DataType::Struct(Fields::from(vec![
670 Arc::new(Field::new(
671 "b",
672 DataType::Struct(Fields::from(vec![Arc::new(Field::new(
673 "x",
674 DataType::Int64,
675 true,
676 ))])),
677 true,
678 )),
679 Arc::new(Field::new("c", DataType::Int64, true)),
680 ])),
681 true,
682 ))])),
683 true,
684 )
685 .with_extension_type(Json2ExtensionType::default())]);
686 let mut aligner = JsonSchemaAligner::new(
687 stream::iter([Ok(input)]),
688 vec![true],
689 output_schema.clone(),
690 AlignMode::AlignToSchema,
691 )
692 .unwrap();
693 let output = aligner.next().await.unwrap().unwrap();
694
695 assert_eq!(output_schema, output.schema());
696 let j = output
697 .column(0)
698 .as_any()
699 .downcast_ref::<StructArray>()
700 .unwrap();
701 let a = j.column(0).as_any().downcast_ref::<StructArray>().unwrap();
702 let b = a.column(0).as_any().downcast_ref::<StructArray>().unwrap();
703 assert_eq!(
704 &[Some(1), Some(2)],
705 b.column(0)
706 .as_any()
707 .downcast_ref::<Int64Array>()
708 .unwrap()
709 .iter()
710 .collect::<Vec<_>>()
711 .as_slice()
712 );
713 assert_eq!(
714 &[Some(10), Some(20)],
715 a.column(1)
716 .as_any()
717 .downcast_ref::<Int64Array>()
718 .unwrap()
719 .iter()
720 .collect::<Vec<_>>()
721 .as_slice()
722 );
723 }
724
725 #[tokio::test]
726 async fn test_rewrite_multiple_columns_and_fill_missing_roots() {
727 let logical_settings = JsonSettings::default();
728 let target_layout = JsonSettings::try_new(vec![], Some(0)).unwrap();
729 let target_type = json2_physical_data_type(&target_layout);
730 let output_schema = schema([
731 Field::new("j", target_type.clone(), true)
732 .with_extension_type(Json2ExtensionType::default()),
733 Field::new("missing", target_type.clone(), true)
734 .with_extension_type(Json2ExtensionType::default()),
735 Field::new("k", target_type.clone(), true)
736 .with_extension_type(Json2ExtensionType::default()),
737 Field::new("a", DataType::Int64, true),
738 ]);
739 let values = [Some(parse_string_to_jsonb(r#"{"x":1}"#).unwrap()), None];
740 let source = Arc::new(BinaryArray::from_iter(
741 values.iter().map(|value| value.as_deref()),
742 )) as ArrayRef;
743 let source_field = Field::new("j", DataType::Binary, true)
744 .with_extension_type(Json2ExtensionType::default());
745 let expected = JsonArray::from(&source)
746 .rewrite_to_v2(&source_field, &logical_settings, &target_layout)
747 .unwrap();
748 let input = RecordBatch::try_new(
749 schema([
750 source_field,
751 Field::new("k", DataType::Binary, true)
752 .with_extension_type(Json2ExtensionType::default()),
753 Field::new("a", DataType::Int64, true),
754 ]),
755 vec![source.clone(), source, int_array([10, 20])],
756 )
757 .unwrap();
758 let columns = ["j", "missing", "k"]
759 .into_iter()
760 .map(|name| {
761 (
762 name.to_string(),
763 Json2TargetLayout {
764 extension_metadata: serde_json::to_string(&JsonMetadata::new(
765 logical_settings.clone(),
766 ))
767 .unwrap(),
768 target_layout: target_layout.clone(),
769 },
770 )
771 })
772 .collect();
773 let mut aligner = JsonSchemaAligner::new(
774 stream::iter([Ok(input)]),
775 vec![true, false, true, true],
776 output_schema.clone(),
777 AlignMode::Rewrite { columns },
778 )
779 .unwrap();
780 let output = aligner.next().await.unwrap().unwrap();
781 assert_eq!(output_schema, output.schema());
782 assert_eq!(expected.as_ref(), output.column(0).as_ref());
783 assert_eq!(expected.as_ref(), output.column(2).as_ref());
784 assert_eq!(&target_type, output.column(1).data_type());
785 assert_eq!(2, output.column(1).null_count());
786 assert_eq!(int_array([10, 20]).as_ref(), output.column(3).as_ref());
787 }
788
789 #[tokio::test]
790 async fn test_rewrite_keeps_rows_with_invalid_json2_settings() {
791 let settings = JsonSettings::try_new(
792 vec![JsonTypeHint {
793 path: vec!["kind".to_string()],
794 data_type: ConcreteDataType::string_datatype(),
795 inverted_index: false,
796 }],
797 Some(0),
798 )
799 .unwrap();
800 let target_type = json2_physical_data_type(&settings);
801 let output_schema = schema([
802 Field::new("j", target_type.clone(), true).with_extension_type(
803 Json2ExtensionType::new(Arc::new(JsonMetadata::new(settings.clone()))),
804 ),
805 Field::new("value", DataType::Int64, true),
806 ]);
807 let source = Arc::new(BinaryArray::from_iter([
808 Some(parse_string_to_jsonb(r#"{"kind":"valid"}"#).unwrap()),
809 Some(parse_string_to_jsonb(r#"{"kind":1}"#).unwrap()),
810 ])) as ArrayRef;
811 let input = RecordBatch::try_new(
812 schema([
813 Field::new("j", DataType::Binary, true)
814 .with_extension_type(Json2ExtensionType::default()),
815 Field::new("value", DataType::Int64, true),
816 ]),
817 vec![source, int_array([10, 20])],
818 )
819 .unwrap();
820 let columns = HashMap::from([(
821 "j".to_string(),
822 Json2TargetLayout {
823 extension_metadata: serde_json::to_string(&JsonMetadata::new(settings.clone()))
824 .unwrap(),
825 target_layout: settings,
826 },
827 )]);
828 let mut aligner = JsonSchemaAligner::new(
829 stream::iter([Ok(input)]),
830 vec![true, true],
831 output_schema,
832 AlignMode::Rewrite { columns },
833 )
834 .unwrap();
835
836 let output = aligner.next().await.unwrap().unwrap();
837 assert_eq!(2, output.num_rows());
838 assert_eq!(
839 10,
840 output
841 .column(1)
842 .as_any()
843 .downcast_ref::<Int64Array>()
844 .unwrap()
845 .value(0)
846 );
847 assert_eq!(
848 20,
849 output
850 .column(1)
851 .as_any()
852 .downcast_ref::<Int64Array>()
853 .unwrap()
854 .value(1)
855 );
856 }
857
858 #[test]
859 fn test_rewrite_rejects_mismatched_output_layout() {
860 let columns = HashMap::from([(
861 "j".to_string(),
862 Json2TargetLayout {
863 extension_metadata: serde_json::to_string(&JsonMetadata::new(
864 JsonSettings::default(),
865 ))
866 .unwrap(),
867 target_layout: JsonSettings::try_new(vec![], Some(0)).unwrap(),
868 },
869 )]);
870 let result = JsonSchemaAligner::new(
871 stream::empty::<Result<RecordBatch>>(),
872 vec![false],
873 schema([Field::new("j", DataType::Binary, true)
874 .with_extension_type(Json2ExtensionType::default())]),
875 AlignMode::Rewrite { columns },
876 );
877 assert!(
878 result
879 .unwrap_err()
880 .to_string()
881 .contains("does not match output field")
882 );
883 }
884
885 #[tokio::test]
886 async fn test_empty_rewrite_only_fills_missing_roots() {
887 let source = int_array([10, 20]);
888 let input = RecordBatch::try_new(
889 schema([Field::new("a", DataType::Int64, true)]),
890 vec![source.clone()],
891 )
892 .unwrap();
893 let output_schema = schema([
894 Field::new("missing", DataType::Utf8, true),
895 Field::new("a", DataType::Int64, true),
896 ]);
897 let mut aligner = JsonSchemaAligner::new(
898 stream::iter([Ok(input)]),
899 vec![false, true],
900 output_schema.clone(),
901 AlignMode::Rewrite {
902 columns: HashMap::new(),
903 },
904 )
905 .unwrap();
906 let output = aligner.next().await.unwrap().unwrap();
907 assert_eq!(output_schema, output.schema());
908 assert_eq!(2, output.num_rows());
909 assert_eq!(2, output.column(0).null_count());
910 assert!(Arc::ptr_eq(&source, output.column(1)));
911 }
912
913 #[tokio::test]
914 async fn test_empty_rewrite_does_not_align_existing_struct() {
915 let source = Arc::new(StructArray::from(vec![(
916 Arc::new(Field::new("x", DataType::Int64, true)),
917 int_array([1, 2]),
918 )])) as ArrayRef;
919 let input = RecordBatch::try_new(
920 schema([Field::new("j", source.data_type().clone(), true)]),
921 vec![source],
922 )
923 .unwrap();
924 let output_schema = schema([Field::new(
925 "j",
926 DataType::Struct(Fields::from(vec![
927 Field::new("x", DataType::Int64, true),
928 Field::new("y", DataType::Utf8, true),
929 ])),
930 true,
931 )]);
932 let mut aligner = JsonSchemaAligner::new(
933 stream::iter([Ok(input)]),
934 vec![true],
935 output_schema,
936 AlignMode::Rewrite {
937 columns: HashMap::new(),
938 },
939 )
940 .unwrap();
941 assert!(matches!(
942 aligner.next().await.unwrap(),
943 Err(crate::error::Error::NewRecordBatch { .. })
944 ));
945 }
946
947 fn schema(fields: impl IntoIterator<Item = Field>) -> SchemaRef {
948 Arc::new(Schema::new(fields.into_iter().collect::<Vec<_>>()))
949 }
950
951 fn int_array(values: impl IntoIterator<Item = i64>) -> ArrayRef {
952 Arc::new(Int64Array::from_iter_values(values))
953 }
954
955 fn string_array(values: impl IntoIterator<Item = &'static str>) -> ArrayRef {
956 Arc::new(StringArray::from_iter_values(values))
957 }
958}