use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, Data, DeriveInput, Fields, Lit, Meta};
#[proc_macro_derive(SpecShape, attributes(spec))]
pub fn derive_spec_shape(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let mut args_type: Option<String> = None;
let mut quirk_type: Option<String> = None;
let mut args_field = "build_rust_crate_args".to_string();
let mut root_field = "root_crate".to_string();
let mut members_field = "workspace_members".to_string();
let mut crates_field = "crates".to_string();
for attr in &input.attrs {
if !attr.path().is_ident("spec") {
continue;
}
let Meta::List(list) = &attr.meta else { continue };
let _ = list.parse_nested_meta(|meta| {
let Some(ident) = meta.path.get_ident() else {
return Ok(());
};
let value: Lit = meta.value()?.parse()?;
let Lit::Str(s) = value else {
return Ok(());
};
let v = s.value();
match ident.to_string().as_str() {
"args" => args_type = Some(v),
"quirk" => quirk_type = Some(v),
"args_field" => args_field = v,
"root_field" => root_field = v,
"members_field" => members_field = v,
"crates_field" => crates_field = v,
_ => {}
}
Ok(())
});
}
let args_type = match args_type {
Some(t) => syn::parse_str::<syn::Type>(&t).expect("invalid `args` type"),
None => {
return TokenStream::from(quote! {
compile_error!("SpecShape requires `#[spec(args = \"<TypeName>\", quirk = \"<TypeName>\")]`");
});
}
};
let quirk_type = match quirk_type {
Some(t) => syn::parse_str::<syn::Type>(&t).expect("invalid `quirk` type"),
None => {
return TokenStream::from(quote! {
compile_error!("SpecShape requires `#[spec(args = \"<TypeName>\", quirk = \"<TypeName>\")]`");
});
}
};
let args_field_ident = syn::Ident::new(&args_field, proc_macro2::Span::call_site());
let root_field_ident = syn::Ident::new(&root_field, proc_macro2::Span::call_site());
let members_field_ident = syn::Ident::new(&members_field, proc_macro2::Span::call_site());
let crates_field_ident = syn::Ident::new(&crates_field, proc_macro2::Span::call_site());
let expanded = quote! {
impl ::gen_types::Spec for #name {
type Args = #args_type;
type Quirk = #quirk_type;
fn schema_version(&self) -> u32 {
self.version
}
fn root_key(&self) -> &str {
self.#root_field_ident.as_str()
}
fn member_keys(&self) -> ::std::vec::Vec<&str> {
self.#members_field_ident.iter().map(::std::string::String::as_str).collect()
}
fn args_for(&self, key: &str) -> ::std::option::Option<&Self::Args> {
self.#crates_field_ident.get(key).map(|c| &c.#args_field_ident)
}
fn quirks_for(&self, key: &str) -> &[Self::Quirk] {
self.#crates_field_ident
.get(key)
.map(|c| c.quirks.as_slice())
.unwrap_or(&[])
}
}
};
TokenStream::from(expanded)
}
#[proc_macro_derive(QuirkRegistry, attributes(quirks))]
pub fn derive_quirk_registry(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let mut enum_name: Option<String> = None;
let mut registry_fn: Option<String> = None;
for attr in &input.attrs {
if !attr.path().is_ident("quirks") {
continue;
}
let Meta::List(list) = &attr.meta else { continue };
let _ = list.parse_nested_meta(|meta| {
let Some(ident) = meta.path.get_ident() else {
return Ok(());
};
let value: Lit = meta.value()?.parse()?;
let Lit::Str(s) = value else {
return Ok(());
};
let v = s.value();
match ident.to_string().as_str() {
"enum_name" => enum_name = Some(v),
"registry_fn" => registry_fn = Some(v),
_ => {}
}
Ok(())
});
}
let enum_ty = match enum_name {
Some(t) => syn::parse_str::<syn::Type>(&t).expect("invalid `enum_name`"),
None => {
return TokenStream::from(quote! {
compile_error!("QuirkRegistry requires `#[quirks(enum_name = \"<EnumName>\", registry_fn = \"<path>\")]`");
});
}
};
let reg_path = match registry_fn {
Some(t) => syn::parse_str::<syn::Path>(&t).expect("invalid `registry_fn`"),
None => {
return TokenStream::from(quote! {
compile_error!("QuirkRegistry requires `#[quirks(enum_name = \"<EnumName>\", registry_fn = \"<path>\")]`");
});
}
};
let expanded = quote! {
impl ::gen_types::QuirkRegistry for #name {
type Quirk = #enum_ty;
fn registry() -> ::std::vec::Vec<(&'static str, ::std::vec::Vec<Self::Quirk>)> {
#reg_path()
}
}
};
TokenStream::from(expanded)
}
#[proc_macro_derive(TypedDispatcher)]
pub fn derive_typed_dispatcher(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = &input.ident;
let Data::Enum(data) = &input.data else {
return TokenStream::from(quote! {
compile_error!("#[derive(TypedDispatcher)] only works on enums");
});
};
let mut kind_entries: Vec<proc_macro2::TokenStream> = Vec::new();
let mut field_entries: Vec<proc_macro2::TokenStream> = Vec::new();
for variant in &data.variants {
let tag = to_kebab_case(&variant.ident.to_string());
let fields = match &variant.fields {
Fields::Named(named) => named
.named
.iter()
.filter_map(|f| f.ident.as_ref().map(std::string::ToString::to_string))
.collect::<Vec<_>>(),
Fields::Unit => Vec::new(),
Fields::Unnamed(_) => {
let msg = format!(
"#[derive(TypedDispatcher)] variant `{}` uses tuple fields; only named-field and unit variants are supported (matches the serde-tagged-enum shape pleme-io requires)",
variant.ident
);
return TokenStream::from(quote! {
compile_error!(#msg);
});
}
};
kind_entries.push(quote! { #tag });
let field_strs: Vec<proc_macro2::TokenStream> =
fields.iter().map(|f| quote! { #f }).collect();
field_entries.push(quote! {
(#tag, ::std::vec![ #( #field_strs ),* ])
});
}
let expanded = quote! {
impl ::gen_types::TypedDispatcher for #name {
fn variant_kinds() -> ::std::vec::Vec<&'static str> {
::std::vec![ #( #kind_entries ),* ]
}
fn variant_fields() -> ::std::vec::Vec<(&'static str, ::std::vec::Vec<&'static str>)> {
::std::vec![ #( #field_entries ),* ]
}
}
};
TokenStream::from(expanded)
}
#[derive(Clone, Copy)]
enum DiscriminantCase {
Kebab,
Snake,
Lower,
Title,
}
impl DiscriminantCase {
fn apply(self, s: &str) -> String {
match self {
DiscriminantCase::Kebab => to_kebab_case(s),
DiscriminantCase::Snake => discriminant_to_snake(s),
DiscriminantCase::Lower => s.to_ascii_lowercase(),
DiscriminantCase::Title => s.to_string(),
}
}
fn parse(s: &str) -> Option<Self> {
match s {
"kebab" | "kebab-case" => Some(DiscriminantCase::Kebab),
"snake" | "snake_case" => Some(DiscriminantCase::Snake),
"lower" | "lowercase" => Some(DiscriminantCase::Lower),
"title" | "Title" | "TitleCase" => Some(DiscriminantCase::Title),
_ => None,
}
}
}
fn discriminant_to_snake(s: &str) -> String {
let mut out = String::with_capacity(s.len() + 4);
for (i, c) in s.chars().enumerate() {
if c.is_ascii_uppercase() {
if i > 0 {
out.push('_');
}
out.push(c.to_ascii_lowercase());
} else {
out.push(c);
}
}
out
}
fn discriminant_variant_pattern(v: &syn::Variant) -> proc_macro2::TokenStream {
let name = &v.ident;
match &v.fields {
Fields::Unit => quote! { Self::#name },
Fields::Unnamed(_) => quote! { Self::#name(..) },
Fields::Named(_) => quote! { Self::#name { .. } },
}
}
fn discriminant_variant_explicit_name(v: &syn::Variant) -> Option<String> {
for attr in &v.attrs {
if !attr.path().is_ident("discriminant") {
continue;
}
let mut out = None;
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("name") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
out = Some(s.value());
}
Ok(())
});
if out.is_some() {
return out;
}
}
None
}
#[proc_macro_derive(Discriminant, attributes(discriminant))]
pub fn derive_discriminant(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let enum_name = input.ident.clone();
let Data::Enum(de) = input.data.clone() else {
return syn::Error::new_spanned(
&enum_name,
"#[derive(Discriminant)] is only valid on enums",
)
.to_compile_error()
.into();
};
let mut method = "discriminant".to_string();
let mut case = DiscriminantCase::Kebab;
let mut also_display = false;
for attr in &input.attrs {
if !attr.path().is_ident("discriminant") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("method") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
method = s.value();
} else if meta.path.is_ident("case") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
if let Some(c) = DiscriminantCase::parse(&s.value()) {
case = c;
}
} else if meta.path.is_ident("also_display") {
also_display = true;
}
Ok(())
});
}
let method_ident = syn::Ident::new(&method, proc_macro2::Span::call_site());
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let arms: Vec<proc_macro2::TokenStream> = de
.variants
.iter()
.map(|v| {
let pattern = discriminant_variant_pattern(v);
let name_str = discriminant_variant_explicit_name(v)
.unwrap_or_else(|| case.apply(&v.ident.to_string()));
quote! { #pattern => #name_str }
})
.collect();
let display_impl = if also_display {
quote! {
impl #impl_generics ::core::fmt::Display for #enum_name #ty_generics #where_clause {
fn fmt(&self, f: &mut ::core::fmt::Formatter<'_>) -> ::core::fmt::Result {
f.write_str(self.#method_ident())
}
}
}
} else {
quote! {}
};
let expanded = quote! {
impl #impl_generics #enum_name #ty_generics #where_clause {
pub const fn #method_ident(&self) -> &'static str {
match self {
#(#arms),*
}
}
}
#display_impl
};
expanded.into()
}
#[proc_macro_derive(FromStrKind, attributes(from_str_kind))]
pub fn derive_from_str_kind(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let enum_name = input.ident.clone();
let Data::Enum(de) = input.data.clone() else {
return syn::Error::new_spanned(
&enum_name,
"#[derive(FromStrKind)] is only valid on enums",
)
.to_compile_error()
.into();
};
let mut case = DiscriminantCase::Kebab;
let mut error_name = format!("{enum_name}ParseError");
for attr in &input.attrs {
if !attr.path().is_ident("from_str_kind") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("case") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
if let Some(c) = DiscriminantCase::parse(&s.value()) {
case = c;
}
} else if meta.path.is_ident("error") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
error_name = s.value();
}
Ok(())
});
}
let error_ident = syn::Ident::new(&error_name, proc_macro2::Span::call_site());
let mut arms: Vec<proc_macro2::TokenStream> = Vec::new();
let mut known_strings: Vec<String> = Vec::new();
for v in &de.variants {
if !matches!(v.fields, Fields::Unit) {
return syn::Error::new_spanned(
&v.ident,
"#[derive(FromStrKind)] requires all variants to be unit variants (no data payloads)",
)
.to_compile_error()
.into();
}
let v_ident = &v.ident;
let explicit = v.attrs.iter().find_map(|attr| {
if !attr.path().is_ident("from_str_kind") {
return None;
}
let mut out = None;
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("name") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
out = Some(s.value());
}
Ok(())
});
out
});
let name_str = explicit.unwrap_or_else(|| case.apply(&v_ident.to_string()));
known_strings.push(name_str.clone());
arms.push(quote! { #name_str => Ok(Self::#v_ident) });
}
let known_list = known_strings.join(" | ");
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let expanded = quote! {
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct #error_ident {
pub input: ::std::string::String,
}
impl ::core::fmt::Display for #error_ident {
fn fmt(&self, f: &mut ::core::fmt::Formatter<'_>) -> ::core::fmt::Result {
write!(
f,
"unknown variant {input:?}; expected one of: {known}",
input = self.input,
known = #known_list,
)
}
}
impl ::std::error::Error for #error_ident {}
impl #impl_generics ::core::str::FromStr for #enum_name #ty_generics #where_clause {
type Err = #error_ident;
fn from_str(s: &str) -> ::core::result::Result<Self, Self::Err> {
match s {
#(#arms),*,
other => Err(#error_ident { input: other.to_string() }),
}
}
}
};
expanded.into()
}
fn is_variant_method_name(v: &syn::Variant) -> syn::Ident {
let explicit = v.attrs.iter().find_map(|attr| {
if !attr.path().is_ident("is_variant") {
return None;
}
let mut out = None;
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("name") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
out = Some(s.value());
}
Ok(())
});
out
});
let snake = explicit.unwrap_or_else(|| discriminant_to_snake(&v.ident.to_string()));
syn::Ident::new(&format!("is_{snake}"), proc_macro2::Span::call_site())
}
#[proc_macro_derive(IsVariant, attributes(is_variant))]
pub fn derive_is_variant(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let enum_name = input.ident.clone();
let Data::Enum(de) = input.data.clone() else {
return syn::Error::new_spanned(
&enum_name,
"#[derive(IsVariant)] is only valid on enums",
)
.to_compile_error()
.into();
};
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let methods: Vec<proc_macro2::TokenStream> = de
.variants
.iter()
.map(|v| {
let pattern = discriminant_variant_pattern(v);
let method_name = is_variant_method_name(v);
quote! {
pub const fn #method_name(&self) -> bool {
matches!(self, #pattern)
}
}
})
.collect();
let expanded = quote! {
impl #impl_generics #enum_name #ty_generics #where_clause {
#(#methods)*
}
};
expanded.into()
}
#[proc_macro_derive(BackendError, attributes(backend_error))]
pub fn derive_backend_error(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let enum_name = input.ident.clone();
let Data::Enum(de) = input.data.clone() else {
return syn::Error::new_spanned(
&enum_name,
"#[derive(BackendError)] is only valid on enums",
)
.to_compile_error()
.into();
};
let mut trait_path: syn::Path = syn::parse_quote!(BackendError);
let mut kind_method = "discriminant".to_string();
for attr in &input.attrs {
if !attr.path().is_ident("backend_error") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("trait_path") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
if let Ok(p) = syn::parse_str::<syn::Path>(&s.value()) {
trait_path = p;
}
} else if meta.path.is_ident("kind_method") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
kind_method = s.value();
}
Ok(())
});
}
let kind_method_ident = syn::Ident::new(&kind_method, proc_macro2::Span::call_site());
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let mut transient_patterns: Vec<proc_macro2::TokenStream> = Vec::new();
let mut auth_patterns: Vec<proc_macro2::TokenStream> = Vec::new();
for v in &de.variants {
let mut tags: std::collections::HashSet<String> = std::collections::HashSet::new();
for attr in &v.attrs {
if !attr.path().is_ident("backend_error") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if let Some(ident) = meta.path.get_ident() {
tags.insert(ident.to_string());
}
Ok(())
});
}
let pattern = discriminant_variant_pattern(v);
if tags.contains("transient") {
transient_patterns.push(pattern.clone());
}
if tags.contains("auth") {
auth_patterns.push(pattern.clone());
}
}
let is_retryable_body = if transient_patterns.is_empty() {
quote! { false }
} else {
quote! { matches!(self, #(#transient_patterns)|*) }
};
let is_auth_failure_body = if auth_patterns.is_empty() {
quote! { false }
} else {
quote! { matches!(self, #(#auth_patterns)|*) }
};
let expanded = quote! {
impl #impl_generics #trait_path for #enum_name #ty_generics #where_clause {
fn is_retryable(&self) -> bool {
#is_retryable_body
}
fn is_auth_failure(&self) -> bool {
#is_auth_failure_body
}
fn kind(&self) -> &'static str {
self.#kind_method_ident()
}
}
};
expanded.into()
}
#[proc_macro_derive(OutcomeLattice, attributes(outcome_lattice, outcome))]
pub fn derive_outcome_lattice(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let enum_name = input.ident.clone();
let Data::Enum(de) = input.data.clone() else {
return syn::Error::new_spanned(
&enum_name,
"#[derive(OutcomeLattice)] is only valid on enums",
)
.to_compile_error()
.into();
};
let mut trait_path: syn::Path = syn::parse_quote!(OutcomeLattice);
for attr in &input.attrs {
if !attr.path().is_ident("outcome_lattice") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("trait_path") {
let value = meta.value()?;
let s: syn::LitStr = value.parse()?;
if let Ok(p) = syn::parse_str::<syn::Path>(&s.value()) {
trait_path = p;
}
}
Ok(())
});
}
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let mut severity_arms: Vec<proc_macro2::TokenStream> = Vec::new();
let mut baseline_variant: Option<syn::Ident> = None;
for v in &de.variants {
let v_ident = &v.ident;
let pattern = discriminant_variant_pattern(v);
let mut sev: u32 = 0;
let mut is_baseline = false;
for attr in &v.attrs {
if !attr.path().is_ident("outcome") {
continue;
}
let _ = attr.parse_nested_meta(|meta| {
if meta.path.is_ident("severity") {
let value = meta.value()?;
let lit: syn::LitInt = value.parse()?;
sev = lit.base10_parse::<u32>().unwrap_or(0);
} else if meta.path.is_ident("baseline") {
is_baseline = true;
}
Ok(())
});
}
if is_baseline {
if !matches!(v.fields, Fields::Unit) {
return syn::Error::new_spanned(
v_ident,
"#[outcome(baseline)] requires a unit variant",
)
.to_compile_error()
.into();
}
if baseline_variant.is_some() {
return syn::Error::new_spanned(
v_ident,
"exactly one variant may carry #[outcome(baseline)]",
)
.to_compile_error()
.into();
}
baseline_variant = Some(v_ident.clone());
}
let lit = syn::LitInt::new(&sev.to_string(), proc_macro2::Span::call_site());
severity_arms.push(quote! { #pattern => #lit });
}
let Some(baseline_ident) = baseline_variant else {
return syn::Error::new_spanned(
&enum_name,
"exactly one variant must carry #[outcome(baseline)] to derive OutcomeLattice",
)
.to_compile_error()
.into();
};
let expanded = quote! {
impl #impl_generics #trait_path for #enum_name #ty_generics #where_clause {
fn severity(&self) -> u32 {
match self {
#(#severity_arms),*
}
}
fn baseline() -> Self {
Self::#baseline_ident
}
}
};
expanded.into()
}
fn to_kebab_case(s: &str) -> String {
let mut out = String::with_capacity(s.len() + 4);
let mut prev_lower = false;
let mut prev_digit = false;
for ch in s.chars() {
if ch.is_ascii_uppercase() {
if prev_lower || prev_digit {
out.push('-');
}
for c in ch.to_lowercase() {
out.push(c);
}
prev_lower = false;
prev_digit = false;
} else if ch.is_ascii_digit() {
out.push(ch);
prev_lower = false;
prev_digit = true;
} else {
out.push(ch);
prev_lower = true;
prev_digit = false;
}
}
out
}
#[proc_macro_attribute]
pub fn fsm(args: TokenStream, item: TokenStream) -> TokenStream {
let args = parse_macro_input!(args as FsmArgs);
let item = parse_macro_input!(item as syn::ItemEnum);
let name = &item.ident;
let label = &args.label;
quote! {
#[derive(
::core::clone::Clone,
::core::fmt::Debug,
::core::cmp::PartialEq,
::core::cmp::Eq,
::serde::Serialize,
::serde::Deserialize,
::gen_macros::TypedDispatcher,
::gen_macros::Discriminant,
::gen_macros::IsVariant,
)]
#[serde(tag = "kind", rename_all = "kebab-case")]
#item
::gen_platform::register_dispatcher!(#label, #name);
}
.into()
}
struct FsmArgs {
label: String,
}
impl syn::parse::Parse for FsmArgs {
fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
let kw: syn::Ident = input.parse()?;
if kw != "label" {
return Err(syn::Error::new(
kw.span(),
"expected `label = \"<catalog-label>\"`",
));
}
input.parse::<syn::Token![=]>()?;
let lit: syn::LitStr = input.parse()?;
Ok(Self { label: lit.value() })
}
}