Skip to main content

bake_macros/
lib.rs

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