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, SchemaRef};
23use datatypes::arrow::record_batch::RecordBatch;
24use datatypes::extension::json::{JsonMetadata, is_json2_extension_type};
25use datatypes::json::JsonSettings;
26use datatypes::vectors::json::array::JsonArray;
27use futures::Stream;
28use snafu::{ResultExt, ensure};
29
30use crate::error::{
31 CastColumnSnafu, DataTypeMismatchSnafu, NewRecordBatchSnafu, Result, UnexpectedSnafu,
32};
33use crate::sst::parquet::Json2TargetLayout;
34
35#[derive(Debug)]
36struct Json2RewriteSettings {
37 logical_settings: JsonSettings,
38 target_layout: JsonSettings,
39}
40
41#[derive(derive_more::Debug)]
60pub struct NestedSchemaAligner<S> {
61 #[debug(skip)]
62 inner: S,
63 output_schema: SchemaRef,
65 projected_root_presence: Vec<bool>,
68 expected_input_col_num: usize,
70 all_roots_present: bool,
73 json2_rewrite_targets: HashMap<String, Json2RewriteSettings>,
75 is_schema_matched: Option<bool>,
77}
78
79impl<S> NestedSchemaAligner<S>
80where
81 S: Stream<Item = Result<RecordBatch>>,
82{
83 pub fn new(
84 inner: S,
85 projected_root_presence: Vec<bool>,
86 output_schema: SchemaRef,
87 ) -> Result<NestedSchemaAligner<S>> {
88 ensure!(
89 projected_root_presence.len() == output_schema.fields().len(),
90 UnexpectedSnafu {
91 reason: format!(
92 "NestedSchemaAligner projected root presence len {} does not match output schema columns {}",
93 projected_root_presence.len(),
94 output_schema.fields().len()
95 ),
96 }
97 );
98
99 let expected_input_col_num = projected_root_presence
100 .iter()
101 .filter(|matched| **matched)
102 .count();
103 let all_roots_present = projected_root_presence.iter().all(|&m| m);
104 Ok(NestedSchemaAligner {
105 inner,
106 output_schema,
107 projected_root_presence,
108 expected_input_col_num,
109 all_roots_present,
110 json2_rewrite_targets: HashMap::new(),
111 is_schema_matched: None,
112 })
113 }
114
115 pub(crate) fn with_json2_rewrite_targets(
117 mut self,
118 targets: &HashMap<String, Json2TargetLayout>,
119 ) -> Result<Self> {
120 self.json2_rewrite_targets = targets
121 .iter()
122 .map(|(name, layout)| {
123 let metadata = serde_json::from_str::<JsonMetadata>(&layout.extension_metadata)
124 .map_err(|e| {
125 UnexpectedSnafu {
126 reason: format!(
127 "invalid JSON2 extension metadata for column '{name}': {e}"
128 ),
129 }
130 .build()
131 })?;
132 Ok((
133 name.clone(),
134 Json2RewriteSettings {
135 logical_settings: metadata.into_json_settings(),
136 target_layout: layout.target_layout.clone(),
137 },
138 ))
139 })
140 .collect::<Result<_>>()?;
141 Ok(self)
142 }
143}
144
145impl<S> Stream for NestedSchemaAligner<S>
146where
147 S: Stream<Item = Result<RecordBatch>> + Unpin,
148{
149 type Item = Result<RecordBatch>;
150
151 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
152 let this = self.get_mut();
153
154 match Pin::new(&mut this.inner).poll_next(cx) {
155 Poll::Ready(Some(Ok(rb))) => {
156 let is_schema_matched = this.all_roots_present
157 && *this
158 .is_schema_matched
159 .get_or_insert_with(|| rb.schema() == this.output_schema);
160
161 if is_schema_matched {
162 Poll::Ready(Some(Ok(rb)))
163 } else {
164 Poll::Ready(Some(align_projected_batch(
165 rb,
166 &this.output_schema,
167 &this.projected_root_presence,
168 this.expected_input_col_num,
169 &this.json2_rewrite_targets,
170 )))
171 }
172 }
173 Poll::Ready(Some(Err(err))) => Poll::Ready(Some(Err(err))),
174 Poll::Ready(None) => Poll::Ready(None),
175 Poll::Pending => Poll::Pending,
176 }
177 }
178}
179
180fn align_projected_batch(
181 rb: RecordBatch,
182 output_schema: &SchemaRef,
183 projected_root_presence: &[bool],
184 expected_input_col_num: usize,
185 json2_rewrite_targets: &HashMap<String, Json2RewriteSettings>,
186) -> Result<RecordBatch> {
187 ensure!(
188 rb.columns().len() == expected_input_col_num,
189 UnexpectedSnafu {
190 reason: format!(
191 "NestedSchemaAligner expected {} input columns but got {}",
192 expected_input_col_num,
193 rb.columns().len()
194 ),
195 }
196 );
197
198 let mut cols = Vec::with_capacity(projected_root_presence.len());
199 let mut idx = 0;
200 let input_schema = rb.schema_ref();
201
202 for (field, present) in output_schema.fields().iter().zip(projected_root_presence) {
203 if !present {
204 cols.push(new_null_array(field.data_type(), rb.num_rows()));
205 continue;
206 }
207
208 cols.push(align_array(
209 rb.column(idx),
210 input_schema.field(idx),
211 field,
212 json2_rewrite_targets.get(field.name()),
213 )?);
214 idx += 1;
215 }
216
217 RecordBatch::try_new(output_schema.clone(), cols).context(NewRecordBatchSnafu)
218}
219
220fn align_array(
221 array: &ArrayRef,
222 source: &Field,
223 field: &FieldRef,
224 rewrite_settings: Option<&Json2RewriteSettings>,
225) -> Result<ArrayRef> {
226 if let Some(settings) = rewrite_settings {
227 return JsonArray::from(array)
228 .rewrite_to_v2(source, &settings.logical_settings, &settings.target_layout)
229 .context(DataTypeMismatchSnafu);
230 }
231 if array.data_type() == field.data_type() {
232 return Ok(array.clone());
233 }
234
235 if is_json2_extension_type(field) {
236 if is_json2_extension_type(source) {
237 return JsonArray::from(array)
238 .project_to_v2(source, field.data_type())
239 .context(DataTypeMismatchSnafu);
240 }
241 return JsonArray::from(array)
242 .project_to(field.data_type())
243 .context(DataTypeMismatchSnafu);
244 }
245
246 if !matches!(field.data_type(), DataType::Struct(_)) {
247 return Ok(array.clone());
248 }
249
250 cast_column(array, field.as_ref(), &DEFAULT_CAST_OPTIONS).context(CastColumnSnafu)
251}
252
253#[cfg(test)]
254mod tests {
255 use std::collections::HashMap;
256 use std::sync::Arc;
257
258 use datatypes::arrow::array::{
259 Array, ArrayRef, BinaryArray, Int64Array, StringArray, StringViewArray, StructArray,
260 };
261 use datatypes::arrow::datatypes::{DataType, Field, Fields, Schema};
262 use datatypes::extension::json::Json2ExtensionType;
263 use datatypes::types::parse_string_to_jsonb;
264 use futures::{StreamExt, stream};
265
266 use super::*;
267
268 #[test]
269 fn test_aligner_resolves_json2_rewrite_settings()
270 -> std::result::Result<(), Box<dyn std::error::Error>> {
271 let logical_settings = JsonSettings::default();
272 let target_layout = JsonSettings::try_new(vec![], Some(0))?;
273 let rewrite_targets = HashMap::from([(
274 "j".to_string(),
275 Json2TargetLayout {
276 extension_metadata: serde_json::to_string(&JsonMetadata::new(
277 logical_settings.clone(),
278 ))?,
279 target_layout: target_layout.clone(),
280 },
281 )]);
282 let aligner = NestedSchemaAligner::new(
283 stream::empty::<Result<RecordBatch>>(),
284 vec![],
285 schema(Vec::<Field>::new()),
286 )?
287 .with_json2_rewrite_targets(&rewrite_targets)?;
288
289 let settings = &aligner.json2_rewrite_targets["j"];
290 assert_eq!(logical_settings, settings.logical_settings);
291 assert_eq!(target_layout, settings.target_layout);
292 Ok(())
293 }
294
295 #[tokio::test]
296 async fn test_aligner_with_all_projected_roots_match() {
297 let output_schema = schema([
298 Field::new("a", DataType::Int64, true),
299 Field::new("b", DataType::Utf8, true),
300 ]);
301 let input = RecordBatch::try_new(
302 output_schema.clone(),
303 vec![int_array([1, 2, 3]), string_array(["x", "y", "z"])],
304 )
305 .unwrap();
306 let stream = stream::iter([Ok(input.clone())]);
307
308 let mut aligner =
309 NestedSchemaAligner::new(stream, vec![true, true], output_schema.clone()).unwrap();
310 let output = aligner.next().await.unwrap().unwrap();
311
312 assert_eq!(input, output);
313 assert!(aligner.next().await.is_none());
314 }
315
316 #[tokio::test]
317 async fn test_aligner_with_fills_null_root_columns() {
318 let input_schema = schema([Field::new("a", DataType::Int64, true)]);
319 let output_schema = schema([
320 Field::new("a", DataType::Int64, true),
321 Field::new("missing", DataType::Utf8, true),
322 Field::new("c", DataType::Int64, true),
323 ]);
324 let input = RecordBatch::try_new(input_schema, vec![int_array([10, 20])]).unwrap();
325 let stream = stream::iter([Ok(input)]);
326
327 let mut aligner =
328 NestedSchemaAligner::new(stream, vec![true, false, false], output_schema.clone())
329 .unwrap();
330 let output = aligner.next().await.unwrap().unwrap();
331
332 assert_eq!(output_schema, output.schema());
333 assert_eq!(3, output.num_columns());
334 assert_eq!(
335 &[Some(10), Some(20)],
336 output
337 .column(0)
338 .as_any()
339 .downcast_ref::<Int64Array>()
340 .unwrap()
341 .iter()
342 .collect::<Vec<_>>()
343 .as_slice()
344 );
345 assert_eq!(DataType::Utf8, *output.column(1).data_type());
346 assert_eq!(output.num_rows(), output.column(1).null_count());
347 assert_eq!(DataType::Int64, *output.column(2).data_type());
348 assert_eq!(output.num_rows(), output.column(2).null_count());
349 }
350
351 #[tokio::test]
352 async fn test_aligner_with_fills_missing_struct_root_column() {
353 let input_schema = schema([Field::new("a", DataType::Int64, true)]);
354 let struct_type = DataType::Struct(Fields::from(vec![
355 Field::new("x", DataType::Int64, true),
356 Field::new("y", DataType::Utf8, true),
357 ]));
358 let output_schema = schema([
359 Field::new("a", DataType::Int64, true),
360 Field::new("missing_struct", struct_type.clone(), true),
361 ]);
362 let input = RecordBatch::try_new(input_schema, vec![int_array([10, 20])]).unwrap();
363 let stream = stream::iter([Ok(input)]);
364
365 let mut aligner =
366 NestedSchemaAligner::new(stream, vec![true, false], output_schema.clone()).unwrap();
367 let output = aligner.next().await.unwrap().unwrap();
368
369 assert_eq!(output_schema, output.schema());
370 assert_eq!(2, output.num_columns());
371 assert_eq!(struct_type, output.column(1).data_type().clone());
372 assert_eq!(output.num_rows(), output.column(1).null_count());
373 }
374
375 #[tokio::test]
376 async fn test_aligner_reject_projection_len_mismatch() {
377 let output_schema = schema([Field::new("a", DataType::Int64, true)]);
378 let stream = stream::iter([]);
379
380 let err = match NestedSchemaAligner::new(stream, vec![true, false], output_schema) {
381 Ok(_) => panic!("NestedSchemaAligner should reject projection length mismatch"),
382 Err(err) => err,
383 };
384
385 assert!(
386 err.to_string()
387 .contains("projected root presence len 2 does not match output schema columns 1")
388 );
389 }
390
391 #[tokio::test]
392 async fn test_aligner_reject_with_input_column_mismatch() {
393 let input_schema = schema([Field::new("a", DataType::Int64, true)]);
394 let output_schema = schema([
395 Field::new("a", DataType::Int64, true),
396 Field::new("b", DataType::Int64, true),
397 Field::new("missing", DataType::Int64, true),
398 ]);
399 let input = RecordBatch::try_new(input_schema, vec![int_array([1, 2])]).unwrap();
400 let stream = stream::iter([Ok(input)]);
401
402 let mut aligner =
403 NestedSchemaAligner::new(stream, vec![true, true, false], output_schema).unwrap();
404 let err = aligner.next().await.unwrap().unwrap_err();
405
406 assert!(
407 err.to_string()
408 .contains("expected 2 input columns but got 1")
409 );
410 }
411
412 #[tokio::test]
413 async fn test_nested_schema_aligner_aligns_struct_field() {
414 let output_schema = schema([Field::new(
415 "nested",
416 DataType::Struct(Fields::from(vec![
417 Field::new("x", DataType::Int64, true),
418 Field::new("y", DataType::Utf8, true),
419 ])),
420 true,
421 )]);
422 let input = RecordBatch::try_new(
423 schema([Field::new(
424 "nested",
425 DataType::Struct(Fields::from(vec![Field::new("x", DataType::Int64, true)])),
426 true,
427 )]),
428 vec![Arc::new(StructArray::from(vec![(
429 Arc::new(Field::new("x", DataType::Int64, true)),
430 int_array([1, 2]),
431 )]))],
432 )
433 .unwrap();
434
435 let mut aligner =
436 NestedSchemaAligner::new(stream::iter([Ok(input)]), vec![true], output_schema.clone())
437 .unwrap();
438 let output = aligner.next().await.unwrap().unwrap();
439
440 assert_eq!(output_schema, output.schema());
441 let nested = output
442 .column(0)
443 .as_any()
444 .downcast_ref::<StructArray>()
445 .unwrap();
446 assert_eq!(2, nested.columns().len());
447 assert_eq!(2, nested.column(1).null_count());
448 }
449
450 #[tokio::test]
451 async fn test_nested_schema_aligner_decodes_variant_to_struct() {
452 let source_values = [
453 Some(parse_string_to_jsonb("1").unwrap()),
454 Some(parse_string_to_jsonb(r#"{"b":2}"#).unwrap()),
455 None,
456 ];
457 let source = Arc::new(BinaryArray::from_iter(
458 source_values.iter().map(|value| value.as_deref()),
459 )) as ArrayRef;
460 let input_fields = Fields::from(vec![Arc::new(Field::new("a", DataType::Binary, true))]);
461 let input = RecordBatch::try_new(
462 schema([Field::new(
463 "j",
464 DataType::Struct(input_fields.clone()),
465 true,
466 )]),
467 vec![Arc::new(StructArray::new(input_fields, vec![source], None))],
468 )
469 .unwrap();
470
471 let output_schema = schema([Field::new(
472 "j",
473 DataType::Struct(Fields::from(vec![Arc::new(Field::new(
474 "a",
475 DataType::Struct(Fields::from(vec![
476 Arc::new(Field::new("b", DataType::UInt64, true)),
477 Arc::new(Field::new("c", DataType::Utf8View, true)),
478 ])),
479 true,
480 ))])),
481 true,
482 )
483 .with_extension_type(Json2ExtensionType::default())]);
484 let mut aligner =
485 NestedSchemaAligner::new(stream::iter([Ok(input)]), vec![true], output_schema.clone())
486 .unwrap();
487 let output = aligner.next().await.unwrap().unwrap();
488
489 assert_eq!(output_schema, output.schema());
490 let j = output
491 .column(0)
492 .as_any()
493 .downcast_ref::<StructArray>()
494 .unwrap();
495 let a = j.column(0).as_any().downcast_ref::<StructArray>().unwrap();
496 assert_eq!(
497 &[None, Some(2), None],
498 a.column(0)
499 .as_any()
500 .downcast_ref::<datatypes::arrow::array::UInt64Array>()
501 .unwrap()
502 .iter()
503 .collect::<Vec<_>>()
504 .as_slice()
505 );
506 assert_eq!(
507 &[None, None, None],
508 a.column(1)
509 .as_any()
510 .downcast_ref::<StringViewArray>()
511 .unwrap()
512 .iter()
513 .collect::<Vec<_>>()
514 .as_slice()
515 );
516 }
517
518 #[tokio::test]
519 async fn test_nested_schema_aligner_preserves_struct_siblings() {
520 let source_values = [
521 Some(parse_string_to_jsonb(r#"{"x":1}"#).unwrap()),
522 Some(parse_string_to_jsonb(r#"{"x":2}"#).unwrap()),
523 ];
524 let source = Arc::new(BinaryArray::from_iter(
525 source_values.iter().map(|value| value.as_deref()),
526 )) as ArrayRef;
527 let c = Arc::new(Int64Array::from_iter_values([10, 20])) as ArrayRef;
528
529 let a_fields = Fields::from(vec![
530 Arc::new(Field::new("b", DataType::Binary, true)),
531 Arc::new(Field::new("c", DataType::Int64, true)),
532 ]);
533 let input_fields = Fields::from(vec![Arc::new(Field::new(
534 "a",
535 DataType::Struct(a_fields.clone()),
536 true,
537 ))]);
538 let input = RecordBatch::try_new(
539 schema([Field::new(
540 "j",
541 DataType::Struct(input_fields.clone()),
542 true,
543 )]),
544 vec![Arc::new(StructArray::new(
545 input_fields,
546 vec![Arc::new(StructArray::new(a_fields, vec![source, c], None))],
547 None,
548 ))],
549 )
550 .unwrap();
551
552 let output_schema = schema([Field::new(
553 "j",
554 DataType::Struct(Fields::from(vec![Arc::new(Field::new(
555 "a",
556 DataType::Struct(Fields::from(vec![
557 Arc::new(Field::new(
558 "b",
559 DataType::Struct(Fields::from(vec![Arc::new(Field::new(
560 "x",
561 DataType::Int64,
562 true,
563 ))])),
564 true,
565 )),
566 Arc::new(Field::new("c", DataType::Int64, true)),
567 ])),
568 true,
569 ))])),
570 true,
571 )
572 .with_extension_type(Json2ExtensionType::default())]);
573 let mut aligner =
574 NestedSchemaAligner::new(stream::iter([Ok(input)]), vec![true], output_schema.clone())
575 .unwrap();
576 let output = aligner.next().await.unwrap().unwrap();
577
578 assert_eq!(output_schema, output.schema());
579 let j = output
580 .column(0)
581 .as_any()
582 .downcast_ref::<StructArray>()
583 .unwrap();
584 let a = j.column(0).as_any().downcast_ref::<StructArray>().unwrap();
585 let b = a.column(0).as_any().downcast_ref::<StructArray>().unwrap();
586 assert_eq!(
587 &[Some(1), Some(2)],
588 b.column(0)
589 .as_any()
590 .downcast_ref::<Int64Array>()
591 .unwrap()
592 .iter()
593 .collect::<Vec<_>>()
594 .as_slice()
595 );
596 assert_eq!(
597 &[Some(10), Some(20)],
598 a.column(1)
599 .as_any()
600 .downcast_ref::<Int64Array>()
601 .unwrap()
602 .iter()
603 .collect::<Vec<_>>()
604 .as_slice()
605 );
606 }
607
608 fn schema(fields: impl IntoIterator<Item = Field>) -> SchemaRef {
609 Arc::new(Schema::new(fields.into_iter().collect::<Vec<_>>()))
610 }
611
612 fn int_array(values: impl IntoIterator<Item = i64>) -> ArrayRef {
613 Arc::new(Int64Array::from_iter_values(values))
614 }
615
616 fn string_array(values: impl IntoIterator<Item = &'static str>) -> ArrayRef {
617 Arc::new(StringArray::from_iter_values(values))
618 }
619}