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::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/// Aligns projected batches to the expected output schema for nested projections.
33///
34/// Background
35/// ----------
36/// Nested projection may ask parquet to read leaves under a root column. If none
37/// of the requested leaves exists in the current parquet file, parquet decoding
38/// omits the whole root from the physical [`RecordBatch`].
39///
40/// In addition, after nested-path filtering, returned struct arrays may contain
41/// only a subset of fields. The current output schema is not pruned by nested
42/// paths, so physical struct fields can be a subset of the expected struct
43/// fields, and their nested schema can differ from the expected output schema.
44///
45/// To keep projected batches schema-consistent before entering upper readers:
46/// - Root-column presence alignment restores missing projected root columns by
47///   inserting root-level null arrays.
48/// - Nested struct alignment aligns struct arrays to the expected nested field
49///   layout.
50#[derive(derive_more::Debug)]
51pub struct NestedSchemaAligner<S> {
52    #[debug(skip)]
53    inner: S,
54    /// Output schema expected by the upper reader.
55    output_schema: SchemaRef,
56    /// Whether each projected root exists in the physical batch returned by
57    /// parquet.
58    projected_root_presence: Vec<bool>,
59    /// Number of columns expected from the physical batch returned by parquet.
60    expected_input_col_num: usize,
61    /// Whether all projected roots are present and the stream can pass batches
62    /// through.
63    all_roots_present: bool,
64    /// The cache for whether incoming batches already match output schema.
65    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}