use serde::de::{self, SeqAccess, Visitor};
use serde::{Deserializer, Serialize, Serializer};
use std::fmt;
const MAX_SEQ_PREALLOC: usize = 4096;
struct ByteBuf<'a>(&'a [u8]);
impl Serialize for ByteBuf<'_> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_bytes(self.0)
}
}
struct ByteVecVisitor;
impl<'de> Visitor<'de> for ByteVecVisitor {
type Value = Vec<u8>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a byte string or an array of bytes")
}
fn visit_bytes<E: de::Error>(self, v: &[u8]) -> Result<Self::Value, E> {
Ok(v.to_vec())
}
fn visit_byte_buf<E: de::Error>(self, v: Vec<u8>) -> Result<Self::Value, E> {
Ok(v)
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
let mut out = Vec::with_capacity(seq.size_hint().unwrap_or(0).min(MAX_SEQ_PREALLOC));
while let Some(byte) = seq.next_element::<u8>()? {
out.push(byte);
}
Ok(out)
}
}
pub(crate) mod bin_bytes {
use super::*;
pub fn serialize<S: Serializer>(value: &[u8], serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_bytes(value)
}
pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Vec<u8>, D::Error> {
deserializer.deserialize_any(ByteVecVisitor)
}
}
pub(crate) mod opt_bin_bytes {
use super::*;
pub fn serialize<S: Serializer>(
value: &Option<Vec<u8>>,
serializer: S,
) -> Result<S::Ok, S::Error> {
match value {
Some(bytes) => serializer.serialize_some(&ByteBuf(bytes)),
None => serializer.serialize_none(),
}
}
pub fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Option<Vec<u8>>, D::Error> {
struct OptVisitor;
impl<'de> Visitor<'de> for OptVisitor {
type Value = Option<Vec<u8>>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("an optional CBOR byte string")
}
fn visit_none<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(None)
}
fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(None)
}
fn visit_some<D: Deserializer<'de>>(
self,
deserializer: D,
) -> Result<Self::Value, D::Error> {
bin_bytes::deserialize(deserializer).map(Some)
}
}
deserializer.deserialize_option(OptVisitor)
}
}
#[cfg(all(test, feature = "cbor"))]
mod tests {
use crate::framing::{decode_named, encode_named};
use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize, PartialEq, Debug)]
struct WithBytes {
#[serde(with = "super::bin_bytes")]
payload: Vec<u8>,
#[serde(default, with = "super::opt_bin_bytes")]
opt: Option<Vec<u8>>,
}
#[test]
fn given_byte_fields_when_round_tripped_through_cbor_then_should_preserve_bytes() {
let value = WithBytes {
payload: vec![0xff, 0x00, 0x10],
opt: Some(vec![1, 2, 3]),
};
let bytes = encode_named(&value).expect("encodes");
let back: WithBytes = decode_named(&bytes).expect("decodes");
assert_eq!(back, value);
}
#[test]
fn given_byte_fields_when_round_tripped_through_json_then_should_preserve_bytes() {
let value = WithBytes {
payload: vec![0xff, 0x00, 0x10],
opt: Some(vec![1, 2, 3]),
};
let json = serde_json::to_string(&value).expect("json encodes");
assert!(json.contains("[255,0,16]"), "bytes render as a JSON array");
let back: WithBytes = serde_json::from_str(&json).expect("json decodes");
assert_eq!(back, value);
}
#[test]
fn given_absent_option_when_round_tripped_then_should_stay_none() {
let value = WithBytes {
payload: Vec::new(),
opt: None,
};
let bytes = encode_named(&value).expect("encodes");
let back: WithBytes = decode_named(&bytes).expect("decodes");
assert_eq!(back.opt, None);
}
#[test]
fn given_a_byte_array_with_a_huge_declared_length_when_decoded_then_should_reject_without_oom()
{
let bytes = [
0xa1, 0x67, b'p', b'a', b'y', b'l', b'o', b'a', b'd', 0x9a, 0xff, 0xff, 0xff, 0xff, ];
assert!(decode_named::<WithBytes>(&bytes).is_err());
}
#[test]
fn given_a_wrong_typed_byte_field_when_decoded_then_should_reject() {
#[derive(Serialize)]
struct AsString {
payload: &'static str,
opt: Option<Vec<u8>>,
}
let bad = AsString {
payload: "not bytes",
opt: None,
};
let bytes = encode_named(&bad).expect("encodes");
assert!(decode_named::<WithBytes>(&bytes).is_err());
}
}