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