1#![cfg_attr(docsrs, feature(doc_cfg))]
2
3use proc_macro::TokenStream;
4use proc_macro2::TokenStream as TokenStream2;
5use quote::quote;
6use syn::{parse_macro_input, parse_quote, Data, DeriveInput, Member, Result, Type};
7
8struct StructField<'a> {
16 member: Member,
17 ty: &'a Type,
18 skip: bool,
21}
22
23fn struct_fields<'a>(input: &'a DeriveInput, derive: &str) -> Result<Vec<StructField<'a>>> {
28 let Data::Struct(data) = &input.data else {
29 return Err(syn::Error::new_spanned(
30 &input.ident,
31 format!("{derive} can only be derived for structs"),
32 ));
33 };
34
35 data.fields
36 .iter()
37 .enumerate()
38 .map(|(index, field)| {
39 let member = field
40 .ident
41 .clone()
42 .map_or_else(|| Member::Unnamed(syn::Index::from(index)), Member::Named);
43 Ok(StructField {
44 member,
45 ty: &field.ty,
46 skip: has_skip_attribute(&field.attrs)?,
47 })
48 })
49 .collect()
50}
51
52fn bounded_types<'a>(fields: &'a [StructField<'a>]) -> Vec<&'a Type> {
58 fields
59 .iter()
60 .filter(|field| !field.skip)
61 .map(|field| field.ty)
62 .collect()
63}
64
65fn impl_block(
68 input: &DeriveInput,
69 trait_path: &TokenStream2,
70 bounded: &[&Type],
71 items: &TokenStream2,
72) -> TokenStream2 {
73 let name = &input.ident;
74 let mut generics = input.generics.clone();
75
76 if !bounded.is_empty() {
77 let where_clause = generics.make_where_clause();
78 for ty in bounded {
79 where_clause.predicates.push(parse_quote!(#ty: #trait_path));
80 }
81 }
82
83 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
84 quote! {
85 impl #impl_generics #trait_path for #name #ty_generics #where_clause {
86 #items
87 }
88 }
89}
90
91fn field_repr_size(field_type: &Type) -> TokenStream2 {
99 quote! {
100 ::core::mem::size_of::<<#field_type as ::spongefish::FromUniform>::Repr>()
101 }
102}
103
104fn from_uniform_field_expr(field_type: &Type) -> TokenStream2 {
112 quote! {
113 {
114 let mut field_buf = <#field_type as ::spongefish::FromUniform>::Repr::default();
115 let field_size = ::core::convert::AsMut::<[u8]>::as_mut(&mut field_buf).len();
116 let start = offset;
117 let end = start + field_size;
118 assert!(
119 end <= bytes.len(),
120 "`FromUniform` derive: field representation is wider than the derived buffer; \
121 `Repr` must satisfy `size_of::<Repr>() == Repr::default().as_mut().len()`"
122 );
123 ::core::convert::AsMut::<[u8]>::as_mut(&mut field_buf)
124 .copy_from_slice(&bytes[start..end]);
125 offset = end;
126 <#field_type as ::spongefish::FromUniform>::from_uniform(field_buf)
127 }
128 }
129}
130
131fn generate_encoding_impl(input: &DeriveInput) -> Result<TokenStream2> {
132 let fields = struct_fields(input, "Encoding")?;
133 let bounded = bounded_types(&fields);
134
135 let field_encodings = fields.iter().filter(|field| !field.skip).map(|field| {
136 let member = &field.member;
137 quote! {
138 output.extend_from_slice(self.#member.encode().as_ref());
139 }
140 });
141
142 let trait_path = quote!(::spongefish::Encoding);
143 Ok(impl_block(
144 input,
145 &trait_path,
146 &bounded,
147 "e! {
148 fn encode(&self) -> impl AsRef<[u8]> {
149 let mut output = ::spongefish::__private::Vec::with_capacity(
156 ::core::mem::size_of::<Self>(),
157 );
158 #(#field_encodings)*
159 output
160 }
161 },
162 ))
163}
164
165fn generate_from_uniform_impl(input: &DeriveInput) -> Result<TokenStream2> {
166 let fields = struct_fields(input, "FromUniform")?;
167 let bounded = bounded_types(&fields);
168
169 let field_inits = fields.iter().map(|field| {
170 let member = &field.member;
171 if field.skip {
172 return quote!(#member: Default::default(),);
173 }
174 let from_uniform_field = from_uniform_field_expr(field.ty);
175 quote!(#member: #from_uniform_field,)
176 });
177
178 let size_components = bounded.iter().copied().map(field_repr_size);
179 let size_calc = if bounded.is_empty() {
180 quote!(0usize)
181 } else {
182 quote!(#(#size_components)+*)
183 };
184
185 let body = if bounded.is_empty() {
189 quote! {
190 let _ = buf;
191 Self { #(#field_inits)* }
192 }
193 } else {
194 quote! {
195 let bytes = buf.as_ref();
198 let mut offset = 0usize;
199 let value = Self { #(#field_inits)* };
200 assert_eq!(
201 offset,
202 bytes.len(),
203 "`FromUniform` derive: field representations do not cover the derived buffer; \
204 every `Repr` must satisfy `size_of::<Repr>() == Repr::default().as_mut().len()`"
205 );
206 value
207 }
208 };
209
210 let trait_path = quote!(::spongefish::FromUniform);
211 Ok(impl_block(
212 input,
213 &trait_path,
214 &bounded,
215 "e! {
216 type Repr = ::spongefish::ByteArray<{ #size_calc }>;
217
218 fn from_uniform(buf: Self::Repr) -> Self {
219 #body
220 }
221 },
222 ))
223}
224
225fn generate_from_narg_impl(input: &DeriveInput) -> Result<TokenStream2> {
226 let fields = struct_fields(input, "FromNarg")?;
227 let bounded = bounded_types(&fields);
228
229 let field_inits = fields.iter().map(|field| {
230 let member = &field.member;
231 if field.skip {
232 return quote!(#member: Default::default(),);
233 }
234 let field_type = field.ty;
235 quote! {
236 #member: ::spongefish::NargReader::read::<#field_type>(reader)?,
237 }
238 });
239
240 let trait_path = quote!(::spongefish::FromNarg);
241 Ok(impl_block(
242 input,
243 &trait_path,
244 &bounded,
245 "e! {
246 fn from_narg(
247 reader: &mut ::spongefish::NargReader<'_>,
248 ) -> ::core::result::Result<Self, ::spongefish::VerificationError> {
249 let _ = &reader;
252 Ok(Self { #(#field_inits)* })
253 }
254 },
255 ))
256}
257
258fn generate_unit_impl(input: &DeriveInput) -> Result<TokenStream2> {
259 let fields = struct_fields(input, "Unit")?;
260 let bounded = bounded_types(&fields);
261
262 let zero_fields = fields.iter().map(|field| {
263 let member = &field.member;
264 if field.skip {
265 return quote!(#member: ::core::default::Default::default(),);
266 }
267 let field_type = field.ty;
268 quote!(#member: <#field_type as ::spongefish::Unit>::ZERO,)
269 });
270
271 let trait_path = quote!(::spongefish::Unit);
272 Ok(impl_block(
273 input,
274 &trait_path,
275 &bounded,
276 "e! {
277 const ZERO: Self = Self { #(#zero_fields)* };
278 },
279 ))
280}
281
282fn expand(generated: Result<TokenStream2>) -> TokenStream {
285 TokenStream::from(generated.unwrap_or_else(syn::Error::into_compile_error))
286}
287
288#[proc_macro_derive(Encoding, attributes(spongefish))]
316pub fn derive_encoding(input: TokenStream) -> TokenStream {
317 let input = parse_macro_input!(input as DeriveInput);
318 expand(generate_encoding_impl(&input))
319}
320
321#[proc_macro_derive(FromUniform, attributes(spongefish))]
326pub fn derive_from_uniform(input: TokenStream) -> TokenStream {
327 let input = parse_macro_input!(input as DeriveInput);
328 expand(generate_from_uniform_impl(&input))
329}
330
331#[proc_macro_derive(FromNarg, attributes(spongefish))]
338pub fn derive_from_narg(input: TokenStream) -> TokenStream {
339 let input = parse_macro_input!(input as DeriveInput);
340 expand(generate_from_narg_impl(&input))
341}
342
343#[proc_macro_derive(Codec, attributes(spongefish))]
347pub fn derive_codec(input: TokenStream) -> TokenStream {
348 let input = parse_macro_input!(input as DeriveInput);
349 expand((|| {
350 let encoding = generate_encoding_impl(&input)?;
351 let from_uniform = generate_from_uniform_impl(&input)?;
352 let from_narg = generate_from_narg_impl(&input)?;
353 Ok(quote! {
354 #encoding
355 #from_uniform
356 #from_narg
357 })
358 })())
359}
360
361#[proc_macro_derive(Unit, attributes(spongefish))]
377pub fn derive_unit(input: TokenStream) -> TokenStream {
378 let input = parse_macro_input!(input as DeriveInput);
379 expand(generate_unit_impl(&input))
380}
381
382fn has_skip_attribute(attrs: &[syn::Attribute]) -> Result<bool> {
397 for attr in attrs {
398 if !attr.path().is_ident("spongefish") {
399 continue;
400 }
401
402 let mut skip = false;
403 attr.parse_nested_meta(|meta| {
404 if meta.path.is_ident("skip") {
405 skip = true;
406 Ok(())
407 } else {
408 Err(meta.error("unknown `spongefish` option; expected `skip`"))
409 }
410 })?;
411
412 if !skip {
413 return Err(syn::Error::new_spanned(
414 attr,
415 "empty `#[spongefish(..)]`: write `#[spongefish(skip)]` or remove the attribute",
416 ));
417 }
418 return Ok(true);
419 }
420
421 Ok(false)
422}