Skip to main content

bauer_macros/
lib.rs

1#![doc = include_str!("../README.md")]
2
3use proc_macro::TokenStream;
4use proc_macro2::TokenStream as TokenStream2;
5use quote::{ToTokens, format_ident, quote, quote_spanned};
6use syn::{
7    DeriveInput, Ident, Pat, parse::ParseStream, parse_macro_input, parse_quote_spanned,
8    spanned::Spanned,
9};
10use util::stricter_visibility;
11
12use crate::{
13    attr::builder::{BuilderAttr, Kind},
14    attr::field::{BuilderField, Len, Repeat, WrappedType},
15    util::parallel_assign,
16};
17
18mod attr;
19mod type_state;
20mod util;
21
22/// A very minimal builder implementation where everything panics
23fn failed_builder(
24    mut builder_attr: BuilderAttr,
25    input: &DeriveInput,
26    fields: Vec<BuilderField>,
27    errors: &[syn::Error],
28) -> TokenStream2 {
29    assert!(!errors.is_empty());
30
31    let is_type_state = builder_attr.kind == Kind::TypeState;
32    // Some of the simpler functions require that the kind is not type-state and since we're
33    // failing, it doesn't really matter.
34    builder_attr.kind = Kind::Owned;
35
36    let ident = &input.ident;
37    let assert_crate = builder_attr.assert_crate();
38    let builder_attributes = &builder_attr.attributes;
39    let builder_vis = &builder_attr.vis;
40    let builder = format_ident!("{}Builder", ident);
41    let build_err = builder_attr.error.name(ident);
42    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
43
44    let konst = builder_attr.konst_kw();
45    let self_param = builder_attr.self_param();
46
47    let functions: TokenStream2 = fields
48        .iter()
49        .filter(|f| !f.should_skip())
50        .map(|f| f.fail_fn(&builder_attr))
51        .collect();
52
53    let (build_err_variants, _) = gen_error_enum(&fields);
54
55    let infallible = (is_type_state || build_err_variants.is_empty()) && !builder_attr.error.force;
56
57    let error_vis =
58        stricter_visibility(builder_attr.build_fn.vis(&builder_attr), &builder_attr.vis);
59
60    let build_err_enum = if infallible {
61        quote! {}
62    } else {
63        let attributes = &builder_attr.error.attributes;
64        quote! {
65            #(#attributes)*
66            #[derive(::std::fmt::Debug, ::std::cmp::PartialEq, ::std::cmp::Eq)]
67            #error_vis enum #build_err {}
68
69            impl ::core::fmt::Display for #build_err {
70                fn fmt(&self, f: &mut ::core::fmt::Formatter<'_>) -> ::core::fmt::Result {
71                    panic!("Invalid Builder");
72                }
73            }
74
75            impl ::core::error::Error for #build_err {}
76        }
77    };
78
79    let ret_ty = if infallible {
80        quote! { #ident #ty_generics }
81    } else {
82        quote! { ::core::result::Result<#ident #ty_generics, #build_err> }
83    };
84
85    let build_fn = {
86        let attributes = &builder_attr.build_fn.attributes;
87        let name = &builder_attr.build_fn.name;
88        let vis = builder_attr.build_fn.vis(&builder_attr);
89        quote! {
90            #(#attributes)*
91            #vis #konst fn #name(#self_param) -> #ret_ty {
92                panic!("Invalid Builder")
93            }
94        }
95    };
96
97    let into_impl = into_impl(
98        &builder_attr,
99        input,
100        &builder,
101        (!infallible).then_some(build_err),
102    );
103
104    let builder_fn = builder_fn(input, &builder_attr, &builder, &[]);
105
106    let errors = errors.iter().map(syn::Error::to_compile_error);
107
108    quote! {
109        #assert_crate
110
111        #build_err_enum
112
113        #(#builder_attributes)*
114        #[must_use = "The builder doesn't construct its type until `.build()` is called"]
115        #builder_vis struct #builder #impl_generics #where_clause {}
116
117        impl #impl_generics #builder #ty_generics #where_clause {
118            #functions
119
120            #build_fn
121        }
122
123        impl #impl_generics #builder #ty_generics #where_clause {
124            #konst fn new() -> Self {
125                panic!("Invalid Builder")
126            }
127        }
128
129        impl #impl_generics ::core::default::Default for #builder #ty_generics #where_clause {
130            fn default() -> Self {
131                Self::new()
132            }
133        }
134
135        #builder_fn
136
137        #into_impl
138
139        #(#errors)*
140    }
141}
142
143fn into_impl(
144    builder_attr: &BuilderAttr,
145    input: &DeriveInput,
146    builder: &Ident,
147    error: Option<impl ToTokens>,
148) -> TokenStream2 {
149    let ident = &input.ident;
150    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
151    let build_fn_name = &builder_attr.build_fn.name;
152
153    if let Some(build_err) = error {
154        quote! {
155            #[allow(clippy::infallible_try_from)]
156            impl #impl_generics ::core::convert::TryFrom<#builder #ty_generics> for #ident #ty_generics #where_clause {
157                type Error = #build_err;
158
159                fn try_from(mut builder: #builder #ty_generics) -> Result<Self, Self::Error> {
160                    builder.#build_fn_name()
161                }
162            }
163        }
164    } else {
165        quote! {
166            impl #impl_generics ::core::convert::From<#builder #ty_generics> for #ident #ty_generics #where_clause {
167                fn from(mut builder: #builder #ty_generics) -> Self {
168                    builder.#build_fn_name()
169                }
170            }
171        }
172    }
173}
174
175fn builder_args(
176    fields: &[BuilderField],
177) -> (
178    Vec<&Ident>,       // field name
179    Vec<TokenStream2>, // function arguments
180    Vec<TokenStream2>, // expanded value
181) {
182    fields
183        .iter()
184        .filter(|f| f.is_associated())
185        .map(|f| {
186            let name = f.arg_name();
187            let (args, value) = f.attr.to_args_and_value(&f.ty, name);
188            (name, args, value)
189        })
190        .collect()
191}
192
193fn builder_fn(
194    input: &DeriveInput,
195    builder_attr: &BuilderAttr,
196    builder: &Ident,
197    fields: &[BuilderField],
198) -> TokenStream2 {
199    let ident = &input.ident;
200    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
201    let konst = builder_attr.konst_kw();
202
203    let name = &builder_attr.builder_fn.name;
204    let attributes = &builder_attr.builder_fn.attributes;
205    let vis = builder_attr.builder_fn.vis(builder_attr);
206
207    let (associated_names, arguments, _) = builder_args(fields);
208
209    quote! {
210        impl #impl_generics #ident #ty_generics #where_clause {
211            #(#attributes)*
212            #vis #konst fn #name(#(#arguments),*) -> #builder #ty_generics {
213                #builder::new(#(#associated_names),*)
214            }
215        }
216    }
217}
218
219fn parse_build_attr(input: &DeriveInput, errors: &mut Vec<syn::Error>) -> BuilderAttr {
220    let mut out = BuilderAttr::new(input.vis.clone());
221    for attr in input.attrs.iter().filter(|a| a.path().is_ident("builder")) {
222        if let Err(e) = attr.parse_args_with(|ps: ParseStream| out.parse(ps)) {
223            errors.push(e);
224        }
225    }
226    out
227}
228
229fn gen_error_enum(fields: &[BuilderField]) -> (Vec<TokenStream2>, Vec<TokenStream2>) {
230    fields
231        .iter()
232        .filter(|f| !f.should_skip())
233        .flat_map(|f| {
234            let mut variants = Vec::new();
235            if let Some(err) = &f.missing_err {
236                let msg = format!("Missing required field '{}'", f.ident);
237                variants.push((
238                    err.to_token_stream(),
239                    quote! { Self::#err => write!(f, #msg) },
240                ));
241            }
242
243            if let WrappedType::Repeat(
244                _,
245                Repeat {
246                    len: Len::Raw { pattern, error },
247                    ..
248                },
249            ) = &f.wrapped_ty
250            {
251                let error_msg = format!(
252                    "Invalid number of repeat arguments provided.  Expected {}, got {{}}",
253                    pattern.to_token_stream()
254                );
255                variants.push((
256                    quote! {
257                        #error(usize)
258                    },
259                    quote! {
260                        Self::#error(n) => write!(f, #error_msg, n)
261                    },
262                ));
263            }
264
265            variants.into_iter()
266        })
267        .collect()
268}
269
270#[proc_macro_derive(Builder, attributes(builder))]
271pub fn builder(input: TokenStream) -> TokenStream {
272    let input = parse_macro_input!(input as DeriveInput);
273    let ident = &input.ident;
274
275    let mut errors = Vec::new();
276
277    let builder_attr: BuilderAttr = parse_build_attr(&input, &mut errors);
278
279    let data_struct = match input.data {
280        syn::Data::Struct(ref data_struct) => data_struct,
281        syn::Data::Enum(data_enum) => {
282            return syn::Error::new(data_enum.enum_token.span(), "Enums are not supported.")
283                .to_compile_error()
284                .into();
285        }
286        syn::Data::Union(data_union) => {
287            return syn::Error::new(data_union.union_token.span(), "Unions are not supported.")
288                .to_compile_error()
289                .into();
290        }
291    };
292
293    let self_param = builder_attr.self_param();
294    let builder_vis = &builder_attr.vis;
295
296    let builder = format_ident!("{}Builder", ident);
297    let build_err = builder_attr.error.name(ident);
298    let inner = format_ident!("__unsafe_builder_content");
299
300    let mut tuple_index = 0;
301    let fields = match data_struct.fields {
302        syn::Fields::Named(ref fields_named) => {
303            //
304            fields_named
305                .named
306                .iter()
307                .map(|f| {
308                    BuilderField::parse(f, &builder_attr, ident, &mut tuple_index, &mut errors)
309                })
310                .collect::<Vec<_>>()
311        }
312        syn::Fields::Unnamed(_) => {
313            return syn::Error::new(ident.span(), "Unnamed fields are not supported.")
314                .to_compile_error()
315                .into();
316        }
317        syn::Fields::Unit => {
318            return syn::Error::new(ident.span(), "Unit structs are not supported.")
319                .to_compile_error()
320                .into();
321        }
322    };
323
324    let private_module = builder_attr.private_module();
325
326    if !errors.is_empty() {
327        return failed_builder(builder_attr, &input, fields, &errors).into();
328    }
329
330    if builder_attr.kind == Kind::TypeState {
331        return type_state::type_state_builder(&builder_attr, &input, fields).into();
332    }
333
334    let (field_types, init): (Vec<_>, Vec<_>) = fields
335        .iter()
336        .filter(|f| !f.should_skip())
337        .map(|f| {
338            if f.is_associated() {
339                let (_, value) = f.attr.to_args_and_value(&f.ty, f.arg_name());
340                return (f.ty.to_token_stream(), value);
341            }
342
343            match &f.wrapped_ty {
344                WrappedType::None => {
345                    let ty = &f.ty;
346                    (
347                        quote! { ::core::option::Option<#ty> },
348                        quote! { ::core::option::Option::None },
349                    )
350                }
351                WrappedType::Flag => (quote! { bool }, quote! { false }),
352                WrappedType::Option(ty) => (
353                    quote! { ::core::option::Option<#ty> },
354                    quote! { ::core::option::Option::None },
355                ),
356                WrappedType::Repeat(
357                    ty,
358                    Repeat {
359                        array: true, len, ..
360                    },
361                ) => {
362                    let pattern = match &len {
363                        Len::Raw { pattern, .. } => pattern.to_token_stream(),
364                        Len::Int { len } => len.to_token_stream(),
365                        _ => {
366                            unreachable!("If array, then Len::Raw set");
367                        }
368                    };
369                    (
370                        quote! { #private_module::PushableArray<#pattern, #ty> },
371                        quote! { #private_module::PushableArray::new() },
372                    )
373                }
374                WrappedType::Repeat(inner_ty, Repeat { array: false, .. }) => (
375                    quote! { ::std::vec::Vec<#inner_ty> },
376                    quote! { ::std::vec::Vec::new() },
377                ),
378            }
379        })
380        .collect();
381
382    let functions: TokenStream2 = fields
383        .iter()
384        .filter(|f| !f.should_skip() && !f.is_associated())
385        .map(|f| f.function(&builder_attr, &inner))
386        .collect();
387
388    let (build_err_variants, build_err_messages) = gen_error_enum(&fields);
389
390    let not_skipped_field_values = fields.iter().filter(|f| !f.should_skip()).map(|field| {
391        let name = &field.ident;
392        let wrapped_ty = &field.ty;
393        let field_i = field.tuple_index();
394
395        let value = if field.is_associated() {
396            let clone = if builder_attr.kind == Kind::Borrowed {
397                quote! { .clone() }
398            } else {
399                quote! {}
400            };
401
402            quote! {
403                inner.#field_i #clone
404            }
405        } else if !field.wrapped_ty.is_none() {
406            match &field.wrapped_ty {
407                WrappedType::None => unreachable!("Checked in if branch"),
408                WrappedType::Flag => quote! { inner.#field_i },
409                WrappedType::Option(_) => quote! { inner.#field_i.take() },
410                WrappedType::Repeat(inner_ty, rep @ Repeat { collector, .. }) => {
411                    if let Len::Raw { pattern, error } = &rep.len {
412                        let value = if rep.array {
413                            quote_spanned! { inner_ty.span()=> {
414                                let arr = ::core::mem::replace(&mut inner.#field_i, #private_module::PushableArray::new());
415                                arr.into_array()
416                                    .expect("The match ensures the length of this array is correct")
417                            }}
418                        } else {
419                            assert!(!rep.array);
420                            assert!(!builder_attr.konst);
421
422                            collector.collect(parse_quote_spanned! {inner_ty.span()=>
423                                inner.#field_i.drain(..)
424                            })
425                        };
426
427                        if let Pat::Ident(_) = pattern {
428                            quote_spanned! { pattern.span()=>
429                                if inner.#field_i.len() == #pattern {
430                                    #value
431                                } else {
432                                    return Err(#build_err::#error(self.#inner.#field_i.len()));
433                                }
434                            }
435                        } else {
436                            quote_spanned! { pattern.span()=>
437                                match inner.#field_i.len() {
438                                    #pattern => #value,
439                                    len => return Err(#build_err::#error(len)),
440                                }
441                            }
442                        }
443                    } else {
444                        assert!(!rep.array);
445                        assert!(!builder_attr.konst);
446                        collector.collect(parse_quote_spanned! {inner_ty.span()=>
447                            inner.#field_i.drain(..)
448                        })
449                    }
450                },
451            }
452        } else if let Some(default) = &field.attr.default {
453            let default = default.to_value(field.attr.into);
454            quote! {
455                // NOTE: not using Option::unwrap_or_else, since it's not stable in const
456                match inner.#field_i.take() {
457                    Some(v) => v,
458                    None => #default
459                }
460            }
461        } else {
462            let err = field
463                .missing_err
464                .as_ref()
465                .expect("missing_err is set when default is none");
466            quote! {
467                // NOTE: not using Option::ok_or, since it's not stable in const
468                match inner.#field_i.take() {
469                    Some(v) => v,
470                    None => return Err(#build_err::#err),
471                }
472            }
473        };
474
475        quote! {{
476            let #name: #wrapped_ty = #value;
477            #name
478        }}
479    });
480
481    let not_skipped_fields: Vec<_> = fields
482        .iter()
483        .filter(|f| !f.should_skip())
484        .map(|f| &f.ident)
485        .collect();
486
487    let set_not_skipped_fields = parallel_assign(
488        not_skipped_fields.iter().copied(),
489        not_skipped_field_values,
490        if builder_attr.kind == Kind::Borrowed {
491            quote! {
492                let inner = &mut self.#inner;
493            }
494        } else {
495            quote! {
496                let mut inner = self.#inner;
497            }
498        },
499    );
500
501    let set_skipped_fields = parallel_assign(
502        fields.iter().filter(|f| f.should_skip()).map(|f| &f.ident),
503        fields.iter().filter_map(BuilderField::skipped_field_value),
504        quote! {
505            #[allow(unused)]
506            let (#(#not_skipped_fields),*) = (#(&#not_skipped_fields),*);
507        },
508    );
509
510    let finish_fields = fields.iter().map(|field| &field.ident);
511
512    let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
513
514    let konst = builder_attr.konst_kw();
515
516    let (mut ret_ty, mut ret_val) = (quote! { #ident #ty_generics }, quote! { ret });
517
518    if let Some((param, ty, body)) = &builder_attr.build_fn.mapper {
519        ret_val = quote! {{
520            #konst fn __private_mapper(#param: #ret_ty) -> #ty {
521                #[allow(unused_braces)]
522                #body
523            }
524            __private_mapper(#ret_val)
525        }};
526        ret_ty = ty.to_token_stream();
527    }
528
529    if !build_err_variants.is_empty() || builder_attr.error.force {
530        ret_ty = quote! { ::core::result::Result<#ret_ty, #build_err> };
531        ret_val = quote! { Ok(#ret_val) };
532    }
533
534    let build_fn = {
535        let attributes = &builder_attr.build_fn.attributes;
536        let name = &builder_attr.build_fn.name;
537        let vis = builder_attr.build_fn.vis(&builder_attr);
538        quote! {
539            #(#attributes)*
540            #vis #konst fn #name(#self_param) -> #ret_ty {
541                #[allow(deprecated)] // #inner is set to deprecated
542                let ret = {
543                    #set_not_skipped_fields
544                    #set_skipped_fields
545
546                    #ident {
547                        #(#finish_fields),*
548                    }
549                };
550                #ret_val
551            }
552        }
553    };
554
555    let build_err_enum = if build_err_variants.is_empty() && !builder_attr.error.force {
556        quote! {}
557    } else {
558        let attributes = &builder_attr.error.attributes;
559
560        let error_vis =
561            stricter_visibility(builder_attr.build_fn.vis(&builder_attr), &builder_attr.vis);
562
563        quote! {
564            #(#attributes)*
565            #[derive(::std::fmt::Debug, ::std::cmp::PartialEq, ::std::cmp::Eq)]
566            #[allow(enum_variant_names)]
567            #error_vis enum #build_err {
568                #(#build_err_variants),*
569            }
570
571            impl ::core::fmt::Display for #build_err {
572                fn fmt(&self, f: &mut ::core::fmt::Formatter<'_>) -> ::core::fmt::Result {
573                    use ::core::fmt::Write;
574                    match *self {
575                        #(#build_err_messages),*
576                    }
577                }
578            }
579
580            impl ::core::error::Error for #build_err {}
581        }
582    };
583
584    let into_impl = if builder_attr.build_fn.mapper.is_some() {
585        quote! {} // can't do into if the build_fn is a mapper
586    } else {
587        into_impl(
588            &builder_attr,
589            &input,
590            &builder,
591            (!build_err_variants.is_empty() || builder_attr.error.force).then_some(build_err),
592        )
593    };
594
595    let builder_attributes = &builder_attr.attributes;
596
597    let builder_fn = builder_fn(&input, &builder_attr, &builder, &fields);
598
599    let new_fn = {
600        let (_, arguments, _) = builder_args(&fields);
601        quote! {
602            impl #impl_generics #builder #ty_generics #where_clause {
603                #konst fn new(#(#arguments),*) -> Self {
604                    Self {
605                        #inner: (#(#init,)*),
606                    }
607                }
608            }
609        }
610    };
611    let default_fn = {
612        fields.iter().all(|f| !f.is_associated()).then(|| quote! {
613            impl #impl_generics ::core::default::Default for #builder #ty_generics #where_clause {
614                fn default() -> Self {
615                    Self::new()
616                }
617            }
618        })
619    };
620
621    let assert_crate = builder_attr.assert_crate();
622    quote! {
623        #assert_crate
624
625        #build_err_enum
626
627        #(#builder_attributes)*
628        #[must_use = "The builder doesn't construct its type until `.build()` is called"]
629        #builder_vis struct #builder #impl_generics #where_clause {
630            #[deprecated = "This field is for internal use only; You almost certainly don't need to touch this. If you encounter a bug or missing feature, file an issue on the repo."]
631            #[doc(hidden)]
632            #inner: (#(#field_types,)*),
633        }
634
635        impl #impl_generics #builder #ty_generics #where_clause {
636            #functions
637
638            #build_fn
639        }
640
641        #new_fn
642        #default_fn
643
644        #builder_fn
645
646        #into_impl
647    }
648    .into()
649}