Skip to main content

subtest_impl/
lib.rs

1use proc_macro2::TokenStream;
2use quote::{format_ident, quote};
3use syn::punctuated::Punctuated;
4use syn::{Attribute, Block, FnArg, Item, ItemFn, ReturnType, Signature, Stmt, Token};
5
6pub fn expand_subtest_main_fn(args: TokenStream, input: TokenStream) -> TokenStream {
7    expand_subtest_main_fn_fallible(args, input).unwrap_or_else(|err| err.to_compile_error())
8}
9
10fn expand_subtest_main_fn_fallible(
11    args: TokenStream,
12    input: TokenStream,
13) -> Result<TokenStream, syn::Error> {
14    if !args.is_empty() {
15        return Err(syn::Error::new_spanned(args, "expected no arguments"));
16    }
17
18    let input_fn: ItemFn = syn::parse2(input)?;
19
20    let main_subtest = Subtest::new(
21        input_fn,
22        vec![],
23        &[],
24        &Punctuated::new(),
25        &ReturnType::Default,
26    )?;
27
28    Ok(main_subtest.render())
29}
30
31struct Subtest {
32    function: ItemFn,
33    subtests: Vec<Subtest>,
34}
35
36impl Subtest {
37    fn new(
38        input_fn: ItemFn,
39        parent_fn_statements: Vec<Stmt>,
40        parent_fn_attrs: &[Attribute],
41        parent_fn_params: &Punctuated<FnArg, Token![,]>,
42        parent_fn_return_type: &ReturnType,
43    ) -> Result<Self, syn::Error> {
44        // If the subtest fn does not specify any attributes (#[subtest] itself excluded),
45        // inherit attributes from the parent test fn
46        let attrs = if input_fn.attrs.is_empty() {
47            parent_fn_attrs.to_vec()
48        } else {
49            input_fn.attrs
50        };
51
52        // Inherit function parameters if the subtest fn does not specify any
53        let fn_params = if input_fn.sig.inputs.is_empty() {
54            parent_fn_params.clone()
55        } else {
56            input_fn.sig.inputs
57        };
58
59        // Inherit function return type if the subtest fn does not specify any
60        let fn_return_type = if matches!(input_fn.sig.output, ReturnType::Default) {
61            parent_fn_return_type.clone()
62        } else {
63            input_fn.sig.output
64        };
65
66        let mut function = ItemFn {
67            attrs,
68            vis: input_fn.vis,
69            sig: Signature {
70                inputs: fn_params,
71                output: fn_return_type,
72                ..input_fn.sig
73            },
74            block: Box::new(Block {
75                brace_token: input_fn.block.brace_token,
76                // inherit all preceding statements from the parent
77                stmts: parent_fn_statements,
78            }),
79        };
80
81        let mut subtests = Vec::new();
82
83        for statement in input_fn.block.stmts {
84            match statement {
85                Stmt::Item(Item::Fn(nested_fn)) => {
86                    subtests.push(Subtest::new(
87                        check_and_remove_subtest_attr(nested_fn)?,
88                        function.block.stmts.clone(),
89                        &function.attrs,
90                        &function.sig.inputs,
91                        &function.sig.output,
92                    )?);
93                }
94                other => {
95                    function.block.stmts.push(other);
96                }
97            }
98        }
99
100        Ok(Self { function, subtests })
101    }
102
103    fn render(self) -> TokenStream {
104        let Self { function, subtests } = self;
105
106        let subtest_module = if subtests.is_empty() {
107            None
108        } else {
109            let module_name = format_ident!("{}_subtests", function.sig.ident);
110            let rendered_subtests = subtests.into_iter().map(Subtest::render);
111
112            Some(quote! {
113                mod #module_name {
114                    use super::*;
115                    #(#rendered_subtests)*
116                }
117            })
118        };
119
120        quote! {
121            #function
122            #subtest_module
123        }
124    }
125}
126
127fn check_and_remove_subtest_attr(mut from_fn: ItemFn) -> Result<ItemFn, syn::Error> {
128    let mut subtest_attr_found = false;
129    let mut validation_error = None;
130
131    from_fn.attrs.retain(|attr| {
132        if attr.meta.path().is_ident("subtest") {
133            subtest_attr_found = true;
134            if validation_error.is_none() {
135                validation_error = attr.meta.require_path_only().err().map(|_| {
136                    syn::Error::new_spanned(attr, "expected #[subtest] with no arguments")
137                });
138            }
139            false
140        } else {
141            true
142        }
143    });
144
145    if !subtest_attr_found {
146        return Err(syn::Error::new_spanned(
147            from_fn,
148            "function is missing the #[subtest] attribute",
149        ));
150    }
151
152    if let Some(validation_error) = validation_error {
153        return Err(validation_error);
154    }
155
156    Ok(from_fn)
157}