Skip to main content

lapdog_derive/
lib.rs

1use std::{collections::HashMap, ops::BitOr};
2
3use proc_macro2::TokenStream;
4use quote::{format_ident, quote};
5use syn::{DataStruct, DeriveInput, Field, Fields, Ident, parse_quote};
6
7#[proc_macro_derive(Entry, attributes(lapdog))]
8pub fn implement_from_entry(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
9    let input = syn::parse_macro_input!(item as DeriveInput);
10    let name = input.ident;
11    let (fields, object_name_field) = match parse_fields(
12        match input.data {
13            syn::Data::Struct(DataStruct { fields, .. }) => match fields {
14                Fields::Named(f) => f,
15                _ => panic!("Structs fields/attributes must be named to be derivable"),
16            },
17            _ => unimplemented!("non-struct derives are not supported"),
18        }
19        .named,
20    ) {
21        Ok(f) => f,
22        Err(e) => return e.into_compile_error().into(),
23    };
24    let (impl_generics, type_generics, where_clause) = input.generics.split_for_impl();
25
26    // Generics' type parameters
27    let generic_params: Vec<syn::Ident> = input
28        .generics
29        .params
30        .iter()
31        .filter_map(|param| match param {
32            syn::GenericParam::Type(type_param) => Some(type_param.ident.clone()),
33            _ => None,
34        })
35        .collect();
36
37    // If a field has a generic parameter
38    let mut generic_bounds = HashMap::<syn::Ident, NeedsBound>::new();
39    for field in &fields {
40        if let syn::Type::Path(type_path) = &field.field.ty {
41            if let Some(ident) = type_path.path.get_ident() {
42                if generic_params.contains(ident) {
43                    let this_field = if field.multiple {
44                        NeedsBound::Multiple
45                    } else {
46                        NeedsBound::Octet
47                    };
48                    generic_bounds
49                        .entry(ident.clone())
50                        .and_modify(|x| *x = *x | this_field)
51                        .or_insert(this_field);
52                }
53            }
54        }
55    }
56
57    let mut where_preds: Vec<syn::WherePredicate> = where_clause
58        .map(|wc| wc.predicates.clone().into_iter().collect())
59        .unwrap_or_default();
60
61    for (ident, needs_bound) in generic_bounds {
62        let multi = || {
63            [
64                parse_quote!(#ident: lapdog::search::FromMultipleOctetStrings),
65                parse_quote!(<#ident as lapdog::search::FromMultipleOctetStrings>::Err: 'static),
66            ]
67        };
68        let single = || {
69            [
70                parse_quote!(#ident: lapdog::search::FromOctetString),
71                parse_quote!(<#ident as lapdog::search::FromOctetString>::Err: 'static),
72            ]
73        };
74        match needs_bound {
75            NeedsBound::Both => {
76                where_preds.extend(multi());
77                where_preds.extend(single());
78            }
79            NeedsBound::Multiple => {
80                where_preds.extend(multi());
81            }
82            NeedsBound::Octet => {
83                where_preds.extend(single());
84            }
85        }
86    }
87
88    let where_clause = if where_preds.is_empty() {
89        quote!()
90    } else {
91        quote!(where #(#where_preds),*)
92    };
93
94    let insert_object_name = object_name_field.as_ref().map(insert_object_name);
95    let field_quotes = fields.iter().map(field_line);
96    let field_names = fields.iter().map(|x| x.ident());
97    let attribute_names = fields.iter().map(|x| x.attribute_name.clone());
98    quote!(
99        impl #impl_generics lapdog::search::FromEntry for #name #type_generics #where_clause {
100            fn from_entry(entry: lapdog::search::RawEntry) -> Result<#name #type_generics, lapdog::search::FailedToGetFromEntry> {
101                #( #field_quotes )*
102                Ok(#name { #(#field_names,)* #insert_object_name })
103            }
104
105            fn attributes() -> Option<impl Iterator<Item = &'static str>> {
106                Some(vec![#(#attribute_names,)*].into_iter())
107            }
108        }
109    )
110    .into()
111}
112
113#[derive(Clone, Copy, PartialEq, Eq, Hash)]
114enum NeedsBound {
115    Octet,
116    Multiple,
117    Both,
118}
119impl BitOr for NeedsBound {
120    type Output = NeedsBound;
121
122    fn bitor(self, rhs: Self) -> Self::Output {
123        match (self, rhs) {
124            (a, b) if a == b => a,
125            _ => NeedsBound::Both,
126        }
127    }
128}
129
130fn insert_object_name(field: &Field) -> TokenStream {
131    let field_name = field.ident.as_ref().expect("checked to be named field");
132    let ty = &field.ty;
133    quote! {
134        #field_name: <#ty as From<String>>::from(entry.object_name)
135    }
136}
137
138struct AttributeField {
139    attribute_name: String,
140    multiple: bool,
141    default: bool,
142    field: Field,
143}
144impl AttributeField {
145    fn ident(&self) -> Ident {
146        self.field.ident.clone().expect("checked to be named field")
147    }
148}
149fn parse_fields(
150    raw_fields: impl IntoIterator<Item = Field>,
151) -> Result<(Vec<AttributeField>, Option<Field>), syn::Error> {
152    let mut fields: Vec<AttributeField> = Vec::new();
153    let mut object_name_field = None;
154    'fields: for field in raw_fields {
155        let mut multiple = false;
156        let mut default = false;
157        let mut replaced_attribute_name = None;
158        for attr in &field.attrs {
159            let mut has_set_object_name_field = false;
160            attr.parse_nested_meta(|meta| {
161                if meta.path.is_ident("object_name") {
162                    if object_name_field.replace(field.clone()).is_some() {
163                        return Err(meta.error("\"object_name\" can only be declared on one field"));
164                    };
165                    has_set_object_name_field = true;
166                    return Ok(());
167                }
168                if meta.path.require_ident()? == "rename" {
169                    let lookahead = meta.input.lookahead1();
170                    if lookahead.peek(syn::Token![=]) {
171                        let expr = meta
172                            .value()
173                            .expect("Meta has no value")
174                            .parse()
175                            .expect("Meta is no expression");
176                        let mut value = &expr;
177                        while let syn::Expr::Group(e) = value {
178                            value = &e.expr;
179                        }
180                        if let syn::Expr::Lit(syn::ExprLit {
181                            lit: syn::Lit::Str(lit),
182                            ..
183                        }) = value
184                        {
185                            replaced_attribute_name = Some(lit.value());
186                        } else {
187                            return Err(meta.error("rename argument must be a string literal"));
188                        }
189                    } else {
190                        return Err(meta.error("rename must be used like \"rename = <LDAP NAME>\""));
191                    }
192                }
193                if meta.path.require_ident()? == "multiple" {
194                    multiple = true;
195                }
196                if meta.path.require_ident()? == "default" {
197                    default = true;
198                }
199                Ok(())
200            })?;
201            if has_set_object_name_field {
202                continue 'fields;
203            }
204        }
205        let attribute_name = replaced_attribute_name
206            .unwrap_or_else(|| field.ident.as_ref().expect("checked as named field").to_string());
207        fields.push(AttributeField {
208            attribute_name,
209            multiple,
210            default,
211            field,
212        })
213    }
214    Ok((fields, object_name_field))
215}
216
217fn field_line(data: &AttributeField) -> TokenStream {
218    let lookup_name = &data.attribute_name;
219    let field_type = &data.field.ty;
220    let varname = format_ident!("{}", data.ident());
221    let fallback = if data.default {
222        quote! { <#field_type as Default>::default() }
223    } else {
224        quote! { return Err(lapdog::search::FailedToGetFromEntry::MissingField(#lookup_name)) }
225    };
226    if data.multiple {
227        quote! {
228            let #varname = match entry.attributes.iter().find(|x| x.r#type == #lookup_name) {
229                Some(attrs) => <#field_type as lapdog::search::FromMultipleOctetStrings>::from_multiple_octet_strings(attrs.values.iter().map(|x| x.as_ref()))
230                    .map_err(|b| lapdog::search::FailedToGetFromEntry::FailedToParseField(#lookup_name, Box::new(b)))?,
231                None => {#fallback},
232            };
233        }
234    } else {
235        quote! {
236            let #varname = match entry.attributes.iter().find(|x| x.r#type == #lookup_name).map(|x| x.values.as_slice()) {
237                Some([attr]) => <#field_type as lapdog::search::FromOctetString>::from_octet_string(attr).map_err(|b| lapdog::search::FailedToGetFromEntry::FailedToParseField(#lookup_name, Box::new(b)))?,
238                Some([]) | None => {#fallback},
239                Some(_) => {return Err(lapdog::search::FailedToGetFromEntry::TooManyValues(#lookup_name))}
240            };
241        }
242    }
243}