Skip to main content

bake_macros/
lib.rs

1//! Compile ordinary Rust functions into explicit Bake task descriptors.
2use proc_macro::TokenStream;
3use quote::{format_ident, quote};
4use syn::{Attribute, Expr, FnArg, ItemFn, LitStr, Pat, Path, Type, parse_macro_input};
5
6/// Preserve a synchronous function and generate `function_task()` beside it.
7/// Task functions return a Result whose successful value implements Serialize.
8#[proc_macro_attribute]
9pub fn task(attributes: TokenStream, input: TokenStream) -> TokenStream {
10    let function = parse_macro_input!(input as ItemFn);
11    match expand(attributes.into(), function) {
12        Ok(output) => output.into(),
13        Err(error) => error.into_compile_error().into(),
14    }
15}
16
17#[derive(Default)]
18struct ParameterOptions {
19    default: Option<Expr>,
20    named: bool,
21    context: bool,
22    input: bool,
23    help: Option<LitStr>,
24}
25
26fn options(attributes: &mut Vec<Attribute>) -> syn::Result<ParameterOptions> {
27    let mut options = ParameterOptions::default();
28    let mut retained = Vec::new();
29    for attribute in attributes.drain(..) {
30        if !attribute.path().is_ident("bake") {
31            retained.push(attribute);
32            continue;
33        }
34        attribute.parse_nested_meta(|meta| {
35            if meta.path.is_ident("default") {
36                if options.default.is_some() {
37                    return Err(meta.error("duplicate default"));
38                }
39                options.default = Some(meta.value()?.parse()?);
40            } else if meta.path.is_ident("named") {
41                options.named = true;
42            } else if meta.path.is_ident("context") {
43                options.context = true;
44            } else if meta.path.is_ident("input") {
45                options.input = true;
46            } else if meta.path.is_ident("help") {
47                options.help = Some(meta.value()?.parse()?);
48            } else {
49                return Err(meta.error("expected default, named, context, input, or help"));
50            }
51            Ok(())
52        })?;
53    }
54    *attributes = retained;
55    Ok(options)
56}
57
58fn inner_type<'a>(kind: &str, value: &'a Type) -> Option<&'a Type> {
59    let Type::Path(path) = value else { return None };
60    let segment = path.path.segments.last()?;
61    if segment.ident != kind {
62        return None;
63    }
64    let syn::PathArguments::AngleBracketed(arguments) = &segment.arguments else {
65        return None;
66    };
67    if arguments.args.len() != 1 {
68        return None;
69    }
70    match arguments.args.first()? {
71        syn::GenericArgument::Type(value) => Some(value),
72        _ => None,
73    }
74}
75
76fn expand(
77    attributes: proc_macro2::TokenStream,
78    mut function: ItemFn,
79) -> syn::Result<proc_macro2::TokenStream> {
80    use syn::parse::Parser;
81    let mut name = None::<LitStr>;
82    let mut handles_output = false;
83    let mut output_option_seen = false;
84    let mut builtin = false;
85    let mut runtime: Path = syn::parse_quote!(::bake);
86    let parser = syn::meta::parser(|meta| {
87        if meta.path.is_ident("name") {
88            name = Some(meta.value()?.parse()?);
89        } else if meta.path.is_ident("runtime") {
90            runtime = meta.value()?.parse()?;
91        } else if meta.path.is_ident("output") {
92            if output_option_seen {
93                return Err(meta.error("duplicate output option"));
94            }
95            output_option_seen = true;
96            handles_output = true;
97        } else if meta.path.is_ident("builtin") {
98            builtin = true;
99        } else {
100            return Err(meta.error("expected name, runtime, output, or builtin"));
101        }
102        Ok(())
103    });
104    parser.parse2(attributes)?;
105    if function.sig.asyncness.is_some()
106        || function.sig.unsafety.is_some()
107        || function.sig.abi.is_some()
108        || function.sig.variadic.is_some()
109        || !function.sig.generics.params.is_empty()
110        || function.sig.generics.where_clause.is_some()
111    {
112        return Err(syn::Error::new_spanned(
113            &function.sig,
114            "tasks must be safe, synchronous, non-generic Rust functions",
115        ));
116    }
117    let function_name = &function.sig.ident;
118    let descriptor_name = format_ident!("{}_task", function_name);
119    let registration_name = format_ident!(
120        "__BAKE_REGISTER_{}",
121        function_name
122            .to_string()
123            .trim_start_matches("r#")
124            .to_uppercase()
125    );
126    let command_name = name.unwrap_or_else(|| {
127        LitStr::new(
128            function_name.to_string().trim_start_matches("r#"),
129            function_name.span(),
130        )
131    });
132    let visibility = &function.vis;
133    let mut documentation = Vec::new();
134    for attribute in &function.attrs {
135        if attribute.path().is_ident("doc")
136            && let syn::Meta::NameValue(value) = &attribute.meta
137            && let Expr::Lit(value) = &value.value
138            && let syn::Lit::Str(value) = &value.lit
139        {
140            documentation.push(value.value().trim().to_owned());
141        }
142    }
143    let documentation = documentation.join("\n");
144    let conditional_attributes: Vec<_> = function
145        .attrs
146        .iter()
147        .filter(|attribute| {
148            attribute.path().is_ident("cfg") || attribute.path().is_ident("cfg_attr")
149        })
150        .cloned()
151        .collect();
152    let mut parameters = Vec::new();
153    let mut bindings = Vec::new();
154    let mut call_arguments = Vec::new();
155    let invocation_context =
156        format_ident!("__bake_context", span = proc_macro2::Span::mixed_site());
157    let invocation_arguments =
158        format_ident!("__bake_arguments", span = proc_macro2::Span::mixed_site());
159    let mut has_context = false;
160    let mut has_input = false;
161    for input in &mut function.sig.inputs {
162        let FnArg::Typed(argument) = input else {
163            return Err(syn::Error::new_spanned(
164                input,
165                "task methods with self are not supported",
166            ));
167        };
168        let Pat::Ident(pattern) = argument.pat.as_ref() else {
169            return Err(syn::Error::new_spanned(
170                &argument.pat,
171                "task parameters must have simple names",
172            ));
173        };
174        if pattern.by_ref.is_some() || pattern.subpat.is_some() {
175            return Err(syn::Error::new_spanned(
176                pattern,
177                "task parameters must have simple names",
178            ));
179        }
180        let identifier = &pattern.ident;
181        let parameter_name = identifier.to_string().trim_start_matches("r#").to_owned();
182        let settings = options(&mut argument.attrs)?;
183        let parameter_type = argument.ty.as_ref();
184        if settings.input {
185            let is_value = matches!(parameter_type, Type::Path(path)
186                if path.path.segments.last().is_some_and(|segment| {
187                    segment.ident == "Value" && matches!(segment.arguments, syn::PathArguments::None)
188                })
189            );
190            if !is_value
191                || has_input
192                || settings.default.is_some()
193                || settings.named
194                || settings.context
195                || settings.help.is_some()
196            {
197                return Err(syn::Error::new_spanned(
198                    argument,
199                    "use one #[bake(input)] owned bake::Value parameter without other options",
200                ));
201            }
202            has_input = true;
203            bindings.push(
204                quote!(let #identifier: #parameter_type = #invocation_context.previous().clone();),
205            );
206            call_arguments.push(quote!(#identifier));
207            continue;
208        }
209        let is_context = settings.context
210            || (parameter_name == "context" && matches!(parameter_type, Type::Reference(_)));
211        if is_context {
212            if has_context
213                || settings.default.is_some()
214                || settings.named
215                || settings.help.is_some()
216                || settings.input
217            {
218                return Err(syn::Error::new_spanned(
219                    argument,
220                    "use one context parameter without argument options",
221                ));
222            }
223            has_context = true;
224            call_arguments.push(quote!(#invocation_context));
225            continue;
226        }
227        if matches!(parameter_type, Type::Reference(_)) {
228            return Err(syn::Error::new_spanned(
229                parameter_type,
230                "use owned task arguments such as String or PathBuf",
231            ));
232        }
233        let optional_type = inner_type("Option", parameter_type);
234        let repeated_type = inner_type("Vec", parameter_type);
235        if settings.default.is_some() && (optional_type.is_some() || repeated_type.is_some()) {
236            return Err(syn::Error::new_spanned(
237                parameter_type,
238                "Option and Vec already default to None and empty; use a scalar for an explicit default",
239            ));
240        }
241        let parsed_type = optional_type.or(repeated_type).unwrap_or(parameter_type);
242        let mut parameter = quote!(#runtime::Parameter::new::<#parsed_type>(#parameter_name));
243        let binding = if optional_type.is_some() {
244            parameter = quote!(#parameter.named().optional());
245            quote!(#invocation_arguments.optional::<#parsed_type>(#parameter_name)?)
246        } else if repeated_type.is_some() {
247            parameter = quote!(#parameter.repeated());
248            quote!(#invocation_arguments.repeated::<#parsed_type>(#parameter_name)?)
249        } else if let Some(default) = &settings.default {
250            let description = quote!(#default).to_string();
251            parameter = quote!(#parameter.default(#description));
252            // String literals conveniently initialize String/PathBuf/custom types.
253            // Other expressions retain the parameter's type inference (e.g. 1
254            // for a usize) and must produce that parameter's type.
255            let default = if matches!(default, Expr::Lit(value) if matches!(value.lit, syn::Lit::Str(_)))
256            {
257                quote!((#default).into())
258            } else {
259                quote!(#default)
260            };
261            quote!(#invocation_arguments.optional::<#parsed_type>(#parameter_name)?.unwrap_or_else(|| #default))
262        } else {
263            quote!(#invocation_arguments.required::<#parsed_type>(#parameter_name)?)
264        };
265        if settings.named {
266            parameter = quote!(#parameter.named());
267        }
268        if let Some(help) = settings.help {
269            parameter = quote!(#parameter.help(#help));
270        }
271        parameters.push(parameter);
272        bindings.push(quote!(let #identifier: #parameter_type = #binding;));
273        call_arguments.push(quote!(#identifier));
274    }
275    let task = quote! {
276            #runtime::Task::new(#command_name, #documentation, vec![#(#parameters),*], |#invocation_context, #invocation_arguments| {
277                #(#bindings)*
278                let output = #function_name(#(#call_arguments),*).map_err(|error| #runtime::Error::new(error.to_string()))?;
279                #runtime::value(output)
280            })
281    };
282    let task = if handles_output {
283        quote!(#task.handles_output())
284    } else {
285        task
286    };
287    let builtin = syn::LitBool::new(builtin, proc_macro2::Span::call_site());
288    Ok(quote! {
289        #function
290
291        #(#conditional_attributes)*
292        #[doc = "Generated Bake task descriptor."]
293        #visibility fn #descriptor_name() -> #runtime::Task {
294            #task
295        }
296
297        #(#conditional_attributes)*
298        #[doc(hidden)]
299        #[#runtime::__private::linkme::distributed_slice(#runtime::__private::TASK_REGISTRATIONS)]
300        #[linkme(crate = #runtime::__private::linkme)]
301        static #registration_name: #runtime::__private::TaskRegistration =
302            #runtime::__private::TaskRegistration {
303                factory: #descriptor_name,
304                module_path: module_path!(),
305                builtin: #builtin,
306            };
307    })
308}
309
310#[cfg(test)]
311mod expansion_tests;