use proc_macro::TokenStream;
use quote::{format_ident, quote};
use syn::{parse_macro_input, Attribute, Data, DeriveInput, Fields, Meta};
fn get_crate_path(attrs: &[Attribute]) -> proc_macro2::TokenStream {
for attr in attrs {
if attr.path().is_ident("graph") {
if let Meta::List(list) = &attr.meta {
let tokens = list.tokens.to_string();
for part in tokens.split(',') {
let part = part.trim();
if let Some(rest) = part.strip_prefix("crate") {
let rest = rest.trim();
if let Some(rest) = rest.strip_prefix('=') {
let rest = rest.trim();
if rest.len() >= 2 && rest.starts_with('"') && rest.ends_with('"') {
let path_str = &rest[1..rest.len() - 1];
let path: syn::Path =
syn::parse_str(path_str).expect("Invalid crate path");
return quote! { #path };
}
}
}
}
}
}
}
quote! { packr_abi }
}
#[proc_macro_derive(GraphValue, attributes(graph))]
pub fn derive_graph_value(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let crate_path = get_crate_path(&input.attrs);
let expanded = match &input.data {
Data::Struct(data) => derive_struct(&input, data, &crate_path),
Data::Enum(data) => derive_enum(&input, data, &crate_path),
Data::Union(_) => {
return syn::Error::new_spanned(&input, "GraphValue cannot be derived for unions")
.to_compile_error()
.into();
}
};
expanded.into()
}
fn has_forward_compatible(attrs: &[Attribute]) -> bool {
for attr in attrs {
if attr.path().is_ident("graph") {
if let Meta::List(list) = &attr.meta {
let tokens = list.tokens.to_string();
if tokens.split(',').any(|t| t.trim() == "forward_compatible") {
return true;
}
}
}
}
false
}
fn augmented_generics(generics: &syn::Generics, krate: &proc_macro2::TokenStream) -> syn::Generics {
let mut generics = generics.clone();
let type_idents: Vec<syn::Ident> = generics.type_params().map(|tp| tp.ident.clone()).collect();
if type_idents.is_empty() {
return generics;
}
let where_clause = generics.make_where_clause();
for ident in type_idents {
where_clause
.predicates
.push(syn::parse_quote!(#ident: ::core::convert::Into<#krate::Value>));
where_clause.predicates.push(syn::parse_quote!(
#ident: #krate::__private::TryFrom<#krate::Value, Error = #krate::ConversionError>
));
where_clause
.predicates
.push(syn::parse_quote!(#ident: #krate::KnownValueType));
}
generics
}
fn box_inner(ty: &syn::Type) -> Option<&syn::Type> {
let syn::Type::Path(tp) = ty else {
return None;
};
let seg = tp.path.segments.last()?;
if seg.ident != "Box" {
return None;
}
let syn::PathArguments::AngleBracketed(args) = &seg.arguments else {
return None;
};
args.args.iter().find_map(|a| match a {
syn::GenericArgument::Type(t) => Some(t),
_ => None,
})
}
fn decode_field(
field_type: &syn::Type,
value: proc_macro2::TokenStream,
krate: &proc_macro2::TokenStream,
) -> proc_macro2::TokenStream {
if let Some(inner) = box_inner(field_type) {
quote! {
<#inner as #krate::FromValue>::from_value(#value).map(#krate::__private::Box::new)
}
} else {
quote! {
<#field_type as #krate::FromValue>::from_value(#value)
}
}
}
fn derive_struct(
input: &DeriveInput,
data: &syn::DataStruct,
krate: &proc_macro2::TokenStream,
) -> proc_macro2::TokenStream {
let name = &input.ident;
let generics = augmented_generics(&input.generics, krate);
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let forward_compatible = has_forward_compatible(&input.attrs);
match &data.fields {
Fields::Named(fields) => {
let field_from_value: Vec<_> = fields
.named
.iter()
.map(|f| {
let field_name = f.ident.as_ref().unwrap();
let field_name_str =
get_rename(&f.attrs).unwrap_or_else(|| field_name.to_string());
let field_type = &f.ty;
let decode = decode_field(field_type, quote! { field_value }, krate);
if forward_compatible {
quote! {
#field_name: match fields.iter()
.find(|(name, _)| name == #field_name_str)
.map(|(_, v)| v.clone())
{
#krate::__private::Some(field_value) =>
#decode
.map_err(|e| #krate::ConversionError::FieldError(
#krate::__private::String::from(#field_name_str),
#krate::__private::Box::new(e)
))?,
#krate::__private::None =>
<#field_type as ::core::default::Default>::default(),
}
}
} else {
quote! {
#field_name: {
let field_value = fields.iter()
.find(|(name, _)| name == #field_name_str)
.map(|(_, v)| v.clone())
.ok_or_else(|| #krate::ConversionError::MissingField(
#krate::__private::String::from(#field_name_str)
))?;
#decode
.map_err(|e| #krate::ConversionError::FieldError(
#krate::__private::String::from(#field_name_str),
#krate::__private::Box::new(e)
))?
}
}
}
})
.collect();
let field_count = fields.named.len();
let field_accessors: Vec<_> = fields
.named
.iter()
.map(|f| {
let field_name = f.ident.as_ref().unwrap();
let field_name_str =
get_rename(&f.attrs).unwrap_or_else(|| field_name.to_string());
quote! {
(
#krate::__private::String::from(#field_name_str),
::core::convert::Into::<#krate::Value>::into(value.#field_name)
)
}
})
.collect();
let type_name_str = name.to_string();
let count_check = if forward_compatible {
quote! {}
} else {
quote! {
if fields.len() != #field_count {
return #krate::__private::Err(#krate::ConversionError::WrongFieldCount {
expected: #field_count,
got: fields.len(),
});
}
}
};
quote! {
impl #impl_generics #krate::__private::From<#name #ty_generics> for #krate::Value #where_clause {
fn from(value: #name #ty_generics) -> #krate::Value {
#krate::Value::Record {
type_name: #krate::__private::String::from(#type_name_str),
fields: #krate::__private::vec![
#(#field_accessors),*
],
}
}
}
impl #impl_generics #krate::__private::TryFrom<#krate::Value> for #name #ty_generics #where_clause {
type Error = #krate::ConversionError;
fn try_from(value: #krate::Value) -> #krate::__private::Result<Self, Self::Error> {
match value {
#krate::Value::Record { fields, .. } => {
#count_check
#krate::__private::Ok(Self {
#(#field_from_value),*
})
}
other => #krate::__private::Err(#krate::ConversionError::ExpectedRecord(
#krate::__private::format!("{:?}", other)
)),
}
}
}
impl #impl_generics #krate::KnownValueType for #name #ty_generics #where_clause {
fn known_value_type() -> #krate::ValueType {
#krate::ValueType::Record(
#krate::__private::String::from(#type_name_str)
)
}
}
}
}
Fields::Unnamed(fields) => {
let field_indices: Vec<_> = (0..fields.unnamed.len()).map(syn::Index::from).collect();
let field_from_value: Vec<_> = fields.unnamed.iter().enumerate().map(|(i, f)| {
let field_type = &f.ty;
if forward_compatible {
let decode = decode_field(field_type, quote! { field_value }, krate);
quote! {
match fields.get(#i).cloned() {
#krate::__private::Some(field_value) =>
#decode
.map_err(|e| #krate::ConversionError::IndexError(#i, #krate::__private::Box::new(e)))?,
#krate::__private::None =>
<#field_type as ::core::default::Default>::default(),
}
}
} else {
let decode = decode_field(
field_type,
quote! { fields.get(#i).cloned().ok_or_else(|| #krate::ConversionError::MissingIndex(#i))? },
krate,
);
quote! {
#decode.map_err(|e| #krate::ConversionError::IndexError(#i, #krate::__private::Box::new(e)))?
}
}
}).collect();
let field_count = fields.unnamed.len();
let field_types: Vec<_> = fields.unnamed.iter().map(|f| &f.ty).collect();
let count_check = if forward_compatible {
quote! {}
} else {
quote! {
if fields.len() != #field_count {
return #krate::__private::Err(#krate::ConversionError::WrongFieldCount {
expected: #field_count,
got: fields.len(),
});
}
}
};
quote! {
impl #impl_generics #krate::__private::From<#name #ty_generics> for #krate::Value #where_clause {
fn from(value: #name #ty_generics) -> #krate::Value {
#krate::Value::Tuple(#krate::__private::vec![
#(::core::convert::Into::<#krate::Value>::into(value.#field_indices)),*
])
}
}
impl #impl_generics #krate::__private::TryFrom<#krate::Value> for #name #ty_generics #where_clause {
type Error = #krate::ConversionError;
fn try_from(value: #krate::Value) -> #krate::__private::Result<Self, Self::Error> {
match value {
#krate::Value::Tuple(fields) => {
#count_check
#krate::__private::Ok(Self(
#(#field_from_value),*
))
}
other => #krate::__private::Err(#krate::ConversionError::ExpectedTuple(
#krate::__private::format!("{:?}", other)
)),
}
}
}
impl #impl_generics #krate::KnownValueType for #name #ty_generics #where_clause {
fn known_value_type() -> #krate::ValueType {
#krate::ValueType::Tuple(#krate::__private::vec![
#(<#field_types as #krate::KnownValueType>::known_value_type()),*
])
}
}
}
}
Fields::Unit => {
quote! {
impl #impl_generics #krate::__private::From<#name #ty_generics> for #krate::Value #where_clause {
fn from(_: #name #ty_generics) -> #krate::Value {
#krate::Value::Tuple(#krate::__private::vec![])
}
}
impl #impl_generics #krate::__private::TryFrom<#krate::Value> for #name #ty_generics #where_clause {
type Error = #krate::ConversionError;
fn try_from(value: #krate::Value) -> #krate::__private::Result<Self, Self::Error> {
match value {
#krate::Value::Tuple(fields) if fields.is_empty() => {
#krate::__private::Ok(Self)
}
#krate::Value::Tuple(fields) => {
#krate::__private::Err(#krate::ConversionError::WrongFieldCount {
expected: 0,
got: fields.len(),
})
}
other => #krate::__private::Err(#krate::ConversionError::ExpectedTuple(
#krate::__private::format!("{:?}", other)
)),
}
}
}
impl #impl_generics #krate::KnownValueType for #name #ty_generics #where_clause {
fn known_value_type() -> #krate::ValueType {
#krate::ValueType::Tuple(#krate::__private::vec![])
}
}
}
}
}
}
fn derive_enum(
input: &DeriveInput,
data: &syn::DataEnum,
krate: &proc_macro2::TokenStream,
) -> proc_macro2::TokenStream {
let name = &input.ident;
let type_name_str = name.to_string();
let generics = augmented_generics(&input.generics, krate);
let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
let to_value_arms: Vec<_> = data
.variants
.iter()
.enumerate()
.map(|(default_tag, variant)| {
let variant_name = &variant.ident;
let case_name_str = variant_name.to_string();
let tag = get_tag(&variant.attrs).unwrap_or(default_tag);
match &variant.fields {
Fields::Named(fields) => {
let field_names: Vec<_> = fields
.named
.iter()
.map(|f| f.ident.as_ref().unwrap())
.collect();
let field_to_value: Vec<_> = fields
.named
.iter()
.map(|f| {
let field_name = f.ident.as_ref().unwrap();
let field_name_str =
get_rename(&f.attrs).unwrap_or_else(|| field_name.to_string());
quote! {
(
#krate::__private::String::from(#field_name_str),
::core::convert::Into::<#krate::Value>::into(#field_name)
)
}
})
.collect();
quote! {
#name::#variant_name { #(#field_names),* } => {
#krate::Value::Variant {
type_name: #krate::__private::String::from(#type_name_str),
case_name: #krate::__private::String::from(#case_name_str),
tag: #tag,
payload: #krate::__private::vec![
#krate::Value::Record {
type_name: #krate::__private::String::from(#case_name_str),
fields: #krate::__private::vec![#(#field_to_value),*],
}
],
}
}
}
}
Fields::Unnamed(fields) => {
let field_names: Vec<_> = (0..fields.unnamed.len())
.map(|i| format_ident!("f{}", i))
.collect();
quote! {
#name::#variant_name(#(#field_names),*) => {
#krate::Value::Variant {
type_name: #krate::__private::String::from(#type_name_str),
case_name: #krate::__private::String::from(#case_name_str),
tag: #tag,
payload: #krate::__private::vec![
#(::core::convert::Into::<#krate::Value>::into(#field_names)),*
],
}
}
}
}
Fields::Unit => {
quote! {
#name::#variant_name => {
#krate::Value::Variant {
type_name: #krate::__private::String::from(#type_name_str),
case_name: #krate::__private::String::from(#case_name_str),
tag: #tag,
payload: #krate::__private::vec![],
}
}
}
}
}
})
.collect();
let from_value_arms: Vec<_> = data.variants.iter().enumerate().map(|(default_tag, variant)| {
let variant_name = &variant.ident;
let tag = get_tag(&variant.attrs).unwrap_or(default_tag);
match &variant.fields {
Fields::Named(fields) => {
let field_from_value: Vec<_> = fields.named.iter().map(|f| {
let field_name = f.ident.as_ref().unwrap();
let field_name_str = get_rename(&f.attrs).unwrap_or_else(|| field_name.to_string());
let field_type = &f.ty;
let decode = decode_field(field_type, quote! { field_value }, krate);
quote! {
#field_name: {
let field_value = record_fields.iter()
.find(|(name, _)| name == #field_name_str)
.map(|(_, v)| v.clone())
.ok_or_else(|| #krate::ConversionError::MissingField(
#krate::__private::String::from(#field_name_str)
))?;
#decode
.map_err(|e| #krate::ConversionError::FieldError(
#krate::__private::String::from(#field_name_str),
#krate::__private::Box::new(e)
))?
}
}
}).collect();
quote! {
#tag => {
if payload.len() != 1 {
return #krate::__private::Err(#krate::ConversionError::WrongFieldCount {
expected: 1,
got: payload.len(),
});
}
match &payload[0] {
#krate::Value::Record { fields: record_fields, .. } => {
#krate::__private::Ok(#name::#variant_name {
#(#field_from_value),*
})
}
other => #krate::__private::Err(#krate::ConversionError::ExpectedRecord(
#krate::__private::format!("{:?}", other)
)),
}
}
}
}
Fields::Unnamed(fields) => {
let field_count = fields.unnamed.len();
let field_conversions: Vec<_> = fields.unnamed.iter().enumerate().map(|(i, f)| {
let field_type = &f.ty;
let decode = decode_field(
field_type,
quote! { payload.get(#i).cloned().ok_or_else(|| #krate::ConversionError::MissingIndex(#i))? },
krate,
);
quote! {
#decode.map_err(|e| #krate::ConversionError::IndexError(#i, #krate::__private::Box::new(e)))?
}
}).collect();
quote! {
#tag => {
if payload.len() != #field_count {
return #krate::__private::Err(#krate::ConversionError::WrongFieldCount {
expected: #field_count,
got: payload.len(),
});
}
#krate::__private::Ok(#name::#variant_name(
#(#field_conversions),*
))
}
}
}
Fields::Unit => {
quote! {
#tag => {
if !payload.is_empty() {
return #krate::__private::Err(#krate::ConversionError::UnexpectedPayload);
}
#krate::__private::Ok(#name::#variant_name)
}
}
}
}
}).collect();
let variant_count = data.variants.len();
quote! {
impl #impl_generics #krate::__private::From<#name #ty_generics> for #krate::Value #where_clause {
fn from(value: #name #ty_generics) -> #krate::Value {
match value {
#(#to_value_arms),*
}
}
}
impl #impl_generics #krate::__private::TryFrom<#krate::Value> for #name #ty_generics #where_clause {
type Error = #krate::ConversionError;
fn try_from(value: #krate::Value) -> #krate::__private::Result<Self, Self::Error> {
match value {
#krate::Value::Variant { tag, payload, .. } => {
match tag {
#(#from_value_arms),*
other => #krate::__private::Err(#krate::ConversionError::UnknownTag {
tag: other,
max: #variant_count,
}),
}
}
other => #krate::__private::Err(#krate::ConversionError::ExpectedVariant(
#krate::__private::format!("{:?}", other)
)),
}
}
}
impl #impl_generics #krate::KnownValueType for #name #ty_generics #where_clause {
fn known_value_type() -> #krate::ValueType {
#krate::ValueType::Variant(#krate::__private::String::from(#type_name_str))
}
}
}
}
fn get_rename(attrs: &[Attribute]) -> Option<String> {
for attr in attrs {
if attr.path().is_ident("graph") {
if let Meta::List(list) = &attr.meta {
let tokens = list.tokens.to_string();
if let Some(rest) = tokens.strip_prefix("rename") {
let rest = rest.trim();
if let Some(rest) = rest.strip_prefix('=') {
let rest = rest.trim();
if rest.starts_with('"') && rest.ends_with('"') {
return Some(rest[1..rest.len() - 1].to_string());
}
}
}
}
}
}
None
}
fn get_tag(attrs: &[Attribute]) -> Option<usize> {
for attr in attrs {
if attr.path().is_ident("graph") {
if let Meta::List(list) = &attr.meta {
let tokens = list.tokens.to_string();
if let Some(rest) = tokens.strip_prefix("tag") {
let rest = rest.trim();
if let Some(rest) = rest.strip_prefix('=') {
let rest = rest.trim();
return rest.parse().ok();
}
}
}
}
}
None
}
#[cfg(test)]
mod tests {
use super::{get_crate_path, has_forward_compatible};
#[test]
fn combined_crate_and_forward_compatible() {
let attr: syn::Attribute =
syn::parse_quote!(#[graph(crate = "foo::bar", forward_compatible)]);
let attrs = [attr];
assert_eq!(
get_crate_path(&attrs).to_string(),
quote::quote!(foo::bar).to_string(),
"combined form must parse the crate path, not default to packr_abi"
);
assert!(has_forward_compatible(&attrs));
}
#[test]
fn crate_only_parses_and_no_flag() {
let attrs = [syn::parse_quote!(#[graph(crate = "foo::bar")])];
assert_eq!(
get_crate_path(&attrs).to_string(),
quote::quote!(foo::bar).to_string()
);
assert!(!has_forward_compatible(&attrs));
}
#[test]
fn flag_only_defaults_crate() {
let attrs = [syn::parse_quote!(#[graph(forward_compatible)])];
assert_eq!(
get_crate_path(&attrs).to_string(),
quote::quote!(packr_abi).to_string()
);
assert!(has_forward_compatible(&attrs));
}
}