use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{FnArg, ItemFn, Pat, ReturnType, Type};
use crate::task_attr::TaskAttr;
pub fn generate_task(attr: &TaskAttr, func: &ItemFn) -> syn::Result<TokenStream> {
let fn_name = &func.sig.ident;
let vis = &func.vis;
let struct_name = format_ident!("{}", to_pascal_case(&fn_name.to_string()));
if func.sig.asyncness.is_none() {
return Err(syn::Error::new_spanned(
func.sig.fn_token,
"#[task] function must be async",
));
}
let mut params = func.sig.inputs.iter();
let first = params.next().ok_or_else(|| {
syn::Error::new_spanned(
&func.sig,
"#[task] function must have at least one parameter: ctx: &TaskContext",
)
})?;
match first {
FnArg::Typed(pat_type) => {
if let Type::Reference(_) = pat_type.ty.as_ref() {
} else {
return Err(syn::Error::new_spanned(
pat_type,
"first parameter must be a reference to TaskContext (e.g., ctx: &TaskContext)",
));
}
}
FnArg::Receiver(_) => {
return Err(syn::Error::new_spanned(
first,
"#[task] function cannot have self parameter",
));
}
}
let mut field_names = Vec::new();
let mut field_types = Vec::new();
for param in params {
match param {
FnArg::Typed(pat_type) => {
if let Pat::Ident(ident) = pat_type.pat.as_ref() {
field_names.push(ident.ident.clone());
field_types.push(pat_type.ty.as_ref().clone());
} else {
return Err(syn::Error::new_spanned(
pat_type,
"expected named parameter",
));
}
}
FnArg::Receiver(_) => {
return Err(syn::Error::new_spanned(param, "unexpected self parameter"));
}
}
}
let task_name = fn_name.to_string();
let queue = attr.queue.as_deref().unwrap_or("default");
let max_retries = attr.max_retries.unwrap_or(3);
let output_type = match &func.sig.output {
ReturnType::Default => quote! { () },
ReturnType::Type(_, ty) => {
extract_result_inner(ty)
}
};
let body = &func.block;
let output = quote! {
#[derive(Debug, ::serde::Serialize, ::serde::Deserialize)]
#vis struct #struct_name {
#( pub #field_names: #field_types, )*
}
#[::async_trait::async_trait]
impl ::kojin_core::Task for #struct_name {
const NAME: &'static str = #task_name;
const QUEUE: &'static str = #queue;
const MAX_RETRIES: u32 = #max_retries;
type Output = #output_type;
async fn run(&self, ctx: &::kojin_core::TaskContext) -> ::kojin_core::TaskResult<Self::Output> {
let Self { #( ref #field_names, )* } = *self;
#( let #field_names = #field_names.clone(); )*
#body
}
}
impl #struct_name {
pub fn new(#( #field_names: #field_types ),*) -> Self {
Self { #( #field_names, )* }
}
}
};
Ok(output)
}
fn to_pascal_case(s: &str) -> String {
s.split('_')
.map(|word| {
let mut chars = word.chars();
match chars.next() {
None => String::new(),
Some(c) => c.to_uppercase().to_string() + chars.as_str(),
}
})
.collect()
}
fn extract_result_inner(ty: &Type) -> TokenStream {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
let name = segment.ident.to_string();
if name == "TaskResult" || name == "Result" {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(syn::GenericArgument::Type(inner)) = args.args.first() {
return quote! { #inner };
}
}
}
}
}
quote! { #ty }
}