use std::marker::PhantomData;
use bytes::BytesMut;
use deser_core::Error;
use deser_core::de::StreamDeserializer;
use deser_core::de::{DeserializeOwned, OwnedDriver};
use deser_core::ser::{Serialize, StreamSerializer};
use deser_core::stream::{InputBuffer, Status};
#[cfg_attr(docsrs, doc(cfg(feature = "codec")))]
pub struct Codec<D: StreamDeserializer, S: StreamSerializer, T> {
buffer: InputBuffer<D>,
serializer: S,
pending: Option<OwnedDriver<'static, T>>,
_marker: PhantomData<fn() -> T>,
}
impl<D: StreamDeserializer, S: StreamSerializer, T> Codec<D, S, T> {
pub fn new(deserializer: D, serializer: S) -> Codec<D, S, T> {
Codec {
buffer: InputBuffer::new(deserializer),
serializer,
pending: None,
_marker: PhantomData,
}
}
pub fn deserializer(&self) -> &D {
self.buffer.deserializer()
}
pub fn serializer(&self) -> &S {
&self.serializer
}
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 decode_buffered(&mut self, src: &mut BytesMut) -> Result<Option<T>, Error>
where
T: DeserializeOwned,
{
if !src.is_empty() {
self.buffer.extend_from_slice(src);
src.clear();
}
if !self.buffer.supports_partial() {
return match self.buffer.poll()? {
Status::Ready => self.buffer.deserialize().map(Some),
Status::NeedInput | Status::End => Ok(None),
};
}
let mut driver = self.pending.take().unwrap_or_default();
match driver.with(|driver| self.buffer.drive_partial(driver))? {
Status::Ready => driver.finish().map(Some),
Status::End => Ok(None),
Status::NeedInput => {
self.pending = Some(driver);
Ok(None)
}
}
}
}
impl<D: StreamDeserializer, S: StreamSerializer, T: DeserializeOwned> tokio_util::codec::Decoder
for Codec<D, S, T>
{
type Item = T;
type Error = Error;
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<T>, Error> {
self.decode_buffered(src)
}
fn decode_eof(&mut self, src: &mut BytesMut) -> Result<Option<T>, Error> {
if !src.is_empty() {
self.buffer.extend_from_slice(src);
src.clear();
}
if !self.buffer.is_eof() {
self.buffer.set_eof();
}
self.decode_buffered(src)
}
}
impl<D: StreamDeserializer, S: StreamSerializer, T, V: Serialize> tokio_util::codec::Encoder<V>
for Codec<D, S, T>
{
type Error = Error;
fn encode(&mut self, item: V, dst: &mut BytesMut) -> Result<(), Error> {
let mut driver = deser_core::ser::SerializeDriver::new(&item);
driver.set_default_context(self.buffer.context().clone());
self.serializer.drive(&mut driver)?;
dst.extend_from_slice(self.serializer.output());
self.serializer.clear_output();
Ok(())
}
}