1use 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#[derive(Clone)]
36pub(crate) struct Json2Aligner {
37 schema: SchemaRef,
39 json_columns: Vec<(usize, ArrowDataType)>,
41}
42
43impl Json2Aligner {
44 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 let base_schema = input_schemas.next().context(UnexpectedSnafu {
55 reason: "Json2Aligner requires at least one input schema",
56 })?;
57
58 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 if merged_types.is_empty() {
73 return Ok(Self {
74 schema: base_schema,
75 json_columns: Vec::new(),
76 });
77 }
78
79 for schema in input_schemas {
81 #[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 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 pub(crate) fn schema(&self) -> &SchemaRef {
127 &self.schema
128 }
129
130 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 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 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}