use proc_macro::TokenStream as CompilerTokenStream;
use proc_macro2::{Span, TokenStream};
use quote::{format_ident, quote};
use syn::{parse_macro_input, parse_quote, Data, DeriveInput, Fields, Ident, Path};
#[proc_macro_derive(AsStd140)]
pub fn derive_as_std140(input: CompilerTokenStream) -> CompilerTokenStream {
let input = parse_macro_input!(input as DeriveInput);
let expanded = EmitOptions::new("Std140", "std140", 16).emit(input);
CompilerTokenStream::from(expanded)
}
#[proc_macro_derive(AsStd430)]
pub fn derive_as_std430(input: CompilerTokenStream) -> CompilerTokenStream {
let input = parse_macro_input!(input as DeriveInput);
let expanded = EmitOptions::new("Std430", "std430", 0).emit(input);
CompilerTokenStream::from(expanded)
}
struct EmitOptions {
layout_name: Ident,
min_struct_alignment: usize,
mod_path: Path,
trait_path: Path,
as_trait_path: Path,
as_trait_assoc: Ident,
as_trait_method: Ident,
}
impl EmitOptions {
fn new(layout_name: &'static str, mod_name: &'static str, min_struct_alignment: usize) -> Self {
let mod_name = Ident::new(mod_name, Span::call_site());
let layout_name = Ident::new(layout_name, Span::call_site());
let mod_path = parse_quote!(::crevice::#mod_name);
let trait_path = parse_quote!(#mod_path::#layout_name);
let as_trait_name = format_ident!("As{}", layout_name);
let as_trait_path = parse_quote!(#mod_path::#as_trait_name);
let as_trait_assoc = format_ident!("{}Type", layout_name);
let as_trait_method = format_ident!("as_{}", mod_name);
Self {
layout_name,
min_struct_alignment,
mod_path,
trait_path,
as_trait_path,
as_trait_assoc,
as_trait_method,
}
}
fn emit(&self, input: DeriveInput) -> TokenStream {
let min_struct_alignment = self.min_struct_alignment;
let layout_name = &self.layout_name;
let mod_path = &self.mod_path;
let trait_path = &self.trait_path;
let as_trait_path = &self.as_trait_path;
let as_trait_assoc = &self.as_trait_assoc;
let as_trait_method = &self.as_trait_method;
let visibility = input.vis;
let name = input.ident;
let generated_name = format_ident!("{}{}", layout_name, name);
let alignment_mod_name = format_ident!("{}{}Alignment", layout_name, name);
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let fields = match &input.data {
Data::Struct(data) => match &data.fields {
Fields::Named(fields) => fields,
Fields::Unnamed(_) => panic!("Tuple structs are not supported"),
Fields::Unit => panic!("Unit structs are not supported"),
},
Data::Enum(_) | Data::Union(_) => panic!("Only structs are supported"),
};
let align_names: Vec<_> = fields
.named
.iter()
.map(|field| format_ident!("_{}_align", field.ident.as_ref().unwrap()))
.collect();
let alignment_calculators: Vec<_> = fields
.named
.iter()
.enumerate()
.map(|(index, field)| {
let align_name = &align_names[index];
let offset_accumulation =
fields
.named
.iter()
.zip(&align_names)
.take(index)
.map(|(field, align_name)| {
let field_ty = &field.ty;
quote! {
offset += #align_name();
offset += ::std::mem::size_of::<#field_ty>();
}
});
let field_ty = &field.ty;
quote! {
pub const fn #align_name() -> usize {
let mut offset = 0;
#( #offset_accumulation )*
::crevice::internal::align_offset(
offset,
<<#field_ty as #as_trait_path>::#as_trait_assoc as #mod_path::#layout_name>::ALIGNMENT
)
}
}
})
.collect();
let generated_fields: Vec<_> = fields
.named
.iter()
.zip(&align_names)
.map(|(field, align_name)| {
let field_ty = &field.ty;
let field_name = field.ident.as_ref().unwrap();
quote! {
#align_name: [u8; #alignment_mod_name::#align_name()],
#field_name: <#field_ty as #as_trait_path>::#as_trait_assoc,
}
})
.collect();
let field_initializers: Vec<_> = fields
.named
.iter()
.map(|field| {
let field_name = field.ident.as_ref().unwrap();
quote!(#field_name: self.#field_name.#as_trait_method())
})
.collect();
let struct_alignment = fields.named.iter().fold(
quote!(#min_struct_alignment),
|last, field| {
let field_ty = &field.ty;
quote! {
::crevice::internal::max(
<<#field_ty as #as_trait_path>::#as_trait_assoc as #trait_path>::ALIGNMENT,
#last,
)
}
},
);
let type_layout_derive = if cfg!(feature = "test_type_layout") {
quote!(#[derive(::type_layout::TypeLayout)])
} else {
quote!()
};
quote! {
#[allow(non_snake_case)]
mod #alignment_mod_name {
use super::*;
#( #alignment_calculators )*
}
#[derive(Debug, Clone, Copy)]
#type_layout_derive
#[repr(C)]
#visibility struct #generated_name #ty_generics #where_clause {
#( #generated_fields )*
}
unsafe impl #impl_generics ::crevice::internal::bytemuck::Zeroable for #generated_name #ty_generics #where_clause {}
unsafe impl #impl_generics ::crevice::internal::bytemuck::Pod for #generated_name #ty_generics #where_clause {}
unsafe impl #impl_generics #mod_path::#layout_name for #generated_name #ty_generics #where_clause {
const ALIGNMENT: usize = #struct_alignment;
}
impl #impl_generics #as_trait_path for #name #ty_generics #where_clause {
type #as_trait_assoc = #generated_name;
fn #as_trait_method(&self) -> Self::#as_trait_assoc {
Self::#as_trait_assoc {
#( #field_initializers, )*
..::crevice::internal::bytemuck::Zeroable::zeroed()
}
}
}
}
}
}