use std::mem;
use proc_macro2::{TokenStream, TokenTree};
use quote::quote;
use syn::{
Attribute, Fields, FieldsNamed, Ident, Item, ItemEnum, ItemStruct, parse_quote,
punctuated::Punctuated, token::Comma,
};
use crate::generated::{derived_traits::get_trait_crate_and_generics, structs::STRUCTS};
pub fn ast(item: &mut Item, args: TokenStream) -> TokenStream {
match item {
Item::Enum(item) => modify_enum(item),
Item::Struct(item) => modify_struct(item, args),
_ => unreachable!(),
}
}
fn modify_enum(item: &ItemEnum) -> TokenStream {
let repr = if item.variants.iter().any(|var| !matches!(var.fields, Fields::Unit)) {
quote!(#[repr(C, u8)])
} else {
quote!(#[repr(u8)])
};
let assertions = assert_generated_derives(&item.attrs);
quote! {
#repr
#[derive(::oxc_ast_macros::Ast)]
#item
#assertions
}
}
pub struct StructDetails {
pub field_order: Option<&'static [u8]>,
}
fn modify_struct(item: &mut ItemStruct, args: TokenStream) -> TokenStream {
let assertions = assert_generated_derives(&item.attrs);
let item = reorder_struct_fields(item, args).unwrap_or_else(|| quote!(#item));
quote! {
#[repr(C)]
#[derive(::oxc_ast_macros::Ast)]
#item
#assertions
}
}
fn reorder_struct_fields(item: &mut ItemStruct, args: TokenStream) -> Option<TokenStream> {
if let Some(TokenTree::Ident(ident)) = args.into_iter().next() {
if ident == "foreign" {
return None;
}
}
let struct_name = item.ident.to_string();
let field_order = STRUCTS[&struct_name].field_order?;
let fields = mem::replace(&mut item.fields, Fields::Unit);
let Fields::Named(FieldsNamed { brace_token, mut named }) = fields else { unreachable!() };
assert!(
named.len() == field_order.len(),
"Wrong number of fields for `{struct_name}` in `STRUCTS`"
);
let mut fields = named.clone().into_pairs().zip(field_order).collect::<Vec<_>>();
fields.sort_unstable_by_key(|(_, index)| **index);
for field in &mut named {
field.attrs.insert(0, parse_quote!( #[cfg(doc)]));
}
named.extend(fields.into_iter().map(|(mut pair, _)| {
pair.value_mut().attrs.insert(0, parse_quote!( #[cfg(not(doc))]));
pair
}));
item.fields = Fields::Named(FieldsNamed { brace_token, named });
Some(quote!( #item ))
}
fn assert_generated_derives(attrs: &[Attribute]) -> TokenStream {
let assertions = attrs
.iter()
.filter(|attr| attr.path().is_ident("generate_derive"))
.flat_map(parse_attr)
.map(|trait_ident| {
let trait_name = trait_ident.to_string();
let Some((trait_path, generics)) = get_trait_crate_and_generics(&trait_name) else {
panic!("Invalid derive trait(generate_derive): {trait_name}");
};
quote! {{
trait AssertionTrait: #trait_path #generics {}
impl<T: #trait_ident #generics> AssertionTrait for T {}
}}
});
quote!( const _: () = { #(#assertions)* }; )
}
#[inline]
fn parse_attr(attr: &Attribute) -> impl Iterator<Item = Ident> + use<> {
attr.parse_args_with(Punctuated::<Ident, Comma>::parse_terminated)
.expect("`#[generate_derive]` only accepts traits as single segment paths. Found an invalid argument.")
.into_iter()
}