1use 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
101pub 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 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
168pub struct StreamingJsonArrayDecoder {
174 record_buf: Vec<u8>,
177 completed: Vec<Vec<(String, LoraValue)>>,
179 state: StreamState,
180 depth: u32,
184 in_string: bool,
186 string_escape: bool,
189 bytes_fed: u64,
190 rows_emitted: u64,
191 record_index: u64,
195 permissive: bool,
196 errors: Vec<RowParseError>,
197}
198
199#[derive(Debug, Clone, Copy, PartialEq, Eq)]
200enum StreamState {
201 Pre,
203 BetweenRecords,
205 InRecord,
207 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 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 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 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 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 let err = dec.feed(br#"[{"a":1},1,{"c":3}]"#).unwrap_err();
548 let parse = super::super::format::downcast_row_parse_error(&err);
549 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 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 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}