#![doc = include_str!("../README.md")]
use proc_macro::TokenStream;
use quote::{ToTokens as _, format_ident, quote};
use syn::{
AngleBracketedGenericArguments, Data, DataStruct, Error, Fields, FieldsUnnamed, PathArguments,
Type, TypePath, spanned::Spanned as _,
};
#[cfg(test)]
mod tests;
#[proc_macro_attribute]
pub fn derive(attr: TokenStream, item: TokenStream) -> TokenStream {
do_output(do_derive(attr.into(), item.into()))
}
fn do_output(res: Result<(proc_macro2::TokenStream, Vec<Error>), Error>) -> TokenStream {
match res {
Err(err) => err.to_compile_error().into(),
Ok((out, errors)) => {
let compiler_errors = errors.iter().map(Error::to_compile_error);
let output = quote! {
#out
#( #compiler_errors )*
};
output.into()
}
}
}
fn do_derive(
attr: proc_macro2::TokenStream,
item: proc_macro2::TokenStream,
) -> Result<(proc_macro2::TokenStream, Vec<Error>), Error> {
let mut errors = Vec::new();
let attr = syn::parse2::<syn::Ident>(attr)?;
if attr != "From" {
errors.push(Error::new(attr.span(), "expected `From`"));
}
let item = syn::parse2::<syn::DeriveInput>(item)?;
let Data::Struct(DataStruct { fields, .. }) = &item.data else {
return Err(Error::new(item.span(), "expected a `struct`"));
};
let Fields::Unnamed(FieldsUnnamed { unnamed, .. }) = fields else {
return Err(Error::new(
fields.span(),
"expected a tuple struct `struct Foo(_)`",
));
};
let unnamed_error = Error::new(
unnamed.span(),
"expected newtype to contain a single type, like: `struct Foo(Bar)`",
);
let Some(field) = unnamed.first() else {
return Err(unnamed_error);
};
if unnamed.len() != 1 {
errors.push(unnamed_error);
}
let is_try_new = item
.attrs
.iter()
.flat_map(|attr| attr.meta.require_list())
.find(|list| {
list.path
.segments
.first()
.is_some_and(|first_segment| first_segment.ident == "nutype")
})
.is_some_and(|list| {
list.tokens.clone().into_iter().any(|token| {
if let proc_macro2::TokenTree::Ident(ident) = token {
ident == proc_macro2::Ident::new("validate", ident.span())
} else {
false
}
})
});
let (passed_to_constructor, from_ty) = if let Type::Path(TypePath { path, .. }) = &field.ty
&& let Some(last_segment) = path.segments.last()
&& last_segment.ident == "NonZero"
&& let PathArguments::AngleBracketed(AngleBracketedGenericArguments { args, .. }) =
&last_segment.arguments
&& args.len() == 1
&& let Some(syn::GenericArgument::Type(Type::Path(TypePath { path, .. }))) = args.first()
&& path.segments.len() == 1
&& let Some(path) = path.segments.first()
&& let Some(non_zero) = parse_numeric_primitive(&path.ident.to_string(), true)
{
let non_zero = format_ident!("{non_zero}");
(
quote! { ::core::num::NonZero::<#non_zero>::new(value).unwrap() },
quote! { #non_zero },
)
} else if let Type::Path(TypePath { path, .. }) = &field.ty
&& let Some(last_segment) = path.segments.last()
&& let Some(non_zero_ty_upper) = last_segment.ident.to_string().strip_prefix("NonZero")
&& let Some(non_zero_ty_lower) = parse_numeric_primitive(non_zero_ty_upper, false)
{
let non_zero_ty = format_ident!("NonZero{non_zero_ty_upper}");
let non_zero_ty_lower = format_ident!("{}", non_zero_ty_lower);
(
quote! { ::core::num::#non_zero_ty::new(value).unwrap() },
quote! { #non_zero_ty_lower },
)
} else {
(quote!(value), field.ty.to_token_stream())
};
let into_ty = &item.ident;
let constructor = if is_try_new {
quote! { Self::try_new(#passed_to_constructor).unwrap() }
} else {
quote! { Self::new(#passed_to_constructor) }
};
let from_impl = quote! {
#[cfg(test)]
impl<T: Into<#from_ty>> ::core::convert::From<T> for #into_ty {
fn from(value: T) -> Self {
let value: #from_ty = value.into();
#constructor
}
}
};
Ok((
quote! {
#item
#from_impl
},
errors,
))
}
fn parse_numeric_primitive(ty: &str, is_lower: bool) -> Option<String> {
fn convert(ty: &str, (signed_in, unsigned_in): (char, char)) -> Option<String> {
if let Some(suffix) = ty.strip_prefix(signed_in) {
if is_valid_suffix(suffix) {
return Some(format!("i{suffix}"));
}
} else if let Some(suffix) = ty.strip_prefix(unsigned_in) {
if is_valid_suffix(suffix) {
return Some(format!("u{suffix}"));
}
}
None
}
fn is_valid_suffix(s: &str) -> bool {
matches!(s, "8" | "16" | "32" | "64" | "128" | "size")
}
let in_prefixes = if is_lower { ('i', 'u') } else { ('I', 'U') };
convert(ty, in_prefixes)
}