Skip to main content

mito2/memtable/bulk/
json_align.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::sync::Arc;
17
18use datatypes::arrow::datatypes::{DataType as ArrowDataType, Schema, SchemaRef};
19use datatypes::arrow::record_batch::RecordBatch;
20use datatypes::extension::json::is_json2_extension_type;
21use datatypes::types::json_type::JsonNativeType;
22use datatypes::vectors::json::array::JsonArray;
23use snafu::{OptionExt, ResultExt};
24
25use crate::error::{
26    ConvertValueSnafu, DataTypeMismatchSnafu, NewRecordBatchSnafu, Result, UnexpectedSnafu,
27};
28use crate::memtable::BoxedRecordBatchIterator;
29
30/// Aligns concrete JSON2 Arrow types across record batches.
31///
32/// JSON2 column concrete Arrow types are derived from data. Different memtable
33/// parts may therefore have different concrete types for the same JSON2 column.
34/// This helper merges those concrete types and aligns batches to the merged schema.
35#[derive(Clone)]
36pub(crate) struct Json2Aligner {
37    /// Schema after merging all JSON2 column concrete types.
38    schema: SchemaRef,
39    /// JSON2 columns that may need per-batch alignment.
40    json_columns: Vec<(usize, ArrowDataType)>,
41}
42
43impl Json2Aligner {
44    /// Builds an aligner from input schemas.
45    ///
46    /// Note: except for JSON2 columns, all input schemas must be identical.
47    pub(crate) fn try_new<I>(input_schemas: I) -> Result<Self>
48    where
49        I: IntoIterator<Item = SchemaRef>,
50    {
51        let mut input_schemas = input_schemas.into_iter();
52
53        // Use first schema as base: it defines column order and non-JSON types.
54        let base_schema = input_schemas.next().context(UnexpectedSnafu {
55            reason: "Json2Aligner requires at least one input schema",
56        })?;
57
58        // Init merged types from base schema.
59        let mut merged_types = base_schema
60            .fields()
61            .iter()
62            .enumerate()
63            .filter(|&(_idx, field)| is_json2_extension_type(field))
64            .map(|(idx, field)| {
65                let json_type =
66                    JsonNativeType::try_from(field.data_type()).context(DataTypeMismatchSnafu)?;
67                Ok((idx, json_type))
68            })
69            .collect::<Result<HashMap<usize, JsonNativeType>>>()?;
70
71        // No JSON2 columns, no alignment needed.
72        if merged_types.is_empty() {
73            return Ok(Self {
74                schema: base_schema,
75                json_columns: Vec::new(),
76            });
77        }
78
79        // Merge JSON2 types from remaining schemas.
80        for schema in input_schemas {
81            // Input schemas should only differ in JSON2 concrete types.
82            #[cfg(debug_assertions)]
83            assert_columns_match_except_json2(&base_schema, &schema);
84
85            for (idx, merged) in &mut merged_types {
86                if *idx >= schema.fields().len() {
87                    continue;
88                }
89                let json_type = JsonNativeType::try_from(schema.field(*idx).data_type())
90                    .context(DataTypeMismatchSnafu)?;
91                merged.merge(&json_type);
92            }
93        }
94
95        // Build output schema with merged JSON2 types.
96        let mut json_columns = Vec::with_capacity(merged_types.len());
97        let fields: Vec<_> = base_schema
98            .fields()
99            .iter()
100            .enumerate()
101            .map(|(idx, field)| {
102                if let Some(merged) = merged_types.get(&idx) {
103                    let data_type = merged.as_arrow_type();
104                    json_columns.push((idx, data_type.clone()));
105                    let mut field = (**field).clone();
106                    field.set_data_type(data_type);
107                    Arc::new(field)
108                } else {
109                    field.clone()
110                }
111            })
112            .collect();
113
114        let schema = Arc::new(Schema::new_with_metadata(
115            fields,
116            base_schema.metadata().clone(),
117        ));
118
119        Ok(Self {
120            schema,
121            json_columns,
122        })
123    }
124
125    /// Returns the aligned output schema.
126    pub(crate) fn schema(&self) -> &SchemaRef {
127        &self.schema
128    }
129
130    /// Aligns a [`RecordBatch`] to [`Self::schema`].
131    pub(crate) fn align_batch(&self, batch: RecordBatch) -> Result<RecordBatch> {
132        if self.json_columns.is_empty() {
133            return Ok(batch);
134        }
135        let mut cols = batch.columns().to_vec();
136        for (idx, expected_type) in &self.json_columns {
137            if batch.schema_ref().field(*idx).data_type() != expected_type {
138                cols[*idx] = JsonArray::from(batch.column(*idx))
139                    .widen_to(expected_type)
140                    .context(ConvertValueSnafu)?;
141            }
142        }
143        RecordBatch::try_new(self.schema.clone(), cols).context(NewRecordBatchSnafu)
144    }
145
146    /// Aligns [`RecordBatch`]s to [`Self::schema`].
147    pub(crate) fn align_batches<I>(&self, batches: I) -> Result<Vec<RecordBatch>>
148    where
149        I: IntoIterator<Item = RecordBatch>,
150    {
151        batches
152            .into_iter()
153            .map(|batch| self.align_batch(batch))
154            .collect()
155    }
156
157    /// Wraps an iterator so each yielded [`RecordBatch`] is lazily aligned.
158    pub(crate) fn wrap_iter(&self, iter: BoxedRecordBatchIterator) -> BoxedRecordBatchIterator {
159        let aligner = self.clone();
160        Box::new(iter.map(move |batch| aligner.align_batch(batch?)))
161    }
162}
163
164#[cfg(debug_assertions)]
165fn assert_columns_match_except_json2(base_schema: &Schema, schema: &Schema) {
166    debug_assert_eq!(
167        base_schema.fields().len(),
168        schema.fields().len(),
169        "input schemas for Json2Aligner must have the same column count"
170    );
171    for (idx, (base_field, field)) in base_schema.fields().iter().zip(schema.fields()).enumerate() {
172        let base_is_json2 = is_json2_extension_type(base_field);
173        let is_json2 = is_json2_extension_type(field);
174        debug_assert_eq!(
175            base_is_json2, is_json2,
176            "column {idx} must be JSON2 in all input schemas or none"
177        );
178        if !base_is_json2 && !is_json2 {
179            debug_assert_eq!(
180                base_field, field,
181                "non-JSON2 column {idx} must be identical across input schemas"
182            );
183        }
184    }
185}
186
187#[cfg(test)]
188mod tests {
189    use std::sync::Arc;
190
191    use datatypes::arrow::array::{
192        Array, ArrayRef, AsArray, Int64Array, StringViewArray, StructArray, UInt64Array,
193    };
194    use datatypes::arrow::datatypes::{DataType, Field, Fields, Schema};
195    use datatypes::extension::json::{Json2ExtensionType, JsonExtensionType};
196    use serde_json::json;
197
198    use super::*;
199
200    #[test]
201    fn test_try_new_rejects_empty_input() {
202        let err = match Json2Aligner::try_new([]) {
203            Ok(_) => panic!("expected empty input to fail"),
204            Err(err) => err,
205        };
206        assert!(
207            err.to_string()
208                .contains("Json2Aligner requires at least one input schema")
209        );
210    }
211
212    #[test]
213    fn test_try_new_keeps_non_json_schema_unchanged() {
214        let schema = Arc::new(Schema::new(vec![
215            Arc::new(Field::new("ts", DataType::Int64, false)),
216            Arc::new(Field::new("value", DataType::UInt64, true)),
217        ]));
218        let batch = RecordBatch::try_new(
219            schema.clone(),
220            vec![
221                Arc::new(Int64Array::from_iter_values([1, 2])) as ArrayRef,
222                Arc::new(UInt64Array::from(vec![Some(10), None])) as ArrayRef,
223            ],
224        )
225        .unwrap();
226
227        let aligner = Json2Aligner::try_new([schema.clone()]).unwrap();
228        assert!(Arc::ptr_eq(aligner.schema(), &schema));
229
230        let aligned = aligner.align_batch(batch).unwrap();
231        assert!(Arc::ptr_eq(aligned.schema_ref(), &schema));
232    }
233
234    #[test]
235    fn test_try_new_ignores_legacy_jsonb_extension_field() {
236        let legacy_jsonb_field = Arc::new(
237            Field::new("data", DataType::Binary, true).with_extension_type(JsonExtensionType),
238        );
239        let schema = Arc::new(Schema::new(vec![
240            Arc::new(Field::new("ts", DataType::Int64, false)),
241            legacy_jsonb_field,
242        ]));
243
244        let aligner = Json2Aligner::try_new([schema.clone()]).unwrap();
245
246        assert!(Arc::ptr_eq(aligner.schema(), &schema));
247        assert!(aligner.json_columns.is_empty());
248    }
249
250    #[test]
251    fn test_try_new_merges_json2_object_fields() {
252        let id_fields = Fields::from(vec![id_field()]);
253        let name_fields = Fields::from(vec![name_field()]);
254        let schema_with_id = schema_with_json_field(json_field("data", id_fields));
255        let schema_with_name = schema_with_json_field(json_field("data", name_fields));
256
257        let aligner = Json2Aligner::try_new([schema_with_id, schema_with_name]).unwrap();
258        let data_field = aligner.schema().field(1);
259        let DataType::Struct(fields) = data_field.data_type() else {
260            panic!("expected JSON2 field to be a struct");
261        };
262
263        assert_eq!(2, fields.len());
264        assert_eq!("id", fields[0].name());
265        assert_eq!(&DataType::Int64, fields[0].data_type());
266        assert_eq!("name", fields[1].name());
267        assert_eq!(&DataType::Utf8View, fields[1].data_type());
268        assert!(is_json2_extension_type(&aligner.schema().fields()[1]));
269    }
270
271    #[test]
272    fn test_align_batch_fills_missing_json2_fields() {
273        let id_fields = Fields::from(vec![id_field()]);
274        let name_fields = Fields::from(vec![name_field()]);
275        let schema_with_id = schema_with_json_field(json_field("data", id_fields.clone()));
276        let schema_with_name = schema_with_json_field(json_field("data", name_fields.clone()));
277
278        let batch_with_id = RecordBatch::try_new(
279            schema_with_id.clone(),
280            vec![
281                Arc::new(Int64Array::from_iter_values([1, 2])) as ArrayRef,
282                struct_array(
283                    id_fields,
284                    vec![Arc::new(Int64Array::from_iter_values([10, 20])) as ArrayRef],
285                ),
286            ],
287        )
288        .unwrap();
289        let batch_with_name = RecordBatch::try_new(
290            schema_with_name.clone(),
291            vec![
292                Arc::new(Int64Array::from_iter_values([3, 4])) as ArrayRef,
293                struct_array(
294                    name_fields,
295                    vec![
296                        Arc::new(StringViewArray::from(vec![Some("alice"), Some("bob")]))
297                            as ArrayRef,
298                    ],
299                ),
300            ],
301        )
302        .unwrap();
303
304        let aligner = Json2Aligner::try_new([schema_with_id, schema_with_name]).unwrap();
305        let aligned_with_id = aligner.align_batch(batch_with_id).unwrap();
306        let aligned_with_name = aligner.align_batch(batch_with_name).unwrap();
307
308        let data_with_id = aligned_with_id
309            .column(1)
310            .as_any()
311            .downcast_ref::<StructArray>()
312            .unwrap();
313        let id_values = data_with_id
314            .column(0)
315            .as_any()
316            .downcast_ref::<Int64Array>()
317            .unwrap();
318        let missing_names = data_with_id.column(1);
319        assert_eq!(10, id_values.value(0));
320        assert_eq!(20, id_values.value(1));
321        assert!(missing_names.is_null(0));
322        assert!(missing_names.is_null(1));
323
324        let data_with_name = aligned_with_name
325            .column(1)
326            .as_any()
327            .downcast_ref::<StructArray>()
328            .unwrap();
329        let missing_ids = data_with_name.column(0);
330        let name_values = data_with_name
331            .column(1)
332            .as_any()
333            .downcast_ref::<StringViewArray>()
334            .unwrap();
335        assert!(missing_ids.is_null(0));
336        assert!(missing_ids.is_null(1));
337        assert_eq!("alice", name_values.value(0));
338        assert_eq!("bob", name_values.value(1));
339    }
340
341    #[test]
342    fn test_align_conflicting_number_types_as_variant() {
343        let u64_fields = Fields::from(vec![Arc::new(Field::new("value", DataType::UInt64, true))]);
344        let i64_fields = Fields::from(vec![Arc::new(Field::new("value", DataType::Int64, true))]);
345        let u64_schema = schema_with_json_field(json_field("data", u64_fields.clone()));
346        let i64_schema = schema_with_json_field(json_field("data", i64_fields.clone()));
347        let u64_batch = RecordBatch::try_new(
348            u64_schema.clone(),
349            vec![
350                Arc::new(Int64Array::from_iter_values([1])) as ArrayRef,
351                struct_array(
352                    u64_fields,
353                    vec![Arc::new(UInt64Array::from_iter_values([u64::MAX])) as ArrayRef],
354                ),
355            ],
356        )
357        .unwrap();
358        let i64_batch = RecordBatch::try_new(
359            i64_schema.clone(),
360            vec![
361                Arc::new(Int64Array::from_iter_values([2])) as ArrayRef,
362                struct_array(
363                    i64_fields,
364                    vec![Arc::new(Int64Array::from_iter_values([i64::MIN])) as ArrayRef],
365                ),
366            ],
367        )
368        .unwrap();
369
370        let aligner = Json2Aligner::try_new([u64_schema, i64_schema]).unwrap();
371        let DataType::Struct(fields) = aligner.schema().field(1).data_type() else {
372            panic!("expected JSON2 field to be a struct");
373        };
374        assert_eq!(&DataType::Binary, fields[0].data_type());
375
376        for (batch, expected) in [(u64_batch, json!(u64::MAX)), (i64_batch, json!(i64::MIN))] {
377            let aligned = aligner.align_batch(batch).unwrap();
378            let data = aligned.column(1).as_struct();
379            assert_eq!(
380                expected,
381                JsonArray::from(data.column(0)).try_get_value(0).unwrap()
382            );
383        }
384    }
385
386    #[test]
387    fn test_wrap_iter_aligns_each_batch() {
388        let id_fields = Fields::from(vec![id_field()]);
389        let name_fields = Fields::from(vec![name_field()]);
390        let schema_with_id = schema_with_json_field(json_field("data", id_fields.clone()));
391        let schema_with_name = schema_with_json_field(json_field("data", name_fields.clone()));
392
393        let batch_with_id = RecordBatch::try_new(
394            schema_with_id.clone(),
395            vec![
396                Arc::new(Int64Array::from_iter_values([1])) as ArrayRef,
397                struct_array(
398                    id_fields,
399                    vec![Arc::new(Int64Array::from_iter_values([10])) as ArrayRef],
400                ),
401            ],
402        )
403        .unwrap();
404        let batch_with_name = RecordBatch::try_new(
405            schema_with_name.clone(),
406            vec![
407                Arc::new(Int64Array::from_iter_values([2])) as ArrayRef,
408                struct_array(
409                    name_fields,
410                    vec![Arc::new(StringViewArray::from(vec![Some("alice")])) as ArrayRef],
411                ),
412            ],
413        )
414        .unwrap();
415
416        let aligner = Json2Aligner::try_new([schema_with_id, schema_with_name]).unwrap();
417        let iter: BoxedRecordBatchIterator =
418            Box::new(vec![Ok(batch_with_id), Ok(batch_with_name)].into_iter());
419        let aligned = aligner.wrap_iter(iter).collect::<Result<Vec<_>>>().unwrap();
420
421        assert_eq!(2, aligned.len());
422        assert!(Arc::ptr_eq(aligned[0].schema_ref(), aligner.schema()));
423        assert!(Arc::ptr_eq(aligned[1].schema_ref(), aligner.schema()));
424    }
425
426    fn json_field(name: &str, fields: Fields) -> Arc<Field> {
427        Arc::new(
428            Field::new(name, DataType::Struct(fields), true)
429                .with_extension_type(Json2ExtensionType::default()),
430        )
431    }
432
433    fn schema_with_json_field(json_field: Arc<Field>) -> SchemaRef {
434        Arc::new(Schema::new(vec![
435            Arc::new(Field::new("ts", DataType::Int64, false)),
436            json_field,
437        ]))
438    }
439
440    fn id_field() -> Arc<Field> {
441        Arc::new(Field::new("id", DataType::Int64, true))
442    }
443
444    fn name_field() -> Arc<Field> {
445        Arc::new(Field::new("name", DataType::Utf8View, true))
446    }
447
448    fn struct_array(fields: Fields, columns: Vec<ArrayRef>) -> ArrayRef {
449        Arc::new(StructArray::new(fields, columns, None))
450    }
451}