diff --git a/Cargo.toml b/Cargo.toml index 9ea81a2..91582ae 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -15,7 +15,8 @@ categories = ["asynchronous", "network-programming"] [features] json = [ "serde", "serde_json" ] cbor = [ "serde", "serde_cbor" ] -default = [ "json", "cbor" ] +serde_bincode = [ "serde", "bincode" ] +default = [ "json", "cbor", "serde_bincode" ] [dependencies] @@ -36,3 +37,7 @@ optional = true [dependencies.serde_cbor] version = '0.10.2' optional = true + +[dependencies.bincode] +version = "1.2.1" +optional = true diff --git a/src/codec/length.rs b/src/codec/length.rs index e90cb74..2494559 100644 --- a/src/codec/length.rs +++ b/src/codec/length.rs @@ -67,11 +67,36 @@ impl Decoder for LengthCodec { return Ok(None); } - let len = src.get_u64() as usize; + let mut len_bytes = [0u8; U64_LENGTH]; + len_bytes.copy_from_slice(&src[..U64_LENGTH]); + let len = u64::from_be_bytes(len_bytes) as usize; + if src.len() - U64_LENGTH >= len { + // Skip the length header we already read. + src.advance(U64_LENGTH); Ok(Some(src.split_to(len).freeze())) } else { Ok(None) } } } + +#[cfg(test)] +mod tests { + use super::*; + + mod decode { + use super::*; + + #[test] + fn it_returns_bytes_withouth_length_header() { + let mut codec = LengthCodec{ }; + + let mut src = BytesMut::with_capacity(5); + src.put(&[0, 0, 0, 0, 0, 0, 0, 3u8, 1, 2, 3, 4][..]); + let item = codec.decode(&mut src).unwrap(); + + assert!(item == Some(Bytes::from(&[1u8, 2, 3][..]))); + } + } +} diff --git a/src/codec/mod.rs b/src/codec/mod.rs index f39f7b9..0eaca3b 100644 --- a/src/codec/mod.rs +++ b/src/codec/mod.rs @@ -12,3 +12,6 @@ pub use self::lines::LinesCodec; #[cfg(feature = "cbor")] mod cbor; #[cfg(feature = "cbor")] pub use self::cbor::{CborCodec, CborCodecError}; + +#[cfg(feature = "bincode")] mod serde; +#[cfg(feature = "bincode")] pub use self::serde::SerdeCodec; diff --git a/src/codec/serde.rs b/src/codec/serde.rs new file mode 100644 index 0000000..840df2f --- /dev/null +++ b/src/codec/serde.rs @@ -0,0 +1,60 @@ +use std::io; +use std::marker::PhantomData; +use bytes::{BytesMut, Bytes}; +use serde::Serialize; +use serde::de::DeserializeOwned; +use bincode; + +use crate::encoder::Encoder; +use crate::decoder::Decoder; +use crate::codec::LengthCodec; + +/// Encodes/decodes types implementing Serde Serialize/Deserialize traits. +/// It is built on top of `LengthCodec`. +pub struct SerdeCodec { + inner: LengthCodec, + phantom: PhantomData, +} + +impl Default for SerdeCodec { + fn default() -> Self { + Self { + inner: LengthCodec {}, + phantom: PhantomData, + } + } +} + +impl Encoder for SerdeCodec { + type Item = T; + type Error = io::Error; + + fn encode(&mut self, src: Self::Item, dst: &mut BytesMut) -> Result<(), Self::Error> { + let data = bincode::serialize(&src).map_err(to_io_err)?; + let bytes = Bytes::from(data); + self.inner.encode(bytes, dst) + } +} + +impl Decoder for SerdeCodec { + type Item = T; + type Error = io::Error; + + fn decode(&mut self, src: &mut BytesMut) -> Result, Self::Error> { + match self.inner.decode(src)? { + Some(bytes) => { + bincode::deserialize(&bytes) + .map_err(to_io_err) + .map(|item| Some(item)) + } + None => Ok(None), + } + } +} + +fn to_io_err(err: Box) -> io::Error { + match *err { + bincode::ErrorKind::Io(e) => e, + other => io::Error::new(io::ErrorKind::Other, other), + } +} diff --git a/src/lib.rs b/src/lib.rs index 9c2309f..8096099 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -22,7 +22,7 @@ //! ``` mod codec; -pub use codec::{BytesCodec, LengthCodec, LinesCodec}; +pub use codec::{BytesCodec, LengthCodec, LinesCodec, SerdeCodec}; #[cfg(feature = "json")] pub use codec::{JsonCodec, JsonCodecError}; #[cfg(feature = "cbor")] pub use codec::{CborCodec, CborCodecError}; diff --git a/tests/length_delimited.rs b/tests/length_delimited.rs new file mode 100644 index 0000000..ea8d188 --- /dev/null +++ b/tests/length_delimited.rs @@ -0,0 +1,29 @@ +use bytes::Bytes; +use futures::io::Cursor; +use futures::{executor, SinkExt, StreamExt}; +use futures_codec::{Framed, LengthCodec}; + +#[test] +fn same_msgs_are_received_as_were_sent() { + let cur = Cursor::new(vec![0; 256]); + let mut framed = Framed::new(cur, LengthCodec {}); + + let send_msgs = async { + framed.send(Bytes::from("msg1")).await.unwrap(); + framed.send(Bytes::from("msg2")).await.unwrap(); + framed.send(Bytes::from("msg3")).await.unwrap(); + }; + executor::block_on(send_msgs); + + let (mut cur, _) = framed.release(); + cur.set_position(0); + let framed = Framed::new(cur, LengthCodec {}); + + let recv_msgs = framed.take(3) + .map(|res| res.unwrap()) + .map(|buf| String::from_utf8(buf.to_vec()).unwrap()) + .collect::>(); + let msgs: Vec = executor::block_on(recv_msgs); + + assert!(msgs == vec!["msg1", "msg2", "msg3"]); +} diff --git a/tests/serde.rs b/tests/serde.rs new file mode 100644 index 0000000..36d8a42 --- /dev/null +++ b/tests/serde.rs @@ -0,0 +1,47 @@ +use futures::io::Cursor; +use futures::{executor, SinkExt, StreamExt}; +use futures_codec::{Framed, SerdeCodec}; +use serde::{Serialize, Deserialize}; + +#[derive(Serialize, Deserialize, Debug, PartialEq)] +struct Person { + name: String, + age: u8, +} + +impl Person { + fn new(name: &str, age: u8) -> Self { + Self { + name: name.into(), + age, + } + } +} + +#[test] +fn serializes_serde_enabled_structures() { + let cur = Cursor::new(vec![0; 4096]); + let mut framed = Framed::new(cur, SerdeCodec::default()); + + let send_msgs = async { + framed.send(Person::new("John", 11)).await.unwrap(); + framed.send(Person::new("Paul", 12)).await.unwrap(); + framed.send(Person::new("Mike", 13)).await.unwrap(); + }; + executor::block_on(send_msgs); + + let (mut cur, _) = framed.release(); + cur.set_position(0); + let framed = Framed::new(cur, SerdeCodec::default()); + + let recv_msgs = framed.take(3) + .map(|res| res.unwrap()) + .collect::>(); + let items: Vec = executor::block_on(recv_msgs); + + assert!(items == vec![ + Person::new("John", 11), + Person::new("Paul", 12), + Person::new("Mike", 13), + ]) +}