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
22fn 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 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>, Vec<TokenStream2>, Vec<TokenStream2>, ) {
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 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 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 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)] 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! {} } 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}