vmi-macros 0.9.0

Procedural macros for VMI
Documentation
use proc_macro2::{Span, TokenStream};
use quote::{format_ident, quote};
use syn::{
    CapturedParam, Error, GenericParam, Generics, Ident, ItemImpl, Lifetime, Result, Visibility,
    parse_macro_input,
};

use crate::{
    common::{self, __vmi_lifetime},
    method::{FnArgExt, ItemExt, ItemFnExt},
    transform,
};

struct TraitFn {
    os_context_sig: TokenStream,
    os_context_fn: TokenStream,
}

fn generate_trait_fn(
    item_fn: impl ItemFnExt,
    helper_trait_name: &Ident,
    trait_capture_extras: &[CapturedParam],
    impl_capture_extras: &[CapturedParam],
) -> Option<TraitFn> {
    let sig = item_fn.sig();
    let ident = &sig.ident;
    let generics = &sig.generics;
    let return_type = &sig.output;
    //let receiver = sig.receiver()?;

    let (args, arg_names) = common::build_args(sig)?;
    let where_clause = common::build_where_clause(sig);

    // The trait declaration and trait impl introduce in-scope generics
    // beyond the source method (`Self`, `'__vmi`) that Rust 2024 requires
    // to be listed in every `use<...>` bound. Append them here before
    // emitting each variant.
    let mut trait_return_type = return_type.clone();
    transform::extend_precise_captures(&mut trait_return_type, trait_capture_extras);

    let mut impl_return_type = return_type.clone();
    transform::extend_precise_captures(&mut impl_return_type, impl_capture_extras);

    let doc = item_fn.doc();
    let os_context_sig = quote! {
        #(#doc)*
        fn #ident #generics(&self, #(#args),*) #trait_return_type
            #where_clause;
    };

    let doc = item_fn.doc();
    let os_context_fn = quote! {
        #(#doc)*
        fn #ident #generics(&self, #(#args),*) #impl_return_type
            #where_clause
        {
            <<Self as #helper_trait_name>::Os>::#ident(self.state(), #(#arg_names),*)
        }
    };

    Some(TraitFn {
        os_context_sig,
        os_context_fn,
    })
}

fn filter_pub(item: impl ItemFnExt) -> Option<impl ItemFnExt> {
    match item.vis() {
        Some(Visibility::Public(_)) => Some(item),
        _ => None,
    }
}

fn transform_fn_to_trait_fn(
    item_fn: impl ItemFnExt,
    helper_trait_name: &Ident,
    trait_capture_extras: &[CapturedParam],
    impl_capture_extras: &[CapturedParam],
) -> Option<TraitFn> {
    let sig = item_fn.sig();
    let mut inputs = sig.inputs.iter();

    // First argument must be a receiver (`self`).
    //inputs.next()?.receiver()?;

    // Second argument _might_ be of type `&VmiCore<Driver>`.
    if !inputs
        .next()
        .map(|fn_arg| fn_arg.contains("VmiState"))
        .unwrap_or(false)
    {
        return None;
    }

    generate_trait_fn(
        item_fn,
        helper_trait_name,
        trait_capture_extras,
        impl_capture_extras,
    )
}

fn verify_generics(generics: &Generics) -> Result<()> {
    let mut params = generics.params.iter();
    let param = match params.next() {
        Some(param) => param,
        None => {
            return Err(Error::new(
                Span::call_site(),
                "missing generic `Driver` parameter",
            ));
        }
    };

    let ident = match param {
        GenericParam::Type(ty) => &ty.ident,
        _ => return Err(Error::new(Span::call_site(), "expected type parameter")),
    };

    if ident != "Driver" {
        return Err(Error::new(
            Span::call_site(),
            "expected `Driver` type parameter",
        ));
    }

    Ok(())
}

pub fn derive_trait_from_impl(
    args: proc_macro::TokenStream,
    item: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
    let os_context_name = parse_macro_input!(args as Ident);
    let input = parse_macro_input!(item as ItemImpl);

    match verify_generics(&input.generics) {
        Ok(_) => {}
        Err(err) => return proc_macro::TokenStream::from(err.to_compile_error()),
    }

    let struct_type = &input.self_ty;
    let struct_type_raw = quote! { #struct_type }.to_string();
    let struct_type_raw = match struct_type_raw.find(' ') {
        Some(index) => &struct_type_raw[..index],
        None => &struct_type_raw,
    };
    let where_clause = input.generics.where_clause.as_ref();
    let helper_trait_name = format_ident!("__{os_context_name}VmiOsStateExt");

    // The generated trait and its impl introduce in-scope generics beyond
    // what the source method declares. Rust 2024 requires every `use<...>`
    // bound to list them, so we append them when copying the return type.
    //
    // The trait (`trait Foo<'__vmi, Driver>`) adds the `'__vmi` lifetime
    // and carries an implicit `Self` type parameter.
    //
    // The trait impl (`impl<'__vmi, Driver> Foo for VmiOsState<'__vmi, _>`)
    // adds `'__vmi`; `Self` there is an alias for the concrete impl type
    // and must NOT appear in `use<...>`.
    let vmi_param = CapturedParam::Lifetime(Lifetime::new("'__vmi", Span::call_site()));
    let self_param = CapturedParam::Ident(Ident::new("Self", Span::call_site()));

    let trait_capture_extras = vec![vmi_param.clone(), self_param];
    let impl_capture_extras = vec![vmi_param];

    let fns = input
        .items
        .iter()
        .filter_map(ItemExt::as_fn)
        .filter_map(filter_pub)
        .filter_map(|item_fn| {
            transform_fn_to_trait_fn(
                item_fn,
                &helper_trait_name,
                &trait_capture_extras,
                &impl_capture_extras,
            )
        })
        .collect::<Vec<_>>();

    let os_context_sigs = fns.iter().map(|m| &m.os_context_sig);
    let os_context_fns = fns.iter().map(|m| &m.os_context_fn);

    let vmi_lifetime = __vmi_lifetime();

    let expanded = quote! {
        #input

        // A helper trait to expose the Os type from VmiOsState
        trait #helper_trait_name {
            type Os;
        }

        impl<'__vmi, TOs> #helper_trait_name for vmi_core::VmiOsState<'__vmi, TOs>
        where
            TOs: vmi_core::VmiOs,
        {
            type Os = TOs;
        }

        //
        // OS Context
        //

        #[doc = concat!("[`", #struct_type_raw, "`] extensions for the [`VmiContext`].")]
        #[doc = ""]
        #[doc = "[`VmiContext`]: vmi_core::VmiContext"]
        pub trait #os_context_name <#vmi_lifetime, Driver>
            #where_clause
        {
            #(#os_context_sigs)*
        }

        impl<#vmi_lifetime, Driver> #os_context_name <#vmi_lifetime, Driver>
            for vmi_core::VmiOsState<#vmi_lifetime, #struct_type>
            #where_clause
        {
            #(#os_context_fns)*
        }
    };

    proc_macro::TokenStream::from(expanded)
}