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