#[proc_macro_attribute]
pub fn command(attr: TokenStream, item: TokenStream) -> TokenStream {
let attr = parse_macro_input!(attr as CommandAttr);
let func = parse_macro_input!(item as ItemFn);
let doc_comment = command_doc_comment(&func);
let is_async = func.sig.asyncness.is_some();
let fn_name = &func.sig.ident;
let vis = &func.vis;
let output_type = match &func.sig.output {
ReturnType::Type(_, ty) => match extract_result_inner(ty) {
Some(inner) => inner,
None => {
return syn::Error::new_spanned(
ty,
"#[command] must return `Result<O>` where O: Serialize + JsonSchema",
)
.to_compile_error()
.into();
}
},
_ => {
return syn::Error::new_spanned(
&func.sig,
"#[command] must have an explicit return type `Result<O>`",
)
.to_compile_error()
.into();
}
};
struct ParamInfo {
pat: syn::Pat,
ty: Type,
is_state: bool,
state_inner: Option<Type>,
}
let mut params = Vec::new();
for input in &func.sig.inputs {
match input {
syn::FnArg::Receiver(_) => {
return syn::Error::new_spanned(input, "#[command] functions cannot accept `self`")
.to_compile_error()
.into();
}
syn::FnArg::Typed(pat_type) => {
let is_state_inner = extract_state_inner(&pat_type.ty);
let is_state = is_state_inner.is_some();
params.push(ParamInfo {
pat: (*pat_type.pat).clone(),
ty: (*pat_type.ty).clone(),
is_state,
state_inner: is_state_inner,
});
}
}
}
let data_params: Vec<&ParamInfo> = params.iter().filter(|p| !p.is_state).collect();
if data_params.len() > 1 {
return syn::Error::new_spanned(
&func.sig.inputs,
"#[command] supports at most one input data parameter (plus optional State<T> parameters)",
)
.to_compile_error()
.into();
}
let input_type = if let Some(data) = data_params.first() {
let ty = &data.ty;
quote! { #ty }
} else {
quote! { () }
};
let inner_fn_name = Ident::new(
&format!("__rustra_inner_{}", fn_name),
proc_macro2::Span::call_site(),
);
let mut inner_func = func.clone();
inner_func.sig.ident = inner_fn_name.clone();
inner_func.vis = syn::Visibility::Inherited;
let command_name = attr.name.unwrap_or_else(|| {
let raw = fn_name.to_string();
snake_to_lower_camel(raw.trim_end_matches("_command"))
});
let meta_ident = Ident::new(
&format!("__RUstra_meta_{}", fn_name),
proc_macro2::Span::call_site(),
);
let execution_ident = Ident::new(
&format!("__RUstra_execution_{}", fn_name),
proc_macro2::Span::call_site(),
);
let execution = if is_async {
quote! { rustra::CommandExecution::Async }
} else {
quote! { rustra::CommandExecution::Sync }
};
let doc_ident = Ident::new(
&format!("__RUstra_doc_{}", fn_name),
proc_macro2::Span::call_site(),
);
let capability_ident = Ident::new(
&format!("__RUstra_cap_{}", fn_name),
proc_macro2::Span::call_site(),
);
let capability_const = meta_opt_const(
&capability_ident,
quote! { Option<&str> },
attr.capability.as_ref().map(|cap| quote! { #cap }),
);
let platforms_ident = Ident::new(
&format!("__RUstra_platforms_{}", fn_name),
proc_macro2::Span::call_site(),
);
let platforms_const_ty = quote! { Option<&'static [rustra::platform::Platform]> };
let (supported_cfg, unsupported_cfg, platform_paths, platforms_const): (
TokenStream2,
TokenStream2,
Vec<TokenStream2>,
TokenStream2,
) = if let Some(platforms) = &attr.platforms {
let mut os_checks = Vec::new();
let mut paths = Vec::new();
for platform in platforms {
let (os, variant_name) = match platform.as_str() {
"windows" => ("windows", "Windows"),
"macos" => ("macos", "Macos"),
"linux" => ("linux", "Linux"),
"android" => ("android", "Android"),
"ios" => ("ios", "Ios"),
other => {
return syn::Error::new_spanned(
&func.sig.ident,
format!(
"unknown platform '{other}'; supported: windows, macos, linux, android, ios"
),
)
.to_compile_error()
.into();
}
};
os_checks.push(quote! { target_os = #os });
let variant = Ident::new(variant_name, proc_macro2::Span::call_site());
paths.push(quote! { rustra::platform::Platform::#variant });
}
(
quote! { any(#(#os_checks),*) },
quote! { not(any(#(#os_checks),*)) },
paths.clone(),
meta_opt_const(
&platforms_ident,
platforms_const_ty.clone(),
Some(quote! { &[#(#paths),*] }),
),
)
} else {
(
quote! {},
quote! {},
Vec::new(),
meta_opt_const(&platforms_ident, platforms_const_ty, None),
)
};
let errors_ident = Ident::new(
&format!("__RUstra_errors_{}", fn_name),
proc_macro2::Span::call_site(),
);
let errors_const = meta_opt_const(
&errors_ident,
quote! { Option<&'static [rustra::CommandErrorVariant]> },
attr.errors.as_ref().map(|errors| {
let variants = errors
.iter()
.map(|code| quote! { rustra::CommandErrorVariant::new(#code) });
quote! { &[#(#variants),*] }
}),
);
let devices_ident = Ident::new(
&format!("__RUstra_devices_{}", fn_name),
proc_macro2::Span::call_site(),
);
let devices_const = meta_opt_const(
&devices_ident,
quote! { Option<&'static [rustra::device_capabilities::DeviceCapability]> },
attr.devices.as_ref().map(|devices| {
let capabilities = devices
.iter()
.map(|token| quote! { rustra::device_capabilities::DeviceCapability::new(#token) });
quote! { &[#(#capabilities),*] }
}),
);
let wrapper_unsafety: TokenStream2 = if attr.capability.is_some() {
quote! { unsafe }
} else {
quote! {}
};
let register_ident = Ident::new(
&format!("__rustra_register_{}", fn_name),
proc_macro2::Span::call_site(),
);
let register_call: TokenStream2 = if attr.capability.is_some() {
quote! { unsafe { #fn_name(__rustra_input) } }
} else {
quote! { #fn_name(__rustra_input) }
};
let mut state_bindings = Vec::new();
let mut call_args = Vec::new();
for param in ¶ms {
if param.is_state {
let pat = ¶m.pat;
let ty = ¶m.ty;
let inner_ty = param.state_inner.as_ref().unwrap();
state_bindings.push(quote! {
let #pat: #ty = rustra::get_state::<#inner_ty>()
.ok_or_else(|| rustra::RustraError::internal(concat!("State<", stringify!(#inner_ty), "> not managed in package")))?;
});
call_args.push(quote! { #pat });
} else {
call_args.push(quote! { __rustra_input });
}
}
let outer_input_arg = if data_params.is_empty() {
quote! { _: () }
} else {
quote! { __rustra_input: #input_type }
};
let inner_invocation = if is_async {
quote! {
rustra::__private::block_on(async move {
#inner_fn_name(#(#call_args),*).await
})
}
} else {
quote! {
#inner_fn_name(#(#call_args),*)
}
};
let mut stub_func = func.clone();
stub_func.sig.ident = inner_fn_name.clone();
stub_func.vis = syn::Visibility::Inherited;
stub_func.block = syn::parse_quote! {
{
Err(rustra::RustraError::platform_unavailable(
#command_name,
&[#(#platform_paths),*],
))
}
};
let (real_inner, stub_inner): (TokenStream2, TokenStream2) = if attr.platforms.is_some() {
(
quote! { #[cfg(#supported_cfg)] #inner_func },
quote! { #[cfg(#unsupported_cfg)] #[allow(unused_variables)] #stub_func },
)
} else {
(quote! { #inner_func }, quote! {})
};
let expanded = quote! {
#real_inner
#stub_inner
#vis #wrapper_unsafety fn #fn_name(#outer_input_arg) -> rustra::Result<#output_type> {
#(#state_bindings)*
#inner_invocation
}
#[allow(non_upper_case_globals, dead_code)]
const #execution_ident: rustra::CommandExecution = #execution;
#capability_const
#platforms_const
#errors_const
#devices_const
#[doc(hidden)]
fn #register_ident(__rustra_input: #input_type) -> rustra::Result<#output_type> {
#register_call
}
#[allow(non_upper_case_globals, dead_code)]
const #meta_ident: &str = #command_name;
#[allow(non_upper_case_globals, dead_code)]
const #doc_ident: &str = #doc_comment;
#[allow(dead_code)]
const _: () = {
fn _assert_command_bounds<
__I: rustra::__private::CommandInput,
__O: rustra::__private::CommandOutput,
>() {
}
fn _check_command_bounds() {
_assert_command_bounds::<#input_type, #output_type>();
}
};
};
expanded.into()
}