subtest-impl 0.0.1

Implementation detail of the `subtest` crate
Documentation
use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::punctuated::Punctuated;
use syn::{Attribute, Block, FnArg, Item, ItemFn, ReturnType, Signature, Stmt, Token};

pub fn expand_subtest_main_fn(args: TokenStream, input: TokenStream) -> TokenStream {
    expand_subtest_main_fn_fallible(args, input).unwrap_or_else(|err| err.to_compile_error())
}

fn expand_subtest_main_fn_fallible(
    args: TokenStream,
    input: TokenStream,
) -> Result<TokenStream, syn::Error> {
    if !args.is_empty() {
        return Err(syn::Error::new_spanned(args, "expected no arguments"));
    }

    let input_fn: ItemFn = syn::parse2(input)?;

    let main_subtest = Subtest::new(
        input_fn,
        vec![],
        &[],
        &Punctuated::new(),
        &ReturnType::Default,
    )?;

    Ok(main_subtest.render())
}

struct Subtest {
    function: ItemFn,
    subtests: Vec<Subtest>,
}

impl Subtest {
    fn new(
        input_fn: ItemFn,
        parent_fn_statements: Vec<Stmt>,
        parent_fn_attrs: &[Attribute],
        parent_fn_params: &Punctuated<FnArg, Token![,]>,
        parent_fn_return_type: &ReturnType,
    ) -> Result<Self, syn::Error> {
        // If the subtest fn does not specify any attributes (#[subtest] itself excluded),
        // inherit attributes from the parent test fn
        let attrs = if input_fn.attrs.is_empty() {
            parent_fn_attrs.to_vec()
        } else {
            input_fn.attrs
        };

        // Inherit function parameters if the subtest fn does not specify any
        let fn_params = if input_fn.sig.inputs.is_empty() {
            parent_fn_params.clone()
        } else {
            input_fn.sig.inputs
        };

        // Inherit function return type if the subtest fn does not specify any
        let fn_return_type = if matches!(input_fn.sig.output, ReturnType::Default) {
            parent_fn_return_type.clone()
        } else {
            input_fn.sig.output
        };

        let mut function = ItemFn {
            attrs,
            vis: input_fn.vis,
            sig: Signature {
                inputs: fn_params,
                output: fn_return_type,
                ..input_fn.sig
            },
            block: Box::new(Block {
                brace_token: input_fn.block.brace_token,
                // inherit all preceding statements from the parent
                stmts: parent_fn_statements,
            }),
        };

        let mut subtests = Vec::new();

        for statement in input_fn.block.stmts {
            match statement {
                Stmt::Item(Item::Fn(nested_fn)) => {
                    subtests.push(Subtest::new(
                        check_and_remove_subtest_attr(nested_fn)?,
                        function.block.stmts.clone(),
                        &function.attrs,
                        &function.sig.inputs,
                        &function.sig.output,
                    )?);
                }
                other => {
                    function.block.stmts.push(other);
                }
            }
        }

        Ok(Self { function, subtests })
    }

    fn render(self) -> TokenStream {
        let Self { function, subtests } = self;

        let subtest_module = if subtests.is_empty() {
            None
        } else {
            let module_name = format_ident!("{}_subtests", function.sig.ident);
            let rendered_subtests = subtests.into_iter().map(Subtest::render);

            Some(quote! {
                mod #module_name {
                    use super::*;
                    #(#rendered_subtests)*
                }
            })
        };

        quote! {
            #function
            #subtest_module
        }
    }
}

fn check_and_remove_subtest_attr(mut from_fn: ItemFn) -> Result<ItemFn, syn::Error> {
    let mut subtest_attr_found = false;
    let mut validation_error = None;

    from_fn.attrs.retain(|attr| {
        if attr.meta.path().is_ident("subtest") {
            subtest_attr_found = true;
            if validation_error.is_none() {
                validation_error = attr.meta.require_path_only().err().map(|_| {
                    syn::Error::new_spanned(attr, "expected #[subtest] with no arguments")
                });
            }
            false
        } else {
            true
        }
    });

    if !subtest_attr_found {
        return Err(syn::Error::new_spanned(
            from_fn,
            "function is missing the #[subtest] attribute",
        ));
    }

    if let Some(validation_error) = validation_error {
        return Err(validation_error);
    }

    Ok(from_fn)
}