Skip to main content

validated_struct_macros/
lib.rs

1#![allow(clippy::too_many_arguments)]
2
3use proc_macro::TokenStream;
4use quote::quote;
5use syn::{Attribute, Ident, Type};
6
7use structure::{FieldSpec, FieldType, StructSpec};
8
9mod structure;
10
11const SEPARATOR: char = if cfg!(feature = "dot_separator") {
12    '.'
13} else {
14    '/'
15};
16impl StructSpec {
17    fn structure(&self) -> impl quote::ToTokens {
18        unzip_n::unzip_n!(10);
19        let ident = &self.ident;
20        let mut notifying = false;
21        let sattrs: Vec<_> = self
22            .attrs
23            .iter()
24            .filter_map(|a| {
25                if a.path().is_ident("notifying") {
26                    notifying = true;
27                    None
28                } else {
29                    Some(a.clone())
30                }
31            })
32            .collect();
33        let (
34            fields,
35            args,
36            associations,
37            accessors,
38            constructor_validations,
39            constructor_rec_validations,
40            serde_match,
41            get_match,
42            json_get_match,
43            get_keys,
44        ) = self
45            .fields
46            .iter()
47            .map(|spec| {
48                let id = &spec.ident;
49                let field_name = id;
50                let ty = spec.ty.ty();
51                let predicate = spec.constraint.as_ref().map(|e| quote! {#e(&value)});
52                let str_id = format!("{}", id);
53                let set_id = quote::format_ident!("set_{}", id);
54                let validate_id = quote::format_ident!("validate_{}", id);
55                let validate_id_rec = quote::format_ident!("validate_{}_rec", id);
56                (
57                    field(spec, field_name),
58                    quote! {#id: #ty},
59                    quote! {#field_name},
60                    accessors(
61                        spec,
62                        id,
63                        field_name,
64                        &ty,
65                        &set_id,
66                        &validate_id,
67                        &validate_id_rec,
68                        &predicate,
69                    ),
70                    if predicate.is_some() {
71                        Some(quote! {Self::#validate_id(self.#id())})
72                    } else {
73                        None
74                    },
75                    quote! {Self::#validate_id_rec(self.#id())},
76                    serde_match(id, spec, &str_id, set_id, field_name),
77                    get_match(spec, id, &str_id, field_name),
78                    json_get_match(spec, id, &str_id, field_name),
79                    keys_match(spec, field_name, &str_id),
80                )
81            })
82            .collect::<Vec<_>>()
83            .into_iter()
84            .unzip_n_vec();
85        let serde_access =
86            serde_access(ident, &serde_match, &get_match, &get_keys, &json_get_match);
87        let constructor_validations = constructor_validations
88            .into_iter()
89            .flatten()
90            .collect::<Vec<_>>();
91        main_implementation(
92            &sattrs,
93            ident,
94            &fields,
95            &constructor_validations,
96            &constructor_rec_validations,
97            &args,
98            &associations,
99            &accessors,
100            &serde_access,
101        )
102    }
103}
104
105fn main_implementation(
106    sattrs: &[Attribute],
107    ident: &Ident,
108    fields: &[proc_macro2::TokenStream],
109    constructor_validations: &[proc_macro2::TokenStream],
110    constructor_rec_validations: &[proc_macro2::TokenStream],
111    args: &[proc_macro2::TokenStream],
112    associations: &[proc_macro2::TokenStream],
113    accessors: &[proc_macro2::TokenStream],
114    serde_access: &Option<proc_macro2::TokenStream>,
115) -> proc_macro2::TokenStream {
116    quote! {
117        #(#sattrs)*
118        pub struct #ident {
119            #(#fields),*
120        }
121        impl #ident {
122            pub fn validate(&self) -> bool {
123                true #(&& #constructor_validations)*
124            }
125            fn validate_rec(&self) -> bool {
126                true #(&& #constructor_rec_validations)*
127            }
128            #[allow(clippy::too_many_arguments)]
129            pub fn new(#(#args),*) -> Result<Self, Self> {
130                let constructed = #ident {
131                    #(#associations),*
132                };
133                if constructed.validate() {Ok(constructed)} else {Err(constructed)}
134            }
135            #(#accessors)*
136        }
137        #serde_access
138    }
139}
140
141fn field(spec: &FieldSpec, field: &Ident) -> proc_macro2::TokenStream {
142    let ty = spec.ty.ty();
143    let attrs = &spec.attributes;
144    let vis = &spec.vis;
145    quote! {#(#attrs)* #vis #field: #ty}
146}
147
148fn serde_access(
149    ident: &Ident,
150    serde_match: &[proc_macro2::TokenStream],
151    get_match: &[proc_macro2::TokenStream],
152    get_keys: &[proc_macro2::TokenStream],
153    json_get_match: &[proc_macro2::TokenStream],
154) -> Option<proc_macro2::TokenStream> {
155    let get_json = cfg!(feature = "serde_json").then(|| {
156        quote! {
157            fn get_json(& self, key: &str) -> Result<String, validated_struct::GetError>{
158                match validated_struct::split_once(key, #SEPARATOR) {
159                    #(#json_get_match)*
160                    ("", key) if !key.is_empty() => self.get_json(key),
161                    _ => Err(validated_struct::GetError::NoMatchingKey),
162                }
163            }
164        }
165    });
166    cfg!(feature = "serde").then(|| quote! {
167        impl #ident {
168            pub fn from_deserializer<'d, D: serde::Deserializer<'d>>(
169                d: D,
170            ) -> Result<Self, Result<Self, D::Error>>
171            where
172                Self: serde::Deserialize<'d>,
173            {
174                match <Self as serde::Deserialize>::deserialize(d) {
175                    Ok(value) => {
176                        if value.validate_rec() {
177                            Ok(value)
178                        } else {
179                            Err(Ok(value))
180                        }
181                    }
182                    Err(e) => Err(Err(e)),
183                }
184            }
185        }
186        impl<'a> validated_struct::ValidatedMapAssociatedTypes<'a> for #ident {
187            type Accessor = &'a dyn std::any::Any;
188        }
189        impl validated_struct::ValidatedMap for #ident {
190            fn insert<'d, D: serde::Deserializer<'d>>(&mut self, key: &str, value: D) -> Result<(), validated_struct::InsertionError>
191            where
192                validated_struct::InsertionError: From<D::Error> {
193                if let Some(e) = match validated_struct::split_once(key, #SEPARATOR) {
194                    #(#serde_match)*
195                    ("", key) if !key.is_empty() => self.insert(key, value).err(),
196                    _ => Some("unknown key".into())
197                } {return Err(e)};
198                Ok(())
199            }
200            fn get<'a>(&'a self, key: &str) -> Result<&dyn std::any::Any, validated_struct::GetError>{
201                match validated_struct::split_once(key, #SEPARATOR) {
202                    #(#get_match)*
203                    ("", key) if !key.is_empty() => self.get(key),
204                    _ => Err(validated_struct::GetError::NoMatchingKey),
205                }
206            }
207            #get_json
208            type Keys = std::vec::Vec<String>;
209            fn keys(&self) -> Self::Keys {
210                let mut keys = std::vec::Vec::new();
211                #(#get_keys)*
212                keys
213            }
214        }
215    })
216}
217
218fn accessors(
219    spec: &FieldSpec,
220    id: &Ident,
221    field: &Ident,
222    ty: &Type,
223    set_id: &Ident,
224    validate_id: &Ident,
225    validate_id_rec: &Ident,
226    predicate: &Option<proc_macro2::TokenStream>,
227) -> proc_macro2::TokenStream {
228    let doc_attrs: Vec<_> = spec
229        .attributes
230        .iter()
231        .filter(|&attr| attr.path().is_ident("doc"))
232        .cloned()
233        .collect();
234    let validate_id_rec_impl =
235        implement_validation(spec, ty, predicate, validate_id_rec, validate_id);
236    match predicate {
237        Some(predicate) => quote! {
238            #[inline(always)]
239            #(#doc_attrs)*
240            pub fn #id(&self) -> & #ty {
241                &self.#field
242            }
243            #[allow(clippy::ptr_arg)]
244            pub fn #validate_id(value: &#ty) -> bool {
245                #predicate
246            }
247            #validate_id_rec_impl
248            #(#doc_attrs)*
249            pub fn #set_id(&mut self, mut value: #ty) -> Result<#ty, #ty> {
250                if Self::#validate_id(&value) {
251                    std::mem::swap(&mut self.#field, &mut value);
252                    Ok(value)
253                } else {
254                    Err(value)
255                }
256            }
257        },
258        None => quote! {
259            #[inline(always)]
260            #(#doc_attrs)*
261            pub fn #id(&self) -> & #ty {
262                &self.#field
263            }
264            #validate_id_rec_impl
265            #(#doc_attrs)*
266            pub fn #set_id(&mut self, mut value: #ty) -> Result<#ty, #ty> {
267                std::mem::swap(&mut self.#field, &mut value);
268                Ok(value)
269            }
270        },
271    }
272}
273
274fn keys_match(spec: &FieldSpec, field: &Ident, str_id: &str) -> proc_macro2::TokenStream {
275    match spec.ty {
276        FieldType::Concrete(_) => quote! {keys.push(#str_id.into());},
277        FieldType::Structure(_) => quote! {
278            keys.push(#str_id.into());
279            keys.extend(self.#field.keys().into_iter().map(|s|format!("{}{}{}",#str_id, #SEPARATOR, s.as_str())));
280        },
281    }
282}
283
284fn get_match(
285    spec: &FieldSpec,
286    id: &Ident,
287    str_id: &str,
288    field: &Ident,
289) -> proc_macro2::TokenStream {
290    let get_exact = quote! {(#str_id, "") => Ok(self.#id() as &dyn std::any::Any),};
291    if spec.recursive_accessors() {
292        quote! {
293            #get_exact
294            (#str_id, key) => self.#field.get(key),
295        }
296    } else {
297        get_exact
298    }
299}
300
301fn json_get_match(
302    spec: &FieldSpec,
303    id: &Ident,
304    str_id: &str,
305    field: &Ident,
306) -> proc_macro2::TokenStream {
307    let get_exact = quote! {(#str_id, "") => serde_json::to_string(self.#id()).map_err(|e| validated_struct::GetError::Other(e.into())),};
308    if spec.recursive_accessors() {
309        quote! {
310            #get_exact
311            (#str_id, key) => self.#field.get_json(key),
312        }
313    } else {
314        get_exact
315    }
316}
317
318fn serde_match(
319    id: &Ident,
320    spec: &FieldSpec,
321    str_id: &str,
322    set_id: Ident,
323    field: &Ident,
324) -> proc_macro2::TokenStream {
325    let serde_set_err = format!("Predicate rejected value for {}", id);
326    let set_exact = quote! {
327        (#str_id, "") => self.#set_id(serde::Deserialize::deserialize(value)?).is_err().then(||#serde_set_err.into()),
328    };
329    if spec.recursive_accessors() {
330        quote! {
331            #set_exact
332            (#str_id, key) => self.#field.insert(key, value).err(),
333        }
334    } else {
335        set_exact
336    }
337}
338
339fn implement_validation(
340    f: &FieldSpec,
341    ty: &Type,
342    predicate: &Option<proc_macro2::TokenStream>,
343    validate_id_rec: &Ident,
344    validate_id: &Ident,
345) -> proc_macro2::TokenStream {
346    if let FieldType::Structure(_) = f.ty {
347        match predicate {
348            Some(predicate) => quote! {
349                fn #validate_id_rec(value: &#ty) -> bool {
350                    value.validate_rec() && #predicate
351                }
352            },
353            None => quote! {
354                fn #validate_id_rec(value: &#ty) -> bool {
355                    value.validate_rec()
356                }
357            },
358        }
359    } else {
360        let validate_rec_inner = match *predicate {
361            Some(_) => quote! {Self::#validate_id(value)},
362            None => quote! {true},
363        };
364        quote! {
365            #[allow(clippy::ptr_arg)]
366            fn #validate_id_rec(value: &#ty) -> bool {
367                #validate_rec_inner
368            }
369        }
370    }
371}
372
373#[proc_macro]
374pub fn validator(stream: TokenStream) -> TokenStream {
375    let spec: StructSpec = syn::parse(stream).unwrap();
376    let structure: Vec<_> = spec.flatten().iter().map(StructSpec::structure).collect();
377    (quote! {
378        #(#structure)*
379    })
380    .into()
381}
382
383#[cfg(test)]
384mod constructor_tests {
385    use super::StructSpec;
386    use quote::ToTokens;
387    use syn::{parse_quote, ExprStruct, File, ImplItem, Item, Stmt};
388    #[test]
389    fn constructor_spelling_preserves_complete_field_associations() {
390        let source: StructSpec = parse_quote! { Service { port: u16, label: String } };
391        let mut syntax: File = syn::parse2(source.structure().to_token_stream()).unwrap();
392        let Item::Impl(mut implementation) = syntax.items.remove(1) else {
393            panic!("expected owning implementation")
394        };
395        let ImplItem::Fn(mut constructor) = implementation.items.remove(2) else {
396            panic!("expected constructor")
397        };
398        let Stmt::Local(constructed) = constructor.block.stmts.remove(0) else {
399            panic!("expected constructed value")
400        };
401        let mut actual: ExprStruct =
402            syn::parse2(constructed.init.unwrap().expr.to_token_stream()).unwrap();
403        let mut expected: ExprStruct = parse_quote! { Service { port: port, label: label } };
404        for field in actual.fields.iter_mut().chain(expected.fields.iter_mut()) {
405            field.colon_token = None;
406        }
407        assert_eq!(
408            actual.to_token_stream().to_string(),
409            expected.to_token_stream().to_string()
410        );
411    }
412}