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},
};
#[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(),
}
}
fn make_error(span: proc_macro2::Span, msg: &str) -> TokenStream {
syn::Error::new(span, msg).to_compile_error().into()
}
fn builtin_future_path(ident: &Ident) -> Option<Path> {
if ident == "Sendable" {
Some(parse_quote!(::futures::future::BoxFuture))
} else if ident == "Local" {
Some(parse_quote!(::futures::future::LocalBoxFuture))
} else {
None
}
}
fn is_likely_variant(ident: &Ident) -> bool {
let s = ident.to_string();
if !s.chars().next().is_some_and(char::is_uppercase) {
return false;
}
if s == "where" || s == "Self" {
return false;
}
if s.len() == 1 {
return false;
}
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
}
#[allow(clippy::expect_used)] fn parse_kinds(attr: &TokenStream) -> Result<Vec<FutureFormVariant>, TokenStream> {
use proc_macro2::TokenTree;
let attr2: TokenStream2 = attr.clone().into();
let tokens: Vec<TokenTree> = 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() {
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;
}
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);
let path: Path = 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;
let extra_bounds = if i < tokens.len() {
if let TokenTree::Ident(ident) = tokens.get(i).expect("bounds checked") {
if ident == "where" {
i += 1;
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)
}
#[allow(clippy::expect_used)] fn collect_where_clause(
tokens: &[proc_macro2::TokenTree],
start: usize,
span: proc_macro2::Span,
) -> Result<(Vec<WherePredicate>, usize), TokenStream> {
use proc_macro2::TokenTree;
let mut i = start;
let mut current_predicate_tokens: Vec<TokenTree> = Vec::new();
let mut predicates = Vec::new();
let mut in_bound = false;
while i < tokens.len() {
let token = tokens.get(i).expect("bounds checked");
if let TokenTree::Punct(p) = token {
if p.as_char() == ':' {
in_bound = true;
} else if p.as_char() == ',' {
in_bound = false;
}
}
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));
}
if let TokenTree::Punct(p) = token
&& p.as_char() == ','
{
let next_is_variant = tokens.get(i + 1).is_some_and(|t| {
if let TokenTree::Ident(id) = t {
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 {
if !current_predicate_tokens.is_empty() {
predicates.push(parse_predicate_tokens(¤t_predicate_tokens, span)?);
}
i += 1; return Ok((predicates, i));
}
if !current_predicate_tokens.is_empty() {
predicates.push(parse_predicate_tokens(¤t_predicate_tokens, span)?);
current_predicate_tokens.clear();
}
i += 1;
continue;
}
current_predicate_tokens.push(token.clone());
i += 1;
}
if !current_predicate_tokens.is_empty() {
predicates.push(parse_predicate_tokens(¤t_predicate_tokens, span)?);
}
Ok((predicates, i))
}
fn parse_predicate_tokens(
tokens: &[proc_macro2::TokenTree],
span: proc_macro2::Span,
) -> Result<WherePredicate, TokenStream> {
let token_stream: TokenStream2 = tokens.iter().cloned().collect();
let token_str = token_stream.to_string();
syn::parse2::<WherePredicate>(token_stream).map_err(|e| {
make_error(span, &format!("malformed where clause `{token_str}`: {e}"))
})
}
fn generate_impls(input: &ItemImpl, kinds: &[FutureFormVariant]) -> syn::Result<TokenStream2> {
let k_param = find_k_param(input)?;
let impls: Vec<TokenStream2> = kinds
.iter()
.map(|kind| generate_impl_for_kind(input, &k_param, kind))
.collect();
Ok(quote! {
#(#impls)*
})
}
#[derive(Clone)]
struct FutureFormVariant {
kind_path: Path,
future_path: Option<Path>,
extra_bounds: Vec<WherePredicate>,
}
struct KindReplacer {
from_ident: Ident,
to_path: Path,
future_path: Option<Path>,
}
impl VisitMut for KindReplacer {
fn visit_path_mut(&mut self, path: &mut Path) {
visit_mut::visit_path_mut(self, path);
if let Some(first) = path.segments.first()
&& first.ident == self.from_ident
&& first.arguments.is_empty()
{
if path.segments.len() == 1 {
*path = self.to_path.clone();
} else if let Some(second) = path.segments.get(1) {
if second.ident == "Future" {
if let Some(ref concrete_future) = self.future_path {
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 {
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 {
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;
}
}
}
}
}
struct IdentFinder {
target: Ident,
found: bool,
}
impl<'ast> Visit<'ast> for IdentFinder {
fn visit_path(&mut self, path: &'ast Path) {
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<Ident> {
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<Path> = variant.future_path.clone();
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();
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();
for bound in &variant.extra_bounds {
where_clause.predicates.push(bound.clone());
}
} else if !variant.extra_bounds.is_empty() {
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,
});
}
let new_self_ty = replace_ident_in_type(&input.self_ty, k_param, &kind_path, &future_path);
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)
});
let new_items: Vec<ImplItem> = 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)] fn replace_ident_in_type(
ty: &Type,
from: &Ident,
kind_path: &Path,
future_path: &Option<Path>,
) -> 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>,
) -> 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)] #[allow(clippy::ref_option)]
fn transform_impl_item(
item: &ImplItem,
k_param: &Ident,
kind_path: &Path,
future_path: &Option<Path>,
) -> 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<Path>,
) -> ImplItemFn {
let mut new_method = method.clone();
let mut replacer = KindReplacer {
from_ident: k_param.clone(),
to_path: kind_path.clone(),
future_path: future_path.clone(),
};
if let ReturnType::Type(_, ref mut ty) = new_method.sig.output {
replacer.visit_type_mut(ty);
}
replacer.visit_block_mut(&mut new_method.block);
new_method
}