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#[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#[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 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 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 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 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 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 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 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 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 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 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 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 let expanded = quote! {
418 pub struct #struct_ident;
419
420 impl #struct_ident {
421 #(#stripped_methods)*
422
423 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 #[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 let user_attrs = &ast.attrs;
464
465 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 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 #[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}