use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::{format_ident, quote};
use syn::{
parse_macro_input, Data, DeriveInput, Fields, GenericArgument, Meta, PathArguments, Type,
};
#[proc_macro_derive(Pod)]
pub fn derive_pod(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let has_repr_c = input.attrs.iter().any(|attr| {
if !attr.path().is_ident("repr") {
return false;
}
let mut found = false;
if let Meta::List(list) = &attr.meta {
let _ = list.parse_nested_meta(|nested| {
if nested.path.is_ident("C") {
found = true;
}
Ok(())
});
}
found
});
if !has_repr_c {
return syn::Error::new_spanned(
&input.ident,
"Pod can only be derived for #[repr(C)] structs",
)
.to_compile_error()
.into();
}
let fields = match &input.data {
Data::Struct(s) => match &s.fields {
Fields::Named(f) => f.named.iter().collect::<Vec<_>>(),
Fields::Unnamed(f) => f.unnamed.iter().collect::<Vec<_>>(),
Fields::Unit => vec![],
},
_ => {
return syn::Error::new_spanned(&input.ident, "Pod can only be derived for structs")
.to_compile_error()
.into();
}
};
let field_assertions = fields.iter().map(|f| {
let ty = &f.ty;
quote! {
const _: () = {
fn _assert_pod<T: photon_ring::Pod>() {}
fn _check() { _assert_pod::<#ty>(); }
};
}
});
let field_types: Vec<_> = fields.iter().map(|f| &f.ty).collect();
let expanded = quote! {
#(#field_assertions)*
const _: () = {
let sum = 0usize #( + core::mem::size_of::<#field_types>() )*;
assert!(
core::mem::size_of::<#name>() == sum,
"Pod cannot be derived for a type with padding: add explicit \
padding fields, or reorder fields so none is inserted",
);
};
unsafe impl #impl_generics photon_ring::Pod for #name #ty_generics #where_clause {}
};
TokenStream::from(expanded)
}
enum FieldKind {
Passthrough,
Bool,
Usize,
Isize,
Option {
wire_ty: proc_macro2::TokenStream,
to_value: proc_macro2::TokenStream,
from_value: proc_macro2::TokenStream,
is_usize_isize: bool,
},
Enum,
Unsupported,
UnsupportedOption(String),
}
fn type_name(ty: &Type) -> Option<String> {
if let Type::Path(p) = ty {
if let Some(seg) = p.path.segments.last() {
return Some(seg.ident.to_string());
}
}
None
}
fn align_rank(ty: &Type) -> u8 {
if let Type::Array(a) = ty {
return align_rank(&a.elem);
}
match type_name(ty).as_deref() {
Some("u128") | Some("i128") => 0,
Some("u64") | Some("i64") | Some("f64") | Some("usize") | Some("isize") => 1,
Some("u32") | Some("f32") => 2,
Some("u16") => 3,
Some("u8") => 4,
_ => 0,
}
}
fn classify(ty: &Type) -> FieldKind {
match ty {
Type::Array(_) => FieldKind::Passthrough,
Type::Path(p) => {
let seg = match p.path.segments.last() {
Some(s) => s,
None => return FieldKind::Unsupported,
};
let id = seg.ident.to_string();
match id.as_str() {
"u8" | "u16" | "u32" | "u64" | "u128" | "i8" | "i16" | "i32" | "i64" | "i128"
| "f32" | "f64" => FieldKind::Passthrough,
"bool" => FieldKind::Bool,
"usize" => FieldKind::Usize,
"isize" => FieldKind::Isize,
"Option" => {
if let PathArguments::AngleBracketed(args) = &seg.arguments {
if let Some(GenericArgument::Type(inner)) = args.args.first() {
let name = type_name(inner).unwrap_or_default();
let opt =
|wire_ty, to_value, from_value, is_usize_isize| FieldKind::Option {
wire_ty,
to_value,
from_value,
is_usize_isize,
};
return match name.as_str() {
"bool" => opt(
quote!(u8),
quote!(if v { 1 } else { 0 }),
quote!(raw != 0),
false,
),
"f32" => opt(
quote!(u32),
quote!(v.to_bits()),
quote!(f32::from_bits(raw)),
false,
),
"f64" => opt(
quote!(u64),
quote!(v.to_bits()),
quote!(f64::from_bits(raw)),
false,
),
"u128" => opt(quote!(u128), quote!(v), quote!(raw), false),
"i128" => {
opt(quote!(u128), quote!(v as u128), quote!(raw as i128), false)
}
"usize" => {
opt(quote!(u64), quote!(v as u64), quote!(raw as usize), true)
}
"isize" => {
opt(quote!(i64), quote!(v as i64), quote!(raw as isize), true)
}
"u8" | "u16" | "u32" | "u64" => {
opt(quote!(u64), quote!(v as u64), quote!(raw as #inner), false)
}
"i8" | "i16" | "i32" | "i64" => {
opt(quote!(i64), quote!(v as i64), quote!(raw as #inner), false)
}
_ => FieldKind::UnsupportedOption(name),
};
}
}
FieldKind::UnsupportedOption(String::new())
}
_ => FieldKind::Unsupported,
}
}
_ => FieldKind::Unsupported,
}
}
#[proc_macro_derive(Message, attributes(photon))]
pub fn derive_message(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let wire_name = format_ident!("{}Wire", name);
let fields = match &input.data {
Data::Struct(s) => match &s.fields {
Fields::Named(f) => f.named.iter().collect::<Vec<_>>(),
_ => {
return syn::Error::new_spanned(
&input.ident,
"Message can only be derived for structs with named fields",
)
.to_compile_error()
.into();
}
},
_ => {
return syn::Error::new_spanned(
&input.ident,
"Message can only be derived for structs",
)
.to_compile_error()
.into();
}
};
let mut wire_fields = Vec::new();
let mut wire_types: Vec<proc_macro2::TokenStream> = Vec::new();
let mut wire_ranks: Vec<u8> = Vec::new();
let mut to_wire = Vec::new();
let mut from_wire = Vec::new();
let mut assertions = Vec::new();
let mut has_enum_fields = false;
let mut has_usize_isize = false;
for field in &fields {
let fname = field.ident.as_ref().unwrap();
let fty = &field.ty;
let is_explicit_enum = field.attrs.iter().any(|attr| {
if attr.path().is_ident("photon") {
if let Ok(meta) = attr.parse_args::<syn::Ident>() {
return meta == "as_enum";
}
}
false
});
let kind = if is_explicit_enum {
FieldKind::Enum
} else {
classify(fty)
};
match kind {
FieldKind::Passthrough => {
assertions.push(quote! {
const _: () = {
fn _assert_pod<T: photon_ring::Pod>() {}
fn _check() { _assert_pod::<#fty>(); }
};
});
wire_fields.push(quote! { pub #fname: #fty });
wire_types.push(quote!(#fty));
wire_ranks.push(align_rank(fty));
to_wire.push(quote! { #fname: src.#fname });
from_wire.push(quote! { #fname: src.#fname });
}
FieldKind::Bool => {
wire_fields.push(quote! { pub #fname: u8 });
wire_types.push(quote!(u8));
wire_ranks.push(4);
to_wire.push(quote! { #fname: if src.#fname { 1 } else { 0 } });
from_wire.push(quote! { #fname: src.#fname != 0 });
}
FieldKind::Usize => {
has_usize_isize = true;
wire_fields.push(quote! { pub #fname: u64 });
wire_types.push(quote!(u64));
wire_ranks.push(1);
to_wire.push(quote! { #fname: src.#fname as u64 });
from_wire.push(quote! { #fname: src.#fname as usize });
}
FieldKind::Isize => {
has_usize_isize = true;
wire_fields.push(quote! { pub #fname: i64 });
wire_types.push(quote!(i64));
wire_ranks.push(1);
to_wire.push(quote! { #fname: src.#fname as i64 });
from_wire.push(quote! { #fname: src.#fname as isize });
}
FieldKind::Option {
wire_ty,
to_value,
from_value,
is_usize_isize,
} => {
if is_usize_isize {
has_usize_isize = true;
}
let value_field = format_ident!("{}_value", fname);
let has_field = format_ident!("{}_has", fname);
wire_fields.push(quote! { pub #value_field: #wire_ty });
wire_types.push(quote!(#wire_ty));
wire_ranks.push(match wire_ty.to_string().as_str() {
"u128" => 0,
"u32" => 2,
"u16" => 3,
"u8" => 4,
_ => 1,
});
wire_fields.push(quote! { pub #has_field: u8 });
wire_types.push(quote!(u8));
wire_ranks.push(4);
to_wire.push(quote! {
#value_field: match src.#fname {
Some(v) => #to_value,
None => 0,
}
});
to_wire.push(quote! {
#has_field: if src.#fname.is_some() { 1 } else { 0 }
});
from_wire.push(quote! {
#fname: if src.#has_field != 0 {
let raw = src.#value_field;
Some(#from_value)
} else {
None
}
});
}
FieldKind::Enum => {
has_enum_fields = true;
wire_fields.push(quote! { pub #fname: u8 });
wire_types.push(quote!(u8));
wire_ranks.push(4);
to_wire.push(quote! { #fname: src.#fname as u8 });
from_wire.push(quote! {
#fname: unsafe { core::mem::transmute::<u8, #fty>(src.#fname) }
});
let msg = format!(
"Message derive: field `{}` has type `{}` which is not 1 byte. \
Enum fields must have #[repr(u8)].",
fname,
quote! { #fty },
);
let msg_lit = syn::LitStr::new(&msg, Span::call_site());
assertions.push(quote! {
const _: () = {
assert!(
core::mem::size_of::<#fty>() == 1,
#msg_lit,
);
};
});
}
FieldKind::Unsupported => {
let msg = format!(
"Unsupported field type `{}`. Use #[photon(as_enum)] for #[repr(u8)] enum fields, \
or convert to a numeric type manually.",
quote!(#fty),
);
return syn::Error::new_spanned(fty, msg).to_compile_error().into();
}
FieldKind::UnsupportedOption(inner_name) => {
let msg = format!(
"Message derive: field `{}` has unsupported type `Option<{}>`. \
Only Option<bool>, Option<integer>, Option<f32>, and Option<f64> \
are supported.",
fname, inner_name,
);
return syn::Error::new_spanned(fty, msg).to_compile_error().into();
}
}
}
if has_usize_isize {
assertions.push(quote! {
const _: () = assert!(
core::mem::size_of::<usize>() <= core::mem::size_of::<u64>(),
"photon-ring Message derive requires usize to fit in u64",
);
});
}
let from_wire_impl = if has_enum_fields {
quote! {
impl #wire_name {
#[inline]
pub unsafe fn into_domain(self) -> #name {
let src = self;
#name {
#(#from_wire),*
}
}
}
}
} else {
quote! {
impl From<#wire_name> for #name {
#[inline]
fn from(src: #wire_name) -> Self {
#name {
#(#from_wire),*
}
}
}
}
};
let wire_struct_doc = if has_enum_fields {
quote! {
}
} else {
quote! {
}
};
let mut ordered: Vec<usize> = (0..wire_fields.len()).collect();
ordered.sort_by_key(|&i| wire_ranks[i]);
let wire_fields: Vec<_> = ordered.iter().map(|&i| wire_fields[i].clone()).collect();
let wire_types: Vec<_> = ordered.iter().map(|&i| wire_types[i].clone()).collect();
let pad_const = format_ident!("__{}_TAIL_PAD", wire_name.to_string().to_uppercase());
let expanded = quote! {
#(#assertions)*
#[doc(hidden)]
const #pad_const: usize = {
let sum = 0usize #( + core::mem::size_of::<#wire_types>() )*;
let align = { let mut a = 1usize; #( { let f = core::mem::align_of::<#wire_types>(); if f > a { a = f; } } )* a };
(align - sum % align) % align
};
#wire_struct_doc
#[repr(C)]
#[derive(Clone, Copy)]
pub struct #wire_name {
#(#wire_fields,)*
pub _pad: [u8; #pad_const],
}
const _: () = {
let sum = 0usize #( + core::mem::size_of::<#wire_types>() )* + #pad_const;
assert!(
core::mem::size_of::<#wire_name>() == sum,
"photon-ring: the generated wire struct has internal padding, which \
is not a valid Pod. The macro orders fields by width, but cannot \
see through a type alias, so one of them landed out of order. \
Use a concrete primitive or array type for that field, or add an \
explicit padding field to close the gap.",
);
};
unsafe impl photon_ring::Pod for #wire_name {}
impl From<#name> for #wire_name {
#[inline]
fn from(src: #name) -> Self {
#wire_name {
#(#to_wire,)*
_pad: [0; #pad_const],
}
}
}
#from_wire_impl
};
TokenStream::from(expanded)
}