#[cfg(feature = "io")]
use std::io::Read;
use alloc::vec::Vec;
#[cfg(feature = "io")]
use deser_core::de::DeserializeOwned;
use deser_core::de::{self, DeserializeDriver, Frame, Progress};
use deser_core::{Error, ErrorKind, State};
use crate::de::{Deserializer, DeserializerConfig};
use crate::head::{Head, HeadError, decode_head};
use crate::parser::{Copying, Discard, Parser, Progress as ParseProgress};
#[derive(Default)]
struct StreamState {
pos: usize,
stack: Vec<u64>,
failed: bool,
parser: Parser,
started: bool,
skipping: Option<usize>,
feed_failed: bool,
ended: bool,
}
impl core::fmt::Debug for StreamState {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("StreamState").finish_non_exhaustive()
}
}
enum Scan {
Complete(usize),
Incomplete,
Malformed(usize),
}
impl StreamState {
fn scan(&mut self, input: &[u8]) -> Scan {
loop {
let head_start = self.pos;
let (head, len) = match decode_head(&input[head_start..]) {
Ok(rv) => rv,
Err(HeadError::Incomplete) => return Scan::Incomplete,
Err(HeadError::Reserved) => return Scan::Malformed(head_start),
};
let mut pos = head_start + len;
let items = match head {
Head::Str(len) | Head::Bin(len) | Head::Ext(_, len) => {
match pos.checked_add(len as usize) {
Some(end) if end <= input.len() => pos = end,
_ => return Scan::Incomplete,
}
0
}
Head::Array(len) => u64::from(len),
Head::Map(len) => u64::from(len) * 2,
_ => 0,
};
self.pos = pos;
if items > 0 {
self.stack.push(items);
continue;
}
loop {
match self.stack.last_mut() {
None => return Scan::Complete(pos),
Some(remaining) => {
*remaining -= 1;
if *remaining > 0 {
break;
}
self.stack.pop();
}
}
}
}
}
}
#[derive(Debug)]
pub struct StreamDeserializer {
config: DeserializerConfig,
state: StreamState,
}
impl Default for StreamDeserializer {
fn default() -> StreamDeserializer {
StreamDeserializer::new()
}
}
impl StreamDeserializer {
pub fn new() -> StreamDeserializer {
StreamDeserializer::with_config(&DeserializerConfig::new())
}
pub fn with_config(config: &DeserializerConfig) -> StreamDeserializer {
StreamDeserializer {
config: config.clone(),
state: StreamState::default(),
}
}
pub fn config(&self) -> &DeserializerConfig {
&self.config
}
fn skip_to_item(
&mut self,
input: &[u8],
offset: usize,
eof: bool,
) -> Result<Result<usize, Progress>, Error> {
let state = &mut self.state;
if state.ended {
return Ok(Err(Progress::End));
}
if state.feed_failed {
return Err(Error::new(
ErrorKind::Unexpected,
"cannot continue after an error",
));
}
let mut pos = 0;
if let Some(skip) = state.skipping {
let mut discard = Discard(State::new());
match state.parser.parse(input, skip, eof, offset, &mut discard) {
Ok(ParseProgress::Done(end)) => {
state.skipping = None;
pos = end;
}
Ok(ParseProgress::NeedMore(consumed)) => {
state.skipping = Some(0);
return Ok(Err(Progress::NeedMore { consumed }));
}
Err(err) => {
state.skipping = None;
return Err(fail(state, err, eof));
}
}
}
if !state.started && pos == input.len() {
return Ok(Err(if eof {
Progress::End
} else {
Progress::NeedMore { consumed: pos }
}));
}
Ok(Ok(pos))
}
}
impl de::StreamDeserializer for StreamDeserializer {
fn frame(&mut self, input: &[u8], eof: bool) -> Result<Frame, Error> {
let state = &mut self.state;
if state.failed {
return Err(Error::new(
ErrorKind::Unexpected,
"cannot continue after an item that is not well-formed",
));
}
if input.is_empty() && eof {
return Ok(Frame::End);
}
let end = match state.scan(input) {
Scan::Complete(end) => end,
Scan::Incomplete if !eof => return Ok(Frame::Incomplete { consumed: 0 }),
Scan::Incomplete => input.len(),
Scan::Malformed(offset) => {
state.failed = true;
offset + 1
}
};
state.pos = 0;
state.stack.clear();
Ok(Frame::Value {
start: 0,
end,
consumed: end,
})
}
fn drive_frame<'de>(
&mut self,
frame: &'de [u8],
driver: &mut DeserializeDriver<'_, 'de>,
) -> Result<(), Error> {
let mut de = Deserializer::from_slice_with_config(frame, &self.config);
de.drive(driver)?;
de.end()
}
fn supports_feed(&self) -> bool {
true
}
fn feed(
&mut self,
input: &[u8],
offset: usize,
eof: bool,
driver: &mut DeserializeDriver<'_, '_>,
) -> Result<Progress, Error> {
let pos = match self.skip_to_item(input, offset, eof)? {
Ok(pos) => pos,
Err(progress) => return Ok(progress),
};
let state = &mut self.state;
state.started = true;
match state
.parser
.parse(input, pos, eof, offset, &mut Copying(driver))
{
Ok(ParseProgress::Done(end)) => {
state.started = false;
Ok(Progress::Done { consumed: end })
}
Ok(ParseProgress::NeedMore(consumed)) => Ok(Progress::NeedMore { consumed }),
Err(err) => {
state.started = false;
match state.parser.recoverable() {
Some(resume) => {
state.skipping = Some(resume);
Err(err)
}
None => Err(fail(state, err, eof)),
}
}
}
}
fn peek(&mut self, input: &[u8], eof: bool) -> Result<Option<Progress>, Error> {
Ok(Some(match self.skip_to_item(input, 0, eof)? {
Ok(pos) => Progress::Done { consumed: pos },
Err(progress) => progress,
}))
}
}
#[cfg(feature = "io")]
impl DeserializerConfig {
pub fn reader<R: Read>(&self, reader: R) -> deser_core::io::Reader<R, StreamDeserializer> {
deser_core::io::Reader::new(reader, StreamDeserializer::with_config(self))
}
pub fn from_reader<T: DeserializeOwned, R: Read>(&self, reader: R) -> Result<T, Error> {
deser_core::io::from_reader(reader, StreamDeserializer::with_config(self))
}
}
#[cfg(feature = "io")]
pub fn from_reader<T: DeserializeOwned, R: Read>(reader: R) -> Result<T, Error> {
DeserializerConfig::new().from_reader(reader)
}
fn fail(state: &mut StreamState, err: Error, eof: bool) -> Error {
state.parser.reset();
state.feed_failed = true;
state.ended = eof && err.kind() == ErrorKind::EndOfFile;
err
}