use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, Fields, GenericArgument, Ident, ItemStruct, PathArguments, Type};
pub fn make_chain_push(attr: TokenStream, item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as ItemStruct);
let field_name = parse_macro_input!(attr as Ident);
let struct_name = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let field = if let Fields::Named(fields) = &input.fields {
fields
.named
.iter()
.find(|f| f.ident.as_ref() == Some(&field_name))
} else {
return syn::Error::new_spanned(
input,
"chain_push can only be used on structs with named fields",
)
.to_compile_error()
.into();
};
let field = match field {
Some(f) => f,
None => {
return syn::Error::new_spanned(
field_name.clone(),
format!("Field '{}' not found in struct", field_name),
)
.to_compile_error()
.into();
}
};
let inner_type = extract_vec_inner_type(&field.ty);
let inner_type = match inner_type {
Some(t) => t,
None => {
return syn::Error::new_spanned(
&field.ty,
"chain_push can only be used on Vec<T> fields",
)
.to_compile_error()
.into();
}
};
let method_name = Ident::new(&format!("push_{}", field_name), field_name.span());
let expanded = quote! {
#input
impl #impl_generics #struct_name #ty_generics #where_clause {
pub fn #method_name(mut self, value: #inner_type) -> Self {
self.#field_name.push(value);
self
}
}
};
TokenStream::from(expanded)
}
pub fn make_chain_insert(attr: TokenStream, item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as ItemStruct);
let field_name = parse_macro_input!(attr as Ident);
let struct_name = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let field = if let Fields::Named(fields) = &input.fields {
fields
.named
.iter()
.find(|f| f.ident.as_ref() == Some(&field_name))
} else {
return syn::Error::new_spanned(
input,
"chain_insert can only be used on structs with named fields",
)
.to_compile_error()
.into();
};
let field = match field {
Some(f) => f,
None => {
return syn::Error::new_spanned(
field_name.clone(),
format!("Field '{}' not found in struct", field_name),
)
.to_compile_error()
.into();
}
};
let (key_type, value_type) = extract_hashmap_types(&field.ty);
let (key_type, value_type) = match (key_type, value_type) {
(Some(k), Some(v)) => (k, v),
_ => {
return syn::Error::new_spanned(
&field.ty,
"chain_insert can only be used on HashMap<K, V> fields",
)
.to_compile_error()
.into();
}
};
let method_name = Ident::new(&format!("insert_{}", field_name), field_name.span());
let expanded = quote! {
#input
impl #impl_generics #struct_name #ty_generics #where_clause {
pub fn #method_name(mut self, key: #key_type, value: #value_type) -> Self {
self.#field_name.insert(key, value);
self
}
}
};
TokenStream::from(expanded)
}
fn extract_vec_inner_type(ty: &Type) -> Option<&Type> {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "Vec" {
if let PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(GenericArgument::Type(inner_type)) = args.args.first() {
return Some(inner_type);
}
}
}
}
}
None
}
fn extract_hashmap_types(ty: &Type) -> (Option<&Type>, Option<&Type>) {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "HashMap" {
if let PathArguments::AngleBracketed(args) = &segment.arguments {
if args.args.len() == 2 {
let key_type = if let Some(GenericArgument::Type(t)) = args.args.first() {
Some(t)
} else {
None
};
let value_type = if let Some(GenericArgument::Type(t)) = args.args.get(1) {
Some(t)
} else {
None
};
return (key_type, value_type);
}
}
}
}
}
(None, None)
}