#![doc(html_logo_url = "https://raw.githubusercontent.com/mitsuhiko/deser/main/artwork/logo.svg")]
#![cfg_attr(docsrs, feature(doc_cfg))]
use std::any::Any;
use std::future::poll_fn;
use std::marker::PhantomData;
use std::pin::Pin;
use std::task::{Context, Poll, ready};
use deser_core::de::{
Deserialize, DeserializeDriver, DeserializeOwned, OwnedDriver, StreamDeserializer,
};
use deser_core::ser::{Serialize, SerializeDriver, StreamSerializer};
use deser_core::stream::{
DEFAULT_BUFFER_LIMIT, ElementReader, ElementStatus, InputBuffer, Part, Status,
};
use deser_core::{Error, ErrorKind};
use futures_core::Stream;
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, ReadBuf};
#[cfg(feature = "codec")]
mod codec;
#[cfg(feature = "codec")]
pub use self::codec::Codec;
pub struct Reader<R, D: StreamDeserializer> {
reader: R,
buffer: InputBuffer<D>,
pending: Option<Box<dyn Any + Send>>,
}
impl<R, D: StreamDeserializer> Unpin for Reader<R, D> {}
impl<R: AsyncRead + Unpin, D: StreamDeserializer> Reader<R, D> {
pub fn new(reader: R, deserializer: D) -> Reader<R, D> {
Reader {
reader,
buffer: InputBuffer::new(deserializer),
pending: None,
}
}
pub fn set_context(&mut self, context: deser_core::Context) {
self.buffer.set_context(context);
}
pub fn context(&self) -> &deser_core::Context {
self.buffer.context()
}
fn poll_read_more(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
let mut buf = ReadBuf::new(self.buffer.read_buf());
ready!(Pin::new(&mut self.reader).poll_read(cx, &mut buf))?;
let read = buf.filled().len();
if read == 0 {
self.buffer.set_eof();
} else {
self.buffer.filled(read);
}
Poll::Ready(Ok(()))
}
fn poll_fill(&mut self, cx: &mut Context<'_>) -> Poll<Result<bool, Error>> {
loop {
match self.buffer.poll()? {
Status::Ready => return Poll::Ready(Ok(true)),
Status::End => return Poll::Ready(Ok(false)),
Status::NeedInput => ready!(self.poll_read_more(cx))?,
}
}
}
fn poll_read_setup<T, F>(
&mut self,
cx: &mut Context<'_>,
setup: &mut Option<F>,
) -> Poll<Result<Option<T>, Error>>
where
T: DeserializeOwned + 'static,
F: FnOnce(&mut DeserializeDriver<'_, '_>),
{
if !self.buffer.supports_partial() {
if !ready!(self.poll_fill(cx))? {
return Poll::Ready(Ok(None));
}
let setup = setup.take();
return Poll::Ready(
self.buffer
.deserialize_with(|driver| {
if let Some(setup) = setup {
setup(driver);
}
})
.map(Some),
);
}
let mut driver = match self.pending.take() {
Some(pending) => match pending.downcast::<OwnedDriver<'static, T>>() {
Ok(driver) => *driver,
Err(pending) => {
self.pending = Some(pending);
return Poll::Ready(Err(Error::new(
ErrorKind::InvalidState,
"a value of another type is being read",
)));
}
},
None => {
let mut driver = OwnedDriver::<'static, T>::new();
if let Some(setup) = setup.take() {
driver.with(|driver| setup(driver));
}
driver
}
};
loop {
match driver.with(|driver| self.buffer.drive_partial(driver))? {
Status::Ready => return Poll::Ready(driver.finish().map(Some)),
Status::End => return Poll::Ready(Ok(None)),
Status::NeedInput => match self.poll_read_more(cx) {
Poll::Ready(Ok(())) => {}
rv => {
self.pending = Some(Box::new(driver));
return rv.map(|rv| rv.map(|_| None));
}
},
}
}
}
pub fn poll_read<T: DeserializeOwned + 'static>(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<T>, Error>> {
self.poll_read_setup(cx, &mut None::<fn(&mut DeserializeDriver<'_, '_>)>)
}
pub async fn read<T: DeserializeOwned + 'static>(&mut self) -> Result<Option<T>, Error> {
poll_fn(|cx| self.poll_read(cx)).await
}
pub async fn read_with<T, F>(&mut self, setup: F) -> Result<Option<T>, Error>
where
T: DeserializeOwned + 'static,
F: FnOnce(&mut DeserializeDriver<'_, '_>),
{
let mut setup = Some(setup);
poll_fn(|cx| self.poll_read_setup(cx, &mut setup)).await
}
pub fn poll_read_next<T, E>(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<Option<Part<E, T>>, Error>>
where
T: DeserializeOwned + 'static,
E: Send + 'static,
{
let mut reader = match self.pending.take() {
Some(pending) => match pending.downcast::<ElementReader<T, E>>() {
Ok(reader) => reader,
Err(pending) => {
self.pending = Some(pending);
return Poll::Ready(Err(Error::new(
ErrorKind::InvalidState,
"a value of another type is being read",
)));
}
},
None => Box::new(ElementReader::<T, E>::new()),
};
loop {
match reader.poll(&mut self.buffer)? {
ElementStatus::Ready(next) => {
if reader.is_reading() {
self.pending = Some(reader);
}
return Poll::Ready(Ok(Some(next)));
}
ElementStatus::End => return Poll::Ready(Ok(None)),
ElementStatus::NeedInput => match self.poll_read_more(cx) {
Poll::Ready(Ok(())) => {}
rv => {
self.pending = Some(reader);
return rv.map(|rv| rv.map(|_| None));
}
},
}
}
}
pub async fn read_next<T, E>(&mut self) -> Result<Option<Part<E, T>>, Error>
where
T: DeserializeOwned + 'static,
E: Send + 'static,
{
poll_fn(|cx| self.poll_read_next(cx)).await
}
pub fn into_element_stream<T, E>(self) -> ElementStream<R, D, T, E>
where
T: DeserializeOwned + 'static,
E: Send + 'static,
{
ElementStream {
reader: self,
failed: false,
_marker: PhantomData,
}
}
pub async fn read_borrowed<'a, T: Deserialize<'a>>(&'a mut self) -> Result<Option<T>, Error> {
if !poll_fn(|cx| self.poll_fill(cx)).await? {
return Ok(None);
}
self.buffer.deserialize().map(Some)
}
pub async fn is_end(&mut self) -> Result<bool, Error> {
poll_fn(|cx| {
loop {
match self.buffer.peek()? {
Status::Ready => return Poll::Ready(Ok(false)),
Status::End => return Poll::Ready(Ok(true)),
Status::NeedInput => ready!(self.poll_read_more(cx))?,
}
}
})
.await
}
pub async fn end(&mut self) -> Result<(), Error> {
if poll_fn(|cx| self.poll_fill(cx)).await? {
return Err(self.buffer.trailing_error());
}
Ok(())
}
pub fn into_stream<T: DeserializeOwned + 'static>(self) -> ReaderStream<R, D, T> {
ReaderStream {
reader: self,
failed: false,
_marker: PhantomData,
}
}
pub fn deserializer(&self) -> &D {
self.buffer.deserializer()
}
pub fn get_ref(&self) -> &R {
&self.reader
}
pub fn get_mut(&mut self) -> &mut R {
&mut self.reader
}
pub fn into_inner(self) -> R {
self.reader
}
}
pub struct ReaderStream<R, D: StreamDeserializer, T> {
reader: Reader<R, D>,
failed: bool,
_marker: PhantomData<fn() -> T>,
}
impl<R, D: StreamDeserializer, T> Unpin for ReaderStream<R, D, T> {}
impl<R, D: StreamDeserializer, T> ReaderStream<R, D, T> {
pub fn into_inner(self) -> Reader<R, D> {
self.reader
}
}
impl<R: AsyncRead + Unpin, D: StreamDeserializer, T: DeserializeOwned + 'static> Stream
for ReaderStream<R, D, T>
{
type Item = Result<T, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if self.failed {
return Poll::Ready(None);
}
match ready!(self.reader.poll_read(cx)) {
Ok(value) => Poll::Ready(value.map(Ok)),
Err(err) => {
self.failed = true;
Poll::Ready(Some(Err(err)))
}
}
}
}
pub struct ElementStream<R, D: StreamDeserializer, T, E> {
reader: Reader<R, D>,
failed: bool,
_marker: PhantomData<fn() -> (T, E)>,
}
impl<R, D: StreamDeserializer, T, E> Unpin for ElementStream<R, D, T, E> {}
impl<R, D: StreamDeserializer, T, E> ElementStream<R, D, T, E> {
pub fn into_inner(self) -> Reader<R, D> {
self.reader
}
}
impl<R, D, T, E> Stream for ElementStream<R, D, T, E>
where
R: AsyncRead + Unpin,
D: StreamDeserializer,
T: DeserializeOwned + 'static,
E: Send + 'static,
{
type Item = Result<Part<E, T>, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if self.failed {
return Poll::Ready(None);
}
match ready!(self.reader.poll_read_next(cx)) {
Ok(next) => Poll::Ready(next.map(Ok)),
Err(err) => {
self.failed = true;
Poll::Ready(Some(Err(err)))
}
}
}
}
pub struct Writer<W, S: StreamSerializer> {
writer: W,
serializer: S,
limit: usize,
writing: bool,
context: deser_core::Context,
}
impl<W: AsyncWrite + Unpin, S: StreamSerializer> Writer<W, S> {
pub fn new(writer: W, serializer: S) -> Writer<W, S> {
Writer {
writer,
serializer,
limit: DEFAULT_BUFFER_LIMIT,
writing: false,
context: deser_core::Context::new(),
}
}
pub fn set_context(&mut self, context: deser_core::Context) {
self.context = context;
}
pub fn context(&self) -> &deser_core::Context {
&self.context
}
pub fn set_buffer_limit(&mut self, limit: usize) {
self.limit = limit;
}
pub fn buffer_limit(&self) -> usize {
self.limit
}
pub async fn write<T: Serialize + ?Sized>(&mut self, value: &T) -> Result<(), Error> {
self.write_with(value, |_| {}).await
}
pub async fn write_with<T, F>(&mut self, value: &T, setup: F) -> Result<(), Error>
where
T: Serialize + ?Sized,
F: FnOnce(&mut SerializeDriver<'_>),
{
if self.writing || self.serializer.in_progress() {
return Err(Error::new(
ErrorKind::InvalidState,
"a value was only partially written, the stream cannot continue",
));
}
let mut driver = SerializeDriver::new(&value);
setup(&mut driver);
if !self.context.is_empty() {
driver.set_default_context(self.context.clone());
}
self.write_output().await?;
let limit = match self.serializer.supports_partial() {
true => self.limit.max(1),
false => usize::MAX,
};
loop {
let done = self.serializer.drive_partial(&mut driver, limit)?;
self.write_output().await?;
if done {
return Ok(());
}
}
}
async fn write_output(&mut self) -> Result<(), Error> {
if self.serializer.output().is_empty() {
return Ok(());
}
self.writing = true;
self.writer.write_all(self.serializer.output()).await?;
self.serializer.clear_output();
self.writing = false;
Ok(())
}
pub async fn flush(&mut self) -> Result<(), Error> {
self.writer.flush().await?;
Ok(())
}
pub async fn shutdown(&mut self) -> Result<(), Error> {
self.writer.shutdown().await?;
Ok(())
}
pub fn serializer(&self) -> &S {
&self.serializer
}
pub fn get_ref(&self) -> &W {
&self.writer
}
pub fn get_mut(&mut self) -> &mut W {
&mut self.writer
}
pub fn into_inner(self) -> W {
self.writer
}
pub fn into_parts(self) -> (W, S) {
(self.writer, self.serializer)
}
}
pub async fn from_reader<T, R, D>(reader: R, deserializer: D) -> Result<T, Error>
where
T: DeserializeOwned + 'static,
R: AsyncRead + Unpin,
D: StreamDeserializer,
{
let mut reader = Reader::new(reader, deserializer);
let value = reader
.read()
.await?
.ok_or_else(|| Error::new(ErrorKind::EndOfFile, "empty input"))?;
reader.end().await?;
Ok(value)
}
pub async fn to_writer<W, S, T>(writer: W, serializer: S, value: &T) -> Result<(), Error>
where
W: AsyncWrite + Unpin,
S: StreamSerializer,
T: Serialize + ?Sized,
{
let mut writer = Writer::new(writer, serializer);
writer.write(value).await?;
writer.flush().await
}