diff options
Diffstat (limited to 'macros/src/encode.rs')
| -rw-r--r-- | macros/src/encode.rs | 516 |
1 files changed, 516 insertions, 0 deletions
diff --git a/macros/src/encode.rs b/macros/src/encode.rs new file mode 100644 index 0000000..d8df604 --- /dev/null +++ b/macros/src/encode.rs @@ -0,0 +1,516 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use quote::quote; +use syn::{Data, DataEnum, DataStruct, DeriveInput, Generics, Ident, Type}; + +// ── Shared token helpers ──────────────────────────────────────────── + +pub(crate) fn next_ident(iter: &mut proc_macro2::token_stream::IntoIter) -> Option<Ident> { + if let Some(proc_macro2::TokenTree::Ident(ident)) = iter.clone().next() { + iter.next(); + return Some(ident); + } + None +} + +pub(crate) fn next_group( + iter: &mut proc_macro2::token_stream::IntoIter, +) -> Option<proc_macro2::Group> { + if let Some(proc_macro2::TokenTree::Group(group)) = iter.clone().next() { + iter.next(); + return Some(group); + } + None +} + +pub(crate) fn skip_commas(iter: &mut proc_macro2::token_stream::IntoIter) { + loop { + let peek = iter.clone().next(); + match peek { + Some(proc_macro2::TokenTree::Punct(p)) if p.as_char() == ',' => { + iter.next(); + } + _ => break, + } + } +} + +// ── Field attribute parsing ───────────────────────────────────────── + +enum EncodeAttr { + Normal, + Skip, + Count(syn::Type), + Len(syn::Type), + Custom(syn::Path), +} + +struct FieldAttrs { + encode: EncodeAttr, + condition: Option<syn::Expr>, +} + +fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { + let mut encode: Option<EncodeAttr> = None; + let mut condition: Option<syn::Expr> = None; + + for attr in &field.attrs { + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + encode = Some(EncodeAttr::Skip); + } + "count" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected count(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in count(...)"); + encode = Some(EncodeAttr::Count(ty)); + } + "len" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected len(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in len(...)"); + encode = Some(EncodeAttr::Len(ty)); + } + "with" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected with(path)"); + let path: syn::Path = + syn::parse2(group.stream()).expect("expected path in with(...)"); + encode = Some(EncodeAttr::Custom(path)); + } + "if" => { + assert!(condition.is_none(), "multiple codec(if(...)) on one field"); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper: `{other}`"), + } + } + } + + FieldAttrs { + encode: encode.unwrap_or(EncodeAttr::Normal), + condition, + } +} + +// ── Variant attribute parsing ─────────────────────────────────────── + +pub(crate) struct VariantAttrs { + pub id: Option<u8>, + pub condition: Option<syn::Expr>, + pub skip: bool, +} + +pub(crate) fn parse_variant_attrs(variant: &syn::Variant) -> VariantAttrs { + let mut id: Option<u8> = None; + let mut condition: Option<syn::Expr> = None; + let mut skip = false; + + for attr in &variant.attrs { + // #[id(0x00)] + if attr.path().is_ident("id") { + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + let lit: syn::LitInt = + syn::parse2(list.tokens.clone()).expect("expected u8 literal in #[id(...)]"); + let val: u8 = lit.base10_parse().expect("id must be a u8 value"); + assert!(id.is_none(), "multiple #[id(...)] on one variant"); + id = Some(val); + continue; + } + + // #[codec(if(...)), codec(skip)] + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + skip = true; + } + "if" => { + assert!( + condition.is_none(), + "multiple codec(if(...)) on one variant" + ); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper on variant: `{other}`"), + } + } + } + + VariantAttrs { + id, + condition, + skip, + } +} + +// ── Context type extraction ───────────────────────────────────────── + +fn extract_context_ty(attrs: &[syn::Attribute]) -> Option<Type> { + attrs.iter().find_map(|attr| { + if attr.path().is_ident("context") { + let syn::Meta::List(list) = &attr.meta else { + return None; + }; + syn::parse2::<Type>(list.tokens.clone()).ok() + } else { + None + } + }) +} + +// ── Build impl generics ───────────────────────────────────────────── + +fn build_encode_impl_generics( + params: &syn::punctuated::Punctuated<syn::GenericParam, syn::Token![,]>, + where_clause: &Option<syn::WhereClause>, + context_ty: &Option<Type>, +) -> (TokenStream2, TokenStream2, TokenStream2, TokenStream2) { + let params = params.iter().cloned().collect::<Vec<_>>(); + let has_params = !params.is_empty(); + + if let Some(ctx_ty) = context_ty { + // With context type: impl<#params> Encode<Ctx> for Ident + let impl_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { #ctx_ty }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } else { + // Without context: impl<Data, #params> Encode<Data> for Ident + let mut all_params = vec![syn::parse_quote! { Data }]; + all_params.extend(params.iter().cloned()); + let impl_generics = quote! { <#(#all_params),*> }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { Data }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } +} + +// ── Encode code generation ────────────────────────────────────────── + +pub fn derive_encode( + DeriveInput { + ident, + generics, + data, + attrs, + .. + }: DeriveInput, +) -> TokenStream { + let context_ty = extract_context_ty(&attrs); + match data { + Data::Struct(data_struct) => impl_encode_struct(ident, generics, data_struct, context_ty), + Data::Enum(data_enum) => impl_encode_enum(ident, generics, data_enum, context_ty), + Data::Union(_) => panic!("Not implemented for Union"), + } +} + +fn gen_encode_field( + field_access: TokenStream2, + attrs: &FieldAttrs, + _data_param: &TokenStream2, +) -> TokenStream2 { + let inner = match &attrs.encode { + EncodeAttr::Normal => quote! { + len += #field_access.encode(buffer, ctx)?; + }, + EncodeAttr::Skip => quote! {}, + EncodeAttr::Count(len_type) => quote! { + let count: #len_type = #field_access.len() as #len_type; + len += count.encode(buffer, ctx)?; + for item in &#field_access { + len += item.encode(buffer, ctx)?; + } + }, + EncodeAttr::Len(len_type) => quote! { + let mut __buf = Vec::with_capacity(2048); + #field_access.encode(&mut __buf, ctx)?; + let __len: #len_type = zr_protocol::types::size::Size::from_size(__buf.len()); + len += __len.encode(buffer, ctx)?; + buffer.extend_from_slice(&__buf); + len += __buf.len(); + }, + EncodeAttr::Custom(fn_path) => quote! { + len += #fn_path(&#field_access, buffer, ctx)?; + }, + }; + + match &attrs.condition { + Some(expr) => quote! { + if #expr { + #inner + } + }, + None => inner, + } +} + +fn gen_encode_body( + enum_ident: &Ident, + variant_ident: &Ident, + pattern: &TokenStream2, + encode_fields: &TokenStream2, + discriminant: u8, + condition: &Option<syn::Expr>, +) -> TokenStream2 { + let encode_disc_and_fields = quote! { + len += #discriminant.encode(buffer, ctx)?; + #encode_fields + }; + + match condition { + Some(expr) => quote! { + #enum_ident::#variant_ident #pattern => { + if #expr { + #encode_disc_and_fields + } else { + return Err(zr_protocol::codec::error::CodecError::Custom( + concat!("variant ", stringify!(#variant_ident), " not valid in current context").into() + )); + } + } + }, + None => quote! { + #enum_ident::#variant_ident #pattern => { + #encode_disc_and_fields + } + }, + } +} + +fn impl_encode_struct( + ident: Ident, + generics: Generics, + data_struct: DataStruct, + context_ty: Option<Type>, +) -> TokenStream { + let (impl_generics, ty_generics, where_clause, data_param) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + let encode_fields = match &data_struct.fields { + syn::Fields::Named(fields) => { + let bodies: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let attrs = parse_field_attrs(field); + gen_encode_field(quote! { self.#name }, &attrs, &data_param) + }) + .collect(); + quote! { #(#bodies)* } + } + syn::Fields::Unnamed(fields) => { + let bodies: Vec<_> = fields + .unnamed + .iter() + .enumerate() + .map(|(i, field)| { + let idx = syn::Index::from(i); + let attrs = parse_field_attrs(field); + gen_encode_field(quote! { self.#idx }, &attrs, &data_param) + }) + .collect(); + quote! { #(#bodies)* } + } + syn::Fields::Unit => quote! {}, + }; + + quote! { + impl #impl_generics zr_protocol::codec::encode::Encode<#data_param> for #ident #ty_generics #where_clause { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<usize> { + let mut len = 0; + #encode_fields + Ok(len) + } + } + } + .into() +} + +fn impl_encode_enum( + ident: Ident, + generics: Generics, + data_enum: DataEnum, + context_ty: Option<Type>, +) -> TokenStream { + let (_, _, _, data_param) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + // First pass: collect explicit IDs and count auto-IDs needed + let mut explicit_ids = std::collections::HashSet::new(); + + // Pre-parse all variant attributes + let parsed: Vec<_> = data_enum + .variants + .iter() + .map(|v| { + let attrs = parse_variant_attrs(v); + if let Some(id) = attrs.id { + explicit_ids.insert(id); + } + attrs + }) + .collect(); + + let match_arms: Vec<_> = { + let mut auto_counter: u8 = 0; + data_enum + .variants + .iter() + .zip(parsed.iter()) + .map(|(variant, attrs)| { + let variant_ident = &variant.ident; + + if attrs.skip { + return quote! { + #ident::#variant_ident { .. } => { + unreachable!("skipped variant") + } + }; + } + + // Resolve discriminant: explicit #[id(...)] or auto-increment + let discriminant = if let Some(id) = attrs.id { + id + } else { + while explicit_ids.contains(&auto_counter) { + auto_counter = auto_counter.wrapping_add(1); + } + let d = auto_counter; + auto_counter = auto_counter.wrapping_add(1); + d + }; + + let (pattern, encode_fields) = match &variant.fields { + syn::Fields::Named(fields) => { + let names: Vec<_> = fields + .named + .iter() + .map(|f| f.ident.as_ref().unwrap()) + .collect(); + let pattern = quote! { { #(#names),* } }; + let bodies: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let fa = parse_field_attrs(field); + gen_encode_field(quote! { #name }, &fa, &data_param) + }) + .collect(); + (pattern, quote! { #(#bodies)* }) + } + syn::Fields::Unnamed(fields) => { + let names: Vec<_> = (0..fields.unnamed.len()) + .map(|i| Ident::new(&format!("_{}", i), proc_macro2::Span::call_site())) + .collect(); + let pattern = quote! { ( #(#names),* ) }; + let bodies: Vec<_> = fields + .unnamed + .iter() + .enumerate() + .map(|(i, field)| { + let name = &names[i]; + let fa = parse_field_attrs(field); + gen_encode_field(quote! { #name }, &fa, &data_param) + }) + .collect(); + (pattern, quote! { #(#bodies)* }) + } + syn::Fields::Unit => (quote! {}, quote! {}), + }; + + gen_encode_body( + &ident, + variant_ident, + &pattern, + &encode_fields, + discriminant, + &attrs.condition, + ) + }) + .collect() + }; + + let (impl_generics, ty_generics, where_clause, _) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + quote! { + impl #impl_generics zr_protocol::codec::encode::Encode<#data_param> for #ident #ty_generics #where_clause { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<usize> { + let mut len = 0; + match self { + #(#match_arms),* + } + Ok(len) + } + } + } + .into() +} |
