1#![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#[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
52enum 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 let pack_val = pack_value(&kind, "e!(__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 let from_ref = unpack_from_ref(&kind, "e!(__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 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
274fn 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
283fn 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 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 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 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
366fn 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, "e!(__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
383fn 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, "e!(__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}