method_chaining 0.1.1

A Rust procedural macro that automatically makes functions and structs chainable.
Documentation
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();
        }
    };

    // 提取 Vec<T> 中的 T 类型
    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();
        }
    };

    // 提取 HashMap<K, V> 中的 K 和 V 类型
    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)
}