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 let attrs = if input_fn.attrs.is_empty() {
47 parent_fn_attrs.to_vec()
48 } else {
49 input_fn.attrs
50 };
51
52 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 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 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}