1use 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 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
147pub 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
172pub 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
498enum 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}