Skip to main content

cranpose_macros/
lib.rs

1use proc_macro::TokenStream;
2use proc_macro_crate::{FoundCrate, crate_name};
3use proc_macro2::{Span, TokenStream as TokenStream2};
4use quote::quote;
5use syn::{FnArg, Ident, ItemFn, Pat, PatType, ReturnType, Type, parse_macro_input};
6
7mod branch_groups;
8
9fn is_fn_like_type(ty: &Type) -> bool {
10    match ty {
11        Type::ImplTrait(impl_trait) => impl_trait.bounds.iter().any(|bound| {
12            if let syn::TypeParamBound::Trait(trait_bound) = bound {
13                let path = &trait_bound.path;
14                if let Some(segment) = path.segments.last() {
15                    let ident_str = segment.ident.to_string();
16                    return ident_str == "FnMut" || ident_str == "Fn" || ident_str == "FnOnce";
17                }
18            }
19            false
20        }),
21        Type::Path(type_path) => {
22            if let Some(segment) = type_path.path.segments.last()
23                && segment.ident == "Box"
24                && let syn::PathArguments::AngleBracketed(args) = &segment.arguments
25                && let Some(syn::GenericArgument::Type(Type::TraitObject(trait_obj))) =
26                    args.args.first()
27            {
28                return trait_obj.bounds.iter().any(|bound| {
29                    if let syn::TypeParamBound::Trait(trait_bound) = bound {
30                        let path = &trait_bound.path;
31                        if let Some(segment) = path.segments.last() {
32                            let ident_str = segment.ident.to_string();
33                            return ident_str == "FnMut"
34                                || ident_str == "Fn"
35                                || ident_str == "FnOnce";
36                        }
37                    }
38                    false
39                });
40            }
41            false
42        }
43        Type::FnPtr(_) => true,
44        _ => false,
45    }
46}
47
48fn is_generic_fn_like(ty: &Type, generics: &syn::Generics) -> bool {
49    let type_ident = match ty {
50        Type::Path(type_path) if type_path.path.segments.len() == 1 => {
51            &type_path.path.segments[0].ident
52        }
53        _ => return false,
54    };
55
56    for param in &generics.params {
57        if let syn::GenericParam::Type(type_param) = param
58            && type_param.ident == *type_ident
59        {
60            for bound in &type_param.bounds {
61                if let syn::TypeParamBound::Trait(trait_bound) = bound
62                    && let Some(segment) = trait_bound.path.segments.last()
63                {
64                    let ident_str = segment.ident.to_string();
65                    if ident_str == "FnMut" || ident_str == "Fn" || ident_str == "FnOnce" {
66                        return true;
67                    }
68                }
69            }
70        }
71    }
72
73    if let Some(where_clause) = &generics.where_clause {
74        for predicate in &where_clause.predicates {
75            if let syn::WherePredicate::Type(pred) = predicate
76                && let Type::Path(bounded_type) = &pred.bounded_ty
77                && bounded_type.path.segments.len() == 1
78                && bounded_type.path.segments[0].ident == *type_ident
79            {
80                for bound in &pred.bounds {
81                    if let syn::TypeParamBound::Trait(trait_bound) = bound
82                        && let Some(segment) = trait_bound.path.segments.last()
83                    {
84                        let ident_str = segment.ident.to_string();
85                        if ident_str == "FnMut" || ident_str == "Fn" || ident_str == "FnOnce" {
86                            return true;
87                        }
88                    }
89                }
90            }
91        }
92    }
93
94    false
95}
96
97fn is_fn_param(ty: &Type, generics: &syn::Generics) -> bool {
98    is_fn_like_type(ty) || is_generic_fn_like(ty, generics)
99}
100
101fn is_zero_arg_fn_impl_trait(ty: &Type) -> bool {
102    if let Type::ImplTrait(impl_trait) = ty {
103        impl_trait.bounds.iter().any(|bound| {
104            if let syn::TypeParamBound::Trait(trait_bound) = bound
105                && let Some(segment) = trait_bound.path.segments.last()
106            {
107                let ident_str = segment.ident.to_string();
108                if (ident_str == "Fn" || ident_str == "FnMut")
109                    && let syn::PathArguments::Parenthesized(args) = &segment.arguments
110                {
111                    return args.inputs.is_empty();
112                }
113            }
114            false
115        })
116    } else {
117        false
118    }
119}
120
121fn type_bare_generic_ident(ty: &Type) -> Option<&Ident> {
122    match ty {
123        Type::Path(type_path)
124            if type_path.qself.is_none()
125                && type_path.path.segments.len() == 1
126                && type_path.path.segments[0].arguments.is_none() =>
127        {
128            Some(&type_path.path.segments[0].ident)
129        }
130        _ => None,
131    }
132}
133
134fn stream_mentions_ident(tokens: &TokenStream2, name: &str) -> bool {
135    tokens.clone().into_iter().any(|tt| match tt {
136        proc_macro2::TokenTree::Ident(ident) => ident == name,
137        proc_macro2::TokenTree::Group(group) => stream_mentions_ident(&group.stream(), name),
138        _ => false,
139    })
140}
141
142fn filter_generics(
143    generics: &syn::Generics,
144    strip: &std::collections::HashSet<String>,
145) -> syn::Generics {
146    let mut filtered = generics.clone();
147    filtered.params = filtered
148        .params
149        .into_iter()
150        .filter(|param| match param {
151            syn::GenericParam::Type(type_param) => !strip.contains(&type_param.ident.to_string()),
152            _ => true,
153        })
154        .collect();
155    if let Some(where_clause) = &mut filtered.where_clause {
156        where_clause.predicates = where_clause
157            .predicates
158            .clone()
159            .into_iter()
160            .filter(|predicate| {
161                if let syn::WherePredicate::Type(pred) = predicate
162                    && let Some(ident) = type_bare_generic_ident(&pred.bounded_ty)
163                {
164                    return !strip.contains(&ident.to_string());
165                }
166                true
167            })
168            .collect();
169        if where_clause.predicates.is_empty() {
170            filtered.where_clause = None;
171        }
172    }
173    filtered
174}
175
176fn is_node_id_return(ty: &Type) -> bool {
177    matches!(
178        ty,
179        Type::Path(type_path)
180            if type_path
181                .path
182                .segments
183                .last()
184                .is_some_and(|segment| segment.ident == "NodeId")
185    )
186}
187
188fn core_crate_path() -> TokenStream2 {
189    let crate_name = crate_name("cranpose")
190        .ok()
191        .or_else(|| crate_name("cranpose-core").ok());
192
193    match crate_name {
194        Some(FoundCrate::Itself) => quote!(crate),
195        Some(FoundCrate::Name(name)) => {
196            let ident = Ident::new(&name, Span::call_site());
197            quote!(#ident)
198        }
199        None => quote!(cranpose_core),
200    }
201}
202
203fn definition_key_stmt(core_path: &TokenStream2, caller_key_ident: &Ident) -> TokenStream2 {
204    quote! {
205        let #caller_key_ident = #core_path::composable_identity_key({
206            struct __CranposeDefinitionMarker;
207            static __CRANPOSE_DEFINITION_KEY: ::std::sync::OnceLock<#core_path::Key> =
208                ::std::sync::OnceLock::new();
209            #core_path::cached_composable_definition_key(
210                &__CRANPOSE_DEFINITION_KEY,
211                file!(),
212                line!(),
213                column!(),
214                ::std::any::TypeId::of::<__CranposeDefinitionMarker>(),
215            )
216        });
217    }
218}
219
220/// Turns a function into a composable: its body runs inside a group keyed by
221/// the call site, and recomposition skips it while its arguments are
222/// unchanged. `#[composable(no_skip)]` always re-runs the body.
223///
224/// Composables are named in CamelCase, as in Jetpack Compose, so the
225/// generated function carries `#[allow(non_snake_case)]` and callers need no
226/// lint allowance of their own.
227#[proc_macro_attribute]
228pub fn composable(attr: TokenStream, item: TokenStream) -> TokenStream {
229    let attr_tokens = TokenStream2::from(attr);
230    let mut enable_skip = true;
231    let core_path = core_crate_path();
232    if !attr_tokens.is_empty() {
233        match syn::parse2::<Ident>(attr_tokens) {
234            Ok(ident) if ident == "no_skip" => enable_skip = false,
235            Ok(other) => {
236                return syn::Error::new_spanned(other, "unsupported composable attribute")
237                    .to_compile_error()
238                    .into();
239            }
240            Err(err) => {
241                return err.to_compile_error().into();
242            }
243        }
244    }
245
246    let mut func = parse_macro_input!(item as ItemFn);
247
248    struct ParamInfo {
249        ident: Ident,
250        pat: Box<Pat>,
251        ty: Type,
252        pat_is_mut: bool,
253        is_impl_trait: bool,
254    }
255
256    let mut param_info: Vec<ParamInfo> = Vec::new();
257
258    for (index, arg) in func.sig.inputs.iter_mut().enumerate() {
259        if let FnArg::Typed(PatType { pat, ty, .. }) = arg {
260            if let Some(reserved) = find_reserved_pattern_ident(pat) {
261                let name = reserved.to_string();
262                return syn::Error::new(
263                    reserved.span(),
264                    format!("`{name}` is reserved by #[composable]"),
265                )
266                .to_compile_error()
267                .into();
268            }
269            let pat_is_mut = matches!(
270                pat.as_ref(),
271                Pat::Ident(pat_ident) if pat_ident.mutability.is_some()
272            );
273            let is_impl_trait = matches!(**ty, Type::ImplTrait(_));
274
275            if is_impl_trait {
276                let original_pat: Box<Pat> = pat.clone();
277                if let Pat::Ident(pat_ident) = &**pat {
278                    param_info.push(ParamInfo {
279                        ident: pat_ident.ident.clone(),
280                        pat: original_pat,
281                        ty: ty.as_ref().clone(),
282                        pat_is_mut,
283                        is_impl_trait: true,
284                    });
285                } else {
286                    param_info.push(ParamInfo {
287                        ident: Ident::new(&format!("__arg{index}"), Span::mixed_site()),
288                        pat: original_pat,
289                        ty: ty.as_ref().clone(),
290                        pat_is_mut,
291                        is_impl_trait: true,
292                    });
293                }
294            } else {
295                let ident = Ident::new(&format!("__arg{index}"), Span::mixed_site());
296                let original_pat: Box<Pat> = pat.clone();
297                **pat = syn::parse_quote! { #ident };
298                param_info.push(ParamInfo {
299                    ident,
300                    pat: original_pat,
301                    ty: ty.as_ref().clone(),
302                    pat_is_mut,
303                    is_impl_trait: false,
304                });
305            }
306        }
307    }
308
309    branch_groups::inject_branch_groups(&core_path, &mut func.block);
310    let has_rust_abi = match &func.sig.abi {
311        None => true,
312        Some(abi) => abi.name.as_ref().is_some_and(|name| name.value() == "Rust"),
313    };
314    if has_rust_abi {
315        func.attrs.push(syn::parse_quote!(#[track_caller]));
316    }
317    func.attrs.push(syn::parse_quote!(#[allow(non_snake_case)]));
318
319    let scope_label_ident = func.sig.ident.clone();
320    let original_block = func.block.clone();
321    let helper_block = original_block.clone();
322    let recompose_block = original_block.clone();
323    let composer_ident = Ident::new("__composer", Span::mixed_site());
324    let outer_composer_ident = Ident::new("__outer_composer", Span::mixed_site());
325    let caller_key_ident = Ident::new("__cranpose_caller_key", Span::mixed_site());
326    let current_scope_ident = Ident::new("__current_scope", Span::mixed_site());
327    let result_slot_index_ident = Ident::new("__result_slot_index", Span::mixed_site());
328    let has_previous_ident = Ident::new("__has_previous", Span::mixed_site());
329    let result_ident = Ident::new("__result", Span::mixed_site());
330    let value_ident = Ident::new("__value", Span::mixed_site());
331    let key_expr = quote! { #caller_key_ident };
332    let caller_key_stmt = definition_key_stmt(&core_path, &caller_key_ident);
333
334    let rebinds_for_no_skip: Vec<_> = param_info
335        .iter()
336        .map(|info| {
337            let ident = &info.ident;
338            let pat = &info.pat;
339            quote! { let #pat = #ident; }
340        })
341        .collect();
342
343    let return_ty: syn::Type = match &func.sig.output {
344        ReturnType::Default => syn::parse_quote! { () },
345        ReturnType::Type(_, ty) => ty.as_ref().clone(),
346    };
347    let returns_unit = match &func.sig.output {
348        ReturnType::Default => true,
349        ReturnType::Type(_, ty) => {
350            matches!(ty.as_ref(), Type::Tuple(tuple) if tuple.elems.is_empty())
351        }
352    };
353    let invalidate_return_consumer = if returns_unit || is_node_id_return(&return_ty) {
354        quote! {}
355    } else {
356        quote! { #composer_ident.__invalidate_return_consumer_scope(); }
357    };
358    let _helper_ident = Ident::new(
359        &format!("__cranpose_impl_{}", func.sig.ident),
360        Span::mixed_site(),
361    );
362    let generics = func.sig.generics.clone();
363    let (_impl_generics, _ty_generics, _where_clause) = generics.split_for_impl();
364
365    let _helper_inputs: Vec<TokenStream2> = param_info
366        .iter()
367        .map(|info| {
368            let ident = &info.ident;
369            let ty = &info.ty;
370            quote! { #ident: #ty }
371        })
372        .collect();
373
374    let has_unhandled_impl_trait = param_info
375        .iter()
376        .any(|info| info.is_impl_trait && !is_zero_arg_fn_impl_trait(&info.ty));
377
378    if enable_skip && !has_unhandled_impl_trait {
379        let helper_ident = Ident::new(
380            &format!("__cranpose_impl_{}", func.sig.ident),
381            Span::mixed_site(),
382        );
383        let generics = func.sig.generics.clone();
384
385        let param_erased: Vec<bool> = param_info
386            .iter()
387            .map(|info| {
388                (info.is_impl_trait && is_zero_arg_fn_impl_trait(&info.ty))
389                    || (!info.is_impl_trait
390                        && type_bare_generic_ident(&info.ty).is_some()
391                        && is_generic_fn_like(&info.ty, &generics))
392            })
393            .collect();
394
395        let mut strippable: std::collections::HashSet<String> = param_info
396            .iter()
397            .zip(&param_erased)
398            .filter(|(info, erased)| **erased && !info.is_impl_trait)
399            .filter_map(|(info, _)| type_bare_generic_ident(&info.ty))
400            .map(Ident::to_string)
401            .collect();
402        loop {
403            use quote::ToTokens;
404            let mut used_elsewhere: Vec<TokenStream2> = Vec::new();
405            for (info, erased) in param_info.iter().zip(&param_erased) {
406                if !*erased {
407                    used_elsewhere.push(info.ty.to_token_stream());
408                }
409            }
410            used_elsewhere.push(return_ty.to_token_stream());
411            for param in &generics.params {
412                match param {
413                    syn::GenericParam::Type(type_param) => {
414                        if !strippable.contains(&type_param.ident.to_string()) {
415                            used_elsewhere.push(type_param.bounds.to_token_stream());
416                            if let Some((_, default)) = &type_param.default {
417                                used_elsewhere.push(default.to_token_stream());
418                            }
419                        }
420                    }
421                    syn::GenericParam::Const(const_param) => {
422                        used_elsewhere.push(const_param.ty.to_token_stream());
423                    }
424                    syn::GenericParam::Lifetime(_) => {}
425                }
426            }
427            if let Some(where_clause) = &generics.where_clause {
428                for predicate in &where_clause.predicates {
429                    if let syn::WherePredicate::Type(pred) = predicate
430                        && let Some(ident) = type_bare_generic_ident(&pred.bounded_ty)
431                        && strippable.contains(&ident.to_string())
432                    {
433                        continue;
434                    }
435                    used_elsewhere.push(predicate.to_token_stream());
436                }
437            }
438            let before = strippable.len();
439            strippable.retain(|name| {
440                !used_elsewhere
441                    .iter()
442                    .any(|tokens| stream_mentions_ident(tokens, name))
443            });
444            if strippable.len() == before {
445                break;
446            }
447        }
448
449        let helper_generics = filter_generics(&generics, &strippable);
450        let (impl_generics, ty_generics, where_clause) = helper_generics.split_for_impl();
451        let ty_generics_turbofish = ty_generics.as_turbofish();
452
453        let helper_inputs: Vec<TokenStream2> = param_info
454            .iter()
455            .zip(&param_erased)
456            .filter_map(|(info, erased)| {
457                if info.is_impl_trait && !is_zero_arg_fn_impl_trait(&info.ty) {
458                    None
459                } else if *erased {
460                    let ident = &info.ident;
461                    Some(quote! { #ident: ::std::boxed::Box<dyn ::core::ops::FnMut() + 'static> })
462                } else {
463                    let ident = &info.ident;
464                    let ty = &info.ty;
465                    Some(quote! { #ident: #ty })
466                }
467            })
468            .collect();
469
470        let param_state_slots: Vec<Ident> = (0..param_info.len())
471            .map(|index| Ident::new(&format!("__param_state_slot{index}"), Span::mixed_site()))
472            .collect();
473
474        let param_setup: Vec<TokenStream2> = param_info
475            .iter()
476            .zip(param_state_slots.iter())
477            .zip(&param_erased)
478            .map(|((info, slot_ident), erased)| {
479                if (info.is_impl_trait && is_zero_arg_fn_impl_trait(&info.ty))
480                    || (!info.is_impl_trait && is_fn_param(&info.ty, &generics))
481                {
482                    let ident = &info.ident;
483                    let update = if *erased {
484                        quote! { holder.update_boxed(#ident); }
485                    } else {
486                        quote! { holder.update(#ident); }
487                    };
488                    quote! {
489                        let #slot_ident = #composer_ident
490                            .__use_param_slot(|| #core_path::CallbackHolder::new());
491                        #composer_ident.with_slot_value::<#core_path::CallbackHolder, _>(
492                            #slot_ident,
493                            |holder| {
494                                #update
495                            },
496                        );
497                        __changed = true;
498                    }
499                } else if info.is_impl_trait {
500                    quote! { __changed = true; }
501                } else {
502                    let ident = &info.ident;
503                    let ty = &info.ty;
504                    quote! {
505                        let #slot_ident = #composer_ident
506                            .__use_param_slot(|| #core_path::ParamState::<#ty>::default());
507                        if #composer_ident.with_slot_value_mut::<#core_path::ParamState<#ty>, _>(
508                            #slot_ident,
509                            |state| state.update(&#ident),
510                        )
511                        {
512                            __changed = true;
513                        }
514                    }
515                }
516            })
517            .collect();
518
519        let param_setup_recompose: Vec<TokenStream2> = param_info
520            .iter()
521            .zip(param_state_slots.iter())
522            .map(|(info, slot_ident)| {
523                if (info.is_impl_trait && is_zero_arg_fn_impl_trait(&info.ty))
524                    || (!info.is_impl_trait && is_fn_param(&info.ty, &generics))
525                {
526                    quote! {
527                        let #slot_ident = #composer_ident
528                            .__use_param_slot(|| #core_path::CallbackHolder::new());
529                    }
530                } else if info.is_impl_trait {
531                    quote! {}
532                } else {
533                    let ty = &info.ty;
534                    quote! {
535                        let #slot_ident = #composer_ident
536                            .__use_param_slot(|| #core_path::ParamState::<#ty>::default());
537                    }
538                }
539            })
540            .collect();
541
542        let rebinds: Vec<TokenStream2> = param_info
543            .iter()
544            .zip(param_state_slots.iter())
545            .map(|(info, slot_ident)| {
546                if (info.is_impl_trait && is_zero_arg_fn_impl_trait(&info.ty))
547                    || (!info.is_impl_trait && is_fn_param(&info.ty, &generics))
548                {
549                    let pat = &info.pat;
550                    let can_add_mut = matches!(pat.as_ref(), Pat::Ident(_));
551                    if can_add_mut && !info.pat_is_mut {
552                        quote! {
553                            #[allow(unused_mut)]
554                            let mut #pat = #composer_ident
555                                .with_slot_value::<#core_path::CallbackHolder, _>(
556                                    #slot_ident,
557                                    |holder| holder.clone_rc(),
558                                );
559                        }
560                    } else {
561                        quote! {
562                            #[allow(unused_mut)]
563                            let #pat = #composer_ident
564                                .with_slot_value::<#core_path::CallbackHolder, _>(
565                                    #slot_ident,
566                                    |holder| holder.clone_rc(),
567                                );
568                        }
569                    }
570                } else if info.is_impl_trait {
571                    quote! {}
572                } else {
573                    let pat = &info.pat;
574                    let ident = &info.ident;
575                    quote! {
576                        let #pat = #ident;
577                    }
578                }
579            })
580            .collect();
581
582        let rebinds_for_recompose: Vec<TokenStream2> = param_info
583            .iter()
584            .zip(param_state_slots.iter())
585            .map(|(info, slot_ident)| {
586                if (info.is_impl_trait && is_zero_arg_fn_impl_trait(&info.ty))
587                    || (!info.is_impl_trait && is_fn_param(&info.ty, &generics))
588                {
589                    let pat = &info.pat;
590                    let can_add_mut = matches!(pat.as_ref(), Pat::Ident(_));
591                    if can_add_mut && !info.pat_is_mut {
592                        quote! {
593                            #[allow(unused_mut)]
594                            let mut #pat = #composer_ident
595                                .with_slot_value::<#core_path::CallbackHolder, _>(
596                                    #slot_ident,
597                                    |holder| holder.clone_rc(),
598                                );
599                        }
600                    } else {
601                        quote! {
602                            #[allow(unused_mut)]
603                            let #pat = #composer_ident
604                                .with_slot_value::<#core_path::CallbackHolder, _>(
605                                    #slot_ident,
606                                    |holder| holder.clone_rc(),
607                                );
608                        }
609                    }
610                } else if info.is_impl_trait {
611                    quote! {}
612                } else {
613                    let pat = &info.pat;
614                    let ty = &info.ty;
615                    quote! {
616                        let #pat = #composer_ident
617                            .with_slot_value::<#core_path::ParamState<#ty>, _>(
618                                #slot_ident,
619                                |state| {
620                                    state
621                                        .value()
622                                        .expect("composable parameter missing for recomposition")
623                                },
624                            );
625                    }
626                }
627            })
628            .collect();
629
630        let recompose_fn_ident = Ident::new(
631            &format!("__cranpose_recompose_{}", func.sig.ident),
632            Span::mixed_site(),
633        );
634
635        let recompose_setter = quote! {
636            {
637                #composer_ident.set_recompose_callback(move |
638                    #composer_ident: &#core_path::Composer|
639                {
640                    let _ = #recompose_fn_ident #ty_generics_turbofish (
641                        #composer_ident
642                    );
643                });
644            }
645        };
646
647        let helper_body = if returns_unit {
648            quote! {
649                #core_path::debug_label_current_scope(stringify!(#scope_label_ident));
650                let #current_scope_ident = #composer_ident
651                    .current_recompose_scope()
652                    .expect("missing recompose scope");
653                let mut __changed = #current_scope_ident.should_recompose();
654                #(#param_setup)*
655                #recompose_setter
656                if !__changed && #current_scope_ident.has_composed_once() {
657                    #composer_ident.skip_current_group();
658                    return;
659                }
660                #(#rebinds)*
661                #helper_block
662            }
663        } else {
664            quote! {
665                #core_path::debug_label_current_scope(stringify!(#scope_label_ident));
666                let #current_scope_ident = #composer_ident
667                    .current_recompose_scope()
668                    .expect("missing recompose scope");
669                let mut __changed = #current_scope_ident.should_recompose();
670                #(#param_setup)*
671                #recompose_setter
672                let #result_slot_index_ident = #composer_ident
673                    .__use_return_slot(|| #core_path::ReturnSlot::<#return_ty>::default());
674                let #has_previous_ident = #composer_ident
675                    .with_slot_value::<#core_path::ReturnSlot<#return_ty>, _>(
676                        #result_slot_index_ident,
677                        |slot| slot.get().is_some(),
678                    );
679                if !__changed && #has_previous_ident {
680                    #composer_ident.skip_current_group();
681                    let #result_ident = #composer_ident
682                        .with_slot_value::<#core_path::ReturnSlot<#return_ty>, _>(
683                            #result_slot_index_ident,
684                            |slot| {
685                                slot.get()
686                                    .expect("composable return value missing during skip")
687                            },
688                        );
689                    return #result_ident;
690                }
691                let #value_ident: #return_ty = {
692                    #(#rebinds)*
693                    #helper_block
694                };
695                #composer_ident.with_slot_value_mut::<#core_path::ReturnSlot<#return_ty>, _>(
696                    #result_slot_index_ident,
697                    |slot| {
698                        slot.store(#value_ident.clone());
699                    },
700                );
701                #value_ident
702            }
703        };
704
705        let recompose_fn_body = if returns_unit {
706            quote! {
707                #(#param_setup_recompose)*
708                #(#rebinds_for_recompose)*
709                #recompose_block
710                #recompose_setter
711            }
712        } else {
713            quote! {
714                #(#param_setup_recompose)*
715                let #result_slot_index_ident = #composer_ident
716                    .__use_return_slot(|| #core_path::ReturnSlot::<#return_ty>::default());
717                #(#rebinds_for_recompose)*
718                let #value_ident: #return_ty = {
719                    #recompose_block
720                };
721                #composer_ident.with_slot_value_mut::<#core_path::ReturnSlot<#return_ty>, _>(
722                    #result_slot_index_ident,
723                    |slot| {
724                        slot.store(#value_ident.clone());
725                    },
726                );
727                #recompose_setter
728                #invalidate_return_consumer
729                #value_ident
730            }
731        };
732
733        let recompose_fn = quote! {
734            #[allow(non_snake_case)]
735            fn #recompose_fn_ident #impl_generics (
736                #composer_ident: &#core_path::Composer
737            ) -> #return_ty #where_clause {
738                #recompose_fn_body
739            }
740        };
741
742        let helper_fn = quote! {
743            #[allow(non_snake_case, clippy::too_many_arguments)]
744            fn #helper_ident #impl_generics (
745                #composer_ident: &#core_path::Composer
746                #(, #helper_inputs)*
747            ) -> #return_ty #where_clause {
748                #helper_body
749            }
750        };
751
752        let wrapper_args: Vec<TokenStream2> = param_info
753            .iter()
754            .zip(&param_erased)
755            .filter_map(|(info, erased)| {
756                if info.is_impl_trait && !is_zero_arg_fn_impl_trait(&info.ty) {
757                    None
758                } else if *erased {
759                    let ident = &info.ident;
760                    Some(quote! { ::std::boxed::Box::new(#ident) })
761                } else {
762                    let ident = &info.ident;
763                    Some(quote! { #ident })
764                }
765            })
766            .collect();
767
768        let wrapped = quote!({
769            #caller_key_stmt
770            #core_path::with_current_composer(|#composer_ident: &#core_path::Composer| {
771                #composer_ident.with_group(#key_expr, |#composer_ident: &#core_path::Composer| {
772                    #helper_ident(#composer_ident #(, #wrapper_args)*)
773                })
774            })
775        });
776        *func.block = syn::parse2(wrapped).expect("failed to build block");
777        TokenStream::from(quote! {
778            #recompose_fn
779            #helper_fn
780            #func
781        })
782    } else {
783        let wrapped = quote!({
784            #caller_key_stmt
785            #core_path::with_current_composer(|#outer_composer_ident: &#core_path::Composer| {
786                #outer_composer_ident.with_group(#key_expr, |#composer_ident: &#core_path::Composer| {
787                    #core_path::debug_label_current_scope(stringify!(#scope_label_ident));
788                    #(#rebinds_for_no_skip)*
789                    #original_block
790                })
791            })
792        });
793        *func.block = syn::parse2(wrapped).expect("failed to build block");
794        TokenStream::from(quote! { #func })
795    }
796}
797
798fn find_reserved_pattern_ident(pat: &Pat) -> Option<&Ident> {
799    use syn::visit::Visit;
800
801    struct Scan<'ast> {
802        found: Option<&'ast Ident>,
803    }
804    impl<'ast> syn::visit::Visit<'ast> for Scan<'ast> {
805        fn visit_pat_ident(&mut self, node: &'ast syn::PatIdent) {
806            if self.found.is_none() {
807                let name = node.ident.to_string();
808                if name == "__composer" || name.starts_with("__cranpose") {
809                    self.found = Some(&node.ident);
810                }
811            }
812            syn::visit::visit_pat_ident(self, node);
813        }
814    }
815    let mut scan = Scan { found: None };
816    scan.visit_pat(pat);
817    scan.found
818}
819
820#[cfg(test)]
821mod tests {
822    use super::*;
823
824    #[test]
825    fn definition_key_does_not_monomorphise_the_once_lock_initializer() {
826        let core_path = quote!(::cranpose_core);
827        let ident = Ident::new("__cranpose_caller_key", Span::mixed_site());
828        let tokens = definition_key_stmt(&core_path, &ident).to_string();
829
830        assert!(
831            tokens.contains("cached_composable_definition_key"),
832            "the definition key must be latched through the outlined core \
833             helper, got: {tokens}"
834        );
835        assert!(
836            !tokens.contains("get_or_init"),
837            "no initializer closure may reach the expansion site, got: {tokens}"
838        );
839    }
840}