use proc_macro2::{TokenStream, TokenTree};
use quote::{quote, quote_spanned};
use syn::{
Attribute, Fields, FieldsNamed, Ident, Item, ItemEnum, ItemStruct, parse_quote,
punctuated::Punctuated, spanned::Spanned, 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 reorder_result = reorder_struct_fields(item, args);
let error = reorder_result.err().map(|message| compile_error(&item.ident, message));
let field_count = item.fields.len();
let repr = if field_count == 1 { quote!(#[repr(transparent)]) } else { quote!(#[repr(C)]) };
quote! {
#repr
#[derive(::oxc_ast_macros::Ast)]
#item
#error
#assertions
}
}
fn reorder_struct_fields(item: &mut ItemStruct, args: TokenStream) -> Result<(), &'static str> {
if let Some(TokenTree::Ident(ident)) = args.into_iter().next()
&& ident == "foreign"
{
return Ok(());
}
let struct_name = item.ident.to_string();
let Some(struct_details) = STRUCTS.get(&struct_name) else {
return Err("Struct is unknown. Run `just ast` to re-run the codegen.");
};
let Some(field_order) = struct_details.field_order else {
return Ok(());
};
let named = match &mut item.fields {
Fields::Named(FieldsNamed { named, .. }) if named.len() == field_order.len() => named,
_ => {
return Err("Struct has been altered. Run `just ast` to re-run the codegen.");
}
};
let mut fields = named.clone().into_pairs().zip(field_order).collect::<Vec<_>>();
fields.sort_unstable_by_key(|(_, index)| **index);
for field in named.iter_mut() {
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
}));
Ok(())
}
fn assert_generated_derives(attrs: &[Attribute]) -> TokenStream {
let mut assertions = quote!();
for attr in attrs {
if !attr.path().is_ident("generate_derive") {
continue;
}
let Ok(parsed) = attr.parse_args_with(Punctuated::<Ident, Comma>::parse_terminated) else {
continue;
};
for trait_ident in parsed {
let trait_name = trait_ident.to_string();
let Some((trait_path, generics)) = get_trait_crate_and_generics(&trait_name) else {
continue;
};
assertions.extend(quote! {{
trait AssertionTrait: #trait_path #generics {}
impl<T: #trait_ident #generics> AssertionTrait for T {}
}});
}
}
quote! {
const _: () = { #assertions };
}
}
fn compile_error<S: Spanned>(spanned: &S, message: &str) -> TokenStream {
quote_spanned! { spanned.span() => compile_error!(#message); }
}