#![forbid(unsafe_code)]
use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::{format_ident, quote};
use syn::{parse_macro_input, Data, DeriveInput, Fields, GenericArgument, PathArguments, Type};
#[proc_macro_derive(Verit, attributes(verit))]
pub fn derive_verit(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
expand(input)
.unwrap_or_else(syn::Error::into_compile_error)
.into()
}
enum Kind {
Scalar(&'static str),
Str,
Bytes,
List(Box<Kind>),
Nested(Box<Type>),
}
const SCALARS: &[&str] = &[
"bool", "u8", "u16", "u32", "u64", "i8", "i16", "i32", "i64", "f32", "f64",
];
fn expand(input: DeriveInput) -> syn::Result<TokenStream2> {
let ident = &input.ident;
let name_str = ident.to_string();
if !input.generics.params.is_empty() {
return Err(syn::Error::new_spanned(
&input.generics,
"#[derive(Verit)] does not support generic types",
));
}
let mode = parse_mode(&input)?;
let mode_expr = match mode {
Mode::Sparse => quote!(::verit::StructMode::Sparse),
Mode::Dense => quote!(::verit::StructMode::Dense),
Mode::Packed => quote!(::verit::StructMode::Packed),
};
let fields = match &input.data {
Data::Struct(s) => match &s.fields {
Fields::Named(named) => &named.named,
_ => {
return Err(syn::Error::new_spanned(
ident,
"#[derive(Verit)] requires a struct with named fields",
))
}
},
_ => {
return Err(syn::Error::new_spanned(
ident,
"#[derive(Verit)] can only be applied to structs",
))
}
};
let mut dt_entries = Vec::new();
let mut pack_stmts = Vec::new();
let mut unpack_inits = Vec::new();
let mut nested_types: Vec<Type> = Vec::new();
for f in fields {
let fname = f.ident.as_ref().unwrap();
let fname_str = fname.to_string();
let id = parse_field_id(f)?;
let (optional, core_ty) = strip_option(&f.ty);
if optional && mode == Mode::Dense {
return Err(syn::Error::new_spanned(
&f.ty,
"a `dense` struct has no presence bitmap, so its fields cannot be \
`Option<…>`; use the default `sparse` mode (or `packed`)",
));
}
let kind = classify(core_ty)?;
collect_nested(&kind, &mut nested_types);
let dt = dt_expr(&kind);
dt_entries.push(quote!((#id, #fname_str, #dt)));
let pack_val = pack_value(&kind, "e!(__v));
if optional {
pack_stmts.push(quote! {
if let ::core::option::Option::Some(__v) = &self.#fname {
entries.push((#id, #pack_val));
}
});
} else {
pack_stmts.push(quote! {
{ let __v = &self.#fname; entries.push((#id, #pack_val)); }
});
}
let from_ref = unpack_from_ref(&kind, "e!(__r));
let read = if optional {
quote! {
match reader.get(#id)? {
::core::option::Option::Some(__r) => ::core::option::Option::Some(#from_ref),
::core::option::Option::None => ::core::option::Option::None,
}
}
} else {
quote! {
match reader.get(#id)? {
::core::option::Option::Some(__r) => #from_ref,
::core::option::Option::None => return ::core::result::Result::Err(::verit::Error::MissingField(#id)),
}
}
};
unpack_inits.push(quote!(#fname: #read));
}
let mut seen_nested = std::collections::BTreeSet::new();
let nested_registers: Vec<TokenStream2> = nested_types
.iter()
.filter(|t| seen_nested.insert(quote!(#t).to_string()))
.map(|t| quote!(let builder = <#t as ::verit::VeritType>::verit_register(builder, seen);))
.collect();
Ok(quote! {
impl ::verit::VeritType for #ident {
const VERIT_NAME: &'static str = #name_str;
const VERIT_MODE: ::verit::StructMode = #mode_expr;
fn verit_register(
builder: ::verit::SchemaBuilder,
seen: &mut ::std::collections::BTreeSet<&'static str>,
) -> ::verit::SchemaBuilder {
if !seen.insert(<Self as ::verit::VeritType>::VERIT_NAME) {
return builder;
}
let fields = ::std::vec![ #(#dt_entries),* ];
let builder = match <Self as ::verit::VeritType>::VERIT_MODE {
::verit::StructMode::Dense => builder.add_dense_struct(<Self as ::verit::VeritType>::VERIT_NAME, fields),
::verit::StructMode::Packed => builder.add_packed_struct(<Self as ::verit::VeritType>::VERIT_NAME, fields),
::verit::StructMode::Sparse => builder.add_struct(<Self as ::verit::VeritType>::VERIT_NAME, fields),
};
#(#nested_registers)*
builder
}
fn verit_schema() -> &'static ::verit::Schema {
static SCHEMA: ::std::sync::OnceLock<::verit::Schema> = ::std::sync::OnceLock::new();
SCHEMA.get_or_init(|| {
let mut seen = ::std::collections::BTreeSet::new();
<Self as ::verit::VeritType>::verit_register(::verit::SchemaBuilder::new(), &mut seen)
.build(<Self as ::verit::VeritType>::VERIT_NAME)
.expect("derived Veritate schema is valid")
})
}
fn verit_pack(&self) -> ::verit::Value {
let mut entries: ::std::vec::Vec<(u16, ::verit::Value)> = ::std::vec::Vec::new();
#(#pack_stmts)*
::verit::Value::Struct(entries)
}
fn verit_unpack(reader: &::verit::StructReader) -> ::verit::Result<Self> {
::core::result::Result::Ok(Self {
#(#unpack_inits),*
})
}
}
})
}
#[derive(PartialEq, Clone, Copy)]
enum Mode {
Sparse,
Dense,
Packed,
}
fn parse_mode(input: &DeriveInput) -> syn::Result<Mode> {
let mut mode = Mode::Sparse;
for attr in &input.attrs {
if !attr.path().is_ident("verit") {
continue;
}
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("mode") {
let value = meta.value()?;
let lit: syn::LitStr = value.parse()?;
mode = match lit.value().as_str() {
"sparse" => Mode::Sparse,
"dense" => Mode::Dense,
"packed" => Mode::Packed,
other => {
return Err(meta.error(format!(
"unknown verit mode {other:?} (expected sparse, dense, or packed)"
)))
}
};
Ok(())
} else {
Err(meta.error("unknown #[verit(…)] container option (expected `mode`)"))
}
})?;
}
Ok(mode)
}
fn parse_field_id(f: &syn::Field) -> syn::Result<u16> {
let mut id: Option<u16> = None;
for attr in &f.attrs {
if !attr.path().is_ident("verit") {
continue;
}
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("id") {
let value = meta.value()?;
let lit: syn::LitInt = value.parse()?;
id = Some(lit.base10_parse()?);
Ok(())
} else {
Err(meta.error("unknown #[verit(…)] field option (expected `id`)"))
}
})?;
}
id.ok_or_else(|| {
syn::Error::new_spanned(f, "every field needs a Veritate id: add `#[verit(id = N)]`")
})
}
fn strip_option(ty: &Type) -> (bool, &Type) {
if let Some(inner) = path_generic(ty, "Option") {
(true, inner)
} else {
(false, ty)
}
}
fn path_generic<'a>(ty: &'a Type, name: &str) -> Option<&'a Type> {
let Type::Path(tp) = ty else { return None };
let seg = tp.path.segments.last()?;
if seg.ident != name {
return None;
}
let PathArguments::AngleBracketed(args) = &seg.arguments else {
return None;
};
for a in &args.args {
if let GenericArgument::Type(t) = a {
return Some(t);
}
}
None
}
fn classify(ty: &Type) -> syn::Result<Kind> {
if let Some(inner) = path_generic(ty, "Vec") {
if type_is_ident(inner, "u8") {
return Ok(Kind::Bytes);
}
return Ok(Kind::List(Box::new(classify(inner)?)));
}
if let Type::Path(tp) = ty {
if let Some(seg) = tp.path.segments.last() {
let id = seg.ident.to_string();
if id == "String" {
return Ok(Kind::Str);
}
if let Some(&s) = SCALARS.iter().find(|&&s| s == id) {
return Ok(Kind::Scalar(s));
}
}
return Ok(Kind::Nested(Box::new(ty.clone())));
}
Err(syn::Error::new_spanned(
ty,
"unsupported #[derive(Verit)] field type (expected a scalar, String, \
Vec<u8>, Vec<T>, Option<T>, or a nested #[derive(Verit)] struct)",
))
}
fn type_is_ident(ty: &Type, name: &str) -> bool {
matches!(ty, Type::Path(tp) if tp.path.is_ident(name))
}
fn collect_nested(kind: &Kind, out: &mut Vec<Type>) {
match kind {
Kind::Nested(t) => out.push((**t).clone()),
Kind::List(inner) => collect_nested(inner, out),
_ => {}
}
}
fn scalar_variant(s: &str) -> proc_macro2::Ident {
let mut c = s.chars();
let first = c.next().unwrap().to_ascii_uppercase();
format_ident!("{}{}", first, c.as_str())
}
fn dt_expr(kind: &Kind) -> TokenStream2 {
match kind {
Kind::Scalar(s) => {
let v = scalar_variant(s);
quote!(::verit::Dt::#v)
}
Kind::Str => quote!(::verit::Dt::Str),
Kind::Bytes => quote!(::verit::Dt::Bytes),
Kind::List(inner) => {
let e = dt_expr(inner);
quote!(::verit::Dt::list(#e))
}
Kind::Nested(t) => {
quote!(::verit::Dt::named(<#t as ::verit::VeritType>::VERIT_NAME))
}
}
}
fn pack_value(kind: &Kind, expr: &TokenStream2) -> TokenStream2 {
match kind {
Kind::Scalar(s) => {
let v = scalar_variant(s);
quote!(::verit::Value::#v(*#expr))
}
Kind::Str => quote!(::verit::Value::str(#expr)),
Kind::Bytes => quote!(::verit::Value::Bytes((#expr).to_vec())),
Kind::List(inner) => {
let e = pack_value(inner, "e!(__e));
quote!(::verit::Value::List((#expr).iter().map(|__e| #e).collect()))
}
Kind::Nested(_) => quote!(::verit::VeritType::verit_pack(#expr)),
}
}
fn unpack_from_ref(kind: &Kind, expr: &TokenStream2) -> TokenStream2 {
let mismatch = |want: &str| {
let want = want.to_string();
quote! {
__other => return ::core::result::Result::Err(::verit::Error::TypeMismatch {
expected: #want.into(),
got: __other.kind().into(),
}),
}
};
match kind {
Kind::Scalar(s) => {
let v = scalar_variant(s);
let m = mismatch(s);
quote! {
match #expr {
::verit::Ref::#v(__x) => __x,
#m
}
}
}
Kind::Str => {
let m = mismatch("string");
quote! {
match #expr {
::verit::Ref::Str(__s) => __s.to_string(),
#m
}
}
}
Kind::Bytes => {
let m = mismatch("bytes");
quote! {
match #expr {
::verit::Ref::Bytes(__b) => __b.to_vec(),
#m
}
}
}
Kind::List(inner) => {
let elem = unpack_from_ref(inner, "e!(__list.get(__i)?));
let m = mismatch("list");
quote! {
match #expr {
::verit::Ref::List(__list) => {
let mut __out = ::std::vec::Vec::with_capacity(__list.len() as usize);
for __i in 0..__list.len() {
__out.push(#elem);
}
__out
}
#m
}
}
}
Kind::Nested(t) => {
let m = mismatch("struct");
quote! {
match #expr {
::verit::Ref::Struct(__sr) => <#t as ::verit::VeritType>::verit_unpack(&__sr)?,
#m
}
}
}
}
}