diff --git a/tls_codec/src/quic_vec.rs b/tls_codec/src/quic_vec.rs index d477dd9fb..90de9bf75 100644 --- a/tls_codec/src/quic_vec.rs +++ b/tls_codec/src/quic_vec.rs @@ -112,6 +112,36 @@ impl DeserializeBytes for Vec { } } +impl SerializeBytes for VLBytes { + #[inline(always)] + fn tls_serialize(&self) -> Result, Error> { + let content_length = self.as_slice().len(); + let length = ContentLength::from_usize(content_length)?; + let len_len = length.0.bytes_len(); + + let mut out = Vec::with_capacity(content_length + len_len); + out.resize(len_len, 0); + length.0.write_bytes(&mut out)?; + + // Extend with the data + out.extend(self.as_slice()); + + #[cfg(debug_assertions)] + if out.len() - len_len != content_length { + return Err(Error::LibraryError); + } + + Ok(out) + } +} + +impl SerializeBytes for &VLBytes { + #[inline(always)] + fn tls_serialize(&self) -> Result, Error> { + (*self).tls_serialize() + } +} + impl SerializeBytes for &[T] { #[inline(always)] fn tls_serialize(&self) -> Result, Error> { @@ -672,7 +702,7 @@ mod rw_bytes { impl Serialize for &VLBytes { #[inline(always)] fn tls_serialize(&self, writer: &mut W) -> Result { - (*self).tls_serialize(writer) + Serialize::tls_serialize(*self, writer) } } @@ -802,7 +832,7 @@ mod secret_bytes { impl Serialize for SecretVLBytes { fn tls_serialize(&self, writer: &mut W) -> Result { - self.0.tls_serialize(writer) + Serialize::tls_serialize(&self.0, writer) } } diff --git a/tls_codec/tests/encode.rs b/tls_codec/tests/encode.rs index 94cda264f..8cb213c00 100644 --- a/tls_codec/tests/encode.rs +++ b/tls_codec/tests/encode.rs @@ -87,3 +87,15 @@ fn serialize_var_len_boundaries() { let serialized = v.tls_serialize_detached().expect("Error encoding vector"); assert_eq!(&serialized[0..5], &[0x80, 0, 0x40, 0, 99]); } + +#[test] +/// Test that serializations of [`VLBytes`] match for [`tls_codec::Serialize`] +/// and [`tls_codec::SerializeBytes`]. +fn test_matching_vl_bytes_serialization() { + use tls_codec::SerializeBytes; + let byte_vec = VLBytes::new(vec![99u8; 16384]); + assert_eq!( + Serialize::tls_serialize_detached(&byte_vec).expect("Error encoding vector"), + SerializeBytes::tls_serialize(&byte_vec).expect("Error encoding byte vector") + ); +} diff --git a/tls_codec/tests/encode_bytes.rs b/tls_codec/tests/encode_bytes.rs index 50e3a33b6..cb6ff9167 100644 --- a/tls_codec/tests/encode_bytes.rs +++ b/tls_codec/tests/encode_bytes.rs @@ -1,4 +1,8 @@ -use tls_codec::{SerializeBytes, TlsByteVecU8, TlsByteVecU16, TlsByteVecU24, TlsByteVecU32, U24}; +#![allow(deprecated)] + +use tls_codec::{ + SerializeBytes, TlsByteVecU8, TlsByteVecU16, TlsByteVecU24, TlsByteVecU32, U24, VLBytes, +}; #[test] fn serialize_primitives() { @@ -82,3 +86,12 @@ fn serialize_tls_byte_vec_u32() { .expect("Error encoding byte vector"); assert_eq!(actual_result, vec![0, 0, 0, 3, 1, 2, 3]); } + +#[test] +fn serialize_vlbytes() { + let byte_vec = VLBytes::new(vec![1, 2, 3]); + let actual_result = byte_vec + .tls_serialize() + .expect("Error encoding byte vector"); + assert_eq!(actual_result, vec![3, 1, 2, 3]); +}