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