From b38c0dcd9fa59a527e1c26ef90ec128b2f94290c Mon Sep 17 00:00:00 2001 From: Nicolas Stalder Date: Thu, 18 Feb 2021 01:44:13 +0100 Subject: [PATCH] Add `slice` attribute to derive macros --- derive/src/decodable.rs | 15 ++++++--- derive/src/encodable.rs | 10 ++++-- derive/src/lib.rs | 74 +++++++++++++++++++++++++---------------- src/decoder.rs | 7 ++++ src/traits.rs | 60 ++++++++++++++++++--------------- tests/derive.rs | 37 ++++++++++++++++++--- 6 files changed, 137 insertions(+), 66 deletions(-) diff --git a/derive/src/decodable.rs b/derive/src/decodable.rs index 6cb52c8..98a8510 100644 --- a/derive/src/decodable.rs +++ b/derive/src/decodable.rs @@ -3,7 +3,7 @@ use quote::{quote, ToTokens}; use syn::{Attribute, DataStruct, Field, Ident}; use synstructure::Structure; -use crate::{extract_tag, FieldAttrs}; +use crate::{extract_attrs, FieldAttrs}; /// Derive Decodable on a struct pub(crate) struct DeriveDecodableStruct { @@ -17,7 +17,7 @@ pub(crate) struct DeriveDecodableStruct { impl DeriveDecodableStruct { pub fn derive(s: Structure<'_>, data: &DataStruct, name: &Ident, attrs: &[Attribute]) -> TokenStream { - let tag = extract_tag(name, attrs); + let (tag, _) = extract_attrs(name, attrs); let mut state = Self { decode_fields: TokenStream::new(), @@ -41,7 +41,14 @@ impl DeriveDecodableStruct { fn derive_field_decoder(&mut self, field: &FieldAttrs) { let field_name = &field.name; let field_tag = field.tag; - let field_decoder = quote! { let #field_name = decoder.decode_tagged_value(::simple_tlv::Tag::try_from(#field_tag).unwrap())?; }; + let field_decoder = if field.slice { + quote! { let #field_name = + decoder.decode_tagged_slice(::simple_tlv::Tag::try_from(#field_tag).unwrap())?.try_into() + .map_err(|_| simple_tlv::ErrorKind::Length { tag: simple_tlv::Tag::try_from(#field_tag).unwrap() })?; + } + } else { + quote! { let #field_name = decoder.decode_tagged_value(::simple_tlv::Tag::try_from(#field_tag).unwrap())?; } + }; field_decoder.to_tokens(&mut self.decode_fields); let field_result = quote!(#field_name,); @@ -58,7 +65,7 @@ impl DeriveDecodableStruct { gen impl<'a> core::convert::TryFrom> for @Self { type Error = simple_tlv::Error; - fn try_from(tagged_slice: simple_tlv::TaggedSlice<'a>) -> simple_tlv::Result { + fn try_from(tagged_slice: simple_tlv::TaggedSlice<'a>) -> simple_tlv::Result { use core::convert::TryInto; tagged_slice.tag().assert_eq(simple_tlv::Tag::try_from(#tag).unwrap())?; tagged_slice.decode_nested(|decoder| { diff --git a/derive/src/encodable.rs b/derive/src/encodable.rs index ec4ae8f..e7f9ecc 100644 --- a/derive/src/encodable.rs +++ b/derive/src/encodable.rs @@ -3,7 +3,7 @@ use quote::{quote, ToTokens}; use syn::{Attribute, DataStruct, Field, Ident}; use synstructure::Structure; -use crate::{extract_tag, FieldAttrs}; +use crate::{extract_attrs, FieldAttrs}; /// Derive Encodable on a struct pub(crate) struct DeriveEncodableStruct { @@ -14,7 +14,7 @@ pub(crate) struct DeriveEncodableStruct { impl DeriveEncodableStruct { pub fn derive(s: Structure<'_>, data: &DataStruct, name: &Ident, attrs: &[Attribute]) -> TokenStream { - let tag = extract_tag(name, attrs); + let (tag, _) = extract_attrs(name, attrs); let mut state = Self { encode_fields: TokenStream::new(), @@ -37,7 +37,11 @@ impl DeriveEncodableStruct { fn derive_field_encoder(&mut self, field: &FieldAttrs) { let field_name = &field.name; let field_tag = field.tag; - let field_encoder = quote! { &(::simple_tlv::Tag::try_from(#field_tag).unwrap().with_value(&self.#field_name)), }; + let field_encoder = if field.slice { + quote! { &(::simple_tlv::TaggedSlice::from(simple_tlv::Tag::try_from(#field_tag).unwrap(), &self.#field_name)?), } + } else { + quote! { &(::simple_tlv::Tag::try_from(#field_tag).unwrap().with_value(&self.#field_name)), } + }; field_encoder.to_tokens(&mut self.encode_fields); } diff --git a/derive/src/lib.rs b/derive/src/lib.rs index e78e470..b1e85f4 100644 --- a/derive/src/lib.rs +++ b/derive/src/lib.rs @@ -1,4 +1,8 @@ //! Custom derive support for the `simple-tlv` crate +//! +//! With `#[tlv(slice)]` set, `Encodable` should work for fields implementing `AsRef<[u8]>`, +//! and `Decodable` should work for fields implementing `TryFrom<[u8]>`, even if the field +//! is not `Decodable` or `Encodable`. #![crate_type = "proc-macro"] #![warn(rust_2018_idioms, trivial_casts, unused_qualifications)] @@ -69,8 +73,11 @@ struct FieldAttrs { /// Name of the field pub name: Ident, - /// Value of the `#[asn1(type = "...")]` attribute if provided + /// Value of the `#[tlv(tag = "...")]` attribute if provided pub tag: u8, + + /// Whether the `#[tlv(slice)]` attribute was set + pub slice: bool } impl FieldAttrs { @@ -82,14 +89,15 @@ impl FieldAttrs { .cloned() .expect("no name on struct field i.e. tuple structs unsupported"); - let tag = extract_tag(&name, &field.attrs); + let (tag, slice) = extract_attrs(&name, &field.attrs); - Self { name, tag } + Self { name, tag, slice } } } -fn extract_tag(name: &Ident, attrs: &[Attribute]) -> u8 { +fn extract_attrs(name: &Ident, attrs: &[Attribute]) -> (u8, bool) { let mut tag = None; + let mut slice = false; for attr in attrs { if !attr.path.is_ident("tlv") { @@ -97,45 +105,53 @@ fn extract_tag(name: &Ident, attrs: &[Attribute]) -> u8 { } match attr.parse_meta().expect("error parsing `tlv` attribute") { - Meta::List(MetaList { nested, .. }) if nested.len() == 1 => { - match nested.first() { - Some(NestedMeta::Meta(Meta::NameValue(MetaNameValue { - path, - lit: Lit::Str(lit_str), - .. - }))) => { - // Parse the `type = "..."` attribute - if !path.is_ident("tag") { - panic!("unknown `tlv` attribute for field `{}`: {:?}", name, path); + Meta::List(MetaList { nested, .. }) if !nested.is_empty() => { + for entry in nested { + match entry { + NestedMeta::Meta(Meta::Path(path)) => { + if !path.is_ident("slice") { + panic!("unknown `tlv` attribute for field `{}`: {:?}", name, path); + } + slice = true; } + NestedMeta::Meta(Meta::NameValue(MetaNameValue { + path, + lit: Lit::Str(lit_str), + .. + })) => { + // Parse the `type = "..."` attribute + if !path.is_ident("tag") { + panic!("unknown `tlv` attribute for field `{}`: {:?}", name, path); + } - if tag.is_some() { - panic!("duplicate SIMPLE-TLV `tag` attribute for field: {}", name); - } + if tag.is_some() { + panic!("duplicate SIMPLE-TLV `tag` attribute for field: {}", name); + } - let possibly_with_prefix = lit_str.value(); - let without_prefix = possibly_with_prefix.trim_start_matches("0x"); - let tag_value = u8::from_str_radix(without_prefix, 16).expect("tag values must be between one and 254"); - if tag_value == 0 || tag_value == 255 { - panic!("SIMPLE-TLV tags must not be zero or 255"); + let possibly_with_prefix = lit_str.value(); + let without_prefix = possibly_with_prefix.trim_start_matches("0x"); + let tag_value = u8::from_str_radix(without_prefix, 16).expect("tag values must be between one and 254"); + if tag_value == 0 || tag_value == 255 { + panic!("SIMPLE-TLV tags must not be zero or 255"); + } + tag = Some(tag_value); } - tag = Some(tag_value); + other => panic!( + "a malformed `tlv` attribute for field `{}`: {:?}", + name, other + ), } - other => panic!( - "malformed `tlv` attribute for field `{}`: {:?}", - name, other - ), } } other => panic!( - "malformed `tlv` attribute for field `{}`: {:?}", + "malformed `tlv` attribute for field `{}`: {:#?}", name, other ), } } if let Some(tag) = tag { - tag + (tag, slice) } else { panic!("SIMPLE-TLV tag missing for `{}`", name); } diff --git a/src/decoder.rs b/src/decoder.rs index 6eb5d52..b34669e 100644 --- a/src/decoder.rs +++ b/src/decoder.rs @@ -42,6 +42,13 @@ impl<'a> Decoder<'a> { Self::new(tagged.as_bytes()).decode() } + /// Decode a TaggedSlice with tag checked to be as expected, returning the value + pub fn decode_tagged_slice(&mut self, tag: crate::Tag) -> Result<&'a [u8]> { + let tagged: crate::TaggedSlice = self.decode()?; + tagged.tag().assert_eq(tag)?; + Ok(tagged.as_bytes()) + } + /// Return an error with the given [`ErrorKind`], annotating it with /// context about where the error occurred. pub fn error(&mut self, kind: ErrorKind) -> Result { diff --git a/src/traits.rs b/src/traits.rs index 6e5c0cb..69f384f 100644 --- a/src/traits.rs +++ b/src/traits.rs @@ -1,21 +1,18 @@ // pub use der::{Decodable, Encodable}; //! Trait definitions -use core::convert::TryFrom; +use core::convert::{TryFrom, TryInto}; use crate::{Decoder, Encoder, Error, header::Header, Length, Result, Tag, TaggedSlice, TaggedValue}; #[cfg(feature = "alloc")] use { alloc::vec::Vec, - core::{convert::TryInto, iter}, + core::iter, crate::ErrorKind, }; #[cfg(feature = "heapless")] -use { - core::convert::TryInto, - crate::ErrorKind, -}; +use crate::ErrorKind; /// Decoding trait. /// @@ -244,27 +241,38 @@ where // } // } +impl<'a> Encodable for &'a [u8] { + fn encoded_length(&self) -> Result { + self.len().try_into() + } + + /// Encode this value as SIMPLE-TLV using the provided [`Encoder`]. + fn encode(&self, encoder: &mut Encoder<'_>) -> Result<()> { + encoder.bytes(self) + } +} + macro_rules! impl_array { ($($N:literal),*) => { $( - impl Encodable for [u8; $N] { - fn encoded_length(&self) -> Result { - Ok(($N as u8).into()) + impl Encodable for [u8; $N] { + fn encoded_length(&self) -> Result { + Ok(($N as u8).into()) + } + + /// Encode this value as SIMPLE-TLV using the provided [`Encoder`]. + fn encode(&self, encoder: &mut Encoder<'_>) -> Result<()> { + encoder.bytes(self.as_ref()) + } } - /// Encode this value as SIMPLE-TLV using the provided [`Encoder`]. - fn encode(&self, encoder: &mut Encoder<'_>) -> Result<()> { - encoder.bytes(self.as_ref()) + impl Decodable<'_> for [u8; $N] { + fn decode(decoder: &mut Decoder<'_>) -> Result { + use core::convert::TryInto; + let bytes: &[u8] = decoder.bytes($N as u8)?; + Ok(bytes.try_into().unwrap()) + } } - } - - impl Decodable<'_> for [u8; $N] { - fn decode(decoder: &mut Decoder<'_>) -> Result { - use core::convert::TryInto; - let bytes: &[u8] = decoder.bytes($N as u8)?; - Ok(bytes.try_into().unwrap()) - } - } )* } } @@ -331,10 +339,10 @@ mod tests { { // both approaches equivalent field_encoder(&[ - &(Tag::try_from(0x11).unwrap().with_value(&self.x)), + &(Tag::try_from(0x11).unwrap().with_value(&self.x.as_ref())), // &self.x.tagged(Tag::try_from(0x11).unwrap()), - &self.y.tagged(Tag::try_from(0x22).unwrap()), - &self.z.tagged(Tag::try_from(0x33).unwrap()), + &self.y.as_ref().tagged(Tag::try_from(0x22).unwrap()), + &self.z.as_ref().tagged(Tag::try_from(0x33).unwrap()), ]) } @@ -396,7 +404,7 @@ mod tests { { field_encoder(&[ &self.s.tagged(Tag::try_from(0x1).unwrap()), - &self.t.tagged(Tag::try_from(0x2).unwrap()), + &self.t.as_ref().tagged(Tag::try_from(0x2).unwrap()), ]) } } @@ -464,7 +472,7 @@ mod tests { { field_encoder(&[ &self.s, - &self.t.tagged(Tag::try_from(0x2).unwrap()), + &self.t.as_ref().tagged(Tag::try_from(0x2).unwrap()), ]) } } diff --git a/tests/derive.rs b/tests/derive.rs index 16a9834..8c5bacf 100644 --- a/tests/derive.rs +++ b/tests/derive.rs @@ -5,16 +5,23 @@ use simple_tlv::{Decodable, Encodable}; #[derive(Clone, Copy, Debug, Decodable, Encodable, Eq, PartialEq)] -#[tlv(tag = "AA")] +#[tlv(tag = "0xAA")] struct S { - #[tlv(tag = "11")] + #[tlv(slice, tag = "0x11")] x: [u8; 2], - #[tlv(tag = "22")] + #[tlv(slice, tag = "0x22")] y: [u8; 3], - #[tlv(tag = "33")] + #[tlv(slice, tag = "0x33")] z: [u8; 4], } +#[derive(Clone, Copy, Debug, Decodable, Encodable, Eq, PartialEq)] +#[tlv(tag = "0xBB")] +struct T { + #[tlv(tag = "0x44", slice)] + x: [u8; 1234], +} + #[test] fn derived_reconstruct() { let s = S { x: [1,2], y: [3,4,5], z: [6,7,8,9] }; @@ -35,3 +42,25 @@ fn derived_reconstruct() { assert_eq!(s, s2); } +#[test] +fn pretty_big() { + let mut x = [0u8; 1234]; + for (i, x) in x.iter_mut().enumerate() { + *x = i as _; + }; + + let t = T { x }; + + let mut buf = [0u8; 1024]; + assert!(t.encode_to_slice(&mut buf).is_err()); + + let mut buf = [0u8; 1500]; + let encoded = t.encode_to_slice(&mut buf).unwrap(); + + assert_eq!(&encoded[..8], [ + // 1234 + 4 + 0xBB, 0xFF, 0x04, 0xD6, + // 1234 + 0x44, 0xFF, 0x04, 0xD2]); + assert_eq!(&encoded[8..], x); +}