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 (args, arg_names) = common::build_args(sig)?;
let where_clause = common::build_where_clause(sig);
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();
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");
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
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;
}
#[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)
}