mod error;
use error::BuilderError::*;
use proc_macro2::TokenStream;
use quote::{quote, ToTokens};
use syn::spanned::Spanned;
use syn::*;
#[rustfmt::skip]
fn is_option(ty: &Type) -> Option<Type> {
if let Type::Path(TypePath { path: Path { segments, .. }, .. }) = ty {
if segments[0].ident == "Option" {
return match &segments[0].arguments {
PathArguments::None => None,
PathArguments::Parenthesized(_) => None,
PathArguments::AngleBracketed(AngleBracketedGenericArguments { args, .. }) => {
if let GenericArgument::Type(inner_ty) = &args[0] {
Some(inner_ty.clone())
} else {
None
}
}
};
}
}
None
}
#[derive(Clone)]
enum GenericParamName {
Type(Ident),
Lifetime(Lifetime),
Const(Ident),
}
impl ToTokens for GenericParamName {
fn to_tokens(&self, tokens: &mut TokenStream) {
match self {
GenericParamName::Type(ty) => ty.to_tokens(tokens),
GenericParamName::Lifetime(lt) => lt.to_tokens(tokens),
GenericParamName::Const(ct) => ct.to_tokens(tokens),
}
}
}
fn param_to_name(generics: &Generics) -> Vec<GenericParamName> {
generics
.params
.iter()
.map(|param| match param {
GenericParam::Type(ty) => GenericParamName::Type(ty.ident.clone()),
GenericParam::Lifetime(lt) => GenericParamName::Lifetime(lt.lifetime.clone()),
GenericParam::Const(c) => GenericParamName::Const(c.ident.clone()),
})
.collect()
}
fn split_param_names(
param_names: Vec<GenericParamName>,
) -> (
Vec<GenericParamName>, // Lifetime generic parameters
Vec<GenericParamName>, // Const generic parameters
Vec<GenericParamName>, // Type generic parameters
) {
let mut lifetimes = vec![];
let mut consts = vec![];
let mut types = vec![];
for param_name in param_names {
match param_name {
GenericParamName::Lifetime(_) => lifetimes.push(param_name.clone()),
GenericParamName::Const(_) => consts.push(param_name.clone()),
GenericParamName::Type(_) => types.push(param_name.clone()),
}
}
(lifetimes, consts, types)
}
fn split_params(
params: Vec<GenericParam>,
) -> (
Vec<GenericParam>, // Lifetime generic parameters
Vec<GenericParam>, // Const generic parameters
Vec<GenericParam>, // Type generic parameters
) {
let mut lifetimes = vec![];
let mut consts = vec![];
let mut types = vec![];
for param in params {
match param {
GenericParam::Lifetime(_) => lifetimes.push(param.clone()),
GenericParam::Const(_) => consts.push(param.clone()),
GenericParam::Type(_) => types.push(param.clone()),
}
}
(lifetimes, consts, types)
}
#[proc_macro_derive(Builder)]
pub fn builder(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
let ast = parse_macro_input!(input as DeriveInput);
match ast.data {
Data::Struct(struct_t) => match struct_t.fields {
Fields::Named(FieldsNamed { named, .. }) => {
let fields = named;
let struct_ident = ast.ident.clone();
let (impl_generics, ty_generics, where_clause) = ast.generics.split_for_impl();
let builder_ident =
Ident::new(&format!("{struct_ident}Builder"), struct_ident.span());
let st_param_names = param_to_name(&ast.generics);
let (st_lt_pn, st_ct_pn, st_ty_pn) = split_param_names(st_param_names);
let st_params: Vec<_> = ast.generics.params.iter().cloned().collect();
let (st_lt_p, st_ct_p, st_ty_p) = split_params(st_params);
let (optional_fields, required_fields): (Vec<_>, Vec<_>) = fields
.iter()
.partition(|field| is_option(&field.ty).is_some());
let mut all_false = vec![];
let mut all_true = vec![];
let mut b_ct_pn = vec![];
let mut b_ct_p = vec![];
let mut b_fields = vec![];
let mut b_inits = vec![];
let mut req_moves = vec![];
let mut req_unwraps = vec![];
for (index, field) in required_fields.iter().enumerate() {
let field_ident = &field.ident;
let field_ty = &field.ty;
let ct_param_ident = Ident::new(&format!("P{}", index), field.span());
b_fields.push(quote! { #field_ident: ::std::option::Option<#field_ty> });
b_inits.push(quote! { #field_ident: None });
req_moves.push(quote! { #field_ident: self.#field_ident });
req_unwraps.push(quote! { #field_ident: self.#field_ident.unwrap_unchecked() });
all_false.push(quote! { false });
all_true.push(quote! { true });
b_ct_pn.push(quote! { #ct_param_ident });
b_ct_p.push(quote! { const #ct_param_ident: bool });
}
let mut opt_moves = vec![];
for opt_field in &optional_fields {
let field_ident = &opt_field.ident;
let field_ty = &opt_field.ty;
opt_moves.push(quote! { #field_ident: self.#field_ident });
b_fields.push(quote! { #field_ident: #field_ty });
b_inits.push(quote! { #field_ident: None });
}
let mut opt_setters = vec![];
for opt_field in &optional_fields {
let field_ident = &opt_field.ident;
let field_ty = &opt_field.ty;
let inner_ty = is_option(field_ty).unwrap();
opt_setters.push(
quote! {
pub fn #field_ident(mut self, #field_ident: #inner_ty) ->
#builder_ident<#(#st_lt_pn,)* #(#st_ct_pn,)* #(#b_ct_pn,)* #(#st_ty_pn,)*>
{
self.#field_ident = Some(#field_ident);
self
}
}
);
}
let mut req_setters = vec![];
for (index, req_field) in required_fields.iter().enumerate() {
let field_ident = &req_field.ident;
let field_ty = &req_field.ty;
let before_req_moves = &req_moves[..index];
let after_req_moves = &req_moves[index + 1..];
let before_pn = &b_ct_pn[..index];
let after_pn = &b_ct_pn[index + 1..];
req_setters.push(
quote! {
pub fn #field_ident(self, #field_ident: #field_ty) ->
#builder_ident<#(#st_lt_pn,)* #(#st_ct_pn,)* #(#before_pn,)* true, #(#after_pn,)* #(#st_ty_pn,)*>
{
#builder_ident {
#(#before_req_moves,)*
#field_ident: Some(#field_ident),
#(#after_req_moves,)*
#(#opt_moves,)*
}
}
}
);
}
quote! {
pub struct #builder_ident<#(#st_lt_p,)* #(#st_ct_p,)* #(#b_ct_p,)* #(#st_ty_p,)*> #where_clause {
#(#b_fields),*
}
impl #impl_generics #struct_ident #ty_generics #where_clause {
pub fn builder() -> #builder_ident<#(#st_lt_pn,)* #(#st_ct_pn,)* #(#all_false,)* #(#st_ty_pn,)*> {
#builder_ident {
#(#b_inits),*
}
}
}
impl<#(#st_lt_p,)* #(#st_ct_p,)* #(#b_ct_p,)* #(#st_ty_p,)*>
#builder_ident<#(#st_lt_pn,)* #(#st_ct_pn,)* #(#b_ct_pn,)* #(#st_ty_pn,)* >
#where_clause
{
#(#opt_setters)*
#(#req_setters)*
}
impl<#(#st_lt_p,)* #(#st_ct_p,)* #(#st_ty_p,)*>
#builder_ident<#(#st_lt_pn,)* #(#st_ct_pn,)* #(#all_true,)* #(#st_ty_pn,)* >
#where_clause
{
fn build(self) -> #struct_ident #ty_generics {
unsafe {
#struct_ident {
#(#opt_moves,)*
#(#req_unwraps,)*
}
}
}
}
}
.into()
}
Fields::Unnamed(_) => UnnamedFields(struct_t.fields).into(),
Fields::Unit => UnitStruct(struct_t.fields).into(),
},
Data::Enum(enum_t) => Enum(enum_t).into(),
Data::Union(union_t) => Union(union_t).into(),
}
}