xabi-macros 0.1.2

Procedural macros for xabi
Documentation
use proc_macro2::TokenStream as TokenStream2;
use quote::{format_ident, quote};
use syn::{Error, GenericParam, ItemStruct};

use crate::type_shape::{
    XabiValueContext, replace_lifetimes_with_static, validate_xabi_value_type,
};

pub(crate) fn expand_data(attr: TokenStream2, item: TokenStream2) -> syn::Result<TokenStream2> {
    if !attr.is_empty() {
        return Err(Error::new_spanned(
            attr,
            "`#[xabi::data]` does not accept options",
        ));
    }

    let item_struct = syn::parse2::<ItemStruct>(item)?;
    if item_struct
        .generics
        .params
        .iter()
        .any(|param| !matches!(param, GenericParam::Lifetime(_)))
    {
        return Err(Error::new_spanned(
            &item_struct.generics,
            "xabi data types only support lifetime parameters for borrowed input handles",
        ));
    }

    let syn::Fields::Named(fields) = &item_struct.fields else {
        return Err(Error::new_spanned(
            &item_struct.fields,
            "xabi data types must use named fields",
        ));
    };
    for field in &fields.named {
        validate_xabi_value_type(&field.ty, XabiValueContext::DataField)?;
    }

    let vis = &item_struct.vis;
    let ident = &item_struct.ident;
    let (impl_generics, ty_generics, where_clause) = item_struct.generics.split_for_impl();
    let wire_ident = format_ident!("XabiV1Data{}", ident);
    let field_idents = fields
        .named
        .iter()
        .map(|field| field.ident.as_ref().expect("named field"))
        .collect::<Vec<_>>();
    let wire_field_idents = field_idents
        .iter()
        .map(|ident| wire_field_ident(ident))
        .collect::<Vec<_>>();
    let field_tys = fields
        .named
        .iter()
        .map(|field| &field.ty)
        .collect::<Vec<_>>();
    let wire_field_tys = field_tys
        .iter()
        .map(|ty| replace_lifetimes_with_static(ty))
        .collect::<Vec<_>>();
    let field_available_arms = fields
        .named
        .iter()
        .map(|field| {
            let ident = field.ident.as_ref().expect("named field");
            quote! {
                stringify!(#ident) => {
                    self.size == Self::FULL_SIZE
                }
            }
        })
        .collect::<Vec<_>>();
    let constructor_args = fields
        .named
        .iter()
        .map(|field| {
            let ident = field.ident.as_ref().expect("named field");
            let ty = &field.ty;
            if is_string_type(ty) {
                quote!(#ident: impl Into<#ty>)
            } else {
                quote!(#ident: #ty)
            }
        })
        .collect::<Vec<_>>();
    let constructor_fields = fields
        .named
        .iter()
        .map(|field| {
            let ident = field.ident.as_ref().expect("named field");
            let ty = &field.ty;
            if is_string_type(ty) {
                quote!(#ident: #ident.into())
            } else {
                quote!(#ident)
            }
        })
        .collect::<Vec<_>>();

    Ok(quote! {
        #item_struct

        #[repr(C)]
        #[derive(Clone, Copy)]
        #vis struct #wire_ident {
            pub size: usize,
            pub abi_version: u32,
            #(pub #wire_field_idents: <#wire_field_tys as ::xabi::XabiType>::Wire,)*
        }

        impl #wire_ident {
            pub const ABI_VERSION: u32 = ::xabi::ABI_VERSION;
            pub const FULL_SIZE: usize = std::mem::size_of::<Self>();
            pub const MIN_SIZE: usize = Self::FULL_SIZE;

            pub fn validate(&self) -> ::xabi::Result<()> {
                ::xabi::validate_exact_size(
                    self.size,
                    Self::FULL_SIZE,
                    stringify!(#wire_ident),
                )?;
                ::xabi::validate_abi_version(
                    self.abi_version,
                    Self::ABI_VERSION,
                    stringify!(#wire_ident),
                )?;
                Ok(())
            }

            pub fn field_available(&self, field: &str) -> bool {
                match field {
                    #(#field_available_arms,)*
                    _ => false,
                }
            }
        }

        impl #impl_generics #ident #ty_generics #where_clause {
            #[allow(clippy::too_many_arguments)]
            pub fn new(#(#constructor_args),*) -> Self {
                Self {
                    #(#constructor_fields,)*
                }
            }
        }

        impl #impl_generics ::xabi::XabiType for #ident #ty_generics #where_clause {
            type Wire = #wire_ident;
            const WIRE_TYPE_NAME: &'static str = stringify!(#wire_ident);

            fn into_wire(self) -> Self::Wire {
                let mut wire = std::mem::MaybeUninit::<#wire_ident>::zeroed();
                unsafe {
                    let wire_ptr = wire.as_mut_ptr();
                    std::ptr::addr_of_mut!((*wire_ptr).size)
                        .write(std::mem::size_of::<#wire_ident>());
                    std::ptr::addr_of_mut!((*wire_ptr).abi_version)
                        .write(#wire_ident::ABI_VERSION);
                    #(std::ptr::addr_of_mut!((*wire_ptr).#wire_field_idents)
                        .write(::xabi::XabiType::into_wire(self.#field_idents));)*
                    wire.assume_init()
                }
            }

            unsafe fn from_wire(wire: *const Self::Wire) -> ::xabi::Result<Self> {
                let wire = unsafe {
                    wire.as_ref()
                        .ok_or(::xabi::Error::NullPointer(concat!(stringify!(#wire_ident), " pointer")))?
                };
                wire.validate()?;
                #(
                    if !wire.field_available(stringify!(#field_idents)) {
                        return Err(::xabi::Error::AbiMismatch(format!(
                            "{} is missing required field {}",
                            stringify!(#wire_ident),
                            stringify!(#field_idents),
                        )));
                    }
                )*
                Ok(Self {
                    #(#field_idents: unsafe {
                        <#field_tys as ::xabi::XabiType>::from_wire(
                            std::ptr::addr_of!(wire.#wire_field_idents)
                        )
                    }?,)*
                })
            }

            fn collect_xabi_layout(collector: &mut dyn ::xabi::XabiLayoutCollector) {
                #(<#wire_field_tys as ::xabi::XabiType>::collect_xabi_layout(collector);)*
                const __XABI_FIELDS: &[::xabi::XabiFieldLayout] = &[
                    ::xabi::XabiFieldLayout::new(
                        "size",
                        std::mem::offset_of!(#wire_ident, size),
                        "usize",
                    ),
                    ::xabi::XabiFieldLayout::new(
                        "abi_version",
                        std::mem::offset_of!(#wire_ident, abi_version),
                        "u32",
                    ),
                    #(
                        ::xabi::XabiFieldLayout::new(
                            stringify!(#field_idents),
                            std::mem::offset_of!(#wire_ident, #wire_field_idents),
                            <#wire_field_tys as ::xabi::XabiType>::WIRE_TYPE_NAME,
                        ),
                    )*
                ];
                collector.push(::xabi::XabiLayoutItem::Type(::xabi::XabiTypeLayout::new(
                    concat!(module_path!(), "::", stringify!(#wire_ident)),
                    ::xabi::XabiLayoutStability::Fixed,
                    std::mem::size_of::<#wire_ident>(),
                    std::mem::align_of::<#wire_ident>(),
                    __XABI_FIELDS,
                )));
            }

            fn retain_module(
                &mut self,
                module: &std::sync::Arc<::xabi::ModuleHandle>,
            ) {
                #(<#field_tys as ::xabi::XabiType>::retain_module(
                    &mut self.#field_idents,
                    module,
                );)*
            }
        }
    })
}

fn wire_field_ident(ident: &syn::Ident) -> syn::Ident {
    match ident.to_string().as_str() {
        "size" | "abi_version" => format_ident!("__xabi_field_{}", ident),
        _ => ident.clone(),
    }
}

fn is_string_type(ty: &syn::Type) -> bool {
    let syn::Type::Path(path) = ty else {
        return false;
    };
    path.path
        .segments
        .last()
        .map(|segment| segment.ident == "String")
        .unwrap_or(false)
}