use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::quote;
use syn::{
DeriveInput, GenericArgument, Ident, ItemFn, LitStr, PathArguments, ReturnType, Token, Type,
parse::Parse, parse::ParseStream, parse_macro_input,
};
struct CommandAttr {
name: Option<String>,
capability: Option<String>,
}
impl Parse for CommandAttr {
fn parse(input: ParseStream) -> syn::Result<Self> {
let mut attr = CommandAttr {
name: None,
capability: None,
};
if input.is_empty() {
return Ok(attr);
}
loop {
let key: Ident = input.parse()?;
if key == "name" {
let _: Token![=] = input.parse()?;
let name: LitStr = input.parse()?;
attr.name = Some(name.value());
} else if key == "capability" {
let _: Token![=] = input.parse()?;
let cap: LitStr = input.parse()?;
attr.capability = Some(cap.value());
} else {
return Err(syn::Error::new(
key.span(),
"unsupported `#[command]` key; supported keys: `name`, `capability`",
));
}
if input.parse::<Token![,]>().is_err() {
break;
}
}
Ok(attr)
}
}
#[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 docs: Vec<String> = func
.attrs
.iter()
.filter_map(|attr| {
if attr.path().is_ident("doc")
&& let syn::Meta::NameValue(nv) = &attr.meta
&& let syn::Expr::Lit(syn::ExprLit {
lit: syn::Lit::Str(s),
..
}) = &nv.value
{
return Some(s.value().trim().to_string());
}
None
})
.collect();
let doc_comment = docs.join("\n");
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();
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 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: TokenStream2 = if let Some(cap) = &attr.capability {
quote! {
#[allow(non_upper_case_globals, dead_code)]
const #capability_ident: Option<&str> = Some(#cap);
}
} else {
quote! {
#[allow(non_upper_case_globals, dead_code)]
const #capability_ident: Option<&str> = None;
}
};
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 expanded = quote! {
#inner_func
#vis fn #fn_name(#outer_input_arg) -> rustra::Result<#output_type> {
#(#state_bindings)*
#inner_invocation
}
#capability_const
#[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()
}
fn extract_state_inner(ty: &Type) -> Option<Type> {
let Type::Path(type_path) = ty else {
return None;
};
let segment = type_path.path.segments.last()?;
if segment.ident == "State"
&& let PathArguments::AngleBracketed(args) = &segment.arguments
&& let Some(GenericArgument::Type(inner_ty)) = args.args.first()
{
return Some(inner_ty.clone());
}
None
}
struct RegisterInput {
builder: syn::Expr,
commands: Vec<Ident>,
}
impl Parse for RegisterInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let builder: syn::Expr = input.parse()?;
let _: Token![,] = input.parse()?;
let mut commands = Vec::new();
loop {
let name: Ident = input.parse()?;
commands.push(name);
if input.parse::<Token![,]>().is_err() {
break;
}
}
Ok(RegisterInput { builder, commands })
}
}
#[proc_macro]
pub fn register(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as RegisterInput);
if input.commands.is_empty() {
return syn::Error::new(
proc_macro2::Span::call_site(),
"register! requires at least one command function after the builder expression",
)
.to_compile_error()
.into();
}
let builder = &input.builder;
let chain: TokenStream2 = input
.commands
.iter()
.map(|fn_name| {
let meta_ident = Ident::new(
&format!("__RUstra_meta_{}", fn_name),
proc_macro2::Span::call_site(),
);
let cap_ident = Ident::new(
&format!("__RUstra_cap_{}", fn_name),
proc_macro2::Span::call_site(),
);
quote! {
.command(#meta_ident, #fn_name)
.require_capability_if(#meta_ident, #cap_ident)
}
})
.collect();
let expanded = quote! {
#builder #chain
};
expanded.into()
}
fn extract_result_inner(ty: &Type) -> Option<TokenStream2> {
let Type::Path(type_path) = ty else {
return None;
};
let segment = type_path.path.segments.last()?;
if segment.ident != "Result" {
return None;
}
let PathArguments::AngleBracketed(args) = &segment.arguments else {
return None;
};
let GenericArgument::Type(inner_ty) = args.args.first()? else {
return None;
};
Some(quote! { #inner_ty })
}
#[proc_macro_attribute]
pub fn bridge_type(_attr: TokenStream, item: TokenStream) -> TokenStream {
let mut input = parse_macro_input!(item as DeriveInput);
input.attrs.push(syn::parse_quote! {
#[derive(Debug, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]
});
let has_serde_rename = input.attrs.iter().any(|attr| {
if !attr.path().is_ident("serde") {
return false;
}
let Ok(nested) = attr.parse_args_with(
syn::punctuated::Punctuated::<syn::MetaNameValue, syn::Token![,]>::parse_terminated,
) else {
return false;
};
nested.iter().any(|nv| nv.path.is_ident("rename_all"))
});
if !has_serde_rename {
input.attrs.push(syn::parse_quote! {
#[serde(rename_all = "camelCase")]
});
}
quote! { #input }.into()
}
struct BuildInput {
package_name: LitStr,
commands: Vec<Ident>,
}
impl Parse for BuildInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let package_name: LitStr = input.parse()?;
let _: Token![,] = input.parse()?;
let mut commands = Vec::new();
loop {
if input.is_empty() {
break;
}
let name: Ident = input.parse()?;
commands.push(name);
if input.parse::<Token![,]>().is_err() {
break;
}
}
if commands.is_empty() {
return Err(syn::Error::new(
package_name.span(),
"build! requires at least one command function after the package name",
));
}
Ok(BuildInput {
package_name,
commands,
})
}
}
#[proc_macro]
pub fn build(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as BuildInput);
let package_name = &input.package_name;
let chain: TokenStream2 = input
.commands
.iter()
.map(|fn_name| {
let meta_ident = Ident::new(
&format!("__RUstra_meta_{}", fn_name),
proc_macro2::Span::call_site(),
);
let cap_ident = Ident::new(
&format!("__RUstra_cap_{}", fn_name),
proc_macro2::Span::call_site(),
);
quote! {
.command(#meta_ident, #fn_name)
.require_capability_if(#meta_ident, #cap_ident)
}
})
.collect();
let expanded = quote! {
rustra::Package::builder(#package_name) #chain
};
expanded.into()
}
fn snake_to_lower_camel(name: &str) -> String {
let mut output = String::new();
let mut uppercase_next = false;
for character in name.chars() {
if character == '_' || character == '-' || character == '.' {
uppercase_next = true;
continue;
}
if output.is_empty() {
output.push(character.to_ascii_lowercase());
} else if uppercase_next {
output.push(character.to_ascii_uppercase());
uppercase_next = false;
} else {
output.push(character);
}
}
output
}