//! Proc macros for `future_form`. //! //! This crate provides the `#[future_form]` attribute macro. use proc_macro::TokenStream; use proc_macro2::TokenStream as TokenStream2; use quote::quote; use syn::{ GenericParam, Ident, ImplItem, ImplItemFn, ItemImpl, Path, ReturnType, Type, WherePredicate, parse_macro_input, parse_quote, visit::Visit, visit_mut::{self, VisitMut}, }; /// Generate implementations of a trait for `Sendable` and/or `Local` `FutureForm`s. /// /// This attribute macro allows you to write a single implementation that works for both /// `Send` and `!Send` futures, avoiding code duplication. /// /// # Usage /// /// ```rust,ignore /// // Generate both Sendable and Local impls /// #[future_form(Sendable, Local)] /// impl MyTrait for MyType { ... } /// /// // Generate only Sendable impl /// #[future_form(Sendable)] /// impl MyTrait for MyType { ... } /// /// // Generate only Local impl /// #[future_form(Local)] /// impl MyTrait for MyType { ... } /// /// // Add bounds only for specific variants /// #[future_form(Sendable where T: Send, Local)] /// impl MyTrait for Container { ... } /// // Generates: impl MyTrait for Container /// // impl MyTrait for Container /// ``` /// /// # Example /// /// ```rust,ignore /// use std::marker::PhantomData; /// use future_form::{FutureForm, Sendable, Local, future_form}; /// /// trait Counter { /// fn next(&self) -> F::Future<'_, u32>; /// } /// /// struct Memory { /// val: u32, /// _marker: PhantomData, /// } /// /// #[future_form(Sendable, Local)] /// impl Counter for Memory { /// fn next(&self) -> F::Future<'_, u32> { /// let val = self.val; /// F::from_future(async move { val + 1 }) /// } /// } /// ``` #[proc_macro_attribute] pub fn future_form(attr: TokenStream, item: TokenStream) -> TokenStream { let input = parse_macro_input!(item as ItemImpl); let kinds = match parse_kinds(&attr) { Ok(k) => k, Err(err) => return err, }; match generate_impls(&input, &kinds) { Ok(tokens) => tokens.into(), Err(err) => err.to_compile_error().into(), } } /// Creates an error `TokenStream` with proper span information. fn make_error(span: proc_macro2::Span, msg: &str) -> TokenStream { syn::Error::new(span, msg).to_compile_error().into() } /// Returns the concrete future path for built-in variants, or None for custom types. fn builtin_future_path(ident: &Ident) -> Option { if ident == "Sendable" { Some(parse_quote!(::futures::future::BoxFuture)) } else if ident == "Local" { Some(parse_quote!(::futures::future::LocalBoxFuture)) } else { None } } /// Checks if an ident could be a variant name. /// /// Variants are type names like `Sendable`, `Local`, or custom `FutureForm` types. /// Excludes: `where`, single-letter idents (likely type params like T, K, F), /// and common trait names that appear in bounds. fn is_likely_variant(ident: &Ident) -> bool { let s = ident.to_string(); // Must start uppercase if !s.chars().next().is_some_and(char::is_uppercase) { return false; } // Exclude keywords if s == "where" || s == "Self" { return false; } // Exclude single-letter idents (likely type parameters: T, K, F, U, etc.) if s.len() == 1 { return false; } // Exclude common trait names that appear in bounds let common_traits = [ "Send", "Sync", "Clone", "Copy", "Debug", "Display", "Default", "Fn", "FnMut", "FnOnce", "Future", "Iterator", "IntoIterator", "From", "Into", "TryFrom", "TryInto", "AsRef", "AsMut", "Eq", "PartialEq", "Ord", "PartialOrd", "Hash", "Sized", "Unpin", "Drop", ]; if common_traits.contains(&s.as_str()) { return false; } true } /// Parses the attribute tokens into a list of `FutureFormVariant`s. /// /// Uses token-based parsing to correctly handle all delimiters (parentheses, /// brackets, braces, angle brackets) in where clause bounds. #[allow(clippy::expect_used)] // Indexing is bounds-checked by loop conditions fn parse_kinds(attr: &TokenStream) -> Result, TokenStream> { use proc_macro2::TokenTree; let attr2: TokenStream2 = attr.clone().into(); let tokens: Vec = attr2.into_iter().collect(); if tokens.is_empty() { return Err(make_error( proc_macro2::Span::call_site(), "missing FutureForm variants: expected #[future_form(Sendable)], #[future_form(Local)], or #[future_form(Sendable, Local)]", )); } let mut kinds = Vec::new(); let mut i = 0; while i < tokens.len() { // Skip leading commas while i < tokens.len() { if let TokenTree::Punct(p) = tokens.get(i).expect("bounds checked") && p.as_char() == ',' { i += 1; continue; } break; } if i >= tokens.len() { break; } // Expect a variant name (any FutureForm type: Sendable, Local, or custom) let (kind_path, future_path, variant_span) = match tokens.get(i).expect("bounds checked") { TokenTree::Ident(ident) => { if is_likely_variant(ident) { let future = builtin_future_path(ident); // Qualify built-in variants so users don't need to import them let path: Path = if future.is_some() { parse_quote!(::future_form::#ident) } else { parse_quote!(#ident) }; (path, future, ident.span()) } else { return Err(make_error( ident.span(), &format!("expected FutureForm variant, found `{ident}`"), )); } } other @ (TokenTree::Group(_) | TokenTree::Punct(_) | TokenTree::Literal(_)) => { return Err(make_error( other.span(), "expected FutureForm variant (e.g., `Sendable`, `Local`, or custom type)", )); } }; i += 1; // Check for optional `where` clause let extra_bounds = if i < tokens.len() { if let TokenTree::Ident(ident) = tokens.get(i).expect("bounds checked") { if ident == "where" { i += 1; // Collect tokens until we hit another variant name or end let (predicates, new_i) = collect_where_clause(&tokens, i, variant_span)?; i = new_i; predicates } else { vec![] } } else { vec![] } } else { vec![] }; kinds.push(FutureFormVariant { kind_path, future_path, extra_bounds, }); } if kinds.is_empty() { return Err(make_error( proc_macro2::Span::call_site(), "missing FutureForm variants: expected #[future_form(Sendable)], #[future_form(Local)], or #[future_form(Sendable, Local)]", )); } Ok(kinds) } /// Collects where clause tokens until hitting another variant name or end. /// /// Returns the parsed predicates and the new token index. #[allow(clippy::expect_used)] // Indexing is bounds-checked by loop conditions fn collect_where_clause( tokens: &[proc_macro2::TokenTree], start: usize, span: proc_macro2::Span, ) -> Result<(Vec, usize), TokenStream> { use proc_macro2::TokenTree; let mut i = start; let mut current_predicate_tokens: Vec = Vec::new(); let mut predicates = Vec::new(); // Track whether we're after a `:` (inside a type bound) - variants only appear after `,` let mut in_bound = false; while i < tokens.len() { let token = tokens.get(i).expect("bounds checked"); // Track if we're inside a bound (after `:`) if let TokenTree::Punct(p) = token { if p.as_char() == ':' { in_bound = true; } else if p.as_char() == ',' { in_bound = false; } } // Check if this is a variant name (end of where clause) // Only check at the START of a predicate (not inside a bound) // Must NOT be followed by `:` (which would make it a type param bound, not a variant) if !in_bound && current_predicate_tokens.is_empty() && let TokenTree::Ident(ident) = token && is_likely_variant(ident) && !tokens .get(i + 1) .is_some_and(|t| matches!(t, TokenTree::Punct(p) if p.as_char() == ':')) { return Ok((predicates, i)); } // Check for comma at top level (predicate separator) if let TokenTree::Punct(p) = token && p.as_char() == ',' { // Check if next token is a variant name (starts a new variant, not a predicate) let next_is_variant = tokens.get(i + 1).is_some_and(|t| { if let TokenTree::Ident(id) = t { // It's a variant if it looks like one AND is not followed by `:` is_likely_variant(id) && !tokens.get(i + 2).is_some_and(|t2| { matches!(t2, TokenTree::Punct(p2) if p2.as_char() == ':') }) } else { false } }); if next_is_variant { // This comma separates variants, not predicates if !current_predicate_tokens.is_empty() { predicates.push(parse_predicate_tokens(¤t_predicate_tokens, span)?); } i += 1; // Skip the comma return Ok((predicates, i)); } // This comma separates predicates within the where clause if !current_predicate_tokens.is_empty() { predicates.push(parse_predicate_tokens(¤t_predicate_tokens, span)?); current_predicate_tokens.clear(); } i += 1; continue; } // Add token to current predicate current_predicate_tokens.push(token.clone()); i += 1; } // End of tokens - parse any remaining predicate if !current_predicate_tokens.is_empty() { predicates.push(parse_predicate_tokens(¤t_predicate_tokens, span)?); } Ok((predicates, i)) } /// Parses a sequence of tokens as a where predicate. fn parse_predicate_tokens( tokens: &[proc_macro2::TokenTree], span: proc_macro2::Span, ) -> Result { let token_stream: TokenStream2 = tokens.iter().cloned().collect(); let token_str = token_stream.to_string(); syn::parse2::(token_stream).map_err(|e| { make_error(span, &format!("malformed where clause `{token_str}`: {e}")) }) } fn generate_impls(input: &ItemImpl, kinds: &[FutureFormVariant]) -> syn::Result { // Find the K type parameter let k_param = find_k_param(input)?; // Generate impl for each requested kind let impls: Vec = kinds .iter() .map(|kind| generate_impl_for_kind(input, &k_param, kind)) .collect(); Ok(quote! { #(#impls)* }) } #[derive(Clone)] struct FutureFormVariant { /// Path to the `FutureForm` implementor (e.g., `Sendable`, `Local`, `MyCustomForm`) kind_path: Path, /// Concrete future path for built-in types (e.g., `BoxFuture`), None for custom types future_path: Option, /// Additional where clause bounds for this variant extra_bounds: Vec, } /// Visitor that replaces standalone `K` type parameter references with concrete paths. /// /// Only replaces `Path` nodes where the first segment is exactly the target ident /// (e.g., `K`, `K::Future`, `K::from_future`), not idents containing it as a substring. struct KindReplacer { from_ident: Ident, to_path: Path, /// Concrete future path for built-in types (`BoxFuture`, `LocalBoxFuture`). /// When `None` (custom types), `K::Future` becomes `CustomType::Future`. future_path: Option, } impl VisitMut for KindReplacer { fn visit_path_mut(&mut self, path: &mut Path) { // First, recurse into nested paths (generic arguments, etc.) visit_mut::visit_path_mut(self, path); // Check if first segment is exactly our target ident if let Some(first) = path.segments.first() && first.ident == self.from_ident && first.arguments.is_empty() { if path.segments.len() == 1 { // Standalone K → replace with the target type *path = self.to_path.clone(); } else if let Some(second) = path.segments.get(1) { // K::Future<...> or K::from_future(...) if second.ident == "Future" { if let Some(ref concrete_future) = self.future_path { // Built-in: K::Future<'a, T> → BoxFuture<'a, T> let args = second.arguments.clone(); let mut new_path = concrete_future.clone(); if let Some(last) = new_path.segments.last_mut() { last.arguments = args; } *path = new_path; } else { // Custom: K::Future<'a, T> → CustomType::Future<'a, T> let remaining: Vec<_> = path.segments.iter().skip(1).cloned().collect(); let mut new_path = self.to_path.clone(); new_path.segments.extend(remaining); *path = new_path; } } else { // K::from_future → Type::from_future let remaining: Vec<_> = path.segments.iter().skip(1).cloned().collect(); let mut new_path = self.to_path.clone(); new_path.segments.extend(remaining); *path = new_path; } } } } } /// Visitor that checks if a where predicate references a specific ident as a standalone type. struct IdentFinder { target: Ident, found: bool, } impl<'ast> Visit<'ast> for IdentFinder { fn visit_path(&mut self, path: &'ast Path) { // Only match standalone K, not as part of larger identifiers if let Some(first) = path.segments.first() && path.segments.len() == 1 && first.ident == self.target && first.arguments.is_empty() { self.found = true; } syn::visit::visit_path(self, path); } } fn find_k_param(input: &ItemImpl) -> syn::Result { // Look for a type parameter that has FutureForm bound for param in &input.generics.params { if let GenericParam::Type(type_param) = param { for bound in &type_param.bounds { if let syn::TypeParamBound::Trait(trait_bound) = bound { let path = &trait_bound.path; if path .segments .last() .is_some_and(|s| s.ident == "FutureForm") { return Ok(type_param.ident.clone()); } } } } } Err(syn::Error::new_spanned( &input.generics, "Expected a type parameter with FutureForm bound (e.g., `K: FutureForm`)", )) } fn generate_impl_for_kind( input: &ItemImpl, k_param: &Ident, variant: &FutureFormVariant, ) -> TokenStream2 { let kind_path: Path = variant.kind_path.clone(); let future_path: Option = variant.future_path.clone(); // Clone and modify generics - remove the K parameter let mut new_generics = input.generics.clone(); new_generics.params = new_generics .params .into_iter() .filter(|p| { if let GenericParam::Type(tp) = p { tp.ident != *k_param } else { true } }) .collect(); // Remove where clause predicates that reference K, and add extra bounds if let Some(ref mut where_clause) = new_generics.where_clause { where_clause.predicates = where_clause .predicates .clone() .into_iter() .filter(|pred| !predicate_references_ident(pred, k_param)) .collect(); // Add variant-specific extra bounds for bound in &variant.extra_bounds { where_clause.predicates.push(bound.clone()); } } else if !variant.extra_bounds.is_empty() { // Create a where clause if we have extra bounds but none existed let mut predicates = syn::punctuated::Punctuated::new(); for bound in &variant.extra_bounds { predicates.push(bound.clone()); } new_generics.where_clause = Some(syn::WhereClause { where_token: syn::token::Where::default(), predicates, }); } // Replace K in self_ty let new_self_ty = replace_ident_in_type(&input.self_ty, k_param, &kind_path, &future_path); // Replace K in trait path if present let new_trait = input.trait_.as_ref().map(|(bang, path, for_token)| { let new_path = replace_ident_in_path(path, k_param, &kind_path, &future_path); (*bang, new_path, *for_token) }); // Transform methods let new_items: Vec = input .items .iter() .map(|item| transform_impl_item(item, k_param, &kind_path, &future_path)) .collect(); let (impl_generics, _, where_clause) = new_generics.split_for_impl(); let trait_tokens = new_trait.map(|(bang, path, for_token)| { quote! { #bang #path #for_token } }); quote! { impl #impl_generics #trait_tokens #new_self_ty #where_clause { #(#new_items)* } } } fn predicate_references_ident(pred: &syn::WherePredicate, ident: &Ident) -> bool { let mut finder = IdentFinder { target: ident.clone(), found: false, }; finder.visit_where_predicate(pred); finder.found } #[allow(clippy::ref_option)] // We clone the Option, so &Option is fine fn replace_ident_in_type( ty: &Type, from: &Ident, kind_path: &Path, future_path: &Option, ) -> Type { let mut ty = ty.clone(); let mut replacer = KindReplacer { from_ident: from.clone(), to_path: kind_path.clone(), future_path: future_path.clone(), }; replacer.visit_type_mut(&mut ty); ty } #[allow(clippy::ref_option)] fn replace_ident_in_path( path: &Path, from: &Ident, kind_path: &Path, future_path: &Option, ) -> Path { let mut path = path.clone(); let mut replacer = KindReplacer { from_ident: from.clone(), to_path: kind_path.clone(), future_path: future_path.clone(), }; replacer.visit_path_mut(&mut path); path } #[allow(clippy::wildcard_enum_match_arm)] // We want to pass through unknown variants unchanged #[allow(clippy::ref_option)] fn transform_impl_item( item: &ImplItem, k_param: &Ident, kind_path: &Path, future_path: &Option, ) -> ImplItem { match item { ImplItem::Fn(method) => { ImplItem::Fn(transform_method(method, k_param, kind_path, future_path)) } other => other.clone(), } } #[allow(clippy::ref_option)] fn transform_method( method: &ImplItemFn, k_param: &Ident, kind_path: &Path, future_path: &Option, ) -> ImplItemFn { let mut new_method = method.clone(); // Use the visitor to replace K references in the entire method let mut replacer = KindReplacer { from_ident: k_param.clone(), to_path: kind_path.clone(), future_path: future_path.clone(), }; // Transform return type if let ReturnType::Type(_, ref mut ty) = new_method.sig.output { replacer.visit_type_mut(ty); } // Transform method body replacer.visit_block_mut(&mut new_method.block); new_method }