Skip to main content

deser_tokio/
codec.rs

1use std::marker::PhantomData;
2
3use bytes::BytesMut;
4use deser_core::Error;
5use deser_core::de::StreamDeserializer;
6use deser_core::de::{DeserializeOwned, OwnedDriver};
7use deser_core::ser::{Serialize, StreamSerializer};
8use deser_core::stream::{InputBuffer, Status};
9
10/// Implements the codec traits of [`tokio-util`](https://docs.rs/tokio-util).
11///
12/// The codec decodes values of type `T` with a [`StreamDeserializer`] and
13/// encodes values with a [`StreamSerializer`] (for instance
14/// `deser_json::StreamDeserializer` and `deser_json::Serializer`).  This
15/// makes it usable with `FramedRead`, `FramedWrite` and `Framed`:
16///
17/// ```
18/// # #[tokio::main(flavor = "current_thread")]
19/// # async fn main() {
20/// use futures_util::{SinkExt, StreamExt};
21/// use deser_json::{DeserializerConfig, Serializer, SerializerConfig, StreamDeserializer, Trailing};
22/// use deser_tokio::Codec;
23/// use tokio_util::codec::Framed;
24///
25/// const READ_LINES: DeserializerConfig =
26///     DeserializerConfig::builder().trailing(Trailing::Newline).build();
27/// const WRITE_LINES: SerializerConfig =
28///     SerializerConfig::builder().trailing(Trailing::Newline).build();
29///
30/// let (client, server) = tokio::io::duplex(1024);
31/// let codec = || {
32///     Codec::<_, _, Vec<u32>>::new(
33///         StreamDeserializer::with_config(READ_LINES),
34///         Serializer::with_config(WRITE_LINES),
35///     )
36/// };
37/// let mut client = Framed::new(client, codec());
38/// let mut server = Framed::new(server, codec());
39/// client.send(vec![1, 2]).await.unwrap();
40/// assert_eq!(server.next().await.unwrap().unwrap(), [1, 2]);
41/// # }
42/// ```
43///
44/// The data read by the framed reader is moved into the codec's buffer, so
45/// errors refer to positions in the stream.  If the format supports it
46/// (see [`StreamDeserializer::supports_partial`]), values are deserialized
47/// while their input arrives.
48#[cfg_attr(docsrs, doc(cfg(feature = "codec")))]
49pub struct Codec<D: StreamDeserializer, S: StreamSerializer, T> {
50    buffer: InputBuffer<D>,
51    serializer: S,
52    // the value which is deserialized while its input arrives
53    pending: Option<OwnedDriver<'static, T>>,
54    _marker: PhantomData<fn() -> T>,
55}
56
57impl<D: StreamDeserializer, S: StreamSerializer, T> Codec<D, S, T> {
58    /// Creates a codec.
59    pub fn new(deserializer: D, serializer: S) -> Codec<D, S, T> {
60        Codec {
61            buffer: InputBuffer::new(deserializer),
62            serializer,
63            pending: None,
64            _marker: PhantomData,
65        }
66    }
67
68    /// Returns the stream deserializer.
69    pub fn deserializer(&self) -> &D {
70        self.buffer.deserializer()
71    }
72
73    /// Returns the stream serializer.
74    pub fn serializer(&self) -> &S {
75        &self.serializer
76    }
77
78    /// Sets the context the values are deserialized and serialized in.
79    ///
80    /// This replaces the context of the stream deserializer (see
81    /// [`StreamDeserializer::context`](deser_core::de::StreamDeserializer::context)).
82    /// The values of the context are the defaults of the extension values
83    /// of the state (see [`Context`](deser_core::Context)).
84    pub fn set_context(&mut self, context: deser_core::Context) {
85        self.buffer.set_context(context);
86    }
87
88    /// Returns the context the values are deserialized and serialized in.
89    pub fn context(&self) -> &deser_core::Context {
90        self.buffer.context()
91    }
92
93    fn decode_buffered(&mut self, src: &mut BytesMut) -> Result<Option<T>, Error>
94    where
95        T: DeserializeOwned,
96    {
97        if !src.is_empty() {
98            self.buffer.extend_from_slice(src);
99            src.clear();
100        }
101        if !self.buffer.supports_partial() {
102            return match self.buffer.poll()? {
103                Status::Ready => self.buffer.deserialize().map(Some),
104                Status::NeedInput | Status::End => Ok(None),
105            };
106        }
107        let mut driver = self.pending.take().unwrap_or_default();
108        match driver.with(|driver| self.buffer.drive_partial(driver))? {
109            Status::Ready => driver.finish().map(Some),
110            Status::End => Ok(None),
111            Status::NeedInput => {
112                self.pending = Some(driver);
113                Ok(None)
114            }
115        }
116    }
117}
118
119impl<D: StreamDeserializer, S: StreamSerializer, T: DeserializeOwned> tokio_util::codec::Decoder
120    for Codec<D, S, T>
121{
122    type Item = T;
123    type Error = Error;
124
125    fn decode(&mut self, src: &mut BytesMut) -> Result<Option<T>, Error> {
126        self.decode_buffered(src)
127    }
128
129    fn decode_eof(&mut self, src: &mut BytesMut) -> Result<Option<T>, Error> {
130        if !src.is_empty() {
131            self.buffer.extend_from_slice(src);
132            src.clear();
133        }
134        if !self.buffer.is_eof() {
135            self.buffer.set_eof();
136        }
137        self.decode_buffered(src)
138    }
139}
140
141impl<D: StreamDeserializer, S: StreamSerializer, T, V: Serialize> tokio_util::codec::Encoder<V>
142    for Codec<D, S, T>
143{
144    type Error = Error;
145
146    fn encode(&mut self, item: V, dst: &mut BytesMut) -> Result<(), Error> {
147        // the value is written at once, the output is moved to the
148        // destination
149        let mut driver = deser_core::ser::SerializeDriver::new(&item);
150        driver.set_default_context(self.buffer.context().clone());
151        self.serializer.drive(&mut driver)?;
152        dst.extend_from_slice(self.serializer.output());
153        self.serializer.clear_output();
154        Ok(())
155    }
156}