1#![allow(unused_variables)] use proc_macro::TokenStream;
4use proc_macro2::TokenStream as TokenStream2;
5use quote::quote;
6use std::time::Duration;
7use syn::{
8 FnArg, ItemFn, LitStr, Pat, PatType, Result, ReturnType, Token, Type, parse::Parse,
9 parse::ParseStream, parse_macro_input,
10};
11
12#[derive(Default)]
14struct ProviderArgs {
15 interval: Option<Duration>,
16 cache_expiration: Option<Duration>,
17 stale_time: Option<Duration>,
18 compose: Vec<syn::Ident>, }
20
21#[derive(Default)]
23struct MutationArgs {
24 invalidates: Vec<syn::Ident>, optimistic: Option<syn::ExprClosure>, }
27
28impl Parse for ProviderArgs {
29 fn parse(input: ParseStream) -> Result<Self> {
30 let mut args = ProviderArgs::default();
31
32 while !input.is_empty() {
33 let ident: syn::Ident = input.parse()?;
34 input.parse::<Token![=]>()?;
35
36 match ident.to_string().as_str() {
37 "interval" => {
38 let lit: LitStr = input.parse()?;
39 let duration_str = lit.value();
40 let duration = humantime::parse_duration(&duration_str).map_err(|e| {
41 syn::Error::new_spanned(lit, format!("Invalid duration format: {e}"))
42 })?;
43 args.interval = Some(duration);
44 }
45 "cache_expiration" => {
46 let lit: LitStr = input.parse()?;
47 let duration_str = lit.value();
48 let duration = humantime::parse_duration(&duration_str).map_err(|e| {
49 syn::Error::new_spanned(lit, format!("Invalid duration format: {e}"))
50 })?;
51 args.cache_expiration = Some(duration);
52 }
53 "stale_time" => {
54 let lit: LitStr = input.parse()?;
55 let duration_str = lit.value();
56 let duration = humantime::parse_duration(&duration_str).map_err(|e| {
57 syn::Error::new_spanned(lit, format!("Invalid duration format: {e}"))
58 })?;
59 args.stale_time = Some(duration);
60 }
61 "compose" => {
62 let content;
64 syn::bracketed!(content in input);
65 let providers = content.parse_terminated(syn::Ident::parse, Token![,])?;
66 args.compose = providers.into_iter().collect();
67 }
68 _ => return Err(syn::Error::new_spanned(ident, "Unknown argument")),
69 }
70
71 if input.peek(Token![,]) {
72 input.parse::<Token![,]>()?;
73 }
74 }
75
76 Ok(args)
77 }
78}
79
80impl Parse for MutationArgs {
81 fn parse(input: ParseStream) -> Result<Self> {
82 let mut args = MutationArgs::default();
83
84 while !input.is_empty() {
85 let ident: syn::Ident = input.parse()?;
86 input.parse::<Token![=]>()?;
87
88 match ident.to_string().as_str() {
89 "invalidates" => {
90 let content;
92 syn::bracketed!(content in input);
93 let providers = content.parse_terminated(syn::Ident::parse, Token![,])?;
94 args.invalidates = providers.into_iter().collect();
95 }
96 "optimistic" => {
97 let expr: syn::ExprClosure = input.parse()?;
98 args.optimistic = Some(expr);
99 }
100 _ => return Err(syn::Error::new_spanned(ident, "Unknown argument")),
101 }
102
103 if input.peek(Token![,]) {
104 input.parse::<Token![,]>()?;
105 }
106 }
107
108 Ok(args)
109 }
110}
111
112#[proc_macro_attribute]
186pub fn provider(args: TokenStream, input: TokenStream) -> TokenStream {
187 let provider_args = if args.is_empty() {
188 ProviderArgs::default()
189 } else {
190 match syn::parse(args) {
191 Ok(args) => args,
192 Err(err) => return err.to_compile_error().into(),
193 }
194 };
195
196 let input_fn = parse_macro_input!(input as ItemFn);
197
198 let result = generate_provider(input_fn, provider_args);
199
200 match result {
201 Ok(tokens) => tokens.into(),
202 Err(err) => err.to_compile_error().into(),
203 }
204}
205
206#[proc_macro_attribute]
277pub fn mutation(args: TokenStream, input: TokenStream) -> TokenStream {
278 let mutation_args = if args.is_empty() {
279 MutationArgs::default()
280 } else {
281 match syn::parse(args) {
282 Ok(args) => args,
283 Err(err) => return err.to_compile_error().into(),
284 }
285 };
286
287 let input_fn = parse_macro_input!(input as ItemFn);
288
289 let result = generate_mutation(input_fn, mutation_args);
290
291 match result {
292 Ok(tokens) => tokens.into(),
293 Err(err) => err.to_compile_error().into(),
294 }
295}
296
297fn generate_provider(input_fn: ItemFn, provider_args: ProviderArgs) -> Result<TokenStream2> {
298 let info = extract_provider_info(&input_fn)?;
299
300 let ProviderInfo {
301 fn_vis,
302 fn_block,
303 output_type,
304 error_type,
305 struct_name,
306 ..
307 } = &info;
308
309 let params = extract_all_params(&input_fn)?;
311
312 if !provider_args.compose.is_empty() {
314 validate_composition_requirements(&provider_args.compose, ¶ms)?;
315 }
316
317 let enhanced_fn_block =
319 generate_enhanced_function_body(&provider_args.compose, ¶ms, fn_block);
320
321 let interval_impl = generate_interval_impl(&provider_args);
323 let cache_expiration_impl = generate_cache_expiration_impl(&provider_args);
324 let stale_time_impl = generate_stale_time_impl(&provider_args);
325
326 let common_struct = generate_common_struct_and_const(&info);
328
329 if params.is_empty() {
331 Ok(quote! {
333 #common_struct
334
335 impl #struct_name {
336 #fn_vis async fn call() -> Result<#output_type, #error_type> {
337 #enhanced_fn_block
338 }
339 }
340
341 impl ::dioxus_provider::hooks::Provider<()> for #struct_name {
342 type Output = #output_type;
343 type Error = #error_type;
344
345 fn run(&self, _param: ()) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send {
346 Self::call()
347 }
348
349 #interval_impl
350 #cache_expiration_impl
351 #stale_time_impl
352 }
353 })
354 } else if params.len() == 1 {
355 let param = ¶ms[0];
357 let param_name = ¶m.name;
358 let param_type = ¶m.ty;
359
360 Ok(quote! {
361 #common_struct
362
363 impl #struct_name {
364 #fn_vis async fn call(#param_name: #param_type) -> Result<#output_type, #error_type> {
365 #enhanced_fn_block
366 }
367 }
368
369 impl ::dioxus_provider::hooks::Provider<#param_type> for #struct_name {
370 type Output = #output_type;
371 type Error = #error_type;
372
373 fn run(&self, #param_name: #param_type) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send {
374 Self::call(#param_name)
375 }
376
377 #interval_impl
378 #cache_expiration_impl
379 #stale_time_impl
380 }
381 })
382 } else {
383 let param_names: Vec<_> = params.iter().map(|p| &p.name).collect();
385 let param_types: Vec<_> = params.iter().map(|p| &p.ty).collect();
386 let tuple_type = quote! { (#(#param_types,)*) };
387
388 Ok(quote! {
389 #common_struct
390
391 impl #struct_name {
392 #fn_vis async fn call(#(#param_names: #param_types,)*) -> Result<#output_type, #error_type> {
393 #enhanced_fn_block
394 }
395 }
396
397 impl ::dioxus_provider::hooks::Provider<#tuple_type> for #struct_name {
398 type Output = #output_type;
399 type Error = #error_type;
400
401 fn run(&self, params: #tuple_type) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send {
402 let (#(#param_names,)*) = params;
403 Self::call(#(#param_names,)*)
404 }
405
406 #interval_impl
407 #cache_expiration_impl
408 #stale_time_impl
409 }
410 })
411 }
412}
413
414fn generate_mutation(input_fn: ItemFn, mutation_args: MutationArgs) -> Result<TokenStream2> {
415 let info = extract_provider_info(&input_fn)?;
416
417 let ProviderInfo {
418 fn_vis,
419 fn_block,
420 output_type,
421 error_type,
422 struct_name,
423 fn_name: _fn_name,
424 ..
425 } = &info;
426
427 let enhanced_fn_block = generate_enhanced_function_body(&[], &[], fn_block);
428 let invalidation_impl = generate_invalidation_impl(&mutation_args);
429 let common_struct = generate_common_struct_and_const(&info);
430
431 let raw_params = extract_all_params(&input_fn)?;
432 let has_optimistic = mutation_args.optimistic.is_some();
433 let (input_params, context_param, data_param) =
434 split_mutation_params(raw_params.clone(), output_type, has_optimistic)?;
435
436 let is_auto_apply = has_optimistic && data_param.is_some();
438
439 let call_params: Vec<_> = raw_params
441 .iter()
442 .map(|p| {
443 let name = &p.name;
444 if let Some(ctx) = &context_param && ctx.name == p.name {
445 let data_ty = &ctx.data_ty;
446 let error_ty = &ctx.error_ty;
447 quote! { #name: ::dioxus_provider::mutation::MutationContext<'_, #data_ty, #error_ty> }
448 } else {
449 let ty = &p.ty;
450 quote! { #name: #ty }
451 }
452 })
453 .collect();
454
455 let call_signature = quote! { #fn_vis async fn call(#(#call_params),*) -> Result<#output_type, #error_type> {
456 #enhanced_fn_block
457 } };
458
459 let input_count = input_params.len();
460 let input_type = build_input_type(&input_params);
461
462 let data_param_name = data_param.as_ref().map(|p| &p.name);
463
464 let call_args_builder = |ctx_ident: Option<&syn::Ident>,
465 auto_apply_data_expr: Option<TokenStream2>|
466 -> Vec<TokenStream2> {
467 raw_params
468 .iter()
469 .map(|param| {
470 if let Some(ctx) = ctx_ident {
472 if param.name == *ctx {
473 return quote! { #ctx };
474 }
475 }
476 if let Some(data_name) = data_param_name {
478 if param.name == *data_name {
479 if let Some(ref data_expr) = auto_apply_data_expr {
480 return data_expr.clone();
481 }
482 }
483 }
484 let name = ¶m.name;
486 quote! { #name }
487 })
488 .collect()
489 };
490
491 let context_ident = context_param.as_ref().map(|ctx| ctx.name.clone());
492 let context_data_ty = context_param.as_ref().map(|ctx| ctx.data_ty.clone());
493 let context_error_ty = context_param.as_ref().map(|ctx| ctx.error_ty.clone());
494
495 let optimistic_impl = if let Some(optimistic_expr) = &mutation_args.optimistic {
496 let optimistic_call = match input_params.len() {
498 0 => quote! { (#optimistic_expr)(&mut updated) },
499 1 => quote! { (#optimistic_expr)(&mut updated, input) },
500 _ => {
501 let names: Vec<_> = input_params.iter().map(|p| &p.name).collect();
502 quote! {
503 let (#(ref #names,)*) = *input;
504 (#optimistic_expr)(&mut updated, #(#names,)*)
505 }
506 }
507 };
508
509 quote! {
510 fn optimistic_updates_with_current(
511 &self,
512 input: &#input_type,
513 current_data: Option<&Result<Self::Output, Self::Error>>,
514 ) -> Vec<(String, Result<Self::Output, Self::Error>)> {
515 let keys = self.invalidates();
516 if keys.is_empty() {
517 return Vec::new();
518 }
519
520 if let Some(Ok(current)) = current_data {
521 let mut updated = current.clone();
522 #optimistic_call;
523
524 let mut results = Vec::with_capacity(keys.len());
525 for key in keys {
526 results.push((key, Ok(updated.clone())));
527 }
528 results
529 } else {
530 Vec::new()
531 }
532 }
533 }
534 } else {
535 quote! {}
536 };
537
538 let (mutate_signature, mutate_body) = {
539 let (signature, mut prelude) = match input_count {
540 0 => (
541 quote! { fn mutate(&self, _input: ()) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send },
542 Vec::<TokenStream2>::new(),
543 ),
544 1 => {
545 let param = &input_params[0];
546 let name = ¶m.name;
547 let ty = ¶m.ty;
548 (
549 quote! { fn mutate(&self, #name: #ty) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send },
550 Vec::<TokenStream2>::new(),
551 )
552 }
553 _ => {
554 let names: Vec<_> = input_params.iter().map(|p| &p.name).collect();
555 (
556 quote! { fn mutate(&self, input: #input_type) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send },
557 vec![quote! { let (#(#names),*) = input; }],
558 )
559 }
560 };
561
562 if !is_auto_apply {
564 if let (Some(ctx_ident), Some(data_ty), Some(err_ty)) = (
565 context_ident.as_ref(),
566 context_data_ty.as_ref(),
567 context_error_ty.as_ref(),
568 ) {
569 prelude.push(quote! { let #ctx_ident = ::dioxus_provider::mutation::MutationContext::<'static, #data_ty, #err_ty>::new(None); });
570 }
571 }
572
573 let call_args = if is_auto_apply {
574 call_args_builder(None, Some(quote! { Default::default() }))
576 } else {
577 call_args_builder(context_ident.as_ref(), None)
579 };
580
581 let call_expr = quote! { Self::call(#(#call_args),*) };
582 let body = quote! { async move { #(#prelude)* #call_expr.await } };
583 (signature, body)
584 };
585
586 let (mutate_with_current_signature, mutate_with_current_body) = {
587 let (signature, mut prelude) = match input_count {
588 0 => (
589 quote! { fn mutate_with_current(
590 &self,
591 _input: (),
592 current_data: Option<&Result<Self::Output, Self::Error>>,
593 ) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send },
594 Vec::<TokenStream2>::new(),
595 ),
596 1 => {
597 let param = &input_params[0];
598 let name = ¶m.name;
599 let ty = ¶m.ty;
600 (
601 quote! { fn mutate_with_current(
602 &self,
603 #name: #ty,
604 current_data: Option<&Result<Self::Output, Self::Error>>,
605 ) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send },
606 Vec::<TokenStream2>::new(),
607 )
608 }
609 _ => {
610 let names: Vec<_> = input_params.iter().map(|p| &p.name).collect();
611 (
612 quote! { fn mutate_with_current(
613 &self,
614 input: #input_type,
615 current_data: Option<&Result<Self::Output, Self::Error>>,
616 ) -> impl ::std::future::Future<Output = Result<Self::Output, Self::Error>> + Send },
617 vec![quote! { let (#(#names),*) = input; }],
618 )
619 }
620 };
621
622 let call_args = if is_auto_apply {
623 prelude.push(quote! {
626 let __auto_apply_data = if let Some(Ok(current)) = current_data {
627 current.clone()
628 } else {
629 Default::default()
631 };
632 });
633
634 call_args_builder(None, Some(quote! { __auto_apply_data }))
635 } else {
636 if let Some(ctx_ident) = context_ident.as_ref() {
638 prelude.push(quote! { let #ctx_ident = ::dioxus_provider::mutation::MutationContext::new(current_data); });
639 }
640 call_args_builder(context_ident.as_ref(), None)
641 };
642
643 let call_expr = quote! { Self::call(#(#call_args),*) };
644 let body = quote! { async move { #(#prelude)* #call_expr.await } };
645 (signature, body)
646 };
647
648 let has_optimistic_impl = if has_optimistic {
649 quote! {
650 fn has_optimistic(&self) -> bool {
651 true
652 }
653 }
654 } else {
655 quote! {}
656 };
657
658 let mutation_impl = quote! {
659 impl ::dioxus_provider::mutation::Mutation<#input_type> for #struct_name {
660 type Output = #output_type;
661 type Error = #error_type;
662
663 #mutate_signature {
664 #mutate_body
665 }
666
667 #mutate_with_current_signature {
668 #mutate_with_current_body
669 }
670
671 #optimistic_impl
672
673 #invalidation_impl
674
675 #has_optimistic_impl
676 }
677 };
678
679 Ok(quote! {
680 #common_struct
681
682 impl #struct_name {
683 #call_signature
684 }
685
686 #mutation_impl
687 })
688}
689
690fn generate_duration_impl(method_name: &str, duration: Option<Duration>) -> TokenStream2 {
692 if let Some(duration) = duration {
693 let duration_secs = duration.as_secs();
694 let method_ident = syn::Ident::new(method_name, proc_macro2::Span::call_site());
695
696 quote! {
697 fn #method_ident(&self) -> Option<::std::time::Duration> {
698 Some(::std::time::Duration::from_secs(#duration_secs))
699 }
700 }
701 } else {
702 quote! {}
703 }
704}
705
706fn generate_interval_impl(provider_args: &ProviderArgs) -> TokenStream2 {
708 generate_duration_impl("interval", provider_args.interval)
709}
710
711fn generate_cache_expiration_impl(provider_args: &ProviderArgs) -> TokenStream2 {
713 generate_duration_impl("cache_expiration", provider_args.cache_expiration)
714}
715
716fn generate_stale_time_impl(provider_args: &ProviderArgs) -> TokenStream2 {
718 generate_duration_impl("stale_time", provider_args.stale_time)
719}
720
721fn generate_invalidation_impl(mutation_args: &MutationArgs) -> TokenStream2 {
723 if mutation_args.invalidates.is_empty() {
724 quote! {}
725 } else {
726 let provider_calls: Vec<_> = mutation_args
727 .invalidates
728 .iter()
729 .map(|provider_fn| {
730 quote! {
731 ::dioxus_provider::mutation::provider_cache_key_simple(#provider_fn())
732 }
733 })
734 .collect();
735
736 quote! {
737 fn invalidates(&self) -> Vec<String> {
738 vec![#(#provider_calls,)*]
739 }
740 }
741 }
742}
743
744struct ProviderInfo {
746 fn_vis: syn::Visibility,
747 fn_attrs: Vec<syn::Attribute>,
748 fn_block: Box<syn::Block>,
749 output_type: Type,
750 error_type: Type,
751 struct_name: syn::Ident,
752 fn_name: syn::Ident,
753}
754
755#[derive(Clone)]
757struct ParamInfo {
758 name: syn::Ident,
759 ty: Type,
760}
761
762#[derive(Clone)]
763struct ContextInfo {
764 name: syn::Ident,
765 data_ty: Type,
766 error_ty: Type,
767}
768
769fn parse_context_type(ty: &Type) -> Option<(Type, Type)> {
770 if let Type::Path(type_path) = ty {
771 if let Some(segment) = type_path.path.segments.last() {
772 if segment.ident == "MutationContext" {
773 if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
774 if args.args.len() == 2 {
775 let mut iter = args.args.iter();
776 let data_ty = match iter.next()? {
777 syn::GenericArgument::Type(ty) => ty.clone(),
778 _ => return None,
779 };
780 let error_ty = match iter.next()? {
781 syn::GenericArgument::Type(ty) => ty.clone(),
782 _ => return None,
783 };
784 return Some((data_ty, error_ty));
785 }
786 }
787 }
788 }
789 }
790 None
791}
792
793#[allow(dead_code)]
794fn split_params(params: Vec<ParamInfo>) -> Result<(Vec<ParamInfo>, Option<ContextInfo>)> {
795 let mut input_params = Vec::new();
796 let mut context_param = None;
797
798 for param in params {
799 if let Some((data_ty, error_ty)) = parse_context_type(¶m.ty) {
800 if context_param.is_some() {
801 return Err(syn::Error::new_spanned(
802 param.ty,
803 "Only one MutationContext parameter is allowed",
804 ));
805 }
806 context_param = Some(ContextInfo {
807 name: param.name,
808 data_ty,
809 error_ty,
810 });
811 } else {
812 input_params.push(param);
813 }
814 }
815
816 Ok((input_params, context_param))
817}
818
819fn split_mutation_params(
821 params: Vec<ParamInfo>,
822 output_type: &Type,
823 has_optimistic: bool,
824) -> Result<(Vec<ParamInfo>, Option<ContextInfo>, Option<ParamInfo>)> {
825 let mut input_params = Vec::new();
826 let mut context_param = None;
827 let mut data_param = None;
828
829 for param in params {
830 if let Some((data_ty, error_ty)) = parse_context_type(¶m.ty) {
831 if context_param.is_some() {
832 return Err(syn::Error::new_spanned(
833 param.ty,
834 "Only one MutationContext parameter is allowed",
835 ));
836 }
837 context_param = Some(ContextInfo {
838 name: param.name,
839 data_ty,
840 error_ty,
841 });
842 } else {
843 input_params.push(param);
844 }
845 }
846
847 if has_optimistic && context_param.is_none() && !input_params.is_empty() {
849 if let Some(last_param) = input_params.last() {
851 if types_equal(&last_param.ty, output_type) {
852 data_param = input_params.pop();
853 }
854 }
855 }
856
857 Ok((input_params, context_param, data_param))
858}
859
860fn types_equal(ty1: &Type, ty2: &Type) -> bool {
862 ty1 == ty2
863}
864
865fn extract_provider_info(input_fn: &ItemFn) -> Result<ProviderInfo> {
867 let fn_name = input_fn.sig.ident.clone();
868 let fn_vis = input_fn.vis.clone();
869 let fn_attrs = input_fn.attrs.clone();
870 let fn_block = input_fn.block.clone();
871
872 let (output_type, error_type) = extract_result_types(&input_fn.sig.output)?;
873 let struct_name = syn::Ident::new(
874 &to_pascal_case(&fn_name.to_string()),
875 proc_macro2::Span::call_site(),
876 );
877
878 Ok(ProviderInfo {
879 fn_vis,
880 fn_attrs,
881 fn_block,
882 output_type,
883 error_type,
884 struct_name,
885 fn_name,
886 })
887}
888
889fn generate_common_struct_and_const(info: &ProviderInfo) -> TokenStream2 {
891 let struct_name = &info.struct_name;
892 let fn_attrs = &info.fn_attrs;
893 let fn_name = &info.fn_name;
894
895 quote! {
896 #[derive(Clone, PartialEq)]
897 #(#fn_attrs)*
898 pub struct #struct_name;
899
900 impl Default for #struct_name {
901 fn default() -> Self {
902 Self
903 }
904 }
905
906 pub fn #fn_name() -> #struct_name {
908 #struct_name
909 }
910 }
911}
912
913fn extract_all_params(input_fn: &ItemFn) -> Result<Vec<ParamInfo>> {
915 let mut params = Vec::new();
916
917 for input in &input_fn.sig.inputs {
918 match input {
919 FnArg::Typed(PatType { pat, ty, .. }) => {
920 if let Pat::Ident(pat_ident) = &**pat {
921 params.push(ParamInfo {
922 name: pat_ident.ident.clone(),
923 ty: (**ty).clone(),
924 });
925 } else {
926 return Err(syn::Error::new_spanned(
927 pat,
928 "Only simple parameter names are supported",
929 ));
930 }
931 }
932 FnArg::Receiver(_) => {
933 return Err(syn::Error::new_spanned(
934 input,
935 "Methods with self parameter are not supported",
936 ));
937 }
938 }
939 }
940
941 Ok(params)
942}
943
944fn build_input_type(params: &[ParamInfo]) -> TokenStream2 {
946 match params.len() {
947 0 => quote! { () },
948 1 => {
949 let ty = ¶ms[0].ty;
950 quote! { #ty }
951 }
952 _ => {
953 let types: Vec<_> = params.iter().map(|p| &p.ty).collect();
954 quote! { (#(#types,)*) }
955 }
956 }
957}
958
959fn extract_result_types(return_type: &ReturnType) -> Result<(Type, Type)> {
961 match return_type {
962 ReturnType::Default => Err(syn::Error::new_spanned(
963 return_type,
964 "Provider functions must return Result<T, E>",
965 )),
966 ReturnType::Type(_, ty) => {
967 if let Type::Path(type_path) = &**ty {
968 if let Some(segment) = type_path.path.segments.last() {
969 if segment.ident == "Result" {
970 if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
971 if args.args.len() == 2 {
972 let mut args_iter = args.args.iter();
973
974 let output_type = match args_iter.next().unwrap() {
975 syn::GenericArgument::Type(ty) => ty.clone(),
976 _ => {
977 return Err(syn::Error::new_spanned(
978 args,
979 "Result must have type arguments",
980 ));
981 }
982 };
983
984 let error_type = match args_iter.next().unwrap() {
985 syn::GenericArgument::Type(ty) => ty.clone(),
986 _ => {
987 return Err(syn::Error::new_spanned(
988 args,
989 "Result must have type arguments",
990 ));
991 }
992 };
993
994 return Ok((output_type, error_type));
995 }
996 }
997 }
998 }
999 }
1000
1001 Err(syn::Error::new_spanned(
1002 return_type,
1003 "Provider functions must return Result<T, E>",
1004 ))
1005 }
1006 }
1007}
1008
1009fn to_pascal_case(s: &str) -> String {
1011 let mut result = String::new();
1012 let mut capitalize_next = true;
1013
1014 for c in s.chars() {
1015 if c == '_' {
1016 capitalize_next = true;
1017 } else if capitalize_next {
1018 result.push(c.to_ascii_uppercase());
1019 capitalize_next = false;
1020 } else {
1021 result.push(c);
1022 }
1023 }
1024
1025 result
1026}
1027
1028fn validate_composition_requirements(
1030 compose_providers: &[syn::Ident],
1031 params: &[ParamInfo],
1032) -> Result<()> {
1033 if !params.is_empty() {
1035 validate_clone_requirements(params)?;
1036 }
1037
1038 validate_provider_existence(compose_providers)?;
1040
1041 Ok(())
1042}
1043
1044fn validate_clone_requirements(params: &[ParamInfo]) -> Result<()> {
1046 for param in params {
1047 let param_type = ¶m.ty;
1048 let param_name = ¶m.name;
1049
1050 let _clone_check = quote! {
1053 const _: fn() = || {
1054 fn assert_clone<T: Clone>() {}
1055 assert_clone::<#param_type>();
1056 };
1057 };
1058
1059 }
1063
1064 Ok(())
1065}
1066
1067fn validate_provider_existence(compose_providers: &[syn::Ident]) -> Result<()> {
1069 for provider in compose_providers {
1074 let _existence_check = quote! {
1076 const _: fn() = || {
1077 let _ = #provider;
1079 };
1080 };
1081 }
1082
1083 Ok(())
1084}
1085
1086fn generate_enhanced_function_body(
1088 compose_providers: &[syn::Ident],
1089 params: &[ParamInfo],
1090 original_block: &syn::Block,
1091) -> syn::Block {
1092 let mut statements = Vec::new();
1093
1094 if !compose_providers.is_empty() {
1096 let composition_statements = generate_composition_statements(compose_providers, params);
1097 statements.extend(composition_statements);
1098 }
1099
1100 statements.extend(original_block.stmts.clone());
1102
1103 syn::Block {
1104 brace_token: original_block.brace_token,
1105 stmts: statements,
1106 }
1107}
1108
1109fn generate_composition_statements(
1111 compose_providers: &[syn::Ident],
1112 params: &[ParamInfo],
1113) -> Vec<syn::Stmt> {
1114 if compose_providers.is_empty() {
1115 return vec![];
1116 }
1117
1118 let mut statements = Vec::new();
1119
1120 statements.extend(generate_validation_statements(compose_providers, params));
1122
1123 let result_vars: Vec<_> = compose_providers
1125 .iter()
1126 .map(|provider| {
1127 syn::Ident::new(
1128 &format!("__dioxus_composed_{provider}_result"),
1129 proc_macro2::Span::call_site(),
1130 )
1131 })
1132 .collect();
1133
1134 if params.is_empty() {
1136 let provider_calls: Vec<_> = compose_providers
1138 .iter()
1139 .map(|provider| {
1140 quote! {
1141 async { #provider().run(()).await }
1142 }
1143 })
1144 .collect();
1145
1146 let join_stmt: syn::Stmt = syn::parse_quote! {
1147 let (#(#result_vars,)*) = ::futures::join!(
1148 #(#provider_calls,)*
1149 );
1150 };
1151 statements.push(join_stmt);
1152 } else if params.len() == 1 {
1153 let param_name = ¶ms[0].name;
1155 let param_type = ¶ms[0].ty;
1156
1157 let provider_calls: Vec<_> = compose_providers
1158 .iter()
1159 .map(|provider| {
1160 quote! {
1161 async {
1162 let param: #param_type = #param_name.clone();
1164 #provider().run(param).await
1165 }
1166 }
1167 })
1168 .collect();
1169
1170 let join_stmt: syn::Stmt = syn::parse_quote! {
1171 let (#(#result_vars,)*) = ::futures::join!(
1172 #(#provider_calls,)*
1173 );
1174 };
1175 statements.push(join_stmt);
1176 } else {
1177 let param_names: Vec<_> = params.iter().map(|p| &p.name).collect();
1179 let param_types: Vec<_> = params.iter().map(|p| &p.ty).collect();
1180
1181 let provider_calls: Vec<_> = compose_providers
1182 .iter()
1183 .map(|provider| {
1184 quote! {
1185 async {
1186 let params: (#(#param_types,)*) = (#(#param_names.clone(),)*);
1188 #provider().run(params).await
1189 }
1190 }
1191 })
1192 .collect();
1193
1194 let join_stmt: syn::Stmt = syn::parse_quote! {
1195 let (#(#result_vars,)*) = ::futures::join!(
1196 #(#provider_calls,)*
1197 );
1198 };
1199 statements.push(join_stmt);
1200 }
1201
1202 statements
1203}
1204
1205fn generate_validation_statements(
1207 compose_providers: &[syn::Ident],
1208 params: &[ParamInfo],
1209) -> Vec<syn::Stmt> {
1210 let mut statements = Vec::new();
1211
1212 if !params.is_empty() {
1214 for param in params {
1215 let param_type = ¶m.ty;
1216 let param_name = ¶m.name;
1217
1218 let clone_check: syn::Stmt = syn::parse_quote! {
1220 const _: () = {
1221 fn __dioxus_provider_assert_clone<T: ::std::clone::Clone>() {}
1222 fn __dioxus_provider_validate_parameter_clone() {
1223 __dioxus_provider_assert_clone::<#param_type>();
1224 }
1225 };
1226 };
1227 statements.push(clone_check);
1228 }
1229 }
1230
1231 for provider in compose_providers {
1233 let existence_check: syn::Stmt = syn::parse_quote! {
1235 const _: () = {
1236 fn __dioxus_provider_validate_existence() {
1237 let _provider_exists = #provider;
1239 }
1240 };
1241 };
1242 statements.push(existence_check);
1243 }
1244
1245 statements
1246}