Skip to main content

mnesis_macros/
lib.rs

1use std::collections::{HashMap, HashSet};
2
3use proc_macro::TokenStream;
4use quote::quote;
5use syn::{Data, DeriveInput, Error, Result, Type, parse_macro_input};
6
7/// Attribute macro that transforms a unit struct into an aggregate newtype.
8///
9/// Usage:
10/// ```ignore
11/// #[mnesis::aggregate(state = MyState, error = MyError, id = MyId)]
12/// struct MyAggregate;
13/// ```
14///
15/// Generates `impl Aggregate` for the unit struct (a type-level marker) plus a
16/// convenience `Name::new(id) -> AggregateRoot<Self>` constructor. Implement
17/// `Handle<C>` on the marker as `handle(state, cmd) -> events`.
18#[proc_macro_attribute]
19pub fn aggregate(attr: TokenStream, item: TokenStream) -> TokenStream {
20    let ast = parse_macro_input!(item as DeriveInput);
21    let args = attr;
22    match parse_aggregate(&ast, args.into()) {
23        Ok(code) => code,
24        Err(e) => e.to_compile_error(),
25    }
26    .into()
27}
28
29#[proc_macro_derive(DomainEvent)]
30pub fn domain_event(input: TokenStream) -> TokenStream {
31    let ast = parse_macro_input!(input as DeriveInput);
32    match parse_domain_event(&ast) {
33        Ok(code) => code,
34        Err(e) => e.to_compile_error(),
35    }
36    .into()
37}
38
39fn parse_domain_event(ast: &DeriveInput) -> Result<proc_macro2::TokenStream> {
40    let name = &ast.ident;
41    match &ast.data {
42        Data::Enum(data_enum) => {
43            if data_enum.variants.is_empty() {
44                return Err(Error::new(
45                    name.span(),
46                    "DomainEvent enum must have at least one variant.",
47                ));
48            }
49
50            let variant_arms: Vec<_> = data_enum
51                .variants
52                .iter()
53                .map(|variant| {
54                    let variant_ident = &variant.ident;
55                    let variant_name = variant_ident.to_string();
56                    match &variant.fields {
57                        syn::Fields::Unit => {
58                            quote! { #name::#variant_ident => #variant_name }
59                        }
60                        syn::Fields::Unnamed(_) => {
61                            quote! { #name::#variant_ident(..) => #variant_name }
62                        }
63                        syn::Fields::Named(_) => {
64                            quote! { #name::#variant_ident { .. } => #variant_name }
65                        }
66                    }
67                })
68                .collect();
69
70            let expanded = quote! {
71                impl ::mnesis::Message for #name {}
72
73                impl ::mnesis::DomainEvent for #name {
74                    fn name(&self) -> &'static str {
75                        match self {
76                            #(#variant_arms),*
77                        }
78                    }
79                }
80            };
81
82            Ok(expanded)
83        }
84        Data::Struct(_) => Err(Error::new(
85            name.span(),
86            "DomainEvent derive requires an enum. Wrap event structs in an enum: `enum MyEvent { Created(Created), ... }`",
87        )),
88        Data::Union(_) => Err(Error::new(name.span(), "Unions are not supported.")),
89    }
90}
91
92/// Generates a unit struct with inherent `upcast` and `current_version`
93/// functions from annotated transform functions.
94///
95/// # Attributes
96///
97/// - `aggregate = Type` — the aggregate type these transforms belong to
98/// - `error = Type` — the error type returned by transform functions
99///
100/// Each method must be annotated with `#[transform(...)]`:
101/// - `event = "EventName"` — the event type this transform handles
102/// - `from = N` — source schema version (>= 1)
103/// - `to = N` — target schema version (must be `from + 1`)
104/// - `rename = "NewName"` — optional event type rename
105///
106/// # Compile-time validation
107///
108/// - `from >= 1`
109/// - `to == from + 1` for each transform (contiguity per step)
110/// - No duplicate `(event, from)` pairs
111/// - **Chain coverage**: for each event type, every schema version in
112///   `[1, current_version]` is reachable via a contiguous chain — gaps
113///   produce a compile error naming the missing step
114///
115/// # Emitted output
116///
117/// The macro emits a `pub struct <Name>;` plus an inherent impl block
118/// carrying the user's transform functions (with `#[transform]` attrs
119/// stripped) and two associated functions:
120///
121/// - `pub fn upcast<'a>(EventMorsel<'a>) -> Result<EventMorsel<'a>, Error>` —
122///   runs the chain to current schema version. Associated (no `&self`)
123///   so call sites are `OrderTransforms::upcast(morsel)` — a `'static`
124///   function pointer pluggable into [`EventStore::load_with`].
125/// - `pub fn current_version(event_type: &str) -> Option<Version>` — the
126///   write-path schema-version stamp lookup. Also associated; call as
127///   `OrderTransforms::current_version("EventName")`.
128///
129/// # Example
130///
131/// ```ignore
132/// #[mnesis::transforms(aggregate = Order, error = MyError)]
133/// impl OrderTransforms {
134///     #[transform(event = "OrderCreated", from = 1, to = 2)]
135///     fn v1_to_v2(payload: &[u8]) -> Result<Vec<u8>, MyError> {
136///         Ok(payload.to_vec())
137///     }
138/// }
139///
140/// // Direct call:
141/// let upgraded = OrderTransforms::upcast(morsel)?;
142///
143/// // Plugged into the facade:
144/// let root = store.load_with(id, OrderTransforms::upcast).await?;
145/// ```
146#[proc_macro_attribute]
147pub fn transforms(attr: TokenStream, item: TokenStream) -> TokenStream {
148    let args = attr;
149    let ast = parse_macro_input!(item as syn::ItemImpl);
150    match parse_transforms(&ast, args.into()) {
151        Ok(code) => code,
152        Err(e) => e.to_compile_error(),
153    }
154    .into()
155}
156
157struct TransformDef {
158    fn_name: syn::Ident,
159    event_type: String,
160    from_version: u64,
161    to_version: u64,
162    rename: Option<String>,
163}
164
165fn parse_transform_attr(method: &syn::ImplItemFn) -> Result<Option<TransformDef>> {
166    let mut transform_attr = None;
167
168    for attr in &method.attrs {
169        if attr.path().is_ident("transform") {
170            if transform_attr.is_some() {
171                return Err(Error::new_spanned(attr, "duplicate #[transform] attribute"));
172            }
173
174            let mut event_type: Option<String> = None;
175            let mut from_version: Option<u64> = None;
176            let mut to_version: Option<u64> = None;
177            let mut rename: Option<String> = None;
178
179            attr.parse_nested_meta(|meta| {
180                if meta.path.is_ident("event") {
181                    let value = meta.value()?;
182                    let lit: syn::LitStr = value.parse()?;
183                    event_type = Some(lit.value());
184                } else if meta.path.is_ident("from") {
185                    let value = meta.value()?;
186                    let lit: syn::LitInt = value.parse()?;
187                    from_version = Some(lit.base10_parse()?);
188                } else if meta.path.is_ident("to") {
189                    let value = meta.value()?;
190                    let lit: syn::LitInt = value.parse()?;
191                    to_version = Some(lit.base10_parse()?);
192                } else if meta.path.is_ident("rename") {
193                    let value = meta.value()?;
194                    let lit: syn::LitStr = value.parse()?;
195                    rename = Some(lit.value());
196                } else {
197                    return Err(meta.error("expected `event`, `from`, `to`, or `rename`"));
198                }
199                Ok(())
200            })?;
201
202            let event_type = event_type.ok_or_else(|| {
203                Error::new_spanned(attr, "`event` is required in #[transform(...)]")
204            })?;
205            let from_version = from_version.ok_or_else(|| {
206                Error::new_spanned(attr, "`from` is required in #[transform(...)]")
207            })?;
208            let to_version = to_version
209                .ok_or_else(|| Error::new_spanned(attr, "`to` is required in #[transform(...)]"))?;
210
211            transform_attr = Some(TransformDef {
212                fn_name: method.sig.ident.clone(),
213                event_type,
214                from_version,
215                to_version,
216                rename,
217            });
218        }
219    }
220
221    Ok(transform_attr)
222}
223
224fn parse_transforms(
225    ast: &syn::ItemImpl,
226    args: proc_macro2::TokenStream,
227) -> Result<proc_macro2::TokenStream> {
228    // 1. Parse aggregate = Type, error = Type from outer attributes
229    let mut aggregate_type: Option<Type> = None;
230    let mut error_type: Option<Type> = None;
231    let parser = syn::meta::parser(|meta| {
232        if meta.path.is_ident("aggregate") {
233            aggregate_type = Some(meta.value()?.parse()?);
234        } else if meta.path.is_ident("error") {
235            error_type = Some(meta.value()?.parse()?);
236        } else {
237            return Err(meta.error("expected `aggregate` or `error`"));
238        }
239        Ok(())
240    });
241    syn::parse::Parser::parse2(parser, args)?;
242    let _aggregate_type = aggregate_type
243        .ok_or_else(|| Error::new(proc_macro2::Span::call_site(), "`aggregate` is required"))?;
244    let error_type = error_type
245        .ok_or_else(|| Error::new(proc_macro2::Span::call_site(), "`error` is required"))?;
246
247    // 2. Get the struct name from the impl block
248    let struct_ident = match &*ast.self_ty {
249        syn::Type::Path(p) => {
250            &p.path
251                .segments
252                .last()
253                .ok_or_else(|| Error::new_spanned(&ast.self_ty, "expected a type name"))?
254                .ident
255        }
256        _ => return Err(Error::new_spanned(&ast.self_ty, "expected a type name")),
257    };
258
259    // 3. Parse each method's #[transform] attributes
260    let mut transforms = Vec::new();
261    for item in &ast.items {
262        let method = match item {
263            syn::ImplItem::Fn(m) => m,
264            _ => continue,
265        };
266        if let Some(def) = parse_transform_attr(method)? {
267            transforms.push(def);
268        }
269    }
270
271    // 4. Validate from >= 1
272    for t in &transforms {
273        if t.from_version < 1 {
274            return Err(Error::new_spanned(&t.fn_name, "from version must be >= 1"));
275        }
276    }
277
278    // 5. Validate to == from + 1
279    for t in &transforms {
280        if t.to_version != t.from_version + 1 {
281            return Err(Error::new_spanned(
282                &t.fn_name,
283                format!(
284                    "non-contiguous version: to ({}) must equal from + 1 ({})",
285                    t.to_version,
286                    t.from_version + 1,
287                ),
288            ));
289        }
290    }
291
292    // 6. Validate no duplicate (event, from)
293    let mut seen = HashSet::new();
294    for t in &transforms {
295        let key = (t.event_type.clone(), t.from_version);
296        if !seen.insert(key) {
297            return Err(Error::new_spanned(
298                &t.fn_name,
299                format!(
300                    "duplicate transform for event '{}' at source version {}",
301                    t.event_type, t.from_version,
302                ),
303            ));
304        }
305    }
306
307    // 6b. Validate chain coverage per event type.
308    //
309    // For each event type, every schema version in [1, current_version]
310    // must be reachable through a contiguous chain. Find the smallest gap
311    // (i.e. the smallest `from` in [1, max_from] for which no transform
312    // exists) and reject with a diagnostic naming the missing step.
313    let mut from_versions_by_event: HashMap<String, HashSet<u64>> = HashMap::new();
314    let mut max_from_by_event: HashMap<String, u64> = HashMap::new();
315    for t in &transforms {
316        from_versions_by_event
317            .entry(t.event_type.clone())
318            .or_default()
319            .insert(t.from_version);
320        let entry = max_from_by_event.entry(t.event_type.clone()).or_insert(0);
321        if t.from_version > *entry {
322            *entry = t.from_version;
323        }
324    }
325    for t in &transforms {
326        let Some(from_set) = from_versions_by_event.get(&t.event_type) else {
327            continue;
328        };
329        let Some(&max_from) = max_from_by_event.get(&t.event_type) else {
330            continue;
331        };
332        for v in 1..max_from {
333            if !from_set.contains(&v) {
334                return Err(Error::new_spanned(
335                    &t.fn_name,
336                    format!(
337                        "transform chain gap for event '{}': missing step from version {} to version {} (chain must cover every version in [1, {}])",
338                        t.event_type,
339                        v,
340                        v + 1,
341                        max_from + 1,
342                    ),
343                ));
344            }
345        }
346    }
347
348    // 7. Build the original impl block with #[transform] attrs stripped
349    let stripped_methods: Vec<_> = ast
350        .items
351        .iter()
352        .map(|item| match item {
353            syn::ImplItem::Fn(m) => {
354                let mut method = m.clone();
355                method.attrs.retain(|a| !a.path().is_ident("transform"));
356                syn::ImplItem::Fn(method)
357            }
358            other => other.clone(),
359        })
360        .collect();
361
362    // 8. Generate match arms for upcast()
363    let match_arms: Vec<_> = transforms
364        .iter()
365        .map(|t| {
366            let fn_name = &t.fn_name;
367            let event_type = &t.event_type;
368            let from_version = t.from_version;
369            let to_version = t.to_version;
370            let output_event_type = t.rename.as_deref().unwrap_or(&t.event_type);
371
372            quote! {
373                (#event_type, v) if v == ::mnesis::Version::new(#from_version).expect("nonzero") => {
374                    let payload = Self::#fn_name(morsel.payload())?;
375                    ::mnesis_store::upcasting::EventMorsel::new(
376                        #output_event_type,
377                        ::mnesis::Version::new(#to_version).expect("nonzero"),
378                        payload,
379                    )
380                }
381            }
382        })
383        .collect();
384
385    // 9. Compute max version per event type for current_version()
386    let mut max_versions: HashMap<String, u64> = HashMap::new();
387    for t in &transforms {
388        let entry = max_versions.entry(t.event_type.clone()).or_insert(1);
389        if t.to_version > *entry {
390            *entry = t.to_version;
391        }
392        // Track renamed destination event types too
393        if let Some(ref rename) = t.rename {
394            let entry = max_versions.entry(rename.clone()).or_insert(1);
395            if t.to_version > *entry {
396                *entry = t.to_version;
397            }
398        }
399    }
400
401    let version_arms: Vec<_> = max_versions
402        .iter()
403        .map(|(event_type, version)| {
404            quote! {
405                #event_type => ::core::option::Option::Some(
406                    ::mnesis::Version::new(#version).expect("nonzero")
407                )
408            }
409        })
410        .collect();
411
412    // 10. Emit: pub struct + impl block carrying user methods + the two
413    //     generated associated functions. No trait impl — call sites use
414    //     path syntax (`X::upcast(...)`, `X::current_version(...)`) which
415    //     yields `'static` function pointers pluggable into the facade's
416    //     `load_with` / `save_with` methods.
417    let expanded = quote! {
418        pub struct #struct_ident;
419
420        impl #struct_ident {
421            #(#stripped_methods)*
422
423            /// Run all matching transforms until the morsel reaches the
424            /// current schema version. Generated by `#[mnesis::transforms]`.
425            pub fn upcast<'a>(
426                mut morsel: ::mnesis_store::upcasting::EventMorsel<'a>,
427            ) -> ::core::result::Result<
428                ::mnesis_store::upcasting::EventMorsel<'a>,
429                #error_type,
430            > {
431                loop {
432                    morsel = match (morsel.event_type(), morsel.schema_version()) {
433                        #(#match_arms,)*
434                        _ => break,
435                    };
436                }
437                ::core::result::Result::Ok(morsel)
438            }
439
440            /// Current schema version for `event_type` (stamped on new
441            /// events). `None` when the event type has no transforms.
442            /// Generated by `#[mnesis::transforms]`.
443            #[must_use]
444            pub fn current_version(event_type: &str) -> ::core::option::Option<::mnesis::Version> {
445                match event_type {
446                    #(#version_arms,)*
447                    _ => ::core::option::Option::None,
448                }
449            }
450        }
451    };
452
453    Ok(expanded)
454}
455
456fn parse_aggregate(
457    ast: &DeriveInput,
458    args: proc_macro2::TokenStream,
459) -> Result<proc_macro2::TokenStream> {
460    let name = &ast.ident;
461    let vis = &ast.vis;
462    // Preserve user attributes (#[cfg(...)], #[doc = "..."], etc.)
463    let user_attrs = &ast.attrs;
464
465    // Only unit structs allowed
466    match &ast.data {
467        Data::Struct(data) => {
468            if !data.fields.is_empty() {
469                return Err(Error::new(
470                    name.span(),
471                    "aggregate macro requires a unit struct (no fields).",
472                ));
473            }
474        }
475        _ => {
476            return Err(Error::new(
477                name.span(),
478                "aggregate macro only works on unit structs.",
479            ));
480        }
481    }
482
483    // Parse state = ..., error = ..., id = ... from attribute args
484    let mut state_type: Option<Type> = None;
485    let mut error_type: Option<Type> = None;
486    let mut id_type: Option<Type> = None;
487
488    let parser = syn::meta::parser(|meta| {
489        if meta.path.is_ident("state") {
490            state_type = Some(meta.value()?.parse()?);
491        } else if meta.path.is_ident("error") {
492            error_type = Some(meta.value()?.parse()?);
493        } else if meta.path.is_ident("id") {
494            id_type = Some(meta.value()?.parse()?);
495        } else {
496            return Err(meta.error("expected `state`, `error`, or `id`"));
497        }
498        Ok(())
499    });
500
501    syn::parse::Parser::parse2(parser, args)?;
502
503    let state_type = state_type.ok_or_else(|| Error::new(name.span(), "`state` is required"))?;
504    let error_type = error_type.ok_or_else(|| Error::new(name.span(), "`error` is required"))?;
505    let id_type = id_type.ok_or_else(|| Error::new(name.span(), "`id` is required"))?;
506
507    let expanded = quote! {
508        #(#user_attrs)*
509        #vis struct #name;
510
511        impl ::mnesis::Aggregate for #name {
512            type State = #state_type;
513            type Error = #error_type;
514            type Id = #id_type;
515        }
516
517        impl #name {
518            /// Create a fresh aggregate at initial state (version `None`).
519            ///
520            /// Returns the live [`AggregateRoot`](::mnesis::AggregateRoot) — the
521            /// stateful container. The aggregate type itself is a marker; this
522            /// is a convenience entry point for
523            /// `AggregateRoot::<Self>::new(id)`.
524            #[must_use]
525            #vis fn new(id: #id_type) -> ::mnesis::AggregateRoot<Self> {
526                ::mnesis::AggregateRoot::new(id)
527            }
528        }
529    };
530
531    Ok(expanded)
532}