1use proc_macro::TokenStream; 2use quote::quote; 3use syn::{Attribute, Data, DeriveInput, Fields}; 4 5pub fn enum_display(input: TokenStream) -> TokenStream { 6 let ast: DeriveInput = syn::parse(input).expect("unable to parse input"); 7 8 let name = &ast.ident; 9 10 let Data::Enum(data) = &ast.data else { 11 panic!("EnumDisplay derive macro can only be applied to enums; {name} is not an enum"); 12 }; 13 14 let match_arms: Vec<_> = data 15 .variants 16 .iter() 17 .map(|variant| { 18 let variant_name = &variant.ident; 19 assert!(variant.fields.is_empty(), "EnumDisplay macro only supports enums with only fieldless variants; {name}::{variant_name} has fields"); 20 21 let variant_name_str = variant_name.to_string(); 22 quote! { 23 Self::#variant_name => #variant_name_str 24 } 25 }) 26 .collect(); 27 28 let expanded = quote! { 29 impl #name { 30 pub fn to_str(&self) -> &'static str { 31 match self { 32 #(#match_arms,)* 33 } 34 } 35 } 36 37 impl ::std::fmt::Display for #name { 38 fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result { 39 ::std::write!(f, "{}", self.to_str()) 40 } 41 } 42 }; 43 44 expanded.into() 45} 46 47pub fn enum_from_str(input: TokenStream) -> TokenStream { 48 let ast: DeriveInput = syn::parse(input).expect("unable to parse input"); 49 50 let name = &ast.ident; 51 52 let Data::Enum(data) = &ast.data else { 53 panic!("EnumFromStr derive macro can only be applied to enums; {name} is not an enum"); 54 }; 55 56 let match_arms: Vec<_> = data 57 .variants 58 .iter() 59 .map(|variant| { 60 let variant_name = &variant.ident; 61 assert!(variant.fields.is_empty(), "EnumFromStr macro only supports enums with only fieldless variants; {name}::{variant_name} has fields"); 62 63 let variant_name_lowercase = variant_name.to_string().to_ascii_lowercase(); 64 quote! { 65 #variant_name_lowercase => ::std::result::Result::Ok(Self::#variant_name) 66 } 67 }) 68 .collect(); 69 70 let err_fmt_string = format!("invalid {name} string: '{{}}'"); 71 let expanded = quote! { 72 impl ::std::str::FromStr for #name { 73 type Err = ::std::string::String; 74 75 fn from_str(s: &str) -> ::std::result::Result<Self, Self::Err> { 76 match s.to_ascii_lowercase().as_str() { 77 #(#match_arms,)* 78 _ => ::std::result::Result::Err(::std::format!(#err_fmt_string, s)) 79 } 80 } 81 } 82 }; 83 84 expanded.into() 85} 86 87pub fn enum_all(input: TokenStream) -> TokenStream { 88 let input: DeriveInput = syn::parse(input).expect("Unable to parse input"); 89 90 let type_ident = &input.ident; 91 let Data::Enum(data) = &input.data else { 92 panic!("EnumAll only supports enums; {type_ident} is not an enum"); 93 }; 94 95 let variant_constructors: Vec<_> = data.variants.iter().map(|variant| { 96 let variant_ident = &variant.ident; 97 assert!( 98 matches!(variant.fields, Fields::Unit), 99 "EnumAll only supports enums with fieldless variants; {type_ident}::{variant_ident} is not a fieldless variant", 100 ); 101 102 quote! { 103 Self::#variant_ident 104 } 105 }).collect(); 106 107 let num_variants = data.variants.len(); 108 let expanded = quote! { 109 impl #type_ident { 110 pub const ALL: [Self; #num_variants] = [#(#variant_constructors,)*]; 111 } 112 }; 113 114 expanded.into() 115} 116 117pub fn custom_value_enum(input: TokenStream) -> TokenStream { 118 let input: DeriveInput = syn::parse(input).expect("Unable to parse input"); 119 120 let type_ident = &input.ident; 121 122 let Data::Enum(data) = &input.data else { 123 panic!("CustomValueEnum only supports enums"); 124 }; 125 126 let included_fields: Vec<_> = data 127 .variants 128 .iter() 129 .filter_map(|variant| { 130 let skip = value_enum_should_skip(&variant.attrs); 131 (!skip).then_some(&variant.ident) 132 }) 133 .collect(); 134 135 let expanded = quote! { 136 impl ::clap::ValueEnum for #type_ident { 137 fn value_variants<'a>() -> &'a [Self] { 138 const ALL: &[#type_ident] = &[ 139 #(#type_ident::#included_fields,)* 140 ]; 141 142 ALL 143 } 144 145 fn to_possible_value(&self) -> ::std::option::Option<::clap::builder::PossibleValue> { 146 match self { 147 #( 148 Self::#included_fields => ::std::option::Option::Some( 149 ::clap::builder::PossibleValue::new(self.to_str()) 150 ), 151 )* 152 _ => ::std::option::Option::None 153 } 154 } 155 } 156 }; 157 158 expanded.into() 159} 160 161fn value_enum_should_skip(attrs: &[Attribute]) -> bool { 162 attrs.iter().any(|attr| { 163 if !attr.path().is_ident("value_enum") { 164 return false; 165 } 166 167 let mut skip = false; 168 attr.parse_nested_meta(|meta| { 169 if meta.path.is_ident("skip") { 170 skip = true; 171 Ok(()) 172 } else { 173 Err(meta.error("Invalid value_enum meta")) 174 } 175 }) 176 .expect("Failed to parse value_enum attribute"); 177 178 skip 179 }) 180} 181 182pub fn match_each_variant_macro(input: TokenStream) -> TokenStream { 183 let input: DeriveInput = syn::parse(input).expect("Unable to parse input"); 184 185 let ident = &input.ident; 186 187 let Data::Enum(data) = &input.data else { 188 panic!("{ident} is not an enum"); 189 }; 190 191 let match_arms: Vec<_> = data 192 .variants 193 .iter() 194 .map(|variant| { 195 let variant_ident = &variant.ident; 196 197 let Fields::Unnamed(fields) = &variant.fields else { 198 panic!("{ident}::{variant_ident} should have unnamed fields"); 199 }; 200 201 assert_eq!( 202 fields.unnamed.len(), 203 1, 204 "{ident}::{variant_ident} has {} unnamed fields, expected 1", 205 fields.unnamed.len() 206 ); 207 208 quote! { 209 #ident::#variant_ident($field) => $match_arm 210 } 211 }) 212 .collect(); 213 214 let variant_match_arms: Vec<_> = data 215 .variants 216 .iter() 217 .map(|variant| { 218 let variant_ident = &variant.ident; 219 220 // No need to re-validate the enum variants, that was done above 221 222 quote! { 223 #ident::#variant_ident($field) => #ident::#variant_ident($match_arm) 224 } 225 }) 226 .collect(); 227 228 let expanded = quote! { 229 macro_rules! match_each_variant { 230 ($value:expr, $field:ident => $match_arm:expr) => { 231 match $value { 232 #(#match_arms,)* 233 } 234 }; 235 ($value:expr, $field:ident => :variant($match_arm:expr)) => { 236 match $value { 237 #(#variant_match_arms,)* 238 } 239 }; 240 } 241 }; 242 243 expanded.into() 244}