Skip to main content

mito2/sst/parquet/json_align/
stream.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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/// Aligns projected batches to the expected output schema for nested projections.
42///
43/// Background
44/// ----------
45/// Nested projection may ask parquet to read leaves under a root column. If none
46/// of the requested leaves exists in the current parquet file, parquet decoding
47/// omits the whole root from the physical [`RecordBatch`].
48///
49/// In addition, after nested-path filtering, returned struct arrays may contain
50/// only a subset of fields. The current output schema is not pruned by nested
51/// paths, so physical struct fields can be a subset of the expected struct
52/// fields, and their nested schema can differ from the expected output schema.
53///
54/// To keep projected batches schema-consistent before entering upper readers:
55/// - Root-column presence alignment restores missing projected root columns by
56///   inserting root-level null arrays.
57/// - Nested struct alignment aligns struct arrays to the expected nested field
58///   layout.
59#[derive(derive_more::Debug)]
60pub struct NestedSchemaAligner<S> {
61    #[debug(skip)]
62    inner: S,
63    /// Output schema expected by the upper reader.
64    output_schema: SchemaRef,
65    /// Whether each projected root exists in the physical batch returned by
66    /// parquet.
67    projected_root_presence: Vec<bool>,
68    /// Number of columns expected from the physical batch returned by parquet.
69    expected_input_col_num: usize,
70    /// Whether all projected roots are present and the stream can pass batches
71    /// through.
72    all_roots_present: bool,
73    /// JSON2 columns that require semantic source-to-target layout rewriting.
74    json2_rewrite_targets: HashMap<String, Json2RewriteSettings>,
75    /// The cache for whether incoming batches already match output schema.
76    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    /// Sets JSON2 columns that must be rewritten into the output field layout.
116    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}