Skip to main content

nautilus_serialization/arrow/
json.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
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// -------------------------------------------------------------------------------------------------
15
16use std::{collections::HashMap, sync::Arc};
17
18use arrow::{
19    array::{
20        Array, ArrayRef, BooleanArray, BooleanBuilder, Float64Array, Float64Builder, StringBuilder,
21        UInt64Array, UInt64Builder,
22    },
23    datatypes::{DataType, Field, Schema},
24    error::ArrowError,
25    record_batch::RecordBatch,
26};
27use serde::{Serialize, de::DeserializeOwned};
28use serde_json::{Map, Number, Value};
29
30use super::{EncodingError, StringColumnRef, extract_column, extract_column_string};
31
32#[derive(Clone, Copy, Debug, PartialEq, Eq)]
33pub enum JsonFieldEncoding {
34    Utf8,
35    Utf8Json,
36    /// Exact decimal written as `Utf8`, read back from `Utf8`, `Utf8View`, or `Float64`.
37    ///
38    /// The `Float64` case is what lets catalog files written before a field moved from `f64` to
39    /// `Decimal` keep decoding: `Decimal`'s `Deserialize` accepts both a JSON string and a JSON
40    /// number, so no version discriminator is needed.
41    DecimalStr,
42    UInt64,
43    Float64,
44    Boolean,
45}
46
47#[derive(Clone, Copy, Debug, PartialEq, Eq)]
48pub struct JsonFieldSpec {
49    pub name: &'static str,
50    pub encoding: JsonFieldEncoding,
51    pub nullable: bool,
52}
53
54impl JsonFieldSpec {
55    #[must_use]
56    pub const fn utf8(name: &'static str, nullable: bool) -> Self {
57        Self {
58            name,
59            encoding: JsonFieldEncoding::Utf8,
60            nullable,
61        }
62    }
63
64    #[must_use]
65    pub const fn utf8_json(name: &'static str, nullable: bool) -> Self {
66        Self {
67            name,
68            encoding: JsonFieldEncoding::Utf8Json,
69            nullable,
70        }
71    }
72
73    #[must_use]
74    pub const fn decimal_str(name: &'static str, nullable: bool) -> Self {
75        Self {
76            name,
77            encoding: JsonFieldEncoding::DecimalStr,
78            nullable,
79        }
80    }
81
82    #[must_use]
83    pub const fn u64(name: &'static str, nullable: bool) -> Self {
84        Self {
85            name,
86            encoding: JsonFieldEncoding::UInt64,
87            nullable,
88        }
89    }
90
91    #[must_use]
92    pub const fn f64(name: &'static str, nullable: bool) -> Self {
93        Self {
94            name,
95            encoding: JsonFieldEncoding::Float64,
96            nullable,
97        }
98    }
99
100    #[must_use]
101    pub const fn boolean(name: &'static str, nullable: bool) -> Self {
102        Self {
103            name,
104            encoding: JsonFieldEncoding::Boolean,
105            nullable,
106        }
107    }
108
109    fn field(self) -> Field {
110        let data_type = match self.encoding {
111            JsonFieldEncoding::Utf8
112            | JsonFieldEncoding::Utf8Json
113            | JsonFieldEncoding::DecimalStr => DataType::Utf8,
114            JsonFieldEncoding::UInt64 => DataType::UInt64,
115            JsonFieldEncoding::Float64 => DataType::Float64,
116            JsonFieldEncoding::Boolean => DataType::Boolean,
117        };
118
119        Field::new(self.name, data_type, self.nullable)
120    }
121}
122
123#[must_use]
124pub fn metadata_for_type(type_name: &'static str) -> HashMap<String, String> {
125    HashMap::from([("type".to_string(), type_name.to_string())])
126}
127
128#[must_use]
129pub fn schema_for_type(
130    type_name: &'static str,
131    metadata: Option<HashMap<String, String>>,
132    fields: &[JsonFieldSpec],
133) -> Schema {
134    let mut merged = metadata.unwrap_or_default();
135    merged.insert("type".to_string(), type_name.to_string());
136
137    Schema::new_with_metadata(
138        fields
139            .iter()
140            .copied()
141            .map(JsonFieldSpec::field)
142            .collect::<Vec<_>>(),
143        merged,
144    )
145}
146
147/// Encodes typed records into an Arrow record batch with the supplied schema metadata.
148///
149/// # Errors
150///
151/// Returns an error if JSON serialization fails or if a field cannot be encoded into
152/// the requested Arrow column type.
153pub fn encode_batch<T: Serialize>(
154    type_name: &'static str,
155    metadata: &HashMap<String, String>,
156    data: &[T],
157    fields: &[JsonFieldSpec],
158) -> Result<RecordBatch, ArrowError> {
159    let rows = serialize_rows(data)?;
160    let arrays: Result<Vec<ArrayRef>, ArrowError> = fields
161        .iter()
162        .copied()
163        .map(|field| encode_column(field, &rows))
164        .collect();
165
166    RecordBatch::try_new(
167        Arc::new(schema_for_type(type_name, Some(metadata.clone()), fields)),
168        arrays?,
169    )
170}
171
172/// Decodes typed records from an Arrow record batch produced by encode_batch.
173///
174/// # Errors
175///
176/// Returns an error if a required column is missing, has the wrong type, contains
177/// invalid JSON, or cannot be deserialized into the target type.
178pub fn decode_batch<T: DeserializeOwned>(
179    metadata: &HashMap<String, String>,
180    record_batch: &RecordBatch,
181    fields: &[JsonFieldSpec],
182    fallback_type_name: Option<&'static str>,
183) -> Result<Vec<T>, EncodingError> {
184    let columns: Result<Vec<_>, EncodingError> = fields
185        .iter()
186        .enumerate()
187        .map(|(index, field)| decode_column_ref(record_batch.columns(), *field, index))
188        .collect();
189    let columns = columns?;
190
191    let mut decoded = Vec::with_capacity(record_batch.num_rows());
192    let type_name = metadata
193        .get("type")
194        .cloned()
195        .or_else(|| fallback_type_name.map(str::to_string));
196
197    for row in 0..record_batch.num_rows() {
198        let mut value = Map::new();
199        if let Some(type_name) = &type_name {
200            value.insert("type".to_string(), Value::String(type_name.clone()));
201        }
202
203        for column in &columns {
204            value.insert(column.name().to_string(), column.to_json(row)?);
205        }
206
207        let json = serde_json::to_vec(&Value::Object(value))
208            .map_err(|e| EncodingError::ParseError("record_batch", format!("row {row}: {e}")))?;
209        decoded.push(
210            serde_json::from_slice(&json).map_err(|e| {
211                EncodingError::ParseError("record_batch", format!("row {row}: {e}"))
212            })?,
213        );
214    }
215
216    Ok(decoded)
217}
218
219fn serialize_rows<T: Serialize>(data: &[T]) -> Result<Vec<Map<String, Value>>, ArrowError> {
220    data.iter()
221        .map(|item| match serde_json::to_value(item) {
222            Ok(Value::Object(map)) => Ok(map),
223            Ok(_) => Err(invalid_argument(
224                "Expected serialized value to be a JSON object".to_string(),
225            )),
226            Err(e) => Err(invalid_argument(e.to_string())),
227        })
228        .collect()
229}
230
231fn encode_column(
232    field: JsonFieldSpec,
233    rows: &[Map<String, Value>],
234) -> Result<ArrayRef, ArrowError> {
235    match field.encoding {
236        JsonFieldEncoding::Utf8 | JsonFieldEncoding::DecimalStr => encode_utf8_column(field, rows),
237        JsonFieldEncoding::Utf8Json => encode_utf8_json_column(field, rows),
238        JsonFieldEncoding::UInt64 => encode_u64_column(field, rows),
239        JsonFieldEncoding::Float64 => encode_f64_column(field, rows),
240        JsonFieldEncoding::Boolean => encode_bool_column(field, rows),
241    }
242}
243
244fn encode_utf8_column(
245    field: JsonFieldSpec,
246    rows: &[Map<String, Value>],
247) -> Result<ArrayRef, ArrowError> {
248    let mut builder = StringBuilder::new();
249
250    for row in rows {
251        match require_value(field, row.get(field.name))? {
252            Some(value) => builder.append_value(value_to_string(value)?),
253            None => builder.append_null(),
254        }
255    }
256
257    Ok(Arc::new(builder.finish()))
258}
259
260fn encode_utf8_json_column(
261    field: JsonFieldSpec,
262    rows: &[Map<String, Value>],
263) -> Result<ArrayRef, ArrowError> {
264    let mut builder = StringBuilder::new();
265
266    for row in rows {
267        match require_value(field, row.get(field.name))? {
268            Some(value) => builder.append_value(
269                serde_json::to_string(value).map_err(|e| invalid_argument(e.to_string()))?,
270            ),
271            None => builder.append_null(),
272        }
273    }
274
275    Ok(Arc::new(builder.finish()))
276}
277
278fn encode_u64_column(
279    field: JsonFieldSpec,
280    rows: &[Map<String, Value>],
281) -> Result<ArrayRef, ArrowError> {
282    let mut builder = UInt64Builder::new();
283
284    for row in rows {
285        match require_value(field, row.get(field.name))? {
286            Some(value) => builder.append_value(parse_u64(value)?),
287            None => builder.append_null(),
288        }
289    }
290
291    Ok(Arc::new(builder.finish()))
292}
293
294fn encode_f64_column(
295    field: JsonFieldSpec,
296    rows: &[Map<String, Value>],
297) -> Result<ArrayRef, ArrowError> {
298    let mut builder = Float64Builder::new();
299
300    for row in rows {
301        match require_value(field, row.get(field.name))? {
302            Some(value) => builder.append_value(parse_f64(value)?),
303            None => builder.append_null(),
304        }
305    }
306
307    Ok(Arc::new(builder.finish()))
308}
309
310fn encode_bool_column(
311    field: JsonFieldSpec,
312    rows: &[Map<String, Value>],
313) -> Result<ArrayRef, ArrowError> {
314    let mut builder = BooleanBuilder::new();
315
316    for row in rows {
317        match require_value(field, row.get(field.name))? {
318            Some(value) => builder.append_value(parse_bool(value)?),
319            None => builder.append_null(),
320        }
321    }
322
323    Ok(Arc::new(builder.finish()))
324}
325
326fn require_value(
327    field: JsonFieldSpec,
328    value: Option<&Value>,
329) -> Result<Option<&Value>, ArrowError> {
330    match value {
331        Some(Value::Null) | None if !field.nullable => Err(invalid_argument(format!(
332            "Missing required field `{}`",
333            field.name
334        ))),
335        Some(Value::Null) | None => Ok(None),
336        Some(value) => Ok(Some(value)),
337    }
338}
339
340fn value_to_string(value: &Value) -> Result<String, ArrowError> {
341    match value {
342        Value::String(value) => Ok(value.clone()),
343        Value::Null => Err(invalid_argument("Unexpected null value".to_string())),
344        Value::Bool(_) | Value::Number(_) => Ok(value.to_string()),
345        Value::Array(_) | Value::Object(_) => {
346            serde_json::to_string(value).map_err(|e| invalid_argument(e.to_string()))
347        }
348    }
349}
350
351fn parse_u64(value: &Value) -> Result<u64, ArrowError> {
352    match value {
353        Value::Number(number) => number
354            .as_u64()
355            .ok_or_else(|| invalid_argument(format!("Expected u64, found `{number}`"))),
356        Value::String(value) => value
357            .parse::<u64>()
358            .map_err(|e| invalid_argument(format!("Failed to parse u64 from `{value}`: {e}"))),
359        _ => Err(invalid_argument(format!(
360            "Expected u64-compatible value, found `{value}`"
361        ))),
362    }
363}
364
365fn parse_f64(value: &Value) -> Result<f64, ArrowError> {
366    match value {
367        Value::Number(number) => number
368            .as_f64()
369            .ok_or_else(|| invalid_argument(format!("Expected f64, found `{number}`"))),
370        Value::String(value) => value
371            .parse::<f64>()
372            .map_err(|e| invalid_argument(format!("Failed to parse f64 from `{value}`: {e}"))),
373        _ => Err(invalid_argument(format!(
374            "Expected f64-compatible value, found `{value}`"
375        ))),
376    }
377}
378
379fn parse_bool(value: &Value) -> Result<bool, ArrowError> {
380    match value {
381        Value::Bool(value) => Ok(*value),
382        Value::String(value) => value
383            .parse::<bool>()
384            .map_err(|e| invalid_argument(format!("Failed to parse bool from `{value}`: {e}"))),
385        _ => Err(invalid_argument(format!(
386            "Expected bool-compatible value, found `{value}`"
387        ))),
388    }
389}
390
391enum ColumnRef<'a> {
392    Utf8 {
393        name: &'static str,
394        values: StringColumnRef<'a>,
395    },
396    Utf8Json {
397        name: &'static str,
398        values: StringColumnRef<'a>,
399    },
400    DecimalStr {
401        name: &'static str,
402        values: DecimalColumnRef<'a>,
403    },
404    UInt64 {
405        name: &'static str,
406        values: &'a UInt64Array,
407    },
408    Float64 {
409        name: &'static str,
410        values: &'a Float64Array,
411    },
412    Boolean {
413        name: &'static str,
414        values: &'a BooleanArray,
415    },
416}
417
418impl ColumnRef<'_> {
419    fn name(&self) -> &'static str {
420        match self {
421            Self::Utf8 { name, .. }
422            | Self::Utf8Json { name, .. }
423            | Self::DecimalStr { name, .. }
424            | Self::UInt64 { name, .. }
425            | Self::Float64 { name, .. }
426            | Self::Boolean { name, .. } => name,
427        }
428    }
429
430    fn to_json(&self, row: usize) -> Result<Value, EncodingError> {
431        match self {
432            Self::Utf8 { values, .. } => Ok(string_to_json(values, row)),
433            Self::Utf8Json { values, .. } => {
434                if values_is_null(values, row) {
435                    Ok(Value::Null)
436                } else {
437                    serde_json::from_str(values.value(row)).map_err(|e| {
438                        EncodingError::ParseError(self.name(), format!("row {row}: {e}"))
439                    })
440                }
441            }
442            Self::DecimalStr { values, .. } => match values {
443                DecimalColumnRef::Str(values) => Ok(string_to_json(values, row)),
444                DecimalColumnRef::Float64(values) => f64_to_json(self.name(), values, row),
445            },
446            Self::UInt64 { values, .. } => {
447                if values.is_null(row) {
448                    Ok(Value::Null)
449                } else {
450                    Ok(Value::Number(Number::from(values.value(row))))
451                }
452            }
453            Self::Float64 { values, .. } => f64_to_json(self.name(), values, row),
454            Self::Boolean { values, .. } => {
455                if values.is_null(row) {
456                    Ok(Value::Null)
457                } else {
458                    Ok(Value::Bool(values.value(row)))
459                }
460            }
461        }
462    }
463}
464
465fn decode_column_ref(
466    columns: &[ArrayRef],
467    field: JsonFieldSpec,
468    index: usize,
469) -> Result<ColumnRef<'_>, EncodingError> {
470    match field.encoding {
471        JsonFieldEncoding::Utf8 => Ok(ColumnRef::Utf8 {
472            name: field.name,
473            values: extract_column_string(columns, field.name, index)?,
474        }),
475        JsonFieldEncoding::Utf8Json => Ok(ColumnRef::Utf8Json {
476            name: field.name,
477            values: extract_column_string(columns, field.name, index)?,
478        }),
479        JsonFieldEncoding::DecimalStr => Ok(ColumnRef::DecimalStr {
480            name: field.name,
481            values: extract_column_decimal(columns, field.name, index)?,
482        }),
483        JsonFieldEncoding::UInt64 => Ok(ColumnRef::UInt64 {
484            name: field.name,
485            values: extract_column::<UInt64Array>(columns, field.name, index, DataType::UInt64)?,
486        }),
487        JsonFieldEncoding::Float64 => Ok(ColumnRef::Float64 {
488            name: field.name,
489            values: extract_column::<Float64Array>(columns, field.name, index, DataType::Float64)?,
490        }),
491        JsonFieldEncoding::Boolean => Ok(ColumnRef::Boolean {
492            name: field.name,
493            values: extract_column::<BooleanArray>(columns, field.name, index, DataType::Boolean)?,
494        }),
495    }
496}
497
498// Reference to a decimal column, either the current `Utf8`/`Utf8View` form or the `Float64`
499// form written before the field became exact.
500enum DecimalColumnRef<'a> {
501    Str(StringColumnRef<'a>),
502    Float64(&'a Float64Array),
503}
504
505fn extract_column_decimal<'a>(
506    columns: &'a [ArrayRef],
507    column_key: &'static str,
508    column_index: usize,
509) -> Result<DecimalColumnRef<'a>, EncodingError> {
510    extract_column_string(columns, column_key, column_index)
511        .map(DecimalColumnRef::Str)
512        .or_else(|e| {
513            extract_column::<Float64Array>(columns, column_key, column_index, DataType::Float64)
514                .map(DecimalColumnRef::Float64)
515                .map_err(|_| e)
516        })
517}
518
519fn string_to_json(values: &StringColumnRef<'_>, row: usize) -> Value {
520    if values_is_null(values, row) {
521        Value::Null
522    } else {
523        Value::String(values.value(row).to_string())
524    }
525}
526
527fn f64_to_json(
528    name: &'static str,
529    values: &Float64Array,
530    row: usize,
531) -> Result<Value, EncodingError> {
532    if values.is_null(row) {
533        return Ok(Value::Null);
534    }
535
536    Number::from_f64(values.value(row))
537        .map(Value::Number)
538        .ok_or_else(|| EncodingError::ParseError(name, format!("row {row}: invalid f64 value")))
539}
540
541fn values_is_null(values: &StringColumnRef<'_>, row: usize) -> bool {
542    match values {
543        StringColumnRef::Utf8(values) => values.is_null(row),
544        StringColumnRef::Utf8View(values) => values.is_null(row),
545    }
546}
547
548fn invalid_argument(message: String) -> ArrowError {
549    ArrowError::InvalidArgumentError(message)
550}