Skip to main content

verit_derive/
lib.rs

1//! `#[derive(Verit)]` — generate a Veritate schema, encoder, and decoder from a
2//! plain Rust struct. The Rust peer of Python's `@verit` decorator.
3//!
4//! ```ignore
5//! use verit::{Verit, VeritType, SchemaMode};
6//!
7//! #[derive(Verit)]
8//! #[verit(mode = "dense")]
9//! struct Point { #[verit(id = 1)] x: f64, #[verit(id = 2)] y: f64 }
10//!
11//! #[derive(Verit)]
12//! struct Order {
13//!     #[verit(id = 1)] id: u64,
14//!     #[verit(id = 2)] item: String,
15//!     #[verit(id = 3)] tags: Vec<String>,
16//!     #[verit(id = 4)] origin: Point,          // nested #[derive(Verit)]
17//!     #[verit(id = 5)] note: Option<String>,   // optional (absent = None)
18//! }
19//!
20//! let bytes = order.to_verit(SchemaMode::Inline)?;
21//! let back = Order::from_verit(&bytes)?;
22//! ```
23//!
24//! ## Field type mapping
25//! - `bool`, `u8..u64`, `i8..i64`, `f32`, `f64` → the matching scalar.
26//! - `String` → `string`; `Vec<u8>` → `bytes`.
27//! - `Vec<T>` (T ≠ u8) → `list<T>` (recursively).
28//! - `Option<T>` → an optional field (absent reads back as `None`; only valid
29//!   on `sparse`/`packed` structs, never `dense`).
30//! - any other path type `T` → a nested `struct` (T must also `#[derive(Verit)]`).
31//!
32//! The generated code references the `verit` crate as `::verit`, so a downstream
33//! crate needs `verit = { version = "…", features = ["derive"] }`.
34
35// The proc-macro carries no `unsafe` either — compiler-enforced.
36#![forbid(unsafe_code)]
37
38use proc_macro::TokenStream;
39use proc_macro2::TokenStream as TokenStream2;
40use quote::{format_ident, quote};
41use syn::{parse_macro_input, Data, DeriveInput, Fields, GenericArgument, PathArguments, Type};
42
43/// Derive [`VeritType`](trait@verit::VeritType) for a named-field struct.
44#[proc_macro_derive(Verit, attributes(verit))]
45pub fn derive_verit(input: TokenStream) -> TokenStream {
46    let input = parse_macro_input!(input as DeriveInput);
47    expand(input)
48        .unwrap_or_else(syn::Error::into_compile_error)
49        .into()
50}
51
52/// One classified field type. `Nested` carries the original type tokens so the
53/// generated code can name the type in trait-qualified paths.
54enum Kind {
55    Scalar(&'static str),
56    Str,
57    Bytes,
58    List(Box<Kind>),
59    Nested(Box<Type>),
60}
61
62const SCALARS: &[&str] = &[
63    "bool", "u8", "u16", "u32", "u64", "i8", "i16", "i32", "i64", "f32", "f64",
64];
65
66fn expand(input: DeriveInput) -> syn::Result<TokenStream2> {
67    let ident = &input.ident;
68    let name_str = ident.to_string();
69
70    if !input.generics.params.is_empty() {
71        return Err(syn::Error::new_spanned(
72            &input.generics,
73            "#[derive(Verit)] does not support generic types",
74        ));
75    }
76
77    let mode = parse_mode(&input)?;
78    let mode_expr = match mode {
79        Mode::Sparse => quote!(::verit::StructMode::Sparse),
80        Mode::Dense => quote!(::verit::StructMode::Dense),
81        Mode::Packed => quote!(::verit::StructMode::Packed),
82    };
83
84    let fields = match &input.data {
85        Data::Struct(s) => match &s.fields {
86            Fields::Named(named) => &named.named,
87            _ => {
88                return Err(syn::Error::new_spanned(
89                    ident,
90                    "#[derive(Verit)] requires a struct with named fields",
91                ))
92            }
93        },
94        _ => {
95            return Err(syn::Error::new_spanned(
96                ident,
97                "#[derive(Verit)] can only be applied to structs",
98            ))
99        }
100    };
101
102    let mut dt_entries = Vec::new();
103    let mut pack_stmts = Vec::new();
104    let mut unpack_inits = Vec::new();
105    let mut nested_types: Vec<Type> = Vec::new();
106
107    for f in fields {
108        let fname = f.ident.as_ref().unwrap();
109        let fname_str = fname.to_string();
110        let id = parse_field_id(f)?;
111
112        let (optional, core_ty) = strip_option(&f.ty);
113        if optional && mode == Mode::Dense {
114            return Err(syn::Error::new_spanned(
115                &f.ty,
116                "a `dense` struct has no presence bitmap, so its fields cannot be \
117                 `Option<…>`; use the default `sparse` mode (or `packed`)",
118            ));
119        }
120        let kind = classify(core_ty)?;
121        collect_nested(&kind, &mut nested_types);
122
123        let dt = dt_expr(&kind);
124        dt_entries.push(quote!((#id, #fname_str, #dt)));
125
126        // Pack: push (id, Value) — Option fields only when Some.
127        let pack_val = pack_value(&kind, &quote!(__v));
128        if optional {
129            pack_stmts.push(quote! {
130                if let ::core::option::Option::Some(__v) = &self.#fname {
131                    entries.push((#id, #pack_val));
132                }
133            });
134        } else {
135            pack_stmts.push(quote! {
136                { let __v = &self.#fname; entries.push((#id, #pack_val)); }
137            });
138        }
139
140        // Unpack: read the field by id from the dynamic reader.
141        let from_ref = unpack_from_ref(&kind, &quote!(__r));
142        let read = if optional {
143            quote! {
144                match reader.get(#id)? {
145                    ::core::option::Option::Some(__r) => ::core::option::Option::Some(#from_ref),
146                    ::core::option::Option::None => ::core::option::Option::None,
147                }
148            }
149        } else {
150            quote! {
151                match reader.get(#id)? {
152                    ::core::option::Option::Some(__r) => #from_ref,
153                    ::core::option::Option::None => return ::core::result::Result::Err(::verit::Error::MissingField(#id)),
154                }
155            }
156        };
157        unpack_inits.push(quote!(#fname: #read));
158    }
159
160    // De-duplicate nested types by their token string so a type referenced by
161    // several fields is registered once.
162    let mut seen_nested = std::collections::BTreeSet::new();
163    let nested_registers: Vec<TokenStream2> = nested_types
164        .iter()
165        .filter(|t| seen_nested.insert(quote!(#t).to_string()))
166        .map(|t| quote!(let builder = <#t as ::verit::VeritType>::verit_register(builder, seen);))
167        .collect();
168
169    Ok(quote! {
170        impl ::verit::VeritType for #ident {
171            const VERIT_NAME: &'static str = #name_str;
172            const VERIT_MODE: ::verit::StructMode = #mode_expr;
173
174            fn verit_register(
175                builder: ::verit::SchemaBuilder,
176                seen: &mut ::std::collections::BTreeSet<&'static str>,
177            ) -> ::verit::SchemaBuilder {
178                if !seen.insert(<Self as ::verit::VeritType>::VERIT_NAME) {
179                    return builder;
180                }
181                let fields = ::std::vec![ #(#dt_entries),* ];
182                let builder = match <Self as ::verit::VeritType>::VERIT_MODE {
183                    ::verit::StructMode::Dense => builder.add_dense_struct(<Self as ::verit::VeritType>::VERIT_NAME, fields),
184                    ::verit::StructMode::Packed => builder.add_packed_struct(<Self as ::verit::VeritType>::VERIT_NAME, fields),
185                    ::verit::StructMode::Sparse => builder.add_struct(<Self as ::verit::VeritType>::VERIT_NAME, fields),
186                };
187                #(#nested_registers)*
188                builder
189            }
190
191            fn verit_schema() -> &'static ::verit::Schema {
192                static SCHEMA: ::std::sync::OnceLock<::verit::Schema> = ::std::sync::OnceLock::new();
193                SCHEMA.get_or_init(|| {
194                    let mut seen = ::std::collections::BTreeSet::new();
195                    <Self as ::verit::VeritType>::verit_register(::verit::SchemaBuilder::new(), &mut seen)
196                        .build(<Self as ::verit::VeritType>::VERIT_NAME)
197                        .expect("derived Veritate schema is valid")
198                })
199            }
200
201            fn verit_pack(&self) -> ::verit::Value {
202                let mut entries: ::std::vec::Vec<(u16, ::verit::Value)> = ::std::vec::Vec::new();
203                #(#pack_stmts)*
204                ::verit::Value::Struct(entries)
205            }
206
207            fn verit_unpack(reader: &::verit::StructReader) -> ::verit::Result<Self> {
208                ::core::result::Result::Ok(Self {
209                    #(#unpack_inits),*
210                })
211            }
212        }
213    })
214}
215
216#[derive(PartialEq, Clone, Copy)]
217enum Mode {
218    Sparse,
219    Dense,
220    Packed,
221}
222
223fn parse_mode(input: &DeriveInput) -> syn::Result<Mode> {
224    let mut mode = Mode::Sparse;
225    for attr in &input.attrs {
226        if !attr.path().is_ident("verit") {
227            continue;
228        }
229        attr.parse_nested_meta(|meta| {
230            if meta.path.is_ident("mode") {
231                let value = meta.value()?;
232                let lit: syn::LitStr = value.parse()?;
233                mode = match lit.value().as_str() {
234                    "sparse" => Mode::Sparse,
235                    "dense" => Mode::Dense,
236                    "packed" => Mode::Packed,
237                    other => {
238                        return Err(meta.error(format!(
239                            "unknown verit mode {other:?} (expected sparse, dense, or packed)"
240                        )))
241                    }
242                };
243                Ok(())
244            } else {
245                Err(meta.error("unknown #[verit(…)] container option (expected `mode`)"))
246            }
247        })?;
248    }
249    Ok(mode)
250}
251
252fn parse_field_id(f: &syn::Field) -> syn::Result<u16> {
253    let mut id: Option<u16> = None;
254    for attr in &f.attrs {
255        if !attr.path().is_ident("verit") {
256            continue;
257        }
258        attr.parse_nested_meta(|meta| {
259            if meta.path.is_ident("id") {
260                let value = meta.value()?;
261                let lit: syn::LitInt = value.parse()?;
262                id = Some(lit.base10_parse()?);
263                Ok(())
264            } else {
265                Err(meta.error("unknown #[verit(…)] field option (expected `id`)"))
266            }
267        })?;
268    }
269    id.ok_or_else(|| {
270        syn::Error::new_spanned(f, "every field needs a Veritate id: add `#[verit(id = N)]`")
271    })
272}
273
274/// Peel one `Option<T>`. Returns `(is_option, inner_type)`.
275fn strip_option(ty: &Type) -> (bool, &Type) {
276    if let Some(inner) = path_generic(ty, "Option") {
277        (true, inner)
278    } else {
279        (false, ty)
280    }
281}
282
283/// If `ty` is `Name<Inner>` (last path segment), return `Inner`.
284fn path_generic<'a>(ty: &'a Type, name: &str) -> Option<&'a Type> {
285    let Type::Path(tp) = ty else { return None };
286    let seg = tp.path.segments.last()?;
287    if seg.ident != name {
288        return None;
289    }
290    let PathArguments::AngleBracketed(args) = &seg.arguments else {
291        return None;
292    };
293    for a in &args.args {
294        if let GenericArgument::Type(t) = a {
295            return Some(t);
296        }
297    }
298    None
299}
300
301fn classify(ty: &Type) -> syn::Result<Kind> {
302    // Vec<u8> => bytes; Vec<T> => list<T>.
303    if let Some(inner) = path_generic(ty, "Vec") {
304        if type_is_ident(inner, "u8") {
305            return Ok(Kind::Bytes);
306        }
307        return Ok(Kind::List(Box::new(classify(inner)?)));
308    }
309    if let Type::Path(tp) = ty {
310        if let Some(seg) = tp.path.segments.last() {
311            let id = seg.ident.to_string();
312            if id == "String" {
313                return Ok(Kind::Str);
314            }
315            if let Some(&s) = SCALARS.iter().find(|&&s| s == id) {
316                return Ok(Kind::Scalar(s));
317            }
318        }
319        // Anything else that is a bare path: treat as a nested Verit struct.
320        return Ok(Kind::Nested(Box::new(ty.clone())));
321    }
322    Err(syn::Error::new_spanned(
323        ty,
324        "unsupported #[derive(Verit)] field type (expected a scalar, String, \
325         Vec<u8>, Vec<T>, Option<T>, or a nested #[derive(Verit)] struct)",
326    ))
327}
328
329fn type_is_ident(ty: &Type, name: &str) -> bool {
330    matches!(ty, Type::Path(tp) if tp.path.is_ident(name))
331}
332
333fn collect_nested(kind: &Kind, out: &mut Vec<Type>) {
334    match kind {
335        Kind::Nested(t) => out.push((**t).clone()),
336        Kind::List(inner) => collect_nested(inner, out),
337        _ => {}
338    }
339}
340
341fn scalar_variant(s: &str) -> proc_macro2::Ident {
342    // "u8" -> U8, "bool" -> Bool, "f64" -> F64.
343    let mut c = s.chars();
344    let first = c.next().unwrap().to_ascii_uppercase();
345    format_ident!("{}{}", first, c.as_str())
346}
347
348fn dt_expr(kind: &Kind) -> TokenStream2 {
349    match kind {
350        Kind::Scalar(s) => {
351            let v = scalar_variant(s);
352            quote!(::verit::Dt::#v)
353        }
354        Kind::Str => quote!(::verit::Dt::Str),
355        Kind::Bytes => quote!(::verit::Dt::Bytes),
356        Kind::List(inner) => {
357            let e = dt_expr(inner);
358            quote!(::verit::Dt::list(#e))
359        }
360        Kind::Nested(t) => {
361            quote!(::verit::Dt::named(<#t as ::verit::VeritType>::VERIT_NAME))
362        }
363    }
364}
365
366/// Build a `::verit::Value` from `expr`, an expression of type `&Inner`.
367fn pack_value(kind: &Kind, expr: &TokenStream2) -> TokenStream2 {
368    match kind {
369        Kind::Scalar(s) => {
370            let v = scalar_variant(s);
371            quote!(::verit::Value::#v(*#expr))
372        }
373        Kind::Str => quote!(::verit::Value::str(#expr)),
374        Kind::Bytes => quote!(::verit::Value::Bytes((#expr).to_vec())),
375        Kind::List(inner) => {
376            let e = pack_value(inner, &quote!(__e));
377            quote!(::verit::Value::List((#expr).iter().map(|__e| #e).collect()))
378        }
379        Kind::Nested(_) => quote!(::verit::VeritType::verit_pack(#expr)),
380    }
381}
382
383/// Build an `Inner` from `expr`, an expression of type `::verit::Ref`. May use
384/// `?` / `return Err(…)`, so it must be spliced inside the generated
385/// `verit_unpack` (which returns `Result`).
386fn unpack_from_ref(kind: &Kind, expr: &TokenStream2) -> TokenStream2 {
387    let mismatch = |want: &str| {
388        let want = want.to_string();
389        quote! {
390            __other => return ::core::result::Result::Err(::verit::Error::TypeMismatch {
391                expected: #want.into(),
392                got: __other.kind().into(),
393            }),
394        }
395    };
396    match kind {
397        Kind::Scalar(s) => {
398            let v = scalar_variant(s);
399            let m = mismatch(s);
400            quote! {
401                match #expr {
402                    ::verit::Ref::#v(__x) => __x,
403                    #m
404                }
405            }
406        }
407        Kind::Str => {
408            let m = mismatch("string");
409            quote! {
410                match #expr {
411                    ::verit::Ref::Str(__s) => __s.to_string(),
412                    #m
413                }
414            }
415        }
416        Kind::Bytes => {
417            let m = mismatch("bytes");
418            quote! {
419                match #expr {
420                    ::verit::Ref::Bytes(__b) => __b.to_vec(),
421                    #m
422                }
423            }
424        }
425        Kind::List(inner) => {
426            let elem = unpack_from_ref(inner, &quote!(__list.get(__i)?));
427            let m = mismatch("list");
428            quote! {
429                match #expr {
430                    ::verit::Ref::List(__list) => {
431                        let mut __out = ::std::vec::Vec::with_capacity(__list.len() as usize);
432                        for __i in 0..__list.len() {
433                            __out.push(#elem);
434                        }
435                        __out
436                    }
437                    #m
438                }
439            }
440        }
441        Kind::Nested(t) => {
442            let m = mismatch("struct");
443            quote! {
444                match #expr {
445                    ::verit::Ref::Struct(__sr) => <#t as ::verit::VeritType>::verit_unpack(&__sr)?,
446                    #m
447                }
448            }
449        }
450    }
451}