Skip to main content

lora_io/
json_array.rs

1//! JSON-array codec: `[ {...}, {...}, ... ]`.
2//!
3//! Both an eager [`JsonArrayDecoder`] (one-shot, loads the whole
4//! document) and a push-based [`StreamingJsonArrayDecoder`] are
5//! available. The streaming variant walks the byte stream with a
6//! depth-and-string-state tokenizer and emits a record each time
7//! the brace nesting returns to zero inside the top-level array.
8
9use std::io::{BufRead, Write};
10
11use lora_executor::{LoraValue, Row};
12use serde_json::Value as J;
13
14use super::format::{
15    invalid_data, row_parse_io_error, RowDecoder, RowEncoder, RowParseError, StreamingRowDecoder,
16};
17use super::value_json::{lora_value_from_json, lora_value_to_json};
18
19pub struct JsonArrayEncoder<W: Write> {
20    writer: W,
21    needs_comma: bool,
22    started: bool,
23    finished: bool,
24}
25
26impl<W: Write> JsonArrayEncoder<W> {
27    pub fn new(writer: W) -> Self {
28        Self {
29            writer,
30            needs_comma: false,
31            started: false,
32            finished: false,
33        }
34    }
35
36    pub fn into_inner(self) -> W {
37        self.writer
38    }
39
40    fn write_separator(&mut self) -> std::io::Result<()> {
41        if self.needs_comma {
42            self.writer.write_all(b",\n")?;
43        } else {
44            self.writer.write_all(b"\n")?;
45        }
46        Ok(())
47    }
48}
49
50impl<W: Write> RowEncoder for JsonArrayEncoder<W> {
51    fn begin(&mut self, _columns: &[String]) -> std::io::Result<()> {
52        if self.started {
53            return Ok(());
54        }
55        self.writer.write_all(b"[")?;
56        self.started = true;
57        Ok(())
58    }
59
60    fn write_row(&mut self, row: &Row) -> std::io::Result<()> {
61        if !self.started {
62            self.begin(&[])?;
63        }
64        self.write_separator()?;
65        let mut obj = serde_json::Map::with_capacity(row.len());
66        for (_, name, value) in row.iter_named() {
67            obj.insert(name.into_owned(), lora_value_to_json(value));
68        }
69        serde_json::to_writer(&mut self.writer, &J::Object(obj))?;
70        self.needs_comma = true;
71        Ok(())
72    }
73
74    fn write_named_row(&mut self, columns: &[(String, LoraValue)]) -> std::io::Result<()> {
75        if !self.started {
76            self.begin(&[])?;
77        }
78        self.write_separator()?;
79        let mut obj = serde_json::Map::with_capacity(columns.len());
80        for (name, value) in columns {
81            obj.insert(name.clone(), lora_value_to_json(value));
82        }
83        serde_json::to_writer(&mut self.writer, &J::Object(obj))?;
84        self.needs_comma = true;
85        Ok(())
86    }
87
88    fn finish(&mut self) -> std::io::Result<()> {
89        if self.finished {
90            return Ok(());
91        }
92        if !self.started {
93            self.writer.write_all(b"[")?;
94        }
95        self.writer.write_all(b"\n]\n")?;
96        self.finished = true;
97        self.writer.flush()
98    }
99}
100
101/// Decoder that parses the entire array up-front, then yields one
102/// element per `next_row` call. Not streaming — use [`super::JsonlDecoder`]
103/// for that.
104pub struct JsonArrayDecoder<R: BufRead> {
105    state: State<R>,
106}
107
108enum State<R: BufRead> {
109    Pending(Option<R>),
110    Loaded(std::vec::IntoIter<J>),
111}
112
113impl<R: BufRead> JsonArrayDecoder<R> {
114    pub fn new(reader: R) -> Self {
115        Self {
116            state: State::Pending(Some(reader)),
117        }
118    }
119
120    fn ensure_loaded(&mut self) -> std::io::Result<()> {
121        if let State::Pending(slot) = &mut self.state {
122            let reader = slot
123                .take()
124                .ok_or_else(|| invalid_data("JSON array reader was already consumed"))?;
125            let value: J = serde_json::from_reader(reader).map_err(invalid_data)?;
126            let J::Array(items) = value else {
127                return Err(invalid_data("expected a JSON array at the top level"));
128            };
129            self.state = State::Loaded(items.into_iter());
130        }
131        Ok(())
132    }
133}
134
135impl<R: BufRead> RowDecoder for JsonArrayDecoder<R> {
136    fn header(&mut self) -> std::io::Result<Option<Vec<String>>> {
137        Ok(None)
138    }
139
140    fn next_row(&mut self) -> std::io::Result<Option<Vec<(String, LoraValue)>>> {
141        self.ensure_loaded()?;
142        // Belt-and-braces: `ensure_loaded` is what advances `Pending`
143        // into `Loaded`, but treating the invariant as an error rather
144        // than `unreachable!()` means a future code change that doesn't
145        // perform the transition surfaces as an io error on untrusted
146        // input instead of a panic.
147        let State::Loaded(iter) = &mut self.state else {
148            return Err(invalid_data(
149                "JSON array decoder state was not loaded after ensure_loaded",
150            ));
151        };
152        let Some(v) = iter.next() else {
153            return Ok(None);
154        };
155        let J::Object(obj) = v else {
156            return Err(invalid_data(
157                "expected JSON object per array element".to_string(),
158            ));
159        };
160        let mut out = Vec::with_capacity(obj.len());
161        for (k, raw) in obj {
162            out.push((k, lora_value_from_json(raw).map_err(invalid_data)?));
163        }
164        Ok(Some(out))
165    }
166}
167
168/// Push-based JSON-array decoder. Tracks brace/bracket depth and
169/// JSON string state byte-by-byte so the input can be fed in
170/// arbitrarily sized chunks. Memory bound: at most one in-progress
171/// record (between its opening `{` and matching `}`) plus a single
172/// chunk of incoming bytes.
173pub struct StreamingJsonArrayDecoder {
174    /// Bytes of the in-progress record, including its opening `{`.
175    /// Cleared every time a record finishes parsing.
176    record_buf: Vec<u8>,
177    /// Completed records waiting to be drained.
178    completed: Vec<Vec<(String, LoraValue)>>,
179    state: StreamState,
180    /// Brace/bracket nesting depth within the current record.
181    /// Records start at depth 1 (after consuming their opening `{`)
182    /// and finish when depth returns to 0.
183    depth: u32,
184    /// True while inside a JSON string literal.
185    in_string: bool,
186    /// True when the previous byte inside a string was `\`. Persisted
187    /// across `feed` calls so cross-chunk escapes are handled.
188    string_escape: bool,
189    bytes_fed: u64,
190    rows_emitted: u64,
191    /// 1-indexed counter of records seen so far. Advances when a `{`
192    /// is observed at depth 0; persists across parse failures so
193    /// permissive-mode errors carry the correct row number.
194    record_index: u64,
195    permissive: bool,
196    errors: Vec<RowParseError>,
197}
198
199#[derive(Debug, Clone, Copy, PartialEq, Eq)]
200enum StreamState {
201    /// Skipping whitespace until the opening `[` is seen.
202    Pre,
203    /// Inside the top-level array, between records.
204    BetweenRecords,
205    /// Buffering bytes for an in-progress record.
206    InRecord,
207    /// Saw the closing `]`. Trailing whitespace is OK; anything else
208    /// is an error.
209    Post,
210}
211
212impl Default for StreamingJsonArrayDecoder {
213    fn default() -> Self {
214        Self::new()
215    }
216}
217
218impl StreamingJsonArrayDecoder {
219    pub fn new() -> Self {
220        Self {
221            record_buf: Vec::with_capacity(4 * 1024),
222            completed: Vec::new(),
223            state: StreamState::Pre,
224            depth: 0,
225            in_string: false,
226            string_escape: false,
227            bytes_fed: 0,
228            rows_emitted: 0,
229            record_index: 0,
230            permissive: false,
231            errors: Vec::new(),
232        }
233    }
234
235    fn parse_record(&mut self) -> std::io::Result<()> {
236        // UTF-8 errors are fatal — the stream is desynced at the byte
237        // level and there's no clean place to resume.
238        let s = std::str::from_utf8(&self.record_buf).map_err(invalid_data)?;
239        match parse_json_object(s) {
240            Ok(record) => {
241                self.record_buf.clear();
242                self.completed.push(record);
243                self.rows_emitted += 1;
244                Ok(())
245            }
246            Err(message) => self.report_error(message),
247        }
248    }
249
250    fn report_error(&mut self, message: String) -> std::io::Result<()> {
251        let err = RowParseError {
252            row: self.record_index,
253            column: None,
254            raw_sample: RowParseError::make_sample_from_bytes(&self.record_buf),
255            message,
256        };
257        self.record_buf.clear();
258        if self.permissive {
259            self.errors.push(err);
260            Ok(())
261        } else {
262            Err(row_parse_io_error(err))
263        }
264    }
265}
266
267fn parse_json_object(s: &str) -> Result<Vec<(String, LoraValue)>, String> {
268    let v: J = serde_json::from_str(s).map_err(|e| e.to_string())?;
269    let J::Object(obj) = v else {
270        return Err("expected JSON object per array element".to_string());
271    };
272    let mut record = Vec::with_capacity(obj.len());
273    for (k, raw) in obj {
274        let value = lora_value_from_json(raw).map_err(|e| format!("key `{k}`: {e}"))?;
275        record.push((k, value));
276    }
277    Ok(record)
278}
279
280impl StreamingRowDecoder for StreamingJsonArrayDecoder {
281    fn feed(&mut self, chunk: &[u8]) -> std::io::Result<()> {
282        if chunk.is_empty() {
283            return Ok(());
284        }
285        self.bytes_fed += chunk.len() as u64;
286        for &b in chunk {
287            match self.state {
288                StreamState::Pre => {
289                    if b.is_ascii_whitespace() {
290                        continue;
291                    }
292                    if b == b'[' {
293                        self.state = StreamState::BetweenRecords;
294                    } else {
295                        return Err(invalid_data(format!(
296                            "expected `[` at the top level, found byte 0x{b:02x}"
297                        )));
298                    }
299                }
300                StreamState::BetweenRecords => {
301                    if b.is_ascii_whitespace() || b == b',' {
302                        continue;
303                    }
304                    if b == b']' {
305                        self.state = StreamState::Post;
306                        continue;
307                    }
308                    if b == b'{' {
309                        self.record_buf.clear();
310                        self.record_buf.push(b);
311                        self.depth = 1;
312                        self.in_string = false;
313                        self.string_escape = false;
314                        self.record_index += 1;
315                        self.state = StreamState::InRecord;
316                    } else {
317                        return Err(invalid_data(format!(
318                            "expected JSON object inside array, found byte 0x{b:02x}"
319                        )));
320                    }
321                }
322                StreamState::InRecord => {
323                    self.record_buf.push(b);
324                    if self.in_string {
325                        if self.string_escape {
326                            self.string_escape = false;
327                        } else if b == b'\\' {
328                            self.string_escape = true;
329                        } else if b == b'"' {
330                            self.in_string = false;
331                        }
332                        continue;
333                    }
334                    match b {
335                        b'"' => self.in_string = true,
336                        b'{' | b'[' => self.depth += 1,
337                        b'}' | b']' => {
338                            self.depth -= 1;
339                            if self.depth == 0 {
340                                self.parse_record()?;
341                                self.state = StreamState::BetweenRecords;
342                            }
343                        }
344                        _ => {}
345                    }
346                }
347                StreamState::Post => {
348                    if !b.is_ascii_whitespace() {
349                        return Err(invalid_data(format!(
350                            "unexpected byte 0x{b:02x} after closing `]`"
351                        )));
352                    }
353                }
354            }
355        }
356        Ok(())
357    }
358
359    fn drain(&mut self) -> std::io::Result<Vec<Vec<(String, LoraValue)>>> {
360        Ok(std::mem::take(&mut self.completed))
361    }
362
363    fn finish(&mut self) -> std::io::Result<Vec<Vec<(String, LoraValue)>>> {
364        match self.state {
365            StreamState::Pre => {
366                return Err(invalid_data(
367                    "unexpected end of input: never saw opening `[`",
368                ));
369            }
370            StreamState::BetweenRecords => {
371                return Err(invalid_data(
372                    "unexpected end of input: array was never closed",
373                ));
374            }
375            StreamState::InRecord => {
376                return Err(invalid_data("unexpected end of input mid-record"));
377            }
378            StreamState::Post => {}
379        }
380        Ok(std::mem::take(&mut self.completed))
381    }
382
383    fn header(&self) -> Option<&[String]> {
384        None
385    }
386
387    fn bytes_fed(&self) -> u64 {
388        self.bytes_fed
389    }
390
391    fn rows_emitted(&self) -> u64 {
392        self.rows_emitted
393    }
394
395    fn set_permissive(&mut self, on: bool) {
396        self.permissive = on;
397    }
398
399    fn take_errors(&mut self) -> Vec<RowParseError> {
400        std::mem::take(&mut self.errors)
401    }
402}
403
404#[cfg(test)]
405mod tests {
406    use super::*;
407    use std::io::Cursor;
408
409    #[test]
410    fn encode_then_decode() {
411        let mut buf = Vec::new();
412        {
413            let mut enc = JsonArrayEncoder::new(&mut buf);
414            enc.begin(&[]).unwrap();
415            enc.write_named_row(&[("name".into(), LoraValue::String("alice".into()))])
416                .unwrap();
417            enc.write_named_row(&[("name".into(), LoraValue::String("bob".into()))])
418                .unwrap();
419            enc.finish().unwrap();
420        }
421        let text = std::str::from_utf8(&buf).unwrap();
422        assert!(text.trim_start().starts_with('['));
423        assert!(text.trim_end().ends_with(']'));
424
425        let mut dec = JsonArrayDecoder::new(Cursor::new(buf));
426        let r1 = dec.next_row().unwrap().unwrap();
427        assert_eq!(r1[0], ("name".into(), LoraValue::String("alice".into())));
428        let r2 = dec.next_row().unwrap().unwrap();
429        assert_eq!(r2[0], ("name".into(), LoraValue::String("bob".into())));
430        assert!(dec.next_row().unwrap().is_none());
431    }
432
433    #[test]
434    fn empty_array() {
435        let mut dec = JsonArrayDecoder::new(Cursor::new("[]"));
436        assert!(dec.next_row().unwrap().is_none());
437    }
438
439    #[test]
440    fn rejects_non_object_elements() {
441        let mut dec = JsonArrayDecoder::new(Cursor::new("[1, 2]"));
442        let err = dec.next_row().unwrap_err();
443        assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
444    }
445
446    #[test]
447    fn streaming_basic_round_trip() {
448        let mut dec = StreamingJsonArrayDecoder::new();
449        dec.feed(br#"[{"a":1},{"b":2}]"#).unwrap();
450        let rows = dec.drain().unwrap();
451        assert_eq!(rows.len(), 2);
452        assert_eq!(rows[0][0], ("a".into(), LoraValue::Int(1)));
453        assert_eq!(rows[1][0], ("b".into(), LoraValue::Int(2)));
454        assert!(dec.finish().unwrap().is_empty());
455        assert_eq!(dec.rows_emitted(), 2);
456    }
457
458    #[test]
459    fn streaming_empty_array() {
460        let mut dec = StreamingJsonArrayDecoder::new();
461        dec.feed(b"[]").unwrap();
462        assert!(dec.drain().unwrap().is_empty());
463        assert!(dec.finish().unwrap().is_empty());
464        assert_eq!(dec.rows_emitted(), 0);
465    }
466
467    #[test]
468    fn streaming_split_across_chunks_inside_string() {
469        // The split lands inside a string literal — the string-state
470        // flag must persist so `}` inside the string isn't treated as
471        // a record terminator.
472        let mut dec = StreamingJsonArrayDecoder::new();
473        dec.feed(br#"[{"name":"al"#).unwrap();
474        assert!(dec.drain().unwrap().is_empty());
475        dec.feed(br#"ice}"},{"name":"bob"}]"#).unwrap();
476        let rows = dec.drain().unwrap();
477        assert_eq!(rows.len(), 2);
478        assert_eq!(
479            rows[0][0],
480            ("name".into(), LoraValue::String("alice}".into()))
481        );
482        assert_eq!(rows[1][0], ("name".into(), LoraValue::String("bob".into())));
483    }
484
485    #[test]
486    fn streaming_split_inside_escape() {
487        // The split lands between `\` and `"`. The escape-state flag
488        // has to persist so the second chunk's `"` is treated as a
489        // literal character, not the string-closing quote.
490        let mut dec = StreamingJsonArrayDecoder::new();
491        dec.feed(br#"[{"a":"x\"#).unwrap();
492        dec.feed(br#"""}]"#).unwrap();
493        let rows = dec.drain().unwrap();
494        assert_eq!(rows.len(), 1);
495        assert_eq!(rows[0][0], ("a".into(), LoraValue::String("x\"".into())));
496    }
497
498    #[test]
499    fn streaming_nested_objects_and_arrays() {
500        let mut dec = StreamingJsonArrayDecoder::new();
501        dec.feed(br#"[{"arr":[1,2,3],"obj":{"k":"v"}}]"#).unwrap();
502        let rows = dec.drain().unwrap();
503        assert_eq!(rows.len(), 1);
504        let by_key: std::collections::BTreeMap<_, _> = rows[0].iter().cloned().collect();
505        assert!(matches!(by_key.get("arr"), Some(LoraValue::List(_))));
506        assert!(matches!(by_key.get("obj"), Some(LoraValue::Map(_))));
507    }
508
509    #[test]
510    fn streaming_whitespace_between_records() {
511        let mut dec = StreamingJsonArrayDecoder::new();
512        dec.feed(b"[\n  {\"a\":1},\n  {\"b\":2}\n]\n").unwrap();
513        let rows = dec.drain().unwrap();
514        assert_eq!(rows.len(), 2);
515        assert!(dec.finish().unwrap().is_empty());
516    }
517
518    #[test]
519    fn streaming_rejects_truncated_input() {
520        let mut dec = StreamingJsonArrayDecoder::new();
521        dec.feed(br#"[{"a":1}"#).unwrap();
522        // Record parsed; array never closed.
523        assert_eq!(dec.drain().unwrap().len(), 1);
524        let err = dec.finish().unwrap_err();
525        assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
526    }
527
528    #[test]
529    fn streaming_rejects_missing_open_bracket() {
530        let mut dec = StreamingJsonArrayDecoder::new();
531        let err = dec.feed(b"{\"a\":1}").unwrap_err();
532        assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
533    }
534
535    #[test]
536    fn streaming_rejects_non_object_element() {
537        let mut dec = StreamingJsonArrayDecoder::new();
538        let err = dec.feed(b"[1,2]").unwrap_err();
539        assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
540    }
541
542    #[test]
543    fn streaming_strict_mode_attributes_row() {
544        let mut dec = StreamingJsonArrayDecoder::new();
545        // Second record is a non-object — the structured error must
546        // identify it as record 2.
547        let err = dec.feed(br#"[{"a":1},1,{"c":3}]"#).unwrap_err();
548        let parse = super::super::format::downcast_row_parse_error(&err);
549        // Non-object elements are caught at the byte level by the
550        // state machine; the structured-error path covers object-shape
551        // failures. Either error path is acceptable as long as the
552        // input is rejected.
553        assert!(parse.is_none() || parse.unwrap().row >= 1);
554    }
555
556    #[test]
557    fn streaming_permissive_skips_bad_records() {
558        let mut dec = StreamingJsonArrayDecoder::new();
559        dec.set_permissive(true);
560        // Second record has a key whose value is unrepresentable —
561        // serde_json rejects e.g. a duplicate or malformed escape.
562        // Here we use a malformed inner value via `lora_value_from_json`
563        // by constructing an object with an unsupported tagged shape.
564        dec.feed(br#"[{"a":1},{"bad":{"kind":"date","iso":"not-a-date"}},{"c":3}]"#)
565            .unwrap();
566        let rows = dec.drain().unwrap();
567        assert_eq!(rows.len(), 2);
568        assert_eq!(rows[0][0], ("a".into(), LoraValue::Int(1)));
569        assert_eq!(rows[1][0], ("c".into(), LoraValue::Int(3)));
570        let errors = dec.take_errors();
571        assert_eq!(errors.len(), 1);
572        assert_eq!(errors[0].row, 2);
573        assert!(errors[0].column.is_none());
574    }
575
576    #[test]
577    fn streaming_one_byte_at_a_time() {
578        // Worst-case chunking: every byte arrives in its own feed call.
579        let input = br#"[{"a":1,"b":"hi"},{"c":[1,2]}]"#;
580        let mut dec = StreamingJsonArrayDecoder::new();
581        for &b in input {
582            dec.feed(&[b]).unwrap();
583        }
584        let rows = dec.drain().unwrap();
585        assert_eq!(rows.len(), 2);
586        assert!(dec.finish().unwrap().is_empty());
587    }
588}