|
| 1 | +use proc_macro2::{Span, TokenStream}; |
| 2 | +use quote::{format_ident, quote}; |
| 3 | +use syn::{ |
| 4 | + parse::{Parse, ParseStream}, |
| 5 | + punctuated::Punctuated, |
| 6 | + spanned::Spanned, |
| 7 | + Attribute, Expr, Field, Fields, Ident, ItemStruct, Token, Visibility |
| 8 | +}; |
| 9 | + |
| 10 | +// --- Input types --- |
| 11 | + |
| 12 | +struct VersionVariant { |
| 13 | + name: Ident, |
| 14 | + condition: Option<Expr> |
| 15 | +} |
| 16 | + |
| 17 | +struct VersionedArgs { |
| 18 | + variants: Vec<VersionVariant> |
| 19 | +} |
| 20 | + |
| 21 | +// --- Field annotation types --- |
| 22 | + |
| 23 | +enum FieldVersionInfo { |
| 24 | + AllVersions, |
| 25 | + OnlyIn(Ident) |
| 26 | +} |
| 27 | + |
| 28 | +struct VersionedField { |
| 29 | + field: Field, |
| 30 | + version_info: FieldVersionInfo |
| 31 | +} |
| 32 | + |
| 33 | +// --- Parsing --- |
| 34 | + |
| 35 | +impl Parse for VersionVariant { |
| 36 | + fn parse(input: ParseStream) -> syn::Result<Self> { |
| 37 | + let name: Ident = input.parse()?; |
| 38 | + let condition = if input.peek(Token![if]) { |
| 39 | + input.parse::<Token![if]>()?; |
| 40 | + Some(input.parse()?) |
| 41 | + } else { |
| 42 | + None |
| 43 | + }; |
| 44 | + Ok(Self { name, condition }) |
| 45 | + } |
| 46 | +} |
| 47 | + |
| 48 | +impl Parse for VersionedArgs { |
| 49 | + fn parse(input: ParseStream) -> syn::Result<Self> { |
| 50 | + let variants = Punctuated::<VersionVariant, Token![,]>::parse_terminated(input)?.into_iter().collect(); |
| 51 | + Ok(Self { variants }) |
| 52 | + } |
| 53 | +} |
| 54 | + |
| 55 | +// --- Field processing --- |
| 56 | + |
| 57 | +fn extract_field_version_info(field: &mut Field) -> FieldVersionInfo { |
| 58 | + let mut result = FieldVersionInfo::AllVersions; |
| 59 | + field.attrs.retain(|attr| { |
| 60 | + if attr.path().is_ident("only_in") { |
| 61 | + if let Ok(variant) = attr.parse_args::<Ident>() { |
| 62 | + result = FieldVersionInfo::OnlyIn(variant); |
| 63 | + return false; |
| 64 | + } |
| 65 | + } |
| 66 | + true |
| 67 | + }); |
| 68 | + result |
| 69 | +} |
| 70 | + |
| 71 | +fn extract_versioned_fields(input: &mut ItemStruct) -> Result<Vec<VersionedField>, syn::Error> { |
| 72 | + let Fields::Named(ref mut fields) = input.fields else { |
| 73 | + return Err(syn::Error::new( |
| 74 | + input.fields.span(), |
| 75 | + "#[versioned] only supports structs with named fields" |
| 76 | + )); |
| 77 | + }; |
| 78 | + Ok(fields |
| 79 | + .named |
| 80 | + .iter_mut() |
| 81 | + .map(|field| { |
| 82 | + let version_info = extract_field_version_info(field); |
| 83 | + VersionedField { |
| 84 | + field: field.clone(), |
| 85 | + version_info |
| 86 | + } |
| 87 | + }) |
| 88 | + .collect()) |
| 89 | +} |
| 90 | + |
| 91 | +// --- Validation --- |
| 92 | + |
| 93 | +fn validate_variants(variants: &[VersionVariant]) -> Option<syn::Error> { |
| 94 | + if variants.len() < 2 { |
| 95 | + return Some(syn::Error::new(Span::call_site(), "#[versioned] requires at least 2 variants")); |
| 96 | + } |
| 97 | + if variants.last().unwrap().condition.is_some() { |
| 98 | + return Some(syn::Error::new( |
| 99 | + variants.last().unwrap().name.span(), |
| 100 | + "the last variant must be the fallback (no `if` condition)" |
| 101 | + )); |
| 102 | + } |
| 103 | + for v in &variants[..variants.len() - 1] { |
| 104 | + if v.condition.is_none() { |
| 105 | + return Some(syn::Error::new(v.name.span(), "only the last variant can omit the `if` condition")); |
| 106 | + } |
| 107 | + } |
| 108 | + None |
| 109 | +} |
| 110 | + |
| 111 | +// --- Code generation --- |
| 112 | + |
| 113 | +fn field_to_tokens(field: &Field) -> TokenStream { |
| 114 | + let attrs = &field.attrs; |
| 115 | + let vis = &field.vis; |
| 116 | + let ident = &field.ident; |
| 117 | + let ty = &field.ty; |
| 118 | + quote! { |
| 119 | + #(#attrs)* |
| 120 | + #vis #ident: #ty |
| 121 | + } |
| 122 | +} |
| 123 | + |
| 124 | +fn variant_inner_name(struct_ident: &Ident, variant: &VersionVariant) -> Ident { |
| 125 | + format_ident!("{}{}", struct_ident, variant.name) |
| 126 | +} |
| 127 | + |
| 128 | +fn variant_field_name(variant: &VersionVariant) -> Ident { |
| 129 | + format_ident!("{}", variant.name.to_string().to_lowercase()) |
| 130 | +} |
| 131 | + |
| 132 | +fn generate_variant_struct(struct_ident: &Ident, struct_attrs: &[Attribute], variant: &VersionVariant, fields: &[VersionedField]) -> TokenStream { |
| 133 | + let inner_name = variant_inner_name(struct_ident, variant); |
| 134 | + let variant_name = &variant.name; |
| 135 | + |
| 136 | + let variant_fields = fields.iter().filter_map(|vf| match &vf.version_info { |
| 137 | + FieldVersionInfo::AllVersions => Some(field_to_tokens(&vf.field)), |
| 138 | + FieldVersionInfo::OnlyIn(v) if v == variant_name => Some(field_to_tokens(&vf.field)), |
| 139 | + FieldVersionInfo::OnlyIn(_) => None |
| 140 | + }); |
| 141 | + |
| 142 | + quote! { |
| 143 | + #[allow(dead_code)] |
| 144 | + #[derive(Copy, Clone)] |
| 145 | + #(#struct_attrs)* |
| 146 | + struct #inner_name { |
| 147 | + #(#variant_fields,)* |
| 148 | + } |
| 149 | + } |
| 150 | +} |
| 151 | + |
| 152 | +fn generate_union(struct_ident: &Ident, vis: &Visibility, attrs: &[Attribute], variants: &[VersionVariant]) -> TokenStream { |
| 153 | + let union_fields = variants.iter().map(|v| { |
| 154 | + let field_name = variant_field_name(v); |
| 155 | + let inner_name = variant_inner_name(struct_ident, v); |
| 156 | + quote! { #field_name: #inner_name } |
| 157 | + }); |
| 158 | + |
| 159 | + quote! { |
| 160 | + #(#attrs)* |
| 161 | + #vis union #struct_ident { |
| 162 | + #(#union_fields,)* |
| 163 | + } |
| 164 | + } |
| 165 | +} |
| 166 | + |
| 167 | +fn build_dispatch_chain(struct_ident: &Ident, variants: &[VersionVariant], helper_names: &[Ident]) -> TokenStream { |
| 168 | + let fallback = helper_names.last().unwrap(); |
| 169 | + let mut chain = quote! { #struct_ident::#fallback }; |
| 170 | + |
| 171 | + for (variant, helper) in variants[..variants.len() - 1].iter().zip(&helper_names[..helper_names.len() - 1]).rev() { |
| 172 | + let cond = variant.condition.as_ref().unwrap(); |
| 173 | + chain = quote! { |
| 174 | + if #cond { #struct_ident::#helper } else { #chain } |
| 175 | + }; |
| 176 | + } |
| 177 | + |
| 178 | + chain |
| 179 | +} |
| 180 | + |
| 181 | +fn generate_field_statics(struct_ident: &Ident, field: &Field, variants: &[VersionVariant]) -> TokenStream { |
| 182 | + let field_name = field.ident.as_ref().unwrap(); |
| 183 | + let field_ty = &field.ty; |
| 184 | + let static_name = format_ident!("__VERSIONED_{}_{}", struct_ident, field_name); |
| 185 | + let static_name_mut = format_ident!("__VERSIONED_MUT_{}_{}", struct_ident, field_name); |
| 186 | + let fallback = variants.last().unwrap(); |
| 187 | + let fallback_fn = format_ident!("__versioned_{}_{}", field_name, variant_field_name(fallback)); |
| 188 | + let fallback_fn_mut = format_ident!("__versioned_{}_{}_mut", field_name, variant_field_name(fallback)); |
| 189 | + |
| 190 | + quote! { |
| 191 | + #[allow(non_upper_case_globals)] |
| 192 | + static mut #static_name: fn(&#struct_ident) -> &#field_ty = #struct_ident::#fallback_fn; |
| 193 | + #[allow(non_upper_case_globals)] |
| 194 | + static mut #static_name_mut: fn(&mut #struct_ident) -> &mut #field_ty = #struct_ident::#fallback_fn_mut; |
| 195 | + } |
| 196 | +} |
| 197 | + |
| 198 | +fn generate_init_fn(struct_ident: &Ident, pub_fields: &[&VersionedField], variants: &[VersionVariant]) -> TokenStream { |
| 199 | + let init_fn_name = format_ident!("__versioned_init_{}", struct_ident.to_string().to_lowercase()); |
| 200 | + |
| 201 | + let assignments: Vec<_> = pub_fields |
| 202 | + .iter() |
| 203 | + .map(|vf| { |
| 204 | + let field_name = vf.field.ident.as_ref().unwrap(); |
| 205 | + let static_name = format_ident!("__VERSIONED_{}_{}", struct_ident, field_name); |
| 206 | + let static_name_mut = format_ident!("__VERSIONED_MUT_{}_{}", struct_ident, field_name); |
| 207 | + |
| 208 | + let helper_names: Vec<Ident> = variants |
| 209 | + .iter() |
| 210 | + .map(|v| format_ident!("__versioned_{}_{}", field_name, variant_field_name(v))) |
| 211 | + .collect(); |
| 212 | + let helper_names_mut: Vec<Ident> = variants |
| 213 | + .iter() |
| 214 | + .map(|v| format_ident!("__versioned_{}_{}_mut", field_name, variant_field_name(v))) |
| 215 | + .collect(); |
| 216 | + |
| 217 | + let dispatch = build_dispatch_chain(struct_ident, variants, &helper_names); |
| 218 | + let dispatch_mut = build_dispatch_chain(struct_ident, variants, &helper_names_mut); |
| 219 | + |
| 220 | + quote! { |
| 221 | + #static_name = #dispatch; |
| 222 | + #static_name_mut = #dispatch_mut; |
| 223 | + } |
| 224 | + }) |
| 225 | + .collect(); |
| 226 | + |
| 227 | + quote! { |
| 228 | + fn #init_fn_name() -> Result<(), String> { |
| 229 | + unsafe { |
| 230 | + #(#assignments)* |
| 231 | + } |
| 232 | + Ok(()) |
| 233 | + } |
| 234 | + crate::inventory::submit!(crate::init::PartialInitFunc(#init_fn_name)); |
| 235 | + } |
| 236 | +} |
| 237 | + |
| 238 | +fn generate_impl(struct_ident: &Ident, variants: &[VersionVariant], pub_fields: &[&VersionedField]) -> TokenStream { |
| 239 | + let methods = pub_fields.iter().map(|vf| { |
| 240 | + let field = &vf.field; |
| 241 | + let field_name = field.ident.as_ref().unwrap(); |
| 242 | + let field_ty = &field.ty; |
| 243 | + let field_name_mut = format_ident!("{}_mut", field_name); |
| 244 | + let static_name = format_ident!("__VERSIONED_{}_{}", struct_ident, field_name); |
| 245 | + let static_name_mut = format_ident!("__VERSIONED_MUT_{}_{}", struct_ident, field_name); |
| 246 | + |
| 247 | + let helpers = variants.iter().map(|v| { |
| 248 | + let helper_name = format_ident!("__versioned_{}_{}", field_name, variant_field_name(v)); |
| 249 | + let helper_name_mut = format_ident!("__versioned_{}_{}_mut", field_name, variant_field_name(v)); |
| 250 | + let union_field = variant_field_name(v); |
| 251 | + quote! { |
| 252 | + fn #helper_name(this: &Self) -> &#field_ty { |
| 253 | + unsafe { &this.#union_field.#field_name } |
| 254 | + } |
| 255 | + fn #helper_name_mut(this: &mut Self) -> &mut #field_ty { |
| 256 | + unsafe { &mut this.#union_field.#field_name } |
| 257 | + } |
| 258 | + } |
| 259 | + }); |
| 260 | + |
| 261 | + quote! { |
| 262 | + #(#helpers)* |
| 263 | + pub fn #field_name(&self) -> &#field_ty { |
| 264 | + unsafe { #static_name(self) } |
| 265 | + } |
| 266 | + pub fn #field_name_mut(&mut self) -> &mut #field_ty { |
| 267 | + unsafe { #static_name_mut(self) } |
| 268 | + } |
| 269 | + } |
| 270 | + }); |
| 271 | + |
| 272 | + quote! { |
| 273 | + impl #struct_ident { |
| 274 | + #(#methods)* |
| 275 | + } |
| 276 | + } |
| 277 | +} |
| 278 | + |
| 279 | +// --- Entry point --- |
| 280 | + |
| 281 | +pub fn versioned(attr: TokenStream, item: TokenStream) -> TokenStream { |
| 282 | + let args = match syn::parse2::<VersionedArgs>(attr) { |
| 283 | + Ok(a) => a, |
| 284 | + Err(e) => return e.to_compile_error() |
| 285 | + }; |
| 286 | + let mut input = match syn::parse2::<ItemStruct>(item) { |
| 287 | + Ok(i) => i, |
| 288 | + Err(e) => return e.to_compile_error() |
| 289 | + }; |
| 290 | + |
| 291 | + if let Some(err) = validate_variants(&args.variants) { |
| 292 | + return err.to_compile_error(); |
| 293 | + } |
| 294 | + |
| 295 | + let fields = match extract_versioned_fields(&mut input) { |
| 296 | + Ok(f) => f, |
| 297 | + Err(e) => return e.to_compile_error() |
| 298 | + }; |
| 299 | + |
| 300 | + let pub_all_fields: Vec<&VersionedField> = fields |
| 301 | + .iter() |
| 302 | + .filter(|vf| matches!(vf.version_info, FieldVersionInfo::AllVersions) && matches!(vf.field.vis, Visibility::Public(_))) |
| 303 | + .collect(); |
| 304 | + |
| 305 | + let variant_structs = args |
| 306 | + .variants |
| 307 | + .iter() |
| 308 | + .map(|v| generate_variant_struct(&input.ident, &input.attrs, v, &fields)); |
| 309 | + let union_def = generate_union(&input.ident, &input.vis, &input.attrs, &args.variants); |
| 310 | + let field_statics = pub_all_fields |
| 311 | + .iter() |
| 312 | + .map(|vf| generate_field_statics(&input.ident, &vf.field, &args.variants)); |
| 313 | + let init_fn = if pub_all_fields.is_empty() { |
| 314 | + quote! {} |
| 315 | + } else { |
| 316 | + generate_init_fn(&input.ident, &pub_all_fields, &args.variants) |
| 317 | + }; |
| 318 | + let impl_block = generate_impl(&input.ident, &args.variants, &pub_all_fields); |
| 319 | + |
| 320 | + quote! { |
| 321 | + #(#variant_structs)* |
| 322 | + #union_def |
| 323 | + #(#field_statics)* |
| 324 | + #init_fn |
| 325 | + #impl_block |
| 326 | + } |
| 327 | +} |
0 commit comments