diff options
Diffstat (limited to 'mingling_macros/src/attr/chain.rs')
| -rw-r--r-- | mingling_macros/src/attr/chain.rs | 378 |
1 files changed, 378 insertions, 0 deletions
diff --git a/mingling_macros/src/attr/chain.rs b/mingling_macros/src/attr/chain.rs new file mode 100644 index 0000000..dbd87dd --- /dev/null +++ b/mingling_macros/src/attr/chain.rs @@ -0,0 +1,378 @@ +#![allow(clippy::too_many_arguments)] + +use crate::res_injection::{ + ResourceInjection, extract_args_info, generate_immut_resource_bindings, + wrap_body_with_mut_resources, wrap_body_with_mut_resources_async, +}; +use proc_macro::TokenStream; +use quote::{ToTokens, quote}; +use syn::spanned::Spanned; +use syn::{Ident, ItemFn, Pat, ReturnType, Signature, Type, TypePath, parse_macro_input}; + +/// Checks whether the return type is `()` +fn is_unit_return_type(sig: &Signature) -> bool { + match &sig.output { + ReturnType::Type(_, ty) => match &**ty { + Type::Tuple(tuple) => tuple.elems.is_empty(), + _ => false, + }, + ReturnType::Default => true, + } +} + +/// Validates that the return type is acceptable. +/// Accepts `()`, `Next`, `ChainProcess<...>`, or any type that can +/// be converted to `ChainProcess` via `.into()` (i.e. any pack type). +fn validate_return_type(sig: &Signature) -> Result<(), proc_macro2::TokenStream> { + // `()` or omitted is always valid + if is_unit_return_type(sig) { + return Ok(()); + } + + Ok(()) +} + +/// Builds the `proc` function implementation inside the generated `Chain` impl. +/// +/// Instead of inlining the user's body, the trait method calls the original +/// function by name, with resources injected from the application context. +fn generate_proc_fn( + fn_name: &Ident, + has_resources: bool, + resources: &[ResourceInjection], + program_type: &proc_macro2::TokenStream, + previous_type: &TypePath, + is_async_fn: bool, + is_unit_return: bool, + origin_return_type: &proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + let immut_resource_stmts = generate_immut_resource_bindings(resources.iter(), program_type); + let mut_resources: Vec<_> = resources.iter().filter(|r| r.is_mut).collect(); + + // Use a fixed parameter name `prev` for the trait method, regardless of + // the user's original parameter name (which may be `_` and cannot be + // referenced in expression position). + let fixed_prev: Pat = syn::parse_quote!(prev); + + // Build the call to the original function with resource arguments injected. + // The variable names come from the resource injection bindings + // (immut bindings are `let #name = …`, mut closures receive `|#name: &mut T|`), + // which match the original function's parameter names. + let resource_args: Vec<_> = resources + .iter() + .map(|res| { + let var_name = &res.var_name; + quote! { #var_name } + }) + .collect(); + + let fn_call = if has_resources { + let call = quote! { #fn_name(#fixed_prev, #(#resource_args),*) }; + if is_async_fn { + quote! { #call.await } + } else { + call + } + } else { + let call = quote! { #fn_name(#fixed_prev) }; + if is_async_fn { + quote! { #call.await } + } else { + call + } + }; + + // Convert the function call to a syn::Stmt so existing wrapping functions can use it + let fn_call_expr: syn::Expr = syn::parse_quote! { #fn_call }; + let fn_call_stmt = syn::Stmt::Expr(fn_call_expr, None); + + let wrapped_body = if is_async_fn && !mut_resources.is_empty() { + wrap_body_with_mut_resources_async(&[fn_call_stmt], &mut_resources, program_type) + } else { + wrap_body_with_mut_resources( + &[fn_call_stmt], + &mut_resources, + program_type, + is_unit_return, + ) + }; + + let proc_body = if is_unit_return { + let body_with_ending = if has_resources { + quote! { + #(#immut_resource_stmts)* + #wrapped_body; + <crate::ResultEmpty as ::mingling::Routable::<crate::ThisProgram>> + ::to_chain(crate::ResultEmpty) + } + } else { + quote! { + #wrapped_body; + <crate::ResultEmpty as ::mingling::Routable::<crate::ThisProgram>> + ::to_chain(crate::ResultEmpty) + } + }; + quote! { #body_with_ending } + } else { + let body = if has_resources { + quote! { + #(#immut_resource_stmts)* + #wrapped_body + } + } else { + quote! { #wrapped_body } + }; + quote! { + let __chain_result = { #body }; + <#origin_return_type as ::std::convert::Into< + ::mingling::ChainProcess<#program_type> + >>::into(__chain_result) + } + }; + + #[cfg(feature = "async")] + { + quote! { + async fn proc(#fixed_prev: #previous_type) -> ::mingling::ChainProcess<#program_type> { + #proc_body + } + } + } + + #[cfg(not(feature = "async"))] + { + quote! { + fn proc(#fixed_prev: #previous_type) -> ::mingling::ChainProcess<#program_type> { + #proc_body + } + } + } +} + +/// Assembles the final expanded output: hidden struct, `register_chain!` invocation, +/// `Chain` impl with the `proc` method, and the preserved original function. +fn generate_struct_and_impl( + fn_attrs: &[syn::Attribute], + vis: &syn::Visibility, + struct_name: &Ident, + previous_type: &TypePath, + previous_type_str: &proc_macro2::TokenStream, + program_type: &proc_macro2::TokenStream, + proc_fn: &proc_macro2::TokenStream, + original_fn: &proc_macro2::TokenStream, +) -> proc_macro2::TokenStream { + quote! { + #(#fn_attrs)* + #[doc(hidden)] + #[allow(non_camel_case_types)] + #vis struct #struct_name; + + ::mingling::macros::register_chain!(#previous_type_str, #struct_name); + + impl ::mingling::Chain<#program_type> for #struct_name { + type Previous = #previous_type; + + #proc_fn + } + + // Keep the original function unchanged + #original_fn + } +} + +/// Ensures the function is not async when the `async` feature is disabled. +#[cfg(not(feature = "async"))] +fn reject_async(sig: &Signature) -> Result<(), proc_macro2::TokenStream> { + if sig.asyncness.is_some() { + return Err(syn::Error::new( + sig.span(), + "Chain function cannot be async when async feature is disabled", + ) + .to_compile_error()); + } + Ok(()) +} + +pub(crate) fn chain_attr(attr: TokenStream, item: TokenStream) -> TokenStream { + // Reject non-empty attribute arguments; #[chain] must be bare + if !attr.is_empty() { + return syn::Error::new( + attr.into_iter().next().unwrap().span().into(), + "#[chain] does not accept arguments", + ) + .to_compile_error() + .into(); + } + + // Parse the function item + let input_fn = parse_macro_input!(item as ItemFn); + + // Handle async feature gate + #[cfg(feature = "async")] + let is_async_fn = input_fn.sig.asyncness.is_some(); + + #[cfg(not(feature = "async"))] + { + if let Err(err) = reject_async(&input_fn.sig) { + return err.into(); + } + } + + // Check if return type is unit + let is_unit_return = is_unit_return_type(&input_fn.sig); + + // Validate return type + if let Err(err) = validate_return_type(&input_fn.sig) { + return err.into(); + } + + // Extract the previous type, parameter name, and resource injection params + let (_, previous_type, resources) = match extract_args_info(&input_fn.sig) { + Ok(info) => info, + Err(e) => return e.to_compile_error().into(), + }; + + // Prepare building blocks + let mut fn_attrs = input_fn.attrs.clone(); + fn_attrs.retain(|attr| !attr.path().is_ident("chain")); + let vis = &input_fn.vis; + let fn_name = &input_fn.sig.ident; + let has_resources = !resources.is_empty(); + + // Generate struct name + let internal_name = format!( + "__internal_chain_{}", + just_fmt::snake_case!(fn_name.to_string()) + ); + let struct_name = Ident::new(&internal_name, fn_name.span()); + + // Always use the default crate-defined program path + let program_type = crate::default_program_path(); + + // Extract the user's return type for the explicit Into turbofish + let origin_return_type = match &input_fn.sig.output { + ReturnType::Type(_, ty) => quote! { #ty }, + ReturnType::Default => quote! { () }, + }; + + // Generate the `proc` function for the Chain impl + let proc_fn = generate_proc_fn( + fn_name, + has_resources, + &resources, + &program_type, + &previous_type, + #[cfg(feature = "async")] + is_async_fn, + #[cfg(not(feature = "async"))] + false, + is_unit_return, + &origin_return_type, + ); + + // Preserve the original function untouched + // Note: do NOT add `#vis` here — `input_fn` (ItemFn) already contains its own visibility. + let original_fn = quote! { + #(#fn_attrs)* + #input_fn + }; + + // Assemble the final output + let previous_type_str = quote! { #previous_type }; + let expanded = generate_struct_and_impl( + &fn_attrs, + vis, + &struct_name, + &previous_type, + &previous_type_str, + &program_type, + &proc_fn, + &original_fn, + ); + + expanded.into() +} + +/// Builds a match arm for chain mapping +pub(crate) fn build_chain_arm( + struct_name: &Ident, + previous_type: &TypePath, +) -> proc_macro2::TokenStream { + let enum_variant = &previous_type.path.segments.last().unwrap().ident; + quote! { + #struct_name => #enum_variant, + } +} + +/// Builds a match arm for chain existence check +pub(crate) fn build_chain_exist_arm(previous_type: &TypePath) -> proc_macro2::TokenStream { + let enum_variant = &previous_type.path.segments.last().unwrap().ident; + quote! { + Self::#enum_variant => true, + } +} + +pub(crate) fn register_chain(input: TokenStream) -> TokenStream { + // Parse the input as a comma-separated list of arguments + let input_parsed = syn::parse_macro_input!(input with syn::punctuated::Punctuated<syn::Expr, syn::Token![,]>::parse_terminated); + + // Check that there are exactly two elements + if input_parsed.len() != 2 { + return syn::Error::new( + input_parsed.span(), + "Expected exactly two comma-separated arguments: `PreviousType, StructName`", + ) + .to_compile_error() + .into(); + } + + // Extract the two elements + let previous_type_expr = &input_parsed[0]; + let struct_name_expr = &input_parsed[1]; + + // Convert expressions to TypePath and Ident + let previous_type = match syn::parse2::<TypePath>(previous_type_expr.to_token_stream()) { + Ok(ty) => ty, + Err(e) => return e.to_compile_error().into(), + }; + + let struct_name = match syn::parse2::<syn::Ident>(struct_name_expr.to_token_stream()) { + Ok(ident) => ident, + Err(e) => return e.to_compile_error().into(), + }; + + // Record the chain mapping: previous_type => struct_name + let chain_entry = build_chain_arm(&struct_name, &previous_type); + + // Record the chain existence check + let chain_exist_entry = build_chain_exist_arm(&previous_type); + + let mut chains = crate::get_global_set(&crate::CHAINS).lock().unwrap(); + let mut chain_exist = crate::get_global_set(&crate::CHAINS_EXIST).lock().unwrap(); + + let chain_entry_str = chain_entry.to_string(); + let chain_exist_entry_str = chain_exist_entry.to_string(); + + // Check for duplicate variant before inserting + let variant_name = previous_type + .path + .segments + .last() + .unwrap() + .ident + .to_string(); + if let Err(err) = crate::check_duplicate_variant( + &chains, + &chain_entry_str, + &variant_name, + "chain", + previous_type.span(), + ) { + return err.into(); + } + + chains.insert(chain_entry_str); + chain_exist.insert(chain_exist_entry_str); + + quote! {}.into() +} |
