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.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
159pub struct StreamingJsonArrayDecoder {
165 record_buf: Vec<u8>,
168 completed: Vec<Vec<(String, LoraValue)>>,
170 state: StreamState,
171 depth: u32,
175 in_string: bool,
177 string_escape: bool,
180 bytes_fed: u64,
181 rows_emitted: u64,
182 record_index: u64,
186 permissive: bool,
187 errors: Vec<RowParseError>,
188}
189
190#[derive(Debug, Clone, Copy, PartialEq, Eq)]
191enum StreamState {
192 Pre,
194 BetweenRecords,
196 InRecord,
198 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 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 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 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 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 let err = dec.feed(br#"[{"a":1},1,{"c":3}]"#).unwrap_err();
539 let parse = super::super::format::downcast_row_parse_error(&err);
540 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 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 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}