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