use std::error::Error as StdError;
use std::fmt;
use serde::de::DeserializeOwned;
use serde::Serialize;
#[derive(Debug)]
pub struct EncodeError(::bincode::Error);
impl fmt::Display for EncodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "bincode encode failed: {}", self.0)
}
}
impl StdError for EncodeError {
fn source(&self) -> Option<&(dyn StdError + 'static)> {
Some(&*self.0)
}
}
#[derive(Debug)]
pub struct DecodeError(::bincode::Error);
impl fmt::Display for DecodeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "bincode decode failed: {}", self.0)
}
}
impl StdError for DecodeError {
fn source(&self) -> Option<&(dyn StdError + 'static)> {
Some(&*self.0)
}
}
pub fn encode<T: Serialize>(value: &T, out: &mut Vec<u8>) -> Result<(), EncodeError> {
use ::bincode::{DefaultOptions, Serializer};
value
.serialize(&mut Serializer::new(out, DefaultOptions::new()))
.map_err(EncodeError)
}
pub fn decode_stream<T: DeserializeOwned>(
bytes: &[u8],
max_items: usize,
) -> Result<Vec<T>, DecodeError> {
use ::bincode::{DefaultOptions, Deserializer};
let mut deserializer = Deserializer::from_slice(bytes, DefaultOptions::new());
let mut out = Vec::new();
while out.len() < max_items {
match T::deserialize(&mut deserializer) {
Ok(value) => out.push(value),
Err(err) => {
if let ::bincode::ErrorKind::Io(io_err) = err.as_ref() {
if io_err.kind() == std::io::ErrorKind::UnexpectedEof {
break;
}
}
return Err(DecodeError(err));
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trips_a_stream_of_messages() {
let mut buf = Vec::new();
encode(&1u32, &mut buf).unwrap();
encode(&2u32, &mut buf).unwrap();
encode(&3u32, &mut buf).unwrap();
let decoded: Vec<u32> = decode_stream(&buf, 100).unwrap();
assert_eq!(decoded, vec![1, 2, 3]);
}
#[test]
fn max_items_caps_the_decoded_count() {
let mut buf = Vec::new();
for i in 0..10u32 {
encode(&i, &mut buf).unwrap();
}
let decoded: Vec<u32> = decode_stream(&buf, 4).unwrap();
assert_eq!(decoded, vec![0, 1, 2, 3]);
}
#[test]
fn empty_input_decodes_to_nothing() {
let decoded: Vec<u32> = decode_stream(&[], 100).unwrap();
assert!(decoded.is_empty());
}
#[test]
fn a_malformed_message_is_an_error() {
let decoded: Result<Vec<bool>, _> = decode_stream(&[2u8], 100);
let err = match decoded {
Ok(_) => panic!("a malformed message must be rejected, not silently accepted"),
Err(e) => e,
};
assert!(
err.to_string().contains("bincode decode failed"),
"unexpected Display: {err}"
);
assert!(
StdError::source(&err).is_some(),
"expected source() to chain to the underlying bincode error"
);
}
struct AlwaysFailsToSerialize;
impl serde::Serialize for AlwaysFailsToSerialize {
fn serialize<S: serde::Serializer>(&self, _serializer: S) -> Result<S::Ok, S::Error> {
Err(serde::ser::Error::custom("deliberately unserializable"))
}
}
#[test]
fn encode_failure_reports_display_and_source() {
let mut buf = Vec::new();
let err = match encode(&AlwaysFailsToSerialize, &mut buf) {
Ok(()) => panic!("AlwaysFailsToSerialize must fail to encode"),
Err(e) => e,
};
assert!(
err.to_string().contains("bincode encode failed"),
"unexpected Display: {err}"
);
assert!(
StdError::source(&err).is_some(),
"expected source() to chain to the underlying bincode error"
);
}
#[test]
fn a_truncated_trailing_message_ends_the_stream_leniently() {
let mut buf = Vec::new();
encode(&[1u8, 2, 3, 4], &mut buf).unwrap();
encode(&[5u8, 6, 7, 8], &mut buf).unwrap();
buf.pop();
let decoded: Vec<[u8; 4]> = decode_stream(&buf, 100)
.expect("a truncated trailing message ends the stream, not an error");
assert_eq!(decoded, vec![[1, 2, 3, 4]]);
}
}