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